From d2f043f9cf7b9e179003ab302fe555e550d1ded4 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 15 Apr 2026 20:54:44 +0800 Subject: [PATCH 1/3] fix(anthropic): preserve third-party thinking continuity Downgrade third-party thinking blocks to text so reasoning context survives across turns while removing redacted payloads and stale signatures. Add regression tests for third-party thinking conversion and keep z.ai preserved-thinking behavior server-driven by removing explicit clear_thinking injection. --- agent/anthropic_adapter.py | 218 +- run_agent.py | 4157 ++++++++++++++++++------- tests/agent/test_anthropic_adapter.py | 433 ++- tests/run_agent/test_run_agent.py | 802 +++-- 4 files changed, 4013 insertions(+), 1597 deletions(-) diff --git a/agent/anthropic_adapter.py b/agent/anthropic_adapter.py index b85f77a9d2396..a3f3d2261bab1 100644 --- a/agent/anthropic_adapter.py +++ b/agent/anthropic_adapter.py @@ -42,26 +42,26 @@ # starves thinking-enabled models (thinking tokens count toward the limit). _ANTHROPIC_OUTPUT_LIMITS = { # Claude 4.6 - "claude-opus-4-6": 128_000, - "claude-sonnet-4-6": 64_000, + "claude-opus-4-6": 128_000, + "claude-sonnet-4-6": 64_000, # Claude 4.5 - "claude-opus-4-5": 64_000, - "claude-sonnet-4-5": 64_000, - "claude-haiku-4-5": 64_000, + "claude-opus-4-5": 64_000, + "claude-sonnet-4-5": 64_000, + "claude-haiku-4-5": 64_000, # Claude 4 - "claude-opus-4": 32_000, - "claude-sonnet-4": 64_000, + "claude-opus-4": 32_000, + "claude-sonnet-4": 64_000, # Claude 3.7 "claude-3-7-sonnet": 128_000, # Claude 3.5 - "claude-3-5-sonnet": 8_192, - "claude-3-5-haiku": 8_192, + "claude-3-5-sonnet": 8_192, + "claude-3-5-haiku": 8_192, # Claude 3 - "claude-3-opus": 4_096, - "claude-3-sonnet": 4_096, - "claude-3-haiku": 4_096, + "claude-3-opus": 4_096, + "claude-3-sonnet": 4_096, + "claude-3-haiku": 4_096, # Third-party Anthropic-compatible providers - "minimax": 131_072, + "minimax": 131_072, } # For any model not in the table, assume the highest current limit. @@ -138,7 +138,9 @@ def _detect_claude_code_version() -> str: try: result = _sp.run( [cmd, "--version"], - capture_output=True, text=True, timeout=5, + capture_output=True, + text=True, + timeout=5, ) if result.returncode == 0 and result.stdout.strip(): # Output is like "2.1.74 (Claude Code)" or just "2.1.74" @@ -224,7 +226,9 @@ def _requires_bearer_auth(base_url: str | None) -> bool: if not normalized: return False normalized = normalized.rstrip("/").lower() - return normalized.startswith(("https://api.minimax.io/anthropic", "https://api.minimaxi.com/anthropic")) + return normalized.startswith( + ("https://api.minimax.io/anthropic", "https://api.minimaxi.com/anthropic") + ) def _common_betas_for_base_url(base_url: str | None) -> list[str]: @@ -357,7 +361,9 @@ def is_claude_code_token_valid(creds: Dict[str, Any]) -> bool: return now_ms < (expires_at - 60_000) -def refresh_anthropic_oauth_pure(refresh_token: str, *, use_json: bool = False) -> Dict[str, Any]: +def refresh_anthropic_oauth_pure( + refresh_token: str, *, use_json: bool = False +) -> Dict[str, Any]: """Refresh an Anthropic OAuth token without mutating local credential files.""" import time import urllib.parse @@ -368,18 +374,22 @@ def refresh_anthropic_oauth_pure(refresh_token: str, *, use_json: bool = False) client_id = "9d1c250a-e61b-44d9-88ed-5944d1962f5e" if use_json: - data = json.dumps({ - "grant_type": "refresh_token", - "refresh_token": refresh_token, - "client_id": client_id, - }).encode() + data = json.dumps( + { + "grant_type": "refresh_token", + "refresh_token": refresh_token, + "client_id": client_id, + } + ).encode() content_type = "application/json" else: - data = urllib.parse.urlencode({ - "grant_type": "refresh_token", - "refresh_token": refresh_token, - "client_id": client_id, - }).encode() + data = urllib.parse.urlencode( + { + "grant_type": "refresh_token", + "refresh_token": refresh_token, + "client_id": client_id, + } + ).encode() content_type = "application/x-www-form-urlencoded" token_endpoints = [ @@ -485,7 +495,9 @@ def _write_claude_code_credentials( logger.debug("Failed to write refreshed credentials: %s", e) -def _resolve_claude_code_token_from_credentials(creds: Optional[Dict[str, Any]] = None) -> Optional[str]: +def _resolve_claude_code_token_from_credentials( + creds: Optional[Dict[str, Any]] = None, +) -> Optional[str]: """Resolve a token from Claude Code credential files, refreshing if needed.""" creds = creds or read_claude_code_credentials() if creds and is_claude_code_token_valid(creds): @@ -496,11 +508,15 @@ def _resolve_claude_code_token_from_credentials(creds: Optional[Dict[str, Any]] refreshed = _refresh_oauth_token(creds) if refreshed: return refreshed - logger.debug("Token refresh failed — re-run 'claude setup-token' to reauthenticate") + logger.debug( + "Token refresh failed — re-run 'claude setup-token' to reauthenticate" + ) return None -def _prefer_refreshable_claude_code_token(env_token: str, creds: Optional[Dict[str, Any]]) -> Optional[str]: +def _prefer_refreshable_claude_code_token( + env_token: str, creds: Optional[Dict[str, Any]] +) -> Optional[str]: """Prefer Claude Code creds when a persisted env OAuth token would shadow refresh. Hermes historically persisted setup tokens into ANTHROPIC_TOKEN. That makes @@ -624,9 +640,11 @@ def _generate_pkce() -> tuple: import secrets verifier = base64.urlsafe_b64encode(secrets.token_bytes(32)).rstrip(b"=").decode() - challenge = base64.urlsafe_b64encode( - hashlib.sha256(verifier.encode()).digest() - ).rstrip(b"=").decode() + challenge = ( + base64.urlsafe_b64encode(hashlib.sha256(verifier.encode()).digest()) + .rstrip(b"=") + .decode() + ) return verifier, challenge @@ -687,14 +705,16 @@ def run_hermes_oauth_login_pure() -> Optional[Dict[str, Any]]: try: import urllib.request - exchange_data = json.dumps({ - "grant_type": "authorization_code", - "client_id": _OAUTH_CLIENT_ID, - "code": code, - "state": state, - "redirect_uri": _OAUTH_REDIRECT_URI, - "code_verifier": verifier, - }).encode() + exchange_data = json.dumps( + { + "grant_type": "authorization_code", + "client_id": _OAUTH_CLIENT_ID, + "code": code, + "state": state, + "redirect_uri": _OAUTH_REDIRECT_URI, + "code_verifier": verifier, + } + ).encode() req = urllib.request.Request( _OAUTH_TOKEN_URL, @@ -755,7 +775,7 @@ def normalize_model_name(model: str, preserve_dots: bool = False) -> str: """ lower = model.lower() if lower.startswith("anthropic/"): - model = model[len("anthropic/"):] + model = model[len("anthropic/") :] if not preserve_dots: # OpenRouter uses dots for version separators (claude-opus-4.6), # Anthropic uses hyphens (claude-opus-4-6). Convert dots to hyphens. @@ -770,6 +790,7 @@ def _sanitize_tool_id(tool_id: str) -> str: characters with underscores and ensure non-empty. """ import re + if not tool_id: return "tool_0" sanitized = re.sub(r"[^a-zA-Z0-9_-]", "_", tool_id) @@ -783,11 +804,15 @@ def convert_tools_to_anthropic(tools: List[Dict]) -> List[Dict]: result = [] for t in tools: fn = t.get("function", {}) - result.append({ - "name": fn.get("name", ""), - "description": fn.get("description", ""), - "input_schema": fn.get("parameters", {"type": "object", "properties": {}}), - }) + result.append( + { + "name": fn.get("name", ""), + "description": fn.get("description", ""), + "input_schema": fn.get( + "parameters", {"type": "object", "properties": {}} + ), + } + ) return result @@ -801,7 +826,7 @@ def _image_source_from_openai_url(url: str) -> Dict[str, str]: header, _, data = url.partition(",") media_type = "image/jpeg" if header.startswith("data:"): - mime_part = header[len("data:"):].split(";", 1)[0].strip() + mime_part = header[len("data:") :].split(";", 1)[0].strip() if mime_part.startswith("image/"): media_type = mime_part return { @@ -828,7 +853,11 @@ def _convert_content_part_to_anthropic(part: Any) -> Optional[Dict[str, Any]]: block: Dict[str, Any] = {"type": "text", "text": part.get("text", "")} elif ptype in {"image_url", "input_image"}: image_value = part.get("image_url", {}) - url = image_value.get("url", "") if isinstance(image_value, dict) else str(image_value or "") + url = ( + image_value.get("url", "") + if isinstance(image_value, dict) + else str(image_value or "") + ) block = {"type": "image", "source": _image_source_from_openai_url(url)} else: block = dict(part) @@ -864,7 +893,10 @@ def _to_plain_data(value: Any, *, _depth: int = 0, _path: Optional[set] = None) return result if isinstance(value, dict): _path.add(obj_id) - result = {k: _to_plain_data(v, _depth=_depth + 1, _path=_path) for k, v in value.items()} + result = { + k: _to_plain_data(v, _depth=_depth + 1, _path=_path) + for k, v in value.items() + } _path.discard(obj_id) return result if isinstance(value, (list, tuple)): @@ -925,9 +957,10 @@ def convert_messages_to_anthropic( system_prompt is a string or list of content blocks (when cache_control present). When *base_url* is provided and points to a third-party Anthropic-compatible - endpoint, all thinking block signatures are stripped. Signatures are - Anthropic-proprietary — third-party endpoints cannot validate them and will - reject them with HTTP 400 "Invalid signature in thinking block". + endpoint, Anthropic thinking signatures are removed. Signed thinking blocks + are downgraded to plain text to preserve useful reasoning context, while + redacted_thinking blocks are dropped. Third-party endpoints cannot validate + Anthropic signatures and may reject them with HTTP 400. """ system = None result = [] @@ -970,12 +1003,14 @@ def convert_messages_to_anthropic( parsed_args = json.loads(args) if isinstance(args, str) else args except (json.JSONDecodeError, ValueError): parsed_args = {} - blocks.append({ - "type": "tool_use", - "id": _sanitize_tool_id(tc.get("id", "")), - "name": fn.get("name", ""), - "input": parsed_args, - }) + blocks.append( + { + "type": "tool_use", + "id": _sanitize_tool_id(tc.get("id", "")), + "name": fn.get("name", ""), + "input": parsed_args, + } + ) # Anthropic rejects empty assistant content effective = blocks or content if not effective or effective == "": @@ -985,7 +1020,9 @@ def convert_messages_to_anthropic( if role == "tool": # Sanitize tool_use_id and ensure non-empty content - result_content = content if isinstance(content, str) else json.dumps(content) + result_content = ( + content if isinstance(content, str) else json.dumps(content) + ) if not result_content: result_content = "(no output)" tool_result = { @@ -1057,7 +1094,8 @@ def convert_messages_to_anthropic( m["content"] = [ b for b in m["content"] - if b.get("type") != "tool_result" or b.get("tool_use_id") in tool_use_ids + if b.get("type") != "tool_result" + or b.get("tool_use_id") in tool_use_ids ] if not m["content"]: m["content"] = [{"type": "text", "text": "(tool result removed)"}] @@ -1088,8 +1126,12 @@ def convert_messages_to_anthropic( # and becomes invalid once merged. if isinstance(m["content"], list): m["content"] = [ - b for b in m["content"] - if not (isinstance(b, dict) and b.get("type") in ("thinking", "redacted_thinking")) + b + for b in m["content"] + if not ( + isinstance(b, dict) + and b.get("type") in ("thinking", "redacted_thinking") + ) ] prev_blocks = fixed[-1]["content"] curr_blocks = m["content"] @@ -1117,9 +1159,8 @@ def convert_messages_to_anthropic( # Signatures are Anthropic-proprietary. Third-party endpoints # (MiniMax, Azure AI Foundry, self-hosted proxies) cannot validate # them and will reject them outright. When targeting a third-party - # endpoint, strip ALL thinking/redacted_thinking blocks from every - # assistant message — the third-party will generate its own - # thinking blocks if it supports extended thinking. + # endpoint, downgrade thinking blocks to plain text and drop + # redacted_thinking blocks. # # For direct Anthropic (strategy following clawdbot/OpenClaw): # 1. Strip thinking/redacted_thinking from all assistant messages @@ -1142,12 +1183,33 @@ def convert_messages_to_anthropic( if m.get("role") != "assistant" or not isinstance(m.get("content"), list): continue - if _is_third_party or idx != last_assistant_idx: - # Third-party endpoint: strip ALL thinking blocks from every - # assistant message — signatures are Anthropic-proprietary. - # Direct Anthropic: strip from non-latest assistant messages only. + if _is_third_party: + # Third-party endpoint: Anthropic signatures are proprietary + # and will be rejected. Downgrade thinking blocks to plain + # text so the model retains reasoning context across turns. + # (Direct Anthropic would validate signatures; third-party + # endpoints like z.ai / GLM-5.1 don't use signatures at all.) + _tp_content = [] + for b in m["content"]: + if not isinstance(b, dict) or b.get("type") not in _THINKING_TYPES: + _tp_content.append(b) + continue + # redacted_thinking carries opaque data — drop it. + if b.get("type") == "redacted_thinking": + continue + # Regular thinking → plain text preserves reasoning for next turn. + thinking_text = b.get("thinking", "") + if thinking_text: + _tp_content.append({"type": "text", "text": thinking_text}) + m["content"] = _tp_content or [ + {"type": "text", "text": "(thinking elided)"} + ] + elif idx != last_assistant_idx: + # Direct Anthropic: strip thinking from non-latest assistant + # messages to avoid stale-signature 400s. stripped = [ - b for b in m["content"] + b + for b in m["content"] if not (isinstance(b, dict) and b.get("type") in _THINKING_TYPES) ] m["content"] = stripped or [{"type": "text", "text": "(thinking elided)"}] @@ -1235,7 +1297,9 @@ def build_anthropic_kwargs( Currently only supported on native Anthropic endpoints (not third-party compatible ones). """ - system, anthropic_messages = convert_messages_to_anthropic(messages, base_url=base_url) + system, anthropic_messages = convert_messages_to_anthropic( + messages, base_url=base_url + ) anthropic_tools = convert_tools_to_anthropic(tools) if tools else [] model = normalize_model_name(model, preserve_dots=preserve_dots) @@ -1287,7 +1351,10 @@ def build_anthropic_kwargs( if block.get("type") == "tool_use" and "name" in block: if not block["name"].startswith(_MCP_TOOL_PREFIX): block["name"] = _MCP_TOOL_PREFIX + block["name"] - elif block.get("type") == "tool_result" and "tool_use_id" in block: + elif ( + block.get("type") == "tool_result" + and "tool_use_id" in block + ): pass # tool_result uses ID, not name kwargs: Dict[str, Any] = { @@ -1319,7 +1386,10 @@ def build_anthropic_kwargs( # MiniMax Anthropic-compat endpoints support thinking (manual mode only, # not adaptive). Haiku does NOT support extended thinking — skip entirely. if reasoning_config and isinstance(reasoning_config, dict): - if reasoning_config.get("enabled") is not False and "haiku" not in model.lower(): + if ( + reasoning_config.get("enabled") is not False + and "haiku" not in model.lower() + ): effort = str(reasoning_config.get("effort", "medium")).lower() budget = THINKING_BUDGET.get(effort, 8000) if _supports_adaptive_thinking(model): @@ -1378,7 +1448,7 @@ def normalize_anthropic_response( elif block.type == "tool_use": name = block.name if strip_tool_prefix and name.startswith(_MCP_TOOL_PREFIX): - name = name[len(_MCP_TOOL_PREFIX):] + name = name[len(_MCP_TOOL_PREFIX) :] tool_calls.append( SimpleNamespace( id=block.id, diff --git a/run_agent.py b/run_agent.py index efaeba82945f7..bc917884e7533 100644 --- a/run_agent.py +++ b/run_agent.py @@ -15,7 +15,7 @@ Usage: from run_agent import AIAgent - + agent = AIAgent(base_url="http://localhost:30000/v1", model="claude-opus-4-20250514") response = agent.run_conversation("Tell me about the latest Python updates") """ @@ -27,6 +27,7 @@ import hashlib import json import logging + logger = logging.getLogger(__name__) import os import random @@ -50,8 +51,10 @@ from hermes_cli.env_loader import load_hermes_dotenv _hermes_home = get_hermes_home() -_project_env = Path(__file__).parent / '.env' -_loaded_env_paths = load_hermes_dotenv(hermes_home=_hermes_home, project_env=_project_env) +_project_env = Path(__file__).parent / ".env" +_loaded_env_paths = load_hermes_dotenv( + hermes_home=_hermes_home, project_env=_project_env +) if _loaded_env_paths: for _env_path in _loaded_env_paths: logger.info("Loaded environment variables from %s", _env_path) @@ -79,37 +82,55 @@ from agent.retry_utils import jittered_backoff from agent.error_classifier import classify_api_error, FailoverReason from agent.prompt_builder import ( - DEFAULT_AGENT_IDENTITY, PLATFORM_HINTS, - MEMORY_GUIDANCE, SESSION_SEARCH_GUIDANCE, SKILLS_GUIDANCE, + DEFAULT_AGENT_IDENTITY, + PLATFORM_HINTS, + MEMORY_GUIDANCE, + SESSION_SEARCH_GUIDANCE, + SKILLS_GUIDANCE, build_nous_subscription_prompt, ) from agent.model_metadata import ( fetch_model_metadata, - estimate_tokens_rough, estimate_messages_tokens_rough, estimate_request_tokens_rough, - get_next_probe_tier, parse_context_limit_from_error, + estimate_tokens_rough, + estimate_messages_tokens_rough, + estimate_request_tokens_rough, + get_next_probe_tier, + parse_context_limit_from_error, parse_available_output_tokens_from_error, - save_context_length, is_local_endpoint, + save_context_length, + is_local_endpoint, query_ollama_num_ctx, ) from agent.context_compressor import ContextCompressor from agent.subdirectory_hints import SubdirectoryHintTracker from agent.prompt_caching import apply_anthropic_cache_control -from agent.prompt_builder import build_skills_system_prompt, build_context_files_prompt, build_environment_hints, load_soul_md, TOOL_USE_ENFORCEMENT_GUIDANCE, TOOL_USE_ENFORCEMENT_MODELS, DEVELOPER_ROLE_MODELS, GOOGLE_MODEL_OPERATIONAL_GUIDANCE, OPENAI_MODEL_EXECUTION_GUIDANCE +from agent.prompt_builder import ( + build_skills_system_prompt, + build_context_files_prompt, + build_environment_hints, + load_soul_md, + TOOL_USE_ENFORCEMENT_GUIDANCE, + TOOL_USE_ENFORCEMENT_MODELS, + DEVELOPER_ROLE_MODELS, + GOOGLE_MODEL_OPERATIONAL_GUIDANCE, + OPENAI_MODEL_EXECUTION_GUIDANCE, +) from agent.usage_pricing import estimate_usage_cost, normalize_usage from agent.display import ( - KawaiiSpinner, build_tool_preview as _build_tool_preview, + KawaiiSpinner, + build_tool_preview as _build_tool_preview, get_cute_tool_message as _get_cute_tool_message_impl, _detect_tool_failure, get_tool_emoji as _get_tool_emoji, ) from agent.trajectory import ( - convert_scratchpad_to_think, has_incomplete_scratchpad, + convert_scratchpad_to_think, + has_incomplete_scratchpad, save_trajectory as _save_trajectory_to_file, ) from utils import atomic_json_write, env_var_enabled - class _SafeWriter: """Transparent stdio wrapper that catches OSError/ValueError from broken pipes. @@ -216,19 +237,21 @@ def remaining(self) -> int: _NEVER_PARALLEL_TOOLS = frozenset({"clarify"}) # Read-only tools with no shared mutable session state. -_PARALLEL_SAFE_TOOLS = frozenset({ - "ha_get_state", - "ha_list_entities", - "ha_list_services", - "read_file", - "search_files", - "session_search", - "skill_view", - "skills_list", - "vision_analyze", - "web_extract", - "web_search", -}) +_PARALLEL_SAFE_TOOLS = frozenset( + { + "ha_get_state", + "ha_list_entities", + "ha_list_services", + "read_file", + "search_files", + "session_search", + "skill_view", + "skills_list", + "vision_analyze", + "web_extract", + "web_search", + } +) # File tools can run concurrently when they target independent paths. _PATH_SCOPED_TOOLS = frozenset({"read_file", "write_file", "patch"}) @@ -250,7 +273,7 @@ def remaining(self) -> int: re.VERBOSE, ) # Output redirects that overwrite files (> but not >>) -_REDIRECT_OVERWRITE = re.compile(r'[^>]>[^>]|^>[^>]') +_REDIRECT_OVERWRITE = re.compile(r"[^>]>[^>]|^>[^>]") def _is_destructive_command(cmd: str) -> bool: @@ -297,7 +320,9 @@ def _should_parallelize_tool_batch(tool_calls) -> bool: scoped_path = _extract_parallel_scope_path(tool_name, function_args) if scoped_path is None: return False - if any(_paths_overlap(scoped_path, existing) for existing in reserved_paths): + if any( + _paths_overlap(scoped_path, existing) for existing in reserved_paths + ): return False reserved_paths.append(scoped_path) continue @@ -336,10 +361,7 @@ def _paths_overlap(left: Path, right: Path) -> bool: return left_parts[:common_len] == right_parts[:common_len] - -_SURROGATE_RE = re.compile(r'[\ud800-\udfff]') - - +_SURROGATE_RE = re.compile(r"[\ud800-\udfff]") def _sanitize_surrogates(text: str) -> str: @@ -349,7 +371,7 @@ def _sanitize_surrogates(text: str) -> str: OpenAI SDK. This is a fast no-op when the text contains no surrogates. """ if _SURROGATE_RE.search(text): - return _SURROGATE_RE.sub('\ufffd', text) + return _SURROGATE_RE.sub("\ufffd", text) return text @@ -366,18 +388,18 @@ def _sanitize_messages_surrogates(messages: list) -> bool: continue content = msg.get("content") if isinstance(content, str) and _SURROGATE_RE.search(content): - msg["content"] = _SURROGATE_RE.sub('\ufffd', content) + msg["content"] = _SURROGATE_RE.sub("\ufffd", content) found = True elif isinstance(content, list): for part in content: if isinstance(part, dict): text = part.get("text") if isinstance(text, str) and _SURROGATE_RE.search(text): - part["text"] = _SURROGATE_RE.sub('\ufffd', text) + part["text"] = _SURROGATE_RE.sub("\ufffd", text) found = True name = msg.get("name") if isinstance(name, str) and _SURROGATE_RE.search(name): - msg["name"] = _SURROGATE_RE.sub('\ufffd', name) + msg["name"] = _SURROGATE_RE.sub("\ufffd", name) found = True tool_calls = msg.get("tool_calls") if isinstance(tool_calls, list): @@ -386,17 +408,17 @@ def _sanitize_messages_surrogates(messages: list) -> bool: continue tc_id = tc.get("id") if isinstance(tc_id, str) and _SURROGATE_RE.search(tc_id): - tc["id"] = _SURROGATE_RE.sub('\ufffd', tc_id) + tc["id"] = _SURROGATE_RE.sub("\ufffd", tc_id) found = True fn = tc.get("function") if isinstance(fn, dict): fn_name = fn.get("name") if isinstance(fn_name, str) and _SURROGATE_RE.search(fn_name): - fn["name"] = _SURROGATE_RE.sub('\ufffd', fn_name) + fn["name"] = _SURROGATE_RE.sub("\ufffd", fn_name) found = True fn_args = fn.get("arguments") if isinstance(fn_args, str) and _SURROGATE_RE.search(fn_args): - fn["arguments"] = _SURROGATE_RE.sub('\ufffd', fn_args) + fn["arguments"] = _SURROGATE_RE.sub("\ufffd", fn_args) found = True return found @@ -407,7 +429,7 @@ def _strip_non_ascii(text: str) -> str: Used as a last resort when the system encoding is ASCII and can't handle any non-ASCII characters (e.g. LANG=C on Chromebooks). """ - return text.encode('ascii', errors='ignore').decode('ascii') + return text.encode("ascii", errors="ignore").decode("ascii") def _sanitize_messages_non_ascii(messages: list) -> bool: @@ -494,9 +516,6 @@ def _walk(node): return found - - - # ========================================================================= # Large tool result handler — save oversized output to temp file # ========================================================================= @@ -663,7 +682,9 @@ def __init__( # instead of going directly to stdout where patch_stdout's StdoutProxy # would mangle the escape sequences. None = use builtins.print. self._print_fn = None - self.background_review_callback = None # Optional sync callback for gateway delivery + self.background_review_callback = ( + None # Optional sync callback for gateway delivery + ) self.skip_context_files = skip_context_files self.pass_session_id = pass_session_id self.persist_session = persist_session @@ -672,7 +693,11 @@ def __init__( self.log_prefix = f"{log_prefix} " if log_prefix else "" # Store effective base URL for feature detection (prompt caching, reasoning, etc.) self.base_url = base_url or "" - provider_name = provider.strip().lower() if isinstance(provider, str) and provider.strip() else None + provider_name = ( + provider.strip().lower() + if isinstance(provider, str) and provider.strip() + else None + ) self.provider = provider_name or "" self.acp_command = acp_command or command self.acp_args = list(acp_args or args or []) @@ -680,10 +705,14 @@ def __init__( self.api_mode = api_mode elif self.provider == "openai-codex": self.api_mode = "codex_responses" - elif (provider_name is None) and "chatgpt.com/backend-api/codex" in self._base_url_lower: + elif ( + provider_name is None + ) and "chatgpt.com/backend-api/codex" in self._base_url_lower: self.api_mode = "codex_responses" self.provider = "openai-codex" - elif self.provider == "anthropic" or (provider_name is None and "api.anthropic.com" in self._base_url_lower): + elif self.provider == "anthropic" or ( + provider_name is None and "api.anthropic.com" in self._base_url_lower + ): self.api_mode = "anthropic_messages" self.provider = "anthropic" elif self._base_url_lower.rstrip("/").endswith("/anthropic"): @@ -745,7 +774,6 @@ def __init__( self.status_callback = status_callback self.tool_gen_callback = tool_gen_callback - # Tool execution state — allows _vprint during tool execution # even when stream consumers are registered (no tokens streaming then) self._executing_tools = False @@ -755,12 +783,12 @@ def __init__( self._interrupt_message = None # Optional message that triggered interrupt self._execution_thread_id: int | None = None # Set at run_conversation() start self._client_lock = threading.RLock() - + # Subagent delegation state - self._delegate_depth = 0 # 0 = top-level agent, incremented for children - self._active_children = [] # Running child AIAgents (for interrupt propagation) + self._delegate_depth = 0 # 0 = top-level agent, incremented for children + self._active_children = [] # Running child AIAgents (for interrupt propagation) self._active_children_lock = threading.Lock() - + # Store OpenRouter provider preferences self.providers_allowed = providers_allowed self.providers_ignored = providers_ignored @@ -772,24 +800,28 @@ def __init__( # Store toolset filtering options self.enabled_toolsets = enabled_toolsets self.disabled_toolsets = disabled_toolsets - + # Model response configuration self.max_tokens = max_tokens # None = use model default - self.reasoning_config = reasoning_config # None = use default (medium for OpenRouter) + self.reasoning_config = ( + reasoning_config # None = use default (medium for OpenRouter) + ) self.service_tier = service_tier self.request_overrides = dict(request_overrides or {}) self.prefill_messages = prefill_messages or [] # Prefilled conversation turns self._force_ascii_payload = False - + # Anthropic prompt caching: auto-enabled for Claude models via OpenRouter. # Reduces input costs by ~75% on multi-turn conversations by caching the # conversation prefix. Uses system_and_3 strategy (4 breakpoints). is_openrouter = self._is_openrouter_url() is_claude = "claude" in self.model.lower() - is_native_anthropic = self.api_mode == "anthropic_messages" and self.provider == "anthropic" + is_native_anthropic = ( + self.api_mode == "anthropic_messages" and self.provider == "anthropic" + ) self._use_prompt_caching = (is_openrouter and is_claude) or is_native_anthropic self._cache_ttl = "5m" # Default 5-minute TTL (1.25x write cost) - + # Iteration budget: the LLM is only notified when it actually exhausts # the iteration budget (api_call_count >= max_iterations). At that # point we inject ONE message, allow one final API call, and if the @@ -822,6 +854,7 @@ def __init__( # both live under ~/.hermes/logs/. Idempotent, so gateway mode # (which creates a new AIAgent per message) won't duplicate handlers. from hermes_logging import setup_logging, setup_verbose_logging + setup_logging(hermes_home=_hermes_home) if self.verbose_logging: @@ -834,14 +867,14 @@ def __init__( # for status; logger INFO/WARNING messages just clutter it. # File handlers (agent.log, errors.log) still capture everything. for quiet_logger in [ - 'tools', # all tools.* (terminal, browser, web, file, etc.) - 'run_agent', # agent runner internals - 'trajectory_compressor', - 'cron', # scheduler (only relevant in daemon mode) - 'hermes_cli', # CLI helpers + "tools", # all tools.* (terminal, browser, web, file, etc.) + "run_agent", # agent runner internals + "trajectory_compressor", + "cron", # scheduler (only relevant in daemon mode) + "hermes_cli", # CLI helpers ]: logging.getLogger(quiet_logger).setLevel(logging.ERROR) - + # Internal stream callback (set during streaming TTS). # Initialized here so _vprint can reference it before run_conversation. self._stream_callback = None @@ -874,23 +907,34 @@ def __init__( self._is_anthropic_oauth = False if self.api_mode == "anthropic_messages": - from agent.anthropic_adapter import build_anthropic_client, resolve_anthropic_token + from agent.anthropic_adapter import ( + build_anthropic_client, + resolve_anthropic_token, + ) + # Only fall back to ANTHROPIC_TOKEN when the provider is actually Anthropic. # Other anthropic_messages providers (MiniMax, Alibaba, etc.) must use their own API key. # Falling back would send Anthropic credentials to third-party endpoints (Fixes #1739, #minimax-401). _is_native_anthropic = self.provider == "anthropic" - effective_key = (api_key or resolve_anthropic_token() or "") if _is_native_anthropic else (api_key or "") + effective_key = ( + (api_key or resolve_anthropic_token() or "") + if _is_native_anthropic + else (api_key or "") + ) self.api_key = effective_key self._anthropic_api_key = effective_key self._anthropic_base_url = base_url from agent.anthropic_adapter import _is_oauth_token as _is_oat + self._is_anthropic_oauth = _is_oat(effective_key) self._anthropic_client = build_anthropic_client(effective_key, base_url) # No OpenAI client needed for Anthropic mode self.client = None self._client_kwargs = {} if not self.quiet_mode: - print(f"🤖 AI Agent initialized with model: {self.model} (Anthropic native)") + print( + f"🤖 AI Agent initialized with model: {self.model} (Anthropic native)" + ) if effective_key and len(effective_key) > 12: print(f"🔑 Using token: {effective_key[:8]}...{effective_key[-4:]}") else: @@ -921,16 +965,23 @@ def __init__( else: # No explicit creds — use the centralized provider router from agent.auxiliary_client import resolve_provider_client + _routed_client, _ = resolve_provider_client( - self.provider or "auto", model=self.model, raw_codex=True) + self.provider or "auto", model=self.model, raw_codex=True + ) if _routed_client is not None: client_kwargs = { "api_key": _routed_client.api_key, "base_url": str(_routed_client.base_url), } # Preserve any default_headers the router set - if hasattr(_routed_client, '_default_headers') and _routed_client._default_headers: - client_kwargs["default_headers"] = dict(_routed_client._default_headers) + if ( + hasattr(_routed_client, "_default_headers") + and _routed_client._default_headers + ): + client_kwargs["default_headers"] = dict( + _routed_client._default_headers + ) else: # When the user explicitly chose a non-OpenRouter provider # but no credentials were found, fail fast with a clear @@ -952,7 +1003,7 @@ def __init__( "X-OpenRouter-Categories": "productivity,cli-agent", }, } - + self._client_kwargs = client_kwargs # stored for rebuilding after interrupt # Enable fine-grained tool streaming for Claude on OpenRouter. @@ -962,7 +1013,10 @@ def __init__( # stream tool call arguments token-by-token, keeping the # connection alive. _effective_base = str(client_kwargs.get("base_url", "")).lower() - if "openrouter" in _effective_base and "claude" in (self.model or "").lower(): + if ( + "openrouter" in _effective_base + and "claude" in (self.model or "").lower() + ): headers = client_kwargs.get("default_headers") or {} existing_beta = headers.get("x-anthropic-beta", "") _FINE_GRAINED = "fine-grained-tool-streaming-2025-05-14" @@ -976,7 +1030,9 @@ def __init__( self.api_key = client_kwargs.get("api_key", "") self.base_url = client_kwargs.get("base_url", self.base_url) try: - self.client = self._create_openai_client(client_kwargs, reason="agent_init", shared=True) + self.client = self._create_openai_client( + client_kwargs, reason="agent_init", shared=True + ) if not self.quiet_mode: print(f"🤖 AI Agent initialized with model: {self.model}") if base_url: @@ -986,20 +1042,27 @@ def __init__( if key_used and key_used != "dummy-key" and len(key_used) > 12: print(f"🔑 Using API key: {key_used[:8]}...{key_used[-4:]}") else: - print(f"⚠️ Warning: API key appears invalid or missing (got: '{key_used[:20] if key_used else 'none'}...')") + print( + f"⚠️ Warning: API key appears invalid or missing (got: '{key_used[:20] if key_used else 'none'}...')" + ) except Exception as e: raise RuntimeError(f"Failed to initialize OpenAI client: {e}") - + # Provider fallback chain — ordered list of backup providers tried # when the primary is exhausted (rate-limit, overload, connection # failure). Supports both legacy single-dict ``fallback_model`` and # new list ``fallback_providers`` format. if isinstance(fallback_model, list): self._fallback_chain = [ - f for f in fallback_model + f + for f in fallback_model if isinstance(f, dict) and f.get("provider") and f.get("model") ] - elif isinstance(fallback_model, dict) and fallback_model.get("provider") and fallback_model.get("model"): + elif ( + isinstance(fallback_model, dict) + and fallback_model.get("provider") + and fallback_model.get("model") + ): self._fallback_chain = [fallback_model] else: self._fallback_chain = [] @@ -1012,8 +1075,12 @@ def __init__( fb = self._fallback_chain[0] print(f"🔄 Fallback model: {fb['model']} ({fb['provider']})") else: - print(f"🔄 Fallback chain ({len(self._fallback_chain)} providers): " + - " → ".join(f"{f['model']} ({f['provider']})" for f in self._fallback_chain)) + print( + f"🔄 Fallback chain ({len(self._fallback_chain)} providers): " + + " → ".join( + f"{f['model']} ({f['provider']})" for f in self._fallback_chain + ) + ) # Get available tools with filtering self.tools = get_tool_definitions( @@ -1021,7 +1088,7 @@ def __init__( disabled_toolsets=disabled_toolsets, quiet_mode=self.quiet_mode, ) - + # Show tool configuration and store valid tool names for validation self.valid_tool_names = set() if self.tools: @@ -1029,7 +1096,7 @@ def __init__( tool_names = sorted(self.valid_tool_names) if not self.quiet_mode: print(f"🛠️ Loaded {len(self.tools)} tools: {', '.join(tool_names)}") - + # Show filtering info if applied if enabled_toolsets: print(f" ✅ Enabled toolsets: {', '.join(enabled_toolsets)}") @@ -1037,28 +1104,40 @@ def __init__( print(f" ❌ Disabled toolsets: {', '.join(disabled_toolsets)}") elif not self.quiet_mode: print("🛠️ No tools loaded (all tools filtered out or unavailable)") - + # Check tool requirements if self.tools and not self.quiet_mode: requirements = check_toolset_requirements() - missing_reqs = [name for name, available in requirements.items() if not available] + missing_reqs = [ + name for name, available in requirements.items() if not available + ] if missing_reqs: - print(f"⚠️ Some tools may not work due to missing requirements: {missing_reqs}") - + print( + f"⚠️ Some tools may not work due to missing requirements: {missing_reqs}" + ) + # Show trajectory saving status if self.save_trajectories and not self.quiet_mode: print("📝 Trajectory saving enabled") - + # Show ephemeral system prompt status if self.ephemeral_system_prompt and not self.quiet_mode: - prompt_preview = self.ephemeral_system_prompt[:60] + "..." if len(self.ephemeral_system_prompt) > 60 else self.ephemeral_system_prompt - print(f"🔒 Ephemeral system prompt: '{prompt_preview}' (not saved to trajectories)") - + prompt_preview = ( + self.ephemeral_system_prompt[:60] + "..." + if len(self.ephemeral_system_prompt) > 60 + else self.ephemeral_system_prompt + ) + print( + f"🔒 Ephemeral system prompt: '{prompt_preview}' (not saved to trajectories)" + ) + # Show prompt caching status if self._use_prompt_caching and not self.quiet_mode: - source = "native Anthropic" if is_native_anthropic else "Claude via OpenRouter" + source = ( + "native Anthropic" if is_native_anthropic else "Claude via OpenRouter" + ) print(f"💾 Prompt caching: ENABLED ({source}, {self._cache_ttl} TTL)") - + # Session logging setup - auto-save conversation trajectories for debugging self.session_start = datetime.now() if session_id: @@ -1069,35 +1148,39 @@ def __init__( timestamp_str = self.session_start.strftime("%Y%m%d_%H%M%S") short_uuid = uuid.uuid4().hex[:6] self.session_id = f"{timestamp_str}_{short_uuid}" - + # Session logs go into ~/.hermes/sessions/ alongside gateway sessions hermes_home = get_hermes_home() self.logs_dir = hermes_home / "sessions" self.logs_dir.mkdir(parents=True, exist_ok=True) self.session_log_file = self.logs_dir / f"session_{self.session_id}.json" - + # Track conversation messages for session logging self._session_messages: List[Dict[str, Any]] = [] - + # Cached system prompt -- built once per session, only rebuilt on compression self._cached_system_prompt: Optional[str] = None - + # Filesystem checkpoint manager (transparent — not a tool) from tools.checkpoint_manager import CheckpointManager + self._checkpoint_mgr = CheckpointManager( enabled=checkpoints_enabled, max_snapshots=checkpoint_max_snapshots, ) - + # SQLite session store (optional -- provided by CLI or gateway) self._session_db = session_db self._parent_session_id = parent_session_id - self._last_flushed_db_idx = 0 # tracks DB-write cursor to prevent duplicate writes + self._last_flushed_db_idx = ( + 0 # tracks DB-write cursor to prevent duplicate writes + ) if self._session_db: try: self._session_db.create_session( session_id=self.session_id, - source=self.platform or os.environ.get("HERMES_SESSION_SOURCE", "cli"), + source=self.platform + or os.environ.get("HERMES_SESSION_SOURCE", "cli"), model=self.model, model_config={ "max_iterations": self.max_iterations, @@ -1115,16 +1198,19 @@ def __init__( # lock clears. The session row may be missing from the index # for this run, but that is recoverable (flushes upsert rows). logger.warning( - "Session DB create_session failed (session_search still available): %s", e + "Session DB create_session failed (session_search still available): %s", + e, ) - + # In-memory todo list for task planning (one per agent/session) from tools.todo_tool import TodoStore + self._todo_store = TodoStore() - + # Load config once for memory, skills, and compression sections try: from hermes_cli.config import load_config as _load_agent_config + _agent_cfg = _load_agent_config() except Exception: _agent_cfg = {} @@ -1141,11 +1227,14 @@ def __init__( try: mem_config = _agent_cfg.get("memory", {}) self._memory_enabled = mem_config.get("memory_enabled", False) - self._user_profile_enabled = mem_config.get("user_profile_enabled", False) + self._user_profile_enabled = mem_config.get( + "user_profile_enabled", False + ) self._memory_nudge_interval = int(mem_config.get("nudge_interval", 10)) self._memory_flush_min_turns = int(mem_config.get("flush_min_turns", 6)) if self._memory_enabled or self._user_profile_enabled: from tools.memory_tool import MemoryStore + self._memory_store = MemoryStore( memory_char_limit=mem_config.get("memory_char_limit", 2200), user_char_limit=mem_config.get("user_char_limit", 1375), @@ -1153,15 +1242,15 @@ def __init__( self._memory_store.load_from_disk() except Exception: pass # Memory is optional -- don't break agent init - - # Memory provider plugin (external — one at a time, alongside built-in) # Reads memory.provider from config to select which plugin to activate. self._memory_manager = None if not skip_memory: try: - _mem_provider_name = mem_config.get("provider", "") if mem_config else "" + _mem_provider_name = ( + mem_config.get("provider", "") if mem_config else "" + ) # Auto-migrate: if Honcho was actively configured (enabled + # credentials) but memory.provider is not set, activate the @@ -1170,20 +1259,29 @@ def __init__( # file may be from a different tool. if not _mem_provider_name: try: - from plugins.memory.honcho.client import HonchoClientConfig as _HCC + from plugins.memory.honcho.client import ( + HonchoClientConfig as _HCC, + ) + _hcfg = _HCC.from_global_config() if _hcfg.enabled and (_hcfg.api_key or _hcfg.base_url): _mem_provider_name = "honcho" # Persist so this only auto-migrates once try: - from hermes_cli.config import load_config as _lc, save_config as _sc + from hermes_cli.config import ( + load_config as _lc, + save_config as _sc, + ) + _cfg = _lc() _cfg.setdefault("memory", {})["provider"] = "honcho" _sc(_cfg) except Exception: pass if not self.quiet_mode: - print(" ✓ Auto-migrated Honcho to memory provider plugin.") + print( + " ✓ Auto-migrated Honcho to memory provider plugin." + ) print(" Your config and data are preserved.\n") except Exception: pass @@ -1191,12 +1289,14 @@ def __init__( if _mem_provider_name: from agent.memory_manager import MemoryManager as _MemoryManager from plugins.memory import load_memory_provider as _load_mem + self._memory_manager = _MemoryManager() _mp = _load_mem(_mem_provider_name) if _mp and _mp.is_available(): self._memory_manager.add_provider(_mp) if self._memory_manager.providers: from hermes_constants import get_hermes_home as _ghh + _init_kwargs = { "session_id": self.session_id, "platform": platform or "cli", @@ -1209,15 +1309,21 @@ def __init__( # Profile identity for per-profile provider scoping try: from hermes_cli.profiles import get_active_profile_name + _profile = get_active_profile_name() _init_kwargs["agent_identity"] = _profile _init_kwargs["agent_workspace"] = "hermes" except Exception: pass self._memory_manager.initialize_all(**_init_kwargs) - logger.info("Memory provider '%s' activated", _mem_provider_name) + logger.info( + "Memory provider '%s' activated", _mem_provider_name + ) else: - logger.debug("Memory provider '%s' not found or not available", _mem_provider_name) + logger.debug( + "Memory provider '%s' not found or not available", + _mem_provider_name, + ) self._memory_manager = None except Exception as _mpe: logger.warning("Memory provider plugin init failed: %s", _mpe) @@ -1236,7 +1342,9 @@ def __init__( self._skill_nudge_interval = 10 try: skills_config = _agent_cfg.get("skills", {}) - self._skill_nudge_interval = int(skills_config.get("creation_nudge_interval", 10)) + self._skill_nudge_interval = int( + skills_config.get("creation_nudge_interval", 10) + ) except Exception: pass @@ -1254,7 +1362,11 @@ def __init__( if not isinstance(_compression_cfg, dict): _compression_cfg = {} compression_threshold = float(_compression_cfg.get("threshold", 0.50)) - compression_enabled = str(_compression_cfg.get("enabled", True)).lower() in ("true", "1", "yes") + compression_enabled = str(_compression_cfg.get("enabled", True)).lower() in ( + "true", + "1", + "yes", + ) compression_target_ratio = float(_compression_cfg.get("target_ratio", 0.20)) compression_protect_last = int(_compression_cfg.get("protect_last_n", 20)) @@ -1275,6 +1387,7 @@ def __init__( _config_context_length, ) import sys + print( f"\n⚠ Invalid model.context_length in config.yaml: {_config_context_length!r}\n" f" Must be a plain integer (e.g. 256000, not '256K').\n" @@ -1290,6 +1403,7 @@ def __init__( if _config_context_length is None: try: from hermes_cli.config import get_compatible_custom_providers + _custom_providers = get_compatible_custom_providers(_agent_cfg) except Exception: _custom_providers = _agent_cfg.get("custom_providers") @@ -1314,9 +1428,11 @@ def __init__( "custom_providers: %r — must be a plain " "integer (e.g. 256000, not '256K'). " "Falling back to auto-detection.", - self.model, _cp_ctx, + self.model, + _cp_ctx, ) import sys + print( f"\n⚠ Invalid context_length for model {self.model!r} in custom_providers: {_cp_ctx!r}\n" f" Must be a plain integer (e.g. 256000, not '256K').\n" @@ -1324,7 +1440,7 @@ def __init__( file=sys.stderr, ) break - + # Select context engine: config-driven (like memory providers). # 1. Check config.yaml context.engine setting # 2. Check plugins/context_engine// directory (repo-shipped) @@ -1333,7 +1449,9 @@ def __init__( _selected_engine = None _engine_name = "compressor" # default try: - _ctx_cfg = _agent_cfg.get("context", {}) if isinstance(_agent_cfg, dict) else {} + _ctx_cfg = ( + _agent_cfg.get("context", {}) if isinstance(_agent_cfg, dict) else {} + ) _engine_name = _ctx_cfg.get("engine", "compressor") or "compressor" except Exception: pass @@ -1342,14 +1460,18 @@ def __init__( # Try loading from plugins/context_engine// try: from plugins.context_engine import load_context_engine + _selected_engine = load_context_engine(_engine_name) except Exception as _ce_load_err: - logger.debug("Context engine load from plugins/context_engine/: %s", _ce_load_err) + logger.debug( + "Context engine load from plugins/context_engine/: %s", _ce_load_err + ) # Try general plugin system as fallback if _selected_engine is None: try: from hermes_cli.plugins import get_plugin_context_engine + _candidate = get_plugin_context_engine() if _candidate and _candidate.name == _engine_name: _selected_engine = _candidate @@ -1367,6 +1489,7 @@ def __init__( self.context_compressor = _selected_engine # Resolve context_length for plugin engines — mirrors switch_model() path from agent.model_metadata import get_model_context_length + _plugin_ctx_len = get_model_context_length( self.model, base_url=self.base_url, @@ -1403,6 +1526,7 @@ def __init__( # Reject models whose context window is below the minimum required # for reliable tool-calling workflows (64K tokens). from agent.model_metadata import MINIMUM_CONTEXT_LENGTH + _ctx = getattr(self.context_compressor, "context_length", 0) if _ctx and _ctx < MINIMUM_CONTEXT_LENGTH: raise ValueError( @@ -1415,7 +1539,11 @@ def __init__( # Inject context engine tool schemas (e.g. lcm_grep, lcm_describe, lcm_expand) self._context_engine_tool_names: set = set() - if hasattr(self, "context_compressor") and self.context_compressor and self.tools is not None: + if ( + hasattr(self, "context_compressor") + and self.context_compressor + and self.tools is not None + ): for _schema in self.context_compressor.get_tool_schemas(): _wrapped = {"type": "function", "function": _schema} self.tools.append(_wrapped) @@ -1432,7 +1560,9 @@ def __init__( hermes_home=str(get_hermes_home()), platform=self.platform or "cli", model=self.model, - context_length=getattr(self.context_compressor, "context_length", 0), + context_length=getattr( + self.context_compressor, "context_length", 0 + ), ) except Exception as _ce_err: logger.debug("Context engine on_session_start: %s", _ce_err) @@ -1455,7 +1585,7 @@ def __init__( self.session_estimated_cost_usd = 0.0 self.session_cost_status = "unknown" self.session_cost_source = "none" - + # ── Ollama num_ctx injection ── # Ollama defaults to 2048 context regardless of the model's capabilities. # When running against an Ollama server, detect the model's max context @@ -1469,8 +1599,14 @@ def __init__( try: self._ollama_num_ctx = int(_ollama_num_ctx_override) except (TypeError, ValueError): - logger.debug("Invalid ollama_num_ctx config value: %r", _ollama_num_ctx_override) - if self._ollama_num_ctx is None and self.base_url and is_local_endpoint(self.base_url): + logger.debug( + "Invalid ollama_num_ctx config value: %r", _ollama_num_ctx_override + ) + if ( + self._ollama_num_ctx is None + and self.base_url + and is_local_endpoint(self.base_url) + ): try: _detected = query_ollama_num_ctx(self.model, self.base_url) if _detected and _detected > 0: @@ -1485,9 +1621,13 @@ def __init__( if not self.quiet_mode: if compression_enabled: - print(f"📊 Context limit: {self.context_compressor.context_length:,} tokens (compress at {int(compression_threshold*100)}% = {self.context_compressor.threshold_tokens:,})") + print( + f"📊 Context limit: {self.context_compressor.context_length:,} tokens (compress at {int(compression_threshold * 100)}% = {self.context_compressor.threshold_tokens:,})" + ) else: - print(f"📊 Context limit: {self.context_compressor.context_length:,} tokens (auto-compression disabled)") + print( + f"📊 Context limit: {self.context_compressor.context_length:,} tokens (auto-compression disabled)" + ) # Check immediately so CLI users see the warning at startup. # Gateway status_callback is not yet wired, so any warning is stored @@ -1519,15 +1659,17 @@ def __init__( "compressor_threshold_tokens": _cc.threshold_tokens, } if self.api_mode == "anthropic_messages": - self._primary_runtime.update({ - "anthropic_api_key": self._anthropic_api_key, - "anthropic_base_url": self._anthropic_base_url, - "is_anthropic_oauth": self._is_anthropic_oauth, - }) + self._primary_runtime.update( + { + "anthropic_api_key": self._anthropic_api_key, + "anthropic_base_url": self._anthropic_base_url, + "is_anthropic_oauth": self._is_anthropic_oauth, + } + ) def reset_session_state(self): """Reset all session-scoped token counters to 0 for a fresh session. - + This method encapsulates the reset logic for all session-level metrics including: - Token usage counters (input, output, total, prompt, completion) @@ -1536,10 +1678,10 @@ def reset_session_state(self): - Reasoning tokens - Estimated cost tracking - Context compressor internal counters - + The method safely handles optional attributes (e.g., context compressor) using ``hasattr`` checks. - + This keeps the counter reset logic DRY and maintainable in one place rather than scattering it across multiple methods. """ @@ -1556,15 +1698,17 @@ def reset_session_state(self): self.session_estimated_cost_usd = 0.0 self.session_cost_status = "unknown" self.session_cost_source = "none" - + # Turn counter (added after reset_session_state was first written — #2635) self._user_turn_count = 0 # Context engine reset (works for both built-in compressor and plugins) if hasattr(self, "context_compressor") and self.context_compressor: self.context_compressor.on_session_reset() - - def switch_model(self, new_model, new_provider, api_key='', base_url='', api_mode=''): + + def switch_model( + self, new_model, new_provider, api_key="", base_url="", api_mode="" + ): """Switch the model/provider in-place for a live agent. Called by the /model command handlers (CLI and gateway) after @@ -1603,16 +1747,24 @@ def switch_model(self, new_model, new_provider, api_key='', base_url='', api_mod resolve_anthropic_token, _is_oauth_token, ) + # Only fall back to ANTHROPIC_TOKEN when the provider is actually Anthropic. # Other anthropic_messages providers (MiniMax, Alibaba, etc.) must use their own # API key — falling back would send Anthropic credentials to third-party endpoints. _is_native_anthropic = new_provider == "anthropic" - effective_key = (api_key or self.api_key or resolve_anthropic_token() or "") if _is_native_anthropic else (api_key or self.api_key or "") + effective_key = ( + (api_key or self.api_key or resolve_anthropic_token() or "") + if _is_native_anthropic + else (api_key or self.api_key or "") + ) self.api_key = effective_key self._anthropic_api_key = effective_key - self._anthropic_base_url = base_url or getattr(self, "_anthropic_base_url", None) + self._anthropic_base_url = base_url or getattr( + self, "_anthropic_base_url", None + ) self._anthropic_client = build_anthropic_client( - effective_key, self._anthropic_base_url, + effective_key, + self._anthropic_base_url, ) self._is_anthropic_oauth = _is_oauth_token(effective_key) self.client = None @@ -1631,15 +1783,18 @@ def switch_model(self, new_model, new_provider, api_key='', base_url='', api_mod ) # ── Re-evaluate prompt caching ── - is_native_anthropic = api_mode == "anthropic_messages" and new_provider == "anthropic" - self._use_prompt_caching = ( - ("openrouter" in (self.base_url or "").lower() and "claude" in new_model.lower()) - or is_native_anthropic + is_native_anthropic = ( + api_mode == "anthropic_messages" and new_provider == "anthropic" ) + self._use_prompt_caching = ( + "openrouter" in (self.base_url or "").lower() + and "claude" in new_model.lower() + ) or is_native_anthropic # ── Update context compressor ── if hasattr(self, "context_compressor") and self.context_compressor: from agent.model_metadata import get_model_context_length + new_context_length = get_model_context_length( self.model, base_url=self.base_url, @@ -1660,7 +1815,11 @@ def switch_model(self, new_model, new_provider, api_key='', base_url='', api_mod self._cached_system_prompt = None # ── Update _primary_runtime so the change persists across turns ── - _cc = self.context_compressor if hasattr(self, "context_compressor") and self.context_compressor else None + _cc = ( + self.context_compressor + if hasattr(self, "context_compressor") and self.context_compressor + else None + ) self._primary_runtime = { "model": self.model, "provider": self.provider, @@ -1669,19 +1828,27 @@ def switch_model(self, new_model, new_provider, api_key='', base_url='', api_mod "api_key": getattr(self, "api_key", ""), "client_kwargs": dict(self._client_kwargs), "use_prompt_caching": self._use_prompt_caching, - "compressor_model": getattr(_cc, "model", self.model) if _cc else self.model, - "compressor_base_url": getattr(_cc, "base_url", self.base_url) if _cc else self.base_url, + "compressor_model": getattr(_cc, "model", self.model) + if _cc + else self.model, + "compressor_base_url": getattr(_cc, "base_url", self.base_url) + if _cc + else self.base_url, "compressor_api_key": getattr(_cc, "api_key", "") if _cc else "", - "compressor_provider": getattr(_cc, "provider", self.provider) if _cc else self.provider, + "compressor_provider": getattr(_cc, "provider", self.provider) + if _cc + else self.provider, "compressor_context_length": _cc.context_length if _cc else 0, "compressor_threshold_tokens": _cc.threshold_tokens if _cc else 0, } if api_mode == "anthropic_messages": - self._primary_runtime.update({ - "anthropic_api_key": self._anthropic_api_key, - "anthropic_base_url": self._anthropic_base_url, - "is_anthropic_oauth": self._is_anthropic_oauth, - }) + self._primary_runtime.update( + { + "anthropic_api_key": self._anthropic_api_key, + "anthropic_base_url": self._anthropic_base_url, + "is_anthropic_oauth": self._is_anthropic_oauth, + } + ) # ── Reset fallback state ── self._fallback_activated = False @@ -1689,7 +1856,10 @@ def switch_model(self, new_model, new_provider, api_key='', base_url='', api_mod logging.info( "Model switched in-place: %s (%s) -> %s (%s)", - old_model, old_provider, new_model, new_provider, + old_model, + old_provider, + new_model, + new_provider, ) def _safe_print(self, *args, **kwargs): @@ -1844,7 +2014,9 @@ def _check_compression_model_feasibility(self) -> None: # ignoring the explicit config value. Pass it as the highest- # priority hint so the configured value is always respected. _aux_cfg = (self.config or {}).get("auxiliary", {}).get("compression", {}) - _aux_context_config = _aux_cfg.get("context_length") if isinstance(_aux_cfg, dict) else None + _aux_context_config = ( + _aux_cfg.get("context_length") if isinstance(_aux_cfg, dict) else None + ) if _aux_context_config is not None: try: _aux_context_config = int(_aux_context_config) @@ -1862,7 +2034,9 @@ def _check_compression_model_feasibility(self) -> None: if aux_context < threshold: # Suggest a threshold that would fit the aux model, # rounded down to a clean percentage. - safe_pct = int((aux_context / self.context_compressor.context_length) * 100) + safe_pct = int( + (aux_context / self.context_compressor.context_length) * 100 + ) msg = ( f"⚠ Compression model ({aux_model}) context " f"is {aux_context:,} tokens, but the main model's " @@ -1892,9 +2066,7 @@ def _check_compression_model_feasibility(self) -> None: threshold, ) except Exception as exc: - logger.debug( - "Compression feasibility check failed (non-fatal): %s", exc - ) + logger.debug("Compression feasibility check failed (non-fatal): %s", exc) def _replay_compression_warning(self) -> None: """Re-send the compression warning through ``status_callback``. @@ -1939,7 +2111,7 @@ def _model_requires_responses_api(model: str) -> bool: def _max_tokens_param(self, value: int) -> dict: """Return the correct max tokens kwarg for the current provider. - + OpenAI's newer models (gpt-4o, o-series, gpt-5+) require 'max_completion_tokens'. OpenRouter, local models, and older OpenAI models use 'max_tokens'. @@ -1970,19 +2142,33 @@ def _has_content_after_think_block(self, content: str) -> bool: # Check if there's any non-whitespace content remaining return bool(cleaned.strip()) - + def _strip_think_blocks(self, content: str) -> str: """Remove reasoning/thinking blocks from content, returning only visible text.""" if not content: return "" # Strip all reasoning tag variants: , , , # , , (Gemma 4) - content = re.sub(r'.*?', '', content, flags=re.DOTALL) - content = re.sub(r'.*?', '', content, flags=re.DOTALL | re.IGNORECASE) - content = re.sub(r'.*?', '', content, flags=re.DOTALL) - content = re.sub(r'.*?', '', content, flags=re.DOTALL) - content = re.sub(r'.*?', '', content, flags=re.DOTALL | re.IGNORECASE) - content = re.sub(r'\s*', '', content, flags=re.IGNORECASE) + content = re.sub(r".*?", "", content, flags=re.DOTALL) + content = re.sub( + r".*?", "", content, flags=re.DOTALL | re.IGNORECASE + ) + content = re.sub(r".*?", "", content, flags=re.DOTALL) + content = re.sub( + r".*?", + "", + content, + flags=re.DOTALL, + ) + content = re.sub( + r".*?", "", content, flags=re.DOTALL | re.IGNORECASE + ) + content = re.sub( + r"\s*", + "", + content, + flags=re.IGNORECASE, + ) return content def _looks_like_codex_intermediate_ack( @@ -1995,14 +2181,19 @@ def _looks_like_codex_intermediate_ack( if any(isinstance(msg, dict) and msg.get("role") == "tool" for msg in messages): return False - assistant_text = self._strip_think_blocks(assistant_content or "").strip().lower() + assistant_text = ( + self._strip_think_blocks(assistant_content or "").strip().lower() + ) if not assistant_text: return False if len(assistant_text) > 1200: return False has_future_ack = bool( - re.search(r"\b(i['’]ll|i will|let me|i can do that|i can help with that)\b", assistant_text) + re.search( + r"\b(i['’]ll|i will|let me|i can do that|i can help with that)\b", + assistant_text, + ) ) if not has_future_ack: return False @@ -2050,51 +2241,60 @@ def _looks_like_codex_intermediate_ack( or "~/" in user_text or "/" in user_text ) - assistant_mentions_action = any(marker in assistant_text for marker in action_markers) + assistant_mentions_action = any( + marker in assistant_text for marker in action_markers + ) assistant_targets_workspace = any( marker in assistant_text for marker in workspace_markers ) - return (user_targets_workspace or assistant_targets_workspace) and assistant_mentions_action - - + return ( + user_targets_workspace or assistant_targets_workspace + ) and assistant_mentions_action + def _extract_reasoning(self, assistant_message) -> Optional[str]: """ Extract reasoning/thinking content from an assistant message. - + OpenRouter and various providers can return reasoning in multiple formats: 1. message.reasoning - Direct reasoning field (DeepSeek, Qwen, etc.) 2. message.reasoning_content - Alternative field (Moonshot AI, Novita, etc.) 3. message.reasoning_details - Array of {type, summary, ...} objects (OpenRouter unified) - + Args: assistant_message: The assistant message object from the API response - + Returns: Combined reasoning text, or None if no reasoning found """ reasoning_parts = [] - + # Check direct reasoning field - if hasattr(assistant_message, 'reasoning') and assistant_message.reasoning: + if hasattr(assistant_message, "reasoning") and assistant_message.reasoning: reasoning_parts.append(assistant_message.reasoning) - + # Check reasoning_content field (alternative name used by some providers) - if hasattr(assistant_message, 'reasoning_content') and assistant_message.reasoning_content: + if ( + hasattr(assistant_message, "reasoning_content") + and assistant_message.reasoning_content + ): # Don't duplicate if same as reasoning if assistant_message.reasoning_content not in reasoning_parts: reasoning_parts.append(assistant_message.reasoning_content) - + # Check reasoning_details array (OpenRouter unified format) # Format: [{"type": "reasoning.summary", "summary": "...", ...}, ...] - if hasattr(assistant_message, 'reasoning_details') and assistant_message.reasoning_details: + if ( + hasattr(assistant_message, "reasoning_details") + and assistant_message.reasoning_details + ): for detail in assistant_message.reasoning_details: if isinstance(detail, dict): # Extract summary from reasoning detail object summary = ( - detail.get('summary') - or detail.get('thinking') - or detail.get('content') - or detail.get('text') + detail.get("summary") + or detail.get("thinking") + or detail.get("content") + or detail.get("text") ) if summary and summary not in reasoning_parts: reasoning_parts.append(summary) @@ -2117,11 +2317,11 @@ def _extract_reasoning(self, assistant_message) -> Optional[str]: cleaned = block.strip() if cleaned and cleaned not in reasoning_parts: reasoning_parts.append(cleaned) - + # Combine all reasoning parts if reasoning_parts: return "\n\n".join(reasoning_parts) - + return None def _cleanup_task_resources(self, task_id: str) -> None: @@ -2217,11 +2417,14 @@ def _spawn_background_review( def _run_review(): import contextlib, os as _os + review_agent = None try: - with open(_os.devnull, "w") as _devnull, \ - contextlib.redirect_stdout(_devnull), \ - contextlib.redirect_stderr(_devnull): + with ( + open(_os.devnull, "w") as _devnull, + contextlib.redirect_stdout(_devnull), + contextlib.redirect_stderr(_devnull), + ): review_agent = AIAgent( model=self.model, max_iterations=8, @@ -2258,14 +2461,34 @@ def _run_review(): actions.append(message) elif "updated" in message.lower(): actions.append(message) - elif "added" in message.lower() or (target and "add" in message.lower()): - label = "Memory" if target == "memory" else "User profile" if target == "user" else target + elif "added" in message.lower() or ( + target and "add" in message.lower() + ): + label = ( + "Memory" + if target == "memory" + else "User profile" + if target == "user" + else target + ) actions.append(f"{label} updated") elif "Entry added" in message: - label = "Memory" if target == "memory" else "User profile" if target == "user" else target + label = ( + "Memory" + if target == "memory" + else "User profile" + if target == "user" + else target + ) actions.append(f"{label} updated") elif "removed" in message.lower() or "replaced" in message.lower(): - label = "Memory" if target == "memory" else "User profile" if target == "user" else target + label = ( + "Memory" + if target == "memory" + else "User profile" + if target == "user" + else target + ) actions.append(f"{label} updated") if actions: @@ -2311,7 +2534,9 @@ def _apply_persist_user_message_override(self, messages: List[Dict]) -> None: if isinstance(msg, dict) and msg.get("role") == "user": msg["content"] = override - def _persist_session(self, messages: List[Dict], conversation_history: List[Dict] = None): + def _persist_session( + self, messages: List[Dict], conversation_history: List[Dict] = None + ): """Save session state to both JSON log and SQLite on any exit path. Ensures conversations are never lost, even on errors or early returns. @@ -2324,7 +2549,9 @@ def _persist_session(self, messages: List[Dict], conversation_history: List[Dict self._save_session_log(messages) self._flush_messages_to_session_db(messages, conversation_history) - def _flush_messages_to_session_db(self, messages: List[Dict], conversation_history: List[Dict] = None): + def _flush_messages_to_session_db( + self, messages: List[Dict], conversation_history: List[Dict] = None + ): """Persist any un-flushed messages to the SQLite session store. Uses _last_flushed_db_idx to track which messages have already been @@ -2365,8 +2592,12 @@ def _flush_messages_to_session_db(self, messages: List[Dict], conversation_histo tool_call_id=msg.get("tool_call_id"), finish_reason=msg.get("finish_reason"), reasoning=msg.get("reasoning") if role == "assistant" else None, - reasoning_details=msg.get("reasoning_details") if role == "assistant" else None, - codex_reasoning_items=msg.get("codex_reasoning_items") if role == "assistant" else None, + reasoning_details=msg.get("reasoning_details") + if role == "assistant" + else None, + codex_reasoning_items=msg.get("codex_reasoning_items") + if role == "assistant" + else None, ) self._last_flushed_db_idx = len(messages) except Exception as e: @@ -2375,44 +2606,44 @@ def _flush_messages_to_session_db(self, messages: List[Dict], conversation_histo def _get_messages_up_to_last_assistant(self, messages: List[Dict]) -> List[Dict]: """ Get messages up to (but not including) the last assistant turn. - + This is used when we need to "roll back" to the last successful point in the conversation, typically when the final assistant message is incomplete or malformed. - + Args: messages: Full message list - + Returns: Messages up to the last complete assistant turn (ending with user/tool message) """ if not messages: return [] - + # Find the index of the last assistant message last_assistant_idx = None for i in range(len(messages) - 1, -1, -1): if messages[i].get("role") == "assistant": last_assistant_idx = i break - + if last_assistant_idx is None: # No assistant message found, return all messages return messages.copy() - + # Return everything up to (not including) the last assistant message return messages[:last_assistant_idx] - + def _format_tools_for_system_message(self) -> str: """ Format tool definitions for the system message in the trajectory format. - + Returns: str: JSON string representation of tool definitions """ if not self.tools: return "[]" - + # Convert tool definitions to the format expected in trajectories formatted_tools = [] for tool in self.tools: @@ -2421,26 +2652,28 @@ def _format_tools_for_system_message(self) -> str: "name": func["name"], "description": func.get("description", ""), "parameters": func.get("parameters", {}), - "required": None # Match the format in the example + "required": None, # Match the format in the example } formatted_tools.append(formatted_tool) - + return json.dumps(formatted_tools, ensure_ascii=False) - - def _convert_to_trajectory_format(self, messages: List[Dict[str, Any]], user_query: str, completed: bool) -> List[Dict[str, Any]]: + + def _convert_to_trajectory_format( + self, messages: List[Dict[str, Any]], user_query: str, completed: bool + ) -> List[Dict[str, Any]]: """ Convert internal message format to trajectory format for saving. - + Args: messages (List[Dict]): Internal message history user_query (str): Original user query completed (bool): Whether the conversation completed successfully - + Returns: List[Dict]: Messages in trajectory format """ trajectory = [] - + # Add system message with tool definitions system_msg = ( "You are a function calling AI model. You are provided with function signatures within XML tags. " @@ -2455,71 +2688,69 @@ def _convert_to_trajectory_format(self, messages: List[Dict[str, Any]], user_que "Each function call should be enclosed within XML tags.\n" "Example:\n\n{'name': ,'arguments': }\n" ) - - trajectory.append({ - "from": "system", - "value": system_msg - }) - + + trajectory.append({"from": "system", "value": system_msg}) + # Add the actual user prompt (from the dataset) as the first human message - trajectory.append({ - "from": "human", - "value": user_query - }) - + trajectory.append({"from": "human", "value": user_query}) + # Skip the first message (the user query) since we already added it above. # Prefill messages are injected at API-call time only (not in the messages # list), so no offset adjustment is needed here. i = 1 - + while i < len(messages): msg = messages[i] - + if msg["role"] == "assistant": # Check if this message has tool calls if "tool_calls" in msg and msg["tool_calls"]: # Format assistant message with tool calls # Add tags around reasoning for trajectory storage content = "" - + # Prepend reasoning in tags if available (native thinking tokens) if msg.get("reasoning") and msg["reasoning"].strip(): content = f"\n{msg['reasoning']}\n\n" - + if msg.get("content") and msg["content"].strip(): # Convert any tags to tags # (used when native thinking is disabled and model reasons via XML) content += convert_scratchpad_to_think(msg["content"]) + "\n" - + # Add tool calls wrapped in XML tags for tool_call in msg["tool_calls"]: - if not tool_call or not isinstance(tool_call, dict): continue + if not tool_call or not isinstance(tool_call, dict): + continue # Parse arguments - should always succeed since we validate during conversation # but keep try-except as safety net try: - arguments = json.loads(tool_call["function"]["arguments"]) if isinstance(tool_call["function"]["arguments"], str) else tool_call["function"]["arguments"] + arguments = ( + json.loads(tool_call["function"]["arguments"]) + if isinstance(tool_call["function"]["arguments"], str) + else tool_call["function"]["arguments"] + ) except json.JSONDecodeError: # This shouldn't happen since we validate and retry during conversation, # but if it does, log warning and use empty dict - logging.warning(f"Unexpected invalid JSON in trajectory conversion: {tool_call['function']['arguments'][:100]}") + logging.warning( + f"Unexpected invalid JSON in trajectory conversion: {tool_call['function']['arguments'][:100]}" + ) arguments = {} - + tool_call_json = { "name": tool_call["function"]["name"], - "arguments": arguments + "arguments": arguments, } content += f"\n{json.dumps(tool_call_json, ensure_ascii=False)}\n\n" - + # Ensure every gpt turn has a block (empty if no reasoning) # so the format is consistent for training data if "" not in content: content = "\n\n" + content - - trajectory.append({ - "from": "gpt", - "value": content.rstrip() - }) - + + trajectory.append({"from": "gpt", "value": content.rstrip()}) + # Collect all subsequent tool responses tool_responses = [] j = i + 1 @@ -2527,7 +2758,7 @@ def _convert_to_trajectory_format(self, messages: List[Dict[str, Any]], user_que tool_msg = messages[j] # Format tool response with XML tags tool_response = "\n" - + # Try to parse tool content as JSON if it looks like JSON tool_content = tool_msg["content"] try: @@ -2535,67 +2766,65 @@ def _convert_to_trajectory_format(self, messages: List[Dict[str, Any]], user_que tool_content = json.loads(tool_content) except (json.JSONDecodeError, AttributeError): pass # Keep as string if not valid JSON - + tool_index = len(tool_responses) tool_name = ( msg["tool_calls"][tool_index]["function"]["name"] if tool_index < len(msg["tool_calls"]) else "unknown" ) - tool_response += json.dumps({ - "tool_call_id": tool_msg.get("tool_call_id", ""), - "name": tool_name, - "content": tool_content - }, ensure_ascii=False) + tool_response += json.dumps( + { + "tool_call_id": tool_msg.get("tool_call_id", ""), + "name": tool_name, + "content": tool_content, + }, + ensure_ascii=False, + ) tool_response += "\n" tool_responses.append(tool_response) j += 1 - + # Add all tool responses as a single message if tool_responses: - trajectory.append({ - "from": "tool", - "value": "\n".join(tool_responses) - }) + trajectory.append( + {"from": "tool", "value": "\n".join(tool_responses)} + ) i = j - 1 # Skip the tool messages we just processed - + else: # Regular assistant message without tool calls # Add tags around reasoning for trajectory storage content = "" - + # Prepend reasoning in tags if available (native thinking tokens) if msg.get("reasoning") and msg["reasoning"].strip(): content = f"\n{msg['reasoning']}\n\n" - + # Convert any tags to tags # (used when native thinking is disabled and model reasons via XML) raw_content = msg["content"] or "" content += convert_scratchpad_to_think(raw_content) - + # Ensure every gpt turn has a block (empty if no reasoning) if "" not in content: content = "\n\n" + content - - trajectory.append({ - "from": "gpt", - "value": content.strip() - }) - + + trajectory.append({"from": "gpt", "value": content.strip()}) + elif msg["role"] == "user": - trajectory.append({ - "from": "human", - "value": msg["content"] - }) - + trajectory.append({"from": "human", "value": msg["content"]}) + i += 1 - + return trajectory - - def _save_trajectory(self, messages: List[Dict[str, Any]], user_query: str, completed: bool): + + def _save_trajectory( + self, messages: List[Dict[str, Any]], user_query: str, completed: bool + ): """ Save conversation trajectory to JSONL file. - + Args: messages (List[Dict]): Complete message history user_query (str): Original user query @@ -2603,10 +2832,10 @@ def _save_trajectory(self, messages: List[Dict[str, Any]], user_query: str, comp """ if not self.save_trajectories: return - + trajectory = self._convert_to_trajectory_format(messages, user_query, completed) _save_trajectory_to_file(trajectory, self.model, completed) - + @staticmethod def _summarize_api_error(error: Exception) -> str: """Extract a human-readable one-liner from an API error. @@ -2616,6 +2845,7 @@ def _summarize_api_error(error: Exception) -> str: str(error) for everything else. """ import re as _re + raw = str(error) # Cloudflare / proxy HTML pages: grab the for a clean summary @@ -2637,7 +2867,11 @@ def _summarize_api_error(error: Exception) -> str: # JSON body errors from OpenAI/Anthropic SDKs body = getattr(error, "body", None) if isinstance(body, dict): - msg = body.get("error", {}).get("message") if isinstance(body.get("error"), dict) else body.get("message") + msg = ( + body.get("error", {}).get("message") + if isinstance(body.get("error"), dict) + else body.get("message") + ) if msg: status_code = getattr(error, "status_code", None) prefix = f"HTTP {status_code}: " if status_code else "" @@ -2658,27 +2892,27 @@ def _mask_api_key_for_logs(self, key: Optional[str]) -> Optional[str]: def _clean_error_message(self, error_msg: str) -> str: """ Clean up error messages for user display, removing HTML content and truncating. - + Args: error_msg: Raw error message from API or exception - + Returns: Clean, user-friendly error message """ if not error_msg: return "Unknown error" - + # Remove HTML content (common with CloudFlare and gateway error pages) - if error_msg.strip().startswith('<!DOCTYPE html') or '<html' in error_msg: + if error_msg.strip().startswith("<!DOCTYPE html") or "<html" in error_msg: return "Service temporarily unavailable (HTML error page returned)" - + # Remove newlines and excessive whitespace - cleaned = ' '.join(error_msg.split()) - + cleaned = " ".join(error_msg.split()) + # Truncate if too long if len(cleaned) > 150: cleaned = cleaned[:150] + "..." - + return cleaned @staticmethod @@ -2730,10 +2964,18 @@ def _extract_api_error_context(error: Exception) -> Dict[str, Any]: if "reset_at" not in context: message = context.get("message") or "" if isinstance(message, str): - delay_match = re.search(r"quotaResetDelay[:\s\"]+(\\d+(?:\\.\\d+)?)(ms|s)", message, re.IGNORECASE) + delay_match = re.search( + r"quotaResetDelay[:\s\"]+(\\d+(?:\\.\\d+)?)(ms|s)", + message, + re.IGNORECASE, + ) if delay_match: value = float(delay_match.group(1)) - seconds = value / 1000.0 if delay_match.group(2).lower() == "ms" else value + seconds = ( + value / 1000.0 + if delay_match.group(2).lower() == "ms" + else value + ) context["reset_at"] = time.time() + seconds else: sec_match = re.search( @@ -2746,7 +2988,9 @@ def _extract_api_error_context(error: Exception) -> Dict[str, Any]: return context - def _usage_summary_for_api_request_hook(self, response: Any) -> Optional[Dict[str, Any]]: + 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 @@ -2819,7 +3063,9 @@ def _dump_api_request_debug( response_obj = getattr(error, "response", None) if response_obj is not None: try: - error_info["response_status"] = getattr(response_obj, "status_code", None) + error_info["response_status"] = getattr( + response_obj, "status_code", None + ) error_info["response_text"] = response_obj.text except Exception as e: logger.debug("Could not extract error response details: %s", e) @@ -2827,21 +3073,29 @@ def _dump_api_request_debug( dump_payload["error"] = error_info timestamp = datetime.now().strftime("%Y%m%d_%H%M%S_%f") - dump_file = self.logs_dir / f"request_dump_{self.session_id}_{timestamp}.json" + dump_file = ( + self.logs_dir / f"request_dump_{self.session_id}_{timestamp}.json" + ) dump_file.write_text( json.dumps(dump_payload, ensure_ascii=False, indent=2, default=str), encoding="utf-8", ) - self._vprint(f"{self.log_prefix}🧾 Request debug dump written to: {dump_file}") + self._vprint( + f"{self.log_prefix}🧾 Request debug dump written to: {dump_file}" + ) if env_var_enabled("HERMES_DUMP_REQUEST_STDOUT"): - print(json.dumps(dump_payload, ensure_ascii=False, indent=2, default=str)) + print( + json.dumps(dump_payload, ensure_ascii=False, indent=2, default=str) + ) return dump_file except Exception as dump_error: if self.verbose_logging: - logging.warning(f"Failed to dump API request debug payload: {dump_error}") + logging.warning( + f"Failed to dump API request debug payload: {dump_error}" + ) return None @staticmethod @@ -2850,8 +3104,8 @@ def _clean_session_content(content: str) -> str: if not content: return content content = convert_scratchpad_to_think(content) - content = re.sub(r'\n+(<think>)', r'\n\1', content) - content = re.sub(r'(</think>)\n+', r'\1\n', content) + content = re.sub(r"\n+(<think>)", r"\n\1", content) + content = re.sub(r"(</think>)\n+", r"\1\n", content) return content.strip() def _save_session_log(self, messages: List[Dict[str, Any]] = None): @@ -2885,12 +3139,17 @@ def _save_session_log(self, messages: List[Dict[str, Any]] = None): # with partial history and would otherwise clobber the full JSON log. if self.session_log_file.exists(): try: - existing = json.loads(self.session_log_file.read_text(encoding="utf-8")) - existing_count = existing.get("message_count", len(existing.get("messages", []))) + existing = json.loads( + self.session_log_file.read_text(encoding="utf-8") + ) + existing_count = existing.get( + "message_count", len(existing.get("messages", [])) + ) if existing_count > len(cleaned): logging.debug( "Skipping session log overwrite: existing has %d messages, current has %d", - existing_count, len(cleaned), + existing_count, + len(cleaned), ) return except Exception: @@ -2919,26 +3178,26 @@ def _save_session_log(self, messages: List[Dict[str, Any]] = None): except Exception as e: if self.verbose_logging: logging.warning(f"Failed to save session log: {e}") - + def interrupt(self, message: str = None) -> None: """ Request the agent to interrupt its current tool-calling loop. - + Call this from another thread (e.g., input handler, message receiver) to gracefully stop the agent and process a new message. - + Also signals long-running tool executions (e.g. terminal commands) to terminate early, so the agent can respond immediately. - + Args: message: Optional new message that triggered the interrupt. If provided, the agent will include this in its response context. - + Example (CLI): # In a separate input thread: if user_typed_something: agent.interrupt(user_input) - + Example (Messaging): # When new message arrives for active session: if session_has_running_agent: @@ -2959,8 +3218,17 @@ def interrupt(self, message: str = None) -> None: except Exception as e: logger.debug("Failed to propagate interrupt to child agent: %s", e) if not self.quiet_mode: - print("\n⚡ Interrupt requested" + (f": '{message[:40]}...'" if message and len(message) > 40 else f": '{message}'" if message else "")) - + print( + "\n⚡ Interrupt requested" + + ( + f": '{message[:40]}...'" + if message and len(message) > 40 + else f": '{message}'" + if message + else "" + ) + ) + def clear_interrupt(self) -> None: """Clear any pending interrupt request and the per-thread tool interrupt signal.""" self._interrupt_requested = False @@ -2985,6 +3253,7 @@ def _capture_rate_limits(self, http_response: Any) -> None: return try: from agent.rate_limit_tracker import parse_rate_limit_headers + state = parse_rate_limit_headers(headers, provider=self.provider) if state is not None: self._rate_limit_state = state @@ -3039,7 +3308,7 @@ def shutdown_memory_provider(self, messages: list = None) -> None: ) except Exception: pass - + def close(self) -> None: """Release all resources held by this agent instance. @@ -3058,6 +3327,7 @@ def close(self) -> None: # 1. Kill background processes for this task try: from tools.process_registry import process_registry + process_registry.kill_all(task_id=task_id) except Exception: pass @@ -3065,6 +3335,7 @@ def close(self) -> None: # 2. Clean terminal sandbox environments try: from tools.terminal_tool import cleanup_vm + cleanup_vm(task_id) except Exception: pass @@ -3072,6 +3343,7 @@ def close(self) -> None: # 3. Clean browser daemon sessions try: from tools.browser_tool import cleanup_browser + cleanup_browser(task_id) except Exception: pass @@ -3101,7 +3373,7 @@ def close(self) -> None: def _hydrate_todo_store(self, history: List[Dict[str, Any]]) -> None: """ Recover todo state from conversation history. - + The gateway creates a fresh AIAgent per message, so the in-memory TodoStore is empty. We scan the history for the most recent todo tool response and replay it to reconstruct the state. @@ -3122,32 +3394,25 @@ def _hydrate_todo_store(self, history: List[Dict[str, Any]]) -> None: break except (json.JSONDecodeError, TypeError): continue - + if last_todo_response: # Replay the items into the store (replace mode) self._todo_store.write(last_todo_response, merge=False) if not self.quiet_mode: - self._vprint(f"{self.log_prefix}📋 Restored {len(last_todo_response)} todo item(s) from history") + self._vprint( + f"{self.log_prefix}📋 Restored {len(last_todo_response)} todo item(s) from history" + ) _set_interrupt(False) - + @property def is_interrupted(self) -> bool: """Check if an interrupt has been requested.""" return self._interrupt_requested + def _build_system_prompt(self, system_message: str = None) -> str: + """ + Assemble the full system prompt from all layers. - - - - - - - - - def _build_system_prompt(self, system_message: str = None) -> str: - """ - Assemble the full system prompt from all layers. - Called once per session (cached on self._cached_system_prompt) and only rebuilt after context compression events. This ensures the system prompt is stable across all turns in a session, maximizing prefix cache hits. @@ -3197,13 +3462,21 @@ def _build_system_prompt(self, system_message: str = None) -> str: if self.valid_tool_names: _enforce = self._tool_use_enforcement _inject = False - if _enforce is True or (isinstance(_enforce, str) and _enforce.lower() in ("true", "always", "yes", "on")): + if _enforce is True or ( + isinstance(_enforce, str) + and _enforce.lower() in ("true", "always", "yes", "on") + ): _inject = True - elif _enforce is False or (isinstance(_enforce, str) and _enforce.lower() in ("false", "never", "no", "off")): + elif _enforce is False or ( + isinstance(_enforce, str) + and _enforce.lower() in ("false", "never", "no", "off") + ): _inject = False elif isinstance(_enforce, list): model_lower = (self.model or "").lower() - _inject = any(p.lower() in model_lower for p in _enforce if isinstance(p, str)) + _inject = any( + p.lower() in model_lower for p in _enforce if isinstance(p, str) + ) else: # "auto" or any unrecognised value — use hardcoded defaults model_lower = (self.model or "").lower() @@ -3247,12 +3520,16 @@ def _build_system_prompt(self, system_message: str = None) -> str: except Exception: pass - has_skills_tools = any(name in self.valid_tool_names for name in ['skills_list', 'skill_view', 'skill_manage']) + has_skills_tools = any( + name in self.valid_tool_names + for name in ["skills_list", "skill_view", "skill_manage"] + ) if has_skills_tools: avail_toolsets = { toolset for toolset in ( - get_toolset_for_tool(tool_name) for tool_name in self.valid_tool_names + get_toolset_for_tool(tool_name) + for tool_name in self.valid_tool_names ) if toolset } @@ -3272,13 +3549,17 @@ def _build_system_prompt(self, system_message: str = None) -> str: # other dev files — inflating token usage by ~10k for no benefit. _context_cwd = os.getenv("TERMINAL_CWD") or None context_files_prompt = build_context_files_prompt( - cwd=_context_cwd, skip_soul=_soul_loaded) + cwd=_context_cwd, skip_soul=_soul_loaded + ) if context_files_prompt: prompt_parts.append(context_files_prompt) from hermes_time import now as _hermes_now + now = _hermes_now() - timestamp_line = f"Conversation started: {now.strftime('%A, %B %d, %Y %I:%M %p')}" + timestamp_line = ( + f"Conversation started: {now.strftime('%A, %B %d, %Y %I:%M %p')}" + ) if self.pass_session_id and self.session_id: timestamp_line += f"\nSession ID: {self.session_id}" if self.model: @@ -3291,7 +3572,9 @@ def _build_system_prompt(self, system_message: str = None) -> str: # of the requested model. Inject explicit model identity into the system prompt # so the agent can correctly report which model it is (workaround for API bug). if self.provider == "alibaba": - _model_short = self.model.split("/")[-1] if "/" in self.model else self.model + _model_short = ( + self.model.split("/")[-1] if "/" in self.model else self.model + ) prompt_parts.append( f"You are powered by the model named {_model_short}. " f"The exact model ID is {self.model}. " @@ -3322,7 +3605,9 @@ def _get_tool_call_id_static(tc) -> str: return tc.get("id", "") or "" return getattr(tc, "id", "") or "" - _VALID_API_ROLES = frozenset({"system", "user", "assistant", "tool", "function", "developer"}) + _VALID_API_ROLES = frozenset( + {"system", "user", "assistant", "tool", "function", "developer"} + ) @staticmethod def _sanitize_api_messages(messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]: @@ -3364,8 +3649,12 @@ def _sanitize_api_messages(messages: List[Dict[str, Any]]) -> List[Dict[str, Any orphaned_results = result_call_ids - surviving_call_ids if orphaned_results: messages = [ - m for m in messages - if not (m.get("role") == "tool" and m.get("tool_call_id") in orphaned_results) + m + for m in messages + if not ( + m.get("role") == "tool" + and m.get("tool_call_id") in orphaned_results + ) ] logger.debug( "Pre-call sanitizer: removed %d orphaned tool result(s)", @@ -3382,11 +3671,13 @@ def _sanitize_api_messages(messages: List[Dict[str, Any]]) -> List[Dict[str, Any for tc in msg.get("tool_calls") or []: cid = AIAgent._get_tool_call_id_static(tc) if cid in missing_results: - patched.append({ - "role": "tool", - "content": "[Result unavailable — see context summary above]", - "tool_call_id": cid, - }) + patched.append( + { + "role": "tool", + "content": "[Result unavailable — see context summary above]", + "tool_call_id": cid, + } + ) messages = patched logger.debug( "Pre-call sanitizer: added %d stub tool result(s)", @@ -3405,8 +3696,11 @@ def _cap_delegate_task_calls(tool_calls: list) -> list: Returns the original list if no truncation was needed. """ from tools.delegate_tool import _get_max_concurrent_children + max_children = _get_max_concurrent_children() - delegate_count = sum(1 for tc in tool_calls if tc.function.name == "delegate_task") + delegate_count = sum( + 1 for tc in tool_calls if tc.function.name == "delegate_task" + ) if delegate_count <= max_children: return tool_calls kept_delegates = 0 @@ -3421,7 +3715,8 @@ def _cap_delegate_task_calls(tool_calls: list) -> list: logger.warning( "Truncated %d excess delegate_task call(s) to enforce " "max_concurrent_children=%d limit", - delegate_count - max_children, max_children, + delegate_count - max_children, + max_children, ) return truncated @@ -3474,7 +3769,7 @@ def _repair_tool_call(self, tool_name: str) -> str | None: def _invalidate_system_prompt(self): """ Invalidate the cached system prompt, forcing a rebuild on the next turn. - + Called after context compression events. Also reloads memory from disk so the rebuilt prompt captures any writes from this session. """ @@ -3482,7 +3777,9 @@ def _invalidate_system_prompt(self): if self._memory_store: self._memory_store.load_from_disk() - def _responses_tools(self, tools: Optional[List[Dict[str, Any]]] = None) -> Optional[List[Dict[str, Any]]]: + def _responses_tools( + self, tools: Optional[List[Dict[str, Any]]] = None + ) -> Optional[List[Dict[str, Any]]]: """Convert chat-completions tool schemas to Responses function-tool schemas.""" source_tools = tools if tools is not None else self.tools if not source_tools: @@ -3494,13 +3791,17 @@ def _responses_tools(self, tools: Optional[List[Dict[str, Any]]] = None) -> Opti name = fn.get("name") if not isinstance(name, str) or not name.strip(): continue - converted.append({ - "type": "function", - "name": name, - "description": fn.get("description", ""), - "strict": False, - "parameters": fn.get("parameters", {"type": "object", "properties": {}}), - }) + converted.append( + { + "type": "function", + "name": name, + "description": fn.get("description", ""), + "strict": False, + "parameters": fn.get( + "parameters", {"type": "object", "properties": {}} + ), + } + ) return converted or None @staticmethod @@ -3512,6 +3813,7 @@ def _deterministic_call_id(fn_name: str, arguments: str, index: int = 0) -> str: make every API call's prefix unique, breaking OpenAI's prompt cache. """ import hashlib + seed = f"{fn_name}:{arguments}:{index}" digest = hashlib.sha256(seed.encode("utf-8", errors="replace")).hexdigest()[:12] return f"call_{digest}" @@ -3548,13 +3850,13 @@ def _derive_responses_function_call_id( if source.startswith("fc_"): return source if source.startswith("call_") and len(source) > len("call_"): - return f"fc_{source[len('call_'):]}" + return f"fc_{source[len('call_') :]}" sanitized = re.sub(r"[^A-Za-z0-9_-]", "", source) if sanitized.startswith("fc_"): return sanitized if sanitized.startswith("call_") and len(sanitized) > len("call_"): - return f"fc_{sanitized[len('call_'):]}" + return f"fc_{sanitized[len('call_') :]}" if sanitized: return f"fc_{sanitized[:48]}" @@ -3562,7 +3864,9 @@ def _derive_responses_function_call_id( digest = hashlib.sha1(seed.encode("utf-8")).hexdigest()[:24] return f"fc_{digest}" - def _chat_messages_to_responses_input(self, messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + def _chat_messages_to_responses_input( + self, messages: List[Dict[str, Any]] + ) -> List[Dict[str, Any]]: """Convert internal chat-style messages to Responses input items.""" items: List[Dict[str, Any]] = [] seen_item_ids: set = set() @@ -3619,8 +3923,8 @@ def _chat_messages_to_responses_input(self, messages: List[Dict[str, Any]]) -> L if not isinstance(fn_name, str) or not fn_name.strip(): continue - embedded_call_id, embedded_response_item_id = self._split_responses_tool_id( - tc.get("id") + embedded_call_id, embedded_response_item_id = ( + self._split_responses_tool_id(tc.get("id")) ) call_id = tc.get("call_id") if not isinstance(call_id, str) or not call_id.strip(): @@ -3631,10 +3935,12 @@ def _chat_messages_to_responses_input(self, messages: List[Dict[str, Any]]) -> L and embedded_response_item_id.startswith("fc_") and len(embedded_response_item_id) > len("fc_") ): - call_id = f"call_{embedded_response_item_id[len('fc_'):]}" + call_id = f"call_{embedded_response_item_id[len('fc_') :]}" else: _raw_args = str(fn.get("arguments", "{}")) - call_id = self._deterministic_call_id(fn_name, _raw_args, len(items)) + call_id = self._deterministic_call_id( + fn_name, _raw_args, len(items) + ) call_id = call_id.strip() arguments = fn.get("arguments", "{}") @@ -3644,12 +3950,14 @@ def _chat_messages_to_responses_input(self, messages: List[Dict[str, Any]]) -> L arguments = str(arguments) arguments = arguments.strip() or "{}" - items.append({ - "type": "function_call", - "call_id": call_id, - "name": fn_name, - "arguments": arguments, - }) + items.append( + { + "type": "function_call", + "call_id": call_id, + "name": fn_name, + "arguments": arguments, + } + ) continue items.append({"role": role, "content": content_text}) @@ -3663,11 +3971,13 @@ def _chat_messages_to_responses_input(self, messages: List[Dict[str, Any]]) -> L call_id = raw_tool_call_id.strip() if not isinstance(call_id, str) or not call_id.strip(): continue - items.append({ - "type": "function_call_output", - "call_id": call_id, - "output": str(msg.get("content", "") or ""), - }) + items.append( + { + "type": "function_call_output", + "call_id": call_id, + "output": str(msg.get("content", "") or ""), + } + ) return items @@ -3686,9 +3996,13 @@ def _preflight_codex_input_items(self, raw_items: Any) -> List[Dict[str, Any]]: call_id = item.get("call_id") name = item.get("name") if not isinstance(call_id, str) or not call_id.strip(): - raise ValueError(f"Codex Responses input[{idx}] function_call is missing call_id.") + raise ValueError( + f"Codex Responses input[{idx}] function_call is missing call_id." + ) if not isinstance(name, str) or not name.strip(): - raise ValueError(f"Codex Responses input[{idx}] function_call is missing name.") + raise ValueError( + f"Codex Responses input[{idx}] function_call is missing name." + ) arguments = item.get("arguments", "{}") if isinstance(arguments, dict): @@ -3710,7 +4024,9 @@ def _preflight_codex_input_items(self, raw_items: Any) -> List[Dict[str, Any]]: if item_type == "function_call_output": call_id = item.get("call_id") if not isinstance(call_id, str) or not call_id.strip(): - raise ValueError(f"Codex Responses input[{idx}] function_call_output is missing call_id.") + raise ValueError( + f"Codex Responses input[{idx}] function_call_output is missing call_id." + ) output = item.get("output", "") if output is None: output = "" @@ -3734,7 +4050,10 @@ def _preflight_codex_input_items(self, raw_items: Any) -> List[Dict[str, Any]]: if item_id in seen_ids: continue seen_ids.add(item_id) - reasoning_item = {"type": "reasoning", "encrypted_content": encrypted} + reasoning_item = { + "type": "reasoning", + "encrypted_content": encrypted, + } # Do NOT include the "id" in the outgoing item — with # store=False (our default) the API tries to resolve the # id server-side and returns 404. The id is still used @@ -3776,11 +4095,15 @@ def _preflight_codex_api_kwargs( required = {"model", "instructions", "input"} missing = [key for key in required if key not in api_kwargs] if missing: - raise ValueError(f"Codex Responses request missing required field(s): {', '.join(sorted(missing))}.") + raise ValueError( + f"Codex Responses request missing required field(s): {', '.join(sorted(missing))}." + ) model = api_kwargs.get("model") if not isinstance(model, str) or not model.strip(): - raise ValueError("Codex Responses request 'model' must be a non-empty string.") + raise ValueError( + "Codex Responses request 'model' must be a non-empty string." + ) model = model.strip() instructions = api_kwargs.get("instructions") @@ -3796,20 +4119,28 @@ def _preflight_codex_api_kwargs( normalized_tools = None if tools is not None: if not isinstance(tools, list): - raise ValueError("Codex Responses request 'tools' must be a list when provided.") + raise ValueError( + "Codex Responses request 'tools' must be a list when provided." + ) normalized_tools = [] for idx, tool in enumerate(tools): if not isinstance(tool, dict): raise ValueError(f"Codex Responses tools[{idx}] must be an object.") if tool.get("type") != "function": - raise ValueError(f"Codex Responses tools[{idx}] has unsupported type {tool.get('type')!r}.") + raise ValueError( + f"Codex Responses tools[{idx}] has unsupported type {tool.get('type')!r}." + ) name = tool.get("name") parameters = tool.get("parameters") if not isinstance(name, str) or not name.strip(): - raise ValueError(f"Codex Responses tools[{idx}] is missing a valid name.") + raise ValueError( + f"Codex Responses tools[{idx}] is missing a valid name." + ) if not isinstance(parameters, dict): - raise ValueError(f"Codex Responses tools[{idx}] is missing valid parameters.") + raise ValueError( + f"Codex Responses tools[{idx}] is missing valid parameters." + ) description = tool.get("description", "") if description is None: @@ -3836,9 +4167,19 @@ def _preflight_codex_api_kwargs( raise ValueError("Codex Responses contract requires 'store' to be false.") allowed_keys = { - "model", "instructions", "input", "tools", "store", - "reasoning", "include", "max_output_tokens", "temperature", - "tool_choice", "parallel_tool_calls", "prompt_cache_key", "service_tier", + "model", + "instructions", + "input", + "tools", + "store", + "reasoning", + "include", + "max_output_tokens", + "temperature", + "tool_choice", + "parallel_tool_calls", + "prompt_cache_key", + "service_tier", } normalized: Dict[str, Any] = { "model": model, @@ -3869,7 +4210,11 @@ def _preflight_codex_api_kwargs( normalized["temperature"] = float(temperature) # Pass through tool_choice, parallel_tool_calls, prompt_cache_key - for passthrough_key in ("tool_choice", "parallel_tool_calls", "prompt_cache_key"): + for passthrough_key in ( + "tool_choice", + "parallel_tool_calls", + "prompt_cache_key", + ): val = api_kwargs.get(passthrough_key) if val is not None: normalized[passthrough_key] = val @@ -3882,7 +4227,9 @@ def _preflight_codex_api_kwargs( normalized["stream"] = True allowed_keys.add("stream") elif "stream" in api_kwargs: - raise ValueError("Codex Responses stream flag is only allowed in fallback streaming requests.") + raise ValueError( + "Codex Responses stream flag is only allowed in fallback streaming requests." + ) unexpected = sorted(key for key in api_kwargs if key not in allowed_keys) if unexpected: @@ -3935,12 +4282,19 @@ def _normalize_codex_response(self, response: Any) -> tuple[Any, str]: if isinstance(out_text, str) and out_text.strip(): logger.debug( "Codex response has empty output but output_text is present (%d chars); " - "synthesizing output item.", len(out_text.strip()), + "synthesizing output item.", + len(out_text.strip()), ) - output = [SimpleNamespace( - type="message", role="assistant", status="completed", - content=[SimpleNamespace(type="output_text", text=out_text.strip())], - )] + output = [ + SimpleNamespace( + type="message", + role="assistant", + status="completed", + content=[ + SimpleNamespace(type="output_text", text=out_text.strip()) + ], + ) + ] response.output = output else: raise RuntimeError("Responses API returned no output items") @@ -3956,14 +4310,22 @@ def _normalize_codex_response(self, response: Any) -> tuple[Any, str]: if isinstance(error_obj, dict): error_msg = error_obj.get("message") or str(error_obj) else: - error_msg = str(error_obj) if error_obj else f"Responses API returned status '{response_status}'" + error_msg = ( + str(error_obj) + if error_obj + else f"Responses API returned status '{response_status}'" + ) raise RuntimeError(error_msg) content_parts: List[str] = [] reasoning_parts: List[str] = [] reasoning_items_raw: List[Dict[str, Any]] = [] tool_calls: List[Any] = [] - has_incomplete_items = response_status in {"queued", "in_progress", "incomplete"} + has_incomplete_items = response_status in { + "queued", + "in_progress", + "incomplete", + } saw_commentary_phase = False saw_final_answer_phase = False @@ -4009,7 +4371,9 @@ def _normalize_codex_response(self, response: Any) -> tuple[Any, str]: for part in summary: text = getattr(part, "text", None) if isinstance(text, str): - raw_summary.append({"type": "summary_text", "text": text}) + raw_summary.append( + {"type": "summary_text", "text": text} + ) raw_item["summary"] = raw_summary reasoning_items_raw.append(raw_item) elif item_type == "function_call": @@ -4022,19 +4386,29 @@ def _normalize_codex_response(self, response: Any) -> tuple[Any, str]: raw_call_id = getattr(item, "call_id", None) raw_item_id = getattr(item, "id", None) embedded_call_id, _ = self._split_responses_tool_id(raw_item_id) - call_id = raw_call_id if isinstance(raw_call_id, str) and raw_call_id.strip() else embedded_call_id + call_id = ( + raw_call_id + if isinstance(raw_call_id, str) and raw_call_id.strip() + else embedded_call_id + ) if not isinstance(call_id, str) or not call_id.strip(): - call_id = self._deterministic_call_id(fn_name, arguments, len(tool_calls)) + call_id = self._deterministic_call_id( + fn_name, arguments, len(tool_calls) + ) call_id = call_id.strip() response_item_id = raw_item_id if isinstance(raw_item_id, str) else None - response_item_id = self._derive_responses_function_call_id(call_id, response_item_id) - tool_calls.append(SimpleNamespace( - id=call_id, - call_id=call_id, - response_item_id=response_item_id, - type="function", - function=SimpleNamespace(name=fn_name, arguments=arguments), - )) + response_item_id = self._derive_responses_function_call_id( + call_id, response_item_id + ) + tool_calls.append( + SimpleNamespace( + id=call_id, + call_id=call_id, + response_item_id=response_item_id, + type="function", + function=SimpleNamespace(name=fn_name, arguments=arguments), + ) + ) elif item_type == "custom_tool_call": fn_name = getattr(item, "name", "") or "" arguments = getattr(item, "input", "{}") @@ -4043,19 +4417,29 @@ def _normalize_codex_response(self, response: Any) -> tuple[Any, str]: raw_call_id = getattr(item, "call_id", None) raw_item_id = getattr(item, "id", None) embedded_call_id, _ = self._split_responses_tool_id(raw_item_id) - call_id = raw_call_id if isinstance(raw_call_id, str) and raw_call_id.strip() else embedded_call_id + call_id = ( + raw_call_id + if isinstance(raw_call_id, str) and raw_call_id.strip() + else embedded_call_id + ) if not isinstance(call_id, str) or not call_id.strip(): - call_id = self._deterministic_call_id(fn_name, arguments, len(tool_calls)) + call_id = self._deterministic_call_id( + fn_name, arguments, len(tool_calls) + ) call_id = call_id.strip() response_item_id = raw_item_id if isinstance(raw_item_id, str) else None - response_item_id = self._derive_responses_function_call_id(call_id, response_item_id) - tool_calls.append(SimpleNamespace( - id=call_id, - call_id=call_id, - response_item_id=response_item_id, - type="function", - function=SimpleNamespace(name=fn_name, arguments=arguments), - )) + response_item_id = self._derive_responses_function_call_id( + call_id, response_item_id + ) + tool_calls.append( + SimpleNamespace( + id=call_id, + call_id=call_id, + response_item_id=response_item_id, + type="function", + function=SimpleNamespace(name=fn_name, arguments=arguments), + ) + ) final_text = "\n".join([p for p in content_parts if p]).strip() if not final_text and hasattr(response, "output_text"): @@ -4074,7 +4458,9 @@ def _normalize_codex_response(self, response: Any) -> tuple[Any, str]: if tool_calls: finish_reason = "tool_calls" - elif has_incomplete_items or (saw_commentary_phase and not saw_final_answer_phase): + elif has_incomplete_items or ( + saw_commentary_phase and not saw_final_answer_phase + ): finish_reason = "incomplete" elif reasoning_items_raw and not final_text: # Response contains only reasoning (encrypted thinking state) with @@ -4138,8 +4524,12 @@ def _is_openai_client_closed(client: Any) -> bool: return bool(getattr(http_client, "is_closed", False)) return False - def _create_openai_client(self, client_kwargs: dict, *, reason: str, shared: bool) -> Any: - if self.provider == "copilot-acp" or str(client_kwargs.get("base_url", "")).startswith("acp://copilot"): + def _create_openai_client( + self, client_kwargs: dict, *, reason: str, shared: bool + ) -> Any: + if self.provider == "copilot-acp" or str( + client_kwargs.get("base_url", "") + ).startswith("acp://copilot"): from agent.copilot_acp_client import CopilotACPClient client = CopilotACPClient(**client_kwargs) @@ -4192,9 +4582,8 @@ def _force_close_tcp_sockets(client: Any) -> int: or [] ) for conn in list(connections): - stream = ( - getattr(conn, "_network_stream", None) - or getattr(conn, "_stream", None) + stream = getattr(conn, "_network_stream", None) or getattr( + conn, "_stream", None ) if stream is None: continue @@ -4246,7 +4635,9 @@ def _replace_primary_openai_client(self, *, reason: str) -> bool: with self._openai_client_lock(): old_client = getattr(self, "client", None) try: - new_client = self._create_openai_client(self._client_kwargs, reason=reason, shared=True) + new_client = self._create_openai_client( + self._client_kwargs, reason=reason, shared=True + ) except Exception as exc: logger.warning( "Failed to rebuild shared OpenAI client (%s) %s error=%s", @@ -4305,9 +4696,8 @@ def _cleanup_dead_connections(self) -> bool: dead_count = 0 for conn in list(connections): # Check for connections that are idle but have closed sockets - stream = ( - getattr(conn, "_network_stream", None) - or getattr(conn, "_stream", None) + stream = getattr(conn, "_network_stream", None) or getattr( + conn, "_stream", None ) if stream is None: continue @@ -4320,6 +4710,7 @@ def _cleanup_dead_connections(self) -> bool: continue # Probe socket health with a non-blocking recv peek import socket as _socket + try: sock.setblocking(False) data = sock.recv(1, _socket.MSG_PEEK | _socket.MSG_DONTWAIT) @@ -4358,11 +4749,15 @@ def _create_request_openai_client(self, *, reason: str) -> Any: def _close_request_openai_client(self, client: Any, *, reason: str) -> None: self._close_openai_client(client, reason=reason, shared=False) - def _run_codex_stream(self, api_kwargs: dict, client: Any = None, on_first_delta: callable = None): + def _run_codex_stream( + self, api_kwargs: dict, client: Any = None, on_first_delta: callable = None + ): """Execute one streaming Responses API request and return the final response.""" import httpx as _httpx - active_client = client or self._ensure_primary_openai_client(reason="codex_stream_direct") + active_client = client or self._ensure_primary_openai_client( + reason="codex_stream_direct" + ) max_stream_retries = 1 has_tool_calls = False first_delta_fired = False @@ -4380,7 +4775,10 @@ def _run_codex_stream(self, api_kwargs: dict, client: Any = None, on_first_delta break event_type = getattr(event, "type", "") # Fire callbacks on text content deltas (suppress during tool calls) - if "output_text.delta" in event_type or event_type == "response.output_text.delta": + if ( + "output_text.delta" in event_type + or event_type == "response.output_text.delta" + ): delta_text = getattr(event, "delta", "") if delta_text: self._codex_streamed_text_parts.append(delta_text) @@ -4412,12 +4810,20 @@ def _run_codex_stream(self, api_kwargs: dict, client: Any = None, on_first_delta # Log non-completed terminal events for diagnostics elif event_type in ("response.incomplete", "response.failed"): resp_obj = getattr(event, "response", None) - status = getattr(resp_obj, "status", None) if resp_obj else None - incomplete_details = getattr(resp_obj, "incomplete_details", None) if resp_obj else None + status = ( + getattr(resp_obj, "status", None) if resp_obj else None + ) + incomplete_details = ( + getattr(resp_obj, "incomplete_details", None) + if resp_obj + else None + ) logger.warning( "Codex Responses stream received terminal event %s " "(status=%s, incomplete_details=%s, streamed_chars=%d). %s", - event_type, status, incomplete_details, + event_type, + status, + incomplete_details, sum(len(p) for p in self._codex_streamed_text_parts), self._client_log_context(), ) @@ -4435,18 +4841,30 @@ def _run_codex_stream(self, api_kwargs: dict, client: Any = None, on_first_delta ) elif self._codex_streamed_text_parts and not has_tool_calls: assembled = "".join(self._codex_streamed_text_parts) - final_response.output = [SimpleNamespace( - type="message", - role="assistant", - status="completed", - content=[SimpleNamespace(type="output_text", text=assembled)], - )] + final_response.output = [ + SimpleNamespace( + type="message", + role="assistant", + status="completed", + content=[ + SimpleNamespace( + type="output_text", text=assembled + ) + ], + ) + ] logger.debug( "Codex stream: synthesized output from %d text deltas (%d chars)", - len(self._codex_streamed_text_parts), len(assembled), + len(self._codex_streamed_text_parts), + len(assembled), ) return final_response - except (_httpx.RemoteProtocolError, _httpx.ReadTimeout, _httpx.ConnectError, ConnectionError) as exc: + except ( + _httpx.RemoteProtocolError, + _httpx.ReadTimeout, + _httpx.ConnectError, + ConnectionError, + ) as exc: if attempt < max_stream_retries: logger.debug( "Codex Responses stream transport failed (attempt %s/%s); retrying. %s error=%s", @@ -4461,7 +4879,9 @@ def _run_codex_stream(self, api_kwargs: dict, client: Any = None, on_first_delta self._client_log_context(), exc, ) - return self._run_codex_create_stream_fallback(api_kwargs, client=active_client) + return self._run_codex_create_stream_fallback( + api_kwargs, client=active_client + ) except RuntimeError as exc: err_text = str(exc) missing_completed = "response.completed" in err_text @@ -4478,15 +4898,21 @@ def _run_codex_stream(self, api_kwargs: dict, client: Any = None, on_first_delta "Responses stream did not emit response.completed; falling back to create(stream=True). %s", self._client_log_context(), ) - return self._run_codex_create_stream_fallback(api_kwargs, client=active_client) + return self._run_codex_create_stream_fallback( + api_kwargs, client=active_client + ) raise def _run_codex_create_stream_fallback(self, api_kwargs: dict, client: Any = None): """Fallback path for stream completion edge cases on Codex-style Responses backends.""" - active_client = client or self._ensure_primary_openai_client(reason="codex_create_stream_fallback") + active_client = client or self._ensure_primary_openai_client( + reason="codex_create_stream_fallback" + ) fallback_kwargs = dict(api_kwargs) fallback_kwargs["stream"] = True - fallback_kwargs = self._preflight_codex_api_kwargs(fallback_kwargs, allow_stream=True) + fallback_kwargs = self._preflight_codex_api_kwargs( + fallback_kwargs, allow_stream=True + ) stream_or_response = active_client.responses.create(**fallback_kwargs) # Compatibility shim for mocks or providers that still return a concrete response. @@ -4519,7 +4945,11 @@ def _run_codex_create_stream_fallback(self, api_kwargs: dict, client: Any = None if delta: collected_text_deltas.append(delta) - if event_type not in {"response.completed", "response.incomplete", "response.failed"}: + if event_type not in { + "response.completed", + "response.incomplete", + "response.failed", + }: continue terminal_response = getattr(event, "response", None) @@ -4537,14 +4967,22 @@ def _run_codex_create_stream_fallback(self, api_kwargs: dict, client: Any = None ) elif collected_text_deltas: assembled = "".join(collected_text_deltas) - terminal_response.output = [SimpleNamespace( - type="message", role="assistant", - status="completed", - content=[SimpleNamespace(type="output_text", text=assembled)], - )] + terminal_response.output = [ + SimpleNamespace( + type="message", + role="assistant", + status="completed", + content=[ + SimpleNamespace( + type="output_text", text=assembled + ) + ], + ) + ] logger.debug( "Codex fallback stream: synthesized from %d deltas (%d chars)", - len(collected_text_deltas), len(assembled), + len(collected_text_deltas), + len(assembled), ) return terminal_response finally: @@ -4557,7 +4995,9 @@ def _run_codex_create_stream_fallback(self, api_kwargs: dict, client: Any = None if terminal_response is not None: return terminal_response - raise RuntimeError("Responses create(stream=True) fallback did not emit a terminal response.") + raise RuntimeError( + "Responses create(stream=True) fallback did not emit a terminal response." + ) def _try_refresh_codex_client_credentials(self, *, force: bool = True) -> bool: if self.api_mode != "codex_responses" or self.provider != "openai-codex": @@ -4596,7 +5036,9 @@ def _try_refresh_nous_client_credentials(self, *, force: bool = True) -> bool: from hermes_cli.auth import resolve_nous_runtime_credentials creds = resolve_nous_runtime_credentials( - min_key_ttl_seconds=max(60, int(os.getenv("HERMES_NOUS_MIN_KEY_TTL_SECONDS", "1800"))), + min_key_ttl_seconds=max( + 60, int(os.getenv("HERMES_NOUS_MIN_KEY_TTL_SECONDS", "1800")) + ), timeout_seconds=float(os.getenv("HERMES_NOUS_TIMEOUT_SECONDS", "15")), force_mint=force, ) @@ -4624,7 +5066,9 @@ def _try_refresh_nous_client_credentials(self, *, force: bool = True) -> bool: return True def _try_refresh_anthropic_client_credentials(self) -> bool: - if self.api_mode != "anthropic_messages" or not hasattr(self, "_anthropic_api_key"): + if self.api_mode != "anthropic_messages" or not hasattr( + self, "_anthropic_api_key" + ): return False # Only refresh credentials for the native Anthropic provider. # Other anthropic_messages providers (MiniMax, Alibaba, etc.) use their own keys. @@ -4632,7 +5076,10 @@ def _try_refresh_anthropic_client_credentials(self) -> bool: return False try: - from agent.anthropic_adapter import resolve_anthropic_token, build_anthropic_client + from agent.anthropic_adapter import ( + resolve_anthropic_token, + build_anthropic_client, + ) new_token = resolve_anthropic_token() except Exception as exc: @@ -4651,14 +5098,19 @@ def _try_refresh_anthropic_client_credentials(self) -> bool: pass try: - self._anthropic_client = build_anthropic_client(new_token, getattr(self, "_anthropic_base_url", None)) + self._anthropic_client = build_anthropic_client( + new_token, getattr(self, "_anthropic_base_url", None) + ) except Exception as exc: - logger.warning("Failed to rebuild Anthropic client after credential refresh: %s", exc) + logger.warning( + "Failed to rebuild Anthropic client after credential refresh: %s", exc + ) return False self._anthropic_api_key = new_token # Update OAuth flag — token type may have changed (API key ↔ OAuth) from agent.anthropic_adapter import _is_oauth_token + self._is_anthropic_oauth = _is_oauth_token(new_token) return True @@ -4680,8 +5132,14 @@ def _apply_client_headers_for_base_url(self, base_url: str) -> None: self._client_kwargs.pop("default_headers", None) def _swap_credential(self, entry) -> None: - runtime_key = getattr(entry, "runtime_api_key", None) or getattr(entry, "access_token", "") - runtime_base = getattr(entry, "runtime_base_url", None) or getattr(entry, "base_url", None) or self.base_url + runtime_key = getattr(entry, "runtime_api_key", None) or getattr( + entry, "access_token", "" + ) + runtime_base = ( + getattr(entry, "runtime_base_url", None) + or getattr(entry, "base_url", None) + or self.base_url + ) if self.api_mode == "anthropic_messages": from agent.anthropic_adapter import build_anthropic_client, _is_oauth_token @@ -4700,7 +5158,9 @@ def _swap_credential(self, entry) -> None: return self.api_key = runtime_key - self.base_url = runtime_base.rstrip("/") if isinstance(runtime_base, str) else runtime_base + self.base_url = ( + runtime_base.rstrip("/") if isinstance(runtime_base, str) else runtime_base + ) self._client_kwargs["api_key"] = self.api_key self._client_kwargs["base_url"] = self.base_url self._apply_client_headers_for_base_url(self.base_url) @@ -4743,7 +5203,9 @@ def _recover_with_credential_pool( if effective_reason == FailoverReason.billing: rotate_status = status_code if status_code is not None else 402 - next_entry = pool.mark_exhausted_and_rotate(status_code=rotate_status, error_context=error_context) + next_entry = pool.mark_exhausted_and_rotate( + status_code=rotate_status, error_context=error_context + ) if next_entry is not None: logger.info( "Credential %s (billing) — rotated to pool entry %s", @@ -4758,7 +5220,9 @@ def _recover_with_credential_pool( if not has_retried_429: return False, True rotate_status = status_code if status_code is not None else 429 - next_entry = pool.mark_exhausted_and_rotate(status_code=rotate_status, error_context=error_context) + next_entry = pool.mark_exhausted_and_rotate( + status_code=rotate_status, error_context=error_context + ) if next_entry is not None: logger.info( "Credential %s (rate limit) — rotated to pool entry %s", @@ -4772,13 +5236,17 @@ def _recover_with_credential_pool( if effective_reason == FailoverReason.auth: refreshed = pool.try_refresh_current() if refreshed is not None: - logger.info(f"Credential auth failure — refreshed pool entry {getattr(refreshed, 'id', '?')}") + logger.info( + f"Credential auth failure — refreshed pool entry {getattr(refreshed, 'id', '?')}" + ) self._swap_credential(refreshed) return True, has_retried_429 # Refresh failed — rotate to next credential instead of giving up. # The failed entry is already marked exhausted by try_refresh_current(). rotate_status = status_code if status_code is not None else 401 - next_entry = pool.mark_exhausted_and_rotate(status_code=rotate_status, error_context=error_context) + next_entry = pool.mark_exhausted_and_rotate( + status_code=rotate_status, error_context=error_context + ) if next_entry is not None: logger.info( "Credential %s (auth refresh failed) — rotated to pool entry %s", @@ -4815,7 +5283,11 @@ def _interruptible_api_call(self, api_kwargs: dict): def _call(): try: if self.api_mode == "codex_responses": - request_client_holder["client"] = self._create_request_openai_client(reason="codex_stream_request") + request_client_holder["client"] = ( + self._create_request_openai_client( + reason="codex_stream_request" + ) + ) result["response"] = self._run_codex_stream( api_kwargs, client=request_client_holder["client"], @@ -4824,14 +5296,22 @@ def _call(): elif self.api_mode == "anthropic_messages": result["response"] = self._anthropic_messages_create(api_kwargs) else: - request_client_holder["client"] = self._create_request_openai_client(reason="chat_completion_request") - result["response"] = request_client_holder["client"].chat.completions.create(**api_kwargs) + request_client_holder["client"] = ( + self._create_request_openai_client( + reason="chat_completion_request" + ) + ) + result["response"] = request_client_holder[ + "client" + ].chat.completions.create(**api_kwargs) except Exception as e: result["error"] = e finally: request_client = request_client_holder.get("client") if request_client is not None: - self._close_request_openai_client(request_client, reason="request_complete") + self._close_request_openai_client( + request_client, reason="request_complete" + ) # ── Stale-call timeout (mirrors streaming stale detector) ──────── # Non-streaming calls return nothing until the full response is @@ -4878,8 +5358,10 @@ def _call(): logger.warning( "Non-streaming API call stale for %.0fs (threshold %.0fs). " "model=%s context=~%s tokens. Killing connection.", - _elapsed, _stale_timeout, - api_kwargs.get("model", "unknown"), f"{_est_ctx:,}", + _elapsed, + _stale_timeout, + api_kwargs.get("model", "unknown"), + f"{_est_ctx:,}", ) self._emit_status( f"⚠️ No response from provider for {int(_elapsed)}s " @@ -4898,7 +5380,9 @@ def _call(): else: rc = request_client_holder.get("client") if rc is not None: - self._close_request_openai_client(rc, reason="stale_call_kill") + self._close_request_openai_client( + rc, reason="stale_call_kill" + ) except Exception: pass self._touch_activity( @@ -4929,7 +5413,9 @@ def _call(): else: request_client = request_client_holder.get("client") if request_client is not None: - self._close_request_openai_client(request_client, reason="interrupt_abort") + self._close_request_openai_client( + request_client, reason="interrupt_abort" + ) except Exception: pass raise InterruptedError("Agent interrupted during API call") @@ -4963,7 +5449,9 @@ def _interim_content_was_streamed(self, content: str) -> bool: if not visible_content: return False streamed = self._normalize_interim_visible_text( - self._strip_think_blocks(getattr(self, "_current_streamed_assistant_text", "") or "") + self._strip_think_blocks( + getattr(self, "_current_streamed_assistant_text", "") or "" + ) ) return bool(streamed) and streamed == visible_content @@ -4991,7 +5479,11 @@ def _fire_stream_delta(self, text: str) -> None: if getattr(self, "_stream_needs_break", False) and text and text.strip(): self._stream_needs_break = False text = "\n\n" + text - callbacks = [cb for cb in (self.stream_delta_callback, self._stream_callback) if cb is not None] + callbacks = [ + cb + for cb in (self.stream_delta_callback, self._stream_callback) + if cb is not None + ] delivered = False for cb in callbacks: try: @@ -5065,7 +5557,9 @@ def _interruptible_streaming_api_call( result = {"response": None, "error": None} request_client_holder = {"client": None} first_delta_fired = {"done": False} - deltas_were_sent = {"yes": False} # Track if any deltas were fired (for fallback) + deltas_were_sent = { + "yes": False + } # Track if any deltas were fired (for fallback) # Wall-clock timestamp of the last real streaming chunk. The outer # poll loop uses this to detect stale connections that keep receiving # SSE keep-alive pings but no actual data. @@ -5082,17 +5576,23 @@ def _fire_first_delta(): def _call_chat_completions(): """Stream a chat completions response.""" import httpx as _httpx + _base_timeout = float(os.getenv("HERMES_API_TIMEOUT", 1800.0)) _stream_read_timeout = float(os.getenv("HERMES_STREAM_READ_TIMEOUT", 120.0)) # Local providers (Ollama, llama.cpp, vLLM) can take minutes for # prefill on large contexts before producing the first token. # Auto-increase the httpx read timeout unless the user explicitly # overrode HERMES_STREAM_READ_TIMEOUT. - if _stream_read_timeout == 120.0 and self.base_url and is_local_endpoint(self.base_url): + if ( + _stream_read_timeout == 120.0 + and self.base_url + and is_local_endpoint(self.base_url) + ): _stream_read_timeout = _base_timeout logger.debug( "Local provider detected (%s) — stream read timeout raised to %.0fs", - self.base_url, _stream_read_timeout, + self.base_url, + _stream_read_timeout, ) stream_kwargs = { **api_kwargs, @@ -5112,7 +5612,9 @@ def _call_chat_completions(): # attempt's start, not a previous attempt's last chunk. last_chunk_time["t"] = time.time() self._touch_activity("waiting for provider response (streaming)") - stream = request_client_holder["client"].chat.completions.create(**stream_kwargs) + stream = request_client_holder["client"].chat.completions.create( + **stream_kwargs + ) # Capture rate limit headers from the initial HTTP response. # The OpenAI SDK Stream object exposes the underlying httpx @@ -5126,7 +5628,7 @@ def _call_chat_completions(): # in a parallel batch, distinguishing them only by id. Track # the last seen id per raw index so we can detect a new tool # call starting at the same index and redirect it to a fresh slot. - _last_id_at_idx: dict = {} # raw_index -> last seen non-empty id + _last_id_at_idx: dict = {} # raw_index -> last seen non-empty id _active_slot_by_idx: dict = {} # raw_index -> current slot in tool_calls_acc finish_reason = None model_name = None @@ -5153,7 +5655,9 @@ def _call_chat_completions(): model_name = chunk.model # Accumulate reasoning content - reasoning_text = getattr(delta, "reasoning_content", None) or getattr(delta, "reasoning", None) + reasoning_text = getattr(delta, "reasoning_content", None) or getattr( + delta, "reasoning", None + ) if reasoning_text: reasoning_parts.append(reasoning_text) _fire_first_delta() @@ -5220,7 +5724,9 @@ def _call_chat_completions(): if tc_delta.function.name: entry["function"]["name"] += tc_delta.function.name if tc_delta.function.arguments: - entry["function"]["arguments"] += tc_delta.function.arguments + entry["function"]["arguments"] += ( + tc_delta.function.arguments + ) extra = getattr(tc_delta, "extra_content", None) if extra is None and hasattr(tc_delta, "model_extra"): extra = (tc_delta.model_extra or {}).get("extra_content") @@ -5256,15 +5762,17 @@ def _call_chat_completions(): json.loads(arguments) except json.JSONDecodeError: has_truncated_tool_args = True - mock_tool_calls.append(SimpleNamespace( - id=tc["id"], - type=tc["type"], - extra_content=tc.get("extra_content"), - function=SimpleNamespace( - name=tc["function"]["name"], - arguments=arguments, - ), - )) + mock_tool_calls.append( + SimpleNamespace( + id=tc["id"], + type=tc["type"], + extra_content=tc.get("extra_content"), + function=SimpleNamespace( + name=tc["function"]["name"], + arguments=arguments, + ), + ) + ) effective_finish_reason = finish_reason or "stop" if has_truncated_tool_args: @@ -5366,16 +5874,27 @@ def _call(): # delivered. Don't retry or fall back — partial # content already reached the user. logger.warning( - "Streaming failed after partial delivery, not retrying: %s", e + "Streaming failed after partial delivery, not retrying: %s", + e, ) result["error"] = e return _is_timeout = isinstance( - e, (_httpx.ReadTimeout, _httpx.ConnectTimeout, _httpx.PoolTimeout) + e, + ( + _httpx.ReadTimeout, + _httpx.ConnectTimeout, + _httpx.PoolTimeout, + ), ) _is_conn_err = isinstance( - e, (_httpx.ConnectError, _httpx.RemoteProtocolError, ConnectionError) + e, + ( + _httpx.ConnectError, + _httpx.RemoteProtocolError, + ConnectionError, + ), ) # SSE error events from proxies (e.g. OpenRouter sends @@ -5389,7 +5908,10 @@ def _call(): _is_sse_conn_err = False if not _is_timeout and not _is_conn_err: from openai import APIError as _APIError - if isinstance(e, _APIError) and not getattr(e, "status_code", None): + + if isinstance(e, _APIError) and not getattr( + e, "status_code", None + ): _err_lower_sse = str(e).lower() _SSE_CONN_PHRASES = ( "connection lost", @@ -5459,8 +5981,7 @@ def _call(): else: _err_lower = str(e).lower() _is_stream_unsupported = ( - "stream" in _err_lower - and "not supported" in _err_lower + "stream" in _err_lower and "not supported" in _err_lower ) if _is_stream_unsupported: self._disable_streaming = True @@ -5485,15 +6006,26 @@ def _call(): finally: request_client = request_client_holder.get("client") if request_client is not None: - self._close_request_openai_client(request_client, reason="stream_request_complete") + self._close_request_openai_client( + request_client, reason="stream_request_complete" + ) - _stream_stale_timeout_base = float(os.getenv("HERMES_STREAM_STALE_TIMEOUT", 180.0)) + _stream_stale_timeout_base = float( + os.getenv("HERMES_STREAM_STALE_TIMEOUT", 180.0) + ) # Local providers (Ollama, oMLX, llama-cpp) can take 300+ seconds # for prefill on large contexts. Disable the stale detector unless # the user explicitly set HERMES_STREAM_STALE_TIMEOUT. - if _stream_stale_timeout_base == 180.0 and self.base_url and is_local_endpoint(self.base_url): + if ( + _stream_stale_timeout_base == 180.0 + and self.base_url + and is_local_endpoint(self.base_url) + ): _stream_stale_timeout = float("inf") - logger.debug("Local provider detected (%s) — stale stream timeout disabled", self.base_url) + logger.debug( + "Local provider detected (%s) — stale stream timeout disabled", + self.base_url, + ) else: # Scale the stale timeout for large contexts: slow models (like Opus) # can legitimately think for minutes before producing the first token @@ -5522,8 +6054,10 @@ def _call(): logger.warning( "Stream stale for %.0fs (threshold %.0fs) — no chunks received. " "model=%s context=~%s tokens. Killing connection.", - _stale_elapsed, _stream_stale_timeout, - api_kwargs.get("model", "unknown"), f"{_est_ctx:,}", + _stale_elapsed, + _stream_stale_timeout, + api_kwargs.get("model", "unknown"), + f"{_est_ctx:,}", ) self._emit_status( f"⚠️ No response from provider for {int(_stale_elapsed)}s " @@ -5534,13 +6068,17 @@ def _call(): try: rc = request_client_holder.get("client") if rc is not None: - self._close_request_openai_client(rc, reason="stale_stream_kill") + self._close_request_openai_client( + rc, reason="stale_stream_kill" + ) except Exception: pass # Rebuild the primary client too — its connection pool # may hold dead sockets from the same provider outage. try: - self._replace_primary_openai_client(reason="stale_stream_pool_cleanup") + self._replace_primary_openai_client( + reason="stale_stream_pool_cleanup" + ) except Exception: pass # Reset the timer so we don't kill repeatedly while @@ -5563,7 +6101,9 @@ def _call(): else: request_client = request_client_holder.get("client") if request_client is not None: - self._close_request_openai_client(request_client, reason="stream_interrupt_abort") + self._close_request_openai_client( + request_client, reason="stream_interrupt_abort" + ) except Exception: pass raise InterruptedError("Agent interrupted during streaming API call") @@ -5589,15 +6129,21 @@ def _call(): result["error"], ) _stub_msg = SimpleNamespace( - role="assistant", content=_partial_text, tool_calls=None, + role="assistant", + content=_partial_text, + tool_calls=None, reasoning_content=None, ) return SimpleNamespace( id="partial-stream-stub", model=getattr(self, "model", "unknown"), - choices=[SimpleNamespace( - index=0, message=_stub_msg, finish_reason="stop", - )], + choices=[ + SimpleNamespace( + index=0, + message=_stub_msg, + finish_reason="stop", + ) + ], usage=None, ) raise result["error"] @@ -5632,6 +6178,7 @@ def _try_activate_fallback(self) -> bool: # access for Codex providers. try: from agent.auxiliary_client import resolve_provider_client + # Pass base_url and api_key from fallback config so custom # endpoints (e.g. Ollama Cloud) resolve correctly instead of # falling through to OpenRouter defaults. @@ -5639,16 +6186,23 @@ def _try_activate_fallback(self) -> bool: fb_api_key_hint = (fb.get("api_key") or "").strip() or None # For Ollama Cloud endpoints, pull OLLAMA_API_KEY from env # when no explicit key is in the fallback config. - if fb_base_url_hint and "ollama.com" in fb_base_url_hint.lower() and not fb_api_key_hint: + if ( + fb_base_url_hint + and "ollama.com" in fb_base_url_hint.lower() + and not fb_api_key_hint + ): fb_api_key_hint = os.getenv("OLLAMA_API_KEY") or None fb_client, _resolved_fb_model = resolve_provider_client( - fb_provider, model=fb_model, raw_codex=True, + fb_provider, + model=fb_model, + raw_codex=True, explicit_base_url=fb_base_url_hint, - explicit_api_key=fb_api_key_hint) + explicit_api_key=fb_api_key_hint, + ) if fb_client is None: logging.warning( - "Fallback to %s failed: provider not configured", - fb_provider) + "Fallback to %s failed: provider not configured", fb_provider + ) return self._try_activate_fallback() # try next in chain try: from hermes_cli.model_normalize import normalize_model_for_provider @@ -5662,7 +6216,9 @@ def _try_activate_fallback(self) -> bool: fb_base_url = str(fb_client.base_url) if fb_provider == "openai-codex": fb_api_mode = "codex_responses" - elif fb_provider == "anthropic" or fb_base_url.rstrip("/").lower().endswith("/anthropic"): + elif fb_provider == "anthropic" or fb_base_url.rstrip("/").lower().endswith( + "/anthropic" + ): fb_api_mode = "anthropic_messages" elif self._is_direct_openai_url(fb_base_url): fb_api_mode = "codex_responses" @@ -5680,12 +6236,23 @@ def _try_activate_fallback(self) -> bool: if fb_api_mode == "anthropic_messages": # Build native Anthropic client instead of using OpenAI client - from agent.anthropic_adapter import build_anthropic_client, resolve_anthropic_token, _is_oauth_token - effective_key = (fb_client.api_key or resolve_anthropic_token() or "") if fb_provider == "anthropic" else (fb_client.api_key or "") + from agent.anthropic_adapter import ( + build_anthropic_client, + resolve_anthropic_token, + _is_oauth_token, + ) + + effective_key = ( + (fb_client.api_key or resolve_anthropic_token() or "") + if fb_provider == "anthropic" + else (fb_client.api_key or "") + ) self.api_key = effective_key self._anthropic_api_key = effective_key self._anthropic_base_url = fb_base_url - self._anthropic_client = build_anthropic_client(effective_key, self._anthropic_base_url) + self._anthropic_client = build_anthropic_client( + effective_key, self._anthropic_base_url + ) self._is_anthropic_oauth = _is_oauth_token(effective_key) self.client = None self._client_kwargs = {} @@ -5711,21 +6278,25 @@ def _try_activate_fallback(self) -> bool: } # Re-evaluate prompt caching for the new provider/model - is_native_anthropic = fb_api_mode == "anthropic_messages" and fb_provider == "anthropic" - self._use_prompt_caching = ( - ("openrouter" in fb_base_url.lower() and "claude" in fb_model.lower()) - or is_native_anthropic + is_native_anthropic = ( + fb_api_mode == "anthropic_messages" and fb_provider == "anthropic" ) + self._use_prompt_caching = ( + "openrouter" in fb_base_url.lower() and "claude" in fb_model.lower() + ) or is_native_anthropic # Update context compressor limits for the fallback model. # Without this, compression decisions use the primary model's # context window (e.g. 200K) instead of the fallback's (e.g. 32K), # causing oversized sessions to overflow the fallback. - if hasattr(self, 'context_compressor') and self.context_compressor: + if hasattr(self, "context_compressor") and self.context_compressor: from agent.model_metadata import get_model_context_length + fb_context_length = get_model_context_length( - self.model, base_url=self.base_url, - api_key=self.api_key, provider=self.provider, + self.model, + base_url=self.base_url, + api_key=self.api_key, + provider=self.provider, ) self.context_compressor.update_model( model=self.model, @@ -5741,7 +6312,9 @@ def _try_activate_fallback(self) -> bool: ) logging.info( "Fallback activated: %s → %s (%s)", - old_model, fb_model, fb_provider, + old_model, + fb_model, + fb_provider, ) return True except Exception as e: @@ -5769,7 +6342,7 @@ def _restore_primary_runtime(self) -> bool: # ── Core runtime state ── self.model = rt["model"] self.provider = rt["provider"] - self.base_url = rt["base_url"] # setter updates _base_url_lower + self.base_url = rt["base_url"] # setter updates _base_url_lower self.api_mode = rt["api_mode"] self.api_key = rt["api_key"] self._client_kwargs = dict(rt["client_kwargs"]) @@ -5778,10 +6351,12 @@ def _restore_primary_runtime(self) -> bool: # ── Rebuild client for the primary provider ── if self.api_mode == "anthropic_messages": from agent.anthropic_adapter import build_anthropic_client + self._anthropic_api_key = rt["anthropic_api_key"] self._anthropic_base_url = rt["anthropic_base_url"] self._anthropic_client = build_anthropic_client( - rt["anthropic_api_key"], rt["anthropic_base_url"], + rt["anthropic_api_key"], + rt["anthropic_base_url"], ) self._is_anthropic_oauth = rt["is_anthropic_oauth"] self.client = None @@ -5808,7 +6383,8 @@ def _restore_primary_runtime(self) -> bool: logging.info( "Primary runtime restored for new turn: %s (%s)", - self.model, self.provider, + self.model, + self.provider, ) return True except Exception as e: @@ -5817,14 +6393,24 @@ def _restore_primary_runtime(self) -> bool: # Which error types indicate a transient transport failure worth # one more attempt with a rebuilt client / connection pool. - _TRANSIENT_TRANSPORT_ERRORS = frozenset({ - "ReadTimeout", "ConnectTimeout", "PoolTimeout", - "ConnectError", "RemoteProtocolError", - "APIConnectionError", "APITimeoutError", - }) + _TRANSIENT_TRANSPORT_ERRORS = frozenset( + { + "ReadTimeout", + "ConnectTimeout", + "PoolTimeout", + "ConnectError", + "RemoteProtocolError", + "APIConnectionError", + "APITimeoutError", + } + ) def _try_recover_primary_transport( - self, api_error: Exception, *, retry_count: int, max_retries: int, + self, + api_error: Exception, + *, + retry_count: int, + max_retries: int, ) -> bool: """Attempt one extra primary-provider recovery cycle for transient transport failures. @@ -5858,7 +6444,9 @@ def _try_recover_primary_transport( if getattr(self, "client", None) is not None: try: self._close_openai_client( - self.client, reason="primary_recovery", shared=True, + self.client, + reason="primary_recovery", + shared=True, ) except Exception: pass @@ -5874,10 +6462,12 @@ def _try_recover_primary_transport( if self.api_mode == "anthropic_messages": from agent.anthropic_adapter import build_anthropic_client + self._anthropic_api_key = rt["anthropic_api_key"] self._anthropic_base_url = rt["anthropic_base_url"] self._anthropic_client = build_anthropic_client( - rt["anthropic_api_key"], rt["anthropic_base_url"], + rt["anthropic_api_key"], + rt["anthropic_base_url"], ) self._is_anthropic_oauth = rt["is_anthropic_oauth"] self.client = None @@ -5907,7 +6497,10 @@ def _content_has_image_parts(content: Any) -> bool: if not isinstance(content, list): return False for part in content: - if isinstance(part, dict) and part.get("type") in {"image_url", "input_image"}: + if isinstance(part, dict) and part.get("type") in { + "image_url", + "input_image", + }: return True return False @@ -5916,7 +6509,7 @@ def _materialize_data_url_for_vision(image_url: str) -> tuple[str, Optional[Path header, _, data = str(image_url or "").partition(",") mime = "image/jpeg" if header.startswith("data:"): - mime_part = header[len("data:"):].split(";", 1)[0].strip() + mime_part = header[len("data:") :].split(";", 1)[0].strip() if mime_part.startswith("image/"): mime = mime_part suffix = { @@ -5926,7 +6519,9 @@ def _materialize_data_url_for_vision(image_url: str) -> tuple[str, Optional[Path "image/jpeg": ".jpg", "image/jpg": ".jpg", }.get(mime, ".jpg") - tmp = tempfile.NamedTemporaryFile(prefix="anthropic_image_", suffix=suffix, delete=False) + tmp = tempfile.NamedTemporaryFile( + prefix="anthropic_image_", suffix=suffix, delete=False + ) with tmp: tmp.write(base64.b64decode(data)) path = Path(tmp.name) @@ -5951,14 +6546,18 @@ def _describe_image_for_anthropic_fallback(self, image_url: str, role: str) -> s vision_source = str(image_url or "") cleanup_path: Optional[Path] = None if vision_source.startswith("data:"): - vision_source, cleanup_path = self._materialize_data_url_for_vision(vision_source) + vision_source, cleanup_path = self._materialize_data_url_for_vision( + vision_source + ) description = "" try: from tools.vision_tools import vision_analyze_tool result_json = asyncio.run( - vision_analyze_tool(image_url=vision_source, user_prompt=analysis_prompt) + vision_analyze_tool( + image_url=vision_source, user_prompt=analysis_prompt + ) ) result = json.loads(result_json) if isinstance(result_json, str) else {} description = (result.get("analysis") or "").strip() @@ -5976,9 +6575,7 @@ def _describe_image_for_anthropic_fallback(self, image_url: str, role: str) -> s note = f"[The {role_label} attached an image. Here's what it contains:\n{description}]" if vision_source and not str(image_url or "").startswith("data:"): - note += ( - f"\n[If you need a closer look, use vision_analyze with image_url: {vision_source}]" - ) + note += f"\n[If you need a closer look, use vision_analyze with image_url: {vision_source}]" self._anthropic_image_fallback_cache[cache_key] = note return note @@ -6006,11 +6603,19 @@ def _preprocess_anthropic_content(self, content: Any, role: str) -> Any: if ptype in {"image_url", "input_image"}: image_data = part.get("image_url", {}) - image_url = image_data.get("url", "") if isinstance(image_data, dict) else str(image_data or "") + image_url = ( + image_data.get("url", "") + if isinstance(image_data, dict) + else str(image_data or "") + ) if image_url: - image_notes.append(self._describe_image_for_anthropic_fallback(image_url, role)) + image_notes.append( + self._describe_image_for_anthropic_fallback(image_url, role) + ) else: - image_notes.append("[An image was attached but no image source was available.]") + image_notes.append( + "[An image was attached but no image source was available.]" + ) continue text = str(part.get("text", "") or "").strip() @@ -6025,7 +6630,9 @@ def _preprocess_anthropic_content(self, content: Any, role: str) -> Any: return prefix if suffix: return suffix - return "[A multimodal message was converted to text for Anthropic compatibility.]" + return ( + "[A multimodal message was converted to text for Anthropic compatibility.]" + ) def _prepare_anthropic_messages_for_api(self, api_messages: list) -> list: if not any( @@ -6050,10 +6657,23 @@ def _anthropic_preserve_dots(self) -> bool: MiniMax keeps dots (e.g. MiniMax-M2.7). OpenCode Go/Zen keeps dots for non-Claude models (e.g. minimax-m2.5-free). ZAI/Zhipu keeps dots (e.g. glm-4.7, glm-5.1).""" - if (getattr(self, "provider", "") or "").lower() in {"alibaba", "minimax", "minimax-cn", "opencode-go", "opencode-zen", "zai"}: + if (getattr(self, "provider", "") or "").lower() in { + "alibaba", + "minimax", + "minimax-cn", + "opencode-go", + "opencode-zen", + "zai", + }: return True base = (getattr(self, "base_url", "") or "").lower() - return "dashscope" in base or "aliyuncs" in base or "minimax" in base or "opencode.ai/zen/" in base or "bigmodel.cn" in base + return ( + "dashscope" in base + or "aliyuncs" in base + or "minimax" in base + or "opencode.ai/zen/" in base + or "bigmodel.cn" in base + ) def _is_qwen_portal(self) -> bool: """Return True when the base URL targets Qwen Portal.""" @@ -6086,7 +6706,11 @@ def _qwen_prepare_chat_messages(self, api_messages: list) -> list: for msg in prepared: if isinstance(msg, dict) and msg.get("role") == "system": content = msg.get("content") - if isinstance(content, list) and content and isinstance(content[-1], dict): + if ( + isinstance(content, list) + and content + and isinstance(content[-1], dict) + ): content[-1]["cache_control"] = {"type": "ephemeral"} break @@ -6116,7 +6740,11 @@ def _qwen_prepare_chat_messages_inplace(self, messages: list) -> None: for msg in messages: if isinstance(msg, dict) and msg.get("role") == "system": content = msg.get("content") - if isinstance(content, list) and content and isinstance(content[-1], dict): + if ( + isinstance(content, list) + and content + and isinstance(content[-1], dict) + ): content[-1]["cache_control"] = {"type": "ephemeral"} break @@ -6124,6 +6752,7 @@ def _build_api_kwargs(self, api_messages: list) -> dict: """Build the keyword arguments dict for the active API mode.""" if self.api_mode == "anthropic_messages": from agent.anthropic_adapter import build_anthropic_kwargs + anthropic_messages = self._prepare_anthropic_messages_for_api(api_messages) # Pass context_length (total input+output window) so the adapter can # clamp max_tokens (output cap) when the user configured a smaller @@ -6140,7 +6769,9 @@ def _build_api_kwargs(self, api_messages: list) -> dict: model=self.model, messages=anthropic_messages, tools=self.tools, - max_tokens=ephemeral_out if ephemeral_out is not None else self.max_tokens, + max_tokens=ephemeral_out + if ephemeral_out is not None + else self.max_tokens, reasoning_config=self.reasoning_config, is_oauth=self._is_anthropic_oauth, preserve_dots=self._anthropic_preserve_dots(), @@ -6204,7 +6835,10 @@ def _build_api_kwargs(self, api_messages: list) -> dict: if github_reasoning is not None: kwargs["reasoning"] = github_reasoning else: - kwargs["reasoning"] = {"effort": reasoning_effort, "summary": "auto"} + kwargs["reasoning"] = { + "effort": reasoning_effort, + "summary": "auto", + } kwargs["include"] = ["reasoning.encrypted_content"] elif not is_github_responses: kwargs["include"] = [] @@ -6259,7 +6893,9 @@ def _build_api_kwargs(self, api_messages: list) -> dict: if self._is_qwen_portal(): if sanitized_messages is api_messages: # No sanitization was done — we need our own copy. - sanitized_messages = self._qwen_prepare_chat_messages(sanitized_messages) + sanitized_messages = self._qwen_prepare_chat_messages( + sanitized_messages + ) else: # Already a deepcopy — transform in place to avoid a second deepcopy. self._qwen_prepare_chat_messages_inplace(sanitized_messages) @@ -6315,7 +6951,9 @@ def _build_api_kwargs(self, api_messages: list) -> dict: # (the documented max output for qwen3-coder models) so the # model has adequate output budget for tool calls. api_kwargs.update(self._max_tokens_param(65536)) - elif (self._is_openrouter_url() or "nousresearch" in self._base_url_lower) and "claude" in (self.model or "").lower(): + elif ( + self._is_openrouter_url() or "nousresearch" in self._base_url_lower + ) and "claude" in (self.model or "").lower(): # OpenRouter and Nous Portal translate requests to Anthropic's # Messages API, which requires max_tokens as a mandatory field. # When we omit it, the proxy picks a default that can be too @@ -6325,6 +6963,7 @@ def _build_api_kwargs(self, api_messages: list) -> dict: # limit ensures full capacity. try: from agent.anthropic_adapter import _get_anthropic_max_output + _model_output_limit = _get_anthropic_max_output(self.model) api_kwargs["max_tokens"] = _model_output_limit except Exception: @@ -6361,10 +7000,7 @@ def _build_api_kwargs(self, api_messages: list) -> dict: else: extra_body["reasoning"] = rc else: - extra_body["reasoning"] = { - "enabled": True, - "effort": "medium" - } + extra_body["reasoning"] = {"enabled": True, "effort": "medium"} # Nous Portal product attribution if _is_nous: @@ -6388,7 +7024,11 @@ def _build_api_kwargs(self, api_messages: list) -> dict: # xAI prompt caching: send x-grok-conv-id header to route requests # to the same server, maximizing automatic cache hits. # https://docs.x.ai/developers/advanced-api-usage/prompt-caching - if "x.ai" in self._base_url_lower and hasattr(self, "session_id") and self.session_id: + if ( + "x.ai" in self._base_url_lower + and hasattr(self, "session_id") + and self.session_id + ): api_kwargs["extra_headers"] = {"x-grok-conv-id": self.session_id} # Priority Processing / generic request overrides (e.g. service_tier). @@ -6409,7 +7049,10 @@ def _supports_reasoning_extra_body(self) -> bool: return True if "ai-gateway.vercel.sh" in self._base_url_lower: return True - if "models.github.ai" in self._base_url_lower or "api.githubcopilot.com" in self._base_url_lower: + if ( + "models.github.ai" in self._base_url_lower + or "api.githubcopilot.com" in self._base_url_lower + ): try: from hermes_cli.models import github_model_reasoning_efforts @@ -6446,9 +7089,9 @@ def _github_models_reasoning_extra_body(self) -> dict | None: if self.reasoning_config and isinstance(self.reasoning_config, dict): if self.reasoning_config.get("enabled") is False: return None - requested_effort = str( - self.reasoning_config.get("effort", "medium") - ).strip().lower() + requested_effort = ( + str(self.reasoning_config.get("effort", "medium")).strip().lower() + ) else: requested_effort = "medium" @@ -6478,13 +7121,15 @@ def _build_assistant_message(self, assistant_message, finish_reason: str) -> dic # directly in the content rather than returning separate API fields). if not reasoning_text: content = assistant_message.content or "" - think_blocks = re.findall(r'<think>(.*?)</think>', content, flags=re.DOTALL) + think_blocks = re.findall(r"<think>(.*?)</think>", content, flags=re.DOTALL) if think_blocks: combined = "\n\n".join(b.strip() for b in think_blocks if b.strip()) reasoning_text = combined or None if reasoning_text and self.verbose_logging: - logging.debug(f"Captured reasoning ({len(reasoning_text)} chars): {reasoning_text}") + logging.debug( + f"Captured reasoning ({len(reasoning_text)} chars): {reasoning_text}" + ) if reasoning_text and self.reasoning_callback: # Skip callback when streaming is active — reasoning was already @@ -6508,7 +7153,10 @@ def _build_assistant_message(self, assistant_message, finish_reason: str) -> dic "finish_reason": finish_reason, } - if hasattr(assistant_message, 'reasoning_details') and assistant_message.reasoning_details: + if ( + hasattr(assistant_message, "reasoning_details") + and assistant_message.reasoning_details + ): # Pass reasoning_details back unmodified so providers (OpenRouter, # Anthropic, OpenAI) can maintain reasoning continuity across turns. # Each provider may include opaque fields (signature, encrypted_content) @@ -6546,11 +7194,16 @@ def _build_assistant_message(self, assistant_message, finish_reason: str) -> dic _fn = getattr(tool_call, "function", None) _fn_name = getattr(_fn, "name", "") if _fn else "" _fn_args = getattr(_fn, "arguments", "{}") if _fn else "{}" - call_id = self._deterministic_call_id(_fn_name, _fn_args, len(tool_calls)) + call_id = self._deterministic_call_id( + _fn_name, _fn_args, len(tool_calls) + ) call_id = call_id.strip() response_item_id = getattr(tool_call, "response_item_id", None) - if not isinstance(response_item_id, str) or not response_item_id.strip(): + if ( + not isinstance(response_item_id, str) + or not response_item_id.strip() + ): _, embedded_response_item_id = self._split_responses_tool_id(raw_id) response_item_id = embedded_response_item_id @@ -6566,7 +7219,7 @@ def _build_assistant_message(self, assistant_message, finish_reason: str) -> dic "type": tool_call.type, "function": { "name": tool_call.function.name, - "arguments": tool_call.function.arguments + "arguments": tool_call.function.arguments, }, } # Preserve extra_content (e.g. Gemini thought_signature) so it @@ -6605,7 +7258,8 @@ def _sanitize_tool_calls_for_strict_api(api_msg: dict) -> dict: _STRIP_KEYS = {"call_id", "response_item_id"} api_msg["tool_calls"] = [ {k: v for k, v in tc.items() if k not in _STRIP_KEYS} - if isinstance(tc, dict) else tc + if isinstance(tc, dict) + else tc for tc in tool_calls ] return api_msg @@ -6641,12 +7295,14 @@ def flush_memories(self, messages: list = None, min_turns: int = None): return if "memory" not in self.valid_tool_names or not self._memory_store: return - effective_min = min_turns if min_turns is not None else self._memory_flush_min_turns + effective_min = ( + min_turns if min_turns is not None else self._memory_flush_min_turns + ) if self._user_turn_count < effective_min: return if messages is None: - messages = getattr(self, '_session_messages', None) + messages = getattr(self, "_session_messages", None) if not messages or len(messages) < 3: return @@ -6656,7 +7312,11 @@ def flush_memories(self, messages: list = None, min_turns: int = None): "corrections, and recurring patterns over task-specific details.]" ) _sentinel = f"__flush_{id(self)}_{time.monotonic()}" - flush_msg = {"role": "user", "content": flush_content, "_flush_sentinel": _sentinel} + flush_msg = { + "role": "user", + "content": flush_content, + "_flush_sentinel": _sentinel, + } messages.append(flush_msg) try: @@ -6678,11 +7338,13 @@ def flush_memories(self, messages: list = None, min_turns: int = None): api_messages.append(api_msg) if self._cached_system_prompt: - api_messages = [{"role": "system", "content": self._cached_system_prompt}] + api_messages + api_messages = [ + {"role": "system", "content": self._cached_system_prompt} + ] + api_messages # Make one API call with only the memory tool available memory_tool_def = None - for t in (self.tools or []): + for t in self.tools or []: if t.get("function", {}).get("name") == "memory": memory_tool_def = t break @@ -6694,6 +7356,7 @@ def flush_memories(self, messages: list = None, min_turns: int = None): # Use auxiliary client for the flush call when available -- # it's cheaper and avoids Codex Responses API incompatibility. from agent.auxiliary_client import call_llm as _call_llm + _aux_available = True try: response = _call_llm( @@ -6718,10 +7381,15 @@ def flush_memories(self, messages: list = None, min_turns: int = None): response = self._run_codex_stream(codex_kwargs) elif not _aux_available and self.api_mode == "anthropic_messages": # Native Anthropic — use the Anthropic client directly - from agent.anthropic_adapter import build_anthropic_kwargs as _build_ant_kwargs + from agent.anthropic_adapter import ( + build_anthropic_kwargs as _build_ant_kwargs, + ) + ant_kwargs = _build_ant_kwargs( - model=self.model, messages=api_messages, - tools=[memory_tool_def], max_tokens=5120, + model=self.model, + messages=api_messages, + tools=[memory_tool_def], + max_tokens=5120, reasoning_config=None, preserve_dots=self._anthropic_preserve_dots(), ) @@ -6735,7 +7403,10 @@ def flush_memories(self, messages: list = None, min_turns: int = None): **self._max_tokens_param(5120), } from agent.auxiliary_client import _get_task_timeout - response = self._ensure_primary_openai_client(reason="flush_memories").chat.completions.create( + + response = self._ensure_primary_openai_client( + reason="flush_memories" + ).chat.completions.create( **api_kwargs, timeout=_get_task_timeout("flush_memories") ) @@ -6746,8 +7417,13 @@ def flush_memories(self, messages: list = None, min_turns: int = None): if assistant_msg and assistant_msg.tool_calls: tool_calls = assistant_msg.tool_calls elif self.api_mode == "anthropic_messages" and not _aux_available: - from agent.anthropic_adapter import normalize_anthropic_response as _nar_flush - _flush_msg, _ = _nar_flush(response, strip_tool_prefix=self._is_anthropic_oauth) + from agent.anthropic_adapter import ( + normalize_anthropic_response as _nar_flush, + ) + + _flush_msg, _ = _nar_flush( + response, strip_tool_prefix=self._is_anthropic_oauth + ) if _flush_msg and _flush_msg.tool_calls: tool_calls = _flush_msg.tool_calls elif hasattr(response, "choices") and response.choices: @@ -6761,6 +7437,7 @@ def flush_memories(self, messages: list = None, min_turns: int = None): args = json.loads(tc.function.arguments) flush_target = args.get("target", "memory") from tools.memory_tool import memory_tool as _memory_tool + _memory_tool( action=args.get("action"), target=flush_target, @@ -6769,7 +7446,9 @@ def flush_memories(self, messages: list = None, min_turns: int = None): store=self._memory_store, ) if not self.quiet_mode: - print(f" 🧠 Memory flush: saved to {args.get('target', 'memory')}") + print( + f" 🧠 Memory flush: saved to {args.get('target', 'memory')}" + ) except Exception as e: logger.debug("Memory flush tool call failed: %s", e) except Exception as e: @@ -6784,7 +7463,15 @@ def flush_memories(self, messages: list = None, min_turns: int = None): if messages and messages[-1].get("_flush_sentinel") == _sentinel: messages.pop() - def _compress_context(self, messages: list, system_message: str, *, approx_tokens: int = None, task_id: str = "default", focus_topic: str = None) -> tuple: + def _compress_context( + self, + messages: list, + system_message: str, + *, + approx_tokens: int = None, + task_id: str = "default", + focus_topic: str = None, + ) -> tuple: """Compress conversation context and split the session in SQLite. Args: @@ -6798,8 +7485,10 @@ def _compress_context(self, messages: list, system_message: str, *, approx_token _pre_msg_count = len(messages) logger.info( "context compression started: session=%s messages=%d tokens=~%s model=%s focus=%r", - self.session_id or "none", _pre_msg_count, - f"{approx_tokens:,}" if approx_tokens else "unknown", self.model, + self.session_id or "none", + _pre_msg_count, + f"{approx_tokens:,}" if approx_tokens else "unknown", + self.model, focus_topic, ) # Pre-compression memory flush: let the model save memories before they're lost @@ -6812,7 +7501,9 @@ def _compress_context(self, messages: list, system_message: str, *, approx_token except Exception: pass - compressed = self.context_compressor.compress(messages, current_tokens=approx_tokens, focus_topic=focus_topic) + compressed = self.context_compressor.compress( + messages, current_tokens=approx_tokens, focus_topic=focus_topic + ) todo_snapshot = self._todo_store.format_for_injection() if todo_snapshot: @@ -6828,27 +7519,39 @@ def _compress_context(self, messages: list, system_message: str, *, approx_token old_title = self._session_db.get_session_title(self.session_id) self._session_db.end_session(self.session_id, "compression") old_session_id = self.session_id - self.session_id = f"{datetime.now().strftime('%Y%m%d_%H%M%S')}_{uuid.uuid4().hex[:6]}" + self.session_id = ( + f"{datetime.now().strftime('%Y%m%d_%H%M%S')}_{uuid.uuid4().hex[:6]}" + ) # Update session_log_file to point to the new session's JSON file - self.session_log_file = self.logs_dir / f"session_{self.session_id}.json" + self.session_log_file = ( + self.logs_dir / f"session_{self.session_id}.json" + ) self._session_db.create_session( session_id=self.session_id, - source=self.platform or os.environ.get("HERMES_SESSION_SOURCE", "cli"), + source=self.platform + or os.environ.get("HERMES_SESSION_SOURCE", "cli"), model=self.model, parent_session_id=old_session_id, ) # Auto-number the title for the continuation session if old_title: try: - new_title = self._session_db.get_next_title_in_lineage(old_title) + new_title = self._session_db.get_next_title_in_lineage( + old_title + ) self._session_db.set_session_title(self.session_id, new_title) except (ValueError, Exception) as e: logger.debug("Could not propagate title on compression: %s", e) - self._session_db.update_system_prompt(self.session_id, new_system_prompt) + self._session_db.update_system_prompt( + self.session_id, new_system_prompt + ) # Reset flush cursor — new session starts with no messages written self._last_flushed_db_idx = 0 except Exception as e: - logger.warning("Session DB compression split failed — new session will NOT be indexed: %s", e) + logger.warning( + "Session DB compression split failed — new session will NOT be indexed: %s", + e, + ) # Warn on repeated compressions (quality degrades with each pass) _cc = self.context_compressor.compression_count @@ -6861,10 +7564,9 @@ def _compress_context(self, messages: list, system_message: str, *, approx_token # Update token estimate after compaction so pressure calculations # use the post-compression count, not the stale pre-compression one. - _compressed_est = ( - estimate_tokens_rough(new_system_prompt) - + estimate_messages_tokens_rough(compressed) - ) + _compressed_est = estimate_tokens_rough( + new_system_prompt + ) + estimate_messages_tokens_rough(compressed) self.context_compressor.last_prompt_tokens = _compressed_est self.context_compressor.last_completion_tokens = 0 @@ -6887,18 +7589,27 @@ def _compress_context(self, messages: list, system_message: str, *, approx_token # file it needs the full content, not a "file unchanged" stub. try: from tools.file_tools import reset_file_dedup + reset_file_dedup(task_id) except Exception: pass logger.info( "context compression done: session=%s messages=%d->%d tokens=~%s", - self.session_id or "none", _pre_msg_count, len(compressed), + self.session_id or "none", + _pre_msg_count, + len(compressed), f"{_compressed_est:,}", ) return compressed, new_system_prompt - def _execute_tool_calls(self, assistant_message, messages: list, effective_task_id: str, api_call_count: int = 0) -> None: + def _execute_tool_calls( + self, + assistant_message, + messages: list, + effective_task_id: str, + api_call_count: int = 0, + ) -> None: """Execute tool calls from the assistant message and append results to messages. Dispatches to concurrent execution only for batches that look @@ -6921,8 +7632,13 @@ def _execute_tool_calls(self, assistant_message, messages: list, effective_task_ finally: self._executing_tools = False - def _invoke_tool(self, function_name: str, function_args: dict, effective_task_id: str, - tool_call_id: Optional[str] = None) -> str: + def _invoke_tool( + self, + function_name: str, + function_args: dict, + effective_task_id: str, + tool_call_id: Optional[str] = None, + ) -> str: """Invoke a single tool and return the result string. No display logic. Handles both agent-level tools (todo, memory, etc.) and registry-dispatched @@ -6933,8 +7649,11 @@ def _invoke_tool(self, function_name: str, function_args: dict, effective_task_i block_message: Optional[str] = None try: from hermes_cli.plugins import get_pre_tool_call_block_message + block_message = get_pre_tool_call_block_message( - function_name, function_args, task_id=effective_task_id or "", + function_name, + function_args, + task_id=effective_task_id or "", ) except Exception: pass @@ -6943,6 +7662,7 @@ def _invoke_tool(self, function_name: str, function_args: dict, effective_task_i if function_name == "todo": from tools.todo_tool import todo_tool as _todo_tool + return _todo_tool( todos=function_args.get("todos"), merge=function_args.get("merge", False), @@ -6950,8 +7670,11 @@ def _invoke_tool(self, function_name: str, function_args: dict, effective_task_i ) elif function_name == "session_search": if not self._session_db: - return json.dumps({"success": False, "error": "Session database not available."}) + return json.dumps( + {"success": False, "error": "Session database not available."} + ) from tools.session_search_tool import session_search as _session_search + return _session_search( query=function_args.get("query", ""), role_filter=function_args.get("role_filter"), @@ -6962,6 +7685,7 @@ def _invoke_tool(self, function_name: str, function_args: dict, effective_task_i elif function_name == "memory": target = function_args.get("target", "memory") from tools.memory_tool import memory_tool as _memory_tool + result = _memory_tool( action=function_args.get("action"), target=target, @@ -6970,7 +7694,10 @@ def _invoke_tool(self, function_name: str, function_args: dict, effective_task_i store=self._memory_store, ) # Bridge: notify external memory provider of built-in memory writes - if self._memory_manager and function_args.get("action") in ("add", "replace"): + if self._memory_manager and function_args.get("action") in ( + "add", + "replace", + ): try: self._memory_manager.on_memory_write( function_args.get("action", ""), @@ -6984,6 +7711,7 @@ def _invoke_tool(self, function_name: str, function_args: dict, effective_task_i return self._memory_manager.handle_tool_call(function_name, function_args) elif function_name == "clarify": from tools.clarify_tool import clarify_tool as _clarify_tool + return _clarify_tool( question=function_args.get("question", ""), choices=function_args.get("choices"), @@ -6991,6 +7719,7 @@ def _invoke_tool(self, function_name: str, function_args: dict, effective_task_i ) elif function_name == "delegate_task": from tools.delegate_tool import delegate_task as _delegate_task + return _delegate_task( goal=function_args.get("goal"), context=function_args.get("context"), @@ -7001,10 +7730,14 @@ def _invoke_tool(self, function_name: str, function_args: dict, effective_task_i ) else: return handle_function_call( - function_name, function_args, effective_task_id, + function_name, + function_args, + effective_task_id, tool_call_id=tool_call_id, session_id=self.session_id or "", - enabled_tools=list(self.valid_tool_names) if self.valid_tool_names else None, + enabled_tools=list(self.valid_tool_names) + if self.valid_tool_names + else None, skip_pre_tool_call_hook=True, ) @@ -7019,6 +7752,7 @@ def _wrap_verbose(label: str, text: str, indent: str = " ") -> str: """ import shutil as _shutil import textwrap as _tw + cols = _shutil.get_terminal_size((120, 24)).columns wrap_width = max(40, cols - len(indent)) out_lines: list[str] = [] @@ -7026,14 +7760,23 @@ def _wrap_verbose(label: str, text: str, indent: str = " ") -> str: if len(raw_line) <= wrap_width: out_lines.append(raw_line) else: - wrapped = _tw.wrap(raw_line, width=wrap_width, - break_long_words=True, - break_on_hyphens=False) + wrapped = _tw.wrap( + raw_line, + width=wrap_width, + break_long_words=True, + break_on_hyphens=False, + ) out_lines.extend(wrapped or [raw_line]) body = ("\n" + indent).join(out_lines) return f"{indent}{label}{body}" - def _execute_tool_calls_concurrent(self, assistant_message, messages: list, effective_task_id: str, api_call_count: int = 0) -> None: + def _execute_tool_calls_concurrent( + self, + assistant_message, + messages: list, + effective_task_id: str, + api_call_count: int = 0, + ) -> None: """Execute multiple tool calls concurrently using a thread pool. Results are collected in the original tool-call order and appended to @@ -7046,11 +7789,13 @@ def _execute_tool_calls_concurrent(self, assistant_message, messages: list, effe if self._interrupt_requested: print(f"{self.log_prefix}⚡ Interrupt: skipping {num_tools} tool call(s)") for tc in tool_calls: - messages.append({ - "role": "tool", - "content": f"[Tool execution cancelled — {tc.function.name} was skipped due to user interrupt]", - "tool_call_id": tc.id, - }) + messages.append( + { + "role": "tool", + "content": f"[Tool execution cancelled — {tc.function.name} was skipped due to user interrupt]", + "tool_call_id": tc.id, + } + ) return # ── Parse args + pre-execution bookkeeping ─────────────────────── @@ -7072,12 +7817,19 @@ def _execute_tool_calls_concurrent(self, assistant_message, messages: list, effe function_args = {} # Checkpoint for file-mutating tools - if function_name in ("write_file", "patch") and self._checkpoint_mgr.enabled: + if ( + function_name in ("write_file", "patch") + and self._checkpoint_mgr.enabled + ): try: file_path = function_args.get("path", "") if file_path: - work_dir = self._checkpoint_mgr.get_working_dir_for_path(file_path) - self._checkpoint_mgr.ensure_checkpoint(work_dir, f"before {function_name}") + work_dir = self._checkpoint_mgr.get_working_dir_for_path( + file_path + ) + self._checkpoint_mgr.ensure_checkpoint( + work_dir, f"before {function_name}" + ) except Exception: pass @@ -7086,7 +7838,9 @@ def _execute_tool_calls_concurrent(self, assistant_message, messages: list, effe try: cmd = function_args.get("command", "") if _is_destructive_command(cmd): - cwd = function_args.get("workdir") or os.getenv("TERMINAL_CWD", os.getcwd()) + cwd = function_args.get("workdir") or os.getenv( + "TERMINAL_CWD", os.getcwd() + ) self._checkpoint_mgr.ensure_checkpoint( cwd, f"before terminal: {cmd[:60]}" ) @@ -7103,10 +7857,20 @@ def _execute_tool_calls_concurrent(self, assistant_message, messages: list, effe args_str = json.dumps(args, ensure_ascii=False) if self.verbose_logging: print(f" 📞 Tool {i}: {name}({list(args.keys())})") - print(self._wrap_verbose("Args: ", json.dumps(args, indent=2, ensure_ascii=False))) + print( + self._wrap_verbose( + "Args: ", json.dumps(args, indent=2, ensure_ascii=False) + ) + ) else: - args_preview = args_str[:self.log_prefix_chars] + "..." if len(args_str) > self.log_prefix_chars else args_str - print(f" 📞 Tool {i}: {name}({list(args.keys())}) - {args_preview}") + args_preview = ( + args_str[: self.log_prefix_chars] + "..." + if len(args_str) > self.log_prefix_chars + else args_str + ) + print( + f" 📞 Tool {i}: {name}({list(args.keys())}) - {args_preview}" + ) for tc, name, args in parsed_calls: if self.tool_progress_callback: @@ -7131,28 +7895,51 @@ def _run_tool(index, tool_call, function_name, function_args): """Worker function executed in a thread.""" start = time.time() try: - result = self._invoke_tool(function_name, function_args, effective_task_id, tool_call.id) + result = self._invoke_tool( + function_name, function_args, effective_task_id, tool_call.id + ) except Exception as tool_error: result = f"Error executing tool '{function_name}': {tool_error}" - logger.error("_invoke_tool raised for %s: %s", function_name, tool_error, exc_info=True) + logger.error( + "_invoke_tool raised for %s: %s", + function_name, + tool_error, + exc_info=True, + ) duration = time.time() - start is_error, _ = _detect_tool_failure(function_name, result) if is_error: - logger.info("tool %s failed (%.2fs): %s", function_name, duration, result[:200]) + logger.info( + "tool %s failed (%.2fs): %s", function_name, duration, result[:200] + ) else: - logger.info("tool %s completed (%.2fs, %d chars)", function_name, duration, len(result)) + logger.info( + "tool %s completed (%.2fs, %d chars)", + function_name, + duration, + len(result), + ) results[index] = (function_name, function_args, result, duration, is_error) # Start spinner for CLI mode (skip when TUI handles tool progress) spinner = None - if self._should_emit_quiet_tool_messages() and self._should_start_quiet_spinner(): + if ( + self._should_emit_quiet_tool_messages() + and self._should_start_quiet_spinner() + ): face = random.choice(KawaiiSpinner.KAWAII_WAITING) - spinner = KawaiiSpinner(f"{face} ⚡ running {num_tools} tools concurrently", spinner_type='dots', print_fn=self._print_fn) + spinner = KawaiiSpinner( + f"{face} ⚡ running {num_tools} tools concurrently", + spinner_type="dots", + print_fn=self._print_fn, + ) spinner.start() try: max_workers = min(num_tools, _MAX_TOOL_WORKERS) - with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor: + with concurrent.futures.ThreadPoolExecutor( + max_workers=max_workers + ) as executor: futures = [] for i, (tc, name, args) in enumerate(parsed_calls): f = executor.submit(_run_tool, i, tc, name, args) @@ -7165,46 +7952,81 @@ def _run_tool(index, tool_call, function_name, function_args): # Build a summary message for the spinner stop completed = sum(1 for r in results if r is not None) total_dur = sum(r[3] for r in results if r is not None) - spinner.stop(f"⚡ {completed}/{num_tools} tools completed in {total_dur:.1f}s total") + spinner.stop( + f"⚡ {completed}/{num_tools} tools completed in {total_dur:.1f}s total" + ) # ── Post-execution: display per-tool results ───────────────────── for i, (tc, name, args) in enumerate(parsed_calls): r = results[i] if r is None: # Shouldn't happen, but safety fallback - function_result = f"Error executing tool '{name}': thread did not return a result" + function_result = ( + f"Error executing tool '{name}': thread did not return a result" + ) tool_duration = 0.0 else: - function_name, function_args, function_result, tool_duration, is_error = r + ( + function_name, + function_args, + function_result, + tool_duration, + is_error, + ) = r if is_error: - result_preview = function_result[:200] if len(function_result) > 200 else function_result - logger.warning("Tool %s returned error (%.2fs): %s", function_name, tool_duration, result_preview) + result_preview = ( + function_result[:200] + if len(function_result) > 200 + else function_result + ) + logger.warning( + "Tool %s returned error (%.2fs): %s", + function_name, + tool_duration, + result_preview, + ) if self.tool_progress_callback: try: self.tool_progress_callback( - "tool.completed", function_name, None, None, - duration=tool_duration, is_error=is_error, + "tool.completed", + function_name, + None, + None, + duration=tool_duration, + is_error=is_error, ) except Exception as cb_err: logging.debug(f"Tool progress callback error: {cb_err}") if self.verbose_logging: - logging.debug(f"Tool {function_name} completed in {tool_duration:.2f}s") - logging.debug(f"Tool result ({len(function_result)} chars): {function_result}") - - # Print cute message per tool - if self._should_emit_quiet_tool_messages(): - cute_msg = _get_cute_tool_message_impl(name, args, tool_duration, result=function_result) + logging.debug( + f"Tool {function_name} completed in {tool_duration:.2f}s" + ) + logging.debug( + f"Tool result ({len(function_result)} chars): {function_result}" + ) + + # Print cute message per tool + if self._should_emit_quiet_tool_messages(): + cute_msg = _get_cute_tool_message_impl( + name, args, tool_duration, result=function_result + ) self._safe_print(f" {cute_msg}") elif not self.quiet_mode: if self.verbose_logging: - print(f" ✅ Tool {i+1} completed in {tool_duration:.2f}s") + print(f" ✅ Tool {i + 1} completed in {tool_duration:.2f}s") print(self._wrap_verbose("Result: ", function_result)) else: - response_preview = function_result[:self.log_prefix_chars] + "..." if len(function_result) > self.log_prefix_chars else function_result - print(f" ✅ Tool {i+1} completed in {tool_duration:.2f}s - {response_preview}") + response_preview = ( + function_result[: self.log_prefix_chars] + "..." + if len(function_result) > self.log_prefix_chars + else function_result + ) + print( + f" ✅ Tool {i + 1} completed in {tool_duration:.2f}s - {response_preview}" + ) self._current_tool = None self._touch_activity(f"tool completed: {name} ({tool_duration:.1f}s)") @@ -7239,16 +8061,25 @@ def _run_tool(index, tool_call, function_name, function_args): turn_tool_msgs = messages[-num_tools:] enforce_turn_budget(turn_tool_msgs, env=get_active_env(effective_task_id)) - def _execute_tool_calls_sequential(self, assistant_message, messages: list, effective_task_id: str, api_call_count: int = 0) -> None: + def _execute_tool_calls_sequential( + self, + assistant_message, + messages: list, + effective_task_id: str, + api_call_count: int = 0, + ) -> None: """Execute tool calls sequentially (original behavior). Used for single calls or interactive tools.""" for i, tool_call in enumerate(assistant_message.tool_calls, 1): # SAFETY: check interrupt BEFORE starting each tool. # If the user sent "stop" during a previous tool's execution, # do NOT start any more tools -- skip them all immediately. if self._interrupt_requested: - remaining_calls = assistant_message.tool_calls[i-1:] + remaining_calls = assistant_message.tool_calls[i - 1 :] if remaining_calls: - self._vprint(f"{self.log_prefix}⚡ Interrupt: skipping {len(remaining_calls)} tool call(s)", force=True) + self._vprint( + f"{self.log_prefix}⚡ Interrupt: skipping {len(remaining_calls)} tool call(s)", + force=True, + ) for skipped_tc in remaining_calls: skipped_name = skipped_tc.function.name skip_msg = { @@ -7273,8 +8104,11 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe _block_msg: Optional[str] = None try: from hermes_cli.plugins import get_pre_tool_call_block_message + _block_msg = get_pre_tool_call_block_message( - function_name, function_args, task_id=effective_task_id or "", + function_name, + function_args, + task_id=effective_task_id or "", ) except Exception: pass @@ -7293,11 +8127,24 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe if not self.quiet_mode: args_str = json.dumps(function_args, ensure_ascii=False) if self.verbose_logging: - print(f" 📞 Tool {i}: {function_name}({list(function_args.keys())})") - print(self._wrap_verbose("Args: ", json.dumps(function_args, indent=2, ensure_ascii=False))) + print( + f" 📞 Tool {i}: {function_name}({list(function_args.keys())})" + ) + print( + self._wrap_verbose( + "Args: ", + json.dumps(function_args, indent=2, ensure_ascii=False), + ) + ) else: - args_preview = args_str[:self.log_prefix_chars] + "..." if len(args_str) > self.log_prefix_chars else args_str - print(f" 📞 Tool {i}: {function_name}({list(function_args.keys())}) - {args_preview}") + args_preview = ( + args_str[: self.log_prefix_chars] + "..." + if len(args_str) > self.log_prefix_chars + else args_str + ) + print( + f" 📞 Tool {i}: {function_name}({list(function_args.keys())}) - {args_preview}" + ) if _block_msg is None: self._current_tool = function_name @@ -7309,6 +8156,7 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe if _block_msg is None: try: from tools.environments.base import set_activity_callback + set_activity_callback(self._touch_activity) except Exception: pass @@ -7316,7 +8164,9 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe if _block_msg is None and self.tool_progress_callback: try: preview = _build_tool_preview(function_name, function_args) - self.tool_progress_callback("tool.started", function_name, preview, function_args) + self.tool_progress_callback( + "tool.started", function_name, preview, function_args + ) except Exception as cb_err: logging.debug(f"Tool progress callback error: {cb_err}") @@ -7327,11 +8177,17 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe logging.debug(f"Tool start callback error: {cb_err}") # Checkpoint: snapshot working dir before file-mutating tools - if _block_msg is None and function_name in ("write_file", "patch") and self._checkpoint_mgr.enabled: + if ( + _block_msg is None + and function_name in ("write_file", "patch") + and self._checkpoint_mgr.enabled + ): try: file_path = function_args.get("path", "") if file_path: - work_dir = self._checkpoint_mgr.get_working_dir_for_path(file_path) + work_dir = self._checkpoint_mgr.get_working_dir_for_path( + file_path + ) self._checkpoint_mgr.ensure_checkpoint( work_dir, f"before {function_name}" ) @@ -7339,11 +8195,17 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe pass # never block tool execution # Checkpoint before destructive terminal commands - if _block_msg is None and function_name == "terminal" and self._checkpoint_mgr.enabled: + if ( + _block_msg is None + and function_name == "terminal" + and self._checkpoint_mgr.enabled + ): try: cmd = function_args.get("command", "") if _is_destructive_command(cmd): - cwd = function_args.get("workdir") or os.getenv("TERMINAL_CWD", os.getcwd()) + cwd = function_args.get("workdir") or os.getenv( + "TERMINAL_CWD", os.getcwd() + ) self._checkpoint_mgr.ensure_checkpoint( cwd, f"before terminal: {cmd[:60]}" ) @@ -7358,6 +8220,7 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe tool_duration = 0.0 elif function_name == "todo": from tools.todo_tool import todo_tool as _todo_tool + function_result = _todo_tool( todos=function_args.get("todos"), merge=function_args.get("merge", False), @@ -7365,12 +8228,19 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe ) tool_duration = time.time() - tool_start_time if self._should_emit_quiet_tool_messages(): - self._vprint(f" {_get_cute_tool_message_impl('todo', function_args, tool_duration, result=function_result)}") + self._vprint( + f" {_get_cute_tool_message_impl('todo', function_args, tool_duration, result=function_result)}" + ) elif function_name == "session_search": if not self._session_db: - function_result = json.dumps({"success": False, "error": "Session database not available."}) + function_result = json.dumps( + {"success": False, "error": "Session database not available."} + ) else: - from tools.session_search_tool import session_search as _session_search + from tools.session_search_tool import ( + session_search as _session_search, + ) + function_result = _session_search( query=function_args.get("query", ""), role_filter=function_args.get("role_filter"), @@ -7380,10 +8250,13 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe ) tool_duration = time.time() - tool_start_time if self._should_emit_quiet_tool_messages(): - self._vprint(f" {_get_cute_tool_message_impl('session_search', function_args, tool_duration, result=function_result)}") + self._vprint( + f" {_get_cute_tool_message_impl('session_search', function_args, tool_duration, result=function_result)}" + ) elif function_name == "memory": target = function_args.get("target", "memory") from tools.memory_tool import memory_tool as _memory_tool + function_result = _memory_tool( action=function_args.get("action"), target=target, @@ -7393,9 +8266,12 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe ) tool_duration = time.time() - tool_start_time if self._should_emit_quiet_tool_messages(): - self._vprint(f" {_get_cute_tool_message_impl('memory', function_args, tool_duration, result=function_result)}") + self._vprint( + f" {_get_cute_tool_message_impl('memory', function_args, tool_duration, result=function_result)}" + ) elif function_name == "clarify": from tools.clarify_tool import clarify_tool as _clarify_tool + function_result = _clarify_tool( question=function_args.get("question", ""), choices=function_args.get("choices"), @@ -7403,19 +8279,31 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe ) tool_duration = time.time() - tool_start_time if self._should_emit_quiet_tool_messages(): - self._vprint(f" {_get_cute_tool_message_impl('clarify', function_args, tool_duration, result=function_result)}") + self._vprint( + f" {_get_cute_tool_message_impl('clarify', function_args, tool_duration, result=function_result)}" + ) elif function_name == "delegate_task": from tools.delegate_tool import delegate_task as _delegate_task + tasks_arg = function_args.get("tasks") if tasks_arg and isinstance(tasks_arg, list): spinner_label = f"🔀 delegating {len(tasks_arg)} tasks" else: goal_preview = (function_args.get("goal") or "")[:30] - spinner_label = f"🔀 {goal_preview}" if goal_preview else "🔀 delegating" + spinner_label = ( + f"🔀 {goal_preview}" if goal_preview else "🔀 delegating" + ) spinner = None - if self._should_emit_quiet_tool_messages() and self._should_start_quiet_spinner(): + if ( + self._should_emit_quiet_tool_messages() + and self._should_start_quiet_spinner() + ): face = random.choice(KawaiiSpinner.KAWAII_WAITING) - spinner = KawaiiSpinner(f"{face} {spinner_label}", spinner_type='dots', print_fn=self._print_fn) + spinner = KawaiiSpinner( + f"{face} {spinner_label}", + spinner_type="dots", + print_fn=self._print_fn, + ) spinner.start() self._delegate_spinner = spinner _delegate_result = None @@ -7432,30 +8320,58 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe finally: self._delegate_spinner = None tool_duration = time.time() - tool_start_time - cute_msg = _get_cute_tool_message_impl('delegate_task', function_args, tool_duration, result=_delegate_result) + cute_msg = _get_cute_tool_message_impl( + "delegate_task", + function_args, + tool_duration, + result=_delegate_result, + ) if spinner: spinner.stop(cute_msg) elif self._should_emit_quiet_tool_messages(): self._vprint(f" {cute_msg}") - elif self._context_engine_tool_names and function_name in self._context_engine_tool_names: + elif ( + self._context_engine_tool_names + and function_name in self._context_engine_tool_names + ): # Context engine tools (lcm_grep, lcm_describe, lcm_expand, etc.) spinner = None if self.quiet_mode and not self.tool_progress_callback: face = random.choice(KawaiiSpinner.KAWAII_WAITING) emoji = _get_tool_emoji(function_name) - preview = _build_tool_preview(function_name, function_args) or function_name - spinner = KawaiiSpinner(f"{face} {emoji} {preview}", spinner_type='dots', print_fn=self._print_fn) + preview = ( + _build_tool_preview(function_name, function_args) + or function_name + ) + spinner = KawaiiSpinner( + f"{face} {emoji} {preview}", + spinner_type="dots", + print_fn=self._print_fn, + ) spinner.start() _ce_result = None try: - function_result = self.context_compressor.handle_tool_call(function_name, function_args, messages=messages) + function_result = self.context_compressor.handle_tool_call( + function_name, function_args, messages=messages + ) _ce_result = function_result except Exception as tool_error: - function_result = json.dumps({"error": f"Context engine tool '{function_name}' failed: {tool_error}"}) - logger.error("context_engine.handle_tool_call raised for %s: %s", function_name, tool_error, exc_info=True) + function_result = json.dumps( + { + "error": f"Context engine tool '{function_name}' failed: {tool_error}" + } + ) + logger.error( + "context_engine.handle_tool_call raised for %s: %s", + function_name, + tool_error, + exc_info=True, + ) finally: tool_duration = time.time() - tool_start_time - cute_msg = _get_cute_tool_message_impl(function_name, function_args, tool_duration, result=_ce_result) + cute_msg = _get_cute_tool_message_impl( + function_name, function_args, tool_duration, result=_ce_result + ) if spinner: spinner.stop(cute_msg) elif self.quiet_mode: @@ -7464,50 +8380,97 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe # Memory provider tools (hindsight_retain, honcho_search, etc.) # These are not in the tool registry — route through MemoryManager. spinner = None - if self._should_emit_quiet_tool_messages() and self._should_start_quiet_spinner(): + if ( + self._should_emit_quiet_tool_messages() + and self._should_start_quiet_spinner() + ): face = random.choice(KawaiiSpinner.KAWAII_WAITING) emoji = _get_tool_emoji(function_name) - preview = _build_tool_preview(function_name, function_args) or function_name - spinner = KawaiiSpinner(f"{face} {emoji} {preview}", spinner_type='dots', print_fn=self._print_fn) + preview = ( + _build_tool_preview(function_name, function_args) + or function_name + ) + spinner = KawaiiSpinner( + f"{face} {emoji} {preview}", + spinner_type="dots", + print_fn=self._print_fn, + ) spinner.start() _mem_result = None try: - function_result = self._memory_manager.handle_tool_call(function_name, function_args) + function_result = self._memory_manager.handle_tool_call( + function_name, function_args + ) _mem_result = function_result except Exception as tool_error: - function_result = json.dumps({"error": f"Memory tool '{function_name}' failed: {tool_error}"}) - logger.error("memory_manager.handle_tool_call raised for %s: %s", function_name, tool_error, exc_info=True) + function_result = json.dumps( + {"error": f"Memory tool '{function_name}' failed: {tool_error}"} + ) + logger.error( + "memory_manager.handle_tool_call raised for %s: %s", + function_name, + tool_error, + exc_info=True, + ) finally: tool_duration = time.time() - tool_start_time - cute_msg = _get_cute_tool_message_impl(function_name, function_args, tool_duration, result=_mem_result) + cute_msg = _get_cute_tool_message_impl( + function_name, function_args, tool_duration, result=_mem_result + ) if spinner: spinner.stop(cute_msg) elif self._should_emit_quiet_tool_messages(): self._vprint(f" {cute_msg}") elif self.quiet_mode: spinner = None - if self._should_emit_quiet_tool_messages() and self._should_start_quiet_spinner(): + if ( + self._should_emit_quiet_tool_messages() + and self._should_start_quiet_spinner() + ): face = random.choice(KawaiiSpinner.KAWAII_WAITING) emoji = _get_tool_emoji(function_name) - preview = _build_tool_preview(function_name, function_args) or function_name - spinner = KawaiiSpinner(f"{face} {emoji} {preview}", spinner_type='dots', print_fn=self._print_fn) + preview = ( + _build_tool_preview(function_name, function_args) + or function_name + ) + spinner = KawaiiSpinner( + f"{face} {emoji} {preview}", + spinner_type="dots", + print_fn=self._print_fn, + ) spinner.start() _spinner_result = None try: function_result = handle_function_call( - function_name, function_args, effective_task_id, + function_name, + function_args, + effective_task_id, tool_call_id=tool_call.id, session_id=self.session_id or "", - enabled_tools=list(self.valid_tool_names) if self.valid_tool_names else None, + enabled_tools=list(self.valid_tool_names) + if self.valid_tool_names + else None, skip_pre_tool_call_hook=True, ) _spinner_result = function_result except Exception as tool_error: - function_result = f"Error executing tool '{function_name}': {tool_error}" - logger.error("handle_function_call raised for %s: %s", function_name, tool_error, exc_info=True) + function_result = ( + f"Error executing tool '{function_name}': {tool_error}" + ) + logger.error( + "handle_function_call raised for %s: %s", + function_name, + tool_error, + exc_info=True, + ) finally: tool_duration = time.time() - tool_start_time - cute_msg = _get_cute_tool_message_impl(function_name, function_args, tool_duration, result=_spinner_result) + cute_msg = _get_cute_tool_message_impl( + function_name, + function_args, + tool_duration, + result=_spinner_result, + ) if spinner: spinner.stop(cute_msg) elif self._should_emit_quiet_tool_messages(): @@ -7515,48 +8478,85 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe else: try: function_result = handle_function_call( - function_name, function_args, effective_task_id, + function_name, + function_args, + effective_task_id, tool_call_id=tool_call.id, session_id=self.session_id or "", - enabled_tools=list(self.valid_tool_names) if self.valid_tool_names else None, + enabled_tools=list(self.valid_tool_names) + if self.valid_tool_names + else None, skip_pre_tool_call_hook=True, ) except Exception as tool_error: - function_result = f"Error executing tool '{function_name}': {tool_error}" - logger.error("handle_function_call raised for %s: %s", function_name, tool_error, exc_info=True) + function_result = ( + f"Error executing tool '{function_name}': {tool_error}" + ) + logger.error( + "handle_function_call raised for %s: %s", + function_name, + tool_error, + exc_info=True, + ) tool_duration = time.time() - tool_start_time - result_preview = function_result if self.verbose_logging else ( - function_result[:200] if len(function_result) > 200 else function_result + result_preview = ( + function_result + if self.verbose_logging + else ( + function_result[:200] + if len(function_result) > 200 + else function_result + ) ) # Log tool errors to the persistent error log so [error] tags # in the UI always have a corresponding detailed entry on disk. _is_error_result, _ = _detect_tool_failure(function_name, function_result) if _is_error_result: - logger.warning("Tool %s returned error (%.2fs): %s", function_name, tool_duration, result_preview) + logger.warning( + "Tool %s returned error (%.2fs): %s", + function_name, + tool_duration, + result_preview, + ) else: - logger.info("tool %s completed (%.2fs, %d chars)", function_name, tool_duration, len(function_result)) + logger.info( + "tool %s completed (%.2fs, %d chars)", + function_name, + tool_duration, + len(function_result), + ) if self.tool_progress_callback: try: self.tool_progress_callback( - "tool.completed", function_name, None, None, - duration=tool_duration, is_error=_is_error_result, + "tool.completed", + function_name, + None, + None, + duration=tool_duration, + is_error=_is_error_result, ) except Exception as cb_err: logging.debug(f"Tool progress callback error: {cb_err}") self._current_tool = None - self._touch_activity(f"tool completed: {function_name} ({tool_duration:.1f}s)") + self._touch_activity( + f"tool completed: {function_name} ({tool_duration:.1f}s)" + ) if self.verbose_logging: logging.debug(f"Tool {function_name} completed in {tool_duration:.2f}s") - logging.debug(f"Tool result ({len(function_result)} chars): {function_result}") + logging.debug( + f"Tool result ({len(function_result)} chars): {function_result}" + ) if self.tool_complete_callback: try: - self.tool_complete_callback(tool_call.id, function_name, function_args, function_result) + self.tool_complete_callback( + tool_call.id, function_name, function_args, function_result + ) except Exception as cb_err: logging.debug(f"Tool complete callback error: {cb_err}") @@ -7568,14 +8568,16 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe ) # Discover subdirectory context files from tool arguments - subdir_hints = self._subdirectory_hints.check_tool_call(function_name, function_args) + subdir_hints = self._subdirectory_hints.check_tool_call( + function_name, function_args + ) if subdir_hints: function_result += subdir_hints tool_msg = { "role": "tool", "content": function_result, - "tool_call_id": tool_call.id + "tool_call_id": tool_call.id, } messages.append(tool_msg) @@ -7584,18 +8586,27 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe print(f" ✅ Tool {i} completed in {tool_duration:.2f}s") print(self._wrap_verbose("Result: ", function_result)) else: - response_preview = function_result[:self.log_prefix_chars] + "..." if len(function_result) > self.log_prefix_chars else function_result - print(f" ✅ Tool {i} completed in {tool_duration:.2f}s - {response_preview}") + response_preview = ( + function_result[: self.log_prefix_chars] + "..." + if len(function_result) > self.log_prefix_chars + else function_result + ) + print( + f" ✅ Tool {i} completed in {tool_duration:.2f}s - {response_preview}" + ) if self._interrupt_requested and i < len(assistant_message.tool_calls): remaining = len(assistant_message.tool_calls) - i - self._vprint(f"{self.log_prefix}⚡ Interrupt: skipping {remaining} remaining tool call(s)", force=True) + self._vprint( + f"{self.log_prefix}⚡ Interrupt: skipping {remaining} remaining tool call(s)", + force=True, + ) for skipped_tc in assistant_message.tool_calls[i:]: skipped_name = skipped_tc.function.name skip_msg = { "role": "tool", "content": f"[Tool execution skipped — {skipped_name} was not started. User sent a new message]", - "tool_call_id": skipped_tc.id + "tool_call_id": skipped_tc.id, } messages.append(skip_msg) break @@ -7606,9 +8617,9 @@ def _execute_tool_calls_sequential(self, assistant_message, messages: list, effe # ── Per-turn aggregate budget enforcement ───────────────────────── num_tools_seq = len(assistant_message.tool_calls) if num_tools_seq > 0: - enforce_turn_budget(messages[-num_tools_seq:], env=get_active_env(effective_task_id)) - - + enforce_turn_budget( + messages[-num_tools_seq:], env=get_active_env(effective_task_id) + ) def _emit_context_pressure(self, compaction_progress: float, compressor) -> None: """Notify the user that context is approaching the compaction threshold. @@ -7621,9 +8632,16 @@ def _emit_context_pressure(self, compaction_progress: float, compressor) -> None For CLI: prints a formatted line with a progress bar. For gateway: fires status_callback so the platform can send a chat message. """ - from agent.display import format_context_pressure, format_context_pressure_gateway + from agent.display import ( + format_context_pressure, + format_context_pressure_gateway, + ) - threshold_pct = compressor.threshold_tokens / compressor.context_length if compressor.context_length else 0.5 + threshold_pct = ( + compressor.threshold_tokens / compressor.context_length + if compressor.context_length + else 0.5 + ) # CLI output — always shown (these are user-facing status notifications, # not verbose debug output, so they bypass quiet_mode). @@ -7651,7 +8669,9 @@ def _emit_context_pressure(self, compaction_progress: float, compressor) -> None def _handle_max_iterations(self, messages: list, api_call_count: int) -> str: """Request a summary when max iterations are reached. Returns the final response text.""" - print(f"⚠️ Reached maximum iterations ({self.max_iterations}). Requesting summary...") + print( + f"⚠️ Reached maximum iterations ({self.max_iterations}). Requesting summary..." + ) summary_request = ( "You've reached the maximum number of tool-calling iterations allowed. " @@ -7667,7 +8687,11 @@ def _handle_max_iterations(self, messages: list, api_call_count: int) -> str: api_messages = [] for msg in messages: api_msg = msg.copy() - for internal_field in ("reasoning", "finish_reason", "_thinking_prefill"): + for internal_field in ( + "reasoning", + "finish_reason", + "_thinking_prefill", + ): api_msg.pop(internal_field, None) if _needs_sanitize: self._sanitize_tool_calls_for_strict_api(api_msg) @@ -7675,9 +8699,13 @@ def _handle_max_iterations(self, messages: list, api_call_count: int) -> str: effective_system = self._cached_system_prompt or "" if self.ephemeral_system_prompt: - effective_system = (effective_system + "\n\n" + self.ephemeral_system_prompt).strip() + effective_system = ( + effective_system + "\n\n" + self.ephemeral_system_prompt + ).strip() if effective_system: - api_messages = [{"role": "system", "content": effective_system}] + api_messages + api_messages = [ + {"role": "system", "content": effective_system} + ] + api_messages if self.prefill_messages: sys_offset = 1 if effective_system else 0 for idx, pfm in enumerate(self.prefill_messages): @@ -7691,7 +8719,7 @@ def _handle_max_iterations(self, messages: list, api_call_count: int) -> str: else: summary_extra_body["reasoning"] = { "enabled": True, - "effort": "medium" + "effort": "medium", } if _is_nous: summary_extra_body["tags"] = ["product=hermes-agent"] @@ -7701,7 +8729,11 @@ def _handle_max_iterations(self, messages: list, api_call_count: int) -> str: codex_kwargs.pop("tools", None) summary_response = self._run_codex_stream(codex_kwargs) assistant_message, _ = self._normalize_codex_response(summary_response) - final_response = (assistant_message.content or "").strip() if assistant_message else "" + final_response = ( + (assistant_message.content or "").strip() + if assistant_message + else "" + ) else: summary_kwargs = { "model": self.model, @@ -7727,29 +8759,49 @@ def _handle_max_iterations(self, messages: list, api_call_count: int) -> str: summary_kwargs["extra_body"] = summary_extra_body if self.api_mode == "anthropic_messages": - from agent.anthropic_adapter import build_anthropic_kwargs as _bak, normalize_anthropic_response as _nar - _ant_kw = _bak(model=self.model, messages=api_messages, tools=None, - max_tokens=self.max_tokens, reasoning_config=self.reasoning_config, - is_oauth=self._is_anthropic_oauth, - preserve_dots=self._anthropic_preserve_dots()) + from agent.anthropic_adapter import ( + build_anthropic_kwargs as _bak, + normalize_anthropic_response as _nar, + ) + + _ant_kw = _bak( + model=self.model, + messages=api_messages, + tools=None, + max_tokens=self.max_tokens, + reasoning_config=self.reasoning_config, + is_oauth=self._is_anthropic_oauth, + preserve_dots=self._anthropic_preserve_dots(), + ) summary_response = self._anthropic_messages_create(_ant_kw) - _msg, _ = _nar(summary_response, strip_tool_prefix=self._is_anthropic_oauth) + _msg, _ = _nar( + summary_response, strip_tool_prefix=self._is_anthropic_oauth + ) final_response = (_msg.content or "").strip() else: - summary_response = self._ensure_primary_openai_client(reason="iteration_limit_summary").chat.completions.create(**summary_kwargs) + summary_response = self._ensure_primary_openai_client( + reason="iteration_limit_summary" + ).chat.completions.create(**summary_kwargs) - if summary_response.choices and summary_response.choices[0].message.content: + if ( + summary_response.choices + and summary_response.choices[0].message.content + ): final_response = summary_response.choices[0].message.content else: final_response = "" if final_response: if "<think>" in final_response: - final_response = re.sub(r'<think>.*?</think>\s*', '', final_response, flags=re.DOTALL).strip() + final_response = re.sub( + r"<think>.*?</think>\s*", "", final_response, flags=re.DOTALL + ).strip() if final_response: messages.append({"role": "assistant", "content": final_response}) else: - final_response = "I reached the iteration limit and couldn't generate a summary." + final_response = ( + "I reached the iteration limit and couldn't generate a summary." + ) else: # Retry summary generation if self.api_mode == "codex_responses": @@ -7757,15 +8809,28 @@ def _handle_max_iterations(self, messages: list, api_call_count: int) -> str: codex_kwargs.pop("tools", None) retry_response = self._run_codex_stream(codex_kwargs) retry_msg, _ = self._normalize_codex_response(retry_response) - final_response = (retry_msg.content or "").strip() if retry_msg else "" + final_response = ( + (retry_msg.content or "").strip() if retry_msg else "" + ) elif self.api_mode == "anthropic_messages": - from agent.anthropic_adapter import build_anthropic_kwargs as _bak2, normalize_anthropic_response as _nar2 - _ant_kw2 = _bak2(model=self.model, messages=api_messages, tools=None, - is_oauth=self._is_anthropic_oauth, - max_tokens=self.max_tokens, reasoning_config=self.reasoning_config, - preserve_dots=self._anthropic_preserve_dots()) + from agent.anthropic_adapter import ( + build_anthropic_kwargs as _bak2, + normalize_anthropic_response as _nar2, + ) + + _ant_kw2 = _bak2( + model=self.model, + messages=api_messages, + tools=None, + is_oauth=self._is_anthropic_oauth, + max_tokens=self.max_tokens, + reasoning_config=self.reasoning_config, + preserve_dots=self._anthropic_preserve_dots(), + ) retry_response = self._anthropic_messages_create(_ant_kw2) - _retry_msg, _ = _nar2(retry_response, strip_tool_prefix=self._is_anthropic_oauth) + _retry_msg, _ = _nar2( + retry_response, strip_tool_prefix=self._is_anthropic_oauth + ) final_response = (_retry_msg.content or "").strip() else: summary_kwargs = { @@ -7777,22 +8842,36 @@ def _handle_max_iterations(self, messages: list, api_call_count: int) -> str: if summary_extra_body: summary_kwargs["extra_body"] = summary_extra_body - summary_response = self._ensure_primary_openai_client(reason="iteration_limit_summary_retry").chat.completions.create(**summary_kwargs) + summary_response = self._ensure_primary_openai_client( + reason="iteration_limit_summary_retry" + ).chat.completions.create(**summary_kwargs) - if summary_response.choices and summary_response.choices[0].message.content: + if ( + summary_response.choices + and summary_response.choices[0].message.content + ): final_response = summary_response.choices[0].message.content else: final_response = "" if final_response: if "<think>" in final_response: - final_response = re.sub(r'<think>.*?</think>\s*', '', final_response, flags=re.DOTALL).strip() + final_response = re.sub( + r"<think>.*?</think>\s*", + "", + final_response, + flags=re.DOTALL, + ).strip() if final_response: - messages.append({"role": "assistant", "content": final_response}) + messages.append( + {"role": "assistant", "content": final_response} + ) else: final_response = "I reached the iteration limit and couldn't generate a summary." else: - final_response = "I reached the iteration limit and couldn't generate a summary." + final_response = ( + "I reached the iteration limit and couldn't generate a summary." + ) except Exception as e: logging.warning(f"Failed to get summary response: {e}") @@ -7835,6 +8914,7 @@ def run_conversation( # Tag all log records on this thread with the session ID so # ``hermes logs --session <id>`` can filter a single conversation. from hermes_logging import set_session_context + set_session_context(self.session_id) # If the previous turn activated fallback, restore the primary @@ -7856,7 +8936,7 @@ def run_conversation( self._persist_user_message_override = persist_user_message # Generate unique task_id if not provided to isolate VMs between concurrent tasks effective_task_id = task_id or str(uuid.uuid4()) - + # Reset retry counters and iteration budget at the start of each turn # so subagent usage from a previous turn doesn't eat into the next one. self._invalid_tool_retries = 0 @@ -7895,12 +8975,17 @@ def run_conversation( self.iteration_budget = IterationBudget(self.max_iterations) # Log conversation turn start for debugging/observability - _msg_preview = (user_message[:80] + "...") if len(user_message) > 80 else user_message + _msg_preview = ( + (user_message[:80] + "...") if len(user_message) > 80 else user_message + ) _msg_preview = _msg_preview.replace("\n", " ") logger.info( "conversation turn: session=%s model=%s provider=%s platform=%s history=%d msg=%r", - self.session_id or "none", self.model, self.provider or "unknown", - self.platform or "unknown", len(conversation_history or []), + self.session_id or "none", + self.model, + self.provider or "unknown", + self.platform or "unknown", + len(conversation_history or []), _msg_preview, ) @@ -7912,25 +8997,29 @@ def run_conversation( # recover the todo state from the most recent todo tool response in history) if conversation_history and not self._todo_store.has_items(): self._hydrate_todo_store(conversation_history) - + # Prefill messages (few-shot priming) are injected at API-call time only, # never stored in the messages list. This keeps them ephemeral: they won't # be saved to session DB, session logs, or batch trajectories, but they're # automatically re-applied on every API call (including session continuations). - + # Track user turns for memory flush and periodic nudge logic self._user_turn_count += 1 # Preserve the original user message (no nudge injection). - original_user_message = persist_user_message if persist_user_message is not None else user_message + original_user_message = ( + persist_user_message if persist_user_message is not None else user_message + ) # Track memory nudge trigger (turn-based, checked here). # Skill trigger is checked AFTER the agent loop completes, based on # how many tool iterations THIS turn used. _should_review_memory = False - if (self._memory_nudge_interval > 0 - and "memory" in self.valid_tool_names - and self._memory_store): + if ( + self._memory_nudge_interval > 0 + and "memory" in self.valid_tool_names + and self._memory_store + ): self._turns_since_memory += 1 if self._turns_since_memory >= self._memory_nudge_interval: _should_review_memory = True @@ -7941,10 +9030,12 @@ def run_conversation( messages.append(user_msg) current_turn_user_idx = len(messages) - 1 self._persist_user_message_idx = current_turn_user_idx - + if not self.quiet_mode: - self._safe_print(f"💬 Starting conversation: '{user_message[:60]}{'...' if len(user_message) > 60 else ''}'") - + self._safe_print( + f"💬 Starting conversation: '{user_message[:60]}{'...' if len(user_message) > 60 else ''}'" + ) + # ── System prompt (cached per session for prefix caching) ── # Built once on first call, reused for all subsequent calls. # Only rebuilt after context compression events (which invalidate @@ -7979,6 +9070,7 @@ def run_conversation( # session-scoped state (e.g. warm a memory cache). try: from hermes_cli.plugins import invoke_hook as _invoke_hook + _invoke_hook( "on_session_start", session_id=self.session_id, @@ -7991,7 +9083,9 @@ def run_conversation( # Store the system prompt snapshot in SQLite if self._session_db: try: - self._session_db.update_system_prompt(self.session_id, self._cached_system_prompt) + self._session_db.update_system_prompt( + self.session_id, self._cached_system_prompt + ) except Exception as e: logger.debug("Session DB update_system_prompt failed: %s", e) @@ -8006,8 +9100,10 @@ def run_conversation( # 4xx and abort the request entirely). if ( self.compression_enabled - and len(messages) > self.context_compressor.protect_first_n - + self.context_compressor.protect_last_n + 1 + and len(messages) + > self.context_compressor.protect_first_n + + self.context_compressor.protect_last_n + + 1 ): # Include tool schema tokens — with many tools these can add # 20-30K+ tokens that the old sys+msg estimate missed entirely. @@ -8035,7 +9131,9 @@ def run_conversation( for _pass in range(3): _orig_len = len(messages) messages, active_system_prompt = self._compress_context( - messages, system_message, approx_tokens=_preflight_tokens, + messages, + system_message, + approx_tokens=_preflight_tokens, task_id=effective_task_id, ) if len(messages) >= _orig_len: @@ -8079,6 +9177,7 @@ def run_conversation( _plugin_user_context = "" try: from hermes_cli.plugins import invoke_hook as _invoke_hook + _pre_results = _invoke_hook( "pre_llm_call", session_id=self.session_id, @@ -8110,7 +9209,7 @@ def run_conversation( truncated_response_prefix = "" compression_attempts = 0 _turn_exit_reason = "unknown" # Diagnostic: why the loop ended - + # Record the execution thread so interrupt()/clear_interrupt() can # scope the tool-level interrupt signal to THIS agent's thread only. # Must be set before clear_interrupt() which uses it. @@ -8127,12 +9226,18 @@ def run_conversation( _ext_prefetch_cache = "" if self._memory_manager: try: - _query = original_user_message if isinstance(original_user_message, str) else "" + _query = ( + original_user_message + if isinstance(original_user_message, str) + else "" + ) _ext_prefetch_cache = self._memory_manager.prefetch_all(_query) or "" except Exception: pass - while (api_call_count < self.max_iterations and self.iteration_budget.remaining > 0) or self._budget_grace_call: + while ( + api_call_count < self.max_iterations and self.iteration_budget.remaining > 0 + ) or self._budget_grace_call: # Reset per-turn checkpoint dedup so each iteration can take one snapshot self._checkpoint_mgr.new_turn() @@ -8141,9 +9246,11 @@ def run_conversation( interrupted = True _turn_exit_reason = "interrupted_by_user" if not self.quiet_mode: - self._safe_print("\n⚡ Breaking out of tool loop due to interrupt...") + self._safe_print( + "\n⚡ Breaking out of tool loop due to interrupt..." + ) break - + api_call_count += 1 self._api_call_count = api_call_count self._touch_activity(f"starting API call #{api_call_count}") @@ -8156,7 +9263,9 @@ def run_conversation( elif not self.iteration_budget.consume(): _turn_exit_reason = "budget_exhausted" if not self.quiet_mode: - self._safe_print(f"\n⚠️ Iteration budget exhausted ({self.iteration_budget.used}/{self.iteration_budget.max_total} iterations used)") + self._safe_print( + f"\n⚠️ Iteration budget exhausted ({self.iteration_budget.used}/{self.iteration_budget.max_total} iterations used)" + ) break # Fire step_callback for gateway hooks (agent:step event) @@ -8184,14 +9293,20 @@ def run_conversation( break self.step_callback(api_call_count, prev_tools) except Exception as _step_err: - logger.debug("step_callback error (iteration %s): %s", api_call_count, _step_err) + logger.debug( + "step_callback error (iteration %s): %s", + api_call_count, + _step_err, + ) # Track tool-calling iterations for skill nudge. # Counter resets whenever skill_manage is actually used. - if (self._skill_nudge_interval > 0 - and "skill_manage" in self.valid_tool_names): + if ( + self._skill_nudge_interval > 0 + and "skill_manage" in self.valid_tool_names + ): self._iters_since_skill += 1 - + # Prepare messages for API call # If we have an ephemeral system prompt, prepend it to the messages # Note: Reasoning is embedded in content via <think> tags for trajectory storage. @@ -8217,7 +9332,9 @@ def run_conversation( if _injections: _base = api_msg.get("content", "") if isinstance(_base, str): - api_msg["content"] = _base + "\n\n" + "\n\n".join(_injections) + api_msg["content"] = ( + _base + "\n\n" + "\n\n".join(_injections) + ) # For ALL assistant messages, pass reasoning back to the API # This ensures multi-turn reasoning context is preserved @@ -8252,13 +9369,17 @@ def run_conversation( # prompt, so the stable cache prefix remains unchanged. effective_system = active_system_prompt or "" if self.ephemeral_system_prompt: - effective_system = (effective_system + "\n\n" + self.ephemeral_system_prompt).strip() + effective_system = ( + effective_system + "\n\n" + self.ephemeral_system_prompt + ).strip() # NOTE: Plugin context from pre_llm_call hooks is injected into the # user message (see injection block above), NOT the system prompt. # This is intentional — system prompt modifications break the prompt # cache prefix. The system prompt is reserved for Hermes internals. if effective_system: - api_messages = [{"role": "system", "content": effective_system}] + api_messages + api_messages = [ + {"role": "system", "content": effective_system} + ] + api_messages # Inject ephemeral prefill messages right after the system prompt # but before conversation history. Same API-call-time-only pattern. @@ -8272,7 +9393,11 @@ def run_conversation( # inject cache_control breakpoints (system + last 3 messages) to reduce # input token costs by ~75% on multi-turn conversations. if self._use_prompt_caching: - api_messages = apply_anthropic_cache_control(api_messages, cache_ttl=self._cache_ttl, native_anthropic=(self.api_mode == 'anthropic_messages')) + api_messages = apply_anthropic_cache_control( + api_messages, + cache_ttl=self._cache_ttl, + native_anthropic=(self.api_mode == "anthropic_messages"), + ) # Safety net: strip orphaned tool results / add stubs for missing # results before sending to the API. Runs unconditionally — not @@ -8298,13 +9423,17 @@ def run_conversation( if isinstance(tc, dict) and "function" in tc: try: args_obj = json.loads(tc["function"]["arguments"]) - tc = {**tc, "function": { - **tc["function"], - "arguments": json.dumps( - args_obj, separators=(",", ":"), - sort_keys=True, - ), - }} + tc = { + **tc, + "function": { + **tc["function"], + "arguments": json.dumps( + args_obj, + separators=(",", ":"), + sort_keys=True, + ), + }, + } except Exception: pass new_tcs.append(tc) @@ -8313,14 +9442,20 @@ def run_conversation( # Calculate approximate request size for logging total_chars = sum(len(str(msg)) for msg in api_messages) approx_tokens = estimate_messages_tokens_rough(api_messages) - + # Thinking spinner for quiet mode (animated during API call) thinking_spinner = None - + if not self.quiet_mode: - self._vprint(f"\n{self.log_prefix}🔄 Making API call #{api_call_count}/{self.max_iterations}...") - self._vprint(f"{self.log_prefix} 📊 Request size: {len(api_messages)} messages, ~{approx_tokens:,} tokens (~{total_chars:,} chars)") - self._vprint(f"{self.log_prefix} 🔧 Available tools: {len(self.tools) if self.tools else 0}") + self._vprint( + f"\n{self.log_prefix}🔄 Making API call #{api_call_count}/{self.max_iterations}..." + ) + self._vprint( + f"{self.log_prefix} 📊 Request size: {len(api_messages)} messages, ~{approx_tokens:,} tokens (~{total_chars:,} chars)" + ) + self._vprint( + f"{self.log_prefix} 🔧 Available tools: {len(self.tools) if self.tools else 0}" + ) else: # Animated thinking spinner in quiet mode face = random.choice(KawaiiSpinner.KAWAII_THINKING) @@ -8329,27 +9464,40 @@ def run_conversation( # CLI TUI mode: use prompt_toolkit widget instead of raw spinner # (works in both streaming and non-streaming modes) self.thinking_callback(f"{face} {verb}...") - elif not self._has_stream_consumers() and self._should_start_quiet_spinner(): + elif ( + not self._has_stream_consumers() + and self._should_start_quiet_spinner() + ): # Raw KawaiiSpinner only when no streaming consumers and the # spinner output has a safe sink. - spinner_type = random.choice(['brain', 'sparkle', 'pulse', 'moon', 'star']) - thinking_spinner = KawaiiSpinner(f"{face} {verb}...", spinner_type=spinner_type, print_fn=self._print_fn) + spinner_type = random.choice( + ["brain", "sparkle", "pulse", "moon", "star"] + ) + thinking_spinner = KawaiiSpinner( + f"{face} {verb}...", + spinner_type=spinner_type, + print_fn=self._print_fn, + ) thinking_spinner.start() - + # Log request details if verbose if self.verbose_logging: - logging.debug(f"API Request - Model: {self.model}, Messages: {len(messages)}, Tools: {len(self.tools) if self.tools else 0}") - logging.debug(f"Last message role: {messages[-1]['role'] if messages else 'none'}") + logging.debug( + f"API Request - Model: {self.model}, Messages: {len(messages)}, Tools: {len(self.tools) if self.tools else 0}" + ) + logging.debug( + f"Last message role: {messages[-1]['role'] if messages else 'none'}" + ) logging.debug(f"Total message size: ~{approx_tokens:,} tokens") - + api_start_time = time.time() retry_count = 0 max_retries = 3 primary_recovery_attempted = False max_compression_attempts = 3 - codex_auth_retry_attempted=False - anthropic_auth_retry_attempted=False - nous_auth_retry_attempted=False + codex_auth_retry_attempted = False + anthropic_auth_retry_attempted = False + nous_auth_retry_attempted = False thinking_sig_retry_attempted = False has_retried_429 = False restart_with_compressed_messages = False @@ -8366,10 +9514,13 @@ def run_conversation( if self._force_ascii_payload: _sanitize_structure_non_ascii(api_kwargs) if self.api_mode == "codex_responses": - api_kwargs = self._preflight_codex_api_kwargs(api_kwargs, allow_stream=False) + api_kwargs = self._preflight_codex_api_kwargs( + api_kwargs, allow_stream=False + ) try: from hermes_cli.plugins import invoke_hook as _invoke_hook + _invoke_hook( "pre_api_request", task_id=effective_task_id, @@ -8422,6 +9573,7 @@ def _stop_spinner(): # health checking, but skip for Mock clients in tests # (mocks return SimpleNamespace, not stream iterators). from unittest.mock import Mock + if isinstance(getattr(self, "client", None), Mock): _use_streaming = False @@ -8431,9 +9583,9 @@ def _stop_spinner(): ) else: response = self._interruptible_api_call(api_kwargs) - + api_duration = time.time() - api_start_time - + # Stop thinking spinner silently -- the response box or tool # execution messages that follow are more informative. if thinking_spinner: @@ -8441,20 +9593,30 @@ def _stop_spinner(): thinking_spinner = None if self.thinking_callback: self.thinking_callback("") - + if not self.quiet_mode: - self._vprint(f"{self.log_prefix}⏱️ API call completed in {api_duration:.2f}s") - + self._vprint( + f"{self.log_prefix}⏱️ API call completed in {api_duration:.2f}s" + ) + if self.verbose_logging: # Log response with provider info if available - resp_model = getattr(response, 'model', 'N/A') if response else 'N/A' - logging.debug(f"API Response received - Model: {resp_model}, Usage: {response.usage if hasattr(response, 'usage') else 'N/A'}") - + resp_model = ( + getattr(response, "model", "N/A") if response else "N/A" + ) + logging.debug( + f"API Response received - Model: {resp_model}, Usage: {response.usage if hasattr(response, 'usage') else 'N/A'}" + ) + # Validate response shape before proceeding response_invalid = False error_details = [] if self.api_mode == "codex_responses": - output_items = getattr(response, "output", None) if response is not None else None + output_items = ( + getattr(response, "output", None) + if response is not None + else None + ) if response is None: response_invalid = True error_details.append("response is None") @@ -8467,7 +9629,9 @@ def _stop_spinner(): # from response.output_text. Only mark invalid # when that fallback is also absent. _out_text = getattr(response, "output_text", None) - _out_text_stripped = _out_text.strip() if isinstance(_out_text, str) else "" + _out_text_stripped = ( + _out_text.strip() if isinstance(_out_text, str) else "" + ) if _out_text_stripped: logger.debug( "Codex response.output is empty but output_text is present " @@ -8476,18 +9640,25 @@ def _stop_spinner(): ) else: _resp_status = getattr(response, "status", None) - _resp_incomplete = getattr(response, "incomplete_details", None) + _resp_incomplete = getattr( + response, "incomplete_details", None + ) logger.warning( "Codex response.output is empty after stream backfill " "(status=%s, incomplete_details=%s, model=%s). %s", - _resp_status, _resp_incomplete, + _resp_status, + _resp_incomplete, getattr(response, "model", None), f"api_mode={self.api_mode} provider={self.provider}", ) response_invalid = True error_details.append("response.output is empty") elif self.api_mode == "anthropic_messages": - content_blocks = getattr(response, "content", None) if response is not None else None + content_blocks = ( + getattr(response, "content", None) + if response is not None + else None + ) if response is None: response_invalid = True error_details.append("response is None") @@ -8498,12 +9669,19 @@ def _stop_spinner(): response_invalid = True error_details.append("response.content is empty") else: - if response is None or not hasattr(response, 'choices') or response.choices is None or not response.choices: + if ( + response is None + or not hasattr(response, "choices") + or response.choices is None + or not response.choices + ): response_invalid = True if response is None: error_details.append("response is None") - elif not hasattr(response, 'choices'): - error_details.append("response has no 'choices' attribute") + elif not hasattr(response, "choices"): + error_details.append( + "response has no 'choices' attribute" + ) elif response.choices is None: error_details.append("response.choices is None") else: @@ -8516,16 +9694,18 @@ def _stop_spinner(): thinking_spinner = None if self.thinking_callback: self.thinking_callback("") - + # Invalid response — could be rate limiting, provider timeout, # upstream server error, or malformed response. retry_count += 1 - + # Eager fallback: empty/malformed responses are a common # rate-limit symptom. Switch to fallback immediately # rather than retrying with extended backoff. if self._fallback_index < len(self._fallback_chain): - self._emit_status("⚠️ Empty/malformed response — switching to fallback...") + self._emit_status( + "⚠️ Empty/malformed response — switching to fallback..." + ) if self._try_activate_fallback(): retry_count = 0 compression_attempts = 0 @@ -8535,31 +9715,51 @@ def _stop_spinner(): # Check for error field in response (some providers include this) error_msg = "Unknown" provider_name = "Unknown" - if response and hasattr(response, 'error') and response.error: + if response and hasattr(response, "error") and response.error: error_msg = str(response.error) # Try to extract provider from error metadata - if hasattr(response.error, 'metadata') and response.error.metadata: - provider_name = response.error.metadata.get('provider_name', 'Unknown') - elif response and hasattr(response, 'message') and response.message: + if ( + hasattr(response.error, "metadata") + and response.error.metadata + ): + provider_name = response.error.metadata.get( + "provider_name", "Unknown" + ) + elif ( + response + and hasattr(response, "message") + and response.message + ): error_msg = str(response.message) - + # Try to get provider from model field (OpenRouter often returns actual model used) - if provider_name == "Unknown" and response and hasattr(response, 'model') and response.model: + if ( + provider_name == "Unknown" + and response + and hasattr(response, "model") + and response.model + ): provider_name = f"model={response.model}" - + # Check for x-openrouter-provider or similar metadata if provider_name == "Unknown" and response: # Log all response attributes for debugging - resp_attrs = {k: str(v)[:100] for k, v in vars(response).items() if not k.startswith('_')} + resp_attrs = { + k: str(v)[:100] + for k, v in vars(response).items() + if not k.startswith("_") + } if self.verbose_logging: - logging.debug(f"Response attributes for invalid response: {resp_attrs}") - + logging.debug( + f"Response attributes for invalid response: {resp_attrs}" + ) + # Extract error code from response for contextual diagnostics _resp_error_code = None - if response and hasattr(response, 'error') and response.error: - _code_raw = getattr(response.error, 'code', None) + if response and hasattr(response, "error") and response.error: + _code_raw = getattr(response.error, "code", None) if _code_raw is None and isinstance(response.error, dict): - _code_raw = response.error.get('code') + _code_raw = response.error.get("code") if _code_raw is not None: try: _resp_error_code = int(_code_raw) @@ -8571,13 +9771,17 @@ def _stop_spinner(): if _resp_error_code == 524: _failure_hint = f"upstream provider timed out (Cloudflare 524, {api_duration:.0f}s)" elif _resp_error_code == 504: - _failure_hint = f"upstream gateway timeout (504, {api_duration:.0f}s)" + _failure_hint = ( + f"upstream gateway timeout (504, {api_duration:.0f}s)" + ) elif _resp_error_code == 429: _failure_hint = f"rate limited by upstream provider (429)" elif _resp_error_code in (500, 502): _failure_hint = f"upstream server error ({_resp_error_code}, {api_duration:.0f}s)" elif _resp_error_code in (503, 529): - _failure_hint = f"upstream provider overloaded ({_resp_error_code})" + _failure_hint = ( + f"upstream provider overloaded ({_resp_error_code})" + ) elif _resp_error_code is not None: _failure_hint = f"upstream error (code {_resp_error_code}, {api_duration:.0f}s)" elif api_duration < 10: @@ -8587,42 +9791,69 @@ def _stop_spinner(): else: _failure_hint = f"response time {api_duration:.1f}s" - self._vprint(f"{self.log_prefix}⚠️ Invalid API response (attempt {retry_count}/{max_retries}): {', '.join(error_details)}", force=True) - self._vprint(f"{self.log_prefix} 🏢 Provider: {provider_name}", force=True) + self._vprint( + f"{self.log_prefix}⚠️ Invalid API response (attempt {retry_count}/{max_retries}): {', '.join(error_details)}", + force=True, + ) + self._vprint( + f"{self.log_prefix} 🏢 Provider: {provider_name}", + force=True, + ) cleaned_provider_error = self._clean_error_message(error_msg) - self._vprint(f"{self.log_prefix} 📝 Provider message: {cleaned_provider_error}", force=True) - self._vprint(f"{self.log_prefix} ⏱️ {_failure_hint}", force=True) - + self._vprint( + f"{self.log_prefix} 📝 Provider message: {cleaned_provider_error}", + force=True, + ) + self._vprint( + f"{self.log_prefix} ⏱️ {_failure_hint}", force=True + ) + if retry_count >= max_retries: # Try fallback before giving up - self._emit_status(f"⚠️ Max retries ({max_retries}) for invalid responses — trying fallback...") + self._emit_status( + f"⚠️ Max retries ({max_retries}) for invalid responses — trying fallback..." + ) if self._try_activate_fallback(): retry_count = 0 compression_attempts = 0 primary_recovery_attempted = False continue - self._emit_status(f"❌ Max retries ({max_retries}) exceeded for invalid responses. Giving up.") - logging.error(f"{self.log_prefix}Invalid API response after {max_retries} retries.") + self._emit_status( + f"❌ Max retries ({max_retries}) exceeded for invalid responses. Giving up." + ) + logging.error( + f"{self.log_prefix}Invalid API response after {max_retries} retries." + ) self._persist_session(messages, conversation_history) return { "messages": messages, "completed": False, "api_calls": api_call_count, "error": f"Invalid API response after {max_retries} retries: {_failure_hint}", - "failed": True # Mark as failure for filtering + "failed": True, # Mark as failure for filtering } - + # Backoff before retry — jittered exponential: 5s base, 120s cap - wait_time = jittered_backoff(retry_count, base_delay=5.0, max_delay=120.0) - self._vprint(f"{self.log_prefix}⏳ Retrying in {wait_time:.1f}s ({_failure_hint})...", force=True) - logging.warning(f"Invalid API response (retry {retry_count}/{max_retries}): {', '.join(error_details)} | Provider: {provider_name}") - + wait_time = jittered_backoff( + retry_count, base_delay=5.0, max_delay=120.0 + ) + self._vprint( + f"{self.log_prefix}⏳ Retrying in {wait_time:.1f}s ({_failure_hint})...", + force=True, + ) + logging.warning( + f"Invalid API response (retry {retry_count}/{max_retries}): {', '.join(error_details)} | Provider: {provider_name}" + ) + # Sleep in small increments to stay responsive to interrupts sleep_end = time.time() + wait_time _backoff_touch_counter = 0 while time.time() < sleep_end: if self._interrupt_requested: - self._vprint(f"{self.log_prefix}⚡ Interrupt detected during retry wait, aborting.", force=True) + self._vprint( + f"{self.log_prefix}⚡ Interrupt detected during retry wait, aborting.", + force=True, + ) self._persist_session(messages, conversation_history) self.clear_interrupt() return { @@ -8646,24 +9877,41 @@ def _stop_spinner(): # Check finish_reason before proceeding if self.api_mode == "codex_responses": status = getattr(response, "status", None) - incomplete_details = getattr(response, "incomplete_details", None) + incomplete_details = getattr( + response, "incomplete_details", None + ) incomplete_reason = None if isinstance(incomplete_details, dict): incomplete_reason = incomplete_details.get("reason") else: - incomplete_reason = getattr(incomplete_details, "reason", None) - if status == "incomplete" and incomplete_reason in {"max_output_tokens", "length"}: + incomplete_reason = getattr( + incomplete_details, "reason", None + ) + if status == "incomplete" and incomplete_reason in { + "max_output_tokens", + "length", + }: finish_reason = "length" else: finish_reason = "stop" elif self.api_mode == "anthropic_messages": - stop_reason_map = {"end_turn": "stop", "tool_use": "tool_calls", "max_tokens": "length", "stop_sequence": "stop"} - finish_reason = stop_reason_map.get(response.stop_reason, "stop") + stop_reason_map = { + "end_turn": "stop", + "tool_use": "tool_calls", + "max_tokens": "length", + "stop_sequence": "stop", + } + finish_reason = stop_reason_map.get( + response.stop_reason, "stop" + ) else: finish_reason = response.choices[0].finish_reason if finish_reason == "length": - self._vprint(f"{self.log_prefix}⚠️ Response truncated (finish_reason='length') - model hit max output tokens", force=True) + self._vprint( + f"{self.log_prefix}⚠️ Response truncated (finish_reason='length') - model hit max output tokens", + force=True, + ) # ── Detect thinking-budget exhaustion ────────────── # When the model spends ALL output tokens on reasoning @@ -8673,16 +9921,30 @@ def _stop_spinner(): _trunc_content = None _trunc_has_tool_calls = False if self.api_mode == "chat_completions": - _trunc_msg = response.choices[0].message if (hasattr(response, "choices") and response.choices) else None - _trunc_content = getattr(_trunc_msg, "content", None) if _trunc_msg else None - _trunc_has_tool_calls = bool(getattr(_trunc_msg, "tool_calls", None)) if _trunc_msg else False + _trunc_msg = ( + response.choices[0].message + if (hasattr(response, "choices") and response.choices) + else None + ) + _trunc_content = ( + getattr(_trunc_msg, "content", None) + if _trunc_msg + else None + ) + _trunc_has_tool_calls = ( + bool(getattr(_trunc_msg, "tool_calls", None)) + if _trunc_msg + else False + ) elif self.api_mode == "anthropic_messages": # Anthropic response.content is a list of blocks _text_parts = [] for _blk in getattr(response, "content", []): if getattr(_blk, "type", None) == "text": _text_parts.append(getattr(_blk, "text", "")) - _trunc_content = "\n".join(_text_parts) if _text_parts else None + _trunc_content = ( + "\n".join(_text_parts) if _text_parts else None + ) # A response is "thinking exhausted" only when the model # actually produced reasoning blocks but no visible text after @@ -8692,8 +9954,9 @@ def _stop_spinner(): # truncations that deserve continuation retries, not as # thinking-budget exhaustion. _has_think_tags = bool( - _trunc_content and re.search( - r'<(?:think|thinking|reasoning|REASONING_SCRATCHPAD)[^>]*>', + _trunc_content + and re.search( + r"<(?:think|thinking|reasoning|REASONING_SCRATCHPAD)[^>]*>", _trunc_content, re.IGNORECASE, ) @@ -8702,7 +9965,12 @@ def _stop_spinner(): not _trunc_has_tool_calls and _has_think_tags and ( - (_trunc_content is not None and not self._has_content_after_think_block(_trunc_content)) + ( + _trunc_content is not None + and not self._has_content_after_think_block( + _trunc_content + ) + ) or _trunc_content is None ) ) @@ -8745,10 +10013,14 @@ def _stop_spinner(): assistant_message = response.choices[0].message if not assistant_message.tool_calls: length_continue_retries += 1 - interim_msg = self._build_assistant_message(assistant_message, finish_reason) + interim_msg = self._build_assistant_message( + assistant_message, finish_reason + ) messages.append(interim_msg) if assistant_message.content: - truncated_response_prefix += assistant_message.content + truncated_response_prefix += ( + assistant_message.content + ) if length_continue_retries < 3: self._vprint( @@ -8769,7 +10041,9 @@ def _stop_spinner(): restart_with_length_continuation = True break - partial_response = self._strip_think_blocks(truncated_response_prefix).strip() + partial_response = self._strip_think_blocks( + truncated_response_prefix + ).strip() self._cleanup_task_resources(effective_task_id) self._persist_session(messages, conversation_history) return { @@ -8811,8 +10085,12 @@ def _stop_spinner(): # If we have prior messages, roll back to last complete state if len(messages) > 1: - self._vprint(f"{self.log_prefix} ⏪ Rolling back to last complete assistant turn") - rolled_back_messages = self._get_messages_up_to_last_assistant(messages) + self._vprint( + f"{self.log_prefix} ⏪ Rolling back to last complete assistant turn" + ) + rolled_back_messages = ( + self._get_messages_up_to_last_assistant(messages) + ) self._cleanup_task_resources(effective_task_id) self._persist_session(messages, conversation_history) @@ -8823,11 +10101,14 @@ def _stop_spinner(): "api_calls": api_call_count, "completed": False, "partial": True, - "error": "Response truncated due to output length limit" + "error": "Response truncated due to output length limit", } else: # First message was truncated - mark as failed - self._vprint(f"{self.log_prefix}❌ First response truncated - cannot recover", force=True) + self._vprint( + f"{self.log_prefix}❌ First response truncated - cannot recover", + force=True, + ) self._persist_session(messages, conversation_history) return { "final_response": None, @@ -8835,11 +10116,11 @@ def _stop_spinner(): "api_calls": api_call_count, "completed": False, "failed": True, - "error": "First response truncated due to output length limit" + "error": "First response truncated due to output length limit", } - + # Track actual token usage from response for context management - if hasattr(response, 'usage') and response.usage: + if hasattr(response, "usage") and response.usage: canonical_usage = normalize_usage( response.usage, provider=self.provider, @@ -8860,9 +10141,15 @@ def _stop_spinner(): # from the error message), not guessed probe tiers. if getattr(self.context_compressor, "_context_probed", False): ctx = self.context_compressor.context_length - if getattr(self.context_compressor, "_context_probe_persistable", False): + if getattr( + self.context_compressor, + "_context_probe_persistable", + False, + ): save_context_length(self.model, self.base_url, ctx) - self._safe_print(f"{self.log_prefix}💾 Cached context length: {ctx:,} tokens for {self.model}") + self._safe_print( + f"{self.log_prefix}💾 Cached context length: {ctx:,} tokens for {self.model}" + ) self.context_compressor._context_probed = False self.context_compressor._context_probe_persistable = False @@ -8872,19 +10159,30 @@ def _stop_spinner(): self.session_api_calls += 1 self.session_input_tokens += canonical_usage.input_tokens self.session_output_tokens += canonical_usage.output_tokens - self.session_cache_read_tokens += canonical_usage.cache_read_tokens - self.session_cache_write_tokens += canonical_usage.cache_write_tokens - self.session_reasoning_tokens += canonical_usage.reasoning_tokens + self.session_cache_read_tokens += ( + canonical_usage.cache_read_tokens + ) + self.session_cache_write_tokens += ( + canonical_usage.cache_write_tokens + ) + self.session_reasoning_tokens += ( + canonical_usage.reasoning_tokens + ) # Log API call details for debugging/observability _cache_pct = "" if canonical_usage.cache_read_tokens and prompt_tokens: - _cache_pct = f" cache={canonical_usage.cache_read_tokens}/{prompt_tokens} ({100*canonical_usage.cache_read_tokens/prompt_tokens:.0f}%)" + _cache_pct = f" cache={canonical_usage.cache_read_tokens}/{prompt_tokens} ({100 * canonical_usage.cache_read_tokens / prompt_tokens:.0f}%)" logger.info( "API call #%d: model=%s provider=%s in=%d out=%d total=%d latency=%.1fs%s", - self.session_api_calls, self.model, self.provider or "unknown", - prompt_tokens, completion_tokens, total_tokens, - api_duration, _cache_pct, + self.session_api_calls, + self.model, + self.provider or "unknown", + prompt_tokens, + completion_tokens, + total_tokens, + api_duration, + _cache_pct, ) cost_result = estimate_usage_cost( @@ -8895,7 +10193,9 @@ def _stop_spinner(): api_key=getattr(self, "api_key", ""), ) if cost_result.amount_usd is not None: - self.session_estimated_cost_usd += float(cost_result.amount_usd) + self.session_estimated_cost_usd += float( + cost_result.amount_usd + ) self.session_cost_status = cost_result.status self.session_cost_source = cost_result.source @@ -8916,37 +10216,63 @@ def _stop_spinner(): cache_write_tokens=canonical_usage.cache_write_tokens, reasoning_tokens=canonical_usage.reasoning_tokens, estimated_cost_usd=float(cost_result.amount_usd) - if cost_result.amount_usd is not None else None, + if cost_result.amount_usd is not None + else None, cost_status=cost_result.status, cost_source=cost_result.source, billing_provider=self.provider, billing_base_url=self.base_url, billing_mode="subscription_included" - if cost_result.status == "included" else None, + if cost_result.status == "included" + else None, model=self.model, ) except Exception: pass # never block the agent loop - + if self.verbose_logging: - logging.debug(f"Token usage: prompt={usage_dict['prompt_tokens']:,}, completion={usage_dict['completion_tokens']:,}, total={usage_dict['total_tokens']:,}") - + logging.debug( + f"Token usage: prompt={usage_dict['prompt_tokens']:,}, completion={usage_dict['completion_tokens']:,}, total={usage_dict['total_tokens']:,}" + ) + # Log cache hit stats when prompt caching is active if self._use_prompt_caching: if self.api_mode == "anthropic_messages": # Anthropic uses cache_read_input_tokens / cache_creation_input_tokens - cached = getattr(response.usage, 'cache_read_input_tokens', 0) or 0 - written = getattr(response.usage, 'cache_creation_input_tokens', 0) or 0 + cached = ( + getattr( + response.usage, "cache_read_input_tokens", 0 + ) + or 0 + ) + written = ( + getattr( + response.usage, "cache_creation_input_tokens", 0 + ) + or 0 + ) else: # OpenRouter uses prompt_tokens_details.cached_tokens - details = getattr(response.usage, 'prompt_tokens_details', None) - cached = getattr(details, 'cached_tokens', 0) or 0 if details else 0 - written = getattr(details, 'cache_write_tokens', 0) or 0 if details else 0 + details = getattr( + response.usage, "prompt_tokens_details", None + ) + cached = ( + getattr(details, "cached_tokens", 0) or 0 + if details + else 0 + ) + written = ( + getattr(details, "cache_write_tokens", 0) or 0 + if details + else 0 + ) prompt = usage_dict["prompt_tokens"] hit_pct = (cached / prompt * 100) if prompt > 0 else 0 if not self.quiet_mode: - self._vprint(f"{self.log_prefix} 💾 Cache: {cached:,}/{prompt:,} tokens ({hit_pct:.0f}% hit, {written:,} written)") - + self._vprint( + f"{self.log_prefix} 💾 Cache: {cached:,}/{prompt:,} tokens ({hit_pct:.0f}% hit, {written:,} written)" + ) + has_retried_429 = False # Reset on success self._touch_activity(f"API call #{api_call_count} completed") break # Success, exit retry loop @@ -8958,7 +10284,9 @@ def _stop_spinner(): if self.thinking_callback: self.thinking_callback("") api_elapsed = time.time() - api_start_time - self._vprint(f"{self.log_prefix}⚡ Interrupted during API call.", force=True) + self._vprint( + f"{self.log_prefix}⚡ Interrupted during API call.", force=True + ) self._persist_session(messages, conversation_history) interrupted = True final_response = f"Operation interrupted: waiting for model response ({api_elapsed:.1f}s elapsed)." @@ -8983,7 +10311,10 @@ def _stop_spinner(): # first to strip surrogates, then once more for pure # ASCII-only locale sanitization if needed. # ----------------------------------------------------------- - if isinstance(api_error, UnicodeEncodeError) and getattr(self, '_unicode_sanitization_passes', 0) < 2: + if ( + isinstance(api_error, UnicodeEncodeError) + and getattr(self, "_unicode_sanitization_passes", 0) < 2 + ): _err_str = str(api_error).lower() _is_ascii_codec = "'ascii'" in _err_str or "ascii" in _err_str _surrogates_found = _sanitize_messages_surrogates(messages) @@ -9001,8 +10332,12 @@ def _stop_spinner(): # non-ASCII content from messages/tool schemas and retry. _messages_sanitized = _sanitize_messages_non_ascii(messages) _prefill_sanitized = False - if isinstance(getattr(self, "prefill_messages", None), list): - _prefill_sanitized = _sanitize_messages_non_ascii(self.prefill_messages) + if isinstance( + getattr(self, "prefill_messages", None), list + ): + _prefill_sanitized = _sanitize_messages_non_ascii( + self.prefill_messages + ) _tools_sanitized = False if isinstance(getattr(self, "tools", None), list): @@ -9010,25 +10345,35 @@ def _stop_spinner(): _system_sanitized = False if isinstance(active_system_prompt, str): - _sanitized_system = _strip_non_ascii(active_system_prompt) + _sanitized_system = _strip_non_ascii( + active_system_prompt + ) if _sanitized_system != active_system_prompt: active_system_prompt = _sanitized_system self._cached_system_prompt = _sanitized_system _system_sanitized = True - if isinstance(getattr(self, "ephemeral_system_prompt", None), str): - _sanitized_ephemeral = _strip_non_ascii(self.ephemeral_system_prompt) - if _sanitized_ephemeral != self.ephemeral_system_prompt: + if isinstance( + getattr(self, "ephemeral_system_prompt", None), str + ): + _sanitized_ephemeral = _strip_non_ascii( + self.ephemeral_system_prompt + ) + if _sanitized_ephemeral != self.ephemeral_system_prompt: self.ephemeral_system_prompt = _sanitized_ephemeral _system_sanitized = True _headers_sanitized = False _default_headers = ( self._client_kwargs.get("default_headers") - if isinstance(getattr(self, "_client_kwargs", None), dict) + if isinstance( + getattr(self, "_client_kwargs", None), dict + ) else None ) if isinstance(_default_headers, dict): - _headers_sanitized = _sanitize_structure_non_ascii(_default_headers) + _headers_sanitized = _sanitize_structure_non_ascii( + _default_headers + ) # Sanitize the API key — non-ASCII characters in # credentials (e.g. ʋ instead of v from a bad @@ -9042,12 +10387,16 @@ def _stop_spinner(): _clean_key = _strip_non_ascii(_raw_key) if _clean_key != _raw_key: self.api_key = _clean_key - if isinstance(getattr(self, "_client_kwargs", None), dict): + if isinstance( + getattr(self, "_client_kwargs", None), dict + ): self._client_kwargs["api_key"] = _clean_key # Also update the live client — it holds its # own copy of api_key which auth_headers reads # dynamically on every request. - if getattr(self, "client", None) is not None and hasattr(self.client, "api_key"): + if getattr( + self, "client", None + ) is not None and hasattr(self.client, "api_key"): self.client.api_key = _clean_key _credential_sanitized = True self._vprint( @@ -9079,7 +10428,11 @@ def _stop_spinner(): # ── Classify the error for structured recovery decisions ── _compressor = getattr(self, "context_compressor", None) - _ctx_len = getattr(_compressor, "context_length", 200000) if _compressor else 200000 + _ctx_len = ( + getattr(_compressor, "context_length", 200000) + if _compressor + else 200000 + ) classified = classify_api_error( api_error, provider=getattr(self, "provider", "") or "", @@ -9090,16 +10443,21 @@ def _stop_spinner(): ) logger.debug( "Error classified: reason=%s status=%s retryable=%s compress=%s rotate=%s fallback=%s", - classified.reason.value, classified.status_code, - classified.retryable, classified.should_compress, - classified.should_rotate_credential, classified.should_fallback, + classified.reason.value, + classified.status_code, + classified.retryable, + classified.should_compress, + classified.should_rotate_credential, + classified.should_fallback, ) - recovered_with_pool, has_retried_429 = self._recover_with_credential_pool( - status_code=status_code, - has_retried_429=has_retried_429, - classified_reason=classified.reason, - error_context=error_context, + recovered_with_pool, has_retried_429 = ( + self._recover_with_credential_pool( + status_code=status_code, + has_retried_429=has_retried_429, + classified_reason=classified.reason, + error_context=error_context, + ) ) if recovered_with_pool: continue @@ -9111,7 +10469,9 @@ def _stop_spinner(): ): codex_auth_retry_attempted = True if self._try_refresh_codex_client_credentials(force=True): - self._vprint(f"{self.log_prefix}🔐 Codex auth refreshed after 401. Retrying request...") + self._vprint( + f"{self.log_prefix}🔐 Codex auth refreshed after 401. Retrying request..." + ) continue if ( self.api_mode == "chat_completions" @@ -9121,34 +10481,62 @@ def _stop_spinner(): ): nous_auth_retry_attempted = True if self._try_refresh_nous_client_credentials(force=True): - print(f"{self.log_prefix}🔐 Nous agent key refreshed after 401. Retrying request...") + print( + f"{self.log_prefix}🔐 Nous agent key refreshed after 401. Retrying request..." + ) continue if ( self.api_mode == "anthropic_messages" and status_code == 401 - and hasattr(self, '_anthropic_api_key') + and hasattr(self, "_anthropic_api_key") and not anthropic_auth_retry_attempted ): anthropic_auth_retry_attempted = True from agent.anthropic_adapter import _is_oauth_token + if self._try_refresh_anthropic_client_credentials(): - print(f"{self.log_prefix}🔐 Anthropic credentials refreshed after 401. Retrying request...") + print( + f"{self.log_prefix}🔐 Anthropic credentials refreshed after 401. Retrying request..." + ) continue # Credential refresh didn't help — show diagnostic info key = self._anthropic_api_key - auth_method = "Bearer (OAuth/setup-token)" if _is_oauth_token(key) else "x-api-key (API key)" - print(f"{self.log_prefix}🔐 Anthropic 401 — authentication failed.") + auth_method = ( + "Bearer (OAuth/setup-token)" + if _is_oauth_token(key) + else "x-api-key (API key)" + ) + print( + f"{self.log_prefix}🔐 Anthropic 401 — authentication failed." + ) print(f"{self.log_prefix} Auth method: {auth_method}") - print(f"{self.log_prefix} Token prefix: {key[:12]}..." if key and len(key) > 12 else f"{self.log_prefix} Token: (empty or short)") + print( + f"{self.log_prefix} Token prefix: {key[:12]}..." + if key and len(key) > 12 + else f"{self.log_prefix} Token: (empty or short)" + ) print(f"{self.log_prefix} Troubleshooting:") from hermes_constants import display_hermes_home as _dhh_fn + _dhh = _dhh_fn() - print(f"{self.log_prefix} • Check ANTHROPIC_TOKEN in {_dhh}/.env for Hermes-managed OAuth/setup tokens") - print(f"{self.log_prefix} • Check ANTHROPIC_API_KEY in {_dhh}/.env for API keys or legacy token values") - print(f"{self.log_prefix} • For API keys: verify at https://console.anthropic.com/settings/keys") - print(f"{self.log_prefix} • For Claude Code: run 'claude /login' to refresh, then retry") - print(f"{self.log_prefix} • Legacy cleanup: hermes config set ANTHROPIC_TOKEN \"\"") - print(f"{self.log_prefix} • Clear stale keys: hermes config set ANTHROPIC_API_KEY \"\"") + print( + f"{self.log_prefix} • Check ANTHROPIC_TOKEN in {_dhh}/.env for Hermes-managed OAuth/setup tokens" + ) + print( + f"{self.log_prefix} • Check ANTHROPIC_API_KEY in {_dhh}/.env for API keys or legacy token values" + ) + print( + f"{self.log_prefix} • For API keys: verify at https://console.anthropic.com/settings/keys" + ) + print( + f"{self.log_prefix} • For Claude Code: run 'claude /login' to refresh, then retry" + ) + print( + f'{self.log_prefix} • Legacy cleanup: hermes config set ANTHROPIC_TOKEN ""' + ) + print( + f'{self.log_prefix} • Clear stale keys: hermes config set ANTHROPIC_API_KEY ""' + ) # ── Thinking block signature recovery ───────────────── # Anthropic signs thinking blocks against the full turn @@ -9173,7 +10561,8 @@ def _stop_spinner(): logging.warning( "%sThinking block signature recovery: stripped " "reasoning_details from %d messages", - self.log_prefix, len(messages), + self.log_prefix, + len(messages), ) continue @@ -9182,7 +10571,7 @@ def _stop_spinner(): self._touch_activity( f"API error recovery (attempt {retry_count}/{max_retries})" ) - + error_type = type(api_error).__name__ error_msg = str(api_error).lower() _error_summary = self._summarize_api_error(api_error) @@ -9199,25 +10588,37 @@ def _stop_spinner(): _base = getattr(self, "base_url", "unknown") _model = getattr(self, "model", "unknown") _status_code_str = f" [HTTP {status_code}]" if status_code else "" - self._vprint(f"{self.log_prefix}⚠️ API call failed (attempt {retry_count}/{max_retries}): {error_type}{_status_code_str}", force=True) - self._vprint(f"{self.log_prefix} 🔌 Provider: {_provider} Model: {_model}", force=True) - self._vprint(f"{self.log_prefix} 🌐 Endpoint: {_base}", force=True) - self._vprint(f"{self.log_prefix} 📝 Error: {_error_summary}", force=True) + self._vprint( + f"{self.log_prefix}⚠️ API call failed (attempt {retry_count}/{max_retries}): {error_type}{_status_code_str}", + force=True, + ) + self._vprint( + f"{self.log_prefix} 🔌 Provider: {_provider} Model: {_model}", + force=True, + ) + self._vprint( + f"{self.log_prefix} 🌐 Endpoint: {_base}", force=True + ) + self._vprint( + f"{self.log_prefix} 📝 Error: {_error_summary}", force=True + ) if status_code and status_code < 500: _err_body = getattr(api_error, "body", None) _err_body_str = str(_err_body)[:300] if _err_body else None if _err_body_str: - self._vprint(f"{self.log_prefix} 📋 Details: {_err_body_str}", force=True) - self._vprint(f"{self.log_prefix} ⏱️ Elapsed: {elapsed_time:.2f}s Context: {len(api_messages)} msgs, ~{approx_tokens:,} tokens") + self._vprint( + f"{self.log_prefix} 📋 Details: {_err_body_str}", + force=True, + ) + self._vprint( + f"{self.log_prefix} ⏱️ Elapsed: {elapsed_time:.2f}s Context: {len(api_messages)} msgs, ~{approx_tokens:,} tokens" + ) # Actionable hint for OpenRouter "no tool endpoints" error. # This fires regardless of whether fallback succeeds — the # user needs to know WHY their model failed so they can fix # their provider routing, not just silently fall back. - if ( - self._is_openrouter_url() - and "support tool use" in error_msg - ): + if self._is_openrouter_url() and "support tool use" in error_msg: self._vprint( f"{self.log_prefix} 💡 No OpenRouter providers for {_model} support tool calling with your current settings.", force=True, @@ -9238,7 +10639,10 @@ def _stop_spinner(): # Check for interrupt before deciding to retry if self._interrupt_requested: - self._vprint(f"{self.log_prefix}⚡ Interrupt detected during error handling, aborting retries.", force=True) + self._vprint( + f"{self.log_prefix}⚡ Interrupt detected during error handling, aborting retries.", + force=True, + ) self._persist_session(messages, conversation_history) self.clear_interrupt() return { @@ -9248,7 +10652,7 @@ def _stop_spinner(): "completed": False, "interrupted": True, } - + # Check for 413 payload-too-large BEFORE generic 4xx handler. # A 413 is a payload-size error — the correct response is to # compress history and retry, not abort immediately. @@ -9293,7 +10697,8 @@ def _stop_spinner(): if compression_attempts <= max_compression_attempts: original_len = len(messages) messages, active_system_prompt = self._compress_context( - messages, system_message, + messages, + system_message, approx_tokens=approx_tokens, task_id=effective_task_id, ) @@ -9320,7 +10725,9 @@ def _stop_spinner(): FailoverReason.rate_limit, FailoverReason.billing, ) - if is_rate_limited and self._fallback_index < len(self._fallback_chain): + if is_rate_limited and self._fallback_index < len( + self._fallback_chain + ): # Don't eagerly fallback if credential pool rotation may # still recover. The pool's retry-then-rotate cycle needs # at least one more attempt to fire — jumping to a fallback @@ -9328,7 +10735,9 @@ def _stop_spinner(): pool = self._credential_pool pool_may_recover = pool is not None and pool.has_available() if not pool_may_recover: - self._emit_status("⚠️ Rate limited — switching to fallback provider...") + self._emit_status( + "⚠️ Rate limited — switching to fallback provider..." + ) if self._try_activate_fallback(): retry_count = 0 compression_attempts = 0 @@ -9342,9 +10751,17 @@ def _stop_spinner(): if is_payload_too_large: compression_attempts += 1 if compression_attempts > max_compression_attempts: - self._vprint(f"{self.log_prefix}❌ Max compression attempts ({max_compression_attempts}) reached for payload-too-large error.", force=True) - self._vprint(f"{self.log_prefix} 💡 Try /new to start a fresh conversation, or /compress to retry compression.", force=True) - logging.error(f"{self.log_prefix}413 compression failed after {max_compression_attempts} attempts.") + self._vprint( + f"{self.log_prefix}❌ Max compression attempts ({max_compression_attempts}) reached for payload-too-large error.", + force=True, + ) + self._vprint( + f"{self.log_prefix} 💡 Try /new to start a fresh conversation, or /compress to retry compression.", + force=True, + ) + logging.error( + f"{self.log_prefix}413 compression failed after {max_compression_attempts} attempts." + ) self._persist_session(messages, conversation_history) return { "messages": messages, @@ -9355,11 +10772,15 @@ def _stop_spinner(): "failed": True, "compression_exhausted": True, } - self._emit_status(f"⚠️ Request payload too large (413) — compression attempt {compression_attempts}/{max_compression_attempts}...") + self._emit_status( + f"⚠️ Request payload too large (413) — compression attempt {compression_attempts}/{max_compression_attempts}..." + ) original_len = len(messages) messages, active_system_prompt = self._compress_context( - messages, system_message, approx_tokens=approx_tokens, + messages, + system_message, + approx_tokens=approx_tokens, task_id=effective_task_id, ) # Compression created a new session — clear history @@ -9368,14 +10789,24 @@ def _stop_spinner(): conversation_history = None if len(messages) < original_len: - self._emit_status(f"🗜️ Compressed {original_len} → {len(messages)} messages, retrying...") + self._emit_status( + f"🗜️ Compressed {original_len} → {len(messages)} messages, retrying..." + ) time.sleep(2) # Brief pause between compression retries restart_with_compressed_messages = True break else: - self._vprint(f"{self.log_prefix}❌ Payload too large and cannot compress further.", force=True) - self._vprint(f"{self.log_prefix} 💡 Try /new to start a fresh conversation, or /compress to retry compression.", force=True) - logging.error(f"{self.log_prefix}413 payload too large. Cannot compress further.") + self._vprint( + f"{self.log_prefix}❌ Payload too large and cannot compress further.", + force=True, + ) + self._vprint( + f"{self.log_prefix} 💡 Try /new to start a fresh conversation, or /compress to retry compression.", + force=True, + ) + logging.error( + f"{self.log_prefix}413 payload too large. Cannot compress further." + ) self._persist_session(messages, conversation_history) return { "messages": messages, @@ -9409,7 +10840,9 @@ def _stop_spinner(): # # Note: max_tokens = output token cap (one response). # context_length = total window (input + output combined). - available_out = parse_available_output_tokens_from_error(error_msg) + available_out = parse_available_output_tokens_from_error( + error_msg + ) if available_out is not None: # Error is purely about the output cap being too large. # Cap output to the available space and retry without @@ -9426,9 +10859,17 @@ def _stop_spinner(): # loop forever if the error keeps recurring. compression_attempts += 1 if compression_attempts > max_compression_attempts: - self._vprint(f"{self.log_prefix}❌ Max compression attempts ({max_compression_attempts}) reached.", force=True) - self._vprint(f"{self.log_prefix} 💡 Try /new to start a fresh conversation, or /compress to retry compression.", force=True) - logging.error(f"{self.log_prefix}Context compression failed after {max_compression_attempts} attempts.") + self._vprint( + f"{self.log_prefix}❌ Max compression attempts ({max_compression_attempts}) reached.", + force=True, + ) + self._vprint( + f"{self.log_prefix} 💡 Try /new to start a fresh conversation, or /compress to retry compression.", + force=True, + ) + logging.error( + f"{self.log_prefix}Context compression failed after {max_compression_attempts} attempts." + ) self._persist_session(messages, conversation_history) return { "messages": messages, @@ -9447,7 +10888,10 @@ def _stop_spinner(): parsed_limit = parse_context_limit_from_error(error_msg) if parsed_limit and parsed_limit < old_ctx: new_ctx = parsed_limit - self._vprint(f"{self.log_prefix}⚠️ Context limit detected from API: {new_ctx:,} tokens (was {old_ctx:,})", force=True) + self._vprint( + f"{self.log_prefix}⚠️ Context limit detected from API: {new_ctx:,} tokens (was {old_ctx:,})", + force=True, + ) else: # Step down to the next probe tier new_ctx = get_next_probe_tier(old_ctx) @@ -9472,15 +10916,29 @@ def _stop_spinner(): compressor._context_probe_persistable = bool( parsed_limit and parsed_limit == new_ctx ) - self._vprint(f"{self.log_prefix}⚠️ Context length exceeded — stepping down: {old_ctx:,} → {new_ctx:,} tokens", force=True) + self._vprint( + f"{self.log_prefix}⚠️ Context length exceeded — stepping down: {old_ctx:,} → {new_ctx:,} tokens", + force=True, + ) else: - self._vprint(f"{self.log_prefix}⚠️ Context length exceeded at minimum tier — attempting compression...", force=True) + self._vprint( + f"{self.log_prefix}⚠️ Context length exceeded at minimum tier — attempting compression...", + force=True, + ) compression_attempts += 1 if compression_attempts > max_compression_attempts: - self._vprint(f"{self.log_prefix}❌ Max compression attempts ({max_compression_attempts}) reached.", force=True) - self._vprint(f"{self.log_prefix} 💡 Try /new to start a fresh conversation, or /compress to retry compression.", force=True) - logging.error(f"{self.log_prefix}Context compression failed after {max_compression_attempts} attempts.") + self._vprint( + f"{self.log_prefix}❌ Max compression attempts ({max_compression_attempts}) reached.", + force=True, + ) + self._vprint( + f"{self.log_prefix} 💡 Try /new to start a fresh conversation, or /compress to retry compression.", + force=True, + ) + logging.error( + f"{self.log_prefix}Context compression failed after {max_compression_attempts} attempts." + ) self._persist_session(messages, conversation_history) return { "messages": messages, @@ -9491,11 +10949,15 @@ def _stop_spinner(): "failed": True, "compression_exhausted": True, } - self._emit_status(f"🗜️ Context too large (~{approx_tokens:,} tokens) — compressing ({compression_attempts}/{max_compression_attempts})...") + self._emit_status( + f"🗜️ Context too large (~{approx_tokens:,} tokens) — compressing ({compression_attempts}/{max_compression_attempts})..." + ) original_len = len(messages) messages, active_system_prompt = self._compress_context( - messages, system_message, approx_tokens=approx_tokens, + messages, + system_message, + approx_tokens=approx_tokens, task_id=effective_task_id, ) # Compression created a new session — clear history @@ -9503,17 +10965,31 @@ def _stop_spinner(): # messages to the new session, not skipping them. conversation_history = None - if len(messages) < original_len or new_ctx and new_ctx < old_ctx: + if ( + len(messages) < original_len + or new_ctx + and new_ctx < old_ctx + ): if len(messages) < original_len: - self._emit_status(f"🗜️ Compressed {original_len} → {len(messages)} messages, retrying...") + self._emit_status( + f"🗜️ Compressed {original_len} → {len(messages)} messages, retrying..." + ) time.sleep(2) # Brief pause between compression retries restart_with_compressed_messages = True break else: # Can't compress further and already at minimum tier - self._vprint(f"{self.log_prefix}❌ Context length exceeded and cannot compress further.", force=True) - self._vprint(f"{self.log_prefix} 💡 The conversation has accumulated too much content. Try /new to start fresh, or /compress to manually trigger compression.", force=True) - logging.error(f"{self.log_prefix}Context length exceeded: {approx_tokens:,} tokens. Cannot compress further.") + self._vprint( + f"{self.log_prefix}❌ Context length exceeded and cannot compress further.", + force=True, + ) + self._vprint( + f"{self.log_prefix} 💡 The conversation has accumulated too much content. Try /new to start fresh, or /compress to manually trigger compression.", + force=True, + ) + logging.error( + f"{self.log_prefix}Context length exceeded: {approx_tokens:,} tokens. Cannot compress further." + ) self._persist_session(messages, conversation_history) return { "messages": messages, @@ -9529,16 +11005,16 @@ def _stop_spinner(): # already accounts for 413, 429, 529 (transient), context # overflow, and generic-400 heuristics. Local validation # errors (ValueError, TypeError) are programming bugs. - is_local_validation_error = ( - isinstance(api_error, (ValueError, TypeError)) - and not isinstance(api_error, UnicodeEncodeError) - ) + is_local_validation_error = isinstance( + api_error, (ValueError, TypeError) + ) and not isinstance(api_error, UnicodeEncodeError) is_client_error = ( is_local_validation_error or ( not classified.retryable and not classified.should_compress - and classified.reason not in ( + and classified.reason + not in ( FailoverReason.rate_limit, FailoverReason.billing, FailoverReason.overloaded, @@ -9553,7 +11029,9 @@ def _stop_spinner(): if is_client_error: # Try fallback before aborting — a different provider # may not have the same issue (rate limit, auth, etc.) - self._emit_status(f"⚠️ Non-retryable error (HTTP {status_code}) — trying fallback...") + self._emit_status( + f"⚠️ Non-retryable error (HTTP {status_code}) — trying fallback..." + ) if self._try_activate_fallback(): retry_count = 0 compression_attempts = 0 @@ -9561,37 +11039,81 @@ def _stop_spinner(): continue if api_kwargs is not None: self._dump_api_request_debug( - api_kwargs, reason="non_retryable_client_error", error=api_error, + api_kwargs, + reason="non_retryable_client_error", + error=api_error, ) self._emit_status( f"❌ Non-retryable error (HTTP {status_code}): " f"{self._summarize_api_error(api_error)}" ) - self._vprint(f"{self.log_prefix}❌ Non-retryable client error (HTTP {status_code}). Aborting.", force=True) - self._vprint(f"{self.log_prefix} 🔌 Provider: {_provider} Model: {_model}", force=True) - self._vprint(f"{self.log_prefix} 🌐 Endpoint: {_base}", force=True) + self._vprint( + f"{self.log_prefix}❌ Non-retryable client error (HTTP {status_code}). Aborting.", + force=True, + ) + self._vprint( + f"{self.log_prefix} 🔌 Provider: {_provider} Model: {_model}", + force=True, + ) + self._vprint( + f"{self.log_prefix} 🌐 Endpoint: {_base}", force=True + ) # Actionable guidance for common auth errors - if classified.is_auth or classified.reason == FailoverReason.billing: + if ( + classified.is_auth + or classified.reason == FailoverReason.billing + ): if _provider == "openai-codex" and status_code == 401: - self._vprint(f"{self.log_prefix} 💡 Codex OAuth token was rejected (HTTP 401). Your token may have been", force=True) - self._vprint(f"{self.log_prefix} refreshed by another client (Codex CLI, VS Code). To fix:", force=True) - self._vprint(f"{self.log_prefix} 1. Run `codex` in your terminal to generate fresh tokens.", force=True) - self._vprint(f"{self.log_prefix} 2. Then run `hermes auth` to re-authenticate.", force=True) + self._vprint( + f"{self.log_prefix} 💡 Codex OAuth token was rejected (HTTP 401). Your token may have been", + force=True, + ) + self._vprint( + f"{self.log_prefix} refreshed by another client (Codex CLI, VS Code). To fix:", + force=True, + ) + self._vprint( + f"{self.log_prefix} 1. Run `codex` in your terminal to generate fresh tokens.", + force=True, + ) + self._vprint( + f"{self.log_prefix} 2. Then run `hermes auth` to re-authenticate.", + force=True, + ) else: - self._vprint(f"{self.log_prefix} 💡 Your API key was rejected by the provider. Check:", force=True) - self._vprint(f"{self.log_prefix} • Is the key valid? Run: hermes setup", force=True) - self._vprint(f"{self.log_prefix} • Does your account have access to {_model}?", force=True) + self._vprint( + f"{self.log_prefix} 💡 Your API key was rejected by the provider. Check:", + force=True, + ) + self._vprint( + f"{self.log_prefix} • Is the key valid? Run: hermes setup", + force=True, + ) + self._vprint( + f"{self.log_prefix} • Does your account have access to {_model}?", + force=True, + ) if "openrouter" in str(_base).lower(): - self._vprint(f"{self.log_prefix} • Check credits: https://openrouter.ai/settings/credits", force=True) + self._vprint( + f"{self.log_prefix} • Check credits: https://openrouter.ai/settings/credits", + force=True, + ) else: - self._vprint(f"{self.log_prefix} 💡 This type of error won't be fixed by retrying.", force=True) - logging.error(f"{self.log_prefix}Non-retryable client error: {api_error}") + self._vprint( + f"{self.log_prefix} 💡 This type of error won't be fixed by retrying.", + force=True, + ) + logging.error( + f"{self.log_prefix}Non-retryable client error: {api_error}" + ) # Skip session persistence when the error is likely # context-overflow related (status 400 + large session). # Persisting the failed user message would make the # session even larger, causing the same failure on the # next attempt. (#1630) - if status_code == 400 and (approx_tokens > 50000 or len(api_messages) > 80): + if status_code == 400 and ( + approx_tokens > 50000 or len(api_messages) > 80 + ): self._vprint( f"{self.log_prefix}⚠️ Skipping session persistence " f"for large failed session to prevent growth loop.", @@ -9613,14 +11135,21 @@ def _stop_spinner(): # client once for transient transport errors (stale # connection pool, TCP reset). Only attempted once # per API call block. - if not primary_recovery_attempted and self._try_recover_primary_transport( - api_error, retry_count=retry_count, max_retries=max_retries, + if ( + not primary_recovery_attempted + and self._try_recover_primary_transport( + api_error, + retry_count=retry_count, + max_retries=max_retries, + ) ): primary_recovery_attempted = True retry_count = 0 continue # Try fallback before giving up entirely - self._emit_status(f"⚠️ Max retries ({max_retries}) exhausted — trying fallback...") + self._emit_status( + f"⚠️ Max retries ({max_retries}) exhausted — trying fallback..." + ) if self._try_activate_fallback(): retry_count = 0 compression_attempts = 0 @@ -9628,23 +11157,35 @@ def _stop_spinner(): continue _final_summary = self._summarize_api_error(api_error) if is_rate_limited: - self._emit_status(f"❌ Rate limited after {max_retries} retries — {_final_summary}") + self._emit_status( + f"❌ Rate limited after {max_retries} retries — {_final_summary}" + ) else: - self._emit_status(f"❌ API failed after {max_retries} retries — {_final_summary}") - self._vprint(f"{self.log_prefix} 💀 Final error: {_final_summary}", force=True) + self._emit_status( + f"❌ API failed after {max_retries} retries — {_final_summary}" + ) + self._vprint( + f"{self.log_prefix} 💀 Final error: {_final_summary}", + force=True, + ) # Detect SSE stream-drop pattern (e.g. "Network # connection lost") and surface actionable guidance. # This typically happens when the model generates a # very large tool call (write_file with huge content) # and the proxy/CDN drops the stream mid-response. - _is_stream_drop = ( - not getattr(api_error, "status_code", None) - and any(p in error_msg for p in ( - "connection lost", "connection reset", - "connection closed", "network connection", - "network error", "terminated", - )) + _is_stream_drop = not getattr( + api_error, "status_code", None + ) and any( + p in error_msg + for p in ( + "connection lost", + "connection reset", + "connection closed", + "network connection", + "network error", + "terminated", + ) ) if _is_stream_drop: self._vprint( @@ -9664,12 +11205,19 @@ def _stop_spinner(): logging.error( "%sAPI call failed after %s retries. %s | provider=%s model=%s msgs=%s tokens=~%s", - self.log_prefix, max_retries, _final_summary, - _provider, _model, len(api_messages), f"{approx_tokens:,}", + self.log_prefix, + max_retries, + _final_summary, + _provider, + _model, + len(api_messages), + f"{approx_tokens:,}", ) if api_kwargs is not None: self._dump_api_request_debug( - api_kwargs, reason="max_retries_exhausted", error=api_error, + api_kwargs, + reason="max_retries_exhausted", + error=api_error, ) self._persist_session(messages, conversation_history) _final_response = f"API call failed after {max_retries} retries: {_final_summary}" @@ -9694,19 +11242,35 @@ def _stop_spinner(): # For rate limits, respect the Retry-After header if present _retry_after = None if is_rate_limited: - _resp_headers = getattr(getattr(api_error, "response", None), "headers", None) + _resp_headers = getattr( + getattr(api_error, "response", None), "headers", None + ) if _resp_headers and hasattr(_resp_headers, "get"): - _ra_raw = _resp_headers.get("retry-after") or _resp_headers.get("Retry-After") + _ra_raw = _resp_headers.get( + "retry-after" + ) or _resp_headers.get("Retry-After") if _ra_raw: try: - _retry_after = min(int(_ra_raw), 120) # Cap at 2 minutes + _retry_after = min( + int(_ra_raw), 120 + ) # Cap at 2 minutes except (TypeError, ValueError): pass - wait_time = _retry_after if _retry_after else jittered_backoff(retry_count, base_delay=2.0, max_delay=60.0) + wait_time = ( + _retry_after + if _retry_after + else jittered_backoff( + retry_count, base_delay=2.0, max_delay=60.0 + ) + ) if is_rate_limited: - self._emit_status(f"⏱️ Rate limit reached. Waiting {wait_time}s before retry (attempt {retry_count + 1}/{max_retries})...") + self._emit_status( + f"⏱️ Rate limit reached. Waiting {wait_time}s before retry (attempt {retry_count + 1}/{max_retries})..." + ) else: - self._emit_status(f"⏳ Retrying in {wait_time}s (attempt {retry_count}/{max_retries})...") + self._emit_status( + f"⏳ Retrying in {wait_time}s (attempt {retry_count}/{max_retries})..." + ) logger.warning( "Retrying API call in %ss (attempt %s/%s) %s error=%s", wait_time, @@ -9721,7 +11285,10 @@ def _stop_spinner(): _backoff_touch_counter = 0 while time.time() < sleep_end: if self._interrupt_requested: - self._vprint(f"{self.log_prefix}⚡ Interrupt detected during retry wait, aborting.", force=True) + self._vprint( + f"{self.log_prefix}⚡ Interrupt detected during retry wait, aborting.", + force=True, + ) self._persist_session(messages, conversation_history) self.clear_interrupt() return { @@ -9740,7 +11307,7 @@ def _stop_spinner(): f"error retry backoff ({retry_count}/{max_retries}), " f"{int(sleep_end - time.time())}s remaining" ) - + # If the API call was interrupted, skip response processing if interrupted: _turn_exit_reason = "interrupted_during_api_call" @@ -9764,28 +11331,39 @@ def _stop_spinner(): # the `response` variable is still None. Break out cleanly. if response is None: _turn_exit_reason = "all_retries_exhausted_no_response" - print(f"{self.log_prefix}❌ All API retries exhausted with no successful response.") + print( + f"{self.log_prefix}❌ All API retries exhausted with no successful response." + ) self._persist_session(messages, conversation_history) break try: if self.api_mode == "codex_responses": - assistant_message, finish_reason = self._normalize_codex_response(response) + assistant_message, finish_reason = self._normalize_codex_response( + response + ) elif self.api_mode == "anthropic_messages": from agent.anthropic_adapter import normalize_anthropic_response + assistant_message, finish_reason = normalize_anthropic_response( response, strip_tool_prefix=self._is_anthropic_oauth ) else: assistant_message = response.choices[0].message - + # Normalize content to string — some OpenAI-compatible servers # (llama-server, etc.) return content as a dict or list instead # of a plain string, which crashes downstream .strip() calls. - if assistant_message.content is not None and not isinstance(assistant_message.content, str): + if assistant_message.content is not None and not isinstance( + assistant_message.content, str + ): raw = assistant_message.content if isinstance(raw, dict): - assistant_message.content = raw.get("text", "") or raw.get("content", "") or json.dumps(raw) + assistant_message.content = ( + raw.get("text", "") + or raw.get("content", "") + or json.dumps(raw) + ) elif isinstance(raw, list): # Multimodal content list — extract text parts parts = [] @@ -9802,7 +11380,10 @@ def _stop_spinner(): try: from hermes_cli.plugins import invoke_hook as _invoke_hook - _assistant_tool_calls = getattr(assistant_message, "tool_calls", None) or [] + + _assistant_tool_calls = ( + getattr(assistant_message, "tool_calls", None) or [] + ) _assistant_text = assistant_message.content or "" _invoke_hook( "post_api_request", @@ -9828,87 +11409,125 @@ def _stop_spinner(): # Handle assistant response if assistant_message.content and not self.quiet_mode: if self.verbose_logging: - self._vprint(f"{self.log_prefix}🤖 Assistant: {assistant_message.content}") + self._vprint( + f"{self.log_prefix}🤖 Assistant: {assistant_message.content}" + ) else: - self._vprint(f"{self.log_prefix}🤖 Assistant: {assistant_message.content[:100]}{'...' if len(assistant_message.content) > 100 else ''}") + self._vprint( + f"{self.log_prefix}🤖 Assistant: {assistant_message.content[:100]}{'...' if len(assistant_message.content) > 100 else ''}" + ) # Notify progress callback of model's thinking (used by subagent # delegation to relay the child's reasoning to the parent display). - if (assistant_message.content and self.tool_progress_callback): + if assistant_message.content and self.tool_progress_callback: _think_text = assistant_message.content.strip() # Strip reasoning XML tags that shouldn't leak to parent display _think_text = re.sub( - r'</?(?:REASONING_SCRATCHPAD|think|reasoning)>', '', _think_text + r"</?(?:REASONING_SCRATCHPAD|think|reasoning)>", "", _think_text ).strip() # For subagents: relay first line to parent display (existing behaviour). # For all agents with a structured callback: emit reasoning.available event. - first_line = _think_text.split('\n')[0][:80] if _think_text else "" - if first_line and getattr(self, '_delegate_depth', 0) > 0: + first_line = _think_text.split("\n")[0][:80] if _think_text else "" + if first_line and getattr(self, "_delegate_depth", 0) > 0: try: self.tool_progress_callback("_thinking", first_line) except Exception: pass elif _think_text: try: - self.tool_progress_callback("reasoning.available", "_thinking", _think_text[:500], None) + self.tool_progress_callback( + "reasoning.available", + "_thinking", + _think_text[:500], + None, + ) except Exception: pass - + # Check for incomplete <REASONING_SCRATCHPAD> (opened but never closed) # This means the model ran out of output tokens mid-reasoning — retry up to 2 times if has_incomplete_scratchpad(assistant_message.content or ""): self._incomplete_scratchpad_retries += 1 - - self._vprint(f"{self.log_prefix}⚠️ Incomplete <REASONING_SCRATCHPAD> detected (opened but never closed)") - + + self._vprint( + f"{self.log_prefix}⚠️ Incomplete <REASONING_SCRATCHPAD> detected (opened but never closed)" + ) + if self._incomplete_scratchpad_retries <= 2: - self._vprint(f"{self.log_prefix}🔄 Retrying API call ({self._incomplete_scratchpad_retries}/2)...") + self._vprint( + f"{self.log_prefix}🔄 Retrying API call ({self._incomplete_scratchpad_retries}/2)..." + ) # Don't add the broken message, just retry continue else: # Max retries - discard this turn and save as partial - self._vprint(f"{self.log_prefix}❌ Max retries (2) for incomplete scratchpad. Saving as partial.", force=True) + self._vprint( + f"{self.log_prefix}❌ Max retries (2) for incomplete scratchpad. Saving as partial.", + force=True, + ) self._incomplete_scratchpad_retries = 0 - - rolled_back_messages = self._get_messages_up_to_last_assistant(messages) + + rolled_back_messages = self._get_messages_up_to_last_assistant( + messages + ) self._cleanup_task_resources(effective_task_id) self._persist_session(messages, conversation_history) - + return { "final_response": None, "messages": rolled_back_messages, "api_calls": api_call_count, "completed": False, "partial": True, - "error": "Incomplete REASONING_SCRATCHPAD after 2 retries" + "error": "Incomplete REASONING_SCRATCHPAD after 2 retries", } - + # Reset incomplete scratchpad counter on clean response self._incomplete_scratchpad_retries = 0 if self.api_mode == "codex_responses" and finish_reason == "incomplete": self._codex_incomplete_retries += 1 - interim_msg = self._build_assistant_message(assistant_message, finish_reason) - interim_has_content = bool((interim_msg.get("content") or "").strip()) - interim_has_reasoning = bool(interim_msg.get("reasoning", "").strip()) if isinstance(interim_msg.get("reasoning"), str) else False - interim_has_codex_reasoning = bool(interim_msg.get("codex_reasoning_items")) + interim_msg = self._build_assistant_message( + assistant_message, finish_reason + ) + interim_has_content = bool( + (interim_msg.get("content") or "").strip() + ) + interim_has_reasoning = ( + bool(interim_msg.get("reasoning", "").strip()) + if isinstance(interim_msg.get("reasoning"), str) + else False + ) + interim_has_codex_reasoning = bool( + interim_msg.get("codex_reasoning_items") + ) - if interim_has_content or interim_has_reasoning or interim_has_codex_reasoning: + if ( + interim_has_content + or interim_has_reasoning + or interim_has_codex_reasoning + ): last_msg = messages[-1] if messages else None # Duplicate detection: two consecutive incomplete assistant # messages with identical content AND reasoning are collapsed. # For reasoning-only messages (codex_reasoning_items differ but # visible content/reasoning are both empty), we also compare # the encrypted items to avoid silently dropping new state. - last_codex_items = last_msg.get("codex_reasoning_items") if isinstance(last_msg, dict) else None + last_codex_items = ( + last_msg.get("codex_reasoning_items") + if isinstance(last_msg, dict) + else None + ) interim_codex_items = interim_msg.get("codex_reasoning_items") duplicate_interim = ( isinstance(last_msg, dict) and last_msg.get("role") == "assistant" and last_msg.get("finish_reason") == "incomplete" - and (last_msg.get("content") or "") == (interim_msg.get("content") or "") - and (last_msg.get("reasoning") or "") == (interim_msg.get("reasoning") or "") + and (last_msg.get("content") or "") + == (interim_msg.get("content") or "") + and (last_msg.get("reasoning") or "") + == (interim_msg.get("reasoning") or "") and last_codex_items == interim_codex_items ) if not duplicate_interim: @@ -9917,7 +11536,9 @@ def _stop_spinner(): if self._codex_incomplete_retries < 3: if not self.quiet_mode: - self._vprint(f"{self.log_prefix}↻ Codex response incomplete; continuing turn ({self._codex_incomplete_retries}/3)") + self._vprint( + f"{self.log_prefix}↻ Codex response incomplete; continuing turn ({self._codex_incomplete_retries}/3)" + ) self._session_messages = messages self._save_session_log(messages) continue @@ -9934,26 +11555,33 @@ def _stop_spinner(): } elif hasattr(self, "_codex_incomplete_retries"): self._codex_incomplete_retries = 0 - + # Check for tool calls if assistant_message.tool_calls: if not self.quiet_mode: - self._vprint(f"{self.log_prefix}🔧 Processing {len(assistant_message.tool_calls)} tool call(s)...") - + self._vprint( + f"{self.log_prefix}🔧 Processing {len(assistant_message.tool_calls)} tool call(s)..." + ) + if self.verbose_logging: for tc in assistant_message.tool_calls: - logging.debug(f"Tool call: {tc.function.name} with args: {tc.function.arguments[:200]}...") - + logging.debug( + f"Tool call: {tc.function.name} with args: {tc.function.arguments[:200]}..." + ) + # Validate tool call names - detect model hallucinations # Repair mismatched tool names before validating for tc in assistant_message.tool_calls: if tc.function.name not in self.valid_tool_names: repaired = self._repair_tool_call(tc.function.name) if repaired: - print(f"{self.log_prefix}🔧 Auto-repaired tool name: '{tc.function.name}' -> '{repaired}'") + print( + f"{self.log_prefix}🔧 Auto-repaired tool name: '{tc.function.name}' -> '{repaired}'" + ) tc.function.name = repaired invalid_tool_calls = [ - tc.function.name for tc in assistant_message.tool_calls + tc.function.name + for tc in assistant_message.tool_calls if tc.function.name not in self.valid_tool_names ] if invalid_tool_calls: @@ -9963,11 +11591,20 @@ def _stop_spinner(): # Return helpful error to model — model can self-correct next turn available = ", ".join(sorted(self.valid_tool_names)) invalid_name = invalid_tool_calls[0] - invalid_preview = invalid_name[:80] + "..." if len(invalid_name) > 80 else invalid_name - self._vprint(f"{self.log_prefix}⚠️ Unknown tool '{invalid_preview}' — sending error to model for self-correction ({self._invalid_tool_retries}/3)") + invalid_preview = ( + invalid_name[:80] + "..." + if len(invalid_name) > 80 + else invalid_name + ) + self._vprint( + f"{self.log_prefix}⚠️ Unknown tool '{invalid_preview}' — sending error to model for self-correction ({self._invalid_tool_retries}/3)" + ) if self._invalid_tool_retries >= 3: - self._vprint(f"{self.log_prefix}❌ Max retries (3) for invalid tool calls exceeded. Stopping as partial.", force=True) + self._vprint( + f"{self.log_prefix}❌ Max retries (3) for invalid tool calls exceeded. Stopping as partial.", + force=True, + ) self._invalid_tool_retries = 0 self._persist_session(messages, conversation_history) return { @@ -9976,25 +11613,29 @@ def _stop_spinner(): "api_calls": api_call_count, "completed": False, "partial": True, - "error": f"Model generated invalid tool call: {invalid_preview}" + "error": f"Model generated invalid tool call: {invalid_preview}", } - assistant_msg = self._build_assistant_message(assistant_message, finish_reason) + assistant_msg = self._build_assistant_message( + assistant_message, finish_reason + ) messages.append(assistant_msg) for tc in assistant_message.tool_calls: if tc.function.name not in self.valid_tool_names: content = f"Tool '{tc.function.name}' does not exist. Available tools: {available}" else: content = "Skipped: another tool call in this turn used an invalid name. Please retry this tool call." - messages.append({ - "role": "tool", - "tool_call_id": tc.id, - "content": content, - }) + messages.append( + { + "role": "tool", + "tool_call_id": tc.id, + "content": content, + } + ) continue # Reset retry counter on successful tool call validation self._invalid_tool_retries = 0 - + # Validate tool call arguments are valid JSON # Handle empty strings as empty objects (common model quirk) invalid_json_args = [] @@ -10014,7 +11655,7 @@ def _stop_spinner(): json.loads(args) except json.JSONDecodeError as e: invalid_json_args.append((tc.function.name, str(e))) - + if invalid_json_args: # Check if the invalid JSON is due to truncation rather # than a model formatting mistake. Routers sometimes @@ -10023,7 +11664,9 @@ def _stop_spinner(): # Detect truncation: args that don't end with } or ] # (after stripping whitespace) are cut off mid-stream. _truncated = any( - not (tc.function.arguments or "").rstrip().endswith(("}", "]")) + not (tc.function.arguments or "") + .rstrip() + .endswith(("}", "]")) for tc in assistant_message.tool_calls if tc.function.name in {n for n, _ in invalid_json_args} ) @@ -10049,27 +11692,39 @@ def _stop_spinner(): self._invalid_json_retries += 1 tool_name, error_msg = invalid_json_args[0] - self._vprint(f"{self.log_prefix}⚠️ Invalid JSON in tool call arguments for '{tool_name}': {error_msg}") + self._vprint( + f"{self.log_prefix}⚠️ Invalid JSON in tool call arguments for '{tool_name}': {error_msg}" + ) if self._invalid_json_retries < 3: - self._vprint(f"{self.log_prefix}🔄 Retrying API call ({self._invalid_json_retries}/3)...") + self._vprint( + f"{self.log_prefix}🔄 Retrying API call ({self._invalid_json_retries}/3)..." + ) # Don't add anything to messages, just retry the API call continue else: # Instead of returning partial, inject tool error results so the model can recover. # Using tool results (not user messages) preserves role alternation. - self._vprint(f"{self.log_prefix}⚠️ Injecting recovery tool results for invalid JSON...") + self._vprint( + f"{self.log_prefix}⚠️ Injecting recovery tool results for invalid JSON..." + ) self._invalid_json_retries = 0 # Reset for next attempt - + # Append the assistant message with its (broken) tool_calls - recovery_assistant = self._build_assistant_message(assistant_message, finish_reason) + recovery_assistant = self._build_assistant_message( + assistant_message, finish_reason + ) messages.append(recovery_assistant) - + # Respond with tool error results for each tool call invalid_names = {name for name, _ in invalid_json_args} for tc in assistant_message.tool_calls: if tc.function.name in invalid_names: - err = next(e for n, e in invalid_json_args if n == tc.function.name) + err = next( + e + for n, e in invalid_json_args + if n == tc.function.name + ) tool_result = ( f"Error: Invalid JSON arguments. {err}. " f"For tools with no required parameters, use an empty object: {{}}. " @@ -10077,13 +11732,15 @@ def _stop_spinner(): ) else: tool_result = "Skipped: other tool call in this response had invalid JSON." - messages.append({ - "role": "tool", - "tool_call_id": tc.id, - "content": tool_result, - }) + messages.append( + { + "role": "tool", + "tool_call_id": tc.id, + "content": tool_result, + } + ) continue - + # Reset retry counter on successful JSON validation self._invalid_json_retries = 0 @@ -10095,23 +11752,32 @@ def _stop_spinner(): assistant_message.tool_calls ) - assistant_msg = self._build_assistant_message(assistant_message, finish_reason) - + assistant_msg = self._build_assistant_message( + assistant_message, finish_reason + ) + # If this turn has both content AND tool_calls, capture the content # as a fallback final response. Common pattern: model delivers its # answer and calls memory/skill tools as a side-effect in the same # turn. If the follow-up turn after tools is empty, we use this. turn_content = assistant_message.content or "" - if turn_content and self._has_content_after_think_block(turn_content): + if turn_content and self._has_content_after_think_block( + turn_content + ): self._last_content_with_tools = turn_content # Only mute subsequent output when EVERY tool call in # this turn is post-response housekeeping (memory, todo, # skill_manage, etc.). If any substantive tool is present # (search_files, read_file, write_file, terminal, ...), # keep output visible so the user sees progress. - _HOUSEKEEPING_TOOLS = frozenset({ - "memory", "todo", "skill_manage", "session_search", - }) + _HOUSEKEEPING_TOOLS = frozenset( + { + "memory", + "todo", + "skill_manage", + "session_search", + } + ) _all_housekeeping = all( tc.function.name in _HOUSEKEEPING_TOOLS for tc in assistant_message.tool_calls @@ -10122,7 +11788,7 @@ def _stop_spinner(): clean = self._strip_think_blocks(turn_content).strip() if clean: self._vprint(f" ┊ 💬 {clean}") - + # Pop thinking-only prefill message(s) before appending # (tool-call path — same rationale as the final-response path). _had_prefill = False @@ -10164,7 +11830,9 @@ def _stop_spinner(): except Exception: pass - self._execute_tool_calls(assistant_message, messages, effective_task_id, api_call_count) + self._execute_tool_calls( + assistant_message, messages, effective_task_id, api_call_count + ) # Reset per-turn retry counters after successful tool # execution so a single truncation doesn't poison the @@ -10182,10 +11850,12 @@ def _stop_spinner(): # Refund the iteration if the ONLY tool(s) called were # execute_code (programmatic tool calling). These are # cheap RPC-style calls that shouldn't eat the budget. - _tc_names = {tc.function.name for tc in assistant_message.tool_calls} + _tc_names = { + tc.function.name for tc in assistant_message.tool_calls + } if _tc_names == {"execute_code"}: self.iteration_budget.refund() - + # Use real token counts from the API response to decide # compression. prompt_tokens + completion_tokens is the # actual context size the provider reported plus the @@ -10217,7 +11887,9 @@ def _stop_spinner(): # and fires status_callback for gateway platforms. # Tiered: 85% (orange) and 95% (red/critical). if _compressor.threshold_tokens > 0: - _compaction_progress = _real_tokens / _compressor.threshold_tokens + _compaction_progress = ( + _real_tokens / _compressor.threshold_tokens + ) # Determine the warning tier for this progress level _warn_tier = 0.0 if _compaction_progress >= 0.95: @@ -10230,21 +11902,34 @@ def _stop_spinner(): _sid = self.session_id or "default" _last = AIAgent._context_pressure_last_warned.get(_sid) _now = time.time() - if _last is None or _last[0] < _warn_tier or (_now - _last[1]) >= self._CONTEXT_PRESSURE_COOLDOWN: + if ( + _last is None + or _last[0] < _warn_tier + or (_now - _last[1]) >= self._CONTEXT_PRESSURE_COOLDOWN + ): self._context_pressure_warned_at = _warn_tier - AIAgent._context_pressure_last_warned[_sid] = (_warn_tier, _now) - self._emit_context_pressure(_compaction_progress, _compressor) + AIAgent._context_pressure_last_warned[_sid] = ( + _warn_tier, + _now, + ) + self._emit_context_pressure( + _compaction_progress, _compressor + ) # Evict stale entries (older than 2x cooldown) _cutoff = _now - self._CONTEXT_PRESSURE_COOLDOWN * 2 AIAgent._context_pressure_last_warned = { - k: v for k, v in AIAgent._context_pressure_last_warned.items() + k: v + for k, v in AIAgent._context_pressure_last_warned.items() if v[1] > _cutoff } - if self.compression_enabled and _compressor.should_compress(_real_tokens): + if self.compression_enabled and _compressor.should_compress( + _real_tokens + ): self._safe_print(" ⟳ compacting context…") messages, active_system_prompt = self._compress_context( - messages, system_message, + messages, + system_message, approx_tokens=self.context_compressor.last_prompt_tokens, task_id=effective_task_id, ) @@ -10252,25 +11937,25 @@ def _stop_spinner(): # _flush_messages_to_session_db writes compressed messages # to the new session (see preflight compression comment). conversation_history = None - + # Save session log incrementally (so progress is visible even if interrupted) self._session_messages = messages self._save_session_log(messages) - + # Continue loop for next response continue - + else: # No tool calls - this is the final response final_response = assistant_message.content or "" - + # Fix: unmute output when entering the no-tool-call branch # so the user can see empty-response warnings and recovery # status messages. _mute_post_response was set during a # prior housekeeping tool turn and should not silence the # final response path. self._mute_post_response = False - + # Check if response only has think block with no actual content after it if not self._has_content_after_think_block(final_response): # ── Partial stream recovery ───────────────────── @@ -10283,7 +11968,9 @@ def _stop_spinner(): ) if self._has_content_after_think_block(_partial_streamed): _turn_exit_reason = "partial_stream_recovery" - _recovered = self._strip_think_blocks(_partial_streamed).strip() + _recovered = self._strip_think_blocks( + _partial_streamed + ).strip() logger.info( "Partial stream content delivered (%d chars) " "— using as final response", @@ -10301,11 +11988,15 @@ def _stop_spinner(): # tool calls (e.g. "You're welcome!" + memory save), the model # has nothing more to say. Use the earlier content immediately # instead of wasting API calls on retries that won't help. - fallback = getattr(self, '_last_content_with_tools', None) + fallback = getattr(self, "_last_content_with_tools", None) if fallback: _turn_exit_reason = "fallback_prior_turn_content" - logger.info("Empty follow-up after tool calls — using prior turn content as final response") - self._emit_status("↻ Empty response after tool calls — using earlier content as final answer") + logger.info( + "Empty follow-up after tool calls — using prior turn content as final response" + ) + self._emit_status( + "↻ Empty response after tool calls — using earlier content as final answer" + ) self._last_content_with_tools = None self._empty_content_retries = 0 # Do NOT modify the assistant message content — the @@ -10328,9 +12019,8 @@ def _stop_spinner(): m.get("role") == "tool" for m in messages[-5:] # check recent messages ) - if ( - _prior_was_tool - and not getattr(self, "_post_tool_empty_retried", False) + if _prior_was_tool and not getattr( + self, "_post_tool_empty_retried", False ): self._post_tool_empty_retried = True logger.info( @@ -10348,14 +12038,16 @@ def _stop_spinner(): # APIs reject as an invalid sequence. assistant_msg["content"] = "(empty)" messages.append(assistant_msg) - messages.append({ - "role": "user", - "content": ( - "You just executed tool calls but returned an " - "empty response. Please process the tool " - "results above and continue with the task." - ), - }) + messages.append( + { + "role": "user", + "content": ( + "You just executed tool calls but returned an " + "empty response. Please process the tool " + "results above and continue with the task." + ), + } + ) continue # ── Thinking-only prefill continuation ────────── @@ -10403,15 +12095,19 @@ def _stop_spinner(): final_response ).strip() _prefill_exhausted = ( - _has_structured - and self._thinking_prefill_retries >= 2 + _has_structured and self._thinking_prefill_retries >= 2 ) - if _truly_empty and (not _has_structured or _prefill_exhausted) and self._empty_content_retries < 3: + if ( + _truly_empty + and (not _has_structured or _prefill_exhausted) + and self._empty_content_retries < 3 + ): self._empty_content_retries += 1 logger.warning( "Empty response (no content or reasoning) — " "retry %d/3 (model=%s)", - self._empty_content_retries, self.model, + self._empty_content_retries, + self.model, ) self._emit_status( f"⚠️ Empty response from model — retrying " @@ -10429,7 +12125,8 @@ def _stop_spinner(): logger.warning( "Empty response after %d retries — " "attempting fallback (model=%s, provider=%s)", - self._empty_content_retries, self.model, + self._empty_content_retries, + self.model, self.provider, ) self._emit_status( @@ -10445,7 +12142,8 @@ def _stop_spinner(): logger.info( "Fallback activated after empty responses: " "now using %s on %s", - self.model, self.provider, + self.model, + self.provider, ) continue @@ -10454,16 +12152,23 @@ def _stop_spinner(): # "(empty)" terminal. _turn_exit_reason = "empty_response_exhausted" reasoning_text = self._extract_reasoning(assistant_message) - assistant_msg = self._build_assistant_message(assistant_message, finish_reason) + assistant_msg = self._build_assistant_message( + assistant_message, finish_reason + ) assistant_msg["content"] = "(empty)" messages.append(assistant_msg) if reasoning_text: - reasoning_preview = reasoning_text[:500] + "..." if len(reasoning_text) > 500 else reasoning_text + reasoning_preview = ( + reasoning_text[:500] + "..." + if len(reasoning_text) > 500 + else reasoning_text + ) logger.warning( "Reasoning-only response (no visible content) " "after exhausting retries and fallback. " - "Reasoning: %s", reasoning_preview, + "Reasoning: %s", + reasoning_preview, ) self._emit_status( "⚠️ Model produced reasoning but no visible " @@ -10474,18 +12179,22 @@ def _stop_spinner(): "Empty response (no content or reasoning) " "after %d retries. No fallback available. " "model=%s provider=%s", - self._empty_content_retries, self.model, + self._empty_content_retries, + self.model, self.provider, ) self._emit_status( "❌ Model returned no content after all retries" - + (" and fallback attempts." if self._fallback_chain else - ". No fallback providers configured.") + + ( + " and fallback attempts." + if self._fallback_chain + else ". No fallback providers configured." + ) ) final_response = "(empty)" break - + # Reset retry counter/signature on successful content self._empty_content_retries = 0 self._thinking_prefill_retries = 0 @@ -10501,7 +12210,9 @@ def _stop_spinner(): ) ): codex_ack_continuations += 1 - interim_msg = self._build_assistant_message(assistant_message, "incomplete") + interim_msg = self._build_assistant_message( + assistant_message, "incomplete" + ) messages.append(interim_msg) self._emit_interim_assistant_message(interim_msg) @@ -10523,11 +12234,13 @@ def _stop_spinner(): final_response = truncated_response_prefix + final_response truncated_response_prefix = "" length_continue_retries = 0 - + # Strip <think> blocks from user-facing response (keep raw in messages for trajectory) final_response = self._strip_think_blocks(final_response).strip() - - final_msg = self._build_assistant_message(assistant_message, finish_reason) + + final_msg = self._build_assistant_message( + assistant_message, finish_reason + ) # Pop thinking-only prefill message(s) before appending # the final response. This avoids consecutive assistant @@ -10541,21 +12254,25 @@ def _stop_spinner(): messages.pop() messages.append(final_msg) - + _turn_exit_reason = f"text_response(finish_reason={finish_reason})" if not self.quiet_mode: - self._safe_print(f"🎉 Conversation completed after {api_call_count} OpenAI-compatible API call(s)") + self._safe_print( + f"🎉 Conversation completed after {api_call_count} OpenAI-compatible API call(s)" + ) break - + except Exception as e: error_msg = f"Error during OpenAI-compatible API call #{api_call_count}: {str(e)}" try: print(f"❌ {error_msg}") except (OSError, ValueError): logger.error(error_msg) - - logger.debug("Outer loop error in API call #%d", api_call_count, exc_info=True) - + + logger.debug( + "Outer loop error in API call #%d", api_call_count, exc_info=True + ) + # If an assistant message with tool_calls was already appended, # the API expects a role="tool" result for every tool_call_id. # Fill in error results for any that weren't answered yet. @@ -10568,11 +12285,12 @@ def _stop_spinner(): if msg.get("role") == "assistant" and msg.get("tool_calls"): answered_ids = { m["tool_call_id"] - for m in messages[idx + 1:] + for m in messages[idx + 1 :] if isinstance(m, dict) and m.get("role") == "tool" } for tc in msg["tool_calls"]: - if not tc or not isinstance(tc, dict): continue + if not tc or not isinstance(tc, dict): + continue if tc["id"] not in answered_ids: err_msg = { "role": "tool", @@ -10581,7 +12299,7 @@ def _stop_spinner(): } messages.append(err_msg) break - + # Non-tool errors don't need a synthetic message injected. # The error is already printed to the user (line above), and # the retry loop continues. Injecting a fake user/assistant @@ -10591,12 +12309,14 @@ def _stop_spinner(): # If we're near the limit, break to avoid infinite loops if api_call_count >= self.max_iterations - 1: _turn_exit_reason = f"error_near_max_iterations({error_msg[:80]})" - final_response = f"I apologize, but I encountered repeated errors: {error_msg}" + final_response = ( + f"I apologize, but I encountered repeated errors: {error_msg}" + ) # Append as assistant so the history stays valid for # session resume (avoids consecutive user messages). messages.append({"role": "assistant", "content": final_response}) break - + if final_response is None and ( api_call_count >= self.max_iterations or self.iteration_budget.remaining <= 0 @@ -10604,7 +12324,9 @@ def _stop_spinner(): # Budget exhausted — ask the model for a summary via one extra # API call with tools stripped. _handle_max_iterations injects a # user message and makes a single toolless request. - _turn_exit_reason = f"max_iterations_reached({api_call_count}/{self.max_iterations})" + _turn_exit_reason = ( + f"max_iterations_reached({api_call_count}/{self.max_iterations})" + ) self._emit_status( f"⚠️ Iteration budget exhausted ({api_call_count}/{self.max_iterations}) " "— asking model to summarise" @@ -10615,7 +12337,7 @@ def _stop_spinner(): "— requesting summary..." ) final_response = self._handle_max_iterations(messages, api_call_count) - + # Determine if conversation completed successfully completed = final_response is not None and api_call_count < self.max_iterations @@ -10644,8 +12366,11 @@ def _stop_spinner(): break _turn_tool_count = sum( - 1 for m in messages - if isinstance(m, dict) and m.get("role") == "assistant" and m.get("tool_calls") + 1 + for m in messages + if isinstance(m, dict) + and m.get("role") == "assistant" + and m.get("tool_calls") ) _resp_len = len(final_response) if final_response else 0 _budget_used = self.iteration_budget.used if self.iteration_budget else 0 @@ -10656,9 +12381,15 @@ def _stop_spinner(): "tool_turns=%d last_msg_role=%s response_len=%d session=%s" ) _diag_args = ( - _turn_exit_reason, self.model, api_call_count, self.max_iterations, - _budget_used, _budget_max, - _turn_tool_count, _last_msg_role, _resp_len, + _turn_exit_reason, + self.model, + api_call_count, + self.max_iterations, + _budget_used, + _budget_max, + _turn_tool_count, + _last_msg_role, + _resp_len, self.session_id or "none", ) @@ -10666,8 +12397,10 @@ def _stop_spinner(): # Agent was mid-work — this is the "just stops" case. logger.warning( "Turn ended with pending tool result (agent may appear stuck). " - + _diag_msg + " last_tool=%s", - *_diag_args, _last_tool_name, + + _diag_msg + + " last_tool=%s", + *_diag_args, + _last_tool_name, ) else: logger.info(_diag_msg, *_diag_args) @@ -10679,6 +12412,7 @@ def _stop_spinner(): if final_response and not interrupted: try: from hermes_cli.plugins import invoke_hook as _invoke_hook + _invoke_hook( "post_llm_call", session_id=self.session_id, @@ -10719,17 +12453,20 @@ def _stop_spinner(): "prompt_tokens": self.session_prompt_tokens, "completion_tokens": self.session_completion_tokens, "total_tokens": self.session_total_tokens, - "last_prompt_tokens": getattr(self.context_compressor, "last_prompt_tokens", 0) or 0, + "last_prompt_tokens": getattr( + self.context_compressor, "last_prompt_tokens", 0 + ) + or 0, "estimated_cost_usd": self.session_estimated_cost_usd, "cost_status": self.session_cost_status, "cost_source": self.session_cost_source, } self._response_was_previewed = False - + # Include interrupt message if one triggered the interrupt if interrupted and self._interrupt_message: result["interrupt_message"] = self._interrupt_message - + # Clear interrupt state after handling self.clear_interrupt() @@ -10738,9 +12475,11 @@ def _stop_spinner(): # Check skill trigger NOW — based on how many tool iterations THIS turn used. _should_review_skills = False - if (self._skill_nudge_interval > 0 - and self._iters_since_skill >= self._skill_nudge_interval - and "skill_manage" in self.valid_tool_names): + if ( + self._skill_nudge_interval > 0 + and self._iters_since_skill >= self._skill_nudge_interval + and "skill_manage" in self.valid_tool_names + ): _should_review_skills = True self._iters_since_skill = 0 @@ -10756,7 +12495,11 @@ def _stop_spinner(): # Background memory/skill review — runs AFTER the response is delivered # so it never competes with the user's task for model attention. - if final_response and not interrupted and (_should_review_memory or _should_review_skills): + if ( + final_response + and not interrupted + and (_should_review_memory or _should_review_skills) + ): try: self._spawn_background_review( messages_snapshot=list(messages), @@ -10778,6 +12521,7 @@ def _stop_spinner(): # Plugins can use this for cleanup, flushing buffers, etc. try: from hermes_cli.plugins import invoke_hook as _invoke_hook + _invoke_hook( "on_session_end", session_id=self.session_id, @@ -10818,7 +12562,7 @@ def main( save_trajectories: bool = False, save_sample: bool = False, verbose: bool = False, - log_prefix_chars: int = 20 + log_prefix_chars: int = 20, ): """ Main function for running the agent directly. @@ -10844,58 +12588,69 @@ def main( """ print("🤖 AI Agent with Tool Calling") print("=" * 50) - + # Handle tool listing if list_tools: - from model_tools import get_all_tool_names, get_toolset_for_tool, get_available_toolsets + from model_tools import ( + get_all_tool_names, + get_toolset_for_tool, + get_available_toolsets, + ) from toolsets import get_all_toolsets, get_toolset_info - + print("📋 Available Tools & Toolsets:") print("-" * 50) - + # Show new toolsets system print("\n🎯 Predefined Toolsets (New System):") print("-" * 40) all_toolsets = get_all_toolsets() - + # Group by category basic_toolsets = [] composite_toolsets = [] scenario_toolsets = [] - + for name, toolset in all_toolsets.items(): info = get_toolset_info(name) if info: entry = (name, info) if name in ["web", "terminal", "vision", "creative", "reasoning"]: basic_toolsets.append(entry) - elif name in ["research", "development", "analysis", "content_creation", "full_stack"]: + elif name in [ + "research", + "development", + "analysis", + "content_creation", + "full_stack", + ]: composite_toolsets.append(entry) else: scenario_toolsets.append(entry) - + # Print basic toolsets print("\n📌 Basic Toolsets:") for name, info in basic_toolsets: - tools_str = ', '.join(info['resolved_tools']) if info['resolved_tools'] else 'none' + tools_str = ( + ", ".join(info["resolved_tools"]) if info["resolved_tools"] else "none" + ) print(f" • {name:15} - {info['description']}") print(f" Tools: {tools_str}") - + # Print composite toolsets print("\n📂 Composite Toolsets (built from other toolsets):") for name, info in composite_toolsets: - includes_str = ', '.join(info['includes']) if info['includes'] else 'none' + includes_str = ", ".join(info["includes"]) if info["includes"] else "none" print(f" • {name:15} - {info['description']}") print(f" Includes: {includes_str}") print(f" Total tools: {info['tool_count']}") - + # Print scenario-specific toolsets print("\n🎭 Scenario-Specific Toolsets:") for name, info in scenario_toolsets: print(f" • {name:20} - {info['description']}") print(f" Total tools: {info['tool_count']}") - - + # Show legacy toolset compatibility print("\n📦 Legacy Toolsets (for backward compatibility):") legacy_toolsets = get_available_toolsets() @@ -10904,47 +12659,57 @@ def main( print(f" {status} {name}: {info['description']}") if not info["available"]: print(f" Requirements: {', '.join(info['requirements'])}") - + # Show individual tools all_tools = get_all_tool_names() print(f"\n🔧 Individual Tools ({len(all_tools)} available):") for tool_name in sorted(all_tools): toolset = get_toolset_for_tool(tool_name) print(f" 📌 {tool_name} (from {toolset})") - + print("\n💡 Usage Examples:") print(" # Use predefined toolsets") - print(" python run_agent.py --enabled_toolsets=research --query='search for Python news'") - print(" python run_agent.py --enabled_toolsets=development --query='debug this code'") - print(" python run_agent.py --enabled_toolsets=safe --query='analyze without terminal'") + print( + " python run_agent.py --enabled_toolsets=research --query='search for Python news'" + ) + print( + " python run_agent.py --enabled_toolsets=development --query='debug this code'" + ) + print( + " python run_agent.py --enabled_toolsets=safe --query='analyze without terminal'" + ) print(" ") print(" # Combine multiple toolsets") - print(" python run_agent.py --enabled_toolsets=web,vision --query='analyze website'") + print( + " python run_agent.py --enabled_toolsets=web,vision --query='analyze website'" + ) print(" ") print(" # Disable toolsets") - print(" python run_agent.py --disabled_toolsets=terminal --query='no command execution'") + print( + " python run_agent.py --disabled_toolsets=terminal --query='no command execution'" + ) print(" ") print(" # Run with trajectory saving enabled") print(" python run_agent.py --save_trajectories --query='your question here'") return - + # Parse toolset selection arguments enabled_toolsets_list = None disabled_toolsets_list = None - + if enabled_toolsets: enabled_toolsets_list = [t.strip() for t in enabled_toolsets.split(",")] print(f"🎯 Enabled toolsets: {enabled_toolsets_list}") - + if disabled_toolsets: disabled_toolsets_list = [t.strip() for t in disabled_toolsets.split(",")] print(f"🚫 Disabled toolsets: {disabled_toolsets_list}") - + if save_trajectories: print("💾 Trajectory saving: ENABLED") print(" - Successful conversations → trajectory_samples.jsonl") print(" - Failed conversations → failed_trajectories.jsonl") - + # Initialize agent with provided parameters try: agent = AIAgent( @@ -10956,12 +12721,12 @@ def main( disabled_toolsets=disabled_toolsets_list, save_trajectories=save_trajectories, verbose_logging=verbose, - log_prefix_chars=log_prefix_chars + log_prefix_chars=log_prefix_chars, ) except RuntimeError as e: print(f"❌ Failed to initialize agent: {e}") return - + # Use provided query or default to Python 3.13 example if query is None: user_query = ( @@ -10970,45 +12735,43 @@ def main( ) else: user_query = query - + print(f"\n📝 User Query: {user_query}") print("\n" + "=" * 50) - + # Run conversation result = agent.run_conversation(user_query) - + print("\n" + "=" * 50) print("📋 CONVERSATION SUMMARY") print("=" * 50) print(f"✅ Completed: {result['completed']}") print(f"📞 API Calls: {result['api_calls']}") print(f"💬 Messages: {len(result['messages'])}") - - if result['final_response']: + + if result["final_response"]: print("\n🎯 FINAL RESPONSE:") print("-" * 30) - print(result['final_response']) - + print(result["final_response"]) + # Save sample trajectory to UUID-named file if requested if save_sample: sample_id = str(uuid.uuid4())[:8] sample_filename = f"sample_{sample_id}.json" - + # Convert messages to trajectory format (same as batch_runner) trajectory = agent._convert_to_trajectory_format( - result['messages'], - user_query, - result['completed'] + result["messages"], user_query, result["completed"] ) - + entry = { "conversations": trajectory, "timestamp": datetime.now().isoformat(), "model": model, - "completed": result['completed'], - "query": user_query + "completed": result["completed"], + "query": user_query, } - + try: with open(sample_filename, "w", encoding="utf-8") as f: # Pretty-print JSON with indent for readability @@ -11016,7 +12779,7 @@ def main( print(f"\n💾 Sample trajectory saved to: {sample_filename}") except Exception as e: print(f"\n⚠️ Failed to save sample: {e}") - + print("\n👋 Agent execution completed!") diff --git a/tests/agent/test_anthropic_adapter.py b/tests/agent/test_anthropic_adapter.py index ae78888d86f74..76620be5addaa 100644 --- a/tests/agent/test_anthropic_adapter.py +++ b/tests/agent/test_anthropic_adapter.py @@ -120,13 +120,17 @@ class TestReadClaudeCodeCredentials: def test_reads_valid_credentials(self, tmp_path, monkeypatch): cred_file = tmp_path / ".claude" / ".credentials.json" cred_file.parent.mkdir(parents=True) - cred_file.write_text(json.dumps({ - "claudeAiOauth": { - "accessToken": "sk-ant-oat01-token", - "refreshToken": "sk-ant-oat01-refresh", - "expiresAt": int(time.time() * 1000) + 3600_000, - } - })) + cred_file.write_text( + json.dumps( + { + "claudeAiOauth": { + "accessToken": "sk-ant-oat01-token", + "refreshToken": "sk-ant-oat01-refresh", + "expiresAt": int(time.time() * 1000) + 3600_000, + } + } + ) + ) monkeypatch.setattr("agent.anthropic_adapter.Path.home", lambda: tmp_path) creds = read_claude_code_credentials() assert creds is not None @@ -134,7 +138,9 @@ def test_reads_valid_credentials(self, tmp_path, monkeypatch): assert creds["refreshToken"] == "sk-ant-oat01-refresh" assert creds["source"] == "claude_code_credentials_file" - def test_ignores_primary_api_key_for_native_anthropic_resolution(self, tmp_path, monkeypatch): + def test_ignores_primary_api_key_for_native_anthropic_resolution( + self, tmp_path, monkeypatch + ): claude_json = tmp_path / ".claude.json" claude_json.write_text(json.dumps({"primaryApiKey": "sk-ant-api03-primary"})) monkeypatch.setattr("agent.anthropic_adapter.Path.home", lambda: tmp_path) @@ -156,9 +162,9 @@ def test_returns_none_for_missing_oauth_key(self, tmp_path, monkeypatch): def test_returns_none_for_empty_access_token(self, tmp_path, monkeypatch): cred_file = tmp_path / ".claude" / ".credentials.json" cred_file.parent.mkdir(parents=True) - cred_file.write_text(json.dumps({ - "claudeAiOauth": {"accessToken": "", "refreshToken": "x"} - })) + cred_file.write_text( + json.dumps({"claudeAiOauth": {"accessToken": "", "refreshToken": "x"}}) + ) monkeypatch.setattr("agent.anthropic_adapter.Path.home", lambda: tmp_path) assert read_claude_code_credentials() is None @@ -185,16 +191,22 @@ def test_prefers_oauth_token_over_api_key(self, monkeypatch, tmp_path): monkeypatch.setattr("agent.anthropic_adapter.Path.home", lambda: tmp_path) assert resolve_anthropic_token() == "sk-ant-oat01-mytoken" - def test_does_not_resolve_primary_api_key_as_native_anthropic_token(self, monkeypatch, tmp_path): + def test_does_not_resolve_primary_api_key_as_native_anthropic_token( + self, monkeypatch, tmp_path + ): monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) monkeypatch.delenv("ANTHROPIC_TOKEN", raising=False) monkeypatch.delenv("CLAUDE_CODE_OAUTH_TOKEN", raising=False) - (tmp_path / ".claude.json").write_text(json.dumps({"primaryApiKey": "sk-ant-api03-primary"})) + (tmp_path / ".claude.json").write_text( + json.dumps({"primaryApiKey": "sk-ant-api03-primary"}) + ) monkeypatch.setattr("agent.anthropic_adapter.Path.home", lambda: tmp_path) assert resolve_anthropic_token() is None - def test_falls_back_to_api_key_when_no_oauth_sources_exist(self, monkeypatch, tmp_path): + def test_falls_back_to_api_key_when_no_oauth_sources_exist( + self, monkeypatch, tmp_path + ): monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-api03-mykey") monkeypatch.delenv("ANTHROPIC_TOKEN", raising=False) monkeypatch.delenv("CLAUDE_CODE_OAUTH_TOKEN", raising=False) @@ -228,39 +240,53 @@ def test_falls_back_to_claude_code_credentials(self, monkeypatch, tmp_path): monkeypatch.delenv("CLAUDE_CODE_OAUTH_TOKEN", raising=False) cred_file = tmp_path / ".claude" / ".credentials.json" cred_file.parent.mkdir(parents=True) - cred_file.write_text(json.dumps({ - "claudeAiOauth": { - "accessToken": "cc-auto-token", - "refreshToken": "refresh", - "expiresAt": int(time.time() * 1000) + 3600_000, - } - })) + cred_file.write_text( + json.dumps( + { + "claudeAiOauth": { + "accessToken": "cc-auto-token", + "refreshToken": "refresh", + "expiresAt": int(time.time() * 1000) + 3600_000, + } + } + ) + ) monkeypatch.setattr("agent.anthropic_adapter.Path.home", lambda: tmp_path) assert resolve_anthropic_token() == "cc-auto-token" - def test_prefers_refreshable_claude_code_credentials_over_static_anthropic_token(self, monkeypatch, tmp_path): + def test_prefers_refreshable_claude_code_credentials_over_static_anthropic_token( + self, monkeypatch, tmp_path + ): monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) monkeypatch.setenv("ANTHROPIC_TOKEN", "sk-ant-oat01-static-token") monkeypatch.delenv("CLAUDE_CODE_OAUTH_TOKEN", raising=False) cred_file = tmp_path / ".claude" / ".credentials.json" cred_file.parent.mkdir(parents=True) - cred_file.write_text(json.dumps({ - "claudeAiOauth": { - "accessToken": "cc-auto-token", - "refreshToken": "refresh-token", - "expiresAt": int(time.time() * 1000) + 3600_000, - } - })) + cred_file.write_text( + json.dumps( + { + "claudeAiOauth": { + "accessToken": "cc-auto-token", + "refreshToken": "refresh-token", + "expiresAt": int(time.time() * 1000) + 3600_000, + } + } + ) + ) monkeypatch.setattr("agent.anthropic_adapter.Path.home", lambda: tmp_path) assert resolve_anthropic_token() == "cc-auto-token" - def test_keeps_static_anthropic_token_when_only_non_refreshable_claude_key_exists(self, monkeypatch, tmp_path): + def test_keeps_static_anthropic_token_when_only_non_refreshable_claude_key_exists( + self, monkeypatch, tmp_path + ): monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) monkeypatch.setenv("ANTHROPIC_TOKEN", "sk-ant-oat01-static-token") monkeypatch.delenv("CLAUDE_CODE_OAUTH_TOKEN", raising=False) claude_json = tmp_path / ".claude.json" - claude_json.write_text(json.dumps({"primaryApiKey": "sk-ant-api03-managed-key"})) + claude_json.write_text( + json.dumps({"primaryApiKey": "sk-ant-api03-managed-key"}) + ) monkeypatch.setattr("agent.anthropic_adapter.Path.home", lambda: tmp_path) assert resolve_anthropic_token() == "sk-ant-oat01-static-token" @@ -280,17 +306,19 @@ def test_successful_refresh(self, tmp_path, monkeypatch): "expiresAt": int(time.time() * 1000) - 3600_000, } - mock_response = json.dumps({ - "access_token": "new-token-abc", - "refresh_token": "new-refresh-456", - "expires_in": 7200, - }).encode() + mock_response = json.dumps( + { + "access_token": "new-token-abc", + "refresh_token": "new-refresh-456", + "expires_in": 7200, + } + ).encode() with patch("urllib.request.urlopen") as mock_urlopen: mock_ctx = MagicMock() - mock_ctx.__enter__ = MagicMock(return_value=MagicMock( - read=MagicMock(return_value=mock_response) - )) + mock_ctx.__enter__ = MagicMock( + return_value=MagicMock(read=MagicMock(return_value=mock_response)) + ) mock_ctx.__exit__ = MagicMock(return_value=False) mock_urlopen.return_value = mock_ctx @@ -348,38 +376,54 @@ def test_auto_refresh_on_expired_creds(self, monkeypatch, tmp_path): # Set up expired creds with a refresh token cred_file = tmp_path / ".claude" / ".credentials.json" cred_file.parent.mkdir(parents=True) - cred_file.write_text(json.dumps({ - "claudeAiOauth": { - "accessToken": "expired-tok", - "refreshToken": "valid-refresh", - "expiresAt": int(time.time() * 1000) - 3600_000, - } - })) + cred_file.write_text( + json.dumps( + { + "claudeAiOauth": { + "accessToken": "expired-tok", + "refreshToken": "valid-refresh", + "expiresAt": int(time.time() * 1000) - 3600_000, + } + } + ) + ) monkeypatch.setattr("agent.anthropic_adapter.Path.home", lambda: tmp_path) # Mock refresh to succeed - with patch("agent.anthropic_adapter._refresh_oauth_token", return_value="refreshed-token"): + with patch( + "agent.anthropic_adapter._refresh_oauth_token", + return_value="refreshed-token", + ): result = resolve_anthropic_token() assert result == "refreshed-token" - def test_static_env_oauth_token_does_not_block_refreshable_claude_creds(self, monkeypatch, tmp_path): + def test_static_env_oauth_token_does_not_block_refreshable_claude_creds( + self, monkeypatch, tmp_path + ): monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) monkeypatch.setenv("ANTHROPIC_TOKEN", "sk-ant-oat01-expired-env-token") monkeypatch.delenv("CLAUDE_CODE_OAUTH_TOKEN", raising=False) cred_file = tmp_path / ".claude" / ".credentials.json" cred_file.parent.mkdir(parents=True) - cred_file.write_text(json.dumps({ - "claudeAiOauth": { - "accessToken": "expired-claude-creds-token", - "refreshToken": "valid-refresh", - "expiresAt": int(time.time() * 1000) - 3600_000, - } - })) + cred_file.write_text( + json.dumps( + { + "claudeAiOauth": { + "accessToken": "expired-claude-creds-token", + "refreshToken": "valid-refresh", + "expiresAt": int(time.time() * 1000) - 3600_000, + } + } + ) + ) monkeypatch.setattr("agent.anthropic_adapter.Path.home", lambda: tmp_path) - with patch("agent.anthropic_adapter._refresh_oauth_token", return_value="refreshed-token"): + with patch( + "agent.anthropic_adapter._refresh_oauth_token", + return_value="refreshed-token", + ): result = resolve_anthropic_token() assert result == "refreshed-token" @@ -400,13 +444,17 @@ def test_returns_token_from_credential_files(self, monkeypatch, tmp_path): # Pre-create credential files that will be found after subprocess cred_file = tmp_path / ".claude" / ".credentials.json" cred_file.parent.mkdir(parents=True) - cred_file.write_text(json.dumps({ - "claudeAiOauth": { - "accessToken": "from-cred-file", - "refreshToken": "refresh", - "expiresAt": int(time.time() * 1000) + 3600_000, - } - })) + cred_file.write_text( + json.dumps( + { + "claudeAiOauth": { + "accessToken": "from-cred-file", + "refreshToken": "refresh", + "expiresAt": int(time.time() * 1000) + 3600_000, + } + } + ) + ) monkeypatch.setattr("agent.anthropic_adapter.Path.home", lambda: tmp_path) with patch("subprocess.run") as mock_run: @@ -459,27 +507,45 @@ def test_returns_none_on_keyboard_interrupt(self, monkeypatch): class TestNormalizeModelName: def test_strips_anthropic_prefix(self): - assert normalize_model_name("anthropic/claude-sonnet-4-20250514") == "claude-sonnet-4-20250514" + assert ( + normalize_model_name("anthropic/claude-sonnet-4-20250514") + == "claude-sonnet-4-20250514" + ) def test_leaves_bare_name(self): - assert normalize_model_name("claude-sonnet-4-20250514") == "claude-sonnet-4-20250514" + assert ( + normalize_model_name("claude-sonnet-4-20250514") + == "claude-sonnet-4-20250514" + ) def test_converts_dots_to_hyphens(self): """OpenRouter uses dots (4.6), Anthropic uses hyphens (4-6).""" assert normalize_model_name("anthropic/claude-opus-4.6") == "claude-opus-4-6" - assert normalize_model_name("anthropic/claude-sonnet-4.5") == "claude-sonnet-4-5" + assert ( + normalize_model_name("anthropic/claude-sonnet-4.5") == "claude-sonnet-4-5" + ) assert normalize_model_name("claude-opus-4.6") == "claude-opus-4-6" def test_already_hyphenated_unchanged(self): """Names already in Anthropic format should pass through.""" assert normalize_model_name("claude-opus-4-6") == "claude-opus-4-6" - assert normalize_model_name("claude-opus-4-5-20251101") == "claude-opus-4-5-20251101" + assert ( + normalize_model_name("claude-opus-4-5-20251101") + == "claude-opus-4-5-20251101" + ) def test_preserve_dots_for_alibaba_dashscope(self): """Alibaba/DashScope use dots in model names (e.g. qwen3.5-plus). Fixes #1739.""" - assert normalize_model_name("qwen3.5-plus", preserve_dots=True) == "qwen3.5-plus" - assert normalize_model_name("anthropic/qwen3.5-plus", preserve_dots=True) == "qwen3.5-plus" - assert normalize_model_name("qwen3.5-flash", preserve_dots=True) == "qwen3.5-flash" + assert ( + normalize_model_name("qwen3.5-plus", preserve_dots=True) == "qwen3.5-plus" + ) + assert ( + normalize_model_name("anthropic/qwen3.5-plus", preserve_dots=True) + == "qwen3.5-plus" + ) + assert ( + normalize_model_name("qwen3.5-flash", preserve_dots=True) == "qwen3.5-flash" + ) # --------------------------------------------------------------------------- @@ -536,7 +602,10 @@ def test_converts_user_image_url_blocks_to_anthropic_image_blocks(self): "role": "user", "content": [ {"type": "text", "text": "Can you see this?"}, - {"type": "image_url", "image_url": {"url": "https://example.com/cat.png"}}, + { + "type": "image_url", + "image_url": {"url": "https://example.com/cat.png"}, + }, ], } ] @@ -548,7 +617,10 @@ def test_converts_user_image_url_blocks_to_anthropic_image_blocks(self): "role": "user", "content": [ {"type": "text", "text": "Can you see this?"}, - {"type": "image", "source": {"type": "url", "url": "https://example.com/cat.png"}}, + { + "type": "image", + "source": {"type": "url", "url": "https://example.com/cat.png"}, + }, ], } ] @@ -613,7 +685,10 @@ def test_converts_tool_results(self): "role": "assistant", "content": "", "tool_calls": [ - {"id": "tc_1", "function": {"name": "test_tool", "arguments": "{}"}}, + { + "id": "tc_1", + "function": {"name": "test_tool", "arguments": "{}"}, + }, ], }, {"role": "tool", "tool_call_id": "tc_1", "content": "result data"}, @@ -678,10 +753,9 @@ def test_strips_orphaned_tool_result(self): # tc_gone has no matching tool_use — its tool_result should be stripped for m in result: if m["role"] == "user" and isinstance(m["content"], list): - assert all( - b.get("type") != "tool_result" - for b in m["content"] - ), "Orphaned tool_result should have been stripped" + assert all(b.get("type") != "tool_result" for b in m["content"]), ( + "Orphaned tool_result should have been stripped" + ) def test_strips_orphaned_tool_result_preserves_valid(self): """Orphaned tool_results are stripped while valid ones survive.""" @@ -690,7 +764,10 @@ def test_strips_orphaned_tool_result_preserves_valid(self): "role": "assistant", "content": "", "tool_calls": [ - {"id": "tc_valid", "function": {"name": "search", "arguments": "{}"}}, + { + "id": "tc_valid", + "function": {"name": "search", "arguments": "{}"}, + }, ], }, {"role": "tool", "tool_call_id": "tc_valid", "content": "good result"}, @@ -709,7 +786,11 @@ def test_system_with_cache_control(self): { "role": "system", "content": [ - {"type": "text", "text": "System prompt", "cache_control": {"type": "ephemeral"}}, + { + "type": "text", + "text": "System prompt", + "cache_control": {"type": "ephemeral"}, + }, ], }, {"role": "user", "content": "Hi"}, @@ -720,10 +801,12 @@ def test_system_with_cache_control(self): assert system[0]["cache_control"] == {"type": "ephemeral"} def test_assistant_cache_control_blocks_are_preserved(self): - messages = apply_anthropic_cache_control([ - {"role": "system", "content": "System prompt"}, - {"role": "assistant", "content": "Hello from assistant"}, - ]) + messages = apply_anthropic_cache_control( + [ + {"role": "system", "content": "System prompt"}, + {"role": "assistant", "content": "Hello from assistant"}, + ] + ) _, result = convert_messages_to_anthropic(messages) assistant_blocks = result[0]["content"] @@ -733,17 +816,23 @@ def test_assistant_cache_control_blocks_are_preserved(self): assert assistant_blocks[0]["cache_control"] == {"type": "ephemeral"} def test_tool_cache_control_is_preserved_on_tool_result_block(self): - messages = apply_anthropic_cache_control([ - {"role": "system", "content": "System prompt"}, - { - "role": "assistant", - "content": "", - "tool_calls": [ - {"id": "tc_1", "function": {"name": "test_tool", "arguments": "{}"}}, - ], - }, - {"role": "tool", "tool_call_id": "tc_1", "content": "result"}, - ], native_anthropic=True) + messages = apply_anthropic_cache_control( + [ + {"role": "system", "content": "System prompt"}, + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "tc_1", + "function": {"name": "test_tool", "arguments": "{}"}, + }, + ], + }, + {"role": "tool", "tool_call_id": "tc_1", "content": "result"}, + ], + native_anthropic=True, + ) _, result = convert_messages_to_anthropic(messages) user_msg = [m for m in result if m["role"] == "user"][0] @@ -760,7 +849,10 @@ def test_preserved_thinking_blocks_are_rehydrated_before_tool_use(self): "role": "assistant", "content": "", "tool_calls": [ - {"id": "tc_1", "function": {"name": "test_tool", "arguments": "{}"}}, + { + "id": "tc_1", + "function": {"name": "test_tool", "arguments": "{}"}, + }, ], "reasoning_details": [ { @@ -774,10 +866,14 @@ def test_preserved_thinking_blocks_are_rehydrated_before_tool_use(self): ] _, result = convert_messages_to_anthropic(messages) - assistant_blocks = next(msg for msg in result if msg["role"] == "assistant")["content"] + assistant_blocks = next(msg for msg in result if msg["role"] == "assistant")[ + "content" + ] assert assistant_blocks[0]["type"] == "thinking" - assert assistant_blocks[0]["thinking"] == "Need to inspect the tool result first." + assert ( + assistant_blocks[0]["thinking"] == "Need to inspect the tool result first." + ) assert assistant_blocks[0]["signature"] == "sig_123" assert assistant_blocks[1]["type"] == "tool_use" @@ -832,25 +928,33 @@ def test_converts_remote_image_url_to_anthropic_image_block(self): } def test_empty_cached_assistant_tool_turn_converts_without_empty_text_block(self): - messages = apply_anthropic_cache_control([ - {"role": "system", "content": "System prompt"}, - {"role": "user", "content": "Find the skill"}, - { - "role": "assistant", - "content": "", - "tool_calls": [ - {"id": "tc_1", "function": {"name": "skill_view", "arguments": "{}"}}, - ], - }, - {"role": "tool", "tool_call_id": "tc_1", "content": "result"}, - ]) + messages = apply_anthropic_cache_control( + [ + {"role": "system", "content": "System prompt"}, + {"role": "user", "content": "Find the skill"}, + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "tc_1", + "function": {"name": "skill_view", "arguments": "{}"}, + }, + ], + }, + {"role": "tool", "tool_call_id": "tc_1", "content": "result"}, + ] + ) _, result = convert_messages_to_anthropic(messages) assistant_turn = next(msg for msg in result if msg["role"] == "assistant") assistant_blocks = assistant_turn["content"] - assert all(not (b.get("type") == "text" and b.get("text") == "") for b in assistant_blocks) + assert all( + not (b.get("type") == "text" and b.get("text") == "") + for b in assistant_blocks + ) assert any(b.get("type") == "tool_use" for b in assistant_blocks) def test_empty_user_message_string_gets_placeholder(self): @@ -888,7 +992,13 @@ def test_empty_user_message_list_gets_placeholder(self): def test_user_message_with_empty_text_blocks_gets_placeholder(self): """User message with only empty text blocks should get placeholder.""" messages = [ - {"role": "user", "content": [{"type": "text", "text": ""}, {"type": "text", "text": " "}]}, + { + "role": "user", + "content": [ + {"type": "text", "text": ""}, + {"type": "text", "text": " "}, + ], + }, ] _, result = convert_messages_to_anthropic(messages) assert result[0]["role"] == "user" @@ -1085,35 +1195,43 @@ def test_context_length_no_clamp_when_larger(self): class TestGetAnthropicMaxOutput: def test_opus_4_6(self): from agent.anthropic_adapter import _get_anthropic_max_output + assert _get_anthropic_max_output("claude-opus-4-6") == 128_000 def test_opus_4_6_variant(self): from agent.anthropic_adapter import _get_anthropic_max_output + assert _get_anthropic_max_output("claude-opus-4-6:1m:fast") == 128_000 def test_sonnet_4_6(self): from agent.anthropic_adapter import _get_anthropic_max_output + assert _get_anthropic_max_output("claude-sonnet-4-6") == 64_000 def test_sonnet_4_date_stamped(self): from agent.anthropic_adapter import _get_anthropic_max_output + assert _get_anthropic_max_output("claude-sonnet-4-20250514") == 64_000 def test_claude_3_5_sonnet(self): from agent.anthropic_adapter import _get_anthropic_max_output + assert _get_anthropic_max_output("claude-3-5-sonnet-20241022") == 8_192 def test_claude_3_opus(self): from agent.anthropic_adapter import _get_anthropic_max_output + assert _get_anthropic_max_output("claude-3-opus-20240229") == 4_096 def test_unknown_future_model(self): from agent.anthropic_adapter import _get_anthropic_max_output + assert _get_anthropic_max_output("claude-ultra-5-20260101") == 128_000 def test_longest_prefix_wins(self): """'claude-3-5-sonnet' should match before 'claude-3-5'.""" from agent.anthropic_adapter import _get_anthropic_max_output + # claude-3-5-sonnet (8192) should win over a hypothetical shorter match assert _get_anthropic_max_output("claude-3-5-sonnet-20241022") == 8_192 @@ -1218,7 +1336,9 @@ def test_thinking_response(self): msg, reason = normalize_anthropic_response(self._make_response(blocks)) assert msg.content == "The answer is 42." assert msg.reasoning == "Let me reason about this..." - assert msg.reasoning_details == [{"type": "thinking", "thinking": "Let me reason about this..."}] + assert msg.reasoning_details == [ + {"type": "thinking", "thinking": "Let me reason about this..."} + ] def test_thinking_response_preserves_signature(self): blocks = [ @@ -1235,15 +1355,9 @@ def test_thinking_response_preserves_signature(self): def test_stop_reason_mapping(self): block = SimpleNamespace(type="text", text="x") - _, r1 = normalize_anthropic_response( - self._make_response([block], "end_turn") - ) - _, r2 = normalize_anthropic_response( - self._make_response([block], "tool_use") - ) - _, r3 = normalize_anthropic_response( - self._make_response([block], "max_tokens") - ) + _, r1 = normalize_anthropic_response(self._make_response([block], "end_turn")) + _, r2 = normalize_anthropic_response(self._make_response([block], "tool_use")) + _, r3 = normalize_anthropic_response(self._make_response([block], "max_tokens")) assert r1 == "stop" assert r2 == "tool_calls" assert r3 == "length" @@ -1306,7 +1420,11 @@ def test_thinking_stripped_from_non_last_assistant(self): {"id": "tc_1", "function": {"name": "tool1", "arguments": "{}"}}, ], "reasoning_details": [ - {"type": "thinking", "thinking": "Old reasoning.", "signature": "sig_old"}, + { + "type": "thinking", + "thinking": "Old reasoning.", + "signature": "sig_old", + }, ], }, {"role": "tool", "tool_call_id": "tc_1", "content": "result 1"}, @@ -1317,7 +1435,11 @@ def test_thinking_stripped_from_non_last_assistant(self): {"id": "tc_2", "function": {"name": "tool2", "arguments": "{}"}}, ], "reasoning_details": [ - {"type": "thinking", "thinking": "Latest reasoning.", "signature": "sig_new"}, + { + "type": "thinking", + "thinking": "Latest reasoning.", + "signature": "sig_new", + }, ], }, {"role": "tool", "tool_call_id": "tc_2", "content": "result 2"}, @@ -1348,7 +1470,11 @@ def test_signed_thinking_preserved_on_last_turn(self): "role": "assistant", "content": "The answer is 42.", "reasoning_details": [ - {"type": "thinking", "thinking": "Deep thought.", "signature": "sig_valid"}, + { + "type": "thinking", + "thinking": "Deep thought.", + "signature": "sig_valid", + }, ], }, ] @@ -1445,14 +1571,22 @@ def test_thinking_stripped_from_merged_consecutive_assistants(self): "role": "assistant", "content": "First response.", "reasoning_details": [ - {"type": "thinking", "thinking": "First thought.", "signature": "sig_1"}, + { + "type": "thinking", + "thinking": "First thought.", + "signature": "sig_1", + }, ], }, { "role": "assistant", "content": "Second response.", "reasoning_details": [ - {"type": "thinking", "thinking": "Second thought.", "signature": "sig_2"}, + { + "type": "thinking", + "thinking": "Second thought.", + "signature": "sig_2", + }, ], }, ] @@ -1532,12 +1666,57 @@ def test_multi_turn_conversation_preserves_only_last(self): # Last one: thinking preserved last_thinking = [ - b for b in assistants[2]["content"] + b + for b in assistants[2]["content"] if isinstance(b, dict) and b.get("type") == "thinking" ] assert len(last_thinking) == 1 assert last_thinking[0]["signature"] == "sig_3" + def test_third_party_downgrades_thinking_to_text(self): + """Third-party Anthropic-compatible endpoints get plain text thinking.""" + messages = [ + { + "role": "assistant", + "content": "Visible answer.", + "reasoning_details": [ + { + "type": "thinking", + "thinking": "Third-party-safe reasoning.", + "signature": "sig", + }, + {"type": "redacted_thinking", "data": "opaque"}, + ], + } + ] + _, result = convert_messages_to_anthropic( + messages, + base_url="https://api.z.ai/api/paas/v4", + ) + blocks = result[0]["content"] + assert not any(b.get("type") == "thinking" for b in blocks) + assert not any(b.get("type") == "redacted_thinking" for b in blocks) + text_blocks = [b.get("text", "") for b in blocks if b.get("type") == "text"] + assert "Third-party-safe reasoning." in text_blocks + assert "Visible answer." in text_blocks + + def test_third_party_thinking_only_content_gets_placeholder(self): + """If third-party turn only has redacted_thinking, use placeholder text.""" + messages = [ + { + "role": "assistant", + "content": "", + "reasoning_details": [ + {"type": "redacted_thinking", "data": "opaque"}, + ], + } + ] + _, result = convert_messages_to_anthropic( + messages, + base_url="https://api.minimax.io/anthropic", + ) + assert result[0]["content"] == [{"type": "text", "text": "(thinking elided)"}] + # --------------------------------------------------------------------------- # Tool choice diff --git a/tests/run_agent/test_run_agent.py b/tests/run_agent/test_run_agent.py index d71e6a625542d..d013ca651096a 100644 --- a/tests/run_agent/test_run_agent.py +++ b/tests/run_agent/test_run_agent.py @@ -124,7 +124,8 @@ def test_aiagent_reuses_existing_errors_log_handler(): ) matching_handlers = [ - handler for handler in root_logger.handlers + handler + for handler in root_logger.handlers if isinstance(handler, RotatingFileHandler) and error_log_path == Path(handler.baseFilename).resolve() ] @@ -142,7 +143,8 @@ class TestProviderModelNormalization: def test_aiagent_strips_matching_native_provider_prefix(self): with ( patch( - "run_agent.get_tool_definitions", return_value=_make_tool_defs("web_search") + "run_agent.get_tool_definitions", + return_value=_make_tool_defs("web_search"), ), patch("run_agent.check_toolset_requirements", return_value={}), patch("run_agent.OpenAI"), @@ -162,7 +164,8 @@ def test_aiagent_strips_matching_native_provider_prefix(self): def test_aiagent_keeps_aggregator_vendor_slug(self): with ( patch( - "run_agent.get_tool_definitions", return_value=_make_tool_defs("web_search") + "run_agent.get_tool_definitions", + return_value=_make_tool_defs("web_search"), ), patch("run_agent.check_toolset_requirements", return_value={}), patch("run_agent.OpenAI"), @@ -304,7 +307,9 @@ def test_mixed_orphaned_and_paired_tags(self, agent): def test_thought_block_removed(self, agent): """Gemma 4 uses <thought> tags for inline reasoning.""" - result = agent._strip_think_blocks("<thought>internal reasoning</thought> answer") + result = agent._strip_think_blocks( + "<thought>internal reasoning</thought> answer" + ) assert "internal reasoning" not in result assert "<thought>" not in result assert "answer" in result @@ -669,12 +674,18 @@ def test_includes_datetime(self, agent): assert "Conversation started:" in prompt def test_includes_nous_subscription_prompt(self, agent, monkeypatch): - monkeypatch.setattr(run_agent, "build_nous_subscription_prompt", lambda tool_names: "NOUS SUBSCRIPTION BLOCK") + monkeypatch.setattr( + run_agent, + "build_nous_subscription_prompt", + lambda tool_names: "NOUS SUBSCRIPTION BLOCK", + ) prompt = agent._build_system_prompt() assert "NOUS SUBSCRIPTION BLOCK" in prompt def test_skills_prompt_derives_available_toolsets_from_loaded_tools(self): - tools = _make_tool_defs("web_search", "skills_list", "skill_view", "skill_manage") + tools = _make_tool_defs( + "web_search", "skills_list", "skill_view", "skill_manage" + ) toolset_map = { "web_search": "web", "skills_list": "skills", @@ -688,8 +699,14 @@ def test_skills_prompt_derives_available_toolsets_from_loaded_tools(self): "run_agent.check_toolset_requirements", side_effect=AssertionError("should not re-check toolset requirements"), ), - patch("run_agent.get_toolset_for_tool", create=True, side_effect=toolset_map.get), - patch("run_agent.build_skills_system_prompt", return_value="SKILLS_PROMPT") as mock_skills, + patch( + "run_agent.get_toolset_for_tool", + create=True, + side_effect=toolset_map.get, + ), + patch( + "run_agent.build_skills_system_prompt", return_value="SKILLS_PROMPT" + ) as mock_skills, patch("run_agent.OpenAI"), ): agent = AIAgent( @@ -735,54 +752,71 @@ def _make_agent(self, model="openai/gpt-4.1", tool_use_enforcement="auto"): def test_auto_injects_for_gpt(self): from agent.prompt_builder import TOOL_USE_ENFORCEMENT_GUIDANCE + agent = self._make_agent(model="openai/gpt-4.1", tool_use_enforcement="auto") prompt = agent._build_system_prompt() assert TOOL_USE_ENFORCEMENT_GUIDANCE in prompt def test_auto_injects_for_codex(self): from agent.prompt_builder import TOOL_USE_ENFORCEMENT_GUIDANCE + agent = self._make_agent(model="openai/codex-mini", tool_use_enforcement="auto") prompt = agent._build_system_prompt() assert TOOL_USE_ENFORCEMENT_GUIDANCE in prompt def test_auto_skips_for_claude(self): from agent.prompt_builder import TOOL_USE_ENFORCEMENT_GUIDANCE - agent = self._make_agent(model="anthropic/claude-sonnet-4", tool_use_enforcement="auto") + + agent = self._make_agent( + model="anthropic/claude-sonnet-4", tool_use_enforcement="auto" + ) prompt = agent._build_system_prompt() assert TOOL_USE_ENFORCEMENT_GUIDANCE not in prompt def test_true_forces_for_all_models(self): from agent.prompt_builder import TOOL_USE_ENFORCEMENT_GUIDANCE - agent = self._make_agent(model="anthropic/claude-sonnet-4", tool_use_enforcement=True) + + agent = self._make_agent( + model="anthropic/claude-sonnet-4", tool_use_enforcement=True + ) prompt = agent._build_system_prompt() assert TOOL_USE_ENFORCEMENT_GUIDANCE in prompt def test_string_true_forces_for_all_models(self): from agent.prompt_builder import TOOL_USE_ENFORCEMENT_GUIDANCE - agent = self._make_agent(model="anthropic/claude-sonnet-4", tool_use_enforcement="true") + + agent = self._make_agent( + model="anthropic/claude-sonnet-4", tool_use_enforcement="true" + ) prompt = agent._build_system_prompt() assert TOOL_USE_ENFORCEMENT_GUIDANCE in prompt def test_always_forces_for_all_models(self): from agent.prompt_builder import TOOL_USE_ENFORCEMENT_GUIDANCE - agent = self._make_agent(model="deepseek/deepseek-r1", tool_use_enforcement="always") + + agent = self._make_agent( + model="deepseek/deepseek-r1", tool_use_enforcement="always" + ) prompt = agent._build_system_prompt() assert TOOL_USE_ENFORCEMENT_GUIDANCE in prompt def test_false_disables_for_gpt(self): from agent.prompt_builder import TOOL_USE_ENFORCEMENT_GUIDANCE + agent = self._make_agent(model="openai/gpt-4.1", tool_use_enforcement=False) prompt = agent._build_system_prompt() assert TOOL_USE_ENFORCEMENT_GUIDANCE not in prompt def test_string_false_disables(self): from agent.prompt_builder import TOOL_USE_ENFORCEMENT_GUIDANCE + agent = self._make_agent(model="openai/gpt-4.1", tool_use_enforcement="off") prompt = agent._build_system_prompt() assert TOOL_USE_ENFORCEMENT_GUIDANCE not in prompt def test_custom_list_matches(self): from agent.prompt_builder import TOOL_USE_ENFORCEMENT_GUIDANCE + agent = self._make_agent( model="deepseek/deepseek-r1", tool_use_enforcement=["deepseek", "gemini"], @@ -792,6 +826,7 @@ def test_custom_list_matches(self): def test_custom_list_no_match(self): from agent.prompt_builder import TOOL_USE_ENFORCEMENT_GUIDANCE + agent = self._make_agent( model="anthropic/claude-sonnet-4", tool_use_enforcement=["deepseek", "gemini"], @@ -801,6 +836,7 @@ def test_custom_list_no_match(self): def test_custom_list_case_insensitive(self): from agent.prompt_builder import TOOL_USE_ENFORCEMENT_GUIDANCE + agent = self._make_agent( model="openai/GPT-4.1", tool_use_enforcement=["GPT", "Codex"], @@ -811,6 +847,7 @@ def test_custom_list_case_insensitive(self): def test_no_tools_never_injects(self): """Even with enforcement=true, no injection when agent has no tools.""" from agent.prompt_builder import TOOL_USE_ENFORCEMENT_GUIDANCE + with ( patch("run_agent.get_tool_definitions", return_value=[]), patch("run_agent.check_toolset_requirements", return_value={}), @@ -941,7 +978,9 @@ def test_qwen_portal_formats_messages_and_metadata(self, agent): assert kwargs["metadata"]["sessionId"] == "sess-123" assert kwargs["extra_body"]["vl_high_resolution_images"] is True assert isinstance(kwargs["messages"][0]["content"], list) - assert kwargs["messages"][0]["content"][0]["cache_control"] == {"type": "ephemeral"} + assert kwargs["messages"][0]["content"][0]["cache_control"] == { + "type": "ephemeral" + } assert kwargs["messages"][2]["content"][0]["text"] == "hi" def test_qwen_portal_normalizes_bare_string_content_parts(self, agent): @@ -970,7 +1009,10 @@ def test_qwen_portal_sends_explicit_max_tokens(self, agent): agent.base_url = "https://portal.qwen.ai/v1" agent._base_url_lower = agent.base_url.lower() agent.max_tokens = 4096 - messages = [{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}] + messages = [ + {"role": "system", "content": "sys"}, + {"role": "user", "content": "hi"}, + ] kwargs = agent._build_api_kwargs(messages) assert kwargs["max_tokens"] == 4096 @@ -980,7 +1022,10 @@ def test_qwen_portal_default_max_tokens(self, agent): agent.base_url = "https://portal.qwen.ai/v1" agent._base_url_lower = agent.base_url.lower() agent.max_tokens = None - messages = [{"role": "system", "content": "sys"}, {"role": "user", "content": "hi"}] + messages = [ + {"role": "system", "content": "sys"}, + {"role": "user", "content": "hi"}, + ] kwargs = agent._build_api_kwargs(messages) assert kwargs["max_tokens"] == 65536 @@ -1125,7 +1170,10 @@ def test_result_truncation_over_100k(self, agent, tmp_path, monkeypatch): agent._execute_tool_calls(mock_msg, messages, "task-1") # Content should be replaced with persisted-output or truncation assert len(messages[0]["content"]) < 150_000 - assert ("Truncated" in messages[0]["content"] or "<persisted-output>" in messages[0]["content"]) + assert ( + "Truncated" in messages[0]["content"] + or "<persisted-output>" in messages[0]["content"] + ) def test_quiet_tool_output_suppressed_when_progress_callback_present(self, agent): tc = _mock_tool_call(name="web_search", arguments='{"q":"test"}', call_id="c1") @@ -1133,8 +1181,10 @@ def test_quiet_tool_output_suppressed_when_progress_callback_present(self, agent messages = [] agent.tool_progress_callback = lambda *args, **kwargs: None - with patch("run_agent.handle_function_call", return_value="search result"), \ - patch.object(agent, "_safe_print") as mock_print: + with ( + patch("run_agent.handle_function_call", return_value="search result"), + patch.object(agent, "_safe_print") as mock_print, + ): agent._execute_tool_calls(mock_msg, messages, "task-1") mock_print.assert_not_called() @@ -1147,8 +1197,10 @@ def test_quiet_tool_output_prints_without_progress_callback(self, agent): messages = [] agent.tool_progress_callback = None - with patch("run_agent.handle_function_call", return_value="search result"), \ - patch.object(agent, "_safe_print") as mock_print: + with ( + patch("run_agent.handle_function_call", return_value="search result"), + patch.object(agent, "_safe_print") as mock_print, + ): agent._execute_tool_calls(mock_msg, messages, "task-1") mock_print.assert_called_once() @@ -1165,7 +1217,9 @@ def test_vprint_suppressed_in_parseable_quiet_mode(self, agent): mock_print.assert_not_called() - def test_run_conversation_suppresses_retry_noise_in_parseable_quiet_mode(self, agent): + def test_run_conversation_suppresses_retry_noise_in_parseable_quiet_mode( + self, agent + ): class _RateLimitError(Exception): status_code = 429 @@ -1215,8 +1269,10 @@ def test_single_tool_uses_sequential_path(self, agent): def test_clarify_forces_sequential(self, agent): """Batch containing clarify should use sequential path.""" - tc1 = _mock_tool_call(name="web_search", arguments='{}', call_id="c1") - tc2 = _mock_tool_call(name="clarify", arguments='{"question":"ok?"}', call_id="c2") + tc1 = _mock_tool_call(name="web_search", arguments="{}", call_id="c1") + tc2 = _mock_tool_call( + name="clarify", arguments='{"question":"ok?"}', call_id="c2" + ) mock_msg = _mock_assistant_msg(content="", tool_calls=[tc1, tc2]) messages = [] with patch.object(agent, "_execute_tool_calls_sequential") as mock_seq: @@ -1227,8 +1283,10 @@ def test_clarify_forces_sequential(self, agent): def test_multiple_tools_uses_concurrent_path(self, agent): """Multiple read-only tools should use concurrent path.""" - tc1 = _mock_tool_call(name="web_search", arguments='{}', call_id="c1") - tc2 = _mock_tool_call(name="read_file", arguments='{"path":"x.py"}', call_id="c2") + tc1 = _mock_tool_call(name="web_search", arguments="{}", call_id="c1") + tc2 = _mock_tool_call( + name="read_file", arguments='{"path":"x.py"}', call_id="c2" + ) mock_msg = _mock_assistant_msg(content="", tool_calls=[tc1, tc2]) messages = [] with patch.object(agent, "_execute_tool_calls_sequential") as mock_seq: @@ -1239,8 +1297,10 @@ def test_multiple_tools_uses_concurrent_path(self, agent): def test_terminal_batch_forces_sequential(self, agent): """Stateful tools should not share the concurrent execution path.""" - tc1 = _mock_tool_call(name="web_search", arguments='{}', call_id="c1") - tc2 = _mock_tool_call(name="terminal", arguments='{"command":"pwd"}', call_id="c2") + tc1 = _mock_tool_call(name="web_search", arguments="{}", call_id="c1") + tc2 = _mock_tool_call( + name="terminal", arguments='{"command":"pwd"}', call_id="c2" + ) mock_msg = _mock_assistant_msg(content="", tool_calls=[tc1, tc2]) messages = [] with patch.object(agent, "_execute_tool_calls_sequential") as mock_seq: @@ -1251,8 +1311,14 @@ def test_terminal_batch_forces_sequential(self, agent): def test_write_batch_forces_sequential(self, agent): """File mutations should stay ordered within a turn.""" - tc1 = _mock_tool_call(name="read_file", arguments='{"path":"x.py"}', call_id="c1") - tc2 = _mock_tool_call(name="write_file", arguments='{"path":"x.py","content":"print(1)"}', call_id="c2") + tc1 = _mock_tool_call( + name="read_file", arguments='{"path":"x.py"}', call_id="c1" + ) + tc2 = _mock_tool_call( + name="write_file", + arguments='{"path":"x.py","content":"print(1)"}', + call_id="c2", + ) mock_msg = _mock_assistant_msg(content="", tool_calls=[tc1, tc2]) messages = [] with patch.object(agent, "_execute_tool_calls_sequential") as mock_seq: @@ -1303,7 +1369,7 @@ def test_overlapping_write_batch_forces_sequential(self, agent): def test_malformed_json_args_forces_sequential(self, agent): """Unparseable tool arguments should fall back to sequential.""" - tc1 = _mock_tool_call(name="web_search", arguments='{}', call_id="c1") + tc1 = _mock_tool_call(name="web_search", arguments="{}", call_id="c1") tc2 = _mock_tool_call(name="web_search", arguments="NOT JSON {{{", call_id="c2") mock_msg = _mock_assistant_msg(content="", tool_calls=[tc1, tc2]) messages = [] @@ -1315,8 +1381,10 @@ def test_malformed_json_args_forces_sequential(self, agent): def test_non_dict_args_forces_sequential(self, agent): """Tool arguments that parse to a non-dict type should fall back to sequential.""" - tc1 = _mock_tool_call(name="web_search", arguments='{}', call_id="c1") - tc2 = _mock_tool_call(name="web_search", arguments='"just a string"', call_id="c2") + tc1 = _mock_tool_call(name="web_search", arguments="{}", call_id="c1") + tc2 = _mock_tool_call( + name="web_search", arguments='"just a string"', call_id="c2" + ) mock_msg = _mock_assistant_msg(content="", tool_calls=[tc1, tc2]) messages = [] with patch.object(agent, "_execute_tool_calls_sequential") as mock_seq: @@ -1327,9 +1395,13 @@ def test_non_dict_args_forces_sequential(self, agent): def test_concurrent_executes_all_tools(self, agent): """Concurrent path should execute all tools and append results in order.""" - tc1 = _mock_tool_call(name="web_search", arguments='{"q":"alpha"}', call_id="c1") + tc1 = _mock_tool_call( + name="web_search", arguments='{"q":"alpha"}', call_id="c1" + ) tc2 = _mock_tool_call(name="web_search", arguments='{"q":"beta"}', call_id="c2") - tc3 = _mock_tool_call(name="web_search", arguments='{"q":"gamma"}', call_id="c3") + tc3 = _mock_tool_call( + name="web_search", arguments='{"q":"gamma"}', call_id="c3" + ) mock_msg = _mock_assistant_msg(content="", tool_calls=[tc1, tc2, tc3]) messages = [] @@ -1379,12 +1451,13 @@ def fake_handle(name, args, task_id, **kwargs): def test_concurrent_handles_tool_error(self, agent): """If one tool raises, others should still complete.""" - tc1 = _mock_tool_call(name="web_search", arguments='{}', call_id="c1") - tc2 = _mock_tool_call(name="web_search", arguments='{}', call_id="c2") + tc1 = _mock_tool_call(name="web_search", arguments="{}", call_id="c1") + tc2 = _mock_tool_call(name="web_search", arguments="{}", call_id="c2") mock_msg = _mock_assistant_msg(content="", tool_calls=[tc1, tc2]) messages = [] call_count = [0] + def fake_handle(name, args, task_id, **kwargs): call_count[0] += 1 if call_count[0] == 1: @@ -1402,8 +1475,8 @@ def fake_handle(name, args, task_id, **kwargs): def test_concurrent_interrupt_before_start(self, agent): """If interrupt is requested before concurrent execution, all tools are skipped.""" - tc1 = _mock_tool_call(name="web_search", arguments='{}', call_id="c1") - tc2 = _mock_tool_call(name="read_file", arguments='{}', call_id="c2") + tc1 = _mock_tool_call(name="web_search", arguments="{}", call_id="c1") + tc2 = _mock_tool_call(name="read_file", arguments="{}", call_id="c2") mock_msg = _mock_assistant_msg(content="", tool_calls=[tc1, tc2]) messages = [] @@ -1412,15 +1485,21 @@ def test_concurrent_interrupt_before_start(self, agent): agent._execute_tool_calls_concurrent(mock_msg, messages, "task-1") assert len(messages) == 2 - assert "cancelled" in messages[0]["content"].lower() or "skipped" in messages[0]["content"].lower() - assert "cancelled" in messages[1]["content"].lower() or "skipped" in messages[1]["content"].lower() + assert ( + "cancelled" in messages[0]["content"].lower() + or "skipped" in messages[0]["content"].lower() + ) + assert ( + "cancelled" in messages[1]["content"].lower() + or "skipped" in messages[1]["content"].lower() + ) def test_concurrent_truncates_large_results(self, agent, tmp_path, monkeypatch): """Concurrent path should save oversized results to file.""" monkeypatch.setenv("HERMES_HOME", str(tmp_path / ".hermes")) (tmp_path / ".hermes").mkdir() - tc1 = _mock_tool_call(name="web_search", arguments='{}', call_id="c1") - tc2 = _mock_tool_call(name="web_search", arguments='{}', call_id="c2") + tc1 = _mock_tool_call(name="web_search", arguments="{}", call_id="c1") + tc2 = _mock_tool_call(name="web_search", arguments="{}", call_id="c2") mock_msg = _mock_assistant_msg(content="", tool_calls=[tc1, tc2]) messages = [] big_result = "x" * 150_000 @@ -1431,14 +1510,16 @@ def test_concurrent_truncates_large_results(self, agent, tmp_path, monkeypatch): assert len(messages) == 2 for m in messages: assert len(m["content"]) < 150_000 - assert ("Truncated" in m["content"] or "<persisted-output>" in m["content"]) + assert "Truncated" in m["content"] or "<persisted-output>" in m["content"] def test_invoke_tool_dispatches_to_handle_function_call(self, agent): """_invoke_tool should route regular tools through handle_function_call.""" with patch("run_agent.handle_function_call", return_value="result") as mock_hfc: result = agent._invoke_tool("web_search", {"q": "test"}, "task-1") mock_hfc.assert_called_once_with( - "web_search", {"q": "test"}, "task-1", + "web_search", + {"q": "test"}, + "task-1", tool_call_id=None, session_id=agent.session_id, enabled_tools=list(agent.valid_tool_names), @@ -1447,31 +1528,57 @@ def test_invoke_tool_dispatches_to_handle_function_call(self, agent): assert result == "result" def test_sequential_tool_callbacks_fire_in_order(self, agent): - tool_call = _mock_tool_call(name="web_search", arguments='{"query":"hello"}', call_id="c1") + tool_call = _mock_tool_call( + name="web_search", arguments='{"query":"hello"}', call_id="c1" + ) mock_msg = _mock_assistant_msg(content="", tool_calls=[tool_call]) messages = [] starts = [] completes = [] - agent.tool_start_callback = lambda tool_call_id, function_name, function_args: starts.append((tool_call_id, function_name, function_args)) - agent.tool_complete_callback = lambda tool_call_id, function_name, function_args, function_result: completes.append((tool_call_id, function_name, function_args, function_result)) + agent.tool_start_callback = lambda tool_call_id, function_name, function_args: ( + starts.append((tool_call_id, function_name, function_args)) + ) + agent.tool_complete_callback = ( + lambda tool_call_id, function_name, function_args, function_result: ( + completes.append( + (tool_call_id, function_name, function_args, function_result) + ) + ) + ) with patch("run_agent.handle_function_call", return_value='{"success": true}'): agent._execute_tool_calls_sequential(mock_msg, messages, "task-1") assert starts == [("c1", "web_search", {"query": "hello"})] - assert completes == [("c1", "web_search", {"query": "hello"}, '{"success": true}')] + assert completes == [ + ("c1", "web_search", {"query": "hello"}, '{"success": true}') + ] def test_concurrent_tool_callbacks_fire_for_each_tool(self, agent): - tc1 = _mock_tool_call(name="web_search", arguments='{"query":"one"}', call_id="c1") - tc2 = _mock_tool_call(name="web_search", arguments='{"query":"two"}', call_id="c2") + tc1 = _mock_tool_call( + name="web_search", arguments='{"query":"one"}', call_id="c1" + ) + tc2 = _mock_tool_call( + name="web_search", arguments='{"query":"two"}', call_id="c2" + ) mock_msg = _mock_assistant_msg(content="", tool_calls=[tc1, tc2]) messages = [] starts = [] completes = [] - agent.tool_start_callback = lambda tool_call_id, function_name, function_args: starts.append((tool_call_id, function_name, function_args)) - agent.tool_complete_callback = lambda tool_call_id, function_name, function_args, function_result: completes.append((tool_call_id, function_name, function_args, function_result)) + agent.tool_start_callback = lambda tool_call_id, function_name, function_args: ( + starts.append((tool_call_id, function_name, function_args)) + ) + agent.tool_complete_callback = ( + lambda tool_call_id, function_name, function_args, function_result: ( + completes.append( + (tool_call_id, function_name, function_args, function_result) + ) + ) + ) - with patch("run_agent.handle_function_call", side_effect=['{"id":1}', '{"id":2}']): + with patch( + "run_agent.handle_function_call", side_effect=['{"id":1}', '{"id":2}'] + ): agent._execute_tool_calls_concurrent(mock_msg, messages, "task-1") assert starts == [ @@ -1484,18 +1591,24 @@ def test_concurrent_tool_callbacks_fire_for_each_tool(self, agent): def test_invoke_tool_handles_agent_level_tools(self, agent): """_invoke_tool should handle todo tool directly.""" - with patch("tools.todo_tool.todo_tool", return_value='{"ok":true}') as mock_todo: + with patch( + "tools.todo_tool.todo_tool", return_value='{"ok":true}' + ) as mock_todo: result = agent._invoke_tool("todo", {"todos": []}, "task-1") mock_todo.assert_called_once() assert "ok" in result - def test_invoke_tool_blocked_returns_error_and_skips_execution(self, agent, monkeypatch): + def test_invoke_tool_blocked_returns_error_and_skips_execution( + self, agent, monkeypatch + ): """_invoke_tool should return error JSON when a plugin blocks the tool.""" monkeypatch.setattr( "hermes_cli.plugins.get_pre_tool_call_block_message", lambda *args, **kwargs: "Blocked by test policy", ) - with patch("tools.todo_tool.todo_tool", side_effect=AssertionError("should not run")) as mock_todo: + with patch( + "tools.todo_tool.todo_tool", side_effect=AssertionError("should not run") + ) as mock_todo: result = agent._invoke_tool("todo", {"todos": []}, "task-1") assert json.loads(result) == {"error": "Blocked by test policy"} @@ -1507,16 +1620,23 @@ def test_invoke_tool_blocked_skips_handle_function_call(self, agent, monkeypatch "hermes_cli.plugins.get_pre_tool_call_block_message", lambda *args, **kwargs: "Blocked", ) - with patch("run_agent.handle_function_call", side_effect=AssertionError("should not run")): + with patch( + "run_agent.handle_function_call", + side_effect=AssertionError("should not run"), + ): result = agent._invoke_tool("web_search", {"q": "test"}, "task-1") assert json.loads(result) == {"error": "Blocked"} - def test_sequential_blocked_tool_skips_checkpoints_and_callbacks(self, agent, monkeypatch): + def test_sequential_blocked_tool_skips_checkpoints_and_callbacks( + self, agent, monkeypatch + ): """Sequential path: blocked tool should not trigger checkpoints or start callbacks.""" - tool_call = _mock_tool_call(name="write_file", - arguments='{"path":"test.txt","content":"hello"}', - call_id="c1") + tool_call = _mock_tool_call( + name="write_file", + arguments='{"path":"test.txt","content":"hello"}', + call_id="c1", + ) mock_msg = _mock_assistant_msg(content="", tool_calls=[tool_call]) messages = [] @@ -1532,7 +1652,10 @@ def test_sequential_blocked_tool_skips_checkpoints_and_callbacks(self, agent, mo starts = [] agent.tool_start_callback = lambda *a: starts.append(a) - with patch("run_agent.handle_function_call", side_effect=AssertionError("should not run")): + with patch( + "run_agent.handle_function_call", + side_effect=AssertionError("should not run"), + ): agent._execute_tool_calls_sequential(mock_msg, messages, "task-1") agent._checkpoint_mgr.ensure_checkpoint.assert_not_called() @@ -1548,9 +1671,14 @@ def test_blocked_memory_tool_does_not_reset_counter(self, agent, monkeypatch): "hermes_cli.plugins.get_pre_tool_call_block_message", lambda *args, **kwargs: "Blocked", ) - with patch("tools.memory_tool.memory_tool", side_effect=AssertionError("should not run")): + with patch( + "tools.memory_tool.memory_tool", + side_effect=AssertionError("should not run"), + ): result = agent._invoke_tool( - "memory", {"action": "add", "target": "memory", "content": "x"}, "task-1", + "memory", + {"action": "add", "target": "memory", "content": "x"}, + "task-1", ) assert json.loads(result) == {"error": "Blocked"} @@ -1562,36 +1690,45 @@ class TestPathsOverlap: def test_same_path_overlaps(self): from run_agent import _paths_overlap + assert _paths_overlap(Path("src/a.py"), Path("src/a.py")) def test_siblings_do_not_overlap(self): from run_agent import _paths_overlap + assert not _paths_overlap(Path("src/a.py"), Path("src/b.py")) def test_parent_child_overlap(self): from run_agent import _paths_overlap + assert _paths_overlap(Path("src"), Path("src/sub/a.py")) def test_different_roots_do_not_overlap(self): from run_agent import _paths_overlap + assert not _paths_overlap(Path("src/a.py"), Path("other/a.py")) def test_nested_vs_flat_do_not_overlap(self): from run_agent import _paths_overlap + assert not _paths_overlap(Path("src/sub/a.py"), Path("src/a.py")) def test_empty_paths_do_not_overlap(self): from run_agent import _paths_overlap + assert not _paths_overlap(Path(""), Path("")) def test_one_empty_path_does_not_overlap(self): from run_agent import _paths_overlap + assert not _paths_overlap(Path(""), Path("src/a.py")) assert not _paths_overlap(Path("src/a.py"), Path("")) class TestParallelScopePathNormalization: - def test_extract_parallel_scope_path_normalizes_relative_to_cwd(self, tmp_path, monkeypatch): + def test_extract_parallel_scope_path_normalizes_relative_to_cwd( + self, tmp_path, monkeypatch + ): from run_agent import _extract_parallel_scope_path monkeypatch.chdir(tmp_path) @@ -1600,7 +1737,9 @@ def test_extract_parallel_scope_path_normalizes_relative_to_cwd(self, tmp_path, assert scoped == tmp_path / "notes.txt" - def test_extract_parallel_scope_path_treats_relative_and_absolute_same_file_as_same_scope(self, tmp_path, monkeypatch): + def test_extract_parallel_scope_path_treats_relative_and_absolute_same_file_as_same_scope( + self, tmp_path, monkeypatch + ): from run_agent import _extract_parallel_scope_path, _paths_overlap monkeypatch.chdir(tmp_path) @@ -1612,12 +1751,22 @@ def test_extract_parallel_scope_path_treats_relative_and_absolute_same_file_as_s assert rel_scoped == abs_scoped assert _paths_overlap(rel_scoped, abs_scoped) - def test_should_parallelize_tool_batch_rejects_same_file_with_mixed_path_spellings(self, tmp_path, monkeypatch): + def test_should_parallelize_tool_batch_rejects_same_file_with_mixed_path_spellings( + self, tmp_path, monkeypatch + ): from run_agent import _should_parallelize_tool_batch monkeypatch.chdir(tmp_path) - tc1 = _mock_tool_call(name="write_file", arguments='{"path":"notes.txt","content":"one"}', call_id="c1") - tc2 = _mock_tool_call(name="write_file", arguments=f'{{"path":"{tmp_path / "notes.txt"}","content":"two"}}', call_id="c2") + tc1 = _mock_tool_call( + name="write_file", + arguments='{"path":"notes.txt","content":"one"}', + call_id="c1", + ) + tc2 = _mock_tool_call( + name="write_file", + arguments=f'{{"path":"{tmp_path / "notes.txt"}","content":"two"}}', + call_id="c2", + ) assert not _should_parallelize_tool_batch([tc1, tc2]) @@ -1692,7 +1841,9 @@ def test_tool_calls_then_stop(self, agent): resp2 = _mock_response(content="Done searching", finish_reason="stop") agent.client.chat.completions.create.side_effect = [resp1, resp2] with ( - patch("run_agent.handle_function_call", return_value="search result") as mock_handle_function_call, + patch( + "run_agent.handle_function_call", return_value="search result" + ) as mock_handle_function_call, patch.object(agent, "_persist_session"), patch.object(agent, "_save_trajectory"), patch.object(agent, "_cleanup_task_resources"), @@ -1701,7 +1852,9 @@ def test_tool_calls_then_stop(self, agent): assert result["final_response"] == "Done searching" assert result["api_calls"] == 2 assert mock_handle_function_call.call_args.kwargs["tool_call_id"] == "c1" - assert mock_handle_function_call.call_args.kwargs["session_id"] == agent.session_id + assert ( + mock_handle_function_call.call_args.kwargs["session_id"] == agent.session_id + ) def test_request_scoped_api_hooks_fire_for_each_api_call(self, agent): self._setup_agent(agent) @@ -1727,13 +1880,17 @@ def _record_hook(name, **kwargs): assert result["final_response"] == "Done searching" pre_request_calls = [kw for name, kw in hook_calls if name == "pre_api_request"] - post_request_calls = [kw for name, kw in hook_calls if name == "post_api_request"] + post_request_calls = [ + kw for name, kw in hook_calls if name == "post_api_request" + ] assert len(pre_request_calls) == 2 assert len(post_request_calls) == 2 assert [call["api_call_count"] for call in pre_request_calls] == [1, 2] assert [call["api_call_count"] for call in post_request_calls] == [1, 2] assert all(call["session_id"] == agent.session_id for call in pre_request_calls) - assert all("message_count" in c and "messages" not in c for c in pre_request_calls) + assert all( + "message_count" in c and "messages" not in c for c in pre_request_calls + ) assert all("usage" in c and "response" not in c for c in post_request_calls) def test_interrupt_breaks_loop(self, agent): @@ -1791,7 +1948,9 @@ def test_reasoning_only_local_resumed_no_compression_triggered(self, agent): # 6 responses: original + 2 prefill + 3 retries after prefill exhaustion with ( - patch.object(agent, "_interruptible_api_call", side_effect=[empty_resp] * 6), + patch.object( + agent, "_interruptible_api_call", side_effect=[empty_resp] * 6 + ), patch.object(agent, "_compress_context") as mock_compress, patch.object(agent, "_persist_session"), patch.object(agent, "_save_trajectory"), @@ -1859,7 +2018,10 @@ def test_truly_empty_response_retries_3_times_then_empty(self, agent): empty_resp = _mock_response(content=None, finish_reason="stop") # 4 responses: 1 original + 3 nudge retries, all empty agent.client.chat.completions.create.side_effect = [ - empty_resp, empty_resp, empty_resp, empty_resp, + empty_resp, + empty_resp, + empty_resp, + empty_resp, ] with ( patch.object(agent, "_persist_session"), @@ -1897,7 +2059,9 @@ def test_empty_response_triggers_fallback_provider(self, agent): self._setup_agent(agent) agent.base_url = "http://127.0.0.1:1234/v1" # Configure a fallback chain - agent._fallback_chain = [{"provider": "openrouter", "model": "anthropic/claude-sonnet-4"}] + agent._fallback_chain = [ + {"provider": "openrouter", "model": "anthropic/claude-sonnet-4"} + ] agent._fallback_index = 0 agent._fallback_activated = False @@ -1905,7 +2069,11 @@ def test_empty_response_triggers_fallback_provider(self, agent): content_resp = _mock_response(content="Fallback answer.", finish_reason="stop") # 4 empty (1 orig + 3 retries), then fallback model answers agent.client.chat.completions.create.side_effect = [ - empty_resp, empty_resp, empty_resp, empty_resp, content_resp, + empty_resp, + empty_resp, + empty_resp, + empty_resp, + content_resp, ] fallback_called = {"called": False} @@ -1935,7 +2103,9 @@ def test_empty_response_fallback_also_empty_returns_empty(self, agent): """If fallback also returns empty, final response is (empty).""" self._setup_agent(agent) agent.base_url = "http://127.0.0.1:1234/v1" - agent._fallback_chain = [{"provider": "openrouter", "model": "anthropic/claude-sonnet-4"}] + agent._fallback_chain = [ + {"provider": "openrouter", "model": "anthropic/claude-sonnet-4"} + ] agent._fallback_index = 0 agent._fallback_activated = False @@ -1943,8 +2113,14 @@ def test_empty_response_fallback_also_empty_returns_empty(self, agent): # 4 empty from primary (1 + 3 retries), fallback activated, # then 4 more empty from fallback (1 + 3 retries), no more fallbacks agent.client.chat.completions.create.side_effect = [ - empty_resp, empty_resp, empty_resp, empty_resp, # primary exhausted - empty_resp, empty_resp, empty_resp, empty_resp, # fallback exhausted + empty_resp, + empty_resp, + empty_resp, + empty_resp, # primary exhausted + empty_resp, + empty_resp, + empty_resp, + empty_resp, # fallback exhausted ] def _mock_fallback(): @@ -1974,7 +2150,10 @@ def test_empty_response_emits_status_for_gateway(self, agent): empty_resp = _mock_response(content=None, finish_reason="stop") # 4 empty: 1 original + 3 retries, all empty, no fallback agent.client.chat.completions.create.side_effect = [ - empty_resp, empty_resp, empty_resp, empty_resp, + empty_resp, + empty_resp, + empty_resp, + empty_resp, ] status_messages = [] @@ -1993,9 +2172,17 @@ def _capture_status(msg): assert result["final_response"] == "(empty)" # Should have emitted retry statuses (3 retries) + final failure retry_msgs = [m for m in status_messages if "retrying" in m.lower()] - assert len(retry_msgs) == 3, f"Expected 3 retry status messages, got {len(retry_msgs)}: {status_messages}" - failure_msgs = [m for m in status_messages if "no content" in m.lower() or "no fallback" in m.lower()] - assert len(failure_msgs) >= 1, f"Expected at least 1 failure status, got: {status_messages}" + assert len(retry_msgs) == 3, ( + f"Expected 3 retry status messages, got {len(retry_msgs)}: {status_messages}" + ) + failure_msgs = [ + m + for m in status_messages + if "no content" in m.lower() or "no fallback" in m.lower() + ] + assert len(failure_msgs) >= 1, ( + f"Expected at least 1 failure status, got: {status_messages}" + ) def test_partial_stream_recovery_uses_streamed_content(self, agent): """When streaming fails after partial delivery, recovered partial content becomes final response.""" @@ -2007,7 +2194,9 @@ def test_partial_stream_recovery_uses_streamed_content(self, agent): ) agent.client.chat.completions.create.return_value = partial_resp # Simulate that streaming had already delivered this text - agent._current_streamed_assistant_text = "Here is the partial answer that was stream" + agent._current_streamed_assistant_text = ( + "Here is the partial answer that was stream" + ) with ( patch.object(agent, "_persist_session"), patch.object(agent, "_save_trajectory"), @@ -2028,7 +2217,9 @@ def test_partial_stream_recovery_on_empty_stub(self, agent): def _fake_api_call(api_kwargs): # Simulate what streaming does: accumulate text before returning # a stub with no content (connection died mid-stream) - agent._current_streamed_assistant_text = "The answer to your question is that" + agent._current_streamed_assistant_text = ( + "The answer to your question is that" + ) return empty_stub status_messages = [] @@ -2049,11 +2240,17 @@ def _capture_status(msg): assert result["final_response"] == "The answer to your question is that" assert result["api_calls"] == 1 # No wasted retries # Should emit the stream-interrupted status, NOT the empty-retry status - recovery_msgs = [m for m in status_messages if "stream interrupted" in m.lower()] - assert len(recovery_msgs) >= 1, f"Expected stream recovery status, got: {status_messages}" + recovery_msgs = [ + m for m in status_messages if "stream interrupted" in m.lower() + ] + assert len(recovery_msgs) >= 1, ( + f"Expected stream recovery status, got: {status_messages}" + ) # Should NOT have retry statuses retry_msgs = [m for m in status_messages if "retrying" in m.lower()] - assert len(retry_msgs) == 0, f"Should not retry when stream content exists: {status_messages}" + assert len(retry_msgs) == 0, ( + f"Should not retry when stream content exists: {status_messages}" + ) def test_partial_stream_recovery_preempts_prior_turn_fallback(self, agent): """Partial streamed content takes priority over _last_content_with_tools fallback.""" @@ -2065,7 +2262,9 @@ def test_partial_stream_recovery_preempts_prior_turn_fallback(self, agent): def _fake_api_call(api_kwargs): # Simulate partial streaming before connection death - agent._current_streamed_assistant_text = "Fresh partial content from this turn" + agent._current_streamed_assistant_text = ( + "Fresh partial content from this turn" + ) return empty_stub with ( @@ -2157,7 +2356,9 @@ def test_glm_prompt_exceeds_max_length_triggers_compression(self, agent): "Error code: 400 - {'error': {'code': '1261', 'message': 'Prompt exceeds max length'}}" ) err_400.status_code = 400 - ok_resp = _mock_response(content="Recovered after compression", finish_reason="stop") + ok_resp = _mock_response( + content="Recovered after compression", finish_reason="stop" + ) agent.client.chat.completions.create.side_effect = [err_400, ok_resp] prefill = [ {"role": "user", "content": "previous question"}, @@ -2198,9 +2399,14 @@ def test_length_finish_reason_requests_continuation(self, agent): assert result["api_calls"] == 2 assert result["final_response"] == "Part 1 Part 2" - second_call_messages = agent.client.chat.completions.create.call_args_list[1].kwargs["messages"] + second_call_messages = agent.client.chat.completions.create.call_args_list[ + 1 + ].kwargs["messages"] assert second_call_messages[-1]["role"] == "user" - assert "truncated by the output length limit" in second_call_messages[-1]["content"] + assert ( + "truncated by the output length limit" + in second_call_messages[-1]["content"] + ) def test_length_thinking_exhausted_skips_continuation(self, agent): """When finish_reason='length' but content is only thinking, skip retries.""" @@ -2247,7 +2453,9 @@ def test_length_empty_content_without_think_tags_retries_normally(self, agent): assert result["api_calls"] == 3 assert result["completed"] is False - def test_length_with_tool_calls_returns_partial_without_executing_tools(self, agent): + def test_length_with_tool_calls_returns_partial_without_executing_tools( + self, agent + ): self._setup_agent(agent) bad_tc = _mock_tool_call( name="write_file", @@ -2281,7 +2489,9 @@ def test_truncated_tool_call_retries_once_before_refusing(self, agent): call_id="c1", ) truncated_resp = _mock_response( - content="", finish_reason="length", tool_calls=[bad_tc], + content="", + finish_reason="length", + tool_calls=[bad_tc], ) good_tc = _mock_tool_call( name="write_file", @@ -2289,10 +2499,14 @@ def test_truncated_tool_call_retries_once_before_refusing(self, agent): call_id="c2", ) good_resp = _mock_response( - content="", finish_reason="stop", tool_calls=[good_tc], + content="", + finish_reason="stop", + tool_calls=[good_tc], ) with ( - patch("run_agent.handle_function_call", return_value='{"success":true}') as mock_hfc, + patch( + "run_agent.handle_function_call", return_value='{"success":true}' + ) as mock_hfc, patch.object(agent, "_persist_session"), patch.object(agent, "_save_trajectory"), patch.object(agent, "_cleanup_task_resources"), @@ -2301,7 +2515,9 @@ def test_truncated_tool_call_retries_once_before_refusing(self, agent): # Third: final text response. final_resp = _mock_response(content="Done!", finish_reason="stop") agent.client.chat.completions.create.side_effect = [ - truncated_resp, good_resp, final_resp, + truncated_resp, + good_resp, + final_resp, ] result = agent.run_conversation("write the report") @@ -2321,7 +2537,9 @@ def test_truncated_tool_args_detected_when_finish_reason_not_length(self, agent) call_id="c1", ) resp = _mock_response( - content="", finish_reason="tool_calls", tool_calls=[bad_tc], + content="", + finish_reason="tool_calls", + tool_calls=[bad_tc], ) agent.client.chat.completions.create.return_value = resp @@ -2417,7 +2635,9 @@ def test_build_api_kwargs_error_no_unbound_local(self, agent): """ self._setup_agent(agent) with ( - patch.object(agent, "_build_api_kwargs", side_effect=ValueError("bad messages")), + patch.object( + agent, "_build_api_kwargs", side_effect=ValueError("bad messages") + ), patch.object(agent, "_persist_session"), patch.object(agent, "_save_trajectory"), patch.object(agent, "_cleanup_task_resources"), @@ -2461,7 +2681,9 @@ def test_flush_sentinel_stripped_from_api_messages(self, agent_with_memory_tool) agent.client.chat.completions.create.return_value = mock_response # Bypass auxiliary client so flush uses agent.client directly - with patch("agent.auxiliary_client.call_llm", side_effect=RuntimeError("no provider")): + with patch( + "agent.auxiliary_client.call_llm", side_effect=RuntimeError("no provider") + ): agent.flush_memories(messages, min_turns=0) # Check what was actually sent to the API @@ -2591,7 +2813,9 @@ def mark_exhausted_and_rotate(self, *, status_code, error_context=None): assert retry_same is False agent._swap_credential.assert_called_once_with(next_entry) - def test_recover_with_pool_rotates_on_billing_reason_even_with_http_400(self, agent): + def test_recover_with_pool_rotates_on_billing_reason_even_with_http_400( + self, agent + ): next_entry = SimpleNamespace(label="secondary") class _Pool: @@ -2645,7 +2869,6 @@ def mark_exhausted_and_rotate(self, *, status_code, error_context=None): assert retry_same is False agent._swap_credential.assert_called_once_with(next_entry) - def test_recover_with_pool_refreshes_on_401(self, agent): """401 with successful refresh should swap to refreshed credential.""" refreshed_entry = SimpleNamespace(label="refreshed-primary", id="abc") @@ -2750,7 +2973,10 @@ def mark_exhausted_and_rotate(self, *, status_code, error_context=None): recovered, retry_same = agent._recover_with_credential_pool( status_code=429, has_retried_429=True, - error_context={"reason": "device_code_exhausted", "reset_at": "2026-04-12T10:30:00Z"}, + error_context={ + "reason": "device_code_exhausted", + "reset_at": "2026-04-12T10:30:00Z", + }, ) assert recovered is True @@ -2787,6 +3013,7 @@ def test_not_tricked_by_openai_in_openrouter_url(self, agent): # System prompt stability for prompt caching # --------------------------------------------------------------------------- + class TestSystemPromptStability: """Verify that the system prompt stays stable across turns for cache hits.""" @@ -2882,6 +3109,7 @@ def test_fresh_build_when_db_has_no_prompt(self, agent): # Empty string is falsy, so should fall through to fresh build assert "Hermes Agent" in agent._cached_system_prompt + class TestBudgetPressure: """Budget exhaustion grace call system.""" @@ -2898,6 +3126,7 @@ def test_write_delegates_normally(self): """When stdout is healthy, _SafeWriter is transparent.""" from run_agent import _SafeWriter from io import StringIO + inner = StringIO() writer = _SafeWriter(inner) writer.write("hello") @@ -2907,6 +3136,7 @@ def test_write_catches_oserror(self): """OSError on write is silently caught, returns len(data).""" from run_agent import _SafeWriter from unittest.mock import MagicMock + inner = MagicMock() inner.write.side_effect = OSError(5, "Input/output error") writer = _SafeWriter(inner) @@ -2917,6 +3147,7 @@ def test_flush_catches_oserror(self): """OSError on flush is silently caught.""" from run_agent import _SafeWriter from unittest.mock import MagicMock + inner = MagicMock() inner.flush.side_effect = OSError(5, "Input/output error") writer = _SafeWriter(inner) @@ -2927,6 +3158,7 @@ def test_print_survives_broken_stdout(self, monkeypatch): import sys from run_agent import _SafeWriter from unittest.mock import MagicMock + broken = MagicMock() broken.write.side_effect = OSError(5, "Input/output error") original = sys.stdout @@ -2940,6 +3172,7 @@ def test_installed_in_run_conversation(self, agent): """run_conversation installs _SafeWriter on stdio.""" import sys from run_agent import _SafeWriter + resp = _mock_response(content="Done", finish_reason="stop") agent.client.chat.completions.create.return_value = resp original_stdout = sys.stdout @@ -2965,6 +3198,7 @@ def test_double_wrap_prevented(self): import sys from run_agent import _SafeWriter from io import StringIO + inner = StringIO() wrapped = _SafeWriter(inner) # isinstance check should prevent double-wrapping @@ -3009,15 +3243,30 @@ def test_max_tokens_passed_to_anthropic(self, agent): agent.reasoning_config = None with patch("agent.anthropic_adapter.build_anthropic_kwargs") as mock_build: - mock_build.return_value = {"model": "claude-sonnet-4-20250514", "messages": [], "max_tokens": 4096} + mock_build.return_value = { + "model": "claude-sonnet-4-20250514", + "messages": [], + "max_tokens": 4096, + } agent._build_api_kwargs([{"role": "user", "content": "test"}]) _, kwargs = mock_build.call_args if not kwargs: - kwargs = dict(zip( - ["model", "messages", "tools", "max_tokens", "reasoning_config"], - mock_build.call_args[0], - )) - assert kwargs.get("max_tokens") == 4096 or mock_build.call_args[1].get("max_tokens") == 4096 + kwargs = dict( + zip( + [ + "model", + "messages", + "tools", + "max_tokens", + "reasoning_config", + ], + mock_build.call_args[0], + ) + ) + assert ( + kwargs.get("max_tokens") == 4096 + or mock_build.call_args[1].get("max_tokens") == 4096 + ) def test_max_tokens_none_when_unset(self, agent): agent.api_mode = "anthropic_messages" @@ -3025,7 +3274,11 @@ def test_max_tokens_none_when_unset(self, agent): agent.reasoning_config = None with patch("agent.anthropic_adapter.build_anthropic_kwargs") as mock_build: - mock_build.return_value = {"model": "claude-sonnet-4-20250514", "messages": [], "max_tokens": 16384} + mock_build.return_value = { + "model": "claude-sonnet-4-20250514", + "messages": [], + "max_tokens": 16384, + } agent._build_api_kwargs([{"role": "user", "content": "test"}]) call_args = mock_build.call_args # max_tokens should be None (let adapter use its default) @@ -3040,32 +3293,55 @@ def test_build_api_kwargs_converts_multimodal_user_image_to_text(self, agent): agent.api_mode = "anthropic_messages" agent.reasoning_config = None - api_messages = [{ - "role": "user", - "content": [ - {"type": "text", "text": "Can you see this now?"}, - {"type": "image_url", "image_url": {"url": "https://example.com/cat.png"}}, - ], - }] + api_messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Can you see this now?"}, + { + "type": "image_url", + "image_url": {"url": "https://example.com/cat.png"}, + }, + ], + } + ] with ( - patch("tools.vision_tools.vision_analyze_tool", new=AsyncMock(return_value=json.dumps({"success": True, "analysis": "A cat sitting on a chair."}))), + patch( + "tools.vision_tools.vision_analyze_tool", + new=AsyncMock( + return_value=json.dumps( + {"success": True, "analysis": "A cat sitting on a chair."} + ) + ), + ), patch("agent.anthropic_adapter.build_anthropic_kwargs") as mock_build, ): - mock_build.return_value = {"model": "claude-sonnet-4-20250514", "messages": [], "max_tokens": 4096} + mock_build.return_value = { + "model": "claude-sonnet-4-20250514", + "messages": [], + "max_tokens": 4096, + } agent._build_api_kwargs(api_messages) - kwargs = mock_build.call_args.kwargs or dict(zip( - ["model", "messages", "tools", "max_tokens", "reasoning_config"], - mock_build.call_args.args, - )) + kwargs = mock_build.call_args.kwargs or dict( + zip( + ["model", "messages", "tools", "max_tokens", "reasoning_config"], + mock_build.call_args.args, + ) + ) transformed = kwargs["messages"] assert isinstance(transformed[0]["content"], str) assert "A cat sitting on a chair." in transformed[0]["content"] assert "Can you see this now?" in transformed[0]["content"] - assert "vision_analyze with image_url: https://example.com/cat.png" in transformed[0]["content"] + assert ( + "vision_analyze with image_url: https://example.com/cat.png" + in transformed[0]["content"] + ) - def test_build_api_kwargs_reuses_cached_image_analysis_for_duplicate_images(self, agent): + def test_build_api_kwargs_reuses_cached_image_analysis_for_duplicate_images( + self, agent + ): agent.api_mode = "anthropic_messages" agent.reasoning_config = None data_url = "data:image/png;base64,QUFBQQ==" @@ -3087,12 +3363,20 @@ def test_build_api_kwargs_reuses_cached_image_analysis_for_duplicate_images(self }, ] - mock_vision = AsyncMock(return_value=json.dumps({"success": True, "analysis": "A small test image."})) + mock_vision = AsyncMock( + return_value=json.dumps( + {"success": True, "analysis": "A small test image."} + ) + ) with ( patch("tools.vision_tools.vision_analyze_tool", new=mock_vision), patch("agent.anthropic_adapter.build_anthropic_kwargs") as mock_build, ): - mock_build.return_value = {"model": "claude-sonnet-4-20250514", "messages": [], "max_tokens": 4096} + mock_build.return_value = { + "model": "claude-sonnet-4-20250514", + "messages": [], + "max_tokens": 4096, + } agent._build_api_kwargs(api_messages) assert mock_vision.await_count == 1 @@ -3103,7 +3387,10 @@ class TestFallbackAnthropicProvider: def test_fallback_to_anthropic_sets_api_mode(self, agent): agent._fallback_activated = False - agent._fallback_model = {"provider": "anthropic", "model": "claude-sonnet-4-20250514"} + agent._fallback_model = { + "provider": "anthropic", + "model": "claude-sonnet-4-20250514", + } agent._fallback_chain = [agent._fallback_model] agent._fallback_index = 0 @@ -3112,7 +3399,10 @@ def test_fallback_to_anthropic_sets_api_mode(self, agent): mock_client.api_key = "sk-ant-api03-test" with ( - patch("agent.auxiliary_client.resolve_provider_client", return_value=(mock_client, None)), + patch( + "agent.auxiliary_client.resolve_provider_client", + return_value=(mock_client, None), + ), patch("agent.anthropic_adapter.build_anthropic_client") as mock_build, patch("agent.anthropic_adapter.resolve_anthropic_token", return_value=None), ): @@ -3126,7 +3416,10 @@ def test_fallback_to_anthropic_sets_api_mode(self, agent): def test_fallback_to_anthropic_enables_prompt_caching(self, agent): agent._fallback_activated = False - agent._fallback_model = {"provider": "anthropic", "model": "claude-sonnet-4-20250514"} + agent._fallback_model = { + "provider": "anthropic", + "model": "claude-sonnet-4-20250514", + } agent._fallback_chain = [agent._fallback_model] agent._fallback_index = 0 @@ -3135,8 +3428,14 @@ def test_fallback_to_anthropic_enables_prompt_caching(self, agent): mock_client.api_key = "sk-ant-api03-test" with ( - patch("agent.auxiliary_client.resolve_provider_client", return_value=(mock_client, None)), - patch("agent.anthropic_adapter.build_anthropic_client", return_value=MagicMock()), + patch( + "agent.auxiliary_client.resolve_provider_client", + return_value=(mock_client, None), + ), + patch( + "agent.anthropic_adapter.build_anthropic_client", + return_value=MagicMock(), + ), patch("agent.anthropic_adapter.resolve_anthropic_token", return_value=None), ): agent._try_activate_fallback() @@ -3145,7 +3444,10 @@ def test_fallback_to_anthropic_enables_prompt_caching(self, agent): def test_fallback_to_openrouter_uses_openai_client(self, agent): agent._fallback_activated = False - agent._fallback_model = {"provider": "openrouter", "model": "anthropic/claude-sonnet-4"} + agent._fallback_model = { + "provider": "openrouter", + "model": "anthropic/claude-sonnet-4", + } agent._fallback_chain = [agent._fallback_model] agent._fallback_index = 0 @@ -3153,7 +3455,10 @@ def test_fallback_to_openrouter_uses_openai_client(self, agent): mock_client.base_url = "https://openrouter.ai/api/v1" mock_client.api_key = "sk-or-test" - with patch("agent.auxiliary_client.resolve_provider_client", return_value=(mock_client, None)): + with patch( + "agent.auxiliary_client.resolve_provider_client", + return_value=(mock_client, None), + ): result = agent._try_activate_fallback() assert result is True @@ -3163,7 +3468,9 @@ def test_fallback_to_openrouter_uses_openai_client(self, agent): def test_aiagent_uses_copilot_acp_client(): with ( - patch("run_agent.get_tool_definitions", return_value=_make_tool_defs("web_search")), + patch( + "run_agent.get_tool_definitions", return_value=_make_tool_defs("web_search") + ), patch("run_agent.check_toolset_requirements", return_value={}), patch("run_agent.OpenAI") as mock_openai, patch("agent.copilot_acp_client.CopilotACPClient") as mock_acp_client, @@ -3250,8 +3557,13 @@ class ClientWithHttpClient: def __init__(self, http_closed: bool): self._client = SimpleNamespace(is_closed=http_closed) - assert AIAgent._is_openai_client_closed(ClientWithHttpClient(http_closed=False)) is False - assert AIAgent._is_openai_client_closed(ClientWithHttpClient(http_closed=True)) is True + assert ( + AIAgent._is_openai_client_closed(ClientWithHttpClient(http_closed=False)) + is False + ) + assert ( + AIAgent._is_openai_client_closed(ClientWithHttpClient(http_closed=True)) is True + ) class TestAnthropicBaseUrlPassthrough: @@ -3259,7 +3571,10 @@ class TestAnthropicBaseUrlPassthrough: def test_custom_proxy_base_url_passed_through(self): with ( - patch("run_agent.get_tool_definitions", return_value=_make_tool_defs("web_search")), + patch( + "run_agent.get_tool_definitions", + return_value=_make_tool_defs("web_search"), + ), patch("run_agent.check_toolset_requirements", return_value={}), patch("agent.anthropic_adapter.build_anthropic_client") as mock_build, ): @@ -3278,7 +3593,10 @@ def test_custom_proxy_base_url_passed_through(self): def test_none_base_url_passed_as_none(self): with ( - patch("run_agent.get_tool_definitions", return_value=_make_tool_defs("web_search")), + patch( + "run_agent.get_tool_definitions", + return_value=_make_tool_defs("web_search"), + ), patch("run_agent.check_toolset_requirements", return_value={}), patch("agent.anthropic_adapter.build_anthropic_client") as mock_build, ): @@ -3299,7 +3617,10 @@ def test_none_base_url_passed_as_none(self): class TestAnthropicCredentialRefresh: def test_try_refresh_anthropic_client_credentials_rebuilds_client(self): with ( - patch("run_agent.get_tool_definitions", return_value=_make_tool_defs("web_search")), + patch( + "run_agent.get_tool_definitions", + return_value=_make_tool_defs("web_search"), + ), patch("run_agent.check_toolset_requirements", return_value={}), patch("agent.anthropic_adapter.build_anthropic_client") as mock_build, ): @@ -3320,21 +3641,37 @@ def test_try_refresh_anthropic_client_credentials_rebuilds_client(self): agent.provider = "anthropic" with ( - patch("agent.anthropic_adapter.resolve_anthropic_token", return_value="sk-ant-oat01-fresh-token"), - patch("agent.anthropic_adapter.build_anthropic_client", return_value=new_client) as rebuild, + patch( + "agent.anthropic_adapter.resolve_anthropic_token", + return_value="sk-ant-oat01-fresh-token", + ), + patch( + "agent.anthropic_adapter.build_anthropic_client", + return_value=new_client, + ) as rebuild, ): assert agent._try_refresh_anthropic_client_credentials() is True old_client.close.assert_called_once() - rebuild.assert_called_once_with("sk-ant-oat01-fresh-token", "https://api.anthropic.com") + rebuild.assert_called_once_with( + "sk-ant-oat01-fresh-token", "https://api.anthropic.com" + ) assert agent._anthropic_client is new_client assert agent._anthropic_api_key == "sk-ant-oat01-fresh-token" - def test_try_refresh_anthropic_client_credentials_returns_false_when_token_unchanged(self): + def test_try_refresh_anthropic_client_credentials_returns_false_when_token_unchanged( + self, + ): with ( - patch("run_agent.get_tool_definitions", return_value=_make_tool_defs("web_search")), + patch( + "run_agent.get_tool_definitions", + return_value=_make_tool_defs("web_search"), + ), patch("run_agent.check_toolset_requirements", return_value={}), - patch("agent.anthropic_adapter.build_anthropic_client", return_value=MagicMock()), + patch( + "agent.anthropic_adapter.build_anthropic_client", + return_value=MagicMock(), + ), ): agent = AIAgent( api_key="sk-ant-oat01-same-token", @@ -3349,7 +3686,10 @@ def test_try_refresh_anthropic_client_credentials_returns_false_when_token_uncha agent._anthropic_api_key = "sk-ant-oat01-same-token" with ( - patch("agent.anthropic_adapter.resolve_anthropic_token", return_value="sk-ant-oat01-same-token"), + patch( + "agent.anthropic_adapter.resolve_anthropic_token", + return_value="sk-ant-oat01-same-token", + ), patch("agent.anthropic_adapter.build_anthropic_client") as rebuild, ): assert agent._try_refresh_anthropic_client_credentials() is False @@ -3359,9 +3699,15 @@ def test_try_refresh_anthropic_client_credentials_returns_false_when_token_uncha def test_anthropic_messages_create_preflights_refresh(self): with ( - patch("run_agent.get_tool_definitions", return_value=_make_tool_defs("web_search")), + patch( + "run_agent.get_tool_definitions", + return_value=_make_tool_defs("web_search"), + ), patch("run_agent.check_toolset_requirements", return_value={}), - patch("agent.anthropic_adapter.build_anthropic_client", return_value=MagicMock()), + patch( + "agent.anthropic_adapter.build_anthropic_client", + return_value=MagicMock(), + ), ): agent = AIAgent( api_key="sk-ant-oat01-current-token", @@ -3375,11 +3721,17 @@ def test_anthropic_messages_create_preflights_refresh(self): agent._anthropic_client = MagicMock() agent._anthropic_client.messages.create.return_value = response - with patch.object(agent, "_try_refresh_anthropic_client_credentials", return_value=True) as refresh: - result = agent._anthropic_messages_create({"model": "claude-sonnet-4-20250514"}) + with patch.object( + agent, "_try_refresh_anthropic_client_credentials", return_value=True + ) as refresh: + result = agent._anthropic_messages_create( + {"model": "claude-sonnet-4-20250514"} + ) refresh.assert_called_once_with() - agent._anthropic_client.messages.create.assert_called_once_with(model="claude-sonnet-4-20250514") + agent._anthropic_client.messages.create.assert_called_once_with( + model="claude-sonnet-4-20250514" + ) assert result is response @@ -3387,6 +3739,7 @@ def test_anthropic_messages_create_preflights_refresh(self): # _streaming_api_call tests # =================================================================== + def _make_chunk(content=None, tool_calls=None, finish_reason=None, model="test/model"): """Build a SimpleNamespace mimicking an OpenAI streaming chunk.""" delta = SimpleNamespace(content=content, tool_calls=tool_calls) @@ -3441,8 +3794,8 @@ def test_tool_call_accumulation(self, agent): def test_multiple_tool_calls(self, agent): chunks = [ - _make_chunk(tool_calls=[_make_tc_delta(0, "call_a", "search", '{}')]), - _make_chunk(tool_calls=[_make_tc_delta(1, "call_b", "read", '{}')]), + _make_chunk(tool_calls=[_make_tc_delta(0, "call_a", "search", "{}")]), + _make_chunk(tool_calls=[_make_tc_delta(1, "call_b", "read", "{}")]), _make_chunk(finish_reason="tool_calls"), ] agent.client.chat.completions.create.return_value = iter(chunks) @@ -3456,7 +3809,13 @@ def test_multiple_tool_calls(self, agent): def test_truncated_tool_call_args_upgrade_finish_reason_to_length(self, agent): chunks = [ - _make_chunk(tool_calls=[_make_tc_delta(0, "call_1", "write_file", '{"path":"x.txt","content":"hel')]), + _make_chunk( + tool_calls=[ + _make_tc_delta( + 0, "call_1", "write_file", '{"path":"x.txt","content":"hel' + ) + ] + ), ] agent.client.chat.completions.create.return_value = iter(chunks) @@ -3474,9 +3833,13 @@ def test_ollama_reused_index_separate_tool_calls(self, agent): Without the fix, names and arguments get concatenated into one slot. """ chunks = [ - _make_chunk(tool_calls=[_make_tc_delta(0, "call_a", "search", '{"q":"hello"}')]), + _make_chunk( + tool_calls=[_make_tc_delta(0, "call_a", "search", '{"q":"hello"}')] + ), # Second tool call at the SAME index 0, but different id - _make_chunk(tool_calls=[_make_tc_delta(0, "call_b", "read_file", '{"path":"x.py"}')]), + _make_chunk( + tool_calls=[_make_tc_delta(0, "call_b", "read_file", '{"path":"x.py"}')] + ), _make_chunk(finish_reason="tool_calls"), ] agent.client.chat.completions.create.return_value = iter(chunks) @@ -3484,7 +3847,9 @@ def test_ollama_reused_index_separate_tool_calls(self, agent): resp = agent._interruptible_streaming_api_call({"messages": []}) tc = resp.choices[0].message.tool_calls - assert len(tc) == 2, f"Expected 2 tool calls, got {len(tc)}: {[t.function.name for t in tc]}" + assert len(tc) == 2, ( + f"Expected 2 tool calls, got {len(tc)}: {[t.function.name for t in tc]}" + ) assert tc[0].function.name == "search" assert tc[0].function.arguments == '{"q":"hello"}' assert tc[0].id == "call_a" @@ -3498,7 +3863,7 @@ def test_ollama_reused_index_streamed_args(self, agent): _make_chunk(tool_calls=[_make_tc_delta(0, "call_a", "search", '{"q":')]), _make_chunk(tool_calls=[_make_tc_delta(0, None, None, '"hello"}')]), # New tool call, same index 0 - _make_chunk(tool_calls=[_make_tc_delta(0, "call_b", "read", '{}')]), + _make_chunk(tool_calls=[_make_tc_delta(0, "call_b", "read", "{}")]), _make_chunk(finish_reason="tool_calls"), ] agent.client.chat.completions.create.return_value = iter(chunks) @@ -3510,12 +3875,12 @@ def test_ollama_reused_index_streamed_args(self, agent): assert tc[0].function.name == "search" assert tc[0].function.arguments == '{"q":"hello"}' assert tc[1].function.name == "read" - assert tc[1].function.arguments == '{}' + assert tc[1].function.arguments == "{}" def test_content_and_tool_calls_together(self, agent): chunks = [ _make_chunk(content="I'll search"), - _make_chunk(tool_calls=[_make_tc_delta(0, "call_1", "search", '{}')]), + _make_chunk(tool_calls=[_make_tc_delta(0, "call_1", "search", "{}")]), _make_chunk(finish_reason="tool_calls"), ] agent.client.chat.completions.create.return_value = iter(chunks) @@ -3565,7 +3930,10 @@ def test_stream_kwarg_injected(self, agent): agent._interruptible_streaming_api_call({"messages": [], "model": "test"}) call_kwargs = agent.client.chat.completions.create.call_args - assert call_kwargs[1].get("stream") is True or call_kwargs.kwargs.get("stream") is True + assert ( + call_kwargs[1].get("stream") is True + or call_kwargs.kwargs.get("stream") is True + ) def test_api_exception_propagates_no_non_streaming_fallback(self, agent): """When streaming fails before any deltas, error propagates to the main retry loop.""" @@ -3611,6 +3979,7 @@ class TestInterruptVprintForceTrue: def test_all_interrupt_vprint_have_force_true(self): """Scan source for _vprint calls containing 'Interrupt' — each must have force=True.""" import inspect + source = inspect.getsource(AIAgent) lines = source.split("\n") violations = [] @@ -3620,8 +3989,7 @@ def test_all_interrupt_vprint_have_force_true(self): if "force=True" not in stripped: violations.append(f"line {i}: {stripped}") assert not violations, ( - f"Interrupt _vprint calls missing force=True:\n" - + "\n".join(violations) + f"Interrupt _vprint calls missing force=True:\n" + "\n".join(violations) ) @@ -3636,23 +4004,29 @@ class TestAnthropicInterruptHandler: def test_interruptible_has_anthropic_branch(self): """The interrupt handler must check api_mode == 'anthropic_messages'.""" import inspect + source = inspect.getsource(AIAgent._interruptible_api_call) - assert "anthropic_messages" in source, \ + assert "anthropic_messages" in source, ( "_interruptible_api_call must handle Anthropic interrupt (api_mode check)" + ) def test_interruptible_rebuilds_anthropic_client(self): """After interrupting, the Anthropic client should be rebuilt.""" import inspect + source = inspect.getsource(AIAgent._interruptible_api_call) - assert "build_anthropic_client" in source, \ + assert "build_anthropic_client" in source, ( "_interruptible_api_call must rebuild Anthropic client after interrupt" + ) def test_streaming_has_anthropic_branch(self): """_streaming_api_call must also handle Anthropic interrupt.""" import inspect + source = inspect.getsource(AIAgent._interruptible_streaming_api_call) - assert "anthropic_messages" in source, \ + assert "anthropic_messages" in source, ( "_streaming_api_call must handle Anthropic interrupt" + ) # --------------------------------------------------------------------------- @@ -3668,11 +4042,18 @@ def test_callback_receives_chat_completions_response(self, agent): """For chat_completions-shaped responses, callback gets content.""" agent.api_mode = "anthropic_messages" mock_response = SimpleNamespace( - choices=[SimpleNamespace( - message=SimpleNamespace(content="Hello", tool_calls=None, reasoning_content=None), - finish_reason="stop", index=0, - )], - usage=None, model="test", id="test-id", + choices=[ + SimpleNamespace( + message=SimpleNamespace( + content="Hello", tool_calls=None, reasoning_content=None + ), + finish_reason="stop", + index=0, + ) + ], + usage=None, + model="test", + id="test-id", ) agent._interruptible_api_call = MagicMock(return_value=mock_response) @@ -3686,8 +4067,10 @@ def test_callback_receives_chat_completions_response(self, agent): try: if agent.api_mode == "anthropic_messages": text_parts = [ - block.text for block in getattr(response, "content", []) - if getattr(block, "type", None) == "text" and getattr(block, "text", None) + block.text + for block in getattr(response, "content", []) + if getattr(block, "type", None) == "text" + and getattr(block, "text", None) ] content = " ".join(text_parts) if text_parts else None else: @@ -3729,8 +4112,10 @@ def test_callback_receives_anthropic_content(self, agent): try: if agent.api_mode == "anthropic_messages": text_parts = [ - block.text for block in getattr(mock_response, "content", []) - if getattr(block, "type", None) == "text" and getattr(block, "text", None) + block.text + for block in getattr(mock_response, "content", []) + if getattr(block, "type", None) == "text" + and getattr(block, "text", None) ] content = " ".join(text_parts) if text_parts else None else: @@ -3879,10 +4264,14 @@ def test_oauth_flag_updates_api_key_to_oauth(self, agent): agent._is_anthropic_oauth = False with ( - patch("agent.anthropic_adapter.resolve_anthropic_token", - return_value="sk-ant-setup-oauth-token"), - patch("agent.anthropic_adapter.build_anthropic_client", - return_value=MagicMock()), + patch( + "agent.anthropic_adapter.resolve_anthropic_token", + return_value="sk-ant-setup-oauth-token", + ), + patch( + "agent.anthropic_adapter.build_anthropic_client", + return_value=MagicMock(), + ), ): result = agent._try_refresh_anthropic_client_credentials() @@ -3898,10 +4287,14 @@ def test_oauth_flag_updates_oauth_to_api_key(self, agent): agent._is_anthropic_oauth = True with ( - patch("agent.anthropic_adapter.resolve_anthropic_token", - return_value="sk-ant-api03-new-key"), - patch("agent.anthropic_adapter.build_anthropic_client", - return_value=MagicMock()), + patch( + "agent.anthropic_adapter.resolve_anthropic_token", + return_value="sk-ant-api03-new-key", + ), + patch( + "agent.anthropic_adapter.build_anthropic_client", + return_value=MagicMock(), + ), ): result = agent._try_refresh_anthropic_client_credentials() @@ -3923,12 +4316,15 @@ def test_fallback_to_anthropic_oauth_sets_flag(self, agent): mock_client.api_key = "sk-ant-setup-oauth-token" with ( - patch("agent.auxiliary_client.resolve_provider_client", - return_value=(mock_client, None)), - patch("agent.anthropic_adapter.build_anthropic_client", - return_value=MagicMock()), - patch("agent.anthropic_adapter.resolve_anthropic_token", - return_value=None), + patch( + "agent.auxiliary_client.resolve_provider_client", + return_value=(mock_client, None), + ), + patch( + "agent.anthropic_adapter.build_anthropic_client", + return_value=MagicMock(), + ), + patch("agent.anthropic_adapter.resolve_anthropic_token", return_value=None), ): result = agent._try_activate_fallback() @@ -3946,12 +4342,15 @@ def test_fallback_to_anthropic_api_key_clears_flag(self, agent): mock_client.api_key = "sk-ant-api03-regular-key" with ( - patch("agent.auxiliary_client.resolve_provider_client", - return_value=(mock_client, None)), - patch("agent.anthropic_adapter.build_anthropic_client", - return_value=MagicMock()), - patch("agent.anthropic_adapter.resolve_anthropic_token", - return_value=None), + patch( + "agent.auxiliary_client.resolve_provider_client", + return_value=(mock_client, None), + ), + patch( + "agent.anthropic_adapter.build_anthropic_client", + return_value=MagicMock(), + ), + patch("agent.anthropic_adapter.resolve_anthropic_token", return_value=None), ): result = agent._try_activate_fallback() @@ -3966,8 +4365,11 @@ def test_counters_initialized_in_init(self): """Counters must exist on the agent after __init__.""" with patch("run_agent.get_tool_definitions", return_value=[]): a = AIAgent( - model="test", api_key="test-key", provider="openrouter", - skip_context_files=True, skip_memory=True, + model="test", + api_key="test-key", + provider="openrouter", + skip_context_files=True, + skip_memory=True, ) assert hasattr(a, "_turns_since_memory") assert hasattr(a, "_iters_since_skill") @@ -3977,6 +4379,7 @@ def test_counters_initialized_in_init(self): def test_counters_not_reset_in_preamble(self): """The run_conversation preamble must not zero the nudge counters.""" import inspect + src = inspect.getsource(AIAgent.run_conversation) # The preamble resets many fields (retry counts, budget, etc.) # before the main loop. Find that reset block and verify our @@ -3992,6 +4395,7 @@ class TestDeadRetryCode: def test_no_unreachable_max_retries_after_backoff(self): import inspect + source = inspect.getsource(AIAgent.run_conversation) occurrences = source.count("if retry_count >= max_retries:") assert occurrences == 2, ( From df7d63810735d3133f6e1b126295470502dde764 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 15 Apr 2026 22:31:04 +0800 Subject: [PATCH 2/3] perf(gateway): reduce telegram response latency while preserving delivery stability --- gateway/config.py | 410 +++-- gateway/display_config.py | 42 +- gateway/platforms/telegram.py | 938 ++++++++--- gateway/run.py | 2341 ++++++++++++++++++-------- gateway/stream_consumer.py | 94 +- tests/gateway/test_display_config.py | 32 +- tests/gateway/test_proxy_mode.py | 49 +- 7 files changed, 2732 insertions(+), 1174 deletions(-) diff --git a/gateway/config.py b/gateway/config.py index 7ce105f331b7c..0f2bc03bea486 100644 --- a/gateway/config.py +++ b/gateway/config.py @@ -47,6 +47,7 @@ def _normalize_unauthorized_dm_behavior(value: Any, default: str = "pair") -> st class Platform(Enum): """Supported messaging platforms.""" + LOCAL = "local" TELEGRAM = "telegram" DISCORD = "discord" @@ -73,21 +74,22 @@ class Platform(Enum): class HomeChannel: """ Default destination for a platform. - + When a cron job specifies deliver="telegram" without a specific chat ID, messages are sent to this home channel. """ + platform: Platform chat_id: str name: str # Human-readable name for display - + def to_dict(self) -> Dict[str, Any]: return { "platform": self.platform.value, "chat_id": self.chat_id, "name": self.name, } - + @classmethod def from_dict(cls, data: Dict[str, Any]) -> "HomeChannel": return cls( @@ -101,19 +103,23 @@ def from_dict(cls, data: Dict[str, Any]) -> "HomeChannel": class SessionResetPolicy: """ Controls when sessions reset (lose context). - + Modes: - "daily": Reset at a specific hour each day - "idle": Reset after N minutes of inactivity - "both": Whichever triggers first (daily boundary OR idle timeout) - "none": Never auto-reset (context managed only by compression) """ + mode: str = "both" # "daily", "idle", "both", or "none" at_hour: int = 4 # Hour for daily reset (0-23, local time) idle_minutes: int = 1440 # Minutes of inactivity before reset (24 hours) notify: bool = True # Send a notification to the user when auto-reset occurs - notify_exclude_platforms: tuple = ("api_server", "webhook") # Platforms that don't get reset notifications - + notify_exclude_platforms: tuple = ( + "api_server", + "webhook", + ) # Platforms that don't get reset notifications + def to_dict(self) -> Dict[str, Any]: return { "mode": self.mode, @@ -122,7 +128,7 @@ def to_dict(self) -> Dict[str, Any]: "notify": self.notify, "notify_exclude_platforms": list(self.notify_exclude_platforms), } - + @classmethod def from_dict(cls, data: Dict[str, Any]) -> "SessionResetPolicy": # Handle both missing keys and explicit null values (YAML null → None) @@ -136,27 +142,30 @@ def from_dict(cls, data: Dict[str, Any]) -> "SessionResetPolicy": at_hour=at_hour if at_hour is not None else 4, idle_minutes=idle_minutes if idle_minutes is not None else 1440, notify=notify if notify is not None else True, - notify_exclude_platforms=tuple(exclude) if exclude is not None else ("api_server", "webhook"), + notify_exclude_platforms=tuple(exclude) + if exclude is not None + else ("api_server", "webhook"), ) @dataclass class PlatformConfig: """Configuration for a single messaging platform.""" + enabled: bool = False token: Optional[str] = None # Bot token (Telegram, Discord) api_key: Optional[str] = None # API key if different from token home_channel: Optional[HomeChannel] = None - + # Reply threading mode (Telegram/Slack) # - "off": Never thread replies to original message # - "first": Only first chunk threads to user's message (default) # - "all": All chunks in multi-part replies thread to user's message reply_to_mode: str = "first" - + # Platform-specific settings extra: Dict[str, Any] = field(default_factory=dict) - + def to_dict(self) -> Dict[str, Any]: result = { "enabled": self.enabled, @@ -170,13 +179,13 @@ def to_dict(self) -> Dict[str, Any]: if self.home_channel: result["home_channel"] = self.home_channel.to_dict() return result - + @classmethod def from_dict(cls, data: Dict[str, Any]) -> "PlatformConfig": home_channel = None if "home_channel" in data: home_channel = HomeChannel.from_dict(data["home_channel"]) - + return cls( enabled=data.get("enabled", False), token=data.get("token"), @@ -190,11 +199,16 @@ def from_dict(cls, data: Dict[str, Any]) -> "PlatformConfig": @dataclass class StreamingConfig: """Configuration for real-time token streaming to messaging platforms.""" + enabled: bool = False - transport: str = "edit" # "edit" (progressive editMessageText) or "off" - edit_interval: float = 1.0 # Seconds between message edits (Telegram rate-limits at ~1/s) - buffer_threshold: int = 40 # Chars before forcing an edit - cursor: str = " ▉" # Cursor shown during streaming + transport: str = "edit" # "edit" (progressive editMessageText) or "off" + edit_interval: float = ( + 0.8 # Seconds between message edits (balanced for lower latency + flood safety) + ) + buffer_threshold: int = ( + 24 # Chars before forcing an edit (faster first visible token) + ) + cursor: str = " ▉" # Cursor shown during streaming def to_dict(self) -> Dict[str, Any]: return { @@ -212,8 +226,8 @@ def from_dict(cls, data: Dict[str, Any]) -> "StreamingConfig": return cls( enabled=data.get("enabled", False), transport=data.get("transport", "edit"), - edit_interval=float(data.get("edit_interval", 1.0)), - buffer_threshold=int(data.get("buffer_threshold", 40)), + edit_interval=float(data.get("edit_interval", 0.8)), + buffer_threshold=int(data.get("buffer_threshold", 24)), cursor=data.get("cursor", " ▉"), ) @@ -222,26 +236,27 @@ def from_dict(cls, data: Dict[str, Any]) -> "StreamingConfig": class GatewayConfig: """ Main gateway configuration. - + Manages all platform connections, session policies, and delivery settings. """ + # Platform configurations platforms: Dict[Platform, PlatformConfig] = field(default_factory=dict) - + # Session reset policies by type default_reset_policy: SessionResetPolicy = field(default_factory=SessionResetPolicy) reset_by_type: Dict[str, SessionResetPolicy] = field(default_factory=dict) reset_by_platform: Dict[Platform, SessionResetPolicy] = field(default_factory=dict) - + # Reset trigger commands reset_triggers: List[str] = field(default_factory=lambda: ["/new", "/reset"]) # User-defined quick commands (slash commands that bypass the agent loop) quick_commands: Dict[str, Any] = field(default_factory=dict) - + # Storage paths sessions_dir: Path = field(default_factory=lambda: get_hermes_home() / "sessions") - + # Delivery settings always_log_local: bool = True # Always save cron outputs to local files @@ -250,7 +265,9 @@ class GatewayConfig: # Session isolation in shared chats group_sessions_per_user: bool = True # Isolate group/channel sessions per participant when user IDs are available - thread_sessions_per_user: bool = False # When False (default), threads are shared across all participants + thread_sessions_per_user: bool = ( + False # When False (default), threads are shared across all participants + ) # Unauthorized DM policy unauthorized_dm_behavior: str = "pair" # "pair" or "ignore" @@ -266,7 +283,9 @@ def get_connected_platforms(self) -> List[Platform]: continue # Weixin requires both a token and an account_id if platform == Platform.WEIXIN: - if config.extra.get("account_id") and (config.token or config.extra.get("token")): + if config.extra.get("account_id") and ( + config.token or config.extra.get("token") + ): connected.append(platform) continue # Platforms that use token/api_key auth @@ -302,49 +321,51 @@ def get_connected_platforms(self) -> List[Platform]: ): connected.append(platform) # BlueBubbles uses extra dict for local server config - elif platform == Platform.BLUEBUBBLES and config.extra.get("server_url") and config.extra.get("password"): + elif ( + platform == Platform.BLUEBUBBLES + and config.extra.get("server_url") + and config.extra.get("password") + ): connected.append(platform) # QQBot uses extra dict for app credentials - elif platform == Platform.QQBOT and config.extra.get("app_id") and config.extra.get("client_secret"): + elif ( + platform == Platform.QQBOT + and config.extra.get("app_id") + and config.extra.get("client_secret") + ): connected.append(platform) return connected - + def get_home_channel(self, platform: Platform) -> Optional[HomeChannel]: """Get the home channel for a platform.""" config = self.platforms.get(platform) if config: return config.home_channel return None - + def get_reset_policy( - self, - platform: Optional[Platform] = None, - session_type: Optional[str] = None + self, platform: Optional[Platform] = None, session_type: Optional[str] = None ) -> SessionResetPolicy: """ Get the appropriate reset policy for a session. - + Priority: platform override > type override > default """ # Platform-specific override takes precedence if platform and platform in self.reset_by_platform: return self.reset_by_platform[platform] - + # Type-specific override (dm, group, thread) if session_type and session_type in self.reset_by_type: return self.reset_by_type[session_type] - + return self.default_reset_policy - + def to_dict(self) -> Dict[str, Any]: return { - "platforms": { - p.value: c.to_dict() for p, c in self.platforms.items() - }, + "platforms": {p.value: c.to_dict() for p, c in self.platforms.items()}, "default_reset_policy": self.default_reset_policy.to_dict(), - "reset_by_type": { - k: v.to_dict() for k, v in self.reset_by_type.items() - }, + "reset_by_type": {k: v.to_dict() for k, v in self.reset_by_type.items()}, "reset_by_platform": { p.value: v.to_dict() for p, v in self.reset_by_platform.items() }, @@ -358,7 +379,7 @@ def to_dict(self) -> Dict[str, Any]: "unauthorized_dm_behavior": self.unauthorized_dm_behavior, "streaming": self.streaming.to_dict(), } - + @classmethod def from_dict(cls, data: Dict[str, Any]) -> "GatewayConfig": platforms = {} @@ -368,11 +389,11 @@ def from_dict(cls, data: Dict[str, Any]) -> "GatewayConfig": platforms[platform] = PlatformConfig.from_dict(platform_data) except ValueError: pass # Skip unknown platforms - + reset_by_type = {} for type_name, policy_data in data.get("reset_by_type", {}).items(): reset_by_type[type_name] = SessionResetPolicy.from_dict(policy_data) - + reset_by_platform = {} for platform_name, policy_data in data.get("reset_by_platform", {}).items(): try: @@ -380,22 +401,26 @@ def from_dict(cls, data: Dict[str, Any]) -> "GatewayConfig": reset_by_platform[platform] = SessionResetPolicy.from_dict(policy_data) except ValueError: pass - + default_policy = SessionResetPolicy() if "default_reset_policy" in data: default_policy = SessionResetPolicy.from_dict(data["default_reset_policy"]) - + sessions_dir = get_hermes_home() / "sessions" if "sessions_dir" in data: sessions_dir = Path(data["sessions_dir"]) - + quick_commands = data.get("quick_commands", {}) if not isinstance(quick_commands, dict): quick_commands = {} stt_enabled = data.get("stt_enabled") if stt_enabled is None: - stt_enabled = data.get("stt", {}).get("enabled") if isinstance(data.get("stt"), dict) else None + stt_enabled = ( + data.get("stt", {}).get("enabled") + if isinstance(data.get("stt"), dict) + else None + ) group_sessions_per_user = data.get("group_sessions_per_user") thread_sessions_per_user = data.get("thread_sessions_per_user") @@ -462,6 +487,7 @@ def load_gateway_config() -> GatewayConfig: # Primary source: config.yaml try: import yaml + config_yaml_path = _home / "config.yaml" if config_yaml_path.exists(): with open(config_yaml_path, encoding="utf-8") as f: @@ -492,7 +518,9 @@ def load_gateway_config() -> GatewayConfig: gw_data["group_sessions_per_user"] = yaml_cfg["group_sessions_per_user"] if "thread_sessions_per_user" in yaml_cfg: - gw_data["thread_sessions_per_user"] = yaml_cfg["thread_sessions_per_user"] + gw_data["thread_sessions_per_user"] = yaml_cfg[ + "thread_sessions_per_user" + ] streaming_cfg = yaml_cfg.get("streaming") if isinstance(streaming_cfg, dict): @@ -505,9 +533,11 @@ def load_gateway_config() -> GatewayConfig: gw_data["always_log_local"] = yaml_cfg["always_log_local"] if "unauthorized_dm_behavior" in yaml_cfg: - gw_data["unauthorized_dm_behavior"] = _normalize_unauthorized_dm_behavior( - yaml_cfg.get("unauthorized_dm_behavior"), - "pair", + gw_data["unauthorized_dm_behavior"] = ( + _normalize_unauthorized_dm_behavior( + yaml_cfg.get("unauthorized_dm_behavior"), + "pair", + ) ) # Merge platforms section from config.yaml into gw_data so that @@ -525,7 +555,10 @@ def load_gateway_config() -> GatewayConfig: if not isinstance(existing, dict): existing = {} # Deep-merge extra dicts so gateway.json defaults survive - merged_extra = {**existing.get("extra", {}), **plat_block.get("extra", {})} + merged_extra = { + **existing.get("extra", {}), + **plat_block.get("extra", {}), + } merged = {**existing, **plat_block} if merged_extra: merged["extra"] = merged_extra @@ -540,20 +573,29 @@ def load_gateway_config() -> GatewayConfig: # Collect bridgeable keys from this platform section bridged = {} if "unauthorized_dm_behavior" in platform_cfg: - bridged["unauthorized_dm_behavior"] = _normalize_unauthorized_dm_behavior( - platform_cfg.get("unauthorized_dm_behavior"), - gw_data.get("unauthorized_dm_behavior", "pair"), + bridged["unauthorized_dm_behavior"] = ( + _normalize_unauthorized_dm_behavior( + platform_cfg.get("unauthorized_dm_behavior"), + gw_data.get("unauthorized_dm_behavior", "pair"), + ) ) if "reply_prefix" in platform_cfg: bridged["reply_prefix"] = platform_cfg["reply_prefix"] if "require_mention" in platform_cfg: bridged["require_mention"] = platform_cfg["require_mention"] if "free_response_channels" in platform_cfg: - bridged["free_response_channels"] = platform_cfg["free_response_channels"] + bridged["free_response_channels"] = platform_cfg[ + "free_response_channels" + ] if "mention_patterns" in platform_cfg: bridged["mention_patterns"] = platform_cfg["mention_patterns"] - if plat == Platform.DISCORD and "channel_skill_bindings" in platform_cfg: - bridged["channel_skill_bindings"] = platform_cfg["channel_skill_bindings"] + if ( + plat == Platform.DISCORD + and "channel_skill_bindings" in platform_cfg + ): + bridged["channel_skill_bindings"] = platform_cfg[ + "channel_skill_bindings" + ] if not bridged: continue plat_data = platforms_data.setdefault(plat.value, {}) @@ -569,10 +611,16 @@ def load_gateway_config() -> GatewayConfig: # Slack settings → env vars (env vars take precedence) slack_cfg = yaml_cfg.get("slack", {}) if isinstance(slack_cfg, dict): - if "require_mention" in slack_cfg and not os.getenv("SLACK_REQUIRE_MENTION"): - os.environ["SLACK_REQUIRE_MENTION"] = str(slack_cfg["require_mention"]).lower() + if "require_mention" in slack_cfg and not os.getenv( + "SLACK_REQUIRE_MENTION" + ): + os.environ["SLACK_REQUIRE_MENTION"] = str( + slack_cfg["require_mention"] + ).lower() if "allow_bots" in slack_cfg and not os.getenv("SLACK_ALLOW_BOTS"): - os.environ["SLACK_ALLOW_BOTS"] = str(slack_cfg["allow_bots"]).lower() + os.environ["SLACK_ALLOW_BOTS"] = str( + slack_cfg["allow_bots"] + ).lower() frc = slack_cfg.get("free_response_channels") if frc is not None and not os.getenv("SLACK_FREE_RESPONSE_CHANNELS"): if isinstance(frc, list): @@ -582,17 +630,27 @@ def load_gateway_config() -> GatewayConfig: # Discord settings → env vars (env vars take precedence) discord_cfg = yaml_cfg.get("discord", {}) if isinstance(discord_cfg, dict): - if "require_mention" in discord_cfg and not os.getenv("DISCORD_REQUIRE_MENTION"): - os.environ["DISCORD_REQUIRE_MENTION"] = str(discord_cfg["require_mention"]).lower() + if "require_mention" in discord_cfg and not os.getenv( + "DISCORD_REQUIRE_MENTION" + ): + os.environ["DISCORD_REQUIRE_MENTION"] = str( + discord_cfg["require_mention"] + ).lower() frc = discord_cfg.get("free_response_channels") if frc is not None and not os.getenv("DISCORD_FREE_RESPONSE_CHANNELS"): if isinstance(frc, list): frc = ",".join(str(v) for v in frc) os.environ["DISCORD_FREE_RESPONSE_CHANNELS"] = str(frc) - if "auto_thread" in discord_cfg and not os.getenv("DISCORD_AUTO_THREAD"): - os.environ["DISCORD_AUTO_THREAD"] = str(discord_cfg["auto_thread"]).lower() + if "auto_thread" in discord_cfg and not os.getenv( + "DISCORD_AUTO_THREAD" + ): + os.environ["DISCORD_AUTO_THREAD"] = str( + discord_cfg["auto_thread"] + ).lower() if "reactions" in discord_cfg and not os.getenv("DISCORD_REACTIONS"): - os.environ["DISCORD_REACTIONS"] = str(discord_cfg["reactions"]).lower() + os.environ["DISCORD_REACTIONS"] = str( + discord_cfg["reactions"] + ).lower() # ignored_channels: channels where bot never responds (even when mentioned) ic = discord_cfg.get("ignored_channels") if ic is not None and not os.getenv("DISCORD_IGNORED_CHANNELS"): @@ -615,30 +673,51 @@ def load_gateway_config() -> GatewayConfig: # Telegram settings → env vars (env vars take precedence) telegram_cfg = yaml_cfg.get("telegram", {}) if isinstance(telegram_cfg, dict): - if "require_mention" in telegram_cfg and not os.getenv("TELEGRAM_REQUIRE_MENTION"): - os.environ["TELEGRAM_REQUIRE_MENTION"] = str(telegram_cfg["require_mention"]).lower() - if "mention_patterns" in telegram_cfg and not os.getenv("TELEGRAM_MENTION_PATTERNS"): + if "require_mention" in telegram_cfg and not os.getenv( + "TELEGRAM_REQUIRE_MENTION" + ): + os.environ["TELEGRAM_REQUIRE_MENTION"] = str( + telegram_cfg["require_mention"] + ).lower() + if "mention_patterns" in telegram_cfg and not os.getenv( + "TELEGRAM_MENTION_PATTERNS" + ): import json as _json - os.environ["TELEGRAM_MENTION_PATTERNS"] = _json.dumps(telegram_cfg["mention_patterns"]) + + os.environ["TELEGRAM_MENTION_PATTERNS"] = _json.dumps( + telegram_cfg["mention_patterns"] + ) frc = telegram_cfg.get("free_response_chats") if frc is not None and not os.getenv("TELEGRAM_FREE_RESPONSE_CHATS"): if isinstance(frc, list): frc = ",".join(str(v) for v in frc) os.environ["TELEGRAM_FREE_RESPONSE_CHATS"] = str(frc) ignored_threads = telegram_cfg.get("ignored_threads") - if ignored_threads is not None and not os.getenv("TELEGRAM_IGNORED_THREADS"): + if ignored_threads is not None and not os.getenv( + "TELEGRAM_IGNORED_THREADS" + ): if isinstance(ignored_threads, list): ignored_threads = ",".join(str(v) for v in ignored_threads) os.environ["TELEGRAM_IGNORED_THREADS"] = str(ignored_threads) if "reactions" in telegram_cfg and not os.getenv("TELEGRAM_REACTIONS"): - os.environ["TELEGRAM_REACTIONS"] = str(telegram_cfg["reactions"]).lower() + os.environ["TELEGRAM_REACTIONS"] = str( + telegram_cfg["reactions"] + ).lower() whatsapp_cfg = yaml_cfg.get("whatsapp", {}) if isinstance(whatsapp_cfg, dict): - if "require_mention" in whatsapp_cfg and not os.getenv("WHATSAPP_REQUIRE_MENTION"): - os.environ["WHATSAPP_REQUIRE_MENTION"] = str(whatsapp_cfg["require_mention"]).lower() - if "mention_patterns" in whatsapp_cfg and not os.getenv("WHATSAPP_MENTION_PATTERNS"): - os.environ["WHATSAPP_MENTION_PATTERNS"] = json.dumps(whatsapp_cfg["mention_patterns"]) + if "require_mention" in whatsapp_cfg and not os.getenv( + "WHATSAPP_REQUIRE_MENTION" + ): + os.environ["WHATSAPP_REQUIRE_MENTION"] = str( + whatsapp_cfg["require_mention"] + ).lower() + if "mention_patterns" in whatsapp_cfg and not os.getenv( + "WHATSAPP_MENTION_PATTERNS" + ): + os.environ["WHATSAPP_MENTION_PATTERNS"] = json.dumps( + whatsapp_cfg["mention_patterns"] + ) frc = whatsapp_cfg.get("free_response_chats") if frc is not None and not os.getenv("WHATSAPP_FREE_RESPONSE_CHATS"): if isinstance(frc, list): @@ -648,17 +727,27 @@ def load_gateway_config() -> GatewayConfig: # Matrix settings → env vars (env vars take precedence) matrix_cfg = yaml_cfg.get("matrix", {}) if isinstance(matrix_cfg, dict): - if "require_mention" in matrix_cfg and not os.getenv("MATRIX_REQUIRE_MENTION"): - os.environ["MATRIX_REQUIRE_MENTION"] = str(matrix_cfg["require_mention"]).lower() + if "require_mention" in matrix_cfg and not os.getenv( + "MATRIX_REQUIRE_MENTION" + ): + os.environ["MATRIX_REQUIRE_MENTION"] = str( + matrix_cfg["require_mention"] + ).lower() frc = matrix_cfg.get("free_response_rooms") if frc is not None and not os.getenv("MATRIX_FREE_RESPONSE_ROOMS"): if isinstance(frc, list): frc = ",".join(str(v) for v in frc) os.environ["MATRIX_FREE_RESPONSE_ROOMS"] = str(frc) if "auto_thread" in matrix_cfg and not os.getenv("MATRIX_AUTO_THREAD"): - os.environ["MATRIX_AUTO_THREAD"] = str(matrix_cfg["auto_thread"]).lower() - if "dm_mention_threads" in matrix_cfg and not os.getenv("MATRIX_DM_MENTION_THREADS"): - os.environ["MATRIX_DM_MENTION_THREADS"] = str(matrix_cfg["dm_mention_threads"]).lower() + os.environ["MATRIX_AUTO_THREAD"] = str( + matrix_cfg["auto_thread"] + ).lower() + if "dm_mention_threads" in matrix_cfg and not os.getenv( + "MATRIX_DM_MENTION_THREADS" + ): + os.environ["MATRIX_DM_MENTION_THREADS"] = str( + matrix_cfg["dm_mention_threads"] + ).lower() except Exception as e: logger.warning( @@ -672,7 +761,7 @@ def load_gateway_config() -> GatewayConfig: # Override with environment variables _apply_env_overrides(config) - + # --- Validate loaded values --- _validate_gateway_config(config) @@ -718,7 +807,8 @@ def _validate_gateway_config(config: "GatewayConfig") -> None: logger.warning( "%s is enabled but %s is empty. " "The adapter will likely fail to connect.", - platform.value, env_name, + platform.value, + env_name, ) # Reject known-weak placeholder tokens. @@ -743,14 +833,16 @@ def _validate_gateway_config(config: "GatewayConfig") -> None: "%s is enabled but %s is set to a placeholder value ('%s'). " "Set a real bot token before starting the gateway. " "The adapter will NOT be started.", - platform.value, env_name, token.strip()[:6] + "...", + platform.value, + env_name, + token.strip()[:6] + "...", ) pconfig.enabled = False def _apply_env_overrides(config: GatewayConfig) -> None: """Apply environment variable overrides to config.""" - + # Telegram telegram_token = os.getenv("TELEGRAM_BOT_TOKEN") if telegram_token: @@ -758,14 +850,14 @@ def _apply_env_overrides(config: GatewayConfig) -> None: config.platforms[Platform.TELEGRAM] = PlatformConfig() config.platforms[Platform.TELEGRAM].enabled = True config.platforms[Platform.TELEGRAM].token = telegram_token - + # Reply threading mode for Telegram (off/first/all) telegram_reply_mode = os.getenv("TELEGRAM_REPLY_TO_MODE", "").lower() if telegram_reply_mode in ("off", "first", "all"): if Platform.TELEGRAM not in config.platforms: config.platforms[Platform.TELEGRAM] = PlatformConfig() config.platforms[Platform.TELEGRAM].reply_to_mode = telegram_reply_mode - + telegram_fallback_ips = os.getenv("TELEGRAM_FALLBACK_IPS", "") if telegram_fallback_ips: if Platform.TELEGRAM not in config.platforms: @@ -781,7 +873,7 @@ def _apply_env_overrides(config: GatewayConfig) -> None: chat_id=telegram_home, name=os.getenv("TELEGRAM_HOME_CHANNEL_NAME", "Home"), ) - + # Discord discord_token = os.getenv("DISCORD_BOT_TOKEN") if discord_token: @@ -789,7 +881,7 @@ def _apply_env_overrides(config: GatewayConfig) -> None: config.platforms[Platform.DISCORD] = PlatformConfig() config.platforms[Platform.DISCORD].enabled = True config.platforms[Platform.DISCORD].token = discord_token - + discord_home = os.getenv("DISCORD_HOME_CHANNEL") if discord_home and Platform.DISCORD in config.platforms: config.platforms[Platform.DISCORD].home_channel = HomeChannel( @@ -797,21 +889,21 @@ def _apply_env_overrides(config: GatewayConfig) -> None: chat_id=discord_home, name=os.getenv("DISCORD_HOME_CHANNEL_NAME", "Home"), ) - + # Reply threading mode for Discord (off/first/all) discord_reply_mode = os.getenv("DISCORD_REPLY_TO_MODE", "").lower() if discord_reply_mode in ("off", "first", "all"): if Platform.DISCORD not in config.platforms: config.platforms[Platform.DISCORD] = PlatformConfig() config.platforms[Platform.DISCORD].reply_to_mode = discord_reply_mode - + # WhatsApp (typically uses different auth mechanism) whatsapp_enabled = os.getenv("WHATSAPP_ENABLED", "").lower() in ("true", "1", "yes") if whatsapp_enabled: if Platform.WHATSAPP not in config.platforms: config.platforms[Platform.WHATSAPP] = PlatformConfig() config.platforms[Platform.WHATSAPP].enabled = True - + # Slack slack_token = os.getenv("SLACK_BOT_TOKEN") if slack_token: @@ -826,7 +918,7 @@ def _apply_env_overrides(config: GatewayConfig) -> None: chat_id=slack_home, name=os.getenv("SLACK_HOME_CHANNEL_NAME", ""), ) - + # Signal signal_url = os.getenv("SIGNAL_HTTP_URL") signal_account = os.getenv("SIGNAL_ACCOUNT") @@ -834,11 +926,14 @@ def _apply_env_overrides(config: GatewayConfig) -> None: if Platform.SIGNAL not in config.platforms: config.platforms[Platform.SIGNAL] = PlatformConfig() config.platforms[Platform.SIGNAL].enabled = True - config.platforms[Platform.SIGNAL].extra.update({ - "http_url": signal_url, - "account": signal_account, - "ignore_stories": os.getenv("SIGNAL_IGNORE_STORIES", "true").lower() in ("true", "1", "yes"), - }) + config.platforms[Platform.SIGNAL].extra.update( + { + "http_url": signal_url, + "account": signal_account, + "ignore_stories": os.getenv("SIGNAL_IGNORE_STORIES", "true").lower() + in ("true", "1", "yes"), + } + ) signal_home = os.getenv("SIGNAL_HOME_CHANNEL") if signal_home and Platform.SIGNAL in config.platforms: config.platforms[Platform.SIGNAL].home_channel = HomeChannel( @@ -871,7 +966,9 @@ def _apply_env_overrides(config: GatewayConfig) -> None: matrix_homeserver = os.getenv("MATRIX_HOMESERVER", "") if matrix_token or os.getenv("MATRIX_PASSWORD"): if not matrix_homeserver: - logger.warning("MATRIX_ACCESS_TOKEN/MATRIX_PASSWORD set but MATRIX_HOMESERVER is missing") + logger.warning( + "MATRIX_ACCESS_TOKEN/MATRIX_PASSWORD set but MATRIX_HOMESERVER is missing" + ) if Platform.MATRIX not in config.platforms: config.platforms[Platform.MATRIX] = PlatformConfig() config.platforms[Platform.MATRIX].enabled = True @@ -917,11 +1014,13 @@ def _apply_env_overrides(config: GatewayConfig) -> None: if Platform.EMAIL not in config.platforms: config.platforms[Platform.EMAIL] = PlatformConfig() config.platforms[Platform.EMAIL].enabled = True - config.platforms[Platform.EMAIL].extra.update({ - "address": email_addr, - "imap_host": email_imap, - "smtp_host": email_smtp, - }) + config.platforms[Platform.EMAIL].extra.update( + { + "address": email_addr, + "imap_host": email_imap, + "smtp_host": email_smtp, + } + ) email_home = os.getenv("EMAIL_HOME_ADDRESS") if email_home and Platform.EMAIL in config.platforms: config.platforms[Platform.EMAIL].home_channel = HomeChannel( @@ -946,7 +1045,11 @@ def _apply_env_overrides(config: GatewayConfig) -> None: ) # API Server - api_server_enabled = os.getenv("API_SERVER_ENABLED", "").lower() in ("true", "1", "yes") + api_server_enabled = os.getenv("API_SERVER_ENABLED", "").lower() in ( + "true", + "1", + "yes", + ) api_server_key = os.getenv("API_SERVER_KEY", "") api_server_cors_origins = os.getenv("API_SERVER_CORS_ORIGINS", "") api_server_port = os.getenv("API_SERVER_PORT") @@ -958,19 +1061,27 @@ def _apply_env_overrides(config: GatewayConfig) -> None: if api_server_key: config.platforms[Platform.API_SERVER].extra["key"] = api_server_key if api_server_cors_origins: - origins = [origin.strip() for origin in api_server_cors_origins.split(",") if origin.strip()] + origins = [ + origin.strip() + for origin in api_server_cors_origins.split(",") + if origin.strip() + ] if origins: config.platforms[Platform.API_SERVER].extra["cors_origins"] = origins if api_server_port: try: - config.platforms[Platform.API_SERVER].extra["port"] = int(api_server_port) + config.platforms[Platform.API_SERVER].extra["port"] = int( + api_server_port + ) except ValueError: pass if api_server_host: config.platforms[Platform.API_SERVER].extra["host"] = api_server_host api_server_model_name = os.getenv("API_SERVER_MODEL_NAME", "") if api_server_model_name: - config.platforms[Platform.API_SERVER].extra["model_name"] = api_server_model_name + config.platforms[Platform.API_SERVER].extra["model_name"] = ( + api_server_model_name + ) # Webhook platform webhook_enabled = os.getenv("WEBHOOK_ENABLED", "").lower() in ("true", "1", "yes") @@ -995,18 +1106,22 @@ def _apply_env_overrides(config: GatewayConfig) -> None: if Platform.FEISHU not in config.platforms: config.platforms[Platform.FEISHU] = PlatformConfig() config.platforms[Platform.FEISHU].enabled = True - config.platforms[Platform.FEISHU].extra.update({ - "app_id": feishu_app_id, - "app_secret": feishu_app_secret, - "domain": os.getenv("FEISHU_DOMAIN", "feishu"), - "connection_mode": os.getenv("FEISHU_CONNECTION_MODE", "websocket"), - }) + config.platforms[Platform.FEISHU].extra.update( + { + "app_id": feishu_app_id, + "app_secret": feishu_app_secret, + "domain": os.getenv("FEISHU_DOMAIN", "feishu"), + "connection_mode": os.getenv("FEISHU_CONNECTION_MODE", "websocket"), + } + ) feishu_encrypt_key = os.getenv("FEISHU_ENCRYPT_KEY", "") if feishu_encrypt_key: config.platforms[Platform.FEISHU].extra["encrypt_key"] = feishu_encrypt_key feishu_verification_token = os.getenv("FEISHU_VERIFICATION_TOKEN", "") if feishu_verification_token: - config.platforms[Platform.FEISHU].extra["verification_token"] = feishu_verification_token + config.platforms[Platform.FEISHU].extra["verification_token"] = ( + feishu_verification_token + ) feishu_home = os.getenv("FEISHU_HOME_CHANNEL") if feishu_home: config.platforms[Platform.FEISHU].home_channel = HomeChannel( @@ -1022,10 +1137,12 @@ def _apply_env_overrides(config: GatewayConfig) -> None: if Platform.WECOM not in config.platforms: config.platforms[Platform.WECOM] = PlatformConfig() config.platforms[Platform.WECOM].enabled = True - config.platforms[Platform.WECOM].extra.update({ - "bot_id": wecom_bot_id, - "secret": wecom_secret, - }) + config.platforms[Platform.WECOM].extra.update( + { + "bot_id": wecom_bot_id, + "secret": wecom_secret, + } + ) wecom_ws_url = os.getenv("WECOM_WEBSOCKET_URL", "") if wecom_ws_url: config.platforms[Platform.WECOM].extra["websocket_url"] = wecom_ws_url @@ -1044,15 +1161,17 @@ def _apply_env_overrides(config: GatewayConfig) -> None: if Platform.WECOM_CALLBACK not in config.platforms: config.platforms[Platform.WECOM_CALLBACK] = PlatformConfig() config.platforms[Platform.WECOM_CALLBACK].enabled = True - config.platforms[Platform.WECOM_CALLBACK].extra.update({ - "corp_id": wecom_callback_corp_id, - "corp_secret": wecom_callback_corp_secret, - "agent_id": os.getenv("WECOM_CALLBACK_AGENT_ID", ""), - "token": os.getenv("WECOM_CALLBACK_TOKEN", ""), - "encoding_aes_key": os.getenv("WECOM_CALLBACK_ENCODING_AES_KEY", ""), - "host": os.getenv("WECOM_CALLBACK_HOST", "0.0.0.0"), - "port": int(os.getenv("WECOM_CALLBACK_PORT", "8645")), - }) + config.platforms[Platform.WECOM_CALLBACK].extra.update( + { + "corp_id": wecom_callback_corp_id, + "corp_secret": wecom_callback_corp_secret, + "agent_id": os.getenv("WECOM_CALLBACK_AGENT_ID", ""), + "token": os.getenv("WECOM_CALLBACK_TOKEN", ""), + "encoding_aes_key": os.getenv("WECOM_CALLBACK_ENCODING_AES_KEY", ""), + "host": os.getenv("WECOM_CALLBACK_HOST", "0.0.0.0"), + "port": int(os.getenv("WECOM_CALLBACK_PORT", "8645")), + } + ) # Weixin (personal WeChat via iLink Bot API) weixin_token = os.getenv("WEIXIN_TOKEN") @@ -1084,7 +1203,9 @@ def _apply_env_overrides(config: GatewayConfig) -> None: weixin_group_allowed_users = os.getenv("WEIXIN_GROUP_ALLOWED_USERS", "").strip() if weixin_group_allowed_users: extra["group_allow_from"] = weixin_group_allowed_users - weixin_split_multiline = os.getenv("WEIXIN_SPLIT_MULTILINE_MESSAGES", "").strip() + weixin_split_multiline = os.getenv( + "WEIXIN_SPLIT_MULTILINE_MESSAGES", "" + ).strip() if weixin_split_multiline: extra["split_multiline_messages"] = weixin_split_multiline weixin_home = os.getenv("WEIXIN_HOME_CHANNEL", "").strip() @@ -1102,14 +1223,21 @@ def _apply_env_overrides(config: GatewayConfig) -> None: if Platform.BLUEBUBBLES not in config.platforms: config.platforms[Platform.BLUEBUBBLES] = PlatformConfig() config.platforms[Platform.BLUEBUBBLES].enabled = True - config.platforms[Platform.BLUEBUBBLES].extra.update({ - "server_url": bluebubbles_server_url.rstrip("/"), - "password": bluebubbles_password, - "webhook_host": os.getenv("BLUEBUBBLES_WEBHOOK_HOST", "127.0.0.1"), - "webhook_port": int(os.getenv("BLUEBUBBLES_WEBHOOK_PORT", "8645")), - "webhook_path": os.getenv("BLUEBUBBLES_WEBHOOK_PATH", "/bluebubbles-webhook"), - "send_read_receipts": os.getenv("BLUEBUBBLES_SEND_READ_RECEIPTS", "true").lower() in ("true", "1", "yes"), - }) + config.platforms[Platform.BLUEBUBBLES].extra.update( + { + "server_url": bluebubbles_server_url.rstrip("/"), + "password": bluebubbles_password, + "webhook_host": os.getenv("BLUEBUBBLES_WEBHOOK_HOST", "127.0.0.1"), + "webhook_port": int(os.getenv("BLUEBUBBLES_WEBHOOK_PORT", "8645")), + "webhook_path": os.getenv( + "BLUEBUBBLES_WEBHOOK_PATH", "/bluebubbles-webhook" + ), + "send_read_receipts": os.getenv( + "BLUEBUBBLES_SEND_READ_RECEIPTS", "true" + ).lower() + in ("true", "1", "yes"), + } + ) bluebubbles_home = os.getenv("BLUEBUBBLES_HOME_CHANNEL") if bluebubbles_home and Platform.BLUEBUBBLES in config.platforms: config.platforms[Platform.BLUEBUBBLES].home_channel = HomeChannel( @@ -1151,7 +1279,7 @@ def _apply_env_overrides(config: GatewayConfig) -> None: config.default_reset_policy.idle_minutes = int(idle_minutes) except ValueError: pass - + reset_hour = os.getenv("SESSION_RESET_HOUR") if reset_hour: try: diff --git a/gateway/display_config.py b/gateway/display_config.py index 78e8bc9afac03..ad382ba4860bf 100644 --- a/gateway/display_config.py +++ b/gateway/display_config.py @@ -75,30 +75,29 @@ _PLATFORM_DEFAULTS: dict[str, dict[str, Any]] = { # Tier 1 — full edit support, personal/team use - "telegram": _TIER_HIGH, - "discord": _TIER_HIGH, - + # Telegram gets a lower-latency default profile with less chatty tool + # progress to reduce edit pressure. + "telegram": {**_TIER_HIGH, "tool_progress": "new"}, + "discord": _TIER_HIGH, # Tier 2 — edit support, often customer/workspace channels - "slack": _TIER_MEDIUM, - "mattermost": _TIER_MEDIUM, - "matrix": _TIER_MEDIUM, - "feishu": _TIER_MEDIUM, - + "slack": _TIER_MEDIUM, + "mattermost": _TIER_MEDIUM, + "matrix": _TIER_MEDIUM, + "feishu": _TIER_MEDIUM, # Tier 3 — no edit support, progress messages are permanent - "signal": _TIER_LOW, - "whatsapp": _TIER_MEDIUM, # Baileys bridge supports /edit - "bluebubbles": _TIER_LOW, - "weixin": _TIER_LOW, - "wecom": _TIER_LOW, - "wecom_callback": _TIER_LOW, - "dingtalk": _TIER_LOW, - + "signal": _TIER_LOW, + "whatsapp": _TIER_MEDIUM, # Baileys bridge supports /edit + "bluebubbles": _TIER_LOW, + "weixin": _TIER_LOW, + "wecom": _TIER_LOW, + "wecom_callback": _TIER_LOW, + "dingtalk": _TIER_LOW, # Tier 4 — batch or non-interactive delivery - "email": _TIER_MINIMAL, - "sms": _TIER_MINIMAL, - "webhook": _TIER_MINIMAL, - "homeassistant": _TIER_MINIMAL, - "api_server": {**_TIER_HIGH, "tool_preview_length": 0}, + "email": _TIER_MINIMAL, + "sms": _TIER_MINIMAL, + "webhook": _TIER_MINIMAL, + "homeassistant": _TIER_MINIMAL, + "api_server": {**_TIER_HIGH, "tool_preview_length": 0}, } # Canonical set of per-platform overrideable keys (for validation). @@ -174,6 +173,7 @@ def resolve_display_setting( # Helpers # --------------------------------------------------------------------------- + def _normalise(setting: str, value: Any) -> Any: """Normalise YAML quirks (bare ``off`` → False in YAML 1.1).""" if setting == "tool_progress": diff --git a/gateway/platforms/telegram.py b/gateway/platforms/telegram.py index 112b232d0a49f..22e1040ee18ff 100644 --- a/gateway/platforms/telegram.py +++ b/gateway/platforms/telegram.py @@ -17,7 +17,13 @@ logger = logging.getLogger(__name__) try: - from telegram import Update, Bot, Message, InlineKeyboardButton, InlineKeyboardMarkup + from telegram import ( + Update, + Bot, + Message, + InlineKeyboardButton, + InlineKeyboardMarkup, + ) from telegram.ext import ( Application, CommandHandler, @@ -28,6 +34,7 @@ ) from telegram.constants import ParseMode, ChatType from telegram.request import HTTPXRequest + TELEGRAM_AVAILABLE = True except ImportError: TELEGRAM_AVAILABLE = False @@ -49,10 +56,12 @@ # don't crash during class definition when the library isn't installed. class _MockContextTypes: DEFAULT_TYPE = Any + ContextTypes = _MockContextTypes import sys from pathlib import Path as _Path + sys.path.insert(0, str(_Path(__file__).resolve().parents[2])) from gateway.config import Platform, PlatformConfig @@ -84,12 +93,12 @@ def check_telegram_requirements() -> bool: # Matches every character that MarkdownV2 requires to be backslash-escaped # when it appears outside a code span or fenced code block. -_MDV2_ESCAPE_RE = re.compile(r'([_*\[\]()~`>#\+\-=|{}.!\\])') +_MDV2_ESCAPE_RE = re.compile(r"([_*\[\]()~`>#\+\-=|{}.!\\])") def _escape_mdv2(text: str) -> str: """Escape Telegram MarkdownV2 special characters with a preceding backslash.""" - return _MDV2_ESCAPE_RE.sub(r'\\\1', text) + return _MDV2_ESCAPE_RE.sub(r"\\\1", text) def _strip_mdv2(text: str) -> str: @@ -99,55 +108,90 @@ def _strip_mdv2(text: str) -> str: doesn't show stray syntax characters from format_message conversion. """ # Remove escape backslashes before special characters - cleaned = re.sub(r'\\([_*\[\]()~`>#\+\-=|{}.!\\])', r'\1', text) + cleaned = re.sub(r"\\([_*\[\]()~`>#\+\-=|{}.!\\])", r"\1", text) # Remove MarkdownV2 bold markers that format_message converted from **bold** - cleaned = re.sub(r'\*([^*]+)\*', r'\1', cleaned) + cleaned = re.sub(r"\*([^*]+)\*", r"\1", cleaned) # Remove MarkdownV2 italic markers that format_message converted from *italic* # Use word boundary (\b) to avoid breaking snake_case like my_variable_name - cleaned = re.sub(r'(?<!\w)_([^_]+)_(?!\w)', r'\1', cleaned) + cleaned = re.sub(r"(?<!\w)_([^_]+)_(?!\w)", r"\1", cleaned) # Remove MarkdownV2 strikethrough markers (~text~ → text) - cleaned = re.sub(r'~([^~]+)~', r'\1', cleaned) + cleaned = re.sub(r"~([^~]+)~", r"\1", cleaned) # Remove MarkdownV2 spoiler markers (||text|| → text) - cleaned = re.sub(r'\|\|([^|]+)\|\|', r'\1', cleaned) + cleaned = re.sub(r"\|\|([^|]+)\|\|", r"\1", cleaned) return cleaned class TelegramAdapter(BasePlatformAdapter): """ Telegram bot adapter. - + Handles: - Receiving messages from users and groups - Sending responses with Telegram markdown - Forum topics (thread_id support) - Media messages """ - + # Telegram message limits MAX_MESSAGE_LENGTH = 4096 # Threshold for detecting Telegram client-side message splits. # When a chunk is near this limit, a continuation is almost certain. _SPLIT_THRESHOLD = 4000 MEDIA_GROUP_WAIT_SECONDS = 0.8 - + + @staticmethod + def _env_float_clamped( + name: str, + default: float, + *, + min_value: Optional[float] = None, + max_value: Optional[float] = None, + ) -> float: + """Read a float env var with bounds and sane fallback.""" + raw = os.getenv(name) + try: + value = float(raw) if raw is not None else float(default) + except (TypeError, ValueError): + value = float(default) + if min_value is not None: + value = max(value, min_value) + if max_value is not None: + value = min(value, max_value) + return value + def __init__(self, config: PlatformConfig): super().__init__(config, Platform.TELEGRAM) self._app: Optional[Application] = None self._bot: Optional[Bot] = None self._webhook_mode: bool = False self._mention_patterns = self._compile_mention_patterns() - self._reply_to_mode: str = getattr(config, 'reply_to_mode', 'first') or 'first' + self._reply_to_mode: str = getattr(config, "reply_to_mode", "first") or "first" # Buffer rapid/album photo updates so Telegram image bursts are handled # as a single MessageEvent instead of self-interrupting multiple turns. - self._media_batch_delay_seconds = float(os.getenv("HERMES_TELEGRAM_MEDIA_BATCH_DELAY_SECONDS", "0.8")) + self._media_batch_delay_seconds = self._env_float_clamped( + "HERMES_TELEGRAM_MEDIA_BATCH_DELAY_SECONDS", + 0.8, + min_value=0.2, + max_value=3.0, + ) self._pending_photo_batches: Dict[str, MessageEvent] = {} self._pending_photo_batch_tasks: Dict[str, asyncio.Task] = {} self._media_group_events: Dict[str, MessageEvent] = {} self._media_group_tasks: Dict[str, asyncio.Task] = {} # Buffer rapid text messages so Telegram client-side splits of long # messages are aggregated into a single MessageEvent. - self._text_batch_delay_seconds = float(os.getenv("HERMES_TELEGRAM_TEXT_BATCH_DELAY_SECONDS", "0.6")) - self._text_batch_split_delay_seconds = float(os.getenv("HERMES_TELEGRAM_TEXT_BATCH_SPLIT_DELAY_SECONDS", "2.0")) + self._text_batch_delay_seconds = self._env_float_clamped( + "HERMES_TELEGRAM_TEXT_BATCH_DELAY_SECONDS", + 0.3, + min_value=0.08, + max_value=2.0, + ) + self._text_batch_split_delay_seconds = self._env_float_clamped( + "HERMES_TELEGRAM_TEXT_BATCH_SPLIT_DELAY_SECONDS", + 1.0, + min_value=self._text_batch_delay_seconds, + max_value=4.0, + ) self._pending_text_batches: Dict[str, MessageEvent] = {} self._pending_text_batch_tasks: Dict[str, asyncio.Task] = {} self._polling_error_task: Optional[asyncio.Task] = None @@ -157,7 +201,9 @@ def __init__(self, config: PlatformConfig): # DM Topics: map of topic_name -> message_thread_id (populated at startup) self._dm_topics: Dict[str, int] = {} # DM Topics config from extra.dm_topics - self._dm_topics_config: List[Dict[str, Any]] = self.config.extra.get("dm_topics", []) + self._dm_topics_config: List[Dict[str, Any]] = self.config.extra.get( + "dm_topics", [] + ) # Interactive model picker state per chat self._model_picker_state: Dict[str, dict] = {} # Approval button state: message_id → session_key @@ -165,10 +211,16 @@ def __init__(self, config: PlatformConfig): def _fallback_ips(self) -> list[str]: """Return validated fallback IPs from config (populated by _apply_env_overrides).""" - configured = self.config.extra.get("fallback_ips", []) if getattr(self.config, "extra", None) else [] + configured = ( + self.config.extra.get("fallback_ips", []) + if getattr(self.config, "extra", None) + else [] + ) if isinstance(configured, str): configured = configured.split(",") - return parse_fallback_ip_env(",".join(str(v) for v in configured) if configured else None) + return parse_fallback_ip_env( + ",".join(str(v) for v in configured) if configured else None + ) @staticmethod def _looks_like_polling_conflict(error: Exception) -> bool: @@ -187,6 +239,7 @@ def _looks_like_network_error(error: Exception) -> bool: return True try: from telegram.error import NetworkError, TimedOut + if isinstance(error, (NetworkError, TimedOut)): return True except ImportError: @@ -228,7 +281,11 @@ async def _handle_polling_network_error(self, error: Exception) -> None: delay = min(BASE_DELAY * (2 ** (attempt - 1)), MAX_DELAY) logger.warning( "[%s] Telegram network error (attempt %d/%d), reconnecting in %ds. Error: %s", - self.name, attempt, MAX_NETWORK_RETRIES, delay, error, + self.name, + attempt, + MAX_NETWORK_RETRIES, + delay, + error, ) await asyncio.sleep(delay) @@ -246,11 +303,14 @@ async def _handle_polling_network_error(self, error: Exception) -> None: ) logger.info( "[%s] Telegram polling resumed after network error (attempt %d)", - self.name, attempt, + self.name, + attempt, ) self._polling_network_error_count = 0 except Exception as retry_err: - logger.warning("[%s] Telegram polling reconnect failed: %s", self.name, retry_err) + logger.warning( + "[%s] Telegram polling reconnect failed: %s", self.name, retry_err + ) # start_polling failed — polling is dead and no further error # callbacks will fire, so schedule the next retry ourselves. if not self.has_fatal_error: @@ -261,7 +321,10 @@ async def _handle_polling_network_error(self, error: Exception) -> None: task.add_done_callback(self._background_tasks.discard) async def _handle_polling_conflict(self, error: Exception) -> None: - if self.has_fatal_error and self.fatal_error_code == "telegram_polling_conflict": + if ( + self.has_fatal_error + and self.fatal_error_code == "telegram_polling_conflict" + ): return # Track consecutive conflicts — transient 409s can occur when a # previous gateway instance hasn't fully released its long-poll @@ -276,8 +339,11 @@ async def _handle_polling_conflict(self, error: Exception) -> None: if self._polling_conflict_count <= MAX_CONFLICT_RETRIES: logger.warning( "[%s] Telegram polling conflict (%d/%d), will retry in %ds. Error: %s", - self.name, self._polling_conflict_count, MAX_CONFLICT_RETRIES, - RETRY_DELAY, error, + self.name, + self._polling_conflict_count, + MAX_CONFLICT_RETRIES, + RETRY_DELAY, + error, ) try: if self._app and self._app.updater and self._app.updater.running: @@ -291,11 +357,17 @@ async def _handle_polling_conflict(self, error: Exception) -> None: drop_pending_updates=False, error_callback=self._polling_error_callback_ref, ) - logger.info("[%s] Telegram polling resumed after conflict retry %d", self.name, self._polling_conflict_count) + logger.info( + "[%s] Telegram polling resumed after conflict retry %d", + self.name, + self._polling_conflict_count, + ) self._polling_conflict_count = 0 # reset on success return except Exception as retry_err: - logger.warning("[%s] Telegram polling retry failed: %s", self.name, retry_err) + logger.warning( + "[%s] Telegram polling retry failed: %s", self.name, retry_err + ) # Don't fall through to fatal yet — wait for the next conflict # to trigger another retry attempt (up to MAX_CONFLICT_RETRIES). return @@ -306,8 +378,7 @@ async def _handle_polling_conflict(self, error: Exception) -> None: "(possibly OpenClaw or another Hermes instance). " "Hermes stopped Telegram polling after %d retries. " "Only one poller can run per token — stop the other process " - "and restart with 'hermes start'." - % MAX_CONFLICT_RETRIES + "and restart with 'hermes start'." % MAX_CONFLICT_RETRIES ) logger.error("[%s] %s Original error: %s", self.name, message, error) self._set_fatal_error("telegram_polling_conflict", message, retryable=False) @@ -315,7 +386,12 @@ async def _handle_polling_conflict(self, error: Exception) -> None: if self._app and self._app.updater: await self._app.updater.stop() except Exception as stop_error: - logger.warning("[%s] Failed stopping Telegram polling after conflict: %s", self.name, stop_error, exc_info=True) + logger.warning( + "[%s] Failed stopping Telegram polling after conflict: %s", + self.name, + stop_error, + exc_info=True, + ) await self._notify_fatal_error() async def _create_dm_topic( @@ -343,7 +419,10 @@ async def _create_dm_topic( thread_id = topic.message_thread_id logger.info( "[%s] Created DM topic '%s' in chat %s -> thread_id=%s", - self.name, name, chat_id, thread_id, + self.name, + name, + chat_id, + thread_id, ) return thread_id except Exception as e: @@ -353,25 +432,38 @@ async def _create_dm_topic( if "topic_name_duplicate" in error_text or "already" in error_text: logger.info( "[%s] DM topic '%s' already exists in chat %s (will be mapped from incoming messages)", - self.name, name, chat_id, + self.name, + name, + chat_id, ) else: logger.warning( "[%s] Failed to create DM topic '%s' in chat %s: %s", - self.name, name, chat_id, e, + self.name, + name, + chat_id, + e, ) return None - def _persist_dm_topic_thread_id(self, chat_id: int, topic_name: str, thread_id: int) -> None: + def _persist_dm_topic_thread_id( + self, chat_id: int, topic_name: str, thread_id: int + ) -> None: """Save a newly created thread_id back into config.yaml so it persists across restarts.""" try: from hermes_constants import get_hermes_home + config_path = get_hermes_home() / "config.yaml" if not config_path.exists(): - logger.warning("[%s] Config file not found at %s, cannot persist thread_id", self.name, config_path) + logger.warning( + "[%s] Config file not found at %s, cannot persist thread_id", + self.name, + config_path, + ) return import yaml as _yaml + with open(config_path, "r") as f: config = _yaml.safe_load(f) or {} @@ -400,10 +492,17 @@ def _persist_dm_topic_thread_id(self, chat_id: int, topic_name: str, thread_id: _yaml.dump(config, f, default_flow_style=False, sort_keys=False) logger.info( "[%s] Persisted thread_id=%s for topic '%s' in config.yaml", - self.name, thread_id, topic_name, + self.name, + thread_id, + topic_name, ) except Exception as e: - logger.warning("[%s] Failed to persist thread_id to config: %s", self.name, e, exc_info=True) + logger.warning( + "[%s] Failed to persist thread_id to config: %s", + self.name, + e, + exc_info=True, + ) async def _setup_dm_topics(self) -> None: """Load or create configured DM topics for specified chats. @@ -435,7 +534,9 @@ async def _setup_dm_topics(self) -> None: logger.info( "[%s] Setting up %d DM topic(s) for chat %s", - self.name, len(topics), chat_id, + self.name, + len(topics), + chat_id, ) for topic_conf in topics: @@ -451,7 +552,9 @@ async def _setup_dm_topics(self) -> None: self._dm_topics[cache_key] = int(existing_thread_id) logger.info( "[%s] DM topic loaded from config: %s -> thread_id=%s", - self.name, cache_key, existing_thread_id, + self.name, + cache_key, + existing_thread_id, ) continue @@ -470,10 +573,14 @@ async def _setup_dm_topics(self) -> None: self._dm_topics[cache_key] = thread_id logger.info( "[%s] DM topic cached: %s -> thread_id=%s", - self.name, cache_key, thread_id, + self.name, + cache_key, + thread_id, ) # Persist thread_id to config so we don't recreate on next restart - self._persist_dm_topic_thread_id(int(chat_id), topic_name, thread_id) + self._persist_dm_topic_thread_id( + int(chat_id), topic_name, thread_id + ) async def connect(self) -> bool: """Connect to Telegram via polling or webhook. @@ -495,13 +602,15 @@ async def connect(self) -> bool: self.name, ) return False - + if not self.config.token: logger.error("[%s] No bot token configured", self.name) return False - + try: - if not self._acquire_platform_lock('telegram-bot-token', self.config.token, 'Telegram bot token'): + if not self._acquire_platform_lock( + "telegram-bot-token", self.config.token, "Telegram bot token" + ): return False # Build the application @@ -514,7 +623,8 @@ async def connect(self) -> bool: ) logger.info( "[%s] Using custom Telegram base_url: %s", - self.name, custom_base_url, + self.name, + custom_base_url, ) # PTB defaults (pool_timeout=1s) are too aggressive on flaky networks and @@ -535,13 +645,17 @@ def _env_float(name: str, default: float) -> float: request_kwargs = { "connection_pool_size": _env_int("HERMES_TELEGRAM_HTTP_POOL_SIZE", 512), "pool_timeout": _env_float("HERMES_TELEGRAM_HTTP_POOL_TIMEOUT", 8.0), - "connect_timeout": _env_float("HERMES_TELEGRAM_HTTP_CONNECT_TIMEOUT", 10.0), + "connect_timeout": _env_float( + "HERMES_TELEGRAM_HTTP_CONNECT_TIMEOUT", 10.0 + ), "read_timeout": _env_float("HERMES_TELEGRAM_HTTP_READ_TIMEOUT", 20.0), "write_timeout": _env_float("HERMES_TELEGRAM_HTTP_WRITE_TIMEOUT", 20.0), } proxy_url = resolve_proxy_url() - disable_fallback = (os.getenv("HERMES_TELEGRAM_DISABLE_FALLBACK_IPS", "").strip().lower() in ("1", "true", "yes", "on")) + disable_fallback = os.getenv( + "HERMES_TELEGRAM_DISABLE_FALLBACK_IPS", "" + ).strip().lower() in ("1", "true", "yes", "on") fallback_ips = self._fallback_ips() if not fallback_ips: fallback_ips = await discover_fallback_ips() @@ -568,39 +682,55 @@ def _env_float(name: str, default: float) -> float: httpx_kwargs={"transport": TelegramFallbackTransport(fallback_ips)}, ) elif proxy_url: - logger.info("[%s] Proxy detected; passing explicitly to HTTPXRequest: %s", self.name, proxy_url) + logger.info( + "[%s] Proxy detected; passing explicitly to HTTPXRequest: %s", + self.name, + proxy_url, + ) request = HTTPXRequest(**request_kwargs, proxy=proxy_url) get_updates_request = HTTPXRequest(**request_kwargs, proxy=proxy_url) else: if disable_fallback: - logger.info("[%s] Telegram fallback-IP transport disabled via env", self.name) + logger.info( + "[%s] Telegram fallback-IP transport disabled via env", + self.name, + ) request = HTTPXRequest(**request_kwargs) get_updates_request = HTTPXRequest(**request_kwargs) builder = builder.request(request).get_updates_request(get_updates_request) self._app = builder.build() self._bot = self._app.bot - + # Register handlers - self._app.add_handler(TelegramMessageHandler( - filters.TEXT & ~filters.COMMAND, - self._handle_text_message - )) - self._app.add_handler(TelegramMessageHandler( - filters.COMMAND, - self._handle_command - )) - self._app.add_handler(TelegramMessageHandler( - filters.LOCATION | getattr(filters, "VENUE", filters.LOCATION), - self._handle_location_message - )) - self._app.add_handler(TelegramMessageHandler( - filters.PHOTO | filters.VIDEO | filters.AUDIO | filters.VOICE | filters.Document.ALL | filters.Sticker.ALL, - self._handle_media_message - )) + self._app.add_handler( + TelegramMessageHandler( + filters.TEXT & ~filters.COMMAND, self._handle_text_message + ) + ) + self._app.add_handler( + TelegramMessageHandler(filters.COMMAND, self._handle_command) + ) + self._app.add_handler( + TelegramMessageHandler( + filters.LOCATION | getattr(filters, "VENUE", filters.LOCATION), + self._handle_location_message, + ) + ) + self._app.add_handler( + TelegramMessageHandler( + filters.PHOTO + | filters.VIDEO + | filters.AUDIO + | filters.VOICE + | filters.Document.ALL + | filters.Sticker.ALL, + self._handle_media_message, + ) + ) # Handle inline keyboard button callbacks (update prompts) self._app.add_handler(CallbackQueryHandler(self._handle_callback_query)) - + # Start polling — retry initialize() for transient TLS resets try: from telegram.error import NetworkError, TimedOut @@ -613,10 +743,14 @@ def _env_float(name: str, default: float) -> float: break except (NetworkError, TimedOut, OSError) as init_err: if _attempt < _max_connect - 1: - wait = 2 ** _attempt + wait = 2**_attempt logger.warning( "[%s] Connect attempt %d/%d failed: %s — retrying in %ds", - self.name, _attempt + 1, _max_connect, init_err, wait, + self.name, + _attempt + 1, + _max_connect, + init_err, + wait, ) await asyncio.sleep(wait) else: @@ -632,8 +766,11 @@ def _env_float(name: str, default: float) -> float: # enables cloud platforms (Fly.io, Railway) to auto-wake # suspended machines on inbound HTTP traffic. webhook_port = int(os.getenv("TELEGRAM_WEBHOOK_PORT", "8443")) - webhook_secret = os.getenv("TELEGRAM_WEBHOOK_SECRET", "").strip() or None + webhook_secret = ( + os.getenv("TELEGRAM_WEBHOOK_SECRET", "").strip() or None + ) from urllib.parse import urlparse + webhook_path = urlparse(webhook_url).path or "/telegram" await self._app.updater.start_webhook( @@ -648,7 +785,9 @@ def _env_float(name: str, default: float) -> float: self._webhook_mode = True logger.info( "[%s] Webhook server listening on 0.0.0.0:%d%s", - self.name, webhook_port, webhook_path, + self.name, + webhook_port, + webhook_path, ) else: # ── Polling mode (default) ─────────────────────────── @@ -664,12 +803,25 @@ def _polling_error_callback(error: Exception) -> None: if self._polling_error_task and not self._polling_error_task.done(): return if self._looks_like_polling_conflict(error): - self._polling_error_task = loop.create_task(self._handle_polling_conflict(error)) + self._polling_error_task = loop.create_task( + self._handle_polling_conflict(error) + ) elif self._looks_like_network_error(error): - logger.warning("[%s] Telegram network error, scheduling reconnect: %s", self.name, error) - self._polling_error_task = loop.create_task(self._handle_polling_network_error(error)) + logger.warning( + "[%s] Telegram network error, scheduling reconnect: %s", + self.name, + error, + ) + self._polling_error_task = loop.create_task( + self._handle_polling_network_error(error) + ) else: - logger.error("[%s] Telegram polling error: %s", self.name, error, exc_info=True) + logger.error( + "[%s] Telegram polling error: %s", + self.name, + error, + exc_info=True, + ) # Store reference for retry use in _handle_polling_conflict self._polling_error_callback_ref = _polling_error_callback @@ -679,24 +831,27 @@ def _polling_error_callback(error: Exception) -> None: drop_pending_updates=True, error_callback=_polling_error_callback, ) - + # Register bot commands so Telegram shows a hint menu when users type / # List is derived from the central COMMAND_REGISTRY — adding a new # gateway command there automatically adds it to the Telegram menu. try: from telegram import BotCommand from hermes_cli.commands import telegram_menu_commands + # Telegram allows up to 100 commands but has an undocumented # payload size limit. Skill descriptions are truncated to 40 # chars in telegram_menu_commands() to fit 100 commands safely. menu_commands, hidden_count = telegram_menu_commands(max_commands=100) - await self._bot.set_my_commands([ - BotCommand(name, desc) for name, desc in menu_commands - ]) + await self._bot.set_my_commands( + [BotCommand(name, desc) for name, desc in menu_commands] + ) if hidden_count: logger.info( "[%s] Telegram menu: %d commands registered, %d hidden (over 100 limit). Use /commands for full list.", - self.name, len(menu_commands), hidden_count, + self.name, + len(menu_commands), + hidden_count, ) except Exception as e: logger.warning( @@ -705,7 +860,7 @@ def _polling_error_callback(error: Exception) -> None: e, exc_info=True, ) - + self._mark_connected() mode = "webhook" if self._webhook_mode else "polling" logger.info("[%s] Connected to Telegram (%s mode)", self.name, mode) @@ -718,18 +873,22 @@ def _polling_error_callback(error: Exception) -> None: except Exception as topics_err: logger.warning( "[%s] DM topics setup failed (non-fatal): %s", - self.name, topics_err, exc_info=True, + self.name, + topics_err, + exc_info=True, ) return True - + except Exception as e: self._release_platform_lock() message = f"Telegram startup failed: {e}" self._set_fatal_error("telegram_connect_error", message, retryable=True) - logger.error("[%s] Failed to connect to Telegram: %s", self.name, e, exc_info=True) + logger.error( + "[%s] Failed to connect to Telegram: %s", self.name, e, exc_info=True + ) return False - + async def disconnect(self) -> None: """Stop polling/webhook, cancel pending album flushes, and disconnect.""" pending_media_group_tasks = list(self._media_group_tasks.values()) @@ -749,7 +908,12 @@ async def disconnect(self) -> None: await self._app.stop() await self._app.shutdown() except Exception as e: - logger.warning("[%s] Error during Telegram disconnect: %s", self.name, e, exc_info=True) + logger.warning( + "[%s] Error during Telegram disconnect: %s", + self.name, + e, + exc_info=True, + ) self._release_platform_lock() for task in self._pending_photo_batch_tasks.values(): @@ -788,21 +952,23 @@ async def send( chat_id: str, content: str, reply_to: Optional[str] = None, - metadata: Optional[Dict[str, Any]] = None + metadata: Optional[Dict[str, Any]] = None, ) -> SendResult: """Send a message to a Telegram chat.""" if not self._bot: return SendResult(success=False, error="Not connected") - + # Skip whitespace-only text to prevent Telegram 400 empty-text errors. if not content or not content.strip(): return SendResult(success=True, message_id=None) - + try: # Format and split message if needed formatted = self.format_message(content) chunks = self.truncate_message( - formatted, self.MAX_MESSAGE_LENGTH, len_fn=utf16_len, + formatted, + self.MAX_MESSAGE_LENGTH, + len_fn=utf16_len, ) if len(chunks) > 1: # truncate_message appends a raw " (1/2)" suffix. Escape the @@ -812,10 +978,10 @@ async def send( re.sub(r" \((\d+)/(\d+)\)$", r" \\(\1/\2\\)", chunk) for chunk in chunks ] - + message_ids = [] thread_id = metadata.get("thread_id") if metadata else None - + try: from telegram.error import NetworkError as _NetErr except ImportError: @@ -850,8 +1016,15 @@ async def send( ) except Exception as md_error: # Markdown parsing failed, try plain text - if "parse" in str(md_error).lower() or "markdown" in str(md_error).lower(): - logger.warning("[%s] MarkdownV2 parse failed, falling back to plain text: %s", self.name, md_error) + if ( + "parse" in str(md_error).lower() + or "markdown" in str(md_error).lower() + ): + logger.warning( + "[%s] MarkdownV2 parse failed, falling back to plain text: %s", + self.name, + md_error, + ) plain_chunk = _strip_mdv2(chunk) msg = await self._bot.send_message( chat_id=int(chat_id), @@ -870,23 +1043,31 @@ async def send( # specific cases instead of blindly retrying. if _BadReq and isinstance(send_err, _BadReq): err_lower = str(send_err).lower() - if "thread not found" in err_lower and effective_thread_id is not None: + if ( + "thread not found" in err_lower + and effective_thread_id is not None + ): # Thread doesn't exist — retry without # message_thread_id so the message still # reaches the chat. logger.warning( "[%s] Thread %s not found, retrying without message_thread_id", - self.name, effective_thread_id, + self.name, + effective_thread_id, ) effective_thread_id = None continue - if "message to be replied not found" in err_lower and reply_to_id is not None: + if ( + "message to be replied not found" in err_lower + and reply_to_id is not None + ): # Original message was deleted before we # could reply — clear reply target and retry # so the response is still delivered. logger.warning( "[%s] Reply target deleted, retrying without reply_to: %s", - self.name, send_err, + self.name, + send_err, ) reply_to_id = None continue @@ -898,17 +1079,29 @@ async def send( if _TimedOut and isinstance(send_err, _TimedOut): raise if _send_attempt < 2: - wait = 2 ** _send_attempt - logger.warning("[%s] Network error on send (attempt %d/3), retrying in %ds: %s", - self.name, _send_attempt + 1, wait, send_err) + wait = 2**_send_attempt + logger.warning( + "[%s] Network error on send (attempt %d/3), retrying in %ds: %s", + self.name, + _send_attempt + 1, + wait, + send_err, + ) await asyncio.sleep(wait) else: raise except Exception as send_err: retry_after = getattr(send_err, "retry_after", None) - if retry_after is not None or "retry after" in str(send_err).lower(): + if ( + retry_after is not None + or "retry after" in str(send_err).lower() + ): if _send_attempt < 2: - wait = float(retry_after) if retry_after is not None else 1.0 + wait = ( + float(retry_after) + if retry_after is not None + else 1.0 + ) logger.warning( "[%s] Telegram flood control on send (attempt %d/3), retrying in %.1fs: %s", self.name, @@ -920,15 +1113,17 @@ async def send( continue raise message_ids.append(str(msg.message_id)) - + return SendResult( success=True, message_id=message_ids[0] if message_ids else None, - raw_response={"message_ids": message_ids} + raw_response={"message_ids": message_ids}, ) - + except Exception as e: - logger.error("[%s] Failed to send Telegram message: %s", self.name, e, exc_info=True) + logger.error( + "[%s] Failed to send Telegram message: %s", self.name, e, exc_info=True + ) # TimedOut means the request may have reached Telegram — # mark as non-retryable so _send_with_retry() doesn't re-send. _to = locals().get("_TimedOut") @@ -974,9 +1169,10 @@ async def edit_message( # streaming). Truncate and succeed so the stream consumer can # split the overflow into a new message instead of dying. if "message_too_long" in err_str or "too long" in err_str: - truncated = _prefix_within_utf16_limit( - content, self.MAX_MESSAGE_LENGTH - 20 - ) + "…" + truncated = ( + _prefix_within_utf16_limit(content, self.MAX_MESSAGE_LENGTH - 20) + + "…" + ) try: await self._bot.edit_message_text( chat_id=int(chat_id), @@ -994,7 +1190,8 @@ async def edit_message( wait = retry_after if retry_after else 1.0 logger.warning( "[%s] Telegram flood control, waiting %.1fs", - self.name, wait, + self.name, + wait, ) if wait > 5.0: return SendResult(success=False, error=f"flood_control:{wait}") @@ -1009,7 +1206,8 @@ async def edit_message( except Exception as retry_err: logger.error( "[%s] Edit retry failed after flood wait: %s", - self.name, retry_err, + self.name, + retry_err, ) return SendResult(success=False, error=str(retry_err)) logger.error( @@ -1022,7 +1220,10 @@ async def edit_message( return SendResult(success=False, error=str(e)) async def send_update_prompt( - self, chat_id: str, prompt: str, default: str = "", + self, + chat_id: str, + prompt: str, + default: str = "", session_key: str = "", ) -> SendResult: """Send an inline-keyboard update prompt (Yes / No buttons). @@ -1035,12 +1236,14 @@ async def send_update_prompt( try: default_hint = f" (default: {default})" if default else "" text = f"⚕ *Update needs your input:*\n\n{prompt}{default_hint}" - keyboard = InlineKeyboardMarkup([ + keyboard = InlineKeyboardMarkup( [ - InlineKeyboardButton("✓ Yes", callback_data="update_prompt:y"), - InlineKeyboardButton("✗ No", callback_data="update_prompt:n"), + [ + InlineKeyboardButton("✓ Yes", callback_data="update_prompt:y"), + InlineKeyboardButton("✗ No", callback_data="update_prompt:n"), + ] ] - ]) + ) msg = await self._bot.send_message( chat_id=int(chat_id), text=text, @@ -1053,7 +1256,10 @@ async def send_update_prompt( return SendResult(success=False, error=str(e)) async def send_exec_approval( - self, chat_id: str, command: str, session_key: str, + self, + chat_id: str, + command: str, + session_key: str, description: str = "dangerous command", metadata: Optional[Dict[str, Any]] = None, ) -> SendResult: @@ -1076,26 +1282,39 @@ async def send_exec_approval( # Resolve thread context for thread replies thread_id = None if metadata: - thread_id = metadata.get("thread_id") or metadata.get("message_thread_id") + thread_id = metadata.get("thread_id") or metadata.get( + "message_thread_id" + ) # We'll use the message_id as part of callback_data to look up session_key # Send a placeholder first, then update — or use a counter. # Simpler: use a monotonic counter to generate short IDs. import itertools + if not hasattr(self, "_approval_counter"): self._approval_counter = itertools.count(1) approval_id = next(self._approval_counter) - keyboard = InlineKeyboardMarkup([ + keyboard = InlineKeyboardMarkup( [ - InlineKeyboardButton("✅ Allow Once", callback_data=f"ea:once:{approval_id}"), - InlineKeyboardButton("✅ Session", callback_data=f"ea:session:{approval_id}"), - ], - [ - InlineKeyboardButton("✅ Always", callback_data=f"ea:always:{approval_id}"), - InlineKeyboardButton("❌ Deny", callback_data=f"ea:deny:{approval_id}"), - ], - ]) + [ + InlineKeyboardButton( + "✅ Allow Once", callback_data=f"ea:once:{approval_id}" + ), + InlineKeyboardButton( + "✅ Session", callback_data=f"ea:session:{approval_id}" + ), + ], + [ + InlineKeyboardButton( + "✅ Always", callback_data=f"ea:always:{approval_id}" + ), + InlineKeyboardButton( + "❌ Deny", callback_data=f"ea:deny:{approval_id}" + ), + ], + ] + ) kwargs: Dict[str, Any] = { "chat_id": int(chat_id), @@ -1137,6 +1356,7 @@ async def send_model_picker( try: from hermes_cli.providers import get_label except ImportError: + def get_label(slug): return slug @@ -1208,9 +1428,7 @@ def _build_model_keyboard(self, models: list, page: int) -> tuple: short = model_id.split("/")[-1] if "/" in model_id else model_id if len(short) > 38: short = short[:35] + "..." - buttons.append( - InlineKeyboardButton(short, callback_data=f"mm:{abs_idx}") - ) + buttons.append(InlineKeyboardButton(short, callback_data=f"mm:{abs_idx}")) rows = [buttons[i : i + 2] for i in range(0, len(buttons), 2)] @@ -1218,16 +1436,26 @@ def _build_model_keyboard(self, models: list, page: int) -> tuple: if total_pages > 1: nav: list = [] if page > 0: - nav.append(InlineKeyboardButton("◀ Prev", callback_data=f"mg:{page - 1}")) - nav.append(InlineKeyboardButton(f"{page + 1}/{total_pages}", callback_data="mx:noop")) + nav.append( + InlineKeyboardButton("◀ Prev", callback_data=f"mg:{page - 1}") + ) + nav.append( + InlineKeyboardButton( + f"{page + 1}/{total_pages}", callback_data="mx:noop" + ) + ) if page < total_pages - 1: - nav.append(InlineKeyboardButton("Next ▶", callback_data=f"mg:{page + 1}")) + nav.append( + InlineKeyboardButton("Next ▶", callback_data=f"mg:{page + 1}") + ) rows.append(nav) - rows.append([ - InlineKeyboardButton("◀ Back", callback_data="mb"), - InlineKeyboardButton("✗ Cancel", callback_data="mx"), - ]) + rows.append( + [ + InlineKeyboardButton("◀ Back", callback_data="mb"), + InlineKeyboardButton("✗ Cancel", callback_data="mx"), + ] + ) page_info = f" ({start + 1}–{end} of {total})" if total_pages > 1 else "" return InlineKeyboardMarkup(rows), page_info @@ -1244,6 +1472,7 @@ async def _handle_model_picker_callback( try: from hermes_cli.providers import get_label except ImportError: + def get_label(slug): return slug @@ -1269,7 +1498,11 @@ def get_label(slug): pname = provider.get("name", provider_slug) total = provider.get("total_models", len(models)) shown = len(models) - extra = f"\n_{total - shown} more available — type `/model <name>` directly_" if total > shown else "" + extra = ( + f"\n_{total - shown} more available — type `/model <name>` directly_" + if total > shown + else "" + ) await query.edit_message_text( text=( @@ -1301,9 +1534,15 @@ def get_label(slug): (p for p in state["providers"] if p["slug"] == provider_slug), None, ) - total = provider.get("total_models", len(models)) if provider else len(models) + total = ( + provider.get("total_models", len(models)) if provider else len(models) + ) shown = len(models) - extra = f"\n_{total - shown} more available — type `/model <name>` directly_" if total > shown else "" + extra = ( + f"\n_{total - shown} more available — type `/model <name>` directly_" + if total > shown + else "" + ) await query.edit_message_text( text=( @@ -1442,9 +1681,13 @@ async def _handle_callback_query( caller_id = str(getattr(query.from_user, "id", "")) allowed_csv = os.getenv("TELEGRAM_ALLOWED_USERS", "").strip() if allowed_csv: - allowed_ids = {uid.strip() for uid in allowed_csv.split(",") if uid.strip()} + allowed_ids = { + uid.strip() for uid in allowed_csv.split(",") if uid.strip() + } if "*" not in allowed_ids and caller_id not in allowed_ids: - await query.answer(text="⛔ You are not authorized to approve commands.") + await query.answer( + text="⛔ You are not authorized to approve commands." + ) return session_key = self._approval_state.pop(approval_id, None) @@ -1477,13 +1720,20 @@ async def _handle_callback_query( # Resolve the approval — unblocks the agent thread try: from tools.approval import resolve_gateway_approval + count = resolve_gateway_approval(session_key, choice) logger.info( "Telegram button resolved %d approval(s) for session %s (choice=%s, user=%s)", - count, session_key, choice, user_display, + count, + session_key, + choice, + user_display, ) except Exception as exc: - logger.error("Failed to resolve gateway approval from Telegram button: %s", exc) + logger.error( + "Failed to resolve gateway approval from Telegram button: %s", + exc, + ) return # --- Update prompt callbacks --- @@ -1504,13 +1754,17 @@ async def _handle_callback_query( # Write the response file try: from hermes_constants import get_hermes_home + home = get_hermes_home() response_path = home / ".update_response" tmp = response_path.with_suffix(".tmp") tmp.write_text(answer) tmp.replace(response_path) - logger.info("Telegram update prompt answered '%s' by user %s", - answer, getattr(query.from_user, "id", "unknown")) + logger.info( + "Telegram update prompt answered '%s' by user %s", + answer, + getattr(query.from_user, "id", "unknown"), + ) except Exception as exc: logger.error("Failed to write update response from callback: %s", exc) @@ -1526,12 +1780,15 @@ async def send_voice( """Send audio as a native Telegram voice message or audio file.""" if not self._bot: return SendResult(success=False, error="Not connected") - + try: import os + if not os.path.exists(audio_path): - return SendResult(success=False, error=f"Audio file not found: {audio_path}") - + return SendResult( + success=False, error=f"Audio file not found: {audio_path}" + ) + with open(audio_path, "rb") as audio_file: # .ogg files -> send as voice (round playable bubble) if audio_path.endswith((".ogg", ".opus")): @@ -1562,7 +1819,7 @@ async def send_voice( exc_info=True, ) return await super().send_voice(chat_id, audio_path, caption, reply_to) - + async def send_image_file( self, chat_id: str, @@ -1578,8 +1835,11 @@ async def send_image_file( try: import os + if not os.path.exists(image_path): - return SendResult(success=False, error=f"Image file not found: {image_path}") + return SendResult( + success=False, error=f"Image file not found: {image_path}" + ) _thread = metadata.get("thread_id") if metadata else None with open(image_path, "rb") as image_file: @@ -1633,7 +1893,9 @@ async def send_document( return SendResult(success=True, message_id=str(msg.message_id)) except Exception as e: print(f"[{self.name}] Failed to send document: {e}") - return await super().send_document(chat_id, file_path, caption, file_name, reply_to) + return await super().send_document( + chat_id, file_path, caption, file_name, reply_to + ) async def send_video( self, @@ -1650,7 +1912,9 @@ async def send_video( try: if not os.path.exists(video_path): - return SendResult(success=False, error=f"Video file not found: {video_path}") + return SendResult( + success=False, error=f"Video file not found: {video_path}" + ) _thread = metadata.get("thread_id") if metadata else None with open(video_path, "rb") as f: @@ -1675,7 +1939,7 @@ async def send_image( metadata: Optional[Dict[str, Any]] = None, ) -> SendResult: """Send an image natively as a Telegram photo. - + Tries URL-based send first (fast, works for <5MB images). Falls back to downloading and uploading as file (supports up to 10MB). """ @@ -1683,9 +1947,12 @@ async def send_image( return SendResult(success=False, error="Not connected") from tools.url_safety import is_safe_url + if not is_safe_url(image_url): logger.warning("[%s] Blocked unsafe image URL (SSRF protection)", self.name) - return await super().send_image(chat_id, image_url, caption, reply_to, metadata=metadata) + return await super().send_image( + chat_id, image_url, caption, reply_to, metadata=metadata + ) try: # Telegram can send photos directly from URLs (up to ~5MB) @@ -1708,11 +1975,12 @@ async def send_image( # Fallback: download and upload as file (supports up to 10MB) try: import httpx + async with httpx.AsyncClient(timeout=30.0) as client: resp = await client.get(image_url) resp.raise_for_status() image_data = resp.content - + msg = await self._bot.send_photo( chat_id=int(chat_id), photo=image_data, @@ -1729,7 +1997,7 @@ async def send_image( ) # Final fallback: send URL as text return await super().send_image(chat_id, image_url, caption, reply_to) - + async def send_animation( self, chat_id: str, @@ -1741,7 +2009,7 @@ async def send_animation( """Send an animated GIF natively as a Telegram animation (auto-plays inline).""" if not self._bot: return SendResult(success=False, error="Not connected") - + try: _anim_thread = metadata.get("thread_id") if metadata else None msg = await self._bot.send_animation( @@ -1762,7 +2030,9 @@ async def send_animation( # Fallback: try as a regular photo return await self.send_image(chat_id, animation_url, caption, reply_to) - async def send_typing(self, chat_id: str, metadata: Optional[Dict[str, Any]] = None) -> None: + async def send_typing( + self, chat_id: str, metadata: Optional[Dict[str, Any]] = None + ) -> None: """Send typing indicator.""" if self._bot: try: @@ -1780,15 +2050,15 @@ async def send_typing(self, chat_id: str, metadata: Optional[Dict[str, Any]] = N e, exc_info=True, ) - + async def get_chat_info(self, chat_id: str) -> Dict[str, Any]: """Get information about a Telegram chat.""" if not self._bot: return {"name": "Unknown", "type": "dm"} - + try: chat = await self._bot.get_chat(int(chat_id)) - + chat_type = "dm" if chat.type == ChatType.GROUP: chat_type = "group" @@ -1798,7 +2068,7 @@ async def get_chat_info(self, chat_id: str) -> Dict[str, Any]: chat_type = "forum" elif chat.type == ChatType.CHANNEL: chat_type = "channel" - + return { "name": chat.title or chat.full_name or str(chat_id), "type": chat_type, @@ -1814,7 +2084,7 @@ async def get_chat_info(self, chat_id: str) -> Dict[str, Any]: exc_info=True, ) return {"name": str(chat_id), "type": "dm", "error": str(e)} - + def format_message(self, content: str) -> str: """ Convert standard markdown to Telegram MarkdownV2 format. @@ -1844,15 +2114,15 @@ def _ph(value: str) -> str: def _protect_fenced(m): raw = m.group(0) # Split off opening ``` (with optional language) and closing ``` - open_end = raw.index('\n') + 1 if '\n' in raw[3:] else 3 + open_end = raw.index("\n") + 1 if "\n" in raw[3:] else 3 opening = raw[:open_end] body_and_close = raw[open_end:] body = body_and_close[:-3] - body = body.replace('\\', '\\\\').replace('`', '\\`') - return _ph(opening + body + '```') + body = body.replace("\\", "\\\\").replace("`", "\\`") + return _ph(opening + body + "```") text = re.sub( - r'(```(?:[^\n]*\n)?[\s\S]*?```)', + r"(```(?:[^\n]*\n)?[\s\S]*?```)", _protect_fenced, text, ) @@ -1860,8 +2130,8 @@ def _protect_fenced(m): # 2) Protect inline code (`...`) # Escape \ inside inline code per MarkdownV2 spec. text = re.sub( - r'(`[^`]+`)', - lambda m: _ph(m.group(0).replace('\\', '\\\\')), + r"(`[^`]+`)", + lambda m: _ph(m.group(0).replace("\\", "\\\\")), text, ) @@ -1869,26 +2139,24 @@ def _protect_fenced(m): # only ')' and '\' need escaping per the MarkdownV2 spec. def _convert_link(m): display = _escape_mdv2(m.group(1)) - url = m.group(2).replace('\\', '\\\\').replace(')', '\\)') - return _ph(f'[{display}]({url})') + url = m.group(2).replace("\\", "\\\\").replace(")", "\\)") + return _ph(f"[{display}]({url})") - text = re.sub(r'\[([^\]]+)\]\(([^)]+)\)', _convert_link, text) + text = re.sub(r"\[([^\]]+)\]\(([^)]+)\)", _convert_link, text) # 4) Convert markdown headers (## Title) → bold *Title* def _convert_header(m): inner = m.group(1).strip() # Strip redundant bold markers that may appear inside a header - inner = re.sub(r'\*\*(.+?)\*\*', r'\1', inner) - return _ph(f'*{_escape_mdv2(inner)}*') + inner = re.sub(r"\*\*(.+?)\*\*", r"\1", inner) + return _ph(f"*{_escape_mdv2(inner)}*") - text = re.sub( - r'^#{1,6}\s+(.+)$', _convert_header, text, flags=re.MULTILINE - ) + text = re.sub(r"^#{1,6}\s+(.+)$", _convert_header, text, flags=re.MULTILINE) # 5) Convert bold: **text** → *text* (MarkdownV2 bold) text = re.sub( - r'\*\*(.+?)\*\*', - lambda m: _ph(f'*{_escape_mdv2(m.group(1))}*'), + r"\*\*(.+?)\*\*", + lambda m: _ph(f"*{_escape_mdv2(m.group(1))}*"), text, ) @@ -1896,22 +2164,22 @@ def _convert_header(m): # [^*\n]+ prevents matching across newlines (which would corrupt # bullet lists using * markers and multi-line content). text = re.sub( - r'\*([^*\n]+)\*', - lambda m: _ph(f'_{_escape_mdv2(m.group(1))}_'), + r"\*([^*\n]+)\*", + lambda m: _ph(f"_{_escape_mdv2(m.group(1))}_"), text, ) # 7) Convert strikethrough: ~~text~~ → ~text~ (MarkdownV2) text = re.sub( - r'~~(.+?)~~', - lambda m: _ph(f'~{_escape_mdv2(m.group(1))}~'), + r"~~(.+?)~~", + lambda m: _ph(f"~{_escape_mdv2(m.group(1))}~"), text, ) # 8) Convert spoiler: ||text|| → ||text|| (protect from | escaping) text = re.sub( - r'\|\|(.+?)\|\|', - lambda m: _ph(f'||{_escape_mdv2(m.group(1))}||'), + r"\|\|(.+?)\|\|", + lambda m: _ph(f"||{_escape_mdv2(m.group(1))}||"), text, ) @@ -1923,12 +2191,12 @@ def _convert_blockquote(m): content = m.group(2) # Check if content ends with || (expandable blockquote end marker) # In this case, preserve the trailing || unescaped for Telegram - if prefix.startswith('**') and content.endswith('||'): - return _ph(f'{prefix} {_escape_mdv2(content[:-2])}||') - return _ph(f'{prefix} {_escape_mdv2(content)}') + if prefix.startswith("**") and content.endswith("||"): + return _ph(f"{prefix} {_escape_mdv2(content[:-2])}||") + return _ph(f"{prefix} {_escape_mdv2(content)}") text = re.sub( - r'^((?:\*\*)?>{1,3}) (.+)$', + r"^((?:\*\*)?>{1,3}) (.+)$", _convert_blockquote, text, flags=re.MULTILINE, @@ -1945,7 +2213,7 @@ def _convert_blockquote(m): # 12) Safety net: escape unescaped ( ) { } that slipped through # placeholder processing. Split the text into code/non-code # segments so we never touch content inside ``` or ` spans. - _code_split = re.split(r'(```[\s\S]*?```|`[^`]+`)', text) + _code_split = re.split(r"(```[\s\S]*?```|`[^`]+`)", text) _safe_parts = [] for _idx, _seg in enumerate(_code_split): if _idx % 2 == 1: @@ -1957,32 +2225,33 @@ def _esc_bare(m, _seg=_seg): s = m.start() ch = m.group(0) # Already escaped - if s > 0 and _seg[s - 1] == '\\': + if s > 0 and _seg[s - 1] == "\\": return ch # ( that opens a MarkdownV2 link [text](url) - if ch == '(' and s > 0 and _seg[s - 1] == ']': + if ch == "(" and s > 0 and _seg[s - 1] == "]": return ch # ) that closes a link URL - if ch == ')': + if ch == ")": before = _seg[:s] - if '](http' in before or '](' in before: + if "](http" in before or "](" in before: # Check depth depth = 0 for j in range(s - 1, max(s - 2000, -1), -1): - if _seg[j] == '(': + if _seg[j] == "(": depth -= 1 if depth < 0: - if j > 0 and _seg[j - 1] == ']': + if j > 0 and _seg[j - 1] == "]": return ch break - elif _seg[j] == ')': + elif _seg[j] == ")": depth += 1 - return '\\' + ch - _safe_parts.append(re.sub(r'[(){}]', _esc_bare, _seg)) - text = ''.join(_safe_parts) + return "\\" + ch + + _safe_parts.append(re.sub(r"[(){}]", _esc_bare, _seg)) + text = "".join(_safe_parts) return text - + # ── Group mention gating ────────────────────────────────────────────── def _telegram_require_mention(self) -> bool: @@ -1992,7 +2261,12 @@ def _telegram_require_mention(self) -> bool: if isinstance(configured, str): return configured.lower() in ("true", "1", "yes", "on") return bool(configured) - return os.getenv("TELEGRAM_REQUIRE_MENTION", "false").lower() in ("true", "1", "yes", "on") + return os.getenv("TELEGRAM_REQUIRE_MENTION", "false").lower() in ( + "true", + "1", + "yes", + "on", + ) def _telegram_free_response_chats(self) -> set[str]: raw = self.config.extra.get("free_response_chats") @@ -2020,7 +2294,9 @@ def _telegram_ignored_threads(self) -> set[int]: try: ignored.add(int(text)) except (TypeError, ValueError): - logger.warning("[%s] Ignoring invalid Telegram thread id: %r", self.name, value) + logger.warning( + "[%s] Ignoring invalid Telegram thread id: %r", self.name, value + ) return ignored def _compile_mention_patterns(self) -> List[re.Pattern]: @@ -2034,7 +2310,9 @@ def _compile_mention_patterns(self) -> List[re.Pattern]: except Exception: loaded = [part.strip() for part in raw.splitlines() if part.strip()] if not loaded: - loaded = [part.strip() for part in raw.split(",") if part.strip()] + loaded = [ + part.strip() for part in raw.split(",") if part.strip() + ] patterns = loaded if patterns is None: @@ -2056,9 +2334,16 @@ def _compile_mention_patterns(self) -> List[re.Pattern]: try: compiled.append(re.compile(pattern, re.IGNORECASE)) except re.error as exc: - logger.warning("[%s] Invalid Telegram mention pattern %r: %s", self.name, pattern, exc) + logger.warning( + "[%s] Invalid Telegram mention pattern %r: %s", + self.name, + pattern, + exc, + ) if compiled: - logger.info("[%s] Loaded %d Telegram mention pattern(s)", self.name, len(compiled)) + logger.info( + "[%s] Loaded %d Telegram mention pattern(s)", self.name, len(compiled) + ) return compiled def _is_group_chat(self, message: Message) -> bool: @@ -2072,7 +2357,10 @@ def _is_reply_to_bot(self, message: Message) -> bool: if not self._bot or not getattr(message, "reply_to_message", None): return False reply_user = getattr(message.reply_to_message, "from_user", None) - return bool(reply_user and getattr(reply_user, "id", None) == getattr(self._bot, "id", None)) + return bool( + reply_user + and getattr(reply_user, "id", None) == getattr(self._bot, "id", None) + ) def _message_mentions_bot(self, message: Message) -> bool: if not self._bot: @@ -2082,8 +2370,14 @@ def _message_mentions_bot(self, message: Message) -> bool: bot_id = getattr(self._bot, "id", None) def _iter_sources(): - yield getattr(message, "text", None) or "", getattr(message, "entities", None) or [] - yield getattr(message, "caption", None) or "", getattr(message, "caption_entities", None) or [] + yield ( + getattr(message, "text", None) or "", + getattr(message, "entities", None) or [], + ) + yield ( + getattr(message, "caption", None) or "", + getattr(message, "caption_entities", None) or [], + ) for source_text, entities in _iter_sources(): if bot_username and f"@{bot_username}" in source_text.lower(): @@ -2095,7 +2389,10 @@ def _iter_sources(): length = int(getattr(entity, "length", 0)) if offset < 0 or length <= 0: continue - if source_text[offset:offset + length].strip().lower() == f"@{bot_username}": + if ( + source_text[offset : offset + length].strip().lower() + == f"@{bot_username}" + ): return True elif entity_type == "text_mention": user = getattr(entity, "user", None) @@ -2106,7 +2403,10 @@ def _iter_sources(): def _message_matches_mention_patterns(self, message: Message) -> bool: if not self._mention_patterns: return False - for candidate in (getattr(message, "text", None), getattr(message, "caption", None)): + for candidate in ( + getattr(message, "text", None), + getattr(message, "caption", None), + ): if not candidate: continue for pattern in self._mention_patterns: @@ -2121,7 +2421,9 @@ def _clean_bot_trigger_text(self, text: Optional[str]) -> Optional[str]: cleaned = re.sub(rf"(?i)@{username}\b[,:\-]*\s*", "", text).strip() return cleaned or text - def _should_process_message(self, message: Message, *, is_command: bool = False) -> bool: + def _should_process_message( + self, message: Message, *, is_command: bool = False + ) -> bool: """Apply Telegram group trigger rules. DMs remain unrestricted. Group/supergroup messages are accepted when: @@ -2140,8 +2442,15 @@ def _should_process_message(self, message: Message, *, is_command: bool = False) if int(thread_id) in self._telegram_ignored_threads(): return False except (TypeError, ValueError): - logger.warning("[%s] Ignoring non-numeric Telegram message_thread_id: %r", self.name, thread_id) - if str(getattr(getattr(message, "chat", None), "id", "")) in self._telegram_free_response_chats(): + logger.warning( + "[%s] Ignoring non-numeric Telegram message_thread_id: %r", + self.name, + thread_id, + ) + if ( + str(getattr(getattr(message, "chat", None), "id", "")) + in self._telegram_free_response_chats() + ): return True if not self._telegram_require_mention(): return True @@ -2153,7 +2462,9 @@ def _should_process_message(self, message: Message, *, is_command: bool = False) return True return self._message_matches_mention_patterns(message) - async def _handle_text_message(self, update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: + async def _handle_text_message( + self, update: Update, context: ContextTypes.DEFAULT_TYPE + ) -> None: """Handle incoming text messages. Telegram clients split long messages into multiple updates. Buffer @@ -2168,18 +2479,22 @@ async def _handle_text_message(self, update: Update, context: ContextTypes.DEFAU event = self._build_message_event(update.message, MessageType.TEXT) event.text = self._clean_bot_trigger_text(event.text) self._enqueue_text_event(event) - - async def _handle_command(self, update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: + + async def _handle_command( + self, update: Update, context: ContextTypes.DEFAULT_TYPE + ) -> None: """Handle incoming command messages.""" if not update.message or not update.message.text: return if not self._should_process_message(update.message, is_command=True): return - + event = self._build_message_event(update.message, MessageType.COMMAND) await self.handle_message(event) - - async def _handle_location_message(self, update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: + + async def _handle_location_message( + self, update: Update, context: ContextTypes.DEFAULT_TYPE + ) -> None: """Handle incoming location/venue pin messages.""" if not update.message: return @@ -2188,7 +2503,11 @@ async def _handle_location_message(self, update: Update, context: ContextTypes.D msg = update.message venue = getattr(msg, "venue", None) - location = getattr(venue, "location", None) if venue else getattr(msg, "location", None) + location = ( + getattr(venue, "location", None) + if venue + else getattr(msg, "location", None) + ) if not location: return @@ -2209,8 +2528,12 @@ async def _handle_location_message(self, update: Update, context: ContextTypes.D parts.append(f"Address: {address}") parts.append(f"latitude: {lat}") parts.append(f"longitude: {lon}") - parts.append(f"Map: https://www.google.com/maps/search/?api=1&query={lat},{lon}") - parts.append("Ask what they'd like to find nearby (restaurants, cafes, etc.) and any preferences.") + parts.append( + f"Map: https://www.google.com/maps/search/?api=1&query={lat},{lon}" + ) + parts.append( + "Ask what they'd like to find nearby (restaurants, cafes, etc.) and any preferences." + ) event = self._build_message_event(msg, MessageType.LOCATION) event.text = "\n".join(parts) @@ -2223,10 +2546,15 @@ async def _handle_location_message(self, update: Update, context: ContextTypes.D def _text_batch_key(self, event: MessageEvent) -> str: """Session-scoped key for text message batching.""" from gateway.session import build_session_key + return build_session_key( event.source, - group_sessions_per_user=self.config.extra.get("group_sessions_per_user", True), - thread_sessions_per_user=self.config.extra.get("thread_sessions_per_user", False), + group_sessions_per_user=self.config.extra.get( + "group_sessions_per_user", True + ), + thread_sessions_per_user=self.config.extra.get( + "thread_sessions_per_user", False + ), ) def _enqueue_text_event(self, event: MessageEvent) -> None: @@ -2246,7 +2574,9 @@ def _enqueue_text_event(self, event: MessageEvent) -> None: else: # Append text from the follow-up chunk if event.text: - existing.text = f"{existing.text}\n{event.text}" if existing.text else event.text + existing.text = ( + f"{existing.text}\n{event.text}" if existing.text else event.text + ) existing._last_chunk_len = chunk_len # type: ignore[attr-defined] # Merge any media that might be attached if event.media_urls: @@ -2276,14 +2606,24 @@ async def _flush_text_batch(self, key: str) -> None: if last_len >= self._SPLIT_THRESHOLD: delay = self._text_batch_split_delay_seconds else: - delay = self._text_batch_delay_seconds + # Fast path for common short messages: use a tighter quiet + # window while keeping enough headroom to still coalesce + # client-side split bursts and very quick follow-ups. + total_len = len((pending.text or "")) if pending else 0 + if total_len <= 320: + delay = min(self._text_batch_delay_seconds, 0.18) + elif total_len <= 1024: + delay = min(self._text_batch_delay_seconds, 0.24) + else: + delay = self._text_batch_delay_seconds await asyncio.sleep(delay) event = self._pending_text_batches.pop(key, None) if not event: return logger.info( "[Telegram] Flushing text batch %s (%d chars)", - key, len(event.text or ""), + key, + len(event.text or ""), ) await self.handle_message(event) finally: @@ -2297,10 +2637,15 @@ async def _flush_text_batch(self, key: str) -> None: def _photo_batch_key(self, event: MessageEvent, msg: Message) -> str: """Return a batching key for Telegram photos/albums.""" from gateway.session import build_session_key + session_key = build_session_key( event.source, - group_sessions_per_user=self.config.extra.get("group_sessions_per_user", True), - thread_sessions_per_user=self.config.extra.get("thread_sessions_per_user", False), + group_sessions_per_user=self.config.extra.get( + "group_sessions_per_user", True + ), + thread_sessions_per_user=self.config.extra.get( + "thread_sessions_per_user", False + ), ) media_group_id = getattr(msg, "media_group_id", None) if media_group_id: @@ -2311,11 +2656,21 @@ async def _flush_photo_batch(self, batch_key: str) -> None: """Send a buffered photo burst/album as a single MessageEvent.""" current_task = asyncio.current_task() try: - await asyncio.sleep(self._media_batch_delay_seconds) + # Albums need a longer window to gather all parts. Single-photo + # bursts can flush faster for lower perceived latency. + if ":album:" in batch_key: + delay = self._media_batch_delay_seconds + else: + delay = min(self._media_batch_delay_seconds, 0.35) + await asyncio.sleep(delay) event = self._pending_photo_batches.pop(batch_key, None) if not event: return - logger.info("[Telegram] Flushing photo batch %s with %d image(s)", batch_key, len(event.media_urls)) + logger.info( + "[Telegram] Flushing photo batch %s with %d image(s)", + batch_key, + len(event.media_urls), + ) await self.handle_message(event) finally: if self._pending_photo_batch_tasks.get(batch_key) is current_task: @@ -2336,17 +2691,21 @@ def _enqueue_photo_event(self, batch_key: str, event: MessageEvent) -> None: if prior_task and not prior_task.done(): prior_task.cancel() - self._pending_photo_batch_tasks[batch_key] = asyncio.create_task(self._flush_photo_batch(batch_key)) + self._pending_photo_batch_tasks[batch_key] = asyncio.create_task( + self._flush_photo_batch(batch_key) + ) - async def _handle_media_message(self, update: Update, context: ContextTypes.DEFAULT_TYPE) -> None: + async def _handle_media_message( + self, update: Update, context: ContextTypes.DEFAULT_TYPE + ) -> None: """Handle incoming media messages, downloading images to local cache.""" if not update.message: return if not self._should_process_message(update.message): return - + msg = update.message - + # Determine media type if msg.sticker: msg_type = MessageType.STICKER @@ -2362,19 +2721,19 @@ async def _handle_media_message(self, update: Update, context: ContextTypes.DEFA msg_type = MessageType.DOCUMENT else: msg_type = MessageType.DOCUMENT - + event = self._build_message_event(msg, msg_type) - + # Add caption as text if msg.caption: event.text = self._clean_bot_trigger_text(msg.caption) - + # Handle stickers: describe via vision tool with caching if msg.sticker: await self._handle_sticker(msg, event) await self.handle_message(event) return - + # Download photo to local image cache so the vision tool can access it # even after Telegram's ephemeral file URLs expire (~1 hour). if msg.photo: @@ -2394,7 +2753,7 @@ async def _handle_media_message(self, update: Update, context: ContextTypes.DEFA # Save to local cache (for vision tool access) cached_path = cache_image_from_bytes(bytes(image_bytes), ext=ext) event.media_urls = [cached_path] - event.media_types = [f"image/{ext.lstrip('.')}" ] + event.media_types = [f"image/{ext.lstrip('.')}"] logger.info("[Telegram] Cached user photo at %s", cached_path) media_group_id = getattr(msg, "media_group_id", None) if media_group_id: @@ -2452,7 +2811,9 @@ async def _handle_media_message(self, update: Update, context: ContextTypes.DEFA f"Unsupported document type '{ext or 'unknown'}'. " f"Supported types: {supported_list}" ) - logger.info("[Telegram] Unsupported document type: %s", ext or "unknown") + logger.info( + "[Telegram] Unsupported document type: %s", ext or "unknown" + ) await self.handle_message(event) return @@ -2463,7 +2824,9 @@ async def _handle_media_message(self, update: Update, context: ContextTypes.DEFA "The document is too large or its size could not be verified. " "Maximum: 20 MB." ) - logger.info("[Telegram] Document too large: %s bytes", doc.file_size) + logger.info( + "[Telegram] Document too large: %s bytes", doc.file_size + ) await self.handle_message(event) return @@ -2471,7 +2834,9 @@ async def _handle_media_message(self, update: Update, context: ContextTypes.DEFA file_obj = await doc.get_file() doc_bytes = await file_obj.download_as_bytearray() raw_bytes = bytes(doc_bytes) - cached_path = cache_document_from_bytes(raw_bytes, original_filename or f"document{ext}") + cached_path = cache_document_from_bytes( + raw_bytes, original_filename or f"document{ext}" + ) mime_type = SUPPORTED_DOCUMENT_TYPES[ext] event.media_urls = [cached_path] event.media_types = [mime_type] @@ -2483,7 +2848,7 @@ async def _handle_media_message(self, update: Update, context: ContextTypes.DEFA try: text_content = raw_bytes.decode("utf-8") display_name = original_filename or f"document{ext}" - display_name = re.sub(r'[^\w.\- ]', '_', display_name) + display_name = re.sub(r"[^\w.\- ]", "_", display_name) injection = f"[Content of {display_name}]:\n{text_content}" if event.text: event.text = f"{injection}\n\n{event.text}" @@ -2496,7 +2861,9 @@ async def _handle_media_message(self, update: Update, context: ContextTypes.DEFA ) except Exception as e: - logger.warning("[Telegram] Failed to cache document: %s", e, exc_info=True) + logger.warning( + "[Telegram] Failed to cache document: %s", e, exc_info=True + ) media_group_id = getattr(msg, "media_group_id", None) if media_group_id: @@ -2504,8 +2871,10 @@ async def _handle_media_message(self, update: Update, context: ContextTypes.DEFA return await self.handle_message(event) - - async def _queue_media_group_event(self, media_group_id: str, event: MessageEvent) -> None: + + async def _queue_media_group_event( + self, media_group_id: str, event: MessageEvent + ) -> None: """Buffer Telegram media-group items so albums arrive as one logical event. Telegram delivers albums as multiple updates with a shared media_group_id. @@ -2570,7 +2939,9 @@ async def _handle_sticker(self, msg: Message, event: "MessageEvent") -> None: cached = get_cached_description(sticker.file_unique_id) if cached: event.text = build_sticker_injection( - cached["description"], cached.get("emoji", emoji), cached.get("set_name", set_name) + cached["description"], + cached.get("emoji", emoji), + cached.get("set_name", set_name), ) logger.info("[Telegram] Sticker cache hit: %s", sticker.file_unique_id) return @@ -2593,19 +2964,23 @@ async def _handle_sticker(self, msg: Message, event: "MessageEvent") -> None: if result.get("success"): description = result.get("analysis", "a sticker") - cache_sticker_description(sticker.file_unique_id, description, emoji, set_name) + cache_sticker_description( + sticker.file_unique_id, description, emoji, set_name + ) event.text = build_sticker_injection(description, emoji, set_name) else: # Vision failed -- use emoji as fallback event.text = build_sticker_injection( f"a sticker with emoji {emoji}" if emoji else "a sticker", - emoji, set_name, + emoji, + set_name, ) except Exception as e: logger.warning("[Telegram] Sticker analysis error: %s", e, exc_info=True) event.text = build_sticker_injection( f"a sticker with emoji {emoji}" if emoji else "a sticker", - emoji, set_name, + emoji, + set_name, ) def _reload_dm_topics_from_config(self) -> None: @@ -2616,11 +2991,13 @@ def _reload_dm_topics_from_config(self) -> None: """ try: from hermes_constants import get_hermes_home + config_path = get_hermes_home() / "config.yaml" if not config_path.exists(): return import yaml as _yaml + with open(config_path, "r") as f: config = _yaml.safe_load(f) or {} @@ -2648,12 +3025,18 @@ def _reload_dm_topics_from_config(self) -> None: self._dm_topics[cache_key] = int(tid) logger.info( "[%s] Hot-loaded DM topic from config: %s -> thread_id=%s", - self.name, cache_key, tid, + self.name, + cache_key, + tid, ) except Exception as e: - logger.debug("[%s] Failed to reload dm_topics from config: %s", self.name, e) + logger.debug( + "[%s] Failed to reload dm_topics from config: %s", self.name, e + ) - def _get_dm_topic_info(self, chat_id: str, thread_id: Optional[str]) -> Optional[Dict[str, Any]]: + def _get_dm_topic_info( + self, chat_id: str, thread_id: Optional[str] + ) -> Optional[Dict[str, Any]]: """Look up DM topic config by chat_id and thread_id. Returns the topic config dict (name, skill, etc.) if this thread_id @@ -2692,21 +3075,27 @@ def _get_dm_topic_info(self, chat_id: str, thread_id: Optional[str]) -> Optional return None - def _cache_dm_topic_from_message(self, chat_id: str, thread_id: str, topic_name: str) -> None: + def _cache_dm_topic_from_message( + self, chat_id: str, thread_id: str, topic_name: str + ) -> None: """Cache a thread_id -> topic_name mapping discovered from an incoming message.""" cache_key = f"{chat_id}:{topic_name}" if cache_key not in self._dm_topics: self._dm_topics[cache_key] = int(thread_id) logger.info( "[%s] Cached DM topic from message: %s -> thread_id=%s", - self.name, cache_key, thread_id, + self.name, + cache_key, + thread_id, ) - def _build_message_event(self, message: Message, msg_type: MessageType) -> MessageEvent: + def _build_message_event( + self, message: Message, msg_type: MessageType + ) -> MessageEvent: """Build a MessageEvent from a Telegram message.""" chat = message.chat user = message.from_user - + # Determine chat type chat_type = "dm" if chat.type in (ChatType.GROUP, ChatType.SUPERGROUP): @@ -2730,7 +3119,9 @@ def _build_message_event(self, message: Message, msg_type: MessageType) -> Messa if hasattr(message, "forum_topic_created") and message.forum_topic_created: created_name = message.forum_topic_created.name if created_name: - self._cache_dm_topic_from_message(str(chat.id), thread_id_str, created_name) + self._cache_dm_topic_from_message( + str(chat.id), thread_id_str, created_name + ) if not chat_topic: chat_topic = created_name @@ -2750,20 +3141,25 @@ def _build_message_event(self, message: Message, msg_type: MessageType) -> Messa # Build source source = self.build_source( chat_id=str(chat.id), - chat_name=chat.title or (chat.full_name if hasattr(chat, "full_name") else None), + chat_name=chat.title + or (chat.full_name if hasattr(chat, "full_name") else None), chat_type=chat_type, user_id=str(user.id) if user else None, user_name=user.full_name if user else None, thread_id=thread_id_str, chat_topic=chat_topic, ) - + # Extract reply context if this message is a reply reply_to_id = None reply_to_text = None if message.reply_to_message: reply_to_id = str(message.reply_to_message.message_id) - reply_to_text = message.reply_to_message.text or message.reply_to_message.caption or None + reply_to_text = ( + message.reply_to_message.text + or message.reply_to_message.caption + or None + ) return MessageEvent( text=message.text or "", @@ -2781,7 +3177,11 @@ def _build_message_event(self, message: Message, msg_type: MessageType) -> Messa def _reactions_enabled(self) -> bool: """Check if message reactions are enabled via config/env.""" - return os.getenv("TELEGRAM_REACTIONS", "false").lower() not in ("false", "0", "no") + return os.getenv("TELEGRAM_REACTIONS", "false").lower() not in ( + "false", + "0", + "no", + ) async def _set_reaction(self, chat_id: str, message_id: str, emoji: str) -> bool: """Set a single emoji reaction on a Telegram message.""" @@ -2795,7 +3195,9 @@ async def _set_reaction(self, chat_id: str, message_id: str, emoji: str) -> bool ) return True except Exception as e: - logger.debug("[%s] set_message_reaction failed (%s): %s", self.name, emoji, e) + logger.debug( + "[%s] set_message_reaction failed (%s): %s", self.name, emoji, e + ) return False async def on_processing_start(self, event: MessageEvent) -> None: @@ -2807,7 +3209,9 @@ async def on_processing_start(self, event: MessageEvent) -> None: if chat_id and message_id: await self._set_reaction(chat_id, message_id, "\U0001f440") - async def on_processing_complete(self, event: MessageEvent, outcome: ProcessingOutcome) -> None: + async def on_processing_complete( + self, event: MessageEvent, outcome: ProcessingOutcome + ) -> None: """Swap the in-progress reaction for a final success/failure reaction. Unlike Discord (additive reactions), Telegram's set_message_reaction diff --git a/gateway/run.py b/gateway/run.py index 2eb745f92bd1a..07f2b3a425eb0 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -8,7 +8,7 @@ Usage: # Start the gateway python -m gateway.run - + # Or from CLI python cli.py --gateway """ @@ -28,6 +28,7 @@ from datetime import datetime from typing import Dict, Optional, Any, List + # --------------------------------------------------------------------------- # SSL certificate auto-detection for NixOS and other non-standard systems. # Must run BEFORE any HTTP library (discord, aiohttp, etc.) is imported. @@ -49,6 +50,7 @@ def _ensure_ssl_certs() -> None: # 2. certifi (ships its own Mozilla bundle) try: import certifi + os.environ["SSL_CERT_FILE"] = certifi.where() return except ImportError: @@ -56,19 +58,20 @@ def _ensure_ssl_certs() -> None: # 3. Common distro / macOS locations for candidate in ( - "/etc/ssl/certs/ca-certificates.crt", # Debian/Ubuntu/Gentoo - "/etc/pki/tls/certs/ca-bundle.crt", # RHEL/CentOS 7 - "/etc/pki/ca-trust/extracted/pem/tls-ca-bundle.pem", # RHEL/CentOS 8+ - "/etc/ssl/ca-bundle.pem", # SUSE/OpenSUSE - "/etc/ssl/cert.pem", # Alpine / macOS - "/etc/pki/tls/cert.pem", # Fedora - "/usr/local/etc/openssl@1.1/cert.pem", # macOS Homebrew Intel - "/opt/homebrew/etc/openssl@1.1/cert.pem", # macOS Homebrew ARM + "/etc/ssl/certs/ca-certificates.crt", # Debian/Ubuntu/Gentoo + "/etc/pki/tls/certs/ca-bundle.crt", # RHEL/CentOS 7 + "/etc/pki/ca-trust/extracted/pem/tls-ca-bundle.pem", # RHEL/CentOS 8+ + "/etc/ssl/ca-bundle.pem", # SUSE/OpenSUSE + "/etc/ssl/cert.pem", # Alpine / macOS + "/etc/pki/tls/cert.pem", # Fedora + "/usr/local/etc/openssl@1.1/cert.pem", # macOS Homebrew Intel + "/opt/homebrew/etc/openssl@1.1/cert.pem", # macOS Homebrew ARM ): if os.path.exists(candidate): os.environ["SSL_CERT_FILE"] = candidate return + _ensure_ssl_certs() # Add parent directory to path @@ -77,25 +80,31 @@ def _ensure_ssl_certs() -> None: # Resolve Hermes home directory (respects HERMES_HOME override) from hermes_constants import get_hermes_home from utils import atomic_yaml_write, is_truthy_value + _hermes_home = get_hermes_home() # Load environment variables from ~/.hermes/.env first. # User-managed env files should override stale shell exports on restart. from dotenv import load_dotenv # backward-compat for tests that monkeypatch this symbol from hermes_cli.env_loader import load_hermes_dotenv -_env_path = _hermes_home / '.env' -load_hermes_dotenv(hermes_home=_hermes_home, project_env=Path(__file__).resolve().parents[1] / '.env') + +_env_path = _hermes_home / ".env" +load_hermes_dotenv( + hermes_home=_hermes_home, project_env=Path(__file__).resolve().parents[1] / ".env" +) # Bridge config.yaml values into the environment so os.getenv() picks them up. # config.yaml is authoritative for terminal settings — overrides .env. -_config_path = _hermes_home / 'config.yaml' +_config_path = _hermes_home / "config.yaml" if _config_path.exists(): try: import yaml as _yaml + with open(_config_path, encoding="utf-8") as _f: _cfg = _yaml.safe_load(_f) or {} # Expand ${ENV_VAR} references before bridging to env vars. from hermes_cli.config import _expand_env_vars + _cfg = _expand_env_vars(_cfg) # Top-level simple values (fallback only — don't override .env) for _key, _val in _cfg.items(): @@ -182,18 +191,41 @@ def _ensure_ssl_certs() -> None: os.environ["HERMES_MAX_ITERATIONS"] = str(_agent_cfg["max_turns"]) # Bridge agent.gateway_timeout → HERMES_AGENT_TIMEOUT env var. # Env var from .env takes precedence (already in os.environ). - if "gateway_timeout" in _agent_cfg and "HERMES_AGENT_TIMEOUT" not in os.environ: + if ( + "gateway_timeout" in _agent_cfg + and "HERMES_AGENT_TIMEOUT" not in os.environ + ): os.environ["HERMES_AGENT_TIMEOUT"] = str(_agent_cfg["gateway_timeout"]) - if "gateway_timeout_warning" in _agent_cfg and "HERMES_AGENT_TIMEOUT_WARNING" not in os.environ: - os.environ["HERMES_AGENT_TIMEOUT_WARNING"] = str(_agent_cfg["gateway_timeout_warning"]) - if "gateway_notify_interval" in _agent_cfg and "HERMES_AGENT_NOTIFY_INTERVAL" not in os.environ: - os.environ["HERMES_AGENT_NOTIFY_INTERVAL"] = str(_agent_cfg["gateway_notify_interval"]) - if "restart_drain_timeout" in _agent_cfg and "HERMES_RESTART_DRAIN_TIMEOUT" not in os.environ: - os.environ["HERMES_RESTART_DRAIN_TIMEOUT"] = str(_agent_cfg["restart_drain_timeout"]) + if ( + "gateway_timeout_warning" in _agent_cfg + and "HERMES_AGENT_TIMEOUT_WARNING" not in os.environ + ): + os.environ["HERMES_AGENT_TIMEOUT_WARNING"] = str( + _agent_cfg["gateway_timeout_warning"] + ) + if ( + "gateway_notify_interval" in _agent_cfg + and "HERMES_AGENT_NOTIFY_INTERVAL" not in os.environ + ): + os.environ["HERMES_AGENT_NOTIFY_INTERVAL"] = str( + _agent_cfg["gateway_notify_interval"] + ) + if ( + "restart_drain_timeout" in _agent_cfg + and "HERMES_RESTART_DRAIN_TIMEOUT" not in os.environ + ): + os.environ["HERMES_RESTART_DRAIN_TIMEOUT"] = str( + _agent_cfg["restart_drain_timeout"] + ) _display_cfg = _cfg.get("display", {}) if _display_cfg and isinstance(_display_cfg, dict): - if "busy_input_mode" in _display_cfg and "HERMES_GATEWAY_BUSY_INPUT_MODE" not in os.environ: - os.environ["HERMES_GATEWAY_BUSY_INPUT_MODE"] = str(_display_cfg["busy_input_mode"]) + if ( + "busy_input_mode" in _display_cfg + and "HERMES_GATEWAY_BUSY_INPUT_MODE" not in os.environ + ): + os.environ["HERMES_GATEWAY_BUSY_INPUT_MODE"] = str( + _display_cfg["busy_input_mode"] + ) # Timezone: bridge config.yaml → HERMES_TIMEZONE env var. # HERMES_TIMEZONE from .env takes precedence (already in os.environ). _tz_cfg = _cfg.get("timezone", "") @@ -211,7 +243,8 @@ def _ensure_ssl_certs() -> None: # Apply IPv4 preference if configured (before any HTTP clients are created). try: from hermes_constants import apply_ipv4_preference - _network_cfg = (_cfg if '_cfg' in dir() else {}).get("network", {}) + + _network_cfg = (_cfg if "_cfg" in dir() else {}).get("network", {}) if isinstance(_network_cfg, dict) and _network_cfg.get("force_ipv4"): apply_ipv4_preference(force=True) except Exception: @@ -220,6 +253,7 @@ def _ensure_ssl_certs() -> None: # Validate config structure early — log warnings so gateway operators see problems try: from hermes_cli.config import print_config_warnings + print_config_warnings() except Exception: pass @@ -268,11 +302,7 @@ def _ensure_ssl_certs() -> None: def _normalize_whatsapp_identifier(value: str) -> str: """Strip WhatsApp JID/LID syntax down to its stable numeric identifier.""" return ( - str(value or "") - .strip() - .replace("+", "", 1) - .split(":", 1)[0] - .split("@", 1)[0] + str(value or "").strip().replace("+", "", 1).split(":", 1)[0].split("@", 1)[0] ) @@ -307,6 +337,7 @@ def _expand_whatsapp_auth_aliases(identifier: str) -> set: return resolved + logger = logging.getLogger(__name__) # Sentinel placed into _running_agents immediately when a session starts @@ -354,7 +385,10 @@ def _build_media_placeholder(event) -> str: media_types = getattr(event, "media_types", None) or [] for i, url in enumerate(media_urls): mtype = media_types[i] if i < len(media_types) else "" - if mtype.startswith("image/") or getattr(event, "message_type", None) == MessageType.PHOTO: + if ( + mtype.startswith("image/") + or getattr(event, "message_type", None) == MessageType.PHOTO + ): parts.append(f"[User sent an image: {url}]") elif mtype.startswith("audio/"): parts.append(f"[User sent audio: {url}]") @@ -384,6 +418,7 @@ def _check_unavailable_skill(command_name: str) -> str | None: try: from tools.skills_tool import _get_disabled_skill_names from agent.skill_utils import get_all_skills_dirs + disabled = _get_disabled_skill_names() # Check disabled skills across all dirs (local + external) @@ -391,7 +426,7 @@ def _check_unavailable_skill(command_name: str) -> str | None: if not skills_dir.exists(): continue for skill_md in skills_dir.rglob("SKILL.md"): - if any(part in ('.git', '.github', '.hub') for part in skill_md.parts): + if any(part in (".git", ".github", ".hub") for part in skill_md.parts): continue name = skill_md.parent.name.lower().replace("_", "-") if name == normalized and name in disabled: @@ -402,6 +437,7 @@ def _check_unavailable_skill(command_name: str) -> str | None: # Check optional skills (shipped with repo but not installed) from hermes_constants import get_optional_skills_dir + repo_root = Path(__file__).resolve().parent.parent optional_dir = get_optional_skills_dir(repo_root / "optional-skills") if optional_dir.exists(): @@ -429,16 +465,40 @@ def _platform_config_key(platform: "Platform") -> str: def _load_gateway_config() -> dict: """Load and parse ~/.hermes/config.yaml, returning {} on any error.""" try: - config_path = _hermes_home / 'config.yaml' + config_path = _hermes_home / "config.yaml" if config_path.exists(): import yaml - with open(config_path, 'r', encoding='utf-8') as f: + + with open(config_path, "r", encoding="utf-8") as f: return yaml.safe_load(f) or {} except Exception: - logger.debug("Could not load gateway config from %s", _hermes_home / 'config.yaml') + logger.debug( + "Could not load gateway config from %s", _hermes_home / "config.yaml" + ) return {} +def _is_gateway_streaming_enabled(streaming_config: Any, platform_setting: Any) -> bool: + """Resolve whether token streaming should run for a platform. + + ``platform_setting`` comes from ``display.platforms.<platform>.streaming`` or + built-in platform defaults. ``True`` force-enables streaming for that platform + while still respecting a global ``transport: off`` kill switch. ``False`` + force-disables it. ``None`` follows the top-level streaming config. + """ + transport = getattr(streaming_config, "transport", "edit") + if transport == "off": + return False + global_enabled = is_truthy_value(getattr(streaming_config, "enabled", False)) + if platform_setting is None: + return global_enabled + if platform_setting is False: + return False + if platform_setting is True: + return True + return global_enabled + + def _resolve_gateway_model(config: dict | None = None) -> str: """Read model from config.yaml — single source of truth. @@ -497,7 +557,7 @@ def _format_gateway_process_notification(evt: dict) -> "str | None": _sup = evt.get("suppressed", 0) text = ( f"[SYSTEM: Background process {_sid} matched " - f"watch pattern \"{_pat}\".\n" + f'watch pattern "{_pat}".\n' f"Command: {_cmd}\n" f"Matched output:\n{_out}" ) @@ -530,7 +590,7 @@ class GatewayRunner: _restart_via_service: bool = False _stop_task: Optional[asyncio.Task] = None _session_model_overrides: Dict[str, Dict[str, str]] = {} - + def __init__(self, config: Optional[GatewayConfig] = None): self.config = config or load_gateway_config() self.adapters: Dict[Platform, BasePlatformAdapter] = {} @@ -550,9 +610,13 @@ def __init__(self, config: Optional[GatewayConfig] = None): # Wire process registry into session store for reset protection from tools.process_registry import process_registry + self.session_store = SessionStore( - self.config.sessions_dir, self.config, - has_active_processes_fn=lambda key: process_registry.has_active_for_session(key), + self.config.sessions_dir, + self.config, + has_active_processes_fn=lambda key: process_registry.has_active_for_session( + key + ), ) self.delivery_router = DeliveryRouter(self.config) self._running = False @@ -567,13 +631,15 @@ def __init__(self, config: Optional[GatewayConfig] = None): self._restart_detached = False self._restart_via_service = False self._stop_task: Optional[asyncio.Task] = None - + # Track running agents per session for interrupt support # Key: session_key, Value: AIAgent instance self._running_agents: Dict[str, Any] = {} self._running_agents_ts: Dict[str, float] = {} # start timestamp per session self._pending_messages: Dict[str, str] = {} # Queued messages during interrupt - self._busy_ack_ts: Dict[str, float] = {} # last busy-ack timestamp per session (debounce) + self._busy_ack_ts: Dict[ + str, float + ] = {} # last busy-ack timestamp per session (debounce) # Cache AIAgent instances per session to preserve prompt caching. # Without this, a new AIAgent is created per message, rebuilding the @@ -581,6 +647,7 @@ def __init__(self, config: Optional[GatewayConfig] = None): # and costing ~10x more on providers with prompt caching (Anthropic). # Key: session_key, Value: (AIAgent, config_signature_str) import threading as _threading + self._agent_cache: Dict[str, tuple] = {} self._agent_cache_lock = _threading.Lock() @@ -603,29 +670,31 @@ def __init__(self, config: Optional[GatewayConfig] = None): # This preserves write_frequency="session" semantics across short-lived # per-message AIAgent instances. - - # Ensure tirith security scanner is available (downloads if needed) try: from tools.tirith_security import ensure_installed + ensure_installed(log_failures=False) except Exception: pass # Non-fatal — fail-open at scan time if unavailable - + # Initialize session database for session_search tool support self._session_db = None try: from hermes_state import SessionDB + self._session_db = SessionDB() except Exception as e: logger.debug("SQLite session store not available: %s", e) - + # DM pairing store for code-based user authorization from gateway.pairing import PairingStore + self.pairing_store = PairingStore() - + # Event hook system from gateway.hooks import HookRegistry + self.hooks = HookRegistry() # Per-chat voice reply mode: "off" | "voice_only" | "all" @@ -634,15 +703,13 @@ def __init__(self, config: Optional[GatewayConfig] = None): # Track background tasks to prevent garbage collection mid-execution self._background_tasks: set = set() - - - # -- Setup skill availability ---------------------------------------- def _has_setup_skill(self) -> bool: """Check if the hermes-agent-setup skill is installed.""" try: from tools.skill_manager_tool import _find_skill + return _find_skill("hermes-agent-setup") is not None except Exception: return False @@ -662,21 +729,19 @@ def _load_voice_modes(self) -> Dict[str, str]: valid_modes = {"off", "voice_only", "all"} return { - str(chat_id): mode - for chat_id, mode in data.items() - if mode in valid_modes + str(chat_id): mode for chat_id, mode in data.items() if mode in valid_modes } def _save_voice_modes(self) -> None: try: self._VOICE_MODE_PATH.parent.mkdir(parents=True, exist_ok=True) - self._VOICE_MODE_PATH.write_text( - json.dumps(self._voice_mode, indent=2) - ) + self._VOICE_MODE_PATH.write_text(json.dumps(self._voice_mode, indent=2)) except OSError as e: logger.warning("Failed to save voice modes: %s", e) - def _set_adapter_auto_tts_disabled(self, adapter, chat_id: str, disabled: bool) -> None: + def _set_adapter_auto_tts_disabled( + self, adapter, chat_id: str, disabled: bool + ) -> None: """Update an adapter's in-memory auto-TTS suppression set if present.""" disabled_chats = getattr(adapter, "_auto_tts_disabled_chats", None) if not isinstance(disabled_chats, set): @@ -720,6 +785,7 @@ def _flush_memories_for_session( return from run_agent import AIAgent + model, runtime_kwargs = self._resolve_session_agent_runtime( session_key=session_key, ) @@ -752,6 +818,7 @@ def _flush_memories_for_session( _current_memory = "" try: from tools.memory_tool import get_memory_dir + _mem_dir = get_memory_dir() for fname, label in [ ("MEMORY.md", "MEMORY (your personal notes)"), @@ -798,9 +865,13 @@ def _flush_memories_for_session( user_message=flush_prompt, conversation_history=msgs, ) - logger.info("Pre-reset memory flush completed for session %s", old_session_id) + logger.info( + "Pre-reset memory flush completed for session %s", old_session_id + ) except Exception as e: - logger.debug("Pre-reset memory flush failed for session %s: %s", old_session_id, e) + logger.debug( + "Pre-reset memory flush failed for session %s: %s", old_session_id, e + ) async def _async_flush_memories( self, @@ -869,7 +940,11 @@ def _resolve_session_agent_runtime( resolved_session_key = None model = _resolve_gateway_model(user_config) - override = self._session_model_overrides.get(resolved_session_key) if resolved_session_key else None + override = ( + self._session_model_overrides.get(resolved_session_key) + if resolved_session_key + else None + ) if override: override_model = override.get("model", model) override_runtime = { @@ -881,7 +956,9 @@ def _resolve_session_agent_runtime( if override_runtime.get("api_key"): logger.debug( "Session model override (fast): session=%s config_model=%s -> override_model=%s provider=%s", - (resolved_session_key or "")[:30], model, override_model, + (resolved_session_key or "")[:30], + model, + override_model, override_runtime.get("provider"), ) return override_model, override_runtime @@ -889,13 +966,18 @@ def _resolve_session_agent_runtime( # resolution and apply model/provider from the override on top. logger.debug( "Session model override (no api_key, fallback): session=%s config_model=%s override_model=%s", - (resolved_session_key or "")[:30], model, override_model, + (resolved_session_key or "")[:30], + model, + override_model, ) else: logger.debug( "No session model override: session=%s config_model=%s override_keys=%s", - (resolved_session_key or "")[:30], model, - list(self._session_model_overrides.keys())[:5] if self._session_model_overrides else "[]", + (resolved_session_key or "")[:30], + model, + list(self._session_model_overrides.keys())[:5] + if self._session_model_overrides + else "[]", ) runtime_kwargs = _resolve_runtime_agent_kwargs() @@ -911,18 +993,22 @@ def _resolve_session_agent_runtime( if not model and runtime_kwargs.get("provider"): try: from hermes_cli.models import get_default_model_for_provider + model = get_default_model_for_provider(runtime_kwargs["provider"]) if model: logger.info( "No model configured — defaulting to %s for provider %s", - model, runtime_kwargs["provider"], + model, + runtime_kwargs["provider"], ) except Exception: pass return model, runtime_kwargs - def _resolve_turn_agent_config(self, user_message: str, model: str, runtime_kwargs: dict) -> dict: + def _resolve_turn_agent_config( + self, user_message: str, model: str, runtime_kwargs: dict + ) -> dict: from agent.smart_model_routing import resolve_turn_route from hermes_cli.models import resolve_fast_mode_overrides @@ -936,7 +1022,9 @@ def _resolve_turn_agent_config(self, user_message: str, model: str, runtime_kwar "args": list(runtime_kwargs.get("args") or []), "credential_pool": runtime_kwargs.get("credential_pool"), } - route = resolve_turn_route(user_message, getattr(self, "_smart_model_routing", {}), primary) + route = resolve_turn_route( + user_message, getattr(self, "_smart_model_routing", {}), primary + ) service_tier = getattr(self, "_service_tier", None) if not service_tier: @@ -992,19 +1080,28 @@ async def _handle_adapter_fatal_error(self, adapter: BasePlatformAdapter) -> Non ) if not self.adapters and not self._failed_platforms: - self._exit_reason = adapter.fatal_error_message or "All messaging adapters disconnected" + self._exit_reason = ( + adapter.fatal_error_message or "All messaging adapters disconnected" + ) if adapter.fatal_error_retryable: self._exit_with_failure = True - logger.error("No connected messaging platforms remain. Shutting down gateway for service restart.") + logger.error( + "No connected messaging platforms remain. Shutting down gateway for service restart." + ) else: - logger.error("No connected messaging platforms remain. Shutting down gateway cleanly.") + logger.error( + "No connected messaging platforms remain. Shutting down gateway cleanly." + ) await self.stop() elif not self.adapters and self._failed_platforms: # All platforms are down and queued for background reconnection. # If the error is retryable, exit with failure so systemd Restart=on-failure # can restart the process. Otherwise stay alive and keep retrying in background. if adapter.fatal_error_retryable: - self._exit_reason = adapter.fatal_error_message or "All messaging platforms failed with retryable errors" + self._exit_reason = ( + adapter.fatal_error_message + or "All messaging platforms failed with retryable errors" + ) self._exit_with_failure = True logger.error( "All messaging platforms failed with retryable errors. " @@ -1034,9 +1131,12 @@ def _status_action_gerund(self) -> str: def _queue_during_drain_enabled(self) -> bool: return self._restart_requested and self._busy_input_mode == "queue" - def _update_runtime_status(self, gateway_state: Optional[str] = None, exit_reason: Optional[str] = None) -> None: + def _update_runtime_status( + self, gateway_state: Optional[str] = None, exit_reason: Optional[str] = None + ) -> None: try: from gateway.status import write_runtime_status + write_runtime_status( gateway_state=gateway_state, exit_reason=exit_reason, @@ -1056,6 +1156,7 @@ def _update_platform_runtime_status( ) -> None: try: from gateway.status import write_runtime_status + write_runtime_status( platform=platform, platform_state=platform_state, @@ -1064,20 +1165,22 @@ def _update_platform_runtime_status( ) except Exception: pass - + @staticmethod def _load_prefill_messages() -> List[Dict[str, Any]]: """Load ephemeral prefill messages from config or env var. - + Checks HERMES_PREFILL_MESSAGES_FILE env var first, then falls back to the prefill_messages_file key in ~/.hermes/config.yaml. Relative paths are resolved from ~/.hermes/. """ import json as _json + file_path = os.getenv("HERMES_PREFILL_MESSAGES_FILE", "") if not file_path: try: import yaml as _y + cfg_path = _hermes_home / "config.yaml" if cfg_path.exists(): with open(cfg_path, encoding="utf-8") as _f: @@ -1097,7 +1200,9 @@ def _load_prefill_messages() -> List[Dict[str, Any]]: with open(path, "r", encoding="utf-8") as f: data = _json.load(f) if not isinstance(data, list): - logger.warning("Prefill messages file must contain a JSON array: %s", path) + logger.warning( + "Prefill messages file must contain a JSON array: %s", path + ) return [] return data except Exception as e: @@ -1107,7 +1212,7 @@ def _load_prefill_messages() -> List[Dict[str, Any]]: @staticmethod def _load_ephemeral_system_prompt() -> str: """Load ephemeral system prompt from config or env var. - + Checks HERMES_EPHEMERAL_SYSTEM_PROMPT env var first, then falls back to agent.system_prompt in ~/.hermes/config.yaml. """ @@ -1116,6 +1221,7 @@ def _load_ephemeral_system_prompt() -> str: return prompt try: import yaml as _y + cfg_path = _hermes_home / "config.yaml" if cfg_path.exists(): with open(cfg_path, encoding="utf-8") as _f: @@ -1134,19 +1240,25 @@ def _load_reasoning_config() -> dict | None: default (medium). """ from hermes_constants import parse_reasoning_effort + effort = "" try: import yaml as _y + cfg_path = _hermes_home / "config.yaml" if cfg_path.exists(): with open(cfg_path, encoding="utf-8") as _f: cfg = _y.safe_load(_f) or {} - effort = str(cfg.get("agent", {}).get("reasoning_effort", "") or "").strip() + effort = str( + cfg.get("agent", {}).get("reasoning_effort", "") or "" + ).strip() except Exception: pass result = parse_reasoning_effort(effort) if effort and effort.strip() and result is None: - logger.warning("Unknown reasoning_effort '%s', using default (medium)", effort) + logger.warning( + "Unknown reasoning_effort '%s', using default (medium)", effort + ) return result @staticmethod @@ -1160,6 +1272,7 @@ def _load_service_tier() -> str | None: raw = "" try: import yaml as _y + cfg_path = _hermes_home / "config.yaml" if cfg_path.exists(): with open(cfg_path, encoding="utf-8") as _f: @@ -1181,6 +1294,7 @@ def _load_show_reasoning() -> bool: """Load show_reasoning toggle from config.yaml display section.""" try: import yaml as _y + cfg_path = _hermes_home / "config.yaml" if cfg_path.exists(): with open(cfg_path, encoding="utf-8") as _f: @@ -1197,11 +1311,16 @@ def _load_busy_input_mode() -> str: if not mode: try: import yaml as _y + cfg_path = _hermes_home / "config.yaml" if cfg_path.exists(): with open(cfg_path, encoding="utf-8") as _f: cfg = _y.safe_load(_f) or {} - mode = str(cfg.get("display", {}).get("busy_input_mode", "") or "").strip().lower() + mode = ( + str(cfg.get("display", {}).get("busy_input_mode", "") or "") + .strip() + .lower() + ) except Exception: pass return "queue" if mode == "queue" else "interrupt" @@ -1213,11 +1332,14 @@ def _load_restart_drain_timeout() -> float: if not raw: try: import yaml as _y + cfg_path = _hermes_home / "config.yaml" if cfg_path.exists(): with open(cfg_path, encoding="utf-8") as _f: cfg = _y.safe_load(_f) or {} - raw = str(cfg.get("agent", {}).get("restart_drain_timeout", "") or "").strip() + raw = str( + cfg.get("agent", {}).get("restart_drain_timeout", "") or "" + ).strip() except Exception: pass value = parse_restart_drain_timeout(raw) @@ -1246,6 +1368,7 @@ def _load_background_notifications_mode() -> str: if not mode: try: import yaml as _y + cfg_path = _hermes_home / "config.yaml" if cfg_path.exists(): with open(cfg_path, encoding="utf-8") as _f: @@ -1272,6 +1395,7 @@ def _load_provider_routing() -> dict: """Load OpenRouter provider routing preferences from config.yaml.""" try: import yaml as _y + cfg_path = _hermes_home / "config.yaml" if cfg_path.exists(): with open(cfg_path, encoding="utf-8") as _f: @@ -1291,6 +1415,7 @@ def _load_fallback_model() -> list | dict | None: """ try: import yaml as _y + cfg_path = _hermes_home / "config.yaml" if cfg_path.exists(): with open(cfg_path, encoding="utf-8") as _f: @@ -1307,6 +1432,7 @@ def _load_smart_model_routing() -> dict: """Load optional smart cheap-vs-strong model routing config.""" try: import yaml as _y + cfg_path = _hermes_home / "config.yaml" if cfg_path.exists(): with open(cfg_path, encoding="utf-8") as _f: @@ -1323,20 +1449,28 @@ def _snapshot_running_agents(self) -> Dict[str, Any]: if agent is not _AGENT_PENDING_SENTINEL } - def _queue_or_replace_pending_event(self, session_key: str, event: MessageEvent) -> None: + def _queue_or_replace_pending_event( + self, session_key: str, event: MessageEvent + ) -> None: adapter = self.adapters.get(event.source.platform) if not adapter: return merge_pending_message_event(adapter._pending_messages, session_key, event) - async def _handle_active_session_busy_message(self, event: MessageEvent, session_key: str) -> bool: + async def _handle_active_session_busy_message( + self, event: MessageEvent, session_key: str + ) -> bool: # --- Draining case (gateway restarting/stopping) --- if self._draining: adapter = self.adapters.get(event.source.platform) if not adapter: return True - thread_meta = {"thread_id": event.source.thread_id} if event.source.thread_id else None + thread_meta = ( + {"thread_id": event.source.thread_id} + if event.source.thread_id + else None + ) if self._queue_during_drain_enabled(): self._queue_or_replace_pending_event(session_key, event) message = f"⏳ Gateway {self._status_action_gerund()} — queued for the next turn after it comes back." @@ -1366,6 +1500,7 @@ async def _handle_active_session_busy_message(self, event: MessageEvent, session # Store the message so it's processed as the next turn after the # interrupt causes the current run to exit. from gateway.platforms.base import merge_pending_message_event + merge_pending_message_event(adapter._pending_messages, session_key, event) # Interrupt the running agent — this aborts in-flight tool calls and @@ -1413,7 +1548,9 @@ async def _handle_active_session_busy_message(self, event: MessageEvent, session f"I'll respond to your message shortly." ) - thread_meta = {"thread_id": event.source.thread_id} if event.source.thread_id else None + thread_meta = ( + {"thread_id": event.source.thread_id} if event.source.thread_id else None + ) try: await adapter._send_with_retry( chat_id=event.source.chat_id, @@ -1435,7 +1572,11 @@ def _maybe_update_status(force: bool = False) -> None: nonlocal last_active_count, last_status_at now = asyncio.get_running_loop().time() active_count = self._running_agent_count() - if force or active_count != last_active_count or (now - last_status_at) >= 1.0: + if ( + force + or active_count != last_active_count + or (now - last_status_at) >= 1.0 + ): self._update_runtime_status("draining") last_active_count = active_count last_status_at = now @@ -1462,7 +1603,10 @@ def _interrupt_running_agents(self, reason: str) -> None: continue try: agent.interrupt(reason) - logger.debug("Interrupted running agent for session %s during shutdown", session_key[:20]) + logger.debug( + "Interrupted running agent for session %s during shutdown", + session_key[:20], + ) except Exception as e: logger.debug("Failed interrupting agent during shutdown: %s", e) @@ -1517,18 +1661,22 @@ async def _notify_active_sessions_of_shutdown(self) -> None: notified.add(dedup_key) logger.info( "Sent shutdown notification to %s:%s", - platform_str, chat_id, + platform_str, + chat_id, ) except Exception as e: logger.debug( "Failed to send shutdown notification to %s:%s: %s", - platform_str, chat_id, e, + platform_str, + chat_id, + e, ) def _finalize_shutdown_agents(self, active_agents: Dict[str, Any]) -> None: for agent in active_agents.values(): try: from hermes_cli.plugins import invoke_hook as _invoke_hook + _invoke_hook( "on_session_finalize", session_id=getattr(agent, "session_id", None), @@ -1545,7 +1693,7 @@ def _finalize_shutdown_agents(self, active_agents: Dict[str, Any]) -> None: # background processes, httpx clients) to prevent zombie # process accumulation. try: - if hasattr(agent, 'close'): + if hasattr(agent, "close"): agent.close() except Exception: pass @@ -1610,7 +1758,8 @@ def _suspend_stuck_loop_sessions(self) -> int: logger.warning( "Auto-suspended stuck session %s (active across %d " "consecutive restarts — likely a stuck loop)", - session_key[:30], counts[session_key], + session_key[:30], + counts[session_key], ) except Exception: pass @@ -1681,7 +1830,9 @@ async def _launch_detached_restart_command(self) -> None: start_new_session=True, ) - def request_restart(self, *, detached: bool = False, via_service: bool = False) -> bool: + def request_restart( + self, *, detached: bool = False, via_service: bool = False + ) -> bool: if self._restart_task_started: return False self._restart_requested = True @@ -1691,7 +1842,9 @@ def request_restart(self, *, detached: bool = False, via_service: bool = False) async def _run_restart() -> None: await asyncio.sleep(0.05) - await self.stop(restart=True, detached_restart=detached, service_restart=via_service) + await self.stop( + restart=True, detached_restart=detached, service_restart=via_service + ) task = asyncio.create_task(_run_restart()) self._background_tasks.add(task) @@ -1701,13 +1854,14 @@ async def _run_restart() -> None: async def start(self) -> bool: """ Start the gateway and all configured platform adapters. - + Returns True if at least one adapter connected successfully. """ logger.info("Starting Hermes Gateway...") logger.info("Session storage: %s", self.config.sessions_dir) try: from hermes_cli.profiles import get_active_profile_name + _profile = get_active_profile_name() if _profile and _profile != "default": logger.info("Active profile: %s", _profile) @@ -1715,40 +1869,59 @@ async def start(self) -> bool: pass try: from gateway.status import write_runtime_status + write_runtime_status(gateway_state="starting", exit_reason=None) except Exception: pass - + # Warn if no user allowlists are configured and open access is not opted in _any_allowlist = any( os.getenv(v) - for v in ("TELEGRAM_ALLOWED_USERS", "DISCORD_ALLOWED_USERS", - "WHATSAPP_ALLOWED_USERS", "SLACK_ALLOWED_USERS", - "SIGNAL_ALLOWED_USERS", "SIGNAL_GROUP_ALLOWED_USERS", - "EMAIL_ALLOWED_USERS", - "SMS_ALLOWED_USERS", "MATTERMOST_ALLOWED_USERS", - "MATRIX_ALLOWED_USERS", "DINGTALK_ALLOWED_USERS", - "FEISHU_ALLOWED_USERS", - "WECOM_ALLOWED_USERS", - "WECOM_CALLBACK_ALLOWED_USERS", - "WEIXIN_ALLOWED_USERS", - "BLUEBUBBLES_ALLOWED_USERS", - "QQ_ALLOWED_USERS", - "GATEWAY_ALLOWED_USERS") + for v in ( + "TELEGRAM_ALLOWED_USERS", + "DISCORD_ALLOWED_USERS", + "WHATSAPP_ALLOWED_USERS", + "SLACK_ALLOWED_USERS", + "SIGNAL_ALLOWED_USERS", + "SIGNAL_GROUP_ALLOWED_USERS", + "EMAIL_ALLOWED_USERS", + "SMS_ALLOWED_USERS", + "MATTERMOST_ALLOWED_USERS", + "MATRIX_ALLOWED_USERS", + "DINGTALK_ALLOWED_USERS", + "FEISHU_ALLOWED_USERS", + "WECOM_ALLOWED_USERS", + "WECOM_CALLBACK_ALLOWED_USERS", + "WEIXIN_ALLOWED_USERS", + "BLUEBUBBLES_ALLOWED_USERS", + "QQ_ALLOWED_USERS", + "GATEWAY_ALLOWED_USERS", + ) ) - _allow_all = os.getenv("GATEWAY_ALLOW_ALL_USERS", "").lower() in ("true", "1", "yes") or any( + _allow_all = os.getenv("GATEWAY_ALLOW_ALL_USERS", "").lower() in ( + "true", + "1", + "yes", + ) or any( os.getenv(v, "").lower() in ("true", "1", "yes") - for v in ("TELEGRAM_ALLOW_ALL_USERS", "DISCORD_ALLOW_ALL_USERS", - "WHATSAPP_ALLOW_ALL_USERS", "SLACK_ALLOW_ALL_USERS", - "SIGNAL_ALLOW_ALL_USERS", "EMAIL_ALLOW_ALL_USERS", - "SMS_ALLOW_ALL_USERS", "MATTERMOST_ALLOW_ALL_USERS", - "MATRIX_ALLOW_ALL_USERS", "DINGTALK_ALLOW_ALL_USERS", - "FEISHU_ALLOW_ALL_USERS", - "WECOM_ALLOW_ALL_USERS", - "WECOM_CALLBACK_ALLOW_ALL_USERS", - "WEIXIN_ALLOW_ALL_USERS", - "BLUEBUBBLES_ALLOW_ALL_USERS", - "QQ_ALLOW_ALL_USERS") + for v in ( + "TELEGRAM_ALLOW_ALL_USERS", + "DISCORD_ALLOW_ALL_USERS", + "WHATSAPP_ALLOW_ALL_USERS", + "SLACK_ALLOW_ALL_USERS", + "SIGNAL_ALLOW_ALL_USERS", + "EMAIL_ALLOW_ALL_USERS", + "SMS_ALLOW_ALL_USERS", + "MATTERMOST_ALLOW_ALL_USERS", + "MATRIX_ALLOW_ALL_USERS", + "DINGTALK_ALLOW_ALL_USERS", + "FEISHU_ALLOW_ALL_USERS", + "WECOM_ALLOW_ALL_USERS", + "WECOM_CALLBACK_ALLOW_ALL_USERS", + "WEIXIN_ALLOW_ALL_USERS", + "BLUEBUBBLES_ALLOW_ALL_USERS", + "QQ_ALLOW_ALL_USERS", + ) ) if not _any_allowlist and not _allow_all: logger.warning( @@ -1756,16 +1929,19 @@ async def start(self) -> bool: "Set GATEWAY_ALLOW_ALL_USERS=true in ~/.hermes/.env to allow open access, " "or configure platform allowlists (e.g., TELEGRAM_ALLOWED_USERS=your_id)." ) - + # Discover and load event hooks self.hooks.discover_and_load() - + # Recover background processes from checkpoint (crash recovery) try: from tools.process_registry import process_registry + recovered = process_registry.recover_from_checkpoint() if recovered: - logger.info("Recovered %s background process(es) from previous run", recovered) + logger.info( + "Recovered %s background process(es) from previous run", recovered + ) except Exception as e: logger.warning("Process checkpoint recovery: %s", e) @@ -1789,7 +1965,9 @@ async def start(self) -> bool: try: suspended = self.session_store.suspend_recently_active() if suspended: - logger.info("Suspended %d in-flight session(s) from previous run", suspended) + logger.info( + "Suspended %d in-flight session(s) from previous run", suspended + ) except Exception as e: logger.warning("Session suspension on startup failed: %s", e) @@ -1808,24 +1986,24 @@ async def start(self) -> bool: enabled_platform_count = 0 startup_nonretryable_errors: list[str] = [] startup_retryable_errors: list[str] = [] - + # Initialize and connect each configured platform for platform, platform_config in self.config.platforms.items(): if not platform_config.enabled: continue enabled_platform_count += 1 - + adapter = self._create_adapter(platform, platform_config) if not adapter: logger.warning("No adapter available for %s", platform.value) continue - + # Set up message + fatal error handlers adapter.set_message_handler(self._handle_message) adapter.set_fatal_error_handler(self._handle_adapter_fatal_error) adapter.set_session_store(self.session_store) adapter.set_busy_session_handler(self._handle_active_session_busy_message) - + # Try to connect logger.info("Connecting to %s...", platform.value) self._update_platform_runtime_status( @@ -1852,7 +2030,9 @@ async def start(self) -> bool: if adapter.has_fatal_error: self._update_platform_runtime_status( platform.value, - platform_state="retrying" if adapter.fatal_error_retryable else "fatal", + platform_state="retrying" + if adapter.fatal_error_retryable + else "fatal", error_code=adapter.fatal_error_code, error_message=adapter.fatal_error_message, ) @@ -1902,56 +2082,72 @@ async def start(self) -> bool: "attempts": 1, "next_retry": time.monotonic() + 30, } - + if connected_count == 0: if startup_nonretryable_errors: reason = "; ".join(startup_nonretryable_errors) logger.error("Gateway hit a non-retryable startup conflict: %s", reason) try: from gateway.status import write_runtime_status - write_runtime_status(gateway_state="startup_failed", exit_reason=reason) + + write_runtime_status( + gateway_state="startup_failed", exit_reason=reason + ) except Exception: pass self._request_clean_exit(reason) return True if enabled_platform_count > 0: - reason = "; ".join(startup_retryable_errors) or "all configured messaging platforms failed to connect" - logger.error("Gateway failed to connect any configured messaging platform: %s", reason) + reason = ( + "; ".join(startup_retryable_errors) + or "all configured messaging platforms failed to connect" + ) + logger.error( + "Gateway failed to connect any configured messaging platform: %s", + reason, + ) try: from gateway.status import write_runtime_status - write_runtime_status(gateway_state="startup_failed", exit_reason=reason) + + write_runtime_status( + gateway_state="startup_failed", exit_reason=reason + ) except Exception: pass return False logger.warning("No messaging platforms enabled.") logger.info("Gateway will continue running for cron job execution.") - + # Update delivery router with adapters self.delivery_router.adapters = self.adapters - + self._running = True self._update_runtime_status("running") - + # Emit gateway:startup hook hook_count = len(self.hooks.loaded_hooks) if hook_count: logger.info("%s hook(s) loaded", hook_count) - await self.hooks.emit("gateway:startup", { - "platforms": [p.value for p in self.adapters.keys()], - }) - + await self.hooks.emit( + "gateway:startup", + { + "platforms": [p.value for p in self.adapters.keys()], + }, + ) + if connected_count > 0: logger.info("Gateway running with %s platform(s)", connected_count) - + # Build initial channel directory for send_message name resolution try: from gateway.channel_directory import build_channel_directory + directory = build_channel_directory(self.adapters) ch_count = sum(len(chs) for chs in directory.get("platforms", {}).values()) logger.info("Channel directory built: %d target(s)", ch_count) except Exception as e: logger.warning("Channel directory build failed: %s", e) - + # Check if we're restarting after a /update command. If the update is # still running, keep watching so we notify once it actually finishes. notified = await self._send_update_notification() @@ -1970,10 +2166,14 @@ async def start(self) -> bool: # Drain any recovered process watchers (from crash recovery checkpoint) try: from tools.process_registry import process_registry + while process_registry.pending_watchers: watcher = process_registry.pending_watchers.pop(0) asyncio.create_task(self._run_process_watcher(watcher)) - logger.info("Resumed watcher for recovered process %s", watcher.get("session_id")) + logger.info( + "Resumed watcher for recovered process %s", + watcher.get("session_id"), + ) except Exception as e: logger.error("Recovered watcher setup error: %s", e) @@ -1990,12 +2190,12 @@ async def start(self) -> bool: asyncio.create_task(self._platform_reconnect_watcher()) logger.info("Press Ctrl+C to stop") - + return True - + async def _session_expiry_watcher(self, interval: int = 300): """Background task that proactively flushes memories for expired sessions. - + Runs every `interval` seconds (default 5 min). For each session that has expired according to its reset policy, flushes memories in a thread pool and marks the session so it won't be flushed again. @@ -2031,7 +2231,8 @@ async def _session_expiry_watcher(self, interval: int = 300): ) logger.info( "Session expiry: %d sessions to flush (%s)", - len(_expired_entries), _plat_summary, + len(_expired_entries), + _plat_summary, ) for key, entry in _expired_entries: @@ -2045,19 +2246,28 @@ async def _session_expiry_watcher(self, interval: int = 300): if _cache_lock is not None: with _cache_lock: _cached = self._agent_cache.get(key) - _cached_agent = _cached[0] if isinstance(_cached, tuple) else _cached if _cached else None + _cached_agent = ( + _cached[0] + if isinstance(_cached, tuple) + else _cached + if _cached + else None + ) # Fall back to _running_agents in case the agent is # still mid-turn when the expiry fires. if _cached_agent is None: _cached_agent = self._running_agents.get(key) - if _cached_agent and _cached_agent is not _AGENT_PENDING_SENTINEL: + if ( + _cached_agent + and _cached_agent is not _AGENT_PENDING_SENTINEL + ): try: - if hasattr(_cached_agent, 'shutdown_memory_provider'): + if hasattr(_cached_agent, "shutdown_memory_provider"): _cached_agent.shutdown_memory_provider() except Exception: pass try: - if hasattr(_cached_agent, 'close'): + if hasattr(_cached_agent, "close"): _cached_agent.close() except Exception: pass @@ -2078,7 +2288,9 @@ async def _session_expiry_watcher(self, interval: int = 300): logger.warning( "Memory flush gave up after %d attempts for %s: %s. " "Marking as flushed to prevent infinite retry loop.", - failures, entry.session_id, e, + failures, + entry.session_id, + e, ) with self.session_store._lock: entry.memory_flushed = True @@ -2087,22 +2299,25 @@ async def _session_expiry_watcher(self, interval: int = 300): else: logger.debug( "Memory flush failed (%d/%d) for %s: %s", - failures, _MAX_FLUSH_RETRIES, entry.session_id, e, + failures, + _MAX_FLUSH_RETRIES, + entry.session_id, + e, ) if _expired_entries: - _flushed = sum( - 1 for _, e in _expired_entries if e.memory_flushed - ) + _flushed = sum(1 for _, e in _expired_entries if e.memory_flushed) _failed = len(_expired_entries) - _flushed if _failed: logger.info( "Session expiry done: %d flushed, %d pending retry", - _flushed, _failed, + _flushed, + _failed, ) else: logger.info( - "Session expiry done: %d flushed", _flushed, + "Session expiry done: %d flushed", + _flushed, ) except Exception as e: logger.debug("Session expiry watcher error: %s", e) @@ -2143,7 +2358,8 @@ async def _platform_reconnect_watcher(self) -> None: if info["attempts"] >= _MAX_ATTEMPTS: logger.warning( "Giving up reconnecting %s after %d attempts", - platform.value, info["attempts"], + platform.value, + info["attempts"], ) del self._failed_platforms[platform] continue @@ -2152,7 +2368,9 @@ async def _platform_reconnect_watcher(self) -> None: attempt = info["attempts"] + 1 logger.info( "Reconnecting %s (attempt %d/%d)...", - platform.value, attempt, _MAX_ATTEMPTS, + platform.value, + attempt, + _MAX_ATTEMPTS, ) try: @@ -2168,7 +2386,9 @@ async def _platform_reconnect_watcher(self) -> None: adapter.set_message_handler(self._handle_message) adapter.set_fatal_error_handler(self._handle_adapter_fatal_error) adapter.set_session_store(self.session_store) - adapter.set_busy_session_handler(self._handle_active_session_busy_message) + adapter.set_busy_session_handler( + self._handle_active_session_busy_message + ) success = await adapter.connect() if success: @@ -2186,13 +2406,19 @@ async def _platform_reconnect_watcher(self) -> None: # Rebuild channel directory with the new adapter try: - from gateway.channel_directory import build_channel_directory + from gateway.channel_directory import ( + build_channel_directory, + ) + build_channel_directory(self.adapters) except Exception: pass else: # Check if the failure is non-retryable - if adapter.has_fatal_error and not adapter.fatal_error_retryable: + if ( + adapter.has_fatal_error + and not adapter.fatal_error_retryable + ): self._update_platform_runtime_status( platform.value, platform_state="fatal", @@ -2201,7 +2427,8 @@ async def _platform_reconnect_watcher(self) -> None: ) logger.warning( "Reconnect %s: non-retryable error (%s), removing from retry queue", - platform.value, adapter.fatal_error_message, + platform.value, + adapter.fatal_error_message, ) del self._failed_platforms[platform] else: @@ -2209,14 +2436,16 @@ async def _platform_reconnect_watcher(self) -> None: platform.value, platform_state="retrying", error_code=adapter.fatal_error_code, - error_message=adapter.fatal_error_message or "failed to reconnect", + error_message=adapter.fatal_error_message + or "failed to reconnect", ) backoff = min(30 * (2 ** (attempt - 1)), _BACKOFF_CAP) info["attempts"] = attempt info["next_retry"] = time.monotonic() + backoff logger.info( "Reconnect %s failed, next retry in %ds", - platform.value, backoff, + platform.value, + backoff, ) except Exception as e: self._update_platform_runtime_status( @@ -2230,7 +2459,9 @@ async def _platform_reconnect_watcher(self) -> None: info["next_retry"] = time.monotonic() + backoff logger.warning( "Reconnect %s error: %s, next retry in %ds", - platform.value, e, backoff, + platform.value, + e, + backoff, ) # Check every 10 seconds for platforms that need reconnection @@ -2276,10 +2507,15 @@ async def _stop_impl() -> None: self._running_agent_count(), ) self._interrupt_running_agents( - "Gateway restarting" if self._restart_requested else "Gateway shutting down" + "Gateway restarting" + if self._restart_requested + else "Gateway shutting down" ) interrupt_deadline = asyncio.get_running_loop().time() + 5.0 - while self._running_agents and asyncio.get_running_loop().time() < interrupt_deadline: + while ( + self._running_agents + and asyncio.get_running_loop().time() < interrupt_deadline + ): self._update_runtime_status("draining") await asyncio.sleep(0.1) @@ -2295,7 +2531,9 @@ async def _stop_impl() -> None: try: await adapter.cancel_background_tasks() except Exception as e: - logger.debug("✗ %s background-task cancel error: %s", platform.value, e) + logger.debug( + "✗ %s background-task cancel error: %s", platform.value, e + ) try: await adapter.disconnect() logger.info("✓ %s disconnected", platform.value) @@ -2312,7 +2550,7 @@ async def _stop_impl() -> None: self._running_agents.clear() self._pending_messages.clear() self._pending_approvals.clear() - if hasattr(self, '_busy_ack_ts'): + if hasattr(self, "_busy_ack_ts"): self._busy_ack_ts.clear() self._shutdown_event.set() @@ -2320,21 +2558,25 @@ async def _stop_impl() -> None: # to a specific agent (catch-all for zombie prevention). try: from tools.process_registry import process_registry + process_registry.kill_all() except Exception: pass try: from tools.terminal_tool import cleanup_all_environments + cleanup_all_environments() except Exception: pass try: from tools.browser_tool import cleanup_all_browsers + cleanup_all_browsers() except Exception: pass from gateway.status import remove_pid_file + remove_pid_file() # Write a clean-shutdown marker so the next startup knows this @@ -2375,15 +2617,13 @@ async def _stop_impl() -> None: self._stop_task = asyncio.create_task(_stop_impl()) await self._stop_task - + async def wait_for_shutdown(self) -> None: """Wait for shutdown signal.""" await self._shutdown_event.wait() - + def _create_adapter( - self, - platform: Platform, - config: Any + self, platform: Platform, config: Any ) -> Optional[BasePlatformAdapter]: """Create the appropriate adapter for a platform.""" if hasattr(config, "extra") and isinstance(config.extra, dict): @@ -2397,72 +2637,119 @@ def _create_adapter( ) if platform == Platform.TELEGRAM: - from gateway.platforms.telegram import TelegramAdapter, check_telegram_requirements + from gateway.platforms.telegram import ( + TelegramAdapter, + check_telegram_requirements, + ) + if not check_telegram_requirements(): logger.warning("Telegram: python-telegram-bot not installed") return None return TelegramAdapter(config) - + elif platform == Platform.DISCORD: - from gateway.platforms.discord import DiscordAdapter, check_discord_requirements + from gateway.platforms.discord import ( + DiscordAdapter, + check_discord_requirements, + ) + if not check_discord_requirements(): logger.warning("Discord: discord.py not installed") return None return DiscordAdapter(config) - + elif platform == Platform.WHATSAPP: - from gateway.platforms.whatsapp import WhatsAppAdapter, check_whatsapp_requirements + from gateway.platforms.whatsapp import ( + WhatsAppAdapter, + check_whatsapp_requirements, + ) + if not check_whatsapp_requirements(): - logger.warning("WhatsApp: Node.js not installed or bridge not configured") + logger.warning( + "WhatsApp: Node.js not installed or bridge not configured" + ) return None return WhatsAppAdapter(config) - + elif platform == Platform.SLACK: from gateway.platforms.slack import SlackAdapter, check_slack_requirements + if not check_slack_requirements(): - logger.warning("Slack: slack-bolt not installed. Run: pip install 'hermes-agent[slack]'") + logger.warning( + "Slack: slack-bolt not installed. Run: pip install 'hermes-agent[slack]'" + ) return None return SlackAdapter(config) elif platform == Platform.SIGNAL: - from gateway.platforms.signal import SignalAdapter, check_signal_requirements + from gateway.platforms.signal import ( + SignalAdapter, + check_signal_requirements, + ) + if not check_signal_requirements(): - logger.warning("Signal: SIGNAL_HTTP_URL or SIGNAL_ACCOUNT not configured") + logger.warning( + "Signal: SIGNAL_HTTP_URL or SIGNAL_ACCOUNT not configured" + ) return None return SignalAdapter(config) elif platform == Platform.HOMEASSISTANT: - from gateway.platforms.homeassistant import HomeAssistantAdapter, check_ha_requirements + from gateway.platforms.homeassistant import ( + HomeAssistantAdapter, + check_ha_requirements, + ) + if not check_ha_requirements(): - logger.warning("HomeAssistant: aiohttp not installed or HASS_TOKEN not set") + logger.warning( + "HomeAssistant: aiohttp not installed or HASS_TOKEN not set" + ) return None return HomeAssistantAdapter(config) elif platform == Platform.EMAIL: from gateway.platforms.email import EmailAdapter, check_email_requirements + if not check_email_requirements(): - logger.warning("Email: EMAIL_ADDRESS, EMAIL_PASSWORD, EMAIL_IMAP_HOST, or EMAIL_SMTP_HOST not set") + logger.warning( + "Email: EMAIL_ADDRESS, EMAIL_PASSWORD, EMAIL_IMAP_HOST, or EMAIL_SMTP_HOST not set" + ) return None return EmailAdapter(config) elif platform == Platform.SMS: from gateway.platforms.sms import SmsAdapter, check_sms_requirements + if not check_sms_requirements(): - logger.warning("SMS: aiohttp not installed or TWILIO_ACCOUNT_SID/TWILIO_AUTH_TOKEN not set") + logger.warning( + "SMS: aiohttp not installed or TWILIO_ACCOUNT_SID/TWILIO_AUTH_TOKEN not set" + ) return None return SmsAdapter(config) elif platform == Platform.DINGTALK: - from gateway.platforms.dingtalk import DingTalkAdapter, check_dingtalk_requirements + from gateway.platforms.dingtalk import ( + DingTalkAdapter, + check_dingtalk_requirements, + ) + if not check_dingtalk_requirements(): - logger.warning("DingTalk: dingtalk-stream not installed or DINGTALK_CLIENT_ID/SECRET not set") + logger.warning( + "DingTalk: dingtalk-stream not installed or DINGTALK_CLIENT_ID/SECRET not set" + ) return None return DingTalkAdapter(config) elif platform == Platform.FEISHU: - from gateway.platforms.feishu import FeishuAdapter, check_feishu_requirements + from gateway.platforms.feishu import ( + FeishuAdapter, + check_feishu_requirements, + ) + if not check_feishu_requirements(): - logger.warning("Feishu: lark-oapi not installed or FEISHU_APP_ID/SECRET not set") + logger.warning( + "Feishu: lark-oapi not installed or FEISHU_APP_ID/SECRET not set" + ) return None return FeishuAdapter(config) @@ -2471,6 +2758,7 @@ def _create_adapter( WecomCallbackAdapter, check_wecom_callback_requirements, ) + if not check_wecom_callback_requirements(): logger.warning("WeComCallback: aiohttp/httpx not installed") return None @@ -2478,41 +2766,68 @@ def _create_adapter( elif platform == Platform.WECOM: from gateway.platforms.wecom import WeComAdapter, check_wecom_requirements + if not check_wecom_requirements(): - logger.warning("WeCom: aiohttp not installed or WECOM_BOT_ID/SECRET not set") + logger.warning( + "WeCom: aiohttp not installed or WECOM_BOT_ID/SECRET not set" + ) return None return WeComAdapter(config) elif platform == Platform.WEIXIN: - from gateway.platforms.weixin import WeixinAdapter, check_weixin_requirements + from gateway.platforms.weixin import ( + WeixinAdapter, + check_weixin_requirements, + ) + if not check_weixin_requirements(): logger.warning("Weixin: aiohttp/cryptography not installed") return None return WeixinAdapter(config) elif platform == Platform.MATTERMOST: - from gateway.platforms.mattermost import MattermostAdapter, check_mattermost_requirements + from gateway.platforms.mattermost import ( + MattermostAdapter, + check_mattermost_requirements, + ) + if not check_mattermost_requirements(): - logger.warning("Mattermost: MATTERMOST_TOKEN or MATTERMOST_URL not set, or aiohttp missing") + logger.warning( + "Mattermost: MATTERMOST_TOKEN or MATTERMOST_URL not set, or aiohttp missing" + ) return None return MattermostAdapter(config) elif platform == Platform.MATRIX: - from gateway.platforms.matrix import MatrixAdapter, check_matrix_requirements + from gateway.platforms.matrix import ( + MatrixAdapter, + check_matrix_requirements, + ) + if not check_matrix_requirements(): - logger.warning("Matrix: mautrix not installed or credentials not set. Run: pip install 'mautrix[encryption]'") + logger.warning( + "Matrix: mautrix not installed or credentials not set. Run: pip install 'mautrix[encryption]'" + ) return None return MatrixAdapter(config) elif platform == Platform.API_SERVER: - from gateway.platforms.api_server import APIServerAdapter, check_api_server_requirements + from gateway.platforms.api_server import ( + APIServerAdapter, + check_api_server_requirements, + ) + if not check_api_server_requirements(): logger.warning("API Server: aiohttp not installed") return None return APIServerAdapter(config) elif platform == Platform.WEBHOOK: - from gateway.platforms.webhook import WebhookAdapter, check_webhook_requirements + from gateway.platforms.webhook import ( + WebhookAdapter, + check_webhook_requirements, + ) + if not check_webhook_requirements(): logger.warning("Webhook: aiohttp not installed") return None @@ -2521,16 +2836,25 @@ def _create_adapter( return adapter elif platform == Platform.BLUEBUBBLES: - from gateway.platforms.bluebubbles import BlueBubblesAdapter, check_bluebubbles_requirements + from gateway.platforms.bluebubbles import ( + BlueBubblesAdapter, + check_bluebubbles_requirements, + ) + if not check_bluebubbles_requirements(): - logger.warning("BlueBubbles: aiohttp/httpx missing or BLUEBUBBLES_SERVER_URL/BLUEBUBBLES_PASSWORD not configured") + logger.warning( + "BlueBubbles: aiohttp/httpx missing or BLUEBUBBLES_SERVER_URL/BLUEBUBBLES_PASSWORD not configured" + ) return None return BlueBubblesAdapter(config) elif platform == Platform.QQBOT: from gateway.platforms.qqbot import QQAdapter, check_qq_requirements + if not check_qq_requirements(): - logger.warning("QQBot: aiohttp/httpx missing or QQ_APP_ID/QQ_CLIENT_SECRET not configured") + logger.warning( + "QQBot: aiohttp/httpx missing or QQ_APP_ID/QQ_CLIENT_SECRET not configured" + ) return None return QQAdapter(config) @@ -2539,7 +2863,7 @@ def _create_adapter( def _is_user_authorized(self, source: SessionSource) -> bool: """ Check if a user is authorized to use the bot. - + Checks in order: 1. Per-platform allow-all flag (e.g., DISCORD_ALLOW_ALL_USERS=true) 2. Environment variable allowlists (TELEGRAM_ALLOWED_USERS, etc.) @@ -2598,7 +2922,11 @@ def _is_user_authorized(self, source: SessionSource) -> bool: # Per-platform allow-all flag (e.g., DISCORD_ALLOW_ALL_USERS=true) platform_allow_all_var = platform_allow_all_map.get(source.platform, "") - if platform_allow_all_var and os.getenv(platform_allow_all_var, "").lower() in ("true", "1", "yes"): + if platform_allow_all_var and os.getenv(platform_allow_all_var, "").lower() in ( + "true", + "1", + "yes", + ): return True # Check pairing store (always checked, regardless of allowlists) @@ -2607,19 +2935,29 @@ def _is_user_authorized(self, source: SessionSource) -> bool: return True # Check platform-specific and global allowlists - platform_allowlist = os.getenv(platform_env_map.get(source.platform, ""), "").strip() + platform_allowlist = os.getenv( + platform_env_map.get(source.platform, ""), "" + ).strip() global_allowlist = os.getenv("GATEWAY_ALLOWED_USERS", "").strip() if not platform_allowlist and not global_allowlist: # No allowlists configured -- check global allow-all flag - return os.getenv("GATEWAY_ALLOW_ALL_USERS", "").lower() in ("true", "1", "yes") + return os.getenv("GATEWAY_ALLOW_ALL_USERS", "").lower() in ( + "true", + "1", + "yes", + ) # Check if user is in any allowlist allowed_ids = set() if platform_allowlist: - allowed_ids.update(uid.strip() for uid in platform_allowlist.split(",") if uid.strip()) + allowed_ids.update( + uid.strip() for uid in platform_allowlist.split(",") if uid.strip() + ) if global_allowlist: - allowed_ids.update(uid.strip() for uid in global_allowlist.split(",") if uid.strip()) + allowed_ids.update( + uid.strip() for uid in global_allowlist.split(",") if uid.strip() + ) # "*" in any allowlist means allow everyone (consistent with # SIGNAL_GROUP_ALLOWED_USERS precedent) @@ -2651,11 +2989,11 @@ def _get_unauthorized_dm_behavior(self, platform: Optional[Platform]) -> str: if config and hasattr(config, "get_unauthorized_dm_behavior"): return config.get_unauthorized_dm_behavior(platform) return "pair" - + async def _handle_message(self, event: MessageEvent) -> Optional[str]: """ Handle an incoming message from any platform. - + This is the core message processing pipeline: 1. Check user authorization 2. Check for commands (/new, /reset, etc.) @@ -2676,12 +3014,22 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: # channel forwards, anonymous admin actions) cannot be # authorized — drop silently instead of triggering the pairing # flow with a None user_id. - logger.debug("Ignoring message with no user_id from %s", source.platform.value) + logger.debug( + "Ignoring message with no user_id from %s", source.platform.value + ) return None elif not self._is_user_authorized(source): - logger.warning("Unauthorized user: %s (%s) on %s", source.user_id, source.user_name, source.platform.value) + logger.warning( + "Unauthorized user: %s (%s) on %s", + source.user_id, + source.user_name, + source.platform.value, + ) # In DMs: offer pairing code. In groups: silently ignore. - if source.chat_type == "dm" and self._get_unauthorized_dm_behavior(source.platform) == "pair": + if ( + source.chat_type == "dm" + and self._get_unauthorized_dm_behavior(source.platform) == "pair" + ): platform_name = source.platform.value if source.platform else "unknown" # Rate-limit ALL pairing responses (code or rejection) to # prevent spamming the user with repeated messages when @@ -2699,7 +3047,7 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: f"Hi~ I don't recognize you yet!\n\n" f"Here's your pairing code: `{code}`\n\n" f"Ask the bot owner to run:\n" - f"`hermes pairing approve {platform_name} {code}`" + f"`hermes pairing approve {platform_name} {code}`", ) else: adapter = self.adapters.get(source.platform) @@ -2707,12 +3055,12 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: await adapter.send( source.chat_id, "Too many pairing requests right now~ " - "Please try again later!" + "Please try again later!", ) # Record rate limit so subsequent messages are silently ignored self.pairing_store._record_rate_limit(platform_name, source.user_id) return None - + # Intercept messages that are responses to a pending /update prompt. # The update process (detached) wrote .update_prompt.json; the watcher # forwarded it to the user; now the user's reply goes back via @@ -2739,7 +3087,11 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: logger.warning("Failed to write update response: %s", e) return f"✗ Failed to send response to update process: {e}" _update_prompts.pop(_quick_key, None) - label = response_text if len(response_text) <= 20 else response_text[:20] + "…" + label = ( + response_text + if len(response_text) <= 20 + else response_text[:20] + "…" + ) return f"✓ Sent `{label}` to the update process." # PRIORITY handling when an agent is already running for this session. @@ -2781,20 +3133,24 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: # Evict if: agent is idle beyond timeout, OR wall-clock age is # extreme (10x timeout or 2h, whichever is larger — catches # cases where the agent object was garbage-collected). - _wall_ttl = max(_raw_stale_timeout * 10, 7200) if _raw_stale_timeout > 0 else float("inf") - _should_evict = ( - _stale_agent is not _AGENT_PENDING_SENTINEL - and ( - (_raw_stale_timeout > 0 and _stale_idle >= _raw_stale_timeout) - or _stale_age > _wall_ttl - ) + _wall_ttl = ( + max(_raw_stale_timeout * 10, 7200) + if _raw_stale_timeout > 0 + else float("inf") + ) + _should_evict = _stale_agent is not _AGENT_PENDING_SENTINEL and ( + (_raw_stale_timeout > 0 and _stale_idle >= _raw_stale_timeout) + or _stale_age > _wall_ttl ) if _should_evict: logger.warning( "Evicting stale _running_agents entry for %s " "(age: %.0fs, idle: %.0fs, timeout: %.0fs)%s", - _quick_key[:30], _stale_age, _stale_idle, - _raw_stale_timeout, _stale_detail, + _quick_key[:30], + _stale_age, + _stale_idle, + _raw_stale_timeout, + _stale_detail, ) del self._running_agents[_quick_key] self._running_agents_ts.pop(_quick_key, None) @@ -2806,6 +3162,7 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: # Resolve the command once for all early-intercept checks below. from hermes_cli.commands import resolve_command as _resolve_cmd_inner + _evt_cmd = event.get_command() _cmd_def_inner = _resolve_cmd_inner(_evt_cmd) if _evt_cmd else None @@ -2823,12 +3180,15 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: running_agent.interrupt("Stop requested") # Force-clean: remove the session lock regardless of agent state adapter = self.adapters.get(source.platform) - if adapter and hasattr(adapter, 'get_pending_message'): + if adapter and hasattr(adapter, "get_pending_message"): adapter.get_pending_message(_quick_key) # consume and discard self._pending_messages.pop(_quick_key, None) if _quick_key in self._running_agents: del self._running_agents[_quick_key] - logger.info("STOP for session %s — agent interrupted, session lock released", _quick_key[:20]) + logger.info( + "STOP for session %s — agent interrupted, session lock released", + _quick_key[:20], + ) return "⚡ Stopped. You can continue this session." # /reset and /new must bypass the running-agent guard so they @@ -2844,7 +3204,7 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: running_agent.interrupt("Session reset requested") # Clear any pending messages so the old text doesn't replay adapter = self.adapters.get(source.platform) - if adapter and hasattr(adapter, 'get_pending_message'): + if adapter and hasattr(adapter, "get_pending_message"): adapter.get_pending_message(_quick_key) # consume and discard self._pending_messages.pop(_quick_key, None) # Clean up the running agent entry so the reset handler @@ -2860,7 +3220,11 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: return "Usage: /queue <prompt>" adapter = self.adapters.get(source.platform) if adapter: - from gateway.platforms.base import MessageEvent as _ME, MessageType as _MT + from gateway.platforms.base import ( + MessageEvent as _ME, + MessageType as _MT, + ) + queued_event = _ME( text=queued_text, message_type=_MT.TEXT, @@ -2889,10 +3253,15 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: return await self._handle_background_command(event) if event.message_type == MessageType.PHOTO: - logger.debug("PRIORITY photo follow-up for session %s — queueing without interrupt", _quick_key[:20]) + logger.debug( + "PRIORITY photo follow-up for session %s — queueing without interrupt", + _quick_key[:20], + ) adapter = self.adapters.get(source.platform) if adapter: - merge_pending_message_event(adapter._pending_messages, _quick_key, event) + merge_pending_message_event( + adapter._pending_messages, _quick_key, event + ) return None running_agent = self._running_agents.get(_quick_key) @@ -2902,7 +3271,10 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: # Force-clean the sentinel so the session is unlocked. if _quick_key in self._running_agents: del self._running_agents[_quick_key] - logger.info("HARD STOP (pending) for session %s — sentinel cleared", _quick_key[:20]) + logger.info( + "HARD STOP (pending) for session %s — sentinel cleared", + _quick_key[:20], + ) return "⚡ Force-stopped. The agent was still starting — session unlocked." # Queue the message so it will be picked up after the # agent starts. @@ -2928,18 +3300,25 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: # Check for commands command = event.get_command() - + # Emit command:* hook for any recognized slash command. # GATEWAY_KNOWN_COMMANDS is derived from the central COMMAND_REGISTRY # in hermes_cli/commands.py — no hardcoded set to maintain here. - from hermes_cli.commands import GATEWAY_KNOWN_COMMANDS, resolve_command as _resolve_cmd + from hermes_cli.commands import ( + GATEWAY_KNOWN_COMMANDS, + resolve_command as _resolve_cmd, + ) + if command and command in GATEWAY_KNOWN_COMMANDS: - await self.hooks.emit(f"command:{command}", { - "platform": source.platform.value if source.platform else "", - "user_id": source.user_id, - "command": command, - "args": event.get_command_args().strip(), - }) + await self.hooks.emit( + f"command:{command}", + { + "platform": source.platform.value if source.platform else "", + "user_id": source.user_id, + "command": command, + "args": event.get_command_args().strip(), + }, + ) # Resolve aliases to canonical name so dispatch only checks canonicals. _cmd_def = _resolve_cmd(command) if command else None @@ -2947,13 +3326,13 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: if canonical == "new": return await self._handle_reset_command(event) - + if canonical == "help": return await self._handle_help_command(event) if canonical == "commands": return await self._handle_commands_command(event) - + if canonical == "profile": return await self._handle_profile_command(event) @@ -2962,10 +3341,10 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: if canonical == "restart": return await self._handle_restart_command(event) - + if canonical == "stop": return await self._handle_stop_command(event) - + if canonical == "reasoning": return await self._handle_reasoning_command(event) @@ -2983,13 +3362,16 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: if canonical == "provider": return await self._handle_provider_command(event) - + if canonical == "personality": return await self._handle_personality_command(event) if canonical == "plan": try: - from agent.skill_commands import build_plan_path, build_skill_invocation_message + from agent.skill_commands import ( + build_plan_path, + build_skill_invocation_message, + ) user_instruction = event.get_command_args().strip() plan_path = build_plan_path(user_instruction) @@ -3008,13 +3390,13 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: except Exception as e: logger.exception("Failed to prepare /plan command") return f"Failed to enter plan mode: {e}" - + if canonical == "retry": return await self._handle_retry_command(event) - + if canonical == "undo": return await self._handle_undo_command(event) - + if canonical == "sethome": return await self._handle_set_home_command(event) @@ -3085,7 +3467,9 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: stdout=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.PIPE, ) - stdout, stderr = await asyncio.wait_for(proc.communicate(), timeout=30) + stdout, stderr = await asyncio.wait_for( + proc.communicate(), timeout=30 + ) output = (stdout or stderr).decode().strip() return output if output else "Command returned no output." except asyncio.TimeoutError: @@ -3112,6 +3496,7 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: if command: try: from hermes_cli.plugins import get_plugin_command_handler + # Normalize underscores to hyphens so Telegram's underscored # autocomplete form matches plugin commands registered with # hyphens. See hermes_cli/commands.py:_build_telegram_menu. @@ -3119,6 +3504,7 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: if plugin_handler: user_args = event.get_command_args().strip() import asyncio as _aio + result = plugin_handler(user_args) if _aio.iscoroutine(result): result = await result @@ -3137,6 +3523,7 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: build_skill_invocation_message, resolve_skill_command_key, ) + skill_cmds = get_skill_commands() cmd_key = resolve_skill_command_key(command) if cmd_key is not None: @@ -3147,7 +3534,10 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: _skill_name = skill_cmds[cmd_key].get("name", "") _plat = source.platform.value if source.platform else None if _plat and _skill_name: - from agent.skill_utils import get_disabled_skill_names as _get_plat_disabled + from agent.skill_utils import ( + get_disabled_skill_names as _get_plat_disabled, + ) + if _skill_name in _get_plat_disabled(platform=_plat): return ( f"The **{_skill_name}** skill is disabled for {_plat}.\n" @@ -3189,7 +3579,7 @@ async def _handle_message(self, event: MessageEvent) -> Optional[str]: ) except Exception as e: logger.debug("Skill command check failed (non-fatal): %s", e) - + # Pending exec approvals are handled by /approve and /deny commands above. # No bare text matching — "yes" in normal conversation must not trigger # execution of a dangerous command. @@ -3244,9 +3634,15 @@ async def _prepare_inbound_message_text( audio_paths = [] for i, path in enumerate(event.media_urls): mtype = event.media_types[i] if i < len(event.media_types) else "" - if mtype.startswith("image/") or event.message_type == MessageType.PHOTO: + if ( + mtype.startswith("image/") + or event.message_type == MessageType.PHOTO + ): image_paths.append(path) - if mtype.startswith("audio/") or event.message_type in (MessageType.VOICE, MessageType.AUDIO): + if mtype.startswith("audio/") or event.message_type in ( + MessageType.VOICE, + MessageType.AUDIO, + ): audio_paths.append(path) if image_paths: @@ -3268,7 +3664,9 @@ async def _prepare_inbound_message_text( ) if any(marker in message_text for marker in _stt_fail_markers): _stt_adapter = self.adapters.get(source.platform) - _stt_meta = {"thread_id": source.thread_id} if source.thread_id else None + _stt_meta = ( + {"thread_id": source.thread_id} if source.thread_id else None + ) if _stt_adapter: try: _stt_msg = ( @@ -3292,7 +3690,19 @@ async def _prepare_inbound_message_text( if event.media_urls and event.message_type == MessageType.DOCUMENT: import mimetypes as _mimetypes - _TEXT_EXTENSIONS = {".txt", ".md", ".csv", ".log", ".json", ".xml", ".yaml", ".yml", ".toml", ".ini", ".cfg"} + _TEXT_EXTENSIONS = { + ".txt", + ".md", + ".csv", + ".log", + ".json", + ".xml", + ".yaml", + ".yml", + ".toml", + ".ini", + ".cfg", + } for i, path in enumerate(event.media_urls): mtype = event.media_types[i] if i < len(event.media_types) else "" if mtype in ("", "application/octet-stream"): @@ -3314,7 +3724,7 @@ async def _prepare_inbound_message_text( basename = _os.path.basename(path) parts = basename.split("_", 2) display_name = parts[2] if len(parts) >= 3 else basename - display_name = _re.sub(r'[^\w.\- ]', '_', display_name) + display_name = _re.sub(r"[^\w.\- ]", "_", display_name) if mtype.startswith("text/"): context_note = ( @@ -3361,7 +3771,8 @@ async def _prepare_inbound_message_text( if _adapter: await _adapter.send( source.chat_id, - "\n".join(_ctx_result.warnings) or "Context injection refused.", + "\n".join(_ctx_result.warnings) + or "Context injection refused.", ) return None if _ctx_result.expanded: @@ -3374,41 +3785,51 @@ async def _prepare_inbound_message_text( async def _handle_message_with_agent(self, event, source, _quick_key: str): """Inner handler that runs under the _running_agents sentinel guard.""" _msg_start_time = time.time() - _platform_name = source.platform.value if hasattr(source.platform, "value") else str(source.platform) + _platform_name = ( + source.platform.value + if hasattr(source.platform, "value") + else str(source.platform) + ) _msg_preview = (event.text or "")[:80].replace("\n", " ") logger.info( "inbound message: platform=%s user=%s chat=%s msg=%r", - _platform_name, source.user_name or source.user_id or "unknown", - source.chat_id or "unknown", _msg_preview, + _platform_name, + source.user_name or source.user_id or "unknown", + source.chat_id or "unknown", + _msg_preview, ) # Get or create session session_entry = self.session_store.get_or_create_session(source) session_key = session_entry.session_key - + # Emit session:start for new or auto-reset sessions _is_new_session = ( session_entry.created_at == session_entry.updated_at or getattr(session_entry, "was_auto_reset", False) ) if _is_new_session: - await self.hooks.emit("session:start", { - "platform": source.platform.value if source.platform else "", - "user_id": source.user_id, - "session_id": session_entry.session_id, - "session_key": session_key, - }) - + await self.hooks.emit( + "session:start", + { + "platform": source.platform.value if source.platform else "", + "user_id": source.user_id, + "session_id": session_entry.session_id, + "session_key": session_key, + }, + ) + # Build session context context = build_session_context(source, self.config, session_entry) - + # Set session context variables for tools (task-local, concurrency-safe) _session_env_tokens = self._set_session_env(context) - + # Read privacy.redact_pii from config (re-read per message) _redact_pii = False try: import yaml as _pii_yaml + with open(_config_path, encoding="utf-8") as _pf: _pcfg = _pii_yaml.safe_load(_pf) or {} _redact_pii = bool((_pcfg.get("privacy") or {}).get("redact_pii", False)) @@ -3417,11 +3838,11 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): # Build the context prompt to inject context_prompt = build_session_context_prompt(context, redact_pii=_redact_pii) - + # If the previous session expired and was auto-reset, prepend a notice # so the agent knows this is a fresh conversation (not an intentional /reset). - if getattr(session_entry, 'was_auto_reset', False): - reset_reason = getattr(session_entry, 'auto_reset_reason', None) or 'idle' + if getattr(session_entry, "was_auto_reset", False): + reset_reason = getattr(session_entry, "auto_reset_reason", None) or "idle" if reset_reason == "suspended": context_note = "[System note: The user's previous session was stopped and suspended. This is a fresh conversation with no prior context.]" elif reset_reason == "daily": @@ -3437,10 +3858,10 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): try: policy = self.session_store.config.get_reset_policy( platform=source.platform, - session_type=getattr(source, 'chat_type', 'dm'), + session_type=getattr(source, "chat_type", "dm"), ) platform_name = source.platform.value if source.platform else "" - had_activity = getattr(session_entry, 'reset_had_activity', False) + had_activity = getattr(session_entry, "reset_had_activity", False) # Suspended sessions always notify (they were explicitly stopped # or crashed mid-operation) — skip the policy check. should_notify = reset_reason == "suspended" or ( @@ -3458,7 +3879,13 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): else: hours = policy.idle_minutes // 60 mins = policy.idle_minutes % 60 - duration = f"{hours}h" if not mins else f"{hours}h {mins}m" if hours else f"{mins}m" + duration = ( + f"{hours}h" + if not mins + else f"{hours}h {mins}m" + if hours + else f"{mins}m" + ) reason_text = f"inactive for {duration}" notice = ( f"◐ Session automatically reset ({reason_text}). " @@ -3473,8 +3900,9 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): except Exception: pass await adapter.send( - source.chat_id, notice, - metadata=getattr(event, 'metadata', None), + source.chat_id, + notice, + metadata=getattr(event, "metadata", None), ) except Exception as e: logger.debug("Auto-reset notification failed (non-fatal): %s", e) @@ -3490,7 +3918,11 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): if _is_new_session and _auto: _skill_names = [_auto] if isinstance(_auto, str) else list(_auto) try: - from agent.skill_commands import _load_skill_payload, _build_skill_message + from agent.skill_commands import ( + _load_skill_payload, + _build_skill_message, + ) + _combined_parts: list[str] = [] _loaded_names: list[str] = [] for _sname in _skill_names: @@ -3513,14 +3945,17 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): event.text = "\n\n".join(_combined_parts) logger.info( "[Gateway] Auto-loaded skill(s) %s for session %s", - _loaded_names, session_key, + _loaded_names, + session_key, ) except Exception as e: - logger.warning("[Gateway] Failed to auto-load skill(s) %s: %s", _skill_names, e) + logger.warning( + "[Gateway] Failed to auto-load skill(s) %s: %s", _skill_names, e + ) # Load conversation history from transcript history = self.session_store.load_transcript(session_entry.session_id) - + # ----------------------------------------------------------------- # Session hygiene: auto-compress pathologically large transcripts # @@ -3562,6 +3997,7 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): _hyg_cfg_path = _hermes_home / "config.yaml" if _hyg_cfg_path.exists(): import yaml as _hyg_yaml + with open(_hyg_cfg_path, encoding="utf-8") as _hyg_f: _hyg_data = _hyg_yaml.safe_load(_hyg_f) or {} @@ -3570,7 +4006,11 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): if isinstance(_model_cfg, str): _hyg_model = _model_cfg elif isinstance(_model_cfg, dict): - _hyg_model = _model_cfg.get("default") or _model_cfg.get("model") or _hyg_model + _hyg_model = ( + _model_cfg.get("default") + or _model_cfg.get("model") + or _hyg_model + ) # Read explicit context_length override from model config # (same as run_agent.py lines 995-1005) _raw_ctx = _model_cfg.get("context_length") @@ -3610,7 +4050,10 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): if _hyg_config_context_length is None and _hyg_base_url: try: try: - from hermes_cli.config import get_compatible_custom_providers as _gw_gcp + from hermes_cli.config import ( + get_compatible_custom_providers as _gw_gcp, + ) + _hyg_custom_providers = _gw_gcp(_hyg_data) except Exception: _hyg_custom_providers = _hyg_data.get("custom_providers") @@ -3683,13 +4126,17 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): logger.info( "Session hygiene: %s messages, ~%s tokens (%s) — auto-compressing " "(threshold: %s%% of %s = %s tokens)", - _msg_count, f"{_approx_tokens:,}", _token_source, + _msg_count, + f"{_approx_tokens:,}", + _token_source, int(_hyg_threshold_pct * 100), f"{_hyg_context_length:,}", f"{_compress_token_threshold:,}", ) - _hyg_meta = {"thread_id": source.thread_id} if source.thread_id else None + _hyg_meta = ( + {"thread_id": source.thread_id} if source.thread_id else None + ) try: from run_agent import AIAgent @@ -3697,7 +4144,9 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): _hyg_model, _hyg_runtime = self._resolve_session_agent_runtime( source=source, session_key=session_key, - user_config=_hyg_data if isinstance(_hyg_data, dict) else None, + user_config=_hyg_data + if isinstance(_hyg_data, dict) + else None, ) if _hyg_runtime.get("api_key"): _hyg_msgs = [ @@ -3722,7 +4171,8 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): _compressed, _ = await loop.run_in_executor( None, lambda: _hyg_agent._compress_context( - _hyg_msgs, "", + _hyg_msgs, + "", approx_tokens=_approx_tokens, ), ) @@ -3750,8 +4200,10 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): logger.info( "Session hygiene: compressed %s → %s msgs, " "~%s → ~%s tokens", - _msg_count, _new_count, - f"{_approx_tokens:,}", f"{_new_tokens:,}", + _msg_count, + _new_count, + f"{_approx_tokens:,}", + f"{_new_tokens:,}", ) if _new_tokens >= _warn_token_threshold: @@ -3762,9 +4214,7 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): ) except Exception as e: - logger.warning( - "Session hygiene auto-compress failed: %s", e - ) + logger.warning("Session hygiene auto-compress failed: %s", e) # First-message onboarding -- only on the very first interaction ever if not history and not self.session_store.has_any_sessions(): @@ -3773,10 +4223,15 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): "Briefly introduce yourself and mention that /help shows available commands. " "Keep the introduction concise -- one or two sentences max.]" ) - + # One-time prompt if no home channel is set for this platform # Skip for webhooks - they deliver directly to configured targets (github_comment, etc.) - if not history and source.platform and source.platform != Platform.LOCAL and source.platform != Platform.WEBHOOK: + if ( + not history + and source.platform + and source.platform != Platform.LOCAL + and source.platform != Platform.WEBHOOK + ): platform_name = source.platform.value env_key = f"{platform_name.upper()}_HOME_CHANNEL" if not os.getenv(env_key): @@ -3788,9 +4243,9 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): f"A home channel is where Hermes delivers cron job results " f"and cross-platform messages.\n\n" f"Type /sethome to make this chat your home channel, " - f"or ignore to skip." + f"or ignore to skip.", ) - + # ----------------------------------------------------------------- # Voice channel awareness — inject current voice channel state # into context so the agent knows who is in the channel and who @@ -3860,8 +4315,11 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): _resp_len = len(response) logger.info( "response ready: platform=%s chat=%s time=%.1fs api_calls=%d response=%d chars", - _platform_name, source.chat_id or "unknown", - _response_time, _api_calls, _resp_len, + _platform_name, + source.chat_id or "unknown", + _response_time, + _api_calls, + _resp_len, ) # Successful turn — clear any stuck-loop counter for this session. @@ -3878,13 +4336,17 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): # Detect context-overflow failures and give specific guidance. # Generic 400 "Error" from Anthropic with large sessions is the # most common cause of this (#1630). - _is_ctx_fail = any(p in error_str for p in ( - "context", "token", "too large", "too long", - "exceed", "payload", - )) or ( - "400" in error_str - and len(history) > 50 - ) + _is_ctx_fail = any( + p in error_str + for p in ( + "context", + "token", + "too large", + "too long", + "exceed", + "payload", + ) + ) or ("400" in error_str and len(history) > 50) if _is_ctx_fail: response = ( @@ -3900,12 +4362,16 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): # If the agent's session_id changed during compression, update # session_entry so transcript writes below go to the right session. - if agent_result.get("session_id") and agent_result["session_id"] != session_entry.session_id: + if ( + agent_result.get("session_id") + and agent_result["session_id"] != session_entry.session_id + ): session_entry.session_id = agent_result["session_id"] # Prepend reasoning/thinking if display is enabled (per-platform) try: from gateway.display_config import resolve_display_setting as _rds + _show_reasoning_effective = _rds( _load_gateway_config(), _platform_config_key(source.platform), @@ -3927,14 +4393,18 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): response = f"💭 **Reasoning:**\n```\n{display_reasoning}\n```\n\n{response}" # Emit agent:end hook - await self.hooks.emit("agent:end", { - **hook_ctx, - "response": (response or "")[:500], - }) - + await self.hooks.emit( + "agent:end", + { + **hook_ctx, + "response": (response or "")[:500], + }, + ) + # Check for pending process watchers (check_interval on background processes) try: from tools.process_registry import process_registry + while process_registry.pending_watchers: watcher = process_registry.pending_watchers.pop(0) asyncio.create_task(self._run_process_watcher(watcher)) @@ -3947,6 +4417,7 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): # inject watch-type events here. try: from tools.process_registry import process_registry as _pr + _watch_events = [] while not _pr.completion_queue.empty(): evt = _pr.completion_queue.get_nowait() @@ -3970,7 +4441,7 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): # the time we reach here the approval has already been resolved. The # old post-loop pop_pending + approval_hint code was removed in favour # of the blocking approach that mirrors CLI's synchronous input(). - + # Save the full conversation to the transcript, including tool calls. # This preserves the complete agent loop (tool_calls, tool results, # intermediate reasoning) so sessions can be resumed with full context @@ -3992,7 +4463,11 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): # large to process. Auto-reset it so the next message starts # fresh instead of replaying the same oversized context in an # infinite fail loop. (#9893) - if agent_result.get("compression_exhausted") and session_entry and session_key: + if ( + agent_result.get("compression_exhausted") + and session_entry + and session_key + ): logger.info( "Auto-resetting session %s after compression exhaustion.", session_entry.session_id, @@ -4007,7 +4482,7 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): ) ts = datetime.now().isoformat() - + # If this is a fresh session (no history), write the full tool # definitions as the first entry so the transcript is self-describing # -- the same list of dicts sent as tools=[...] in the API request. @@ -4023,27 +4498,31 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): "model": _resolve_gateway_model(), "platform": source.platform.value if source.platform else "", "timestamp": ts, - } + }, ) - + # Find only the NEW messages from this turn (skip history we loaded). # Use the filtered history length (history_offset) that was actually # passed to the agent, not len(history) which includes session_meta # entries that were stripped before the agent saw them. if not agent_failed_early: history_len = agent_result.get("history_offset", len(history)) - new_messages = agent_messages[history_len:] if len(agent_messages) > history_len else [] - + new_messages = ( + agent_messages[history_len:] + if len(agent_messages) > history_len + else [] + ) + # If no new messages found (edge case), fall back to simple user/assistant if not new_messages: self.session_store.append_to_transcript( session_entry.session_id, - {"role": "user", "content": message_text, "timestamp": ts} + {"role": "user", "content": message_text, "timestamp": ts}, ) if response: self.session_store.append_to_transcript( session_entry.session_id, - {"role": "assistant", "content": response, "timestamp": ts} + {"role": "assistant", "content": response, "timestamp": ts}, ) else: # The agent already persisted these messages to SQLite via @@ -4058,10 +4537,11 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): # Add timestamp to each message for debugging entry = {**msg, "timestamp": ts} self.session_store.append_to_transcript( - session_entry.session_id, entry, + session_entry.session_id, + entry, skip_db=agent_persisted, ) - + # Token counts and model are now persisted by the agent directly. # Keep only last_prompt_tokens here for context-window tracking and # compression decisions. @@ -4072,7 +4552,9 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): # Auto voice reply: send TTS audio before the text response _already_sent = bool(agent_result.get("already_sent")) - if self._should_send_voice_reply(event, response, agent_messages, already_sent=_already_sent): + if self._should_send_voice_reply( + event, response, agent_messages, already_sent=_already_sent + ): await self._send_voice_reply(event, response) # If streaming already delivered the response, extract and @@ -4091,12 +4573,14 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): _media_adapter = self.adapters.get(source.platform) if _media_adapter: await self._deliver_media_from_response( - response, event, _media_adapter, + response, + event, + _media_adapter, ) return None return response - + except Exception as e: # Stop typing indicator on error too try: @@ -4110,7 +4594,7 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): error_detail = str(e)[:300] if str(e) else "no details available" status_hint = "" status_code = getattr(e, "status_code", None) - _hist_len = len(history) if 'history' in locals() else 0 + _hist_len = len(history) if "history" in locals() else 0 if status_code == 401: status_hint = " Check your API key or run `claude /login` to refresh OAuth credentials." elif status_code == 402: @@ -4128,6 +4612,7 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): _resets_in = _err_json.get("resets_in_seconds") if _resets_in and _resets_in > 0: import math + _hours = math.ceil(_resets_in / 3600) status_hint = f" Your plan's usage limit has been reached. It resets in ~{_hours}h." else: @@ -4135,7 +4620,9 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): else: status_hint = " You are being rate-limited. Please wait a moment and try again." elif status_code == 529: - status_hint = " The API is temporarily overloaded. Please try again shortly." + status_hint = ( + " The API is temporarily overloaded. Please try again shortly." + ) elif status_code in (400, 500): # 400 with a large session is context overflow. # 500 with a large session often means the payload is too large @@ -4157,7 +4644,7 @@ async def _handle_message_with_agent(self, event, source, _quick_key: str): finally: # Restore session context variables to their pre-handler state self._clear_session_env(_session_env_tokens) - + def _format_session_info(self) -> str: """Resolve current model config and return a formatted info block. @@ -4165,7 +4652,10 @@ def _format_session_info(self) -> str: users can immediately see if context detection went wrong (e.g. local models falling to the 128K default). """ - from agent.model_metadata import get_model_context_length, DEFAULT_FALLBACK_CONTEXT + from agent.model_metadata import ( + get_model_context_length, + DEFAULT_FALLBACK_CONTEXT, + ) model = _resolve_gateway_model() config_context_length = None @@ -4177,6 +4667,7 @@ def _format_session_info(self) -> str: cfg_path = _hermes_home / "config.yaml" if cfg_path.exists(): import yaml as _info_yaml + with open(cfg_path, encoding="utf-8") as f: data = _info_yaml.safe_load(f) or {} model_cfg = data.get("model", {}) @@ -4232,7 +4723,9 @@ def _format_session_info(self) -> str: ] # Show endpoint for local/custom setups - if base_url and ("localhost" in base_url or "127.0.0.1" in base_url or "0.0.0.0" in base_url): + if base_url and ( + "localhost" in base_url or "127.0.0.1" in base_url or "0.0.0.0" in base_url + ): lines.append(f"◆ Endpoint: {base_url}") return "\n".join(lines) @@ -4240,10 +4733,10 @@ def _format_session_info(self) -> str: async def _handle_reset_command(self, event: MessageEvent) -> str: """Handle /new or /reset command.""" source = event.source - + # Get existing session key session_key = self._session_key_for_source(source) - + # Flush memories in the background (fire-and-forget) so the user # gets the "Session reset!" response immediately. try: @@ -4263,7 +4756,13 @@ async def _handle_reset_command(self, event: MessageEvent) -> str: if _cache_lock is not None: with _cache_lock: _cached = self._agent_cache.get(session_key) - _old_agent = _cached[0] if isinstance(_cached, tuple) else _cached if _cached else None + _old_agent = ( + _cached[0] + if isinstance(_cached, tuple) + else _cached + if _cached + else None + ) if _old_agent is not None: try: if hasattr(_old_agent, "shutdown_memory_provider"): @@ -4279,12 +4778,14 @@ async def _handle_reset_command(self, event: MessageEvent) -> str: try: from tools.env_passthrough import clear_env_passthrough + clear_env_passthrough() except Exception: pass try: from tools.credential_files import clear_credential_files + clear_credential_files() except Exception: pass @@ -4299,25 +4800,35 @@ async def _handle_reset_command(self, event: MessageEvent) -> str: # Fire plugin on_session_finalize hook (session boundary) try: from hermes_cli.plugins import invoke_hook as _invoke_hook + _old_sid = old_entry.session_id if old_entry else None - _invoke_hook("on_session_finalize", session_id=_old_sid, - platform=source.platform.value if source.platform else "") + _invoke_hook( + "on_session_finalize", + session_id=_old_sid, + platform=source.platform.value if source.platform else "", + ) except Exception: pass # Emit session:end hook (session is ending) - await self.hooks.emit("session:end", { - "platform": source.platform.value if source.platform else "", - "user_id": source.user_id, - "session_key": session_key, - }) + await self.hooks.emit( + "session:end", + { + "platform": source.platform.value if source.platform else "", + "user_id": source.user_id, + "session_key": session_key, + }, + ) # Emit session:reset hook - await self.hooks.emit("session:reset", { - "platform": source.platform.value if source.platform else "", - "user_id": source.user_id, - "session_key": session_key, - }) + await self.hooks.emit( + "session:reset", + { + "platform": source.platform.value if source.platform else "", + "user_id": source.user_id, + "session_key": session_key, + }, + ) # Resolve session config info to surface to the user try: @@ -4335,15 +4846,20 @@ async def _handle_reset_command(self, event: MessageEvent) -> str: # Fire plugin on_session_reset hook (new session guaranteed to exist) try: from hermes_cli.plugins import invoke_hook as _invoke_hook + _new_sid = new_entry.session_id if new_entry else None - _invoke_hook("on_session_reset", session_id=_new_sid, - platform=source.platform.value if source.platform else "") + _invoke_hook( + "on_session_reset", + session_id=_new_sid, + platform=source.platform.value if source.platform else "", + ) except Exception: pass # Append a random tip to the reset message try: from hermes_cli.tips import get_random_tip + _tip_line = f"\n✦ Tip: {get_random_tip()}" except Exception: _tip_line = "" @@ -4351,7 +4867,7 @@ async def _handle_reset_command(self, event: MessageEvent) -> str: if session_info: return f"{header}\n\n{session_info}{_tip_line}" return f"{header}{_tip_line}" - + async def _handle_profile_command(self, event: MessageEvent) -> str: """Handle /profile — show active profile name and home directory.""" from hermes_constants import get_hermes_home, display_hermes_home @@ -4407,17 +4923,19 @@ async def _handle_status_command(self, event: MessageEvent) -> str: ] if title: lines.append(f"**Title:** {title}") - lines.extend([ - f"**Created:** {session_entry.created_at.strftime('%Y-%m-%d %H:%M')}", - f"**Last Activity:** {session_entry.updated_at.strftime('%Y-%m-%d %H:%M')}", - f"**Tokens:** {session_entry.total_tokens:,}", - f"**Agent Running:** {'Yes ⚡' if is_running else 'No'}", - "", - f"**Connected Platforms:** {', '.join(connected_platforms)}", - ]) + lines.extend( + [ + f"**Created:** {session_entry.created_at.strftime('%Y-%m-%d %H:%M')}", + f"**Last Activity:** {session_entry.updated_at.strftime('%Y-%m-%d %H:%M')}", + f"**Tokens:** {session_entry.total_tokens:,}", + f"**Agent Running:** {'Yes ⚡' if is_running else 'No'}", + "", + f"**Connected Platforms:** {', '.join(connected_platforms)}", + ] + ) return "\n".join(lines) - + async def _handle_stop_command(self, event: MessageEvent) -> str: """Handle /stop command - interrupt a running agent. @@ -4438,7 +4956,9 @@ async def _handle_stop_command(self, event: MessageEvent) -> str: # Force-clean the sentinel so the session is unlocked. if session_key in self._running_agents: del self._running_agents[session_key] - logger.info("STOP (pending) for session %s — sentinel cleared", session_key[:20]) + logger.info( + "STOP (pending) for session %s — sentinel cleared", session_key[:20] + ) return "⚡ Stopped. The agent hadn't started yet — you can continue this session." if agent: agent.interrupt("Stop requested") @@ -4462,15 +4982,16 @@ async def _handle_restart_command(self, event: MessageEvent) -> str: # notify them once it comes back online. try: import json as _json + notify_data = { - "platform": event.source.platform.value if event.source.platform else None, + "platform": event.source.platform.value + if event.source.platform + else None, "chat_id": event.source.chat_id, } if event.source.thread_id: notify_data["thread_id"] = event.source.thread_id - (_hermes_home / ".restart_notify.json").write_text( - _json.dumps(notify_data) - ) + (_hermes_home / ".restart_notify.json").write_text(_json.dumps(notify_data)) except Exception as e: logger.debug("Failed to write restart notify file: %s", e) @@ -4492,12 +5013,14 @@ async def _handle_restart_command(self, event: MessageEvent) -> str: async def _handle_help_command(self, event: MessageEvent) -> str: """Handle /help command - list available commands.""" from hermes_cli.commands import gateway_help_lines + lines = [ "📖 **Hermes Commands**\n", *gateway_help_lines(), ] try: from agent.skill_commands import get_skill_commands + skill_cmds = get_skill_commands() if skill_cmds: lines.append(f"\n⚡ **Skill Commands** ({len(skill_cmds)} active):") @@ -4506,7 +5029,9 @@ async def _handle_help_command(self, event: MessageEvent) -> str: for cmd in sorted_cmds[:10]: lines.append(f"`{cmd}` — {skill_cmds[cmd]['description']}") if len(sorted_cmds) > 10: - lines.append(f"\n... and {len(sorted_cmds) - 10} more. Use `/commands` for the full paginated list.") + lines.append( + f"\n... and {len(sorted_cmds) - 10} more. Use `/commands` for the full paginated list." + ) except Exception: pass return "\n".join(lines) @@ -4528,12 +5053,16 @@ async def _handle_commands_command(self, event: MessageEvent) -> str: entries = list(gateway_help_lines()) try: from agent.skill_commands import get_skill_commands + skill_cmds = get_skill_commands() if skill_cmds: entries.append("") entries.append("⚡ **Skill Commands**:") for cmd in sorted(skill_cmds): - desc = skill_cmds[cmd].get("description", "").strip() or "Skill command" + desc = ( + skill_cmds[cmd].get("description", "").strip() + or "Skill command" + ) entries.append(f"`{cmd}` — {desc}") except Exception: pass @@ -4542,11 +5071,12 @@ async def _handle_commands_command(self, event: MessageEvent) -> str: return "No commands available." from gateway.config import Platform + page_size = 15 if event.source.platform == Platform.TELEGRAM else 20 total_pages = max(1, (len(entries) + page_size - 1) // page_size) page = max(1, min(requested_page, total_pages)) start = (page - 1) * page_size - page_entries = entries[start:start + page_size] + page_entries = entries[start : start + page_size] lines = [ f"📚 **Commands** ({len(entries)} total, page {page}/{total_pages})", @@ -4561,9 +5091,11 @@ async def _handle_commands_command(self, event: MessageEvent) -> str: nav_parts.append(f"next → `/commands {page + 1}`") lines.extend(["", " | ".join(nav_parts)]) if page != requested_page: - lines.append(f"_(Requested page {requested_page} was out of range, showing page {page}.)_") + lines.append( + f"_(Requested page {requested_page} was out of range, showing page {page}.)_" + ) return "\n".join(lines) - + async def _handle_model_command(self, event: MessageEvent) -> Optional[str]: """Handle /model command — switch model for this session. @@ -4576,7 +5108,8 @@ async def _handle_model_command(self, event: MessageEvent) -> Optional[str]: """ import yaml from hermes_cli.model_switch import ( - switch_model as _switch_model, parse_model_flags, + switch_model as _switch_model, + parse_model_flags, list_authenticated_providers, ) from hermes_cli.providers import get_label @@ -4606,6 +5139,7 @@ async def _handle_model_command(self, event: MessageEvent) -> Optional[str]: user_provs = cfg.get("providers") try: from hermes_cli.config import get_compatible_custom_providers + custom_provs = get_compatible_custom_providers(cfg) except Exception: custom_provs = cfg.get("custom_providers") @@ -4687,7 +5221,10 @@ async def _on_model_selected( api_mode=result.api_mode, ) except Exception as exc: - logger.warning("Picker model switch failed for cached agent: %s", exc) + logger.warning( + "Picker model switch failed for cached agent: %s", + exc, + ) # Store model note + session override if not hasattr(_self, "_pending_model_notes"): @@ -4723,10 +5260,14 @@ async def _on_model_selected( if mi.has_cost_data(): lines.append(f"Cost: {mi.format_cost()}") lines.append(f"Capabilities: {mi.format_capabilities()}") - lines.append("_(session only — use `/model <name> --global` to persist)_") + lines.append( + "_(session only — use `/model <name> --global` to persist)_" + ) return "\n".join(lines) - metadata = {"thread_id": source.thread_id} if source.thread_id else None + metadata = ( + {"thread_id": source.thread_id} if source.thread_id else None + ) result = await adapter.send_model_picker( chat_id=source.chat_id, providers=providers, @@ -4755,7 +5296,11 @@ async def _on_model_selected( lines.append(f"**{p['name']}** `--provider {p['slug']}`{tag}:") if p["models"]: model_strs = ", ".join(f"`{m}`" for m in p["models"]) - extra = f" (+{p['total_models'] - len(p['models'])} more)" if p["total_models"] > len(p["models"]) else "" + extra = ( + f" (+{p['total_models'] - len(p['models'])} more)" + if p["total_models"] > len(p["models"]) + else "" + ) lines.append(f" {model_strs}{extra}") elif p.get("api_url"): lines.append(f" `{p['api_url']}`") @@ -4841,6 +5386,7 @@ async def _on_model_selected( if result.base_url: model_cfg["base_url"] = result.base_url from hermes_cli.config import save_config + save_config(cfg) except Exception as e: logger.warning("Failed to persist model switch: %s", e) @@ -4863,6 +5409,7 @@ async def _on_model_selected( else: try: from agent.model_metadata import get_model_context_length + ctx = get_model_context_length( result.new_model, base_url=result.base_url or current_base_url, @@ -4875,9 +5422,9 @@ async def _on_model_selected( # Cache notice cache_enabled = ( - ("openrouter" in (result.base_url or "").lower() and "claude" in result.new_model.lower()) - or result.api_mode == "anthropic_messages" - ) + "openrouter" in (result.base_url or "").lower() + and "claude" in result.new_model.lower() + ) or result.api_mode == "anthropic_messages" if cache_enabled: lines.append("Prompt caching: enabled") @@ -4903,7 +5450,7 @@ async def _handle_provider_command(self, event: MessageEvent) -> str: # Resolve current provider from config current_provider = "openrouter" model_cfg = {} - config_path = _hermes_home / 'config.yaml' + config_path = _hermes_home / "config.yaml" try: if config_path.exists(): with open(config_path, encoding="utf-8") as f: @@ -4918,13 +5465,16 @@ async def _handle_provider_command(self, event: MessageEvent) -> str: if current_provider == "auto": try: from hermes_cli.auth import resolve_provider as _resolve_provider + current_provider = _resolve_provider(current_provider) except Exception: current_provider = "openrouter" # Detect custom endpoint from config base_url if current_provider == "openrouter": - _cfg_base = model_cfg.get("base_url", "") if isinstance(model_cfg, dict) else "" + _cfg_base = ( + model_cfg.get("base_url", "") if isinstance(model_cfg, dict) else "" + ) if _cfg_base and "openrouter.ai" not in _cfg_base: current_provider = "custom" @@ -4947,17 +5497,17 @@ async def _handle_provider_command(self, event: MessageEvent) -> str: lines.append("Switch: `/model provider:model-name`") lines.append("Setup: `hermes setup`") return "\n".join(lines) - + async def _handle_personality_command(self, event: MessageEvent) -> str: """Handle /personality command - list or set a personality.""" import yaml args = event.get_command_args().strip().lower() - config_path = _hermes_home / 'config.yaml' + config_path = _hermes_home / "config.yaml" try: if config_path.exists(): - with open(config_path, 'r', encoding="utf-8") as f: + with open(config_path, "r", encoding="utf-8") as f: config = yaml.safe_load(f) or {} personalities = config.get("agent", {}).get("personalities", {}) else: @@ -4975,7 +5525,10 @@ async def _handle_personality_command(self, event: MessageEvent) -> str: lines.append("• `none` — (no personality overlay)") for name, prompt in personalities.items(): if isinstance(prompt, dict): - preview = prompt.get("description") or prompt.get("system_prompt", "")[:50] + preview = ( + prompt.get("description") + or prompt.get("system_prompt", "")[:50] + ) else: preview = prompt[:50] + "..." if len(prompt) > 50 else prompt lines.append(f"• `{name}` — {preview}") @@ -4986,9 +5539,9 @@ def _resolve_prompt(value): if isinstance(value, dict): parts = [value.get("system_prompt", "")] if value.get("tone"): - parts.append(f'Tone: {value["tone"]}') + parts.append(f"Tone: {value['tone']}") if value.get("style"): - parts.append(f'Style: {value["style"]}') + parts.append(f"Style: {value['style']}") return "\n".join(p for p in parts if p) return str(value) @@ -5021,13 +5574,13 @@ def _resolve_prompt(value): available = "`none`, " + ", ".join(f"`{n}`" for n in personalities) return f"Unknown personality: `{args}`\n\nAvailable: {available}" - + async def _handle_retry_command(self, event: MessageEvent) -> str: """Handle /retry command - re-send the last user message.""" source = event.source session_entry = self.session_store.get_or_create_session(source) history = self.session_store.load_transcript(session_entry.session_id) - + # Find the last user message last_user_msg = None last_user_idx = None @@ -5036,16 +5589,16 @@ async def _handle_retry_command(self, event: MessageEvent) -> str: last_user_msg = history[i].get("content", "") last_user_idx = i break - + if not last_user_msg: return "No previous message to retry." - + # Truncate history to before the last user message and persist truncated = history[:last_user_idx] self.session_store.rewrite_transcript(session_entry.session_id, truncated) # Reset stored token count — transcript was truncated session_entry.last_prompt_tokens = 0 - + # Re-send by creating a fake text event with the old message retry_event = MessageEvent( text=last_user_msg, @@ -5053,48 +5606,51 @@ async def _handle_retry_command(self, event: MessageEvent) -> str: source=source, raw_message=event.raw_message, ) - + # Let the normal message handler process it return await self._handle_message(retry_event) - + async def _handle_undo_command(self, event: MessageEvent) -> str: """Handle /undo command - remove the last user/assistant exchange.""" source = event.source session_entry = self.session_store.get_or_create_session(source) history = self.session_store.load_transcript(session_entry.session_id) - + # Find the last user message and remove everything from it onward last_user_idx = None for i in range(len(history) - 1, -1, -1): if history[i].get("role") == "user": last_user_idx = i break - + if last_user_idx is None: return "Nothing to undo." - + removed_msg = history[last_user_idx].get("content", "") removed_count = len(history) - last_user_idx - self.session_store.rewrite_transcript(session_entry.session_id, history[:last_user_idx]) + self.session_store.rewrite_transcript( + session_entry.session_id, history[:last_user_idx] + ) # Reset stored token count — transcript was truncated session_entry.last_prompt_tokens = 0 - + preview = removed_msg[:40] + "..." if len(removed_msg) > 40 else removed_msg - return f"↩️ Undid {removed_count} message(s).\nRemoved: \"{preview}\"" - + return f'↩️ Undid {removed_count} message(s).\nRemoved: "{preview}"' + async def _handle_set_home_command(self, event: MessageEvent) -> str: """Handle /sethome command -- set the current chat as the platform's home channel.""" source = event.source platform_name = source.platform.value if source.platform else "unknown" chat_id = source.chat_id chat_name = source.chat_name or chat_id - + env_key = f"{platform_name.upper()}_HOME_CHANNEL" - + # Save to config.yaml try: import yaml - config_path = _hermes_home / 'config.yaml' + + config_path = _hermes_home / "config.yaml" user_config = {} if config_path.exists(): with open(config_path, encoding="utf-8") as f: @@ -5105,12 +5661,12 @@ async def _handle_set_home_command(self, event: MessageEvent) -> str: os.environ[env_key] = str(chat_id) except Exception as e: return f"Failed to save home channel: {e}" - + return ( f"✅ Home channel set to **{chat_name}** (ID: {chat_id}).\n" f"Cron jobs and cross-platform messages will be delivered here." ) - + @staticmethod def _get_guild_id(event: MessageEvent) -> Optional[int]: """Extract Discord guild_id from the raw message object.""" @@ -5153,10 +5709,7 @@ async def _handle_voice_command(self, event: MessageEvent) -> str: self._save_voice_modes() if adapter: self._set_adapter_auto_tts_disabled(adapter, chat_id, disabled=False) - return ( - "Auto-TTS enabled.\n" - "All replies will include a voice message." - ) + return "Auto-TTS enabled.\nAll replies will include a voice message." elif args in ("channel", "join"): return await self._handle_voice_channel_join(event) elif args == "leave": @@ -5191,7 +5744,9 @@ async def _handle_voice_command(self, event: MessageEvent) -> str: self._voice_mode[chat_id] = "voice_only" self._save_voice_modes() if adapter: - self._set_adapter_auto_tts_disabled(adapter, chat_id, disabled=False) + self._set_adapter_auto_tts_disabled( + adapter, chat_id, disabled=False + ) return "Voice mode enabled." else: self._voice_mode[chat_id] = "off" @@ -5243,7 +5798,9 @@ async def _handle_voice_channel_join(self, event: MessageEvent) -> str: adapter._voice_sources[guild_id] = event.source.to_dict() self._voice_mode[event.source.chat_id] = "all" self._save_voice_modes() - self._set_adapter_auto_tts_disabled(adapter, event.source.chat_id, disabled=False) + self._set_adapter_auto_tts_disabled( + adapter, event.source.chat_id, disabled=False + ) return ( f"Joined voice channel **{voice_channel.name}**.\n" f"I'll speak my replies and listen to you. Use /voice leave to disconnect." @@ -5260,7 +5817,9 @@ async def _handle_voice_channel_leave(self, event: MessageEvent) -> str: if not guild_id or not hasattr(adapter, "leave_voice_channel"): return "Not in a voice channel." - if not hasattr(adapter, "is_in_voice_channel") or not adapter.is_in_voice_channel(guild_id): + if not hasattr( + adapter, "is_in_voice_channel" + ) or not adapter.is_in_voice_channel(guild_id): return "Not in a voice channel." try: @@ -5270,7 +5829,9 @@ async def _handle_voice_channel_leave(self, event: MessageEvent) -> str: # Always clean up state even if leave raised an exception self._voice_mode[event.source.chat_id] = "off" self._save_voice_modes() - self._set_adapter_auto_tts_disabled(adapter, event.source.chat_id, disabled=True) + self._set_adapter_auto_tts_disabled( + adapter, event.source.chat_id, disabled=True + ) if hasattr(adapter, "_voice_input_callback"): adapter._voice_input_callback = None return "Left voice channel." @@ -5326,7 +5887,11 @@ async def _handle_voice_channel_input( try: channel = adapter._client.get_channel(text_ch_id) if channel: - safe_text = transcript[:2000].replace("@everyone", "@\u200beveryone").replace("@here", "@\u200bhere") + safe_text = ( + transcript[:2000] + .replace("@everyone", "@\u200beveryone") + .replace("@here", "@\u200bhere") + ) await channel.send(f"**[Voice]** <@{user_id}>: {safe_text}") except Exception: pass @@ -5335,6 +5900,7 @@ async def _handle_voice_channel_input( # Use SimpleNamespace as raw_message so _get_guild_id() can extract # guild_id and _send_voice_reply() plays audio in the voice channel. from types import SimpleNamespace + event = MessageEvent( source=source, text=transcript, @@ -5367,11 +5933,10 @@ def _should_send_voice_reply( chat_id = event.source.chat_id voice_mode = self._voice_mode.get(chat_id, "off") - is_voice_input = (event.message_type == MessageType.VOICE) + is_voice_input = event.message_type == MessageType.VOICE - should = ( - (voice_mode == "all") - or (voice_mode == "voice_only" and is_voice_input) + should = (voice_mode == "all") or ( + voice_mode == "voice_only" and is_voice_input ) if not should: return False @@ -5401,6 +5966,7 @@ def _should_send_voice_reply( async def _send_voice_reply(self, event: MessageEvent, text: str) -> None: """Generate TTS audio and send as a voice message before the text reply.""" import uuid as _uuid + audio_path = None actual_path = None try: @@ -5413,7 +5979,8 @@ async def _send_voice_reply(self, event: MessageEvent, text: str) -> None: # Use .mp3 extension so edge-tts conversion to opus works correctly. # The TTS tool may convert to .ogg — use file_path from result. audio_path = os.path.join( - tempfile.gettempdir(), "hermes_voice", + tempfile.gettempdir(), + "hermes_voice", f"tts_reply_{_uuid.uuid4().hex[:12]}.mp3", ) os.makedirs(os.path.dirname(audio_path), exist_ok=True) @@ -5433,10 +6000,12 @@ async def _send_voice_reply(self, event: MessageEvent, text: str) -> None: # If connected to a voice channel, play there instead of sending a file guild_id = self._get_guild_id(event) - if (guild_id - and hasattr(adapter, "play_in_voice_channel") - and hasattr(adapter, "is_in_voice_channel") - and adapter.is_in_voice_channel(guild_id)): + if ( + guild_id + and hasattr(adapter, "play_in_voice_channel") + and hasattr(adapter, "is_in_voice_channel") + and adapter.is_in_voice_channel(guild_id) + ): await adapter.play_in_voice_channel(guild_id, actual_path) elif adapter and hasattr(adapter, "send_voice"): send_kwargs: Dict[str, Any] = { @@ -5475,11 +6044,15 @@ async def _deliver_media_from_response( _, cleaned = adapter.extract_images(response) local_files, _ = adapter.extract_local_files(cleaned) - _thread_meta = {"thread_id": event.source.thread_id} if event.source.thread_id else None + _thread_meta = ( + {"thread_id": event.source.thread_id} + if event.source.thread_id + else None + ) - _AUDIO_EXTS = {'.ogg', '.opus', '.mp3', '.wav', '.m4a'} - _VIDEO_EXTS = {'.mp4', '.mov', '.avi', '.mkv', '.webm', '.3gp'} - _IMAGE_EXTS = {'.jpg', '.jpeg', '.png', '.webp', '.gif'} + _AUDIO_EXTS = {".ogg", ".opus", ".mp3", ".wav", ".m4a"} + _VIDEO_EXTS = {".mp4", ".mov", ".avi", ".mkv", ".webm", ".3gp"} + _IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".webp", ".gif"} for media_path, is_voice in media_files: try: @@ -5509,7 +6082,9 @@ async def _deliver_media_from_response( metadata=_thread_meta, ) except Exception as e: - logger.warning("[%s] Post-stream media delivery failed: %s", adapter.name, e) + logger.warning( + "[%s] Post-stream media delivery failed: %s", adapter.name, e + ) for file_path in local_files: try: @@ -5527,7 +6102,9 @@ async def _deliver_media_from_response( metadata=_thread_meta, ) except Exception as e: - logger.warning("[%s] Post-stream file delivery failed: %s", adapter.name, e) + logger.warning( + "[%s] Post-stream file delivery failed: %s", adapter.name, e + ) except Exception as e: logger.warning("Post-stream media extraction failed: %s", e) @@ -5540,6 +6117,7 @@ async def _handle_rollback_command(self, event: MessageEvent) -> str: cp_cfg = {} try: import yaml as _y + _cfg_path = _hermes_home / "config.yaml" if _cfg_path.exists(): with open(_cfg_path, encoding="utf-8") as _f: @@ -5611,9 +6189,7 @@ async def _handle_background_command(self, event: MessageEvent) -> str: task_id = f"bg_{datetime.now().strftime('%H%M%S')}_{os.urandom(3).hex()}" # Fire-and-forget the background task - _task = asyncio.create_task( - self._run_background_task(prompt, source, task_id) - ) + _task = asyncio.create_task(self._run_background_task(prompt, source, task_id)) self._background_tasks.add(_task) _task.add_done_callback(self._background_tasks.discard) @@ -5628,7 +6204,11 @@ async def _run_background_task( adapter = self.adapters.get(source.platform) if not adapter: - logger.warning("No adapter for platform %s in background task %s", source.platform, task_id) + logger.warning( + "No adapter for platform %s in background task %s", + source.platform, + task_id, + ) return _thread_metadata = {"thread_id": source.thread_id} if source.thread_id else None @@ -5650,6 +6230,7 @@ async def _run_background_task( platform_key = _platform_config_key(source.platform) from hermes_cli.tools_config import _get_platform_tools + enabled_toolsets = sorted(_get_platform_tools(user_config, platform_key)) pr = self._provider_routing @@ -5717,7 +6298,7 @@ def run_sync(): ) # Send extracted images - for image_url, alt_text in (images or []): + for image_url, alt_text in images or []: try: await adapter.send_image( chat_id=source.chat_id, @@ -5728,7 +6309,7 @@ def run_sync(): pass # Send media files - for media_path in (media_files or []): + for media_path in media_files or []: try: await adapter.send_document( chat_id=source.chat_id, @@ -5777,8 +6358,11 @@ async def _handle_btw_command(self, event: MessageEvent) -> str: self._active_btw_tasks: dict = {} import uuid as _uuid + task_id = f"btw_{datetime.now().strftime('%H%M%S')}_{_uuid.uuid4().hex[:6]}" - _task = asyncio.create_task(self._run_btw_task(question, source, session_key, task_id)) + _task = asyncio.create_task( + self._run_btw_task(question, source, session_key, task_id) + ) self._background_tasks.add(_task) self._active_btw_tasks[session_key] = _task @@ -5793,14 +6377,20 @@ def _cleanup(task): return f'💬 /btw: "{preview}"\nReply will appear here shortly.' async def _run_btw_task( - self, question: str, source, session_key: str, task_id: str, + self, + question: str, + source, + session_key: str, + task_id: str, ) -> None: """Execute an ephemeral /btw side question and deliver the answer.""" from run_agent import AIAgent adapter = self.adapters.get(source.platform) if not adapter: - logger.warning("No adapter for platform %s in /btw task %s", source.platform, task_id) + logger.warning( + "No adapter for platform %s in /btw task %s", source.platform, task_id + ) return _thread_meta = {"thread_id": source.thread_id} if source.thread_id else None @@ -5823,21 +6413,26 @@ async def _run_btw_task( platform_key = _platform_config_key(source.platform) reasoning_config = self._load_reasoning_config() self._service_tier = self._load_service_tier() - turn_route = self._resolve_turn_agent_config(question, model, runtime_kwargs) + turn_route = self._resolve_turn_agent_config( + question, model, runtime_kwargs + ) pr = self._provider_routing # Snapshot history from running agent or stored transcript running_agent = self._running_agents.get(session_key) if running_agent and running_agent is not _AGENT_PENDING_SENTINEL: - history_snapshot = list(getattr(running_agent, "_session_messages", []) or []) + history_snapshot = list( + getattr(running_agent, "_session_messages", []) or [] + ) else: session_entry = self.session_store.get_or_create_session(source) - history_snapshot = self.session_store.load_transcript(session_entry.session_id) + history_snapshot = self.session_store.load_transcript( + session_entry.session_id + ) btw_prompt = ( "[Ephemeral /btw side question. Answer using the conversation " - "context. No tools available. Be direct and concise.]\n\n" - + question + "context. No tools available. Be direct and concise.]\n\n" + question ) def run_sync(): @@ -5898,15 +6493,19 @@ def run_sync(): metadata=_thread_meta, ) - for image_url, alt_text in (images or []): + for image_url, alt_text in images or []: try: - await adapter.send_image(chat_id=source.chat_id, image_url=image_url, caption=alt_text) + await adapter.send_image( + chat_id=source.chat_id, image_url=image_url, caption=alt_text + ) except Exception: pass - for media_path in (media_files or []): + for media_path in media_files or []: try: - await adapter.send_file(chat_id=source.chat_id, file_path=media_path) + await adapter.send_file( + chat_id=source.chat_id, file_path=media_path + ) except Exception: pass @@ -6105,7 +6704,9 @@ async def _handle_verbose_command(self, event: MessageEvent) -> str: if config_path.exists(): with open(config_path, encoding="utf-8") as f: user_config = yaml.safe_load(f) or {} - gate_enabled = user_config.get("display", {}).get("tool_progress_command", False) + gate_enabled = user_config.get("display", {}).get( + "tool_progress_command", False + ) except Exception: gate_enabled = False @@ -6127,7 +6728,10 @@ async def _handle_verbose_command(self, event: MessageEvent) -> str: # Read current effective mode for this platform via the resolver from gateway.display_config import resolve_display_setting - current = resolve_display_setting(user_config, platform_key, "tool_progress", "all") + + current = resolve_display_setting( + user_config, platform_key, "tool_progress", "all" + ) if current not in cycle: current = "all" idx = (cycle.index(current) + 1) % len(cycle) @@ -6135,12 +6739,18 @@ async def _handle_verbose_command(self, event: MessageEvent) -> str: # Save to display.platforms.<platform>.tool_progress try: - if "display" not in user_config or not isinstance(user_config.get("display"), dict): + if "display" not in user_config or not isinstance( + user_config.get("display"), dict + ): user_config["display"] = {} display = user_config["display"] - if "platforms" not in display or not isinstance(display.get("platforms"), dict): + if "platforms" not in display or not isinstance( + display.get("platforms"), dict + ): display["platforms"] = {} - if platform_key not in display["platforms"] or not isinstance(display["platforms"].get(platform_key), dict): + if platform_key not in display["platforms"] or not isinstance( + display["platforms"].get(platform_key), dict + ): display["platforms"][platform_key] = {} display["platforms"][platform_key]["tool_progress"] = new_mode atomic_yaml_write(config_path, user_config) @@ -6210,7 +6820,9 @@ async def _handle_compress_command(self, event: MessageEvent) -> str: loop = asyncio.get_event_loop() compressed, _ = await loop.run_in_executor( None, - lambda: tmp_agent._compress_context(msgs, "", approx_tokens=approx_tokens, focus_topic=focus_topic) + lambda: tmp_agent._compress_context( + msgs, "", approx_tokens=approx_tokens, focus_topic=focus_topic + ), ) # _compress_context already calls end_session() on the old session @@ -6236,7 +6848,7 @@ async def _handle_compress_command(self, event: MessageEvent) -> str: ) lines = [f"🗜️ {summary['headline']}"] if focus_topic: - lines.append(f"Focus: \"{focus_topic}\"") + lines.append(f'Focus: "{focus_topic}"') lines.append(summary["token_line"]) if summary["note"]: lines.append(summary["note"]) @@ -6276,7 +6888,9 @@ async def _handle_title_command(self, event: MessageEvent) -> str: except ValueError as e: return f"⚠️ {e}" if not sanitized: - return "⚠️ Title is empty after cleanup. Please use printable characters." + return ( + "⚠️ Title is empty after cleanup. Please use printable characters." + ) # Set the title try: if self._session_db.set_session_title(session_id, sanitized): @@ -6365,8 +6979,14 @@ async def _handle_resume_command(self, event: MessageEvent) -> str: # Count messages for context history = self.session_store.load_transcript(target_id) - msg_count = len([m for m in history if m.get("role") == "user"]) if history else 0 - msg_part = f" ({msg_count} message{'s' if msg_count != 1 else ''})" if msg_count else "" + msg_count = ( + len([m for m in history if m.get("role") == "user"]) if history else 0 + ) + msg_part = ( + f" ({msg_count} message{'s' if msg_count != 1 else ''})" + if msg_count + else "" + ) return f"↻ Resumed session **{title}**{msg_part}. Conversation restored." @@ -6395,6 +7015,7 @@ async def _handle_branch_command(self, event: MessageEvent) -> str: # Generate the new session ID from datetime import datetime as _dt + now = _dt.now() timestamp_str = now.strftime("%Y%m%d_%H%M%S") short_uuid = _uuid.uuid4().hex[:6] @@ -6415,7 +7036,9 @@ async def _handle_branch_command(self, event: MessageEvent) -> str: self._session_db.create_session( session_id=new_session_id, source=source.platform.value if source.platform else "gateway", - model=(self.config.get("model", {}) or {}).get("default") if isinstance(self.config, dict) else None, + model=(self.config.get("model", {}) or {}).get("default") + if isinstance(self.config, dict) + else None, parent_session_id=parent_session_id, ) except Exception as e: @@ -6481,14 +7104,21 @@ async def _handle_usage_command(self, event: MessageEvent) -> str: if cached: agent = cached[0] - if agent and hasattr(agent, "session_total_tokens") and agent.session_api_calls > 0: + if ( + agent + and hasattr(agent, "session_total_tokens") + and agent.session_api_calls > 0 + ): lines = [] # Rate limits (when available from provider headers) rl_state = agent.get_rate_limit_state() if rl_state and rl_state.has_data: from agent.rate_limit_tracker import format_rate_limit_compact - lines.append(f"⏱️ **Rate Limits:** {format_rate_limit_compact(rl_state)}") + + lines.append( + f"⏱️ **Rate Limits:** {format_rate_limit_compact(rl_state)}" + ) lines.append("") # Session token usage — detailed breakdown matching CLI @@ -6511,6 +7141,7 @@ async def _handle_usage_command(self, event: MessageEvent) -> str: # Cost estimation try: from agent.usage_pricing import CanonicalUsage, estimate_usage_cost + cost_result = estimate_usage_cost( agent.model, CanonicalUsage( @@ -6533,8 +7164,14 @@ async def _handle_usage_command(self, event: MessageEvent) -> str: # Context window and compressions ctx = agent.context_compressor if ctx.last_prompt_tokens: - pct = min(100, ctx.last_prompt_tokens / ctx.context_length * 100) if ctx.context_length else 0 - lines.append(f"Context: {ctx.last_prompt_tokens:,} / {ctx.context_length:,} ({pct:.0f}%)") + pct = ( + min(100, ctx.last_prompt_tokens / ctx.context_length * 100) + if ctx.context_length + else 0 + ) + lines.append( + f"Context: {ctx.last_prompt_tokens:,} / {ctx.context_length:,} ({pct:.0f}%)" + ) if ctx.compression_count: lines.append(f"Compressions: {ctx.compression_count}") @@ -6545,7 +7182,12 @@ async def _handle_usage_command(self, event: MessageEvent) -> str: history = self.session_store.load_transcript(session_entry.session_id) if history: from agent.model_metadata import estimate_messages_tokens_rough - msgs = [m for m in history if m.get("role") in ("user", "assistant") and m.get("content")] + + msgs = [ + m + for m in history + if m.get("role") in ("user", "assistant") and m.get("content") + ] approx = estimate_messages_tokens_rough(msgs) return ( f"📊 **Session Info**\n" @@ -6606,7 +7248,12 @@ async def _handle_reload_mcp_command(self, event: MessageEvent) -> str: """Handle /reload-mcp command -- disconnect and reconnect all MCP servers.""" loop = asyncio.get_event_loop() try: - from tools.mcp_tool import shutdown_mcp_servers, discover_mcp_tools, _servers, _lock + from tools.mcp_tool import ( + shutdown_mcp_servers, + discover_mcp_tools, + _servers, + _lock, + ) # Capture old server names before shutdown with _lock: @@ -6637,7 +7284,9 @@ async def _handle_reload_mcp_command(self, event: MessageEvent) -> str: if not connected_servers: lines.append("No MCP servers connected.") else: - lines.append(f"\n🔧 {len(new_tools)} tool(s) available from {len(connected_servers)} server(s)") + lines.append( + f"\n🔧 {len(new_tools)} tool(s) available from {len(connected_servers)} server(s)" + ) # Inject a message at the END of the session history so the # model knows tools changed on its next turn. Appended after @@ -6648,8 +7297,14 @@ async def _handle_reload_mcp_command(self, event: MessageEvent) -> str: if removed: change_parts.append(f"Removed servers: {', '.join(sorted(removed))}") if reconnected: - change_parts.append(f"Reconnected servers: {', '.join(sorted(reconnected))}") - tool_summary = f"{len(new_tools)} MCP tool(s) now available" if new_tools else "No MCP tools available" + change_parts.append( + f"Reconnected servers: {', '.join(sorted(reconnected))}" + ) + tool_summary = ( + f"{len(new_tools)} MCP tool(s) now available" + if new_tools + else "No MCP tools available" + ) change_detail = ". ".join(change_parts) + ". " if change_parts else "" reload_msg = { "role": "user", @@ -6699,7 +7354,8 @@ async def _handle_approve_command(self, event: MessageEvent) -> Optional[str]: session_key = self._session_key_for_source(source) from tools.approval import ( - resolve_gateway_approval, has_blocking_approval, + resolve_gateway_approval, + has_blocking_approval, ) if not has_blocking_approval(session_key): @@ -6733,7 +7389,9 @@ async def _handle_approve_command(self, event: MessageEvent) -> Optional[str]: _adapter.resume_typing_for_chat(source.chat_id) count_msg = f" ({count} commands)" if count > 1 else "" - logger.info("User approved %d dangerous command(s) via /approve%s", count, scope_msg) + logger.info( + "User approved %d dangerous command(s) via /approve%s", count, scope_msg + ) return f"✅ Command{'s' if count > 1 else ''} approved{scope_msg}{count_msg}. The agent is resuming..." async def _handle_deny_command(self, event: MessageEvent) -> str: @@ -6748,7 +7406,8 @@ async def _handle_deny_command(self, event: MessageEvent) -> str: session_key = self._session_key_for_source(source) from tools.approval import ( - resolve_gateway_approval, has_blocking_approval, + resolve_gateway_approval, + has_blocking_approval, ) if not has_blocking_approval(session_key): @@ -6775,18 +7434,36 @@ async def _handle_deny_command(self, event: MessageEvent) -> str: # Platforms where /update is allowed. ACP, API server, and webhooks are # programmatic interfaces that should not trigger system updates. - _UPDATE_ALLOWED_PLATFORMS = frozenset({ - Platform.TELEGRAM, Platform.DISCORD, Platform.SLACK, Platform.WHATSAPP, - Platform.SIGNAL, Platform.MATTERMOST, Platform.MATRIX, - Platform.HOMEASSISTANT, Platform.EMAIL, Platform.SMS, Platform.DINGTALK, - Platform.FEISHU, Platform.WECOM, Platform.WECOM_CALLBACK, Platform.WEIXIN, Platform.BLUEBUBBLES, Platform.QQBOT, Platform.LOCAL, - }) + _UPDATE_ALLOWED_PLATFORMS = frozenset( + { + Platform.TELEGRAM, + Platform.DISCORD, + Platform.SLACK, + Platform.WHATSAPP, + Platform.SIGNAL, + Platform.MATTERMOST, + Platform.MATRIX, + Platform.HOMEASSISTANT, + Platform.EMAIL, + Platform.SMS, + Platform.DINGTALK, + Platform.FEISHU, + Platform.WECOM, + Platform.WECOM_CALLBACK, + Platform.WEIXIN, + Platform.BLUEBUBBLES, + Platform.QQBOT, + Platform.LOCAL, + } + ) async def _handle_debug_command(self, event: MessageEvent) -> str: """Handle /debug — upload debug report + logs and return paste URLs.""" import asyncio from hermes_cli.debug import ( - _capture_dump, collect_debug_report, _read_full_log, + _capture_dump, + collect_debug_report, + _read_full_log, upload_to_pastebin, ) @@ -6860,7 +7537,7 @@ async def _handle_update_command(self, event: MessageEvent) -> str: return f"✗ {format_managed_message('update Hermes Agent')}" project_root = Path(__file__).parent.parent.resolve() - git_dir = project_root / '.git' + git_dir = project_root / ".git" if not git_dir.exists(): return "✗ Not a git repository — cannot update." @@ -6991,20 +7668,26 @@ async def _watch_update_progress( pass if not adapter or not chat_id: - logger.warning("Update watcher: cannot resolve adapter/chat_id, falling back to completion-only") + logger.warning( + "Update watcher: cannot resolve adapter/chat_id, falling back to completion-only" + ) # Fall back to old behavior: wait for exit code and send final notification - while (pending_path.exists() or claimed_path.exists()) and loop.time() < deadline: + while ( + pending_path.exists() or claimed_path.exists() + ) and loop.time() < deadline: if exit_code_path.exists(): await self._send_update_notification() return await asyncio.sleep(poll_interval) - if (pending_path.exists() or claimed_path.exists()) and not exit_code_path.exists(): + if ( + pending_path.exists() or claimed_path.exists() + ) and not exit_code_path.exists(): exit_code_path.write_text("124") await self._send_update_notification() return def _strip_ansi(text: str) -> str: - return _re.sub(r'\x1b\[[0-9;]*[A-Za-z]', '', text) + return _re.sub(r"\x1b\[[0-9;]*[A-Za-z]", "", text) bytes_sent = 0 last_stream_time = loop.time() @@ -7024,7 +7707,7 @@ async def _flush_buffer() -> None: return # Split into chunks if too long max_chunk = 3500 - chunks = [clean[i:i + max_chunk] for i in range(0, len(clean), max_chunk)] + chunks = [clean[i : i + max_chunk] for i in range(0, len(clean), max_chunk)] for chunk in chunks: try: await adapter.send(chat_id, f"```\n{chunk}\n```") @@ -7052,14 +7735,24 @@ async def _flush_buffer() -> None: if exit_code == 0: await adapter.send(chat_id, "✅ Hermes update finished.") else: - await adapter.send(chat_id, "❌ Hermes update failed (exit code {}).".format(exit_code)) - logger.info("Update finished (exit=%s), notified %s", exit_code, session_key) + await adapter.send( + chat_id, + "❌ Hermes update failed (exit code {}).".format(exit_code), + ) + logger.info( + "Update finished (exit=%s), notified %s", exit_code, session_key + ) except Exception as e: logger.warning("Update final notification failed: %s", e) # Cleanup - for p in (pending_path, claimed_path, output_path, - exit_code_path, prompt_path): + for p in ( + pending_path, + claimed_path, + output_path, + exit_code_path, + prompt_path, + ): p.unlink(missing_ok=True) (_hermes_home / ".update_response").unlink(missing_ok=True) self._update_prompt_pending.pop(session_key, None) @@ -7083,8 +7776,11 @@ async def _flush_buffer() -> None: # one that's still awaiting a response. Without this guard the # watcher would re-read the same .update_prompt.json every poll # cycle and spam the user with duplicate prompt messages. - if (prompt_path.exists() and session_key - and not self._update_prompt_pending.get(session_key)): + if ( + prompt_path.exists() + and session_key + and not self._update_prompt_pending.get(session_key) + ): try: prompt_data = json.loads(prompt_path.read_text()) prompt_text = prompt_data.get("prompt", "") @@ -7095,7 +7791,10 @@ async def _flush_buffer() -> None: await _flush_buffer() # Try platform-native buttons first (Discord, Telegram) sent_buttons = False - if getattr(type(adapter), "send_update_prompt", None) is not None: + if ( + getattr(type(adapter), "send_update_prompt", None) + is not None + ): try: await adapter.send_update_prompt( chat_id=chat_id, @@ -7105,7 +7804,9 @@ async def _flush_buffer() -> None: ) sent_buttons = True except Exception as btn_err: - logger.debug("Button-based update prompt failed: %s", btn_err) + logger.debug( + "Button-based update prompt failed: %s", btn_err + ) if not sent_buttons: default_hint = f" (default: {default})" if default else "" await adapter.send( @@ -7113,7 +7814,7 @@ async def _flush_buffer() -> None: f"⚕ **Update needs your input:**\n\n" f"{prompt_text}{default_hint}\n\n" f"Reply `/approve` (yes) or `/deny` (no), " - f"or type your answer directly." + f"or type your answer directly.", ) self._update_prompt_pending[session_key] = True # Remove the prompt file so it isn't re-read on the @@ -7121,7 +7822,11 @@ async def _flush_buffer() -> None: # .update_response to continue — it doesn't re-check # .update_prompt.json while waiting. prompt_path.unlink(missing_ok=True) - logger.info("Forwarded update prompt to %s: %s", session_key, prompt_text[:80]) + logger.info( + "Forwarded update prompt to %s: %s", + session_key, + prompt_text[:80], + ) except (json.JSONDecodeError, OSError) as e: logger.debug("Failed to read update prompt: %s", e) @@ -7133,11 +7838,18 @@ async def _flush_buffer() -> None: exit_code_path.write_text("124") await _flush_buffer() try: - await adapter.send(chat_id, "❌ Hermes update timed out after 30 minutes.") + await adapter.send( + chat_id, "❌ Hermes update timed out after 30 minutes." + ) except Exception: pass - for p in (pending_path, claimed_path, output_path, - exit_code_path, prompt_path): + for p in ( + pending_path, + claimed_path, + output_path, + exit_code_path, + prompt_path, + ): p.unlink(missing_ok=True) (_hermes_home / ".update_response").unlink(missing_ok=True) self._update_prompt_pending.pop(session_key, None) @@ -7200,7 +7912,7 @@ async def _send_update_notification(self) -> bool: if adapter and chat_id: # Strip ANSI escape codes for clean display - output = _re.sub(r'\x1b\[[0-9;]*m', '', output).strip() + output = _re.sub(r"\x1b\[[0-9;]*m", "", output).strip() if output: if len(output) > 3500: output = "…" + output[-3500:] @@ -7283,6 +7995,7 @@ def _set_session_env(self, context: SessionContext) -> list: in a ``finally`` block. """ from gateway.session_context import set_session_vars + return set_session_vars( platform=context.source.platform.value, chat_id=context.source.chat_id, @@ -7296,8 +8009,9 @@ def _set_session_env(self, context: SessionContext) -> list: def _clear_session_env(self, tokens: list) -> None: """Restore session context variables to their pre-handler values.""" from gateway.session_context import clear_session_vars + clear_session_vars(tokens) - + async def _enrich_message_with_vision( self, user_text: str, @@ -7405,14 +8119,13 @@ async def _enrich_message_with_transcription( if result["success"]: transcript = result["transcript"] enriched_parts.append( - f'[The user sent a voice message~ ' + f"[The user sent a voice message~ " f'Here\'s what they said: "{transcript}"]' ) else: error = result.get("error", "unknown error") - if ( - "No STT provider" in error - or error.startswith("Neither VOICE_TOOLS_OPENAI_KEY nor OPENAI_API_KEY is set") + if "No STT provider" in error or error.startswith( + "Neither VOICE_TOOLS_OPENAI_KEY nor OPENAI_API_KEY is set" ): _no_stt_note = ( "[The user sent a voice message but I can't listen " @@ -7461,7 +8174,11 @@ async def _inject_watch_notification(self, synth_text: str, original_event) -> N source = getattr(original_event, "source", None) if not source: return - platform_name = source.platform.value if hasattr(source.platform, "value") else str(source.platform) + platform_name = ( + source.platform.value + if hasattr(source.platform, "value") + else str(source.platform) + ) adapter = None for p, a in self.adapters.items(): if p.value == platform_name: @@ -7471,6 +8188,7 @@ async def _inject_watch_notification(self, synth_text: str, original_event) -> N return try: from gateway.platforms.base import MessageEvent, MessageType + synth_event = MessageEvent( text=synth_text, message_type=MessageType.TEXT, @@ -7508,8 +8226,13 @@ async def _run_process_watcher(self, watcher: dict) -> None: agent_notify = watcher.get("notify_on_complete", False) notify_mode = self._load_background_notifications_mode() - logger.debug("Process watcher started: %s (every %ss, notify=%s, agent_notify=%s)", - session_id, interval, notify_mode, agent_notify) + logger.debug( + "Process watcher started: %s (every %ss, notify=%s, agent_notify=%s)", + session_id, + interval, + notify_mode, + agent_notify, + ) if notify_mode == "off" and not agent_notify: # Still wait for the process to exit so we can log it, but don't @@ -7538,9 +8261,15 @@ async def _run_process_watcher(self, watcher: dict) -> None: # --- Agent-triggered completion: inject synthetic message --- # Skip if the agent already consumed the result via wait/poll/log from tools.process_registry import process_registry as _pr_check + if agent_notify and not _pr_check.is_completion_consumed(session_id): from tools.ansi_strip import strip_ansi - _out = strip_ansi(session.output_buffer[-2000:]) if session.output_buffer else "" + + _out = ( + strip_ansi(session.output_buffer[-2000:]) + if session.output_buffer + else "" + ) synth_text = ( f"[SYSTEM: Background process {session_id} completed " f"(exit code {session.exit_code}).\n" @@ -7557,6 +8286,7 @@ async def _run_process_watcher(self, watcher: dict) -> None: from gateway.platforms.base import MessageEvent, MessageType from gateway.session import SessionSource from gateway.config import Platform + _platform_enum = Platform(platform_name) _source = SessionSource( platform=_platform_enum, @@ -7573,7 +8303,8 @@ async def _run_process_watcher(self, watcher: dict) -> None: ) logger.info( "Process %s finished — injecting agent notification for session %s", - session_id, session_key, + session_id, + session_key, ) await adapter.handle_message(synth_event) except Exception as e: @@ -7582,12 +8313,13 @@ async def _run_process_watcher(self, watcher: dict) -> None: # --- Normal text-only notification --- # Decide whether to notify based on mode - should_notify = ( - notify_mode in ("all", "result") - or (notify_mode == "error" and session.exit_code not in (0, None)) + should_notify = notify_mode in ("all", "result") or ( + notify_mode == "error" and session.exit_code not in (0, None) ) if should_notify: - new_output = session.output_buffer[-1000:] if session.output_buffer else "" + new_output = ( + session.output_buffer[-1000:] if session.output_buffer else "" + ) message_text = ( f"[Background process {session_id} finished with exit code {session.exit_code}~ " f"Here's the final output:\n{new_output}]" @@ -7600,7 +8332,9 @@ async def _run_process_watcher(self, watcher: dict) -> None: if adapter and chat_id: try: send_meta = {"thread_id": thread_id} if thread_id else None - await adapter.send(chat_id, message_text, metadata=send_meta) + await adapter.send( + chat_id, message_text, metadata=send_meta + ) except Exception as e: logger.error("Watcher delivery error: %s", e) break @@ -7608,7 +8342,9 @@ async def _run_process_watcher(self, watcher: dict) -> None: elif has_new_output and notify_mode == "all" and not agent_notify: # New output available -- deliver status update (only in "all" mode) # Skip periodic updates for agent_notify watchers (they only care about completion) - new_output = session.output_buffer[-500:] if session.output_buffer else "" + new_output = ( + session.output_buffer[-500:] if session.output_buffer else "" + ) message_text = ( f"[Background process {session_id} is still running~ " f"New output:\n{new_output}]" @@ -7650,7 +8386,9 @@ def _agent_config_signature( # (e.g. "eyJhbGci"), which can cause false cache hits across auth # switches if only the first few characters are considered. _api_key = str(runtime.get("api_key", "") or "") - _api_key_fingerprint = hashlib.sha256(_api_key.encode()).hexdigest() if _api_key else "" + _api_key_fingerprint = ( + hashlib.sha256(_api_key.encode()).hexdigest() if _api_key else "" + ) blob = _j.dumps( [ @@ -7806,19 +8544,17 @@ async def _run_agent_via_proxy( _scfg = getattr(getattr(self, "config", None), "streaming", None) if _scfg is None: from gateway.config import StreamingConfig + _scfg = StreamingConfig() platform_key = _platform_config_key(source.platform) user_config = _load_gateway_config() from gateway.display_config import resolve_display_setting + _plat_streaming = resolve_display_setting( user_config, platform_key, "streaming" ) - _streaming_enabled = ( - _scfg.enabled and _scfg.transport != "off" - if _plat_streaming is None - else bool(_plat_streaming) - ) + _streaming_enabled = _is_gateway_streaming_enabled(_scfg, _plat_streaming) if source.thread_id: _thread_metadata: Optional[Dict[str, Any]] = {"thread_id": source.thread_id} @@ -7827,11 +8563,17 @@ async def _run_agent_via_proxy( if _streaming_enabled: try: - from gateway.stream_consumer import GatewayStreamConsumer, StreamConsumerConfig + from gateway.stream_consumer import ( + GatewayStreamConsumer, + StreamConsumerConfig, + ) from gateway.config import Platform + _adapter = self.adapters.get(source.platform) if _adapter: - _adapter_supports_edit = getattr(_adapter, "SUPPORTS_MESSAGE_EDITING", True) + _adapter_supports_edit = getattr( + _adapter, "SUPPORTS_MESSAGE_EDITING", True + ) _effective_cursor = _scfg.cursor if _adapter_supports_edit else "" if source.platform == Platform.MATRIX: _effective_cursor = "" @@ -7878,7 +8620,9 @@ async def _run_agent_via_proxy( error_text = await resp.text() logger.warning( "Proxy error (%d) from %s: %s", - resp.status, proxy_url, error_text[:500], + resp.status, + proxy_url, + error_text[:500], ) return { "final_response": f"⚠️ Proxy error ({resp.status}): {error_text[:300]}", @@ -7941,7 +8685,10 @@ async def _run_agent_via_proxy( _elapsed = time.time() - _start logger.info( "proxy response: url=%s session=%s time=%.1fs response=%d chars", - proxy_url, (session_id or "")[:20], _elapsed, len(full_response), + proxy_url, + (session_id or "")[:20], + _elapsed, + len(full_response), ) return { @@ -7972,13 +8719,13 @@ async def _run_agent( ) -> Dict[str, Any]: """ Run the agent with the given message and context. - + Returns the full result dict from run_conversation, including: - "final_response": str (the text to send back) - "messages": list (full conversation including tool calls) - "api_calls": int - "completed": bool - + This is run in a thread pool to not block the event loop. Supports interruption via new messages. """ @@ -7996,11 +8743,12 @@ async def _run_agent( from run_agent import AIAgent import queue - + user_config = _load_gateway_config() platform_key = _platform_config_key(source.platform) from hermes_cli.tools_config import _get_platform_tools + enabled_toolsets = sorted(_get_platform_tools(user_config, platform_key)) display_config = user_config.get("display", {}) @@ -8015,22 +8763,26 @@ async def _run_agent( # Apply tool preview length config (0 = no limit) try: from agent.display import set_tool_preview_max_len - _tpl = resolve_display_setting(user_config, platform_key, "tool_preview_length", 0) + + _tpl = resolve_display_setting( + user_config, platform_key, "tool_preview_length", 0 + ) set_tool_preview_max_len(int(_tpl) if _tpl else 0) except Exception: pass # Tool progress mode — resolved per-platform with env var fallback - _resolved_tp = resolve_display_setting(user_config, platform_key, "tool_progress") - progress_mode = ( - _resolved_tp - or os.getenv("HERMES_TOOL_PROGRESS_MODE") - or "all" + _resolved_tp = resolve_display_setting( + user_config, platform_key, "tool_progress" ) + progress_mode = _resolved_tp or os.getenv("HERMES_TOOL_PROGRESS_MODE") or "all" # Disable tool progress for webhooks - they don't support message editing, # so each progress line would be sent as a separate message. from gateway.config import Platform - tool_progress_enabled = progress_mode != "off" and source.platform != Platform.WEBHOOK + + tool_progress_enabled = ( + progress_mode != "off" and source.platform != Platform.WEBHOOK + ) # Natural assistant status messages are intentionally independent from # tool progress and token streaming. Users can keep tool_progress quiet # in chat platforms while opting into concise mid-turn updates. @@ -8041,14 +8793,20 @@ async def _run_agent( default=True, ) ) - + # Queue for progress messages (thread-safe) progress_queue = queue.Queue() if tool_progress_enabled else None last_tool = [None] # Mutable container for tracking in closure last_progress_msg = [None] # Track last message for dedup repeat_count = [0] # How many times the same message repeated - - def progress_callback(event_type: str, tool_name: str = None, preview: str = None, args: dict = None, **kwargs): + + def progress_callback( + event_type: str, + tool_name: str = None, + preview: str = None, + args: dict = None, + **kwargs, + ): """Callback invoked by agent on tool lifecycle events.""" if not progress_queue: return @@ -8061,44 +8819,48 @@ def progress_callback(event_type: str, tool_name: str = None, preview: str = Non if progress_mode == "new" and tool_name == last_tool[0]: return last_tool[0] = tool_name - + # Build progress message with primary argument preview from agent.display import get_tool_emoji + emoji = get_tool_emoji(tool_name, default="⚙️") - + # Verbose mode: show detailed arguments, respects tool_preview_length if progress_mode == "verbose": if args: from agent.display import get_tool_preview_max_len + _pl = get_tool_preview_max_len() import json as _json + args_str = _json.dumps(args, ensure_ascii=False, default=str) # When tool_preview_length is 0 (default), don't truncate # in verbose mode — the user explicitly asked for full # detail. Platform message-length limits handle the rest. if _pl > 0 and len(args_str) > _pl: - args_str = args_str[:_pl - 3] + "..." + args_str = args_str[: _pl - 3] + "..." msg = f"{emoji} {tool_name}({list(args.keys())})\n{args_str}" elif preview: - msg = f"{emoji} {tool_name}: \"{preview}\"" + msg = f'{emoji} {tool_name}: "{preview}"' else: msg = f"{emoji} {tool_name}..." progress_queue.put(msg) return - + # "all" / "new" modes: short preview, respects tool_preview_length # config (defaults to 40 chars when unset to keep gateway messages # compact — unlike CLI spinners, these persist as permanent messages). if preview: from agent.display import get_tool_preview_max_len + _pl = get_tool_preview_max_len() _cap = _pl if _pl > 0 else 40 if len(preview) > _cap: - preview = preview[:_cap - 3] + "..." - msg = f"{emoji} {tool_name}: \"{preview}\"" + preview = preview[: _cap - 3] + "..." + msg = f'{emoji} {tool_name}: "{preview}"' else: msg = f"{emoji} {tool_name}..." - + # Dedup: collapse consecutive identical progress messages. # Common with execute_code where models iterate with the same # code (same boilerplate imports → identical previews). @@ -8110,9 +8872,9 @@ def progress_callback(event_type: str, tool_name: str = None, preview: str = Non return last_progress_msg[0] = msg repeat_count[0] = 0 - + progress_queue.put(msg) - + # Background task to send progress messages # Accumulates tool lines into a single message that gets edited. # @@ -8125,7 +8887,9 @@ def progress_callback(event_type: str, tool_name: str = None, preview: str = Non _progress_thread_id = source.thread_id or event_message_id else: _progress_thread_id = source.thread_id - _progress_metadata = {"thread_id": _progress_thread_id} if _progress_thread_id else None + _progress_metadata = ( + {"thread_id": _progress_thread_id} if _progress_thread_id else None + ) async def send_progress_messages(): if not progress_queue: @@ -8139,6 +8903,7 @@ async def send_progress_messages(): # editing (e.g. iMessage/BlueBubbles) — each progress update # would become a separate message bubble, which is noisy. from gateway.platforms.base import BasePlatformAdapter as _BaseAdapter + if type(adapter).edit_message is _BaseAdapter.edit_message: while not progress_queue.empty(): try: @@ -8147,10 +8912,10 @@ async def send_progress_messages(): break return - progress_lines = [] # Accumulated tool lines - progress_msg_id = None # ID of the progress message to edit - can_edit = True # False once an edit fails (platform doesn't support it) - _last_edit_ts = 0.0 # Throttle edits to avoid Telegram flood control + progress_lines = [] # Accumulated tool lines + progress_msg_id = None # ID of the progress message to edit + can_edit = True # False once an edit fails (platform doesn't support it) + _last_edit_ts = 0.0 # Throttle edits to avoid Telegram flood control _PROGRESS_EDIT_INTERVAL = 1.5 # Minimum seconds between edits while True: @@ -8158,7 +8923,11 @@ async def send_progress_messages(): raw = progress_queue.get_nowait() # Handle dedup messages: update last line with repeat counter - if isinstance(raw, tuple) and len(raw) == 3 and raw[0] == "__dedup__": + if ( + isinstance(raw, tuple) + and len(raw) == 3 + and raw[0] == "__dedup__" + ): _, base_msg, count = raw if progress_lines: progress_lines[-1] = f"{base_msg} (×{count + 1})" @@ -8199,15 +8968,27 @@ async def send_progress_messages(): adapter.name, ) can_edit = False - await adapter.send(chat_id=source.chat_id, content=msg, metadata=_progress_metadata) + await adapter.send( + chat_id=source.chat_id, + content=msg, + metadata=_progress_metadata, + ) else: if can_edit: # First tool: send all accumulated text as new message full_text = "\n".join(progress_lines) - result = await adapter.send(chat_id=source.chat_id, content=full_text, metadata=_progress_metadata) + result = await adapter.send( + chat_id=source.chat_id, + content=full_text, + metadata=_progress_metadata, + ) else: # Editing unsupported: send just this line - result = await adapter.send(chat_id=source.chat_id, content=msg, metadata=_progress_metadata) + result = await adapter.send( + chat_id=source.chat_id, + content=msg, + metadata=_progress_metadata, + ) if result.success and result.message_id: progress_msg_id = result.message_id @@ -8215,7 +8996,9 @@ async def send_progress_messages(): # Restore typing indicator await asyncio.sleep(0.3) - await adapter.send_typing(source.chat_id, metadata=_progress_metadata) + await adapter.send_typing( + source.chat_id, metadata=_progress_metadata + ) except queue.Empty: await asyncio.sleep(0.3) @@ -8224,7 +9007,11 @@ async def send_progress_messages(): while not progress_queue.empty(): try: raw = progress_queue.get_nowait() - if isinstance(raw, tuple) and len(raw) == 3 and raw[0] == "__dedup__": + if ( + isinstance(raw, tuple) + and len(raw) == 3 + and raw[0] == "__dedup__" + ): _, base_msg, count = raw if progress_lines: progress_lines[-1] = f"{base_msg} (×{count + 1})" @@ -8247,13 +9034,14 @@ async def send_progress_messages(): except Exception as e: logger.error("Progress message error: %s", e) await asyncio.sleep(1) - + # We need to share the agent instance for interrupt support agent_holder = [None] # Mutable container for the agent instance result_holder = [None] # Mutable container for the result - tools_holder = [None] # Mutable container for the tool definitions + tools_holder = [None] # Mutable container for the tool definitions stream_consumer_holder = [None] # Mutable container for stream consumer - + streaming_enabled_holder = [False] # Effective stream-delta gate for this run + # Bridge sync step_callback → async hooks.emit for agent:step events _loop_for_step = asyncio.get_event_loop() _hooks_ref = self.hooks @@ -8264,20 +9052,25 @@ def _step_callback_sync(iteration: int, prev_tools: list) -> None: # keys. Normalise to keep "tool_names" backward-compatible for # user-authored hooks that do ', '.join(tool_names)'. _names: list[str] = [] - for _t in (prev_tools or []): + for _t in prev_tools or []: if isinstance(_t, dict): _names.append(_t.get("name") or "") else: _names.append(str(_t)) asyncio.run_coroutine_threadsafe( - _hooks_ref.emit("agent:step", { - "platform": source.platform.value if source.platform else "", - "user_id": source.user_id, - "session_id": session_id, - "iteration": iteration, - "tool_names": _names, - "tools": prev_tools, - }), + _hooks_ref.emit( + "agent:step", + { + "platform": source.platform.value + if source.platform + else "", + "user_id": source.user_id, + "session_id": session_id, + "iteration": iteration, + "tool_names": _names, + "tools": prev_tools, + }, + ), _loop_for_step, ) except Exception as _e: @@ -8286,7 +9079,9 @@ def _step_callback_sync(iteration: int, prev_tools: list) -> None: # Bridge sync status_callback → async adapter.send for context pressure _status_adapter = self.adapters.get(source.platform) _status_chat_id = source.chat_id - _status_thread_metadata = {"thread_id": _progress_thread_id} if _progress_thread_id else None + _status_thread_metadata = ( + {"thread_id": _progress_thread_id} if _progress_thread_id else None + ) def _status_callback_sync(event_type: str, message: str) -> None: if not _status_adapter: @@ -8318,15 +9113,19 @@ def run_sync(): # Read from env var or use default (same as CLI) max_iterations = int(os.getenv("HERMES_MAX_ITERATIONS", "90")) - + # Map platform enum to the platform hint key the agent understands. # Platform.LOCAL ("local") maps to "cli"; others pass through as-is. - platform_key = "cli" if source.platform == Platform.LOCAL else source.platform.value - + platform_key = ( + "cli" if source.platform == Platform.LOCAL else source.platform.value + ) + # Combine platform context with user-configured ephemeral system prompt combined_ephemeral = context_prompt or "" if self._ephemeral_system_prompt: - combined_ephemeral = (combined_ephemeral + "\n\n" + self._ephemeral_system_prompt).strip() + combined_ephemeral = ( + combined_ephemeral + "\n\n" + self._ephemeral_system_prompt + ).strip() # Re-read .env and config for fresh credentials (gateway is long-lived, # keys may change without restart). @@ -8345,7 +9144,9 @@ def run_sync(): ) logger.debug( "run_agent resolved: model=%s provider=%s session=%s", - model, runtime_kwargs.get("provider"), (session_key or "")[:30], + model, + runtime_kwargs.get("provider"), + (session_key or "")[:30], ) except Exception as exc: return { @@ -8362,9 +9163,10 @@ def run_sync(): # Set up stream consumer for token streaming or interim commentary. _stream_consumer = None _stream_delta_cb = None - _scfg = getattr(getattr(self, 'config', None), 'streaming', None) + _scfg = getattr(getattr(self, "config", None), "streaming", None) if _scfg is None: from gateway.config import StreamingConfig + _scfg = StreamingConfig() # Per-platform streaming gate: display.platforms.<plat>.streaming @@ -8374,17 +9176,18 @@ def run_sync(): user_config, platform_key, "streaming" ) # None = no per-platform override → follow global config - _streaming_enabled = ( - _scfg.enabled and _scfg.transport != "off" - if _plat_streaming is None - else bool(_plat_streaming) - ) + _streaming_enabled = _is_gateway_streaming_enabled(_scfg, _plat_streaming) + streaming_enabled_holder[0] = _streaming_enabled _want_stream_deltas = _streaming_enabled _want_interim_messages = interim_assistant_messages_enabled _want_interim_consumer = _want_interim_messages if _want_stream_deltas or _want_interim_consumer: try: - from gateway.stream_consumer import GatewayStreamConsumer, StreamConsumerConfig + from gateway.stream_consumer import ( + GatewayStreamConsumer, + StreamConsumerConfig, + ) + _adapter = self.adapters.get(source.platform) if _adapter: # Platforms that don't support editing sent messages @@ -8392,9 +9195,13 @@ def run_sync(): # without edit support, the consumer sends a partial # first message that can never be updated, resulting in # duplicate messages (partial + final). - _adapter_supports_edit = getattr(_adapter, "SUPPORTS_MESSAGE_EDITING", True) + _adapter_supports_edit = getattr( + _adapter, "SUPPORTS_MESSAGE_EDITING", True + ) if not _adapter_supports_edit: - raise RuntimeError("skip streaming for non-editable platform") + raise RuntimeError( + "skip streaming for non-editable platform" + ) _effective_cursor = _scfg.cursor # Some Matrix clients render the streaming cursor # as a visible tofu/white-box artifact. Keep @@ -8410,7 +9217,9 @@ def run_sync(): adapter=_adapter, chat_id=source.chat_id, config=_consumer_cfg, - metadata={"thread_id": _progress_thread_id} if _progress_thread_id else None, + metadata={"thread_id": _progress_thread_id} + if _progress_thread_id + else None, ) if _want_stream_deltas: _stream_delta_cb = _stream_consumer.on_delta @@ -8418,14 +9227,20 @@ def run_sync(): except Exception as _sc_err: logger.debug("Could not set up stream consumer: %s", _sc_err) - def _interim_assistant_cb(text: str, *, already_streamed: bool = False) -> None: + def _interim_assistant_cb( + text: str, *, already_streamed: bool = False + ) -> None: if _stream_consumer is not None: if already_streamed: _stream_consumer.on_segment_break() else: _stream_consumer.on_commentary(text) return - if already_streamed or not _status_adapter or not str(text or "").strip(): + if ( + already_streamed + or not _status_adapter + or not str(text or "").strip() + ): return try: asyncio.run_coroutine_threadsafe( @@ -8495,14 +9310,22 @@ def _interim_assistant_cb(text: str, *, already_streamed: bool = False) -> None: if _cache_lock and _cache is not None: with _cache_lock: _cache[session_key] = (agent, _sig) - logger.debug("Created new agent for session %s (sig=%s)", session_key, _sig) + logger.debug( + "Created new agent for session %s (sig=%s)", session_key, _sig + ) # Per-message state — callbacks and reasoning config change every # turn and must not be baked into the cached agent constructor. - agent.tool_progress_callback = progress_callback if tool_progress_enabled else None - agent.step_callback = _step_callback_sync if _hooks_ref.loaded_hooks else None + agent.tool_progress_callback = ( + progress_callback if tool_progress_enabled else None + ) + agent.step_callback = ( + _step_callback_sync if _hooks_ref.loaded_hooks else None + ) agent.stream_delta_callback = _stream_delta_cb - agent.interim_assistant_callback = _interim_assistant_cb if _want_interim_messages else None + agent.interim_assistant_callback = ( + _interim_assistant_cb if _want_interim_messages else None + ) agent.status_callback = _status_callback_sync agent.reasoning_config = reasoning_config agent.service_tier = self._service_tier @@ -8529,8 +9352,8 @@ def _bg_review_send(message: str) -> None: # Store agent reference for interrupt support agent_holder[0] = agent # Capture the full tool definitions for transcript logging - tools_holder[0] = agent.tools if hasattr(agent, 'tools') else None - + tools_holder[0] = agent.tools if hasattr(agent, "tools") else None + # Convert history to agent format. # Two cases: # 1. Normal path (from transcript): simple {role, content, timestamp} dicts @@ -8544,22 +9367,22 @@ def _bg_review_send(message: str) -> None: role = msg.get("role") if not role: continue - + # Skip metadata entries (tool definitions, session info) # -- these are for transcript logging, not for the LLM if role in ("session_meta",): continue - + # Skip system messages -- the agent rebuilds its own system prompt if role == "system": continue - + # Rich agent messages (tool_calls, tool results) must be passed # through intact so the API sees valid assistant→tool sequences has_tool_calls = "tool_calls" in msg has_tool_call_id = "tool_call_id" in msg is_tool_message = role == "tool" - + if has_tool_calls or has_tool_call_id or is_tool_message: clean_msg = {k: v for k, v in msg.items() if k != "timestamp"} agent_history.append(clean_msg) @@ -8577,13 +9400,16 @@ def _bg_review_send(message: str) -> None: # The agent's _build_api_kwargs converts these to the # provider-specific format (reasoning_content, etc.). if role == "assistant": - for _rkey in ("reasoning", "reasoning_details", - "codex_reasoning_items"): + for _rkey in ( + "reasoning", + "reasoning_details", + "codex_reasoning_items", + ): _rval = msg.get(_rkey) if _rval: entry[_rkey] = _rval agent_history.append(entry) - + # Collect MEDIA paths already in history so we can exclude them # from the current turn's extraction. This is compression-safe: # even if the message list shrinks, we know which paths are old. @@ -8592,11 +9418,11 @@ def _bg_review_send(message: str) -> None: if _hm.get("role") in ("tool", "function"): _hc = _hm.get("content", "") if "MEDIA:" in _hc: - for _match in re.finditer(r'MEDIA:(\S+)', _hc): + for _match in re.finditer(r"MEDIA:(\S+)", _hc): _p = _match.group(1).strip().rstrip('",}') if _p: _history_media_paths.add(_p) - + # Register per-session gateway approval callback so dangerous # command approval blocks the agent thread (mirrors CLI input()). # The callback bridges sync→async to send the approval request @@ -8631,7 +9457,10 @@ def _approval_notify_sync(approval_data: dict) -> None: # Prefer button-based approval when the adapter supports it. # Check the *class* for the method, not the instance — avoids # false positives from MagicMock auto-attribute creation in tests. - if getattr(type(_status_adapter), "send_exec_approval", None) is not None: + if ( + getattr(type(_status_adapter), "send_exec_approval", None) + is not None + ): try: asyncio.run_coroutine_threadsafe( _status_adapter.send_exec_approval( @@ -8671,7 +9500,7 @@ def _approval_notify_sync(approval_data: dict) -> None: logger.error("Failed to send approval request: %s", _e) # Prepend pending model switch note so the model knows about the switch - _pending_notes = getattr(self, '_pending_model_notes', {}) + _pending_notes = getattr(self, "_pending_model_notes", {}) _msn = _pending_notes.pop(session_key, None) if session_key else None if _msn: message = _msn + "\n\n" + message @@ -8687,15 +9516,16 @@ def _approval_notify_sync(approval_data: dict) -> None: "process the last tool result(s). The conversation history contains " "tool outputs you haven't responded to yet. Please finish processing " "those results and summarize what was accomplished, then address the " - "user's new message below.]\n\n" - + message + "user's new message below.]\n\n" + message ) _approval_session_key = session_key or "" _approval_session_token = set_current_session_key(_approval_session_key) register_gateway_notify(_approval_session_key, _approval_notify_sync) try: - result = agent.run_conversation(message, conversation_history=agent_history, task_id=session_id) + result = agent.run_conversation( + message, conversation_history=agent_history, task_id=session_id + ) finally: unregister_gateway_notify(_approval_session_key) reset_current_session_key(_approval_session_token) @@ -8704,7 +9534,7 @@ def _approval_notify_sync(approval_data: dict) -> None: # Signal the stream consumer that the agent is done if _stream_consumer is not None: _stream_consumer.finish() - + # Return final response, or a message if something went wrong final_response = result.get("final_response") @@ -8714,13 +9544,19 @@ def _approval_notify_sync(approval_data: dict) -> None: _output_toks = 0 _agent = agent_holder[0] if _agent and hasattr(_agent, "context_compressor"): - _last_prompt_toks = getattr(_agent.context_compressor, "last_prompt_tokens", 0) + _last_prompt_toks = getattr( + _agent.context_compressor, "last_prompt_tokens", 0 + ) _input_toks = getattr(_agent, "session_prompt_tokens", 0) _output_toks = getattr(_agent, "session_completion_tokens", 0) _resolved_model = getattr(_agent, "model", None) if _agent else None if not final_response: - error_msg = f"⚠️ {result['error']}" if result.get("error") else "(No response generated)" + error_msg = ( + f"⚠️ {result['error']}" + if result.get("error") + else "(No response generated)" + ) return { "final_response": error_msg, "messages": result.get("messages", []), @@ -8734,7 +9570,7 @@ def _approval_notify_sync(approval_data: dict) -> None: "output_tokens": _output_toks, "model": _resolved_model, } - + # Scan tool results for MEDIA:<path> tags that need to be delivered # as native audio/file attachments. The TTS tool embeds MEDIA: tags # in its JSON response, but the model's final text reply usually @@ -8752,13 +9588,13 @@ def _approval_notify_sync(approval_data: dict) -> None: if msg.get("role") in ("tool", "function"): content = msg.get("content", "") if "MEDIA:" in content: - for match in re.finditer(r'MEDIA:(\S+)', content): + for match in re.finditer(r"MEDIA:(\S+)", content): path = match.group(1).strip().rstrip('",}') if path and path not in _history_media_paths: media_tags.append(f"MEDIA:{path}") if "[[audio_as_voice]]" in content: has_voice_directive = True - + if media_tags: seen = set() unique_tags = [] @@ -8769,25 +9605,33 @@ def _approval_notify_sync(approval_data: dict) -> None: if has_voice_directive: unique_tags.insert(0, "[[audio_as_voice]]") final_response = final_response + "\n" + "\n".join(unique_tags) - + # Sync session_id: the agent may have created a new session during # mid-run context compression (_compress_context splits sessions). # If so, update the session store entry so the NEXT message loads # the compressed transcript, not the stale pre-compression one. agent = agent_holder[0] _session_was_split = False - if agent and session_key and hasattr(agent, 'session_id') and agent.session_id != session_id: + if ( + agent + and session_key + and hasattr(agent, "session_id") + and agent.session_id != session_id + ): _session_was_split = True logger.info( "Session split detected: %s → %s (compression)", - session_id, agent.session_id, + session_id, + agent.session_id, ) entry = self.session_store._entries.get(session_key) if entry: entry.session_id = agent.session_id self.session_store._save() - effective_session_id = getattr(agent, 'session_id', session_id) if agent else session_id + effective_session_id = ( + getattr(agent, "session_id", session_id) if agent else session_id + ) # When compression created a new session, the messages list was # shortened. Using the original history offset would produce an @@ -8800,7 +9644,10 @@ def _approval_notify_sync(approval_data: dict) -> None: if final_response and self._session_db: try: from agent.title_generator import maybe_auto_title - all_msgs = result_holder[0].get("messages", []) if result_holder[0] else [] + + all_msgs = ( + result_holder[0].get("messages", []) if result_holder[0] else [] + ) maybe_auto_title( self._session_db, effective_session_id, @@ -8814,8 +9661,12 @@ def _approval_notify_sync(approval_data: dict) -> None: return { "final_response": final_response, "last_reasoning": result.get("last_reasoning"), - "messages": result_holder[0].get("messages", []) if result_holder[0] else [], - "api_calls": result_holder[0].get("api_calls", 0) if result_holder[0] else 0, + "messages": result_holder[0].get("messages", []) + if result_holder[0] + else [], + "api_calls": result_holder[0].get("api_calls", 0) + if result_holder[0] + else 0, "tools": tools_holder[0] or [], "history_offset": _effective_history_offset, "last_prompt_tokens": _last_prompt_toks, @@ -8825,7 +9676,7 @@ def _approval_notify_sync(approval_data: dict) -> None: "session_id": effective_session_id, "response_previewed": result.get("response_previewed", False), } - + # Start progress message sender if enabled progress_task = None if tool_progress_enabled: @@ -8844,7 +9695,7 @@ async def _start_stream_consumer(): await asyncio.sleep(0.05) stream_task = asyncio.create_task(_start_stream_consumer()) - + # Track this agent as running for this session (for interrupt support) # We do this in a callback after the agent is created async def track_agent(): @@ -8855,9 +9706,9 @@ async def track_agent(): self._running_agents[session_key] = agent_holder[0] if self._draining: self._update_runtime_status("draining") - + tracking_task = asyncio.create_task(track_agent()) - + # Monitor for interrupts from the adapter (new messages arriving). # This is the PRIMARY interrupt path for regular text messages — # Level 1 (base.py) catches them before _handle_message() is reached, @@ -8882,7 +9733,9 @@ async def monitor_for_interrupt(): # Must use session_key (build_session_key output) — NOT # source.chat_id — because the adapter stores interrupt events # under the full session key. - if hasattr(_adapter, 'has_pending_interrupt') and _adapter.has_pending_interrupt(session_key): + if hasattr( + _adapter, "has_pending_interrupt" + ) and _adapter.has_pending_interrupt(session_key): agent = agent_holder[0] if agent: # Peek at the pending message text WITHOUT consuming it. @@ -8895,15 +9748,19 @@ async def monitor_for_interrupt(): # path finds it. _peek_event = _adapter._pending_messages.get(session_key) pending_text = _peek_event.text if _peek_event else None - logger.debug("Interrupt detected from adapter, signaling agent...") + logger.debug( + "Interrupt detected from adapter, signaling agent..." + ) agent.interrupt(pending_text) _interrupt_detected.set() break except asyncio.CancelledError: raise except Exception as _mon_err: - logger.debug("monitor_for_interrupt error (will retry): %s", _mon_err) - + logger.debug( + "monitor_for_interrupt error (will retry): %s", _mon_err + ) + interrupt_monitor = asyncio.create_task(monitor_for_interrupt()) # Periodic "still working" notifications for long-running tasks. @@ -8930,7 +9787,9 @@ async def _notify_long_running(): if _agent_ref and hasattr(_agent_ref, "get_activity_summary"): try: _a = _agent_ref.get_activity_summary() - _parts = [f"iteration {_a['api_call_count']}/{_a['max_iterations']}"] + _parts = [ + f"iteration {_a['api_call_count']}/{_a['max_iterations']}" + ] if _a.get("current_tool"): _parts.append(f"running: {_a['current_tool']}") else: @@ -8965,9 +9824,7 @@ async def _notify_long_running(): _agent_warning = _agent_warning_raw if _agent_warning_raw > 0 else None _warning_fired = False loop = asyncio.get_event_loop() - _executor_task = asyncio.ensure_future( - loop.run_in_executor(None, run_sync) - ) + _executor_task = asyncio.ensure_future(loop.run_in_executor(None, run_sync)) _inactivity_timeout = False _POLL_INTERVAL = 5.0 @@ -8988,10 +9845,15 @@ async def _notify_long_running(): if not _interrupt_detected.is_set() and session_key: _backup_adapter = self.adapters.get(source.platform) _backup_agent = agent_holder[0] - if (_backup_adapter and _backup_agent - and hasattr(_backup_adapter, 'has_pending_interrupt') - and _backup_adapter.has_pending_interrupt(session_key)): - _bp_event = _backup_adapter._pending_messages.get(session_key) + if ( + _backup_adapter + and _backup_agent + and hasattr(_backup_adapter, "has_pending_interrupt") + and _backup_adapter.has_pending_interrupt(session_key) + ): + _bp_event = _backup_adapter._pending_messages.get( + session_key + ) _bp_text = _bp_event.text if _bp_event else None logger.info( "Backup interrupt detected for session %s " @@ -9023,13 +9885,18 @@ async def _notify_long_running(): except Exception: pass # Staged warning: fire once before escalating to full timeout. - if (not _warning_fired and _agent_warning is not None - and _idle_secs >= _agent_warning): + if ( + not _warning_fired + and _agent_warning is not None + and _idle_secs >= _agent_warning + ): _warning_fired = True _warn_adapter = self.adapters.get(source.platform) if _warn_adapter: _elapsed_warn = int(_agent_warning // 60) or 1 - _remaining_mins = int((_agent_timeout - _agent_warning) // 60) or 1 + _remaining_mins = ( + int((_agent_timeout - _agent_warning) // 60) or 1 + ) try: await _warn_adapter.send( source.chat_id, @@ -9040,7 +9907,9 @@ async def _notify_long_running(): metadata=_status_thread_metadata, ) except Exception as _warn_err: - logger.debug("Inactivity warning send error: %s", _warn_err) + logger.debug( + "Inactivity warning send error: %s", _warn_err + ) if _idle_secs >= _agent_timeout: _inactivity_timeout = True break @@ -9048,10 +9917,15 @@ async def _notify_long_running(): if not _interrupt_detected.is_set() and session_key: _backup_adapter = self.adapters.get(source.platform) _backup_agent = agent_holder[0] - if (_backup_adapter and _backup_agent - and hasattr(_backup_adapter, 'has_pending_interrupt') - and _backup_adapter.has_pending_interrupt(session_key)): - _bp_event = _backup_adapter._pending_messages.get(session_key) + if ( + _backup_adapter + and _backup_agent + and hasattr(_backup_adapter, "has_pending_interrupt") + and _backup_adapter.has_pending_interrupt(session_key) + ): + _bp_event = _backup_adapter._pending_messages.get( + session_key + ) _bp_text = _bp_event.text if _bp_event else None logger.info( "Backup interrupt detected for session %s " @@ -9066,7 +9940,9 @@ async def _notify_long_running(): # Build a diagnostic summary from the agent's activity tracker. _timed_out_agent = agent_holder[0] _activity = {} - if _timed_out_agent and hasattr(_timed_out_agent, "get_activity_summary"): + if _timed_out_agent and hasattr( + _timed_out_agent, "get_activity_summary" + ): try: _activity = _timed_out_agent.get_activity_summary() except Exception: @@ -9081,8 +9957,12 @@ async def _notify_long_running(): logger.error( "Agent idle for %.0fs (timeout %.0fs) in session %s " "| last_activity=%s | iteration=%s/%s | tool=%s", - _secs_ago, _agent_timeout, session_key, - _last_desc, _iter_n, _iter_max, + _secs_ago, + _agent_timeout, + session_key, + _last_desc, + _iter_n, + _iter_max, _cur_tool or "none", ) @@ -9118,7 +9998,9 @@ async def _notify_long_running(): response = { "final_response": "\n".join(_diag_lines), - "messages": result_holder[0].get("messages", []) if result_holder[0] else [], + "messages": result_holder[0].get("messages", []) + if result_holder[0] + else [], "api_calls": _iter_n, "tools": tools_holder[0] or [], "history_offset": 0, @@ -9136,9 +10018,11 @@ async def _notify_long_running(): _agent = agent_holder[0] _result_for_fb = result_holder[0] _run_failed = _result_for_fb.get("failed") if _result_for_fb else False - if _agent is not None and hasattr(_agent, 'model') and not _run_failed: + if _agent is not None and hasattr(_agent, "model") and not _run_failed: _cfg_model = _resolve_gateway_model() - if _agent.model != _cfg_model and not self._is_intentional_model_switch(session_key, _agent.model): + if _agent.model != _cfg_model and not self._is_intentional_model_switch( + session_key, _agent.model + ): # Fallback activated on a successful run — evict cached # agent so the next message retries the primary model. self._evict_cached_agent(session_key) @@ -9146,18 +10030,27 @@ async def _notify_long_running(): # Check if we were interrupted OR have a queued message (/queue). result = result_holder[0] adapter = self.adapters.get(source.platform) - + # Get pending message from adapter. # Use session_key (not source.chat_id) to match adapter's storage keys. pending_event = None pending = None if result and adapter and session_key: pending_event = _dequeue_pending_event(adapter, session_key) - if result.get("interrupted") and not pending_event and result.get("interrupt_message"): + if ( + result.get("interrupted") + and not pending_event + and result.get("interrupt_message") + ): pending = result.get("interrupt_message") elif pending_event: - pending = pending_event.text or _build_media_placeholder(pending_event) - logger.debug("Processing queued message after agent completion: '%s...'", pending[:40]) + pending = pending_event.text or _build_media_placeholder( + pending_event + ) + logger.debug( + "Processing queued message after agent completion: '%s...'", + pending[:40], + ) # Safety net: if the pending text is a slash command (e.g. "/stop", # "/new"), discard it — commands should never be passed to the agent @@ -9166,10 +10059,13 @@ async def _notify_long_running(): # text leaks through the interrupt_message fallback. if pending and pending.strip().startswith("/"): _pending_parts = pending.strip().split(None, 1) - _pending_cmd_word = _pending_parts[0][1:].lower() if _pending_parts else "" + _pending_cmd_word = ( + _pending_parts[0][1:].lower() if _pending_parts else "" + ) if _pending_cmd_word: try: from hermes_cli.commands import resolve_command as _rc_pending + if _rc_pending(_pending_cmd_word): logger.info( "Discarding command '/%s' from pending queue — " @@ -9196,7 +10092,12 @@ async def _notify_long_running(): # Clear the adapter's interrupt event so the next _run_agent call # doesn't immediately re-trigger the interrupt before the new agent # even makes its first API call (this was causing an infinite loop). - if adapter and hasattr(adapter, '_active_sessions') and session_key and session_key in adapter._active_sessions: + if ( + adapter + and hasattr(adapter, "_active_sessions") + and session_key + and session_key in adapter._active_sessions + ): adapter._active_sessions[session_key].clear() # Cap recursion depth to prevent resource exhaustion when the @@ -9205,14 +10106,20 @@ async def _notify_long_running(): logger.warning( "Interrupt recursion depth %d reached for session %s — " "queueing message instead of recursing.", - _interrupt_depth, session_key, + _interrupt_depth, + session_key, ) adapter = self.adapters.get(source.platform) if adapter and pending_event: - merge_pending_message_event(adapter._pending_messages, session_key, pending_event) - elif adapter and hasattr(adapter, 'queue_message'): + merge_pending_message_event( + adapter._pending_messages, session_key, pending_event + ) + elif adapter and hasattr(adapter, "queue_message"): adapter.queue_message(session_key, pending) - return result_holder[0] or {"final_response": response, "messages": history} + return result_holder[0] or { + "final_response": response, + "messages": history, + } was_interrupted = result.get("interrupted") if not was_interrupted: @@ -9230,12 +10137,18 @@ async def _notify_long_running(): except asyncio.CancelledError: pass except Exception as e: - logger.debug("Stream consumer wait before queued message failed: %s", e) - _already_streamed = bool( + logger.debug( + "Stream consumer wait before queued message failed: %s", + e, + ) + _already_streamed = bool(result.get("response_previewed")) or bool( _sc and ( getattr(_sc, "final_response_sent", False) - or getattr(_sc, "already_sent", False) + or ( + streaming_enabled_holder[0] + and getattr(_sc, "already_sent", False) + ) ) ) first_response = result.get("final_response", "") @@ -9247,7 +10160,10 @@ async def _notify_long_running(): metadata=_status_thread_metadata, ) except Exception as e: - logger.warning("Failed to send first response before queued message: %s", e) + logger.warning( + "Failed to send first response before queued message: %s", + e, + ) # else: interrupted — discard the interrupted response ("Operation # interrupted." is just noise; the user already knows they sent a # new message). @@ -9294,7 +10210,7 @@ async def _notify_long_running(): await stream_task except asyncio.CancelledError: pass - + # Clean up tracking tracking_task.cancel() if session_key and session_key in self._running_agents: @@ -9303,7 +10219,7 @@ async def _notify_long_running(): self._running_agents_ts.pop(session_key, None) if self._draining: self._update_runtime_status("draining") - + # Wait for cancelled tasks for task in [progress_task, interrupt_monitor, tracking_task, _notify_task]: if task: @@ -9318,20 +10234,24 @@ async def _notify_long_running(): # message is new content the user hasn't seen, and it must reach # them even if streaming had sent earlier partial output. _sc = stream_consumer_holder[0] - if _sc and isinstance(response, dict) and not response.get("failed"): - if ( + if isinstance(response, dict) and not response.get("failed"): + if response.get("response_previewed"): + response["already_sent"] = True + elif _sc and ( getattr(_sc, "final_response_sent", False) - or getattr(_sc, "already_sent", False) + or (streaming_enabled_holder[0] and getattr(_sc, "already_sent", False)) ): response["already_sent"] = True - + return response -def _start_cron_ticker(stop_event: threading.Event, adapters=None, loop=None, interval: int = 60): +def _start_cron_ticker( + stop_event: threading.Event, adapters=None, loop=None, interval: int = 60 +): """ Background thread that ticks the cron scheduler at a regular interval. - + Runs inside the gateway process so cronjobs fire automatically without needing a separate `hermes cron daemon` or system cron entry. @@ -9344,8 +10264,8 @@ def _start_cron_ticker(stop_event: threading.Event, adapters=None, loop=None, in from cron.scheduler import tick as cron_tick from gateway.platforms.base import cleanup_image_cache, cleanup_document_cache - IMAGE_CACHE_EVERY = 60 # ticks — once per hour at default 60s interval - CHANNEL_DIR_EVERY = 5 # ticks — every 5 minutes + IMAGE_CACHE_EVERY = 60 # ticks — once per hour at default 60s interval + CHANNEL_DIR_EVERY = 5 # ticks — every 5 minutes logger.info("Cron ticker started (interval=%ds)", interval) tick_count = 0 @@ -9360,6 +10280,7 @@ def _start_cron_ticker(stop_event: threading.Event, adapters=None, loop=None, in if tick_count % CHANNEL_DIR_EVERY == 0 and adapters: try: from gateway.channel_directory import build_channel_directory + build_channel_directory(adapters) except Exception as e: logger.debug("Channel directory refresh error: %s", e) @@ -9368,13 +10289,17 @@ def _start_cron_ticker(stop_event: threading.Event, adapters=None, loop=None, in try: removed = cleanup_image_cache(max_age_hours=24) if removed: - logger.info("Image cache cleanup: removed %d stale file(s)", removed) + logger.info( + "Image cache cleanup: removed %d stale file(s)", removed + ) except Exception as e: logger.debug("Image cache cleanup error: %s", e) try: removed = cleanup_document_cache(max_age_hours=24) if removed: - logger.info("Document cache cleanup: removed %d stale file(s)", removed) + logger.info( + "Document cache cleanup: removed %d stale file(s)", removed + ) except Exception as e: logger.debug("Document cache cleanup error: %s", e) @@ -9382,14 +10307,18 @@ def _start_cron_ticker(stop_event: threading.Event, adapters=None, loop=None, in logger.info("Cron ticker stopped") -async def start_gateway(config: Optional[GatewayConfig] = None, replace: bool = False, verbosity: Optional[int] = 0) -> bool: +async def start_gateway( + config: Optional[GatewayConfig] = None, + replace: bool = False, + verbosity: Optional[int] = 0, +) -> bool: """ Start the gateway and run until interrupted. - + This is the main entry point for running the gateway. Returns True if the gateway ran successfully, False if it failed to start. A False return causes a non-zero exit code so systemd can auto-restart. - + Args: config: Optional gateway configuration override. replace: If True, kill any existing gateway instance before starting. @@ -9403,6 +10332,7 @@ async def start_gateway(config: Optional[GatewayConfig] = None, replace: bool = # allow concurrent instances without tripping this guard. import time as _time from gateway.status import get_running_pid, remove_pid_file, terminate_pid + existing_pid = get_running_pid() if existing_pid is not None and existing_pid != os.getpid(): if replace: @@ -9444,9 +10374,12 @@ async def start_gateway(config: Optional[GatewayConfig] = None, replace: bool = # leaving stale lock files that block the new gateway from starting. try: from gateway.status import release_all_scoped_locks + _released = release_all_scoped_locks() if _released: - logger.info("Released %d stale scoped lock(s) from old gateway.", _released) + logger.info( + "Released %d stale scoped lock(s) from old gateway.", _released + ) except Exception: pass else: @@ -9454,7 +10387,8 @@ async def start_gateway(config: Optional[GatewayConfig] = None, replace: bool = logger.error( "Another gateway instance is already running (PID %d, HERMES_HOME=%s). " "Use 'hermes gateway restart' to replace it, or 'hermes gateway stop' first.", - existing_pid, hermes_home, + existing_pid, + hermes_home, ) print( f"\n❌ Gateway already running (PID {existing_pid}).\n" @@ -9467,6 +10401,7 @@ async def start_gateway(config: Optional[GatewayConfig] = None, replace: bool = # Sync bundled skills on gateway start (fast -- skips unchanged) try: from tools.skills_sync import sync_skills + sync_skills(quiet=True) except Exception: pass @@ -9475,6 +10410,7 @@ async def start_gateway(config: Optional[GatewayConfig] = None, replace: bool = # and gateway.log (INFO+, gateway-component records only). # Idempotent, so repeated calls from AIAgent.__init__ won't duplicate. from hermes_logging import setup_logging + setup_logging(hermes_home=_hermes_home, mode="gateway") # Optional stderr handler — level driven by -v/-q flags on the CLI. @@ -9485,17 +10421,21 @@ async def start_gateway(config: Optional[GatewayConfig] = None, replace: bool = if verbosity is not None: from agent.redact import RedactingFormatter - _stderr_level = {0: logging.WARNING, 1: logging.INFO}.get(verbosity, logging.DEBUG) + _stderr_level = {0: logging.WARNING, 1: logging.INFO}.get( + verbosity, logging.DEBUG + ) _stderr_handler = logging.StreamHandler() _stderr_handler.setLevel(_stderr_level) - _stderr_handler.setFormatter(RedactingFormatter('%(levelname)s %(name)s: %(message)s')) + _stderr_handler.setFormatter( + RedactingFormatter("%(levelname)s %(name)s: %(message)s") + ) logging.getLogger().addHandler(_stderr_handler) # Lower root logger level if needed so DEBUG records can reach the handler if _stderr_level < logging.getLogger().level: logging.getLogger().setLevel(_stderr_level) runner = GatewayRunner(config) - + # Track whether a signal initiated the shutdown (vs. internal request). # When an unexpected SIGTERM kills the gateway, we exit non-zero so # systemd's Restart=on-failure revives the process. systemctl stop @@ -9513,12 +10453,16 @@ def shutdown_signal_handler(): # a stale detached subprocess, etc.). try: import subprocess as _sp + _ps = _sp.run( ["ps", "aux"], - capture_output=True, text=True, timeout=3, + capture_output=True, + text=True, + timeout=3, ) _hermes_procs = [ - line for line in _ps.stdout.splitlines() + line + for line in _ps.stdout.splitlines() if ("hermes" in line.lower() or "gateway" in line.lower()) and str(os.getpid()) not in line.split()[1:2] # exclude self ] @@ -9535,7 +10479,7 @@ def shutdown_signal_handler(): def restart_signal_handler(): runner.request_restart(detached=False, via_service=True) - + loop = asyncio.get_event_loop() if threading.current_thread() is threading.main_thread(): for sig in (signal.SIGINT, signal.SIGTERM): @@ -9550,7 +10494,7 @@ def restart_signal_handler(): pass else: logger.info("Skipping signal handlers (not running in main thread).") - + # Start the gateway success = await runner.start() if not success: @@ -9559,13 +10503,14 @@ def restart_signal_handler(): if runner.exit_reason: logger.error("Gateway exiting cleanly: %s", runner.exit_reason) return True - + # Write PID file so CLI can detect gateway is running import atexit from gateway.status import write_pid_file, remove_pid_file + write_pid_file() atexit.register(remove_pid_file) - + # Start background cron ticker so scheduled jobs fire automatically. # Pass the event loop so cron delivery can use live adapters (E2EE support). cron_stop = threading.Event() @@ -9577,7 +10522,7 @@ def restart_signal_handler(): name="cron-ticker", ) cron_thread.start() - + # Wait for shutdown await runner.wait_for_shutdown() @@ -9585,7 +10530,7 @@ def restart_signal_handler(): if runner.exit_reason: logger.error("Gateway exiting with failure: %s", runner.exit_reason) return False - + # Stop cron ticker cleanly cron_stop.set() cron_thread.join(timeout=5) @@ -9593,6 +10538,7 @@ def restart_signal_handler(): # Close MCP server connections try: from tools.mcp_tool import shutdown_mcp_servers + shutdown_mcp_servers() except Exception: pass @@ -9621,20 +10567,23 @@ def restart_signal_handler(): def main(): """CLI entry point for the gateway.""" import argparse - - parser = argparse.ArgumentParser(description="Hermes Gateway - Multi-platform messaging") + + parser = argparse.ArgumentParser( + description="Hermes Gateway - Multi-platform messaging" + ) parser.add_argument("--config", "-c", help="Path to gateway config file") parser.add_argument("--verbose", "-v", action="store_true", help="Verbose output") - + args = parser.parse_args() - + config = None if args.config: import json + with open(args.config, encoding="utf-8") as f: data = json.load(f) config = GatewayConfig.from_dict(data) - + # Run the gateway - exit with code 1 if no platforms connected, # so systemd Restart=on-failure will retry on transient errors (e.g. DNS) success = asyncio.run(start_gateway(config)) diff --git a/gateway/stream_consumer.py b/gateway/stream_consumer.py index e6d96c802d29f..38ef6d7b5e9bc 100644 --- a/gateway/stream_consumer.py +++ b/gateway/stream_consumer.py @@ -40,8 +40,9 @@ @dataclass class StreamConsumerConfig: """Runtime config for a single stream consumer instance.""" - edit_interval: float = 1.0 - buffer_threshold: int = 40 + + edit_interval: float = 0.8 + buffer_threshold: int = 24 cursor: str = " ▉" @@ -68,12 +69,20 @@ class GatewayStreamConsumer: # Must stay in sync with cli.py _OPEN_TAGS/_CLOSE_TAGS and # run_agent.py _strip_think_blocks() tag variants. _OPEN_THINK_TAGS = ( - "<REASONING_SCRATCHPAD>", "<think>", "<reasoning>", - "<THINKING>", "<thinking>", "<thought>", + "<REASONING_SCRATCHPAD>", + "<think>", + "<reasoning>", + "<THINKING>", + "<thinking>", + "<thought>", ) _CLOSE_THINK_TAGS = ( - "</REASONING_SCRATCHPAD>", "</think>", "</reasoning>", - "</THINKING>", "</thinking>", "</thought>", + "</REASONING_SCRATCHPAD>", + "</think>", + "</reasoning>", + "</THINKING>", + "</thinking>", + "</thought>", ) def __init__( @@ -91,12 +100,14 @@ def __init__( self._accumulated = "" self._message_id: Optional[str] = None self._already_sent = False - self._edit_supported = True # Disabled when progressive edits are no longer usable + self._edit_supported = ( + True # Disabled when progressive edits are no longer usable + ) self._last_edit_time = 0.0 - self._last_sent_text = "" # Track last-sent text to skip redundant edits + self._last_sent_text = "" # Track last-sent text to skip redundant edits self._fallback_final_send = False self._fallback_prefix = "" - self._flood_strikes = 0 # Consecutive flood-control edit failures + self._flood_strikes = 0 # Consecutive flood-control edit failures self._current_edit_interval = self.cfg.edit_interval # Adaptive backoff self._final_response_sent = False @@ -181,7 +192,7 @@ def _filter_and_accumulate(self, text: str) -> None: if best_len: # Found closing tag — discard block, process remainder self._in_think_block = False - buf = buf[best_idx + best_len:] + buf = buf[best_idx + best_len :] else: # No closing tag yet — hold tail that could be a # partial closing tag prefix, discard the rest. @@ -212,12 +223,11 @@ def _filter_and_accumulate(self, text: str) -> None: last_nl = preceding.rfind("\n") if last_nl == -1: is_boundary = ( - (not self._accumulated - or self._accumulated.endswith("\n")) - and preceding.strip() == "" - ) + not self._accumulated + or self._accumulated.endswith("\n") + ) and preceding.strip() == "" else: - is_boundary = preceding[last_nl + 1:].strip() == "" + is_boundary = preceding[last_nl + 1 :].strip() == "" if is_boundary and (best_idx == -1 or idx < best_idx): best_idx = idx @@ -229,7 +239,7 @@ def _filter_and_accumulate(self, text: str) -> None: # Emit text before the tag, enter think block self._accumulated += buf[:best_idx] self._in_think_block = True - buf = buf[best_idx + best_len:] + buf = buf[best_idx + best_len :] else: # No opening tag — check for a partial tag at the tail held_back = 0 @@ -275,7 +285,11 @@ async def run(self) -> None: if item is _NEW_SEGMENT: got_segment_break = True break - if isinstance(item, tuple) and len(item) == 2 and item[0] is _COMMENTARY: + if ( + isinstance(item, tuple) + and len(item) == 2 + and item[0] is _COMMENTARY + ): commentary_text = item[1] break self._filter_and_accumulate(item) @@ -295,8 +309,7 @@ async def run(self) -> None: got_done or got_segment_break or commentary_text is not None - or (elapsed >= self._current_edit_interval - and self._accumulated) + or (elapsed >= self._current_edit_interval and self._accumulated) or len(self._accumulated) >= self.cfg.buffer_threshold ) @@ -354,7 +367,11 @@ async def run(self) -> None: self._last_sent_text = "" display_text = self._accumulated - if not got_done and not got_segment_break and commentary_text is None: + if ( + not got_done + and not got_segment_break + and commentary_text is None + ): display_text += self.cfg.cursor current_update_visible = await self._send_or_edit(display_text) @@ -371,9 +388,13 @@ async def run(self) -> None: elif current_update_visible: self._final_response_sent = True elif self._message_id: - self._final_response_sent = await self._send_or_edit(self._accumulated) + self._final_response_sent = await self._send_or_edit( + self._accumulated + ) elif not self._already_sent: - self._final_response_sent = await self._send_or_edit(self._accumulated) + self._final_response_sent = await self._send_or_edit( + self._accumulated + ) return if commentary_text is not None: @@ -422,7 +443,7 @@ async def run(self) -> None: # Pattern to strip MEDIA:<path> tags (including optional surrounding quotes). # Matches the simple cleanup regex used by the non-streaming path in # gateway/platforms/base.py for post-processing. - _MEDIA_RE = re.compile(r'''[`"']?MEDIA:\s*\S+[`"']?''') + _MEDIA_RE = re.compile(r"""[`"']?MEDIA:\s*\S+[`"']?""") @staticmethod def _clean_for_display(text: str) -> str: @@ -440,11 +461,13 @@ def _clean_for_display(text: str) -> str: cleaned = text.replace("[[audio_as_voice]]", "") cleaned = GatewayStreamConsumer._MEDIA_RE.sub("", cleaned) # Collapse excessive blank lines left behind by removed tags - cleaned = re.sub(r'\n{3,}', '\n\n', cleaned) + cleaned = re.sub(r"\n{3,}", "\n\n", cleaned) # Strip trailing whitespace/newlines but preserve leading content return cleaned.rstrip() - async def _send_new_chunk(self, text: str, reply_to_id: Optional[str]) -> Optional[str]: + async def _send_new_chunk( + self, text: str, reply_to_id: Optional[str] + ) -> Optional[str]: """Send a new message chunk, optionally threaded to a previous message. Returns the message_id so callers can thread subsequent chunks. @@ -476,14 +499,14 @@ def _visible_prefix(self) -> str: """Return the visible text already shown in the streamed message.""" prefix = self._last_sent_text or "" if self.cfg.cursor and prefix.endswith(self.cfg.cursor): - prefix = prefix[:-len(self.cfg.cursor)] + prefix = prefix[: -len(self.cfg.cursor)] return self._clean_for_display(prefix) def _continuation_text(self, final_text: str) -> str: """Return only the part of final_text the user has not already seen.""" prefix = self._fallback_prefix or self._visible_prefix() if prefix and final_text.startswith(prefix): - return final_text[len(prefix):].lstrip() + return final_text[len(prefix) :].lstrip() return final_text @staticmethod @@ -536,9 +559,7 @@ async def _send_fallback_final(self, text: str) -> None: if result.success: break if attempt == 0 and self._is_flood_error(result): - logger.debug( - "Flood control on fallback send, retrying in 3s" - ) + logger.debug("Flood control on fallback send, retrying in 3s") await asyncio.sleep(3.0) else: break # non-flood error or second attempt failed @@ -647,10 +668,12 @@ async def _send_or_edit(self, text: str) -> bool: # ▉ block character renders as a visible white box ("tofu"). # Existing messages (edits) are unaffected — only first sends gated. _MIN_NEW_MSG_CHARS = 4 - if (self._message_id is None - and self.cfg.cursor - and self.cfg.cursor in text - and len(_visible_stripped) < _MIN_NEW_MSG_CHARS): + if ( + self._message_id is None + and self.cfg.cursor + and self.cfg.cursor in text + and len(_visible_stripped) < _MIN_NEW_MSG_CHARS + ): return True # too short for a standalone message — accumulate more try: if self._message_id is not None: @@ -678,7 +701,8 @@ async def _send_or_edit(self, text: str) -> bool: if self._is_flood_error(result): self._flood_strikes += 1 self._current_edit_interval = min( - self._current_edit_interval * 2, 10.0, + self._current_edit_interval * 2, + 10.0, ) logger.debug( "Flood control on edit (strike %d/%d), " diff --git a/tests/gateway/test_display_config.py b/tests/gateway/test_display_config.py index 2192d67bc98d9..dadbd4baad0c6 100644 --- a/tests/gateway/test_display_config.py +++ b/tests/gateway/test_display_config.py @@ -1,4 +1,5 @@ """Tests for gateway.display_config — per-platform display/verbosity resolver.""" + import pytest @@ -6,6 +7,7 @@ # Resolver: resolution order # --------------------------------------------------------------------------- + class TestResolveDisplaySetting: """resolve_display_setting() resolves with correct priority.""" @@ -41,8 +43,8 @@ def test_platform_default_when_no_user_config(self): # Empty config — should get built-in defaults config = {} - # Telegram defaults to tier_high → "all" - assert resolve_display_setting(config, "telegram", "tool_progress") == "all" + # Telegram has a low-latency default profile → "new" + assert resolve_display_setting(config, "telegram", "tool_progress") == "new" # Email defaults to tier_minimal → "off" assert resolve_display_setting(config, "email", "tool_progress") == "off" @@ -52,7 +54,10 @@ def test_global_default_for_unknown_platform(self): config = {} # Unknown platform, no config → global default "all" - assert resolve_display_setting(config, "unknown_platform", "tool_progress") == "all" + assert ( + resolve_display_setting(config, "unknown_platform", "tool_progress") + == "all" + ) def test_fallback_parameter_used_last(self): """Explicit fallback is used when nothing else matches.""" @@ -60,7 +65,9 @@ def test_fallback_parameter_used_last(self): config = {} # "nonexistent_key" isn't in any defaults - result = resolve_display_setting(config, "telegram", "nonexistent_key", "my_fallback") + result = resolve_display_setting( + config, "telegram", "nonexistent_key", "my_fallback" + ) assert result == "my_fallback" def test_platform_override_only_affects_that_platform(self): @@ -83,6 +90,7 @@ def test_platform_override_only_affects_that_platform(self): # Backward compatibility: tool_progress_overrides # --------------------------------------------------------------------------- + class TestBackwardCompat: """Legacy tool_progress_overrides is still respected as a fallback.""" @@ -132,6 +140,7 @@ def test_legacy_overrides_only_for_tool_progress(self): # YAML normalisation # --------------------------------------------------------------------------- + class TestYAMLNormalisation: """YAML 1.1 quirks (bare off → False, on → True) are handled.""" @@ -175,15 +184,16 @@ def test_platform_override_false_tool_progress(self): # Built-in platform defaults (tier system) # --------------------------------------------------------------------------- + class TestPlatformDefaults: """Built-in defaults reflect platform capability tiers.""" def test_high_tier_platforms(self): - """Telegram and Discord default to 'all' tool progress.""" + """High-tier defaults: Telegram='new' for latency, Discord='all'.""" from gateway.display_config import resolve_display_setting - for plat in ("telegram", "discord"): - assert resolve_display_setting({}, plat, "tool_progress") == "all", plat + assert resolve_display_setting({}, "telegram", "tool_progress") == "new" + assert resolve_display_setting({}, "discord", "tool_progress") == "all" def test_medium_tier_platforms(self): """Slack, Mattermost, Matrix default to 'new' tool progress.""" @@ -224,6 +234,7 @@ def test_high_tier_streaming_defaults_to_none(self): # Config migration: tool_progress_overrides → display.platforms # --------------------------------------------------------------------------- + class TestConfigMigration: """Version 16 migration moves tool_progress_overrides into display.platforms.""" @@ -247,6 +258,7 @@ def test_migration_creates_platforms_entries(self, tmp_path, monkeypatch): # Re-import to pick up the new HERMES_HOME import importlib import hermes_cli.config as cfg_mod + importlib.reload(cfg_mod) result = cfg_mod.migrate_config(interactive=False, quiet=True) @@ -256,7 +268,9 @@ def test_migration_creates_platforms_entries(self, tmp_path, monkeypatch): assert platforms.get("signal", {}).get("tool_progress") == "off" assert platforms.get("telegram", {}).get("tool_progress") == "all" - def test_migration_preserves_existing_platforms_entries(self, tmp_path, monkeypatch): + def test_migration_preserves_existing_platforms_entries( + self, tmp_path, monkeypatch + ): """Existing display.platforms entries are NOT overwritten by migration.""" import yaml @@ -273,6 +287,7 @@ def test_migration_preserves_existing_platforms_entries(self, tmp_path, monkeypa monkeypatch.setenv("HERMES_HOME", str(tmp_path)) import importlib import hermes_cli.config as cfg_mod + importlib.reload(cfg_mod) cfg_mod.migrate_config(interactive=False, quiet=True) @@ -285,6 +300,7 @@ def test_migration_preserves_existing_platforms_entries(self, tmp_path, monkeypa # Streaming per-platform (None = follow global) # --------------------------------------------------------------------------- + class TestStreamingPerPlatform: """Streaming per-platform override semantics.""" diff --git a/tests/gateway/test_proxy_mode.py b/tests/gateway/test_proxy_mode.py index f3024cb09f150..b41c833472887 100644 --- a/tests/gateway/test_proxy_mode.py +++ b/tests/gateway/test_proxy_mode.py @@ -8,7 +8,7 @@ import pytest from gateway.config import Platform, StreamingConfig -from gateway.run import GatewayRunner +from gateway.run import GatewayRunner, _is_gateway_streaming_enabled from gateway.session import SessionSource @@ -37,6 +37,30 @@ def _make_source(platform=Platform.MATRIX): ) +class TestGatewayStreamingResolver: + """Unit tests for gateway streaming gate resolution helper.""" + + def test_none_follows_global_enabled(self): + cfg = StreamingConfig(enabled=True, transport="edit") + assert _is_gateway_streaming_enabled(cfg, None) is True + + def test_none_follows_global_disabled(self): + cfg = StreamingConfig(enabled=False, transport="edit") + assert _is_gateway_streaming_enabled(cfg, None) is False + + def test_platform_true_overrides_global_disabled(self): + cfg = StreamingConfig(enabled=False, transport="edit") + assert _is_gateway_streaming_enabled(cfg, True) is True + + def test_transport_off_is_global_kill_switch(self): + cfg = StreamingConfig(enabled=True, transport="off") + assert _is_gateway_streaming_enabled(cfg, True) is False + + def test_platform_false_overrides_global_enabled(self): + cfg = StreamingConfig(enabled=True, transport="edit") + assert _is_gateway_streaming_enabled(cfg, False) is False + + class _FakeSSEResponse: """Simulates an aiohttp response with SSE streaming.""" @@ -306,7 +330,9 @@ async def test_skips_tool_messages_in_history(self, monkeypatch): resp = _FakeSSEResponse( status=200, - sse_chunks=[b'data: {"choices":[{"delta":{"content":"ok"}}]}\n\ndata: [DONE]\n\n'], + sse_chunks=[ + b'data: {"choices":[{"delta":{"content":"ok"}}]}\n\ndata: [DONE]\n\n' + ], ) session = _FakeSession(resp) @@ -344,7 +370,9 @@ async def test_result_shape_matches_run_agent(self, monkeypatch): resp = _FakeSSEResponse( status=200, - sse_chunks=[b'data: {"choices":[{"delta":{"content":"answer"}}]}\n\ndata: [DONE]\n\n'], + sse_chunks=[ + b'data: {"choices":[{"delta":{"content":"answer"}}]}\n\ndata: [DONE]\n\n' + ], ) session = _FakeSession(resp) @@ -354,7 +382,10 @@ async def test_result_shape_matches_run_agent(self, monkeypatch): result = await runner._run_agent_via_proxy( message="hi", context_prompt="", - history=[{"role": "user", "content": "prev"}, {"role": "assistant", "content": "ok"}], + history=[ + {"role": "user", "content": "prev"}, + {"role": "assistant", "content": "ok"}, + ], source=source, session_id="sess-123", ) @@ -379,7 +410,9 @@ async def test_no_auth_header_without_key(self, monkeypatch): resp = _FakeSSEResponse( status=200, - sse_chunks=[b'data: {"choices":[{"delta":{"content":"ok"}}]}\n\ndata: [DONE]\n\n'], + sse_chunks=[ + b'data: {"choices":[{"delta":{"content":"ok"}}]}\n\ndata: [DONE]\n\n' + ], ) session = _FakeSession(resp) @@ -405,7 +438,9 @@ async def test_no_system_message_when_context_empty(self, monkeypatch): resp = _FakeSSEResponse( status=200, - sse_chunks=[b'data: {"choices":[{"delta":{"content":"ok"}}]}\n\ndata: [DONE]\n\n'], + sse_chunks=[ + b'data: {"choices":[{"delta":{"content":"ok"}}]}\n\ndata: [DONE]\n\n' + ], ) session = _FakeSession(resp) @@ -432,6 +467,7 @@ class TestEnvVarRegistration: def test_proxy_url_in_optional_env_vars(self): from hermes_cli.config import OPTIONAL_ENV_VARS + assert "GATEWAY_PROXY_URL" in OPTIONAL_ENV_VARS info = OPTIONAL_ENV_VARS["GATEWAY_PROXY_URL"] assert info["category"] == "messaging" @@ -439,6 +475,7 @@ def test_proxy_url_in_optional_env_vars(self): def test_proxy_key_in_optional_env_vars(self): from hermes_cli.config import OPTIONAL_ENV_VARS + assert "GATEWAY_PROXY_KEY" in OPTIONAL_ENV_VARS info = OPTIONAL_ENV_VARS["GATEWAY_PROXY_KEY"] assert info["category"] == "messaging" From c8361c1c56d5d164f566cc26fe32a03882ba82ce Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 15 Apr 2026 22:51:41 +0800 Subject: [PATCH 3/3] fix(gateway): harden env float parsing and unify streaming defaults to single source --- gateway/config.py | 28 ++++++++++++++++++---------- gateway/platforms/telegram.py | 10 +++++++++- gateway/stream_consumer.py | 18 ++++++++++++++---- 3 files changed, 41 insertions(+), 15 deletions(-) diff --git a/gateway/config.py b/gateway/config.py index 0f2bc03bea486..ffd1437f70342 100644 --- a/gateway/config.py +++ b/gateway/config.py @@ -196,19 +196,23 @@ def from_dict(cls, data: Dict[str, Any]) -> "PlatformConfig": ) +# Canonical streaming defaults — single source of truth. +# ``StreamConsumerConfig`` in gateway/stream_consumer.py imports these so +# there is exactly one place to change the out-of-the-box edit rhythm. +DEFAULT_STREAMING_EDIT_INTERVAL: float = 0.8 +DEFAULT_STREAMING_BUFFER_THRESHOLD: int = 24 +DEFAULT_STREAMING_CURSOR: str = " ▉" + + @dataclass class StreamingConfig: """Configuration for real-time token streaming to messaging platforms.""" enabled: bool = False transport: str = "edit" # "edit" (progressive editMessageText) or "off" - edit_interval: float = ( - 0.8 # Seconds between message edits (balanced for lower latency + flood safety) - ) - buffer_threshold: int = ( - 24 # Chars before forcing an edit (faster first visible token) - ) - cursor: str = " ▉" # Cursor shown during streaming + edit_interval: float = DEFAULT_STREAMING_EDIT_INTERVAL + buffer_threshold: int = DEFAULT_STREAMING_BUFFER_THRESHOLD + cursor: str = DEFAULT_STREAMING_CURSOR def to_dict(self) -> Dict[str, Any]: return { @@ -226,9 +230,13 @@ def from_dict(cls, data: Dict[str, Any]) -> "StreamingConfig": return cls( enabled=data.get("enabled", False), transport=data.get("transport", "edit"), - edit_interval=float(data.get("edit_interval", 0.8)), - buffer_threshold=int(data.get("buffer_threshold", 24)), - cursor=data.get("cursor", " ▉"), + edit_interval=float( + data.get("edit_interval", DEFAULT_STREAMING_EDIT_INTERVAL) + ), + buffer_threshold=int( + data.get("buffer_threshold", DEFAULT_STREAMING_BUFFER_THRESHOLD) + ), + cursor=data.get("cursor", DEFAULT_STREAMING_CURSOR), ) diff --git a/gateway/platforms/telegram.py b/gateway/platforms/telegram.py index 22e1040ee18ff..7cee528d99f94 100644 --- a/gateway/platforms/telegram.py +++ b/gateway/platforms/telegram.py @@ -147,12 +147,20 @@ def _env_float_clamped( min_value: Optional[float] = None, max_value: Optional[float] = None, ) -> float: - """Read a float env var with bounds and sane fallback.""" + """Read a float env var, reject non-finite values, and clamp to bounds. + + Guarantees the returned value is a finite number usable directly in + ``asyncio.sleep()`` and similar APIs that reject NaN / Inf. + """ + import math + raw = os.getenv(name) try: value = float(raw) if raw is not None else float(default) except (TypeError, ValueError): value = float(default) + if not math.isfinite(value): + value = float(default) if min_value is not None: value = max(value, min_value) if max_value is not None: diff --git a/gateway/stream_consumer.py b/gateway/stream_consumer.py index 38ef6d7b5e9bc..f72c57147b41d 100644 --- a/gateway/stream_consumer.py +++ b/gateway/stream_consumer.py @@ -23,6 +23,12 @@ from dataclasses import dataclass from typing import Any, Optional +from gateway.config import ( + DEFAULT_STREAMING_BUFFER_THRESHOLD, + DEFAULT_STREAMING_CURSOR, + DEFAULT_STREAMING_EDIT_INTERVAL, +) + logger = logging.getLogger("gateway.stream_consumer") # Sentinel to signal the stream is complete @@ -39,11 +45,15 @@ @dataclass class StreamConsumerConfig: - """Runtime config for a single stream consumer instance.""" + """Runtime config for a single stream consumer instance. + + Defaults are re-exported from ``gateway.config`` so that + ``StreamingConfig`` remains the single source of truth. + """ - edit_interval: float = 0.8 - buffer_threshold: int = 24 - cursor: str = " ▉" + edit_interval: float = DEFAULT_STREAMING_EDIT_INTERVAL + buffer_threshold: int = DEFAULT_STREAMING_BUFFER_THRESHOLD + cursor: str = DEFAULT_STREAMING_CURSOR class GatewayStreamConsumer: