From 64cc8cf0d35e50a35eee3cac83ee64aa114fc5b8 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 11:52:14 -0700 Subject: [PATCH] refactor(run_agent): extract ApiRequestHooksMixin (agent/api_request_hooks.py) --- agent/api_request_hooks.py | 281 +++++++++++++++++++++++++++++++++++++ run_agent.py | 266 +---------------------------------- 2 files changed, 283 insertions(+), 264 deletions(-) create mode 100644 agent/api_request_hooks.py diff --git a/agent/api_request_hooks.py b/agent/api_request_hooks.py new file mode 100644 index 0000000000..04e8d78977 --- /dev/null +++ b/agent/api_request_hooks.py @@ -0,0 +1,281 @@ +"""Lifecycle-hook payloads for ``AIAgent`` API requests. + +JSON-safe coercion, secret-key redaction, size caps, and the ``api_request_error`` hook dispatch. +Extracted from ``run_agent.py``; every method resolves through ``AIAgent``'s MRO unchanged. +""" +import json +import logging +import os +import time +from types import SimpleNamespace +from typing import Any, Dict, Optional + +from agent.usage_pricing import normalize_usage + +# Same logger name as the origin module so log records / caplog filters are unchanged. +logger = logging.getLogger("run_agent") + + +class ApiRequestHooksMixin: + """Hook payload sanitising + ``api_request_error`` dispatch (see module docstring).""" + + 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] = "" + 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 diff --git a/run_agent.py b/run_agent.py index b596439300..1767af84ea 100644 --- a/run_agent.py +++ b/run_agent.py @@ -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] = "" - 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