refactor(run_agent): extract ApiRequestHooksMixin (agent/api_request_hooks.py)
This commit is contained in:
+2
-264
@@ -49,7 +49,6 @@ from typing import List, Dict, Any, Optional, Callable
|
||||
# it.
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
from hermes_constants import get_hermes_home
|
||||
|
||||
@@ -176,6 +175,7 @@ from agent.client_lifecycle import ( # noqa: F401 # _routermint_headers/_qwen_
|
||||
)
|
||||
from agent.stream_delivery import StreamDeliveryMixin
|
||||
from agent.status_output import StatusOutputMixin
|
||||
from agent.api_request_hooks import ApiRequestHooksMixin
|
||||
from agent.lazy_forward import forward as _forward, forward_static as _forward_static
|
||||
from agent.redact import redact_sensitive_text
|
||||
from agent.session_activity import ActivityProvenance
|
||||
@@ -183,7 +183,6 @@ from agent.model_metadata import (
|
||||
estimate_request_tokens_rough, # noqa: F401 # re-exported for tests that mock.patch("run_agent.estimate_request_tokens_rough")
|
||||
is_local_endpoint,
|
||||
)
|
||||
from agent.usage_pricing import normalize_usage
|
||||
# Re-exported for tests that monkeypatch these symbols on run_agent.
|
||||
from agent.context_compressor import ( # noqa: F401
|
||||
COMPRESSED_SUMMARY_METADATA_KEY,
|
||||
@@ -354,7 +353,7 @@ class _StreamErrorEvent(Exception):
|
||||
}
|
||||
|
||||
|
||||
class AIAgent(ClientLifecycleMixin, StreamDeliveryMixin, StatusOutputMixin):
|
||||
class AIAgent(ClientLifecycleMixin, StreamDeliveryMixin, StatusOutputMixin, ApiRequestHooksMixin):
|
||||
"""AI Agent with tool calling capabilities."""
|
||||
|
||||
_TOOL_CALL_ARGUMENTS_CORRUPTION_MARKER = (
|
||||
@@ -2058,267 +2057,6 @@ class AIAgent(ClientLifecycleMixin, StreamDeliveryMixin, StatusOutputMixin):
|
||||
|
||||
_extract_api_error_context = _forward_static("agent.agent_runtime_helpers", "extract_api_error_context")
|
||||
|
||||
def _usage_summary_for_api_request_hook(self, response: Any) -> Optional[Dict[str, Any]]:
|
||||
"""Token buckets for ``post_api_request`` plugins (no raw ``response`` object)."""
|
||||
if response is None:
|
||||
return None
|
||||
raw_usage = getattr(response, "usage", None)
|
||||
if not raw_usage:
|
||||
return None
|
||||
from dataclasses import asdict
|
||||
|
||||
cu = normalize_usage(raw_usage, provider=self.provider, api_mode=self.api_mode)
|
||||
summary = asdict(cu)
|
||||
summary.pop("raw_usage", None)
|
||||
summary["prompt_tokens"] = cu.prompt_tokens
|
||||
summary["total_tokens"] = cu.total_tokens
|
||||
return summary
|
||||
|
||||
@staticmethod
|
||||
def _hook_payload_max_chars() -> int:
|
||||
raw = os.getenv("HERMES_PLUGIN_PAYLOAD_MAX_CHARS", "50000")
|
||||
try:
|
||||
return max(1000, int(raw))
|
||||
except (TypeError, ValueError):
|
||||
return 50000
|
||||
|
||||
@staticmethod
|
||||
def _is_sensitive_hook_key(key: Any) -> bool:
|
||||
if not isinstance(key, str):
|
||||
return False
|
||||
lowered = key.lower().replace("-", "_")
|
||||
exact = {
|
||||
"api_key",
|
||||
"authorization",
|
||||
"proxy_authorization",
|
||||
"cookie",
|
||||
"set_cookie",
|
||||
}
|
||||
return lowered in exact or lowered.endswith("_api_key")
|
||||
|
||||
@classmethod
|
||||
def _hook_jsonable(
|
||||
cls,
|
||||
value: Any,
|
||||
*,
|
||||
depth: int = 0,
|
||||
max_depth: int = 8,
|
||||
max_string: int = 8000,
|
||||
max_sequence: int = 200,
|
||||
) -> Any:
|
||||
if depth > max_depth:
|
||||
return f"<{type(value).__name__} depth limit>"
|
||||
if value is None or isinstance(value, (bool, int, float)):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
if len(value) > max_string:
|
||||
return value[:max_string] + f"...[truncated {len(value) - max_string} chars]"
|
||||
return value
|
||||
if isinstance(value, (bytes, bytearray)):
|
||||
return f"<{len(value)} bytes>"
|
||||
if isinstance(value, dict):
|
||||
out: Dict[str, Any] = {}
|
||||
for idx, (key, item) in enumerate(value.items()):
|
||||
if idx >= max_sequence:
|
||||
out["_truncated_items"] = len(value) - max_sequence
|
||||
break
|
||||
str_key = str(key)
|
||||
if cls._is_sensitive_hook_key(str_key):
|
||||
out[str_key] = "<redacted>"
|
||||
else:
|
||||
out[str_key] = cls._hook_jsonable(
|
||||
item,
|
||||
depth=depth + 1,
|
||||
max_depth=max_depth,
|
||||
max_string=max_string,
|
||||
max_sequence=max_sequence,
|
||||
)
|
||||
return out
|
||||
if isinstance(value, (list, tuple, set)):
|
||||
seq = list(value)
|
||||
out = [
|
||||
cls._hook_jsonable(
|
||||
item,
|
||||
depth=depth + 1,
|
||||
max_depth=max_depth,
|
||||
max_string=max_string,
|
||||
max_sequence=max_sequence,
|
||||
)
|
||||
for item in seq[:max_sequence]
|
||||
]
|
||||
if len(seq) > max_sequence:
|
||||
out.append({"_truncated_items": len(seq) - max_sequence})
|
||||
return out
|
||||
try:
|
||||
if hasattr(value, "model_dump"):
|
||||
try:
|
||||
# warnings=False: pydantic UserWarnings on generic-union SDK models would leak to the
|
||||
# terminal.
|
||||
dumped = value.model_dump(mode="json", warnings=False)
|
||||
except TypeError:
|
||||
try:
|
||||
dumped = value.model_dump(mode="json")
|
||||
except TypeError:
|
||||
dumped = value.model_dump()
|
||||
return cls._hook_jsonable(
|
||||
dumped,
|
||||
depth=depth + 1,
|
||||
max_depth=max_depth,
|
||||
max_string=max_string,
|
||||
max_sequence=max_sequence,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
from dataclasses import asdict, is_dataclass
|
||||
if is_dataclass(value):
|
||||
return cls._hook_jsonable(
|
||||
asdict(value),
|
||||
depth=depth + 1,
|
||||
max_depth=max_depth,
|
||||
max_string=max_string,
|
||||
max_sequence=max_sequence,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
if isinstance(value, SimpleNamespace):
|
||||
return cls._hook_jsonable(
|
||||
vars(value),
|
||||
depth=depth + 1,
|
||||
max_depth=max_depth,
|
||||
max_string=max_string,
|
||||
max_sequence=max_sequence,
|
||||
)
|
||||
if hasattr(value, "__dict__"):
|
||||
try:
|
||||
public_attrs = {
|
||||
k: v
|
||||
for k, v in vars(value).items()
|
||||
if not str(k).startswith("_")
|
||||
}
|
||||
return cls._hook_jsonable(
|
||||
public_attrs,
|
||||
depth=depth + 1,
|
||||
max_depth=max_depth,
|
||||
max_string=max_string,
|
||||
max_sequence=max_sequence,
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return str(value)[:max_string]
|
||||
|
||||
@classmethod
|
||||
def _sanitize_hook_payload(cls, value: Any) -> Any:
|
||||
payload = cls._hook_jsonable(value)
|
||||
limit = cls._hook_payload_max_chars()
|
||||
try:
|
||||
encoded = json.dumps(payload, ensure_ascii=False, default=str)
|
||||
except Exception:
|
||||
return str(payload)[:limit]
|
||||
if len(encoded) <= limit:
|
||||
return payload
|
||||
payload = cls._hook_jsonable(value, max_string=1000, max_sequence=50)
|
||||
try:
|
||||
encoded = json.dumps(payload, ensure_ascii=False, default=str)
|
||||
except Exception:
|
||||
return str(payload)[:limit]
|
||||
if len(encoded) <= limit:
|
||||
return payload
|
||||
return {
|
||||
"_truncated": True,
|
||||
"original_type": type(value).__name__,
|
||||
"preview": encoded[:limit],
|
||||
}
|
||||
|
||||
def _api_request_payload_for_hook(self, api_kwargs: Optional[Dict[str, Any]]) -> Dict[str, Any]:
|
||||
body = {
|
||||
key: value
|
||||
for key, value in (api_kwargs or {}).items()
|
||||
if key not in {"timeout", "http_client"}
|
||||
}
|
||||
return self._sanitize_hook_payload(
|
||||
{
|
||||
"method": "POST",
|
||||
"body": body,
|
||||
}
|
||||
)
|
||||
|
||||
def _api_response_payload_for_hook(
|
||||
self,
|
||||
response: Any,
|
||||
assistant_message: Any,
|
||||
*,
|
||||
finish_reason: Optional[str],
|
||||
) -> Dict[str, Any]:
|
||||
# Raw provider SDK tool_call objects are handed to the sanitizer on purpose; `_hook_jsonable` must
|
||||
# keep normalising them (model_dump / __dict__ / dataclass) or subscribers get str() blobs.
|
||||
tool_calls = getattr(assistant_message, "tool_calls", None) or []
|
||||
return self._sanitize_hook_payload(
|
||||
{
|
||||
"model": getattr(response, "model", None),
|
||||
"finish_reason": finish_reason,
|
||||
"assistant_message": {
|
||||
"role": getattr(assistant_message, "role", "assistant"),
|
||||
"content": getattr(assistant_message, "content", None),
|
||||
"tool_calls": tool_calls,
|
||||
},
|
||||
"usage": self._usage_summary_for_api_request_hook(response),
|
||||
}
|
||||
)
|
||||
|
||||
def _invoke_api_request_error_hook(
|
||||
self,
|
||||
*,
|
||||
task_id: str,
|
||||
turn_id: str,
|
||||
api_request_id: str,
|
||||
api_call_count: int,
|
||||
api_start_time: float,
|
||||
api_kwargs: Optional[Dict[str, Any]],
|
||||
error_type: str,
|
||||
error_message: str,
|
||||
status_code: Optional[int] = None,
|
||||
retry_count: Optional[int] = None,
|
||||
max_retries: Optional[int] = None,
|
||||
retryable: Optional[bool] = None,
|
||||
reason: Optional[str] = None,
|
||||
) -> None:
|
||||
# Lazy module import (not from-import) so tests can replace lifecycle dispatch at this call site.
|
||||
try:
|
||||
from hermes_cli import lifecycle as _lifecycle
|
||||
|
||||
if not _lifecycle.has_hook("api_request_error"):
|
||||
return
|
||||
ended_at = time.time()
|
||||
_lifecycle.invoke_hook(
|
||||
"api_request_error",
|
||||
task_id=task_id,
|
||||
turn_id=turn_id,
|
||||
api_request_id=api_request_id,
|
||||
session_id=self.session_id or "",
|
||||
platform=self.platform or "",
|
||||
model=self.model,
|
||||
provider=self.provider,
|
||||
base_url=self.base_url,
|
||||
api_mode=self.api_mode,
|
||||
api_call_count=api_call_count,
|
||||
api_duration=ended_at - api_start_time,
|
||||
started_at=api_start_time,
|
||||
ended_at=ended_at,
|
||||
status_code=status_code,
|
||||
retry_count=retry_count,
|
||||
max_retries=max_retries,
|
||||
retryable=retryable,
|
||||
reason=reason,
|
||||
error={
|
||||
"type": error_type,
|
||||
"message": error_message,
|
||||
},
|
||||
request=self._api_request_payload_for_hook(api_kwargs),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
_dump_api_request_debug = _forward("agent.agent_runtime_helpers", "dump_api_request_debug")
|
||||
|
||||
@staticmethod
|
||||
|
||||
Reference in New Issue
Block a user