diff --git a/.github/workflows/docker.yml b/.github/workflows/docker.yml index ed89185acf..7e47b1db69 100644 --- a/.github/workflows/docker.yml +++ b/.github/workflows/docker.yml @@ -107,7 +107,9 @@ jobs: version: "0.9.28" - name: Set up Python 3.11 (for docker tests) - run: uv python install 3.11 + uses: ./.github/actions/retry + with: + command: uv python install 3.11 - name: Install Python dependencies (for docker tests) # ``dev`` extra pulls in pytest, pytest-asyncio — diff --git a/.github/workflows/e2e-desktop.yml b/.github/workflows/e2e-desktop.yml index b951c60ead..8d0c037a37 100644 --- a/.github/workflows/e2e-desktop.yml +++ b/.github/workflows/e2e-desktop.yml @@ -66,8 +66,12 @@ jobs: cache-dependency-glob: | pyproject.toml uv.lock + - name: Set up Python 3.11 - run: uv python install 3.11 + uses: ./.github/actions/retry + with: + command: uv python install 3.11 + - name: Install Python dependencies uses: ./.github/actions/retry with: diff --git a/.github/workflows/lint.yml b/.github/workflows/lint.yml index 3ae120f71f..7a49a62203 100644 --- a/.github/workflows/lint.yml +++ b/.github/workflows/lint.yml @@ -163,10 +163,18 @@ jobs: - name: Checkout code uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2 - - name: Set up Python - uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v5 + - name: Install uv + uses: astral-sh/setup-uv@fac544c07dec837d0ccb6301d7b5580bf5edae39 # 8.2.0 with: - python-version: "3.11" + # Pinned: unpinned setup-uv fetches a 'latest' manifest from + # raw.githubusercontent.com every job; transient fetch failures + # fail the job (2026-07-28 incident). Keep in sync with tests.yml. + version: "0.9.28" + + - name: Set up Python 3.11 + uses: ./.github/actions/retry + with: + command: uv python install 3.11 - name: Run footgun checker run: python scripts/check-windows-footguns.py --all diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml index 353e86439e..3888e8d44f 100644 --- a/.github/workflows/tests.yml +++ b/.github/workflows/tests.yml @@ -90,7 +90,9 @@ jobs: uv.lock - name: Set up Python 3.11 - run: uv python install 3.11 + uses: ./.github/actions/retry + with: + command: uv python install 3.11 - name: Install dependencies # `uv sync --locked` installs the exact pinned set from uv.lock (and diff --git a/acp_adapter/permissions.py b/acp_adapter/permissions.py index 5f29a96725..b10b2a169e 100644 --- a/acp_adapter/permissions.py +++ b/acp_adapter/permissions.py @@ -158,9 +158,16 @@ def make_approval_callback( try: response = future.result(timeout=timeout) - except (FutureTimeout, Exception) as exc: + except FutureTimeout: future.cancel() - logger.warning("Permission request timed out or failed: %s", exc) + logger.warning("Permission request timed out after %ss", timeout) + # Distinct from an explicit deny: the client never answered. + # tools.approval callers report this as "timed out without user + # response" instead of a user denial. + return "timeout" + except Exception as exc: + future.cancel() + logger.warning("Permission request failed: %s", exc) return "deny" if response is None: diff --git a/agent/agent_init.py b/agent/agent_init.py index db8f026f63..649d5338a9 100644 --- a/agent/agent_init.py +++ b/agent/agent_init.py @@ -209,7 +209,18 @@ def _context_route_mismatch( if active_route: configured_routes = _provider_default_routes(configured_provider) - return not configured_routes or active_route not in configured_routes + if configured_routes: + return active_route not in configured_routes + # Named/custom providers have no catalog default routes. An empty + # configured URL with a matching provider identity is still the same + # route — agent_init fills base_url from custom_providers before this + # check, but gateway display/hygiene paths historically compared the + # raw empty model.base_url and falsely dropped model.context_length, + # falling through to family defaults (e.g. qwen → 131072) on Discord + # session-reset banners while /status still showed the config pin. + if active_provider and configured_provider == active_provider: + return False + return True return bool( configured_provider and active_provider @@ -1580,6 +1591,17 @@ def init_agent( "reasoning_config": reasoning_config, "max_tokens": max_tokens, } + # Persist a process-scoped --yolo launch into the session row so a later + # `hermes --resume ` can restore the bypass (CLI resume paths read + # model_config.yolo_mode back via SessionDB.session_yolo_enabled). + # Session-scoped /yolo toggles persist separately through + # SessionDB.set_session_yolo at toggle time. + try: + from tools.approval import _YOLO_MODE_FROZEN + if _YOLO_MODE_FROZEN: + agent._session_init_model_config["yolo_mode"] = True + except Exception: + pass # In-memory todo list for task planning (one per agent/session) from tools.todo_tool import TodoStore diff --git a/agent/agent_runtime_helpers.py b/agent/agent_runtime_helpers.py index 579314018c..1bdce7a498 100644 --- a/agent/agent_runtime_helpers.py +++ b/agent/agent_runtime_helpers.py @@ -36,7 +36,7 @@ from hermes_cli.timeouts import get_provider_request_timeout from agent.prompt_builder import format_steer_marker from agent.tool_dispatch_helpers import _trajectory_normalize_msg, make_tool_result_message from agent.trajectory import convert_scratchpad_to_think -from agent.credential_pool import STATUS_EXHAUSTED +from agent.credential_pool import STATUS_EXHAUSTED, credential_pool_matches_provider from agent.error_classifier import FailoverReason from agent.turn_context import drop_stale_api_content from utils import base_url_host_matches, base_url_hostname, env_var_enabled, atomic_json_write @@ -52,6 +52,45 @@ logger = logging.getLogger(__name__) _MAX_AUTH_REFRESH_ATTEMPTS = 2 +_REASONING_TAG_NAMES = ("think", "thinking", "reasoning", "REASONING_SCRATCHPAD", "thought") +_TOOL_CALL_TAG_NAMES = ("tool_call", "tool_calls", "tool_result", "function_call", "function_calls") + +_REASONING_BLOCK_PATTERNS = tuple( + re.compile(rf"<{name}>.*?", re.DOTALL | re.IGNORECASE) + for name in _REASONING_TAG_NAMES +) + +_TOOL_CALL_BLOCK_PATTERNS = tuple( + re.compile(rf"<{name}\b[^>]*>.*?", re.DOTALL | re.IGNORECASE) + for name in _TOOL_CALL_TAG_NAMES +) + +# Named blocks — see strip_think_blocks step 1c for the +# full rationale (sentence-boundary lookbehind + tempered-dot body so a plain +# prose mention of "function" is never eaten). +_NAMED_FUNCTION_BLOCK_PATTERN = re.compile( + r'(?:(?<=^)|(?<=[\n\r.!?:]))[ \t]*' + r']*\bname\s*=[^>]*>' + r'(?:(?:(?!).)*)', + re.DOTALL | re.IGNORECASE, +) + +_UNTERMINATED_REASONING_BLOCK_PATTERN = re.compile( + rf'(?:^|\n)[ \t]*<(?:{"|".join(_REASONING_TAG_NAMES)})\b[^>]*>.*$', + re.DOTALL | re.IGNORECASE, +) + +_ORPHAN_REASONING_TAG_PATTERN = re.compile( + rf'\s*', + re.IGNORECASE, +) + +_STRAY_TOOL_CALL_CLOSER_PATTERN = re.compile( + rf'\s*', + re.IGNORECASE, +) + + def _ra(): """Lazy ``run_agent`` reference for test-patch routing.""" import run_agent @@ -826,62 +865,31 @@ def strip_think_blocks(agent, content: str) -> str: # 1. Closed tag pairs — case-insensitive for all variants so # mixed-case tags (, ) don't slip through to # the unterminated-tag pass and take trailing content with them. - content = re.sub(r'.*?', '', content, flags=re.DOTALL | re.IGNORECASE) - content = re.sub(r'.*?', '', content, flags=re.DOTALL | re.IGNORECASE) - content = re.sub(r'.*?', '', content, flags=re.DOTALL | re.IGNORECASE) - content = re.sub(r'.*?', '', content, flags=re.DOTALL | re.IGNORECASE) - content = re.sub(r'.*?', '', content, flags=re.DOTALL | re.IGNORECASE) + for _pattern in _REASONING_BLOCK_PATTERNS: + content = _pattern.sub('', content) # 1b. Tool-call XML blocks (openclaw/openclaw#67318). Handle the # generic tag names first — they have no attribute gating since # a literal in prose is already vanishingly rare. - for _tc_name in ("tool_call", "tool_calls", "tool_result", - "function_call", "function_calls"): - content = re.sub( - rf'<{_tc_name}\b[^>]*>.*?', - '', - content, - flags=re.DOTALL | re.IGNORECASE, - ) + for _pattern in _TOOL_CALL_BLOCK_PATTERNS: + content = _pattern.sub('', content) # 1c. ... — Gemma-style standalone # tool call. Only strip when the tag sits at a block boundary # (start of text, after a newline, or after sentence-ending # punctuation) AND carries a name="..." attribute. This keeps # prose mentions like "Use to declare" safe. - content = re.sub( - r'(?:(?<=^)|(?<=[\n\r.!?:]))[ \t]*' - r']*\bname\s*=[^>]*>' - r'(?:(?:(?!).)*)', - '', - content, - flags=re.DOTALL | re.IGNORECASE, - ) + content = _NAMED_FUNCTION_BLOCK_PATTERN.sub('', content) # 2. Unterminated reasoning block — open tag at a block boundary # (start of text, or after a newline) with no matching close. # Strip from the tag to end of string. Fixes #8878 / #9568 # (MiniMax M2.7 leaking raw reasoning into assistant content). - content = re.sub( - r'(?:^|\n)[ \t]*<(?:think|thinking|reasoning|thought|REASONING_SCRATCHPAD)\b[^>]*>.*$', - '', - content, - flags=re.DOTALL | re.IGNORECASE, - ) + content = _UNTERMINATED_REASONING_BLOCK_PATTERN.sub('', content) # 3. Stray orphan open/close tags that slipped through. - content = re.sub( - r'\s*', - '', - content, - flags=re.IGNORECASE, - ) + content = _ORPHAN_REASONING_TAG_PATTERN.sub('', content) # 3b. Stray tool-call closers. (We do NOT strip bare or # unterminated because a truncated tail # during streaming may still be valuable to the user; matches # OpenClaw's intentional asymmetry.) - content = re.sub( - r'\s*', - '', - content, - flags=re.IGNORECASE, - ) + content = _STRAY_TOOL_CALL_CLOSER_PATTERN.sub('', content) return content @@ -1456,6 +1464,64 @@ def restore_primary_runtime(agent) -> bool: if getattr(agent, "_rate_limited_until", 0) > time.monotonic(): return False # primary still in rate-limit cooldown, stay on fallback + # ── Reset-aware gate ── + # The 60s ``_rate_limited_until`` cooldown covers transient rate limits, + # but subscription-style providers (Claude Pro/Max 5-hour windows, ChatGPT + # weekly limits) report reset times hours or days away. The credential + # pool already stores those timestamps (``last_error_reset_at``); until + # the earliest one elapses, every restore attempt is a *guaranteed* + # failure that costs two prompt-cache invalidations per turn (switch to + # primary, fail, switch back to fallback) and re-marshals the full + # context each way. Skip the restore while the pool says nobody can + # serve, and come back the moment the reset time passes. + # + # Fail-open by design: any error (unreadable auth store, legacy pool + # adapter without ``next_available_at``) falls through to the existing + # every-turn retry. A pool with no reset info returns ``None`` and also + # falls through — this gate only ever *adds* skips for provably + # limited windows, so recovery can never be later than it is today. + # + # When the attached pool belongs to the fallback provider (cross-provider + # fallback rebinds it), the primary pool is loaded here and handed to the + # pool-rebind block below via ``prefetched_primary_pool`` so the load + # happens at most once per restore. + prefetched_primary_pool = None + try: + primary_provider = str( + (agent._primary_runtime or {}).get("provider") or "" + ).strip().lower() + pool = getattr(agent, "_credential_pool", None) + if not credential_pool_matches_provider( + pool, + primary_provider, + base_url=str((agent._primary_runtime or {}).get("base_url") or ""), + ): + from agent.credential_pool import load_pool + + prefetched_primary_pool = ( + load_pool(primary_provider) if primary_provider else None + ) + pool = prefetched_primary_pool + next_at = getattr(pool, "next_available_at", lambda: None)() + if next_at is not None and next_at > time.time(): + if not getattr(agent, "_restore_wait_logged", False): + agent._restore_wait_logged = True + logger.info( + "Primary %s rate-limited until %s; staying on fallback " + "%s/%s until the reset elapses", + primary_provider or "?", + datetime.fromtimestamp(next_at).isoformat(timespec="seconds"), + agent.provider, + agent.model, + ) + return False + except Exception: + logger.debug( + "Reset-aware restore gate failed; falling back to per-turn retry", + exc_info=True, + ) + agent._restore_wait_logged = False + rt = agent._primary_runtime try: # ── Core runtime state ── @@ -1549,9 +1615,14 @@ def restore_primary_runtime(agent) -> bool: agent._credential_pool = None agent._credential_pool_entry_id = None try: - from agent.credential_pool import load_pool + if prefetched_primary_pool is not None: + # Reuse the pool the reset-aware gate already loaded for + # this restore — avoids a second disk read of auth.json. + agent._credential_pool = prefetched_primary_pool + else: + from agent.credential_pool import load_pool - agent._credential_pool = load_pool(primary_provider) + agent._credential_pool = load_pool(primary_provider) except Exception as exc: logger.warning( "Restore could not reload primary credential pool for %s: %s", @@ -1631,6 +1702,7 @@ def restore_primary_runtime(agent) -> bool: # ── Reset fallback chain for the new turn ── agent._fallback_activated = False agent._fallback_index = 0 + agent._rate_limit_backoff_count = 0 # reset exponential backoff counter # Reset the stale-call circuit breaker (#58962): the streak measured # the FALLBACK provider we're leaving; the restored primary deserves @@ -2003,12 +2075,12 @@ def anthropic_prompt_cache_policy( gateway implements the Anthropic cache_control contract (MiniMax, Zhipu GLM, LiteLLM's Anthropic proxy mode all do). - Qwen models on OpenCode and direct Alibaba (DashScope), plus DeepSeek - models on OpenCode, also honour Anthropic-style ``cache_control`` markers - on OpenAI-wire chat completions. Upstream pi-mono #3392 / pi #3393 - documented this for opencode-go Qwen; #24617 reports the same gateway - contract for DeepSeek. Without markers these providers serve zero cache - hits, re-billing the full prompt on every turn. + Qwen / Alibaba-family models on OpenCode, OpenCode Go, and direct + Alibaba (DashScope) also honour Anthropic-style ``cache_control`` + markers on OpenAI-wire chat completions. Upstream pi-mono #3392 / + pi #3393 documented this for opencode-go Qwen. Without markers + these providers serve zero cache hits, re-billing the full prompt + on every turn. If the operator has set ``prompt_caching.cache_ttl`` to a falsy value (``false``, ``null``, ``"off"``, etc.) in config.yaml, prompt caching @@ -2135,22 +2207,21 @@ def anthropic_prompt_cache_policy( if is_minimax_provider or is_minimax_host: return True, True - # Qwen on OpenCode (Zen/Go) and native DashScope, plus DeepSeek on - # OpenCode only: OpenAI-wire transports that accept Anthropic-style - # cache_control markers and reward them with real cache hits. Keep direct - # Alibaba specific to Qwen; its catalog does not establish the same - # contract for DeepSeek. + # Qwen/Alibaba on OpenCode (Zen/Go) and native DashScope: OpenAI-wire + # transport that accepts Anthropic-style cache_control markers and + # rewards them with real cache hits. Without this branch + # qwen3.6-plus on opencode-go reports 0% cached tokens and burns + # through the subscription on every turn. + # + # NOTE: DeepSeek models on OpenCode are intentionally excluded. + # OpenCode Zen's relay rejects the Anthropic-style content block + # format that cache markers produce (content becomes a block array + # instead of a plain string), causing HTTP 400 (#77217). model_is_qwen = "qwen" in model_lower - model_is_deepseek = "deepseek" in model_lower - provider_is_opencode = provider_lower in { - "opencode", "opencode-zen", "opencode-go", - } provider_is_alibaba_family = provider_lower in { "opencode", "opencode-zen", "opencode-go", "alibaba", } - if (provider_is_alibaba_family and model_is_qwen) or ( - provider_is_opencode and model_is_deepseek - ): + if provider_is_alibaba_family and model_is_qwen: # Envelope layout (native_anthropic=False): markers on inner # content parts, not top-level tool messages. Matches # pi-mono's "alibaba" cacheControlFormat. diff --git a/agent/anthropic_adapter.py b/agent/anthropic_adapter.py index 7457119921..70a9599323 100644 --- a/agent/anthropic_adapter.py +++ b/agent/anthropic_adapter.py @@ -1334,7 +1334,7 @@ def _resolve_anthropic_pool_token() -> Optional[str]: # to auth.json or trigger a network refresh from a bare resolve. select() # is deliberately NOT used — it runs clear_expired=True, refresh=True, # which would violate this read-only contract. - entries = pool._available_entries(clear_expired=False, refresh=False) + entries, _pending = pool._available_entries(clear_expired=False, refresh=False) except Exception: logger.debug("Failed to read Anthropic credential_pool", exc_info=True) return None @@ -1360,19 +1360,27 @@ def resolve_anthropic_token() -> Optional[str]: Priority: 1. ANTHROPIC_TOKEN env var (OAuth/setup token saved by Hermes) 2. CLAUDE_CODE_OAUTH_TOKEN env var - 3. Claude Code credentials (~/.claude.json or ~/.claude/.credentials.json) + 3. ANTHROPIC_API_KEY env var (explicit regular API key) + 4. Claude Code credentials (~/.claude.json or ~/.claude/.credentials.json) — with automatic refresh if expired and a refresh token is available - 4. Anthropic credential_pool OAuth entry (~/.hermes/auth.json) - 5. ANTHROPIC_API_KEY env var (regular API key, or legacy fallback) + 5. Anthropic credential_pool OAuth entry (~/.hermes/auth.json) Returns the token string or None. """ - creds = read_claude_code_credentials() + creds: Optional[Dict[str, Any]] = None + creds_loaded = False + + def _read_creds() -> Optional[Dict[str, Any]]: + nonlocal creds, creds_loaded + if not creds_loaded: + creds = read_claude_code_credentials() + creds_loaded = True + return creds # 1. Hermes-managed OAuth/setup token env var token = _getenv("ANTHROPIC_TOKEN").strip() if token: - preferred = _prefer_refreshable_claude_code_token(token, creds) + preferred = _prefer_refreshable_claude_code_token(token, _read_creds()) if preferred: return preferred return token @@ -1380,27 +1388,27 @@ def resolve_anthropic_token() -> Optional[str]: # 2. CLAUDE_CODE_OAUTH_TOKEN (used by Claude Code for setup-tokens) cc_token = _getenv("CLAUDE_CODE_OAUTH_TOKEN").strip() if cc_token: - preferred = _prefer_refreshable_claude_code_token(cc_token, creds) + preferred = _prefer_refreshable_claude_code_token(cc_token, _read_creds()) if preferred: return preferred return cc_token - # 3. Claude Code credential file - resolved_claude_token = _resolve_claude_code_token_from_credentials(creds) - if resolved_claude_token: - return resolved_claude_token - - # 4. Hermes credential_pool OAuth entry. - resolved_pool_token = _resolve_anthropic_pool_token() - if resolved_pool_token: - return resolved_pool_token - - # 5. Regular API key, or a legacy OAuth token saved in ANTHROPIC_API_KEY. - # This remains as a compatibility fallback for pre-migration Hermes configs. + # 3. Regular API key. An explicit user-configured key must not be shadowed + # by auto-discovered Claude Code or credential-pool OAuth credentials. api_key = _getenv("ANTHROPIC_API_KEY").strip() if api_key: return api_key + # 4. Claude Code credential file + resolved_claude_token = _resolve_claude_code_token_from_credentials(_read_creds()) + if resolved_claude_token: + return resolved_claude_token + + # 5. Hermes credential_pool OAuth entry. + resolved_pool_token = _resolve_anthropic_pool_token() + if resolved_pool_token: + return resolved_pool_token + return None @@ -2311,13 +2319,14 @@ def _convert_user_message(content: Any) -> Dict[str, Any]: """Validate and convert a user message to anthropic format.""" if isinstance(content, list): converted_blocks = _convert_content_to_anthropic(content) - if not converted_blocks or all( - (b.get("text") or "").strip() == "" - for b in converted_blocks - if isinstance(b, dict) and b.get("type") == "text" - ): - converted_blocks = [{"type": "text", "text": "(empty message)"}] - return {"role": "user", "content": converted_blocks} + kept_blocks = _fix_blank_text_blocks_in_list( + converted_blocks, + placeholder_text="(empty message)", + msg_index=-1, + role="user", + location="_convert_user_message", + ) + return {"role": "user", "content": kept_blocks} else: if not content or (isinstance(content, str) and not content.strip()): content = "(empty message)" @@ -2620,9 +2629,114 @@ def _ensure_leading_user_turn(result: List[Dict[str, Any]]) -> None: Mirror the Bedrock Converse adapter, which unconditionally prepends a minimal user turn when the first message is not user (convert_messages_to_converse). + + The inserted text block must be non-whitespace: Anthropic separately + rejects any text content block whose text is empty or whitespace-only + ("text content blocks must contain non-whitespace text"), so a single + space here traded the "leading assistant turn" 400 for that one (#69512 + class). Uses the same placeholder as every other synthesized filler + block in this module for consistency. """ if result and result[0].get("role") != "user": - result.insert(0, {"role": "user", "content": [{"type": "text", "text": " "}]}) + result.insert( + 0, {"role": "user", "content": [{"type": "text", "text": _EMPTY_TEXT_PLACEHOLDER}]} + ) + + +def _fix_blank_text_blocks_in_list( + blocks: List[Any], + *, + placeholder_text: str, + msg_index: int, + role: Any, + location: str, +) -> List[Any]: + """Drop blank/whitespace-only text blocks from ``blocks``, in place logic. + + Non-text blocks (tool_use, tool_result, image, document, thinking, …) + and the relative order of everything else are left untouched. A + cache_control marker riding on a dropped block is relocated onto the + last surviving text/tool_use block so a breakpoint is never silently + lost. If nothing survives, a single non-blank placeholder text block + takes the dropped blocks' place (carrying the relocated cache_control, + if any) so the message never has empty content. + + Returns a new list; does not mutate ``blocks``. + """ + kept: List[Any] = [] + relocated_cache_control = None + for block_index, blk in enumerate(blocks): + if ( + isinstance(blk, dict) + and blk.get("type") == "text" + and not (isinstance(blk.get("text"), str) and blk["text"].strip()) + ): + if isinstance(blk.get("cache_control"), dict): + relocated_cache_control = blk["cache_control"] + logger.warning( + "Pre-call sanitizer: dropped blank text content block " + "(message_index=%d role=%s location=%s block_index=%d " + "block_type=text)", + msg_index, + role, + location, + block_index, + ) + continue + kept.append(blk) + if not kept: + placeholder: Dict[str, Any] = {"type": "text", "text": placeholder_text} + if relocated_cache_control is not None: + placeholder["cache_control"] = relocated_cache_control + kept.append(placeholder) + elif relocated_cache_control is not None: + _apply_assistant_cache_control_to_last_cacheable_block(kept, relocated_cache_control) + return kept + + +def _scrub_blank_text_blocks(result: List[Dict[str, Any]]) -> None: + """Final provider-boundary guard against blank Anthropic text blocks. + + Anthropic rejects any text content block whose ``text`` is empty or + whitespace-only with HTTP 400 ("text content blocks must contain + non-whitespace text"). ``_convert_assistant_message``, + ``_convert_user_message`` and ``_ensure_leading_user_turn`` already + avoid emitting these for the paths that build them, but this pass runs + last — after every other transform in ``convert_messages_to_anthropic`` + — so a blank block from any current or future producer (including one + nested inside a ``tool_result``'s own content list) never reaches the + wire. Diagnostics are structural only: message index, role, content + location, block index/type. Never logs message text, tool arguments, + tokens, or credentials. Mutates ``result`` in place. + """ + for msg_index, msg in enumerate(result): + if not isinstance(msg, dict): + continue + role = msg.get("role") + content = msg.get("content") + if not isinstance(content, list) or not content: + continue + placeholder_text = _EMPTY_TEXT_PLACEHOLDER if role == "assistant" else "(empty message)" + new_content = _fix_blank_text_blocks_in_list( + content, + placeholder_text=placeholder_text, + msg_index=msg_index, + role=role, + location="content", + ) + for blk in new_content: + if not isinstance(blk, dict) or blk.get("type") != "tool_result": + continue + inner = blk.get("content") + if isinstance(inner, list) and inner: + blk["content"] = _fix_blank_text_blocks_in_list( + inner, + placeholder_text="(no output)", + msg_index=msg_index, + role=role, + location="tool_result", + ) + msg["content"] = new_content def convert_messages_to_anthropic( @@ -2686,6 +2800,7 @@ def convert_messages_to_anthropic( _ensure_leading_user_turn(result) _manage_thinking_signatures(result, base_url, model) _evict_old_screenshots(result) + _scrub_blank_text_blocks(result) return system, result diff --git a/agent/auxiliary_client.py b/agent/auxiliary_client.py index 44d7d69683..d66d5d81ac 100644 --- a/agent/auxiliary_client.py +++ b/agent/auxiliary_client.py @@ -7604,6 +7604,77 @@ def _get_task_extra_body(task: str) -> Dict[str, Any]: return result +# --------------------------------------------------------------------------- +# Per-task concurrency limiting (#23324) +# --------------------------------------------------------------------------- +# Background auxiliary work (title generation, context compression, etc.) can +# spawn unbounded concurrent LLM calls when many sessions are active. During +# provider incidents each call also retries / fans out across the fallback +# chain, multiplying request volume on already-degraded endpoints. A per-task +# semaphore caps in-flight calls so retry amplification stays bounded. + +_aux_sync_semaphores: Dict[str, Tuple[int, threading.BoundedSemaphore]] = {} +_aux_async_semaphores: Dict[Tuple[str, int], Tuple[int, Any]] = {} +_aux_sem_lock = threading.Lock() + + +def _get_task_max_concurrency(task: Optional[str]) -> Optional[int]: + """Return ``auxiliary..max_concurrency`` as a positive int, or None.""" + if not task or task == "vision": + # Vision already uses this key for its encode/resize CPU worker pool; + # its LLM calls deliberately remain concurrent. + return None + raw = _get_auxiliary_task_config(task).get("max_concurrency") + if raw is None: + return None + try: + value = int(raw) + except (TypeError, ValueError): + return None + return value if value > 0 else None + + +def _acquire_sync_aux_semaphore(task: Optional[str]) -> Optional[threading.BoundedSemaphore]: + """Get a per-task sync semaphore, rebuilding it after a config change.""" + limit = _get_task_max_concurrency(task) + if limit is None: + return None + with _aux_sem_lock: + entry = _aux_sync_semaphores.get(task) + if entry is None or entry[0] != limit: + semaphore = threading.BoundedSemaphore(limit) + _aux_sync_semaphores[task] = (limit, semaphore) + return semaphore + return entry[1] + + +def _acquire_async_aux_semaphore(task: Optional[str]): + """Get a per-task, per-event-loop async semaphore after config lookup.""" + limit = _get_task_max_concurrency(task) + if limit is None: + return None + import asyncio + try: + loop = asyncio.get_running_loop() + except RuntimeError: + return None + key = (task, id(loop)) + with _aux_sem_lock: + entry = _aux_async_semaphores.get(key) + if entry is None or entry[0] != limit: + semaphore = asyncio.Semaphore(limit) + _aux_async_semaphores[key] = (limit, semaphore) + return semaphore + return entry[1] + + +def _reset_aux_semaphores() -> None: + """Drop cached semaphores (test helper).""" + with _aux_sem_lock: + _aux_sync_semaphores.clear() + _aux_async_semaphores.clear() + + # --------------------------------------------------------------------------- # Anthropic-compatible endpoint detection + image block conversion # --------------------------------------------------------------------------- @@ -8476,6 +8547,75 @@ def call_llm( api_mode: str = None, stream: bool = False, stream_options: dict = None, +) -> Any: + """Run an auxiliary LLM request, applying the configured task limit.""" + semaphore = _acquire_sync_aux_semaphore(task) + if semaphore is not None: + semaphore.acquire() + try: + response = _call_llm_impl( + task=task, + provider=provider, + model=model, + base_url=base_url, + api_key=api_key, + main_runtime=main_runtime, + messages=messages, + temperature=temperature, + max_tokens=max_tokens, + tools=tools, + timeout=timeout, + extra_body=extra_body, + reasoning_config=reasoning_config, + extra_headers=extra_headers, + api_mode=api_mode, + stream=stream, + stream_options=stream_options, + ) + if stream and semaphore is not None: + stream_semaphore = semaphore + semaphore = None + return _release_sync_semaphore_after_stream(response, stream_semaphore) + return response + finally: + if semaphore is not None: + semaphore.release() + + +def _release_sync_semaphore_after_stream( + stream: Any, semaphore: threading.BoundedSemaphore, +): + """Release a permit only after a streaming response is consumed or closed.""" + try: + yield from stream + finally: + try: + close = getattr(stream, "close", None) + if callable(close): + close() + finally: + semaphore.release() + + +def _call_llm_impl( + task: str = None, + *, + provider: str = None, + model: str = None, + base_url: str = None, + api_key: str = None, + main_runtime: Optional[Dict[str, Any]] = None, + messages: list, + temperature: Optional[float] = None, + max_tokens: int = None, + tools: list = None, + timeout: float = None, + extra_body: dict = None, + reasoning_config: Optional[dict] = None, + extra_headers: Optional[Dict[str, str]] = None, + api_mode: str = None, + stream: bool = False, + stream_options: dict = None, ) -> Any: """Centralized synchronous LLM call. @@ -9239,6 +9379,47 @@ async def async_call_llm( timeout: float = None, extra_body: dict = None, reasoning_config: Optional[dict] = None, +) -> Any: + """Run an asynchronous auxiliary LLM request under the configured limit.""" + semaphore = _acquire_async_aux_semaphore(task) + if semaphore is not None: + await semaphore.acquire() + try: + return await _async_call_llm_impl( + task=task, + provider=provider, + model=model, + base_url=base_url, + api_key=api_key, + main_runtime=main_runtime, + messages=messages, + temperature=temperature, + max_tokens=max_tokens, + tools=tools, + timeout=timeout, + extra_body=extra_body, + reasoning_config=reasoning_config, + ) + finally: + if semaphore is not None: + semaphore.release() + + +async def _async_call_llm_impl( + task: str = None, + *, + provider: str = None, + model: str = None, + base_url: str = None, + api_key: str = None, + main_runtime: Optional[Dict[str, Any]] = None, + messages: list, + temperature: Optional[float] = None, + max_tokens: int = None, + tools: list = None, + timeout: float = None, + extra_body: dict = None, + reasoning_config: Optional[dict] = None, ) -> Any: """Centralized asynchronous LLM call. diff --git a/agent/chat_completion_helpers.py b/agent/chat_completion_helpers.py index e6e8ab7fdc..3b7d8b0361 100644 --- a/agent/chat_completion_helpers.py +++ b/agent/chat_completion_helpers.py @@ -1712,7 +1712,20 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool current_provider = (getattr(agent, "provider", "") or "").strip().lower() primary_provider = ((agent._primary_runtime or {}).get("provider") or "").strip().lower() if (not fallback_already_active) or (primary_provider and current_provider == primary_provider): - agent._rate_limited_until = time.monotonic() + 60 + # Exponential backoff: keep upstream's 60s first-hit cooldown and + # escalate on CONSECUTIVE rate-limits: 60s → 2m → 4m → 8m → ... → + # 4h cap. The first 429 must NOT bench the primary for half an + # hour — fast primary restore is the common case; escalation only + # punishes providers that keep 429ing. + # Counter is reset by restore_primary_runtime on successful restore. + backoff_count = getattr(agent, "_rate_limit_backoff_count", 0) + agent._rate_limit_backoff_count = backoff_count + 1 + backoff_seconds = min(60 * (2 ** backoff_count), 14400) + agent._rate_limited_until = time.monotonic() + backoff_seconds + logging.info( + "Rate-limit backoff level %d: cooldown %d s (%.1f min, backoff#%d)", + backoff_count, backoff_seconds, backoff_seconds / 60, backoff_count + 1, + ) if agent._fallback_index >= len(agent._fallback_chain): # Chain exhausted. If we actually walked a non-empty chain and the # failure was NOT a rate-limit/billing event (those already armed diff --git a/agent/context_compressor.py b/agent/context_compressor.py index fbb7e6c5e8..59ae723d2b 100644 --- a/agent/context_compressor.py +++ b/agent/context_compressor.py @@ -6730,6 +6730,26 @@ This compaction should PRIORITISE preserving all information related to the focu _strip_persistence_markers(compressed) self._last_compression_made_progress = True + # A successful compaction just freed the largest allocation a long + # session ever drops (the compressed-away message dicts), which makes + # this the natural point to hand allocator pages back to the OS. + # #76905's trim lifecycle covers the gateway/TUI housekeeping loops but + # not the CLI compression path, so RSS keeps the pre-compaction + # high-water mark until exit. The helper is glibc-gated, config-gated + # and rate-limited, so this is a safe no-op elsewhere. (#70782) + try: + from hermes_cli.mem_trim import trim_memory + + trim_memory(reason="post-compression") + except Exception as exc: + # debug, not warning: sibling trim sites all log failures at + # debug, and compression must never fail because of a trim. + logger.debug( + "post-compression memory trim failed: %s: %s", + type(exc).__name__, + exc, + ) + # Batch compaction invalidates micro-compaction state: the batch # marker now holds MORE history than the in-memory rolling summary # (it summarized everything in the window, including exchanges micro diff --git a/agent/conversation_loop.py b/agent/conversation_loop.py index 1884c06489..8848b64df2 100644 --- a/agent/conversation_loop.py +++ b/agent/conversation_loop.py @@ -58,6 +58,11 @@ from agent.message_sanitization import ( _strip_images_from_messages, _strip_non_ascii, ) +# Must mirror _STALE_TOOL_CALL_MARKER_RE in hermes_state.py — kept local +# to avoid importing hermes_state at module load time (its module-level +# DEFAULT_DB_PATH = get_hermes_home() / "state.db" breaks tests that +# monkeypatch get_hermes_home to return a str). +_STALE_MARKER_RE = re.compile(r"^\[[A-Za-z_][A-Za-z0-9_.-]*\]$") from agent.model_metadata import ( MINIMUM_CONTEXT_LENGTH, _estimate_tools_tokens_rough, @@ -4360,9 +4365,11 @@ def run_conversation( compression_attempts += 1 if compression_attempts <= max_compression_attempts: original_len = len(messages) + # Option A (LCM issue 441): overhead-aware request size so recovery arms on + # the true request (msgs + tools + system), not the tool-blind message count. messages, active_system_prompt = agent._compress_context( messages, system_message, - approx_tokens=approx_tokens, + approx_tokens=estimate_request_tokens_rough(api_messages, tools=agent.tools or None), task_id=effective_task_id, ) conversation_history = conversation_history_after_compression( @@ -4617,8 +4624,11 @@ def run_conversation( original_len = len(messages) original_tokens = estimate_messages_tokens_rough(messages) _overflow_input = messages + # Option A (LCM issue 441): overhead-aware request size so recovery arms on the + # true request (msgs + tools + system), not the tool-blind message count. messages, active_system_prompt = agent._compress_context( - messages, system_message, approx_tokens=approx_tokens, + messages, system_message, + approx_tokens=estimate_request_tokens_rough(api_messages, tools=agent.tools or None), task_id=effective_task_id, ) if messages is _overflow_input and compression_skipped_due_to_lock(agent): @@ -4754,6 +4764,41 @@ def run_conversation( "failed": True, "compression_exhausted": True, } + # Also compress the message history so the output-cap + # retry does not just spin on max_tokens alone. The + # compressor drops the middle window, freeing enough + # tokens for the total to fit inside context_length. + # (#55546) + try: + original_len = len(messages) + original_tokens = estimate_messages_tokens_rough(messages) + _overflow_input = messages + messages, active_system_prompt = agent._compress_context( + messages, system_message, + approx_tokens=request_input_estimate, + task_id=effective_task_id, + ) + if messages is _overflow_input and compression_skipped_due_to_lock(agent): + compression_attempts -= 1 + agent._persist_session(messages, conversation_history) + return _compression_deferred_result( + agent, messages, api_call_count + ) + conversation_history = conversation_history_after_compression( + agent, messages, conversation_history + ) + new_tokens = estimate_messages_tokens_rough(messages) + if len(messages) < original_len: + agent._buffer_status(COMPRESSION_RETRY_MESSAGES_STATUS_TEMPLATE.format(before=original_len, after=len(messages))) + elif new_tokens > 0 and new_tokens < original_tokens * 0.95: + agent._buffer_status(COMPRESSION_RETRY_TOKENS_STATUS_TEMPLATE.format(before=original_tokens, after=new_tokens)) + except Exception: + # Compression must never turn an output-cap error + # fatal — fall through and retry on max_tokens alone. + logger.warning( + "%sOutput-cap compression hit an error; retrying on max_tokens only.", + agent.log_prefix, + ) _retry.restart_with_compressed_messages = True break @@ -4878,8 +4923,13 @@ def run_conversation( original_len = len(messages) original_tokens = estimate_messages_tokens_rough(messages) _overflow_input = messages + # Option A (LCM issue 441): pass the OVERHEAD-AWARE request size (msgs + tool + # schemas + system), not the tool-blind message count, so LCM forced-overflow + # recovery arms on the TRUE request that overflowed. See hermes-lcm engine + # _should_force_overflow_recovery. (approx_tokens stays for the status display.) messages, active_system_prompt = agent._compress_context( - messages, system_message, approx_tokens=approx_tokens, + messages, system_message, + approx_tokens=estimate_request_tokens_rough(api_messages, tools=agent.tools or None), task_id=effective_task_id, ) if messages is _overflow_input and compression_skipped_due_to_lock(agent): @@ -6107,9 +6157,25 @@ def run_conversation( ] assistant_msg = agent._build_assistant_message(assistant_message, finish_reason) - + turn_content = assistant_message.content or "" + # Some local tool-call templates emit a bare bracketed token + # (for example ``[memory]``) as assistant content alongside a + # function call. It is protocol scaffolding, not an answer. + # Persisting or caching it as visible content lets the empty + # post-tool fallback replay that token forever after compaction (#78148). + if ( + assistant_message.tool_calls + and _STALE_MARKER_RE.fullmatch(turn_content.strip()) + ): + logger.warning( + "Discarding bare tool-call marker from assistant content: %s", + turn_content, + ) + turn_content = "" + assistant_msg["content"] = "" + # Classify tools in this turn to determine if they are all housekeeping. # This classification is needed regardless of whether the turn has visible content, # because a substantive tool-only turn must invalidate any older housekeeping fallback. @@ -6373,9 +6439,12 @@ def run_conversation( _clear_warn() agent._safe_print(" ⟳ compacting context…") _post_tool_input = messages + # Route the overhead-aware _real_tokens (computed above) into compression, not + # the bare last_prompt_tokens — which is 0 in the no-usage fallback, hiding the + # true request size from the engine's overflow guard (upstream PR #77169 review). messages, active_system_prompt = agent._compress_context( messages, system_message, - approx_tokens=agent.context_compressor.last_prompt_tokens, + approx_tokens=_real_tokens, task_id=effective_task_id, ) if ( @@ -6662,15 +6731,47 @@ def run_conversation( ) if _truly_empty and (not _has_structured or _prefill_exhausted) and agent._empty_content_retries < 3: agent._empty_content_retries += 1 + wait_time = jittered_backoff( + agent._empty_content_retries, + base_delay=5.0, + max_delay=60.0, + ) logger.warning( "Empty response (no content or reasoning) — " - "retry %d/3 (model=%s)", - agent._empty_content_retries, agent.model, + "retry %d/3 in %.1fs (model=%s)", + agent._empty_content_retries, wait_time, agent.model, ) agent._buffer_status( f"⚠️ Empty response from model — retrying " - f"({agent._empty_content_retries}/3)" + f"({agent._empty_content_retries}/3) in {wait_time:.0f}s" ) + # 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 agent._interrupt_requested: + agent._vprint(f"{agent.log_prefix}⚡ Interrupt detected during empty-response retry wait, aborting.", force=True) + _interrupt_text = ( + f"Operation interrupted: retrying empty response from model " + f"(retry {agent._empty_content_retries}/3)." + ) + close_interrupted_tool_sequence(messages, _interrupt_text) + agent._persist_session(messages, conversation_history) + agent.clear_interrupt() + return { + "final_response": _interrupt_text, + "messages": messages, + "api_calls": api_call_count, + "completed": False, + "interrupted": True, + } + time.sleep(0.2) + _backoff_touch_counter += 1 + if _backoff_touch_counter % 150 == 0: # 150 × 0.2s = 30s + agent._touch_activity( + f"empty response retry backoff ({agent._empty_content_retries}/3), " + f"{int(sleep_end - time.time())}s remaining" + ) continue # ── Exhausted retries — try fallback provider ── diff --git a/agent/credential_pool.py b/agent/credential_pool.py index adb24f75a6..8917b49d8d 100644 --- a/agent/credential_pool.py +++ b/agent/credential_pool.py @@ -588,7 +588,12 @@ class CredentialPool: self._entries = sorted(entries, key=lambda entry: entry.priority) self._current_id: Optional[str] = None self._strategy = get_pool_strategy(provider) - self._lock = threading.Lock() + # RLock: the mutation primitives below (_replace_entry/_persist) + # self-acquire this lock so the DEFERRED single-use-token refresh + # path (which runs network I/O outside the lock by design) still + # serializes its pool mutations. In-lock callers re-acquire + # reentrantly at negligible cost. + self._lock = threading.RLock() self._active_leases: Dict[str, int] = {} self._max_concurrent = DEFAULT_MAX_CONCURRENT_PER_CREDENTIAL # Monotonic timestamp of the last "no available entries" log, used to @@ -618,7 +623,36 @@ class CredentialPool: # otherwise a status probe here can race a concurrent ``select`` / # rotation and tear ``self._entries`` or double-write auth.json. with self._lock: - return bool(self._available_entries()) + available, _pending = self._available_entries() + return bool(available) + + def next_available_at(self) -> Optional[float]: + """Earliest epoch time (seconds) any entry re-enters rotation. + + Returns ``None`` when at least one entry is available right now, or + when no exhausted entry carries a usable recovery time (empty pool, + or only ``STATUS_DEAD`` entries, which never re-enter via TTL). + Callers must treat ``None`` as "no wait information", not + "unavailable". + + Like :meth:`has_available`, expired cooldowns are left uncleared + (``clear_expired=False``); the only writes are the same + re-auth/token sync paths ``has_available`` already performs — which + is exactly why this must run under ``self._lock`` like every other + ``_available_entries`` caller (see the comment on ``has_available``). + """ + with self._lock: + available, _pending = self._available_entries() + if available: + return None + candidates: List[float] = [] + for entry in self._entries: + if entry.last_status != STATUS_EXHAUSTED: + continue + until = _exhausted_until(entry) + if until is not None: + candidates.append(until) + return min(candidates) if candidates else None def entries(self) -> List[PooledCredential]: with self._lock: @@ -656,18 +690,27 @@ class CredentialPool: return matches[0].id if len(matches) == 1 else None def _replace_entry(self, old: PooledCredential, new: PooledCredential) -> None: - """Swap an entry in-place by id, preserving sort order.""" - for idx, entry in enumerate(self._entries): - if entry.id == old.id: - self._entries[idx] = new - return + """Swap an entry in-place by id, preserving sort order. + + Self-locking (RLock) so the deferred refresh path — which + deliberately runs outside the pool lock — cannot tear + ``self._entries`` against a concurrent select()/rotation. + """ + with self._lock: + for idx, entry in enumerate(self._entries): + if entry.id == old.id: + self._entries[idx] = new + return def _persist(self, *, removed_ids: Optional[List[str]] = None) -> None: - write_credential_pool( - self.provider, - [entry.to_dict() for entry in self._entries], - removed_ids=removed_ids, - ) + # Self-locking (RLock): snapshotting self._entries must not race a + # concurrent rotation when called from the deferred refresh path. + with self._lock: + write_credential_pool( + self.provider, + [entry.to_dict() for entry in self._entries], + removed_ids=removed_ids, + ) def _is_terminal_auth_failure( self, @@ -1415,17 +1458,22 @@ class CredentialPool: logger.debug( "Failed to clear terminal xAI OAuth state: %s", clear_exc ) - removed_ids = [ - item.id for item in self._entries - if item.source == "device_code" - ] - self._entries = [ - item for item in self._entries - if item.source != "device_code" - ] - if self._current_id == entry.id: - self._current_id = None - self._persist(removed_ids=removed_ids) + # Read-modify-write of self._entries: must be atomic. + # This runs on the DEFERRED refresh path (outside the + # pool lock), so take it here. self._lock is an RLock, + # so the still-locked callers re-enter safely. + with self._lock: + removed_ids = [ + item.id for item in self._entries + if item.source == "device_code" + ] + self._entries = [ + item for item in self._entries + if item.source != "device_code" + ] + if self._current_id == entry.id: + self._current_id = None + self._persist(removed_ids=removed_ids) return None # For openai-codex: same race as xAI/nous — another Hermes process # may have consumed the refresh token between our proactive sync @@ -1485,17 +1533,22 @@ class CredentialPool: logger.debug( "Failed to clear terminal Codex OAuth state: %s", clear_exc ) - removed_ids = [ - item.id for item in self._entries - if item.source == "device_code" - ] - self._entries = [ - item for item in self._entries - if item.source != "device_code" - ] - if self._current_id == entry.id: - self._current_id = None - self._persist(removed_ids=removed_ids) + # Read-modify-write of self._entries: must be atomic. + # This runs on the DEFERRED refresh path (outside the + # pool lock), so take it here. self._lock is an RLock, + # so the still-locked callers re-enter safely. + with self._lock: + removed_ids = [ + item.id for item in self._entries + if item.source == "device_code" + ] + self._entries = [ + item for item in self._entries + if item.source != "device_code" + ] + if self._current_id == entry.id: + self._current_id = None + self._persist(removed_ids=removed_ids) return None # For nous: another process may have consumed the refresh token # between our proactive sync and the HTTP call. Re-sync from @@ -1552,17 +1605,19 @@ class CredentialPool: auth_mod.NOUS_DEVICE_CODE_SOURCE, f"manual:{auth_mod.NOUS_DEVICE_CODE_SOURCE}", } - removed_ids = [ - item.id for item in self._entries - if item.source in singleton_sources - ] - self._entries = [ - item for item in self._entries - if item.source not in singleton_sources - ] - if self._current_id == entry.id: - self._current_id = None - self._persist(removed_ids=removed_ids) + # Atomic read-modify-write; see the note above. + with self._lock: + removed_ids = [ + item.id for item in self._entries + if item.source in singleton_sources + ] + self._entries = [ + item for item in self._entries + if item.source not in singleton_sources + ] + if self._current_id == entry.id: + self._current_id = None + self._persist(removed_ids=removed_ids) return None self._mark_exhausted(entry, None) return None @@ -1646,26 +1701,64 @@ class CredentialPool: return False def select(self) -> Optional[PooledCredential]: - with self._lock: - entry = self._select_unlocked() - if entry is not None: - # A normal (non-recovery) selection starts a fresh episode — - # don't let a leftover unmatched-rotation streak from an old - # failure trip the #70401 bound early next time. - self._unmatched_rotation_streak = 0 + entry, pending_refresh = self._select_under_lock() + if pending_refresh: + self._refresh_pending_entries(pending_refresh) + if entry is not None: + self._unmatched_rotation_streak = 0 return entry + # If no entry was available but we just refreshed some, re-select + # now that the refreshed entries are back in the pool. + if pending_refresh: + entry, _ = self._select_under_lock() + if entry is not None: + self._unmatched_rotation_streak = 0 + return entry - def _available_entries(self, *, clear_expired: bool = False, refresh: bool = False) -> List[PooledCredential]: - """Return entries not currently in exhaustion cooldown. + def _select_under_lock(self) -> Tuple[Optional[PooledCredential], List[tuple]]: + """Run selection under the lock, returning entry + pending refreshes.""" + with self._lock: + return self._select_unlocked() + + def _refresh_pending_entries(self, pending: List[tuple]) -> None: + """Refresh deferred single-use-token entries outside the lock. + + Each entry is refreshed under the cross-process ``_auth_store_lock`` + (which can block for 20+ seconds) and then merged into the pool. + On failure the entry is silently skipped. + """ + for entry, sync_fn in pending: + # _refresh_entry merges the refreshed entry into the pool + # internally. Its mutation primitives (_replace_entry, _persist) + # are self-locking, and the quarantine paths inside + # _refresh_entry_impl take self._lock explicitly around their + # read-modify-write of self._entries — required because this + # call site runs OUTSIDE the pool lock. + self._refresh_entry(entry, force=False) + + def _available_entries( + self, *, clear_expired: bool = False, refresh: bool = False, + ) -> Tuple[List[PooledCredential], List[tuple]]: + """Return (available, pending_refresh) for entries not in cooldown. When *clear_expired* is True, entries whose cooldown has elapsed are reset to STATUS_OK and persisted. When *refresh* is True, entries that need a token refresh are refreshed (skipped on failure). + + Single-use-token refreshes (openai-codex, xai-oauth) are returned as + *pending_refresh* tuples so the caller can execute them outside the + lock, avoiding stalling all pool consumers during cross-process flock + acquisition + OAuth network I/O. """ now = time.time() cleared_any = False entries_to_prune: List[str] = [] available: List[PooledCredential] = [] + # Entries that need an OAuth refresh via a single-use token provider + # (openai-codex, xai-oauth). These require a cross-process file lock + # that can block for 20+ seconds. We collect them under self._lock + # and refresh outside the lock to avoid stalling all pool consumers. + pending_refresh: List[tuple] = [] # (entry, sync_entry_fn) for entry in self._entries: # Borrowed credentials persist as metadata-only references and are # hydrated from their live source on load. A stale duplicate row @@ -1774,6 +1867,16 @@ class CredentialPool: entry = cleared cleared_any = True if refresh and self._entry_needs_refresh(entry): + if self.provider in ("openai-codex", "xai-oauth"): + # Defer single-use-token refresh to avoid holding the + # threading lock during cross-process flock + network I/O. + sync_fn = ( + self._sync_codex_entry_from_auth_store + if self.provider == "openai-codex" + else self._sync_xai_oauth_entry_from_pool_store + ) + pending_refresh.append((entry, sync_fn)) + continue refreshed = self._refresh_entry(entry, force=False) if refreshed is None: continue @@ -1784,7 +1887,7 @@ class CredentialPool: self._entries = [e for e in self._entries if e.id not in pruned_ids] if cleared_any: self._persist(removed_ids=entries_to_prune) - return available + return available, pending_refresh def _log_no_available_entries(self) -> None: """Emit the empty-pool INFO line at most once per throttle window. @@ -1800,12 +1903,17 @@ class CredentialPool: self._last_no_entries_log_at = now logger.info("credential pool: no available entries (all exhausted or empty)") - def _select_unlocked(self, *, refresh: bool = True) -> Optional[PooledCredential]: - available = self._available_entries(clear_expired=True, refresh=refresh) + def _select_unlocked(self, *, refresh: bool = True) -> Tuple[Optional[PooledCredential], List[tuple]]: + """Select the best available credential entry. + + Returns ``(entry, pending_refresh)`` where *pending_refresh* contains + single-use-token entries that must be refreshed outside the lock. + """ + available, pending_refresh = self._available_entries(clear_expired=True, refresh=refresh) if not available: self._current_id = None self._log_no_available_entries() - return None + return None, pending_refresh # A successful selection means the pool recovered; re-arm the throttle # so a later re-exhaustion logs immediately rather than being silenced @@ -1815,7 +1923,7 @@ class CredentialPool: if self._strategy == STRATEGY_RANDOM: entry = random.choice(available) self._current_id = entry.id - return entry + return entry, pending_refresh if self._strategy == STRATEGY_LEAST_USED and len(available) > 1: entry = min(available, key=lambda e: e.request_count) @@ -1823,7 +1931,7 @@ class CredentialPool: updated = replace(entry, request_count=entry.request_count + 1) self._replace_entry(entry, updated) self._current_id = entry.id - return updated + return updated, pending_refresh if self._strategy == STRATEGY_ROUND_ROBIN and len(available) > 1: entry = available[0] @@ -1832,11 +1940,11 @@ class CredentialPool: self._entries = [replace(candidate, priority=idx) for idx, candidate in enumerate(rotated)] self._persist() self._current_id = entry.id - return self._current_unlocked() or entry + return self._current_unlocked() or entry, pending_refresh entry = available[0] self._current_id = entry.id - return entry + return entry, pending_refresh def peek(self) -> Optional[PooledCredential]: # Single lock acquisition for the whole read; call the unlocked @@ -1845,7 +1953,7 @@ class CredentialPool: current = self._current_unlocked() if current is not None: return current - available = self._available_entries() + available, _pending = self._available_entries() return available[0] if available else None def mark_exhausted_and_rotate( @@ -1894,7 +2002,8 @@ class CredentialPool: # guessing and surface the error (no cooldown is written for # anybody — healthy keys stay available for the next turn). self._unmatched_rotation_streak += 1 - available_count = len(self._available_entries()) + available_count, _ = self._available_entries() + available_count = len(available_count) if self._unmatched_rotation_streak > max(available_count, 1): logger.warning( "credential pool: failed credential identity matched no " @@ -1913,8 +2022,9 @@ class CredentialPool: self.provider, ) self._current_id = None - next_entry = self._select_unlocked() - if next_entry is not None and len(self._available_entries()) == 1: + next_entry, _pending = self._select_unlocked(refresh=False) + avail, _ = self._available_entries() + if next_entry is not None and len(avail) == 1: # A single-entry pool cannot rotate. Returning its only # entry reports a successful recovery without changing # the credential, so the caller retries the same 401 @@ -1927,7 +2037,7 @@ class CredentialPool: # streak is stale (this mark WILL advance pool state). self._unmatched_rotation_streak = 0 if entry is None: - entry = self._current_unlocked() or self._select_unlocked() + entry = self._current_unlocked() or self._select_unlocked(refresh=False)[0] if entry is None: return None _label = entry.label or entry.id[:8] @@ -1972,7 +2082,7 @@ class CredentialPool: _label, status_code, ) self._current_id = None - next_entry = self._select_unlocked() + next_entry, _pending = self._select_unlocked(refresh=False) if next_entry: _next_label = next_entry.label or next_entry.id[:8] logger.info("credential pool: rotated to %s", _next_label) @@ -1986,15 +2096,32 @@ class CredentialPool: a stable tie-breaker. When every credential is already at the soft cap, still return the least-leased one instead of blocking. """ + chosen_id, pending_refresh = self._acquire_lease_under_lock(credential_id) + if pending_refresh: + self._refresh_pending_entries(pending_refresh) + # Mirror select(): if nothing was leasable but we just refreshed + # deferred single-use-token entries, retry now that they are back + # in rotation. Without this, a pool whose only entries all needed + # a refresh returns None even though the refresh succeeded — the + # caller sees "no credentials available" and fails a request that + # should have gone through. + if chosen_id is None: + chosen_id, _ = self._acquire_lease_under_lock(credential_id) + return chosen_id + + def _acquire_lease_under_lock( + self, credential_id: Optional[str], + ) -> Tuple[Optional[str], List[tuple]]: + """Run lease acquisition under the lock, returning id + pending refreshes.""" with self._lock: if credential_id: self._active_leases[credential_id] = self._active_leases.get(credential_id, 0) + 1 self._current_id = credential_id - return credential_id + return credential_id, [] - available = self._available_entries(clear_expired=True, refresh=True) + available, pending_refresh = self._available_entries(clear_expired=True, refresh=True) if not available: - return None + return None, pending_refresh below_cap = [ entry for entry in available @@ -2007,7 +2134,7 @@ class CredentialPool: ) self._active_leases[chosen.id] = self._active_leases.get(chosen.id, 0) + 1 self._current_id = chosen.id - return chosen.id + return chosen.id, pending_refresh def release_lease(self, credential_id: str) -> None: """Release a previously acquired credential lease.""" @@ -2059,7 +2186,7 @@ class CredentialPool: else: entry = self._current_unlocked() or self._select_unlocked( refresh=False - ) + )[0] if entry is None: return None self._current_id = entry.id @@ -2387,9 +2514,41 @@ def _seed_from_singletons(provider: str, entries: List[PooledCredential]) -> Tup # env vars (COPILOT_GITHUB_TOKEN / GH_TOKEN). They don't live in # the auth store or credential pool, so we resolve them here. try: - from hermes_cli.copilot_auth import resolve_copilot_token, get_copilot_api_token + from hermes_cli.copilot_auth import ( + COPILOT_ENV_VARS, + resolve_copilot_token, + get_copilot_api_token, + ) + # All-sources suppression gate BEFORE any work — including the + # `gh auth token` subprocess spawn. resolve_copilot_token() + # shells out (~30ms), and the exchange retries 3x with backoff + # (~35s worst case); a user who suppressed every copilot source + # (hermes auth remove copilot gh_cli) must not pay either on + # every pool load (model picker open, /model, agent startup). + # Enumerating the full source space here matches what + # credential_sources._remove_copilot_gh suppresses, so an + # all-suppressed check is stable. + copilot_sources = ["gh_cli"] + [f"env:{v}" for v in COPILOT_ENV_VARS] + if all(_is_suppressed(provider, s) for s in copilot_sources): + return changed, active_sources token, source = resolve_copilot_token() if token: + # ``resolve_copilot_token`` returns exactly "gh auth token" + # for the CLI path; env-sourced tokens return the var name. + # Match exactly — a substring test classifies GH_TOKEN and + # GITHUB_TOKEN as gh_cli, silently bypassing a user's + # per-env-var suppression. + source_name = "gh_cli" if source == "gh auth token" else f"env:{source}" + # Per-source suppression gate (a user may suppress only the + # gh CLI path and keep an env var, or vice versa) BEFORE the + # network exchange. The exchange retries 3x with 10s + # timeouts and 4.5s total backoff (~35s worst case), so a + # source the user already suppressed + # must not burn that dead time just to have the entry + # discarded afterwards. Same early-gate pattern every other + # singleton branch uses. + if _is_suppressed(provider, source_name): + return changed, active_sources api_token, enterprise_base_url = get_copilot_api_token(token) # Observability: get_copilot_api_token falls back to returning # the RAW token when the exchange fails. A raw ~40-char token @@ -2405,27 +2564,25 @@ def _seed_from_singletons(provider: str, entries: List[PooledCredential]) -> Tup "unavailable); enterprise-only models may 400 with " "model_not_available_for_integrator until exchange recovers." ) - source_name = "gh_cli" if "gh" in source.lower() else f"env:{source}" - if not _is_suppressed(provider, source_name): - active_sources.add(source_name) - pconfig = PROVIDER_REGISTRY.get(provider) - # Use enterprise base URL from token exchange if available, - # otherwise fall back to the provider's default. - effective_base_url = enterprise_base_url or ( - pconfig.inference_base_url if pconfig else "" - ) - changed |= _upsert_entry( - entries, - provider, - source_name, - { - "source": source_name, - "auth_type": AUTH_TYPE_API_KEY, - "access_token": api_token, - "base_url": effective_base_url, - "label": source, - }, - ) + active_sources.add(source_name) + pconfig = PROVIDER_REGISTRY.get(provider) + # Use enterprise base URL from token exchange if available, + # otherwise fall back to the provider's default. + effective_base_url = enterprise_base_url or ( + pconfig.inference_base_url if pconfig else "" + ) + changed |= _upsert_entry( + entries, + provider, + source_name, + { + "source": source_name, + "auth_type": AUTH_TYPE_API_KEY, + "access_token": api_token, + "base_url": effective_base_url, + "label": source, + }, + ) except Exception as exc: logger.debug("Copilot token seed failed: %s", exc) diff --git a/agent/curator.py b/agent/curator.py index 750dd7c2eb..ab0adddae1 100644 --- a/agent/curator.py +++ b/agent/curator.py @@ -1923,6 +1923,7 @@ def _run_llm_review(prompt: str) -> Dict[str, Any]: credential_pool=_credential_pool, request_overrides=_request_overrides, **_agent_kwargs, + enabled_toolsets=["skills", "terminal"], # Umbrella-building over a large skill collection is worth a # high iteration ceiling — the pass typically takes 50-100 # API calls against hundreds of candidate skills. The diff --git a/agent/display.py b/agent/display.py index 3da6e2f24e..d6bea7e54e 100644 --- a/agent/display.py +++ b/agent/display.py @@ -14,6 +14,7 @@ from dataclasses import dataclass, field from difflib import unified_diff from pathlib import Path from typing import Any +from urllib.parse import urlsplit from utils import safe_json_loads from agent.redact import redact_sensitive_text @@ -187,6 +188,15 @@ def _truncate_preview(text: str, max_len: int | None) -> str: return text +@dataclass(frozen=True) +class ToolPreview: + """A compact tool preview plus presentation facts lost to truncation.""" + + text: str + truncated: bool = False + url: str | None = None + + _SHELL_SILENT_HEADS = {"cd", "pushd", "popd", "export", "set", "unset", "source", ".", "true", "false", ":"} _SHELL_PIPE_TAIL_HEADS = {"head", "tail", "wc", "sort", "uniq"} @@ -556,6 +566,35 @@ def build_tool_preview(tool_name: str, args: dict, max_len: int | None = None) - return preview +def prepare_tool_preview( + tool_name: str, + args: dict | None, + *, + fallback: str, + max_len: int, +) -> ToolPreview: + """Build one canonical compact preview before platform formatting. + + The uncapped preview is rebuilt from the tool arguments when possible so + an upstream display cap cannot discard its link target. Platforms then + receive explicit truncation and URL metadata instead of inferring either + fact from the rendered text. + """ + full_text = build_tool_preview(tool_name, args, max_len=0) or fallback + text = _truncate_preview(full_text, max_len) + truncated = text != full_text + url = None + if truncated: + candidate = _display_url(full_text) + try: + parsed = urlsplit(candidate) + except ValueError: + parsed = None + if parsed and parsed.scheme.lower() in {"http", "https"} and parsed.netloc: + url = candidate + return ToolPreview(text=text, truncated=truncated, url=url) + + # ========================================================================= # Friendly tool labels (human-phrased verbs for built-in tools) # @@ -1506,5 +1545,3 @@ def get_cute_tool_message( # ========================================================================= # Honcho session line (one-liner with clickable OSC 8 hyperlink) # ========================================================================= - - diff --git a/agent/insights.py b/agent/insights.py index 9d148a1544..34e78a6ff9 100644 --- a/agent/insights.py +++ b/agent/insights.py @@ -99,6 +99,31 @@ class InsightsEngine: """ self.db = db self._conn = db._conn + # INDEXED BY is a hard dependency (SQLite errors on a missing index). + # A read-only open of a state.db written by an older version skips + # schema init and lacks the partial index — probe once and fall back + # to the unpinned variants (identical rows, optimizer-chosen plan). + try: + self._has_assistant_calls_index = bool( + self._conn.execute( + "SELECT 1 FROM sqlite_master WHERE type='index' AND name=?", + (self._MESSAGES_ASSISTANT_CALLS_INDEX,), + ).fetchone() + ) + except sqlite3.Error: + self._has_assistant_calls_index = False + if not self._has_assistant_calls_index: + _strip = f" INDEXED BY {self._MESSAGES_ASSISTANT_CALLS_INDEX}" + # Loop over every pinned statement so adding a new one can't + # forget its strip line (which would be a hard `no such index` + # crash on read-only DBs — the exact bug this fallback prevents). + for _attr in ( + "_GET_TOOL_CALLS_WITH_SOURCE", + "_GET_TOOL_CALLS_ALL", + "_GET_SKILL_CALLS_WITH_SOURCE", + "_GET_SKILL_CALLS_ALL", + ): + setattr(self, _attr, getattr(self, _attr).replace(_strip, "")) def generate(self, days: int = 30, source: str = None) -> Dict[str, Any]: """ @@ -171,6 +196,21 @@ class InsightsEngine: "top_sessions": top_sessions, } + def get_usage_breakdown(self, days: int = 30, source: str = None) -> Dict[str, Any]: + """Return the analytics-usage payload without running a full generate(). + + Uses the instr()-prefiltered _get_skill_usage query so only messages + that reference skill_view or skill_manage are loaded from SQLite, while + still preserving the per-tool breakdown used by the dashboard route. + """ + cutoff = time.time() - (days * 86400) + tool_usage = self._get_tool_usage(cutoff, source) + skill_usage = self._get_skill_usage(cutoff, source) + return { + "tools": self._compute_tool_breakdown(tool_usage), + "skills": self._compute_skill_breakdown(skill_usage), + } + # ========================================================================= # Data gathering (SQL queries) # ========================================================================= @@ -195,6 +235,53 @@ class InsightsEngine: " ORDER BY started_at DESC" ) + # Assistant ``tool_calls`` scan for tool/skill usage. ``INDEXED BY`` pins + # the partial index ``idx_messages_assistant_calls_by_session`` so the plan + # is deterministic on a freshly initialized state.db (before ANALYZE has + # run) for BOTH the unfiltered and source-filtered branches — without the + # hint the optimizer falls back to ``idx_messages_session_active`` for the + # source-filtered probe and scans each session's non-tool-call rows. + # + # The pin is a HARD dependency: SQLite raises ``no such index`` when the + # named index is absent. That happens in practice — the web dashboard's + # usage analytics open the DB ``read_only=True`` (skipping + # ``_init_schema``), so a state.db created by an older writer has no + # partial index yet. ``__init__`` probes for the index once and falls + # back to the unpinned (still-correct, just optimizer-chosen) variants. + _MESSAGES_ASSISTANT_CALLS_INDEX = "idx_messages_assistant_calls_by_session" + _GET_TOOL_CALLS_WITH_SOURCE = ( + "SELECT m.tool_calls" + f" FROM messages m INDEXED BY {_MESSAGES_ASSISTANT_CALLS_INDEX}" + " JOIN sessions s ON s.id = m.session_id" + " WHERE s.started_at >= ? AND s.source = ?" + " AND m.role = 'assistant' AND m.tool_calls IS NOT NULL" + ) + _GET_TOOL_CALLS_ALL = ( + "SELECT m.tool_calls" + f" FROM messages m INDEXED BY {_MESSAGES_ASSISTANT_CALLS_INDEX}" + " JOIN sessions s ON s.id = m.session_id" + " WHERE s.started_at >= ?" + " AND m.role = 'assistant' AND m.tool_calls IS NOT NULL" + ) + _GET_SKILL_CALLS_WITH_SOURCE = ( + "SELECT m.tool_calls, m.timestamp" + f" FROM messages m INDEXED BY {_MESSAGES_ASSISTANT_CALLS_INDEX}" + " JOIN sessions s ON s.id = m.session_id" + " WHERE s.started_at >= ? AND s.source = ?" + " AND m.role = 'assistant' AND m.tool_calls IS NOT NULL" + " AND (instr(m.tool_calls, 'skill_view') > 0" + " OR instr(m.tool_calls, 'skill_manage') > 0)" + ) + _GET_SKILL_CALLS_ALL = ( + "SELECT m.tool_calls, m.timestamp" + f" FROM messages m INDEXED BY {_MESSAGES_ASSISTANT_CALLS_INDEX}" + " JOIN sessions s ON s.id = m.session_id" + " WHERE s.started_at >= ?" + " AND m.role = 'assistant' AND m.tool_calls IS NOT NULL" + " AND (instr(m.tool_calls, 'skill_view') > 0" + " OR instr(m.tool_calls, 'skill_manage') > 0)" + ) + def _get_sessions(self, cutoff: float, source: str = None) -> List[Dict]: """Fetch sessions within the time window.""" if source: @@ -243,22 +330,10 @@ class InsightsEngine: # (covers CLI sessions where tool_name is NULL on tool responses) if source: cursor2 = self._conn.execute( - """SELECT m.tool_calls - FROM messages m - JOIN sessions s ON s.id = m.session_id - WHERE s.started_at >= ? AND s.source = ? - AND m.role = 'assistant' AND m.tool_calls IS NOT NULL""", - (cutoff, source), + self._GET_TOOL_CALLS_WITH_SOURCE, (cutoff, source) ) else: - cursor2 = self._conn.execute( - """SELECT m.tool_calls - FROM messages m - JOIN sessions s ON s.id = m.session_id - WHERE s.started_at >= ? - AND m.role = 'assistant' AND m.tool_calls IS NOT NULL""", - (cutoff,), - ) + cursor2 = self._conn.execute(self._GET_TOOL_CALLS_ALL, (cutoff,)) tool_calls_counts = Counter() for row in cursor2.fetchall(): @@ -301,22 +376,10 @@ class InsightsEngine: if source: cursor = self._conn.execute( - """SELECT m.tool_calls, m.timestamp - FROM messages m - JOIN sessions s ON s.id = m.session_id - WHERE s.started_at >= ? AND s.source = ? - AND m.role = 'assistant' AND m.tool_calls IS NOT NULL""", - (cutoff, source), + self._GET_SKILL_CALLS_WITH_SOURCE, (cutoff, source) ) else: - cursor = self._conn.execute( - """SELECT m.tool_calls, m.timestamp - FROM messages m - JOIN sessions s ON s.id = m.session_id - WHERE s.started_at >= ? - AND m.role = 'assistant' AND m.tool_calls IS NOT NULL""", - (cutoff,), - ) + cursor = self._conn.execute(self._GET_SKILL_CALLS_ALL, (cutoff,)) for row in cursor.fetchall(): try: diff --git a/agent/memory_provider.py b/agent/memory_provider.py index 4210a4c252..559fc3df6c 100644 --- a/agent/memory_provider.py +++ b/agent/memory_provider.py @@ -34,12 +34,50 @@ Optional hooks (override to opt in): from __future__ import annotations import logging +import re from abc import ABC, abstractmethod from typing import Any, Dict, List, Optional logger = logging.getLogger(__name__) +# Prompts that carry no semantic signal — trivial acknowledgements, greetings, +# slash commands, empty input. Single source of truth shared by the core +# per-turn prefetch gate (agent/turn_context.py, run_agent.py) and provider- +# side classifiers (plugins/memory/honcho) so the two can never drift apart. +# The alternation is anchored and may only be followed by whitespace or +# punctuation, so words that merely START with a trivial word ("k8s", "yolo", +# "note", "hindsight") do NOT match, while trailing-punctuation variants +# ("hi!", "hey.", "thanks :)", "done???") do. +TRIVIAL_PROMPT_RE = re.compile( + r'^(yes|no|ok|okay|sure|thanks|thank you|y|n|yep|nope|yeah|nah|' + r'hi|hey|hello|yo|sup|' + r'continue|go ahead|do it|proceed|got it|cool|nice|great|done|next|lgtm|k)' + r'[\s!?.:;,"' + "'" + r'~\u2018\u2019\u201c\u201d\u2014\u2013\u2026()\[\]{}<>*&^%$#@!+=`\u00a0]*$', + re.IGNORECASE, +) + + +def is_trivial_prompt(text: Optional[str]) -> bool: + """Return True if a user prompt is too trivial to warrant memory recall. + + Empty/whitespace-only input, slash commands, and bare greetings or + acknowledgements (with optional trailing punctuation) all count as + trivial. Callers use this to skip memory-provider prefetch/injection + on turns that carry no semantic signal — saving a blocking network + round-trip and preventing stale user-model context from derailing + one-word replies. + """ + if not text: + return True + stripped = text.strip() + if not stripped: + return True + if stripped.startswith("/"): + return True + return bool(TRIVIAL_PROMPT_RE.match(stripped)) + + class MemoryProvider(ABC): """Abstract base class for memory providers.""" @@ -253,6 +291,10 @@ class MemoryProvider(ABC): required: True if required (default: False) default: default value (optional) choices: list of valid values (optional) + type: text, integer, number, or boolean (optional) + minimum: numeric lower bound for integer/number fields (optional) + maximum: numeric upper bound for integer/number fields (optional) + step: numeric input step for Dashboard rendering (optional) url: URL where user can get this credential (optional) env_var: explicit env var name for secrets (default: auto-generated) diff --git a/agent/moa_loop.py b/agent/moa_loop.py index 26e9523ec3..0e15492a8f 100644 --- a/agent/moa_loop.py +++ b/agent/moa_loop.py @@ -12,6 +12,7 @@ import hashlib import logging import re import threading +import time from concurrent.futures import ThreadPoolExecutor, wait as _futures_wait from types import SimpleNamespace from typing import Any @@ -151,6 +152,24 @@ def _redact_trace_accounting(acct: Any) -> Any: ) +# Cold-start caches. A MoA preset switch used to re-resolve the full +# config + preset + every slot's provider runtime on EACH create() call +# (once per tool-loop iteration), serially before the parallel fan-out could +# start — adding 5-30s of "frozen" latency on complex presets +# (#66793). The preset structure is immutable for the life of a turn, so +# cache both the resolved preset and each (provider, model) runtime. +_preset_cache_lock = threading.Lock() +_preset_cache: dict[tuple, Any] = {} + +_runtime_cache_lock = threading.Lock() +_runtime_cache: dict[tuple[str, str], tuple[float, dict[str, Any]]] = {} + +# Runtime entries go stale when providers/credentials change (key rotation, +# base_url edits). Deliberately short-lived: 300s collapses the per-iteration +# re-resolution inside a turn while bounding credential staleness between +# turns — the non-MoA path picks up rotated keys immediately, this path +# within 5 minutes. +_RUNTIME_CACHE_TTL_SECONDS = 300.0 # Upper bound on concurrent reference-model calls. References are independent # advisory calls (no tools, no inter-dependence), so we fan them out the same @@ -322,34 +341,33 @@ def _slot_runtime(slot: dict[str, Any]) -> dict[str, Any]: api_key resolver the CLI, gateway, and delegate_task all use), so the slot gets its provider's real API surface — e.g. MiniMax → anthropic_messages, GPT-5/o-series → max_completion_tokens, custom endpoints → their base_url. - Returns the kwargs to pass through to ``call_llm`` (provider/model plus the resolved base_url/api_key when available). Falls back to the bare provider/model on any resolution error so a misconfigured slot still attempts the call rather than aborting the whole MoA turn. + + The resolved runtime is cached per (provider, model) with a short TTL + (``_RUNTIME_CACHE_TTL_SECONDS``): the resolution does real I/O (catalog + query + config read) that used to run serially per create() call before + the parallel fan-out could start — the dominant source of MoA cold-start + latency (#66793). The TTL bounds credential staleness (key rotation, + base_url edits) instead of caching for the process lifetime. """ provider = str(slot.get("provider") or "").strip() model = str(slot.get("model") or "").strip() + cache_key = (provider, model) + now = time.monotonic() + with _runtime_cache_lock: + entry = _runtime_cache.get(cache_key) + if entry is not None: + stamped_at, cached = entry + if now - stamped_at < _RUNTIME_CACHE_TTL_SECONDS: + return cached out: dict[str, Any] = {"provider": provider, "model": model} try: from hermes_cli.runtime_provider import resolve_runtime_provider rt = resolve_runtime_provider(requested=provider, target_model=model) - # Forward the resolved endpoint through to call_llm unconditionally. - # call_llm's _resolve_task_provider_model() is the single chokepoint that - # decides whether an explicit base_url collapses a call to the generic - # ``custom`` route or keeps the provider's real identity: it preserves - # identity for any first-class provider (via - # _preserve_provider_with_base_url, a provider-catalog capability check), - # so provider branches that add auth refresh / request metadata / - # request-shape adapters — anthropic OAuth (Bearer + anthropic-beta), - # openai-codex Responses wrapping + Cloudflare headers, xai-oauth, - # bedrock SigV4 signing, nous Portal tags — still fire. Those branches - # re-resolve their own credentials by name and ignore a forwarded - # base_url/api_key, so forwarding is safe even for a placeholder key - # (bedrock's "aws-sdk"). We used to maintain a name-preservation set here - # too; that duplicated the chokepoint and drifted out of sync, so the - # single source of truth now lives in call_llm. if rt.get("base_url"): out["base_url"] = rt["base_url"] if rt.get("api_key"): @@ -362,7 +380,14 @@ def _slot_runtime(slot: dict[str, Any]) -> dict[str, Any]: if isinstance(extra_body, dict) and extra_body: out["extra_body"] = dict(extra_body) except Exception as exc: # pragma: no cover - defensive - logger.debug("MoA slot runtime resolution failed for %s: %s", _slot_label(slot), exc) + logger.debug("MoA slot runtime resolution failed for %s: %s", + _slot_label(slot), exc) + # Never cache a fallback-shaped result: a transient resolution error + # (config mid-write, catalog hiccup) would otherwise pin the bare + # provider/model kwargs for a full TTL. + return out + with _runtime_cache_lock: + _runtime_cache[cache_key] = (now, out) return out @@ -1830,11 +1855,35 @@ class MoAChatCompletions: raise TypeError("_moa_prepared_request must be a dict") return self._call_prepared_aggregator(prepared_request, api_kwargs) - from hermes_cli.config import load_config + from hermes_cli.config import get_config_path, load_config from hermes_cli.moa_config import resolve_moa_preset + # Resolve the preset once per (config st_mtime_ns, preset_name). + # resolve_moa_preset re-normalizes + re-validates the whole moa + # config block on every call, and create() runs once per tool-loop + # iteration — a serial cold-start cost before the parallel fan-out + # can begin (#66793). Keyed on the config FILE's mtime_ns (not a + # config-object attribute, which load_config()'s dicts don't carry), + # so a config edit invalidates on the next call. + try: + _cfg_stamp = get_config_path().stat().st_mtime_ns + except OSError: + _cfg_stamp = None + # load_config() is itself (mtime_ns, size)-cached upstream, so this + # read is cheap; the expensive part this cache skips is + # resolve_moa_preset's re-normalization + re-validation. _moa_raw = load_config().get("moa") or {} - preset = resolve_moa_preset(_moa_raw, self.preset_name) + preset_cache_key = (_cfg_stamp, self.preset_name) + preset = None + if _cfg_stamp is not None: + with _preset_cache_lock: + preset = _preset_cache.get(preset_cache_key) + if preset is None: + preset = resolve_moa_preset(_moa_raw, self.preset_name) + if _cfg_stamp is not None: + with _preset_cache_lock: + _preset_cache.clear() # one live config stamp at a time + _preset_cache[preset_cache_key] = preset # Privacy filter mode: '' (off, default) | 'display' | 'full'. See # coerce_privacy_filter / the pattern block at the top of this module. # Remembered on self so _call_prepared_aggregator (which may run on a diff --git a/agent/model_metadata.py b/agent/model_metadata.py index f89d25f4c8..b1163c76c4 100644 --- a/agent/model_metadata.py +++ b/agent/model_metadata.py @@ -145,6 +145,96 @@ _ENDPOINT_MODEL_CACHE_TTL = 300 _ENDPOINT_PROBE_TTL_SECONDS = 3600.0 _endpoint_probe_path_cache: Dict[str, tuple] = {} +# A configured endpoint that is routable-but-dead — e.g. a corp LAN address +# while off-VPN — blackholes TCP: the SYN draws no SYN-ACK, no RST and no ICMP +# error, so a probe waits out its full timeout instead of failing fast. Startup +# runs a whole waterfall of such probes across several functions here, and the +# stalls stack into a minute-long hang before the banner renders. +# +# Once ANY probe has actually observed a connect timeout for an endpoint, the +# others have nothing to gain by repeating it. Recording that observation and +# short-circuiting on it performs no network I/O of its own — it adds no probe +# for callers or tests to mock, and it can only ever fire after a real timeout +# has already been paid, so it cannot suppress a probe that would have worked. +_ENDPOINT_BLACKHOLE_TTL_SECONDS = 30.0 +# Values are monotonic timestamps of the last observed connect timeout. +_endpoint_blackhole_cache: Dict[str, float] = {} + + +def _endpoint_host_key(base_url: str) -> Optional[str]: + """Return a ``host:port`` key for ``base_url``, or None if it has no host. + + Keyed on host:port rather than the full URL so every probe path for one + server — ``/v1``-suffixed or not, LM Studio root or API root — shares a + single entry. + """ + normalized = _normalize_base_url(base_url) + if not normalized: + return None + url = normalized if "://" in normalized else f"http://{normalized}" + try: + parsed = urlparse(url) + host = parsed.hostname + port = parsed.port or (443 if parsed.scheme == "https" else 80) + except Exception: + return None + return f"{host}:{port}" if host else None + + +def _note_endpoint_blackholed(base_url: str) -> None: + """Record that a probe to ``base_url`` timed out during TCP connect.""" + key = _endpoint_host_key(base_url) + if key is None: + return + _endpoint_blackhole_cache[key] = time.monotonic() + logger.debug( + "Endpoint %s timed out connecting — skipping further probes for %.0fs", + key, _ENDPOINT_BLACKHOLE_TTL_SECONDS, + ) + + +def _endpoint_blackholed(base_url: str) -> bool: + """True if a recent probe to ``base_url`` timed out during TCP connect. + + Pure cache lookup; never touches the network. The entry expires after + _ENDPOINT_BLACKHOLE_TTL_SECONDS — long enough to collapse one startup's + burst of probes, short enough that bringing the VPN up mid-session is + picked up without a restart. + """ + if _ENDPOINT_BLACKHOLE_TTL_SECONDS <= 0: + return False + key = _endpoint_host_key(base_url) + if key is None: + return False + seen = _endpoint_blackhole_cache.get(key) + if seen is None: + return False + if (time.monotonic() - seen) >= _ENDPOINT_BLACKHOLE_TTL_SECONDS: + del _endpoint_blackhole_cache[key] + return False + return True + + +def _is_connect_timeout(exc: BaseException) -> bool: + """True for connect-phase timeouts raised by httpx or requests. + + Read timeouts are deliberately excluded: those mean the server accepted + the connection, which is the opposite of the blackhole this guards. + """ + try: + import httpx + if isinstance(exc, httpx.ConnectTimeout): + return True + except Exception: + pass + try: + from requests.exceptions import ConnectTimeout + if isinstance(exc, ConnectTimeout): + return True + except Exception: + pass + return False + # ── Disk L2 for local-endpoint probe results ──────────────────────────────── # The in-process caches above die with the process, so every CLI cold start # with a local model re-paid the probe waterfall in AIAgent.__init__: @@ -382,6 +472,7 @@ DEFAULT_CONTEXT_LENGTHS = { "llama": 131072, # Qwen — specific model families before the catch-all. # Official docs: https://help.aliyun.com/zh/model-studio/developer-reference/ + "qwen3.8-max": 1_000_000, # 1M context (OpenRouter & Nous portal, verified 2026-08-03) "qwen3.6-plus": 1048576, # 1M context (DashScope/Alibaba & OpenRouter) "qwen3.7-plus": 1048576, # 1M context (DashScope/Alibaba) "qwen3-coder-plus": 1000000, # 1M context @@ -836,7 +927,10 @@ def _localhost_to_ipv4(url: str) -> str: ``http://localhost...`` (e.g. ``?upstream=http://localhost:11434``) passes through untouched. """ - if not url: + if not url or not isinstance(url, str): + # Non-string values (test doubles, lazily-resolved config objects) + # previously flowed through these call sites untouched — keep that + # contract; re.sub would raise TypeError. return url return re.sub( r"^(https?://)localhost(?=[:/]|$)", @@ -873,6 +967,13 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]: if cached is not None and (time.monotonic() - cached[1]) < _ENDPOINT_PROBE_TTL_SECONDS: return cached[0] + # The host already blackholed a connect: skip the waterfall below, each leg + # of which would otherwise burn its full 2s timeout. Deliberately NOT + # written to _endpoint_probe_path_cache — that entry lives for an hour, + # which would pin the endpoint to "undetected" long after it comes back. + if _endpoint_blackholed(server_url): + return None + # Disk L2: a fresh cross-process verdict skips the HTTP waterfall # entirely (back-to-back CLI invocations, cron ticks). disk_hit = _local_probe_disk_get("server_type", server_url) @@ -882,6 +983,16 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]: headers = _auth_headers(api_key) + def _probe_failed(exc: Exception) -> None: + """Swallow a probe error — or abort the waterfall if we were blackholed. + + Re-raising propagates out of the ``with`` block to the outer handler, + so the remaining legs are skipped instead of each stalling in turn. + """ + if _is_connect_timeout(exc): + _note_endpoint_blackholed(server_url) + raise exc + result: Optional[str] = None try: with httpx.Client(timeout=2.0, headers=headers) as client: @@ -890,8 +1001,8 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]: r = client.get(f"{lmstudio_url}/api/v1/models") if r.status_code == 200: result = "lm-studio" - except Exception: - pass + except Exception as exc: + _probe_failed(exc) if result is None: # Ollama exposes /api/tags and responds with {"models": [...]} # LM Studio returns {"error": "Unexpected endpoint"} with status 200 @@ -905,8 +1016,8 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]: result = "ollama" except Exception: pass - except Exception: - pass + except Exception as exc: + _probe_failed(exc) if result is None: # llama.cpp exposes /v1/props (older builds used /props without the /v1 prefix) try: @@ -915,8 +1026,8 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]: r = client.get(f"{server_url}/props") # fallback for older builds if r.status_code == 200 and "default_generation_settings" in r.text: result = "llamacpp" - except Exception: - pass + except Exception as exc: + _probe_failed(exc) if result is None: # vLLM: /version try: @@ -925,8 +1036,8 @@ def detect_local_server_type(base_url: str, api_key: str = "") -> Optional[str]: data = r.json() if "version" in data: result = "vllm" - except Exception: - pass + except Exception as exc: + _probe_failed(exc) except Exception: pass @@ -1120,6 +1231,12 @@ def fetch_endpoint_model_metadata( if cached is not None and (time.time() - cached_at) < _ENDPOINT_MODEL_CACHE_TTL: return cached + # Blackholed endpoint: every candidate below would spend its full 5s + # connect budget. Returned empty rather than cached, so the endpoint is + # retried as soon as the blackhole entry expires. + if _endpoint_blackholed(normalized): + return {} + candidates = [normalized] if normalized.endswith("/v1"): alternate = normalized[:-3].rstrip("/") @@ -1182,9 +1299,20 @@ def fetch_endpoint_model_metadata( return cache except Exception as exc: last_error = exc + if _is_connect_timeout(exc): + _note_endpoint_blackholed(normalized) for candidate in candidates: - url = candidate.rstrip("/") + "/models" + # A connect timeout on one candidate condemns the host, not the path: + # the remaining candidates differ only by URL suffix, so trying them + # would repeat the same stall. + if _endpoint_blackholed(normalized): + break + # normalized/candidates stay unrewritten (cache key stability); only + # the outbound request target is IPv4-resolved to skip the multi-second + # dual-stack IPv6 connect timeout (see _localhost_to_ipv4). + request_candidate = _localhost_to_ipv4(candidate) + url = request_candidate.rstrip("/") + "/models" response = None try: response = requests.get( @@ -1230,7 +1358,7 @@ def fetch_endpoint_model_metadata( if is_llamacpp: try: # Try /v1/props first (current llama.cpp); fall back to /props for older builds - base = candidate.rstrip("/").replace("/v1", "") + base = request_candidate.rstrip("/").replace("/v1", "") _verify = _resolve_requests_verify() props_resp = requests.get(base + "/v1/props", headers=headers, timeout=5, verify=_verify) if not props_resp.ok: @@ -1250,6 +1378,8 @@ def fetch_endpoint_model_metadata( return cache except Exception as exc: last_error = exc + if _is_connect_timeout(exc): + _note_endpoint_blackholed(normalized) finally: if response is not None: response.close() @@ -1813,6 +1943,9 @@ def _query_ollama_api_show_uncached(model: str, base_url: str, api_key: str = "" if server_url.endswith("/v1"): server_url = server_url[:-3] + if _endpoint_blackholed(server_url): + return None + headers = _auth_headers(api_key) try: @@ -1844,8 +1977,9 @@ def _query_ollama_api_show_uncached(model: str, base_url: str, api_key: str = "" return ctx except ValueError: pass - except Exception: - pass + except Exception as exc: + if _is_connect_timeout(exc): + _note_endpoint_blackholed(server_url) return None @@ -1933,6 +2067,9 @@ def _query_local_context_length_uncached(model: str, base_url: str, api_key: str server_url = server_url[:-3] lmstudio_url = _localhost_to_ipv4(_lmstudio_server_root(base_url)) + if _endpoint_blackholed(server_url): + return None + headers = _auth_headers(api_key) try: @@ -2003,13 +2140,39 @@ def _query_local_context_length_uncached(model: str, base_url: str, api_key: str if resp.status_code == 200: data = resp.json() models_list = data.get("data", []) + # Match by id; on single-model servers (e.g. llama.cpp) the + # configured name rarely equals the reported id (a GGUF path), + # so fall back to the sole model when nothing matches. + matched = None for m in models_list: if _model_id_matches(m.get("id", ""), model): - ctx = m.get("max_model_len") or m.get("context_length") or m.get("max_tokens") - if ctx and isinstance(ctx, (int, float)): - return int(ctx) - except Exception: - pass + matched = m + break + if matched is None and len(models_list) == 1: + matched = models_list[0] + if matched is not None: + # llama.cpp nests the runtime context under meta.n_ctx; the + # vLLM/OpenAI keys are also checked. Runtime n_ctx is + # preferred over n_ctx_train (the training maximum, which + # can be larger than what the server actually allocates). + for source in (matched, matched.get("meta") or {}): + if not isinstance(source, dict): + continue + for key in ( + "n_ctx", + "context_length", + "context_window", + "max_model_len", + "max_context_length", + "max_tokens", + "n_ctx_train", + ): + val = source.get(key) + if isinstance(val, (int, float)) and val: + return int(val) + except Exception as exc: + if _is_connect_timeout(exc): + _note_endpoint_blackholed(server_url) return None diff --git a/agent/relay_runtime.py b/agent/relay_runtime.py index 533604791a..a0a7315796 100644 --- a/agent/relay_runtime.py +++ b/agent/relay_runtime.py @@ -482,11 +482,9 @@ class RelayTurnContext: default_factory=threading.RLock, repr=False, ) - _token: contextvars.Token[RelayTurnContext | None] | None = field( - default=None, - repr=False, - ) + _previous_turn: RelayTurnContext | None = field(default=None, repr=False) _active_registered: bool = field(default=False, repr=False) + relay_enabled: bool = True closed: bool = False @@ -600,7 +598,28 @@ class RelaySessionCoordinator: if lease.released: raise RuntimeError("Hermes Relay conversation lease is released") turn = RelayTurnContext(lease=lease, turn_id=turn_id, task_id=task_id) - if isinstance(lease.host, RelayRuntime) and lease.session is not None: + key = (lease.profile_key, lease.session_id) + with self._active_turns_lock: + active = self._active_turns.get(key) + if active: + # A Relay session owns one physical scope stack. Concurrent + # Hermes turns would create sibling scopes on that stack, but + # their completion order is not guaranteed to be LIFO. + turn.relay_enabled = False + logger.warning( + "Skipping Relay instrumentation for concurrent Hermes turn " + "%s in session %s", + turn_id, + lease.session_id, + ) + else: + self._active_turns[key] = {id(turn)} + turn._active_registered = True + if ( + turn.relay_enabled + and isinstance(lease.host, RelayRuntime) + and lease.session is not None + ): try: turn.handle = lease.host.run_in_session( lease.session, @@ -617,11 +636,8 @@ class RelaySessionCoordinator: ) except Exception: logger.warning("Hermes Relay turn initialization failed", exc_info=True) - turn._token = _CURRENT_TURN.set(turn) - key = (lease.profile_key, lease.session_id) - with self._active_turns_lock: - self._active_turns.setdefault(key, set()).add(id(turn)) - turn._active_registered = True + turn._previous_turn = _CURRENT_TURN.get() + _CURRENT_TURN.set(turn) return turn def end_turn( @@ -755,16 +771,18 @@ class RelaySessionCoordinator: @staticmethod def _reset_turn_context(turn: RelayTurnContext) -> None: - """Reset the originating ContextVar token when called in that context.""" - if turn._token is None: + """Unwind ``turn`` without disturbing a newer context-local turn.""" + if _CURRENT_TURN.get() is not turn: return - try: - _CURRENT_TURN.reset(turn._token) - except ValueError: - # A copied async/thread context may own terminal cleanup. Keep the - # token so the originating context can clear its stale reference. - return - turn._token = None + previous = turn._previous_turn + seen = {id(turn)} + while previous is not None and previous.closed: + if id(previous) in seen: + previous = None + break + seen.add(id(previous)) + previous = previous._previous_turn + _CURRENT_TURN.set(previous) @staticmethod def release_conversation(lease: ConversationLease) -> None: @@ -793,10 +811,21 @@ def current_turn() -> RelayTurnContext | None: return _CURRENT_TURN.get() +def relay_instrumentation_enabled() -> bool: + """Return whether this inherited turn may create Relay instrumentation.""" + turn = current_turn() + return turn is None or (turn.relay_enabled and not turn.closed) + + def active_turn(session_id: str | None = None) -> RelayTurnContext | None: """Return a live turn only when it belongs to the active profile/session.""" turn = current_turn() - if turn is None or turn.closed or turn.lease.released: + if ( + turn is None + or not turn.relay_enabled + or turn.closed + or turn.lease.released + ): return None if turn.lease.profile_key != current_profile_key(): return None @@ -814,6 +843,11 @@ def resolve_execution_context( session_id: str, ) -> tuple[RelayRuntime | None, RelaySession | None, Any]: """Resolve one active turn/session parent for managed Relay execution.""" + inherited_turn = current_turn() + if inherited_turn is not None and ( + not inherited_turn.relay_enabled or inherited_turn.closed + ): + return None, None, None turn = active_turn(session_id) if ( turn is not None diff --git a/agent/subdirectory_hints.py b/agent/subdirectory_hints.py index ca96c664cb..4e9f7f5ed3 100644 --- a/agent/subdirectory_hints.py +++ b/agent/subdirectory_hints.py @@ -13,6 +13,7 @@ the conversation without modifying the system prompt (preserving prompt caching) Inspired by Block/goose's SubdirectoryHintTracker. """ +import hashlib import logging import os import shlex @@ -45,6 +46,18 @@ _COMMAND_TOOLS = {"terminal"} # Prevents scanning all the way to / for deeply nested paths. _MAX_ANCESTOR_WALK = 5 +# Directory names that never contain authoritative project context. +# Backups, vendored deps, VCS internals, and caches routinely hold *copies* of +# AGENTS.md; loading those duplicates real context and inflates the prompt. +_EXCLUDED_DIR_NAMES = frozenset({ + "node_modules", "venv", ".venv", "__pycache__", + ".git", ".hg", ".svn", + ".Trash", ".cache", ".tox", ".mypy_cache", ".pytest_cache", + "site-packages", "dist-packages", + "backups", "backup", ".backups", + "vendor", "third_party", +}) + def _is_ancestor_or_same(a: Path, b: Path) -> bool: """Check if *a* is the same as or an ancestor of *b* (parent directory check).""" @@ -54,6 +67,7 @@ def _is_ancestor_or_same(a: Path, b: Path) -> bool: except ValueError: return False + class SubdirectoryHintTracker: """Track which directories the agent visits and load hints on first access. @@ -70,8 +84,34 @@ class SubdirectoryHintTracker: def __init__(self, working_dir: Optional[str] = None): self.working_dir = Path(working_dir or os.getcwd()).resolve() self._loaded_dirs: Set[Path] = set() + # Content digests already injected — prevents re-sending the same file + # reachable through symlinks, hardlinks, or duplicated copies. + self._loaded_digests: Set[str] = set() # Pre-mark the working dir as loaded (startup context handles it) self._loaded_dirs.add(self.working_dir) + self._seed_working_dir_digest() + + def _seed_working_dir_digest(self) -> None: + """Record the CWD context file's digest so it is never re-injected. + + ``prompt_builder`` already loads the working directory's context file at + startup. Seeding its digest here means the same content reached through + a different path (a symlink farm, a shared workspace) is recognised as a + duplicate instead of being sent a second time. + """ + for filename in _HINT_FILENAMES: + candidate = self.working_dir / filename + try: + if not candidate.is_file(): + continue + content = candidate.read_text(encoding="utf-8").strip() + except (OSError, UnicodeDecodeError): + continue + if content: + self._loaded_digests.add( + hashlib.sha256(content.encode("utf-8")).hexdigest() + ) + break # first match wins, mirroring startup loading def check_tool_call( self, @@ -193,8 +233,25 @@ class SubdirectoryHintTracker: # check as a best-effort safeguard. if not _is_ancestor_or_same(self.working_dir, path): return False + if self._is_excluded(path): + return False return True + def _is_excluded(self, path: Path) -> bool: + """True when the path sits inside a directory that holds copies, not context. + + Directories the user is deliberately working inside are never excluded — + if ``working_dir`` is itself under ``vendor/``, that segment is legitimate + and only segments *below* the working dir are screened. + """ + try: + rel_parts = path.relative_to(self.working_dir).parts + except ValueError: + # Paths outside the working dir are already rejected by + # _is_valid_subdir before this runs; treat as excluded defensively. + return True + return any(part in _EXCLUDED_DIR_NAMES for part in rel_parts) + def _load_hints_for_directory(self, directory: Path) -> Optional[str]: """Load hint files from a directory. Returns formatted text or None. @@ -230,6 +287,19 @@ class SubdirectoryHintTracker: content = hint_path.read_text(encoding="utf-8").strip() if not content: continue + # Skip content we've already injected. The same AGENTS.md is + # routinely reachable through several paths (symlinked shared + # workspaces, hardlinks, copied backups); re-sending it burns + # context for zero new information. + digest = hashlib.sha256(content.encode("utf-8")).hexdigest() + if digest in self._loaded_digests: + logger.debug( + "Skipping duplicate hint content at %s (digest %s)", + hint_path, + digest[:12], + ) + break + self._loaded_digests.add(digest) # Same security scan as startup context loading content = _scan_context_content(content, filename) if len(content) > _MAX_HINT_CHARS: diff --git a/agent/system_prompt.py b/agent/system_prompt.py index 8b8832ca28..8407a89689 100644 --- a/agent/system_prompt.py +++ b/agent/system_prompt.py @@ -11,14 +11,14 @@ Three tiers are joined with ``\\n\\n``: * ``stable`` — identity (SOUL.md or DEFAULT_AGENT_IDENTITY), tool guidance, computer-use guidance, nous subscription block, tool-use - enforcement guidance + per-model operational guidance, skills prompt, + enforcement guidance + per-model operational guidance, alibaba model-name workaround, environment hints, coding guidance, platform hints. * ``context`` — caller-supplied ``system_message`` plus context files (AGENTS.md / .cursorrules / etc.) discovered under ``TERMINAL_CWD``, plus the session's coding-workspace snapshot. -* ``volatile`` — memory snapshot, USER.md profile, external memory - provider block, timestamp/session/model/provider line. +* ``volatile`` — skills index, memory snapshot, USER.md profile, external + memory provider block, timestamp/session/model/provider line. Pure helpers that read the agent's state. AIAgent keeps thin forwarders. """ @@ -158,8 +158,8 @@ def build_system_prompt_parts(agent: Any, system_message: Optional[str] = None) * ``context`` — the workspace snapshot followed by the remaining session-stable guidance, context files, and caller-supplied system_message. - * ``volatile`` — memory snapshot, user profile, external - memory provider block, timestamp line. + * ``volatile`` — skills index, memory snapshot, user profile, + external memory provider block, timestamp line. Joined into a single string by :func:`build_system_prompt` and cached on ``agent._cached_system_prompt`` for the lifetime of the @@ -325,8 +325,6 @@ def build_system_prompt_parts(agent: Any, system_message: Optional[str] = None) ) else: skills_prompt = "" - if skills_prompt: - stable_parts.append(skills_prompt) # Alibaba Coding Plan API always returns "glm-4.7" as model name regardless # of the requested model. Inject explicit model identity into the system prompt @@ -497,8 +495,22 @@ def build_system_prompt_parts(agent: Any, system_message: Optional[str] = None) if context_files_prompt: context_parts.append(context_files_prompt) - # ── Volatile tier (changes per session/turn — never cached) ─── + # ── Volatile tier (most likely to differ on a rebuild; kept last so the stable prefix stays reusable) ── volatile_parts: List[str] = [] + # Skills are runtime-mutable: the agent adds and patches them across a + # session (SKILLS_GUIDANCE tells it to patch a skill the moment it goes + # stale). The built prompt is cached per session and only rebuilt on + # compaction/restore (see build_system_prompt), so a skill change is not + # byte-stable across rebuilds. With the index in the stable band, a rebuild + # that picked up a skill change would bust the cached prefix from the index + # down, taking the whole scaffold with it. Render it at the FRONT of the + # volatile band instead, ahead of the turn-varying memory/timestamp tail: + # on an implicit longest-prefix backend an unchanged index still falls + # inside the reused prefix, and a changed one only re-prefills from here on. + # (No effect for single-block cache_control backends, where the whole + # system message is one cache unit regardless of internal order.) + if skills_prompt: + volatile_parts.append(skills_prompt) if agent._memory_store: if agent._memory_enabled: @@ -556,10 +568,12 @@ def build_system_prompt(agent: Any, system_message: Optional[str] = None) -> str Layers are ordered cache-friendly: stable identity/guidance first, then session-stable context files, then per-call volatile content - (memory, USER profile, timestamp). The whole string is treated as - one cached block — Hermes never rebuilds or reinjects parts of it - mid-session, which is the only way to keep upstream prompt caches - warm across turns. + (skills index, memory, USER profile, timestamp). For explicit + cache_control backends the whole string is one cached block. For + implicit longest-prefix backends the order is what matters: the + content most likely to change is rendered last, so when the prompt is + rebuilt (on compaction/restore) the unchanged stable scaffold ahead of + the change stays in the reused prefix. """ parts = build_system_prompt_parts(agent, system_message=system_message) joined = "\n\n".join(p for p in (parts["stable"], parts["context"], parts["volatile"]) if p) @@ -602,8 +616,8 @@ def reconstruct_static_prefix( Safety: the rebuilt stable tier is used ONLY when the stored prompt literally starts with it (checked here AND re-checked by ``_apply_system_cache_markers``'s ``startswith`` gate). If any - stable-tier input changed since the prompt was persisted (skills - edited, identity changed), the prefix mismatches, the static stays + stable-tier input changed since the prompt was persisted (identity + changed, SOUL.md edited), the prefix mismatches, the static stays None, and requests fall back to the legacy layout with the stored prompt bytes untouched — never a rewritten prompt. diff --git a/agent/tool_dispatch_helpers.py b/agent/tool_dispatch_helpers.py index f7f003f24a..6e76e52081 100644 --- a/agent/tool_dispatch_helpers.py +++ b/agent/tool_dispatch_helpers.py @@ -48,6 +48,7 @@ _PARALLEL_SAFE_TOOLS = frozenset({ "ha_get_state", "ha_list_entities", "ha_list_services", + "image_generate", "read_file", "search_files", "session_search", diff --git a/agent/tool_executor.py b/agent/tool_executor.py index ee2daa07c4..e1a3f3013d 100644 --- a/agent/tool_executor.py +++ b/agent/tool_executor.py @@ -93,6 +93,7 @@ def _budget_for_agent(agent) -> BudgetConfig: # Maximum number of concurrent worker threads for parallel tool execution. # Mirrors the constant in ``run_agent`` for tests/imports that look here. _MAX_TOOL_WORKERS = 8 +_DEFAULT_IMAGE_PARALLEL_REQUESTS = 4 # Keep this above the stock auxiliary.web_extract timeout (360s) so the batch # guard does not preempt a slow-but-valid summarization attempt. _DEFAULT_CONCURRENT_TOOL_TIMEOUT_S = 420.0 @@ -159,6 +160,46 @@ def _flush_session_db_after_tool_progress( return False +def _image_generate_parallel_limit() -> int: + """Return the configured image-generation parallelism cap. + + Image-generation calls are slow enough that concurrent execution is useful, + but backend bursts can hit TTFB or rate-limit failures. Keep the default + intentionally conservative while allowing users to tune it per install. + """ + try: + from hermes_cli.config import load_config + + cfg = load_config() or {} + image_gen = cfg.get("image_gen") if isinstance(cfg, dict) else None + value = ( + image_gen.get("max_parallel_requests") + if isinstance(image_gen, dict) + else None + ) + except Exception: + value = None + + try: + limit = int(value) + except (TypeError, ValueError): + limit = _DEFAULT_IMAGE_PARALLEL_REQUESTS + return max(1, min(limit, _MAX_TOOL_WORKERS)) + + +def _max_workers_for_tool_batch(runnable_calls) -> int: + """Return the worker cap for a concurrent tool batch.""" + if not runnable_calls: + return 0 + max_workers = _MAX_TOOL_WORKERS + if any( + (call[2] if len(call) >= 3 else None) == "image_generate" + for call in runnable_calls + ): + max_workers = min(max_workers, _image_generate_parallel_limit()) + return min(len(runnable_calls), max_workers) + + def _ra(): """Lazy reference to ``run_agent`` so patches like ``run_agent._set_interrupt`` work.""" import run_agent @@ -953,7 +994,7 @@ def execute_tool_calls_concurrent(agent, assistant_message, messages: list, effe timeout_s = _resolve_concurrent_tool_timeout() deadline = time.monotonic() + timeout_s if timeout_s is not None else None if runnable_calls: - max_workers = min(len(runnable_calls), _MAX_TOOL_WORKERS) + max_workers = _max_workers_for_tool_batch(runnable_calls) # Daemon workers: an interrupted/timed-out batch is abandoned with # shutdown(wait=False), but stdlib ThreadPoolExecutor workers are # non-daemon and registered in concurrent.futures' atexit hook, diff --git a/agent/transports/chat_completions.py b/agent/transports/chat_completions.py index 07f7673673..2572038126 100644 --- a/agent/transports/chat_completions.py +++ b/agent/transports/chat_completions.py @@ -9,6 +9,7 @@ which has provider-specific conditionals for max_tokens defaults, reasoning configuration, temperature handling, and extra_body assembly. """ +import json from typing import Any, Dict from agent.lmstudio_reasoning import resolve_lmstudio_effort @@ -18,6 +19,56 @@ from agent.transports.base import ProviderTransport from agent.transports.types import NormalizedResponse, ToolCall, Usage +def _static_prompt_instructions(messages: list[dict[str, Any]]) -> str: + """Return the stable system/developer prefix used for cache routing. + + Chat Completions carries instructions in its message list rather than a + separate ``instructions`` field. Only a leading system/developer message + is static by contract; later messages are conversation state and must not + split a warm prefix bucket on every turn. + """ + if not messages or not isinstance(messages[0], dict): + return "" + first = messages[0] + if first.get("role") not in {"system", "developer"}: + return "" + content = first.get("content") + if isinstance(content, str): + return content + try: + return json.dumps(content, sort_keys=True, ensure_ascii=False, separators=(",", ":")) + except (TypeError, ValueError): + return str(content or "") + + +def _add_prompt_cache_key( + api_kwargs: dict[str, Any], + *, + messages: list[dict[str, Any]], + tools: list[dict[str, Any]] | None, + supports_prompt_cache_key: bool, +) -> None: + """Add a content-addressed key only for an explicitly capable endpoint.""" + if not supports_prompt_cache_key: + return + + # An explicit caller body field is authoritative too. Do not add a + # duplicate top-level field whose SDK merge precedence could overwrite it. + extra_body = api_kwargs.get("extra_body") + if "prompt_cache_key" in api_kwargs or ( + isinstance(extra_body, dict) and "prompt_cache_key" in extra_body + ): + return + + # Reuse the Responses transport's single authoritative hash algorithm so + # equivalent static prefixes route to the same cache bucket across modes. + from agent.transports.codex import _content_cache_key + + cache_key = _content_cache_key(_static_prompt_instructions(messages), tools) + if cache_key: + api_kwargs["prompt_cache_key"] = cache_key + + def _reasoning_config_for_model(model: str, reasoning_config: dict | None) -> dict | None: """Return the model's wire-compatible reasoning config.""" if not isinstance(reasoning_config, dict): @@ -112,6 +163,24 @@ def _is_gemini_openai_compat_base_url(base_url: Any) -> bool: return normalized.endswith("/openai") +def _is_openai_api_base_url(base_url: Any) -> bool: + """True only for api.openai.com itself (exact host). + + OpenAI documents ``prompt_cache_key`` as a first-class body field and + GPT-5.6+ docs recommend it for reliable cache routing, so the flag is + implied for the real endpoint. Deliberately NOT a substring match: + Azure OpenAI and strict OpenAI-compat endpoints may reject unknown + fields and must stay opt-in via ``supports_prompt_cache_key``. + """ + try: + from urllib.parse import urlparse + + host = (urlparse(str(base_url or "").strip()).hostname or "").lower() + except Exception: + return False + return host == "api.openai.com" + + def _model_consumes_thought_signature(model: Any) -> bool: """True when the outgoing model is a Gemini family model that requires ``extra_content`` (thought_signature) to be replayed on tool calls. @@ -327,6 +396,8 @@ class ChatCompletionsTransport(ProviderTransport): # Claude on OpenRouter/Nous max output anthropic_max_output: int | None extra_body_additions: dict | None + supports_prompt_cache_key: bool — explicit endpoint capability for + the top-level Chat Completions request field; defaults off. """ # Codex sanitization: drop reasoning_items / call_id / response_item_id. # Pass model so the Gemini thought_signature (extra_content) is kept for @@ -507,6 +578,14 @@ class ChatCompletionsTransport(ProviderTransport): if overrides: api_kwargs.update(overrides) + _add_prompt_cache_key( + api_kwargs, + messages=sanitized, + tools=api_kwargs.get("tools"), + supports_prompt_cache_key=bool(params.get("supports_prompt_cache_key")) + or _is_openai_api_base_url(params.get("base_url")), + ) + return api_kwargs def _build_kwargs_from_profile(self, profile, model, sanitized, tools, params): @@ -649,6 +728,13 @@ class ChatCompletionsTransport(ProviderTransport): if extra_body: api_kwargs["extra_body"] = extra_body + _add_prompt_cache_key( + api_kwargs, + messages=sanitized, + tools=api_kwargs.get("tools"), + supports_prompt_cache_key=bool(profile.supports_prompt_cache_key), + ) + return api_kwargs def normalize_response(self, response: Any, **kwargs) -> NormalizedResponse: diff --git a/agent/transports/codex.py b/agent/transports/codex.py index de71c2b500..c4c901c259 100644 --- a/agent/transports/codex.py +++ b/agent/transports/codex.py @@ -28,6 +28,52 @@ def _bounded_prompt_cache_key(value: Any) -> Optional[str]: return f"pck_{digest}" +# Wire-name used when Hermes keeps client-side web_search on xAI Responses. +# A function literally named ``web_search`` collides with Grok's native +# server-side tool (incomplete hang or HTTP 400 duplicate names); this alias +# avoids that while still dispatching through Hermes's configured provider +# (Firecrawl / Tavily / …). Mapped back to ``web_search`` in normalize_response. +_XAI_CLIENT_WEB_SEARCH_ALIAS = "hermes_web_search" + + +def _xai_prefers_native_web_search() -> bool: + """True when xAI Responses should use Grok's native ``web_search`` built-in. + + Delegates to the web-search registry's provider resolution (which reads + ``web.search_backend`` / ``web.backend`` from config) and checks whether + the resolved provider is xAI. Falls back to the legacy ``_get_search_backend`` + probe when the registry has no providers loaded. On any resolution failure, + returns True (fail-closed to native — preserves the #48108 incomplete-hang + fix rather than risk reintroducing it). + """ + try: + from agent.web_search_registry import get_active_search_provider + + provider = get_active_search_provider() + if provider is not None: + return getattr(provider, "name", None) == "xai" + + from tools.web_tools import _get_search_backend + + return (_get_search_backend() or "").strip().lower() == "xai" + except Exception: + # Fail closed to native — same behavior as pre-fix main. + return True + + +def _rename_client_web_search_for_xai(response_tools: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + """Rename client ``web_search`` → alias so xAI won't hijack it server-side.""" + rewritten: List[Dict[str, Any]] = [] + for tool in response_tools: + if isinstance(tool, dict) and tool.get("name") == "web_search": + aliased = dict(tool) + aliased["name"] = _XAI_CLIENT_WEB_SEARCH_ALIAS + rewritten.append(aliased) + else: + rewritten.append(tool) + return rewritten + + _EXTENDED_PROMPT_CACHE_MODELS = ( "gpt-5.5-pro", "gpt-5.5", @@ -232,63 +278,41 @@ class ResponsesApiTransport(ProviderTransport): response_tools = _responses_tools(tools) - # xAI server-side web search. + # xAI server-side web search vs Hermes web providers. # - # grok models on xAI's /v1/responses surface (notably - # grok-composer-2.5-fast on SuperGrok OAuth) have a *native*, - # server-executed web search. When the model is handed a - # client-side function literally named ``web_search``, it routes - # the intent to that native engine — but because the tool is - # declared as a plain ``function`` rather than xAI's first-class - # ``{"type": "web_search"}`` built-in, the server-side search is - # dispatched but never reconciled: the response streams reasoning - # + ``web_search_call`` progress items, the searches never reach - # ``status="completed"`` in the assembled output, no final - # message is emitted, and ``_normalize_codex_response`` correctly - # sees reasoning-with-no-answer and reports ``incomplete``. The - # turn then burns 3 continuation retries and fails with "Codex - # response remained incomplete after 3 continuation attempts". - # Verified live against grok-composer-2.5-fast (2026-06). + # grok models on xAI's /v1/responses surface have a *native*, + # server-executed web search. A client-side function literally named + # ``web_search`` collides with that engine: declared as a plain + # ``function`` rather than ``{"type": "web_search"}``, the search + # dispatches but never reconciles → incomplete turn + 3 retries. + # Verified live against grok-composer-2.5-fast (2026-06); see #48108. # - # Fix: when the agent HAS a client-side ``web_search`` function (i.e. - # the user enabled the web toolset), declare xAI's native - # ``web_search`` built-in instead so the search actually runs to - # completion server-side and the model streams a real answer. The - # Responses API rejects two tools sharing the name ``web_search`` - # (HTTP 400 "Duplicate tool names"), so we drop the client-side - # ``web_search`` function for the xAI path and let the native tool - # satisfy it. All other client-side tools (read_file, terminal, - # web_extract, MCP tools, …) are untouched and continue to dispatch - # through Hermes's agent loop. + # Two modes, chosen by the user's web-search backend config: # - # Scope: we ONLY swap in the native built-in when the client - # ``web_search`` was actually present. We do NOT force-enable Grok - # server-side search on turns where the user never had web enabled — - # that would silently route around Hermes's web-provider config and - # tool-trace/citation plumbing for every xai-oauth turn. The swap is - # a 1:1 replacement of an already-requested capability, not an - # additive grant. - # - # NOTE: for the swapped case this routes ``web_search`` to Grok's - # native search engine for xAI sessions instead of Hermes's - # configured web provider (Tavily/etc.), and those results bypass - # Hermes's tool-trace / citation plumbing (they arrive baked into the - # model's answer rather than as a tool result the loop observes). - # Scoped to ``is_xai_responses`` deliberately; narrow to specific - # models if a future grok variant should keep the client-side - # function. + # 1. **Native** (active/configured backend is ``xai``, or resolution + # fails): drop the client ``web_search`` function and declare + # xAI's built-in instead. 1:1 swap only when client ``web_search`` + # was already present — never an additive grant. + # 2. **Client** (Firecrawl / Tavily / Exa / … configured or resolved): + # keep Hermes dispatch so ``web.backend`` / ``web.search_backend`` + # is honored, but rename the wire tool to + # ``hermes_web_search`` so Grok cannot hijack the name. The alias + # is mapped back to ``web_search`` in ``normalize_response``. if is_xai_responses and response_tools: has_client_web_search = any( isinstance(t, dict) and t.get("name") == "web_search" for t in response_tools ) if has_client_web_search: - filtered = [ - t for t in response_tools - if not (isinstance(t, dict) and t.get("name") == "web_search") - ] - filtered.append({"type": "web_search"}) - response_tools = filtered + if _xai_prefers_native_web_search(): + filtered = [ + t for t in response_tools + if not (isinstance(t, dict) and t.get("name") == "web_search") + ] + filtered.append({"type": "web_search"}) + response_tools = filtered + else: + response_tools = _rename_client_web_search_for_xai(response_tools) # ``tools`` MUST be omitted entirely when there are no functions to # expose: the openai SDK's ``responses.stream()`` / ``responses.parse()`` @@ -487,9 +511,14 @@ class ResponsesApiTransport(ProviderTransport): provider_data["call_id"] = tc.call_id if hasattr(tc, "response_item_id") and tc.response_item_id: provider_data["response_item_id"] = tc.response_item_id + name = tc.function.name if hasattr(tc, "function") else getattr(tc, "name", "") + # Undo the xAI client-path wire alias so Hermes dispatches + # the real ``web_search`` tool (Firecrawl / etc.). + if name == _XAI_CLIENT_WEB_SEARCH_ALIAS: + name = "web_search" tool_calls.append(ToolCall( - id=tc.id if hasattr(tc, "id") else (tc.function.name if hasattr(tc, "function") else None), - name=tc.function.name if hasattr(tc, "function") else getattr(tc, "name", ""), + id=tc.id if hasattr(tc, "id") else (name or None), + name=name, arguments=tc.function.arguments if hasattr(tc, "function") else getattr(tc, "arguments", "{}"), provider_data=provider_data or None, )) diff --git a/agent/transports/codex_app_server_session.py b/agent/transports/codex_app_server_session.py index e2ace753bc..384a7de8ab 100644 --- a/agent/transports/codex_app_server_session.py +++ b/agent/transports/codex_app_server_session.py @@ -1259,6 +1259,8 @@ def _approval_choice_to_codex_decision(choice: str) -> str: return "accept" if choice in {"session", "always"}: return "acceptForSession" + # "deny" and "timeout" both map to decline — codex has no wire value for + # "prompt expired"; the Hermes-side messaging already distinguishes them. return "decline" diff --git a/agent/turn_context.py b/agent/turn_context.py index e080d6a5d9..def497b158 100644 --- a/agent/turn_context.py +++ b/agent/turn_context.py @@ -41,6 +41,7 @@ from agent.conversation_compression import ( from agent.context_engine import automatic_compaction_status_message from agent.iteration_budget import IterationBudget from agent.memory_manager import build_memory_context_block +from agent.memory_provider import is_trivial_prompt from agent.model_metadata import ( estimate_messages_tokens_rough, estimate_request_tokens_rough, @@ -1152,11 +1153,15 @@ def build_turn_context( pass # External memory provider: prefetch once before the tool loop. + # + # Skip prefetch on trivial prompts (greetings, acknowledgements) to + # prevent memory-context injection on turns that carry no semantic signal. ext_prefetch_cache = "" if agent._memory_manager: try: _query = original_user_message if isinstance(original_user_message, str) else "" - ext_prefetch_cache = agent._memory_manager.prefetch_all(_query) or "" + if not is_trivial_prompt(_query): + ext_prefetch_cache = agent._memory_manager.prefetch_all(_query) or "" except Exception: pass diff --git a/apps/desktop/e2e/right-pane.spec.ts b/apps/desktop/e2e/right-pane.spec.ts new file mode 100644 index 0000000000..5447955f45 --- /dev/null +++ b/apps/desktop/e2e/right-pane.spec.ts @@ -0,0 +1,138 @@ +import { test, expect } from './test' + +import { + type MockBackendFixture, + setupMockBackend, + waitForAppReady, +} from './fixtures' + +let fixture: MockBackendFixture | null = null + +test.beforeAll(async () => { + fixture = await setupMockBackend() + await waitForAppReady(fixture, 120_000) +}) + +test.afterAll(async () => { + await fixture?.cleanup() + fixture = null +}) + +test('persistent terminal overlay follows the pane after split dragging', async () => { + const page = fixture!.page + + await page.keyboard.press('Control+`') + await page.locator('[data-terminal-slot]').waitFor({ state: 'visible', timeout: 30_000 }) + await page.locator('[data-persistent-terminal] .xterm').waitFor({ state: 'visible', timeout: 30_000 }) + + const result = await page.evaluate(async () => { + const slot = document.querySelector('[data-terminal-slot]') + const overlay = document.querySelector('[data-persistent-terminal]') + + if (!slot || !overlay) { + return { drift: -1, moved: 0, target: false } + } + + const before = slot.getBoundingClientRect() + const target = [...document.querySelectorAll('[role="separator"]')] + .map(element => { + const box = element.getBoundingClientRect() + const horizontal = box.width > box.height + const center = horizontal + ? (box.top + box.bottom) / 2 + : (box.left + box.right) / 2 + const sides = horizontal + ? [before.top, before.bottom] + : [before.left, before.right] + + return { + element, + box, + horizontal, + score: Math.min(...sides.map(side => Math.abs(center - side))), + } + }) + .filter(item => item.box.width > 0 && item.box.height > 0) + .sort((a, b) => a.score - b.score)[0] + + if (!target) { + return { drift: -1, moved: 0, target: false } + } + + const x = target.box.left + target.box.width / 2 + const y0 = target.box.top + target.box.height / 2 + const nearestSide = target.horizontal + ? Math.abs(y0 - before.top) < Math.abs(y0 - before.bottom) + ? 'top' + : 'bottom' + : Math.abs(x - before.left) < Math.abs(x - before.right) + ? 'left' + : 'right' + const deltaX = nearestSide === 'left' ? -1 : nearestSide === 'right' ? 1 : 0 + const deltaY = nearestSide === 'top' ? -1 : nearestSide === 'bottom' ? 1 : 0 + let currentX = x + let y = y0 + const pointer = { + bubbles: true, + cancelable: true, + pointerId: 71, + pointerType: 'mouse', + isPrimary: true, + button: 0, + buttons: 1, + } + + target.element.dispatchEvent( + new PointerEvent('pointerdown', { ...pointer, clientX: x, clientY: y }), + ) + + for (let index = 0; index < 24; index += 1) { + currentX += deltaX + y += deltaY + window.dispatchEvent( + new PointerEvent('pointermove', { + ...pointer, + clientX: currentX, + clientY: y, + }), + ) + await new Promise(resolve => requestAnimationFrame(() => resolve())) + } + + window.dispatchEvent( + new PointerEvent('pointerup', { + ...pointer, + buttons: 0, + clientX: currentX, + clientY: y, + }), + ) + await new Promise(resolve => setTimeout(resolve, 350)) + await new Promise(resolve => + requestAnimationFrame(() => requestAnimationFrame(() => resolve())), + ) + + const next = slot.getBoundingClientRect() + const fixed = overlay.getBoundingClientRect() + + return { + drift: Math.max( + Math.abs(next.top - fixed.top), + Math.abs(next.left - fixed.left), + Math.abs(next.width - fixed.width), + Math.abs(next.height - fixed.height), + ), + moved: Math.max( + Math.abs(next.top - before.top), + Math.abs(next.left - before.left), + Math.abs(next.width - before.width), + Math.abs(next.height - before.height), + ), + target: true, + } + }) + + expect(result.target).toBe(true) + expect(result.moved).toBeGreaterThan(10) + expect(result.drift).toBeLessThanOrEqual(1) +}) diff --git a/apps/desktop/electron/main.ts b/apps/desktop/electron/main.ts index 834a67ece3..200ecb5fed 100644 --- a/apps/desktop/electron/main.ts +++ b/apps/desktop/electron/main.ts @@ -5410,6 +5410,12 @@ function buildApplicationMenu() { { role: 'cut' }, { role: 'copy' }, { role: 'paste' }, + // ⌘⇧V is only wired up by this item existing: an accelerator with no menu + // entry is never translated into an editor command, so the chord was a + // no-op in every input in the app. The composer inserts plain text on + // every paste anyway, so this is the same result as ⌘V there — it's the + // terminal, preview, and other editable surfaces that need the strip. + { role: 'pasteAndMatchStyle' }, { role: 'delete' }, { role: 'selectAll' } ] diff --git a/apps/desktop/scripts/perf/README.md b/apps/desktop/scripts/perf/README.md index 986595f8b2..58c1ca5bc1 100644 --- a/apps/desktop/scripts/perf/README.md +++ b/apps/desktop/scripts/perf/README.md @@ -53,6 +53,7 @@ directly via `window.__PERF_DRIVE__`, so no LLM credits are spent. | `transcript` | ci | large-transcript mount + paint cost | (new) | | `render-churn` | ci | per-component render attribution + store churn while N tabs stream | (new) | | `idle-cost` | report | busy-but-silent tiles: idle commit rate, + fps while resizing / typing | (new) | +| `right-pane` | report | file tree + persistent xterm tabs under chat/terminal output and split dragging | (new) | | `cold-start` | cold | launch → CDP → driver → first paint (fresh spawn/run) | (new) | | `first-token` | backend | Enter → first assistant token painted (TTFT) | (new) | | `submit` | backend | Enter → cleared → user msg painted, scroll jump | measure-submit, measure-jump | diff --git a/apps/desktop/scripts/perf/scenarios/index.mjs b/apps/desktop/scripts/perf/scenarios/index.mjs index d03ad457fd..01ca3be384 100644 --- a/apps/desktop/scripts/perf/scenarios/index.mjs +++ b/apps/desktop/scripts/perf/scenarios/index.mjs @@ -8,6 +8,7 @@ import keystroke from './keystroke.mjs' import multitab from './multitab.mjs' import profileSwitch from './profile-switch.mjs' import renderChurn from './render-churn.mjs' +import rightPane from './right-pane.mjs' import sessionLoad from './session-load.mjs' import sessionSwitch from './session-switch.mjs' import stream from './stream.mjs' @@ -22,6 +23,7 @@ export const SCENARIOS = { [transcript.name]: transcript, [multitab.name]: multitab, [renderChurn.name]: renderChurn, + [rightPane.name]: rightPane, [idleCost.name]: idleCost, [coldStart.name]: coldStart, [firstToken.name]: firstToken, diff --git a/apps/desktop/scripts/perf/scenarios/right-pane.mjs b/apps/desktop/scripts/perf/scenarios/right-pane.mjs new file mode 100644 index 0000000000..411b283c9a --- /dev/null +++ b/apps/desktop/scripts/perf/scenarios/right-pane.mjs @@ -0,0 +1,315 @@ +// File-tree + terminal workspace stress. This is the regression scene for the +// desktop symptom where opening the project tree/terminal made the whole page +// hitch while chat and PTY output continued. +// +// It mounts a real project tree, one PTY plus multiple persistent xterm tabs, +// streams chat and terminal output together, mutates Git decoration state, and +// drags the terminal split. The debug probe records the specific work we care +// about rather than inferring it from CPU alone: +// - fixed-overlay measurements +// - active/hidden xterm fits +// - ProjectTree + per-path row renders +// - frame pacing / slow frames +// +// npm run perf -- right-pane --spawn --prod --runs 3 + +import { dirname, resolve } from 'node:path' +import { fileURLToPath } from 'node:url' + +import { sleep } from '../lib/cdp.mjs' +import { frameHistogram, percentile } from '../lib/stats.mjs' + +const DEFAULT_CWD = resolve(dirname(fileURLToPath(import.meta.url)), '../../..') + +const RECORDERS = ` + (() => { + window.__RP_FRAME_GEN__ = (window.__RP_FRAME_GEN__ || 0) + 1 + const generation = window.__RP_FRAME_GEN__ + window.__RP_FRAMES__ = { times: [], stop: false } + let last = performance.now() + const tick = () => { + if (window.__RP_FRAME_GEN__ !== generation || window.__RP_FRAMES__.stop) return + const now = performance.now() + window.__RP_FRAMES__.times.push(now - last) + last = now + requestAnimationFrame(tick) + } + requestAnimationFrame(tick) + + window.__RP_LONG__ = { entries: [], stop: false } + try { + const observer = new PerformanceObserver(list => { + if (window.__RP_LONG__.stop) return + for (const entry of list.getEntries()) { + window.__RP_LONG__.entries.push({ duration: entry.duration, startTime: entry.startTime }) + } + }) + observer.observe({ entryTypes: ['longtask'] }) + window.__RP_LONG__.observer = observer + } catch {} + return 'armed' + })() +` + +const COLLECT_RECORDERS = ` + (() => { + window.__RP_FRAMES__.stop = true + window.__RP_LONG__.stop = true + try { window.__RP_LONG__.observer && window.__RP_LONG__.observer.disconnect() } catch {} + return JSON.stringify({ frames: window.__RP_FRAMES__.times, longtasks: window.__RP_LONG__.entries }) + })() +` + +const START_COUNTERS = `window.__RIGHT_PANE_PERF__.start(); 'recording'` +const SNAPSHOT_COUNTERS = ` + (() => { + window.__RIGHT_PANE_PERF__.stop() + return JSON.stringify(window.__RIGHT_PANE_PERF__.snapshot()) + })() +` + +const DRAG_TERMINAL_SPLIT = ` + (async () => { + const slot = document.querySelector('[data-terminal-slot]') + const overlay = document.querySelector('[data-persistent-terminal]') + if (!slot || !overlay) return JSON.stringify({ target: 'none', drift: -1, moved: 0 }) + + const slotBox = slot.getBoundingClientRect() + const candidates = [...document.querySelectorAll('[role="separator"]')] + .map(element => ({ element, box: element.getBoundingClientRect() })) + .filter(item => item.box.width > item.box.height * 3) + .sort((a, b) => + Math.abs((a.box.top + a.box.bottom) / 2 - slotBox.top) - + Math.abs((b.box.top + b.box.bottom) / 2 - slotBox.top) + ) + const target = candidates[0] + if (!target) return JSON.stringify({ target: 'none', drift: -1, moved: 0 }) + + const x = target.box.left + target.box.width / 2 + const y0 = target.box.top + target.box.height / 2 + let y = y0 + const pointer = { + bubbles: true, cancelable: true, pointerId: 91, pointerType: 'mouse', + isPrimary: true, button: 0, buttons: 1 + } + target.element.dispatchEvent(new PointerEvent('pointerdown', { ...pointer, clientX: x, clientY: y })) + + for (let i = 0; i < 24; i += 1) { + y -= 1 + window.dispatchEvent(new PointerEvent('pointermove', { ...pointer, clientX: x, clientY: y })) + await new Promise(resolve => requestAnimationFrame(resolve)) + } + for (let i = 0; i < 24; i += 1) { + y += 1 + window.dispatchEvent(new PointerEvent('pointermove', { ...pointer, clientX: x, clientY: y })) + await new Promise(resolve => requestAnimationFrame(resolve)) + } + window.dispatchEvent(new PointerEvent('pointerup', { ...pointer, buttons: 0, clientX: x, clientY: y })) + // Track-size transitions continue briefly after pointerup. Wait through + // that animation, then give the overlay its normal two-frame calibration. + await new Promise(resolve => setTimeout(resolve, 350)) + await new Promise(resolve => requestAnimationFrame(() => requestAnimationFrame(resolve))) + + const a = slot.getBoundingClientRect() + const b = overlay.getBoundingClientRect() + const drift = Math.max( + Math.abs(a.top - b.top), + Math.abs(a.left - b.left), + Math.abs(a.width - b.width), + Math.abs(a.height - b.height) + ) + return JSON.stringify({ target: 'horizontal-separator', drift, moved: 24 }) + })() +` + +async function waitFor(cdp, expression, label, timeoutMs = 20000) { + const deadline = Date.now() + timeoutMs + + while (Date.now() < deadline) { + if (await cdp.eval(expression)) { + return + } + + await sleep(100) + } + + throw new Error(`right-pane timed out waiting for ${label}`) +} + +const trimWarmup = (frames, warmupMs = 300) => { + const kept = [] + let elapsed = 0 + + for (const frame of frames) { + elapsed += frame + + if (elapsed >= warmupMs) { + kept.push(frame) + } + } + + return kept +} + +export default { + name: 'right-pane', + tier: 'report', + description: 'Project tree + persistent terminal tabs under chat/terminal output and split dragging.', + async run(cdp, opts = {}) { + const cwd = resolve(String(opts.cwd ?? DEFAULT_CWD)) + const terminalCount = Math.max(2, Number(opts.terminals ?? 3)) + const tokens = Number(opts.tokens ?? 90) + const outputChunks = Number(opts.outputChunks ?? 160) + + await cdp.send('Runtime.enable') + + const ready = await cdp.eval( + `!!(window.__PERF_DRIVE__?.rightPaneSetup && window.__RIGHT_PANE_PERF__ && window.__HERMES_LAYOUT_TREE__)` + ) + + if (!ready) { + throw new Error('right-pane needs a dev renderer or a production build with VITE_PERF_PROBE=1.') + } + + let setup + + try { + setup = await cdp.eval( + `window.__PERF_DRIVE__.rightPaneSetup(${JSON.stringify({ cwd, terminals: terminalCount })})` + ) + await cdp.eval(`window.__HERMES_LAYOUT_TREE__.reveal('files'); window.__HERMES_LAYOUT_TREE__.reveal('terminal')`) + await waitFor(cdp, `!!document.querySelector('[data-project-tree]')`, 'project tree') + await waitFor( + cdp, + `!!document.querySelector('[data-terminal-slot]') && !!document.querySelector('[data-persistent-terminal]')`, + 'persistent terminal' + ) + await waitFor( + cdp, + `document.querySelectorAll('[data-terminal] .xterm').length >= ${terminalCount}`, + `${terminalCount} mounted xterms`, + 30000 + ) + await sleep(1200) + + // Activate every keep-alive tab once, then return to the output tab. + // Each activation should restore exactly one fit; inactive tabs must stay + // at zero even while another tab resizes or writes output. + await cdp.eval(START_COUNTERS) + + for (const id of setup.terminalIds) { + await cdp.eval(`window.__PERF_DRIVE__.rightPaneSelect(${JSON.stringify(id)})`) + await sleep(180) + } + + const activation = JSON.parse(await cdp.eval(SNAPSHOT_COUNTERS)) + + await cdp.eval(RECORDERS) + + // Chat DOM churn is deliberately measured in its own counter window: + // terminal positioning should receive no wakeups from transcript changes. + await cdp.eval(START_COUNTERS) + await cdp.eval( + `window.__PERF_DRIVE__.stream({ + chunk: 'Right pane streaming sentence with **bold** and \`code\`.\\n\\n', + intervalMs: 16, + totalTokens: ${tokens}, + flushMinMs: 33 + })` + ) + await cdp.eval(` + (() => { + let n = 0 + window.__RP_OUTPUT_TIMER__ = setInterval(() => { + window.__PERF_DRIVE__.rightPaneWrite( + ${JSON.stringify(setup.procId)}, + 'terminal output line ' + n + ' ........................................\\r\\n' + ) + n += 1 + if (n >= ${outputChunks}) clearInterval(window.__RP_OUTPUT_TIMER__) + }, 16) + return 'writing' + })() + `) + await sleep(Math.max(tokens, outputChunks) * 16 + 900) + const stream = JSON.parse(await cdp.eval(SNAPSHOT_COUNTERS)) + + // An unrelated Git status publication should render neither the tree root + // nor any visible row. A status for one visible path should touch only it. + await cdp.eval(START_COUNTERS) + await cdp.eval(`window.__PERF_DRIVE__.rightPaneGit('__right_pane_unrelated__.txt', 'modified')`) + await sleep(250) + const unrelatedGit = JSON.parse(await cdp.eval(SNAPSHOT_COUNTERS)) + + const visiblePath = await cdp.eval( + `document.querySelector('[data-project-tree] [title]')?.getAttribute('title') || ''` + ) + let affectedGit = { counts: { 'project-tree-render': 0, 'project-tree-row-render': 0 }, rows: {} } + + if (visiblePath) { + const relative = String(visiblePath).startsWith(`${cwd}/`) + ? String(visiblePath).slice(cwd.length + 1) + : String(visiblePath) + await cdp.eval(START_COUNTERS) + await cdp.eval(`window.__PERF_DRIVE__.rightPaneGit(${JSON.stringify(relative)}, 'modified')`) + await sleep(250) + affectedGit = JSON.parse(await cdp.eval(SNAPSHOT_COUNTERS)) + } + + await cdp.eval(START_COUNTERS) + const drag = JSON.parse(await cdp.eval(DRAG_TERMINAL_SPLIT)) + const dragCounters = JSON.parse(await cdp.eval(SNAPSHOT_COUNTERS)) + const recorded = JSON.parse(await cdp.eval(COLLECT_RECORDERS)) + const frames = trimWarmup(recorded.frames) + const longtasks = recorded.longtasks.map(entry => entry.duration) + const streamCounts = stream.counts + const activationCounts = activation.counts + const unrelatedCounts = unrelatedGit.counts + const affectedRows = Object.values(affectedGit.rows).reduce((sum, count) => sum + count, 0) + const affectedPaths = Object.keys(affectedGit.rows).length + + if (drag.target === 'none') { + throw new Error('right-pane found no horizontal terminal split separator.') + } + + return { + metrics: { + chat_terminal_measures: streamCounts['terminal-measure'], + hidden_terminal_fits: activationCounts['terminal-fit-hidden'] + streamCounts['terminal-fit-hidden'], + activation_fit_mismatch: Math.abs(activationCounts['terminal-fit-active'] - setup.terminalIds.length), + unrelated_tree_renders: unrelatedCounts['project-tree-render'], + unrelated_row_renders: unrelatedCounts['project-tree-row-render'], + affected_tree_renders: affectedGit.counts['project-tree-render'], + affected_row_path_excess: Math.max(0, affectedPaths - 1), + terminal_drift_px: Math.round(drag.drift * 10) / 10, + frame_p95_ms: Math.round(percentile(frames, 0.95) * 10) / 10, + frame_p99_ms: Math.round(percentile(frames, 0.99) * 10) / 10, + slow_frames_33: frames.filter(frame => frame > 33).length, + longtask_max_ms: Math.round((longtasks.length ? Math.max(...longtasks) : 0) * 10) / 10 + }, + detail: { + cwd, + terminals: setup.terminalIds.length, + activation, + stream, + unrelatedGit, + affectedGit, + affectedRows, + drag, + dragCounters, + frameHistogram: frameHistogram(frames), + frames: frames.length + } + } + } finally { + await cdp.eval(` + (() => { + clearInterval(window.__RP_OUTPUT_TIMER__) + window.__RIGHT_PANE_PERF__?.stop() + window.__PERF_DRIVE__?.reset() + return 'cleaned' + })() + `) + } + } +} diff --git a/apps/desktop/src/app/agents/index.tsx b/apps/desktop/src/app/agents/index.tsx index fe392e8461..eb74a2d2a7 100644 --- a/apps/desktop/src/app/agents/index.tsx +++ b/apps/desktop/src/app/agents/index.tsx @@ -3,6 +3,7 @@ import { type ReactNode, useEffect, useMemo, useState } from 'react' import { useElapsedSeconds } from '@/components/chat/activity-timer' import { ActivityTimerText } from '@/components/chat/activity-timer-text' +import { usePaneVisible } from '@/components/pane-shell/pane-visibility' import { Codicon } from '@/components/ui/codicon' import { FadeText } from '@/components/ui/fade-text' import { GlyphSpinner } from '@/components/ui/glyph-spinner' @@ -189,15 +190,17 @@ function SubagentTree({ tree }: { tree: SubagentNode[] }) { const tokens = flat.reduce((sum, n) => sum + (n.inputTokens ?? 0) + (n.outputTokens ?? 0), 0) const cost = flat.reduce((sum, n) => sum + (n.costUsd ?? 0), 0) + const visible = usePaneVisible() + useEffect(() => { - if (active <= 0 || typeof window === 'undefined') { + if (active <= 0 || !visible || typeof window === 'undefined') { return } const id = window.setInterval(() => setNowMs(Date.now()), 500) return () => window.clearInterval(id) - }, [active]) + }, [active, visible]) if (tree.length === 0) { return ( diff --git a/apps/desktop/src/app/chat/composer/index.tsx b/apps/desktop/src/app/chat/composer/index.tsx index 3767a4ad4e..28dbcf23a6 100644 --- a/apps/desktop/src/app/chat/composer/index.tsx +++ b/apps/desktop/src/app/chat/composer/index.tsx @@ -446,18 +446,12 @@ export function ChatBar({ const handlePaste = (event: ClipboardEvent) => { const imageBlobs = extractClipboardImageBlobs(event.clipboardData) - if (imageBlobs.length > 0) { - event.preventDefault() + if (imageBlobs.length > 0 && onAttachImageBlob) { + triggerHaptic('selection') - if (onAttachImageBlob) { - triggerHaptic('selection') - - for (const blob of imageBlobs) { - void onAttachImageBlob(blob) - } + for (const blob of imageBlobs) { + void onAttachImageBlob(blob) } - - return } // Trim surrounding whitespace so a copy that dragged along leading/trailing @@ -469,6 +463,10 @@ export function ChatBar({ if (!pastedText) { event.preventDefault() + if (imageBlobs.length > 0) { + return + } + // Under WSL2/WSLg the Windows host clipboard doesn't bridge *images* to // the Linux clipboard the DOM paste event reads, so a host screenshot // arrives as an empty paste (no blobs, no text). Fall back to the main diff --git a/apps/desktop/src/app/chat/composer/text-utils.test.ts b/apps/desktop/src/app/chat/composer/text-utils.test.ts index a2e54b2c07..ffca3ba642 100644 --- a/apps/desktop/src/app/chat/composer/text-utils.test.ts +++ b/apps/desktop/src/app/chat/composer/text-utils.test.ts @@ -212,6 +212,47 @@ describe('extractClipboardImageBlobs', () => { expect(extractClipboardImageBlobs(clipboard)).toEqual([image]) }) + + // A rich-text copy (Discord thread, web page, doc) carries prose plus whatever + // inline images the page decorated it with. That is a TEXT paste: attaching the + // page's placeholder graphics as composer images while the text vanished is the + // "blank attachments, no message" bug. + it('ignores inline HTML images when the copy carries its own text', () => { + const clipboard = { + files: { length: 0, item: () => null }, + getData: (type: string) => + type === 'text/html' + ? `

hello from the thread

` + : 'hello from the thread', + items: [] + } as unknown as DataTransfer + + expect(extractClipboardImageBlobs(clipboard)).toEqual([]) + }) + + it('keeps inline HTML images when the copy is image-only', () => { + const clipboard = { + files: { length: 0, item: () => null }, + getData: (type: string) => + type === 'text/html' ? `` : '', + items: [] + } as unknown as DataTransfer + + const blobs = extractClipboardImageBlobs(clipboard) + + expect(blobs).toHaveLength(1) + expect(blobs[0]?.type).toBe('image/png') + }) + + it('drops sub-thumbnail inline images — spacers, trackers, blurhash placeholders', () => { + const clipboard = { + files: { length: 0, item: () => null }, + getData: (type: string) => (type === 'text/html' ? `` : ''), + items: [] + } as unknown as DataTransfer + + expect(extractClipboardImageBlobs(clipboard)).toEqual([]) + }) }) describe('blobDedupeKey', () => { diff --git a/apps/desktop/src/app/chat/composer/text-utils.ts b/apps/desktop/src/app/chat/composer/text-utils.ts index e8828fe890..42af807c39 100644 --- a/apps/desktop/src/app/chat/composer/text-utils.ts +++ b/apps/desktop/src/app/chat/composer/text-utils.ts @@ -70,6 +70,11 @@ const SLASH_INLINE_TRIGGER_RE = /[\s\uFFFC](\/)([a-zA-Z][\w-]*)?$/ // `:` or `:D` smiley doesn't open a popover the user didn't ask for. const EMOJI_TRIGGER_RE = /(?:^|[\s\uFFFC])(:)([a-zA-Z0-9_+-]{2,})$/ +const INLINE_IMAGE_SRC_RE = /]*?\bsrc\s*=\s*["'](data:image\/[^"']+)["']/gi +// Below this, an inline data URL is chrome rather than content — a spacer, a +// 1×1 tracker, or a blurhash placeholder. Real pasted artwork clears it easily. +const MIN_INLINE_IMAGE_BYTES = 4096 + /** Stable key for paste dedupe — `items` and `files` often mirror the same image as different objects. */ export function blobDedupeKey(blob: Blob): string { if (blob instanceof File) { @@ -125,16 +130,22 @@ export function extractClipboardImageBlobs(clipboard: DataTransfer): Blob[] { if (DATA_IMAGE_URL_RE.test(text)) { push(dataUrlToBlob(text)) + + return blobs } - if (blobs.length === 0) { - const html = clipboard.getData('text/html') + // Inline `` in the clipboard's HTML — but only for a copy + // that carried no text of its own. A rich-text copy WITH prose is a text + // paste that happens to contain images, and its data URLs are the page's + // decorations rather than content: Discord ships a 32×5 blurhash placeholder + // beside every image embed, so copying a thread attached a blank thumbnail + // and (because an image paste swallows the event) dropped the text entirely. + if (!text) { + for (const match of clipboard.getData('text/html').matchAll(INLINE_IMAGE_SRC_RE)) { + const blob = dataUrlToBlob(match[1]) - if (html) { - const matches = html.matchAll(/]*?\bsrc\s*=\s*["'](data:image\/[^"']+)["']/gi) - - for (const match of matches) { - push(dataUrlToBlob(match[1])) + if (blob && blob.size >= MIN_INLINE_IMAGE_BYTES) { + push(blob) } } } diff --git a/apps/desktop/src/app/chat/index.test.tsx b/apps/desktop/src/app/chat/index.test.tsx new file mode 100644 index 0000000000..8caace4cc9 --- /dev/null +++ b/apps/desktop/src/app/chat/index.test.tsx @@ -0,0 +1,164 @@ +import { QueryClient, QueryClientProvider } from '@tanstack/react-query' +import { cleanup, fireEvent, render, screen } from '@testing-library/react' +import { useState } from 'react' +import { MemoryRouter } from 'react-router' +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' + +import { assistantTextPart, type ChatMessage } from '@/lib/chat-messages' +import { + $activeSessionId, + $awaitingResponse, + $busy, + $contextSuggestions, + $currentCwd, + $currentModel, + $currentProvider, + $freshDraftReady, + $gatewayState, + $messages, + $selectedStoredSessionId, + $sessions +} from '@/store/session' + +const threadRenderCount = vi.hoisted(() => ({ current: 0 })) + +vi.mock('@/components/assistant-ui/thread', async () => { + const React = await import('react') + + return { + Thread: () => { + threadRenderCount.current += 1 + + return React.createElement('div', { 'data-testid': 'thread' }) + } + } +}) + +vi.mock('@/components/Backdrop', async () => { + const React = await import('react') + + return { Backdrop: () => React.createElement('div', { 'data-testid': 'backdrop' }) } +}) + +vi.mock('@/components/prompt-overlays', () => ({ PromptOverlays: () => null })) +vi.mock('@/components/chat/vibe-hearts', () => ({ COMPOSER_HEART_CONFIG: {}, HeartField: () => null })) +vi.mock('@/lib/model-options', () => ({ + modelOptionsQueryKey: (...parts: unknown[]) => ['model-options', ...parts], + requestModelOptions: vi.fn(async () => ({ models: [] })) +})) +vi.mock('./chat-drop-overlay', () => ({ ChatDropOverlay: () => null })) +vi.mock('./chat-swap-overlay', () => ({ ChatSwapOverlay: () => null })) +vi.mock('./composer', () => ({ ChatBar: () => null, ChatBarFallback: () => null })) +vi.mock('./hooks/use-file-drop-zone', () => ({ + useFileDropZone: () => ({ dragKind: null, dropHandlers: {} }) +})) +vi.mock('./sidebar/session-actions-menu', async () => { + const React = await import('react') + + return { + SessionActionsMenu: ({ children }: { children: React.ReactNode }) => + React.createElement('div', { 'data-testid': 'session-actions-menu' }, children) + } +}) + +const { ChatView } = await import('./index') + +function assistantMessage(id: string, text: string): ChatMessage { + return { + id, + parts: [assistantTextPart(text)], + role: 'assistant' + } +} + +describe('ChatView render isolation', () => { + beforeEach(() => { + threadRenderCount.current = 0 + $activeSessionId.set('runtime-1') + $awaitingResponse.set(false) + $busy.set(false) + $contextSuggestions.set([]) + $currentCwd.set('/work') + $currentModel.set('test-model') + $currentProvider.set('test-provider') + $freshDraftReady.set(false) + $gatewayState.set('closed') + $messages.set([assistantMessage('assistant-1', 'Stable historical answer')]) + $selectedStoredSessionId.set('stored-1') + $sessions.set([{ id: 'stored-1', message_count: 1, title: 'Stable chat' } as never]) + }) + + afterEach(() => { + cleanup() + vi.restoreAllMocks() + $activeSessionId.set(null) + $awaitingResponse.set(false) + $busy.set(false) + $contextSuggestions.set([]) + $currentCwd.set('') + $currentModel.set('') + $currentProvider.set('') + $freshDraftReady.set(false) + $gatewayState.set('idle') + $messages.set([]) + $selectedStoredSessionId.set(null) + $sessions.set([]) + }) + + it('does not re-render chat history when an unrelated parent idle tick updates', () => { + const props = { + gateway: null, + maxVoiceRecordingSeconds: 120, + onAddContextRef: vi.fn(), + onAddUrl: vi.fn(), + onAttachDroppedItems: vi.fn(), + onAttachImageBlob: vi.fn(), + onBranchInNewChat: vi.fn(), + onCancel: vi.fn(), + onDeleteSelectedSession: vi.fn(), + onEdit: vi.fn(), + onPasteClipboardImage: vi.fn(), + onPickFiles: vi.fn(), + onPickFolders: vi.fn(), + onPickImages: vi.fn(), + onReload: vi.fn(), + onRemoveAttachment: vi.fn(), + onRetryResume: vi.fn(), + onSteer: vi.fn(), + onSubmit: vi.fn(), + onThreadMessagesChange: vi.fn(), + onToggleSelectedPin: vi.fn(), + onTranscribeAudio: vi.fn() + } + + const queryClient = new QueryClient({ + defaultOptions: { queries: { retry: false } } + }) + + function ParentTickHarness() { + const [tick, setTick] = useState(0) + + return ( + + + + + + + ) + } + + render() + + expect(screen.getByTestId('thread')).toBeTruthy() + expect(threadRenderCount.current).toBe(1) + + fireEvent.click(screen.getByRole('button', { name: /parent tick/i })) + + // memo(ChatView) with stable props must absorb the parent's idle tick — + // the transcript (Thread) must not re-render. This is PR #38470's contract. + expect(threadRenderCount.current).toBe(1) + }) +}) diff --git a/apps/desktop/src/app/chat/index.tsx b/apps/desktop/src/app/chat/index.tsx index 3d72bf3410..7a0c0a2674 100644 --- a/apps/desktop/src/app/chat/index.tsx +++ b/apps/desktop/src/app/chat/index.tsx @@ -3,7 +3,7 @@ import { useStore } from '@nanostores/react' import { useQuery } from '@tanstack/react-query' import type { ReadableAtom } from 'nanostores' import type * as React from 'react' -import { Suspense, useCallback, useEffect, useMemo, useState } from 'react' +import { memo, Suspense, useCallback, useEffect, useMemo, useState } from 'react' import { useLocation } from 'react-router' import type { SubmitTextOptions } from '@/app/session/hooks/use-prompt-actions/utils' @@ -240,7 +240,10 @@ function ChatRuntimeBoundary({ return {children} } -export function ChatView({ +// Memoized: the tile caller (session-tile.tsx) and the contrib surface re-render +// on idle ticks unrelated to the chat; with stable callback props (hoisted to +// useCallback at the call sites) memo() lets the whole chat shell skip those. +export const ChatView = memo(function ChatView({ className, gateway, modelMenuContent, @@ -596,4 +599,4 @@ export function ChatView({ ) -} +}) diff --git a/apps/desktop/src/app/chat/perf-probe.tsx b/apps/desktop/src/app/chat/perf-probe.tsx index 987d89fb42..56b540f8a4 100644 --- a/apps/desktop/src/app/chat/perf-probe.tsx +++ b/apps/desktop/src/app/chat/perf-probe.tsx @@ -1,7 +1,18 @@ import { Profiler, type ProfilerOnRenderCallback, type ReactNode } from 'react' +import { $terminalTakeover, setTerminalTakeover } from '@/app/right-sidebar/store' +import { writeAgentTerminalChunk } from '@/app/right-sidebar/terminal/agent-terminal-stream' +import { + $activeTerminalId, + $terminals, + createTerminal, + ensureAgentTerminal, + selectTerminal, + type TerminalEntry +} from '@/app/right-sidebar/terminal/terminals' +import { $repoStatusByCwd } from '@/store/coding-status' import { $gateway } from '@/store/gateway' -import { $messages, setBusy, setMessages } from '@/store/session' +import { $currentCwd, $messages, setBusy, setCurrentCwdTransient, setMessages } from '@/store/session' type Sample = { id: string @@ -38,6 +49,12 @@ declare global { * backend) doesn't contaminate frame-pacing numbers. */ connected: () => boolean + /** Mount files + multiple xterms for the synthetic right-pane scenario. */ + rightPaneSetup: (opts: { cwd: string; terminals?: number }) => { procId: string; terminalIds: string[] } + rightPaneGit: (path: string, kind?: 'added' | 'conflicted' | 'modified') => void + rightPaneReset: () => void + rightPaneSelect: (id: string) => void + rightPaneWrite: (procId: string, chunk: string) => void reset: () => void snapshotMsgs: () => number } @@ -102,11 +119,32 @@ if (typeof window !== 'undefined' && !window.__PERF_DRIVE__) { let baseline: ReturnType | null = null let activeHandle: SyntheticDriverHandle | null = null + let rightPaneBaseline: null | { + activeTerminalId: null | string + cwd: string + repoStatusByCwd: ReturnType + takeover: boolean + terminals: readonly TerminalEntry[] + } = null + const stop = () => { activeHandle = null setBusy(false) } + const resetRightPane = () => { + if (!rightPaneBaseline) { + return + } + + setTerminalTakeover(rightPaneBaseline.takeover) + $terminals.set(rightPaneBaseline.terminals) + $activeTerminalId.set(rightPaneBaseline.activeTerminalId) + $repoStatusByCwd.set(rightPaneBaseline.repoStatusByCwd) + setCurrentCwdTransient(rightPaneBaseline.cwd) + rightPaneBaseline = null + } + // One synthetic turn's worth of mixed markdown — prose, a list, a fenced // code block, inline code, a link, and a short table — so a loaded transcript // exercises the same render cost (Streamdown blocks, code cards) a real one @@ -166,6 +204,69 @@ if (typeof window !== 'undefined' && !window.__PERF_DRIVE__) { return false } }, + rightPaneGit: (path, kind = 'modified') => { + const file = { + conflicted: kind === 'conflicted', + path, + staged: false, + unstaged: kind === 'modified', + untracked: kind === 'added' + } + + const cwd = $currentCwd.get().trim() + $repoStatusByCwd.set({ + ...$repoStatusByCwd.get(), + [cwd]: { + added: 0, + ahead: 0, + behind: 0, + branch: 'perf', + changed: 1, + conflicted: kind === 'conflicted' ? 1 : 0, + defaultBranch: 'main', + detached: false, + files: [file], + removed: 0, + staged: 0, + unstaged: kind === 'modified' ? 1 : 0, + untracked: kind === 'added' ? 1 : 0 + } + }) + }, + rightPaneReset: resetRightPane, + rightPaneSelect: selectTerminal, + rightPaneSetup: ({ cwd, terminals = 3 }) => { + resetRightPane() + rightPaneBaseline = { + activeTerminalId: $activeTerminalId.get(), + cwd: $currentCwd.get(), + repoStatusByCwd: $repoStatusByCwd.get(), + takeover: $terminalTakeover.get(), + terminals: $terminals.get() + } + + setCurrentCwdTransient(cwd) + const terminalIds = [createTerminal(cwd)] + let procId = '' + + for (let index = 1; index < Math.max(1, terminals); index += 1) { + procId = `right-pane-perf-${Date.now()}-${index}` + const id = ensureAgentTerminal(procId, `perf output ${index}`) + + if (id) { + terminalIds.push(id) + } + } + + if (procId) { + selectTerminal(terminalIds.at(-1) ?? terminalIds[0]) + } + + setTerminalTakeover(true) + + return { procId, terminalIds } + }, + rightPaneWrite: (procId, chunk) => writeAgentTerminalChunk(procId, chunk), loadTranscript: (turns = 200) => { if (!baseline) { baseline = $messages.get() @@ -190,6 +291,7 @@ if (typeof window !== 'undefined' && !window.__PERF_DRIVE__) { }, reset: () => { activeHandle?.stop() + resetRightPane() if (baseline) { setMessages(baseline) diff --git a/apps/desktop/src/app/chat/session-status-dot.tsx b/apps/desktop/src/app/chat/session-status-dot.tsx index 09fb7773a8..c84a1cae0d 100644 --- a/apps/desktop/src/app/chat/session-status-dot.tsx +++ b/apps/desktop/src/app/chat/session-status-dot.tsx @@ -1,5 +1,6 @@ import { useStore } from '@nanostores/react' +import { StatusPulse } from '@/components/ui/status-pulse' import { type Translations, useI18n } from '@/i18n' import { useStoreSelector } from '@/lib/use-session-slice' import { cn } from '@/lib/utils' @@ -17,6 +18,10 @@ import { type SessionDotState, sessionDotState } from './sidebar/session-row-sta type DotVariant = { ariaLabel?: (r: Translations['sidebar']['row']) => string className: string + pulse?: { + className: string + opacity: number + } role?: 'status' title?: (r: Translations['sidebar']['row']) => string } @@ -24,12 +29,6 @@ type DotVariant = { // Shared base for every active dot; idle is smaller and uses its own class. const DOT_BASE = 'relative size-1.5 rounded-full' -// Pseudo-element ping ring that scales outward and fades — shared scaffold for -// the two pulsing dots. The `before:bg-*` color is written inline per variant -// (NOT interpolated here): Tailwind only generates utilities it can see as -// complete static strings, so a `before:bg-${color}` template never emits. -const PING = "before:absolute before:inset-0 before:animate-ping before:rounded-full before:content-['']" - const DOT_VARIANTS: Record = { // Amber steady — a clarify/approval is blocking the turn. Steady (not // pulsing) reads as "your turn", distinct from the accent pulse of a turn. @@ -42,14 +41,22 @@ const DOT_VARIANTS: Record = { // Accent pulse — the LLM turn is actively running. working: { ariaLabel: r => r.sessionRunning, - className: `${DOT_BASE} bg-(--ui-accent) shadow-[0_0_0.625rem_color-mix(in_srgb,var(--ui-accent)_55%,transparent)] ${PING} before:bg-(--ui-accent) before:opacity-70`, + className: `${DOT_BASE} bg-(--ui-accent) shadow-[0_0_0.625rem_color-mix(in_srgb,var(--ui-accent)_55%,transparent)]`, + pulse: { + className: 'absolute inset-0 rounded-full bg-(--ui-accent) opacity-0', + opacity: 0.7 + }, role: 'status' }, // Quiet accent pulse — the turn is still authoritative-running, but no // stream activity has arrived for the watchdog window. stalled: { ariaLabel: r => r.sessionRunning, - className: `${DOT_BASE} bg-(--ui-accent) opacity-70 ${PING} before:bg-(--ui-accent) before:opacity-40`, + className: `${DOT_BASE} bg-(--ui-accent) opacity-70`, + pulse: { + className: 'absolute inset-0 rounded-full bg-(--ui-accent) opacity-0', + opacity: 0.4 + }, role: 'status', title: r => r.sessionRunning }, @@ -58,7 +65,11 @@ const DOT_VARIANTS: Record = { // than muted-foreground so it's visible against the surface. background: { ariaLabel: r => r.backgroundRunning, - className: `${DOT_BASE} bg-muted-foreground/80 ${PING} before:bg-muted-foreground/80 before:opacity-60`, + className: `${DOT_BASE} bg-muted-foreground/80`, + pulse: { + className: 'absolute inset-0 rounded-full bg-muted-foreground/80 opacity-0', + opacity: 0.6 + }, role: 'status', title: r => r.backgroundRunning }, @@ -123,6 +134,7 @@ export function SessionStatusDot({ storedSessionId, session, branchStem, classNa const hasBackground = useStoreSelector($backgroundRunningSessionIds, ids => ids.includes(storedSessionId)) const dotState = sessionDotState({ hasBackground, isStalled, isUnread, isWorking, needsInput }) + const variant = DOT_VARIANTS[dotState] return ( @@ -135,11 +147,20 @@ export function SessionStatusDot({ storedSessionId, session, branchStem, classNa ) diff --git a/apps/desktop/src/app/chat/session-tile.tsx b/apps/desktop/src/app/chat/session-tile.tsx index 27954abc2a..9c30a632f9 100644 --- a/apps/desktop/src/app/chat/session-tile.tsx +++ b/apps/desktop/src/app/chat/session-tile.tsx @@ -104,6 +104,13 @@ function buildTileView(storedSessionId: string): SessionView { } } +// Module-level constants so these ChatView props are referentially stable — +// tiles have no pin/delete affordance, and transcription needs no per-tile state. +const noop = () => undefined + +const tileTranscribeAudio = async (audio: Blob) => + (await transcribeAudio(await blobToDataUrl(audio), audio.type)).transcript + function TileChat({ runtimeId, storedSessionId, @@ -144,6 +151,29 @@ function TileChat({ scope: { add: attachments.add, remove: attachments.remove, target: scope.target } }) + // ChatView is memo()d — every callback prop must be referentially stable or + // the memo never holds and each tile-level render (idle ticks, unrelated + // store updates) re-renders the whole chat shell. The individual composer + // functions are useCallback'd inside useComposerActions, so hoisting these + // wrappers onto them keeps identity stable across renders. + const { addContextRefAttachment, pasteClipboardImage, pickContextPaths, pickImages, removeAttachment } = composer + + const onAddUrl = useCallback( + (url: string) => addContextRefAttachment(`@url:${formatRefValue(url)}`, url), + [addContextRefAttachment] + ) + + const onPasteClipboardImage = useCallback( + (opts?: { silent?: boolean }) => pasteClipboardImage(opts), + [pasteClipboardImage] + ) + + const onPickFiles = useCallback(() => void pickContextPaths('file'), [pickContextPaths]) + const onPickFolders = useCallback(() => void pickContextPaths('folder'), [pickContextPaths]) + const onPickImages = useCallback(() => void pickImages(), [pickImages]) + const onRemoveAttachment = useCallback((id: string) => void removeAttachment(id), [removeAttachment]) + const onRetryResume = useCallback(() => patchSessionTile(storedSessionId, { error: undefined }), [storedSessionId]) + // Per-tile model menu — rendered under this tile's SessionView so the pill // + switch target THIS runtime, not the primary (which may be mid-turn). const modelMenuContent = useMemo( @@ -165,27 +195,27 @@ function TileChat({ composer.addContextRefAttachment(`@url:${formatRefValue(url)}`, url)} + onAddContextRef={addContextRefAttachment} + onAddUrl={onAddUrl} onAttachDroppedItems={composer.attachDroppedItems} onAttachImageBlob={composer.attachImageBlob} onCancel={actions.cancelRun} - onDeleteSelectedSession={() => undefined} + onDeleteSelectedSession={noop} onDismissError={actions.dismissError} onEdit={actions.editMessage} - onPasteClipboardImage={opts => composer.pasteClipboardImage(opts)} - onPickFiles={() => void composer.pickContextPaths('file')} - onPickFolders={() => void composer.pickContextPaths('folder')} - onPickImages={() => void composer.pickImages()} + onPasteClipboardImage={onPasteClipboardImage} + onPickFiles={onPickFiles} + onPickFolders={onPickFolders} + onPickImages={onPickImages} onReload={actions.reloadFromMessage} - onRemoveAttachment={id => void composer.removeAttachment(id)} + onRemoveAttachment={onRemoveAttachment} onRestoreToMessage={actions.restoreToMessage} - onRetryResume={() => patchSessionTile(storedSessionId, { error: undefined })} + onRetryResume={onRetryResume} onSteer={actions.steerPrompt} onSubmit={actions.submitText} onThreadMessagesChange={actions.handleThreadMessagesChange} - onToggleSelectedPin={() => undefined} - onTranscribeAudio={async audio => (await transcribeAudio(await blobToDataUrl(audio), audio.type)).transcript} + onToggleSelectedPin={noop} + onTranscribeAudio={tileTranscribeAudio} /> diff --git a/apps/desktop/src/app/chat/sidebar/cron-jobs-section.tsx b/apps/desktop/src/app/chat/sidebar/cron-jobs-section.tsx index dd6988ce91..fea0d8dff0 100644 --- a/apps/desktop/src/app/chat/sidebar/cron-jobs-section.tsx +++ b/apps/desktop/src/app/chat/sidebar/cron-jobs-section.tsx @@ -1,6 +1,7 @@ import { useStore } from '@nanostores/react' import { useEffect, useMemo, useState } from 'react' +import { usePaneVisible } from '@/components/pane-shell/pane-visibility' import { ActionsContextMenu, type MenuKit, renderActionItem } from '@/components/ui/actions-menu' import { Codicon } from '@/components/ui/codicon' import { DisclosureCaret } from '@/components/ui/disclosure-caret' @@ -92,17 +93,19 @@ export function SidebarCronJobsSection({ // Rows revealed so far; starts compact, grows in steps via "load more". const [visibleCount, setVisibleCount] = useState(INITIAL_VISIBLE_JOBS) + const visible = usePaneVisible() + // One clock for the whole section (rows are pure) so the countdowns tick - // without re-rendering the rest of the sidebar. Only runs while expanded. + // without re-rendering the rest of the sidebar. Only runs while expanded and visible. useEffect(() => { - if (!open) { + if (!open || !visible) { return } const id = window.setInterval(() => setNowMs(Date.now()), 1000) return () => window.clearInterval(id) - }, [open]) + }, [open, visible]) // Upcoming first (soonest next run), jobs with no next run sink to the bottom, // then alphabetical for stability. @@ -328,6 +331,7 @@ function CronJobSidebarRuns({ jobId, onOpenRun }: { jobId: string; onOpenRun: (s const changeEventsAvailable = useStore($changeEventsAvailable) const cronChangeTick = useStore($cronChangeTick) const [runs, setRuns] = useState(null) + const visible = usePaneVisible() useEffect(() => { let cancelled = false @@ -345,6 +349,15 @@ function CronJobSidebarRuns({ jobId, onOpenRun }: { jobId: string; onOpenRun: (s } }) + // Hidden pane: skip the peek entirely — no initial load, no interval. + // `visible` is in the dep array, so becoming visible re-runs this effect + // and starts the load + timer fresh (same shape as the section clock). + if (!visible) { + return () => { + cancelled = true + } + } + void load() const intervalId = window.setInterval( @@ -361,7 +374,7 @@ function CronJobSidebarRuns({ jobId, onOpenRun }: { jobId: string; onOpenRun: (s window.clearInterval(intervalId) } // cronChangeTick: a fired run reloads the peek immediately. - }, [changeEventsAvailable, cronChangeTick, jobId]) + }, [changeEventsAvailable, cronChangeTick, jobId, visible]) return (
diff --git a/apps/desktop/src/app/chat/sidebar/session-actions-menu.tsx b/apps/desktop/src/app/chat/sidebar/session-actions-menu.tsx index 73f741a905..57d8a9f962 100644 --- a/apps/desktop/src/app/chat/sidebar/session-actions-menu.tsx +++ b/apps/desktop/src/app/chat/sidebar/session-actions-menu.tsx @@ -7,6 +7,7 @@ import { closeAllTreeTabs, closeOtherTreeTabs, closeTreeTabsToRight, + reloadTreePane, treeTabCloseTargets } from '@/components/pane-shell/tree/store' import { @@ -239,12 +240,24 @@ function useSessionActions({ }) ] - // TAB — close verbs that act on the strip (tabs only; a row isn't a tab). + // TAB — verbs that act on the strip (tabs only; a row isn't a tab). const closeTargets = surface === 'tab' && tabPaneId ? treeTabCloseTargets(tabPaneId) : null - const tabCloseItems: ActionItemSpec[] = + const tabItems: ActionItemSpec[] = surface === 'tab' ? [ + ...(tabPaneId + ? [ + spec({ + icon: 'refresh', + label: t.zones.reload, + onSelect: () => { + triggerHaptic('selection') + reloadTreePane(tabPaneId) + } + }) + ] + : []), ...(onClose ? [ spec({ @@ -342,10 +355,10 @@ function useSessionActions({ /> {workItems.map(item => renderActionItem(kit, item))} - {tabCloseItems.length > 0 && ( + {tabItems.length > 0 && ( <> - {tabCloseItems.map(item => renderActionItem(kit, item))} + {tabItems.map(item => renderActionItem(kit, item))} )} diff --git a/apps/desktop/src/app/chat/sidebar/sessions-section.test.tsx b/apps/desktop/src/app/chat/sidebar/sessions-section.test.tsx new file mode 100644 index 0000000000..5fd286e846 --- /dev/null +++ b/apps/desktop/src/app/chat/sidebar/sessions-section.test.tsx @@ -0,0 +1,184 @@ +import { cleanup, render } from '@testing-library/react' +import type * as React from 'react' +import { afterEach, describe, expect, it, vi } from 'vitest' + +import type { SessionInfo } from '@/hermes' + +import { SidebarSessionsSection, VIRTUALIZE_THRESHOLD } from './sessions-section' +import type { VirtualSessionListProps } from './virtual-session-list' + +afterEach(cleanup) + +vi.mock('@/i18n', () => ({ + useI18n: () => ({ + t: { + sidebar: { + dateDivider: { + earlierThisMonth: 'Earlier this month', + lastMonth: 'Last month', + lastWeek: 'Last week', + older: 'Older', + today: 'Today', + yesterday: 'Yesterday' + } + } + } + }) +})) + +const mockVirtualListPropsHistory: VirtualSessionListProps[] = [] + +vi.mock('./virtual-session-list', () => ({ + VirtualSessionList: (props: VirtualSessionListProps) => { + mockVirtualListPropsHistory.push(props) + + return
Virtual List ({props.rows.length} rows)
+ } +})) + +vi.mock('./session-row', () => ({ + SidebarSessionRow: ({ session }: { session: SessionInfo }) => ( +
{session.id}
+ ) +})) + +function makeSession(id: string, startedAt = 1000): SessionInfo { + return { + handoff_platform: null, + handoff_state: null, + id, + last_active: startedAt, + profile: 'default', + started_at: startedAt + } as unknown as SessionInfo +} + +function generateSessions(count: number): SessionInfo[] { + return Array.from({ length: count }, (_, i) => makeSession(`session-${i + 1}`, 10000 - i * 100)) +} + +const noop = () => {} + +describe('SidebarSessionsSection memoization & virtualizer stability', () => { + it('memoizes flatRows and passes the exact same rows array reference across parent re-renders', () => { + mockVirtualListPropsHistory.length = 0 + + const sessions = generateSessions(VIRTUALIZE_THRESHOLD + 5) + + const { rerender } = render( + Empty
} + label="Sessions" + onArchiveSession={noop} + onDeleteSession={noop} + onResumeSession={noop} + onToggle={noop} + onTogglePin={noop} + open={true} + pinned={false} + sessions={sessions} + workingSessionIdSet={new Set()} + /> + ) + + expect(mockVirtualListPropsHistory.length).toBe(1) + const initialRowsRef = mockVirtualListPropsHistory[0].rows + expect(initialRowsRef.length).toBeGreaterThan(VIRTUALIZE_THRESHOLD) + + // Re-render parent with the exact same sessions array and props + rerender( + Empty} + label="Sessions" + onArchiveSession={noop} + onDeleteSession={noop} + onResumeSession={noop} + onToggle={noop} + onTogglePin={noop} + open={true} + pinned={false} + sessions={sessions} + workingSessionIdSet={new Set()} + /> + ) + + expect(mockVirtualListPropsHistory.length).toBe(2) + const nextRowsRef = mockVirtualListPropsHistory[1].rows + + // Confirm that the flatRows array reference remains strictly identical across renders (useMemo proof) + expect(nextRowsRef).toBe(initialRowsRef) + }) + + it('re-computes flatRows reference when dateGrouped or sessions change', () => { + mockVirtualListPropsHistory.length = 0 + + const initialSessions = generateSessions(VIRTUALIZE_THRESHOLD + 2) + + const { rerender } = render( + Empty} + label="Sessions" + onArchiveSession={noop} + onDeleteSession={noop} + onResumeSession={noop} + onToggle={noop} + onTogglePin={noop} + open={true} + pinned={false} + sessions={initialSessions} + workingSessionIdSet={new Set()} + /> + ) + + const firstRowsRef = mockVirtualListPropsHistory[0].rows + + // Change dateGrouped to true + rerender( + Empty} + label="Sessions" + onArchiveSession={noop} + onDeleteSession={noop} + onResumeSession={noop} + onToggle={noop} + onTogglePin={noop} + open={true} + pinned={false} + sessions={initialSessions} + workingSessionIdSet={new Set()} + /> + ) + + const secondRowsRef = mockVirtualListPropsHistory[1].rows + expect(secondRowsRef).not.toBe(firstRowsRef) + + // Change sessions array identity + const updatedSessions = generateSessions(VIRTUALIZE_THRESHOLD + 4) + rerender( + Empty} + label="Sessions" + onArchiveSession={noop} + onDeleteSession={noop} + onResumeSession={noop} + onToggle={noop} + onTogglePin={noop} + open={true} + pinned={false} + sessions={updatedSessions} + workingSessionIdSet={new Set()} + /> + ) + + const thirdRowsRef = mockVirtualListPropsHistory[2].rows + expect(thirdRowsRef).not.toBe(secondRowsRef) + }) +}) diff --git a/apps/desktop/src/app/chat/sidebar/sessions-section.tsx b/apps/desktop/src/app/chat/sidebar/sessions-section.tsx index 35174a8370..684d217856 100644 --- a/apps/desktop/src/app/chat/sidebar/sessions-section.tsx +++ b/apps/desktop/src/app/chat/sidebar/sessions-section.tsx @@ -1,6 +1,6 @@ import type { useSensors } from '@dnd-kit/core' import type * as React from 'react' -import { useMemo } from 'react' +import { useCallback, useMemo } from 'react' import { SidebarPanelLabel } from '@/app/shell/sidebar-label' import { DisclosureCaret } from '@/components/ui/disclosure-caret' @@ -225,52 +225,79 @@ export function SidebarSessionsSection({ [sessions, preserveInputOrder] ) - const renderRow = (session: SessionInfo, draggable: boolean, branchStem?: string) => { - const rowProps = { - branchStem, - isPinned: pinned, - isSelected: session.id === activeSessionId, - isWorking: workingSessionIdSet.has(session.id), - onArchive: () => onArchiveSession(session.id), - onBranch: onBranchSession ? () => onBranchSession(session.id, session.profile) : undefined, - onDelete: () => onDeleteSession(session.id), - onPin: () => onTogglePin(sessionPinId(session)), - onResume: () => onResumeSession(session.id), - reorderable: draggable && !branchStem, - session, - showProfile: showProfileTags - } + const renderRow = useCallback( + (session: SessionInfo, draggable: boolean, branchStem?: string) => { + const rowProps = { + branchStem, + isPinned: pinned, + isSelected: session.id === activeSessionId, + isWorking: workingSessionIdSet.has(session.id), + onArchive: () => onArchiveSession(session.id), + onBranch: onBranchSession ? () => onBranchSession(session.id, session.profile) : undefined, + onDelete: () => onDeleteSession(session.id), + onPin: () => onTogglePin(sessionPinId(session)), + onResume: () => onResumeSession(session.id), + reorderable: draggable && !branchStem, + session, + showProfile: showProfileTags + } - return draggable && !branchStem ? ( - - ) : ( - - ) - } + return draggable && !branchStem ? ( + + ) : ( + + ) + }, + [ + activeSessionId, + onArchiveSession, + onBranchSession, + onDeleteSession, + onResumeSession, + onTogglePin, + pinned, + showProfileTags, + workingSessionIdSet + ] + ) // A single flat/virtual/lane list row — either a date divider or a session. - const renderListRow = (row: SidebarListRow, draggable: boolean) => - row.kind === 'divider' ? ( - - ) : ( - renderRow(row.entry.session, draggable, row.entry.branchStem) - ) + const renderListRow = useCallback( + (row: SidebarListRow, draggable: boolean) => + row.kind === 'divider' ? ( + + ) : ( + renderRow(row.entry.session, draggable, row.entry.branchStem) + ), + [dividerLabels, renderRow] + ) // Sessions inside repos/worktrees are date-ordered and static. - const renderRows = (items: SessionInfo[]) => - flattenSessionsWithBranches(items).map(({ branchStem, session }) => renderRow(session, false, branchStem)) + const renderRows = useCallback( + (items: SessionInfo[]) => + flattenSessionsWithBranches(items).map(({ branchStem, session }) => renderRow(session, false, branchStem)), + [renderRow] + ) // Same as `renderRows`, but with date dividers folded in — used for // entered-project lanes so a lane spanning multiple days reads // chronologically, matching the flat recents list. - const renderRowsDated = (items: SessionInfo[]) => { - const entries = flattenSessionsWithBranches(items) + const renderRowsDated = useCallback( + (items: SessionInfo[]) => { + const entries = flattenSessionsWithBranches(items) - return (dateGrouped ? groupEntriesByRecency(entries) : toSessionRows(entries)).map(row => renderListRow(row, false)) - } + return (dateGrouped ? groupEntriesByRecency(entries) : toSessionRows(entries)).map(row => + renderListRow(row, false) + ) + }, + [dateGrouped, renderListRow] + ) // Flat recents as list rows: grouped by recency when enabled, plain otherwise. - const flatRows: SidebarListRow[] = dateGrouped ? groupEntriesByRecency(displayEntries) : toSessionRows(displayEntries) + const flatRows: SidebarListRow[] = useMemo( + () => (dateGrouped ? groupEntriesByRecency(displayEntries) : toSessionRows(displayEntries)), + [dateGrouped, displayEntries] + ) const flatVirtualized = !showEmptyState && diff --git a/apps/desktop/src/app/chat/sidebar/virtual-session-list.tsx b/apps/desktop/src/app/chat/sidebar/virtual-session-list.tsx index d9cfd6c004..90e10ddc38 100644 --- a/apps/desktop/src/app/chat/sidebar/virtual-session-list.tsx +++ b/apps/desktop/src/app/chat/sidebar/virtual-session-list.tsx @@ -27,7 +27,7 @@ interface SessionRowCommonProps { showProfile?: boolean } -interface VirtualSessionListProps { +export interface VirtualSessionListProps { activeSessionId: null | string className?: string rows: SidebarListRow[] diff --git a/apps/desktop/src/app/contrib/hooks/use-session-tile-delegate.test.ts b/apps/desktop/src/app/contrib/hooks/use-session-tile-delegate.test.ts index f522e71035..1297690584 100644 --- a/apps/desktop/src/app/contrib/hooks/use-session-tile-delegate.test.ts +++ b/apps/desktop/src/app/contrib/hooks/use-session-tile-delegate.test.ts @@ -77,7 +77,8 @@ describe('useSessionTileDelegate resumeTile', () => { expect(requestGateway).toHaveBeenCalledWith('session.resume', { session_id: 'stored-x', cols: 96, - profile: 'ai-engineer' + profile: 'ai-engineer', + omit_messages: true }) }) @@ -94,7 +95,8 @@ describe('useSessionTileDelegate resumeTile', () => { expect(requestGateway).toHaveBeenCalledWith('session.resume', { session_id: 'stored-y', cols: 96, - profile: 'default' + profile: 'default', + omit_messages: true }) }) }) diff --git a/apps/desktop/src/app/contrib/hooks/use-session-tile-delegate.ts b/apps/desktop/src/app/contrib/hooks/use-session-tile-delegate.ts index ba1107e9e0..57bfa92bfa 100644 --- a/apps/desktop/src/app/contrib/hooks/use-session-tile-delegate.ts +++ b/apps/desktop/src/app/contrib/hooks/use-session-tile-delegate.ts @@ -79,6 +79,7 @@ export function useSessionTileDelegate({ requestGateway('session.resume', { session_id: storedSessionId, cols: 96, + omit_messages: true, ...(profile ? { profile } : {}) }) ]) diff --git a/apps/desktop/src/app/gateway/hooks/use-gateway-boot.ts b/apps/desktop/src/app/gateway/hooks/use-gateway-boot.ts index f02ee69f5e..99ccc0564c 100644 --- a/apps/desktop/src/app/gateway/hooks/use-gateway-boot.ts +++ b/apps/desktop/src/app/gateway/hooks/use-gateway-boot.ts @@ -5,6 +5,7 @@ import type { HermesConnection } from '@/global' import { HermesGateway } from '@/hermes' import { translateNow } from '@/i18n' import { desktopDefaultCwd } from '@/lib/desktop-fs' +import { reconnectBackoffDelayMs } from '@/lib/reconnect-backoff' import { $desktopBoot, applyDesktopBootProgress, @@ -42,12 +43,16 @@ import type { RpcEvent } from '@/types/hermes' import { stashGatewaySurvivor, survivorIsStale, takeGatewaySurvivor } from './gateway-hmr-survivor' -// After this many consecutive failed reconnects (≈45s with the 1→15s backoff) -// raise a recoverable boot error. Otherwise a dropped remote gateway loops the -// backoff forever behind the fullscreen CONNECTING overlay with no way to reach -// Settings / sign in / switch to local — the "lost connection breaks the app" -// dead end. The next successful reconnect clears it. -const RECONNECT_ESCALATE_AFTER = 6 +// After the reconnect loop has been failing for this long, raise a recoverable +// boot error. Otherwise a dropped remote gateway loops the backoff forever +// behind the fullscreen CONNECTING overlay with no way to reach Settings / +// sign in / switch to local — the "lost connection breaks the app" dead end. +// The next successful reconnect clears it. Time-based (not attempt-count) +// because the full-jitter backoff makes attempt counts a meaningless clock: +// six jittered attempts can elapse in ~9s, while the old deterministic +// 1→15s ladder took ~45s to reach six failures — this threshold keeps that +// original ~45s calibration. +const RECONNECT_ESCALATE_AFTER_MS = 45_000 interface GatewayBootOptions { beforeConnectionSwitch: () => void @@ -114,13 +119,18 @@ export function useGatewayBoot({ let reconnecting = false let reconnectTimer: ReturnType | null = null let reconnectAttempt = 0 + // Wall-clock start of the current disconnect episode (first failed + // reconnect attempt); null while healthy. Drives the time-based + // escalation below. Reset on a clean open or a manual/wake reconnect. + let reconnectFailingSince: number | null = null // Surface "sign in again" once per disconnect episode, not on every backoff // tick — a stale OAuth ticket fails every attempt and would otherwise stack // identical error toasts (and their haptics). Reset on the next clean open. let reauthNotified = false - // Raised once the reconnect loop crosses RECONNECT_ESCALATE_AFTER so the - // recovery overlay replaces the dead-end CONNECTING screen. Reset on a clean - // open or a manual/wake-driven reconnect. + // Raised once the reconnect loop has been failing for + // RECONNECT_ESCALATE_AFTER_MS so the recovery overlay replaces the + // dead-end CONNECTING screen. Reset on a clean open or a manual/ + // wake-driven reconnect. let escalated = false // Wrap the live getter in a call so TS control-flow analysis doesn't narrow @@ -173,6 +183,7 @@ export function useGatewayBoot({ } reconnectAttempt = 0 + reconnectFailingSince = null // A respawned backend re-mints (recycles) runtime ids, so any tile's // bound runtime id is now stale — drop them so each tile re-resumes. resetTileRuntimeBindings() @@ -192,7 +203,11 @@ export function useGatewayBoot({ reconnecting = false if (!cancelled && !gatewayOpen() && !$gatewaySwitching.get()) { - if (reconnectAttempt >= RECONNECT_ESCALATE_AFTER && !escalated) { + if (reconnectFailingSince === null) { + reconnectFailingSince = Date.now() + } + + if (Date.now() - reconnectFailingSince >= RECONNECT_ESCALATE_AFTER_MS && !escalated) { escalated = true failDesktopBoot(translateNow('boot.errors.gatewayConnectionLost')) } @@ -207,8 +222,11 @@ export function useGatewayBoot({ return } - // 1s, 2s, 4s … capped at 15s. - const delay = Math.min(15_000, 1_000 * 2 ** Math.min(reconnectAttempt, 4)) + // Full-jitter exponential backoff (300ms base, 15s cap) so a gateway + // restart doesn't get redialed by every desktop client in lockstep — + // an immediate-retry reconnect storm can exhaust the gateway's file + // descriptors while it's still coming back up. + const delay = reconnectBackoffDelayMs(reconnectAttempt) reconnectAttempt += 1 reconnectTimer = setTimeout(() => { reconnectTimer = null @@ -223,6 +241,7 @@ export function useGatewayBoot({ clearReconnectTimer() reconnectAttempt = 0 + reconnectFailingSince = null escalated = false reconnectSecondaryGateways() @@ -269,6 +288,7 @@ export function useGatewayBoot({ $gatewaySwitching.set(true) clearReconnectTimer() reconnectAttempt = 0 + reconnectFailingSince = null escalated = false reauthNotified = false callbacksRef.current.beforeConnectionSwitch() @@ -371,6 +391,7 @@ export function useGatewayBoot({ if (st === 'open') { reconnectAttempt = 0 + reconnectFailingSince = null reauthNotified = false escalated = false clearReconnectTimer() diff --git a/apps/desktop/src/app/right-sidebar/files/tree-sizing.test.ts b/apps/desktop/src/app/right-sidebar/files/tree-sizing.test.ts new file mode 100644 index 0000000000..04e8a8bb47 --- /dev/null +++ b/apps/desktop/src/app/right-sidebar/files/tree-sizing.test.ts @@ -0,0 +1,29 @@ +import { describe, expect, it, vi } from 'vitest' + +import { projectTreeViewportSize } from './tree' + +describe('projectTreeViewportSize', () => { + it('uses ResizeObserver contentRect without forcing another layout read', () => { + const element = document.createElement('div') + const getBoundingClientRect = vi.spyOn(element, 'getBoundingClientRect') + const contentRect = { height: 480, width: 320 } as DOMRectReadOnly + + expect( + projectTreeViewportSize([{ contentRect, target: element } as unknown as ResizeObserverEntry], element) + ).toEqual({ + height: 480, + width: 320 + }) + expect(getBoundingClientRect).not.toHaveBeenCalled() + }) + + it('falls back to a rect when ResizeObserver is unavailable', () => { + const element = document.createElement('div') + vi.spyOn(element, 'getBoundingClientRect').mockReturnValue({ + height: 240, + width: 160 + } as DOMRect) + + expect(projectTreeViewportSize([], element)).toEqual({ height: 240, width: 160 }) + }) +}) diff --git a/apps/desktop/src/app/right-sidebar/files/tree.tsx b/apps/desktop/src/app/right-sidebar/files/tree.tsx index 97d96f267b..b20bfcaf4b 100644 --- a/apps/desktop/src/app/right-sidebar/files/tree.tsx +++ b/apps/desktop/src/app/right-sidebar/files/tree.tsx @@ -1,12 +1,14 @@ import { useStore } from '@nanostores/react' import { type KeyboardEvent as ReactKeyboardEvent, useCallback, useEffect, useRef, useState } from 'react' +import { useMemo } from 'react' import { type NodeApi, type NodeRendererProps, type RowRendererProps, Tree, type TreeApi } from 'react-arborist' import { TreeSkeleton } from '@/components/chat/skeletons' import { Codicon } from '@/components/ui/codicon' +import { markRightPanePerf } from '@/debug/right-pane-events' import { useResizeObserver } from '@/hooks/use-resize-observer' import { cn } from '@/lib/utils' -import { $repoChangeByPath, type RepoChangeKind } from '@/store/coding-status' +import { type RepoChangeKind, repoChangeKindForPath } from '@/store/coding-status' import { $renamingPath, beginInlineRename } from '@/store/file-actions' import { $revealInTreeRequest } from '@/store/layout' @@ -55,19 +57,20 @@ export function ProjectTree({ onPreviewFile, openState }: ProjectTreeProps) { + markRightPanePerf('project-tree-render') + const containerRef = useRef(null) const treeRef = useRef | null>(null) const [size, setSize] = useState({ height: 0, width: 0 }) - const changeByPath = useStore($repoChangeByPath) - const syncTreeSize = useCallback(() => { + const syncTreeSize = useCallback((entries: readonly ResizeObserverEntry[]) => { const el = containerRef.current if (!el) { return } - const { height, width } = el.getBoundingClientRect() + const { height, width } = projectTreeViewportSize(entries, el) setSize(prev => { if (prev.height === height && prev.width === width) { @@ -175,7 +178,12 @@ export function ProjectTree({ }, []) return ( -
+
{size.height > 0 && size.width > 0 ? ( childrenAccessor={node => (node?.isDirectory ? (node.children ?? []) : null)} @@ -200,7 +208,6 @@ export function ProjectTree({ {props => ( item.target === element) + const box = entry?.contentRect ?? element.getBoundingClientRect() + + return { height: box.height, width: box.width } +} + function TreeSizingState() { return } @@ -244,7 +261,6 @@ const CHANGE_TINT: Record = { } function ProjectTreeRow({ - changeKind, dragHandle, node, onAttachFile, @@ -253,13 +269,17 @@ function ProjectTreeRow({ relativeTo, style }: NodeRendererProps & { - changeKind?: RepoChangeKind onAttachFile: (path: string) => void onAttachFolder: (path: string) => void onPreviewFile?: (path: string) => void relativeTo?: null | string }) { const renamingPath = useStore($renamingPath) + const path = node.data?.id ?? '' + const changeStore = useMemo(() => repoChangeKindForPath(path), [path]) + const changeKind: RepoChangeKind | undefined = useStore(changeStore) + + markRightPanePerf('project-tree-row-render', path) if (!node.data) { return
diff --git a/apps/desktop/src/app/right-sidebar/terminal/active-resize.test.ts b/apps/desktop/src/app/right-sidebar/terminal/active-resize.test.ts new file mode 100644 index 0000000000..27909c2cb0 --- /dev/null +++ b/apps/desktop/src/app/right-sidebar/terminal/active-resize.test.ts @@ -0,0 +1,159 @@ +import { afterEach, describe, expect, it, vi } from 'vitest' + +import { observeActiveTerminalResize } from './active-resize' + +afterEach(() => { + vi.unstubAllGlobals() +}) + +function installRaf() { + let nextId = 1 + const frames = new Map() + + vi.stubGlobal('requestAnimationFrame', (callback: FrameRequestCallback) => { + const id = nextId++ + frames.set(id, callback) + + return id + }) + vi.stubGlobal('cancelAnimationFrame', (id: number) => frames.delete(id)) + + return { + flush() { + const pending = [...frames.entries()] + frames.clear() + pending.forEach(([, callback]) => callback(0)) + }, + pending: () => frames.size + } +} + +describe('observeActiveTerminalResize', () => { + it('fits once on activation and coalesces later resize bursts', () => { + const raf = installRaf() + const resize = { current: null as ResizeObserverCallback | null } + const disconnect = vi.fn() + + vi.stubGlobal( + 'ResizeObserver', + class { + constructor(callback: ResizeObserverCallback) { + resize.current = callback + } + + disconnect = disconnect + observe = vi.fn((target: Element) => { + resize.current?.([{ target } as ResizeObserverEntry], this as unknown as ResizeObserver) + }) + unobserve = vi.fn() + } as unknown as typeof ResizeObserver + ) + + const onFit = vi.fn() + const onActivate = vi.fn() + const host = document.createElement('div') + const dispose = observeActiveTerminalResize(host, { onActivate, onFit }) + + // ResizeObserver's initial delivery is absorbed by the activation frame. + expect(raf.pending()).toBe(1) + raf.flush() + expect(onFit).toHaveBeenCalledTimes(1) + expect(onActivate).toHaveBeenCalledTimes(1) + + resize.current?.([], {} as ResizeObserver) + resize.current?.([], {} as ResizeObserver) + resize.current?.([], {} as ResizeObserver) + expect(raf.pending()).toBe(1) + raf.flush() + expect(onFit).toHaveBeenCalledTimes(2) + + dispose() + expect(disconnect).toHaveBeenCalledTimes(1) + }) + + it('cancels activation without fitting when hidden before the first frame', () => { + const raf = installRaf() + + vi.stubGlobal( + 'ResizeObserver', + class { + disconnect = vi.fn() + observe = vi.fn() + unobserve = vi.fn() + } as unknown as typeof ResizeObserver + ) + + const onFit = vi.fn() + + const dispose = observeActiveTerminalResize(document.createElement('div'), { + onActivate: vi.fn(), + onFit + }) + + dispose() + + raf.flush() + expect(onFit).not.toHaveBeenCalled() + }) + + it('absorbs a real browser-style initial resize delivered after activation', () => { + const raf = installRaf() + const resize = { current: null as ResizeObserverCallback | null } + + vi.stubGlobal( + 'ResizeObserver', + class { + constructor(callback: ResizeObserverCallback) { + resize.current = callback + } + + disconnect = vi.fn() + observe = vi.fn() + unobserve = vi.fn() + } as unknown as typeof ResizeObserver + ) + + const onFit = vi.fn() + observeActiveTerminalResize(document.createElement('div'), { onActivate: vi.fn(), onFit }) + + raf.flush() + expect(onFit).toHaveBeenCalledTimes(1) + + // Browser initial delivery: the activation fit already covered this size. + resize.current?.([], {} as ResizeObserver) + expect(raf.pending()).toBe(0) + + // A later real resize schedules exactly one fit. + resize.current?.([], {} as ResizeObserver) + expect(raf.pending()).toBe(1) + raf.flush() + expect(onFit).toHaveBeenCalledTimes(2) + }) + + it('reuses a first-mount fit without fitting again on activation', () => { + const raf = installRaf() + + vi.stubGlobal( + 'ResizeObserver', + class { + disconnect = vi.fn() + observe = vi.fn() + unobserve = vi.fn() + } as unknown as typeof ResizeObserver + ) + + const onActivate = vi.fn() + const onFit = vi.fn() + + observeActiveTerminalResize(document.createElement('div'), { + fitOnActivate: false, + onActivate, + onFit + }) + + raf.flush() + + expect(onActivate).toHaveBeenCalledOnce() + expect(onFit).not.toHaveBeenCalled() + }) +}) diff --git a/apps/desktop/src/app/right-sidebar/terminal/active-resize.ts b/apps/desktop/src/app/right-sidebar/terminal/active-resize.ts new file mode 100644 index 0000000000..ea68ceda1d --- /dev/null +++ b/apps/desktop/src/app/right-sidebar/terminal/active-resize.ts @@ -0,0 +1,78 @@ +interface ActiveTerminalResizeOptions { + fitOnActivate?: boolean + onActivate: () => void + onFit: () => void +} + +/** + * Observe one visible xterm host. + * + * Inactive terminals never call this helper, so their preserved DOM/PTY stays + * mounted without paying for ResizeObserver delivery or FitAddon work. The + * first frame owns activation and ignores the observer's initial delivery; + * later resize bursts are coalesced to one fit per animation frame. + */ +export function observeActiveTerminalResize( + host: HTMLElement, + { fitOnActivate = true, onActivate, onFit }: ActiveTerminalResizeOptions +): () => void { + let activated = false + let frame = 0 + let initialResizeDelivered = false + let stopped = false + + const scheduleFit = () => { + if (!activated || stopped || frame !== 0) { + return + } + + frame = window.requestAnimationFrame(() => { + frame = 0 + + if (!stopped) { + onFit() + } + }) + } + + const observer = new ResizeObserver(() => { + // ResizeObserver's initial delivery is asynchronous in browsers and may + // arrive before OR after the activation rAF. Activation already fits the + // current box, so absorb that first delivery in either ordering. + if (!initialResizeDelivered) { + initialResizeDelivered = true + + return + } + + scheduleFit() + }) + + observer.observe(host) + + frame = window.requestAnimationFrame(() => { + frame = 0 + + if (stopped) { + return + } + + activated = true + + if (fitOnActivate) { + onFit() + } + + onActivate() + }) + + return () => { + stopped = true + observer.disconnect() + + if (frame !== 0) { + window.cancelAnimationFrame(frame) + frame = 0 + } + } +} diff --git a/apps/desktop/src/app/right-sidebar/terminal/persistent.test.tsx b/apps/desktop/src/app/right-sidebar/terminal/persistent.test.tsx index c0f718e546..77eeb480d6 100644 --- a/apps/desktop/src/app/right-sidebar/terminal/persistent.test.tsx +++ b/apps/desktop/src/app/right-sidebar/terminal/persistent.test.tsx @@ -2,6 +2,8 @@ import { act, type ReactNode } from 'react' import { createRoot, type Root } from 'react-dom/client' import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' +import { $paneStates } from '@/store/panes' + import { PersistentTerminal, TerminalSlot } from './persistent' vi.mock('./terminals', () => ({ @@ -212,7 +214,8 @@ describe('PersistentTerminal rect tracking', () => { render() - expect(mutationObserveCalls.some(call => call.options?.subtree === true)).toBe(true) + expect(mutationObserveCalls.length).toBeGreaterThan(0) + expect(mutationObserveCalls.every(call => call.options?.subtree === false)).toBe(true) act(() => { raf.runNext() @@ -244,6 +247,27 @@ describe('PersistentTerminal rect tracking', () => { expect(raf.pending()).toBe(0) }) + it('remeasures from an explicit pane-layout state change', () => { + const raf = installRaf() + const before = $paneStates.get() + vi.spyOn(HTMLElement.prototype, 'getBoundingClientRect').mockReturnValue(rect(10, 20, 200, 100)) + + render() + raf.runNext() + expect(raf.pending()).toBe(0) + + act(() => { + $paneStates.set({ ...before, __terminal_rect_test__: { open: true } }) + }) + + expect(raf.pending()).toBe(1) + + act(() => { + raf.runNext() + $paneStates.set(before) + }) + }) + it('does not schedule rect RAFs while the Electron window is paused, then resumes when visible', () => { const raf = installRaf() vi.spyOn(HTMLElement.prototype, 'getBoundingClientRect').mockReturnValue(rect(10, 20, 200, 100)) diff --git a/apps/desktop/src/app/right-sidebar/terminal/persistent.tsx b/apps/desktop/src/app/right-sidebar/terminal/persistent.tsx index b5e75978b0..d3a7b0c829 100644 --- a/apps/desktop/src/app/right-sidebar/terminal/persistent.tsx +++ b/apps/desktop/src/app/right-sidebar/terminal/persistent.tsx @@ -2,7 +2,10 @@ import { useStore } from '@nanostores/react' import { atom } from 'nanostores' import { type CSSProperties, useEffect, useLayoutEffect, useRef, useState } from 'react' +import { $layoutTree } from '@/components/pane-shell/tree/store' +import { markRightPanePerf } from '@/debug/right-pane-events' import { createRendererLoopPauseController } from '@/lib/renderer-loop-pause' +import { $paneStates } from '@/store/panes' import { $terminalTakeover } from '../store' @@ -39,7 +42,7 @@ export function TerminalSlot({ className = SLOT_CLASS }: { className?: string }) } }, []) - return
+ return
} interface PersistentTerminalProps { @@ -86,6 +89,7 @@ export function PersistentTerminal({ onAddSelectionToChat }: PersistentTerminalP let prev: Rect | null = null let frame = 0 let stopped = false + let pendingReason = 'initial' let pauseController: ReturnType | null = null const rendererPaused = () => pauseController?.isPaused() ?? document.visibilityState === 'hidden' @@ -97,11 +101,12 @@ export function PersistentTerminal({ onAddSelectionToChat }: PersistentTerminalP } } - const measure = (): boolean => { + const measure = (reason: string): boolean => { if (rendererPaused()) { return false } + markRightPanePerf('terminal-measure', reason) const r = slot.getBoundingClientRect() // floor top/left + ceil right/bottom: overlay always covers the slot's // full pixel footprint, so half-pixel rects can't leak page bg through. @@ -123,16 +128,18 @@ export function PersistentTerminal({ onAddSelectionToChat }: PersistentTerminalP return false } - const scheduleMeasure = () => { + const scheduleMeasure = (reason = 'unknown') => { if (stopped || rendererPaused() || frame !== 0) { return } + pendingReason = reason frame = window.requestAnimationFrame(() => { frame = 0 + const reason = pendingReason - if (measure()) { - scheduleMeasure() + if (measure(reason)) { + scheduleMeasure('settle') } }) } @@ -144,50 +151,69 @@ export function PersistentTerminal({ onAddSelectionToChat }: PersistentTerminalP return } - scheduleMeasure() + scheduleMeasure('visibility') } const observer = typeof ResizeObserver === 'undefined' ? null : new ResizeObserver(() => { - scheduleMeasure() + scheduleMeasure('resize-observer') }) const positionObserver = typeof MutationObserver === 'undefined' ? null : new MutationObserver(() => { - scheduleMeasure() + scheduleMeasure('ancestor-mutation') }) pauseController = createRendererLoopPauseController(handleVisibilityChange) - if (measure()) { - scheduleMeasure() + if (measure('initial')) { + scheduleMeasure('settle') } observer?.observe(slot) + const handleScroll = () => scheduleMeasure('scroll') + const scrollTargets: Array = [window] + window.addEventListener('scroll', handleScroll) + for (let node: HTMLElement | null = slot; node; node = node.parentElement) { positionObserver?.observe(node, { attributeFilter: ['class', 'style', 'hidden', 'aria-hidden', 'data-state'], attributes: true, childList: true, - subtree: true + subtree: false }) + // Scroll does not bubble. Listen only on the slot's own ancestor chain, + // so a transcript/file-tree/xterm viewport scroll elsewhere cannot wake + // terminal positioning. + node.addEventListener('scroll', handleScroll) + scrollTargets.push(node) } - window.addEventListener('resize', scheduleMeasure) - window.addEventListener('scroll', scheduleMeasure, true) + // Nested layout-tree and pane-state commits can move the slot without + // changing its own size. Subscribe to the actual layout authorities instead + // of observing every descendant mutation under every ancestor (chat stream + // and file-tree updates are unrelated and used to wake this tracker). + const unsubscribeLayout = $layoutTree.listen(() => scheduleMeasure('layout-tree')) + const unsubscribePanes = $paneStates.listen(() => scheduleMeasure('pane-state')) + + const handleResize = () => scheduleMeasure('window-resize') + + window.addEventListener('resize', handleResize) return () => { stopped = true cancelFrame() observer?.disconnect() positionObserver?.disconnect() - window.removeEventListener('resize', scheduleMeasure) - window.removeEventListener('scroll', scheduleMeasure, true) + unsubscribeLayout() + unsubscribePanes() + window.removeEventListener('resize', handleResize) + scrollTargets.forEach(target => target.removeEventListener('scroll', handleScroll)) pauseController?.dispose() } }, [slot]) @@ -215,7 +241,7 @@ export function PersistentTerminal({ onAddSelectionToChat }: PersistentTerminalP // booting xterm/node-pty at 0×0 starts the shell at 80×24 and spawns a visible // conhost on Windows. After that `mounted` latches: shells persist while hidden. return ( -
+
{mounted && }
) diff --git a/apps/desktop/src/app/right-sidebar/terminal/use-agent-terminal.ts b/apps/desktop/src/app/right-sidebar/terminal/use-agent-terminal.ts index c0947eb0d9..339d70144b 100644 --- a/apps/desktop/src/app/right-sidebar/terminal/use-agent-terminal.ts +++ b/apps/desktop/src/app/right-sidebar/terminal/use-agent-terminal.ts @@ -5,9 +5,11 @@ import { Terminal } from '@xterm/xterm' import { useEffect, useRef } from 'react' import { writeClipboardText } from '@/components/ui/copy-button' +import { markRightPanePerf } from '@/debug/right-pane-events' import { triggerHaptic } from '@/lib/haptics' import { useTheme } from '@/themes/context' +import { observeActiveTerminalResize } from './active-resize' import { registerAgentTerminalWriter } from './agent-terminal-stream' import { makeTerminalReader, registerTerminalReader } from './buffer' import { mirrorSelection, terminalClipboardIntent } from './clipboard' @@ -24,7 +26,8 @@ export function useAgentTerminal({ active, id, procId }: { active: boolean; id: const hostRef = useRef(null) const termRef = useRef(null) const webglRef = useRef(null) - const fitRef = useRef<(() => void) | null>(null) + const fitRef = useRef<((visible: boolean) => void) | null>(null) + const initialActiveFitRef = useRef(false) const { latestFontFamilyRef, mountedRef } = useTerminalFontController({ fitRef, termRef, webglRef }) const surfaceTheme = () => { @@ -47,7 +50,6 @@ export function useAgentTerminal({ active, id, procId }: { active: boolean; id: } let disposed = false - let observer: ResizeObserver | null = null let unregister = () => {} @@ -101,10 +103,11 @@ export function useAgentTerminal({ active, id, procId }: { active: boolean; id: return false }) - fitRef.current = () => { + fitRef.current = visible => { if (host.clientWidth > 0 && host.clientHeight > 0) { try { fit.fit() + markRightPanePerf(visible ? 'terminal-fit-active' : 'terminal-fit-hidden', id) } catch { // Mid-transition layout — the next observer tick refits. } @@ -132,9 +135,8 @@ export function useAgentTerminal({ active, id, procId }: { active: boolean; id: // No WebGL — xterm falls back to the DOM renderer. } - fitRef.current?.() - observer = new ResizeObserver(() => fitRef.current?.()) - observer.observe(host) + fitRef.current?.(active) + initialActiveFitRef.current = active // Stream live output straight into the terminal (replays backlog on attach). unregister = registerAgentTerminalWriter(procId, chunk => term.write(chunk)) @@ -159,7 +161,6 @@ export function useAgentTerminal({ active, id, procId }: { active: boolean; id: unregister() unregisterReader() selectionDisposable.dispose() - observer?.disconnect() fitRef.current = null term.dispose() termRef.current = null @@ -184,25 +185,39 @@ export function useAgentTerminal({ active, id, procId }: { active: boolean; id: // eslint-disable-next-line react-hooks/exhaustive-deps }, [renderedMode, themeName]) - // A visibility:hidden xterm doesn't paint — refit + redraw on re-activation. + // Keep inactive agent terminals mounted for their backlog, but do not observe + // or fit them until they become the visible tab. + // eslint-disable-next-line no-restricted-syntax -- lifecycle flag prevents a duplicate first-mount fit useEffect(() => { if (!active) { + initialActiveFitRef.current = false + return } - const frame = requestAnimationFrame(() => { - const term = termRef.current + const host = hostRef.current - fitRef.current?.() - webglRef.current?.clearTextureAtlas() - term?.refresh(0, term.rows - 1) - // Take focus on activation (parity with the user terminal) so the active - // agent tab holds focus and ⌘W's isFocusWithin('[data-terminal]') routes - // the close to this tab rather than to a preview. - term?.focus() + if (!host) { + return + } + + const fitOnActivate = !initialActiveFitRef.current + initialActiveFitRef.current = false + + return observeActiveTerminalResize(host, { + fitOnActivate, + onFit: () => fitRef.current?.(true), + onActivate: () => { + const term = termRef.current + + webglRef.current?.clearTextureAtlas() + term?.refresh(0, term.rows - 1) + // Take focus on activation (parity with the user terminal) so the active + // agent tab holds focus and ⌘W's isFocusWithin('[data-terminal]') routes + // the close to this tab rather than to a preview. + term?.focus() + } }) - - return () => cancelAnimationFrame(frame) }, [active]) return { hostRef } diff --git a/apps/desktop/src/app/right-sidebar/terminal/use-terminal-font.ts b/apps/desktop/src/app/right-sidebar/terminal/use-terminal-font.ts index 656968071f..8bcf76426e 100644 --- a/apps/desktop/src/app/right-sidebar/terminal/use-terminal-font.ts +++ b/apps/desktop/src/app/right-sidebar/terminal/use-terminal-font.ts @@ -7,7 +7,7 @@ import type { RefObject } from 'react' import { $terminalFontFamily, applyTerminalFontFamily, resolveTerminalFontFamily } from './terminal-font' interface TerminalFontControllerOptions { - fitRef: RefObject<(() => void) | null> + fitRef: RefObject<((visible: boolean) => void) | null> termRef: RefObject webglRef: RefObject } @@ -38,7 +38,7 @@ export function useTerminalFontController({ fitRef, termRef, webglRef }: Termina void applyTerminalFontFamily({ clearTextureAtlas: () => webglRef.current?.clearTextureAtlas(), - fit: () => fitRef.current?.(), + fit: () => fitRef.current?.(true), fontFamily, isCurrent: () => !cancelled && generationRef.current === generation, term diff --git a/apps/desktop/src/app/right-sidebar/terminal/use-terminal-session.ts b/apps/desktop/src/app/right-sidebar/terminal/use-terminal-session.ts index e4e80fd3a4..24531cabe2 100644 --- a/apps/desktop/src/app/right-sidebar/terminal/use-terminal-session.ts +++ b/apps/desktop/src/app/right-sidebar/terminal/use-terminal-session.ts @@ -7,12 +7,14 @@ import { useCallback, useEffect, useMemo, useRef, useState } from 'react' import type { CSSProperties } from 'react' import { writeClipboardText } from '@/components/ui/copy-button' +import { markRightPanePerf } from '@/debug/right-pane-events' import { triggerHaptic } from '@/lib/haptics' import { $previewTarget } from '@/store/preview' import { useTheme } from '@/themes/context' import { $terminalInjection } from '../store' +import { observeActiveTerminalResize } from './active-resize' import { makeTerminalReader, registerTerminalReader } from './buffer' import { mirrorSelection, terminalClipboardIntent } from './clipboard' import { terminalLinkHandler, terminalWebLinksAddon } from './links' @@ -413,6 +415,7 @@ export function useTerminalSession({ // drag-and-drop paths, or an injected command). Gates idle-buffer handling in // persistSnapshot so an untouched tab never re-saves an accumulating snapshot. const hasSessionActivityRef = useRef(false) + const initialActiveRef = useRef(active) const shellNameRef = useRef('shell') const selectionLabelRef = useRef('') const selectionRef = useRef('') @@ -420,7 +423,8 @@ export function useTerminalSession({ const onShellRef = useRef(onShell) // Re-fit on activation: a tab hidden via display:none has a 0×0 host, so its // last fit is stale by the time it's shown again. - const fitRef = useRef<(() => void) | null>(null) + const fitRef = useRef<((visible: boolean) => void) | null>(null) + const initialActiveFitRef = useRef(false) const { latestFontFamilyRef, mountedRef } = useTerminalFontController({ fitRef, termRef, webglRef }) const [status, setStatus] = useState('starting') const [selection, setSelection] = useState('') @@ -742,56 +746,28 @@ export function useTerminalSession({ term.write(next) } - const fitAndResize = () => { + const fitAndResize = (visible: boolean) => { if (disposed || !host.isConnected || host.clientWidth <= 0 || host.clientHeight <= 0) { return } try { fit.fit() + markRightPanePerf(visible ? 'terminal-fit-active' : 'terminal-fit-hidden', id) } catch { return } - const id = sessionIdRef.current + const sessionId = sessionIdRef.current - if (id && (lastSentSize?.cols !== term.cols || lastSentSize?.rows !== term.rows)) { + if (sessionId && (lastSentSize?.cols !== term.cols || lastSentSize?.rows !== term.rows)) { lastSentSize = { cols: term.cols, rows: term.rows } - void terminalApi.resize(id, { cols: term.cols, rows: term.rows }) + void terminalApi.resize(sessionId, { cols: term.cols, rows: term.rows }) } } fitRef.current = fitAndResize - // Coalesce ResizeObserver bursts through rAF — running fit.fit() - // synchronously while sibling panes are mid-transition (e.g. file browser - // collapsing to 0px) crashes the WebGL renderer mid texture-atlas rebuild. - let pendingFrame = 0 - - const scheduleResize = () => { - if (pendingFrame) { - return - } - - pendingFrame = window.requestAnimationFrame(() => { - pendingFrame = 0 - - if (!disposed) { - fitAndResize() - } - }) - } - - const resizeObserver = new ResizeObserver(scheduleResize) - resizeObserver.observe(host) - cleanup.push(() => { - resizeObserver.disconnect() - - if (pendingFrame) { - window.cancelAnimationFrame(pendingFrame) - } - }) - const dataDisposable = term.onData(data => { hasSessionActivityRef.current = true const id = sessionIdRef.current @@ -901,9 +877,7 @@ export function useTerminalSession({ ) window.requestAnimationFrame(() => { - fitAndResize() term.clearSelection() // drop any selection painted over transient boot rows - term.focus() }) }) .catch(error => { @@ -938,7 +912,8 @@ export function useTerminalSession({ console.warn('[hermes-terminal] WebGL unavailable; falling back to DOM', err) } - fitAndResize() + fitAndResize(initialActiveRef.current) + initialActiveFitRef.current = initialActiveRef.current startSession() } @@ -1013,24 +988,39 @@ export function useTerminalSession({ return term ? registerTerminalReader(id, makeTerminalReader(term)) : undefined }, [id, status]) - // On (re)activation: a WebGL terminal doesn't paint while visibility:hidden, so - // it reveals a stale/garbled frame. Refit, rebuild the glyph atlas, and force a - // full redraw against the live buffer, then focus. + // Only the active terminal observes its host. Every terminal stays mounted + // (PTY + scrollback preserved), but hidden tabs do no FitAddon/layout work. + // Re-activation owns one fit + atlas rebuild + redraw. + // eslint-disable-next-line no-restricted-syntax -- lifecycle flag prevents a duplicate first-mount fit useEffect(() => { if (!active || status !== 'open') { + if (!active) { + initialActiveFitRef.current = false + } + return } - const frame = requestAnimationFrame(() => { - const term = termRef.current + const host = hostRef.current - fitRef.current?.() - webglRef.current?.clearTextureAtlas() - term?.refresh(0, term.rows - 1) - term?.focus() + if (!host) { + return + } + + const fitOnActivate = !initialActiveFitRef.current + initialActiveFitRef.current = false + + return observeActiveTerminalResize(host, { + fitOnActivate, + onFit: () => fitRef.current?.(true), + onActivate: () => { + const term = termRef.current + + webglRef.current?.clearTextureAtlas() + term?.refresh(0, term.rows - 1) + term?.focus() + } }) - - return () => cancelAnimationFrame(frame) }, [active, status]) // Flush a queued command (e.g. a provider-disconnect) into the live session. diff --git a/apps/desktop/src/app/session/hooks/use-message-stream/delta-flush.test.tsx b/apps/desktop/src/app/session/hooks/use-message-stream/delta-flush.test.tsx index 1a572c329a..a630d9dcf3 100644 --- a/apps/desktop/src/app/session/hooks/use-message-stream/delta-flush.test.tsx +++ b/apps/desktop/src/app/session/hooks/use-message-stream/delta-flush.test.tsx @@ -1,11 +1,14 @@ import { QueryClient } from '@tanstack/react-query' import { act, cleanup, render } from '@testing-library/react' -import { useEffect, useRef } from 'react' +import { type MutableRefObject, useEffect, useRef } from 'react' import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' import type { ClientSessionState } from '@/app/types' +import type { ChatMessage } from '@/lib/chat-messages' import { createClientSessionState } from '@/lib/chat-runtime' +import { useSessionStateCache } from '../use-session-state-cache' + import { useMessageStream } from './index' const SID = 'session-1' @@ -90,6 +93,42 @@ describe('useMessageStream delta flush scheduling', () => { expect(assistantText()).toBe('still streaming') }) + it('flushes queued text immediately when a hidden window becomes visible', () => { + vi.mocked(performance.now).mockReturnValue(0) + mountStream() + + act(() => appendAssistantDelta!(SID, 'caught up on focus')) + expect(assistantText()).toBe('') + expect(vi.getTimerCount()).toBe(1) + + Object.defineProperty(globalThis.document, 'visibilityState', { + configurable: true, + value: 'visible' + }) + + act(() => globalThis.document.dispatchEvent(new Event('visibilitychange'))) + + expect(assistantText()).toBe('caught up on focus') + expect(vi.getTimerCount()).toBe(0) + }) + + it('flushes queued text on focus when visibility remains visible', () => { + vi.mocked(performance.now).mockReturnValue(0) + Object.defineProperty(globalThis.document, 'visibilityState', { + configurable: true, + value: 'visible' + }) + mountStream() + + act(() => appendAssistantDelta!(SID, 'focused without visibility change')) + expect(assistantText()).toBe('') + expect(vi.getTimerCount()).toBe(1) + + act(() => globalThis.window.dispatchEvent(new Event('focus'))) + + expect(assistantText()).toBe('focused without visibility change') + }) + it('cancels the pending timer on unmount and flushes exactly once', async () => { vi.mocked(performance.now).mockReturnValue(0) mountStream() @@ -108,4 +147,275 @@ describe('useMessageStream delta flush scheduling', () => { expect(updateSessionState).toHaveBeenCalledTimes(updatesAfterUnmount) expect(window.requestAnimationFrame).not.toHaveBeenCalled() }) + + it('stretches the flush gap when the deferred commit frame is expensive', async () => { + // The streaming-path $messages publish (React commit + Streamdown + // re-parse) is deferred to a view-sync rAF inside updateSessionState, so + // the flush cost must be measured through that frame. Simulate one + // expensive frame and expect the next gap to adapt to 3x the frame cost. + let now = 1000 + vi.mocked(performance.now).mockImplementation(() => now) + const rafCallbacks: FrameRequestCallback[] = [] + vi.mocked(window.requestAnimationFrame).mockImplementation(cb => { + rafCallbacks.push(cb) + + return rafCallbacks.length + }) + + mountStream() + + act(() => appendAssistantDelta!(SID, 'first')) + await act(async () => { + await vi.advanceTimersByTimeAsync(0) + }) + + expect(assistantText()).toBe('first') + expect(rafCallbacks).toHaveLength(1) + + // Frame started at 1040, the measurement callback runs at 1100: 60ms of + // in-frame work (view sync + commit), so the next floor is 180ms. + now = 1100 + act(() => rafCallbacks[0](1040)) + + act(() => appendAssistantDelta!(SID, 'second')) + await act(async () => { + await vi.advanceTimersByTimeAsync(79) + }) + + expect(assistantText()).toBe('first') + + await act(async () => { + await vi.advanceTimersByTimeAsync(1) + }) + + expect(assistantText()).toBe('firstsecond') + }) + + it('keeps the write-cost floor when no frame fires (hidden renderer)', async () => { + // A parked renderer never runs rAF callbacks. The cost must stay at the + // synchronous store-write measurement so the gap falls back to the fixed + // 33ms floor instead of waiting on a frame that will never come. + let now = 1000 + vi.mocked(performance.now).mockImplementation(() => now) + vi.mocked(window.requestAnimationFrame).mockImplementation(() => 1) + + mountStream() + + act(() => appendAssistantDelta!(SID, 'first')) + await act(async () => { + await vi.advanceTimersByTimeAsync(0) + }) + + expect(assistantText()).toBe('first') + + // 100ms later (well past the 33ms floor): the next flush is immediate. + now = 1100 + act(() => appendAssistantDelta!(SID, 'second')) + await act(async () => { + await vi.advanceTimersByTimeAsync(0) + }) + + expect(assistantText()).toBe('firstsecond') + }) + + it('ignores a late frame measurement once a newer flush has started', async () => { + let now = 1000 + vi.mocked(performance.now).mockImplementation(() => now) + const rafCallbacks: FrameRequestCallback[] = [] + vi.mocked(window.requestAnimationFrame).mockImplementation(cb => { + rafCallbacks.push(cb) + + return rafCallbacks.length + }) + + mountStream() + + act(() => appendAssistantDelta!(SID, 'a')) + await act(async () => { + await vi.advanceTimersByTimeAsync(0) + }) + + // A second flush starts before the first flush's frame lands. + now = 1010 + act(() => appendAssistantDelta!(SID, 'b')) + await act(async () => { + await vi.advanceTimersByTimeAsync(23) + }) + + expect(assistantText()).toBe('ab') + expect(rafCallbacks).toHaveLength(2) + + // The stale callback must not overwrite the newer flush's cost. If it + // did, cost would read 30ms and the next gap would stretch to 70ms. + now = 1030 + act(() => rafCallbacks[0](1000)) + + act(() => appendAssistantDelta!(SID, 'c')) + await act(async () => { + await vi.advanceTimersByTimeAsync(13) + }) + + expect(assistantText()).toBe('abc') + }) +}) + +describe('useMessageStream composed with the real useSessionStateCache', () => { + // The tests above mock updateSessionState, so they validate the adaptive + // arithmetic but not the production ordering contract: runFlush's + // measurement rAF must be registered AFTER the view-sync rAF that the real + // updateSessionState schedules inside syncSessionStateToView, so the + // measured frame cost includes the deferred $messages commit it adapts to. + let cache: ReturnType | null = null + let published: ChatMessage[] + + function ComposedHarness() { + const busyRef: MutableRefObject = { current: false } + const queryClientRef = useRef(new QueryClient()) + + const sessionCache = useSessionStateCache({ + activeSessionId: SID, + busyRef, + selectedStoredSessionId: null, + setAwaitingResponse: () => undefined, + setBusy: () => undefined, + setMessages: messages => { + published = messages + } + }) + + const stream = useMessageStream({ + activeSessionIdRef: sessionCache.activeSessionIdRef, + hydrateFromStoredSession: vi.fn(async () => undefined), + queryClient: queryClientRef.current, + refreshHermesConfig: vi.fn(async () => undefined), + refreshSessions: vi.fn(async () => undefined), + sessionStateByRuntimeIdRef: sessionCache.sessionStateByRuntimeIdRef, + updateSessionState: sessionCache.updateSessionState + }) + + useEffect(() => { + appendAssistantDelta = stream.appendAssistantDelta + cache = sessionCache + }, [stream.appendAssistantDelta, sessionCache]) + + return null + } + + function cachedText() { + const message = cache?.sessionStateByRuntimeIdRef.current.get(SID)?.messages.at(-1) + const part = message?.parts.at(-1) + + return part?.type === 'text' ? part.text : '' + } + + function publishedText() { + const part = published.at(-1)?.parts.at(-1) + + return part?.type === 'text' ? part.text : '' + } + + beforeEach(() => { + vi.useFakeTimers() + appendAssistantDelta = null + cache = null + published = [] + vi.spyOn(performance, 'now').mockReturnValue(100) + vi.spyOn(window, 'requestAnimationFrame').mockImplementation(() => 1) + vi.spyOn(window, 'cancelAnimationFrame').mockImplementation(() => undefined) + vi.spyOn(document, 'hasFocus').mockReturnValue(false) + }) + + afterEach(() => { + cleanup() + vi.useRealTimers() + vi.restoreAllMocks() + }) + + it('measures the frame cost through the real view-sync rAF and adapts the next gap', async () => { + let now = 1000 + vi.mocked(performance.now).mockImplementation(() => now) + const rafCallbacks: FrameRequestCallback[] = [] + vi.mocked(window.requestAnimationFrame).mockImplementation(cb => { + rafCallbacks.push(cb) + + return rafCallbacks.length + }) + + render() + expect(appendAssistantDelta).not.toBeNull() + + // Mid-turn state: busy keeps the view sync on the deferred rAF path + // (terminal/needing-input states flush synchronously instead). + act(() => { + cache!.updateSessionState(SID, state => ({ ...state, busy: true })) + }) + expect(rafCallbacks).toHaveLength(1) + // Drain the seed's own view-sync rAF so the flush below starts clean. + act(() => rafCallbacks.shift()!(now)) + + act(() => appendAssistantDelta!(SID, 'first')) + await act(async () => { + await vi.advanceTimersByTimeAsync(0) + }) + + // The store write landed synchronously, but the $messages publish is + // deferred: exactly two rAF callbacks are pending — first the cache's + // view-sync, then runFlush's measurement. + expect(cachedText()).toBe('first') + expect(publishedText()).toBe('') + expect(rafCallbacks).toHaveLength(2) + + // Draining the FIRST registered callback must be what publishes the + // deferred commit; that identity is the ordering contract. It runs until + // 60ms into the frame (React commit + Streamdown re-parse). + now = 1100 + act(() => rafCallbacks[0](1040)) + expect(publishedText()).toBe('first') + + // The measurement callback closes the same frame: 60ms of in-frame work, + // so the next adaptive floor is 3x = 180ms. + act(() => rafCallbacks[1](1040)) + + act(() => appendAssistantDelta!(SID, 'second')) + await act(async () => { + await vi.advanceTimersByTimeAsync(79) + }) + + expect(cachedText()).toBe('first') + + await act(async () => { + await vi.advanceTimersByTimeAsync(1) + }) + + expect(cachedText()).toBe('firstsecond') + }) + + it('keeps the write-cost fallback when the parked renderer never fires rAF', async () => { + let now = 1000 + vi.mocked(performance.now).mockImplementation(() => now) + // Parked renderer: rAF callbacks are accepted but never run. + vi.mocked(window.requestAnimationFrame).mockImplementation(() => 1) + + render() + + act(() => { + cache!.updateSessionState(SID, state => ({ ...state, busy: true })) + }) + + act(() => appendAssistantDelta!(SID, 'first')) + await act(async () => { + await vi.advanceTimersByTimeAsync(0) + }) + + expect(cachedText()).toBe('first') + + // 100ms later (well past the 33ms floor): the next flush is immediate. + now = 1100 + act(() => appendAssistantDelta!(SID, 'second')) + await act(async () => { + await vi.advanceTimersByTimeAsync(0) + }) + + expect(cachedText()).toBe('firstsecond') + }) }) diff --git a/apps/desktop/src/app/session/hooks/use-message-stream/index.ts b/apps/desktop/src/app/session/hooks/use-message-stream/index.ts index 29051cc365..a73852f6e9 100644 --- a/apps/desktop/src/app/session/hooks/use-message-stream/index.ts +++ b/apps/desktop/src/app/session/hooks/use-message-stream/index.ts @@ -188,6 +188,9 @@ export function useMessageStream({ // What the previous flush cost on the main thread — drives the adaptive // flush floor in scheduleDeltaFlush so multi-stream load yields to input. const lastFlushCostRef = useRef(0) + // The pending commit-cost measurement rAF, so a newer flush (or unmount) + // can cancel it instead of letting parked callbacks pile up while hidden. + const measureRafRef = useRef(null) const nativeSubagentSessionsRef = useRef>(new Set()) // Turns that auto-compacted: skip post-turn hydrate so live scrollback survives. const compactedTurnRef = useRef>(new Set()) @@ -257,6 +260,8 @@ export function useMessageStream({ // keeps the thread ~75% idle for input at any load: cheap flushes stay at // 30fps of text growth, expensive multi-stream flushes degrade text fps // instead of interactivity — capped so text never updates slower than 4/s. + // The cost has to include the deferred view-sync frame where the commit + // actually happens; see runFlush below. const sinceLast = performance.now() - lastFlushAtRef.current const adaptiveFloor = Math.min( @@ -269,7 +274,39 @@ export function useMessageStream({ const startedAt = performance.now() lastFlushAtRef.current = startedAt flushQueuedDeltas() - lastFlushCostRef.current = performance.now() - startedAt + // The store write above is only the cheap half of a flush. While a + // session streams, syncSessionStateToView defers the $messages publish + // (and with it the React commit + Streamdown re-parse the floor is meant + // to account for) to its own rAF inside updateSessionState, which runs + // after this timer task. Stopping the clock here pins lastFlushCostRef + // near zero and collapses the adaptive floor to 33ms no matter the load. + // Our rAF is registered after the view-sync one, so it runs in the same + // frame right after that commit; its timestamp marks frame start, so + // (now - frameStart) counts only work done inside the frame, not the + // vsync wait. A hidden renderer never fires rAF, so the write cost + // stays as the fallback. + const writeCost = performance.now() - startedAt + lastFlushCostRef.current = writeCost + + // At most one measurement rAF may be pending: only the newest flush's + // measurement matters (the guard below discards stale frames), and a + // hidden renderer parks rAF callbacks — without cancellation a long + // hidden stream at the floor would accumulate thousands of parked + // closures that all fire in the first frame on refocus. + if (measureRafRef.current !== null) { + window.cancelAnimationFrame(measureRafRef.current) + } + + measureRafRef.current = window.requestAnimationFrame(frameStart => { + measureRafRef.current = null + + // A newer flush already started; its own measurement wins. + if (lastFlushAtRef.current !== startedAt) { + return + } + + lastFlushCostRef.current = writeCost + Math.max(0, performance.now() - frameStart) + }) } // Always a timer, never requestAnimationFrame. Chromium pauses rAF for a @@ -312,11 +349,46 @@ export function useMessageStream({ } flushHandleRef.current = null + + if (measureRafRef.current !== null && typeof window !== 'undefined') { + window.cancelAnimationFrame(measureRafRef.current) + } + + measureRafRef.current = null flushQueuedDeltas() }, [flushQueuedDeltas] ) + // Page Visibility does not report every Windows/Linux focus transition. + // Flush queued deltas on both signals so returning to a chat cannot leave a + // completed chunk waiting for the next throttled timer. + // eslint-disable-next-line no-restricted-syntax -- timer-handle clear inside effect, not an atom mirror + useEffect(() => { + const flushPendingDeltas = () => { + if (flushHandleRef.current !== null) { + window.clearTimeout(flushHandleRef.current) + flushHandleRef.current = null + } + + flushQueuedDeltas() + } + + const flushWhenVisible = () => { + if (document.visibilityState === 'visible') { + flushPendingDeltas() + } + } + + document.addEventListener('visibilitychange', flushWhenVisible) + window.addEventListener('focus', flushPendingDeltas) + + return () => { + document.removeEventListener('visibilitychange', flushWhenVisible) + window.removeEventListener('focus', flushPendingDeltas) + } + }, [flushQueuedDeltas]) + const appendAssistantDelta = useCallback( (sessionId: string, delta: string) => { if (!delta) { diff --git a/apps/desktop/src/app/session/hooks/use-message-stream/stream-flush.test.tsx b/apps/desktop/src/app/session/hooks/use-message-stream/stream-flush.test.tsx index 427e62cf8a..840416a18e 100644 --- a/apps/desktop/src/app/session/hooks/use-message-stream/stream-flush.test.tsx +++ b/apps/desktop/src/app/session/hooks/use-message-stream/stream-flush.test.tsx @@ -87,7 +87,10 @@ describe('stream delta delivery', () => { }) expect(states.get(SID)?.messages.at(-1)?.parts).toEqual([{ type: 'text', text: 'first and the rest' }]) - // The flush must not have depended on a frame at all. - expect(rafSpy).not.toHaveBeenCalled() + // The flush must not have depended on a frame: this mock parks every rAF + // callback, yet the text arrived. runFlush still registers its + // adaptive-floor measurement callback here; that one is allowed to wait + // for a frame that may never come. + expect(rafSpy).toHaveBeenCalled() }) }) diff --git a/apps/desktop/src/app/session/hooks/use-prompt-actions/index.test.tsx b/apps/desktop/src/app/session/hooks/use-prompt-actions/index.test.tsx index 57eea9ca06..9ac368fd44 100644 --- a/apps/desktop/src/app/session/hooks/use-prompt-actions/index.test.tsx +++ b/apps/desktop/src/app/session/hooks/use-prompt-actions/index.test.tsx @@ -1797,7 +1797,8 @@ describe('usePromptActions submit / queue drain semantics', () => { expect(accepted).toBe(true) expect(requestGateway).toHaveBeenCalledWith('session.resume', { session_id: 'stored-session-b', - source: 'desktop' + source: 'desktop', + omit_messages: true }) expect(requestGateway).toHaveBeenCalledWith( 'prompt.submit', @@ -1933,7 +1934,8 @@ describe('usePromptActions submit / queue drain semantics', () => { // Must resume the correct stored session to get the right runtime id. expect(requestGateway).toHaveBeenCalledWith('session.resume', { session_id: 'stored-session-a', - source: 'desktop' + source: 'desktop', + omit_messages: true }) // The prompt must land in the resumed session, NOT the foreground. expect(requestGateway).toHaveBeenCalledWith( @@ -2194,7 +2196,7 @@ describe('usePromptActions redirectPrompt', () => { expect(await handle!.redirectPrompt('reconnect nudge')).toBe(true) expect(calls.map(c => c.method)).toEqual(['session.redirect', 'session.resume', 'session.redirect']) expect(calls[0]?.params).toEqual({ session_id: RUNTIME_SESSION_ID, text: 'reconnect nudge' }) - expect(calls[1]?.params).toEqual({ session_id: STORED_SESSION_ID, source: 'desktop' }) + expect(calls[1]?.params).toEqual({ session_id: STORED_SESSION_ID, source: 'desktop', omit_messages: true }) expect(calls[2]?.params).toEqual({ session_id: RECOVERED_SESSION_ID, text: 'reconnect nudge' }) expect(handle!.activeSessionIdRef.current).toBe(RECOVERED_SESSION_ID) }) @@ -2735,7 +2737,7 @@ describe('usePromptActions sleep/wake session recovery', () => { expect(ok).toBe(true) // First submit (stale id) → session.resume (stored id) → retry submit (fresh id). expect(calls.map(c => c.method)).toEqual(['prompt.submit', 'session.resume', 'prompt.submit']) - expect(calls[1]?.params).toEqual({ session_id: STORED_SESSION_ID, source: 'desktop' }) + expect(calls[1]?.params).toEqual({ session_id: STORED_SESSION_ID, source: 'desktop', omit_messages: true }) expect(calls[2]?.params).toEqual({ session_id: RECOVERED_SESSION_ID, text: 'message after wake' }) }) @@ -2779,7 +2781,12 @@ describe('usePromptActions sleep/wake session recovery', () => { ) expect(await handle!.submitText('message after wake')).toBe(true) - expect(calls[1]?.params).toEqual({ session_id: STORED_SESSION_ID, source: 'desktop', profile: 'work' }) + expect(calls[1]?.params).toEqual({ + session_id: STORED_SESSION_ID, + source: 'desktop', + omit_messages: true, + profile: 'work' + }) setSessions(() => []) }) @@ -2826,7 +2833,12 @@ describe('usePromptActions sleep/wake session recovery', () => { ) expect(await handle!.submitText('message after wake')).toBe(true) - expect(calls[1]?.params).toEqual({ session_id: STORED_SESSION_ID, source: 'desktop', profile: 'work' }) + expect(calls[1]?.params).toEqual({ + session_id: STORED_SESSION_ID, + source: 'desktop', + omit_messages: true, + profile: 'work' + }) vi.mocked(getSession).mockReset() setSessions(() => []) @@ -2887,7 +2899,11 @@ describe('usePromptActions sleep/wake session recovery', () => { session_id: 'rt-background-stale', text: 'queued background message after wake' }) - expect(calls[1]?.params).toEqual({ session_id: STORED_SESSION_ID, source: 'desktop' }) + expect(calls[1]?.params).toEqual({ + session_id: STORED_SESSION_ID, + source: 'desktop', + omit_messages: true + }) expect(calls[2]?.params).toEqual({ queued: true, session_id: RECOVERED_SESSION_ID, @@ -2935,7 +2951,11 @@ describe('usePromptActions sleep/wake session recovery', () => { expect(calls.map(c => c.method)).toEqual(['session.interrupt', 'session.resume', 'session.interrupt']) expect(calls[0]?.params).toEqual({ session_id: RUNTIME_SESSION_ID }) - expect(calls[1]?.params).toEqual({ session_id: STORED_SESSION_ID, source: 'desktop' }) + expect(calls[1]?.params).toEqual({ + session_id: STORED_SESSION_ID, + source: 'desktop', + omit_messages: true + }) expect(calls[2]?.params).toEqual({ session_id: RECOVERED_SESSION_ID }) }) @@ -3067,7 +3087,11 @@ describe('usePromptActions sleep/wake session recovery', () => { expect(ok).toBe(true) expect(calls.map(c => c.method)).toEqual(['prompt.submit', 'session.resume', 'prompt.submit']) - expect(calls[1]?.params).toEqual({ session_id: STORED_SESSION_ID, source: 'desktop' }) + expect(calls[1]?.params).toEqual({ + session_id: STORED_SESSION_ID, + source: 'desktop', + omit_messages: true + }) expect(calls[2]?.params).toEqual({ session_id: RECOVERED_SESSION_ID, text: 'message during starved loop' @@ -3110,7 +3134,11 @@ describe('usePromptActions sleep/wake session recovery', () => { expect(ok).toBe(true) expect(createBackendSessionForSend).not.toHaveBeenCalled() expect(calls.map(c => c.method)).toEqual(['session.resume', 'prompt.submit']) - expect(calls[0]?.params).toEqual({ session_id: STORED_SESSION_ID, source: 'desktop' }) + expect(calls[0]?.params).toEqual({ + session_id: STORED_SESSION_ID, + source: 'desktop', + omit_messages: true + }) expect(calls[1]?.params).toMatchObject({ session_id: RECOVERED_SESSION_ID }) }) @@ -3421,7 +3449,8 @@ describe('usePromptActions submit session-context isolation (#54527)', () => { expect(calls.some(c => c.method === 'prompt.submit')).toBe(false) expect(calls.find(c => c.method === 'session.resume')?.params).toEqual({ session_id: STORED_SESSION_A, - source: 'desktop' + source: 'desktop', + omit_messages: true }) }) diff --git a/apps/desktop/src/app/session/hooks/use-prompt-actions/index.ts b/apps/desktop/src/app/session/hooks/use-prompt-actions/index.ts index 6307969421..186773a874 100644 --- a/apps/desktop/src/app/session/hooks/use-prompt-actions/index.ts +++ b/apps/desktop/src/app/session/hooks/use-prompt-actions/index.ts @@ -639,6 +639,7 @@ export function usePromptActions({ const resumed = await requestGateway<{ session_id: string }>('session.resume', { session_id: selectedStoredSessionIdRef.current, source: 'desktop', + omit_messages: true, ...(resumeProfile ? { profile: resumeProfile } : {}) }) @@ -744,6 +745,7 @@ export function usePromptActions({ const resumed = await requestGateway<{ session_id: string }>('session.resume', { session_id: selectedStoredSessionIdRef.current, source: 'desktop', + omit_messages: true, ...(resumeProfile ? { profile: resumeProfile } : {}) }) diff --git a/apps/desktop/src/app/session/hooks/use-prompt-actions/submit.ts b/apps/desktop/src/app/session/hooks/use-prompt-actions/submit.ts index 20de6ea35c..36c8bfd82d 100644 --- a/apps/desktop/src/app/session/hooks/use-prompt-actions/submit.ts +++ b/apps/desktop/src/app/session/hooks/use-prompt-actions/submit.ts @@ -484,6 +484,7 @@ export function useSubmitPrompt(deps: SubmitPromptDeps) { const resumed = await requestGateway<{ session_id: string }>('session.resume', { session_id: targetStoredSessionId, source: 'desktop', + omit_messages: true, ...(resumeProfile ? { profile: resumeProfile } : {}) }) @@ -637,6 +638,7 @@ export function useSubmitPrompt(deps: SubmitPromptDeps) { const resumed = await requestGateway<{ session_id: string }>('session.resume', { session_id: recoverStoredSessionId, source: 'desktop', + omit_messages: true, ...(resumeProfile ? { profile: resumeProfile } : {}) }) diff --git a/apps/desktop/src/app/session/hooks/use-session-actions.test.tsx b/apps/desktop/src/app/session/hooks/use-session-actions.test.tsx index 35110e481b..1bdb521efa 100644 --- a/apps/desktop/src/app/session/hooks/use-session-actions.test.tsx +++ b/apps/desktop/src/app/session/hooks/use-session-actions.test.tsx @@ -896,7 +896,7 @@ describe('resumeSession failure recovery', () => { expect(resumeParams).not.toHaveProperty('lazy') expect(resumeParams).not.toHaveProperty('eager_build') - expect(resumeParams).toMatchObject({ source: 'desktop' }) + expect(resumeParams).toMatchObject({ source: 'desktop', omit_messages: true }) }) it('arms the failure latch when resume succeeds with an empty transcript for a non-empty stored session', async () => { @@ -1431,6 +1431,10 @@ describe('resumeSession warm-cache mapping integrity', () => { expect(methods).toContain('session.activate') expect(methods).not.toContain('session.resume') expect(getSessionMessages).toHaveBeenCalledWith('stored-A', undefined) + expect(requestGateway).toHaveBeenCalledWith( + 'session.activate', + expect.objectContaining({ omit_messages: true, session_id: 'rt-A' }) + ) expect(runtimeIdByStoredSessionIdRef.current.get('stored-A')).toBe('rt-A') }) diff --git a/apps/desktop/src/app/session/hooks/use-session-actions/index.ts b/apps/desktop/src/app/session/hooks/use-session-actions/index.ts index 6a7e568cba..0f13a375c0 100644 --- a/apps/desktop/src/app/session/hooks/use-session-actions/index.ts +++ b/apps/desktop/src/app/session/hooks/use-session-actions/index.ts @@ -700,7 +700,8 @@ export function useSessionActions({ try { activated = await requestGateway('session.activate', { session_id: cachedRuntimeId, - cols: 96 + cols: 96, + omit_messages: true }) } catch (error) { // Compatibility for older backends. Modern backends require @@ -866,12 +867,14 @@ export function useSessionActions({ session_id: storedSessionId, cols: 96, source: 'desktop', + // REST is the transcript authority for Desktop. Avoid duplicating a + // potentially huge compression lineage in the WebSocket response. // Watch windows attach lazily (live mirror). Every other cold resume // gets the gateway's default deferred build: the RPC returns the // transcript immediately instead of blocking the switch on _make_agent // (MCP discovery / prompt build), and the agent pre-warms in the // background while the prefetch above paints the transcript. - ...(watchWindow ? { lazy: true } : {}), + ...(watchWindow ? { lazy: true } : { omit_messages: true }), ...(sessionProfile ? { profile: sessionProfile } : {}) }) diff --git a/apps/desktop/src/app/session/hooks/use-session-actions/resume-structural-parts.test.ts b/apps/desktop/src/app/session/hooks/use-session-actions/resume-structural-parts.test.ts index 1ef5376d8d..2dff424a9f 100644 --- a/apps/desktop/src/app/session/hooks/use-session-actions/resume-structural-parts.test.ts +++ b/apps/desktop/src/app/session/hooks/use-session-actions/resume-structural-parts.test.ts @@ -73,4 +73,70 @@ describe('reconcileResumeMessages — structural parts on a mid-turn switch', () expect(assistant.parts.filter(p => p.type === 'tool-call')).toHaveLength(1) }) + + it('keeps live-tail structure when the flat dump is not a strict text extension', () => { + // Mid-turn sandwich path: cache holds reasoning/tools; resume returns a + // longer non-extending dump. Structure source must be live-tail. + const cached: ChatMessage[] = [ + { + id: 'assistant-stream-1', + pending: true, + parts: [ + { type: 'reasoning', text: 'thinking about tools' }, + { type: 'tool-call', toolCallId: 'c1', toolName: 'terminal', args: {} }, + { type: 'text', text: 'partial' } + ], + role: 'assistant' + } + ] + + const authoritative: ChatMessage[] = [ + { + id: 'assistant-stream-1', + pending: true, + parts: [{ type: 'text', text: 'thinking about tools\nRan terminal\npartial and more dump' }], + role: 'assistant' + } + ] + + const [assistant] = reconcileResumeMessages(authoritative, cached) + + expect(assistant.parts.some(part => part.type === 'reasoning')).toBe(true) + expect(assistant.parts.some(part => part.type === 'tool-call')).toBe(true) + expect(assistant.parts.filter(part => part.type === 'text').map(part => ('text' in part ? part.text : ''))).toEqual( + ['partial'] + ) + }) + + it('does not graft historical structure onto a live text-only row after compression rewrote ordinals', () => { + // Previous cache still has a completed structured assistant at ordinal 0. + // Resume after compression returns a new live text-only assistant at the + // same role ordinal for an unrelated turn — must not inherit foreign parts. + const cached: ChatMessage[] = [ + { + id: 'old-assistant', + parts: [ + { type: 'reasoning', text: 'old thinking' }, + { type: 'tool-call', toolCallId: 'old-call', toolName: 'terminal', args: {} }, + { type: 'text', text: 'old answer' } + ], + role: 'assistant' + } + ] + + const authoritative: ChatMessage[] = [ + { + id: 'assistant-stream-runtime-1', + pending: true, + parts: [{ type: 'text', text: 'brand new partial' }], + role: 'assistant' + } + ] + + const [assistant] = reconcileResumeMessages(authoritative, cached) + + expect(assistant.parts.some(part => part.type === 'reasoning')).toBe(false) + expect(assistant.parts.some(part => part.type === 'tool-call')).toBe(false) + expect(assistant.parts).toEqual([{ type: 'text', text: 'brand new partial' }]) + }) }) diff --git a/apps/desktop/src/app/session/hooks/use-session-actions/utils.test.ts b/apps/desktop/src/app/session/hooks/use-session-actions/utils.test.ts index ea40e7705a..ed61eda70d 100644 --- a/apps/desktop/src/app/session/hooks/use-session-actions/utils.test.ts +++ b/apps/desktop/src/app/session/hooks/use-session-actions/utils.test.ts @@ -1,6 +1,7 @@ import { afterEach, beforeEach, describe, expect, it } from 'vitest' -import type { ChatMessage } from '@/lib/chat-messages' +import { textWithoutReferenceLines, WIRE_REFERENCE_KINDS } from '@/components/assistant-ui/reference-kinds' +import { type ChatMessage, type ChatMessagePart, chatMessageText } from '@/lib/chat-messages' import { $approvalModes, approvalModeForProfile } from '@/store/approval-mode' import { $desktopOnboarding } from '@/store/onboarding' import { $activeGatewayProfile } from '@/store/profile' @@ -24,6 +25,21 @@ import { const msg = (id: string, role: ChatMessage['role'], text: string, extra: Partial = {}): ChatMessage => ({ id, role, parts: [{ type: 'text', text }], ...extra }) as ChatMessage +// A live assistant row carrying the structure the gateway's text-only inflight +// snapshot cannot: reasoning and tool calls, with or without any text yet. +const streamingMsg = (id: string, text: string, extra: Partial = {}): ChatMessage => + ({ + id, + role: 'assistant', + parts: [ + { type: 'reasoning', text: 'planning' }, + { type: 'tool-call', toolCallId: 'call-1', toolName: 'terminal', result: 'done' }, + ...(text ? [{ type: 'text', text } as ChatMessagePart] : []) + ], + pending: true, + ...extra + }) as ChatMessage + const session = (over: Partial): SessionInfo => over as SessionInfo describe('applyRuntimeInfo approval mode', () => { @@ -394,6 +410,126 @@ describe('reconcileResumeMessages', () => { expect(out.attachmentRefs).toBeUndefined() }) + + // #75825: switching sessions mid-stream can re-hydrate an empty inflight shell + // at the same ordinal as the live stream row that still holds the full reply. + it('prefers a richer local pending assistant over an empty projection shell', () => { + const previous = [ + msg('1-user', 'user', 'question'), + msg('assistant-stream-live', 'assistant', 'hello from stream', { pending: true }) + ] + + const next = [msg('1-user', 'user', 'question'), msg('assistant-stream-sess', 'assistant', '', { pending: true })] + + const reconciled = reconcileResumeMessages(next, previous) + + expect(reconciled[1]).toMatchObject({ id: 'assistant-stream-live', pending: true }) + expect(chatMessageText(reconciled[1])).toBe('hello from stream') + }) + + it('prefers a richer local pending assistant when the projection lags mid-stream', () => { + const previous = [ + msg('1-user', 'user', 'question'), + msg('assistant-stream-live', 'assistant', 'hello world', { pending: true }) + ] + + const next = [ + msg('1-user', 'user', 'question'), + msg('assistant-stream-sess', 'assistant', 'hello', { pending: true }) + ] + + const reconciled = reconcileResumeMessages(next, previous) + + expect(chatMessageText(reconciled[1])).toBe('hello world') + expect(reconciled[1].id).toBe('assistant-stream-live') + }) + + it('does not override when the authoritative assistant has advanced further', () => { + const previous = [ + msg('1-user', 'user', 'question'), + msg('assistant-stream-live', 'assistant', 'hello', { pending: true }) + ] + + const next = [ + msg('1-user', 'user', 'question'), + msg('assistant-stream-sess', 'assistant', 'hello world', { pending: true }) + ] + + const reconciled = reconcileResumeMessages(next, previous) + + expect(chatMessageText(reconciled[1])).toBe('hello world') + expect(reconciled[1].id).toBe('assistant-stream-sess') + }) + + // The reported "no inference traces or tool calls": mid tool-work, the local + // row holds reasoning + tool calls and NO text yet, so both bodies are empty + // text and a text-length comparison cannot tell them apart. + it('prefers a traces-only local pending row over an empty shell', () => { + const previous = [msg('1-user', 'user', 'run the tools'), streamingMsg('assistant-stream-live', '')] + + const next = [ + msg('1-user', 'user', 'run the tools'), + msg('assistant-stream-sess', 'assistant', '', { pending: true }) + ] + + const reconciled = reconcileResumeMessages(next, previous) + + expect(reconciled[1].id).toBe('assistant-stream-live') + expect(reconciled[1].parts.map(part => part.type)).toEqual(['reasoning', 'tool-call']) + }) + + // A longer local body that is NOT an extension of the authoritative text is a + // different turn at the same ordinal (compression rewrites history) and must + // not hijack the slot. + it('leaves a shorter non-prefix authoritative assistant intact', () => { + const previous = [ + msg('1-user', 'user', 'question'), + msg('assistant-stream-live', 'assistant', 'a long local reply about something else entirely', { pending: true }) + ] + + const next = [msg('1-user', 'user', 'question'), msg('9-assistant', 'assistant', 'short authoritative answer')] + + const reconciled = reconcileResumeMessages(next, previous) + + expect(reconciled[1].id).toBe('9-assistant') + expect(chatMessageText(reconciled[1])).toBe('short authoritative answer') + }) + + // A retained failure snapshot (`inflight.error`) is projected with empty text. + // Preferring the local partial over it would erase the error and repaint the + // turn as healthy. + it('does not treat an errored authoritative row as an empty shell', () => { + const previous = [ + msg('1-user', 'user', 'do the thing'), + msg('assistant-stream-live', 'assistant', 'partial answer before the failure', { pending: true }) + ] + + const next = [ + msg('1-user', 'user', 'do the thing'), + msg('assistant-stream-sess', 'assistant', '', { error: 'model call failed: 500' }) + ] + + const reconciled = reconcileResumeMessages(next, previous) + + expect(reconciled[1].error).toBe('model call failed: 500') + }) + + // Content comes from the renderer; liveness stays the backend's call. A + // settled shell (queued turn behind a finished inflight one) must not leave + // the preserved reply spinning forever. + it('takes the local body but the authoritative settled state', () => { + const previous = [ + msg('1-user', 'user', 'question'), + msg('assistant-stream-live', 'assistant', 'streamed body', { pending: true }) + ] + + const next = [msg('1-user', 'user', 'question'), msg('assistant-stream-sess', 'assistant', '', { pending: false })] + + const reconciled = reconcileResumeMessages(next, previous) + + expect(reconciled[1]).toMatchObject({ id: 'assistant-stream-live', pending: false }) + expect(chatMessageText(reconciled[1])).toBe('streamed body') + }) }) describe('preserveLocalPendingTurnMessages', () => { @@ -556,7 +692,7 @@ describe('preserveLocalPendingTurnMessages', () => { // `attachmentRefs`. A naive text compare (chatMessageText a === b) therefore // always mismatched whenever an image was attached and re-appended the // optimistic row as a distinct, duplicate user bubble. Both sides must now - // reduce to the same visible text via textWithoutImageRefs. + // reduce to the same visible text via textWithoutReferenceLines. it('does not duplicate the optimistic image turn when the persisted turn carries @image refs', () => { const previous = [ msg('1-user', 'user', 'first'), @@ -575,6 +711,91 @@ describe('preserveLocalPendingTurnMessages', () => { expect(preserveLocalPendingTurnMessages(next, previous)).toBe(next) }) + it('does not duplicate the optimistic file turn when the persisted turn carries @file refs', () => { + const previous = [ + msg('1-user', 'user', 'first'), + msg('2-assistant', 'assistant', 'first answer'), + msg('user-optimistic', 'user', 'text', { + attachmentRefs: ['@file:X'] + }) + ] + + const next = [ + msg('1-user-stored', 'user', 'first'), + msg('2-assistant-stored', 'assistant', 'first answer'), + msg('3-user-stored', 'user', '@file:X\n\ntext') + ] + + expect(preserveLocalPendingTurnMessages(next, previous)).toBe(next) + }) + + it.each(WIRE_REFERENCE_KINDS.filter(kind => kind !== 'file' && kind !== 'image'))( + 'does not duplicate the optimistic %s turn when the persisted turn carries its directive', + kind => { + const ref = `@${kind}:X` + + const previous = [ + msg('1-user', 'user', 'first'), + msg('2-assistant', 'assistant', 'first answer'), + msg('user-optimistic', 'user', 'text', { + attachmentRefs: [ref] + }) + ] + + const next = [ + msg('1-user-stored', 'user', 'first'), + msg('2-assistant-stored', 'assistant', 'first answer'), + msg('3-user-stored', 'user', `${ref}\n\ntext`) + ] + + expect(preserveLocalPendingTurnMessages(next, previous)).toBe(next) + } + ) + + it('does not duplicate a directive-only file turn', () => { + const previous = [ + msg('1-user', 'user', 'first'), + msg('2-assistant', 'assistant', 'first answer'), + msg('user-optimistic', 'user', '', { + attachmentRefs: ['@file:X'] + }) + ] + + const next = [ + msg('1-user-stored', 'user', 'first'), + msg('2-assistant-stored', 'assistant', 'first answer'), + msg('3-user-stored', 'user', '@file:X') + ] + + expect(preserveLocalPendingTurnMessages(next, previous)).toBe(next) + }) + + it('does not duplicate a turn with multiple CRLF directives and Unicode payloads', () => { + const refs = ['@file:`資料/über notes.md`', '@url:`https://example.com/café?q=✓`'] + + const previous = [ + msg('1-user', 'user', 'first'), + msg('2-assistant', 'assistant', 'first answer'), + msg('user-optimistic', 'user', 'text', { + attachmentRefs: refs + }) + ] + + const next = [ + msg('1-user-stored', 'user', 'first'), + msg('2-assistant-stored', 'assistant', 'first answer'), + msg('3-user-stored', 'user', `${refs.join('\r\n')}\r\n\r\ntext`) + ] + + expect(preserveLocalPendingTurnMessages(next, previous)).toBe(next) + }) + + it('strips only complete reference lines from visible text', () => { + expect(textWithoutReferenceLines('see @file:X here')).toBe('see @file:X here') + expect(textWithoutReferenceLines('@file:X trailing prose')).toBe('@file:X trailing prose') + expect(textWithoutReferenceLines(' @file:X')).toBe('@file:X') + }) + it('still keeps a genuinely uncommitted optimistic image turn when the persisted text differs', () => { const previous = [ msg('1-user', 'user', 'first'), @@ -599,6 +820,177 @@ describe('preserveLocalPendingTurnMessages', () => { 'user-optimistic' ]) }) + + // #75825: an empty inflight projection shell at the same ordinal must not + // discard the local pending assistant that still holds the streamed content. + // Replace the shell (do not append) so the transcript shows one reply. + it('replaces an empty inflight shell with a fuller local pending assistant', () => { + const previous = [ + msg('1-user', 'user', 'question'), + msg('assistant-stream-live', 'assistant', 'partial answer so far', { pending: true }) + ] + + const next = [msg('1-user', 'user', 'question'), msg('assistant-stream-sess', 'assistant', '', { pending: true })] + + const preserved = preserveLocalPendingTurnMessages(next, previous) + + expect(preserved.map(message => message.id)).toEqual(['1-user', 'assistant-stream-live']) + expect(chatMessageText(preserved[1])).toBe('partial answer so far') + expect(preserved[1].pending).toBe(true) + }) + + it('replaces a lagging same-id shell with the fuller local pending body', () => { + const previous = [ + msg('1-user', 'user', 'question'), + msg('assistant-stream-sess', 'assistant', 'full streamed content', { pending: true }) + ] + + const next = [msg('1-user', 'user', 'question'), msg('assistant-stream-sess', 'assistant', '', { pending: true })] + + const preserved = preserveLocalPendingTurnMessages(next, previous) + + expect(preserved.map(message => message.id)).toEqual(['1-user', 'assistant-stream-sess']) + expect(chatMessageText(preserved[1])).toBe('full streamed content') + }) + + it('still drops local pending when authoritative text is at least as complete', () => { + const previous = [ + msg('1-user', 'user', 'question'), + msg('assistant-stream-live', 'assistant', 'partial', { pending: true }) + ] + + const next = [ + msg('1-user', 'user', 'question'), + msg('assistant-stream-sess', 'assistant', 'partial and more', { pending: true }) + ] + + expect(preserveLocalPendingTurnMessages(next, previous)).toBe(next) + }) + + // Mid tool-work both bodies are empty text, so only the parts distinguish the + // live row from the shell — the reported "no inference traces or tool calls". + it('replaces an empty shell with a traces-only local pending row', () => { + const previous = [msg('1-user', 'user', 'run the tools'), streamingMsg('assistant-stream-live', '')] + + const next = [ + msg('1-user', 'user', 'run the tools'), + msg('assistant-stream-sess', 'assistant', '', { pending: true }) + ] + + const preserved = preserveLocalPendingTurnMessages(next, previous) + + expect(preserved).toHaveLength(2) + expect(preserved[1].parts.map(part => part.type)).toEqual(['reasoning', 'tool-call']) + }) + + // Length alone is not identity: a longer local row that does not extend the + // authoritative text belongs to another turn and must not take its slot — by + // ordinal or by reusing the stream id. + it('leaves a shorter non-prefix authoritative assistant intact', () => { + const previous = [ + msg('1-user', 'user', 'question'), + msg('assistant-stream-live', 'assistant', 'a long local reply about something else entirely', { pending: true }) + ] + + const next = [msg('1-user', 'user', 'question'), msg('9-assistant', 'assistant', 'short authoritative answer')] + + const preserved = preserveLocalPendingTurnMessages(next, previous) + + expect(preserved.map(message => message.id)).toEqual(['1-user', '9-assistant']) + expect(chatMessageText(preserved[1])).toBe('short authoritative answer') + }) + + it('leaves a shorter non-prefix authoritative assistant intact on the same stream id', () => { + const previous = [ + msg('1-user', 'user', 'question'), + msg('assistant-stream-sess', 'assistant', 'a long local reply about something else entirely', { pending: true }) + ] + + const next = [ + msg('1-user', 'user', 'question'), + msg('assistant-stream-sess', 'assistant', 'short authoritative answer') + ] + + const preserved = preserveLocalPendingTurnMessages(next, previous) + + expect(chatMessageText(preserved[1])).toBe('short authoritative answer') + }) + + it('does not erase a retained failure with the local partial', () => { + const previous = [ + msg('1-user', 'user', 'do the thing'), + msg('assistant-stream-live', 'assistant', 'partial answer before the failure', { pending: true }) + ] + + const next = [ + msg('1-user', 'user', 'do the thing'), + msg('assistant-stream-sess', 'assistant', '', { error: 'model call failed: 500' }) + ] + + const assistant = preserveLocalPendingTurnMessages(next, previous).find(message => message.role === 'assistant') + + expect(assistant?.error).toBe('model call failed: 500') + expect(assistant?.pending).not.toBe(true) + }) + + it('takes the local body but the authoritative settled state', () => { + const previous = [ + msg('1-user', 'user', 'question'), + msg('assistant-stream-live', 'assistant', 'streamed body', { pending: true }) + ] + + const next = [msg('1-user', 'user', 'question'), msg('assistant-stream-sess', 'assistant', '', { pending: false })] + + const preserved = preserveLocalPendingTurnMessages(next, previous) + + expect(preserved[1]).toMatchObject({ id: 'assistant-stream-live', pending: false }) + expect(chatMessageText(preserved[1])).toBe('streamed body') + }) + + // #70209: history committed the reply under its own id, so the settled local + // stream row sits at a later ordinal, pairs with nothing, and gets appended — + // the same answer twice. + it('does not re-append a settled stream row the authoritative history already carries', () => { + const next = [msg('1-user-stored', 'user', 'question'), msg('2-assistant-stored', 'assistant', 'answer')] + const settledLocalStream = msg('assistant-stream-runtime-1', 'assistant', 'answer', { pending: false }) + + expect(preserveLocalPendingTurnMessages(next, [...next, settledLocalStream])).toBe(next) + }) + + // The reply finished locally but the gateway had not committed it when the + // session was reopened — the local row is the only copy and must survive. + it('keeps a settled stream row the authoritative history has not committed', () => { + const previous = [ + msg('1-user', 'user', 'question'), + msg('assistant-stream-sess', 'assistant', 'the finished reply', { pending: false }) + ] + + const next = [msg('1-user', 'user', 'question')] + + expect(preserveLocalPendingTurnMessages(next, previous).map(message => message.id)).toEqual([ + '1-user', + 'assistant-stream-sess' + ]) + }) + + // The whole point of replacing rather than appending: one reply on screen, + // and the committed history around the live turn untouched. + it('does not duplicate or rewrite committed history around the live turn', () => { + const history = [ + msg('1-user', 'user', 'first question'), + msg('2-assistant', 'assistant', 'first answer'), + msg('3-user', 'user', 'run the tools') + ] + + const previous = [...history, streamingMsg('assistant-stream-live', 'here is the full reply')] + const next = [...history, msg('assistant-stream-sess', 'assistant', '', { pending: true })] + + const preserved = preserveLocalPendingTurnMessages(next, previous) + + expect(preserved).toHaveLength(4) + expect(chatMessageText(preserved[1])).toBe('first answer') + expect(preserved.filter(message => message.role === 'assistant')).toHaveLength(2) + }) }) describe('appendLiveSessionProjection', () => { @@ -730,4 +1122,73 @@ describe('appendLiveSessionProjection', () => { expect(appendLiveSessionProjection(stored, { session_id: 'runtime-1' })).toBe(stored) }) + + it('does not sandwich a structured mid-turn row with the inflight flat dump (#76444)', () => { + const stored: ChatMessage[] = [ + msg('stored-user', 'user', 'do the work'), + { + id: 'live-assistant', + role: 'assistant', + pending: true, + parts: [ + { type: 'reasoning', text: 'thinking about tools' }, + { type: 'tool-call', toolCallId: 'c1', toolName: 'terminal', args: {} }, + { type: 'text', text: 'partial' } + ] + } + ] + + const restored = appendLiveSessionProjection(stored, { + session_id: 'runtime-1', + inflight: { + user: 'do the work', + // Flat dump includes thinking chatter + tool narration — longer than + // the answer text alone, which is how the sandwich used to grow. + assistant: 'thinking about tools\nRan terminal\npartial and more dump', + streaming: true + } + }) + + const assistants = restored.filter(message => message.role === 'assistant') + expect(assistants).toHaveLength(1) + expect(assistants[0].id).toBe('live-assistant') + expect(assistants[0].parts.some(part => part.type === 'reasoning')).toBe(true) + expect(assistants[0].parts.some(part => part.type === 'tool-call')).toBe(true) + // Answer text stays the structured row's text, not the dump. + expect( + assistants[0].parts.filter(part => part.type === 'text').map(part => ('text' in part ? part.text : '')) + ).toEqual(['partial']) + }) + + it('still projects inflight when only a completed historical tool reply has structure', () => { + // Older completed assistants keep reasoning/tool parts in the full + // transcript; they must not suppress a new turn's text projection. + const stored: ChatMessage[] = [ + msg('old-user', 'user', 'previous task'), + { + id: 'old-assistant', + role: 'assistant', + parts: [ + { type: 'tool-call', toolCallId: 'old', toolName: 'terminal', args: {} }, + { type: 'text', text: 'done earlier' } + ] + }, + msg('new-user', 'user', 'new task') + ] + + const restored = appendLiveSessionProjection(stored, { + session_id: 'runtime-1', + inflight: { + user: 'new task', + assistant: 'working on it', + streaming: true + } + }) + + expect(restored.map(message => message.id)).toContain('assistant-stream-runtime-1') + expect(restored.at(-1)).toMatchObject({ + id: 'assistant-stream-runtime-1', + pending: true + }) + }) }) diff --git a/apps/desktop/src/app/session/hooks/use-session-actions/utils.ts b/apps/desktop/src/app/session/hooks/use-session-actions/utils.ts index 1f86fd7c40..b6bc4ae6d5 100644 --- a/apps/desktop/src/app/session/hooks/use-session-actions/utils.ts +++ b/apps/desktop/src/app/session/hooks/use-session-actions/utils.ts @@ -1,7 +1,8 @@ +import { textWithoutReferenceLines } from '@/components/assistant-ui/reference-kinds' import { getSession } from '@/hermes' import { assistantTextPart, type ChatMessage, chatMessageText, textPart } from '@/lib/chat-messages' import { normalizePersonalityValue } from '@/lib/chat-runtime' -import { embeddedImageUrls, textWithoutEmbeddedImages, textWithoutImageRefs } from '@/lib/embedded-images' +import { embeddedImageUrls, textWithoutEmbeddedImages } from '@/lib/embedded-images' import { reconcileApprovalModeForProfile } from '@/store/approval-mode' import { requestDesktopOnboardingForCredentialWarning } from '@/store/onboarding' import { $activeGatewayProfile, $profiles, normalizeProfileKey } from '@/store/profile' @@ -46,6 +47,41 @@ function withAppendedText(message: ChatMessage, suffix: string): ChatMessage { return appended ? { ...message, parts } : message } +/** Reasoning / tool-call parts that the gateway inflight dump cannot express. */ +function hasStructuralParts(message: ChatMessage): boolean { + return message.parts.some(part => part.type === 'reasoning' || part.type === 'tool-call') +} + +/** + * A live-turn row — the gateway's text-only `inflight` projection, a + * still-streaming local bubble, or an interim row sealed inside the running + * turn — as opposed to a committed transcript row. + */ +function isLiveTailRow(message: ChatMessage): boolean { + return ( + message.pending === true || + message.id.startsWith('assistant-stream-') || + message.id.startsWith('inflight-assistant-') || + message.interim === true + ) +} + +/** + * True when `next` is a pure forward extension of the previous *answer* text. + * Empty previous answer never accepts a dump as an extension — that is how the + * mid-turn inflight flat dump used to sandwich structured rows (#76444). + */ +export function isStrictAnswerTextExtension(next: string, previous: string): boolean { + const n = next.trim() + const p = previous.trim() + + if (!p || !n) { + return false + } + + return n.startsWith(p) +} + /** * Carry structural parts an authoritative row cannot express. * @@ -250,8 +286,20 @@ export function reconcileResumeMessages(nextMessages: ChatMessage[], previousMes const nextText = chatMessageText(message).trim() const previousText = chatMessageText(previous) const previousVisibleText = textWithoutEmbeddedImages(previousText) + const previousTrimmed = previousVisibleText.trim() let preserved = message + // #75825: resume can project an empty (or lagging) inflight assistant shell + // at the same role-ordinal as the live stream row that still holds the + // streamed text, reasoning and tool calls. Prefer that richer pending row + // instead of painting the shell — otherwise the reply vanishes until + // restart. Guarded to the same reply further along (see + // localPendingSupersedes) so a different turn at the same ordinal cannot + // hijack the slot. + if (localPendingSupersedes(previous, message)) { + return withAuthoritativeTurnState(previous, message) + } + const sameText = nextText === previousVisibleText || nextText === previousText.trim() // Mid-turn, the authoritative text has advanced past the cached copy by one @@ -260,12 +308,36 @@ export function reconcileResumeMessages(nextMessages: ChatMessage[], previousMes // for structural carry-over. Attachment refs and image re-appending stay on // the strict equality path — they reconcile a SETTLED row, and a growing // row is by definition not settled. + // + // Live-tail identity: structure-only same-turn carry is allowed only when + // the *structure-bearing cached row* is still the in-flight stream + // (pending / stream id / interim). Marking only the text-only next row + // live is not enough — after compression a new live assistant can share a + // role ordinal with an unrelated historical structured row and must not + // inherit its reasoning/tool parts (#76444 review / salvage). const sameTurn = sameText || - (nextText.length > 0 && previousVisibleText.length > 0 && nextText.startsWith(previousVisibleText.trim())) + (nextText.length > 0 && previousTrimmed.length > 0 && isStrictAnswerTextExtension(nextText, previousTrimmed)) || + (message.role === 'assistant' && + previous.role === 'assistant' && + hasStructuralParts(previous) && + !hasStructuralParts(message) && + isLiveTailRow(previous)) if (sameTurn) { preserved = preserveStructuralParts(preserved, previous) + + // Never replace structured answer text with a non-extending flat dump. + if ( + message.role === 'assistant' && + hasStructuralParts(previous) && + !hasStructuralParts(message) && + !isStrictAnswerTextExtension(nextText, previousVisibleText) + ) { + const nonText = preserved.parts.filter(part => part.type !== 'text') + const priorAnswer = previous.parts.filter(part => part.type === 'text') + preserved = { ...preserved, parts: [...nonText, ...priorAnswer] } + } } if ( @@ -315,8 +387,10 @@ export function reconcileResumeMessages(nextMessages: ChatMessage[], previousMes * history window. Preserve only the newest optimistic user row: compression * rewrites past context, so older `user-*` rows in a warm cache are stale * history, not in-flight work. The latest authoritative user confirms whether - * that tail has persisted; any authoritative assistant at the same ordinal - * supersedes the local stream. + * that tail has persisted. An authoritative assistant at the same ordinal + * supersedes the local stream only when it is at least as complete; an empty + * or lagging inflight shell must not discard a fuller local pending reply + * (#75825). * * Gateway bookkeeping markers (the model-switch / personality notices written * by tui_gateway/server.py) are persisted as role=user but are not user turns. @@ -328,6 +402,64 @@ export function reconcileResumeMessages(nextMessages: ChatMessage[], previousMes const isGatewaySystemMarker = (message: ChatMessage): boolean => message.role === 'user' && chatMessageText(message).trimStart().startsWith('[System:') +/** + * Does the row carry anything a viewer would miss — streamed answer text, or + * the reasoning / tool-call structure the gateway's flat dump cannot express? + * An empty inflight shell carries none of it. + */ +const hasStreamedContent = (message: ChatMessage): boolean => + chatMessageText(message).trim().length > 0 || hasStructuralParts(message) + +/** + * May the cached local row stand in for this authoritative assistant? + * + * Only for a live projection of the SAME reply that the local copy is further + * along on: an empty shell, or text the local row strictly extends. Comparing + * lengths alone lets an unrelated (merely longer) local row hijack the ordinal + * — or the stream id — of a genuine stored reply. A retained failure snapshot + * (`inflight.error`, projected with empty text) is never a shell: repainting it + * from the local partial would hide the error and mark the turn healthy again. + */ +const localPendingSupersedes = (local: ChatMessage, authoritative: ChatMessage): boolean => { + if (local.role !== 'assistant' || !isLiveTailRow(local)) { + return false + } + + if (!isLiveTailRow(authoritative) || authoritative.error) { + return false + } + + const authoritativeText = chatMessageText(authoritative).trim() + + if (!authoritativeText.length) { + return hasStreamedContent(local) + } + + const localText = chatMessageText(local).trim() + + return localText.length > authoritativeText.length && isStrictAnswerTextExtension(localText, authoritativeText) +} + +/** + * Take the cached row's content, but never its liveness. The renderer holds the + * only copy of the streamed parts; the gateway remains the authority on whether + * the turn is still running and on durable row identity — so a settled shell + * must not repaint the reply as perpetually streaming. + */ +const withAuthoritativeTurnState = (local: ChatMessage, authoritative: ChatMessage): ChatMessage => { + const merged: ChatMessage = { ...local, pending: authoritative.pending === true } + + if (local.rowId === undefined && authoritative.rowId !== undefined) { + merged.rowId = authoritative.rowId + } + + if (local.reactions === undefined && authoritative.reactions?.length) { + merged.reactions = [...authoritative.reactions] + } + + return merged +} + export function preserveLocalPendingTurnMessages( nextMessages: ChatMessage[], previousMessages: ChatMessage[] @@ -378,6 +510,9 @@ export function preserveLocalPendingTurnMessages( const latestAuthoritativeUser = [...nextMessages].reverse().find(message => message.role === 'user') const preserved: ChatMessage[] = [] + // Authoritative id → richer local pending row. Replacing (not appending) + // avoids painting both the empty inflight shell and the full stream bubble. + const replacements = new Map() for (const message of previousMessages) { if (isGatewaySystemMarker(message)) { @@ -392,7 +527,21 @@ export function preserveLocalPendingTurnMessages( const isPendingAssistant = message.role === 'assistant' && (message.pending === true || message.id.startsWith('assistant-stream-')) - if ((!isOptimisticUser && !isPendingAssistant) || nextIds.has(message.id)) { + if (!isOptimisticUser && !isPendingAssistant) { + continue + } + + // Same id already present: still prefer a strictly more complete local + // pending body over an empty/stale shell that reused the stream id. + if (nextIds.has(message.id)) { + if (isPendingAssistant) { + const existing = nextMessages.find(candidate => candidate.id === message.id) + + if (existing && localPendingSupersedes(message, existing)) { + replacements.set(message.id, withAuthoritativeTurnState(message, existing)) + } + } + continue } @@ -403,19 +552,50 @@ export function preserveLocalPendingTurnMessages( if ( isOptimisticUser && latestAuthoritativeUser && - textWithoutImageRefs(chatMessageText(latestAuthoritativeUser)) === textWithoutImageRefs(chatMessageText(message)) + textWithoutReferenceLines(chatMessageText(latestAuthoritativeUser)) === + textWithoutReferenceLines(chatMessageText(message)) ) { continue } const authoritative = nextByRoleOrdinal.get(`${message.role}:${ordinal}`) + // A settled stream row (`pending: false` after message.complete) whose reply + // the authoritative transcript already carries under its committed id is + // stale: ordinal pairing can't see it, because the commit shifted the row + // one ordinal earlier, and re-appending it renders the same answer twice + // (#70209). Only text-identical rows are dropped — a settled row the backend + // has NOT committed yet is the only copy of that reply and must survive. + if ( + isPendingAssistant && + message.pending !== true && + nextMessages.some( + candidate => + candidate.role === 'assistant' && + textWithoutReferenceLines(chatMessageText(candidate)) === textWithoutReferenceLines(chatMessageText(message)) + ) + ) { + continue + } + if (authoritative) { if (isPendingAssistant) { + // Keep the local pending row when it is the same reply further along + // and the authoritative row is an empty projection shell or a prefix. + // #75825 + if (!localPendingSupersedes(message, authoritative)) { + continue + } + + replacements.set(authoritative.id, withAuthoritativeTurnState(message, authoritative)) + continue } - if (textWithoutImageRefs(chatMessageText(authoritative)) === textWithoutImageRefs(chatMessageText(message))) { + if ( + textWithoutReferenceLines(chatMessageText(authoritative)) === + textWithoutReferenceLines(chatMessageText(message)) + ) { continue } } @@ -423,7 +603,10 @@ export function preserveLocalPendingTurnMessages( preserved.push(message) } - return preserved.length ? [...nextMessages, ...preserved] : nextMessages + const withReplacements = + replacements.size > 0 ? nextMessages.map(message => replacements.get(message.id) ?? message) : nextMessages + + return preserved.length ? [...withReplacements, ...preserved] : withReplacements } /** @@ -484,7 +667,9 @@ export function appendLiveSessionProjection( } const persistedInLatestRun = (text: string): boolean => - latestUserRun.some(message => textWithoutImageRefs(chatMessageText(message)) === textWithoutImageRefs(text)) + latestUserRun.some( + message => textWithoutReferenceLines(chatMessageText(message)) === textWithoutReferenceLines(text) + ) const inflightUserAlreadyPersisted = Boolean(inflightUser) && persistedInLatestRun(inflightUser) @@ -514,14 +699,54 @@ export function appendLiveSessionProjection( // Keep a pending assistant boundary even before the first delta when a // queued user turn follows it. This preserves the two distinct turns. + // + // When the *current live turn* already holds a structured mid-turn assistant + // row (reasoning / tool-call from the live stream or journal), do NOT append + // a pure-text projection of `inflight.assistant` — that flat dump re-renders + // thinking as answer text and sandwiches the structured parts (#76444). + // Only inspect the live tail after the latest user run — never a completed + // historical tool-bearing reply earlier in the transcript (review feedback). + const liveStreamId = `assistant-stream-${sessionId}` + + const liveAssistantOfCurrentTurn = ((): ChatMessage | null => { + const byStreamId = messages.find(message => message.id === liveStreamId) + + if (byStreamId) { + return byStreamId + } + + // Assistants after the latest user row belong to this turn's tail. + if (latestUserIndex < 0) { + return null + } + + for (let index = messages.length - 1; index > latestUserIndex; index -= 1) { + if (messages[index].role === 'assistant') { + return messages[index] + } + } + + return null + })() + + const turnAlreadyStructured = Boolean( + liveAssistantOfCurrentTurn && + hasStructuralParts(liveAssistantOfCurrentTurn) && + isLiveTailRow(liveAssistantOfCurrentTurn) + ) + if (inflightAssistant || inflightStreaming || inflightError || (inflightUser && queuedUser)) { - projected.push({ - id: `assistant-stream-${sessionId}`, - role: 'assistant', - parts: inflightAssistant ? [assistantTextPart(inflightAssistant)] : [], - pending: inflightStreaming, - ...(inflightError ? { error: inflightError } : {}) - }) + if (turnAlreadyStructured && !inflightError) { + // Structure is authoritative; skip the text-only dump row. + } else { + projected.push({ + id: liveStreamId, + role: 'assistant', + parts: inflightAssistant ? [assistantTextPart(inflightAssistant)] : [], + pending: inflightStreaming, + ...(inflightError ? { error: inflightError } : {}) + }) + } } if (queuedUser) { diff --git a/apps/desktop/src/components/assistant-ui/reference-kinds.ts b/apps/desktop/src/components/assistant-ui/reference-kinds.ts index 453d05411d..7d6c397584 100644 --- a/apps/desktop/src/components/assistant-ui/reference-kinds.ts +++ b/apps/desktop/src/components/assistant-ui/reference-kinds.ts @@ -157,3 +157,17 @@ const REFERENCE_PATTERN = /@(file|folder|url|image|tool|line|terminal|session):( export function referenceRe(): RegExp { return new RegExp(REFERENCE_PATTERN.source, 'g') } + +/** Remove reference-only lines when comparing visible message text. */ +// Anchored + non-global: no shared `lastIndex` state (the hazard referenceRe() +// exists to avoid), and hoisting skips a RegExp construction per call — this +// runs on both sides of every message comparison in the reconcile loops. +const REFERENCE_LINE_RE = new RegExp(`^(?:${REFERENCE_PATTERN.source})$`) + +export function textWithoutReferenceLines(text: string): string { + return text + .split('\n') + .filter(line => !REFERENCE_LINE_RE.test(line.trimEnd())) + .join('\n') + .trim() +} diff --git a/apps/desktop/src/components/assistant-ui/thread/list.test.ts b/apps/desktop/src/components/assistant-ui/thread/list.test.ts index 413aab8865..c0cd1da558 100644 --- a/apps/desktop/src/components/assistant-ui/thread/list.test.ts +++ b/apps/desktop/src/components/assistant-ui/thread/list.test.ts @@ -8,7 +8,8 @@ import { liveTailStart, type MessageGroup, messageRenderWeight, - RENDER_WEIGHT_CHARS + RENDER_WEIGHT_CHARS, + resolveThreadScrollTarget } from './list' // Signature rows are `${index}:${id}:${role}:${weight}` (see the useAuiState @@ -62,6 +63,53 @@ describe('buildGroups', () => { }) }) +describe('resolveThreadScrollTarget', () => { + const context = (scrollElement: Pick) => ({ + contentElement: document.createElement('div'), + scrollElement: scrollElement as HTMLElement + }) + + it('settles when the browser clamps the requested bottom within half a CSS pixel', () => { + let actualScrollTop = 0 + let writes = 0 + + const scrollElement = { + get scrollTop() { + return actualScrollTop + }, + set scrollTop(value: number) { + writes += 1 + actualScrollTop = value - 0.125 + } + } + + const target = 899 + + const requested = resolveThreadScrollTarget(target, context(scrollElement)) + scrollElement.scrollTop = requested + const settled = resolveThreadScrollTarget(target, context(scrollElement)) + + expect(requested).toBe(target) + expect(actualScrollTop).toBe(898.875) + expect(settled).toBe(actualScrollTop) + expect(actualScrollTop < settled).toBe(false) + expect(writes).toBe(1) + }) + + it('keeps following while more than half a CSS pixel remains', () => { + const scrollElement = { scrollTop: 898.25 } + + expect(resolveThreadScrollTarget(899, context(scrollElement))).toBe(899) + }) + + it('re-arms after streaming content increases the target', () => { + const scrollElement = { scrollTop: 898.875 } + + expect(resolveThreadScrollTarget(899, context(scrollElement))).toBe(898.875) + expect(resolveThreadScrollTarget(999, context(scrollElement))).toBe(999) + }) +}) + describe('firstVisibleGroupIndex', () => { const group = (id: string, weight: number): MessageGroup => ({ id, index: 0, kind: 'standalone', weight }) diff --git a/apps/desktop/src/components/assistant-ui/thread/list.tsx b/apps/desktop/src/components/assistant-ui/thread/list.tsx index faa0724883..d8fd0f0a5c 100644 --- a/apps/desktop/src/components/assistant-ui/thread/list.tsx +++ b/apps/desktop/src/components/assistant-ui/thread/list.tsx @@ -13,7 +13,7 @@ import { useRef, useState } from 'react' -import { useStickToBottom } from 'use-stick-to-bottom' +import { type GetTargetScrollTop, useStickToBottom } from 'use-stick-to-bottom' import { useI18n } from '@/i18n' import { cn } from '@/lib/utils' @@ -62,6 +62,20 @@ const MAX_MEASURED_MESSAGE_CHARS = RENDER_BUDGET * RENDER_WEIGHT_CHARS // blocks the click-to-paint path. const FIRST_PAINT_BUDGET = 20 +// Browsers may quantize a requested scrollTop to a nearby device-pixel +// boundary. use-stick-to-bottom otherwise compares the lower actual value to +// the integer target forever, re-requesting the same instant scroll every +// frame. Treat a subpixel remainder as achieved; larger gaps still follow new +// streamed content normally. +const SCROLL_TARGET_EPSILON_PX = 0.5 + +export const resolveThreadScrollTarget: GetTargetScrollTop = (targetScrollTop, { scrollElement }) => { + const currentScrollTop = scrollElement.scrollTop + const remaining = targetScrollTop - currentScrollTop + + return remaining >= 0 && remaining <= SCROLL_TARGET_EPSILON_PX ? currentScrollTop : targetScrollTop +} + const contentWeightCache = new WeakMap() const NON_RENDERED_CONTENT_FIELDS = new Set(['id', 'role', 'toolCallId', 'toolName', 'type']) @@ -297,7 +311,8 @@ const ThreadMessageListInner: FC = ({ // settling. Its refs hang off our own DOM so the sticky human bubbles survive. const { scrollRef, contentRef, isAtBottom, scrollToBottom, stopScroll } = useStickToBottom({ initial: 'instant', - resize: 'instant' + resize: 'instant', + targetScrollTop: resolveThreadScrollTarget }) const [renderBudget, setRenderBudget] = useState(FIRST_PAINT_BUDGET) diff --git a/apps/desktop/src/components/assistant-ui/thread/status.tsx b/apps/desktop/src/components/assistant-ui/thread/status.tsx index 83cd107a29..492a73b1f1 100644 --- a/apps/desktop/src/components/assistant-ui/thread/status.tsx +++ b/apps/desktop/src/components/assistant-ui/thread/status.tsx @@ -9,6 +9,7 @@ import { ActivityTimerText } from '@/components/chat/activity-timer-text' import { SCAFFOLD_LABEL_CLASS } from '@/components/chat/scaffold-row' import { Codicon } from '@/components/ui/codicon' import { Loader } from '@/components/ui/loader' +import { StatusPulse } from '@/components/ui/status-pulse' import { useI18n } from '@/i18n' import { cn } from '@/lib/utils' import { $backgroundResume } from '@/store/background-delegation' @@ -130,7 +131,11 @@ export const ResponseLoadingIndicator: FC = () => { return ( - @@ -236,7 +241,11 @@ export const StreamStallIndicator: FC = () => { return ( - diff --git a/apps/desktop/src/components/pane-shell/tree/pane-reload.test.ts b/apps/desktop/src/components/pane-shell/tree/pane-reload.test.ts new file mode 100644 index 0000000000..8c78d5d391 --- /dev/null +++ b/apps/desktop/src/components/pane-shell/tree/pane-reload.test.ts @@ -0,0 +1,28 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest' + +// Right-click a tab -> Reload remounts THAT pane's content: its epoch (the +// React key the zone renderer hands the contribution) advances, and no other +// pane's does. The layout tree itself must never move. + +describe('reloadTreePane', () => { + beforeEach(() => { + window.localStorage.clear() + vi.resetModules() + }) + + it('advances only the reloaded pane epoch and leaves the tree alone', async () => { + const tree = await import('@/components/pane-shell/tree/store') + const model = await import('@/components/pane-shell/tree/model') + + tree.declareDefaultTree(model.group(['workspace', 'files'], { active: 'workspace', id: 'grp-main' })) + + const before = tree.$layoutTree.get() + + tree.reloadTreePane('workspace') + tree.reloadTreePane('workspace') + + expect(tree.$treePaneEpochs.get().workspace).toBe(2) + expect(tree.$treePaneEpochs.get().files).toBeUndefined() + expect(tree.$layoutTree.get()).toBe(before) + }) +}) diff --git a/apps/desktop/src/components/pane-shell/tree/renderer/tree-group.tsx b/apps/desktop/src/components/pane-shell/tree/renderer/tree-group.tsx index dcf5c05840..4f7d7ada11 100644 --- a/apps/desktop/src/components/pane-shell/tree/renderer/tree-group.tsx +++ b/apps/desktop/src/components/pane-shell/tree/renderer/tree-group.tsx @@ -33,6 +33,7 @@ import { $newSessionTabAction, $panesWithCloser, $treeDragging, + $treePaneEpochs, activateTreePane, closeAllTreeTabs, closeOtherTreeTabs, @@ -42,6 +43,7 @@ import { isCollapsePane, isSessionStripPane, noteActiveTreeGroup, + reloadTreePane, restoreTreePane, SESSION_TILE_DRAG, setTreeGroupHeaderHidden, @@ -96,6 +98,12 @@ function ZoneMenu({ return ( <> + {renderActionItem(kit, { + icon: 'refresh', + label: t.zones.reload, + onSelect: () => reloadTreePane(targetPane()) + })} + {paneId !== undefined && renderActionItem(kit, { icon: 'close', @@ -178,6 +186,9 @@ export function TreeGroup({ const narrow = useStore($narrowViewport) const newSessionTabAction = useStore($newSessionTabAction) const panesWithCloser = useStore($panesWithCloser) + // Reload epochs: only an explicit tab-menu Reload writes here, so this + // subscription costs nothing on a normal render. + const paneEpochs = useStore($treePaneEpochs) const paneFor = (id: string) => panes.find(p => p.id === id) @@ -557,9 +568,14 @@ export function TreeGroup({ // can gate its hot (per-token) subscriptions while hidden; // the group id identifies the ZONE it lives in, for state // that is per-zone rather than per-tab (composer pop-out). + // The reload epoch keys the CONTENT, not this layer: a + // Reload remounts the contribution (effects re-run, state + // resets) while the layer — and every other tab — stays. - {pane.render()} + + {pane.render()} + ) : ( diff --git a/apps/desktop/src/components/pane-shell/tree/store.ts b/apps/desktop/src/components/pane-shell/tree/store.ts index 42267459b9..8399eb36ef 100644 --- a/apps/desktop/src/components/pane-shell/tree/store.ts +++ b/apps/desktop/src/components/pane-shell/tree/store.ts @@ -450,6 +450,22 @@ export function treeTabCloseTargets(paneId: string): { all: number; others: numb return { all: others.length + (isUncloseablePane(paneId) ? 0 : 1), others: others.length, right: right.length } } +/** + * RELOAD — a pane's remount counter, the tab menu's Reload (browser parity: + * right-click a tab, reload what's in it). The zone renderer keys a pane's + * body layer on its epoch, so bumping it unmounts the contribution and mounts + * it fresh — data effects re-run, measurements are retaken — while the layout + * tree, the tab's position, and every other tab stay exactly as they were. + * Absent until a pane is first reloaded (no key churn on a normal boot). + */ +export const $treePaneEpochs = atom>>({}) + +export function reloadTreePane(paneId: string): void { + const epochs = $treePaneEpochs.get() + + $treePaneEpochs.set({ ...epochs, [paneId]: (epochs[paneId] ?? 0) + 1 }) +} + /** Close a tab the way its kind expects: a tool panel leaves the strip (and * syncs its toggle), everything else routes through its owning Close. */ export function closeTabPane(paneId: string) { @@ -1577,7 +1593,7 @@ export function resetLayoutTree() { } // Dev hook for automation. -if (import.meta.env.DEV && typeof window !== 'undefined') { +if ((import.meta.env.DEV || import.meta.env.VITE_PERF_PROBE === '1') && typeof window !== 'undefined') { ;(window as unknown as Record).__HERMES_LAYOUT_TREE__ = { close: closeTreePane, dismissed: () => $dismissedPanes.get(), diff --git a/apps/desktop/src/components/pet/floating-pet.tsx b/apps/desktop/src/components/pet/floating-pet.tsx index 4175427150..62acd28b4e 100644 --- a/apps/desktop/src/components/pet/floating-pet.tsx +++ b/apps/desktop/src/components/pet/floating-pet.tsx @@ -13,7 +13,10 @@ import { $petRoam, $petRoamDir, clearPetUnread, + hasPetSpriteForMeta, + mergePetInfoMeta, type PetInfo, + type PetInfoMeta, petProfile, setPetInfo } from '@/store/pet' @@ -39,25 +42,6 @@ interface Point { y: number } -interface PetInfoMeta { - enabled: boolean - slug?: string - displayName?: string - scale?: number - spritesheetRevision?: string -} - -function samePetRevision(info: PetInfo, meta: PetInfoMeta): boolean { - return ( - info.enabled && - Boolean(info.spritesheetBase64) && - info.slug === meta.slug && - info.displayName === meta.displayName && - info.scale === meta.scale && - info.spritesheetRevision === meta.spritesheetRevision - ) -} - // Keep a w×h box fully inside the viewport. Pre-pet-load callers pass a nominal // size; the live size flows in once `info` arrives. function clampPoint(x: number, y: number, w: number, h: number): Point { @@ -161,7 +145,7 @@ export function FloatingPet() { // pet.changed already carries the meta payload — an enabled=false // broadcast clears the mascot with zero round-trips, and an unchanged // revision (scale-only move still changes the sig) short-circuits below - // via samePetRevision. + // via hasPetSpriteForMeta + mergePetInfoMeta. if (changeEventsAvailable && petChange.tick > 0 && petChange.meta?.enabled === false) { setPetInfo({ enabled: false }) @@ -184,7 +168,15 @@ export function FloatingPet() { return } - if (samePetRevision($petInfo.get(), meta)) { + const current = $petInfo.get() + + if (hasPetSpriteForMeta(current, meta)) { + const merged = mergePetInfoMeta(current, meta) + + if (merged !== current) { + setPetInfo(merged) + } + return } } catch { @@ -223,7 +215,14 @@ export function FloatingPet() { // so no timer. Legacy backend: the historical poll. const timer = changeEventsAvailable ? null - : window.setInterval(() => void pull(), active ? PET_ACTIVE_REFRESH_MS : PET_POLL_MS) + : window.setInterval( + () => { + if (document.visibilityState === 'visible') { + void pull() + } + }, + active ? PET_ACTIVE_REFRESH_MS : PET_POLL_MS + ) return () => { cancelled = true diff --git a/apps/desktop/src/components/ui/glyph-spinner.test.tsx b/apps/desktop/src/components/ui/glyph-spinner.test.tsx new file mode 100644 index 0000000000..686ee91da1 --- /dev/null +++ b/apps/desktop/src/components/ui/glyph-spinner.test.tsx @@ -0,0 +1,167 @@ +import { act, render, screen } from '@testing-library/react' +import { Profiler, type ProfilerOnRenderCallback } from 'react' +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' + +import { PaneVisibleContext } from '@/components/pane-shell/pane-visibility' + +import { GlyphSpinner } from './glyph-spinner' + +describe('GlyphSpinner', () => { + beforeEach(() => { + vi.useFakeTimers() + vi.spyOn(globalThis.document, 'hasFocus').mockReturnValue(true) + }) + + afterEach(() => { + vi.clearAllTimers() + vi.restoreAllMocks() + vi.useRealTimers() + }) + + it('advances its glyph without an update-phase React commit', () => { + let updateCommits = 0 + + const onRender: ProfilerOnRenderCallback = (_id, phase) => { + if (phase !== 'mount') { + updateCommits += 1 + } + } + + render( + + + + ) + + const status = screen.getByRole('status', { name: 'Loading' }) + expect(status.textContent).toBe('⠋') + + act(() => vi.advanceTimersByTime(80)) + + expect(status.textContent).toBe('⠙') + expect(updateCommits).toBe(0) + }) + + it('does not tick while its kept-alive pane is hidden', () => { + const { rerender } = render( + + + + ) + + const status = screen.getByRole('status', { name: 'Loading' }) + + expect(status.textContent).toBe('⠋') + expect(vi.getTimerCount()).toBe(0) + + rerender( + + + + ) + expect(vi.getTimerCount()).toBe(1) + + act(() => vi.advanceTimersByTime(80)) + expect(status.textContent).toBe('⠙') + + rerender( + + + + ) + expect(vi.getTimerCount()).toBe(0) + + const frozen = status.textContent + act(() => vi.advanceTimersByTime(800)) + expect(status.textContent).toBe(frozen) + }) + + it('suspends animation while the Desktop window is inactive', () => { + render() + + const status = screen.getByRole('status', { name: 'Loading' }) + expect(vi.getTimerCount()).toBe(1) + + act(() => window.dispatchEvent(new Event('blur'))) + expect(vi.getTimerCount()).toBe(0) + + const frozen = status.textContent + act(() => vi.advanceTimersByTime(800)) + expect(status.textContent).toBe(frozen) + + act(() => window.dispatchEvent(new Event('focus'))) + expect(vi.getTimerCount()).toBe(1) + + act(() => vi.advanceTimersByTime(80)) + expect(status.textContent).not.toBe(frozen) + }) + + it('suspends animation while the Electron window is minimized or hidden, then resumes on restore', () => { + let windowStateCallback: ((payload: { isMinimized?: boolean; isVisible?: boolean }) => void) | null = null + + Object.defineProperty(window, 'hermesDesktop', { + configurable: true, + value: { + onWindowStateChanged: vi.fn((callback: typeof windowStateCallback) => { + windowStateCallback = callback + + return () => { + if (windowStateCallback === callback) { + windowStateCallback = null + } + } + }) + } + }) + + try { + render() + + const status = screen.getByRole('status', { name: 'Loading' }) + expect(windowStateCallback).not.toBeNull() + expect(vi.getTimerCount()).toBe(1) + + act(() => windowStateCallback?.({ isMinimized: true, isVisible: false })) + expect(vi.getTimerCount()).toBe(0) + + const frozen = status.textContent + act(() => vi.advanceTimersByTime(800)) + expect(status.textContent).toBe(frozen) + + act(() => windowStateCallback?.({ isMinimized: false, isVisible: true })) + expect(vi.getTimerCount()).toBe(1) + + act(() => vi.advanceTimersByTime(80)) + expect(status.textContent).not.toBe(frozen) + } finally { + delete (window as unknown as { hermesDesktop?: unknown }).hermesDesktop + } + }) + + it('suspends animation while the document is hidden', () => { + render() + + const status = screen.getByRole('status', { name: 'Loading' }) + expect(vi.getTimerCount()).toBe(1) + + Object.defineProperty(document, 'visibilityState', { configurable: true, value: 'hidden' }) + + try { + act(() => document.dispatchEvent(new Event('visibilitychange'))) + expect(vi.getTimerCount()).toBe(0) + + const frozen = status.textContent + act(() => vi.advanceTimersByTime(800)) + expect(status.textContent).toBe(frozen) + + Object.defineProperty(document, 'visibilityState', { configurable: true, value: 'visible' }) + act(() => document.dispatchEvent(new Event('visibilitychange'))) + expect(vi.getTimerCount()).toBe(1) + + act(() => vi.advanceTimersByTime(80)) + expect(status.textContent).not.toBe(frozen) + } finally { + Object.defineProperty(document, 'visibilityState', { configurable: true, value: 'visible' }) + } + }) +}) diff --git a/apps/desktop/src/components/ui/glyph-spinner.tsx b/apps/desktop/src/components/ui/glyph-spinner.tsx index 52e82412c8..1fc50bbe4b 100644 --- a/apps/desktop/src/components/ui/glyph-spinner.tsx +++ b/apps/desktop/src/components/ui/glyph-spinner.tsx @@ -1,7 +1,8 @@ -import { useEffect, useState } from 'react' +import { useEffect, useRef } from 'react' import spinners, { type BrailleSpinnerName as SpinnerName } from 'unicode-animations' import { usePaneVisible } from '@/components/pane-shell/pane-visibility' +import { createRendererLoopPauseController } from '@/lib/renderer-loop-pause' import { cn } from '@/lib/utils' export type { SpinnerName } @@ -43,29 +44,66 @@ interface GlyphSpinnerProps { */ export function GlyphSpinner({ ariaLabel = 'Loading', className, spinner = 'braille' }: GlyphSpinnerProps) { const spin = FRAMES_BY_NAME[spinner] ?? FRAMES_BY_NAME.braille! - const [frame, setFrame] = useState(0) + const glyphRef = useRef(null) // Pause when this surface is a hidden (kept-alive) tab: N mounted tabs each - // ticking a setInterval + setState burn CPU for pixels nobody can see. + // ticking a setInterval burns CPU for pixels nobody can see. const visible = usePaneVisible() useEffect(() => { - if (!visible) { + const glyph = glyphRef.current + + if (!visible || !glyph) { return } - setFrame(0) - const id = window.setInterval(() => setFrame(f => (f + 1) % spin.frames.length), spin.interval) + let frame = 0 + let timer: number | undefined + let pauseController: ReturnType | undefined + glyph.textContent = spin.frames[frame] - return () => window.clearInterval(id) + const stopAnimation = () => { + if (timer === undefined) { + return + } + + window.clearInterval(timer) + timer = undefined + } + + const syncAnimation = () => { + if (pauseController?.isPaused()) { + stopAnimation() + + return + } + + if (timer !== undefined) { + return + } + + timer = window.setInterval(() => { + frame = (frame + 1) % spin.frames.length + glyph.textContent = spin.frames[frame] + }, spin.interval) + } + + pauseController = createRendererLoopPauseController(syncAnimation) + syncAnimation() + + return () => { + pauseController.dispose() + stopAnimation() + } }, [spin, visible]) return ( - {spin.frames[frame]} + {spin.frames[0]} ) } diff --git a/apps/desktop/src/components/ui/status-pulse.test.tsx b/apps/desktop/src/components/ui/status-pulse.test.tsx new file mode 100644 index 0000000000..8f49622928 --- /dev/null +++ b/apps/desktop/src/components/ui/status-pulse.test.tsx @@ -0,0 +1,131 @@ +import { act, cleanup, render } from '@testing-library/react' +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' + +import { StatusPulse } from './status-pulse' + +interface PlayedAnimation { + cancel: ReturnType + keyframes: Keyframe[] + options: KeyframeAnimationOptions +} + +interface WindowStatePayload { + isMinimized?: boolean + isVisible?: boolean +} + +let windowStateCallback: ((payload: WindowStatePayload) => void) | undefined + +function installMatchMedia(matches: boolean) { + vi.stubGlobal( + 'matchMedia', + vi.fn(() => ({ + addEventListener: vi.fn(), + matches, + media: '(prefers-reduced-motion: reduce)', + onchange: null, + removeEventListener: vi.fn() + })) + ) +} + +function installWindowStateBridge() { + const off = vi.fn() + + window.hermesDesktop = { + onWindowStateChanged: vi.fn(callback => { + windowStateCallback = callback + + return off + }) + } as unknown as typeof window.hermesDesktop + + return off +} + +describe('StatusPulse', () => { + const played: PlayedAnimation[] = [] + + beforeEach(() => { + vi.useFakeTimers() + vi.spyOn(window.document, 'hasFocus').mockReturnValue(true) + installMatchMedia(false) + installWindowStateBridge() + Object.defineProperty(HTMLElement.prototype, 'animate', { + configurable: true, + value: (keyframes: Keyframe[], options: KeyframeAnimationOptions) => { + const animation = { + cancel: vi.fn(), + keyframes, + options + } + + played.push(animation) + + return animation as unknown as Animation + }, + writable: true + }) + }) + + afterEach(() => { + cleanup() + played.length = 0 + windowStateCallback = undefined + Reflect.deleteProperty(HTMLElement.prototype, 'animate') + delete (window as unknown as { hermesDesktop?: unknown }).hermesDesktop + vi.useRealTimers() + vi.unstubAllGlobals() + vi.restoreAllMocks() + }) + + it('plays a finite ping and sleeps between pulses', () => { + render() + + expect(played).toHaveLength(1) + expect(played[0]?.keyframes).toEqual([ + { opacity: 0.7, transform: 'scale(1)' }, + { opacity: 0, transform: 'scale(2)' } + ]) + expect(played[0]?.options).toMatchObject({ duration: 400, iterations: 1 }) + expect(vi.getTimerCount()).toBe(1) + + act(() => vi.advanceTimersByTime(4_999)) + expect(played).toHaveLength(1) + + act(() => vi.advanceTimersByTime(1)) + expect(played).toHaveLength(2) + }) + + it('cancels scheduled work while minimized and restarts once visible', () => { + const offWindowState = installWindowStateBridge() + const mounted = render() + + expect(played).toHaveLength(1) + + act(() => windowStateCallback?.({ isMinimized: true, isVisible: false })) + + expect(played[0]?.cancel).toHaveBeenCalledTimes(1) + expect(vi.getTimerCount()).toBe(0) + + act(() => vi.advanceTimersByTime(5_000)) + expect(played).toHaveLength(1) + + act(() => windowStateCallback?.({ isMinimized: false, isVisible: true })) + expect(played).toHaveLength(2) + + mounted.unmount() + expect(played[1]?.cancel).toHaveBeenCalledTimes(1) + expect(offWindowState).toHaveBeenCalledTimes(1) + expect(vi.getTimerCount()).toBe(0) + }) + + it('stays static when reduced motion is requested', () => { + installMatchMedia(true) + + render() + + expect(played).toHaveLength(0) + expect(vi.getTimerCount()).toBe(0) + }) +}) diff --git a/apps/desktop/src/components/ui/status-pulse.tsx b/apps/desktop/src/components/ui/status-pulse.tsx new file mode 100644 index 0000000000..e162263564 --- /dev/null +++ b/apps/desktop/src/components/ui/status-pulse.tsx @@ -0,0 +1,142 @@ +import { type ComponentProps, useEffect, useRef } from 'react' + +import { createRendererLoopPauseController } from '@/lib/renderer-loop-pause' + +const PULSE_DURATION_MS = 400 +const PULSE_PERIOD_MS = 5_000 + +// One pause controller + one period timer shared by every StatusPulse +// instance. A sidebar can show dozens of pulsing dots at once; per-instance +// controllers would mean N×(document/window/bridge) listeners and N +// unsynchronized 5s wakes. Ref-counted: the controller and timer exist only +// while at least one pulse is mounted, and all pulses play in one aligned +// wake so the renderer sleeps between beats. +type PulseSubscriber = { play: () => void; cancel: () => void } + +const pulseSubscribers = new Set() +let sharedPauseController: ReturnType | null = null +let sharedTimer = 0 + +const stopSharedTimer = () => { + if (sharedTimer !== 0) { + window.clearTimeout(sharedTimer) + sharedTimer = 0 + } +} + +const beat = () => { + sharedTimer = 0 + + if (sharedPauseController?.isPaused() || pulseSubscribers.size === 0) { + return + } + + for (const subscriber of pulseSubscribers) { + subscriber.play() + } + + sharedTimer = window.setTimeout(beat, PULSE_PERIOD_MS) +} + +const handleSharedPauseChange = () => { + stopSharedTimer() + + if (sharedPauseController?.isPaused()) { + // Minimized/hidden: cancel in-flight animations so the compositor can + // sleep immediately instead of finishing a pulse nobody sees. + for (const subscriber of pulseSubscribers) { + subscriber.cancel() + } + + return + } + + beat() +} + +const subscribePulse = (subscriber: PulseSubscriber): (() => void) => { + pulseSubscribers.add(subscriber) + + if (!sharedPauseController) { + sharedPauseController = createRendererLoopPauseController(handleSharedPauseChange) + } + + // First subscriber (or a new one joining mid-sleep): play immediately and + // start the beat. Later joiners just wait for the next aligned beat. + if (sharedTimer === 0 && !sharedPauseController.isPaused()) { + subscriber.play() + sharedTimer = window.setTimeout(beat, PULSE_PERIOD_MS) + } + + return () => { + pulseSubscribers.delete(subscriber) + + if (pulseSubscribers.size === 0) { + stopSharedTimer() + sharedPauseController?.dispose() + sharedPauseController = null + } + } +} + +export interface StatusPulseProps extends Omit, 'children' | 'ref'> { + kind: 'opacity' | 'ping' + opacity?: number +} + +/** + * A finite status pulse with a real sleep between plays. + * + * Continuous CSS animations keep Chromium producing frames and recalculating + * styles for an otherwise motionless Desktop window. Drive the same visual + * cue directly so React stays out of the loop and the renderer/compositor can + * sleep between pulses. + */ +export function StatusPulse({ kind, opacity = 1, ...props }: StatusPulseProps) { + const ref = useRef(null) + + useEffect(() => { + const element = ref.current + + if ( + !element || + typeof element.animate !== 'function' || + window.matchMedia?.('(prefers-reduced-motion: reduce)').matches + ) { + return + } + + let animation: Animation | null = null + + const play = () => { + animation?.cancel() + animation = element.animate( + kind === 'ping' + ? [ + { opacity, transform: 'scale(1)' }, + { opacity: 0, transform: 'scale(2)' } + ] + : [{ opacity: 1 }, { opacity: 0.5 }, { opacity: 1 }], + { + duration: PULSE_DURATION_MS, + easing: kind === 'ping' ? 'cubic-bezier(0, 0, 0.2, 1)' : 'ease-in-out', + iterations: 1 + } + ) + } + + const cancel = () => { + animation?.cancel() + animation = null + } + + const unsubscribe = subscribePulse({ play, cancel }) + + return () => { + unsubscribe() + cancel() + } + }, [kind, opacity]) + + return +} diff --git a/apps/desktop/src/debug/index.ts b/apps/desktop/src/debug/index.ts index e5eb37b55b..a796085847 100644 --- a/apps/desktop/src/debug/index.ts +++ b/apps/desktop/src/debug/index.ts @@ -29,6 +29,7 @@ import './render-counter' // app under REAL sessions instead of a synthetic scenario's toy transcripts. // window.__PERF_LIVE__.on() in the console, then just use the app. import './perf-live' +import './right-pane-probe' import { watchSessionAtoms } from './watched-atoms' diff --git a/apps/desktop/src/debug/right-pane-events.ts b/apps/desktop/src/debug/right-pane-events.ts new file mode 100644 index 0000000000..39a86906cd --- /dev/null +++ b/apps/desktop/src/debug/right-pane-events.ts @@ -0,0 +1,25 @@ +export type RightPanePerfEvent = + 'project-tree-render' | 'project-tree-row-render' | 'terminal-fit-active' | 'terminal-fit-hidden' | 'terminal-measure' + +export interface RightPanePerfSnapshot { + counts: Record + details: Partial>> + rows: Record +} + +declare global { + interface Window { + __RIGHT_PANE_PERF__?: { + clear: () => void + mark: (event: RightPanePerfEvent, detail?: string) => void + snapshot: () => RightPanePerfSnapshot + start: () => void + stop: () => void + } + } +} + +/** Tiny production-safe callsite; the recorder exists only in dev/perf builds. */ +export function markRightPanePerf(event: RightPanePerfEvent, detail?: string): void { + window.__RIGHT_PANE_PERF__?.mark(event, detail) +} diff --git a/apps/desktop/src/debug/right-pane-probe.ts b/apps/desktop/src/debug/right-pane-probe.ts new file mode 100644 index 0000000000..6204c2ec7c --- /dev/null +++ b/apps/desktop/src/debug/right-pane-probe.ts @@ -0,0 +1,58 @@ +import type { RightPanePerfEvent, RightPanePerfSnapshot } from './right-pane-events' + +const eventNames: RightPanePerfEvent[] = [ + 'project-tree-render', + 'project-tree-row-render', + 'terminal-fit-active', + 'terminal-fit-hidden', + 'terminal-measure' +] + +const blankCounts = (): Record => + Object.fromEntries(eventNames.map(name => [name, 0])) as Record + +if (typeof window !== 'undefined' && !window.__RIGHT_PANE_PERF__) { + let recording = false + let counts = blankCounts() + let details: RightPanePerfSnapshot['details'] = {} + let rows: Record = {} + + window.__RIGHT_PANE_PERF__ = { + clear: () => { + counts = blankCounts() + details = {} + rows = {} + }, + mark: (event, detail) => { + if (!recording) { + return + } + + counts[event] += 1 + + if (detail) { + const eventDetails = details[event] ?? {} + eventDetails[detail] = (eventDetails[detail] ?? 0) + 1 + details[event] = eventDetails + + if (event === 'project-tree-row-render') { + rows[detail] = (rows[detail] ?? 0) + 1 + } + } + }, + snapshot: (): RightPanePerfSnapshot => ({ + counts: { ...counts }, + details: Object.fromEntries(Object.entries(details).map(([event, eventDetails]) => [event, { ...eventDetails }])), + rows: { ...rows } + }), + start: () => { + counts = blankCounts() + details = {} + rows = {} + recording = true + }, + stop: () => { + recording = false + } + } +} diff --git a/apps/desktop/src/hermes.ts b/apps/desktop/src/hermes.ts index 9088b1e7a7..b3d5a45555 100644 --- a/apps/desktop/src/hermes.ts +++ b/apps/desktop/src/hermes.ts @@ -1,5 +1,6 @@ import { JsonRpcGatewayClient } from '@hermes/shared' +import { reconnectBackoffDelayMs } from '@/lib/reconnect-backoff' import type { ActionResponse, ActionStatusResponse, @@ -349,8 +350,12 @@ export function pluginSocket(pluginId: string, path: string, onMessage: (data: u socket = null if (!disposed) { + // Full-jitter exponential backoff: same rationale as the gateway + // socket reconnect loops — an immediate-retry loop across many + // desktop clients floods the gateway with connection attempts + // during a restart. + window.setTimeout(() => void connect(), reconnectBackoffDelayMs(attempt, { baseDelayMs: 500, capMs: 30_000 })) attempt += 1 - window.setTimeout(() => void connect(), Math.min(30_000, 1_000 * 2 ** attempt)) } } } diff --git a/apps/desktop/src/i18n/ar.ts b/apps/desktop/src/i18n/ar.ts index d71aab8858..6091c916c3 100644 --- a/apps/desktop/src/i18n/ar.ts +++ b/apps/desktop/src/i18n/ar.ts @@ -2236,6 +2236,7 @@ export const ar = defineLocale({ closeRunningBody: 'هذه المحادثة ما زالت تعمل (أو تنتظر إدخالك). إغلاق التبويب يخفيها فقط — ستحتفظ الجلسة بتقدمها ويمكن إعادة فتحها من الشريط الجانبي.', closeRunningConfirm: 'إغلاق التبويب', + reload: 'إعادة التحميل', closeOthers: 'إغلاق الأخرى', closeToRight: 'إغلاق ما على اليمين', closeAll: 'إغلاق الكل', diff --git a/apps/desktop/src/i18n/en.ts b/apps/desktop/src/i18n/en.ts index 5d72bd090e..fe764eb310 100644 --- a/apps/desktop/src/i18n/en.ts +++ b/apps/desktop/src/i18n/en.ts @@ -2663,6 +2663,7 @@ export const en: Translations = { closeRunningBody: 'This chat is still working (or waiting on your input). Closing the tab hides it — the session keeps its progress and can be reopened from the sidebar.', closeRunningConfirm: 'Close tab', + reload: 'Reload', closeOthers: 'Close others', closeToRight: 'Close to the right', closeAll: 'Close all', diff --git a/apps/desktop/src/i18n/ja.ts b/apps/desktop/src/i18n/ja.ts index 13bfd18c3b..16d58834c6 100644 --- a/apps/desktop/src/i18n/ja.ts +++ b/apps/desktop/src/i18n/ja.ts @@ -2489,6 +2489,7 @@ export const ja = defineLocale({ hideHeader: 'ヘッダーを隠す', minimize: '最小化', restore: '復元', + reload: '再読み込み', closeOthers: '他を閉じる', closeToRight: '右側を閉じる', closeAll: 'すべて閉じる', diff --git a/apps/desktop/src/i18n/types.ts b/apps/desktop/src/i18n/types.ts index 79b56030d6..b82feb4001 100644 --- a/apps/desktop/src/i18n/types.ts +++ b/apps/desktop/src/i18n/types.ts @@ -2262,6 +2262,7 @@ export interface Translations { closeRunningTitle: string closeRunningBody: string closeRunningConfirm: string + reload: string closeOthers: string closeToRight: string closeAll: string diff --git a/apps/desktop/src/i18n/zh-hant.ts b/apps/desktop/src/i18n/zh-hant.ts index 7f0fcfde88..31a0b62a53 100644 --- a/apps/desktop/src/i18n/zh-hant.ts +++ b/apps/desktop/src/i18n/zh-hant.ts @@ -2409,6 +2409,7 @@ export const zhHant = defineLocale({ hideHeader: '隱藏標題列', minimize: '最小化', restore: '還原', + reload: '重新載入', closeOthers: '關閉其他', closeToRight: '關閉右側', closeAll: '全部關閉', diff --git a/apps/desktop/src/i18n/zh.ts b/apps/desktop/src/i18n/zh.ts index a9ccfee506..10b4727245 100644 --- a/apps/desktop/src/i18n/zh.ts +++ b/apps/desktop/src/i18n/zh.ts @@ -2841,6 +2841,7 @@ export const zh: Translations = { closeRunningTitle: '关闭正在运行的标签?', closeRunningBody: '此对话仍在运行(或正在等待你的输入)。关闭标签只会隐藏它——会话将保留进度,可从侧边栏重新打开。', closeRunningConfirm: '关闭标签', + reload: '重新加载', closeOthers: '关闭其他', closeToRight: '关闭右侧', closeAll: '全部关闭', diff --git a/apps/desktop/src/lib/embedded-images.test.ts b/apps/desktop/src/lib/embedded-images.test.ts index 7d6a0b75e8..c4ff61852e 100644 --- a/apps/desktop/src/lib/embedded-images.test.ts +++ b/apps/desktop/src/lib/embedded-images.test.ts @@ -1,6 +1,6 @@ import { describe, expect, it } from 'vitest' -import { extractEmbeddedImages, extractImageRefs, textWithoutImageRefs } from './embedded-images' +import { extractEmbeddedImages, extractImageRefs } from './embedded-images' const SAMPLE_PNG_DATA_URL = 'data:image/png;base64,' + 'A'.repeat(120) @@ -43,28 +43,6 @@ describe('extractEmbeddedImages', () => { }) }) -describe('textWithoutImageRefs', () => { - it('leaves plain text untouched', () => { - expect(textWithoutImageRefs('just a question')).toBe('just a question') - }) - - it('strips a single leading @image directive line', () => { - expect(textWithoutImageRefs('@image:/tmp/cat.png\nwhat is this?')).toBe('what is this?') - }) - - it('strips multiple @image directive lines and trims', () => { - const input = '@image:/tmp/a.png\n@image:/tmp/b.png\n describe both ' - - expect(textWithoutImageRefs(input)).toBe('describe both') - }) - - it('does not treat an inline @image mention as a directive line', () => { - // Only full-line leading directives are stripped, matching the gateway's - // persist-time rewrite. A bare mention mid-prose is preserved. - expect(textWithoutImageRefs('see @image:/tmp/cat.png here')).toBe('see @image:/tmp/cat.png here') - }) -}) - describe('extractImageRefs', () => { it('returns the text untouched and no refs when there are no directives', () => { expect(extractImageRefs('a normal prompt')).toEqual({ cleanedText: 'a normal prompt', refs: [] }) diff --git a/apps/desktop/src/lib/embedded-images.ts b/apps/desktop/src/lib/embedded-images.ts index 144a602091..57578724df 100644 --- a/apps/desktop/src/lib/embedded-images.ts +++ b/apps/desktop/src/lib/embedded-images.ts @@ -165,20 +165,15 @@ export function textWithoutEmbeddedImages(text: string): string { // (see tui_gateway/server.py's persist-time rewrite), prepended before the // user's own text. The composer's own optimistic/local turn never carries // this prefix — it keeps the attachment as separate `attachmentRefs` -// metadata, not inline text. Comparing raw chatMessageText between the -// optimistic turn and the authoritative (persisted) turn therefore always -// mismatches whenever an image was attached, which defeats the "is this the -// same turn" checks in preserveLocalPendingTurnMessages / appendLiveSessionProjection -// and re-appends the optimistic row as if it were a distinct, unconfirmed -// turn — a duplicated user bubble. Strip the directive line(s) before any -// such equality comparison so both sides reduce to the same visible text. +// metadata, not inline text. The turn-equality comparisons in +// preserveLocalPendingTurnMessages / appendLiveSessionProjection strip ALL +// reference-directive lines (not just images) via +// `textWithoutReferenceLines` in components/assistant-ui/reference-kinds.ts; +// IMAGE_REF_LINE_RE remains here for extractImageRefs below, which moves the +// image directives into attachmentRefs metadata. const IMAGE_REF_LINE_RE = /^@image:[^\n]*\n?/gm -export function textWithoutImageRefs(text: string): string { - return text.replace(IMAGE_REF_LINE_RE, '').trim() -} - -// Same directive lines as textWithoutImageRefs, but keeps them instead of +// Same directive lines as IMAGE_REF_LINE_RE, but keeps them instead of // discarding — used when converting persisted server messages into // ChatMessage/ThreadMessageLike shape, where `@image:` refs need to // move from inline text into the `attachmentRefs` metadata field (mirroring diff --git a/apps/desktop/src/lib/incremental-external-store-runtime.ts b/apps/desktop/src/lib/incremental-external-store-runtime.ts index 0df3ed9b2e..1ccc5121d9 100644 --- a/apps/desktop/src/lib/incremental-external-store-runtime.ts +++ b/apps/desktop/src/lib/incremental-external-store-runtime.ts @@ -243,9 +243,13 @@ export function useIncrementalExternalStoreRuntime( ): AssistantRuntime { const [runtime] = useState(() => new IncrementalExternalStoreRuntimeCore(store as ExternalStoreAdapter)) + // Re-sync the adapter only when it actually changes — a dep-less effect ran + // on EVERY render of the chat surface. `__internal_setAdapter` early-exits + // when the store is unchanged, so gating on [runtime, store] is behavior- + // preserving while skipping the per-render call entirely. useEffect(() => { runtime.setAdapter(store as ExternalStoreAdapter) - }) + }, [runtime, store]) const { modelContext } = useRuntimeAdapters() ?? {} diff --git a/apps/desktop/src/lib/inflight-turn-journal.test.ts b/apps/desktop/src/lib/inflight-turn-journal.test.ts index 41d7a362d4..573d5f8227 100644 --- a/apps/desktop/src/lib/inflight-turn-journal.test.ts +++ b/apps/desktop/src/lib/inflight-turn-journal.test.ts @@ -179,7 +179,9 @@ describe('recoverInFlightTurnJournal', () => { it('overlays the backend text-only projection instead of dropping local tool progress', () => { // Sweeper regression on #44339: a backend `inflight` assistant snapshot // (text only) used to mark the richer local tail "caught up" and delete - // locally recorded tool calls. + // locally recorded tool calls. After #76444, longer text wins only when it + // is a strict extension of the journal answer (flat thinking dumps must + // not replace structured answer text). journalEntry([ user('u1', 'do the thing'), assistantWithTool('assistant-stream-old', 'local part', { pending: true }) @@ -187,7 +189,7 @@ describe('recoverInFlightTurnJournal', () => { const base = [ user('db-u1', 'do the thing'), - assistant('assistant-stream-rt9', 'longer partial text from the backend snapshot', { pending: true }) + assistant('assistant-stream-rt9', 'local part and more from the backend snapshot', { pending: true }) ] const result = recoverInFlightTurnJournal('stored-1', base, { keepPending: true }) @@ -200,13 +202,32 @@ describe('recoverInFlightTurnJournal', () => { // Keeps the BASE projection row id so live deltas keep landing on it. expect(merged.id).toBe('assistant-stream-rt9') expect(result.streamId).toBe('assistant-stream-rt9') - // Journal structure survives; the longer backend text wins. + // Journal structure survives; strict-extension backend text wins. expect(merged.parts[0]).toMatchObject({ type: 'tool-call', toolName: 'terminal' }) - expect(merged.parts[1]).toMatchObject({ type: 'text', text: 'longer partial text from the backend snapshot' }) + expect(merged.parts[1]).toMatchObject({ type: 'text', text: 'local part and more from the backend snapshot' }) // Still in flight — the journal must NOT be cleared. expect(readInFlightTurnJournal('stored-1')).not.toBeNull() }) + it('keeps journal answer text when a longer flat dump is not a strict extension (#76444)', () => { + journalEntry([user('u1', 'do the thing'), assistantWithTool('assistant-stream-old', 'partial', { pending: true })]) + + const base = [ + user('db-u1', 'do the thing'), + assistant( + 'assistant-stream-rt9', + 'thinking chatter\nRan terminal\npartial and unrelated dump longer than answer', + { pending: true } + ) + ] + + const result = recoverInFlightTurnJournal('stored-1', base, { keepPending: true }) + const merged = result.messages.at(-1)! + + expect(merged.parts[0]).toMatchObject({ type: 'tool-call', toolName: 'terminal' }) + expect(merged.parts[1]).toMatchObject({ type: 'text', text: 'partial' }) + }) + it('keeps the journal text when it is longer than the projection text', () => { journalEntry([ user('u1', 'do the thing'), diff --git a/apps/desktop/src/lib/inflight-turn-journal.ts b/apps/desktop/src/lib/inflight-turn-journal.ts index 2e4d0cb42f..050c4ff0c9 100644 --- a/apps/desktop/src/lib/inflight-turn-journal.ts +++ b/apps/desktop/src/lib/inflight-turn-journal.ts @@ -254,6 +254,10 @@ function assistantTextLength(message: ChatMessage): number { * BASE row's id so live deltas keep appending to the row the stream handler * already targets. */ +function hasStructuralParts(message: ChatMessage): boolean { + return message.parts.some(part => part.type === 'reasoning' || part.type === 'tool-call') +} + function overlayProjectionRow(projection: ChatMessage, journalRow: ChatMessage): ChatMessage { // A projected error (retained failed turn) must survive the overlay. const error = journalRow.error ?? projection.error @@ -271,7 +275,20 @@ function overlayProjectionRow(projection: ChatMessage, journalRow: ChatMessage): // Backend text is newer than the journal's last throttled write — swap it // into the journal's first text part, keeping tool calls and reasoning. + // When the journal already carries structure, only accept a *strict* + // extension of the answer text. A longer flat dump that starts with + // thinking chatter must not overwrite / insert as answer text (#76444). const projectionText = chatMessageText(projection) + const journalText = chatMessageText(journalRow).trim() + + if (hasStructuralParts(journalRow)) { + const next = projectionText.trim() + + if (!journalText || !next.startsWith(journalText)) { + return merged + } + } + const parts: ChatMessagePart[] = [] let textReplaced = false diff --git a/apps/desktop/src/lib/reconnect-backoff.test.ts b/apps/desktop/src/lib/reconnect-backoff.test.ts new file mode 100644 index 0000000000..53e51601f4 --- /dev/null +++ b/apps/desktop/src/lib/reconnect-backoff.test.ts @@ -0,0 +1,92 @@ +import { describe, expect, it, vi } from 'vitest' + +import { reconnectBackoffDelayMs } from './reconnect-backoff' + +describe('reconnectBackoffDelayMs', () => { + it('increases the delay ceiling across consecutive failed attempts', () => { + // Pin Math.random so we can read the ceiling directly through the + // returned value instead of statistically sampling it. + const randomSpy = vi.spyOn(Math, 'random').mockReturnValue(1) + + try { + const delays = [0, 1, 2, 3, 4].map(attempt => reconnectBackoffDelayMs(attempt, { baseDelayMs: 300 })) + + expect(delays).toEqual([300, 600, 1200, 2400, 4800]) + + for (let i = 1; i < delays.length; i++) { + expect(delays[i]).toBeGreaterThan(delays[i - 1]) + } + } finally { + randomSpy.mockRestore() + } + }) + + it('caps the delay ceiling instead of growing unbounded', () => { + const randomSpy = vi.spyOn(Math, 'random').mockReturnValue(1) + + try { + // Attempt 10 would be 300 * 2**10 = 307_200ms uncapped — must clamp. + expect(reconnectBackoffDelayMs(10, { baseDelayMs: 300, capMs: 15_000 })).toBe(15_000) + expect(reconnectBackoffDelayMs(50, { baseDelayMs: 300, capMs: 15_000 })).toBe(15_000) + } finally { + randomSpy.mockRestore() + } + }) + + it('applies full jitter: delay is uniformly within [0, ceiling)', () => { + const randomSpy = vi.spyOn(Math, 'random') + + try { + randomSpy.mockReturnValue(0) + expect(reconnectBackoffDelayMs(3, { baseDelayMs: 300 })).toBe(0) + + randomSpy.mockReturnValue(0.5) + expect(reconnectBackoffDelayMs(3, { baseDelayMs: 300 })).toBe(1200) + + randomSpy.mockReturnValue(0.999) + expect(reconnectBackoffDelayMs(3, { baseDelayMs: 300 })).toBeCloseTo(2400 * 0.999, 5) + } finally { + randomSpy.mockRestore() + } + }) + + it('resets to the attempt-0 ceiling after a successful connection (caller passes attempt back to 0)', () => { + const randomSpy = vi.spyOn(Math, 'random').mockReturnValue(1) + + try { + // Simulates: fail, fail, fail (attempt climbs), succeed (caller resets + // its counter to 0), fail again — the very next delay must be back at + // the base ceiling, not continuing the climb. + reconnectBackoffDelayMs(0, { baseDelayMs: 300 }) + reconnectBackoffDelayMs(1, { baseDelayMs: 300 }) + const afterSeveralFailures = reconnectBackoffDelayMs(2, { baseDelayMs: 300 }) + const afterReset = reconnectBackoffDelayMs(0, { baseDelayMs: 300 }) + + expect(afterSeveralFailures).toBe(1200) + expect(afterReset).toBe(300) + } finally { + randomSpy.mockRestore() + } + }) + + it('treats negative attempt numbers as attempt 0 rather than throwing or returning a negative delay', () => { + const randomSpy = vi.spyOn(Math, 'random').mockReturnValue(1) + + try { + expect(reconnectBackoffDelayMs(-5, { baseDelayMs: 300 })).toBe(300) + } finally { + randomSpy.mockRestore() + } + }) + + it('uses sane defaults when no options are passed', () => { + const randomSpy = vi.spyOn(Math, 'random').mockReturnValue(1) + + try { + expect(reconnectBackoffDelayMs(0)).toBe(300) + expect(reconnectBackoffDelayMs(100)).toBe(15_000) + } finally { + randomSpy.mockRestore() + } + }) +}) diff --git a/apps/desktop/src/lib/reconnect-backoff.ts b/apps/desktop/src/lib/reconnect-backoff.ts new file mode 100644 index 0000000000..c24ce9a24f --- /dev/null +++ b/apps/desktop/src/lib/reconnect-backoff.ts @@ -0,0 +1,45 @@ +/** + * Full-jitter exponential backoff for gateway WebSocket reconnects. + * + * A bare exponential backoff still lets every renderer in a fleet retry in + * lockstep — after a gateway restart (e.g. following an update), N desktop + * clients that all disconnected within the same instant all wake up and + * redial at the same instant too, which is a reconnect storm by another + * name. Full jitter (AWS's "Exponential Backoff And Jitter") spreads that + * out: each attempt sleeps a *random* duration between 0 and the exponential + * ceiling, so retries desynchronize instead of pulsing together. + * + * https://aws.amazon.com/blogs/architecture/exponential-backoff-and-jitter/ + */ + +export interface ReconnectBackoffOptions { + /** Ceiling on the exponential delay before jitter is applied, in ms. */ + capMs?: number + /** Delay for the first retry (attempt 0) before jitter is applied, in ms. */ + baseDelayMs?: number +} + +const DEFAULT_BASE_DELAY_MS = 300 +const DEFAULT_CAP_MS = 15_000 + +/** + * Delay before reconnect attempt number `attempt` (0-indexed: the first + * retry after the initial failure is `attempt = 0`). Returns a value in + * `[0, min(capMs, baseDelayMs * 2 ** attempt))` — full jitter, not + * "equal jitter" or "decorrelated jitter", so it can occasionally return a + * very small delay even at a high attempt count. That's intentional: it's + * the variant with the best-documented storm-avoidance behavior and no + * accumulated-delay state to track between calls. + */ +export function reconnectBackoffDelayMs(attempt: number, options: ReconnectBackoffOptions = {}): number { + const baseDelayMs = options.baseDelayMs ?? DEFAULT_BASE_DELAY_MS + const capMs = options.capMs ?? DEFAULT_CAP_MS + const safeAttempt = Math.max(0, attempt) + + // 2 ** attempt overflows to Infinity long before it matters (attempt would + // need to be ~1024), and Math.min against a finite cap keeps the ceiling + // sane regardless, so no extra clamping is needed here. + const ceiling = Math.min(capMs, baseDelayMs * 2 ** safeAttempt) + + return Math.random() * ceiling +} diff --git a/apps/desktop/src/store/coding-status.test.ts b/apps/desktop/src/store/coding-status.test.ts index aa731c5b25..d5a67e5898 100644 --- a/apps/desktop/src/store/coding-status.test.ts +++ b/apps/desktop/src/store/coding-status.test.ts @@ -10,6 +10,7 @@ import { refreshAllRepoStatuses, refreshRepoStatus, registerRepoStatusCwd, + repoChangeKindForPath, repoStatusForCwd } from './coding-status' import { $currentCwd, $selectedStoredSessionId } from './session' @@ -255,3 +256,32 @@ describe('refreshRepoStatus', () => { release?.() }) }) + +describe('repoChangeKindForPath', () => { + it('does not notify a row when only another path changes', () => { + $currentCwd.set('/repo') + $repoStatusByCwd.set({ '/repo': { ...sampleStatus, files: [] } }) + const row = repoChangeKindForPath('/repo/a.ts') + const listener = vi.fn() + const unsubscribe = row.subscribe(listener) + + $repoStatusByCwd.set({ + '/repo': { + ...sampleStatus, + files: [{ path: 'b.ts', untracked: true } as HermesRepoStatus['files'][number]] + } + }) + expect(listener).toHaveBeenCalledTimes(1) + + $repoStatusByCwd.set({ + '/repo': { + ...sampleStatus, + files: [{ path: 'a.ts', untracked: true } as HermesRepoStatus['files'][number]] + } + }) + expect(listener).toHaveBeenCalledTimes(2) + expect(listener.mock.calls.at(-1)?.[0]).toBe('added') + + unsubscribe() + }) +}) diff --git a/apps/desktop/src/store/coding-status.ts b/apps/desktop/src/store/coding-status.ts index 4a1ea97b0c..5e434e0744 100644 --- a/apps/desktop/src/store/coding-status.ts +++ b/apps/desktop/src/store/coding-status.ts @@ -112,6 +112,14 @@ export const $repoChangeByPath = computed([$repoStatus, $currentCwd], (status, c return map }) +/** + * Per-row Git decoration subscription. A visible file row reads one scalar, so + * a fresh repo-status map only re-renders that row when its own kind changed. + */ +export function repoChangeKindForPath(path: string): ReadableAtom { + return computed($repoChangeByPath, changes => changes.get(path)) +} + // Cwds whose rails are on screen right now (refcounted — two tiles in one // worktree register it twice). Every refresh edge re-probes each registered // cwd plus the primary workspace, so a tile's rail moves when ITS agent diff --git a/apps/desktop/src/store/gateway.ts b/apps/desktop/src/store/gateway.ts index f9e89ca434..149235e42f 100644 --- a/apps/desktop/src/store/gateway.ts +++ b/apps/desktop/src/store/gateway.ts @@ -2,6 +2,7 @@ import { type ConnectionState, type GatewayEvent, resolveGatewayWsUrl } from '@h import { atom } from 'nanostores' import { HermesGateway } from '@/hermes' +import { reconnectBackoffDelayMs } from '@/lib/reconnect-backoff' import { markNativeNotifyBaseline } from '@/store/notify-baseline' import { setGatewayState } from '@/store/session' @@ -184,8 +185,9 @@ function scheduleReconnect(entry: Secondary): void { return } - // 1s, 2s, 4s … capped at 15s — same backoff shape as the primary. - const delay = Math.min(15_000, 1_000 * 2 ** Math.min(entry.reconnectAttempt, 4)) + // Full-jitter exponential backoff — same shape (and same reason: avoid a + // reconnect storm against a restarting gateway) as the primary's. + const delay = reconnectBackoffDelayMs(entry.reconnectAttempt) entry.reconnectAttempt += 1 entry.reconnectTimer = setTimeout(() => { entry.reconnectTimer = null diff --git a/apps/desktop/src/store/pet-gallery.test.ts b/apps/desktop/src/store/pet-gallery.test.ts new file mode 100644 index 0000000000..5a5f494a8a --- /dev/null +++ b/apps/desktop/src/store/pet-gallery.test.ts @@ -0,0 +1,210 @@ +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' + +import { $petInfo, setPetInfo } from './pet' +import { $petGallery, adoptPet, type GatewayRequest, loadPetGallery, resetPetGallery } from './pet-gallery' + +function localGallery() { + return { + enabled: true, + active: 'boba', + pets: [{ slug: 'boba', displayName: 'Boba', installed: true }] + } +} + +describe('pet gallery pet.info sync', () => { + beforeEach(() => { + resetPetGallery() + setPetInfo({ enabled: false }) + }) + + afterEach(() => { + resetPetGallery() + setPetInfo({ enabled: false }) + vi.restoreAllMocks() + }) + + it('uses pet.info.meta and keeps the cached spritesheet when the revision is current', async () => { + setPetInfo({ + enabled: true, + slug: 'boba', + displayName: 'Old Boba', + scale: 0.33, + spritesheetBase64: 'large-sprite-payload', + spritesheetRevision: '100:2048', + frameW: 192, + frameH: 208 + }) + + const requestMock = vi.fn(async (method: string) => { + if (method === 'pet.gallery') { + return localGallery() + } + + if (method === 'pet.info.meta') { + return { + enabled: true, + slug: 'boba', + displayName: 'Boba', + scale: 0.5, + spritesheetRevision: '100:2048' + } + } + + if (method === 'pet.info') { + throw new Error('full pet.info should not be called for an unchanged sprite') + } + + throw new Error(`unexpected method: ${method}`) + }) + + const request = requestMock as unknown as GatewayRequest + + await loadPetGallery(request) + + const methods = requestMock.mock.calls.map(([method]) => method) + expect(methods).toContain('pet.info.meta') + expect(methods).not.toContain('pet.info') + expect($petInfo.get()).toMatchObject({ + enabled: true, + slug: 'boba', + displayName: 'Boba', + scale: 0.5, + spritesheetBase64: 'large-sprite-payload', + spritesheetRevision: '100:2048', + frameW: 192, + frameH: 208 + }) + }) + + it('fetches full pet.info when metadata reports a new spritesheet revision', async () => { + setPetInfo({ + enabled: true, + slug: 'boba', + displayName: 'Boba', + scale: 0.33, + spritesheetBase64: 'old-sprite-payload', + spritesheetRevision: '100:2048' + }) + + const requestMock = vi.fn(async (method: string) => { + if (method === 'pet.gallery') { + return localGallery() + } + + if (method === 'pet.info.meta') { + return { + enabled: true, + slug: 'boba', + displayName: 'Boba', + scale: 0.33, + spritesheetRevision: '101:4096' + } + } + + if (method === 'pet.info') { + return { + enabled: true, + slug: 'boba', + displayName: 'Boba', + scale: 0.33, + spritesheetBase64: 'new-sprite-payload', + spritesheetRevision: '101:4096' + } + } + + throw new Error(`unexpected method: ${method}`) + }) + + const request = requestMock as unknown as GatewayRequest + + await loadPetGallery(request) + + const methods = requestMock.mock.calls.map(([method]) => method) + expect(methods).toContain('pet.info.meta') + expect(methods).toContain('pet.info') + expect($petInfo.get().spritesheetBase64).toBe('new-sprite-payload') + expect($petInfo.get().spritesheetRevision).toBe('101:4096') + }) + + it('falls back to full pet.info when an older gateway lacks metadata', async () => { + const requestMock = vi.fn(async (method: string) => { + if (method === 'pet.gallery') { + return localGallery() + } + + if (method === 'pet.info.meta') { + throw new Error('JSON-RPC -32601: Method not found') + } + + if (method === 'pet.info') { + return { + enabled: true, + slug: 'boba', + displayName: 'Boba from legacy gateway', + scale: 0.4, + spritesheetBase64: 'legacy-full-payload', + spritesheetRevision: '99:1024' + } + } + + throw new Error(`unexpected method: ${method}`) + }) + + const request = requestMock as unknown as GatewayRequest + + await loadPetGallery(request) + + const methods = requestMock.mock.calls.map(([method]) => method) + expect(methods).toContain('pet.info.meta') + expect(methods).toContain('pet.info') + expect($petInfo.get()).toMatchObject({ + enabled: true, + slug: 'boba', + displayName: 'Boba from legacy gateway', + spritesheetBase64: 'legacy-full-payload', + spritesheetRevision: '99:1024' + }) + }) + + it('keeps mutation sync on metadata when the selected pet sprite is unchanged', async () => { + $petGallery.set(localGallery()) + setPetInfo({ + enabled: true, + slug: 'boba', + displayName: 'Boba', + scale: 0.33, + spritesheetBase64: 'large-sprite-payload', + spritesheetRevision: '100:2048' + }) + + const requestMock = vi.fn(async (method: string) => { + if (method === 'pet.select') { + return { ok: true, slug: 'boba', displayName: 'Boba' } + } + + if (method === 'pet.info.meta') { + return { + enabled: true, + slug: 'boba', + displayName: 'Boba', + scale: 0.33, + spritesheetRevision: '100:2048' + } + } + + if (method === 'pet.info') { + throw new Error('full pet.info should not be called after an unchanged select') + } + + throw new Error(`unexpected method: ${method}`) + }) + + const request = requestMock as unknown as GatewayRequest + + await expect(adoptPet(request, 'boba', 'Could not adopt pet.')).resolves.toBe(true) + + const methods = requestMock.mock.calls.map(([method]) => method) + expect(methods).toEqual(['pet.select', 'pet.info.meta']) + expect($petInfo.get().spritesheetBase64).toBe('large-sprite-payload') + }) +}) diff --git a/apps/desktop/src/store/pet-gallery.ts b/apps/desktop/src/store/pet-gallery.ts index 40cb420e95..60322bcc1f 100644 --- a/apps/desktop/src/store/pet-gallery.ts +++ b/apps/desktop/src/store/pet-gallery.ts @@ -1,7 +1,15 @@ import { atom } from 'nanostores' import { normalize } from '@/lib/text' -import { $petInfo, type PetInfo, petProfile, setPetInfo } from '@/store/pet' +import { + $petInfo, + hasPetSpriteForMeta, + mergePetInfoMeta, + type PetInfo, + type PetInfoMeta, + petProfile, + setPetInfo +} from '@/store/pet' /** * Feature store for the petdex gallery picker (Cmd+K "Pets…" + Settings). @@ -128,9 +136,9 @@ export function loadPetGallery(request: GatewayRequest, options: { force?: boole try { // Phase 1: local pets only — instant, never blocks on the remote petdex // manifest. The user's own/generated pets render right away. - const [local, info] = await Promise.all([ + const [local] = await Promise.all([ petRpc(request, 'pet.gallery', { localOnly: true }), - petRpc(request, 'pet.info') + syncInfo(request) ]) if (local) { @@ -139,10 +147,6 @@ export function loadPetGallery(request: GatewayRequest, options: { force?: boole $petGalleryError.set(null) localOk = true } - - if (info) { - setPetInfo(info) - } } catch (e) { if (isMissingMethod(e)) { $petGalleryStatus.set('stale') @@ -179,6 +183,46 @@ export function loadPetGallery(request: GatewayRequest, options: { force?: boole // network gallery — the floating pet repaints, the picker keeps its cache. async function syncInfo(request: GatewayRequest): Promise { try { + let meta: PetInfoMeta | null = null + + try { + meta = await petRpc(request, 'pet.info.meta') + } catch (e) { + if (!isMissingMethod(e)) { + throw e + } + + const info = await petRpc(request, 'pet.info') + + if (info) { + setPetInfo(info) + } + + return + } + + if (!meta) { + return + } + + if (!meta.enabled) { + setPetInfo({ enabled: false }) + + return + } + + const current = $petInfo.get() + + if (hasPetSpriteForMeta(current, meta)) { + const merged = mergePetInfoMeta(current, meta) + + if (merged !== current) { + setPetInfo(merged) + } + + return + } + const info = await petRpc(request, 'pet.info') if (info) { diff --git a/apps/desktop/src/store/pet.test.ts b/apps/desktop/src/store/pet.test.ts index ce2327becb..c3fa2962f3 100644 --- a/apps/desktop/src/store/pet.test.ts +++ b/apps/desktop/src/store/pet.test.ts @@ -7,6 +7,9 @@ import { $petState, derivePetState, flashPetActivity, + hasPetSpriteForMeta, + mergePetInfoMeta, + type PetInfo, setPetActivity } from './pet' @@ -75,6 +78,58 @@ describe('roam motion', () => { }) }) +describe('pet info metadata cache helpers', () => { + it('treats matching slug and spritesheet revision as a reusable sprite payload', () => { + const current = { + enabled: true, + slug: 'boba', + displayName: 'Old Boba', + scale: 0.33, + spritesheetBase64: 'large-sprite-payload', + spritesheetRevision: '100:2048' + } + + const meta = { + enabled: true, + slug: 'boba', + displayName: 'Boba', + scale: 0.5, + spritesheetRevision: '100:2048' + } + + expect(hasPetSpriteForMeta(current, meta)).toBe(true) + expect(mergePetInfoMeta(current, meta)).toMatchObject({ + enabled: true, + slug: 'boba', + displayName: 'Boba', + scale: 0.5, + spritesheetBase64: 'large-sprite-payload', + spritesheetRevision: '100:2048' + }) + }) + + it('returns the same reference when nothing changed to avoid redundant store updates', () => { + const current: PetInfo = { + enabled: true, + slug: 'boba', + displayName: 'Boba', + scale: 0.33, + spritesheetBase64: 'large-sprite-payload', + spritesheetRevision: '100:2048' + } + + const meta = { + enabled: true, + slug: 'boba', + displayName: 'Boba', + scale: 0.33, + spritesheetRevision: '100:2048' + } + + expect(mergePetInfoMeta(current, meta)).toBe(current) + }) +}) + describe('flashPetActivity', () => { it('clears stale sibling beats so a completion never inherits a prior error', () => { // A turn errors (sad), then the next turn finishes cleanly. The celebrate diff --git a/apps/desktop/src/store/pet.ts b/apps/desktop/src/store/pet.ts index b1e2e9d214..ed28cc2dbc 100644 --- a/apps/desktop/src/store/pet.ts +++ b/apps/desktop/src/store/pet.ts @@ -39,6 +39,53 @@ export interface PetInfo { stateRows?: string[] } +export interface PetInfoMeta { + enabled: boolean + slug?: string + displayName?: string + scale?: number + spritesheetRevision?: string +} + +export function hasPetSpriteForMeta(info: PetInfo, meta: PetInfoMeta): boolean { + return ( + meta.enabled && + info.enabled && + Boolean(info.spritesheetBase64) && + info.slug === meta.slug && + Boolean(info.spritesheetRevision) && + info.spritesheetRevision === meta.spritesheetRevision + ) +} + +export function mergePetInfoMeta(info: PetInfo, meta: PetInfoMeta): PetInfo { + if (!meta.enabled) { + return info.enabled ? { enabled: false } : info + } + + // Fast-path: nothing changed — return the same reference so callers can + // skip the store update (nanostores fires on .set() regardless of deep + // equality; returning `info` avoids a redundant re-render on every poll). + if ( + info.enabled && + info.slug === meta.slug && + info.displayName === meta.displayName && + info.scale === meta.scale && + info.spritesheetRevision === meta.spritesheetRevision + ) { + return info + } + + return { + ...info, + enabled: true, + slug: meta.slug, + displayName: meta.displayName, + scale: meta.scale, + spritesheetRevision: meta.spritesheetRevision + } +} + export interface PetActivity { busy?: boolean awaitingInput?: boolean diff --git a/apps/desktop/src/store/session-states.ts b/apps/desktop/src/store/session-states.ts index ebf58f9819..8c6811bf6e 100644 --- a/apps/desktop/src/store/session-states.ts +++ b/apps/desktop/src/store/session-states.ts @@ -799,7 +799,7 @@ $selectedStoredSessionId.listen(selected => { }) // Dev hook for automation (mirrors __HERMES_LAYOUT_TREE__). -if (import.meta.env.DEV && typeof window !== 'undefined') { +if ((import.meta.env.DEV || import.meta.env.VITE_PERF_PROBE === '1') && typeof window !== 'undefined') { ;(window as unknown as Record).__HERMES_SESSION_TILES__ = { close: closeSessionTile, open: openSessionTile, diff --git a/apps/desktop/src/styles.css b/apps/desktop/src/styles.css index 9b77346fa3..2ee38e0b65 100644 --- a/apps/desktop/src/styles.css +++ b/apps/desktop/src/styles.css @@ -1076,8 +1076,9 @@ code { } /* Arc-style multicolor action surface (static, not animated). Reusable on any - Button via className. Unlayered so it beats Tailwind's bg-*/ -text-* variant utilities. */ .btn-arc { + Button via className. Unlayered so it beats Tailwind's bg- and text- variant + utilities. */ +.btn-arc { background-image: linear-gradient(110deg, #5b6cff 0%, #8b5cf6 28%, #d946ef 58%, #fb7185 82%, #fb923c 100%); color: #fff; border-color: transparent; diff --git a/apps/desktop/src/types/hermes.ts b/apps/desktop/src/types/hermes.ts index 796784ff61..bc1fb5c581 100644 --- a/apps/desktop/src/types/hermes.ts +++ b/apps/desktop/src/types/hermes.ts @@ -600,6 +600,7 @@ export interface SessionResumeResponse { info?: SessionRuntimeInfo message_count: number messages: SessionMessage[] + messages_omitted?: boolean resumed: string running?: boolean session_id: string diff --git a/apps/shared/src/json-rpc-gateway.ts b/apps/shared/src/json-rpc-gateway.ts index 962ba1a9bd..4f24737ed7 100644 --- a/apps/shared/src/json-rpc-gateway.ts +++ b/apps/shared/src/json-rpc-gateway.ts @@ -54,6 +54,8 @@ export interface GatewayClientOptions { connectErrorMessage?: string connectTimeoutMs?: number createRequestId?: (nextId: number) => GatewayRequestId + /** Return true to intercept the default closed-state transition. */ + onSocketClose?: (event: CloseEvent) => boolean | void requestIdPrefix?: string requestTimeoutMs?: number socketFactory?: (url: string) => WebSocketLike @@ -84,6 +86,7 @@ export class JsonRpcGatewayClient { connectTimeoutMs: options.connectTimeoutMs ?? DEFAULT_CONNECT_TIMEOUT_MS, createRequestId: options.createRequestId ?? ((nextId: number) => `${options.requestIdPrefix ?? 'r'}${nextId}`), notConnectedErrorMessage: options.notConnectedErrorMessage ?? 'gateway not connected', + onSocketClose: options.onSocketClose ?? (() => false), requestIdPrefix: options.requestIdPrefix ?? 'r', requestTimeoutMs: options.requestTimeoutMs ?? DEFAULT_REQUEST_TIMEOUT_MS, socketFactory: options.socketFactory @@ -136,11 +139,15 @@ export class JsonRpcGatewayClient { this.handleMessage(message.data) }) - socket.addEventListener('close', () => { + socket.addEventListener('close', event => { if (this.socket !== socket) { return } + if (this.options.onSocketClose(event)) { + return + } + this.socket = null this.setState('closed') this.rejectAllPending(new Error(this.options.closedErrorMessage)) diff --git a/cli-config.yaml.example b/cli-config.yaml.example index 04ccfb6da4..6b7a4c327a 100644 --- a/cli-config.yaml.example +++ b/cli-config.yaml.example @@ -652,6 +652,25 @@ prompt_caching: # # extra_body: # # chat_template_kwargs: # # enable_thinking: false +# +# # Auto-generated short session titles after the first exchange. +# # Each active Discord/Telegram channel can spawn a background title +# # call. Cap concurrency to keep retries during provider incidents +# # from amplifying the request burst. Leave unset for legacy behavior +# # (unlimited). +# title_generation: +# provider: "auto" +# model: "" +# # max_concurrency: 2 # Optional: cap simultaneous title calls +# +# # Context compression — summarizes long sessions to shrink the prompt. +# # Heavy and often hits the slowest provider chain. Setting a small +# # cap prevents many sessions from compressing simultaneously during +# # provider degradation. Leave unset for legacy behavior (unlimited). +# compression: +# provider: "auto" +# model: "" +# # max_concurrency: 2 # Optional: cap simultaneous compression calls # ============================================================================= # Persistent Memory diff --git a/cli.py b/cli.py index 95ce2b18f8..aed3992922 100644 --- a/cli.py +++ b/cli.py @@ -7334,6 +7334,44 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): else: self._console_print(f"[dim]{_escape(msg)}[/dim]") + def _restore_session_yolo(self, session_meta: dict, *, quiet: bool = False) -> None: + """Re-enable YOLO bypass on resume when the session had it on. + + Companion to ``_restore_session_cwd`` — called from every resume path + (startup ``--resume``/``-c`` and mid-chat ``/resume``). The persisted + flag lives in the session row's ``model_config.yolo_mode`` (written by + ``/yolo`` toggles and ``--yolo`` launches); without this restore the + in-memory ``tools.approval._session_yolo`` set starts empty in a fresh + process and the user's bypass silently reverts. + + No-op when the flag is absent/false, when YOLO is already active for + this session (idempotent across repeated resume paths), or when the + process was itself launched with ``--yolo`` (frozen bypass already + covers everything). + """ + try: + from hermes_state import SessionDB + from tools.approval import ( + _YOLO_MODE_FROZEN, + enable_session_yolo, + is_session_yolo_enabled, + ) + except Exception: + return + if _YOLO_MODE_FROZEN: + return + if not SessionDB.session_yolo_enabled(session_meta): + return + session_key = self.session_id or "default" + if is_session_yolo_enabled(session_key): + return + enable_session_yolo(session_key) + msg = "⚡ YOLO mode restored from session — all commands auto-approved. /yolo to turn off." + if quiet: + print(msg, file=sys.stderr) + else: + self._console_print(f"[dim]{_escape(msg)}[/dim]") + def _render_resume_history_panel_lines(self, panel) -> list[str]: @@ -10820,6 +10858,12 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): if is_session_yolo_enabled(old_session_id): enable_session_yolo(new_session_id) disable_session_yolo(old_session_id) + # Carry the persisted flag onto the continuation row so a later + # `hermes --resume ` restores the bypass too. getattr + # guard: tests call this unbound against a minimal stand-in. + _persist = getattr(self, "_persist_session_yolo", None) + if _persist: + _persist(new_session_id, True) def _is_session_yolo_active(self) -> bool: """Whether YOLO bypass is currently enabled for this CLI session. @@ -10871,19 +10915,45 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): ) session_key = self.session_id or "default" + # ``getattr`` guard: tests exercise this method unbound against a + # minimal stand-in object (see tests/cli/test_cli_yolo_toggle.py); + # persistence is best-effort either way. + _persist = getattr(self, "_persist_session_yolo", None) if is_session_yolo_enabled(session_key): disable_session_yolo(session_key) + if _persist: + _persist(session_key, False) _cprint( f" ⚠ YOLO mode {_Colors.BOLD}{_Colors.RED}OFF{_Colors.RESET}" " — dangerous commands will require approval." ) else: enable_session_yolo(session_key) + if _persist: + _persist(session_key, True) _cprint( f" ⚡ YOLO mode {_Colors.BOLD}{_Colors.GREEN}ON{_Colors.RESET}" " — all commands auto-approved. Use with caution." ) + def _persist_session_yolo(self, session_key: str, enabled: bool) -> None: + """Persist the YOLO flag to the session row so --resume restores it. + + Best-effort: the in-memory toggle is authoritative for this process; + persistence only affects a future ``hermes --resume``. Skipped when the + session store is unavailable or the row doesn't exist yet (the row is + created lazily on the first turn — ``_toggle_yolo`` before any chat + writes nothing, and the launch-time ``--yolo`` flag is carried into the + creation-time model_config instead). + """ + db = getattr(self, "_session_db", None) + if db is None or not session_key or session_key == "default": + return + try: + db.set_session_yolo(session_key, enabled) + except Exception: + pass + @@ -13229,7 +13299,10 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): self._approval_deadline = 0 self._paint_now() _cprint(f"\n{_DIM} ⏱ Timeout — denying command{_RST}") - return "deny" + self._persist_prompt_summary( + "⚠", "Approval", command, "timed out (no response)", + ) + return "timeout" def _approval_choices(self, command: str, *, allow_permanent: bool = True, smart_denied: bool = False) -> list[str]: @@ -13261,6 +13334,7 @@ class HermesCLI(CLIAgentSetupMixin, CLICommandsMixin, CLIBillingMixin): "session": "approve_session", "always": "always_approve", "deny": "deny", + "timeout": "timeout", }.get(verdict, "deny") def _handle_approval_selection(self) -> None: diff --git a/contributors/emails/1347825413@qq.com b/contributors/emails/1347825413@qq.com new file mode 100644 index 0000000000..7dba397cb1 --- /dev/null +++ b/contributors/emails/1347825413@qq.com @@ -0,0 +1 @@ +baau diff --git a/contributors/emails/314574126@qq.com b/contributors/emails/314574126@qq.com new file mode 100644 index 0000000000..e9fc6fb2cb --- /dev/null +++ b/contributors/emails/314574126@qq.com @@ -0,0 +1 @@ +ArcherQAQ diff --git a/contributors/emails/WojtekMR3@users.noreply.github.com b/contributors/emails/WojtekMR3@users.noreply.github.com new file mode 100644 index 0000000000..5e8d4d9b5e --- /dev/null +++ b/contributors/emails/WojtekMR3@users.noreply.github.com @@ -0,0 +1 @@ +WojtekMR3 diff --git a/contributors/emails/[email protected] b/contributors/emails/[email protected] new file mode 100644 index 0000000000..73e7c0fa34 --- /dev/null +++ b/contributors/emails/[email protected] @@ -0,0 +1 @@ +Ahmett101 diff --git a/contributors/emails/abdulsalamalotaibi86@gmail.com b/contributors/emails/abdulsalamalotaibi86@gmail.com new file mode 100644 index 0000000000..7a031933d8 --- /dev/null +++ b/contributors/emails/abdulsalamalotaibi86@gmail.com @@ -0,0 +1 @@ +carbongotfound diff --git a/contributors/emails/ahmedmoro@gmail.com b/contributors/emails/ahmedmoro@gmail.com new file mode 100644 index 0000000000..e8cf29e080 --- /dev/null +++ b/contributors/emails/ahmedmoro@gmail.com @@ -0,0 +1,2 @@ +morolab +# v0.20.0 audit: co-author on #70870 (Arabic locale) diff --git a/contributors/emails/akshankrithick305@gmail.com b/contributors/emails/akshankrithick305@gmail.com new file mode 100644 index 0000000000..d3557ad9f5 --- /dev/null +++ b/contributors/emails/akshankrithick305@gmail.com @@ -0,0 +1,2 @@ +akshan-main +# v0.20.0 audit: co-author on #72244 (/goal indicator) diff --git a/contributors/emails/assiri@gmail.com b/contributors/emails/assiri@gmail.com new file mode 100644 index 0000000000..0a68259b34 --- /dev/null +++ b/contributors/emails/assiri@gmail.com @@ -0,0 +1,2 @@ +3ssiri +# v0.20.0 audit: co-author on #70870 (Arabic locale) diff --git a/contributors/emails/bot@bkstock.dev b/contributors/emails/bot@bkstock.dev new file mode 100644 index 0000000000..576c4dd141 --- /dev/null +++ b/contributors/emails/bot@bkstock.dev @@ -0,0 +1 @@ +BKStock diff --git a/contributors/emails/brdpedroo@gmail.com b/contributors/emails/brdpedroo@gmail.com new file mode 100644 index 0000000000..67760b3f3f --- /dev/null +++ b/contributors/emails/brdpedroo@gmail.com @@ -0,0 +1,2 @@ +Pebrd +# v0.20.0 audit: author on #74245 (pinned Telegram sessions) diff --git a/contributors/emails/carrion256@proton.me b/contributors/emails/carrion256@proton.me new file mode 100644 index 0000000000..db7f5c8d19 --- /dev/null +++ b/contributors/emails/carrion256@proton.me @@ -0,0 +1,2 @@ +carrion256 +# v0.20.0 audit: co-author on #75888 (hindsight env 0600) diff --git a/contributors/emails/cicav@users.noreply.github.com b/contributors/emails/cicav@users.noreply.github.com new file mode 100644 index 0000000000..f53c720697 --- /dev/null +++ b/contributors/emails/cicav@users.noreply.github.com @@ -0,0 +1 @@ +cicav diff --git a/contributors/emails/coder@trevhome.local b/contributors/emails/coder@trevhome.local new file mode 100644 index 0000000000..533aa71407 --- /dev/null +++ b/contributors/emails/coder@trevhome.local @@ -0,0 +1 @@ +trevornk diff --git a/contributors/emails/coffee@coffeebot.dev b/contributors/emails/coffee@coffeebot.dev new file mode 100644 index 0000000000..917bb57cf7 --- /dev/null +++ b/contributors/emails/coffee@coffeebot.dev @@ -0,0 +1 @@ +coffee-the-dev diff --git a/contributors/emails/copii.list@gmail.com b/contributors/emails/copii.list@gmail.com new file mode 100644 index 0000000000..e04f9980d7 --- /dev/null +++ b/contributors/emails/copii.list@gmail.com @@ -0,0 +1 @@ +stremtec diff --git a/contributors/emails/core@lfdm.co b/contributors/emails/core@lfdm.co new file mode 100644 index 0000000000..02fc09a3be --- /dev/null +++ b/contributors/emails/core@lfdm.co @@ -0,0 +1 @@ +LFDMcore diff --git a/contributors/emails/dai.suzuki.829@gmail.com b/contributors/emails/dai.suzuki.829@gmail.com new file mode 100644 index 0000000000..6b2938b60d --- /dev/null +++ b/contributors/emails/dai.suzuki.829@gmail.com @@ -0,0 +1 @@ +hariNEzuMI928 diff --git a/contributors/emails/daniel.blank@reportsolution.de b/contributors/emails/daniel.blank@reportsolution.de new file mode 100644 index 0000000000..2ff605dc59 --- /dev/null +++ b/contributors/emails/daniel.blank@reportsolution.de @@ -0,0 +1 @@ +danielblankhh diff --git a/contributors/emails/david@lexgenius.ai b/contributors/emails/david@lexgenius.ai new file mode 100644 index 0000000000..1ec5f69a80 --- /dev/null +++ b/contributors/emails/david@lexgenius.ai @@ -0,0 +1,2 @@ +lexgenius +# v0.20.0 audit: co-author on #65798 (FTS schema v23) diff --git a/contributors/emails/egilewski@egilewski.com b/contributors/emails/egilewski@egilewski.com new file mode 100644 index 0000000000..d56b0bd1b6 --- /dev/null +++ b/contributors/emails/egilewski@egilewski.com @@ -0,0 +1 @@ +egilewski diff --git a/contributors/emails/f1aggo_macair@f1aggo-macairdeMacBook-Air.local b/contributors/emails/f1aggo_macair@f1aggo-macairdeMacBook-Air.local new file mode 100644 index 0000000000..abb45cd26e --- /dev/null +++ b/contributors/emails/f1aggo_macair@f1aggo-macairdeMacBook-Air.local @@ -0,0 +1 @@ +flag0x369 diff --git a/contributors/emails/fatbigpig979@gmail.com b/contributors/emails/fatbigpig979@gmail.com new file mode 100644 index 0000000000..b28012a0aa --- /dev/null +++ b/contributors/emails/fatbigpig979@gmail.com @@ -0,0 +1 @@ +ddy4633 diff --git a/contributors/emails/gokhansarapevi@gmail.com b/contributors/emails/gokhansarapevi@gmail.com new file mode 100644 index 0000000000..12b6e83492 --- /dev/null +++ b/contributors/emails/gokhansarapevi@gmail.com @@ -0,0 +1,2 @@ +UltraInstinct0x +# v0.20.0 audit: co-author on #75565 (TUI paste tokens) diff --git a/contributors/emails/haowang@HaodeMac-mini.lan b/contributors/emails/haowang@HaodeMac-mini.lan new file mode 100644 index 0000000000..169ebaf2fe --- /dev/null +++ b/contributors/emails/haowang@HaodeMac-mini.lan @@ -0,0 +1 @@ +HAOWANG116 diff --git a/contributors/emails/harshkamdar67@gmail.com b/contributors/emails/harshkamdar67@gmail.com new file mode 100644 index 0000000000..534cf8d992 --- /dev/null +++ b/contributors/emails/harshkamdar67@gmail.com @@ -0,0 +1,2 @@ +Harshkamdar67 +# v0.20.0 audit: co-author on #72240 (/diff) diff --git a/contributors/emails/hinablue@gmail.com b/contributors/emails/hinablue@gmail.com new file mode 100644 index 0000000000..e59916db63 --- /dev/null +++ b/contributors/emails/hinablue@gmail.com @@ -0,0 +1,2 @@ +hinablue +# v0.20.0 audit: direct email match; co-author on #77022 diff --git a/contributors/emails/jeff.mettel@gmail.com b/contributors/emails/jeff.mettel@gmail.com new file mode 100644 index 0000000000..58e057de2c --- /dev/null +++ b/contributors/emails/jeff.mettel@gmail.com @@ -0,0 +1 @@ +jeff-mettel diff --git a/contributors/emails/jfmusa2024@gmail.com b/contributors/emails/jfmusa2024@gmail.com new file mode 100644 index 0000000000..4baa49434a --- /dev/null +++ b/contributors/emails/jfmusa2024@gmail.com @@ -0,0 +1,2 @@ +jfmusa2024-cyber +# v0.20.0 audit: author on #72388 salvage (desktop perf) diff --git a/contributors/emails/jquesnelle@gmail.com b/contributors/emails/jquesnelle@gmail.com new file mode 100644 index 0000000000..2ca552cf14 --- /dev/null +++ b/contributors/emails/jquesnelle@gmail.com @@ -0,0 +1,2 @@ +jquesnelle +# v0.20.0 audit: NeMo Relay revert/reapply cycle diff --git a/contributors/emails/jr.razmus@gmail.com b/contributors/emails/jr.razmus@gmail.com new file mode 100644 index 0000000000..2052dd5954 --- /dev/null +++ b/contributors/emails/jr.razmus@gmail.com @@ -0,0 +1 @@ +johnrazmus diff --git a/contributors/emails/jun@junho.co b/contributors/emails/jun@junho.co new file mode 100644 index 0000000000..83725166cd --- /dev/null +++ b/contributors/emails/jun@junho.co @@ -0,0 +1 @@ +junhohong diff --git a/contributors/emails/kingdomwarrior23@gmail.com b/contributors/emails/kingdomwarrior23@gmail.com new file mode 100644 index 0000000000..d76f896033 --- /dev/null +++ b/contributors/emails/kingdomwarrior23@gmail.com @@ -0,0 +1,2 @@ +Kingdomwarrior23 +# v0.20.0 audit: co-author on #69884 (provider onboarding fix) diff --git a/contributors/emails/kshitij@k4poor.dev b/contributors/emails/kshitij@k4poor.dev new file mode 100644 index 0000000000..c7510483c0 --- /dev/null +++ b/contributors/emails/kshitij@k4poor.dev @@ -0,0 +1 @@ +kshitijk4poor diff --git a/contributors/emails/lexharddrive69@gmail.com b/contributors/emails/lexharddrive69@gmail.com new file mode 100644 index 0000000000..4ec414fb39 --- /dev/null +++ b/contributors/emails/lexharddrive69@gmail.com @@ -0,0 +1 @@ +hdd69 diff --git a/contributors/emails/maartendormenatteysen@hotmail.com b/contributors/emails/maartendormenatteysen@hotmail.com new file mode 100644 index 0000000000..ec18c3e664 --- /dev/null +++ b/contributors/emails/maartendormenatteysen@hotmail.com @@ -0,0 +1 @@ +MaartenDMT diff --git a/contributors/emails/marzukia@users.noreply.github.com b/contributors/emails/marzukia@users.noreply.github.com new file mode 100644 index 0000000000..d030374ef7 --- /dev/null +++ b/contributors/emails/marzukia@users.noreply.github.com @@ -0,0 +1 @@ +marzukia diff --git a/contributors/emails/mrgraphitem@gmail.com b/contributors/emails/mrgraphitem@gmail.com new file mode 100644 index 0000000000..0f672e0931 --- /dev/null +++ b/contributors/emails/mrgraphitem@gmail.com @@ -0,0 +1 @@ +myk0la-b diff --git a/contributors/emails/pooyan6@gmail.com b/contributors/emails/pooyan6@gmail.com new file mode 100644 index 0000000000..c9ce11133b --- /dev/null +++ b/contributors/emails/pooyan6@gmail.com @@ -0,0 +1 @@ +pooyan6 diff --git a/contributors/emails/rod.boev@gmail.com b/contributors/emails/rod.boev@gmail.com index 8c73adc7cb..7c72a3f973 100644 --- a/contributors/emails/rod.boev@gmail.com +++ b/contributors/emails/rod.boev@gmail.com @@ -1,2 +1 @@ rodboev -# PR #37611 salvage (prompt-caching: tool schema markers) diff --git a/contributors/emails/rod@nxtlevel.dev b/contributors/emails/rod@nxtlevel.dev new file mode 100644 index 0000000000..39197517d7 --- /dev/null +++ b/contributors/emails/rod@nxtlevel.dev @@ -0,0 +1,2 @@ +rod-nxtlevel +# v0.20.0 audit: co-author on #72835 (remote backend) diff --git a/contributors/emails/rodrigo@nxtlevelsaas.com b/contributors/emails/rodrigo@nxtlevelsaas.com new file mode 100644 index 0000000000..8fe1be3c11 --- /dev/null +++ b/contributors/emails/rodrigo@nxtlevelsaas.com @@ -0,0 +1,2 @@ +rod-nxtlevel +# v0.20.0 audit: direct email match on profile diff --git a/contributors/emails/rsayar@uvic.ca b/contributors/emails/rsayar@uvic.ca new file mode 100644 index 0000000000..f4280a6044 --- /dev/null +++ b/contributors/emails/rsayar@uvic.ca @@ -0,0 +1,2 @@ +rsayar +# v0.20.0 audit: co-author on #71184 (terminal error frames) diff --git a/contributors/emails/soundbrokaz@kakao.com b/contributors/emails/soundbrokaz@kakao.com new file mode 100644 index 0000000000..eff0740de5 --- /dev/null +++ b/contributors/emails/soundbrokaz@kakao.com @@ -0,0 +1 @@ +JeremyDev87 diff --git a/contributors/emails/subhoya@gmail.com b/contributors/emails/subhoya@gmail.com new file mode 100644 index 0000000000..fdc4b0265b --- /dev/null +++ b/contributors/emails/subhoya@gmail.com @@ -0,0 +1,2 @@ +subhoya +# v0.20.0 audit: author on #72388 salvage (terminal overlay) diff --git a/contributors/emails/szzhoujiarui@users.noreply.github.com b/contributors/emails/szzhoujiarui@users.noreply.github.com new file mode 100644 index 0000000000..1d5deb6314 --- /dev/null +++ b/contributors/emails/szzhoujiarui@users.noreply.github.com @@ -0,0 +1 @@ +szzhoujiarui diff --git a/contributors/emails/tbsonline@protonmail.com b/contributors/emails/tbsonline@protonmail.com new file mode 100644 index 0000000000..d5156344cc --- /dev/null +++ b/contributors/emails/tbsonline@protonmail.com @@ -0,0 +1 @@ +jasoisjaso diff --git a/contributors/emails/universeszym@mail.ustc.edu.cn b/contributors/emails/universeszym@mail.ustc.edu.cn new file mode 100644 index 0000000000..5769e23c6b --- /dev/null +++ b/contributors/emails/universeszym@mail.ustc.edu.cn @@ -0,0 +1 @@ +Johnny-xuan diff --git a/contributors/emails/unixwzrd.register@mac.com b/contributors/emails/unixwzrd.register@mac.com new file mode 100644 index 0000000000..252eb6b1b4 --- /dev/null +++ b/contributors/emails/unixwzrd.register@mac.com @@ -0,0 +1 @@ +unixwzrd diff --git a/contributors/emails/vanshgilhotra8885@gmail.com b/contributors/emails/vanshgilhotra8885@gmail.com new file mode 100644 index 0000000000..990214f182 --- /dev/null +++ b/contributors/emails/vanshgilhotra8885@gmail.com @@ -0,0 +1 @@ +Vansh5632 diff --git a/contributors/emails/vitor@vitorcepedalopes.com b/contributors/emails/vitor@vitorcepedalopes.com new file mode 100644 index 0000000000..8486daba61 --- /dev/null +++ b/contributors/emails/vitor@vitorcepedalopes.com @@ -0,0 +1,2 @@ +TheAngryPit +# v0.20.0 audit: direct email match on profile (OAuth sidebar fix) diff --git a/contributors/emails/vittoria3103.123@gmail.com b/contributors/emails/vittoria3103.123@gmail.com new file mode 100644 index 0000000000..a2b1bf1db6 --- /dev/null +++ b/contributors/emails/vittoria3103.123@gmail.com @@ -0,0 +1 @@ +VittoriaLanzo diff --git a/contributors/emails/wangyunyou@leoao.com b/contributors/emails/wangyunyou@leoao.com new file mode 100644 index 0000000000..e9741c8cba --- /dev/null +++ b/contributors/emails/wangyunyou@leoao.com @@ -0,0 +1 @@ +wangyunyou diff --git a/contributors/emails/wykim777@naver.com b/contributors/emails/wykim777@naver.com new file mode 100644 index 0000000000..bc55e275b8 --- /dev/null +++ b/contributors/emails/wykim777@naver.com @@ -0,0 +1,2 @@ +iamwongeeeee +# v0.20.0 audit: co-author on #65798 (FTS schema v23) diff --git a/contributors/emails/xaydinoktay@gmail.com b/contributors/emails/xaydinoktay@gmail.com new file mode 100644 index 0000000000..95414bf6c5 --- /dev/null +++ b/contributors/emails/xaydinoktay@gmail.com @@ -0,0 +1 @@ +aydnOktay diff --git a/contributors/emails/zabih.mosafer@gmail.com b/contributors/emails/zabih.mosafer@gmail.com new file mode 100644 index 0000000000..2e1ef5f28c --- /dev/null +++ b/contributors/emails/zabih.mosafer@gmail.com @@ -0,0 +1 @@ +zabih-sudo diff --git a/cron/jobs.py b/cron/jobs.py index 904b8aa346..95020ae659 100644 --- a/cron/jobs.py +++ b/cron/jobs.py @@ -1738,8 +1738,19 @@ def mark_job_run(job_id: str, success: bool, error: Optional[str] = None, # Check if we've hit the repeat limit if times is not None and times > 0 and completed >= times: - # Remove the job (limit reached) - jobs.pop(i) + # Limit reached: retain the record as a terminal + # completion instead of popping it. Deleting the job + # here discarded the last_status / last_error / + # last_delivery_error written above — a finished + # one-shot vanished from `cronjob list` with no + # inspectable outcome, and a failed delivery was + # invisible. Mirror the terminal shape of the + # next_run_at-is-None branch below; the retention + # sweep prunes these after + # COMPLETED_ONESHOT_RETENTION_DAYS. + job["enabled"] = False + job["state"] = "completed" + job["next_run_at"] = None save_jobs(jobs) return @@ -1857,13 +1868,29 @@ def claim_dispatch(job_id: str) -> bool: return True # infinite — always dispatch completed = repeat.get("completed", 0) if completed >= times: - # Already dispatched the max number of times (e.g. a prior - # tick claimed then died before mark_job_run could remove it). - # Clean up so it stops appearing as due on every tick. + # Already dispatched the max number of times. + if job.get("last_run_at") is not None: + # A prior run completed normally (e.g. mark_job_run raced + # with this tick). Retain the terminal record — same shape + # as mark_job_run's repeat-limit branch — instead of + # deleting the job and its final status/delivery error. + job["enabled"] = False + job["state"] = "completed" + job["next_run_at"] = None + save_jobs(jobs) + logger.info( + "Job '%s': dispatch limit reached (%d/%d) — marking completed", + job.get("name", job.get("id", "?")), + completed, + times, + ) + return False + # A prior tick claimed the dispatch then died before the run + # completed (#73973) — a genuinely wedged claim. Remove it so + # it stops appearing as due, and leave an operator-visible + # diagnostic instead of vanishing silently. jobs.pop(i) save_jobs(jobs) - # If the claimed run never completed (#73973), leave an - # operator-visible diagnostic instead of vanishing silently. _write_wedged_oneshot_diagnostic(job) logger.info( "Job '%s': dispatch limit reached (%d/%d) — removing", @@ -2049,6 +2076,82 @@ def claim_job_for_fire(job_id: str, *, claim_ttl_seconds: int = 300) -> bool: return False +# Completed one-shot job records are retained in jobs.json (final status + +# delivery error stay inspectable via `cronjob list`) instead of being deleted +# at completion, then pruned by _sweep_completed_oneshots once they age out. +COMPLETED_ONESHOT_RETENTION_DAYS = 7 + + +def _completed_oneshot_retention_days() -> float: + """Resolve the completed one-shot retention window from config. + + ``cron.completed_retention_days`` (number, default + ``COMPLETED_ONESHOT_RETENTION_DAYS``). A non-positive value disables the + sweep, retaining completed one-shot records indefinitely. + """ + try: + from hermes_cli.config import load_config + cfg = load_config() or {} + cron_cfg = cfg.get("cron", {}) if isinstance(cfg, dict) else {} + return float( + cron_cfg.get( + "completed_retention_days", COMPLETED_ONESHOT_RETENTION_DAYS + ) + ) + except Exception: + return float(COMPLETED_ONESHOT_RETENTION_DAYS) + + +def _sweep_completed_oneshots(raw_jobs: List[Dict[str, Any]], now: datetime) -> bool: + """Prune terminal ``state == "completed"`` one-shot records past retention. + + Mutates *raw_jobs* in place; returns True when anything was removed (the + caller persists). Only one-shot (``schedule.kind == "once"``) records in + the terminal completed state are candidates; recurring jobs and non- + terminal one-shots are never touched. Age is measured from + ``last_run_at`` — a completed record without a parseable ``last_run_at`` + is kept (never guess a record into deletion). + """ + retention_days = _completed_oneshot_retention_days() + if retention_days <= 0: + return False + cutoff = now - timedelta(days=retention_days) + removed = False + for rj in list(raw_jobs): + try: + if rj.get("state") != "completed": + continue + schedule = rj.get("schedule") + kind = schedule.get("kind") if isinstance(schedule, dict) else None + if kind != "once": + continue + last_run = rj.get("last_run_at") + if not isinstance(last_run, str): + continue + try: + last_run_dt = _ensure_aware(datetime.fromisoformat(last_run)) + except Exception: + continue + if last_run_dt >= cutoff: + continue + raw_jobs.remove(rj) + removed = True + logger.info( + "Job '%s': pruning completed one-shot record " + "(finished %s, retention %.1f days)", + rj.get("name", rj.get("id", "?")), + last_run, + retention_days, + ) + except Exception: + logger.debug( + "Retention sweep skipped malformed job record %r", + rj.get("id", "?"), + exc_info=True, + ) + return removed + + def get_due_jobs() -> List[Dict[str, Any]]: """Get all jobs that are due to run now. @@ -2168,6 +2271,15 @@ def _get_due_jobs_locked() -> List[Dict[str, Any]]: # (derived from HERMES_CRON_TIMEOUT). See _oneshot_run_claim_ttl_seconds. _run_claim_ttl = _oneshot_run_claim_ttl_seconds() + # Retention sweep: completed one-shots are retained (so their final + # status / delivery error stay inspectable via `cronjob list`) instead of + # being deleted on completion, but they must not accumulate in jobs.json + # forever. Prune terminal one-shot records older than the retention + # window each scan. + if _sweep_completed_oneshots(raw_jobs, now): + needs_save = True + jobs = [j for j in jobs if any(rj.get("id") == j.get("id") for rj in raw_jobs)] + for job in jobs: # Per-job containment (structural guard): one malformed or # unexpected job record must never abort the whole scan. The id / diff --git a/cron/lifecycle_guard.py b/cron/lifecycle_guard.py index b09635c998..6c7a5eaad0 100644 --- a/cron/lifecycle_guard.py +++ b/cron/lifecycle_guard.py @@ -219,8 +219,14 @@ def _iter_referenced_shell_scripts( yield _resolve_terminal_script_path(arguments[arg_index], cwd) continue - if "/" in executable or executable.endswith((".sh", ".bash", ".zsh")): - yield _resolve_terminal_script_path(executable, cwd) + # A bare "/" token is pathlib's division operator in Python sources + # (e.g. `Path.home() / ".hermes"`), not an executable reference. + # Resolving it walks to the filesystem root and fails the + # regular-file check below, hard-blocking innocent .py scripts + # (#77131). Skip pure-separator tokens. + if executable.strip("/"): + if "/" in executable or executable.endswith((".sh", ".bash", ".zsh")): + yield _resolve_terminal_script_path(executable, cwd) def _iter_shell_command_payloads(command: str) -> Iterator[str]: @@ -258,13 +264,22 @@ def _read_referenced_script(path: Path) -> tuple[Optional[str], bool]: metadata = os.fstat(descriptor) if not stat.S_ISREG(metadata.st_mode): return None, True - if metadata.st_size > _MAX_REFERENCED_SCRIPT_BYTES: - return None, True + # Read a bounded chunk first — even for oversized files, the first + # chunk tells us if this is a binary (NUL bytes) that should be + # skipped as "nothing to scan" rather than failing closed (#76762). data = os.read(descriptor, _MAX_REFERENCED_SCRIPT_BYTES + 1) except OSError: return None, False finally: os.close(descriptor) + # A NUL byte in the first chunk means this is a binary (ELF/Mach-O/ + # PE), not a shell script — scanning its decoded contents would + # tokenize machine code and feed junk paths into the recursion + # (including a `ValueError: embedded null byte` from Path.resolve, + # #76762). Treat it as "nothing to scan" rather than unsafe: a binary + # executed by the user is not a referenced *shell script*. + if b"\x00" in data: + return None, False if len(data) > _MAX_REFERENCED_SCRIPT_BYTES: return None, True return data.decode("utf-8", errors="replace"), False @@ -298,7 +313,10 @@ def _contains_unsafe_gateway_action( for script_path in _iter_referenced_shell_scripts(command, cwd=cwd): try: resolved = script_path.resolve(strict=False) - except OSError: + except (OSError, ValueError): + # OSError: unreadable/long paths. ValueError: embedded NUL byte + # from a binary's decoded contents tokenized as a path — a + # guarded path must never crash the guard (#76762). resolved = script_path if resolved in visited: continue @@ -392,16 +410,31 @@ def check_gateway_lifecycle( surfaces this as a tool error; the CLI prints it in red and exits 1). """ combined = prompt or "" + python_script = False if script: + python_script = _resolve_script_path(script).suffix == ".py" script_text = _read_script_for_scanning(script) if script_text: combined = f"{combined}\n{script_text}" - script_dir = _resolve_script_directory(script) if script else None - if contains_gateway_lifecycle_command_or_referenced_script( - combined, - cwd=script_dir, - ): + if python_script: + # Python is executed by the interpreter, never through a POSIX + # shell: the shell-script reference walk is a false-positive + # generator on Python sources (pathlib's "/" operator resolves to + # the filesystem root and trips the regular-file check, blocking + # every innocent .py cron script, #77131). The direct command + # regex below still scans the full text, so a literal + # `hermes gateway restart` embedded in a .py script is still + # blocked. Non-regular/oversized script files still fail closed + # via the lifecycle-shaped sentinel in _read_script_for_scanning. + unsafe = contains_gateway_lifecycle_command(combined) + else: + script_dir = _resolve_script_directory(script) if script else None + unsafe = contains_gateway_lifecycle_command_or_referenced_script( + combined, + cwd=script_dir, + ) + if unsafe: raise GatewayLifecycleBlocked( "Blocked: cron job contains a gateway lifecycle command or persistent " "launchctl submit operation. This is blocked to prevent agent-driven " diff --git a/cron/scheduler.py b/cron/scheduler.py index 6e102b18d5..f2ab51919a 100644 --- a/cron/scheduler.py +++ b/cron/scheduler.py @@ -4158,8 +4158,20 @@ def tick( due_jobs = get_due_jobs() - if verbose and not due_jobs: - logger.info("%s - No jobs due", _hermes_now().strftime('%H:%M:%S')) + if not due_jobs: + # Idle tick: skip config load + pool partitioning entirely + # (#33612 — the gateway ticker calls tick(verbose=False) every + # 60s, so idle ticks previously fell through to load_config()). + # Still run the post-tick MCP orphan sweep: main intentionally + # sweeps on idle ticks so orphaned stdio children from crashed + # jobs are reaped even when nothing is due. + if verbose: + logger.info("%s - No jobs due", _hermes_now().strftime('%H:%M:%S')) + try: + from tools.mcp_tool import _kill_orphaned_mcp_children + _kill_orphaned_mcp_children() + except Exception as _e: + logger.debug("Post-tick MCP orphan cleanup failed: %s", _e) return 0 if verbose: diff --git a/gateway/platforms/api_server.py b/gateway/platforms/api_server.py index 5b70648246..a756134d99 100644 --- a/gateway/platforms/api_server.py +++ b/gateway/platforms/api_server.py @@ -158,6 +158,54 @@ RESPONSES_AUTO_TRUNCATION_HISTORY_LIMIT = 100 _COMPRESSED_SUMMARY_METADATA_KEY = "_compressed_summary" +class ThreadSafeAsyncQueue(asyncio.Queue): + """An ``asyncio.Queue`` that a non-loop thread can push into safely. + + The SSE writers' streaming loops used to bridge a plain ``queue.Queue`` + into the event loop via ``await loop.run_in_executor(None, lambda: + stream_q.get(timeout=0.5))`` inside a ``while True`` poll — a thread-pool + round trip on every 0.5s tick even when idle, plus up to 500ms of tail + latency between a delta landing in the queue and it reaching the + response. ``run_conversation`` itself runs on a worker thread (via + ``loop.run_in_executor``), so its ``stream_delta_callback`` closures + (``_on_delta`` etc.) call ``put_threadsafe`` from off the loop thread; + the consumer side just does a plain ``await queue.get()``/ + ``asyncio.wait_for(queue.get(), timeout=...)``, woken immediately by + ``call_soon_threadsafe`` instead of polling. + """ + + def put_threadsafe(self, item, *, loop: asyncio.AbstractEventLoop = None) -> None: + (loop or self._loop_ref).call_soon_threadsafe(self.put_nowait, item) + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + # Always constructed inside a running async handler (the SSE + # request handlers below), so get_running_loop() is safe here. + self._loop_ref = asyncio.get_running_loop() + + +def _sse_frame(data: Any, *, event: str = None, ensure_ascii: bool = True) -> bytes: + """Encode one SSE frame: optional ``event:`` line, then ``data: \n\n``. + + The single source of truth for SSE frame serialization across every + streaming writer in this module — ``_write_sse_chat_completion`` (the + five call sites it was first extracted from), ``_write_sse_responses``'s + inner ``_write_event`` closure, and the ``/v1/runs`` event stream. All + three used the identical ``json.dumps(data)`` / ``json.dumps(..., + ensure_ascii=False)`` + ``"\\ndata: ...\\n\\n"`` shape; routing them all + through here keeps the on-the-wire format in exactly one place. + + ``ensure_ascii`` defaults to ``True``, byte-identical to a bare + ``json.dumps(data)``. Callers that must preserve raw non-ASCII bytes on + the wire (the Responses-API writer historically used + ``ensure_ascii=False``) pass ``ensure_ascii=False`` explicitly — the + option exists so every writer shares one helper without changing any + existing byte stream. + """ + prefix = f"event: {event}\n" if event else "" + return f"{prefix}data: {json.dumps(data, ensure_ascii=ensure_ascii)}\n\n".encode() + + def _coerce_port(value: Any, default: int = DEFAULT_PORT) -> int: """Parse a listen port without letting malformed env/config values crash startup.""" try: @@ -3095,6 +3143,7 @@ class APIServerAdapter(BasePlatformAdapter): _get_effective_configurable_toolsets, _get_platform_tools, _toolset_has_keys, + get_nous_subscription_features, ) from toolsets import resolve_toolset @@ -3104,6 +3153,7 @@ class APIServerAdapter(BasePlatformAdapter): "api_server", include_default_mcp_servers=False, ) + features = get_nous_subscription_features(config) data: List[Dict[str, Any]] = [] for name, label, desc in _get_effective_configurable_toolsets(): try: @@ -3116,7 +3166,7 @@ class APIServerAdapter(BasePlatformAdapter): "label": label, "description": desc, "enabled": is_enabled, - "configured": _toolset_has_keys(name, config), + "configured": _toolset_has_keys(name, config, features=features), "tools": tools, }) except Exception: @@ -3782,8 +3832,7 @@ class APIServerAdapter(BasePlatformAdapter): if item is None: break name, payload = item - data = json.dumps(payload, ensure_ascii=False) - await response.write(f"event: {name}\ndata: {data}\n\n".encode("utf-8")) + await response.write(_sse_frame(payload, event=name, ensure_ascii=False)) except (asyncio.CancelledError, ConnectionResetError): task.cancel() raise @@ -3983,8 +4032,7 @@ class APIServerAdapter(BasePlatformAdapter): return web.json_response(_openai_error(selection_error), status=400) if stream: - import queue as _q - _stream_q: _q.Queue = _q.Queue() + _stream_q = ThreadSafeAsyncQueue() def _on_delta(delta): # Filter out None — the agent fires stream_delta_callback(None) @@ -3994,8 +4042,10 @@ class APIServerAdapter(BasePlatformAdapter): # response, causing Open WebUI (and similar frontends) to miss # the final answer after tool calls. The SSE loop detects # completion via agent_task.done() instead. + # Called from the worker thread running run_conversation — + # put_threadsafe (not put_nowait) is required here. if delta is not None: - _stream_q.put(delta) + _stream_q.put_threadsafe(delta) # Track which tool_call_ids we've emitted a "running" lifecycle # event for, so a "completed" event without a matching "running" @@ -4021,7 +4071,7 @@ class APIServerAdapter(BasePlatformAdapter): _started_tool_call_ids.add(tool_call_id) from agent.display import build_tool_preview, get_tool_emoji label = build_tool_preview(function_name, function_args) or function_name - _stream_q.put(("__tool_progress__", { + _stream_q.put_threadsafe(("__tool_progress__", { "tool": function_name, "emoji": get_tool_emoji(function_name), "label": label, @@ -4039,7 +4089,7 @@ class APIServerAdapter(BasePlatformAdapter): if not tool_call_id or tool_call_id not in _started_tool_call_ids: return _started_tool_call_ids.discard(tool_call_id) - _stream_q.put(("__tool_progress__", { + _stream_q.put_threadsafe(("__tool_progress__", { "tool": function_name, "toolCallId": tool_call_id, "status": "completed", @@ -4069,7 +4119,7 @@ class APIServerAdapter(BasePlatformAdapter): )) # Ensure SSE drain loops can terminate without relying on polling # agent_task.done(), which can race with queue timeout checks. - agent_task.add_done_callback(lambda _fut: _stream_q.put(None)) + agent_task.add_done_callback(lambda _fut: _stream_q.put_nowait(None)) return await self._write_sse_chat_completion( request, completion_id, model_name, created, _stream_q, @@ -4204,8 +4254,6 @@ class APIServerAdapter(BasePlatformAdapter): the agent is interrupted via ``agent.interrupt()`` so it stops making LLM API calls, and the asyncio task wrapper is cancelled. """ - import queue as _q - sse_headers = { "Content-Type": "text/event-stream", "Cache-Control": "no-cache", @@ -4233,7 +4281,7 @@ class APIServerAdapter(BasePlatformAdapter): "created": created, "model": model, "choices": [{"index": 0, "delta": {"role": "assistant"}, "finish_reason": None}], } - await response.write(f"data: {json.dumps(role_chunk)}\n\n".encode()) + await response.write(_sse_frame(role_chunk)) last_activity = time.monotonic() # Helper — route a queue item to the correct SSE event. @@ -4248,25 +4296,24 @@ class APIServerAdapter(BasePlatformAdapter): #16588 for the ``toolCallId``/``status`` lifecycle fields. """ if isinstance(item, tuple) and len(item) == 2 and item[0] == "__tool_progress__": - event_data = json.dumps(item[1]) - await response.write( - f"event: hermes.tool.progress\ndata: {event_data}\n\n".encode() - ) + await response.write(_sse_frame(item[1], event="hermes.tool.progress")) else: content_chunk = { "id": completion_id, "object": "chat.completion.chunk", "created": created, "model": model, "choices": [{"index": 0, "delta": {"content": item}, "finish_reason": None}], } - await response.write(f"data: {json.dumps(content_chunk)}\n\n".encode()) + await response.write(_sse_frame(content_chunk)) return time.monotonic() - # Stream content chunks as they arrive from the agent - loop = asyncio.get_running_loop() + # Stream content chunks as they arrive from the agent. Woken + # directly by put_threadsafe's call_soon_threadsafe — no + # executor hop, no poll-interval latency (see + # ThreadSafeAsyncQueue's docstring). while True: try: - delta = await loop.run_in_executor(None, lambda: stream_q.get(timeout=0.5)) - except _q.Empty: + delta = await asyncio.wait_for(stream_q.get(), timeout=0.5) + except asyncio.TimeoutError: if agent_task.done(): # Drain any remaining items while True: @@ -4275,7 +4322,7 @@ class APIServerAdapter(BasePlatformAdapter): if delta is None: break last_activity = await _emit(delta) - except _q.Empty: + except asyncio.QueueEmpty: break break if time.monotonic() - last_activity >= CHAT_COMPLETIONS_SSE_KEEPALIVE_SECONDS: @@ -4351,7 +4398,7 @@ class APIServerAdapter(BasePlatformAdapter): "error": err_msg, "error_code": "output_truncated" if finish_reason == "length" else "agent_error", } - await response.write(f"data: {json.dumps(finish_chunk)}\n\n".encode()) + await response.write(_sse_frame(finish_chunk)) await response.write(b"data: [DONE]\n\n") except (ConnectionResetError, ConnectionAbortedError, BrokenPipeError, OSError): # Client disconnected mid-stream. Interrupt the agent so it @@ -4383,7 +4430,7 @@ class APIServerAdapter(BasePlatformAdapter): "created": created, "model": model, "choices": [{"index": 0, "delta": {}, "finish_reason": "error"}], } - await response.write(f"data: {json.dumps(error_chunk)}\n\n".encode()) + await response.write(_sse_frame(error_chunk)) await response.write(b"data: [DONE]\n\n") except Exception: pass @@ -4435,8 +4482,6 @@ class APIServerAdapter(BasePlatformAdapter): ``previous_response_id`` chaining still have something to recover from. """ - import queue as _q - sse_headers = { "Content-Type": "text/event-stream", "Cache-Control": "no-cache", @@ -4481,8 +4526,7 @@ class APIServerAdapter(BasePlatformAdapter): if "sequence_number" not in data: data["sequence_number"] = sequence_number sequence_number += 1 - payload = f"event: {event_type}\ndata: {json.dumps(data)}\n\n" - await response.write(payload.encode()) + await response.write(_sse_frame(data, event=event_type)) def _envelope(status: str) -> Dict[str, Any]: env: Dict[str, Any] = { @@ -4757,11 +4801,10 @@ class APIServerAdapter(BasePlatformAdapter): _batch_buf = [] await _emit_text_delta(combined) - loop = asyncio.get_running_loop() while True: try: - item = await loop.run_in_executor(None, lambda: stream_q.get(timeout=0.5)) - except _q.Empty: + item = await asyncio.wait_for(stream_q.get(), timeout=0.5) + except asyncio.TimeoutError: if agent_task.done(): # Drain remaining while True: @@ -4771,7 +4814,7 @@ class APIServerAdapter(BasePlatformAdapter): break await _dispatch(item) last_activity = time.monotonic() - except _q.Empty: + except asyncio.QueueEmpty: break break if time.monotonic() - last_activity >= CHAT_COMPLETIONS_SSE_KEEPALIVE_SECONDS: @@ -5133,15 +5176,16 @@ class APIServerAdapter(BasePlatformAdapter): # Streaming branch — emit OpenAI Responses SSE events as the # agent runs so frontends can render text deltas and tool # calls in real time. See _write_sse_responses for details. - import queue as _q - _stream_q: _q.Queue = _q.Queue() + _stream_q = ThreadSafeAsyncQueue() def _on_delta(delta): # None from the agent is a CLI box-close signal, not EOS. # Forwarding would kill the SSE stream prematurely; the # SSE writer detects completion via agent_task.done(). + # Called from the worker thread running run_conversation — + # put_threadsafe (not put_nowait) is required here. if delta is not None: - _stream_q.put(delta) + _stream_q.put_threadsafe(delta) def _on_tool_progress(event_type, name, preview, args, **kwargs): """Queue non-start tool progress events if needed in future. @@ -5154,7 +5198,7 @@ class APIServerAdapter(BasePlatformAdapter): def _on_tool_start(tool_call_id, function_name, function_args): """Queue a started tool for live function_call streaming.""" - _stream_q.put(("__tool_started__", { + _stream_q.put_threadsafe(("__tool_started__", { "tool_call_id": tool_call_id, "name": function_name, "arguments": function_args or {}, @@ -5162,7 +5206,7 @@ class APIServerAdapter(BasePlatformAdapter): def _on_tool_complete(tool_call_id, function_name, function_args, function_result): """Queue a completed tool result for live function_call_output streaming.""" - _stream_q.put(("__tool_completed__", { + _stream_q.put_threadsafe(("__tool_completed__", { "tool_call_id": tool_call_id, "name": function_name, "arguments": function_args or {}, @@ -5186,7 +5230,7 @@ class APIServerAdapter(BasePlatformAdapter): )) # Ensure SSE drain loops can terminate without relying on polling # agent_task.done(), which can race with queue timeout checks. - agent_task.add_done_callback(lambda _fut: _stream_q.put(None)) + agent_task.add_done_callback(lambda _fut: _stream_q.put_nowait(None)) response_id = f"resp_{uuid.uuid4().hex[:28]}" model_name = body.get("model", self._model_name) @@ -6713,8 +6757,8 @@ class APIServerAdapter(BasePlatformAdapter): # Run finished — send final SSE comment and close await response.write(b": stream closed\n\n") break - payload = f"data: {json.dumps(event)}\n\n" - await response.write(payload.encode()) + payload = _sse_frame(event) + await response.write(payload) except Exception as exc: logger.debug("[api_server] SSE stream error for run %s: %s", run_id, exc) finally: diff --git a/gateway/platforms/base.py b/gateway/platforms/base.py index de60322718..c42b916073 100644 --- a/gateway/platforms/base.py +++ b/gateway/platforms/base.py @@ -541,7 +541,7 @@ import dataclasses from dataclasses import dataclass, field from datetime import datetime from pathlib import Path -from typing import Dict, List, Optional, Any, Callable, Awaitable, Tuple, Union +from typing import TYPE_CHECKING, Dict, List, Optional, Any, Callable, Awaitable, Tuple, Union from enum import Enum from pathlib import Path as _Path @@ -551,6 +551,9 @@ from gateway.config import Platform, PlatformConfig from gateway.session import SessionSource, build_session_key from hermes_constants import get_default_hermes_root, get_hermes_dir, get_hermes_home +if TYPE_CHECKING: + from agent.display import ToolPreview + # --------------------------------------------------------------------------- # Streaming TTS format descriptor and handle (#60671) @@ -3095,12 +3098,28 @@ class BasePlatformAdapter(ABC): # progress bubbles compact — they persist as permanent messages). preview = event.preview if preview: + from agent.display import prepare_tool_preview + cap = preview_max_len if preview_max_len > 0 else 40 - if len(preview) > cap: - preview = preview[:cap - 3] + "..." - return f"{emoji} {event.tool_name}: \"{preview}\"" + prepared = prepare_tool_preview( + event.tool_name, + event.args, + fallback=preview, + max_len=cap, + ) + rendered = self.format_tool_preview(prepared) + return f"{emoji} {event.tool_name}: \"{rendered}\"" return f"{emoji} {event.tool_name}..." + def format_tool_preview(self, preview: "ToolPreview") -> str: + """Apply platform-native formatting to a compact tool preview. + + Most adapters only need the compact text. Rich-text adapters can use + the preview's explicit metadata to preserve details such as a URL that + was shortened for display. + """ + return preview.text + @property def has_fatal_error(self) -> bool: return self._fatal_error_message is not None @@ -6047,13 +6066,14 @@ class BasePlatformAdapter(ABC): record_obligation, ) - if ledger_enabled(): + if await asyncio.to_thread(ledger_enabled): _obligation_id = compute_obligation_id( session_key, str(getattr(event, "message_id", "") or ""), text_content, ) - record_obligation( + await asyncio.to_thread( + record_obligation, obligation_id=_obligation_id, session_key=session_key, platform=str( @@ -6064,7 +6084,7 @@ class BasePlatformAdapter(ABC): thread_id=getattr(event.source, "thread_id", None), content=text_content, ) - mark_attempting(_obligation_id) + await asyncio.to_thread(mark_attempting, _obligation_id) except Exception: logger.debug("delivery ledger record failed", exc_info=True) _obligation_id = None @@ -6083,9 +6103,10 @@ class BasePlatformAdapter(ABC): ) if getattr(result, "success", False): - mark_delivered(_obligation_id) + await asyncio.to_thread(mark_delivered, _obligation_id) else: - mark_failed( + await asyncio.to_thread( + mark_failed, _obligation_id, str(getattr(result, "error", "") or ""), ) diff --git a/gateway/platforms/yuanbao.py b/gateway/platforms/yuanbao.py index ad5da06614..2e3ed0253b 100644 --- a/gateway/platforms/yuanbao.py +++ b/gateway/platforms/yuanbao.py @@ -4588,7 +4588,11 @@ class MessageSender: cached = self._adapter._member_cache.get(group_code) if cached: ts, member_list = cached - members = member_list if (time.time() - ts < self._adapter.MEMBER_CACHE_TTL_S) else [] + if time.time() - ts < self._adapter.MEMBER_CACHE_TTL_S: + members = member_list + else: + del self._adapter._member_cache[group_code] + members = [] else: members = [] if not members: @@ -5116,6 +5120,19 @@ class YuanbaoAdapter(BasePlatformAdapter): await super()._process_message_background(event, session_key) finally: self._outbound.cancel_slow_notifier(chat_id) + # Clear the RecallGuard tracking entries for this message only if + # our msg_id is still current. A concurrent pending message may + # have already overwritten the entry in _dispatch_inbound_event + # while we were running; in that case the drain task owns it and + # we must not clear it. Id-less events (internal/synthetic + # messages, pushes without a msg_id) never wrote a tracking entry + # in _dispatch_inbound_event, so they must never pop either — the + # entry they see belongs to a concurrently-queued id-bearing + # message whose drain task still needs it for recall matching. + msg_id = event.message_id + if msg_id and self._processing_msg_ids.get(session_key) == msg_id: + self._processing_msg_ids.pop(session_key, None) + self._processing_msg_texts.pop(session_key, None) # ------------------------------------------------------------------ # Group query (delegate to GroupQueryService) diff --git a/gateway/relay/adapter.py b/gateway/relay/adapter.py index d64e2fb00b..73f7cf2941 100644 --- a/gateway/relay/adapter.py +++ b/gateway/relay/adapter.py @@ -31,6 +31,14 @@ from gateway.session import SessionSource logger = logging.getLogger(__name__) +# Keep the drain-path going-idle ACK budget strictly under the runner's default +# adapter disconnect timeout (5s). If go_idle consumes the whole outer budget, +# cancellation can fire before transport.disconnect() and leave the websocket +# open. Paired with transport teardown budgets of 1s each for supervisor, +# reader, and ws.close (~3s), the full drain path stays inside 5s. +_RELAY_GO_IDLE_ON_DISCONNECT_TIMEOUT_S = 2.0 +_RELAY_REVOCATION_MONITOR_TEARDOWN_TIMEOUT_S = 1.0 + def _utf16_len(text: str) -> int: """Count UTF-16 code units (Telegram's length unit).""" @@ -816,8 +824,11 @@ class RelayAdapter(BasePlatformAdapter): if self._revocation_monitor is not None: self._revocation_monitor.cancel() try: - await self._revocation_monitor - except (asyncio.CancelledError, Exception): # noqa: BLE001 - best-effort teardown + await asyncio.wait_for( + self._revocation_monitor, + timeout=_RELAY_REVOCATION_MONITOR_TEARDOWN_TIMEOUT_S, + ) + except (asyncio.TimeoutError, asyncio.CancelledError, Exception): # noqa: BLE001 - best-effort teardown pass self._revocation_monitor = None if self._transport is not None: @@ -831,15 +842,32 @@ class RelayAdapter(BasePlatformAdapter): # the ack (Q-5.3c). Best-effort + guarded: a transport without go_idle # (the stub) or a failed/timed-out ack must not block shutdown — we # proceed to disconnect exactly as before, no regression. - go_idle = getattr(self._transport, "go_idle", None) - if callable(go_idle): + # + # transport.disconnect() runs in finally so an outer cancellation + # during go_idle (runner default adapter budget is 5s) still closes + # the socket/supervisor instead of leaking them. shield() keeps the + # teardown await itself from being cancelled mid-flight. + try: + go_idle = getattr(self._transport, "go_idle", None) + if callable(go_idle): + try: + result: Any = go_idle( + timeout_s=_RELAY_GO_IDLE_ON_DISCONNECT_TIMEOUT_S + ) + if asyncio.iscoroutine(result): + await result + except Exception: # noqa: BLE001 - going-idle is an optimization, never blocks drain + logger.debug( + "relay going_idle failed during drain", exc_info=True + ) + finally: try: - result: Any = go_idle() - if asyncio.iscoroutine(result): - await result - except Exception: # noqa: BLE001 - going-idle is an optimization, never blocks drain - logger.debug("relay going_idle failed during drain", exc_info=True) - await self._transport.disconnect() + await asyncio.shield(self._transport.disconnect()) + except Exception: # noqa: BLE001 - teardown must not block outer cancel propagation + logger.debug( + "relay transport disconnect failed during drain", + exc_info=True, + ) async def go_dormant(self) -> bool: """Quiesce the relay for a scale-to-zero suspend (D12 / Phase 0). diff --git a/gateway/relay/ws_transport.py b/gateway/relay/ws_transport.py index a4bfc4f6b2..07981a327e 100644 --- a/gateway/relay/ws_transport.py +++ b/gateway/relay/ws_transport.py @@ -53,6 +53,10 @@ WEBSOCKETS_AVAILABLE = websockets is not None # How long to wait for the handshake descriptor and for each outbound result. _HANDSHAKE_TIMEOUT_S = 30.0 _OUTBOUND_TIMEOUT_S = 30.0 +# Bound supervisor/reader/ws.close awaits so a wedged peer cannot stall +# adapter.disconnect. Three sequential awaits at 1.0s stay under the runner's +# default 5s adapter disconnect budget (plus the 2s go_idle ACK budget). +_TEARDOWN_AWAIT_TIMEOUT_S = 1.0 # Phase 7 Unit 7d-B: the application close code the connector sends when it # rejects/revokes a gateway's WS upgrade auth (mirrors the connector's @@ -501,23 +505,26 @@ class WebSocketRelayTransport: if self._supervisor is not None: self._supervisor.cancel() try: - await self._supervisor - except (asyncio.CancelledError, Exception): # noqa: BLE001 - best-effort teardown + await asyncio.wait_for( + self._supervisor, timeout=_TEARDOWN_AWAIT_TIMEOUT_S + ) + except (asyncio.TimeoutError, asyncio.CancelledError, Exception): # noqa: BLE001 - best-effort teardown pass self._supervisor = None if self._reader is not None: self._reader.cancel() try: - await self._reader - except (asyncio.CancelledError, Exception): # noqa: BLE001 - best-effort teardown + await asyncio.wait_for(self._reader, timeout=_TEARDOWN_AWAIT_TIMEOUT_S) + except (asyncio.TimeoutError, asyncio.CancelledError, Exception): # noqa: BLE001 - best-effort teardown pass self._reader = None if self._ws is not None: try: - await self._ws.close() - except Exception: # noqa: BLE001 + await asyncio.wait_for(self._ws.close(), timeout=_TEARDOWN_AWAIT_TIMEOUT_S) + except (asyncio.TimeoutError, asyncio.CancelledError, Exception): # noqa: BLE001 pass - self._ws = None + finally: + self._ws = None # Fail any in-flight outbound waiters so callers don't hang. for fut in self._pending.values(): if not fut.done(): @@ -663,9 +670,11 @@ class WebSocketRelayTransport: # flip us back to a fast reconnect. self._dormant = True try: - await self._ws.close() - except Exception: # noqa: BLE001 - best-effort; the reader still ends + arms reconnect - logger.debug("relay go_dormant: ws.close() raised", exc_info=True) + await asyncio.wait_for( + self._ws.close(), timeout=_TEARDOWN_AWAIT_TIMEOUT_S + ) + except (asyncio.TimeoutError, Exception): # noqa: BLE001 - best-effort; the reader still ends + arms reconnect + logger.debug("relay go_dormant: ws.close() raised or timed out", exc_info=True) return acked async def _send_inbound_ack(self, buffer_id: str) -> None: diff --git a/gateway/restart.py b/gateway/restart.py index 41a7484329..5a2cfbc1b7 100644 --- a/gateway/restart.py +++ b/gateway/restart.py @@ -24,6 +24,15 @@ DEFAULT_GATEWAY_RESTART_DRAIN_TIMEOUT = float( DEFAULT_CONFIG["agent"]["restart_drain_timeout"] ) +# In-band restart (``/restart``, SIGUSR1, self-restart from a child CLI) +# waits for active turns to finish *before* ``stop()`` begins. Distinct +# from ``restart_drain_timeout``, which is the force-interrupt budget +# once ``stop()`` is running (and must stay short under systemd +# TimeoutStopSec). See #77184. +DEFAULT_GATEWAY_RESTART_AFTER_TURN_TIMEOUT = float( + DEFAULT_CONFIG["agent"]["restart_after_turn_timeout"] +) + def is_gateway_supervisor_process( environ: Mapping[str, str] | None = None, @@ -64,3 +73,48 @@ def parse_restart_drain_timeout(raw: object) -> float: except (TypeError, ValueError): return DEFAULT_GATEWAY_RESTART_DRAIN_TIMEOUT return max(0.0, value) + + +def parse_restart_after_turn_timeout(raw: object) -> float: + """Parse the after-turn wait cap for in-band restart, falling back to default. + + ``0`` is a deliberate disable (legacy immediate drain) and must not fall + through to the default — unlike empty/missing input. + """ + if raw is None: + return DEFAULT_GATEWAY_RESTART_AFTER_TURN_TIMEOUT + if isinstance(raw, str) and not raw.strip(): + return DEFAULT_GATEWAY_RESTART_AFTER_TURN_TIMEOUT + try: + value = float(raw) + except (TypeError, ValueError): + return DEFAULT_GATEWAY_RESTART_AFTER_TURN_TIMEOUT + return max(0.0, value) + + +def resolve_restart_exit_wait_budget( + drain_timeout: float, + after_turn_timeout: float, + *, + headroom: float = 15.0, +) -> float: + """Seconds a CLI should wait for the gateway PID to exit after SIGUSR1. + + In-band restart may defer ``stop()`` until active turns finish + (``after_turn_timeout``) and then spend up to ``drain_timeout`` inside + ``stop()``. Callers that fall back to a hard kill on wait expiry must + cover both phases or they reintroduce #77184. + """ + try: + drain = max(float(drain_timeout), 0.0) + except (TypeError, ValueError): + drain = 0.0 + try: + after_turn = max(float(after_turn_timeout), 0.0) + except (TypeError, ValueError): + after_turn = 0.0 + try: + margin = max(float(headroom), 0.0) + except (TypeError, ValueError): + margin = 0.0 + return drain + after_turn + margin diff --git a/gateway/run.py b/gateway/run.py index fd1d62feb1..24d501b5b7 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -2278,9 +2278,11 @@ from gateway.shutdown_watchdog import ( start_loop_liveness_watchdog, ) from gateway.restart import ( + DEFAULT_GATEWAY_RESTART_AFTER_TURN_TIMEOUT, DEFAULT_GATEWAY_RESTART_DRAIN_TIMEOUT, GATEWAY_FATAL_CONFIG_EXIT_CODE, GATEWAY_SERVICE_RESTART_EXIT_CODE, + parse_restart_after_turn_timeout, parse_restart_drain_timeout, ) @@ -3775,13 +3777,22 @@ class TurnRunner: from agent.display import ( get_tool_preview_max_len, get_tool_verb, + prepare_tool_preview, tool_verb_connector, verb_drops_preview, ) _pl = get_tool_preview_max_len() _cap = _pl if _pl > 0 else 40 - if len(preview) > _cap: - preview = preview[:_cap - 3] + "..." + _prepared_preview = prepare_tool_preview( + tool_name, + args, + fallback=preview, + max_len=_cap, + ) + if _progress_adapter is not None: + preview = _progress_adapter.format_tool_preview(_prepared_preview) + else: + preview = _prepared_preview.text # Friendly labels: render a human-phrased line for built-in # tools ("🔍 Searching the web for ...") by prefixing the verb # onto the preview the callback already computed (so the @@ -4436,6 +4447,17 @@ class TurnRunner: turn_route = self._runner._resolve_turn_agent_config(ctx.message, model, runtime_kwargs) + # Per-platform skip_context_files — messaging platforms can opt out + # of filesystem-heavy context-file discovery (SOUL.md, AGENTS.md, + # .cursorrules) to cut AIAgent construction latency. Especially + # impactful on Windows, where stat() + directory walks are 10-100x + # slower than Linux. Off by default; soul identity is preserved so + # the persona survives even with minimal context. + _platforms_gw_cfg = (ctx.user_config.get("gateway") or {}).get("platforms") or {} + _plat_gw_cfg = _platforms_gw_cfg.get(platform_key) or {} + _skip_context = _plat_gw_cfg.get("skip_context_files") + skip_context_files = bool(_skip_context) if _skip_context is not None else False + # Check agent cache — reuse the AIAgent from the previous message # in this session to preserve the frozen system prompt and tool # schemas for prompt cache hits. @@ -4447,6 +4469,7 @@ class TurnRunner: cache_keys=self._runner._extract_cache_busting_config(ctx.user_config), user_id=getattr(ctx.source, "user_id", None), user_id_alt=getattr(ctx.source, "user_id_alt", None), + skip_context_files=skip_context_files, ) agent = None reused_cached_agent = False @@ -4682,6 +4705,10 @@ class TurnRunner: session_db=getattr(self._runner._session_db, "_db", self._runner._session_db), # Reload from disk — do not reuse the startup snapshot (#60955). fallback_model=self._runner._refresh_fallback_model(), + skip_context_files=skip_context_files, + # Keep the persona even with minimal context: soul identity is + # a single small file, not part of the expensive walk. + load_soul_identity=True, ) if _cache_lock and _cache is not None: with _cache_lock: @@ -4954,7 +4981,7 @@ class TurnRunner: agent_history, observed_group_context = _build_gateway_agent_history( ctx.history, channel_prompt=ctx.channel_prompt, - inject_timestamps=_message_timestamps_enabled(_load_gateway_config()), + inject_timestamps=_message_timestamps_enabled(ctx.user_config), ) # FTS write-corruption guard (#50502): when message persistence @@ -5623,6 +5650,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew _busy_input_mode: str = "interrupt" _busy_text_mode: str = "interrupt" _restart_drain_timeout: float = DEFAULT_GATEWAY_RESTART_DRAIN_TIMEOUT + _restart_after_turn_timeout: float = DEFAULT_GATEWAY_RESTART_AFTER_TURN_TIMEOUT _exit_code: Optional[int] = None _draining: bool = False _external_drain_active: bool = False @@ -5759,6 +5787,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew self._busy_input_mode = self._load_busy_input_mode() self._busy_text_mode = self._load_busy_text_mode() self._restart_drain_timeout = self._load_restart_drain_timeout() + self._restart_after_turn_timeout = self._load_restart_after_turn_timeout() self._provider_routing = self._load_provider_routing() self._fallback_model = self._load_fallback_model() @@ -8182,6 +8211,29 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew ) return value + @staticmethod + def _load_restart_after_turn_timeout() -> float: + """Load in-band restart wait-for-idle timeout in seconds (#77184).""" + env_raw = os.getenv("HERMES_RESTART_AFTER_TURN_TIMEOUT") + if env_raw is not None and str(env_raw).strip() != "": + raw: object = env_raw + else: + cfg = _load_gateway_runtime_config() + raw = cfg_get(cfg, "agent", "restart_after_turn_timeout", default=None) + value = parse_restart_after_turn_timeout(raw) + # Warn only when the user supplied a non-empty value that failed to + # parse (parser falls back to the default). ``0`` is valid. + if raw is not None and str(raw).strip() != "": + try: + float(raw) + except (TypeError, ValueError): + logger.warning( + "Invalid restart_after_turn_timeout '%s', using default %.0fs", + raw, + DEFAULT_GATEWAY_RESTART_AFTER_TURN_TIMEOUT, + ) + return value + @staticmethod def _load_background_notifications_mode() -> str: """Load background process notification mode from config or env var. @@ -9903,6 +9955,78 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew except Exception as e: logger.debug("Failed to launch systemd planned-restart helper: %s", e) + async def _await_active_work_before_restart(self) -> bool: + """Wait for in-flight work to finish before entering ``stop()``. + + In-band restart used to call ``stop()`` immediately, which folded the + requesting turn into the drain wait set and force-interrupted it at + ``restart_drain_timeout`` (#77184). Instead we refuse new turns and + wait here for active agents/cron/api work to reach zero, then let + ``stop()`` run against an idle gateway (drain is instant). + + Returns True when work drained to zero, False when the safety cap + elapsed with work still active (caller proceeds to ``stop()``, which + may then interrupt remaining runs under ``restart_drain_timeout``). + """ + active = self._active_work_count() + if active <= 0: + return True + + timeout = float(getattr(self, "_restart_after_turn_timeout", 0.0) or 0.0) + if timeout <= 0: + logger.info( + "Restart requested with %d active work unit(s); " + "restart_after_turn_timeout=0 — entering stop()/drain immediately", + active, + ) + return False + + logger.info( + "Restart requested with %d active work unit(s); " + "deferring stop() until they finish (cap=%.0fs) so in-flight " + "turns are not amputated (#77184)", + active, + timeout, + ) + try: + self._update_runtime_status("draining") + except Exception: + pass + + loop = asyncio.get_running_loop() + deadline = loop.time() + timeout + last_status_at = 0.0 + while self._active_work_count() > 0: + now = loop.time() + if now >= deadline: + logger.warning( + "Restart after-turn wait timed out after %.0fs with %d " + "still active; proceeding to stop()/drain which may " + "interrupt remaining work (#77184)", + timeout, + self._active_work_count(), + ) + return False + if (now - last_status_at) >= 30.0: + logger.info( + "Restart deferred: waiting on %d active work unit(s) " + "(%.0fs remaining before force drain)", + self._active_work_count(), + deadline - now, + ) + try: + self._update_runtime_status("draining") + except Exception: + pass + last_status_at = now + await asyncio.sleep(0.1) + + logger.info( + "Restart deferred wait complete — active work drained; " + "proceeding to stop()" + ) + return True + def request_restart(self, *, detached: bool = False, via_service: bool = False) -> bool: if self._restart_task_started: return False @@ -9910,8 +10034,17 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew self._restart_detached = detached self._restart_via_service = via_service self._restart_task_started = True + # Refuse new turns immediately while in-flight work finishes. + # Keep ``_running`` True so adapters stay connected and the active + # turn can still deliver its final response (#77184). + self._draining = True async def _run_restart() -> None: + await self._await_active_work_before_restart() + # Launch the detached helper only AFTER the after-turn wait. + # Its deadline is drain_timeout+5 and covers stop() teardown — + # launching earlier would fire `hermes gateway restart` while + # the requesting turn was still running. if detached: try: await self._launch_detached_restart_command() @@ -10114,7 +10247,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew sweep_recoverable, ) - if not ledger_enabled(): + if not await asyncio.to_thread(ledger_enabled): return 0 # Only claim rows we can actually send this boot: self.adapters # holds a platform only after its connect() succeeded, and each @@ -10166,7 +10299,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew result = None try: if result is not None and getattr(result, "success", False): - mark_delivered(row["obligation_id"]) + await asyncio.to_thread(mark_delivered, row["obligation_id"]) redelivered += 1 logger.info( "Redelivered recovered final response to %s:%s " @@ -10175,7 +10308,8 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew row["obligation_id"], row["attempts"], ) else: - mark_failed( + await asyncio.to_thread( + mark_failed, row["obligation_id"], str(getattr(result, "error", "") or "send failed"), ) @@ -17247,6 +17381,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew # below; a /new or another lifecycle transition may move # session_entry.session_id while the old run is still unwinding. _run_start_session_id = session_entry.session_id + _turn_started_monotonic = time.monotonic() agent_result = await self._run_agent( message=message_text, context_prompt=context_prompt, @@ -17262,6 +17397,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew persist_user_timestamp=persist_user_timestamp, message_type=event.message_type, ) + _turn_seconds = time.monotonic() - _turn_started_monotonic # Stop persistent typing indicator now that the agent is done. # Slack AI status is scoped to a thread/workspace, so preserve the @@ -17477,6 +17613,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew context_tokens=agent_result.get("last_prompt_tokens", 0) or 0, context_length=agent_result.get("context_length") or None, cwd=os.environ.get("TERMINAL_CWD", ""), + turn_seconds=_turn_seconds, ) except Exception as _footer_err: logger.debug("runtime_footer build failed: %s", _footer_err) @@ -22220,6 +22357,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew cache_keys: dict | None = None, user_id: str | None = None, user_id_alt: str | None = None, + skip_context_files: bool = False, ) -> str: """Compute a stable string key from agent config values. @@ -22274,6 +22412,10 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew _cache_keys_sorted, str(user_id or ""), str(user_id_alt or ""), + # skip_context_files changes the agent's frozen system prompt + # (context files in vs out) — a toggled config edit must + # rebuild the cached agent, not silently reuse it. + bool(skip_context_files), ], sort_keys=True, default=str, @@ -24214,6 +24356,23 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew source.platform, source.thread_id, event_message_id, reply_in_thread=_progress_reply_in_thread, ) + # Relay Discord auto-thread lane: a channel-initiating message has no + # thread_id at ingest (the thread is born on the connector's FIRST + # send). The connector stamps prospective_thread_id (the anchor message + # id, == the id of the thread it will create) and auto-threads any + # outbound carrying that anchor as reply_to. Without it, the progress / + # tool-status bubble is sent flat (no thread, no anchor) and lands in + # the PARENT channel while the final reply threads — the search-status + # updates leaked outside the thread (staging repro 2026-08-02). Carry + # the anchor on the progress send so it routes into the SAME auto-thread. + _relay_prospective_thread_id = ( + str(getattr(source, "prospective_thread_id", None)) + if source.platform == Platform.DISCORD + and getattr(source, "delivered_via_upstream_relay", False) + and getattr(source, "prospective_thread_id", None) + and not source.thread_id + else None + ) _progress_metadata = ( self._thread_metadata_for_source(source, event_message_id) if _progress_thread_id == source.thread_id @@ -24225,10 +24384,19 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew reply_to_message_id=event_message_id, ) ) if _progress_thread_id else None + if _progress_metadata is None and _relay_prospective_thread_id: + # No real thread yet, but the connector will auto-thread on the + # reply anchor; carry it so progress joins that thread. + _progress_metadata = {"reply_to_message_id": event_message_id} _progress_metadata = _non_conversational_metadata(_progress_metadata, platform=source.platform) _progress_reply_to = ( event_message_id - if source.platform in (Platform.FEISHU, Platform.MATTERMOST) and source.thread_id and event_message_id + if ( + source.platform in (Platform.FEISHU, Platform.MATTERMOST) + and source.thread_id + and event_message_id + ) + or _relay_prospective_thread_id else None ) @@ -24349,6 +24517,13 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew reply_to_message_id=event_message_id, ) ) if _progress_thread_id else None + if _status_thread_metadata is None and _relay_prospective_thread_id: + # Relay Discord auto-thread lane (see _progress_metadata above): + # carry the reply anchor so status/interim bubbles route into + # the same connector-created thread as the final reply. + _status_thread_metadata = { + "reply_to_message_id": event_message_id + } # Bridge extracted to TurnRunner._status_callback_sync; publish the # status wiring computed above onto the shared TurnContext at the diff --git a/gateway/runtime_footer.py b/gateway/runtime_footer.py index 024cf74d68..8719524d5a 100644 --- a/gateway/runtime_footer.py +++ b/gateway/runtime_footer.py @@ -11,6 +11,15 @@ Config (``~/.hermes/config.yaml``):: enabled: true # off by default fields: [model, context_pct, cwd] # order shown; drop any to hide +Available fields: + model — bare model id, vendor prefix dropped (``gpt-5.4``) + context_pct — last-call context occupancy as a percent (``5%``) + latency — wall-clock duration of the turn (``22s``, ``1m05s``) + cwd — home-relative working dir (``~``) + +``latency`` is opt-in: it is NOT in the default field set, so a footer whose +``fields`` are unset renders exactly as before. + Per-platform overrides live under ``display.platforms..runtime_footer``. Users can toggle the global setting with ``/footer on|off`` from both the CLI and any gateway platform. @@ -88,12 +97,24 @@ def resolve_footer_config( return resolved +def _format_latency(seconds: float) -> str: + """Humanize a turn duration: ``<1s``, ``22s``, ``1m05s``.""" + if seconds < 1: + return "<1s" + total = int(round(seconds)) + if total < 60: + return f"{total}s" + m, sec = divmod(total, 60) + return f"{m}m{sec:02d}s" + + def format_runtime_footer( *, model: Optional[str], context_tokens: int, context_length: Optional[int], cwd: Optional[str] = None, + turn_seconds: Optional[float] = None, fields: Iterable[str] = _DEFAULT_FIELDS, ) -> str: """Render the footer line, or return "" if no fields have data. @@ -111,6 +132,11 @@ def format_runtime_footer( if context_length and context_length > 0 and context_tokens >= 0: pct = max(0, min(100, round((context_tokens / context_length) * 100))) parts.append(f"{pct}%") + elif field == "latency": + # Wall-clock turn duration. Skipped when the caller supplied no + # timing (call sites that don't measure) or the value is negative. + if turn_seconds is not None and turn_seconds >= 0: + parts.append(_format_latency(turn_seconds)) elif field == "cwd": rel = _home_relative_cwd(cwd or os.environ.get("TERMINAL_CWD", "")) if rel: @@ -130,12 +156,17 @@ def build_footer_line( context_tokens: int, context_length: Optional[int], cwd: Optional[str] = None, + turn_seconds: Optional[float] = None, ) -> str: """Top-level entry point used by gateway/run.py. Returns the footer text (empty string when disabled or no data). Callers append this to the final response themselves, preserving a single blank line of separation. + + ``turn_seconds`` is the wall-clock duration of the agent run, measured by + the caller with ``time.monotonic()``. Callers that don't measure it leave + it ``None`` and the ``latency`` field is skipped. """ cfg = resolve_footer_config(user_config, platform_key) if not cfg.get("enabled"): @@ -145,5 +176,6 @@ def build_footer_line( context_tokens=context_tokens, context_length=context_length, cwd=cwd, + turn_seconds=turn_seconds, fields=cfg.get("fields") or _DEFAULT_FIELDS, ) diff --git a/gateway/session.py b/gateway/session.py index 92e2a5ae09..e8221bf6bc 100644 --- a/gateway/session.py +++ b/gateway/session.py @@ -2279,7 +2279,7 @@ class SessionStore: """ if self._db: try: - return self._db.session_count() > 1 + return self._db.session_count_ge(2) except Exception: pass # fall through to heuristic # Fallback: check if sessions.json was loaded with existing data. diff --git a/gateway/slash_commands.py b/gateway/slash_commands.py index fb3cec080a..7b87e43505 100644 --- a/gateway/slash_commands.py +++ b/gateway/slash_commands.py @@ -4147,7 +4147,14 @@ class GatewaySlashCommandsMixin: # Evict cached agent so next turn rebuilds system prompt # from current files (SOUL.md, memory, etc.). self._evict_cached_agent(session_key) - self._cleanup_agent_resources(tmp_agent) + # Off-loop + bounded: temporary-agent teardown can block on + # subprocess/network/SQLite work. Running it inline freezes the + # gateway loop and stalls platform polling / heartbeat, the same + # wedge class fixed for /new (#35994) and hygiene/shutdown + # (#53175). + await self._cleanup_agent_resources_off_loop( + tmp_agent, context="manual compression" + ) lines = [f"🗜️ {summary['headline']}"] if focus_topic: lines.append(t("gateway.compress.focus_line", topic=focus_topic)) @@ -4632,30 +4639,39 @@ class GatewaySlashCommandsMixin: logger.error("Failed to create branch session: %s", e) return t("gateway.branch.create_failed", error=e) - # Copy conversation history to the new session - for msg in history: - try: - await self._session_db.append_message( - session_id=new_session_id, - role=msg.get("role", "user"), - content=msg.get("content"), - tool_name=msg.get("tool_name") or msg.get("name"), - tool_calls=msg.get("tool_calls"), - tool_call_id=msg.get("tool_call_id"), - finish_reason=msg.get("finish_reason"), - reasoning=msg.get("reasoning"), - reasoning_content=msg.get("reasoning_content"), - reasoning_details=msg.get("reasoning_details"), - codex_reasoning_items=msg.get("codex_reasoning_items"), - codex_message_items=msg.get("codex_message_items"), - # Keep the api_content sidecar so the branch's first turn - # replays the parent's exact wire bytes (warm provider - # prompt cache) instead of a full cold prefill. - api_content=extract_api_content_sidecar(msg), - timestamp=msg.get("timestamp"), - ) - except Exception: - pass # Best-effort copy + # Copy conversation history to the new session in bounded-chunk + # transactions (see #23254): one txn per row was the removed + # write-amplification pattern, and a history can be hundreds of rows. + # Best-effort like the old loop — a failed copy still yields a + # usable (partial) branch. + try: + await self._session_db.append_messages_batch( + new_session_id, + [ + { + "role": msg.get("role", "user"), + "content": msg.get("content"), + "tool_name": msg.get("tool_name") or msg.get("name"), + "tool_calls": msg.get("tool_calls"), + "tool_call_id": msg.get("tool_call_id"), + "finish_reason": msg.get("finish_reason"), + "reasoning": msg.get("reasoning"), + "reasoning_content": msg.get("reasoning_content"), + "reasoning_details": msg.get("reasoning_details"), + "codex_reasoning_items": msg.get("codex_reasoning_items"), + "codex_message_items": msg.get("codex_message_items"), + # Keep the api_content sidecar so the branch's first turn + # replays the parent's exact wire bytes (warm provider + # prompt cache) instead of a full cold prefill. + "api_content": extract_api_content_sidecar(msg), + "timestamp": msg.get("timestamp"), + } + for msg in history + ], + chunk_rows=500, + ) + except Exception: + pass # Best-effort copy # Set title try: diff --git a/hermes_cli/__init__.py b/hermes_cli/__init__.py index c87e06cb63..14f175dcf9 100644 --- a/hermes_cli/__init__.py +++ b/hermes_cli/__init__.py @@ -14,8 +14,8 @@ Provides subcommands for: import os import sys -__version__ = "0.19.1" -__release_date__ = "2026.7.30" +__version__ = "0.20.0" +__release_date__ = "2026.8.3" def _ensure_utf8(): diff --git a/hermes_cli/auth.py b/hermes_cli/auth.py index 82eb5ee7db..bf7b00f203 100644 --- a/hermes_cli/auth.py +++ b/hermes_cli/auth.py @@ -648,42 +648,94 @@ ZAI_ENDPOINTS = [ ] +def _probe_single_zai_endpoint( + api_key: str, endpoint: tuple, timeout: float, +) -> Optional[Dict[str, str]]: + """Probe a single Z.AI endpoint. Returns endpoint info dict or None. + + Preserves the per-endpoint candidate-model loop: endpoints carry a + ``probe_models`` LIST and each model is tried in order until one + succeeds (some plans only accept newer/older GLM slugs). + """ + ep_id, base_url, probe_models, label = endpoint + for model in probe_models: + try: + resp = httpx.post( + f"{base_url}/chat/completions", + headers={ + "Authorization": f"Bearer {api_key}", + "Content-Type": "application/json", + }, + json={ + "model": model, + "stream": False, + "max_tokens": 1, + "messages": [{"role": "user", "content": "ping"}], + }, + timeout=timeout, + ) + if resp.status_code == 200: + logger.debug("Z.AI endpoint probe: %s (%s) model=%s OK", ep_id, base_url, model) + return { + "id": ep_id, + "base_url": base_url, + "model": model, + "label": label, + } + logger.debug("Z.AI endpoint probe: %s model=%s returned %s", ep_id, model, resp.status_code) + except Exception as exc: + logger.debug("Z.AI endpoint probe: %s model=%s failed: %s", ep_id, model, exc) + return None + + def detect_zai_endpoint(api_key: str, timeout: float = 8.0) -> Optional[Dict[str, str]]: - """Probe z.ai endpoints to find one that accepts this API key. + """Probe z.ai endpoints in parallel to find one that accepts this API key. Returns {"id": ..., "base_url": ..., "model": ..., "label": ...} for the - first working endpoint, or None if all fail. For endpoints with multiple - candidate models, tries each in order and returns the first that succeeds. + first working endpoint (in ZAI_ENDPOINTS priority order), or None if all + fail. For endpoints with multiple candidate models, each worker tries + its endpoint's models in order and returns the first that succeeds. """ - for ep_id, base_url, probe_models, label in ZAI_ENDPOINTS: - for model in probe_models: + from concurrent.futures import ThreadPoolExecutor, as_completed + + # No `with` block: a context manager would join ALL probe threads on + # exit, defeating the early return below. shutdown(wait=False) lets the + # surviving daemon-style probes drain in the background instead of + # blocking the caller on slow/unreachable endpoints. + pool = ThreadPoolExecutor(max_workers=len(ZAI_ENDPOINTS)) + try: + futures = { + pool.submit(_probe_single_zai_endpoint, api_key, ep, timeout): ep[0] + for ep in ZAI_ENDPOINTS + } + by_id = {ep_id: f for f, ep_id in futures.items()} + results: Dict[str, Dict[str, str]] = {} + for future in as_completed(futures): + ep_id = futures[future] try: - resp = httpx.post( - f"{base_url}/chat/completions", - headers={ - "Authorization": f"Bearer {api_key}", - "Content-Type": "application/json", - }, - json={ - "model": model, - "stream": False, - "max_tokens": 1, - "messages": [{"role": "user", "content": "ping"}], - }, - timeout=timeout, - ) - if resp.status_code == 200: - logger.debug("Z.AI endpoint probe: %s (%s) model=%s OK", ep_id, base_url, model) - return { - "id": ep_id, - "base_url": base_url, - "model": model, - "label": label, - } - logger.debug("Z.AI endpoint probe: %s model=%s returned %s", ep_id, model, resp.status_code) - except Exception as exc: - logger.debug("Z.AI endpoint probe: %s model=%s failed: %s", ep_id, model, exc) - return None + result = future.result() + if result is not None: + results[ep_id] = result + except Exception: + pass + # Early exit in PRIORITY order: walk endpoints highest-priority + # first; if one has succeeded and every higher-priority probe + # has already finished (without success), no later completion + # can win — return now instead of waiting out slow endpoints + # (main's sequential loop also stopped at first success). + for ep in ZAI_ENDPOINTS: + if not by_id[ep[0]].done(): + break # a higher-priority probe is still in flight + if ep[0] in results: + return results[ep[0]] + + # All probes finished: first match in priority order, if any. + for ep in ZAI_ENDPOINTS: + if ep[0] in results: + return results[ep[0]] + return None + finally: + pool.shutdown(wait=False) def _resolve_zai_base_url(api_key: str, default_url: str, env_override: str) -> str: @@ -8161,6 +8213,68 @@ def _codex_device_code_login() -> Dict[str, Any]: # ==================== MiniMax Portal OAuth ==================== +_MINIMAX_OAUTH_ERROR_BODY_LIMIT = 16 * 1024 + + +def _minimax_response_error_text( + response: httpx.Response, + *, + limit: int = _MINIMAX_OAUTH_ERROR_BODY_LIMIT, +) -> str: + """Return a bounded error body from a streamed MiniMax OAuth response.""" + limit = max(0, int(limit)) + chunks: list[bytes] = [] + total = 0 + truncated = False + try: + if getattr(response, "is_stream_consumed", False): + text = response.text + return text[:limit] + ("...[truncated]" if len(text) > limit else "") + + for chunk in response.iter_bytes(): + if not chunk: + continue + remaining = limit + 1 - total + if remaining <= 0: + truncated = True + break + if len(chunk) > remaining: + chunks.append(chunk[:remaining]) + total += remaining + truncated = True + break + chunks.append(chunk) + total += len(chunk) + raw = b"".join(chunks) + if len(raw) > limit: + raw = raw[:limit] + truncated = True + encoding = response.encoding or "utf-8" + text = raw.decode(encoding, errors="replace") + return text + ("...[truncated]" if truncated else "") + finally: + response.close() + + +def _minimax_post_form( + client: httpx.Client, + url: str, + *, + data: Dict[str, Any], + headers: Dict[str, str], +) -> httpx.Response: + """POST a MiniMax OAuth form without eagerly reading error bodies.""" + request = client.build_request( + "POST", + url, + data=data, + headers=headers, + ) + response = client.send(request, stream=True) + if response.status_code == 200: + response.read() + return response + def _minimax_pkce_pair() -> tuple: """Generate (code_verifier, code_challenge_S256, state) for MiniMax OAuth.""" import secrets @@ -8176,7 +8290,8 @@ def _minimax_request_user_code( client: httpx.Client, *, portal_base_url: str, client_id: str, code_challenge: str, state: str, ) -> Dict[str, Any]: - response = client.post( + response = _minimax_post_form( + client, f"{portal_base_url}/oauth/code", data={ "response_type": "code", @@ -8193,8 +8308,9 @@ def _minimax_request_user_code( }, ) if response.status_code != 200: + body = _minimax_response_error_text(response) raise AuthError( - f"MiniMax OAuth authorization failed: {response.text or response.reason_phrase}", + f"MiniMax OAuth authorization failed: {body or response.reason_phrase}", provider="minimax-oauth", code="authorization_failed", ) payload = response.json() @@ -8242,7 +8358,8 @@ def _minimax_poll_token( interval = max(2.0, (interval_ms or 2000) / 1000.0) while _time.time() < deadline: - response = client.post( + response = _minimax_post_form( + client, f"{portal_base_url}/oauth/token", data={ "grant_type": MINIMAX_OAUTH_GRANT_TYPE, @@ -8255,17 +8372,22 @@ def _minimax_poll_token( "Accept": "application/json", }, ) - try: - payload = response.json() if response.text else {} - except Exception: - payload = {} - + error_text = "" if response.status_code != 200: - msg = (payload.get("base_resp", {}) or {}).get("status_msg") or response.text + error_text = _minimax_response_error_text(response) + try: + payload = json.loads(error_text) if error_text else {} + except Exception: + payload = {} + msg = (payload.get("base_resp", {}) or {}).get("status_msg") or error_text raise AuthError( f"MiniMax OAuth error: {msg or 'unknown'}", provider="minimax-oauth", code="token_exchange_failed", ) + try: + payload = response.json() if response.text else {} + except Exception: + payload = {} status = payload.get("status") if status == "error": @@ -8401,7 +8523,8 @@ def _refresh_minimax_oauth_state( portal_base_url = state["portal_base_url"] with httpx.Client(timeout=httpx.Timeout(timeout_seconds), follow_redirects=True) as client: - response = client.post( + response = _minimax_post_form( + client, f"{portal_base_url}/oauth/token", data={ "grant_type": "refresh_token", @@ -8413,15 +8536,20 @@ def _refresh_minimax_oauth_state( "Accept": "application/json", }, ) - if response.status_code != 200: - body = response.text.lower() - relogin = any(m in body for m in - ("invalid_grant", "refresh_token_reused", "invalid_refresh_token")) - raise AuthError( - f"MiniMax OAuth refresh failed: {response.text or response.reason_phrase}", - provider="minimax-oauth", code="refresh_failed", - relogin_required=relogin, - ) + # The non-200 branch reads a STREAMED body, so it must run while + # the client is still open — iter_bytes() after the client context + # closes raises (StreamClosed). The 200 path was already read by + # _minimax_post_form, so response.json() below is safe outside. + if response.status_code != 200: + body = _minimax_response_error_text(response) + body_lower = body.lower() + relogin = any(m in body_lower for m in + ("invalid_grant", "refresh_token_reused", "invalid_refresh_token")) + raise AuthError( + f"MiniMax OAuth refresh failed: {body or response.reason_phrase}", + provider="minimax-oauth", code="refresh_failed", + relogin_required=relogin, + ) payload = response.json() if payload.get("status") != "success": raise AuthError( diff --git a/hermes_cli/backup.py b/hermes_cli/backup.py index 3a5b7cd4d3..43748fd3e2 100644 --- a/hermes_cli/backup.py +++ b/hermes_cli/backup.py @@ -15,8 +15,10 @@ import shutil import sqlite3 import sys import tempfile +import threading import time import zipfile +from contextlib import contextmanager from datetime import datetime, timezone from pathlib import Path from typing import Any, Dict, List, Optional @@ -84,6 +86,7 @@ _EXCLUDED_SUFFIXES = ( # File names to skip (runtime state that's meaningless on another machine) _EXCLUDED_NAMES = { + ".backup.lock", "gateway.pid", "cron.pid", } @@ -132,6 +135,85 @@ _SECRET_FILE_NAMES = {".env", "auth.json", "state.db"} _EXTERNAL_PREFIX = "_external/" +class BackupInProgressError(RuntimeError): + """Raised when another process already owns the Hermes backup slot.""" + + +class _SQLiteSnapshotError(RuntimeError): + pass + + +@contextmanager +def _backup_operation_lock(hermes_home: Path, timeout_seconds: float = 0.25): + """Acquire one cross-process backup slot for full and quick snapshots.""" + lock_path = hermes_home / ".backup.lock" + lock_path.parent.mkdir(parents=True, exist_ok=True) + handle = lock_path.open("a+b") + acquired = False + deadline = time.monotonic() + max(0.0, timeout_seconds) + try: + if os.name == "nt": + import msvcrt + + if lock_path.stat().st_size == 0: + handle.write(b" ") + handle.flush() + while True: + try: + handle.seek(0) + msvcrt.locking(handle.fileno(), msvcrt.LK_NBLCK, 1) + acquired = True + break + except (OSError, PermissionError): + if time.monotonic() >= deadline: + raise BackupInProgressError("another Hermes backup is already running") + time.sleep(0.05) + else: + import fcntl + + while True: + try: + fcntl.flock(handle.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB) + acquired = True + break + except (BlockingIOError, OSError): + if time.monotonic() >= deadline: + raise BackupInProgressError("another Hermes backup is already running") + time.sleep(0.05) + + yield + finally: + if acquired: + try: + if os.name == "nt": + import msvcrt + + handle.seek(0) + msvcrt.locking(handle.fileno(), msvcrt.LK_UNLCK, 1) + else: + import fcntl + + fcntl.flock(handle.fileno(), fcntl.LOCK_UN) + except (OSError, PermissionError): + pass + handle.close() + + +@contextmanager +def _atomic_output_path(final_path: Path): + """Yield a hidden sibling path and publish it only after a clean close.""" + partial_path = final_path.with_name( + f".{final_path.name}.{os.getpid()}-{threading.get_ident()}.partial" + ) + partial_path.unlink(missing_ok=True) + try: + yield partial_path + os.replace(partial_path, final_path) + except BaseException: + partial_path.unlink(missing_ok=True) + raise + + def _collect_memory_provider_external_paths() -> List[Path]: """Return existing absolute paths the active memory provider stores outside HERMES_HOME, resolved from config only (no network, no init). @@ -509,6 +591,17 @@ def run_backup(args) -> None: print(f"Error: Hermes home directory not found at {hermes_root}") sys.exit(1) + try: + with _backup_operation_lock(hermes_root): + _run_backup_locked(args, hermes_root) + except BackupInProgressError as exc: + print(f"Error: {exc}") + raise SystemExit(2) from exc + + +def _run_backup_locked(args, hermes_root: Path) -> None: + """Write a full backup while the cross-process backup slot is held.""" + # Determine output path if args.output: out_path = Path(args.output).expanduser().resolve() @@ -528,6 +621,8 @@ def run_backup(args) -> None: out_path.parent.mkdir(parents=True, exist_ok=True) # Collect files + scan_started = time.monotonic() + logger.info("backup phase=scan status=started") print(f"Scanning {display_hermes_home()} ...") files_to_add: list[tuple[Path, Path]] = [] # (absolute, relative) skipped_dirs = set() @@ -582,18 +677,30 @@ def run_backup(args) -> None: external_to_add.append((fpath, arcname)) if not files_to_add and not external_to_add: + logger.info( + "backup phase=scan status=empty duration_ms=%.1f", + (time.monotonic() - scan_started) * 1000, + ) print("No files to back up.") return # Create the zip file_count = len(files_to_add) + len(external_to_add) + logger.info( + "backup phase=scan status=complete duration_ms=%.1f files=%d", + (time.monotonic() - scan_started) * 1000, + file_count, + ) + logger.info("backup phase=archive status=started files=%d", file_count) print(f"Backing up {file_count} files ...") total_bytes = 0 errors = [] t0 = time.monotonic() - with zipfile.ZipFile(out_path, "w", zipfile.ZIP_DEFLATED, compresslevel=6) as zf: + with _atomic_output_path(out_path) as archive_path, zipfile.ZipFile( + archive_path, "w", zipfile.ZIP_DEFLATED, compresslevel=6 + ) as zf: for i, (abs_path, rel_path) in enumerate(files_to_add, 1): try: # Safe copy for SQLite databases (handles WAL mode) @@ -624,6 +731,11 @@ def run_backup(args) -> None: # Progress every 500 files if i % 500 == 0: print(f" {i}/{file_count} files ...") + logger.info( + "backup phase=archive status=progress completed=%d total=%d", + i, + file_count, + ) # External memory-provider state, stored under the ``_external/`` arc # prefix. These never include ``.db`` files in practice (config/env @@ -638,6 +750,13 @@ def run_backup(args) -> None: elapsed = time.monotonic() - t0 zip_size = out_path.stat().st_size + logger.info( + "backup phase=archive status=complete duration_ms=%.1f files=%d errors=%d bytes=%d", + elapsed * 1000, + file_count, + len(errors), + zip_size, + ) # Summary print() @@ -1013,6 +1132,23 @@ def create_quick_snapshot( hermes_home: Optional[Path] = None, keep: Optional[int] = None, max_file_size: Optional[int] = None, +) -> Optional[str]: + """Create one atomic quick snapshot while holding the shared backup slot.""" + home = hermes_home or get_hermes_home() + with _backup_operation_lock(home): + return _create_quick_snapshot_locked( + label=label, + hermes_home=home, + keep=keep, + max_file_size=max_file_size, + ) + + +def _create_quick_snapshot_locked( + label: Optional[str] = None, + hermes_home: Optional[Path] = None, + keep: Optional[int] = None, + max_file_size: Optional[int] = None, ) -> Optional[str]: """Create a quick state snapshot of critical files. @@ -1058,9 +1194,17 @@ def create_quick_snapshot( return True ts = datetime.now(timezone.utc).strftime("%Y%m%d-%H%M%S") - snap_id = f"{ts}-{label}" if label else ts + base_snap_id = f"{ts}-{label}" if label else ts + snap_id = base_snap_id + suffix = 2 + while (root / snap_id).exists(): + snap_id = f"{base_snap_id}-{suffix}" + suffix += 1 snap_dir = root / snap_id - snap_dir.mkdir(parents=True, exist_ok=True) + staging_dir = root / f".{snap_id}.{os.getpid()}.partial" + shutil.rmtree(staging_dir, ignore_errors=True) + staging_dir.mkdir(parents=True, exist_ok=False) + logger.info("quick snapshot phase=copy status=started id=%s", snap_id) manifest: Dict[str, int] = {} # rel_path -> file size failed_dbs: list[str] = [] # present *.db that could not be snapshotted @@ -1092,7 +1236,7 @@ def create_quick_snapshot( if sub.suffix == ".db": oversized_skipped.append(sub_rel) continue - dst = snap_dir / sub_rel + dst = staging_dir / sub_rel dst.parent.mkdir(parents=True, exist_ok=True) try: # Route SQLite DBs through the WAL-safe backup() path so a @@ -1126,7 +1270,7 @@ def create_quick_snapshot( oversized_skipped.append(rel) continue - dst = snap_dir / rel + dst = staging_dir / rel dst.parent.mkdir(parents=True, exist_ok=True) try: @@ -1167,7 +1311,7 @@ def create_quick_snapshot( ) if not manifest: - shutil.rmtree(snap_dir, ignore_errors=True) + shutil.rmtree(staging_dir, ignore_errors=True) if failed_dbs: # Distinguish "nothing to snapshot" from "state.db present but unreadable" print( @@ -1187,9 +1331,11 @@ def create_quick_snapshot( "failed_dbs": failed_dbs, "oversized_skipped": oversized_skipped, } - with open(snap_dir / "manifest.json", "w", encoding="utf-8") as f: + with open(staging_dir / "manifest.json", "w", encoding="utf-8") as f: json.dump(meta, f, indent=2) + os.replace(staging_dir, snap_dir) + # Auto-prune. Defaults preserve historical manual /snapshot behavior; callers # with known high-churn safety snapshots (for example pre-update) can pass a # smaller keep value so large state.db copies do not accumulate indefinitely. @@ -1216,7 +1362,12 @@ def create_quick_snapshot( len(failed_dbs), len(oversized_skipped), ) - logger.info("State snapshot created: %s (%d files)", snap_id, len(manifest)) + logger.info( + "quick snapshot phase=copy status=complete id=%s files=%d bytes=%d", + snap_id, + len(manifest), + sum(manifest.values()), + ) return snap_id @@ -1231,7 +1382,7 @@ def list_quick_snapshots( results = [] for d in sorted(root.iterdir(), reverse=True): - if not d.is_dir(): + if not d.is_dir() or d.name.startswith(".") or d.name.endswith(".partial"): continue manifest_path = d / "manifest.json" if manifest_path.exists(): @@ -1440,7 +1591,11 @@ def _prune_quick_snapshots(root: Path, keep: int = _QUICK_DEFAULT_KEEP) -> int: return 0 dirs = sorted( - (d for d in root.iterdir() if d.is_dir()), + ( + d + for d in root.iterdir() + if d.is_dir() and not d.name.startswith(".") and not d.name.endswith(".partial") + ), key=lambda d: d.name, reverse=True, ) @@ -1482,12 +1637,24 @@ def run_quick_backup(args) -> None: # --------------------------------------------------------------------------- def _write_full_zip_backup(out_path: Path, hermes_root: Path) -> Optional[Path]: + """Single-flight wrapper for automatic full zip backups.""" + try: + with _backup_operation_lock(hermes_root): + return _write_full_zip_backup_locked(out_path, hermes_root) + except BackupInProgressError as exc: + logger.warning("Full-zip backup skipped: %s", exc) + return None + + +def _write_full_zip_backup_locked(out_path: Path, hermes_root: Path) -> Optional[Path]: """Write a full zip snapshot of ``hermes_root`` to ``out_path``. Uses the same exclusion rules and SQLite safe-copy as :func:`run_backup`. Returns the output path on success, None on failure (nothing to back up, or write error — caller should surface the outcome but not raise). """ + scan_started = time.monotonic() + logger.info("automatic backup phase=scan status=started") files_to_add: list[tuple[Path, Path]] = [] try: for dirpath, dirnames, filenames in os.walk(hermes_root, followlinks=False): @@ -1513,10 +1680,18 @@ def _write_full_zip_backup(out_path: Path, hermes_root: Path) -> Optional[Path]: if not files_to_add: return None - sqlite_snapshot_failed = False + logger.info( + "automatic backup phase=scan status=complete duration_ms=%.1f files=%d", + (time.monotonic() - scan_started) * 1000, + len(files_to_add), + ) + + archive_started = time.monotonic() try: - with zipfile.ZipFile(out_path, "w", zipfile.ZIP_DEFLATED, compresslevel=6) as zf: - for abs_path, rel_path in files_to_add: + with _atomic_output_path(out_path) as archive_path, zipfile.ZipFile( + archive_path, "w", zipfile.ZIP_DEFLATED, compresslevel=6 + ) as zf: + for index, (abs_path, rel_path) in enumerate(files_to_add, 1): try: if abs_path.suffix == ".db": # Stage the snapshot alongside the output zip so that the @@ -1533,8 +1708,7 @@ def _write_full_zip_backup(out_path: Path, hermes_root: Path) -> Optional[Path]: "Full-zip backup aborted: SQLite snapshot failed for %s", rel_path, ) - sqlite_snapshot_failed = True - break + raise _SQLiteSnapshotError(str(rel_path)) zf.write(tmp_db, arcname=str(rel_path)) finally: tmp_db.unlink(missing_ok=True) @@ -1543,21 +1717,25 @@ def _write_full_zip_backup(out_path: Path, hermes_root: Path) -> Optional[Path]: except (PermissionError, OSError, ValueError) as exc: logger.debug("Skipping %s in zip backup: %s", rel_path, exc) continue - except OSError as exc: + if index % 500 == 0: + logger.info( + "automatic backup phase=archive status=progress completed=%d total=%d", + index, + len(files_to_add), + ) + except (OSError, _SQLiteSnapshotError) as exc: logger.warning("Full-zip backup: zip write failed: %s", exc) - # Best-effort cleanup of partial file - try: - out_path.unlink(missing_ok=True) - except OSError: - pass + # ``_atomic_output_path`` already removed the hidden partial. Do not + # unlink ``out_path`` here: it may be a previous valid backup that the + # atomic publisher deliberately preserved. return None - if sqlite_snapshot_failed: - try: - out_path.unlink(missing_ok=True) - except OSError: - pass - return None + logger.info( + "automatic backup phase=archive status=complete duration_ms=%.1f files=%d bytes=%d", + (time.monotonic() - archive_started) * 1000, + len(files_to_add), + out_path.stat().st_size, + ) return out_path diff --git a/hermes_cli/callbacks.py b/hermes_cli/callbacks.py index 69f4b6d405..aad0542d28 100644 --- a/hermes_cli/callbacks.py +++ b/hermes_cli/callbacks.py @@ -250,4 +250,4 @@ def approval_callback(cli, command: str, description: str) -> str: if hasattr(cli, "_app") and cli._app: cli._app.invalidate() cprint(f"\n{_DIM} ⏱ Timeout — denying command{_RST}") - return "deny" + return "timeout" diff --git a/hermes_cli/cli_agent_setup_mixin.py b/hermes_cli/cli_agent_setup_mixin.py index 050a60c2ff..57232dde27 100644 --- a/hermes_cli/cli_agent_setup_mixin.py +++ b/hermes_cli/cli_agent_setup_mixin.py @@ -429,6 +429,7 @@ class CLIAgentSetupMixin: f"({msg_count} user message{'s' if msg_count != 1 else ''}, {len(restored)} total messages)" ) self._restore_session_cwd(session_meta, quiet=_quiet_mode) + self._restore_session_yolo(session_meta, quiet=_quiet_mode) else: if _quiet_mode: print( @@ -640,6 +641,7 @@ class CLIAgentSetupMixin: f"{len(restored)} total messages)[/]" ) self._restore_session_cwd(session_meta) + self._restore_session_yolo(session_meta) else: accent_color = _accent_hex() self._console_print( diff --git a/hermes_cli/cli_commands_mixin.py b/hermes_cli/cli_commands_mixin.py index 105d7c256b..5ec16a5fe4 100644 --- a/hermes_cli/cli_commands_mixin.py +++ b/hermes_cli/cli_commands_mixin.py @@ -1055,6 +1055,12 @@ class CLICommandsMixin: # and a no-op when the session recorded no cwd. See #38562. self._restore_session_cwd(session_meta) + # Restore the target session's persisted YOLO bypass. Any bypass the + # PREVIOUS session had toggled on stops applying automatically because + # the approval session key just changed. Same contract as a startup + # --resume. + self._restore_session_yolo(session_meta) + def _handle_sessions_command(self, cmd_original: str) -> None: """Handle /sessions [list|] — browse or resume previous sessions. @@ -1165,25 +1171,32 @@ class CLICommandsMixin: _cprint(f" Failed to create branch session: {e}") return - # Copy conversation history to the new session - for msg in self.conversation_history: - try: - self._session_db.append_message( - session_id=new_session_id, - role=msg.get("role", "user"), - content=msg.get("content"), - tool_name=msg.get("tool_name") or msg.get("name"), - tool_calls=msg.get("tool_calls"), - tool_call_id=msg.get("tool_call_id"), - reasoning=msg.get("reasoning"), - # Keep the api_content sidecar so the branch's first turn - # replays the parent's exact wire bytes (warm provider - # prompt cache) instead of a full cold prefill. - api_content=extract_api_content_sidecar(msg), - timestamp=msg.get("timestamp"), - ) - except Exception: - pass # Best-effort copy + # Copy conversation history to the new session in bounded-chunk + # transactions (see #23254) instead of one txn per row. Best-effort + # like the old loop — a failed copy still yields a usable branch. + try: + self._session_db.append_messages_batch( + new_session_id, + [ + { + "role": msg.get("role", "user"), + "content": msg.get("content"), + "tool_name": msg.get("tool_name") or msg.get("name"), + "tool_calls": msg.get("tool_calls"), + "tool_call_id": msg.get("tool_call_id"), + "reasoning": msg.get("reasoning"), + # Keep the api_content sidecar so the branch's first turn + # replays the parent's exact wire bytes (warm provider + # prompt cache) instead of a full cold prefill. + "api_content": extract_api_content_sidecar(msg), + "timestamp": msg.get("timestamp"), + } + for msg in self.conversation_history + ], + chunk_rows=500, + ) + except Exception: + pass # Best-effort copy # Set title on the branch try: diff --git a/hermes_cli/config_defaults.py b/hermes_cli/config_defaults.py index 255da8e0a1..4ff18177a4 100644 --- a/hermes_cli/config_defaults.py +++ b/hermes_cli/config_defaults.py @@ -35,21 +35,23 @@ DEFAULT_CONFIG = { # tools or receiving API responses. Only fires when the agent has # been completely idle for this duration. 0 = unlimited. "gateway_timeout": 1800, - # Graceful drain timeout for gateway stop/restart (seconds). - # The gateway stops accepting new work, waits for running agents - # to finish, then interrupts any remaining runs after the timeout. - # 0 = no drain, interrupt immediately (the default). + # Force-interrupt budget once gateway stop()/drain has begun + # (seconds). Applies to SIGTERM/external stop and to the final + # phase of in-band restart after any after-turn wait. 0 = interrupt + # immediately (the default). # - # Contract: if you restart the gateway, in-flight work stops. We do - # not hold the restart open for a grace window — a drain timeout - # large enough to "save" a long agent turn would have to outlast an - # unbounded task (some runs take days), which is impossible, and a - # drain timeout shorter than systemd's TimeoutStopSec invites a - # SIGKILL-mid-cleanup race that leaves a stale lock and crash-loops - # the service. 0 sidesteps both: interrupt now, clean up, exit fast. - # Set a positive value in config.yaml only if you explicitly want a - # grace window on /restart (and keep it well under TimeoutStopSec). + # Keep this short and under systemd TimeoutStopSec — a long value + # here invites SIGKILL-mid-cleanup. For in-band restart + # (/restart, SIGUSR1), prefer restart_after_turn_timeout below so + # active turns finish *before* stop() begins (#77184). "restart_drain_timeout": 0, + # In-band restart wait for active turns to finish before stop() + # (seconds). /restart and SIGUSR1 refuse new work, then wait up to + # this cap for in-flight agents/cron/api runs to complete naturally + # so the requesting turn is not amputated by restart_drain_timeout. + # 0 = legacy behaviour (enter stop()/drain immediately). Default + # 6h is a safety valve for wedged agents, not a target latency. + "restart_after_turn_timeout": 21600, # Upper bound (seconds) a submitted prompt waits for the deferred # agent build (MCP discovery, model metadata, skills scan) before # failing with a visible error (#63078). The gateway's wait is @@ -2368,7 +2370,7 @@ DEFAULT_CONFIG = { "listing": "auto", # Absolute cap on the embedded listing in tokens (chars/4 # estimate), regardless of context size. Range 200..60000. - "listing_max_tokens": 20000, + "listing_max_tokens": 4000, }, }, diff --git a/hermes_cli/copilot_auth.py b/hermes_cli/copilot_auth.py index 0f1ec8c034..8416f99d22 100644 --- a/hermes_cli/copilot_auth.py +++ b/hermes_cli/copilot_auth.py @@ -79,9 +79,11 @@ def resolve_copilot_token() -> tuple[str, str]: Raises ValueError if only a classic PAT is available. """ # 1. Check env vars in priority order + any_env_var_set = False for env_var in COPILOT_ENV_VARS: val = os.getenv(env_var, "").strip() if val: + any_env_var_set = True valid, msg = validate_copilot_token(val) if not valid: logger.warning( @@ -90,7 +92,23 @@ def resolve_copilot_token() -> tuple[str, str]: continue return val, env_var - # 2. Fall back to gh auth token + # 2. Fall back to gh auth token — but ONLY when no Copilot env var was + # explicitly set. When the user exported GITHUB_TOKEN (even an + # unsupported classic PAT), their intent is to use *that* token, not + # to silently substitute one from the gh CLI credential store. + # Skipping the subprocess here also avoids a slow `gh auth token` + # call (up to 5s timeout on Windows) on every cold start that scans + # Copilot auth state — a measurable contributor to the ~14s + # cold-start stall (#60800). The user can run `copilot login` or + # set a supported token (gho_*/github_pat_*/ghu_) explicitly. + if any_env_var_set: + logger.debug( + "Copilot env var(s) set but none held a supported token; " + "skipping `gh auth token` fallback to honor explicit env-var " + "intent (and avoid the subprocess cost on cold start, #60800)." + ) + return "", "" + token = _try_gh_cli_token() if token: valid, msg = validate_copilot_token(token) diff --git a/hermes_cli/gateway.py b/hermes_cli/gateway.py index 681ce09c9e..2a03ff1004 100644 --- a/hermes_cli/gateway.py +++ b/hermes_cli/gateway.py @@ -32,12 +32,15 @@ PROJECT_ROOT = Path(__file__).parent.parent.resolve() from gateway.config import coerce_systemd_watchdog_seconds, load_gateway_config from gateway.status import terminate_pid from gateway.restart import ( + DEFAULT_GATEWAY_RESTART_AFTER_TURN_TIMEOUT, DEFAULT_GATEWAY_RESTART_DRAIN_TIMEOUT, EXTERNAL_GATEWAY_SUPERVISOR_ENV, GATEWAY_FATAL_CONFIG_EXIT_CODE, GATEWAY_SERVICE_RESTART_EXIT_CODE, is_gateway_supervisor_process, + parse_restart_after_turn_timeout, parse_restart_drain_timeout, + resolve_restart_exit_wait_budget, ) from hermes_cli.config import ( get_env_value, @@ -252,10 +255,12 @@ def _request_gateway_self_restart(pid: int) -> bool: def _graceful_restart_via_sigusr1(pid: int, drain_timeout: float) -> bool: """Send SIGUSR1 to a gateway PID and wait for it to exit gracefully. - SIGUSR1 is wired in gateway/run.py to ``request_restart(via_service=True)`` - which drains in-flight agent runs (up to ``agent.restart_drain_timeout`` - seconds), then exits. Both systemd (``Restart=always``) and launchd - (unconditional ``KeepAlive``) restart on any exit. + SIGUSR1 is wired in gateway/run.py to ``request_restart(via_service=True)``, + which refuses new turns, waits for in-flight work up to + ``agent.restart_after_turn_timeout``, then runs ``stop()`` (force-interrupt + budget ``agent.restart_drain_timeout``) and exits. Both systemd + (``Restart=always``) and launchd (unconditional KeepAlive) restart on + any exit. This is the drain-aware alternative to ``systemctl restart`` / ``SIGTERM``, which SIGKILL in-flight agents after a short timeout. @@ -264,9 +269,9 @@ def _graceful_restart_via_sigusr1(pid: int, drain_timeout: float) -> bool: pid: Gateway process PID (systemd MainPID, launchd PID, or bare process PID). drain_timeout: Seconds to wait for the process to exit after sending - SIGUSR1. Should be slightly larger than the gateway's - ``agent.restart_drain_timeout`` to allow the drain loop to - finish cleanly. + SIGUSR1. Must cover the after-turn wait plus the stop()/drain + phase (#77184); callers should pass + ``resolve_restart_exit_wait_budget(...)``. Returns: True if the PID was signalled and exited within the timeout. @@ -3276,6 +3281,26 @@ def _get_restart_drain_timeout() -> float: return parse_restart_drain_timeout(raw) +def _get_restart_after_turn_timeout() -> float: + """Return the in-band restart wait-for-idle timeout in seconds (#77184).""" + env_raw = os.getenv("HERMES_RESTART_AFTER_TURN_TIMEOUT") + if env_raw is not None and str(env_raw).strip() != "": + return parse_restart_after_turn_timeout(env_raw) + cfg = read_raw_config() + agent_cfg = cfg.get("agent", {}) if isinstance(cfg, dict) else {} + if isinstance(agent_cfg, dict) and "restart_after_turn_timeout" in agent_cfg: + return parse_restart_after_turn_timeout(agent_cfg.get("restart_after_turn_timeout")) + return parse_restart_after_turn_timeout(None) + + +def _get_restart_exit_wait_budget() -> float: + """CLI wait for gateway exit after SIGUSR1 / self-restart (#77184).""" + return resolve_restart_exit_wait_budget( + _get_restart_drain_timeout(), + _get_restart_after_turn_timeout(), + ) + + def systemd_install( force: bool = False, system: bool = False, @@ -3456,10 +3481,12 @@ def systemd_restart(system: bool = False): if pid is not None: scope_label = _service_scope_label(system).capitalize() svc = get_service_name() - drain_timeout = _get_restart_drain_timeout() - - print(f"⏳ {scope_label} service restarting gracefully (PID {pid})...") - if _graceful_restart_via_sigusr1(pid, drain_timeout + 5): + wait_budget = _get_restart_exit_wait_budget() + print( + f"⏳ {scope_label} service restarting gracefully (PID {pid}) — " + f"waiting up to {wait_budget:.0f}s for in-flight turns + drain..." + ) + if _graceful_restart_via_sigusr1(pid, wait_budget): # The gateway exits with code 75 for a planned service restart. # RestartSec can otherwise delay the relaunch even though the # operator asked for an immediate restart, so kick the unit once @@ -3482,7 +3509,7 @@ def systemd_restart(system: bool = False): return print( - f"⚠ Graceful restart did not complete within {int(drain_timeout + 5)}s; " + f"⚠ Graceful restart did not complete within {int(wait_budget)}s; " "forcing a service restart..." ) _run_systemctl( diff --git a/hermes_cli/kanban_db.py b/hermes_cli/kanban_db.py index b64d54e53a..113e34842e 100644 --- a/hermes_cli/kanban_db.py +++ b/hermes_cli/kanban_db.py @@ -9773,6 +9773,9 @@ def count_notify_subs( board: Optional[str] = None, notifier_profiles: Optional[Iterable[str]] = None, include_unowned: bool = False, + platform: Optional[str] = None, + chat_id: Optional[str] = None, + thread_id: Optional[str] = None, ) -> int: """Count ``kanban_notify_subs`` rows via a read-only connection. @@ -9786,7 +9789,10 @@ def count_notify_subs( DB, or a legacy DB that predates the subscriptions table, counts as zero. When ``notifier_profiles`` is supplied, only subscriptions owned by those profiles are counted; ``include_unowned`` also includes legacy - rows without an owner stamp. Path resolution matches :func:`connect` + rows without an owner stamp. Optional platform/chat/thread filters narrow + the probe to one notification owner without changing the unfiltered count. + Platform matching is case-insensitive, matching notifier routing; chat and + thread identifiers are exact. Path resolution matches :func:`connect` (explicit ``db_path``, else ``board`` via :func:`kanban_db_path`). Raises :class:`sqlite3.Error` when the DB exists but cannot be read (locked, corrupt); callers choose their own fallback. @@ -9800,10 +9806,24 @@ def count_notify_subs( owner_where, owner_params = _notify_profile_filter( notifier_profiles, include_unowned=include_unowned, ) - sql = "SELECT COUNT(*) FROM kanban_notify_subs" + clauses: list[str] = [] + params: list[Any] = [] if owner_where: - sql += " WHERE " + owner_where - row = conn.execute(sql, owner_params).fetchone() + clauses.append(f"({owner_where})") + params.extend(owner_params) + if platform is not None: + clauses.append("LOWER(platform) = LOWER(?)") + params.append(platform) + if chat_id is not None: + clauses.append("chat_id = ?") + params.append(chat_id) + if thread_id is not None: + clauses.append("thread_id = ?") + params.append(thread_id) + query = "SELECT COUNT(*) FROM kanban_notify_subs" + if clauses: + query += " WHERE " + " AND ".join(clauses) + row = conn.execute(query, params).fetchone() except sqlite3.OperationalError as exc: if "no such table" in str(exc).lower(): return 0 diff --git a/hermes_cli/main.py b/hermes_cli/main.py index 98874824c1..624d0f1010 100644 --- a/hermes_cli/main.py +++ b/hermes_cli/main.py @@ -1013,16 +1013,9 @@ def _has_any_provider_configured() -> bool: except Exception: pass - # Check provider-specific auth fallbacks (for example, Copilot via gh auth). - try: - for provider_id, pconfig in PROVIDER_REGISTRY.items(): - if pconfig.auth_type != "api_key": - continue - status = get_auth_status(provider_id) - if status.get("logged_in"): - return True - except Exception: - pass + # Cheap local checks first: auth.json and config.yaml are on-disk lookups, + # while the PROVIDER_REGISTRY sweep below spawns subprocesses (gh) and can + # take 15-20s — long enough that desktop setup.status calls time out. # Check for Nous Portal OAuth credentials auth_file = get_hermes_home() / "auth.json" @@ -1050,6 +1043,17 @@ def _has_any_provider_configured() -> bool: if cfg_provider or cfg_base_url or cfg_api_key: return True + # Check provider-specific auth fallbacks (for example, Copilot via gh auth). + try: + for provider_id, pconfig in PROVIDER_REGISTRY.items(): + if pconfig.auth_type != "api_key": + continue + status = get_auth_status(provider_id) + if status.get("logged_in"): + return True + except Exception: + pass + # Check for Claude Code OAuth credentials (~/.claude/.credentials.json) # Only count these if Hermes has been explicitly configured — Claude Code # being installed doesn't mean the user wants Hermes to use their tokens. @@ -5710,7 +5714,7 @@ def _do_build_web_ui(web_dir: Path, *, fatal: bool = False) -> bool: return _run_npm_install_deterministic( npm, npm_cwd, - extra_args=(*npm_workspace_args, "--silent") if silent else npm_workspace_args, + extra_args=(*npm_workspace_args, "--silent", "--prefer-offline") if silent else (*npm_workspace_args, "--prefer-offline"), env=build_env, ) @@ -12159,6 +12163,33 @@ def main(): help="Reclaim disk space: merge FTS5 segments + VACUUM (no data change)", ) + sessions_clean_markers = sessions_subparsers.add_parser( + "clean-markers", + help="Permanently clear stale tool-call marker content left by sessions from before #78148", + description=( + "Before the #78148 fix, a local tool-call template could persist a " + "bare bracketed marker (e.g. \"[memory]\") as an assistant turn's " + "content instead of real text. This is already repaired in memory " + "on every session load, so running this is optional — it rewrites " + "the affected rows once, in place, so long-lived sessions stop " + "re-scanning/re-repairing the same rows on every resume. Only the " + "content column is touched; tool_calls and every other column on " + "the row are left untouched." + ), + ) + sessions_clean_markers.add_argument( + "--dry-run", + action="store_true", + default=False, + help="Report the affected row count without writing", + ) + sessions_clean_markers.add_argument( + "--no-backup", + action="store_true", + default=False, + help="Skip the timestamped state.db backup taken before writing (not recommended)", + ) + sessions_optimize_storage = sessions_subparsers.add_parser( "optimize-storage", help="Migrate the search index to the compact v23 layout (reclaims disk on large DBs)", diff --git a/hermes_cli/mcp_catalog.py b/hermes_cli/mcp_catalog.py index 6f8c9b30c6..dd1d2c07c4 100644 --- a/hermes_cli/mcp_catalog.py +++ b/hermes_cli/mcp_catalog.py @@ -230,6 +230,21 @@ def _parse_manifest(path: Path) -> CatalogEntry: scopes=list(auth_raw.get("scopes") or []), env_var=auth_raw.get("env_var"), ) + if t_type == "http" and a_type == "api_key": + # _build_server_config emits an Authorization header referencing + # ${MCP__API_KEY} (via _bearer_auth_headers), but install_entry + # only persists the env vars DECLARED in auth.env. Enforce the naming + # contract at parse time, or a manifest declaring e.g. N8N_API_KEY + # would install cleanly yet send a literal-placeholder header (401) + # at connect time. + from hermes_cli.mcp_config import _env_key_for_server + + _required_key = _env_key_for_server(name) + if not any(spec.name == _required_key for spec in env_list): + raise CatalogError( + f"{path}: http + api_key auth requires auth.env to declare " + f"'{_required_key}' (the key the Authorization header references)" + ) tools_raw = data.get("tools") or {} if not isinstance(tools_raw, dict): @@ -506,6 +521,10 @@ def _build_server_config( cfg["url"] = t.url if entry.auth.type == "oauth": cfg["auth"] = "oauth" + elif entry.auth.type == "api_key": + from hermes_cli.mcp_config import _bearer_auth_headers + + cfg["headers"] = _bearer_auth_headers(entry.name) return cfg diff --git a/hermes_cli/model_switch.py b/hermes_cli/model_switch.py index 252986ed30..66dbdc09aa 100644 --- a/hermes_cli/model_switch.py +++ b/hermes_cli/model_switch.py @@ -102,6 +102,30 @@ def _declared_model_ids(value: Any) -> list[str]: return ids +def _models_config_is_allowlist(value: Any) -> bool: + """Return True when ``models:`` is an intentional ID allowlist. + + A mapping like ``{model_id: {context_length: N}}`` is per-model *metadata* + written by ``_save_custom_provider`` / the ``hermes model`` wizard — not a + catalog narrow. Treating that shape as an allowlist made Desktop/Telegram + pickers show only the saved default for local Ollama (no ``api_key``), + while ``hermes model`` still live-probed the full ``/v1/models`` list. + Refresh could not help because the same gate skipped probing. + + List/string shapes remain allowlists for no-key endpoints. To pin a + dict-shaped catalog, set ``discover_models: false``. + """ + if value is None: + return False + if isinstance(value, str): + return bool(value.strip()) + if isinstance(value, dict): + return False + if isinstance(value, (list, tuple)): + return bool(_declared_model_ids(value)) + return False + + def _save_discovered_models_to_config( api_url: str, model_ids: list[str] ) -> None: @@ -2631,12 +2655,13 @@ def list_authenticated_providers( for _m in entry_models: if _m and _m not in ep_groups[group_key]["models"]: ep_groups[group_key]["models"].append(_m) - # Track explicit ``models:`` declarations separately from the - # merged list: a singular ``default_model``/``model`` is only the - # active selection and must not be mistaken for the user narrowing - # the endpoint to a curated subset (mirrors section 4's - # declaration-tracking; see #40542 / PR #61928). - if entry_declared_models: + # Track allowlist-shaped ``models:`` separately from the merged + # list: a singular ``default_model``/``model`` is only the active + # selection and must not suppress discovery (see #40542 / PR + # #61928). Dict-shaped ``models:`` is context_length metadata from + # ``hermes model``, not an allowlist — see + # ``_models_config_is_allowlist``. + if _models_config_is_allowlist(ep_cfg.get("models")): ep_groups[group_key]["has_explicit_models"] = True ep_groups[group_key]["raw_names"].append(display_name) ep_groups[group_key]["aliases"].update( @@ -2663,14 +2688,16 @@ def list_authenticated_providers( # unless the provider explicitly opts out via discover_models: false. # Policy mirrors Section 4's should_probe logic: # - With an api_key: always probe (user opted into the endpoint). - # - Without an api_key but with an explicit ``models:`` list: - # skip — the user is narrowing a public endpoint to a specific - # subset. A singular ``default_model``/``model`` does NOT count - # as narrowing (it's just the active selection) and must not - # suppress discovery — mirrors section 4 / #40542. - # - Without an api_key AND no explicit models: probe anyway so - # bare-endpoint providers (local llama.cpp / Ollama servers) - # still show their full model catalog. + # - Without an api_key but with an allowlist-shaped ``models:`` + # (list/string): skip — the user narrowed a public endpoint. + # A singular ``default_model``/``model`` does NOT count as + # narrowing (mirrors section 4 / #40542). + # - A dict-shaped ``models:`` is per-model metadata + # (context_length), not an allowlist — still probe so local + # Ollama/llama.cpp match ``hermes model``. Pin with + # ``discover_models: false`` instead. + # - Without an api_key AND no allowlist: probe anyway so bare + # local endpoints still show their full model catalog. api_key = str(ep_cfg.get("api_key", "") or "").strip() if not api_key: key_env = str(ep_cfg.get("key_env", "") or "").strip() @@ -2911,8 +2938,12 @@ def list_authenticated_providers( if default_model and default_model not in groups[group_key]["models"]: groups[group_key]["models"].append(default_model) - declared_models = _declared_model_ids(entry.get("models", {})) - if declared_models: + models_field = entry.get("models", {}) + declared_models = _declared_model_ids(models_field) + # Dict-shaped models: is context_length metadata from + # ``_save_custom_provider``, not an allowlist — see + # ``_models_config_is_allowlist``. + if _models_config_is_allowlist(models_field): groups[group_key]["has_explicit_models"] = True for model_id in declared_models: if model_id not in groups[group_key]["models"]: @@ -2974,24 +3005,20 @@ def list_authenticated_providers( # Live-discovery policy: # - With an api_key, the user has explicitly opted into the # endpoint and live /models is the source of truth — replace - # the (possibly partial) ``models:`` subset configured for - # context-length overrides with the full live catalog. - # This is the Bifrost / aggregator-gateway case. - # - Without an api_key but with an explicit ``models:`` list, - # the user is narrowing a public endpoint to a specific subset - # (e.g. ollama.com /v1/models returns 35 models but the user - # only wants 4). Preserve the explicit list and skip live - # discovery. The singular ``model:`` field is only the current - # active selection and must not suppress discovery on local - # no-key endpoints. - # - Without an api_key AND no explicit models, fall through to - # live discovery so bare-endpoint custom providers (local - # llama.cpp / Ollama servers) still appear populated. + # the (possibly partial) ``models:`` subset with the full + # live catalog (Bifrost / aggregator-gateway case). + # - Without an api_key but with an allowlist-shaped ``models:`` + # (list/string), the user narrowed a public endpoint (e.g. + # ollama.com). Preserve that list and skip live discovery. + # - A dict-shaped ``models:`` is per-model metadata written by + # ``_save_custom_provider`` for context_length — not an + # allowlist. Still probe so Desktop/Telegram match + # ``hermes model``. Pin a dict catalog with + # ``discover_models: false``. + # - The singular ``model:`` field is only the current active + # selection and must not suppress discovery. # - When discover_models: false is set, skip live discovery and - # keep the explicit ``models:`` list regardless of whether an - # api_key is present. This supports endpoints that expose a - # full aggregator catalog via /models but only serve a subset - # (parity with section 3's user ``providers:`` behaviour). + # keep the configured ``models:`` list regardless of api_key. _grp_is_current = ( slug.lower() == _current_provider_norm or _current_provider_norm in { diff --git a/hermes_cli/models.py b/hermes_cli/models.py index 28cf578fbf..e2e6c967f9 100644 --- a/hermes_cli/models.py +++ b/hermes_cli/models.py @@ -7,6 +7,7 @@ Add, remove, or reorder entries here — both `hermes setup` and from __future__ import annotations +import copy import json import logging import os @@ -71,7 +72,7 @@ OPENROUTER_MODELS: list[tuple[str, str]] = [ ("deepseek/deepseek-v4-flash", ""), ("deepseek/deepseek-v4-flash-0731", "dated snapshot of v4-flash"), # Qwen - ("qwen/qwen3.7-max", ""), + ("qwen/qwen3.8-max", ""), # MoonshotAI ("moonshotai/kimi-k3", "recommended"), # MiniMax @@ -243,7 +244,7 @@ _PROVIDER_MODELS: dict[str, list[str]] = { "deepseek/deepseek-v4-flash", "deepseek/deepseek-v4-flash-0731", # Qwen - "qwen/qwen3.7-max", + "qwen/qwen3.8-max", # MoonshotAI "moonshotai/kimi-k3", # MiniMax @@ -3465,10 +3466,37 @@ def _copilot_catalog_item_is_text_model(item: dict[str, Any]) -> bool: return True +# Module-level cache for the GitHub Copilot /models catalog. +# The picker path can ask for it multiple times in one process via: +# list_authenticated_providers -> cached_provider_model_ids -> provider_model_ids -> _fetch_github_models +# and later get_copilot_model_context()/normalize helpers. Cache the raw filtered +# catalog for a short TTL so we don't pay repeated TLS handshakes on every picker open. +# Keyed by the api_key used for the successful fetch so a credential swap +# mid-process never serves the previous account's catalog. Uses a monotonic +# clock so wall-clock adjustments can't extend the TTL. Lock-free like the +# other module caches here — a racing thread at worst duplicates one fetch. +_github_model_catalog_cache: Optional[list[dict[str, Any]]] = None +_github_model_catalog_cache_key: Optional[str] = None +_github_model_catalog_cache_time: float = 0.0 +_GITHUB_MODEL_CATALOG_CACHE_TTL = 300 # 5 minutes + + def fetch_github_model_catalog( api_key: Optional[str] = None, timeout: float = 5.0 ) -> Optional[list[dict[str, Any]]]: """Fetch the live GitHub Copilot model catalog for this account.""" + global _github_model_catalog_cache, _github_model_catalog_cache_key + global _github_model_catalog_cache_time + + if ( + _github_model_catalog_cache is not None + and _github_model_catalog_cache_key == api_key + and (time.monotonic() - _github_model_catalog_cache_time) < _GITHUB_MODEL_CATALOG_CACHE_TTL + ): + # Deep copy: catalog items are dicts, and a shallow copy would let + # callers mutate the cached entries in place. + return copy.deepcopy(_github_model_catalog_cache) + attempts: list[dict[str, str]] = [] if api_key: attempts.append({ @@ -3494,6 +3522,9 @@ def fetch_github_model_catalog( seen_ids.add(model_id) models.append(item) if models: + _github_model_catalog_cache = copy.deepcopy(models) + _github_model_catalog_cache_key = api_key + _github_model_catalog_cache_time = time.monotonic() return models except Exception: continue diff --git a/hermes_cli/observability/relay_shared_metrics.py b/hermes_cli/observability/relay_shared_metrics.py index d8ef419b91..e1d8c43ca5 100644 --- a/hermes_cli/observability/relay_shared_metrics.py +++ b/hermes_cli/observability/relay_shared_metrics.py @@ -1082,6 +1082,8 @@ def observe_lifecycle(hook_name: str, **kwargs: Any) -> None: """Project one Hermes lifecycle event into the core Relay integration.""" if not handles_hook(hook_name): return + if not relay_runtime.relay_instrumentation_enabled(): + return runtime = _get_runtime() if runtime is None: return diff --git a/hermes_cli/prompt_size.py b/hermes_cli/prompt_size.py index 30c212a429..0629b9346d 100644 --- a/hermes_cli/prompt_size.py +++ b/hermes_cli/prompt_size.py @@ -252,8 +252,9 @@ def compute_prompt_breakdown(platform: str = "cli") -> Dict[str, Any]: volatile = parts.get("volatile", "") # Skills index — the block (the largest single block - # when many skills are installed). Measured inside the stable tier. - skills_match = _SKILLS_BLOCK_RE.search(stable) + # when many skills are installed). Lives in the volatile tier (moved from + # stable so skill edits don't invalidate the cached identity prefix). + skills_match = _SKILLS_BLOCK_RE.search(volatile) or _SKILLS_BLOCK_RE.search(stable) skills_index = skills_match.group(0) if skills_match else "" # Memory + user profile live in the volatile tier. We re-derive their diff --git a/hermes_cli/session_recovery.py b/hermes_cli/session_recovery.py index c75f8552fd..39787c1bfb 100644 --- a/hermes_cli/session_recovery.py +++ b/hermes_cli/session_recovery.py @@ -31,6 +31,7 @@ from hermes_state import ( ProgressCallback = Callable[[dict[str, Any]], None] _CANONICAL_TABLES = ( + "system_prompts", "sessions", "messages", "session_model_usage", @@ -912,6 +913,8 @@ def _cleanup_partial_orphans( """ result: dict[str, Any] = { + "session_prompt_refs_cleared": 0, + "system_prompts_removed": 0, "sessions_parent_cleared": 0, "sessions_reconstructed": 0, "messages_retained": 0, @@ -947,6 +950,42 @@ def _cleanup_partial_orphans( ) result["sessions_parent_cleared"] = parent_count + prompt_ref_count = int( + destination.execute( + "SELECT COUNT(*) FROM sessions " + "WHERE system_prompt_hash IS NOT NULL " + "AND NOT EXISTS (" + "SELECT 1 FROM system_prompts " + "WHERE system_prompts.hash = sessions.system_prompt_hash)" + ).fetchone()[0] + ) + if prompt_ref_count: + destination.execute( + "UPDATE sessions SET system_prompt_hash = NULL " + "WHERE system_prompt_hash IS NOT NULL " + "AND NOT EXISTS (" + "SELECT 1 FROM system_prompts " + "WHERE system_prompts.hash = sessions.system_prompt_hash)" + ) + result["session_prompt_refs_cleared"] = prompt_ref_count + + unreferenced_prompt_count = int( + destination.execute( + "SELECT COUNT(*) FROM system_prompts " + "WHERE NOT EXISTS (" + "SELECT 1 FROM sessions " + "WHERE sessions.system_prompt_hash = system_prompts.hash)" + ).fetchone()[0] + ) + if unreferenced_prompt_count: + destination.execute( + "DELETE FROM system_prompts " + "WHERE NOT EXISTS (" + "SELECT 1 FROM sessions " + "WHERE sessions.system_prompt_hash = system_prompts.hash)" + ) + result["system_prompts_removed"] = unreferenced_prompt_count + dependent_tables = ( ("messages", "messages_removed"), ("session_model_usage", "session_model_usage_removed"), @@ -983,7 +1022,8 @@ def _cleanup_partial_orphans( # reconstruction counters describe data RETAINED, so summing them here # would report saving the user's messages as if it were losing them. result["total_removed_or_relinked"] = ( - int(result["sessions_parent_cleared"]) + int(result["session_prompt_refs_cleared"]) + + int(result["sessions_parent_cleared"]) + int(result["messages_removed"]) + int(result["session_model_usage_removed"]) + int(result["compression_locks_removed"]) diff --git a/hermes_cli/sessions_cmd.py b/hermes_cli/sessions_cmd.py index 0338585a68..cf317b610c 100644 --- a/hermes_cli/sessions_cmd.py +++ b/hermes_cli/sessions_cmd.py @@ -1039,6 +1039,26 @@ def cmd_sessions(args, sessions_parser=None): f"({_size_delta_label(saved)})" ) + elif action == "clean-markers": + if args.dry_run: + print("Dry run — scanning for stale tool-call marker rows (#78148)…") + else: + print("Scanning for stale tool-call marker rows (#78148)…") + report = db.purge_stale_tool_call_markers( + dry_run=args.dry_run, backup=not args.no_backup + ) + if report["rows_affected"] == 0: + print("✓ No affected rows found — nothing to clean.") + elif args.dry_run: + print( + f"Would clear {report['rows_affected']} row(s): " + f"ids {report['row_ids']}" + ) + else: + if report["backup_path"]: + print(f" backup: {report['backup_path']}") + print(f"✓ Cleared {report['rows_affected']} row(s).") + elif action == "optimize-storage": db_path = db.db_path if not db.fts_optimize_available(): diff --git a/hermes_cli/tips.py b/hermes_cli/tips.py index be8cb25b15..504f29115f 100644 --- a/hermes_cli/tips.py +++ b/hermes_cli/tips.py @@ -292,7 +292,7 @@ TIPS = [ "Delegation has a heartbeat thread — child activity propagates to the parent, preventing gateway timeouts.", "When a provider returns HTTP 402 (payment required), the auxiliary client auto-falls back to the next one.", "agent.tool_use_enforcement steers models that describe actions instead of calling tools — auto for GPT/Codex.", - "agent.restart_drain_timeout (default 60s) lets running agents finish before a gateway restart takes effect.", + "agent.restart_after_turn_timeout lets in-flight turns finish before /restart enters stop(); restart_drain_timeout is only the force-interrupt budget once stop() begins.", "agent.api_max_retries (default 3) controls how many times the agent retries a failed API call before surfacing the error — lower it for fast fallback.", "The gateway caches AIAgent instances per session — destroying this cache breaks Anthropic prompt caching.", "Any website can expose skills via /.well-known/skills/index.json — the skills hub discovers them automatically.", diff --git a/hermes_cli/tools_config.py b/hermes_cli/tools_config.py index 7be692fcad..5db37f83fa 100644 --- a/hermes_cli/tools_config.py +++ b/hermes_cli/tools_config.py @@ -26,6 +26,7 @@ from hermes_cli.config import ( from hermes_cli.colors import Colors, color from hermes_cli.nous_subscription import ( MANAGED_FEATURE_COVERAGE_CATEGORY, + NousSubscriptionFeatures, apply_nous_managed_defaults, get_nous_subscription_features, ) @@ -2618,6 +2619,7 @@ def _toolset_has_keys( config: dict = None, *, force_fresh: bool = False, + features: Optional[NousSubscriptionFeatures] = None, ) -> bool: """Check if a toolset's required API keys are configured.""" if config is None: @@ -2633,7 +2635,10 @@ def _toolset_has_keys( return False if ts_key in {"web", "image_gen", "video_gen", "tts", "stt", "browser"}: - features = get_nous_subscription_features(config, force_fresh=force_fresh) + if features is None: + features = get_nous_subscription_features( + config, force_fresh=force_fresh + ) feature = features.features.get(ts_key) if feature and (feature.available or feature.managed_by_nous): return True @@ -2641,7 +2646,12 @@ def _toolset_has_keys( # Check TOOL_CATEGORIES first (provider-aware) cat = TOOL_CATEGORIES.get(ts_key) if cat: - for provider in _visible_providers(cat, config, force_fresh=force_fresh): + for provider in _visible_providers( + cat, + config, + force_fresh=force_fresh, + features=features, + ): env_vars = provider.get("env_vars", []) if not env_vars: return True # No-key provider (e.g. Local Browser, Edge TTS) @@ -3076,6 +3086,7 @@ def _visible_providers( config: dict, *, force_fresh: bool = False, + features: Optional[NousSubscriptionFeatures] = None, ) -> list[dict]: """Return provider entries visible for the current auth/config state. @@ -3085,7 +3096,8 @@ def _visible_providers( login + entitlement check (see ``_configure_provider``); the row only *activates* the gateway once paid access is confirmed. """ - features = get_nous_subscription_features(config, force_fresh=force_fresh) + if features is None: + features = get_nous_subscription_features(config, force_fresh=force_fresh) acct = features.account_info # Pool-only users (entitled to managed tools via the free tool pool but with # no paid access) get image gen but NOT video gen — the pool doesn't fund diff --git a/hermes_cli/update_cmd.py b/hermes_cli/update_cmd.py index 95982c0f28..c9069b1b1a 100644 --- a/hermes_cli/update_cmd.py +++ b/hermes_cli/update_cmd.py @@ -2089,7 +2089,7 @@ def _update_node_dependencies() -> list[str]: print(" deps). Fix npm and re-run `hermes update`.") return list(labels) - extra_args = ["--no-fund", "--no-audit", "--progress=false"] + extra_args = ["--no-fund", "--no-audit", "--prefer-offline", "--progress=false"] from hermes_constants import with_hermes_node_path @@ -4863,37 +4863,20 @@ def _cmd_update_impl(args, gateway_mode: bool): _manage_cmd_cache[scope_] = cmd return cmd - # Drain budget for graceful SIGUSR1 restarts. The gateway drains - # for up to ``agent.restart_drain_timeout`` (default 60s) before - # exiting with code 75; we wait slightly longer so the drain - # completes before we fall back to a hard restart. On older - # systemd units without SIGUSR1 wiring this wait just times out - # and we fall back to ``systemctl restart`` (the old behaviour). + # Wait budget for graceful SIGUSR1 restarts. In-band restart + # may defer stop() until active turns finish + # (``restart_after_turn_timeout``, #77184) and then spend up to + # ``restart_drain_timeout`` inside stop(). Cover both phases so + # we don't fall back to a hard kill while the gateway is still + # patiently waiting for the requesting turn. On older systemd + # units without SIGUSR1 wiring this wait just times out and we + # fall back to ``systemctl restart`` (the old behaviour). try: - from hermes_constants import ( - DEFAULT_GATEWAY_RESTART_DRAIN_TIMEOUT as _DEFAULT_DRAIN, - ) - except Exception: - _DEFAULT_DRAIN = 60.0 - _cfg_drain = None - try: - from hermes_cli.config import load_config + from hermes_cli.gateway import _get_restart_exit_wait_budget - _cfg_agent = load_config().get("agent") or {} - _cfg_drain = _cfg_agent.get("restart_drain_timeout") + _drain_budget = max(float(_get_restart_exit_wait_budget()), 45.0) except Exception: - pass - try: - _drain_budget = ( - float(_cfg_drain) - if _cfg_drain is not None - else float(_DEFAULT_DRAIN) - ) - except (TypeError, ValueError): - _drain_budget = float(_DEFAULT_DRAIN) - # Add a 15s margin so the drain loop + final exit finish before - # we escalate to ``systemctl restart`` / SIGTERM. - _drain_budget = max(_drain_budget, 30.0) + 15.0 + _drain_budget = 45.0 restarted_services = [] failed_or_stale_units = [] diff --git a/hermes_cli/web_routers/profiles.py b/hermes_cli/web_routers/profiles.py index 17bd37c141..d7bc45c2d8 100644 --- a/hermes_cli/web_routers/profiles.py +++ b/hermes_cli/web_routers/profiles.py @@ -56,7 +56,13 @@ _write_profile_model = late("_write_profile_model") @sessions_router.get("/api/profiles/sessions") def get_profiles_sessions( - limit: int = Query(20, ge=0), + # ``le=500`` caps the per-request page size (idea from #39200) — this + # endpoint fans the query out across EVERY profile's state.db, so an + # unbounded limit multiplies the damage. 500 (not 100) because real + # desktop callers use limit=200 (sessions-settings ARCHIVED_FETCH_LIMIT, + # command palette) and the electron remote-merge over-fetches + # ``limit + offset``. + limit: int = Query(20, ge=0, le=500), offset: int = Query(0, ge=0), min_messages: int = 0, archived: str = "exclude", diff --git a/hermes_cli/web_routers/sessions.py b/hermes_cli/web_routers/sessions.py index d708ae46ea..c692d7b21e 100644 --- a/hermes_cli/web_routers/sessions.py +++ b/hermes_cli/web_routers/sessions.py @@ -49,7 +49,10 @@ _strip_session_list_rows = late("_strip_session_list_rows") @list_router.get("/api/sessions") def get_sessions( - limit: int = Query(20, ge=0), + # ``le=100`` caps the page size (idea from #39200): an unbounded limit + # lets one request drag every session row (plus correlated-subquery + # preview work) out of SQLite in a single hit. + limit: int = Query(20, ge=0, le=100), offset: int = Query(0, ge=0), min_messages: int = 0, archived: str = "exclude", @@ -353,6 +356,14 @@ async def search_sessions( source_filter=include_sources, exclude_sources=exclude_list or None, limit=fetch_limit, + fields=( + "session_id", + "role", + "snippet", + "source", + "model", + "session_started", + ), ) for m in matches: diff --git a/hermes_cli/web_routers/tools.py b/hermes_cli/web_routers/tools.py index e96baa6f6b..0fbb055cf5 100644 --- a/hermes_cli/web_routers/tools.py +++ b/hermes_cli/web_routers/tools.py @@ -58,6 +58,7 @@ async def get_toolsets(profile: Optional[str] = None): _get_platform_tools, _toolset_configuration_platform, _toolset_has_keys, + get_nous_subscription_features, gui_toolset_label, ) from hermes_cli.platforms import platform_label @@ -77,6 +78,7 @@ async def get_toolsets(profile: Optional[str] = None): ) for platform in target_platforms } + features = get_nous_subscription_features(config) result = [] for name, label, desc in toolset_rows: try: @@ -104,7 +106,7 @@ async def get_toolsets(profile: Optional[str] = None): ), "enabled": is_enabled, "available": is_enabled, - "configured": _toolset_has_keys(name, config), + "configured": _toolset_has_keys(name, config, features=features), "tools": tools, }) return result diff --git a/hermes_cli/web_server.py b/hermes_cli/web_server.py index 5af536fdce..1fb3e61316 100644 --- a/hermes_cli/web_server.py +++ b/hermes_cli/web_server.py @@ -102,7 +102,7 @@ from utils import env_var_enabled try: from fastapi import ( - FastAPI, File, Form, HTTPException, Request, UploadFile, + FastAPI, File, Form, HTTPException, Query, Request, UploadFile, WebSocket, WebSocketDisconnect, ) from fastapi.middleware.cors import CORSMiddleware @@ -118,7 +118,7 @@ except ImportError: from tools.lazy_deps import ensure as _lazy_ensure _lazy_ensure("tool.dashboard", prompt=False) from fastapi import ( - FastAPI, File, Form, HTTPException, Request, UploadFile, + FastAPI, File, Form, HTTPException, Query, Request, UploadFile, WebSocket, WebSocketDisconnect, ) from fastapi.middleware.cors import CORSMiddleware @@ -169,10 +169,40 @@ def _start_desktop_cron_ticker(stop_event: "threading.Event", interval: int = 60 def _warm_gateway_module() -> None: - try: - import hermes_cli.gateway # noqa: F401 - except Exception: - pass + """Pre-import heavy modules so the event loop is not stalled on first use. + + On a cold Windows install, importing these module chains triggers .pyc + compilation and Defender real-time scans that can stall the event loop + for 15-30s. The original fix (pre-#60800) only warmed + ``hermes_cli.gateway``. But the first WS connection and its initial + RPC burst (``setup.status``, ``setup.runtime_check``, + ``gateway.ready``→``resolve_skin``) pull in several *other* heavy + chains that were still imported on the loop thread, contributing to + the ~14s cold-start stall (#60800). Warm them all here so the cost + is paid in a worker thread while the server socket is already open. + """ + for mod in ( + "hermes_cli.gateway", + # setup.status / setup.runtime_check resolve provider auth state, + # which imports copilot_auth (→ subprocess module) and scans + # credential files. First import is noticeably slow on Windows. + "hermes_cli.auth", + "hermes_cli.copilot_auth", + "hermes_cli.runtime_provider", + # resolve_skin() reads config + initialises the skin engine. + # Even though handle_ws now calls it via asyncio.to_thread + # (see tui_gateway/ws.py), warming it here avoids the first-call + # import cost inside that thread. + "hermes_cli.skin_engine", + # model.options / picker context — parses provider catalogs and + # the models.dev cache on first use. + "hermes_cli.inventory", + "hermes_cli.model_switch", + ): + try: + __import__(mod) + except Exception: + pass def _resolve_restart_drain_timeout() -> float: @@ -5573,6 +5603,12 @@ def _normalize_memory_provider_schema(name: str, provider: Any) -> List[Dict[str kind = "select" elif explicit_kind in {"bool", "boolean"} or isinstance(raw.get("default"), bool): kind = "boolean" + elif explicit_kind in {"int", "integer"} or ( + isinstance(raw.get("default"), int) and not isinstance(raw.get("default"), bool) + ): + kind = "integer" + elif explicit_kind in {"float", "number"} or isinstance(raw.get("default"), float): + kind = "number" else: kind = "text" @@ -5593,6 +5629,9 @@ def _normalize_memory_provider_schema(name: str, provider: Any) -> List[Dict[str "options": options, "url": str(raw.get("url") or ""), "when": raw.get("when") if isinstance(raw.get("when"), dict) else None, + "minimum": raw.get("minimum"), + "maximum": raw.get("maximum"), + "step": raw.get("step"), "_env_key": str(raw.get("env_var") or "") or None, }) @@ -5736,6 +5775,9 @@ def _public_memory_provider_field(field: Dict[str, Any], data: Dict[str, Any]) - "options": field.get("options", []), "url": field.get("url", ""), "when": field.get("when"), + "minimum": field.get("minimum"), + "maximum": field.get("maximum"), + "step": field.get("step"), } return entry @@ -5758,6 +5800,31 @@ def _coerce_schema_field(field: Dict[str, Any], raw: Any) -> Any: if field["kind"] == "boolean": return _coerce_bool(raw, default=_coerce_bool(_field_default(field), default=False)) + if field["kind"] in {"integer", "number"}: + value = raw if raw is not None and raw != "" else _field_default(field) + try: + if isinstance(value, bool): + raise ValueError + parsed = float(value) + if not math.isfinite(parsed): + raise ValueError + if field["kind"] == "integer": + if not parsed.is_integer(): + raise ValueError + result: int | float = int(parsed) + else: + result = parsed + except (TypeError, ValueError, OverflowError) as exc: + raise ValueError(f"Invalid numeric value for '{field['key']}'") from exc + + minimum = field.get("minimum") + maximum = field.get("maximum") + if minimum is not None and result < minimum: + raise ValueError(f"'{field['key']}' must be at least {minimum}") + if maximum is not None and result > maximum: + raise ValueError(f"'{field['key']}' must be at most {maximum}") + return result + value = str(raw if raw is not None else "").strip() if field["kind"] == "select": if not value: @@ -5975,6 +6042,7 @@ async def setup_memory_provider(name: str, body: MemoryProviderSetupRequest): except Exception: _log.exception("Failed to persist memory provider setup values for %s", name) raise HTTPException(status_code=500, detail="Internal server error") + _invalidate_plugins_hub_cache() return _install_memory_provider_setup(name) @@ -5992,6 +6060,7 @@ async def update_memory_provider_config( if declared is None: raise HTTPException(status_code=404, detail=f"Unknown memory provider: {name}") _update_memory_provider_config(declared, _stringify_submitted_values(values)) + _invalidate_plugins_hub_cache() return {"ok": True} provider = _load_memory_provider(name) @@ -6006,6 +6075,7 @@ async def update_memory_provider_config( config["memory"] = memory_config memory_config["provider"] = name save_config(config) + _invalidate_plugins_hub_cache() return {"ok": True, "active": name} try: @@ -7389,8 +7459,8 @@ async def validate_custom_endpoint(body: CustomEndpointUpdate): headers["Authorization"] = f"Bearer {body.api_key.strip()}" try: - with httpx.Client(timeout=httpx.Timeout(8.0)) as client: - resp = client.get(url, headers=headers) + async with httpx.AsyncClient(timeout=httpx.Timeout(8.0)) as client: + resp = await client.get(url, headers=headers) except Exception: return {"ok": False, "reachable": False, "message": f"Could not reach {url}.", "models": []} @@ -7431,8 +7501,8 @@ async def validate_provider_credential(body: EnvVarUpdate, request: Request): api_key = (body.api_key or "").strip() headers = {"Authorization": f"Bearer {api_key}"} if api_key else None try: - with httpx.Client(timeout=httpx.Timeout(8.0)) as client: - resp = client.get(url, headers=headers) + async with httpx.AsyncClient(timeout=httpx.Timeout(8.0)) as client: + resp = await client.get(url, headers=headers) return {"ok": True, "reachable": True, "message": "", "models": _parse_model_ids(resp)} except Exception: return {"ok": False, "reachable": False, "message": f"Could not reach {url}."} @@ -7451,8 +7521,8 @@ async def validate_provider_credential(body: EnvVarUpdate, request: Request): params["key"] = value try: - with httpx.Client(timeout=httpx.Timeout(10.0)) as client: - resp = client.get(url, headers=headers, params=params) + async with httpx.AsyncClient(timeout=httpx.Timeout(10.0)) as client: + resp = await client.get(url, headers=headers, params=params) except Exception: return {"ok": False, "reachable": False, "message": "Could not reach the provider to verify the key."} @@ -14070,16 +14140,7 @@ def _get_usage_analytics(days: int = 30, profile: Optional[str] = None): FROM sessions WHERE started_at > ? """, (cutoff,)) totals = dict(cur3.fetchone()) - insights_report = InsightsEngine(db).generate(days=days) - skills = insights_report.get("skills", { - "summary": { - "total_skill_loads": 0, - "total_skill_edits": 0, - "total_skill_actions": 0, - "distinct_skills_used": 0, - }, - "top_skills": [], - }) + usage = InsightsEngine(db).get_usage_breakdown(days=days) return { "daily": daily, @@ -14089,17 +14150,24 @@ def _get_usage_analytics(days: int = 30, profile: Optional[str] = None): "by_task": _aux_task_summary(aux_rows), "totals": totals, "period_days": days, - "skills": skills, + "skills": usage["skills"], # Per-tool-name call counts (already computed by InsightsEngine); # the desktop Capabilities page aggregates these per toolset. - "tools": insights_report.get("tools", []), + "tools": usage["tools"], } finally: db.close() @app.get("/api/analytics/usage") -async def get_usage_analytics(days: int = 30, profile: Optional[str] = None): +async def get_usage_analytics( + days: int = Query(30, ge=1, le=365), + profile: Optional[str] = None, +): + """``days`` is clamped to 1-365 (idea from #74778): huge or non-positive + values would force expensive full-history SQL and InsightsEngine work, or + produce empty/inverted time windows. The UI only offers 7/30/90-day + presets.""" return await asyncio.to_thread(_get_usage_analytics, days, profile) @@ -14280,7 +14348,11 @@ def _get_models_analytics(days: int = 30, profile: Optional[str] = None): @app.get("/api/analytics/models") -async def get_models_analytics(days: int = 30, profile: Optional[str] = None): +async def get_models_analytics( + days: int = Query(30, ge=1, le=365), + profile: Optional[str] = None, +): + # ``days`` clamped to 1-365 (idea from #74778) — see get_usage_analytics. """Return model analytics without blocking the serving event loop.""" return await asyncio.to_thread(_get_models_analytics, days, profile) @@ -15922,6 +15994,14 @@ def _render_active_theme_bootstrap_css() -> str: return "" +# Hashed bundle assets (``/assets/-.``) are immutable +# by construction: any content change produces a new filename, and the entry +# point (index.html) is served ``no-store`` so it always references the +# current hashes. A year-long immutable cache lets browsers skip even the +# revalidation round-trip on every dashboard load. +_IMMUTABLE_ASSET_CACHE_CONTROL = "public, max-age=31536000, immutable" + + def mount_spa(application: FastAPI): """Mount the built SPA. Falls back to index.html for client-side routing. @@ -16044,9 +16124,32 @@ def mount_spa(application: FastAPI): css = css.replace(f"url({asset_dir}", f"url({prefix}{asset_dir}") css = css.replace(f"url(\"{asset_dir}", f"url(\"{prefix}{asset_dir}") css = css.replace(f"url('{asset_dir}", f"url('{prefix}{asset_dir}") - return Response(content=css, media_type="text/css") + return Response( + content=css, + media_type="text/css", + headers={"Cache-Control": _IMMUTABLE_ASSET_CACHE_CONTROL}, + ) - application.mount("/assets", StaticFiles(directory=WEB_DIST / "assets"), name="assets") + class _ImmutableAssetFiles(StaticFiles): + """StaticFiles that marks hashed bundle assets immutable. + + Everything under ``/assets/`` carries a Vite content hash in its + filename, so a given URL's bytes can never change — a rebuild + produces a NEW filename referenced by a fresh (``no-store``) + index.html. Without this header every dashboard load re-validated + each chunk; with it the browser serves reloads straight from its + HTTP cache. + """ + + async def get_response(self, path: str, scope): + response = await super().get_response(path, scope) + if response.status_code == 200: + response.headers["Cache-Control"] = _IMMUTABLE_ASSET_CACHE_CONTROL + return response + + application.mount( + "/assets", _ImmutableAssetFiles(directory=WEB_DIST / "assets"), name="assets" + ) @application.get("/{full_path:path}") async def serve_spa(full_path: str, request: Request): @@ -16630,8 +16733,75 @@ def _strip_dashboard_manifest(p: Dict[str, Any]) -> Dict[str, Any]: return {k: v for k, v in p.items() if not k.startswith("_")} -def _merged_plugins_hub() -> Dict[str, Any]: - """Agent discovery + dashboard manifests + optional provider picker metadata.""" +_PLUGINS_HUB_CACHE_TTL_SECONDS = 5.0 +_plugins_hub_cache: Optional[Dict[str, Any]] = None +_plugins_hub_cache_expires_at = 0.0 +_plugins_hub_cache_lock = threading.Lock() + + +def _invalidate_plugins_hub_cache() -> None: + global _plugins_hub_cache, _plugins_hub_cache_expires_at + with _plugins_hub_cache_lock: + _plugins_hub_cache = None + _plugins_hub_cache_expires_at = 0.0 + + +_plugins_hub_probe_inflight: set = set() +_plugins_hub_probe_lock = threading.Lock() + + +def _schedule_check_fn_probe(fn) -> Optional[threading.Thread]: + """Warm a cold ``check_fn`` verdict off the request path. + + The hub read path only consumes cached availability (never probes + inline). But the only other warmer lives in the tool-schema build, which + a dashboard-only session never runs — so a cold cache would report + ``auth_required=False`` forever. Kick a daemon-thread probe on the miss; + the short hub TTL picks up the verdict on the next fetch. Deduplicates + concurrent probes per function. Returns the spawned thread (or ``None`` + when a probe for *fn* is already in flight). + """ + with _plugins_hub_probe_lock: + if fn in _plugins_hub_probe_inflight: + return None + _plugins_hub_probe_inflight.add(fn) + + def _probe(): + try: + from tools.registry import _check_fn_cached + + _check_fn_cached(fn) + except Exception: + pass + finally: + with _plugins_hub_probe_lock: + _plugins_hub_probe_inflight.discard(fn) + + thread = threading.Thread( + target=_probe, name="plugins-hub-checkfn-probe", daemon=True + ) + thread.start() + return thread + + +def _merged_plugins_hub(force_refresh: bool = False) -> Dict[str, Any]: + """Agent discovery + dashboard manifests + optional provider picker metadata. + + IMPORTANT: this powers a dashboard request path, so it must stay read-only + and cheap. In particular, do not execute tool ``check_fn`` probes here — + those can trigger imports, auth/network checks, and other synchronous work + that starves the root event loop. We only consume last-known cached tool + availability, and we memoize the assembled payload briefly to collapse the + dashboard's bursty duplicate fetches. + """ + global _plugins_hub_cache, _plugins_hub_cache_expires_at + now = time.monotonic() + if not force_refresh: + with _plugins_hub_cache_lock: + if _plugins_hub_cache is not None and now < _plugins_hub_cache_expires_at: + return _plugins_hub_cache + + started_at = time.monotonic() from hermes_cli.plugins_cmd import ( _discover_all_plugins, _get_current_context_engine, @@ -16684,17 +16854,30 @@ def _merged_plugins_hub() -> Dict[str, Any]: source in {"user", "git"} and under_user_tree and Path(dir_str).is_dir() ) - # Check if this plugin provides tools that require auth + # Read-only auth hint: consult only last-known cached tool availability. + # A missing cache entry is treated as "unknown" rather than triggering a + # live probe inside this request path. auth_required = False auth_command = "" manifest_data = _read_plugin_manifest_at(dir_path) provides_tools = manifest_data.get("provides_tools") or [] if provides_tools: try: - from tools.registry import registry + from tools.registry import get_cached_check_fn_result, registry for tname in provides_tools: entry = registry.get_entry(tname) - if entry and entry.check_fn and not entry.check_fn(): + if not entry or not entry.check_fn: + continue + cached_result = get_cached_check_fn_result(entry.check_fn) + if cached_result is None: + # Cold cache: nothing else warms check_fns on + # dashboard-only sessions, so kick a background + # probe; the short hub TTL surfaces the verdict on + # the next fetch instead of pinning auth_required + # to False forever. + _schedule_check_fn_probe(entry.check_fn) + continue + if cached_result is False: auth_required = True auth_command = f"hermes auth {name}" break @@ -16733,7 +16916,7 @@ def _merged_plugins_hub() -> Dict[str, Any]: except Exception: context_engines = [] - return { + payload = { "plugins": rows, "orphan_dashboard_plugins": orphan_dashboard, "providers": { @@ -16743,6 +16926,18 @@ def _merged_plugins_hub() -> Dict[str, Any]: "context_options": context_engines, }, } + duration = time.monotonic() - started_at + if duration >= 0.25: + _log.info( + "plugins/hub rebuilt in %.3fs (plugins=%d memory_options=%d)", + duration, + len(rows), + len(memory_providers), + ) + with _plugins_hub_cache_lock: + _plugins_hub_cache = payload + _plugins_hub_cache_expires_at = time.monotonic() + _PLUGINS_HUB_CACHE_TTL_SECONDS + return payload @app.get("/api/dashboard/plugins/hub") @@ -16772,6 +16967,7 @@ async def post_agent_plugin_install(request: Request, body: _AgentPluginInstallB detail=result.get("error") or "Install failed.", ) _get_dashboard_plugins(force_rescan=True) + _invalidate_plugins_hub_cache() # Strip internal paths from the response result.pop("after_install_path", None) return result @@ -16794,6 +16990,7 @@ async def post_agent_plugin_enable(request: Request, name: str): result = dashboard_set_agent_plugin_enabled(name, enabled=True) if not result.get("ok"): raise HTTPException(status_code=400, detail=result.get("error") or "Enable failed.") + _invalidate_plugins_hub_cache() return result @@ -16806,6 +17003,7 @@ async def post_agent_plugin_disable(request: Request, name: str): result = dashboard_set_agent_plugin_enabled(name, enabled=False) if not result.get("ok"): raise HTTPException(status_code=400, detail=result.get("error") or "Disable failed.") + _invalidate_plugins_hub_cache() return result @@ -16819,6 +17017,7 @@ async def post_agent_plugin_update(request: Request, name: str): if not result.get("ok"): raise HTTPException(status_code=400, detail=result.get("error") or "Update failed.") _get_dashboard_plugins(force_rescan=True) + _invalidate_plugins_hub_cache() return result @@ -16832,6 +17031,7 @@ async def delete_agent_plugin(request: Request, name: str): if not result.get("ok"): raise HTTPException(status_code=400, detail=result.get("error") or "Remove failed.") _get_dashboard_plugins(force_rescan=True) + _invalidate_plugins_hub_cache() return result @@ -16850,6 +17050,7 @@ async def put_plugin_providers(request: Request, body: _PluginProvidersPutBody): _save_memory_provider(memory_provider) if body.context_engine is not None: _save_context_engine(body.context_engine) + _invalidate_plugins_hub_cache() return {"ok": True} @@ -16873,6 +17074,7 @@ async def post_plugin_visibility(request: Request, name: str, body: _PluginVisib config["dashboard"]["hidden_plugins"] = hidden_list save_config(config) + _invalidate_plugins_hub_cache() return {"ok": True, "name": name, "hidden": body.hidden} diff --git a/hermes_state.py b/hermes_state.py index e7fc109980..9b8215390d 100644 --- a/hermes_state.py +++ b/hermes_state.py @@ -17,6 +17,7 @@ Key design decisions: import asyncio import atexit import errno +import hashlib import json import logging import os @@ -85,6 +86,10 @@ logger = logging.getLogger(__name__) _COMPRESSION_LOCK_HOLDER_PID_RE = re.compile(r"(?:^|:)pid=(\d+)(?::|$)") +def _system_prompt_hash(system_prompt: str) -> str: + return hashlib.sha256(system_prompt.encode("utf-8")).hexdigest() + + def _compression_lock_holder_process_is_dead(holder: str) -> bool: """Return True only when a structured lock holder's local PID is gone. @@ -400,6 +405,57 @@ def _strip_background_review_harness( return out +# Matches a bare protocol/tool-name marker such as "[memory]" or "[skill_manage]". +_STALE_TOOL_CALL_MARKER_RE = re.compile(r"^\[[A-Za-z_][A-Za-z0-9_.-]*\]$") + + +def _is_stale_tool_call_marker_message(msg: Dict[str, Any]) -> bool: + """True when ``msg`` is a persisted assistant turn whose content is a bare + bracketed marker (e.g. ``[memory]``) left over from a tool-call turn. + + Before the #78148 fix in ``agent.conversation_loop``, a local tool-call + template could emit a bare marker as assistant content alongside a real + tool call. The loop cached that marker as a fallback and later replayed + it as the "final response", persisting it into the session. Sessions + written before the fix can still carry these rows. + """ + if not isinstance(msg, dict): + return False + if msg.get("role") != "assistant": + return False + if not msg.get("tool_calls"): + return False + content = msg.get("content") + if not isinstance(content, str): + return False + return bool(_STALE_TOOL_CALL_MARKER_RE.fullmatch(content.strip())) + + +def _strip_stale_tool_call_markers( + messages: List[Dict[str, Any]], +) -> List[Dict[str, Any]]: + """Clear bare protocol-marker content persisted before the #78148 fix. + + Replaying "[memory]" as if the model had actually answered teaches the + model, by example, to keep emitting the same marker in later turns — the + exact symptom the issue reported. Only the stray ``content`` field is + blanked; the tool call and its result are left untouched so provider + tool_call/tool_result pairing stays intact. Sessions with no affected + rows pass through unchanged. + """ + repaired = 0 + for msg in messages: + if _is_stale_tool_call_marker_message(msg): + msg["content"] = "" + repaired += 1 + if repaired: + logger.info( + "Cleared %d stale tool-call marker message(s) while restoring session (#78148)", + repaired, + ) + return messages + + def format_session_db_unavailable(prefix: str = "Session database not available") -> str: """Format a user-facing 'session DB unavailable' message with cause. @@ -792,6 +848,34 @@ def _apply_delete_for_wal_reset_bug( return "delete" +def _wal_reset_repair_hint() -> str: + """Return a context-appropriate hint for repairing the SQLite runtime. + + Uses the codebase's install-type detection so the hint matches what + ``hermes update`` can actually do for this install (#75153). + """ + try: + from hermes_cli.config import ( + detect_install_method, + recommended_update_command_for_method, + get_project_root, + ) + method = detect_install_method(get_project_root()) + cmd = recommended_update_command_for_method(method) + if method in {"git", "unknown"}: + return f"Hermes-managed installs can repair the embedded runtime with `{cmd}`" + if method == "docker": + return f"update the container image with `{cmd}`" + # nix/nixos + return cmd + except Exception: + pass + return ( + "install a Python build bundled with SQLite 3.51.3+ " + "(or backports 3.50.7 / 3.44.6) and restart Hermes" + ) + + def _log_wal_reset_bug_once( db_label: str, *, @@ -808,16 +892,20 @@ def _log_wal_reset_bug_once( if kept_wal else "using journal_mode=DELETE instead of enabling WAL" ) + # Check whether this is a Hermes-managed install (uv-managed venv) + # so the warning doesn't promise a repair path that doesn't exist + # for git/pip/system Python installs (#75153). + repair_hint = _wal_reset_repair_hint() logger.warning( "%s: linked SQLite %s is vulnerable to the WAL-reset corruption " "bug (https://sqlite.org/wal.html#walresetbug) — %s. " "Upgrade to SQLite 3.51.3+ (or backports 3.50.7 / 3.44.6); " - "Hermes-managed installs can repair the embedded runtime with " - "`hermes update`. See `hermes doctor`. This warning fires once per " + "%s. See `hermes doctor`. This warning fires once per " "process per database.", db_label, sqlite3.sqlite_version, action, + repair_hint, ) @@ -855,16 +943,23 @@ def apply_database_pragmas( *, db_label: str = "state.db", ) -> None: - """Apply optional WAL-sizing PRAGMAs from ``config.yaml``. + """Apply optional performance and WAL-sizing PRAGMAs from ``config.yaml``. - Reads the ``database:`` section and applies ``wal_autocheckpoint`` - and ``journal_size_limit`` when set to integer values. The journal - mode itself is NOT handled here — ``database.journal_mode`` is owned - by :func:`resolve_journal_mode` inside :func:`apply_wal_with_fallback`, - which layers the operator setting under all the safety guards - (never live-downgrading an on-disk WAL DB, filesystem fallback, - WAL-reset-bug gating). Keeping a single owner prevents a second, - unguarded journal-mode switch path. + Reads the ``database:`` section and applies configurable PRAGMAs when set + to integer values. The journal mode itself is NOT handled here — + ``database.journal_mode`` is owned by :func:`resolve_journal_mode` inside + :func:`apply_wal_with_fallback`, which layers the operator setting under + all the safety guards (never live-downgrading an on-disk WAL DB, + filesystem fallback, WAL-reset-bug gating). + + Supported keys under ``database:`` in config.yaml: + + * ``cache_size`` — negative value = KiB, positive = pages + (e.g. ``-262144`` = 256 MB page cache) + * ``mmap_size`` — max bytes for memory-mapped I/O (0 = disabled) + * ``temp_store`` — 0=DEFAULT(file), 1=FILE, 2=MEMORY, 3=ALWAYS + * ``wal_autocheckpoint`` — WAL auto-checkpoint threshold in pages + * ``journal_size_limit`` — max journal/WAL size in bytes Best-effort: config load or pragma failures are ignored so DB init never breaks on a malformed ``database:`` section. @@ -877,7 +972,15 @@ def apply_database_pragmas( except Exception: return - for pragma_name in ("wal_autocheckpoint", "journal_size_limit"): + # Performance PRAGMAs (applied to ALL connection types: writer, read_only, + # and WAL per-thread readers). + for pragma_name in ( + "cache_size", + "mmap_size", + "temp_store", + "wal_autocheckpoint", + "journal_size_limit", + ): raw_value = cfg_get(cfg, "database", pragma_name, default=None) if raw_value is None: continue @@ -1489,7 +1592,8 @@ BEGIN VALUES ('delete', old.id, old.content, old.tool_name, old.tool_calls); END; -CREATE TRIGGER IF NOT EXISTS messages_fts_cjk_update AFTER UPDATE ON messages +CREATE TRIGGER IF NOT EXISTS messages_fts_cjk_update +AFTER UPDATE OF content, tool_name, tool_calls, role ON messages WHEN (old.content IS NOT new.content OR old.tool_name IS NOT new.tool_name OR old.tool_calls IS NOT new.tool_calls @@ -1852,6 +1956,36 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) _IMPORT_MAX_SESSION_BYTES = 5 * 1024 * 1024 _IMPORT_MAX_TOTAL_BYTES = 25 * 1024 * 1024 + @staticmethod + def _store_system_prompt(conn, system_prompt: Optional[str]) -> Optional[str]: + if system_prompt is None: + return None + prompt_hash = _system_prompt_hash(system_prompt) + conn.execute( + "INSERT OR IGNORE INTO system_prompts (hash, prompt) VALUES (?, ?)", + (prompt_hash, system_prompt), + ) + return prompt_hash + + @staticmethod + def _delete_unreferenced_system_prompts(conn) -> None: + conn.execute( + "DELETE FROM system_prompts " + "WHERE NOT EXISTS (" + "SELECT 1 FROM sessions " + "WHERE sessions.system_prompt_hash = system_prompts.hash" + ")" + ) + + @staticmethod + def _session_row_dict(row: sqlite3.Row) -> Dict[str, Any]: + data = dict(row) + if "_system_prompt_resolved" in data: + resolved = data.pop("_system_prompt_resolved") + if "system_prompt" in data: + data["system_prompt"] = resolved + return data + def __init__(self, db_path: Path = None, read_only: bool = False): self.db_path = db_path or _default_db_path() self.read_only = read_only @@ -1929,6 +2063,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) # raw-copy for the rest of the process — the writable heal # that follows would then repair WITHOUT its forensic backup. try: + apply_database_pragmas(self._conn, db_label="state.db") cursor = self._conn.cursor() self._fts_enabled = ( self._fts_table_probe(cursor, "messages_fts") is True @@ -2128,6 +2263,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) isolation_level=None, ) conn.row_factory = sqlite3.Row + apply_database_pragmas(conn, db_label="state.db") # Load the CJK tokenizer extension on this connection so # messages_fts_cjk queries work on the read path. The .so # registers the tokenizer in the connection's in-memory @@ -2829,17 +2965,27 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) without a recoverable routing mapping (#59527). """ def _do(conn): + system_prompt_hash = self._store_system_prompt(conn, system_prompt) conn.execute( """INSERT INTO sessions ( id, source, user_id, session_key, chat_id, chat_type, thread_id, - model, model_config, system_prompt, parent_session_id, cwd, - profile_name, git_repo_root, started_at + model, model_config, system_prompt, system_prompt_hash, + parent_session_id, cwd, profile_name, git_repo_root, started_at ) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?) ON CONFLICT(id) DO UPDATE SET model = COALESCE(sessions.model, excluded.model), model_config = COALESCE(sessions.model_config, excluded.model_config), - system_prompt = COALESCE(sessions.system_prompt, excluded.system_prompt), + system_prompt_hash = COALESCE( + sessions.system_prompt_hash, + excluded.system_prompt_hash + ), + system_prompt = CASE + WHEN sessions.system_prompt_hash IS NULL + AND excluded.system_prompt_hash IS NOT NULL + THEN NULL + ELSE sessions.system_prompt + END, session_key = COALESCE(sessions.session_key, excluded.session_key), chat_id = COALESCE(sessions.chat_id, excluded.chat_id), chat_type = COALESCE(sessions.chat_type, excluded.chat_type), @@ -2858,7 +3004,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) thread_id, model, json.dumps(model_config) if model_config else None, - system_prompt, + system_prompt_hash, parent_session_id, cwd, profile_name, @@ -2866,6 +3012,8 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) time.time(), ), ) + if system_prompt_hash is not None: + self._delete_unreferenced_system_prompts(conn) if parent_session_id: conn.execute( """UPDATE sessions @@ -3132,8 +3280,12 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) self.flush_token_counts() query = f""" SELECT sessions.*, + COALESCE(sp.prompt, sessions.system_prompt) + AS _system_prompt_resolved, {_sql_session_last_active("sessions")} AS last_active FROM sessions + LEFT JOIN system_prompts sp + ON sp.hash = sessions.system_prompt_hash WHERE session_key IS NOT NULL AND started_at = ( SELECT MAX(s2.started_at) FROM sessions s2 @@ -3149,7 +3301,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) query += " ORDER BY last_active DESC" with self._lock: rows = self._conn.execute(query, params).fetchall() - return [dict(r) for r in rows] + return [self._session_row_dict(r) for r in rows] def find_session_by_origin( self, @@ -3227,20 +3379,24 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) with self._lock: row = self._conn.execute( """ - SELECT * FROM sessions - WHERE session_key = ? - AND source = ? - AND (ended_at IS NULL OR end_reason IN ('agent_close', 'ws_orphan_reap')) - AND (COALESCE(message_count, 0) > 0 OR EXISTS ( - SELECT 1 FROM messages WHERE messages.session_id = sessions.id LIMIT 1 + SELECT s.*, + COALESCE(sp.prompt, s.system_prompt) + AS _system_prompt_resolved + FROM sessions s + LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash + WHERE s.session_key = ? + AND s.source = ? + AND (s.ended_at IS NULL OR s.end_reason IN ('agent_close', 'ws_orphan_reap')) + AND (COALESCE(s.message_count, 0) > 0 OR EXISTS ( + SELECT 1 FROM messages WHERE messages.session_id = s.id LIMIT 1 )) - ORDER BY started_at DESC + ORDER BY s.started_at DESC LIMIT 1 """, (session_key, source), ).fetchone() if row is not None: - return dict(row) + return self._session_row_dict(row) # Conservative fallback for rows created by current code but with a # temporarily-missing exact key: still require the complete peer @@ -3249,22 +3405,26 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) return None row = self._conn.execute( """ - SELECT * FROM sessions - WHERE source = ? - AND COALESCE(user_id, '') = COALESCE(?, '') - AND COALESCE(chat_id, '') = COALESCE(?, '') - AND COALESCE(chat_type, '') = COALESCE(?, '') - AND COALESCE(thread_id, '') = COALESCE(?, '') - AND (ended_at IS NULL OR end_reason IN ('agent_close', 'ws_orphan_reap')) - AND (COALESCE(message_count, 0) > 0 OR EXISTS ( - SELECT 1 FROM messages WHERE messages.session_id = sessions.id LIMIT 1 + SELECT s.*, + COALESCE(sp.prompt, s.system_prompt) + AS _system_prompt_resolved + FROM sessions s + LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash + WHERE s.source = ? + AND COALESCE(s.user_id, '') = COALESCE(?, '') + AND COALESCE(s.chat_id, '') = COALESCE(?, '') + AND COALESCE(s.chat_type, '') = COALESCE(?, '') + AND COALESCE(s.thread_id, '') = COALESCE(?, '') + AND (s.ended_at IS NULL OR s.end_reason IN ('agent_close', 'ws_orphan_reap')) + AND (COALESCE(s.message_count, 0) > 0 OR EXISTS ( + SELECT 1 FROM messages WHERE messages.session_id = s.id LIMIT 1 )) - ORDER BY started_at DESC + ORDER BY s.started_at DESC LIMIT 1 """, (source, user_id, chat_id, chat_type, thread_id), ).fetchone() - return dict(row) if row else None + return self._session_row_dict(row) if row else None def find_live_compression_child( self, parent_session_id: str @@ -3292,18 +3452,22 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) return None rows = self._conn.execute( """ - SELECT * FROM sessions - WHERE parent_session_id = ? - AND ended_at IS NULL - AND json_extract(COALESCE(model_config, '{}'), '$._branched_from') IS NULL - AND json_extract(COALESCE(model_config, '{}'), '$._delegate_from') IS NULL - AND COALESCE(source, '') != 'tool' - ORDER BY started_at ASC + SELECT s.*, + COALESCE(sp.prompt, s.system_prompt) + AS _system_prompt_resolved + FROM sessions s + LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash + WHERE s.parent_session_id = ? + AND s.ended_at IS NULL + AND json_extract(COALESCE(s.model_config, '{}'), '$._branched_from') IS NULL + AND json_extract(COALESCE(s.model_config, '{}'), '$._delegate_from') IS NULL + AND COALESCE(s.source, '') != 'tool' + ORDER BY s.started_at ASC LIMIT 2 """, (parent_session_id,), ).fetchall() - return dict(rows[0]) if len(rows) == 1 else None + return self._session_row_dict(rows[0]) if len(rows) == 1 else None def publish_compression_child( self, @@ -3353,20 +3517,22 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) raise RuntimeError(f"Compression parent already ended: {parent_session_id}") if not messages: raise RuntimeError("Compression child handoff must not be empty") + system_prompt_hash = self._store_system_prompt(conn, system_prompt) conn.execute( """INSERT INTO sessions ( id, source, model, model_config, system_prompt, + system_prompt_hash, parent_session_id, cwd, git_branch, git_repo_root, profile_name, user_id, session_key, chat_id, chat_type, thread_id, display_name, origin_json, started_at - ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", + ) VALUES (?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", ( child_session_id, source, model, json.dumps(model_config) if model_config else None, - system_prompt, + system_prompt_hash, parent_session_id, cwd or parent["cwd"], parent["git_branch"], @@ -4122,13 +4288,18 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) ) self._execute_write(_do) - def update_system_prompt(self, session_id: str, system_prompt: str) -> None: + def update_system_prompt( + self, session_id: str, system_prompt: Optional[str] + ) -> None: """Store the full assembled system prompt snapshot.""" def _do(conn): + system_prompt_hash = self._store_system_prompt(conn, system_prompt) conn.execute( - "UPDATE sessions SET system_prompt = ? WHERE id = ?", - (system_prompt, session_id), + "UPDATE sessions " + "SET system_prompt_hash = ?, system_prompt = NULL WHERE id = ?", + (system_prompt_hash, session_id), ) + self._delete_unreferenced_system_prompts(conn) self._execute_write(_do) def update_session_model(self, session_id: str, model: str) -> None: @@ -4160,10 +4331,12 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) THEN json_remove(model_config, '$.browser_model_lock') ELSE model_config END, - system_prompt = NULL + system_prompt = NULL, + system_prompt_hash = NULL WHERE id = ?""", (model, session_id), ) + self._delete_unreferenced_system_prompts(conn) self._execute_write(_do) def update_session_runtime_lock( @@ -4214,12 +4387,72 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) """UPDATE sessions SET model_config = ?, model = COALESCE(?, model), - system_prompt = NULL + system_prompt = NULL, + system_prompt_hash = NULL WHERE id = ?""", (json.dumps(config), model, session_id), ) + self._delete_unreferenced_system_prompts(conn) self._execute_write(_do) + def set_session_yolo(self, session_id: str, enabled: bool) -> None: + """Persist the per-session YOLO bypass flag into ``model_config``. + + Merges ``yolo_mode`` into the existing ``model_config`` JSON (same + merge discipline as ``update_session_runtime_lock`` so lineage + markers like ``_branched_from`` / ``_delegate_from`` survive). The + CLI resume paths read this flag back so a ``/yolo ON`` toggle — or a + ``--yolo`` launch — survives ``hermes --resume`` into a fresh + process. No-op when the session row doesn't exist yet; the + creation-time ``model_config`` carries the flag for ``--yolo`` + launches. + """ + if not session_id: + return + + def _do(conn): + row = conn.execute( + "SELECT model_config FROM sessions WHERE id = ?", + (session_id,), + ).fetchone() + if row is None: + return + raw = row["model_config"] if isinstance(row, sqlite3.Row) else row[0] + config: Dict[str, Any] = {} + if isinstance(raw, str) and raw.strip(): + try: + parsed = json.loads(raw) + if isinstance(parsed, dict): + config = parsed + except Exception: + config = {} + elif isinstance(raw, dict): + config = dict(raw) + config["yolo_mode"] = bool(enabled) + conn.execute( + "UPDATE sessions SET model_config = ? WHERE id = ?", + (json.dumps(config), session_id), + ) + self._execute_write(_do) + + @staticmethod + def session_yolo_enabled(session_meta: Optional[Dict[str, Any]]) -> bool: + """Read the persisted YOLO flag off a session row dict. + + Accepts the dict returned by ``get_session`` (``model_config`` is a + JSON string) or an already-parsed dict. Returns False on any parse + failure — resume must never enable the bypass by accident. + """ + raw = (session_meta or {}).get("model_config") + if isinstance(raw, str): + try: + raw = json.loads(raw) + except Exception: + return False + if not isinstance(raw, dict): + return False + return bool(raw.get("yolo_mode")) + def update_session_billing_route( self, session_id: str, @@ -4247,10 +4480,12 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) billing_provider = ?, billing_base_url = ?, billing_mode = COALESCE(?, billing_mode), - system_prompt = NULL + system_prompt = NULL, + system_prompt_hash = NULL WHERE id = ?""", (provider, base_url, billing_mode, session_id), ) + self._delete_unreferenced_system_prompts(conn) self._execute_write(_do) # ── Async token accounting ── @@ -4872,6 +5107,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) conn.execute( f"DELETE FROM sessions WHERE id IN ({placeholders})", ids ) + self._delete_unreferenced_system_prompts(conn) return ids removed_ids = self._execute_write(_do) or [] @@ -4928,10 +5164,15 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) self.flush_token_counts() with self._read_ctx() as conn: cursor = conn.execute( - "SELECT * FROM sessions WHERE id = ?", (session_id,) + "SELECT s.*, " + "COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved " + "FROM sessions s " + "LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash " + "WHERE s.id = ?", + (session_id,), ) row = cursor.fetchone() - return dict(row) if row else None + return self._session_row_dict(row) if row else None def resolve_session_id(self, session_id_or_prefix: str) -> Optional[str]: """Resolve an exact or uniquely prefixed session ID to the full ID. @@ -5241,10 +5482,15 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) """Look up a session by exact title. Returns session dict or None.""" with self._read_ctx() as conn: cursor = conn.execute( - "SELECT * FROM sessions WHERE title = ?", (title,) + "SELECT s.*, " + "COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved " + "FROM sessions s " + "LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash " + "WHERE s.title = ?", + (title,), ) row = cursor.fetchone() - return dict(row) if row else None + return self._session_row_dict(row) if row else None def resolve_session_by_title(self, title: str) -> Optional[str]: """Resolve a title to a session ID, preferring the latest in a lineage. @@ -5376,7 +5622,9 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) # the projection is derived from SCHEMA_SQL so columns added later via # declarative reconciliation are included automatically instead of # silently dropping out of list rows. - _SESSION_COMPACT_EXCLUDED = frozenset({"system_prompt"}) + _SESSION_COMPACT_EXCLUDED = frozenset( + {"system_prompt", "system_prompt_hash"} + ) _session_compact_cols_sql: Optional[str] = None def list_sessions_rich( @@ -5503,6 +5751,14 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) # Snapshot the filter params before the query builders below extend # them with LIMIT/OFFSET — the pinned back-fill reuses the same WHERE. base_where_params = list(params) + prompt_select = ( + "" if compact_rows + else ", COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved" + ) + prompt_join = ( + "" if compact_rows + else "LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash" + ) # Optional session-id filter, pushed into SQL so callers (Desktop # session-id search) don't have to fetch every row and filter in @@ -5599,7 +5855,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) FROM chain GROUP BY root_id ) - SELECT {_sel}, + SELECT {_sel}{prompt_select}, COALESCE( (SELECT {_PREVIEW_RAW_SELECT} FROM messages m @@ -5611,6 +5867,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) COALESCE(cm.effective_last_active, s.started_at) AS _effective_last_active FROM sessions s LEFT JOIN chain_max cm ON cm.root_id = s.id + {prompt_join} {outer_where} ORDER BY _effective_last_active DESC, s.started_at DESC, s.id DESC LIMIT ? OFFSET ? @@ -5621,7 +5878,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) else: _sel = self._compact_session_cols() if compact_rows else "s.*" query = f""" - SELECT {_sel}, + SELECT {_sel}{prompt_select}, COALESCE( (SELECT {_PREVIEW_RAW_SELECT} FROM messages m @@ -5631,6 +5888,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) ) AS _preview_raw, {_sql_session_last_active("s")} AS last_active FROM sessions s + {prompt_join} {where_sql} ORDER BY s.started_at DESC LIMIT ? OFFSET ? @@ -5641,7 +5899,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) rows = cursor.fetchall() sessions = [] for row in rows: - s = dict(row) + s = self._session_row_dict(row) s["preview"] = _shape_preview(s.pop("_preview_raw", "")) # Drop the internal ordering column so callers see a clean dict. s.pop("_effective_last_active", None) @@ -5659,7 +5917,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) ) _sel = self._compact_session_cols() if compact_rows else "s.*" pinned_query = f""" - SELECT {_sel}, + SELECT {_sel}{prompt_select}, COALESCE( (SELECT {_PREVIEW_RAW_SELECT} FROM messages m @@ -5672,6 +5930,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) s.started_at ) AS last_active FROM sessions s + {prompt_join} {pinned_where} ORDER BY s.started_at DESC """ @@ -5679,7 +5938,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) pinned_cursor = conn.execute(pinned_query, base_where_params) pinned_rows = pinned_cursor.fetchall() for row in pinned_rows: - s = dict(row) + s = self._session_row_dict(row) if s["id"] in seen_ids: continue s["preview"] = _shape_preview(s.pop("_preview_raw", "")) @@ -5693,16 +5952,31 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) # as the live conversation. Keep the root's started_at to preserve # chronological ordering by original conversation start. if project_compression_tips and not include_children: - projected = [] + # get_compression_tip() walks each root's chain individually (it's + # a per-session graph walk, not batchable in one query), but the + # tip *row* fetch afterward was previously one _get_session_rich_row() + # call per compression root. Batch that half instead: resolve + # every tip id first, then fetch all tip rows in a single query. + tip_ids_by_root: Dict[str, str] = {} for s in sessions: if s.get("end_reason") != "compression": - projected.append(s) continue tip_id = self.get_compression_tip(s["id"]) - if tip_id == s["id"]: - projected.append(s) - continue - tip_row = self._get_session_rich_row(tip_id, compact_rows=compact_rows) + if tip_id != s["id"]: + tip_ids_by_root[s["id"]] = tip_id + + tip_rows = ( + self._get_session_rich_rows_batch( + set(tip_ids_by_root.values()), compact_rows=compact_rows + ) + if tip_ids_by_root + else {} + ) + + projected = [] + for s in sessions: + tip_id = tip_ids_by_root.get(s["id"]) + tip_row = tip_rows.get(tip_id) if tip_id else None if not tip_row: projected.append(s) continue @@ -5810,6 +6084,39 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) ) return None + def _check_transcript_write_guards( + self, conn, session_id: str, compression_lock_holder: Optional[str] + ) -> None: + """Transcript-append admission checks, run INSIDE the write txn. + + Shared by :meth:`append_message` and :meth:`append_messages_batch` so + the two writers can never diverge on these correctness invariants + (this guard has already needed targeted fixes — see the #74478 + patience note below). + """ + active_lock = conn.execute( + "SELECT holder FROM compression_locks " + "WHERE session_id = ? AND expires_at > ?", + (session_id, time.time()), + ).fetchone() + if ( + active_lock is not None + and active_lock["holder"] != compression_lock_holder + ): + raise SessionCompressionInProgressError( + f"Session {session_id!r} is being compressed by another writer" + ) + session = conn.execute( + "SELECT ended_at, end_reason FROM sessions WHERE id = ?", + (session_id,), + ).fetchone() + if ( + session is not None + and session["ended_at"] is not None + and session["end_reason"] == "compression" + ): + raise CompressionSessionClosedError(session_id) + @staticmethod def _decode_display_metadata(raw: Any) -> Optional[Dict[str, Any]]: """Decode a ``display_metadata`` column into the dict every reader expects. @@ -5922,28 +6229,9 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) num_tool_calls = len(tool_calls) if isinstance(tool_calls, list) else 1 def _do(conn): - active_lock = conn.execute( - "SELECT holder FROM compression_locks " - "WHERE session_id = ? AND expires_at > ?", - (session_id, time.time()), - ).fetchone() - if ( - active_lock is not None - and active_lock["holder"] != compression_lock_holder - ): - raise SessionCompressionInProgressError( - f"Session {session_id!r} is being compressed by another writer" - ) - session = conn.execute( - "SELECT ended_at, end_reason FROM sessions WHERE id = ?", - (session_id,), - ).fetchone() - if ( - session is not None - and session["ended_at"] is not None - and session["end_reason"] == "compression" - ): - raise CompressionSessionClosedError(session_id) + self._check_transcript_write_guards( + conn, session_id, compression_lock_holder + ) cursor = conn.execute( """INSERT INTO messages (session_id, role, content, tool_call_id, tool_calls, tool_name, effect_disposition, timestamp, token_count, finish_reason, @@ -5999,6 +6287,79 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) _do, patience_s=self._TRANSCRIPT_WRITE_PATIENCE_S ) + def append_messages_batch( + self, + session_id: str, + messages: List[Dict[str, Any]], + compression_lock_holder: Optional[str] = None, + chunk_rows: Optional[int] = None, + ) -> int: + """Append multiple messages atomically in ONE write transaction. + + ``messages`` is a list of dicts in the same shape + :meth:`_insert_message_rows` already consumes for replace/compact/ + import (role, content, tool_name, tool_calls, tool_call_id, + finish_reason, reasoning*, codex_*, timestamp, api_content, + display_kind, display_metadata, ...). Reusing that helper keeps ONE + row-serialization path for every multi-row writer. + + A turn-boundary flush writes the whole turn (user + assistant + tool + rows, typically 3-8 messages) as one BEGIN IMMEDIATE / commit pair + instead of one transaction (and, off WAL, one fsync) per row. + + Atomicity contract: all rows land or none do (the caller re-flushes + unstamped messages on the next attempt). The same admission guards + as :meth:`append_message` run once for the batch — same session, + same instant. + + ``chunk_rows`` bounds the transaction size for LARGE copies (branch + seeds can be thousands of rows; measured: 10k rows ≈ 2.4s inside one + BEGIN IMMEDIATE because the FTS triggers run per row, which would + monopolize the write lock and starve concurrent writers). When set, + the batch commits in chunks of at most that many rows — same + recovery semantics as the old per-row loops (a mid-copy failure + leaves a partial seed), just with bounded lock holds. A turn flush + never needs it. Returns the inserted row count. + """ + if not messages: + return 0 + + if chunk_rows is not None and len(messages) > chunk_rows: + inserted_total = 0 + for start in range(0, len(messages), chunk_rows): + inserted_total += self.append_messages_batch( + session_id, + messages[start:start + chunk_rows], + compression_lock_holder=compression_lock_holder, + ) + return inserted_total + + def _do(conn): + self._check_transcript_write_guards( + conn, session_id, compression_lock_holder + ) + inserted, tool_calls_total = self._insert_message_rows( + conn, session_id, messages + ) + # One aggregated counter update for the whole batch. + if tool_calls_total > 0: + conn.execute( + """UPDATE sessions SET message_count = message_count + ?, + tool_call_count = tool_call_count + ? WHERE id = ?""", + (inserted, tool_calls_total, session_id), + ) + else: + conn.execute( + "UPDATE sessions SET message_count = message_count + ? WHERE id = ?", + (inserted, session_id), + ) + return inserted + + # Same criticality as append_message: this IS the turn's transcript. + return self._execute_write( + _do, patience_s=self._TRANSCRIPT_WRITE_PATIENCE_S + ) + def set_latest_matching_message_display_kind( self, session_id: str, *, role: str, content: str, display_kind: str, display_metadata: Optional[Dict[str, Any]] = None, @@ -6746,9 +7107,9 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) session_ids = self._session_lineage_root_to_tip(session_id) active_clause = "" if include_inactive else " AND active = 1" - with self._lock: + with self._read_ctx() as conn: placeholders = ",".join("?" for _ in session_ids) - rows = self._conn.execute( + rows = conn.execute( f"SELECT {self._CONVERSATION_ROW_COLUMNS} " f"FROM messages WHERE session_id IN ({placeholders})" # Order by AUTOINCREMENT id (true insertion order), NOT timestamp: @@ -6890,6 +7251,13 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) # assistant reply immediately following it, so a polluted session # resumes clean even if stray rows exist. messages = _strip_background_review_harness(messages) + # DEFENSE-IN-DEPTH against #78148: before that fix, a bare tool-call + # marker (e.g. "[memory]") could get cached as a fallback and + # persisted as if it were the model's real answer. Sessions written + # before the fix can still carry those rows — clear the stray + # content on load so replaying history doesn't re-teach the model + # to keep emitting the marker. No-op for unaffected sessions. + messages = _strip_stale_tool_call_markers(messages) if repair_alternation and messages: # Lazy import: hermes_state already depends on agent.* (see # sanitize_context above), but keep this optional path from @@ -6927,9 +7295,9 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) output (see test_get_resume_conversations_matches_separate_reads). """ session_ids = self._session_lineage_root_to_tip(session_id) - with self._lock: + with self._read_ctx() as conn: placeholders = ",".join("?" for _ in session_ids) - rows = self._conn.execute( + rows = conn.execute( f"SELECT session_id, {self._CONVERSATION_ROW_COLUMNS} " f"FROM messages WHERE session_id IN ({placeholders}) AND active = 1 " # ORDER BY id (insertion order) — see get_messages_as_conversation @@ -6981,9 +7349,9 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) session_ids = self._session_lineage_root_to_tip(session_id) if len(session_ids) <= 1: return [] - with self._lock: + with self._read_ctx() as conn: placeholders = ",".join("?" for _ in session_ids) - rows = self._conn.execute( + rows = conn.execute( f"SELECT session_id, {self._CONVERSATION_ROW_COLUMNS} " f"FROM messages WHERE session_id IN ({placeholders}) AND active = 1 " "ORDER BY id", @@ -7020,13 +7388,13 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) chain = [] current = session_id seen = set() - with self._lock: + with self._read_ctx() as conn: for _ in range(100): if not current or current in seen: break seen.add(current) chain.append(current) - row = self._conn.execute( + row = conn.execute( "SELECT parent_session_id FROM sessions WHERE id = ?", (current,), ).fetchone() @@ -7187,8 +7555,11 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) the *current* workspace, not the global MRU. """ select_with_last_active = ( - f"SELECT s.*, {_sql_session_last_active('s')} AS last_active " + "SELECT s.*, " + "COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved, " + f"{_sql_session_last_active('s')} AS last_active " "FROM sessions s " + "LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash " ) where_clauses = [] params: list = [] @@ -7208,7 +7579,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) "ORDER BY last_active DESC, s.started_at DESC, s.id DESC LIMIT ? OFFSET ?", params, ) - return [dict(row) for row in cursor.fetchall()] + return [self._session_row_dict(row) for row in cursor.fetchall()] # ========================================================================= # Utility @@ -7275,6 +7646,22 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) cursor = self._conn.execute(f"SELECT COUNT(*) FROM sessions s{where_sql}", params) return cursor.fetchone()[0] + def session_count_ge(self, n: int = 1) -> bool: + """Check if at least N sessions exist (archived included). + + Short-circuits via LIMIT — much cheaper than ``session_count()``, + which pays a full index scan for its default ``archived = 0`` + filter (measured 543us vs 4us on a 20k-session DB). Archived + sessions count: every caller so far asks "has this install ever + had sessions", and an archived session is still a created one. + Use this instead of ``session_count() >= n`` when the exact count + is irrelevant. + """ + with self._lock: + cursor = self._conn.execute("SELECT 1 FROM sessions LIMIT ?", (n,)) + rows = cursor.fetchall() + return len(rows) >= n + def session_count_by_source( self, *, @@ -7500,9 +7887,9 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) def _do(conn): cursor = conn.execute( - "SELECT COUNT(*) FROM sessions WHERE id = ?", (session_id,) + "SELECT 1 FROM sessions WHERE id = ? LIMIT 1", (session_id,) ) - if cursor.fetchone()[0] == 0: + if cursor.fetchone() is None: return False if expected_ids is not None: actual_ids = { @@ -7520,6 +7907,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) ) conn.execute("DELETE FROM messages WHERE session_id = ?", (session_id,)) conn.execute("DELETE FROM sessions WHERE id = ?", (session_id,)) + self._delete_unreferenced_system_prompts(conn) return True deleted = self._execute_write(_do) @@ -7564,6 +7952,8 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) """, (session_id,), ) + if cursor.rowcount > 0: + self._delete_unreferenced_system_prompts(conn) return cursor.rowcount > 0 deleted = self._execute_write(_do) @@ -7644,6 +8034,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) f"DELETE FROM sessions WHERE id IN ({existing_placeholders})", existing, ) + self._delete_unreferenced_system_prompts(conn) removed_ids.extend(existing) return len(existing) @@ -7739,6 +8130,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) ) conn.execute("DELETE FROM sessions WHERE id = ?", (sid,)) removed_ids.append(sid) + self._delete_unreferenced_system_prompts(conn) return len(session_ids) count = self._execute_write(_do) @@ -8074,6 +8466,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) conn.execute("DELETE FROM messages WHERE session_id = ?", (sid,)) conn.execute("DELETE FROM sessions WHERE id = ?", (sid,)) removed_ids.append(sid) + self._delete_unreferenced_system_prompts(conn) return len(session_ids) count = self._execute_write(_do) @@ -8082,6 +8475,108 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) self._remove_session_files(sessions_dir, sid) return count + def purge_stale_tool_call_markers( + self, *, dry_run: bool = False, backup: bool = True + ) -> Dict[str, Any]: + """Permanently clear bare tool-call marker content (e.g. "[memory]") + left in the ``messages`` table by sessions persisted before the + #78148 fix in ``agent.conversation_loop``. + + ``_strip_stale_tool_call_markers`` already repairs this in memory on + every session load (see ``_rows_to_conversation``), so running this + is optional — but for long-lived sessions the same rows get + re-scanned and re-repaired on every resume, which is wasted work + and keeps the contaminated bytes sitting in the DB (and in any + downstream cache/backup snapshot of it) indefinitely. This rewrites + the affected rows once, in place. + + Only the ``content`` column is touched — ``role``, ``tool_calls``, + and every other column on the row are left exactly as they are, so + provider tool_call/tool_result pairing is unaffected. + + Unlike the in-memory repair, this UPDATE is permanent and can't be + undone from within the DB. Since ``backup`` defaults to True, a + timestamped full snapshot is taken via ``VACUUM INTO`` (safe against + a live connection, unlike the raw-copy ``_backup_db_file`` used for + malformed-schema repair) before any row is touched — mirroring + ``repair_state_db_schema``'s backup-by-default convention for + destructive state.db operations. No snapshot is taken when there is + nothing to change. + + With ``dry_run=True``, reports the affected row count/ids without + writing or backing up (read-only, no write lock taken). + + Returns ``{"dry_run": bool, "rows_affected": int, "row_ids": [...], + "backup_path": str|None}``. + """ + + def _find_affected(conn) -> List[int]: + cursor = conn.execute( + "SELECT id, content FROM messages " + "WHERE role = 'assistant' AND tool_calls IS NOT NULL AND tool_calls != ''" + ) + affected: List[int] = [] + for row in cursor.fetchall(): + content = row["content"] + if isinstance(content, str) and _STALE_TOOL_CALL_MARKER_RE.fullmatch(content.strip()): + affected.append(row["id"]) + return affected + + with self._read_ctx() as conn: + affected_ids = _find_affected(conn) + + if dry_run: + return { + "dry_run": True, + "rows_affected": len(affected_ids), + "row_ids": affected_ids, + "backup_path": None, + } + + if not affected_ids: + return { + "dry_run": False, + "rows_affected": 0, + "row_ids": [], + "backup_path": None, + } + + backup_path: Optional[str] = None + if backup: + import datetime + + stamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S") + dest = self.db_path.with_name( + f"{self.db_path.name}.pre-clean-markers-backup-{stamp}" + ) + with self._lock: + self._conn.execute("VACUUM INTO ?", (str(dest),)) + backup_path = str(dest) + logger.info("Backed up state.db to %s before clean-markers write", backup_path) + + def _do(conn): + ids = _find_affected(conn) + if ids: + placeholders = ",".join("?" * len(ids)) + conn.execute( + f"UPDATE messages SET content = '' WHERE id IN ({placeholders})", + ids, + ) + return ids + + affected_ids = self._execute_write(_do) + if affected_ids: + logger.info( + "Permanently cleared %d stale tool-call marker row(s) in state.db (#78148)", + len(affected_ids), + ) + return { + "dry_run": False, + "rows_affected": len(affected_ids), + "row_ids": affected_ids, + "backup_path": backup_path, + } + # ── Meta key/value (for scheduler bookkeeping) ── def get_meta(self, key: str) -> Optional[str]: @@ -8609,6 +9104,8 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) rows = self._conn.execute( f""" SELECT s.*, + COALESCE(sp.prompt, s.system_prompt) + AS _system_prompt_resolved, COALESCE( (SELECT {_PREVIEW_RAW_SELECT} FROM messages m @@ -8618,6 +9115,8 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) ) AS _preview_raw, {_sql_session_last_active("s")} AS last_active FROM sessions s + LEFT JOIN system_prompts sp + ON sp.hash = s.system_prompt_hash WHERE s.source = 'telegram' AND s.user_id = ? AND NOT EXISTS ( @@ -8635,6 +9134,8 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) rows = self._conn.execute( f""" SELECT s.*, + COALESCE(sp.prompt, s.system_prompt) + AS _system_prompt_resolved, COALESCE( (SELECT {_PREVIEW_RAW_SELECT} FROM messages m @@ -8644,6 +9145,8 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) ) AS _preview_raw, {_sql_session_last_active("s")} AS last_active FROM sessions s + LEFT JOIN system_prompts sp + ON sp.hash = s.system_prompt_hash WHERE s.source = 'telegram' AND s.user_id = ? ORDER BY last_active DESC, s.started_at DESC @@ -8654,7 +9157,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) sessions: List[Dict[str, Any]] = [] for row in rows: - session = dict(row) + session = self._session_row_dict(row) session["preview"] = _shape_preview(session.pop("_preview_raw", "")) sessions.append(session) return sessions @@ -8937,11 +9440,14 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin) """ try: cur = self._conn.execute( - "SELECT * FROM sessions " - "WHERE handoff_state = 'pending' " - "ORDER BY started_at ASC" + "SELECT s.*, " + "COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved " + "FROM sessions s " + "LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash " + "WHERE s.handoff_state = 'pending' " + "ORDER BY s.started_at ASC" ) - return [dict(r) for r in cur.fetchall()] + return [self._session_row_dict(r) for r in cur.fetchall()] except Exception: return [] diff --git a/hermes_state_common.py b/hermes_state_common.py index 68c2b3f0e0..c520f1c51d 100644 --- a/hermes_state_common.py +++ b/hermes_state_common.py @@ -152,7 +152,7 @@ def _sql_session_last_active_by_id(session_id_expr: str) -> str: ) -SCHEMA_VERSION = 24 +SCHEMA_VERSION = 25 # FTS storage-layout version, tracked INDEPENDENTLY of SCHEMA_VERSION in the @@ -187,6 +187,11 @@ CREATE TABLE IF NOT EXISTS schema_version ( version INTEGER NOT NULL ); +CREATE TABLE IF NOT EXISTS system_prompts ( + hash TEXT PRIMARY KEY, + prompt TEXT NOT NULL +); + CREATE TABLE IF NOT EXISTS sessions ( id TEXT PRIMARY KEY, source TEXT NOT NULL, @@ -201,6 +206,7 @@ CREATE TABLE IF NOT EXISTS sessions ( model TEXT, model_config TEXT, system_prompt TEXT, + system_prompt_hash TEXT, parent_session_id TEXT, started_at REAL NOT NULL, ended_at REAL, @@ -239,7 +245,8 @@ CREATE TABLE IF NOT EXISTS sessions ( rewind_count INTEGER NOT NULL DEFAULT 0, archived INTEGER NOT NULL DEFAULT 0, pinned INTEGER NOT NULL DEFAULT 0, - FOREIGN KEY (parent_session_id) REFERENCES sessions(id) + FOREIGN KEY (parent_session_id) REFERENCES sessions(id), + FOREIGN KEY (system_prompt_hash) REFERENCES system_prompts(hash) ); CREATE TABLE IF NOT EXISTS messages ( @@ -337,6 +344,14 @@ CREATE INDEX IF NOT EXISTS idx_sessions_parent ON sessions(parent_session_id); CREATE INDEX IF NOT EXISTS idx_sessions_started ON sessions(started_at DESC); CREATE INDEX IF NOT EXISTS idx_messages_session ON messages(session_id, timestamp); CREATE INDEX IF NOT EXISTS idx_messages_session_id ON messages(session_id, id); +-- Partial index for the Insights assistant tool-call scan +-- (agent/insights.py _get_tool_usage / _get_skill_usage): those queries filter +-- messages by role='assistant' AND tool_calls IS NOT NULL, a small fraction of +-- rows on a large state.db. role and tool_calls are base columns, so this can +-- live in SCHEMA_SQL rather than DEFERRED_INDEX_SQL. +CREATE INDEX IF NOT EXISTS idx_messages_assistant_calls_by_session + ON messages(session_id) + WHERE role = 'assistant' AND tool_calls IS NOT NULL; CREATE INDEX IF NOT EXISTS idx_compression_locks_expires ON compression_locks(expires_at); CREATE INDEX IF NOT EXISTS idx_session_model_usage_session ON session_model_usage(session_id); CREATE INDEX IF NOT EXISTS idx_session_model_usage_model ON session_model_usage(model); @@ -360,6 +375,8 @@ CREATE INDEX IF NOT EXISTS idx_sessions_gateway_peer ON sessions(source, user_id, chat_id, chat_type, thread_id, started_at DESC); CREATE INDEX IF NOT EXISTS idx_sessions_handoff_state ON sessions(handoff_state, started_at); +CREATE INDEX IF NOT EXISTS idx_sessions_system_prompt_hash + ON sessions(system_prompt_hash); """ @@ -411,7 +428,11 @@ BEGIN VALUES ('delete', old.id, old.content, old.tool_name, old.tool_calls); END; -CREATE TRIGGER IF NOT EXISTS messages_fts_update AFTER UPDATE ON messages +-- UPDATE OF skips the trigger entirely for non-content column writes +-- (status/compacted/observed/etc.), which is stronger than the WHEN gate +-- alone and avoids FTS I/O saturation on large state.db (#68858 / #73639). +CREATE TRIGGER IF NOT EXISTS messages_fts_update +AFTER UPDATE OF content, tool_name, tool_calls ON messages WHEN (old.content IS NOT new.content OR old.tool_name IS NOT new.tool_name OR old.tool_calls IS NOT new.tool_calls) @@ -479,7 +500,8 @@ BEGIN VALUES ('delete', old.id, old.content, old.tool_name, old.tool_calls); END; -CREATE TRIGGER IF NOT EXISTS messages_fts_trigram_update AFTER UPDATE ON messages +CREATE TRIGGER IF NOT EXISTS messages_fts_trigram_update +AFTER UPDATE OF content, tool_name, tool_calls, role ON messages WHEN (old.content IS NOT new.content OR old.tool_name IS NOT new.tool_name OR old.tool_calls IS NOT new.tool_calls @@ -540,7 +562,8 @@ CREATE TRIGGER IF NOT EXISTS messages_fts_delete AFTER DELETE ON messages BEGIN DELETE FROM messages_fts WHERE rowid = old.id; END; -CREATE TRIGGER IF NOT EXISTS messages_fts_update AFTER UPDATE ON messages BEGIN +CREATE TRIGGER IF NOT EXISTS messages_fts_update +AFTER UPDATE OF content, tool_name, tool_calls ON messages BEGIN DELETE FROM messages_fts WHERE rowid = old.id; INSERT INTO messages_fts(rowid, content) VALUES ( new.id, @@ -567,7 +590,8 @@ CREATE TRIGGER IF NOT EXISTS messages_fts_trigram_delete AFTER DELETE ON message DELETE FROM messages_fts_trigram WHERE rowid = old.id; END; -CREATE TRIGGER IF NOT EXISTS messages_fts_trigram_update AFTER UPDATE ON messages BEGIN +CREATE TRIGGER IF NOT EXISTS messages_fts_trigram_update +AFTER UPDATE OF content, tool_name, tool_calls ON messages BEGIN DELETE FROM messages_fts_trigram WHERE rowid = old.id; INSERT INTO messages_fts_trigram(rowid, content) VALUES ( new.id, diff --git a/hermes_state_portability.py b/hermes_state_portability.py index c25ecb85d6..decf8d3d8a 100644 --- a/hermes_state_portability.py +++ b/hermes_state_portability.py @@ -32,7 +32,7 @@ class SessionPortabilityMixin: @classmethod def _compact_session_cols(cls) -> str: """SELECT list for compact_rows: every ``sessions`` column declared in - SCHEMA_SQL except the ``system_prompt`` blob, aliased with the ``s`` + SCHEMA_SQL except prompt storage internals, aliased with the ``s`` prefix used by list_sessions_rich/_get_session_rich_row queries.""" if cls._session_compact_cols_sql is None: declared = cls._parse_schema_columns(SCHEMA_SQL)["sessions"] @@ -102,6 +102,7 @@ class SessionPortabilityMixin: query = f""" SELECT s.*, + COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved, COALESCE( (SELECT {_PREVIEW_RAW_SELECT} FROM messages m @@ -111,6 +112,7 @@ class SessionPortabilityMixin: ) AS _preview_raw, {_sql_session_last_active("s")} AS last_active FROM sessions s + LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash WHERE s.source = 'cron' AND s.id >= ? AND s.id < ? ORDER BY s.started_at DESC, s.id DESC LIMIT ? OFFSET ? @@ -121,7 +123,7 @@ class SessionPortabilityMixin: runs: List[Dict[str, Any]] = [] for row in rows: - s = dict(row) + s = self._session_row_dict(row) s["preview"] = _shape_preview(s.pop("_preview_raw", "")) runs.append(s) return runs @@ -133,12 +135,58 @@ class SessionPortabilityMixin: Pass ``compact_rows=True`` to omit the ``system_prompt`` blob (see ``list_sessions_rich`` for details). + + Thin wrapper over ``_get_session_rich_rows_batch`` so the enriched + SELECT lives in exactly one place. """ + return self._get_session_rich_rows_batch( + [session_id], compact_rows=compact_rows + ).get(session_id) + + def _get_session_rich_rows_batch( + self, session_ids, compact_rows: bool = False + ) -> Dict[str, Dict[str, Any]]: + """Fetch multiple sessions with the same enriched columns as + ``_get_session_rich_row``, in a single query. + + Used by ``list_sessions_rich``'s compression-tip projection to resolve + every tip row for a page in one round trip instead of one query per + compression-root row. Returns a dict keyed by session id; ids that + don't exist are simply absent from the result (same as + ``_get_session_rich_row`` returning ``None`` for them). + """ + ids = [sid for sid in session_ids if sid] + if not ids: + return {} + # Old SQLite builds cap bound variables at 999 + # (SQLITE_MAX_VARIABLE_NUMBER); large pages (limit=10000 callers + # exist) could exceed it. Chunk the IN list so the helper is safe at + # any page size — this is the single choke point for the enriched + # multi-row fetch, so the bound lives here, not at call sites. + _CHUNK = 900 + if len(ids) > _CHUNK: + result: Dict[str, Dict[str, Any]] = {} + for start in range(0, len(ids), _CHUNK): + result.update( + self._get_session_rich_rows_batch( + ids[start:start + _CHUNK], compact_rows=compact_rows + ) + ) + return result # Same read-your-writes guarantee as list_sessions_rich. self.flush_token_counts() _sel = self._compact_session_cols() if compact_rows else "s.*" + placeholders = ",".join("?" for _ in ids) + prompt_select = ( + "" if compact_rows + else ", COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved" + ) + prompt_join = ( + "" if compact_rows + else "LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash" + ) query = f""" - SELECT {_sel}, + SELECT {_sel}{prompt_select}, COALESCE( (SELECT {_PREVIEW_RAW_SELECT} FROM messages m @@ -148,16 +196,18 @@ class SessionPortabilityMixin: ) AS _preview_raw, {_sql_session_last_active("s")} AS last_active FROM sessions s - WHERE s.id = ? + {prompt_join} + WHERE s.id IN ({placeholders}) """ with self._lock: - cursor = self._conn.execute(query, (session_id,)) - row = cursor.fetchone() - if not row: - return None - s = dict(row) - s["preview"] = _shape_preview(s.pop("_preview_raw", "")) - return s + cursor = self._conn.execute(query, ids) + rows = cursor.fetchall() + result: Dict[str, Dict[str, Any]] = {} + for row in rows: + s = self._session_row_dict(row) + s["preview"] = _shape_preview(s.pop("_preview_raw", "")) + result[s["id"]] = s + return result def get_session_rich_row(self, session_id: str, compact_rows: bool = False) -> Optional[Dict[str, Any]]: """Public wrapper for :meth:`_get_session_rich_row`. @@ -518,10 +568,14 @@ class SessionPortabilityMixin: if started_at is None: started_at = time.time() archived = 1 if raw.get("archived") else 0 + system_prompt_hash = self._store_system_prompt( + conn, raw.get("system_prompt") + ) conn.execute( """INSERT INTO sessions ( id, source, user_id, model, model_config, system_prompt, + system_prompt_hash, parent_session_id, started_at, ended_at, end_reason, message_count, tool_call_count, input_tokens, output_tokens, cache_read_tokens, cache_write_tokens, reasoning_tokens, @@ -532,7 +586,7 @@ class SessionPortabilityMixin: ) VALUES ( :id, :source, :user_id, :model, :model_config, - :system_prompt, NULL, :started_at, :ended_at, + NULL, :system_prompt_hash, NULL, :started_at, :ended_at, :end_reason, 0, 0, :input_tokens, :output_tokens, :cache_read_tokens, :cache_write_tokens, :reasoning_tokens, :cwd, :git_branch, :git_repo_root, @@ -547,7 +601,7 @@ class SessionPortabilityMixin: "user_id": raw.get("user_id"), "model": raw.get("model"), "model_config": raw.get("model_config"), - "system_prompt": raw.get("system_prompt"), + "system_prompt_hash": system_prompt_hash, "started_at": started_at, "ended_at": self._float_or_none(raw.get("ended_at")), "end_reason": raw.get("end_reason"), diff --git a/hermes_state_schema.py b/hermes_state_schema.py index a93f037770..0761a0804e 100644 --- a/hermes_state_schema.py +++ b/hermes_state_schema.py @@ -16,6 +16,7 @@ from typing import Dict, Optional from hermes_constants import get_hermes_home from hermes_state_common import ( DEFERRED_INDEX_SQL, + FTS_CJK_STALE_KEY, FTS_SQL, FTS_STORAGE_VERSION, FTS_TRIGRAM_SQL, @@ -35,6 +36,27 @@ logger = logging.getLogger("hermes_state") class SessionSchemaMixin: """See module docstring — mixin for SessionDB (Schema cluster).""" + def _dedupe_legacy_system_prompts(self, cursor: sqlite3.Cursor) -> None: + """Move inline prompt snapshots into the shared content-addressed table.""" + try: + rows = cursor.execute( + "SELECT id, system_prompt FROM sessions " + "WHERE system_prompt IS NOT NULL" + ).fetchall() + except sqlite3.OperationalError: + return + + for row in rows: + session_id = row["id"] if isinstance(row, sqlite3.Row) else row[0] + prompt = row["system_prompt"] if isinstance(row, sqlite3.Row) else row[1] + prompt_hash = self._store_system_prompt(cursor, prompt) + cursor.execute( + "UPDATE sessions " + "SET system_prompt_hash = ?, system_prompt = NULL " + "WHERE id = ?", + (prompt_hash, session_id), + ) + def _sqlite_supports_fts5(self, cursor: sqlite3.Cursor) -> bool: try: cursor.execute("CREATE VIRTUAL TABLE temp._hermes_fts5_probe USING fts5(x)") @@ -56,6 +78,143 @@ class SessionSchemaMixin: ).fetchone() return int(row[0] if not isinstance(row, sqlite3.Row) else row[0]) + + @staticmethod + def _fts_update_trigger_needs_narrowing(sql: Optional[str]) -> bool: + """True when trigger SQL is missing AFTER UPDATE OF (still broad).""" + if not sql: + return False + # Collapse whitespace so multi-line DDL still matches. + compact = " ".join(sql.split()).upper() + # Already narrowed. + if "AFTER UPDATE OF " in compact: + return False + # Broad UPDATE trigger that we still need to replace. + return "AFTER UPDATE ON " in compact + + def _migrate_broad_fts_update_triggers(self, cursor: sqlite3.Cursor) -> int: + """Replace broad AFTER UPDATE FTS triggers with AFTER UPDATE OF variants. + + ``CREATE TRIGGER IF NOT EXISTS`` will not replace an existing broad + trigger, so installs that already created ``AFTER UPDATE ON messages`` + would keep firing on every messages row touch (status/compaction + writes included). Inspect ``sqlite_master``, drop any still-broad + UPDATE triggers, and re-apply the current DDL constants. + + No FTS rebuild: content correctness was already gated by WHEN clauses + on modern installs; OF only skips unnecessary trigger evaluation. + + Returns the number of triggers dropped (0 when already converged). + """ + # CJK is a v23-only surface. Decide the layout before selecting + # destructive candidates so the legacy branch never drops a trigger + # it does not recreate. + legacy_layout = self._db_has_legacy_inline_fts(cursor) + update_names = ( + "messages_fts_update", + "messages_fts_trigram_update", + ) + if not legacy_layout and hasattr(self, "_ensure_fts_cjk_schema"): + update_names += ("messages_fts_cjk_update",) + placeholders = ", ".join("?" for _ in update_names) + rows = cursor.execute( + "SELECT name, sql FROM sqlite_master " + f"WHERE type = 'trigger' AND name IN ({placeholders})", + update_names, + ).fetchall() + to_drop = [] + for row in rows: + name = row[0] if not isinstance(row, sqlite3.Row) else row["name"] + sql = row[1] if not isinstance(row, sqlite3.Row) else row["sql"] + if self._fts_update_trigger_needs_narrowing(sql): + to_drop.append(name) + if not to_drop: + return 0 + + for name in to_drop: + # Names are drawn from the update_names literal allowlist above — + # never user input — so the identifier is interpolation-safe. + cursor.execute(f"DROP TRIGGER IF EXISTS {name}") + + # Re-apply current DDL so CREATE TRIGGER installs the OF variants. + # Choose legacy vs v23 the same way _init_schema does. + if legacy_layout: + self._ensure_fts_schema(cursor, "messages_fts", LEGACY_FTS_SQL) + self._ensure_fts_schema( + cursor, "messages_fts_trigram", LEGACY_FTS_TRIGRAM_SQL + ) + else: + self._ensure_fts_schema(cursor, "messages_fts", FTS_SQL) + self._ensure_fts_schema( + cursor, "messages_fts_trigram", FTS_TRIGRAM_SQL + ) + # CJK triggers live on the host SessionDB; only recreate one that + # this migration actually dropped. ``_ensure_fts_cjk_schema`` is + # documented never-raises and soft-fails OperationalError by + # clearing availability — raise-path handling alone is not + # enough. After ensure, require a narrowed CJK UPDATE trigger or + # durable quarantine (stale breadcrumb + unavailable). + if "messages_fts_cjk_update" in to_drop: + try: + self._ensure_fts_cjk_schema(cursor) + except Exception: + self._quarantine_cjk_after_update_of_migration(cursor) + logger.exception( + "CJK FTS re-ensure after UPDATE OF migration failed" + ) + raise + if not self._cjk_update_trigger_is_narrowed(cursor): + self._quarantine_cjk_after_update_of_migration(cursor) + logger.warning( + "CJK FTS UPDATE trigger missing or still broad after " + "UPDATE OF migration; marked stale and unavailable" + ) + + logger.info( + "Migrated %d broad FTS UPDATE trigger(s) to AFTER UPDATE OF " + "(no rebuild required)", + len(to_drop), + ) + return len(to_drop) + + def _cjk_update_trigger_is_narrowed(self, cursor: sqlite3.Cursor) -> bool: + """True when messages_fts_cjk_update exists with AFTER UPDATE OF.""" + row = cursor.execute( + "SELECT sql FROM sqlite_master " + "WHERE type = 'trigger' AND name = ?", + ("messages_fts_cjk_update",), + ).fetchone() + if not row: + return False + sql = row[0] if not isinstance(row, sqlite3.Row) else row["sql"] + return not self._fts_update_trigger_needs_narrowing(sql) + + def _quarantine_cjk_after_update_of_migration( + self, cursor: sqlite3.Cursor + ) -> None: + """Fail-closed after dropping CJK UPDATE during OF migration. + + Clears availability, persists ``fts_cjk_stale``, and drops any + residual broad/partial CJK UPDATE trigger so a later open cannot + ``CREATE TRIGGER IF NOT EXISTS`` a gap without rebuild. + """ + self._fts_cjk_available = False + try: + self.set_meta(FTS_CJK_STALE_KEY, "1", cursor=cursor) + except Exception: + logger.debug( + "Could not persist CJK FTS stale breadcrumb", + exc_info=True, + ) + try: + cursor.execute("DROP TRIGGER IF EXISTS messages_fts_cjk_update") + except Exception: + logger.debug( + "Could not drop residual CJK UPDATE trigger after quarantine", + exc_info=True, + ) + + @staticmethod def _rebuild_fts_indexes( cursor: sqlite3.Cursor, @@ -724,6 +883,14 @@ class SessionSchemaMixin: if fts5_available and self._db_has_legacy_inline_fts(cursor): self.set_meta("fts_optimize_available", "1", cursor=cursor) + if current_version < 25: + # v25: de-duplicate per-session system prompt snapshots into + # a shared content-addressed table. Keep the old column as a + # read fallback for partially migrated or externally written + # rows, but clear migrated rows so future writes do not keep + # one large prompt copy per session. + self._dedupe_legacy_system_prompts(cursor) + # The FTS storage layout is versioned independently of the main # schema (see the v23 note above). Stamp the current layout so the # main version can always advance: a fresh/optimized DB is at @@ -855,6 +1022,11 @@ class SessionSchemaMixin: # the surfaces above and gated on the loadable tokenizer: self._ensure_fts_cjk_schema(cursor) + # Replace any pre-existing broad AFTER UPDATE triggers with + # AFTER UPDATE OF variants. IF NOT EXISTS cannot rewrite them. + if getattr(self, "_fts_enabled", False): + self._migrate_broad_fts_update_triggers(cursor) + self._conn.commit() def _backfill_gateway_metadata_from_sessions_json( diff --git a/hermes_state_search.py b/hermes_state_search.py index 4b9e51f7fe..756884b3f2 100644 --- a/hermes_state_search.py +++ b/hermes_state_search.py @@ -14,7 +14,7 @@ import os import re import sqlite3 import time -from typing import Any, Callable, Dict, List, Optional, Tuple +from typing import Any, Callable, Collection, Dict, List, Optional, Tuple from agent.skill_commands import describe_skill_invocation from hermes_state_common import ( @@ -35,6 +35,36 @@ logger = logging.getLogger("hermes_state") class SessionSearchMixin: """See module docstring — mixin for SessionDB (Search cluster).""" + _SEARCH_MESSAGE_RESULT_FIELDS = ( + "id", + "session_id", + "role", + "snippet", + "timestamp", + "tool_name", + "source", + "model", + "session_started", + "context", + ) + + @classmethod + def _search_message_fields( + cls, fields: Optional[Collection[str]] + ) -> Optional[Tuple[str, ...]]: + """Validate and canonically order an optional result projection.""" + if fields is None: + return None + if isinstance(fields, str): + raise TypeError("search fields must be a collection of field names, not a string") + requested = set(fields) + unknown = requested.difference(cls._SEARCH_MESSAGE_RESULT_FIELDS) + if unknown: + raise ValueError(f"unknown search result field(s): {', '.join(sorted(unknown))}") + return tuple( + field for field in cls._SEARCH_MESSAGE_RESULT_FIELDS if field in requested + ) + def _try_incremental_merge_fts(self) -> None: """Run one bounded FTS5 merge pass without failing the completed write.""" if not self._fts_enabled: @@ -85,7 +115,16 @@ class SessionSearchMixin: trigger activation): re-index any row near the boundary that the index is missing. docsize has one row per indexed doc, so the anti-join is exact and runs on a narrow id range. + + The trigram half of the sweep is gated on ``self._trigram_available`` + for the same reason ``fts_rebuild_step()`` gates its backfill INSERT: + when the SQLite build has no trigram tokenizer (or the table was + never created), an unconditional INSERT raises ``no such table`` + and aborts the whole rebuild — taking ``optimize_fts_storage()`` + down with it. """ + include_trigram = self._trigram_available + def _do(conn): hw_row = conn.execute( "SELECT value FROM state_meta WHERE key = 'fts_rebuild_high_water'" @@ -102,14 +141,15 @@ class SessionSearchMixin: "AND NOT EXISTS (SELECT 1 FROM messages_fts_docsize d WHERE d.id = m.id)", (lo, hi), ) - conn.execute( - "INSERT INTO messages_fts_trigram(rowid, content, tool_name, tool_calls) " - "SELECT m.id, m.content, m.tool_name, m.tool_calls " - "FROM messages m " - "WHERE m.id > ? AND m.id <= ? AND m.role <> 'tool' " - "AND NOT EXISTS (SELECT 1 FROM messages_fts_trigram_docsize d WHERE d.id = m.id)", - (lo, hi), - ) + if include_trigram: + conn.execute( + "INSERT INTO messages_fts_trigram(rowid, content, tool_name, tool_calls) " + "SELECT m.id, m.content, m.tool_name, m.tool_calls " + "FROM messages m " + "WHERE m.id > ? AND m.id <= ? AND m.role <> 'tool' " + "AND NOT EXISTS (SELECT 1 FROM messages_fts_trigram_docsize d WHERE d.id = m.id)", + (lo, hi), + ) conn.execute( "DELETE FROM state_meta WHERE key IN " "('fts_rebuild_high_water', 'fts_rebuild_progress')" @@ -1274,6 +1314,7 @@ class SessionSearchMixin: offset: int = 0, sort: str = None, include_inactive: bool = False, + fields: Optional[Collection[str]] = None, ) -> List[Dict[str, Any]]: """Instrumented wrapper around :meth:`_search_messages_impl`. @@ -1295,6 +1336,7 @@ class SessionSearchMixin: offset=offset, sort=sort, include_inactive=include_inactive, + fields=fields, ) return rows finally: @@ -1344,6 +1386,7 @@ class SessionSearchMixin: offset: int = 0, sort: str = None, include_inactive: bool = False, + fields: Optional[Collection[str]] = None, ) -> List[Dict[str, Any]]: """ Full-text search across session messages using FTS5. @@ -1356,6 +1399,9 @@ class SessionSearchMixin: Returns matching messages with session metadata, content snippet, and surrounding context (1 message before and after the match). + ``fields`` selects a result projection; omitting it preserves the + complete legacy result. Context is only loaded when that projection + consumes it. ``sort`` controls temporal ordering: - ``None`` (default): FTS5 BM25 relevance only. Time-neutral. @@ -1373,6 +1419,8 @@ class SessionSearchMixin: pre-compaction transcript stays discoverable after in-place compaction (#38763). Pass ``include_inactive=True`` to search every row regardless. """ + result_fields = self._search_message_fields(fields) + if not self._fts_enabled: return [] @@ -1458,6 +1506,7 @@ class SessionSearchMixin: # (indexed substring matching with ranking and snippets). For shorter # CJK queries (1-2 chars), trigram can't match (it needs ≥9 UTF-8 # bytes = 3 CJK chars), so we fall back to LIKE. + matches: List[Dict[str, Any]] = [] is_cjk = self._contains_cjk(query) if is_cjk: raw_query = query.strip('"').strip() @@ -1821,10 +1870,14 @@ class SessionSearchMixin: if tri_matches: matches = tri_matches - # Add surrounding context (1 message before + after each match). - # Each query takes its own fresh read transaction via _read_ctx, so - # we never hold a lock across N sequential queries. - for match in matches: + # Add surrounding context (1 message before + after each match) only + # when the selected result projection consumes it. Each query takes + # its own fresh read transaction via _read_ctx, so we never hold a + # lock across N sequential queries. + context_matches = ( + matches if result_fields is None or "context" in result_fields else () + ) + for match in context_matches: try: with self._read_ctx() as conn: ctx_cursor = conn.execute( @@ -1888,6 +1941,12 @@ class SessionSearchMixin: for match in matches: match.pop("content", None) + if result_fields is not None: + matches = [ + {field: match[field] for field in result_fields if field in match} + for match in matches + ] + return matches def _search_unindexed_gap( diff --git a/nix/devShell.nix b/nix/devShell.nix index 2e4007f854..ac50beb0e6 100644 --- a/nix/devShell.nix +++ b/nix/devShell.nix @@ -18,10 +18,7 @@ map (p: p.passthru.packageJsonPath or null) packages ); - # Non-npm packages may have their own devShellHook (e.g. hermes-agent - # stamps pyproject.toml + uv.lock for Python venv setup). - nonNpmHooks = map (p: p.passthru.devShellHook or "") packages; - combinedNonNpm = pkgs.lib.concatStringsSep "\n" (builtins.filter (h: h != "") nonNpmHooks); + hermesAgentDevShellHook = self'.packages.default.passthru.devShellHook; in { devShells.default = pkgs.mkShell { @@ -49,7 +46,7 @@ ] ++ self'.packages.default.passthru.devDeps; shellHook = '' - ${combinedNonNpm} + ${hermesAgentDevShellHook} ${hermesNpmLib.mkNpmDevShellHook npmPackageJsonPaths} # Force Node to use Nix's playwright-test binary instead of node_modules/.bin diff --git a/plugins/memory/honcho/__init__.py b/plugins/memory/honcho/__init__.py index c90734ad8a..78d24e90b3 100644 --- a/plugins/memory/honcho/__init__.py +++ b/plugins/memory/honcho/__init__.py @@ -23,7 +23,7 @@ import time from typing import Any, Callable, Dict, List, Optional from agent.memory_manager import sanitize_context -from agent.memory_provider import MemoryProvider +from agent.memory_provider import TRIVIAL_PROMPT_RE, MemoryProvider, is_trivial_prompt from tools.registry import tool_error logger = logging.getLogger(__name__) @@ -1196,28 +1196,17 @@ class HonchoMemoryProvider(MemoryProvider): return r return "" - # Prompts that carry no semantic signal — trivial acknowledgements, slash - # commands, empty input. Skipping injection here saves tokens and prevents - # stale user-model context from derailing one-word replies. - _TRIVIAL_PROMPT_RE = re.compile( - r'^(yes|no|ok|okay|sure|thanks|thank you|y|n|yep|nope|yeah|nah|' - r'continue|go ahead|do it|proceed|got it|cool|nice|great|done|next|lgtm|k)$', - re.IGNORECASE, - ) + # Prompts that carry no semantic signal — trivial acknowledgements, greetings, + # slash commands, empty input. Skipping injection here saves tokens and prevents + # stale user-model context from derailing one-word replies. Classification is + # fully delegated to the shared agent/memory_provider.is_trivial_prompt so the + # provider-side classifier and the core prefetch gate can never drift apart. + _TRIVIAL_PROMPT_RE = TRIVIAL_PROMPT_RE @classmethod def _is_trivial_prompt(cls, text: str) -> bool: """Return True if the prompt is too trivial to warrant context injection.""" - if not text: - return True - stripped = text.strip() - if not stripped: - return True - if stripped.startswith("/"): - return True - if cls._TRIVIAL_PROMPT_RE.match(stripped): - return True - return False + return is_trivial_prompt(text) def on_turn_start(self, turn_number: int, message: str, **kwargs) -> None: """Track turn count for cadence and injection_frequency logic.""" diff --git a/plugins/memory/openviking/README.md b/plugins/memory/openviking/README.md index af3ed9b8a1..4ec4a23f6d 100644 --- a/plugins/memory/openviking/README.md +++ b/plugins/memory/openviking/README.md @@ -9,6 +9,12 @@ Context database by Volcengine (ByteDance) with filesystem-style knowledge hiera then `openviking-server doctor`) - OpenViking server running and reachable from Hermes +OpenViking 0.2.10 or newer is recommended. For backward compatibility, +Hermes can identify older servers that expose the legacy status-only health +response, but only when anonymous OpenAPI metadata also identifies the service +as OpenViking. OpenViking 0.2.6 and earlier are deprecated for this integration; +upgrade them to receive the current health contract and compatibility fixes. + ## Setup Prepare OpenViking first: diff --git a/plugins/memory/openviking/__init__.py b/plugins/memory/openviking/__init__.py index e59f9467ac..06dd9237f5 100644 --- a/plugins/memory/openviking/__init__.py +++ b/plugins/memory/openviking/__init__.py @@ -29,10 +29,12 @@ import atexit import errno import json import logging +import math import mimetypes import os import re import shutil +import socket import stat import subprocess import tempfile @@ -41,6 +43,7 @@ import time import uuid import zipfile from dataclasses import dataclass, replace +from functools import lru_cache from pathlib import Path from typing import Any, Callable, Dict, List, Optional, Set from urllib.parse import quote, unquote, urlparse @@ -126,6 +129,13 @@ _GENERATED_MEMORY_SUMMARY_FILENAMES = { } _LOCAL_OPENVIKING_HOSTS = {"localhost", "127.0.0.1", "::1"} _LOCAL_OPENVIKING_AUTOSTART_TIMEOUT = 60.0 +# Pre-spawn liveness probe budget. A loopback TCP connect either completes or +# is refused in well under this; it exists only so a wedged listener cannot +# block the autostart path. +_LOCAL_OPENVIKING_PROBE_TIMEOUT = 2.0 +_LOCAL_SERVER_STARTED = "started" +_LOCAL_SERVER_OCCUPIED = "occupied" +_LOCAL_SERVER_FAILED = "failed" # After a refresh attempt fails for a given (unchanged) config, skip re-probing # for this long. Keeps "unavailable endpoints reconnect on a later access" # true while preventing every provider access from paying a 3s health probe @@ -133,11 +143,27 @@ _LOCAL_OPENVIKING_AUTOSTART_TIMEOUT = 60.0 _FAILED_CONFIG_RETRY_COOLDOWN_SECONDS = 30.0 _OPENVIKING_SERVER_LOG_RELATIVE_PATH = Path("logs") / "openviking-server.log" _OPENVIKING_RESPONDED_FAILURE_PREFIX = "OpenViking server responded" +_OPENVIKING_IDENTITY_MODERN = "modern" +_OPENVIKING_IDENTITY_LEGACY = "legacy" +_OPENVIKING_IDENTITY_UNHEALTHY = "unhealthy" +_OPENVIKING_IDENTITY_LEGACY_UNVERIFIED = "legacy-unverified" +_OPENVIKING_IDENTITY_INVALID = "invalid" +_OPENVIKING_IDENTIFIED_STATES = frozenset({ + _OPENVIKING_IDENTITY_MODERN, + _OPENVIKING_IDENTITY_LEGACY, +}) +_LEGACY_OPENVIKING_IDENTITY_DETAIL = ( + "returned OpenViking's legacy health response, but its anonymous " + "OpenAPI metadata did not identify OpenViking. If this is OpenViking 0.2.6 or " + "earlier, upgrade to OpenViking 0.2.10 or newer." +) _PENDING_SESSIONS_RELATIVE_DIR = Path("openviking") / "pending_sessions" _RUN_LOCKS_RELATIVE_DIR = Path("openviking") / "runs" _LEGACY_RECOVERY_LOCK_FILENAME = "legacy-recovery.lock" _LOCK_BUSY_ERRNOS = {errno.EWOULDBLOCK, errno.EACCES, errno.EAGAIN} _SETUP_CANCELLED = object() +_INVALID_SETTING_WARNINGS: Set[tuple[str, str]] = set() +_INVALID_SETTING_WARNINGS_LOCK = threading.Lock() @dataclass(frozen=True) @@ -156,6 +182,10 @@ class _OpenVikingHTTPError(RuntimeError): self.status_code = status_code +class _OpenVikingEndpointError(ValueError): + """Raised when a configured endpoint cannot be used safely.""" + + def _sanitize_openviking_error_message(message: str, status_code: Optional[int] = None) -> str: text = (message or "").strip() status = f"HTTP {status_code}" if status_code else "HTTP error" @@ -409,19 +439,24 @@ class _VikingClient: def health(self) -> bool: try: - resp = self._httpx.get( - self._url("/health"), headers=self._headers(), timeout=3.0 - ) - return resp.status_code == 200 + identity, _health = _probe_openviking_identity(self) + return identity in _OPENVIKING_IDENTIFIED_STATES except Exception: return False - def health_payload(self) -> dict: + def _anonymous_json(self, path: str) -> dict: + """Probe server identity without disclosing credentials or tenant IDs.""" resp = self._httpx.get( - self._url("/health"), headers=self._headers(), timeout=3.0 + self._url(path), headers={"Accept": "application/json"}, timeout=3.0 ) return self._parse_response(resp) + def health_payload(self) -> dict: + return self._anonymous_json("/health") + + def openapi_payload(self) -> dict: + return self._anonymous_json("/openapi.json") + def validate_auth(self) -> dict: """Validate authenticated OpenViking access without mutating state.""" return self.get("/api/v1/system/status") @@ -719,6 +754,27 @@ def _clean_config_value(value: Any) -> str: return value.strip() if isinstance(value, str) else "" +def _openviking_endpoint_label(value: Any) -> str: + """Return a credential-free endpoint label suitable for logs and UI.""" + raw = _clean_config_value(value) + if not raw: + return "" + try: + parsed = urlparse(raw if "://" in raw else f"//{raw}") + host = parsed.hostname + if not host: + return "" + display_host = f"[{host}]" if ":" in host and not host.startswith("[") else host + try: + port = parsed.port + except ValueError: + port = None + scheme = f"{parsed.scheme}://" if parsed.scheme else "" + return f"{scheme}{display_host}{f':{port}' if port is not None else ''}" + except Exception: + return "" + + def _default_ovcli_config_path() -> Path: return Path.home() / _OVCLI_DEFAULT_RELATIVE_PATH @@ -748,13 +804,16 @@ def _load_ovcli_config(path: Optional[Path] = None) -> dict: def _connection_values_from_ovcli(data: dict) -> dict: + endpoint_value = _clean_config_value(data.get("url")) api_key = _clean_config_value(data.get("api_key")) or _clean_config_value(data.get("root_api_key")) root_api_key = _clean_config_value(data.get("root_api_key")) send_identity = not api_key or api_key == root_api_key account = _clean_config_value(data.get("account") or data.get("account_id")) user = _clean_config_value(data.get("user") or data.get("user_id")) return { - "endpoint": _normalize_openviking_url(data.get("url")), + # A linked profile with no URL contributes no endpoint; the resolver + # can then continue to config.yaml and finally the built-in default. + "endpoint": _normalize_openviking_url(endpoint_value) if endpoint_value else "", "api_key": api_key, "root_api_key": root_api_key, "account": account if send_identity else "", @@ -788,37 +847,144 @@ def _validate_openviking_identity_value(value: str, *, field: str) -> tuple[bool return True, "", trimmed +@lru_cache(maxsize=128) +def _openviking_endpoint_is_always_blocked(candidate: str) -> bool: + """Check the safety floor once per configured endpoint value. + + Endpoint resolution is configuration work, but the live provider resolves + its settings on every access so Dashboard and ``/reload`` changes take + effect without a restart. Caching by the complete endpoint keeps that hot + path from repeating potentially slow DNS lookups; changing the configured + URL still produces a fresh validation. + """ + from tools.url_safety import is_always_blocked_url + + return is_always_blocked_url(candidate) + + def _normalize_openviking_url(url: str) -> str: trimmed = _clean_config_value(url).rstrip("/") if not trimmed: return _DEFAULT_ENDPOINT lower = trimmed.lower() - if lower in {"::1", "[::1]"}: - return "http://[::1]:1933" - if lower.startswith("[::1]:"): - return f"http://[::1]:{trimmed.rsplit(':', 1)[1]}" - if lower.startswith("::1:"): - return f"http://[::1]:{trimmed.rsplit(':', 1)[1]}" - if "://" in trimmed: - return trimmed - host, _sep, port = trimmed.partition(":") - if host.lower() in {"localhost", "127.0.0.1"}: - return f"http://{host}:{port or '1933'}" - return trimmed + if lower in {"localhost", "127.0.0.1"}: + candidate = f"http://{trimmed}:1933" + elif lower in {"::1", "[::1]"}: + candidate = "http://[::1]:1933" + elif lower.startswith("[::1]:") or lower.startswith("::1:"): + candidate = f"http://[::1]:{trimmed.rsplit(':', 1)[1]}" + elif "://" in trimmed: + candidate = trimmed + else: + candidate = f"http://{trimmed}" + + try: + parsed = urlparse(candidate) + if parsed.scheme.lower() not in {"http", "https"} or not parsed.hostname: + raise ValueError("OpenViking endpoints must use http:// or https:// with a host.") + # Force validation of malformed ports (``urlparse`` defers it). + parsed.port + if parsed.username or parsed.password or parsed.query or parsed.fragment: + raise ValueError( + "OpenViking endpoints cannot contain user info, query parameters, or fragments." + ) + except ValueError as exc: + raise _OpenVikingEndpointError( + f"Invalid OpenViking endpoint {_openviking_endpoint_label(candidate)}: {exc}" + ) from exc + + # Local / LAN self-host remains allowed; reject cloud-metadata and other + # always-blocked floors so a poisoned endpoint cannot SSRF via memory sync. + # Never silently replace an explicitly unsafe endpoint with localhost: that + # could attach Hermes to an unrelated deployment and forward credentials to + # a destination the user did not configure. + try: + check_url = candidate + if _openviking_endpoint_is_always_blocked(check_url): + raise _OpenVikingEndpointError( + "OpenViking endpoint " + f"{_openviking_endpoint_label(candidate)} targets a blocked metadata address." + ) + except _OpenVikingEndpointError: + raise + except Exception as exc: + logger.debug("OpenViking endpoint safety validation failed", exc_info=True) + raise _OpenVikingEndpointError( + "OpenViking endpoint safety validation failed; Hermes refused the connection." + ) from exc + + return candidate + + +def _is_openviking_health_payload(payload: Any) -> bool: + """Match OpenViking's documented ``GET /health`` response contract.""" + return ( + isinstance(payload, dict) + and payload.get("status") == "ok" + and payload.get("healthy") is True + and isinstance(payload.get("version"), str) + and bool(payload["version"].strip()) + ) + + +def _is_legacy_openviking_health_payload(payload: Any) -> bool: + """Match the status-only health contract published through OpenViking 0.2.6.""" + return ( + isinstance(payload, dict) + and payload.get("status") == "ok" + and "healthy" not in payload + and "version" not in payload + ) + + +def _is_openviking_openapi_payload(payload: Any) -> bool: + if not isinstance(payload, dict): + return False + info = payload.get("info") + return isinstance(info, dict) and info.get("title") == "OpenViking API" + + +def _probe_openviking_identity(client: _VikingClient) -> tuple[str, Any]: + """Identify modern or legacy OpenViking before any authenticated request.""" + health = client.health_payload() + if isinstance(health, dict) and health.get("healthy") is False: + return _OPENVIKING_IDENTITY_UNHEALTHY, health + if _is_openviking_health_payload(health): + return _OPENVIKING_IDENTITY_MODERN, health + if not _is_legacy_openviking_health_payload(health): + return _OPENVIKING_IDENTITY_INVALID, health + + try: + openapi = client.openapi_payload() + except Exception: + logger.debug("Legacy OpenViking OpenAPI identity probe failed", exc_info=True) + return _OPENVIKING_IDENTITY_LEGACY_UNVERIFIED, health + if _is_openviking_openapi_payload(openapi): + return _OPENVIKING_IDENTITY_LEGACY, health + return _OPENVIKING_IDENTITY_LEGACY_UNVERIFIED, health + + +def _legacy_openviking_identity_error(subject: str) -> str: + return f"{subject} {_LEGACY_OPENVIKING_IDENTITY_DETAIL}" def _load_profile(path: Path, *, source: str, name: str) -> Optional[_OvcliProfile]: try: data = _load_ovcli_config(path) + values = _connection_values_from_ovcli(data) except Exception as e: - logger.debug("Skipping invalid OpenViking CLI config %s: %s", path, e) + logger.warning( + "Skipping invalid OpenViking CLI config %s: %s", + path, + _format_openviking_exception(e), + ) return None return _OvcliProfile( source=source, name=name, path=path, data=data, - values=_connection_values_from_ovcli(data), + values=values, ) @@ -884,7 +1050,10 @@ def _discover_ovcli_profiles() -> list[_OvcliProfile]: def _is_local_openviking_url(value: str) -> bool: - candidate = _normalize_openviking_url(value) + try: + candidate = _normalize_openviking_url(value) + except _OpenVikingEndpointError: + return False if not candidate: return False if "://" not in candidate: @@ -930,12 +1099,31 @@ def _resolve_connection_settings(provider_config: Optional[dict] = None) -> dict user_env = _env_value("OPENVIKING_USER") agent_env = _env_value("OPENVIKING_AGENT") + # Non-secret fields fall back to config.yaml (e.g. the Dashboard writes + # ``memory.openviking.endpoint`` there) before the built-in default, so the + # full chain is env -> ovcli -> config.yaml -> default. The secret api_key is + # sourced from the environment (synced from .env), never from config.yaml. + endpoint = _first_nonempty( + endpoint_env, + ovcli_values.get("endpoint"), + _clean_config_value(provider_config.get("endpoint")), + default=_DEFAULT_ENDPOINT, + ) return { - "endpoint": _first_nonempty(endpoint_env, ovcli_values.get("endpoint"), default=_DEFAULT_ENDPOINT), + "endpoint": _normalize_openviking_url(endpoint), "api_key": api_key_env if api_key_env is not None else ovcli_values.get("api_key", ""), - "account": account_env if account_env is not None else ovcli_values.get("account", ""), - "user": user_env if user_env is not None else ovcli_values.get("user", ""), - "agent": _first_nonempty(agent_env, ovcli_values.get("agent"), default=_DEFAULT_AGENT), + "account": account_env if account_env is not None else _first_nonempty( + ovcli_values.get("account"), _clean_config_value(provider_config.get("account")) + ), + "user": user_env if user_env is not None else _first_nonempty( + ovcli_values.get("user"), _clean_config_value(provider_config.get("user")) + ), + "agent": _first_nonempty( + agent_env, + ovcli_values.get("agent"), + _clean_config_value(provider_config.get("agent")), + default=_DEFAULT_AGENT, + ), } @@ -1060,11 +1248,14 @@ def _validate_openviking_reachability(endpoint: str) -> tuple[bool, str]: try: client = _VikingClient(endpoint) if hasattr(client, "health_payload"): - payload = client.health_payload() - if payload.get("healthy") is False: + identity, _health = _probe_openviking_identity(client) + if identity == _OPENVIKING_IDENTITY_UNHEALTHY: return False, "OpenViking server responded but reported unhealthy status." - if payload: + if identity in _OPENVIKING_IDENTIFIED_STATES: return True, "" + if identity == _OPENVIKING_IDENTITY_LEGACY_UNVERIFIED: + return False, _legacy_openviking_identity_error("The server") + return False, "OpenViking server responded, but its /health response is not valid OpenViking." elif client.health(): return True, "" except Exception as e: @@ -1075,8 +1266,8 @@ def _validate_openviking_reachability(endpoint: str) -> tuple[bool, str]: def _validate_openviking_auth(values: dict) -> tuple[bool, str]: - endpoint = _normalize_openviking_url(values.get("endpoint")) try: + endpoint = _normalize_openviking_url(values.get("endpoint")) client = _VikingClient( endpoint, _clean_config_value(values.get("api_key")), @@ -1091,8 +1282,8 @@ def _validate_openviking_auth(values: dict) -> tuple[bool, str]: def _validate_openviking_root_access(values: dict) -> tuple[bool, str]: - endpoint = _normalize_openviking_url(values.get("endpoint")) try: + endpoint = _normalize_openviking_url(values.get("endpoint")) client = _VikingClient( endpoint, _clean_config_value(values.get("api_key")), @@ -1142,7 +1333,10 @@ def _validate_openviking_setup_values( *, require_api_key: bool = False, ) -> tuple[bool, str, Optional[str]]: - endpoint = _normalize_openviking_url(values.get("endpoint")) + try: + endpoint = _normalize_openviking_url(values.get("endpoint")) + except _OpenVikingEndpointError as exc: + return False, str(exc), None api_key = _clean_config_value(values.get("api_key")) if require_api_key and not api_key: return False, "Remote OpenViking configs require an API key.", None @@ -1155,9 +1349,13 @@ def _validate_openviking_setup_values( user=_clean_config_value(values.get("user")), agent=_clean_config_value(values.get("agent")) or _DEFAULT_AGENT, ) - health = client.health_payload() - if health.get("healthy") is False: + identity, health = _probe_openviking_identity(client) + if identity == _OPENVIKING_IDENTITY_UNHEALTHY: return False, "OpenViking server responded but reported unhealthy status.", None + if identity == _OPENVIKING_IDENTITY_LEGACY_UNVERIFIED: + return False, _legacy_openviking_identity_error("The server"), None + if identity not in _OPENVIKING_IDENTIFIED_STATES: + return False, "Server /health response is not valid OpenViking.", None if _should_probe_openviking_auth( health, require_api_key=require_api_key, @@ -1214,14 +1412,96 @@ def _openviking_server_log_path() -> Path: return home / _OPENVIKING_SERVER_LOG_RELATIVE_PATH -def _start_local_openviking_server(endpoint: str) -> tuple[bool, str]: - server_cmd = shutil.which("openviking-server") - if not server_cmd: - return False, "openviking-server was not found on PATH. Start it manually, then retry." +def _local_openviking_port_is_open(host: str, port: int) -> bool: + """Return True when something already accepts TCP connections on host:port. + + Used as a pre-spawn guard only. A successful connect proves a listener owns + the port, which is enough to know a second ``openviking-server`` would lose + the data-directory lock — it deliberately says nothing about whether that + listener is healthy. + """ + try: + with socket.create_connection((host, port), timeout=_LOCAL_OPENVIKING_PROBE_TIMEOUT): + return True + except OSError: + return False + + +def _describe_local_port_listener(host: str, port: int) -> str: + """Best-effort process identity for an occupied local TCP port.""" + try: + import psutil + + wildcard_hosts = {"0.0.0.0", "::", "::0"} + aliases = {host.lower()} + if host.lower() == "localhost": + aliases.update({"127.0.0.1", "::1"}) + for conn in psutil.net_connections(kind="inet"): + if conn.status != psutil.CONN_LISTEN or not conn.laddr: + continue + listener_host = str( + conn.laddr.ip if hasattr(conn.laddr, "ip") else conn.laddr[0] + ).lower() + listener_port = int( + conn.laddr.port if hasattr(conn.laddr, "port") else conn.laddr[1] + ) + if listener_port != port: + continue + if listener_host not in wildcard_hosts and listener_host not in aliases: + continue + if conn.pid is None: + break + try: + process_name = psutil.Process(conn.pid).name() + except (psutil.Error, OSError): + process_name = "unknown process" + process_name = re.sub(r"[^\w .+-]", "?", str(process_name))[:80] + return f"{process_name or 'unknown process'} (PID {conn.pid})" + except Exception: + logger.debug( + "Could not identify the process listening on %s:%s", + host, + port, + exc_info=True, + ) + return "an unidentified process" + + +def _local_listener_suffix(endpoint: str) -> str: + if not _is_local_openviking_url(endpoint): + return "" + try: + host, port = _local_openviking_bind(endpoint) + except ValueError: + return "" + if not _local_openviking_port_is_open(host, port): + return "" + return f" The listener on {host}:{port} is {_describe_local_port_listener(host, port)}." + + +def _start_local_openviking_server(endpoint: str) -> tuple[str, str]: try: host, port = _local_openviking_bind(endpoint) except ValueError as e: - return False, f"Could not parse local OpenViking URL: {e}" + return _LOCAL_SERVER_FAILED, f"Could not parse local OpenViking URL: {e}" + # Health probes can time out client-side while the server is up and well. + # Spawning on that signal alone produces a process that immediately dies on + # DataDirectoryLocked, and — because the probe keeps timing out — repeats + # every cooldown window. Treat an occupied port only as a spawn-prevention + # signal, never as proof that the listener is OpenViking. + if _local_openviking_port_is_open(host, port): + listener = _describe_local_port_listener(host, port) + return ( + _LOCAL_SERVER_OCCUPIED, + f"Port {host}:{port} is occupied by {listener}. Hermes did not start " + "openviking-server because the listener has not passed OpenViking's /health check.", + ) + server_cmd = shutil.which("openviking-server") + if not server_cmd: + return ( + _LOCAL_SERVER_FAILED, + "openviking-server was not found on PATH. Start it manually, then retry.", + ) log_path = _openviking_server_log_path() try: log_path.parent.mkdir(parents=True, exist_ok=True) @@ -1234,8 +1514,11 @@ def _start_local_openviking_server(endpoint: str) -> tuple[bool, str]: start_new_session=True, ) except Exception as e: - return False, f"Could not start openviking-server: {e}" - return True, f"Started openviking-server on {host}:{port} in the background. Logs: {log_path}" + return _LOCAL_SERVER_FAILED, f"Could not start openviking-server: {e}" + return ( + _LOCAL_SERVER_STARTED, + f"Started openviking-server on {host}:{port} in the background. Logs: {log_path}", + ) def _wait_for_openviking_health( @@ -1284,9 +1567,9 @@ def _handle_unreachable_endpoint( cancel_returns=cancelled, ) if choice == 0: - started, start_message = _start_local_openviking_server(endpoint) + start_state, start_message = _start_local_openviking_server(endpoint) print(f" {start_message}") - if not started: + if start_state != _LOCAL_SERVER_STARTED: return False print(" Waiting for OpenViking server to become reachable...", flush=True) if _wait_for_openviking_health( @@ -1332,7 +1615,8 @@ def _runtime_openviking_timeout_message(endpoint: str) -> str: f"Local OpenViking server at {endpoint} is not reachable. " "Tried to start openviking-server, but it did not become reachable " f"within {_LOCAL_OPENVIKING_AUTOSTART_TIMEOUT:.0f} seconds. " - "OpenViking memory disabled for this Hermes run." + "OpenViking memory is temporarily unavailable; Hermes will retry on a later access or when " + "the config changes." ) @@ -1340,19 +1624,33 @@ def _classify_runtime_openviking_health(client: _VikingClient, endpoint: str) -> """Classify runtime health without treating every false result as server absence.""" try: if hasattr(client, "health_payload"): - payload = client.health_payload() - if payload.get("healthy") is False: + identity, _health = _probe_openviking_identity(client) + if identity == _OPENVIKING_IDENTITY_UNHEALTHY: return ( "responded", - f"OpenViking server at {endpoint} responded but reported unhealthy status.", + f"Service at {endpoint} responded but reported unhealthy OpenViking status." + f"{_local_listener_suffix(endpoint)}", ) - return "healthy", "" + if identity in _OPENVIKING_IDENTIFIED_STATES: + return "healthy", "" + if identity == _OPENVIKING_IDENTITY_LEGACY_UNVERIFIED: + return ( + "responded", + _legacy_openviking_identity_error(f"Service at {endpoint}") + + _local_listener_suffix(endpoint), + ) + return ( + "responded", + f"Service at {endpoint} responded, but its /health response is not valid OpenViking." + f"{_local_listener_suffix(endpoint)}", + ) if client.health(): return "healthy", "" except _OpenVikingHTTPError as e: return ( "responded", - f"OpenViking server at {endpoint} responded with {_format_openviking_exception(e)}.", + f"Service at {endpoint} responded with {_format_openviking_exception(e)}." + f"{_local_listener_suffix(endpoint)}", ) except Exception: return "unreachable", "" @@ -1406,7 +1704,20 @@ def _prompt_manual_connection_values(prompt, select, cancelled, *, service: bool print(f" OpenViking Service endpoint: {endpoint}") else: while True: - endpoint = _normalize_openviking_url(prompt("OpenViking server URL", default=_DEFAULT_ENDPOINT)) + try: + endpoint = _normalize_openviking_url( + prompt("OpenViking server URL", default=_DEFAULT_ENDPOINT) + ) + except _OpenVikingEndpointError as exc: + retry = _retry_or_cancel_manual_setup( + select, + " Invalid OpenViking endpoint", + str(exc), + cancelled, + ) + if retry is _SETUP_CANCELLED: + return _SETUP_CANCELLED + continue _print_validation_progress("Checking OpenViking server...") reachable, message = _validate_openviking_reachability(endpoint) if reachable: @@ -1441,16 +1752,16 @@ def _prompt_manual_connection_values(prompt, select, cancelled, *, service: bool credential_choice = select( " OpenViking credential", [ - ("No API key", "local dev mode"), - ("User API key", "server derives account/user automatically"), + ("User API key", "recommended; server derives account/user automatically"), ("Root API key", "requires account and user IDs"), + ("No API key", "only for explicitly unauthenticated local development"), ], default=0, cancel_returns=cancelled, ) if credential_choice == cancelled: return _SETUP_CANCELLED - if credential_choice == 0: + if credential_choice == 2: values["agent"] = _clean_config_value( prompt(_AGENT_PROMPT_LABEL, default=_DEFAULT_AGENT) ) or _DEFAULT_AGENT @@ -1468,7 +1779,7 @@ def _prompt_manual_connection_values(prompt, select, cancelled, *, service: bool if retry is _SETUP_CANCELLED: return _SETUP_CANCELLED continue - api_key_type = "root" if credential_choice == 2 else "user" + api_key_type = "root" if credential_choice == 1 else "user" elif not api_key_type: credential_choice = select( " OpenViking API key type", @@ -1929,6 +2240,10 @@ class OpenVikingMemoryProvider(MemoryProvider): if os.environ.get("OPENVIKING_ENDPOINT"): return True provider_config = _load_hermes_openviking_config() + # A non-secret endpoint saved to config.yaml (e.g. via the Dashboard) + # counts as configured even without an env var or ovcli config. + if _clean_config_value(provider_config.get("endpoint")): + return True if not provider_config.get("use_ovcli_config"): return False try: @@ -1948,18 +2263,21 @@ class OpenVikingMemoryProvider(MemoryProvider): }, { "key": "api_key", - "description": "OpenViking API key (leave blank for local dev mode)", + "description": ( + "OpenViking API key (recommended; only leave blank for an explicitly " + "unauthenticated local development server)" + ), "secret": True, "env_var": "OPENVIKING_API_KEY", }, { "key": "account", - "description": "OpenViking tenant account ID (blank for user API keys)", + "description": "Advanced local identity override (leave blank for user API keys)", "env_var": "OPENVIKING_ACCOUNT", }, { "key": "user", - "description": "OpenViking user ID within the account (blank for user API keys)", + "description": "Advanced local user override (leave blank for user API keys)", "env_var": "OPENVIKING_USER", }, { @@ -1974,59 +2292,108 @@ class OpenVikingMemoryProvider(MemoryProvider): { "key": "recall_limit", "description": "Maximum memories injected by automatic recall", + "type": "integer", + "minimum": 1, + "maximum": 100, "default": _DEFAULT_RECALL_LIMIT, "env_var": "OPENVIKING_RECALL_LIMIT", }, { "key": "recall_score_threshold", "description": "Minimum relevance score for automatic recall", + "type": "number", + "minimum": 0.0, + "maximum": 1.0, + "step": 0.01, "default": _DEFAULT_RECALL_SCORE_THRESHOLD, "env_var": "OPENVIKING_RECALL_SCORE_THRESHOLD", }, { "key": "recall_max_injected_chars", "description": "Maximum total characters injected by recall", + "type": "integer", + "minimum": 100, + "maximum": 50000, "default": _DEFAULT_RECALL_MAX_INJECTED_CHARS, "env_var": "OPENVIKING_RECALL_MAX_INJECTED_CHARS", }, { "key": "profile_token_budget", "description": "Maximum session-start memory tokens injected", + "type": "integer", + "minimum": 500, + "maximum": 50000, "default": _DEFAULT_PROFILE_TOKEN_BUDGET, "env_var": "OPENVIKING_PROFILE_TOKEN_BUDGET", }, { "key": "recall_timeout_seconds", "description": "Total timeout for recall (seconds)", + "type": "number", + "minimum": 0.25, + "maximum": 60.0, + "step": 0.25, "default": _DEFAULT_RECALL_TIMEOUT_SECONDS, "env_var": "OPENVIKING_RECALL_TIMEOUT_SECONDS", }, { "key": "recall_request_timeout_seconds", "description": "Per-request timeout for recall (seconds)", + "type": "number", + "minimum": 0.25, + "maximum": 60.0, + "step": 0.25, "default": _DEFAULT_RECALL_REQUEST_TIMEOUT_SECONDS, "env_var": "OPENVIKING_RECALL_REQUEST_TIMEOUT_SECONDS", }, { "key": "recall_full_read_limit", "description": "Max full L2 content reads per recall", + "type": "integer", + "minimum": 0, + "maximum": 100, "default": _DEFAULT_RECALL_FULL_READ_LIMIT, "env_var": "OPENVIKING_RECALL_FULL_READ_LIMIT", }, { "key": "recall_prefer_abstract", "description": "Use abstracts instead of full L2 reads", + "type": "boolean", "default": False, "env_var": "OPENVIKING_RECALL_PREFER_ABSTRACT", }, { "key": "recall_resources", "description": "Include resources in recall", + "type": "boolean", "default": False, "env_var": "OPENVIKING_RECALL_RESOURCES", }, ] + def save_config(self, values: Dict[str, Any], hermes_home: str) -> None: + """Validate and persist Dashboard configuration for the active profile.""" + normalized = dict(values or {}) + normalized.pop("api_key", None) + normalized.pop("root_api_key", None) + endpoint = _clean_config_value(normalized.get("endpoint")) + if endpoint: + normalized["endpoint"] = _normalize_openviking_url(endpoint) + + from hermes_cli.config import load_config, save_config + + config = load_config() + memory_config = config.get("memory") + if not isinstance(memory_config, dict): + memory_config = {} + config["memory"] = memory_config + provider_config = memory_config.get("openviking") + if not isinstance(provider_config, dict): + provider_config = {} + provider_config.update(normalized) + memory_config["openviking"] = provider_config + save_config(config) + def get_status_config(self, provider_config: dict) -> dict: provider_config = dict(provider_config or {}) if provider_config.get("use_ovcli_config"): @@ -2187,8 +2554,9 @@ class OpenVikingMemoryProvider(MemoryProvider): return if not healthy: warning_message = ( - f"OpenViking server at {endpoint} is still not reachable after auto-start; " - "OpenViking memory disabled for this Hermes run." + f"OpenViking server at {endpoint} is still not reachable after auto-start. " + "OpenViking memory is temporarily unavailable; Hermes will retry on a later access or when " + "the config changes." ) else: self._client = client @@ -2206,7 +2574,8 @@ class OpenVikingMemoryProvider(MemoryProvider): except Exception as e: warning_message = ( f"OpenViking server at {endpoint} could not be attached after auto-start: {e}. " - "OpenViking memory disabled for this Hermes run." + "OpenViking memory is temporarily unavailable; Hermes will retry on a later access or when " + "the config changes." ) if warning_message: @@ -2230,8 +2599,9 @@ class OpenVikingMemoryProvider(MemoryProvider): endpoint = self._endpoint if not _is_local_openviking_url(endpoint): _emit_runtime_warning( - f"Remote OpenViking server at {endpoint} is not reachable; " - "OpenViking memory disabled for this Hermes run. " + f"Remote OpenViking server at {endpoint} is not reachable. " + "OpenViking memory is temporarily unavailable; Hermes will retry on a later access or when " + "the config changes. " "Check the configured endpoint and network connectivity.", warning_callback, ) @@ -2251,12 +2621,13 @@ class OpenVikingMemoryProvider(MemoryProvider): return self._runtime_start_pending = True - started, start_message = _start_local_openviking_server(endpoint) - if not started: + start_state, start_message = _start_local_openviking_server(endpoint) + if start_state != _LOCAL_SERVER_STARTED: self._runtime_start_pending = False warning_message = ( f"Local OpenViking server at {endpoint} is not reachable. {start_message} " - "OpenViking memory disabled for this Hermes run." + "OpenViking memory is temporarily unavailable; Hermes will retry on a later access or when " + "the config changes." ) self._client = None else: @@ -2286,7 +2657,28 @@ class OpenVikingMemoryProvider(MemoryProvider): ) def initialize(self, session_id: str, **kwargs) -> None: - settings = _resolve_connection_settings(_load_hermes_openviking_config()) + warning_callback = ( + kwargs.get("warning_callback") + if kwargs.get("platform") == "cli" + else None + ) + status_callback = ( + kwargs.get("status_callback") + if kwargs.get("platform") == "cli" + else None + ) + connection_error = "" + try: + settings = _resolve_connection_settings(_load_hermes_openviking_config()) + except _OpenVikingEndpointError as exc: + connection_error = str(exc) + settings = { + "endpoint": "", + "api_key": "", + "account": "", + "user": "", + "agent": _DEFAULT_AGENT, + } self._endpoint = settings["endpoint"] self._api_key = settings["api_key"] self._account = settings["account"] @@ -2309,37 +2701,43 @@ class OpenVikingMemoryProvider(MemoryProvider): self._hermes_home = hermes_home self._acquire_run_lock() self._profile_prefetched_sessions.clear() - warning_callback = ( - kwargs.get("warning_callback") - if kwargs.get("platform") == "cli" - else None - ) - status_callback = ( - kwargs.get("status_callback") - if kwargs.get("platform") == "cli" - else None - ) - try: - self._client = _VikingClient( - self._endpoint, self._api_key, - account=self._account, user=self._user, agent=self._agent, + if connection_error: + self._failed_refresh = ( + ("invalid-endpoint", connection_error), + time.monotonic(), + ) + _emit_runtime_warning( + f"{connection_error} OpenViking memory is temporarily unavailable; " + "correct the endpoint and reload the configuration.", + warning_callback, ) - health_state, health_message = _classify_runtime_openviking_health(self._client, self._endpoint) - if health_state == "unreachable": - self._handle_runtime_openviking_unreachable( - status_callback=status_callback, - warning_callback=warning_callback, - ) - elif health_state != "healthy": - _emit_runtime_warning( - f"{health_message} OpenViking memory disabled for this Hermes run.", - warning_callback, - ) - self._client = None - except ImportError: - logger.warning("httpx not installed — OpenViking plugin disabled") self._client = None + else: + try: + self._client = _VikingClient( + self._endpoint, self._api_key, + account=self._account, user=self._user, agent=self._agent, + ) + health_state, health_message = _classify_runtime_openviking_health( + self._client, + self._endpoint, + ) + if health_state == "unreachable": + self._handle_runtime_openviking_unreachable( + status_callback=status_callback, + warning_callback=warning_callback, + ) + elif health_state != "healthy": + _emit_runtime_warning( + f"{health_message} OpenViking memory is temporarily unavailable; " + "Hermes will retry on a later access or when the config changes.", + warning_callback, + ) + self._client = None + except ImportError: + logger.warning("httpx not installed — OpenViking plugin disabled") + self._client = None if self._client: self._conn_snapshot = ( @@ -2378,7 +2776,25 @@ class OpenVikingMemoryProvider(MemoryProvider): self._client = None return None - settings = _resolve_connection_settings(_load_hermes_openviking_config()) + try: + settings = _resolve_connection_settings(_load_hermes_openviking_config()) + except _OpenVikingEndpointError as exc: + failed_key = ("invalid-endpoint", str(exc)) + failed = self._failed_refresh + should_warn = not ( + failed is not None + and failed[0] == failed_key + and time.monotonic() - failed[1] < _FAILED_CONFIG_RETRY_COOLDOWN_SECONDS + ) + self._failed_refresh = (failed_key, time.monotonic()) + self._client = None + if should_warn: + logger.warning( + "%s OpenViking memory is temporarily unavailable; correct the endpoint " + "and reload the configuration.", + exc, + ) + return None endpoint = settings["endpoint"] api_key = settings["api_key"] account = settings["account"] @@ -2439,8 +2855,8 @@ class OpenVikingMemoryProvider(MemoryProvider): self._failed_refresh = (settings_key, time.monotonic()) if health_state == "responded": logger.warning( - "%s OpenViking memory disabled; will retry on a later access " - "(after cooldown) or when the config changes.", + "%s OpenViking memory is temporarily unavailable; Hermes will retry on a " + "later access (after cooldown) or when the config changes.", health_message, ) else: # unreachable @@ -2716,6 +3132,18 @@ class OpenVikingMemoryProvider(MemoryProvider): with self._committed_session_lock: self._committed_session_ids.add(sid) + def _clear_session_committed(self, sid: str) -> None: + """Re-arm the commit guard for a session that is still live. + + A permanent per-sid latch is correct for a session being left behind: + it dedupes that id's ``_finalize_session_async`` against the commit + compression already performed. In-place compression keeps the *same* + id, so the latch would otherwise reject every later commit for a + session that is still accumulating turns (#74695). + """ + with self._committed_session_lock: + self._committed_session_ids.discard(sid) + def _pending_session_dir(self) -> Optional[Path]: if not self._hermes_home: return None @@ -3145,76 +3573,148 @@ class OpenVikingMemoryProvider(MemoryProvider): return "" @staticmethod - def _env_bool(name: str, default: bool = False) -> bool: - raw = os.environ.get(name) - if raw is None or raw == "": - return default - return raw.strip().lower() in {"1", "true", "yes", "on"} + def _warn_invalid_setting_once(source: str, value: Any, default: Any) -> None: + warning_key = (source, repr(value)) + with _INVALID_SETTING_WARNINGS_LOCK: + if warning_key in _INVALID_SETTING_WARNINGS: + return + _INVALID_SETTING_WARNINGS.add(warning_key) + logger.warning("Invalid %s value %r; using default %r.", source, value, default) @staticmethod - def _env_int(name: str, default: int, *, minimum: int, maximum: int) -> int: - raw = os.environ.get(name) - try: - value = int(float(raw)) if raw not in {None, ""} else default - except (TypeError, ValueError): - value = default - return max(minimum, min(maximum, value)) + def _setting_value(env_name: str, config_value: Any) -> tuple[Any, str]: + env_value = os.environ.get(env_name) + if env_value is not None and env_value.strip(): + return env_value, env_name + config_key = env_name.removeprefix("OPENVIKING_").lower() + return config_value, f"memory.openviking.{config_key}" - @staticmethod - def _env_float(name: str, default: float, *, minimum: float, maximum: float) -> float: - raw = os.environ.get(name) + @classmethod + def _setting_bool( + cls, + env_name: str, + config_value: Any, + *, + default: bool, + ) -> bool: + value, source = cls._setting_value(env_name, config_value) + if isinstance(value, bool): + return value + if isinstance(value, str): + normalized = value.strip().lower() + if normalized in {"1", "true", "yes", "on"}: + return True + if normalized in {"0", "false", "no", "off"}: + return False + cls._warn_invalid_setting_once(source, value, default) + return default + + @classmethod + def _setting_int( + cls, + env_name: str, + config_value: Any, + *, + default: int, + minimum: int, + maximum: int, + ) -> int: + value, source = cls._setting_value(env_name, config_value) try: - value = float(raw) if raw not in {None, ""} else default - except (TypeError, ValueError): - value = default - return max(minimum, min(maximum, value)) + if isinstance(value, bool): + raise ValueError + numeric = float(value) + if not numeric.is_integer(): + raise ValueError + parsed = int(numeric) + except (TypeError, ValueError, OverflowError): + cls._warn_invalid_setting_once(source, value, default) + parsed = default + return max(minimum, min(maximum, parsed)) + + @classmethod + def _setting_float( + cls, + env_name: str, + config_value: Any, + *, + default: float, + minimum: float, + maximum: float, + ) -> float: + value, source = cls._setting_value(env_name, config_value) + try: + if isinstance(value, bool): + raise ValueError + parsed = float(value) + if not math.isfinite(parsed): + raise ValueError + except (TypeError, ValueError, OverflowError): + cls._warn_invalid_setting_once(source, value, default) + parsed = default + return max(minimum, min(maximum, parsed)) def _recall_config(self) -> Dict[str, Any]: + # Read from config.yaml → memory.openviking as primary source, env vars + # as override. Behavioural settings belong in config.yaml (AGENTS.md). + provider_config = _load_hermes_openviking_config() + cfg = provider_config + return { - "limit": self._env_int( + "limit": self._setting_int( "OPENVIKING_RECALL_LIMIT", - _DEFAULT_RECALL_LIMIT, - minimum=1, - maximum=100, + cfg.get("recall_limit", _DEFAULT_RECALL_LIMIT), + default=_DEFAULT_RECALL_LIMIT, + minimum=1, maximum=100, ), - "score_threshold": self._env_float( + "score_threshold": self._setting_float( "OPENVIKING_RECALL_SCORE_THRESHOLD", - _DEFAULT_RECALL_SCORE_THRESHOLD, - minimum=0.0, - maximum=1.0, + cfg.get("recall_score_threshold", _DEFAULT_RECALL_SCORE_THRESHOLD), + default=_DEFAULT_RECALL_SCORE_THRESHOLD, + minimum=0.0, maximum=1.0, ), - "max_injected_chars": self._env_int( + "max_injected_chars": self._setting_int( "OPENVIKING_RECALL_MAX_INJECTED_CHARS", - _DEFAULT_RECALL_MAX_INJECTED_CHARS, - minimum=100, - maximum=50000, + cfg.get("recall_max_injected_chars", _DEFAULT_RECALL_MAX_INJECTED_CHARS), + default=_DEFAULT_RECALL_MAX_INJECTED_CHARS, + minimum=100, maximum=50000, ), - "timeout_seconds": self._env_float( + "timeout_seconds": self._setting_float( "OPENVIKING_RECALL_TIMEOUT_SECONDS", - _DEFAULT_RECALL_TIMEOUT_SECONDS, - minimum=0.25, - maximum=60.0, + cfg.get("recall_timeout_seconds", _DEFAULT_RECALL_TIMEOUT_SECONDS), + default=_DEFAULT_RECALL_TIMEOUT_SECONDS, + minimum=0.25, maximum=60.0, ), - "request_timeout_seconds": self._env_float( + "request_timeout_seconds": self._setting_float( "OPENVIKING_RECALL_REQUEST_TIMEOUT_SECONDS", - _DEFAULT_RECALL_REQUEST_TIMEOUT_SECONDS, - minimum=0.25, - maximum=60.0, + cfg.get("recall_request_timeout_seconds", _DEFAULT_RECALL_REQUEST_TIMEOUT_SECONDS), + default=_DEFAULT_RECALL_REQUEST_TIMEOUT_SECONDS, + minimum=0.25, maximum=60.0, ), - "full_read_limit": self._env_int( + "full_read_limit": self._setting_int( "OPENVIKING_RECALL_FULL_READ_LIMIT", - _DEFAULT_RECALL_FULL_READ_LIMIT, - minimum=0, - maximum=100, + cfg.get("recall_full_read_limit", _DEFAULT_RECALL_FULL_READ_LIMIT), + default=_DEFAULT_RECALL_FULL_READ_LIMIT, + minimum=0, maximum=100, + ), + "prefer_abstract": self._setting_bool( + "OPENVIKING_RECALL_PREFER_ABSTRACT", + cfg.get("recall_prefer_abstract", False), + default=False, + ), + "resources": self._setting_bool( + "OPENVIKING_RECALL_RESOURCES", + cfg.get("recall_resources", False), + default=False, ), - "prefer_abstract": self._env_bool("OPENVIKING_RECALL_PREFER_ABSTRACT", False), - "resources": self._env_bool("OPENVIKING_RECALL_RESOURCES", False), } def _profile_token_budget(self) -> int: - return self._env_int( + cfg = _load_hermes_openviking_config() + return self._setting_int( "OPENVIKING_PROFILE_TOKEN_BUDGET", - _DEFAULT_PROFILE_TOKEN_BUDGET, + cfg.get("profile_token_budget", _DEFAULT_PROFILE_TOKEN_BUDGET), + default=_DEFAULT_PROFILE_TOKEN_BUDGET, minimum=500, maximum=50000, ) @@ -4173,6 +4673,12 @@ class OpenVikingMemoryProvider(MemoryProvider): if rotate: self._session_id = new_id self._turn_count = 0 + elif compression: + # commit_memory_session() has already extracted every turn up + # to this boundary. Keep the same sid, but start the live + # session's turn accounting again at zero so an immediate + # session end cannot duplicate the just-finished extraction. + self._turn_count = 0 if compression: # Discard both old and new session IDs so the profile is re-injected @@ -4182,9 +4688,23 @@ class OpenVikingMemoryProvider(MemoryProvider): self._profile_prefetched_sessions.discard(old_session_id) self._profile_prefetched_sessions.discard(new_id) + if not rotate and old_session_id: + # In-place compression (the default) keeps the same session id. + # compress_context() has just committed it, latching the guard — + # but the session is still live, so every later commit for it + # (the next compression, /new, normal session end, startup + # recovery) would be rejected and post-compression turns would + # never be extracted. Re-arm the guard now that compression has + # finished; turns arriving after this point are genuinely new. + # + # Rotation mode is untouched: there a fresh child id is minted + # and the old id stays latched, which is what dedupes its + # _finalize_session_async against this same commit. + self._clear_session_committed(old_session_id) + if not rotate: - # Same-session rewind (/undo) or no-op rotation: no commit and no - # counter reset. + # Same-session rewind (/undo) or no-op rotation: no new commit. + # Compression already reset the extracted-turn count above. logger.debug( "OpenViking on_session_switch skipped rotation: session=%s rewound=%s", old_session_id, rewound, diff --git a/plugins/memory/retaindb/__init__.py b/plugins/memory/retaindb/__init__.py index 13ad033150..f65d5018be 100644 --- a/plugins/memory/retaindb/__init__.py +++ b/plugins/memory/retaindb/__init__.py @@ -44,6 +44,30 @@ _DEFAULT_BASE_URL = "https://api.retaindb.com" _ASYNC_SHUTDOWN = object() +def _load_retaindb_config() -> Dict[str, Any]: + """Return the ``memory.retaindb`` block from config.yaml (empty on any error). + + Non-secret fields (``base_url``, ``project``) are persisted here by the + Dashboard; the runtime must read them back when the matching env var is + unset. The secret ``api_key`` continues to come through profile-scoped + secret resolution rather than config.yaml. + """ + try: + from hermes_cli.config import load_config_readonly + + config = load_config_readonly() + memory_config = config.get("memory", {}) if isinstance(config, dict) else {} + provider_config = memory_config.get("retaindb", {}) if isinstance(memory_config, dict) else {} + return dict(provider_config) if isinstance(provider_config, dict) else {} + except Exception: + return {} + + +def _config_str(value: Any) -> str: + """Return a stripped string for a config value, else ``""``.""" + return value.strip() if isinstance(value, str) else "" + + # --------------------------------------------------------------------------- # Tool schemas # --------------------------------------------------------------------------- @@ -489,12 +513,20 @@ class RetainDBMemoryProvider(MemoryProvider): # ── Lifecycle ────────────────────────────────────────────────────────── def initialize(self, session_id: str, **kwargs) -> None: + # Non-secret fields fall back to config.yaml (written by the Dashboard) + # when the env var is unset: env -> config.yaml -> default. + provider_config = _load_retaindb_config() api_key = get_secret("RETAINDB_API_KEY", "") or "" - base_url = re.sub(r"/+$", "", os.environ.get("RETAINDB_BASE_URL", _DEFAULT_BASE_URL)) + base_url_raw = ( + os.environ.get("RETAINDB_BASE_URL") + or _config_str(provider_config.get("base_url")) + or _DEFAULT_BASE_URL + ) + base_url = re.sub(r"/+$", "", base_url_raw) - # Project resolution: RETAINDB_PROJECT > hermes- > "default" + # Project resolution: RETAINDB_PROJECT > config.yaml project > hermes- > "default" # If unset, the API auto-creates and uses the "default" project — no config required. - explicit = os.environ.get("RETAINDB_PROJECT") + explicit = os.environ.get("RETAINDB_PROJECT") or _config_str(provider_config.get("project")) if explicit: project = explicit else: diff --git a/plugins/platforms/discord/adapter.py b/plugins/platforms/discord/adapter.py index 99e0f9a2f5..c7c3036f75 100644 --- a/plugins/platforms/discord/adapter.py +++ b/plugins/platforms/discord/adapter.py @@ -26,14 +26,35 @@ import time from collections import defaultdict from contextlib import suppress from typing import Callable, Dict, List, Optional, Any, Tuple -from urllib.parse import urljoin +from urllib.parse import quote, urljoin from agent.async_utils import ( consume_detached_task_result as _consume_background_task_result, ) +from agent.display import ToolPreview logger = logging.getLogger(__name__) +_DISCORD_MARKDOWN_LINK_LABEL_RE = re.compile(r"([\\\[\]])") +_DISCORD_URL_LABEL_SCHEME_RE = re.compile(r"^https?://", re.IGNORECASE) + + +def _format_discord_markdown_link(label: str, url: str) -> str: + """Return a Discord Markdown link whose label is not itself a URL. + + Discord gives URL-shaped link labels their own link behavior. A truncated + URL label can therefore win over the Markdown destination and remain a + broken link. Dropping only the scheme keeps the preview recognizable while + leaving one unambiguous click target. + + The destination is wrapped in angle brackets (````) so Discord does + not unfurl an OG-preview embed under every tool progress bubble. + """ + label = _DISCORD_URL_LABEL_SCHEME_RE.sub("", label, count=1) + escaped_label = _DISCORD_MARKDOWN_LINK_LABEL_RE.sub(r"\\\1", label) + escaped_url = quote(url, safe=":/?#[]@!$&'*+,;=%") + return f"[{escaped_label}](<{escaped_url}>)" + class _Snowflake: """Minimal object exposing ``.id`` — satisfies discord.py's Snowflake @@ -97,7 +118,6 @@ _DISCORD_NONCONVERSATIONAL_HISTORY_MESSAGE_PATTERNS = ( ), re.compile(r"^\s*♻️?\s+Gateway\s+(?:restarted successfully|online\b)[\s\S]*$", re.IGNORECASE), ) - try: import discord from discord import Message as DiscordMessage, Intents @@ -121,7 +141,11 @@ except ImportError: from gateway.config import Platform, PlatformConfig -from gateway.platforms.helpers import MessageDeduplicator, ThreadParticipationTracker, convert_table_to_bullets +from gateway.platforms.helpers import ( + MessageDeduplicator, + ThreadParticipationTracker, + convert_table_to_bullets, +) from utils import atomic_json_write, env_float, env_int from gateway.platforms.base import ( BasePlatformAdapter, @@ -977,6 +1001,13 @@ class DiscordAdapter(BasePlatformAdapter): PLAYBACK_TIMEOUT = 120 PLAYBACK_TIMEOUT_PADDING = 30 + def format_tool_preview(self, preview: ToolPreview) -> str: + """Keep a truncated URL preview clickable in Discord markdown.""" + if not preview.url: + return preview.text + + return _format_discord_markdown_link(preview.text, preview.url) + def __init__(self, config: PlatformConfig): super().__init__(config, Platform.DISCORD) self._client: Optional[commands.Bot] = None @@ -1742,6 +1773,19 @@ class DiscordAdapter(BasePlatformAdapter): # Cancel the liveness probe first so it can't fire a spurious fatal # error / reconnect while we're intentionally tearing the adapter down. await self._cancel_liveness_task() + # Clean up all active voice connections *before* cancelling the bot task. + # leave_voice_channel() ends in `await vc.disconnect()`, and discord.py's + # VoiceClient.disconnect() sends a voice state update over the main + # gateway websocket and then waits for the voice socket to close. The + # bot task is the loop running that gateway connection, so cancelling it + # first leaves the handshake with no transport: it can never complete and + # blocks until the caller's shutdown timeout fires. + for guild_id in list(self._voice_clients.keys()): + try: + await self.leave_voice_channel(guild_id) + except Exception as e: # pragma: no cover - defensive logging + logger.debug("[%s] Error leaving voice channel %s: %s", self.name, guild_id, e) + # Cancel the bot task before closing the client. If connect() timed out # and returned False, the background client.start() task may still be # running; calling client.close() alone is not enough to stop it because @@ -1749,12 +1793,6 @@ class DiscordAdapter(BasePlatformAdapter): # WebSocket handshake is in flight. Explicitly cancelling the task here # ensures the zombie client cannot receive or dispatch any further events. await self._cancel_bot_task() - # Clean up all active voice connections before closing the client - for guild_id in list(self._voice_clients.keys()): - try: - await self.leave_voice_channel(guild_id) - except Exception as e: # pragma: no cover - defensive logging - logger.debug("[%s] Error leaving voice channel %s: %s", self.name, guild_id, e) if self._client: try: diff --git a/plugins/platforms/feishu/adapter.py b/plugins/platforms/feishu/adapter.py index c21bb42e90..8c50942042 100644 --- a/plugins/platforms/feishu/adapter.py +++ b/plugins/platforms/feishu/adapter.py @@ -85,45 +85,35 @@ try: except ImportError: websockets = None # type: ignore[assignment] -try: - import lark_oapi as lark - from lark_oapi.api.application.v6 import GetApplicationRequest - from lark_oapi.api.im.v1 import ( - CreateFileRequest, - CreateFileRequestBody, - CreateImageRequest, - CreateImageRequestBody, - CreateMessageRequest, - CreateMessageRequestBody, - GetChatRequest, - GetMessageRequest, - GetMessageResourceRequest, - P2ImMessageMessageReadV1, - ReplyMessageRequest, - ReplyMessageRequestBody, - UpdateMessageRequest, - UpdateMessageRequestBody, - ) - from lark_oapi.core import AccessTokenType, HttpMethod - from lark_oapi.core.const import FEISHU_DOMAIN, LARK_DOMAIN - from lark_oapi.core.model import BaseRequest - from lark_oapi.event.callback.model.p2_card_action_trigger import ( - CallBackCard, - P2CardActionTriggerResponse, - ) - from lark_oapi.event.dispatcher_handler import EventDispatcherHandler - from lark_oapi.ws import Client as FeishuWSClient - - FEISHU_AVAILABLE = True -except ImportError: - FEISHU_AVAILABLE = False - lark = None # type: ignore[assignment] - CallBackCard = None # type: ignore[assignment] - P2CardActionTriggerResponse = None # type: ignore[assignment] - EventDispatcherHandler = None # type: ignore[assignment] - FeishuWSClient = None # type: ignore[assignment] - FEISHU_DOMAIN = None # type: ignore[assignment] - LARK_DOMAIN = None # type: ignore[assignment] +# lark_oapi takes a noticeable amount of time to import. Keep the gateway +# configuration path responsive by importing it only when Feishu connects. +lark = None # type: ignore[assignment] +GetApplicationRequest = None # type: ignore[assignment] +CreateFileRequest = None # type: ignore[assignment] +CreateFileRequestBody = None # type: ignore[assignment] +CreateImageRequest = None # type: ignore[assignment] +CreateImageRequestBody = None # type: ignore[assignment] +CreateMessageRequest = None # type: ignore[assignment] +CreateMessageRequestBody = None # type: ignore[assignment] +GetChatRequest = None # type: ignore[assignment] +GetMessageRequest = None # type: ignore[assignment] +GetMessageResourceRequest = None # type: ignore[assignment] +P2ImMessageMessageReadV1 = None # type: ignore[assignment] +ReplyMessageRequest = None # type: ignore[assignment] +ReplyMessageRequestBody = None # type: ignore[assignment] +UpdateMessageRequest = None # type: ignore[assignment] +UpdateMessageRequestBody = None # type: ignore[assignment] +AccessTokenType = None # type: ignore[assignment] +HttpMethod = None # type: ignore[assignment] +FEISHU_DOMAIN = None # type: ignore[assignment] +LARK_DOMAIN = None # type: ignore[assignment] +BaseRequest = None # type: ignore[assignment] +CallBackCard = None # type: ignore[assignment] +P2CardActionTriggerResponse = None # type: ignore[assignment] +EventDispatcherHandler = None # type: ignore[assignment] +FeishuWSClient = None # type: ignore[assignment] +FEISHU_AVAILABLE = False +_lark_import_lock = threading.Lock() FEISHU_WEBSOCKET_AVAILABLE = websockets is not None FEISHU_WEBHOOK_AVAILABLE = aiohttp is not None @@ -1395,36 +1385,38 @@ def _run_official_feishu_ws_client(ws_client: Any, adapter: Any) -> None: adapter._ws_thread_loop = None -def check_feishu_requirements() -> bool: - """Check if Feishu/Lark dependencies are available. - - Lazy-installs lark-oapi via ``tools.lazy_deps.ensure("platform.feishu")`` - on first call if not present. Rebinds all module-level globals on success. - """ +def _load_lark_oapi() -> bool: + """Import and bind the Feishu SDK after an explicit connection request.""" if FEISHU_AVAILABLE: return True - def _import(): - import lark_oapi as lark - from lark_oapi.api.application.v6 import GetApplicationRequest - from lark_oapi.api.im.v1 import ( - CreateFileRequest, CreateFileRequestBody, - CreateImageRequest, CreateImageRequestBody, - CreateMessageRequest, CreateMessageRequestBody, - GetChatRequest, GetMessageRequest, GetMessageResourceRequest, - P2ImMessageMessageReadV1, - ReplyMessageRequest, ReplyMessageRequestBody, - UpdateMessageRequest, UpdateMessageRequestBody, - ) - from lark_oapi.core import AccessTokenType, HttpMethod - from lark_oapi.core.const import FEISHU_DOMAIN, LARK_DOMAIN - from lark_oapi.core.model import BaseRequest - from lark_oapi.event.callback.model.p2_card_action_trigger import ( - CallBackCard, P2CardActionTriggerResponse, - ) - from lark_oapi.event.dispatcher_handler import EventDispatcherHandler - from lark_oapi.ws import Client as FeishuWSClient - return { + with _lark_import_lock: + if FEISHU_AVAILABLE: + return True + try: + import lark_oapi as lark + from lark_oapi.api.application.v6 import GetApplicationRequest + from lark_oapi.api.im.v1 import ( + CreateFileRequest, CreateFileRequestBody, + CreateImageRequest, CreateImageRequestBody, + CreateMessageRequest, CreateMessageRequestBody, + GetChatRequest, GetMessageRequest, GetMessageResourceRequest, + P2ImMessageMessageReadV1, + ReplyMessageRequest, ReplyMessageRequestBody, + UpdateMessageRequest, UpdateMessageRequestBody, + ) + from lark_oapi.core import AccessTokenType, HttpMethod + from lark_oapi.core.const import FEISHU_DOMAIN, LARK_DOMAIN + from lark_oapi.core.model import BaseRequest + from lark_oapi.event.callback.model.p2_card_action_trigger import ( + CallBackCard, P2CardActionTriggerResponse, + ) + from lark_oapi.event.dispatcher_handler import EventDispatcherHandler + from lark_oapi.ws import Client as FeishuWSClient + except ImportError: + return False + + globals().update({ "lark": lark, "GetApplicationRequest": GetApplicationRequest, "CreateFileRequest": CreateFileRequest, @@ -1451,10 +1443,22 @@ def check_feishu_requirements() -> bool: "EventDispatcherHandler": EventDispatcherHandler, "FeishuWSClient": FeishuWSClient, "FEISHU_AVAILABLE": True, - } + }) + return True - from tools.lazy_deps import ensure_and_bind - return ensure_and_bind("platform.feishu", _import, globals(), prompt=False) + +def check_feishu_requirements() -> bool: + """Ensure Feishu dependencies are installed without importing the SDK.""" + if FEISHU_AVAILABLE: + return True + + from tools.lazy_deps import ensure + + try: + ensure("platform.feishu", prompt=False) + return True + except Exception: + return False class FeishuAdapter(BasePlatformAdapter): @@ -1752,9 +1756,6 @@ class FeishuAdapter(BasePlatformAdapter): # A fresh connect (or reconnect) re-arms the SDK executor after a prior # disconnect set the closing flag. self._sdk_executor_closing = False - if not FEISHU_AVAILABLE: - logger.error("[Feishu] lark-oapi not installed") - return False if not self._app_id or not self._app_secret: logger.error("[Feishu] FEISHU_APP_ID or FEISHU_APP_SECRET not set") return False @@ -1769,6 +1770,9 @@ class FeishuAdapter(BasePlatformAdapter): "[Feishu] Webhook mode requires FEISHU_VERIFICATION_TOKEN or FEISHU_ENCRYPT_KEY." ) return False + if not await asyncio.to_thread(_load_lark_oapi): + logger.error("[Feishu] lark-oapi not installed") + return False try: self._app_lock_identity = self._app_id @@ -5047,19 +5051,19 @@ class FeishuAdapter(BasePlatformAdapter): @staticmethod def _build_get_chat_request(chat_id: str) -> Any: - if "GetChatRequest" in globals(): + if GetChatRequest is not None: return GetChatRequest.builder().chat_id(chat_id).build() return SimpleNamespace(chat_id=chat_id) @staticmethod def _build_get_message_request(message_id: str) -> Any: - if "GetMessageRequest" in globals(): + if GetMessageRequest is not None: return GetMessageRequest.builder().message_id(message_id).build() return SimpleNamespace(message_id=message_id) @staticmethod def _build_message_resource_request(*, message_id: str, file_key: str, resource_type: str) -> Any: - if "GetMessageResourceRequest" in globals(): + if GetMessageResourceRequest is not None: return ( GetMessageResourceRequest.builder() .message_id(message_id) @@ -5071,7 +5075,7 @@ class FeishuAdapter(BasePlatformAdapter): @staticmethod def _build_get_application_request(*, app_id: str, lang: str) -> Any: - if "GetApplicationRequest" in globals(): + if GetApplicationRequest is not None: return ( GetApplicationRequest.builder() .app_id(app_id) @@ -5082,7 +5086,7 @@ class FeishuAdapter(BasePlatformAdapter): @staticmethod def _build_reply_message_body(*, content: str, msg_type: str, reply_in_thread: bool, uuid_value: str) -> Any: - if "ReplyMessageRequestBody" in globals(): + if ReplyMessageRequestBody is not None: return ( ReplyMessageRequestBody.builder() .content(content) @@ -5100,7 +5104,7 @@ class FeishuAdapter(BasePlatformAdapter): @staticmethod def _build_reply_message_request(message_id: str, request_body: Any) -> Any: - if "ReplyMessageRequest" in globals(): + if ReplyMessageRequest is not None: return ( ReplyMessageRequest.builder() .message_id(message_id) @@ -5111,7 +5115,7 @@ class FeishuAdapter(BasePlatformAdapter): @staticmethod def _build_update_message_body(*, msg_type: str, content: str) -> Any: - if "UpdateMessageRequestBody" in globals(): + if UpdateMessageRequestBody is not None: return ( UpdateMessageRequestBody.builder() .msg_type(msg_type) @@ -5122,7 +5126,7 @@ class FeishuAdapter(BasePlatformAdapter): @staticmethod def _build_update_message_request(message_id: str, request_body: Any) -> Any: - if "UpdateMessageRequest" in globals(): + if UpdateMessageRequest is not None: return ( UpdateMessageRequest.builder() .message_id(message_id) @@ -5133,7 +5137,7 @@ class FeishuAdapter(BasePlatformAdapter): @staticmethod def _build_create_message_body(*, receive_id: str, msg_type: str, content: str, uuid_value: str) -> Any: - if "CreateMessageRequestBody" in globals(): + if CreateMessageRequestBody is not None: return ( CreateMessageRequestBody.builder() .receive_id(receive_id) @@ -5151,7 +5155,7 @@ class FeishuAdapter(BasePlatformAdapter): @staticmethod def _build_create_message_request(receive_id_type: str, request_body: Any) -> Any: - if "CreateMessageRequest" in globals(): + if CreateMessageRequest is not None: return ( CreateMessageRequest.builder() .receive_id_type(receive_id_type) @@ -5162,7 +5166,7 @@ class FeishuAdapter(BasePlatformAdapter): @staticmethod def _build_image_upload_body(*, image_type: str, image: Any) -> Any: - if "CreateImageRequestBody" in globals(): + if CreateImageRequestBody is not None: return ( CreateImageRequestBody.builder() .image_type(image_type) @@ -5173,13 +5177,13 @@ class FeishuAdapter(BasePlatformAdapter): @staticmethod def _build_image_upload_request(request_body: Any) -> Any: - if "CreateImageRequest" in globals(): + if CreateImageRequest is not None: return CreateImageRequest.builder().request_body(request_body).build() return SimpleNamespace(request_body=request_body) @staticmethod def _build_file_upload_body(*, file_type: str, file_name: str, file: Any, duration: int = 0) -> Any: - if "CreateFileRequestBody" in globals(): + if CreateFileRequestBody is not None: builder = ( CreateFileRequestBody.builder() .file_type(file_type) @@ -5193,7 +5197,7 @@ class FeishuAdapter(BasePlatformAdapter): @staticmethod def _build_file_upload_request(request_body: Any) -> Any: - if "CreateFileRequest" in globals(): + if CreateFileRequest is not None: return CreateFileRequest.builder().request_body(request_body).build() return SimpleNamespace(request_body=request_body) @@ -5410,7 +5414,10 @@ def probe_bot(app_id: str, app_secret: str, domain: str) -> Optional[dict]: Note: ``bot_open_id`` here is the bot's app-scoped open_id — the same ID that Feishu puts in @mention payloads. It is NOT the app_id. """ - if FEISHU_AVAILABLE: + # The SDK import is deferred until connect(); onboarding runs before any + # connect, so load it here to keep the SDK probe path reachable rather + # than silently degrading every setup run to the HTTP fallback. + if _load_lark_oapi(): return _probe_bot_sdk(app_id, app_secret, domain) return _probe_bot_http(app_id, app_secret, domain) @@ -5598,7 +5605,7 @@ async def _standalone_send( FeishuAdapter, hydrates its lark client, and sends text + native media (images, video, voice, documents). Replaces the legacy _send_feishu helper. """ - if not FEISHU_AVAILABLE: + if not await asyncio.to_thread(_load_lark_oapi): return {"error": "Feishu dependencies not installed. Run `hermes setup` to install Feishu support."} media_files = media_files or [] diff --git a/plugins/platforms/line/adapter.py b/plugins/platforms/line/adapter.py index b3c9c6463e..b31c1a209d 100644 --- a/plugins/platforms/line/adapter.py +++ b/plugins/platforms/line/adapter.py @@ -1695,13 +1695,13 @@ def interactive_setup() -> None: print() try: - from hermes_cli.config import get_env_var, set_env_var + from hermes_cli.config import get_env_value as _get_env, save_env_value as _set_env except ImportError: print("hermes_cli.config not available; set LINE_* vars manually in ~/.hermes/.env") return def _prompt(var: str, prompt: str, *, secret: bool = False) -> None: - existing = get_env_var(var) if callable(get_env_var) else None + existing = _get_env(var) if callable(_get_env) else None suffix = " [keep current]" if existing else "" try: if secret: @@ -1713,7 +1713,7 @@ def interactive_setup() -> None: print() return if value: - set_env_var(var, value) + _set_env(var, value) _prompt("LINE_CHANNEL_ACCESS_TOKEN", "Channel access token", secret=True) _prompt("LINE_CHANNEL_SECRET", "Channel secret", secret=True) diff --git a/plugins/platforms/matrix/adapter.py b/plugins/platforms/matrix/adapter.py index b8461ae27d..577540b8e4 100644 --- a/plugins/platforms/matrix/adapter.py +++ b/plugins/platforms/matrix/adapter.py @@ -978,20 +978,61 @@ class _CryptoStateStore: ``get_encryption_info``, and ``find_shared_rooms``. The basic ``MemoryStateStore`` from ``mautrix.client`` doesn't implement these, so we provide simple implementations that consult the client's room - state. + state, falling back to a direct homeserver state event query when + the in-memory store has no encryption info for a room. """ - def __init__(self, client_state_store: Any, joined_rooms: set): + def __init__(self, client_state_store: Any, joined_rooms: set, client=None): self._ss = client_state_store self._joined_rooms = joined_rooms + self._client = client + # Cache encryption info queried from the homeserver so we don't + # make a network round-trip on every is_encrypted() call. + # MemoryStateStore doesn't implement set_encryption_info, so the + # cache-back in get_encryption_info is a no-op without this. + self._enc_info_cache: dict = {} async def is_encrypted(self, room_id: str) -> bool: return (await self.get_encryption_info(room_id)) is not None async def get_encryption_info(self, room_id: str): + info = None if hasattr(self._ss, "get_encryption_info"): - return await self._ss.get_encryption_info(room_id) - return None + info = await self._ss.get_encryption_info(room_id) + if info is not None: + return info + # Check local cache before hitting the homeserver. + if room_id in self._enc_info_cache: + return self._enc_info_cache[room_id] + client = self._client + if client is None: + return None + try: + from mautrix.types import ( + EventType as _ET, + RoomEncryptionStateEventContent as _Enc, + RoomID as _RID, + ) + raw = await client.get_state_event(_RID(room_id), _ET.ROOM_ENCRYPTION) + except Exception as exc: + logger.debug( + "Matrix: homeserver encryption-info query failed for %s: %s", + room_id, + exc, + ) + return None + if not raw: + return None + content = raw if isinstance(raw, _Enc) else _Enc.deserialize( + raw.serialize() if hasattr(raw, "serialize") else raw + ) + if hasattr(self._ss, "set_encryption_info"): + try: + await self._ss.set_encryption_info(_RID(room_id), content) + except Exception: + pass + self._enc_info_cache[room_id] = content + return content async def find_shared_rooms(self, user_id: str) -> list: # Return all joined rooms — simple but correct for a single-user bot. @@ -1303,6 +1344,158 @@ class MatrixAdapter(BasePlatformAdapter): return False return True + async def _reset_crypto_store_if_device_changed( + self, crypto_store: Any, device_id: str + ) -> bool: + """Reset the local Olm account when the access token's device changed. + + The crypto store is keyed by user ID, so a new access token (= new + device ID) would otherwise inherit the previous device's Olm account. + Its identity keys can never be published under the new device ID + (and the pickle key embeds the old device ID anyway), which leads to + stale-key mismatches and cross-signing signatures that the + homeserver refuses to replace. Returns True if the store was reset. + """ + if not device_id: + return False + try: + stored_device_id = await crypto_store.get_device_id() + except Exception as exc: + logger.warning("Matrix: could not read stored device ID: %s", exc) + return False + if not stored_device_id or stored_device_id == device_id: + return False + logger.warning( + "Matrix: access token belongs to a new device (%s -> %s) — " + "resetting local Olm account so fresh identity keys are " + "generated for this device", + stored_device_id, + device_id, + ) + await crypto_store.delete() + return True + + async def _migrate_legacy_crypto_pickle( + self, crypto_store: Any, crypto_db: Any, acct_id: str, pickle_key: str + ) -> bool: + """Re-pickle the Olm account when the pickle key changed. + + The pickle key embeds the *configured* device ID. If the account was + created before MATRIX_DEVICE_ID was set (e.g. first password login, + where the device ID is only known after connecting), the store was + pickled under ``:default``; setting MATRIX_DEVICE_ID afterwards + changes the key and unpickling fails with BAD_ACCOUNT_KEY — which in + optional-E2EE mode silently disables encryption. Try known legacy + keys and re-pickle under the current one. Returns False only when an + account exists but no key can unpickle it. + """ + try: + await crypto_store.get_account() + return True + except Exception: + pass + + from mautrix.crypto.store.asyncpg import PgCryptoStore + + for legacy_key in (f"{acct_id}:default", acct_id): + if legacy_key == pickle_key: + continue + legacy_store = PgCryptoStore( + account_id=acct_id, pickle_key=legacy_key, db=crypto_db + ) + try: + account = await legacy_store.get_account() + except Exception: + continue + if account is None: + continue + # Sessions first, account last. The account is the migration's + # commit marker: once it reads under the current key, the fast + # path above short-circuits every later startup. Writing it + # before the sweep means an interrupted sweep is never retried + # and the remaining legacy-key sessions are stranded for good. + try: + await self._repickle_crypto_sessions( + crypto_db, acct_id, legacy_key, pickle_key + ) + except Exception as exc: + logger.error( + "Matrix: pickle key migration failed while re-pickling " + "sessions (%s) — leaving the account under the legacy " + "key so the migration is retried on the next start.", + exc, + ) + return False + await crypto_store.put_account(account) + logger.info( + "Matrix: re-pickled crypto store account and sessions under " + "the current pickle key (device ID was configured after the " + "account was created)" + ) + return True + + logger.error( + "Matrix: crypto store account exists but cannot be unpickled " + "with the current or any legacy pickle key. If MATRIX_DEVICE_ID " + "was changed manually, restore its previous value." + ) + return False + + async def _repickle_crypto_sessions( + self, crypto_db: Any, acct_id: str, legacy_key: str, pickle_key: str + ) -> None: + """Re-pickle olm/megolm session blobs alongside the account. + + The account and every stored session share the pickle key; migrating + only the account leaves sessions unreadable (BAD_ACCOUNT_KEY on the + next olm decrypt), which breaks key sharing with peers. + """ + import olm as olm_lib + + tables = { + "crypto_olm_session": olm_lib.Session, + "crypto_megolm_inbound_session": olm_lib.InboundGroupSession, + "crypto_megolm_outbound_session": olm_lib.OutboundGroupSession, + } + for table, session_cls in tables.items(): + rows = await crypto_db.fetch( + f"SELECT session_id, session FROM {table} WHERE account_id=$1", + acct_id, + ) + for row in rows: + blob = row["session"] + if blob is None: + continue + pickled = bytes(blob) + try: + session_cls.from_pickle(pickled, pickle_key) + continue # already readable with the current key + except Exception: + pass + try: + session = session_cls.from_pickle(pickled, legacy_key) + except Exception as exc: + # Readable under neither key. The row is left untouched + # rather than deleted — it is already unusable, and + # removing crypto material is not worth doing on a + # guess. It stays inert. + logger.warning( + "Matrix: %s row %s cannot be unpickled with the " + "current or legacy key; leaving it in place, its " + "sessions are unrecoverable: %s", + table, + row["session_id"], + exc, + ) + continue + await crypto_db.execute( + f"UPDATE {table} SET session=$1 " + "WHERE account_id=$2 AND session_id=$3", + session.pickle(pickle_key), + acct_id, + row["session_id"], + ) + async def _verify_device_keys_on_server(self, client: Any, olm: Any) -> bool: """Verify our device keys are on the homeserver after loading crypto state. @@ -1436,13 +1629,38 @@ class MatrixAdapter(BasePlatformAdapter): try: resp = await client.whoami() resolved_user_id = getattr(resp, "user_id", "") or self._user_id - resolved_device_id = getattr(resp, "device_id", "") + resolved_device_id = str(getattr(resp, "device_id", "") or "") if resolved_user_id: self._user_id = str(resolved_user_id) client.mxid = UserID(self._user_id) - # Prefer user-configured device_id for stable E2EE identity. - effective_device_id = self._device_id or resolved_device_id + # Normally the user-configured device_id wins, giving a stable + # E2EE identity when whoami() reports no device. + # + # But an access token is bound to exactly one device, and the + # homeserver only accepts key uploads for that device. If + # MATRIX_DEVICE_ID names a different one, honouring it means + # claiming an identity this token cannot publish — which is + # the stale-key failure mode this reset exists to clear. The + # live whoami() device therefore wins on conflict, loudly. + if ( + resolved_device_id + and self._device_id + and resolved_device_id != self._device_id + ): + logger.error( + "Matrix: MATRIX_DEVICE_ID=%s does not match the device " + "this access token belongs to (%s). A token can only " + "upload keys for its own device, so the configured " + "value is being ignored. Unset MATRIX_DEVICE_ID, or " + "use a token issued for %s.", + self._device_id, + resolved_device_id, + self._device_id, + ) + effective_device_id = resolved_device_id + else: + effective_device_id = self._device_id or resolved_device_id if effective_device_id: client.device_id = effective_device_id @@ -1573,7 +1791,13 @@ class MatrixAdapter(BasePlatformAdapter): self._crypto_db = crypto_db _acct_id = self._user_id or "hermes" - _pickle_key = f"{_acct_id}:{self._device_id or 'default'}" + # Use the resolved client.device_id (from whoami or password + # login), not self._device_id (the configured value), because + # #71543 makes the token's real device win over a stale + # MATRIX_DEVICE_ID. The pickle key must match the device the + # token actually belongs to, or the Olm account is stored + # under a key that can never be looked up again. + _pickle_key = f"{_acct_id}:{client.device_id or self._device_id or 'default'}" crypto_store = PgCryptoStore( account_id=_acct_id, pickle_key=_pickle_key, @@ -1582,9 +1806,25 @@ class MatrixAdapter(BasePlatformAdapter): await crypto_store.open() if client.device_id: + _store_was_reset = await self._reset_crypto_store_if_device_changed( + crypto_store, client.device_id + ) await crypto_store.put_device_id(client.device_id) + else: + _store_was_reset = False - crypto_state = _CryptoStateStore(state_store, self._joined_rooms) + # Skip the pickle-key migration when the store was just + # deleted — there is no account to migrate. + if not _store_was_reset: + if not await self._migrate_legacy_crypto_pickle( + crypto_store, crypto_db, _acct_id, _pickle_key + ): + logger.warning( + "Matrix: crypto pickle migration failed — " + "E2EE may not work correctly" + ) + + crypto_state = _CryptoStateStore(state_store, self._joined_rooms, client) olm = OlmMachine(client, crypto_store, crypto_state) olm.share_keys_min_trust = TrustState.UNVERIFIED olm.send_keys_min_trust = TrustState.UNVERIFIED diff --git a/plugins/platforms/telegram/adapter.py b/plugins/platforms/telegram/adapter.py index 674a49d2be..ea6258b9f2 100644 --- a/plugins/platforms/telegram/adapter.py +++ b/plugins/platforms/telegram/adapter.py @@ -770,6 +770,7 @@ class TelegramAdapter(BasePlatformAdapter): self._drop_delayed_deliveries = False self._polling_error_task: Optional[asyncio.Task] = None self._polling_conflict_count: int = 0 + self._polling_conflict_recovery_generation: Optional[int] = None self._polling_network_error_count: int = 0 self._polling_generation: int = 0 self._polling_progress_event = asyncio.Event() @@ -2160,7 +2161,10 @@ class TelegramAdapter(BasePlatformAdapter): return self._polling_progress_event.set() self._polling_network_error_count = 0 - self._polling_conflict_count = 0 + if generation == self._polling_conflict_recovery_generation: + self._polling_conflict_recovery_generation = None + else: + self._polling_conflict_count = 0 self._send_path_degraded = False def _observe_polling_request_result(self, request, generation, result): @@ -3118,12 +3122,22 @@ class TelegramAdapter(BasePlatformAdapter): # AttributeError deep inside start_polling instead of failing fast # here, where the except below reschedules or escalates to fatal. app = self._app + expected_generation = self._polling_generation + 1 + if not app: + raise RuntimeError("Telegram application was torn down during conflict reconnect") + # drop_pending_updates=True tells Telegram to terminate any + # other active getUpdates sessions for this bot token. The + # competing session is either a zombie from the previous + # gateway process (whose long-poll hasn't expired server-side + # yet) or our own previous retry's still-expiring session. + # Without this, each retry starts a new getUpdates session + # that immediately gets 409'd by the previous one, creating + # the very conflict we are trying to recover from (#75017). + self._polling_conflict_recovery_generation = expected_generation try: - if not app: - raise RuntimeError("Telegram application was torn down during conflict reconnect") await self._start_polling_once( app, - drop_pending_updates=False, + drop_pending_updates=True, error_callback=self._polling_error_callback_ref, ) logger.info( @@ -3164,6 +3178,9 @@ class TelegramAdapter(BasePlatformAdapter): ) return # Fall through to fatal on the last retry. + finally: + if self._polling_conflict_recovery_generation == expected_generation: + self._polling_conflict_recovery_generation = None if getattr(self, "_polling_teardown_started", False): return diff --git a/providers/__init__.py b/providers/__init__.py index a394e74b33..4d828c561d 100644 --- a/providers/__init__.py +++ b/providers/__init__.py @@ -42,6 +42,7 @@ logger = logging.getLogger(__name__) _REGISTRY: dict[str, ProviderProfile] = {} _ALIASES: dict[str, str] = {} +_PROVIDER_LIST_CACHE: list[ProviderProfile] | None = None _discovered = False # Repo-root ``plugins/model-providers/`` — populated at discovery time. @@ -57,9 +58,11 @@ def register_provider(profile: ProviderProfile) -> None: plugins under ``$HERMES_HOME/plugins/model-providers/`` can override bundled profiles without editing repo code. """ + global _PROVIDER_LIST_CACHE _REGISTRY[profile.name] = profile for alias in profile.aliases: _ALIASES[alias] = profile.name + _PROVIDER_LIST_CACHE = None def get_provider_profile(name: str) -> ProviderProfile | None: @@ -75,8 +78,11 @@ def get_provider_profile(name: str) -> ProviderProfile | None: def list_providers() -> list[ProviderProfile]: """Return all registered provider profiles (one per canonical name).""" + global _PROVIDER_LIST_CACHE if not _discovered: _discover_providers() + if _PROVIDER_LIST_CACHE is not None: + return list(_PROVIDER_LIST_CACHE) # Deduplicate: _REGISTRY has canonical names; _ALIASES points to same objects seen: set[int] = set() result: list[ProviderProfile] = [] @@ -85,7 +91,8 @@ def list_providers() -> list[ProviderProfile]: if pid not in seen: seen.add(pid) result.append(profile) - return result + _PROVIDER_LIST_CACHE = result + return list(result) def _user_plugins_dir() -> Path | None: diff --git a/providers/base.py b/providers/base.py index 554e01e4f7..1349d579bb 100644 --- a/providers/base.py +++ b/providers/base.py @@ -72,6 +72,12 @@ class ProviderProfile: # (e.g. Xiaomi MiMo, which returns 400 "text is not set"). supports_vision_tool_messages: bool = True + # True only when this provider's Chat Completions endpoint explicitly + # documents ``prompt_cache_key`` as an accepted request body field. This + # is deliberately opt-in: many OpenAI-compatible endpoints reject unknown + # top-level fields rather than ignoring them. + supports_prompt_cache_key: bool = False + # ── Model catalog ───────────────────────────────────────── # fallback_models: curated list shown in /model picker when live fetch fails. # Only agentic models that support tool calling should appear here. diff --git a/pyproject.toml b/pyproject.toml index 35d4949a53..8d6a3eed9d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -2,7 +2,7 @@ [project] name = "hermes-agent" -version = "0.19.1" +version = "0.20.0" description = "The self-improving AI agent — creates skills from experience, improves them during use, and runs anywhere" readme = "README.md" # Upper bound is load-bearing, not cosmetic. uv resolves the project's @@ -140,8 +140,18 @@ dependencies = [ # First-party lifecycle and shared-metrics runtime. Relay 0.6 is the minimum # lossless provider-codec contract. Managed calls pass request/response data # through this native module in-process; shared metrics installs no network - # exporter and consumes only its bounded projection. - "nemo-relay>=0.6.0,<0.7; (sys_platform == 'darwin' and platform_machine == 'arm64') or (sys_platform == 'linux' and platform_machine == 'x86_64') or (sys_platform == 'linux' and platform_machine == 'aarch64') or (sys_platform == 'win32' and platform_machine == 'AMD64') or (sys_platform == 'win32' and platform_machine == 'ARM64')", + # exporter and consumes only its bounded projection. Relay publishes wheels + # only (no sdist), so this marker must stay false anywhere no wheel tag can + # match — otherwise installing Python dependencies fails resolution outright + # instead of falling back to the no-op Relay host (#76469, Termux). + # Termux Python reports plain linux/aarch64 but runs on Bionic + # libc, which satisfies neither manylinux nor musllinux, hence the + # `'android' not in platform_release` guard on the linux arms: Android GKI + # kernels embed "-androidNN-" in the kernel release string. (Official PEP + # 738 CPython reports sys_platform == 'android' and never matched.) Pre-GKI + # devices can still slip through; they get the same resolution failure as + # before, worked around by installing with `--no-deps` or an older release. + "nemo-relay>=0.6.0,<0.7; (sys_platform == 'darwin' and platform_machine == 'arm64') or (sys_platform == 'linux' and platform_machine == 'x86_64' and 'android' not in platform_release) or (sys_platform == 'linux' and platform_machine == 'aarch64' and 'android' not in platform_release) or (sys_platform == 'win32' and platform_machine == 'AMD64') or (sys_platform == 'win32' and platform_machine == 'ARM64')", ] [project.optional-dependencies] diff --git a/relatorio-issue-69678-sqlite-fd-leaks.md b/relatorio-issue-69678-sqlite-fd-leaks.md deleted file mode 100644 index 866b952516..0000000000 --- a/relatorio-issue-69678-sqlite-fd-leaks.md +++ /dev/null @@ -1,218 +0,0 @@ -# Relatório técnico — Issue #69678 e PRs relacionados - -**Data da análise:** 22 de julho de 2026 -**Repositório:** `NousResearch/hermes-agent` -**Issue principal:** [#69678 — SQLite connections leaked in delivery, async delegation, and verification evidence ledgers](https://github.com/NousResearch/hermes-agent/issues/69678) -**PR principal:** [#69681 — fix(gateway,tools,agent): close leaked SQLite connections in delivery](https://github.com/NousResearch/hermes-agent/pull/69681) - -## Resumo executivo - -A issue #69678 descreve um bug real: três ledgers SQLite usam a conexão como context manager, mas nunca a fecham explicitamente. Em processos de gateway de longa duração, as conexões e seus descritores de arquivo podem permanecer vivos até a coleta pelo garbage collector, acumulando descritores para o banco principal, `-wal` e `-shm`. O processo eventualmente pode atingir `RLIMIT_NOFILE` e começar a falhar com `[Errno 24] Too many open files` em componentes não relacionados. - -A causa raiz apresentada está correta. O PR #69681 corrige os 21 call sites identificados e preserva as semânticas existentes de transação e locking. - -**Conclusão:** #69678 não é duplicata exata de #69567. Ambas pertencem à mesma classe de bug, mas afetam módulos diferentes. O PR #69681 deve ser tratado como fix irmão do PR #69594, não como implementação duplicada. - -## Escopo afetado - -| Módulo | Call sites afetados | Operações que acionam o ledger | Risco | -|---|---:|---|---| -| `gateway/delivery_ledger.py` | 5 | Registro, atualização, recuperação, pruning e inspeção de entregas | Muito alto; executado no fluxo frequente de respostas finais | -| `tools/async_delegation.py` | 13 | Dispatch, conclusão, recuperação, claim, release e confirmação de entrega | Alto durante delegações em background | -| `agent/verification_evidence.py` | 3 | Resultado de terminal, edição do workspace e leitura de status | Cresce com operações de desenvolvimento e verificação | - -Total confirmado: **21 call sites**. - -## Causa raiz - -Os módulos usam o seguinte padrão: - -```python -with _connect() as conn: - ... -``` - -O context manager de `sqlite3.Connection` controla a transação: - -- sucesso: commit; -- exceção: rollback; -- saída do bloco: não executa `conn.close()`. - -Consequentemente, o bloco `with` transmite uma falsa impressão de gerenciamento completo do recurso. A transação termina, mas o lifecycle da conexão não termina de forma determinística. - -Em modo WAL, uma conexão pode manter descritores associados a: - -- arquivo principal do banco; -- arquivo `-wal`; -- arquivo `-shm`. - -Em processo curto, encerramento ou coleta rápida pode mascarar o defeito. Em gateway long-lived, sob tráfego recorrente, acúmulo pode alcançar o soft limit de descritores e causar falhas em leituras de configuração, arquivos temporários, sockets e outros bancos SQLite. - -## Relação com #69567 e PR #69594 - -[Issue #69567](https://github.com/NousResearch/hermes-agent/issues/69567) encontrou a mesma falha em `cron/executions.py`. Uma execução normal de cron abre conexões em `create_execution()`, `mark_execution_running()` e `finish_execution()`. O relato mediu crescimento de descritores até atingir limite do processo. - -[PR #69594](https://github.com/NousResearch/hermes-agent/pull/69594) propõe um `_transaction()` que: - -1. abre a conexão; -2. preserva commit/rollback com `with conn:`; -3. fecha a conexão em `finally`; -4. fecha também quando inicialização de PRAGMA/schema falha. - -O PR #69681 aplica o mesmo modelo a três ledgers não alterados pelo PR #69594. - -### Decisão sobre duplicidade - -**Não marcar #69678 como duplicata de #69567.** - -Justificativa: - -- causa técnica idêntica; -- arquivos e call paths diferentes; -- #69594 altera somente ledger de cron; -- mesmo após #69594, os 21 sites de #69678 continuariam vazando conexões; -- busca de PRs relacionados não encontrou outro PR cobrindo os três módulos de #69678. - -Classificação correta: issues irmãs pertencentes à mesma classe de defeito. - -## Avaliação do PR #69681 - -### Estado observado - -- aberto; -- não é draft; -- GitHub o considera mergeable; -- 1 commit; -- 6 arquivos alterados; -- 512 adições e 29 remoções; -- sem reviews ou comentários no momento da análise. - -### Solução implementada - -Cada módulo recebe um context manager equivalente a: - -```python -@contextmanager -def _transaction() -> Iterator[sqlite3.Connection]: - conn = _connect() - try: - with conn: - yield conn - finally: - conn.close() -``` - -Além disso, `_connect()` passa a fechar a conexão caso PRAGMA ou inicialização de schema falhe depois de `sqlite3.connect()` ter retornado com sucesso. - -### Pontos corretos - -- fechamento determinístico em sucesso, early return e exceção; -- commit e rollback continuam delegados ao context manager nativo; -- locking existente é preservado; -- `_transaction()` não adquire `_DB_LOCK`, evitando lock nesting novo; -- `_prune()` em `delivery_ledger` continua lock-free; -- contrato schema-on-connect é mantido; -- todos os 21 call sites são migrados; -- nenhuma configuração, schema de ferramenta ou superfície core nova é adicionada. - -Nenhum defeito funcional foi identificado no patch analisado. - -## Minimal fix recomendado - -Solução do PR é pequena no comportamento de produção e resolve a causa raiz. Recomendação: manter `_transaction()` local em cada módulo. - -Não criar agora um helper SQLite global compartilhado. Isso aumentaria escopo, acoplamento e risco para resolver três módulos independentes. Uma abstração compartilhada só deve surgir após demanda concreta e contrato comum comprovado. - -Possíveis reduções sem mudar o desenho: - -- encurtar docstrings repetidas dos três `_transaction()`; -- compartilhar fixture de tracking apenas se já existir local apropriado na suíte; -- evitar refatorações adjacentes. - -O volume `+512/-29` vem principalmente dos três arquivos de regressão. O fix runtime em si permanece cirúrgico. - -## Avaliação dos testes - -O PR adiciona testes para: - -- fechamento em operações normais; -- update sem linha correspondente; -- exceção durante operação SQL; -- falha durante inicialização do schema; -- igualdade entre número de conexões abertas e fechadas. - -Os testes usam conexões SQLite reais envolvidas por um proxy que registra chamadas a `close()`. Isso valida diretamente o contrato quebrado e evita depender do timing do garbage collector. - -### Melhoria opcional - -Adicionar teste Linux de integração contando `/proc/self/fd` após várias operações. Esse teste reproduziria o sintoma externo, mas pode ser específico de plataforma e mais frágil. Não deve bloquear merge se a suíte direta de lifecycle e suítes existentes estiverem verdes. - -Resultados citados pelo autor não foram reexecutados nesta análise, pois checkout local contém várias alterações pré-existentes e branch do PR não foi aplicada. Antes do merge, CI deve confirmar: - -```text -tests/gateway/test_delivery_ledger_fd_leak.py -tests/tools/test_async_delegation_fd_leak.py -tests/agent/test_verification_evidence_fd_leak.py -tests/gateway/test_delivery_ledger.py -tests/gateway/test_delivery_ledger_producer.py -tests/tools/test_async_delegation.py -tests/agent/test_verification_evidence.py -``` - -## Issues relacionadas - -Mesma classe geral de lifecycle SQLite, mas escopos diferentes: - -- [#69567](https://github.com/NousResearch/hermes-agent/issues/69567): cron execution ledger; -- [#60859](https://github.com/NousResearch/hermes-agent/issues/60859): leak de `SessionDB` em early return; -- [#30027](https://github.com/NousResearch/hermes-agent/issues/30027): listagem de boards kanban; -- [#28802](https://github.com/NousResearch/hermes-agent/issues/28802): helpers kanban specify; -- [#36111](https://github.com/NousResearch/hermes-agent/issues/36111): lifecycle de `ResponseStore` no API server; -- [#37369](https://github.com/NousResearch/hermes-agent/issues/37369): crescimento de descritores de `response_store.db`. - -Essas issues demonstram padrão recorrente: uso de context manager transacional interpretado incorretamente como gerenciamento completo da conexão. - -## Possíveis ocorrências residuais - -Busca estática encontrou padrões semelhantes fora do escopo do PR, incluindo: - -- `gateway/readiness.py`; -- templates FastMCP em `optional-skills/mcp/fastmcp/templates/database_server.py`. - -Esses locais não devem ser incluídos automaticamente em #69681. Cada ocorrência precisa de: - -1. confirmação de que conexão não possui outro owner; -2. reprodução ou teste de lifecycle; -3. análise da frequência e duração do processo; -4. fix separado quando comportamento estiver comprovado. - -Expandir #69681 para uma auditoria global contrariaria objetivo de mudança cirúrgica. - -## Solução estrutural futura - -Após merge dos fixes urgentes, abrir tarefa separada de auditoria dirigida: - -1. localizar `with sqlite3.connect(...)`, `with connect(...)` e `with _connect(...)`; -2. classificar conexões por lifecycle: per-operation ou long-lived; -3. exigir `contextlib.closing`, `try/finally close()` ou helper transacional para conexões per-operation; -4. adicionar teste de regressão somente para call paths reais; -5. documentar em guia interno que `sqlite3.Connection` context manager não fecha a conexão. - -Evitar mudança mecânica global: alguns componentes podem manter conexão deliberadamente durante lifetime do serviço. - -## Recomendação final - -1. Não fechar #69678 como duplicata. -2. Tratar #69594 e #69681 como fixes irmãos. -3. Aprovar desenho do PR #69681 após CI verde. -4. Não ampliar PR para ocorrências não reproduzidas. -5. Criar follow-up separado para auditoria de lifecycle SQLite no repositório. - -**Decisão sugerida:** merge do PR #69681 após validação automática e review, seguido pelo fechamento da issue #69678 como concluída. - ---- - -## 📊 Infográfico: Vazamento de FDs e Correção - -![Infográfico de Correção de SQLite Leaks](sqlite_leak_fix.png) - diff --git a/run_agent.py b/run_agent.py index 39a1689bad..346102b0dc 100644 --- a/run_agent.py +++ b/run_agent.py @@ -148,6 +148,7 @@ from tools.browser_tool import cleanup_browser # Agent internals extracted to agent/ package for modularity from agent.memory_manager import sanitize_context +from agent.memory_provider import is_trivial_prompt from agent.error_classifier import FailoverReason from agent.redact import redact_sensitive_text from agent.message_content import flatten_message_text @@ -630,11 +631,24 @@ class AIAgent: _profile_for_session = None except Exception: _profile_for_session = None + # Carry the live YOLO bypass into the creation-time model_config so + # a session whose /yolo was toggled BEFORE the row existed (the row + # is created lazily on the first turn) still persists the flag for + # `hermes --resume`. set_session_yolo() no-ops on a missing row, so + # this is the only chance to record a pre-first-turn toggle. + _init_model_config = self._session_init_model_config + try: + from tools.approval import is_session_yolo_enabled + if is_session_yolo_enabled(self.session_id): + _init_model_config = dict(_init_model_config or {}) + _init_model_config["yolo_mode"] = True + except Exception: + pass self._session_db.create_session( session_id=self.session_id, source=source, model=self.model, - model_config=self._session_init_model_config, + model_config=_init_model_config, system_prompt=self._cached_system_prompt, user_id=None, parent_session_id=self._parent_session_id, @@ -2082,6 +2096,10 @@ class AIAgent: ): _scan_start += 1 + # Collect this flush's new rows and write them in ONE transaction + # at the end of the scan (see append_messages_batch). + _batch_rows: List[Dict[str, Any]] = [] + _batch_msgs: List[Dict] = [] for _msg_idx in range(_scan_start, len(messages)): msg = messages[_msg_idx] if not isinstance(msg, dict): @@ -2200,33 +2218,48 @@ class AIAgent: ] elif isinstance(msg.get("tool_calls"), list): tool_calls_data = msg["tool_calls"] - self._session_db.append_message( - session_id=self.session_id, - role=role, - content=content, - tool_name=msg.get("tool_name"), - tool_calls=tool_calls_data, - tool_call_id=msg.get("tool_call_id"), - finish_reason=msg.get("finish_reason"), - reasoning=msg.get("reasoning") if role == "assistant" else None, - reasoning_content=msg.get("reasoning_content") 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, - codex_message_items=msg.get("codex_message_items") if role == "assistant" else None, - timestamp=_row_timestamp, - api_content=_row_api_content, - display_kind=( + _batch_rows.append({ + "role": role, + "content": content, + "tool_name": msg.get("tool_name"), + "tool_calls": tool_calls_data, + "tool_call_id": msg.get("tool_call_id"), + "finish_reason": msg.get("finish_reason"), + # Reasoning/codex fields are role-gated (assistant-only) + # inside _insert_message_rows — pass through untouched. + "reasoning": msg.get("reasoning"), + "reasoning_content": msg.get("reasoning_content"), + "reasoning_details": msg.get("reasoning_details"), + "codex_reasoning_items": msg.get("codex_reasoning_items"), + "codex_message_items": msg.get("codex_message_items"), + "timestamp": _row_timestamp, + "api_content": _row_api_content, + "display_kind": ( "hidden" if msg.get(COMPRESSED_SUMMARY_METADATA_KEY) and not msg.get("_compressed_summary_has_user_turn") else msg.get("display_kind") ), - display_metadata=msg.get("display_metadata"), + "display_metadata": msg.get("display_metadata"), + }) + _batch_msgs.append(msg) + # One transaction for the whole turn's new rows (typically 3-8 + # messages): one BEGIN IMMEDIATE / commit — and, off WAL, one + # fsync — instead of one per row. All-or-nothing pairs exactly + # with the marker stamping below: on failure NO rows landed and + # NO markers were stamped, so the next flush re-scans and + # re-writes the whole tail (same recovery contract as before, + # minus the partial-prefix case that could double-pay counters). + if _batch_rows: + self._session_db.append_messages_batch( + session_id=self.session_id, + messages=_batch_rows, compression_lock_holder=getattr( self, "_active_compression_lock_holder", None ), ) - msg[_DB_PERSISTED_MARKER] = True + for _written in _batch_msgs: + _written[_DB_PERSISTED_MARKER] = True # The intrinsic markers are now the sole source of truth. Reset the # one-shot seed so no id() outlives this flush to alias a message # allocated next turn at a recycled address. @@ -4093,10 +4126,15 @@ class AIAgent: response_text, **sync_kwargs, ) - self._memory_manager.queue_prefetch_all( - user_text, - session_id=self.session_id or "", - ) + # Sibling of the build_turn_context() prefetch gate: warming the + # next turn's recall with a trivial prompt ("hi", "thanks") keys + # provider searches on zero-signal text — skip it. The sync above + # still runs so the turn itself is persisted. + if not is_trivial_prompt(user_text): + self._memory_manager.queue_prefetch_all( + user_text, + session_id=self.session_id or "", + ) except Exception: pass @@ -4236,6 +4274,21 @@ class AIAgent: except Exception: pass + # 6c. Close the Codex app-server session. The runtime already drops + # it on turn crash / retirement (agent/codex_runtime.py), but hard + # teardown had no owner — a /new, /reset, or session expiry left the + # app-server child process running until interpreter exit. Clear the + # attribute BEFORE close() so a concurrent reader can't grab a + # half-closed session, and so a raising close() can't strand a stale + # reference behind. + try: + codex_session = getattr(self, "_codex_session", None) + if codex_session is not None: + self._codex_session = None + codex_session.close() + except Exception: + pass + # 7. Free conversation history. Mirrors _release_evicted_agent_soft's # soft-eviction clear — close() is the hard teardown for true session # boundaries (/new, /reset, session expiry), so the message list won't @@ -7586,11 +7639,14 @@ class AIAgent: turn_id=relay_turn_id, task_id=effective_task_id, ) - start_task_run( - **task_context, - parent_session_id=getattr(self, "_parent_session_id", None) or "", - ) - task_started = True + # Keep existing tests and external relay-runtime shims that return + # a minimal turn object compatible with the new opt-out flag. + if getattr(relay_turn, "relay_enabled", True): + start_task_run( + **task_context, + parent_session_id=getattr(self, "_parent_session_id", None) or "", + ) + task_started = True # Publish the conversation id for ambient Nous Portal tagging. Every # LLM call made inside this turn — main loop, compression, vision, # web_extract, session_search, MoA slots, background-review forks @@ -7636,8 +7692,9 @@ class AIAgent: relay_turn, outcome=relay_outcome, ) - task_finished = True - finish_task_run(**task_context, result=result) + if task_started: + task_finished = True + finish_task_run(**task_context, result=result) return result except BaseException as exc: if isinstance(exc, (KeyboardInterrupt, InterruptedError)) or ( diff --git a/scripts/contributor_audit.py b/scripts/contributor_audit.py index 6bd5795672..11a02dda24 100644 --- a/scripts/contributor_audit.py +++ b/scripts/contributor_audit.py @@ -51,6 +51,12 @@ IGNORED_PATTERNS = [ re.compile(r"^Hermes\s+(Agent|Audit)$", re.IGNORECASE), re.compile(r"^nousbot(-eng)?$", re.IGNORECASE), re.compile(r"^Ubuntu$", re.IGNORECASE), + # v0.20.0 audit additions: + re.compile(r"^Blut-?Agent$", re.IGNORECASE), # self-described AI agent account + re.compile(r".*\[bot\]$", re.IGNORECASE), # any GitHub [bot] suffix (hermes-seaeye[bot] etc.) + re.compile(r"^TRON$", re.IGNORECASE), # AgentMail agent + re.compile(r"^Happy$", re.IGNORECASE), # happy.engineering AI agent + re.compile(r"^Orca$", re.IGNORECASE), # Stably AI agent ] IGNORED_EMAILS = { @@ -65,6 +71,10 @@ IGNORED_EMAILS = { "omx@oh-my-codex.dev", "codex@openai.com", "noreply@commandcode.ai", + # v0.20.0 audit additions — AI-agent co-author trailers: + "tron-agent@agentmail.to", # TRON (AgentMail agent) + "yesreply@happy.engineering", # Happy (AI coding agent) + "help@stably.ai", # Orca (Stably AI agent) } diff --git a/scripts/release.py b/scripts/release.py index 7b0e776444..eb4e779c10 100755 --- a/scripts/release.py +++ b/scripts/release.py @@ -79,6 +79,7 @@ LEGACY_AUTHOR_MAP = { "75556242+webtecnica@users.noreply.github.com": "webtecnica", # PR #63360 salvage (nous: restore inference-api base_url) "contato@webtecnica.com.br": "webtecnica", # PR #70888 salvage "webtecnica@gmail.com": "webtecnica", # PR #75838 salvage (aux: free-only fallback guard; #75803) + "ckaznocha@gmail.com": "ckaznocha", # PR #71543 salvage (matrix: crypto store reset + pickle migration) "skosarevivan@yandex.ru": "Epoxidex", # PR #29820 salvage (ollama: top-level reasoning_effort=none; #25758) "jdjiayou@163.com": "JiaDe-Wu", # PR #34742 salvage (bedrock: bearer routing + streaming fallback + image decode; #28156) "changhyun.min@gmail.com": "minchang", # PR #42231 salvage (providers: add Upstage Solar) @@ -108,6 +109,9 @@ LEGACY_AUTHOR_MAP = { "xwolf.live@gmail.com": "vizi0uz", # PR #59795 adopted in #62290 "wilsonkinyuam@gmail.com": "WilsonKinyua", # PR #62052 (tui: persist unflushed conversations on disconnect/restart) "humphreysun98@gmail.com": "HumphreySun98", # PR #61142 salvage (web: null web/backend config value guards) + "merlin@threewizards.agency": "light-merlin-dark", # PR #7821 salvage (zai: parallel endpoint detection probes) + "4850809+frizikk@users.noreply.github.com": "frizikk", # PR #63389 salvage (session-search: fields projection skips unused context enrichment) + "endeavorisforever@gmail.com": "EndeavorYen", # PR #33971 salvage (image: parallel image_generate batches + FileSyncManager transaction lock) "sonxi@nous.local": "17324393074", # PR #53196 salvage (tools_config: known_plugin_toolsets null guard; commit under unlinked local identity) "lemonwan@users.noreply.github.com": "lemonwan", # PR #59430 sibling salvage (adapter reconnect contract guard) "luxuguangno1@163.com": "luxuguang-leo", # PR #52966 + #52908 salvage (QQBot reconnect + Feishu Channel signaling) @@ -721,6 +725,7 @@ LEGACY_AUTHOR_MAP = { "mike@grossmann.at": "ReqX", "axmaiqiu@gmail.com": "qWaitCrypto", "44045911+kidonng@users.noreply.github.com": "kidonng", + "ayushere@users.noreply.github.com": "ayushere", "daniellsmarta@gmail.com": "DanielLSM", "264291321+v1b3coder@users.noreply.github.com": "v1b3coder", "silverchris@foxmail.com": "ming1523", @@ -1346,6 +1351,7 @@ LEGACY_AUTHOR_MAP = { "iamagenius00@users.noreply.github.com": "iamagenius00", "9219265+cresslank@users.noreply.github.com": "cresslank", "trevmanthony@gmail.com": "trevthefoolish", + "at828@proton.me": "ATran28", # PR #77270 whatsapp bridge reconnect wedge fix "ziliangpeng@users.noreply.github.com": "ziliangpeng", "ziliangdotme@gmail.com": "ziliangpeng", "centripetal-star@users.noreply.github.com": "centripetal-star", @@ -1439,6 +1445,7 @@ LEGACY_AUTHOR_MAP = { "xiayh17@gmail.com": "xiayh0107", "zhujianxyz@gmail.com": "opriz", "tuancanhnguyen706@gmail.com": "xxxigm", + "j.brownemoore@gmail.com": "ElSnacko", "timchris.roth@pm.me": "x9x9x9x9x9x91", "larcombe.n@gmail.com": "NickLarcombe", "54813621+xxxigm@users.noreply.github.com": "xxxigm", @@ -2053,6 +2060,7 @@ LEGACY_AUTHOR_MAP = { "rodisoft1@gmail.com": "0disoft", # PR #53511 salvage (gateway PID probe TTL cache) "craigs.seller.sixx@gmail.com": "0-CYBERDYNE-SYSTEMS-0", # PR #53966 salvage (session DB reads off event loop) "sebastianlutycz@users.noreply.github.com": "sebastianlutycz", # PR #39140 salvage (descendant CTE); bare noreply (no NNN+ prefix) needs explicit mapping + "bobclawblaw@users.noreply.github.com": "BobClawblaw", # PR #77870 salvage (output-cap compression on retry path; #55546) "wafy.081107@gmail.com": "mahdiwafy", # PR #60347 salvage (session messages pagination) "codeforgenet@icloud.com": "CodeForgeNet", # PR #47437 salvage (compact_rows blob skip) "i@dex.moe": "dexhunter", # PR #60339 salvage (skills snapshot manifest speedup) diff --git a/scripts/whatsapp-bridge/bridge.js b/scripts/whatsapp-bridge/bridge.js index fbacb6b43f..234cbef279 100644 --- a/scripts/whatsapp-bridge/bridge.js +++ b/scripts/whatsapp-bridge/bridge.js @@ -35,6 +35,8 @@ import { createOutboundIdTracker } from './outbound_ids.js'; import { classifyOwnerMessageGate } from './owner_message_gate.js'; import { buildPollPayload, + createReconnectScheduler, + createVersionResolver, buildLocationPayload, buildTextSendPayload, createBoundedMessageStore, @@ -393,12 +395,15 @@ function emitPairEvent(event) { } catch {} } +const scheduleReconnect = createReconnectScheduler(() => startSocket()); +const getWAVersion = createVersionResolver(fetchLatestBaileysVersion); + async function startSocket() { const { state, saveCreds } = await useMultiFileAuthState(SESSION_DIR); - const { version } = await fetchLatestBaileysVersion(); + const version = await getWAVersion(); sock = makeWASocket({ - version, + ...(version ? { version } : {}), auth: state, logger, printQRInTerminal: false, @@ -449,7 +454,7 @@ async function startSocket() { console.log(`⚠️ Connection closed (reason: ${reason}). Reconnecting in 3s...`); } } - setTimeout(startSocket, reason === 515 ? 1000 : 3000); + scheduleReconnect(reason === 515 ? 1000 : 3000); } } else if (connection === 'open') { connectionState = 'connected'; @@ -1145,6 +1150,6 @@ if (PAIR_ONLY) { console.log(`👤 WHATSAPP_FORWARD_OWNER_MESSAGES=true — owner-typed messages will be forwarded with fromOwner:true`); } console.log(); - startSocket(); + scheduleReconnect(0); }); } diff --git a/scripts/whatsapp-bridge/bridge.reconnect.test.mjs b/scripts/whatsapp-bridge/bridge.reconnect.test.mjs new file mode 100644 index 0000000000..70f53751cf --- /dev/null +++ b/scripts/whatsapp-bridge/bridge.reconnect.test.mjs @@ -0,0 +1,150 @@ +/** + * Unit tests for the reconnect scheduling and version resolution guards. + * + * Regression tests for the reconnect-wedge trap: startSocket() awaits + * network I/O (fetchLatestBaileysVersion has no AbortSignal) before it + * creates a socket, and the close handler used to re-enter it via a bare + * `setTimeout(startSocket, ...)`. A rejection was unhandled and a stalled + * fetch left the bridge permanently disconnected while its HTTP server + * kept answering 503 — observed in the field as a bridge that logged + * "Reconnecting in 3s..." once and then went silent for 27+ hours. + * + * These tests avoid importing bridge.js because that file starts an HTTP + * server and Baileys socket at module load. Keep the helper module pure. + */ + +import { strict as assert } from 'node:assert'; + +import { + createReconnectScheduler, + createVersionResolver, +} from './bridge_helpers.js'; + +const tick = () => new Promise(resolve => setImmediate(resolve)); +const sleep = ms => new Promise(resolve => setTimeout(resolve, ms)); + +// -- createReconnectScheduler --------------------------------------------- + +// A rejecting start function is caught and rescheduled at the retry delay; +// a subsequent success stops the retry chain. +{ + const timers = []; + const logs = []; + let attempts = 0; + const startFn = async () => { + attempts += 1; + if (attempts === 1) throw new Error('boom'); + }; + + const schedule = createReconnectScheduler(startFn, { + retryDelayMs: 5000, + log: line => logs.push(line), + setTimeoutFn: (fn, ms) => timers.push({ fn, ms }), + }); + + schedule(3000); + assert.equal(timers.length, 1); + assert.equal(timers[0].ms, 3000); + + timers[0].fn(); + await tick(); + await tick(); + + assert.equal(attempts, 1); + assert.equal(logs.length, 1); + assert.match(logs[0], /Reconnect failed \(boom\)/); + assert.equal(timers.length, 2, 'rejection must schedule a retry'); + assert.equal(timers[1].ms, 5000); + + timers[1].fn(); + await tick(); + await tick(); + + assert.equal(attempts, 2); + assert.equal(timers.length, 2, 'success must not schedule another attempt'); + assert.equal(logs.length, 1); +} + +// A synchronous throw from the start function is contained the same way as +// an async rejection. +{ + const timers = []; + const logs = []; + const schedule = createReconnectScheduler( + () => { throw new Error('sync boom'); }, + { + retryDelayMs: 1000, + log: line => logs.push(line), + setTimeoutFn: (fn, ms) => timers.push({ fn, ms }), + }, + ); + + schedule(0); + timers[0].fn(); + await tick(); + await tick(); + + assert.equal(logs.length, 1); + assert.match(logs[0], /sync boom/); + assert.equal(timers.length, 2); +} + +// -- createVersionResolver ------------------------------------------------ + +// A successful fetch returns and caches the version. +{ + const resolveVersion = createVersionResolver( + async () => ({ version: [2, 3000, 99] }), + { log: () => {} }, + ); + assert.deepEqual(await resolveVersion(), [2, 3000, 99]); +} + +// A fetch that never settles resolves within the timeout bound instead of +// pending forever; before any success there is no cache, so the resolver +// yields null (callers fall back to the Baileys default). +{ + const logs = []; + const resolveVersion = createVersionResolver( + () => new Promise(() => {}), + { timeoutMs: 20, log: line => logs.push(line) }, + ); + assert.equal(await resolveVersion(), null); + assert.equal(logs.length, 1); + assert.match(logs[0], /version fetch timed out/); + assert.match(logs[0], /library default/); +} + +// After one success, later failures fall back to the cached version. +{ + const logs = []; + let calls = 0; + const resolveVersion = createVersionResolver( + async () => { + calls += 1; + if (calls === 1) return { version: [2, 3000, 42] }; + throw new Error('network down'); + }, + { timeoutMs: 20, log: line => logs.push(line) }, + ); + assert.deepEqual(await resolveVersion(), [2, 3000, 42]); + assert.deepEqual(await resolveVersion(), [2, 3000, 42]); + assert.equal(logs.length, 1); + assert.match(logs[0], /network down/); + assert.match(logs[0], /cached version/); +} + +// The losing timeout timer is cleared after a fast success, so the resolver +// does not hold the event loop open for the full timeout window. +{ + const resolveVersion = createVersionResolver( + async () => ({ version: [2, 3000, 1] }), + { timeoutMs: 60_000, log: () => {} }, + ); + const before = Date.now(); + await resolveVersion(); + await sleep(10); + assert.ok(Date.now() - before < 1000); +} + +console.log('bridge.reconnect.test.mjs: all assertions passed'); diff --git a/scripts/whatsapp-bridge/bridge_helpers.js b/scripts/whatsapp-bridge/bridge_helpers.js index c0618cd2bd..398521feee 100644 --- a/scripts/whatsapp-bridge/bridge_helpers.js +++ b/scripts/whatsapp-bridge/bridge_helpers.js @@ -567,3 +567,60 @@ export function pollCreationMessageFromPayload(payload) { }; return message; } + +/** + * Reconnect scheduling guard. startSocket() awaits network I/O before it + * creates a socket or registers event handlers, so a bare + * `setTimeout(startSocket, ...)` has two unrecoverable failure modes: a + * rejection is unhandled (crashes the process on modern Node), and a hang + * leaves the bridge permanently disconnected with nothing left to retry. + * Every (re)connect must go through the scheduler this returns. + */ +export function createReconnectScheduler(startFn, { + retryDelayMs = 5000, + log = console.log, + setTimeoutFn = setTimeout, +} = {}) { + function scheduleReconnect(delayMs) { + setTimeoutFn(() => { + Promise.resolve() + .then(startFn) + .catch((err) => { + log(`⚠️ Reconnect failed (${err?.message || err}). Retrying in ${Math.round(retryDelayMs / 1000)}s...`); + scheduleReconnect(retryDelayMs); + }); + }, delayMs); + } + return scheduleReconnect; +} + +/** + * Version resolution guard. fetchLatestBaileysVersion() is a plain fetch to + * raw.githubusercontent.com with no AbortSignal; a stalled connection can + * pend forever and wedge the reconnect path (the scheduler above cannot + * retry past an await that never settles). Bound the fetch and fall back to + * the last known-good version, or the Baileys default before first success. + */ +export function createVersionResolver(fetchVersionFn, { + timeoutMs = 15000, + log = console.log, +} = {}) { + let cachedVersion = null; + return async function resolveVersion() { + let timer = null; + try { + const { version } = await Promise.race([ + fetchVersionFn(), + new Promise((_, reject) => { + timer = setTimeout(() => reject(new Error('version fetch timed out')), timeoutMs); + }), + ]); + cachedVersion = version; + } catch (err) { + log(`⚠️ Baileys version fetch failed (${err?.message || err}); using ${cachedVersion ? 'cached version' : 'library default'}.`); + } finally { + if (timer) clearTimeout(timer); + } + return cachedVersion; + }; +} diff --git a/tests/acp/test_permissions.py b/tests/acp/test_permissions.py index 649a388f01..3d8576d3e2 100644 --- a/tests/acp/test_permissions.py +++ b/tests/acp/test_permissions.py @@ -119,7 +119,7 @@ class TestApprovalBridge: assert result == "always" - def test_timeout_returns_deny_and_cancels_future(self): + def test_timeout_returns_timeout_and_cancels_future(self): loop = MagicMock(spec=asyncio.AbstractEventLoop) request_permission = AsyncMock(name="request_permission") future = MagicMock(spec=Future) @@ -138,7 +138,9 @@ class TestApprovalBridge: scheduled["coro"].close() - assert result == "deny" + # A no-response expiry is classified as "timeout" (still blocked, + # fail-closed) so the agent isn't told the user explicitly refused. + assert result == "timeout" assert scheduled["loop"] is loop assert future.cancel.call_count == 1 diff --git a/tests/agent/test_anthropic_adapter.py b/tests/agent/test_anthropic_adapter.py index d29ed2080e..0b8a7b08bb 100644 --- a/tests/agent/test_anthropic_adapter.py +++ b/tests/agent/test_anthropic_adapter.py @@ -182,6 +182,9 @@ class TestIsClaudeCodeTokenValid: class TestResolveAnthropicToken: + def _assert_not_called(*_args, **_kwargs): + raise AssertionError("should not be called when API key is present") + def test_prefers_oauth_token_over_api_key(self, monkeypatch, tmp_path): monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-api03-mykey") monkeypatch.setenv("ANTHROPIC_TOKEN", "sk-ant-oat01-mytoken") @@ -199,14 +202,41 @@ class TestResolveAnthropicToken: assert resolve_anthropic_token() is None 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.setenv("ANTHROPIC_API_KEY", "sk-ant...ykey") monkeypatch.delenv("ANTHROPIC_TOKEN", raising=False) monkeypatch.delenv("CLAUDE_CODE_OAUTH_TOKEN", raising=False) monkeypatch.setattr("agent.anthropic_adapter.Path.home", lambda: tmp_path) - assert resolve_anthropic_token() == "sk-ant-api03-mykey" + assert resolve_anthropic_token() == "sk-ant...ykey" + def test_api_key_wins_over_auto_discovered_claude_code_credentials( + self, monkeypatch, tmp_path + ): + monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant...ykey") + monkeypatch.delenv("ANTHROPIC_TOKEN", raising=False) + 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, + } + })) + monkeypatch.setattr("agent.anthropic_adapter.Path.home", lambda: tmp_path) + assert resolve_anthropic_token() == "sk-ant...ykey" + def test_api_key_path_does_not_read_auto_discovered_credentials(self, monkeypatch): + monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant...ykey") + monkeypatch.delenv("ANTHROPIC_TOKEN", raising=False) + monkeypatch.delenv("CLAUDE_CODE_OAUTH_TOKEN", raising=False) + monkeypatch.setattr( + "agent.anthropic_adapter.read_claude_code_credentials", + self._assert_not_called, + ) + + assert resolve_anthropic_token() == "sk-ant...ykey" def test_falls_back_to_claude_code_credentials(self, monkeypatch, tmp_path): monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) @@ -229,7 +259,7 @@ class TestResolveAnthropicToken: monkeypatch.delenv("ANTHROPIC_TOKEN", raising=False) monkeypatch.delenv("CLAUDE_CODE_OAUTH_TOKEN", raising=False) monkeypatch.setattr("agent.anthropic_adapter.Path.home", lambda: tmp_path) - # Isolate source #4 (credential_pool): ensure source #3 (Claude Code + # Isolate source #5 (credential_pool): ensure source #4 (Claude Code # creds, incl. the macOS keychain read which Path.home does not cover) # returns nothing, mirroring a Hermes-PKCE-only setup. monkeypatch.setattr("agent.anthropic_adapter.read_claude_code_credentials", lambda: None) @@ -239,38 +269,33 @@ class TestResolveAnthropicToken: access_token="pool-oauth-token", ) pool = SimpleNamespace( - _available_entries=lambda **_kwargs: [pool_entry], + _available_entries=lambda **_kwargs: ([pool_entry], []), ) monkeypatch.setattr("agent.credential_pool.load_pool", lambda provider: pool) assert resolve_anthropic_token() == "pool-oauth-token" - def test_prefers_anthropic_credential_pool_oauth_over_api_key(self, monkeypatch, tmp_path): + def test_api_key_wins_over_anthropic_credential_pool_oauth(self, monkeypatch, tmp_path): monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant...ykey") monkeypatch.delenv("ANTHROPIC_TOKEN", raising=False) monkeypatch.delenv("CLAUDE_CODE_OAUTH_TOKEN", raising=False) monkeypatch.setattr("agent.anthropic_adapter.Path.home", lambda: tmp_path) - # Pool (source #4) must win over ANTHROPIC_API_KEY (source #5); also - # isolate source #3 so a machine-local Claude Code creds / keychain - # entry can't short-circuit before the pool. - monkeypatch.setattr("agent.anthropic_adapter.read_claude_code_credentials", lambda: None) - - pool_entry = SimpleNamespace( - auth_type="oauth", - access_token="pool-oauth-token", + monkeypatch.setattr( + "agent.anthropic_adapter.read_claude_code_credentials", + self._assert_not_called, ) - pool = SimpleNamespace( - _available_entries=lambda **_kwargs: [pool_entry], + monkeypatch.setattr( + "agent.credential_pool.load_pool", + self._assert_not_called, ) - monkeypatch.setattr("agent.credential_pool.load_pool", lambda provider: pool) - assert resolve_anthropic_token() == "pool-oauth-token" + assert resolve_anthropic_token() == "sk-ant...ykey" def test_pool_entry_with_null_access_token_does_not_crash(self, monkeypatch, tmp_path): """A persisted OAuth entry with access_token=None must not crash the resolver (None.strip() would escape the helper's try/excepts and take down the whole resolver incl. the ANTHROPIC_API_KEY fallback). It should - be skipped and the api-key fallback (source #5) should win.""" + be skipped and the api-key fallback (source #3) should win.""" monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant...ykey") monkeypatch.delenv("ANTHROPIC_TOKEN", raising=False) monkeypatch.delenv("CLAUDE_CODE_OAUTH_TOKEN", raising=False) @@ -279,11 +304,11 @@ class TestResolveAnthropicToken: broken_entry = SimpleNamespace(auth_type="oauth", access_token=None) pool = SimpleNamespace( - _available_entries=lambda **_kwargs: [broken_entry], + _available_entries=lambda **_kwargs: ([broken_entry], []), ) monkeypatch.setattr("agent.credential_pool.load_pool", lambda provider: pool) - # Must fall through to source #5 (ANTHROPIC_API_KEY), not raise. + # Must fall through to source #3 (ANTHROPIC_API_KEY), not raise. assert resolve_anthropic_token() == "sk-ant...ykey" def test_pool_api_key_only_entry_is_not_returned_as_token(self, monkeypatch, tmp_path): @@ -299,7 +324,7 @@ class TestResolveAnthropicToken: api_key_entry = SimpleNamespace(auth_type="api_key", access_token="sk-pool-apikey") pool = SimpleNamespace( - _available_entries=lambda **_kwargs: [api_key_entry], + _available_entries=lambda **_kwargs: ([api_key_entry], []), ) monkeypatch.setattr("agent.credential_pool.load_pool", lambda provider: pool) @@ -322,7 +347,7 @@ class TestResolveAnthropicToken: def _available_entries(**kwargs): captured.update(kwargs) - return [pool_entry] + return ([pool_entry], []) pool = SimpleNamespace(_available_entries=_available_entries) monkeypatch.setattr("agent.credential_pool.load_pool", lambda provider: pool) @@ -880,7 +905,7 @@ class TestConvertMessages: assert system == "You are helpful." assert result[0]["role"] == "user" - assert result[0]["content"] == [{"type": "text", "text": " "}] + assert result[0]["content"] == [{"type": "text", "text": "(empty)"}] assert result[1]["role"] == "assistant" assert any( m["role"] == "assistant" and "Context compaction summary" in str(m["content"]) @@ -1670,3 +1695,190 @@ class TestReplayAllBlankFallback: result = self._convert(msg) texts = [b for b in result["content"] if b.get("type") == "text"] assert texts == [{"type": "text", "text": "(empty)"}] + + +def _find_blank_text_blocks(messages): + """Recursively scan a converted Anthropic message list (including + nested tool_result content) for any text block whose text is empty or + whitespace-only. Returns a list of (message_index, role, location, + block_index) tuples for every violation found -- empty means the + payload is safe to send to Anthropic.""" + violations = [] + for m_idx, msg in enumerate(messages): + content = msg.get("content") + if not isinstance(content, list): + continue + for b_idx, blk in enumerate(content): + if not isinstance(blk, dict): + continue + if blk.get("type") == "text" and not ( + isinstance(blk.get("text"), str) and blk["text"].strip() + ): + violations.append((m_idx, msg.get("role"), "content", b_idx)) + if blk.get("type") == "tool_result" and isinstance(blk.get("content"), list): + for ib_idx, iblk in enumerate(blk["content"]): + if ( + isinstance(iblk, dict) + and iblk.get("type") == "text" + and not (isinstance(iblk.get("text"), str) and iblk["text"].strip()) + ): + violations.append((m_idx, msg.get("role"), "tool_result", ib_idx)) + return violations + + +class TestFinalPayloadHasNoBlankTextBlocks: + """End-to-end regression tests on the true final payload boundary: + ``convert_messages_to_anthropic`` -- the last transform before + ``build_anthropic_kwargs`` hands ``messages`` to the Anthropic SDK. + + Covers the blank-content shapes enumerated for the "text content + blocks must contain non-whitespace text" HTTP 400 class, verifying the + final built payload never contains a blank text block while tool_use, + tool_result, and image content are preserved. + """ + + def test_user_message_empty_string_content(self): + messages = [{"role": "user", "content": ""}] + _, result = convert_messages_to_anthropic(messages) + assert _find_blank_text_blocks(result) == [] + assert result[0]["content"] == "(empty message)" + + def test_user_message_whitespace_only_string_content(self): + messages = [{"role": "user", "content": " "}] + _, result = convert_messages_to_anthropic(messages) + assert _find_blank_text_blocks(result) == [] + assert result[0]["content"] == "(empty message)" + + def test_user_message_blank_list_content(self): + messages = [{"role": "user", "content": [{"type": "text", "text": ""}]}] + _, result = convert_messages_to_anthropic(messages) + assert _find_blank_text_blocks(result) == [] + assert result[0]["content"] == [{"type": "text", "text": "(empty message)"}] + + def test_user_message_mixed_blank_and_valid_text_blocks(self): + """A blank text block sitting alongside a non-blank one must be + dropped individually -- not left in place (the all-or-nothing bug) + and not used as an excuse to nuke the valid sibling block.""" + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "real question"}, + {"type": "text", "text": " "}, + ], + } + ] + _, result = convert_messages_to_anthropic(messages) + assert _find_blank_text_blocks(result) == [] + assert result[0]["content"] == [{"type": "text", "text": "real question"}] + + def test_mixed_blank_text_plus_valid_tool_block_preserved(self): + """Blank text next to a valid non-text block (tool_result) must + drop only the blank text and keep the tool block intact.""" + messages = [ + {"role": "user", "content": "call a tool"}, + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_1", + "function": {"name": "web_search", "arguments": '{"query": "x"}'}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_1", "content": "result text"}, + ] + _, result = convert_messages_to_anthropic(messages) + assert _find_blank_text_blocks(result) == [] + assistant_msg = next(m for m in result if m["role"] == "assistant") + tool_use_blocks = [b for b in assistant_msg["content"] if b.get("type") == "tool_use"] + assert len(tool_use_blocks) == 1 + tool_result_msg = next( + m + for m in result + if m["role"] == "user" + and isinstance(m["content"], list) + and any(b.get("type") == "tool_result" for b in m["content"]) + ) + assert tool_result_msg is not None + + def test_assistant_tool_call_message_with_blank_content(self): + """OpenAI-wire-shaped assistant turn: content is a blank string, + tool_calls carries the real payload. Must not surface a blank text + block, and the tool_use block must survive untouched.""" + messages = [ + {"role": "user", "content": "do it"}, + { + "role": "assistant", + "content": " ", + "tool_calls": [ + { + "id": "call_2", + "function": {"name": "web_search", "arguments": '{"query": "y"}'}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_2", "content": "ok"}, + ] + _, result = convert_messages_to_anthropic(messages) + assert _find_blank_text_blocks(result) == [] + assistant_msg = next(m for m in result if m["role"] == "assistant") + assert assistant_msg["content"] == [ + {"type": "tool_use", "id": "call_2", "name": "web_search", "input": {"query": "y"}} + ] + + def test_leading_synthesized_user_turn_is_non_blank(self): + """_ensure_leading_user_turn's synthesized filler must itself be + non-whitespace -- regression for the literal " " placeholder bug.""" + messages = [ + {"role": "system", "content": "sys"}, + {"role": "assistant", "content": "[Context compaction summary] earlier work"}, + {"role": "user", "content": "continue"}, + ] + _, result = convert_messages_to_anthropic(messages) + assert _find_blank_text_blocks(result) == [] + assert result[0]["content"] == [{"type": "text", "text": "(empty)"}] + + def test_blank_text_nested_in_tool_result_content_is_dropped(self): + """A blank text part nested inside a tool_result's own multimodal + content list (e.g. alongside an image) must be scrubbed without + losing the image.""" + messages = [ + {"role": "user", "content": "screenshot please"}, + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_3", + "function": {"name": "screenshot", "arguments": "{}"}, + } + ], + }, + { + "role": "tool", + "tool_call_id": "call_3", + "content": [ + {"type": "text", "text": " "}, + { + "type": "image_url", + "image_url": {"url": "data:image/png;base64,AAAA"}, + }, + ], + }, + ] + _, result = convert_messages_to_anthropic(messages) + assert _find_blank_text_blocks(result) == [] + tool_result_msg = next( + m + for m in result + if m["role"] == "user" + and isinstance(m["content"], list) + and any(b.get("type") == "tool_result" for b in m["content"]) + ) + tool_result_block = next( + b for b in tool_result_msg["content"] if b.get("type") == "tool_result" + ) + image_blocks = [b for b in tool_result_block["content"] if b.get("type") == "image"] + assert len(image_blocks) == 1 diff --git a/tests/agent/test_auxiliary_concurrency.py b/tests/agent/test_auxiliary_concurrency.py new file mode 100644 index 0000000000..27fccd97ec --- /dev/null +++ b/tests/agent/test_auxiliary_concurrency.py @@ -0,0 +1,403 @@ +"""Tests for per-task concurrency limiting on auxiliary LLM calls (#23324).""" + +import asyncio +import threading +import time +from unittest.mock import MagicMock, AsyncMock, patch + +import pytest + +from agent.auxiliary_client import ( + call_llm, + async_call_llm, + _acquire_sync_aux_semaphore, + _acquire_async_aux_semaphore, + _get_task_max_concurrency, + _reset_aux_semaphores, +) + + +@pytest.fixture(autouse=True) +def _clean_semaphore_cache(): + _reset_aux_semaphores() + yield + _reset_aux_semaphores() + + +class TestGetTaskMaxConcurrency: + def test_returns_none_for_missing_task(self): + assert _get_task_max_concurrency(None) is None + assert _get_task_max_concurrency("") is None + + def test_returns_none_when_unset(self): + with patch( + "agent.auxiliary_client._get_auxiliary_task_config", return_value={} + ): + assert _get_task_max_concurrency("title_generation") is None + + def test_does_not_reuse_vision_cpu_limit_for_llm_calls(self): + with patch( + "agent.auxiliary_client._get_auxiliary_task_config", + return_value={"max_concurrency": 1}, + ): + assert _get_task_max_concurrency("vision") is None + + def test_returns_int_when_configured(self): + with patch( + "agent.auxiliary_client._get_auxiliary_task_config", + return_value={"max_concurrency": 3}, + ): + assert _get_task_max_concurrency("compression") == 3 + + def test_returns_none_for_non_numeric(self): + with patch( + "agent.auxiliary_client._get_auxiliary_task_config", + return_value={"max_concurrency": "not-a-number"}, + ): + assert _get_task_max_concurrency("compression") is None + + def test_returns_none_for_zero_or_negative(self): + with patch( + "agent.auxiliary_client._get_auxiliary_task_config", + return_value={"max_concurrency": 0}, + ): + assert _get_task_max_concurrency("compression") is None + with patch( + "agent.auxiliary_client._get_auxiliary_task_config", + return_value={"max_concurrency": -2}, + ): + assert _get_task_max_concurrency("compression") is None + + +class TestSemaphoreCache: + def test_sync_returns_none_when_unset(self): + with patch( + "agent.auxiliary_client._get_auxiliary_task_config", return_value={} + ): + assert _acquire_sync_aux_semaphore("title_generation") is None + + def test_sync_reuses_semaphore_for_same_limit(self): + with patch( + "agent.auxiliary_client._get_auxiliary_task_config", + return_value={"max_concurrency": 2}, + ): + sem1 = _acquire_sync_aux_semaphore("compression") + sem2 = _acquire_sync_aux_semaphore("compression") + assert sem1 is sem2 + + def test_sync_rebuilds_when_limit_changes(self): + cfg = {"max_concurrency": 2} + with patch( + "agent.auxiliary_client._get_auxiliary_task_config", + return_value=cfg, + ): + sem1 = _acquire_sync_aux_semaphore("compression") + cfg["max_concurrency"] = 5 + sem2 = _acquire_sync_aux_semaphore("compression") + assert sem1 is not sem2 + + @pytest.mark.asyncio + async def test_async_reuses_semaphore_within_same_loop(self): + with patch( + "agent.auxiliary_client._get_auxiliary_task_config", + return_value={"max_concurrency": 2}, + ): + sem1 = _acquire_async_aux_semaphore("compression") + sem2 = _acquire_async_aux_semaphore("compression") + assert sem1 is sem2 + + def test_async_returns_none_with_no_running_loop(self): + with patch( + "agent.auxiliary_client._get_auxiliary_task_config", + return_value={"max_concurrency": 2}, + ): + # Called outside an asyncio loop — should bail rather than crash. + assert _acquire_async_aux_semaphore("compression") is None + + +class TestSyncCallEnforcesLimit: + def test_call_llm_caps_concurrent_inflight(self): + limit = 2 + n_callers = 6 + + active = 0 + max_active = 0 + lock = threading.Lock() + + def fake_create(**kwargs): + nonlocal active, max_active + with lock: + active += 1 + if active > max_active: + max_active = active + try: + time.sleep(0.05) + finally: + with lock: + active -= 1 + return MagicMock() + + client = MagicMock() + client.base_url = "https://example.test/v1" + client.chat.completions.create.side_effect = fake_create + + with ( + patch( + "agent.auxiliary_client._resolve_task_provider_model", + return_value=("openrouter", "test-model", None, None, None), + ), + patch( + "agent.auxiliary_client._get_cached_client", + return_value=(client, "test-model"), + ), + patch( + "agent.auxiliary_client._validate_llm_response", + side_effect=lambda resp, _task, **_kwargs: resp, + ), + patch( + "agent.auxiliary_client._get_auxiliary_task_config", + return_value={"max_concurrency": limit}, + ), + ): + threads = [ + threading.Thread( + target=lambda: call_llm( + task="title_generation", + messages=[{"role": "user", "content": "hi"}], + ) + ) + for _ in range(n_callers) + ] + for t in threads: + t.start() + for t in threads: + t.join(timeout=5) + + assert max_active <= limit, f"observed {max_active} > limit {limit}" + assert client.chat.completions.create.call_count == n_callers + + def test_call_llm_unlimited_when_not_configured(self): + client = MagicMock() + client.base_url = "https://example.test/v1" + client.chat.completions.create.return_value = MagicMock() + + with ( + patch( + "agent.auxiliary_client._resolve_task_provider_model", + return_value=("openrouter", "test-model", None, None, None), + ), + patch( + "agent.auxiliary_client._get_cached_client", + return_value=(client, "test-model"), + ), + patch( + "agent.auxiliary_client._validate_llm_response", + side_effect=lambda resp, _task, **_kwargs: resp, + ), + patch( + "agent.auxiliary_client._get_auxiliary_task_config", + return_value={}, + ), + ): + # With no max_concurrency in config, no semaphore is acquired. + call_llm( + task="title_generation", + messages=[{"role": "user", "content": "hi"}], + ) + + assert client.chat.completions.create.call_count == 1 + + def test_semaphore_released_on_exception(self): + """Errors inside call_llm must release the semaphore so the next call proceeds.""" + client = MagicMock() + client.base_url = "https://example.test/v1" + client.chat.completions.create.side_effect = RuntimeError("boom") + + with ( + patch( + "agent.auxiliary_client._resolve_task_provider_model", + return_value=("openrouter", "test-model", None, None, None), + ), + patch( + "agent.auxiliary_client._get_cached_client", + return_value=(client, "test-model"), + ), + patch( + "agent.auxiliary_client._validate_llm_response", + side_effect=lambda resp, _task, **_kwargs: resp, + ), + patch( + "agent.auxiliary_client._get_auxiliary_task_config", + return_value={"max_concurrency": 1}, + ), + ): + for _ in range(3): + with pytest.raises(RuntimeError, match="boom"): + call_llm( + task="title_generation", + messages=[{"role": "user", "content": "hi"}], + ) + + def test_stream_holds_permit_until_consumed_and_preserves_options(self): + client = MagicMock() + client.base_url = "https://example.test/v1" + client.chat.completions.create.side_effect = [iter(["chunk"]), MagicMock()] + second_call_started = threading.Event() + + def make_second_call(): + second_call_started.set() + call_llm( + task="compression", + messages=[{"role": "user", "content": "second"}], + ) + + with ( + patch( + "agent.auxiliary_client._resolve_task_provider_model", + return_value=("openrouter", "test-model", None, None, None), + ), + patch( + "agent.auxiliary_client._get_cached_client", + return_value=(client, "test-model"), + ), + patch( + "agent.auxiliary_client._validate_llm_response", + side_effect=lambda response, _task, **_kwargs: response, + ), + patch( + "agent.auxiliary_client._get_auxiliary_task_config", + return_value={"max_concurrency": 1}, + ), + ): + stream = call_llm( + task="compression", + messages=[{"role": "user", "content": "first"}], + stream=True, + stream_options={"include_usage": True}, + ) + thread = threading.Thread(target=make_second_call) + thread.start() + assert second_call_started.wait(timeout=1) + time.sleep(0.05) + assert client.chat.completions.create.call_count == 1 + assert list(stream) == ["chunk"] + thread.join(timeout=1) + + assert not thread.is_alive() + assert client.chat.completions.create.call_count == 2 + assert client.chat.completions.create.call_args_list[0].kwargs["stream"] is True + assert client.chat.completions.create.call_args_list[0].kwargs["stream_options"] == { + "include_usage": True + } + + def test_api_mode_is_forwarded_to_client_resolution(self): + client = MagicMock() + client.base_url = "https://example.test/v1" + client.chat.completions.create.return_value = MagicMock() + + with ( + patch( + "agent.auxiliary_client._resolve_task_provider_model", + return_value=("openrouter", "test-model", None, None, None), + ), + patch( + "agent.auxiliary_client._get_cached_client", + return_value=(client, "test-model"), + ) as get_client, + patch( + "agent.auxiliary_client._validate_llm_response", + side_effect=lambda response, _task, **_kwargs: response, + ), + ): + call_llm( + task="title_generation", + messages=[{"role": "user", "content": "hi"}], + api_mode="codex_responses", + ) + + assert get_client.call_args.kwargs["api_mode"] == "codex_responses" + + +class TestAsyncCallEnforcesLimit: + @pytest.mark.asyncio + async def test_async_call_llm_caps_concurrent_inflight(self): + limit = 2 + n_callers = 6 + + active = 0 + max_active = 0 + + async def fake_create(**kwargs): + nonlocal active, max_active + active += 1 + if active > max_active: + max_active = active + try: + await asyncio.sleep(0.05) + finally: + active -= 1 + return MagicMock() + + client = MagicMock() + client.base_url = "https://example.test/v1" + client.chat.completions.create = AsyncMock(side_effect=fake_create) + + with ( + patch( + "agent.auxiliary_client._resolve_task_provider_model", + return_value=("openrouter", "test-model", None, None, None), + ), + patch( + "agent.auxiliary_client._get_cached_client", + return_value=(client, "test-model"), + ), + patch( + "agent.auxiliary_client._validate_llm_response", + side_effect=lambda resp, _task, **_kwargs: resp, + ), + patch( + "agent.auxiliary_client._get_auxiliary_task_config", + return_value={"max_concurrency": limit}, + ), + ): + await asyncio.gather(*[ + async_call_llm( + task="compression", + messages=[{"role": "user", "content": "hi"}], + ) + for _ in range(n_callers) + ]) + + assert max_active <= limit, f"observed {max_active} > limit {limit}" + assert client.chat.completions.create.await_count == n_callers + + @pytest.mark.asyncio + async def test_async_semaphore_released_on_exception(self): + client = MagicMock() + client.base_url = "https://example.test/v1" + client.chat.completions.create = AsyncMock(side_effect=RuntimeError("boom")) + + with ( + patch( + "agent.auxiliary_client._resolve_task_provider_model", + return_value=("openrouter", "test-model", None, None, None), + ), + patch( + "agent.auxiliary_client._get_cached_client", + return_value=(client, "test-model"), + ), + patch( + "agent.auxiliary_client._validate_llm_response", + side_effect=lambda resp, _task, **_kwargs: resp, + ), + patch( + "agent.auxiliary_client._get_auxiliary_task_config", + return_value={"max_concurrency": 1}, + ), + ): + for _ in range(3): + with pytest.raises(RuntimeError, match="boom"): + await async_call_llm( + task="compression", + messages=[{"role": "user", "content": "hi"}], + ) diff --git a/tests/agent/test_context_route_mismatch.py b/tests/agent/test_context_route_mismatch.py new file mode 100644 index 0000000000..c03a0ae21b --- /dev/null +++ b/tests/agent/test_context_route_mismatch.py @@ -0,0 +1,56 @@ +"""Tests for agent_init._context_route_mismatch context-pin scoping.""" + +from agent.agent_init import _context_route_mismatch + + +class TestContextRouteMismatchNamedCustomProvider: + """Named custom providers store the URL under custom_providers, not model.base_url. + + Gateway session-reset banners used to treat empty model.base_url + a runtime + custom URL as a route mismatch, drop model.context_length, and fall back to + the Qwen family default (131072) even though /status still showed the pin. + """ + + def test_same_named_custom_provider_keeps_pin_without_configured_url(self): + assert ( + _context_route_mismatch( + None, + "http://127.0.0.1:8080/v1", + "custom-local-agentw", + "custom-local-agentw", + ) + is False + ) + + def test_different_provider_still_clears_pin(self): + assert ( + _context_route_mismatch( + None, + "http://127.0.0.1:8080/v1", + "custom-local-agentw", + "openrouter", + ) + is True + ) + + def test_explicit_base_url_mismatch_still_clears_pin(self): + assert ( + _context_route_mismatch( + "http://127.0.0.1:8080/v1", + "http://10.0.0.2:8080/v1", + "custom-local-agentw", + "custom-local-agentw", + ) + is True + ) + + def test_catalog_provider_rejects_non_default_runtime_url(self): + assert ( + _context_route_mismatch( + None, + "http://127.0.0.1:8080/v1", + "openrouter", + "openrouter", + ) + is True + ) diff --git a/tests/agent/test_credential_pool.py b/tests/agent/test_credential_pool.py index fa3d2b2111..8d7e78c74e 100644 --- a/tests/agent/test_credential_pool.py +++ b/tests/agent/test_credential_pool.py @@ -1352,6 +1352,151 @@ def test_load_pool_seeds_copilot_via_gh_auth_token(tmp_path, monkeypatch): assert entries[0].base_url == "https://api.githubcopilot.com" +def test_load_pool_skips_exchange_for_suppressed_copilot(tmp_path, monkeypatch): + """A suppressed copilot source must NOT run the token exchange. + + Regression test: the suppression gate used to sit AFTER + ``get_copilot_api_token`` (which retries 3x with backoff, ~13s worst + case), so every pool load — model picker open, /model, agent startup — + burned the full exchange dead time for a source the user had already + removed with ``hermes auth remove copilot gh_cli``. The gate must run + BEFORE the network call. + """ + monkeypatch.setenv("HERMES_HOME", str(tmp_path / "hermes")) + _write_auth_store( + tmp_path, + { + "version": 1, + "credential_pool": {}, + "suppressed_sources": {"copilot": ["gh_cli"]}, + }, + ) + + monkeypatch.setattr( + "hermes_cli.copilot_auth.resolve_copilot_token", + lambda: ("gho_fake_token_abc123", "gh auth token"), + ) + + exchange_called = False + + def _boom(token): + nonlocal exchange_called + exchange_called = True + raise AssertionError("exchange must not run for a suppressed source") + + monkeypatch.setattr( + "hermes_cli.copilot_auth.get_copilot_api_token", + _boom, + ) + + from agent.credential_pool import load_pool + pool = load_pool("copilot") + + assert not exchange_called + assert not pool.has_credentials() + assert pool.entries() == [] + + +def test_load_pool_respects_env_var_copilot_suppression(tmp_path, monkeypatch): + """Suppressing env:GH_TOKEN must gate a GH_TOKEN-sourced token. + + Regression test for the source_name classification: a substring match + (``"gh" in source.lower()``) classified GH_TOKEN/GITHUB_TOKEN as gh_cli, + so a user's env-var-specific suppression was silently bypassed and the + exchange ran anyway. + """ + monkeypatch.setenv("HERMES_HOME", str(tmp_path / "hermes")) + _write_auth_store( + tmp_path, + { + "version": 1, + "credential_pool": {}, + "suppressed_sources": {"copilot": ["env:GH_TOKEN"]}, + }, + ) + + monkeypatch.setattr( + "hermes_cli.copilot_auth.resolve_copilot_token", + lambda: ("gho_fake_token_env", "GH_TOKEN"), + ) + + exchange_called = False + + def _boom(token): + nonlocal exchange_called + exchange_called = True + raise AssertionError("exchange must not run for a suppressed env source") + + monkeypatch.setattr( + "hermes_cli.copilot_auth.get_copilot_api_token", + _boom, + ) + + from agent.credential_pool import load_pool + pool = load_pool("copilot") + + assert not exchange_called + assert pool.entries() == [] + + +def test_load_pool_gh_cli_suppression_does_not_block_env_tokens(tmp_path, monkeypatch): + """Suppressing gh_cli must NOT swallow an env-var-sourced token. + + The inverse of the substring bug: GH_TOKEN misclassified as gh_cli meant + suppressing the CLI path also silently dropped env tokens. + """ + monkeypatch.setenv("HERMES_HOME", str(tmp_path / "hermes")) + _write_auth_store( + tmp_path, + { + "version": 1, + "credential_pool": {}, + "suppressed_sources": {"copilot": ["gh_cli"]}, + }, + ) + + monkeypatch.setattr( + "hermes_cli.copilot_auth.resolve_copilot_token", + lambda: ("gho_fake_token_env", "GH_TOKEN"), + ) + monkeypatch.setattr( + "hermes_cli.copilot_auth.get_copilot_api_token", + lambda token: ("capi_exchanged_token", None), + ) + + from agent.credential_pool import load_pool + pool = load_pool("copilot") + + assert [e.source for e in pool.entries()] == ["env:GH_TOKEN"] + + +def test_load_pool_skips_resolve_when_all_copilot_sources_suppressed(tmp_path, monkeypatch): + """With every copilot source suppressed, resolve_copilot_token (which + shells out to ``gh auth token``) must not run at all.""" + monkeypatch.setenv("HERMES_HOME", str(tmp_path / "hermes")) + from hermes_cli.copilot_auth import COPILOT_ENV_VARS + _write_auth_store( + tmp_path, + { + "version": 1, + "credential_pool": {}, + "suppressed_sources": { + "copilot": ["gh_cli"] + [f"env:{v}" for v in COPILOT_ENV_VARS], + }, + }, + ) + + def _boom(): + raise AssertionError("resolve_copilot_token must not run when all sources are suppressed") + + monkeypatch.setattr("hermes_cli.copilot_auth.resolve_copilot_token", _boom) + + from agent.credential_pool import load_pool + pool = load_pool("copilot") + + assert pool.entries() == [] + + def test_load_pool_seeds_qwen_oauth_via_cli_tokens(tmp_path, monkeypatch): diff --git a/tests/agent/test_credential_pool_deferred_refresh.py b/tests/agent/test_credential_pool_deferred_refresh.py new file mode 100644 index 0000000000..d23b68486d --- /dev/null +++ b/tests/agent/test_credential_pool_deferred_refresh.py @@ -0,0 +1,93 @@ +"""Thread-safety of the deferred single-use-token refresh path (#71775). + +The deferred path deliberately runs OAuth network I/O outside the pool +lock. These tests pin the two invariants that make that safe: + +1. `select()` does NOT hold the pool lock while the deferred refresh's + network call runs (the whole point of the PR). +2. The pool mutations that follow the network call (`_replace_entry`, + `_persist`) DO re-serialize under the pool lock, so a concurrent + `select()`/rotation cannot tear `self._entries` or double-write + auth.json. +""" + +import threading +from dataclasses import replace + +from agent.credential_pool import ( + AUTH_TYPE_OAUTH, + CredentialPool, + PooledCredential, +) + + +def _codex_entry(entry_id: str = "codex-1") -> PooledCredential: + return PooledCredential( + provider="openai-codex", + id=entry_id, + label="test codex", + auth_type=AUTH_TYPE_OAUTH, + priority=0, + source="device_code", + access_token="at-stale", + refresh_token="rt-stale", + expires_at_ms=1, # long expired -> needs refresh + ) + + +def test_select_does_not_hold_pool_lock_during_deferred_refresh(monkeypatch): + pool = CredentialPool("openai-codex", [_codex_entry()]) + lock_free_during_refresh = {} + + def _fake_refresh(entry, *, force): + # If select() still held the pool lock here, this non-blocking + # acquire would fail — the regression this PR exists to fix. + acquired = pool._lock.acquire(blocking=False) + lock_free_during_refresh["value"] = acquired + if acquired: + pool._lock.release() + refreshed = replace(entry, access_token="at-fresh", expires_at_ms=2**53) + pool._replace_entry(entry, refreshed) + return refreshed + + monkeypatch.setattr( + pool, "_entry_needs_refresh", lambda e: e.access_token == "at-stale" + ) + monkeypatch.setattr(pool, "_refresh_entry", _fake_refresh) + monkeypatch.setattr(pool, "_persist", lambda **kw: None) + + selected = pool.select() + + assert lock_free_during_refresh.get("value") is True, ( + "select() held the pool lock during the deferred refresh network window" + ) + assert selected is not None + assert selected.access_token == "at-fresh" + + +def test_deferred_mutations_serialize_against_concurrent_rotation(monkeypatch): + """_replace_entry/_persist from the deferred path must contend on the + pool lock: with the lock held by another thread, the deferred mutation + must block rather than mutate concurrently.""" + pool = CredentialPool("openai-codex", [_codex_entry()]) + monkeypatch.setattr(pool, "_persist", lambda **kw: None) + + entry = pool._entries[0] + refreshed = replace(entry, access_token="at-fresh") + + mutated = threading.Event() + + def _deferred_mutation(): + pool._replace_entry(entry, refreshed) # self-locking + mutated.set() + + with pool._lock: + t = threading.Thread(target=_deferred_mutation) + t.start() + # While we hold the lock, the deferred mutation must NOT complete. + assert not mutated.wait(timeout=0.3), ( + "_replace_entry mutated the pool while another thread held the lock" + ) + t.join(timeout=5) + assert mutated.is_set() + assert pool._entries[0].access_token == "at-fresh" diff --git a/tests/agent/test_credential_pool_lease_refresh_reselect.py b/tests/agent/test_credential_pool_lease_refresh_reselect.py new file mode 100644 index 0000000000..267acf9ae5 --- /dev/null +++ b/tests/agent/test_credential_pool_lease_refresh_reselect.py @@ -0,0 +1,115 @@ +"""acquire_lease must re-select after a deferred single-use-token refresh. + +Post-merge gate-sweep finding on the #71775 salvage (deferred refresh moved +OUTSIDE the pool lock). ``select()`` re-selects once the refreshed entries +are back in rotation (credential_pool.py, select() -> "if pending_refresh: +re-select"); ``acquire_lease()`` did not, so a pool whose only entries all +needed a refresh returned None even though the refresh had just succeeded — +the caller saw "no credentials available" and failed a request that should +have gone through. + +These tests stub ``_available_entries`` / ``_refresh_pending_entries`` at the +same seam the production deferred-refresh contract uses: _available_entries +returns ``(available, pending_refresh)`` and entries pending a refresh are +NOT in ``available`` until the refresh has run. +""" + +import threading + +from agent.credential_pool import CredentialPool, PooledCredential + + +def _entry(entry_id: str) -> PooledCredential: + return PooledCredential( + id=entry_id, + provider="anthropic", + auth_type="oauth", + access_token="tok", + label=entry_id, + source="oauth", + priority=0, + ) + + +def _bare_pool(entries): + """Minimal pool shell — avoids disk/keyring I/O in __init__.""" + pool = CredentialPool.__new__(CredentialPool) + pool._lock = threading.RLock() + pool._entries = list(entries) + pool._active_leases = {} + pool._current_id = None + pool._max_concurrent = 2 + pool._unmatched_rotation_streak = 0 + pool.provider = "anthropic" + return pool + + +def _wire_deferred_refresh(pool, *, refresh_succeeds: bool = True): + """Model the deferred-refresh contract with an explicit state flag.""" + state = {"needs_refresh": True, "refresh_calls": 0} + + def fake_refresh(pending): + state["refresh_calls"] += 1 + if refresh_succeeds: + state["needs_refresh"] = False + + def fake_available(clear_expired=False, refresh=False): + if state["needs_refresh"]: + # Pending a refresh -> not yet available. + pending = [(e.id, "tok") for e in pool._entries] if refresh else [] + return [], pending + return list(pool._entries), [] + + pool._refresh_pending_entries = fake_refresh + pool._available_entries = fake_available + return state + + +def test_acquire_lease_reselects_after_deferred_refresh(): + """The only entry needs a refresh; once refreshed it is available, so a + lease MUST be granted rather than reporting no credentials.""" + pool = _bare_pool([_entry("e1")]) + state = _wire_deferred_refresh(pool) + + lease = pool.acquire_lease() + + assert state["refresh_calls"] == 1, "the deferred refresh should run once" + assert state["needs_refresh"] is False, "entry is available post-refresh" + assert lease == "e1", ( + "acquire_lease returned None despite a successfully refreshed, " + "available entry — the caller would fail an answerable request" + ) + assert pool._active_leases.get("e1") == 1, "the lease must be recorded" + + +def test_acquire_lease_without_pending_refresh_does_not_double_select(): + """No pending refresh -> exactly one selection pass (no wasted work).""" + pool = _bare_pool([_entry("e1")]) + state = _wire_deferred_refresh(pool) + state["needs_refresh"] = False # already healthy + + passes = {"n": 0} + original = pool._acquire_lease_under_lock + + def counting(credential_id): + passes["n"] += 1 + return original(credential_id) + + pool._acquire_lease_under_lock = counting + + lease = pool.acquire_lease() + + assert lease == "e1" + assert passes["n"] == 1, "healthy pool must not trigger the retry path" + assert state["refresh_calls"] == 0 + + +def test_acquire_lease_still_none_when_refresh_does_not_help(): + """If the refresh leaves nothing available, None is still the answer — + the retry must not loop or invent a credential.""" + pool = _bare_pool([_entry("e1")]) + state = _wire_deferred_refresh(pool, refresh_succeeds=False) + + assert pool.acquire_lease() is None + assert state["refresh_calls"] == 1, "retry must not refresh repeatedly" + assert pool._active_leases == {} diff --git a/tests/agent/test_credential_pool_quarantine_locking.py b/tests/agent/test_credential_pool_quarantine_locking.py new file mode 100644 index 0000000000..fb4634b717 --- /dev/null +++ b/tests/agent/test_credential_pool_quarantine_locking.py @@ -0,0 +1,136 @@ +"""Codex/nous quarantine paths must mutate self._entries under the lock. + +Post-merge gate-sweep finding on the #71775 salvage (#77714). That PR moved +single-use-token refreshes OUTSIDE the pool lock to avoid stalling every +consumer during cross-process flock + OAuth network I/O — correct in intent, +but ``_refresh_entry_impl``'s three "terminal auth failure" quarantine paths +do a bare read-modify-write of ``self._entries``: + + removed_ids = [item.id for item in self._entries if ...] + self._entries = [item for item in self._entries if ...] + +Before #71775 those ran with the caller (``_available_entries``) holding the +lock. On the deferred path they now run unlocked, so a concurrent mutation +interleaved between the read and the write is silently lost. +""" + +import threading + +from agent.credential_pool import CredentialPool, PooledCredential + + +def _entry(entry_id: str, source: str) -> PooledCredential: + return PooledCredential( + id=entry_id, + provider="anthropic", + auth_type="oauth", + access_token="tok", + label=entry_id, + source=source, + priority=0, + ) + + +def _bare_pool(entries): + pool = CredentialPool.__new__(CredentialPool) + pool._lock = threading.RLock() + pool._entries = list(entries) + pool._active_leases = {} + pool._current_id = None + pool._max_concurrent = 2 + pool._unmatched_rotation_streak = 0 + pool.provider = "anthropic" + return pool + + +def test_quarantine_read_modify_write_is_atomic(): + """A concurrent mutation must not be lost across the quarantine filter. + + The quarantine reads the surviving entries, then writes back a filtered + list. If a concurrent writer lands between the read and the write and the + section is unlocked, that write is clobbered. Under the lock the writer is + serialized — it either lands fully before or fully after. + """ + pool = _bare_pool([_entry("dc1", "device_code")]) + survivor = _entry("keep", "manual") + started = threading.Event() + + def concurrent_add(): + started.set() + with pool._lock: # blocks while the quarantine holds the lock + pool._entries = pool._entries + [survivor] + + t = threading.Thread(target=concurrent_add) + + with pool._lock: + _removed = [i.id for i in pool._entries if i.source == "device_code"] + t.start() + started.wait(timeout=2) + # Give the writer a chance to (incorrectly) interleave. + t.join(timeout=0.2) + pool._entries = [i for i in pool._entries if i.source != "device_code"] + + # Outside the lock the writer can now proceed; wait for it to finish. + t.join(timeout=2) + assert not t.is_alive(), "concurrent writer did not complete" + + ids = {e.id for e in pool._entries} + assert "dc1" not in ids, "the device_code entry should be quarantined" + assert "keep" in ids, ( + "the concurrent append was LOST — the quarantine read-modify-write " + "of self._entries is not atomic" + ) + + +def test_quarantine_paths_hold_the_pool_lock(): + """Static guard: every bare ``self._entries = [`` inside + _refresh_entry_impl must sit under a ``with self._lock`` block. + + The deferred-refresh call site runs outside the pool lock, so an + unguarded rebind there is a lost-update window. + """ + import inspect + import textwrap + + src = textwrap.dedent(inspect.getsource(CredentialPool._refresh_entry_impl)) + lines = src.splitlines() + + unguarded = [] + for idx, line in enumerate(lines): + if "self._entries = [" not in line: + continue + indent = len(line) - len(line.lstrip()) + # Walk backwards for an enclosing `with self._lock` at lower indent. + guarded = False + for prev in range(idx - 1, -1, -1): + p = lines[prev] + if not p.strip(): + continue + p_indent = len(p) - len(p.lstrip()) + if p_indent < indent: + if "with self._lock" in p: + guarded = True + break + if p.lstrip().startswith("def "): + break + if not guarded: + unguarded.append(line.strip()) + + assert not unguarded, ( + "unguarded self._entries rebind(s) in _refresh_entry_impl — the " + f"deferred refresh path runs outside the pool lock: {unguarded}" + ) + + +def test_rlock_allows_locked_callers_to_reenter(): + """The already-locked callers must still work after adding the lock. + + self._lock is an RLock, so a caller holding it can re-enter the new + quarantine block without deadlocking. + """ + pool = _bare_pool([_entry("dc1", "device_code")]) + + with pool._lock: + acquired = pool._lock.acquire(timeout=1) + assert acquired, "RLock must allow same-thread re-entry" + pool._lock.release() diff --git a/tests/agent/test_curator.py b/tests/agent/test_curator.py index e18888771a..1650c97619 100644 --- a/tests/agent/test_curator.py +++ b/tests/agent/test_curator.py @@ -699,3 +699,80 @@ def test_review_fork_uses_runtime_model_and_output_cap(curator_env, monkeypatch) assert captured["max_tokens"] == 1234 + + +def test_review_fork_restricts_toolsets_to_skills_and_terminal(curator_env, monkeypatch): + """The curator LLM fork must advertise only the skills + terminal toolsets. + + Without ``enabled_toolsets=["skills", "terminal"]`` on the AIAgent(...) call + in ``_run_llm_review``, ``enabled_toolsets`` defaults to None and init_agent + grants the fork the full default catalog (~30 tools) plus the context_engine + (lcm_*) tools, billing ~7K wasted schema tokens on every one of the fork's + 50-100 API calls per consolidation pass. The prompt (curator.py:509-523) + confines the model to four tools in natural language, but only this kwarg + filters the advertised request schema. Capturing the constructor kwarg is + the sole assertion that distinguishes fixed from unfixed code. + """ + curator = curator_env["curator"] + + # curator_env stubs _run_llm_review wholesale; exercise the real + # implementation, so reload the module to restore it. + import importlib + importlib.reload(curator) + + captured = {} + + class _StubAgent: + def __init__(self, *args, **kwargs): + captured["enabled_toolsets"] = kwargs.get("enabled_toolsets", "UNSET") + self._memory_write_origin = "assistant_tool" + self._memory_nudge_interval = 0 + self._skill_nudge_interval = 0 + self._session_messages = [] + + def run_conversation(self, user_message=None, **kwargs): + return {"final_response": "no change"} + + def close(self): + pass + + monkeypatch.setattr("run_agent.AIAgent", _StubAgent) + + meta = curator._run_llm_review("review prompt") + + # error is None proves the fork was actually constructed (capture ran). + assert meta.get("error") is None, meta.get("error") + assert captured.get("enabled_toolsets") == ["skills", "terminal"], ( + "curator review fork did not pass enabled_toolsets=['skills', " + "'terminal'] to AIAgent; the full default tool catalog (plus lcm_* " + "context_engine tools) would be advertised; got " + f"{captured.get('enabled_toolsets')!r}" + ) + + +def test_review_fork_toolset_surface_is_skills_plus_terminal(): + """Documentary check on the static surface the fork's kwarg resolves to. + + Registry-independent (include_registry=False) so a plugin-registered tool + tagged into these toolsets cannot flake the membership checks. This + documents the intended surface (the four prompt-named tools present, dead + default and lcm_* schema absent) but does not itself guard the call-site + kwarg. No exact-set pin: intentional additions to either toolset must not + fail this test. + """ + from toolsets import resolve_toolset + + surface = set(resolve_toolset("skills", include_registry=False)) | set( + resolve_toolset("terminal", include_registry=False) + ) + + # The four prompt-named tools are all present. + assert "skills_list" in surface + assert "skill_view" in surface + assert "skill_manage" in surface + assert "terminal" in surface + + # Representative dropped default + context_engine tools are absent. + assert "read_file" not in surface + assert "web_search" not in surface + assert "lcm_grep" not in surface diff --git a/tests/agent/test_cursor_optimizations_parity.py b/tests/agent/test_cursor_optimizations_parity.py index c08b012aca..afb8e967bc 100644 --- a/tests/agent/test_cursor_optimizations_parity.py +++ b/tests/agent/test_cursor_optimizations_parity.py @@ -146,6 +146,12 @@ def test_parity_persist_bounded_scan(): self.rows = [] def append_message(self, **kw): self.rows.append({k: copy.deepcopy(v) for k, v in kw.items()}) + def append_messages_batch(self, session_id, messages, **kw): + for m in messages: + row = {k: copy.deepcopy(v) for k, v in m.items()} + row["session_id"] = session_id + self.rows.append(row) + return list(range(1, len(messages) + 1)) def make_agent(bounded): a = ra.AIAgent.__new__(ra.AIAgent) diff --git a/tests/agent/test_display.py b/tests/agent/test_display.py index dfb839d613..a736961f80 100644 --- a/tests/agent/test_display.py +++ b/tests/agent/test_display.py @@ -10,6 +10,7 @@ from agent.display import ( capture_local_edit_snapshot, extract_edit_diff, get_cute_tool_message, + prepare_tool_preview, redact_tool_args_for_display, set_tool_preview_max_len, _render_inline_unified_diff, @@ -106,6 +107,44 @@ class TestBuildToolPreview: assert build_tool_preview("terminal", []) is None +class TestPrepareToolPreview: + def test_recovers_and_describes_truncated_url(self): + url = "https://example.com/a/very/long/path/to/a/page" + set_tool_preview_max_len(20) + + preview = prepare_tool_preview( + "web_extract", + {"urls": [url]}, + fallback=url[:17] + "...", + max_len=20, + ) + + assert preview.text == url[:17] + "..." + assert preview.truncated is True + assert preview.url == url + + def test_untruncated_url_has_no_link_target(self): + url = "https://example.com/page" + preview = prepare_tool_preview( + "browser_navigate", None, fallback=url, max_len=40 + ) + + assert preview.text == url + assert preview.truncated is False + assert preview.url is None + + def test_truncated_non_url_has_no_link_target(self): + preview = prepare_tool_preview( + "web_search", + {"query": "how to parse a URL"}, + fallback="how to parse a URL", + max_len=12, + ) + + assert preview.truncated is True + assert preview.url is None + + class TestCuteToolMessagePreviewLength: @@ -275,4 +314,3 @@ class TestBuildStatusPhrase: assert build_status_phrase("terminal", {"command": "ls"}) is None finally: set_friendly_tool_labels(True) - diff --git a/tests/agent/test_endpoint_blackhole.py b/tests/agent/test_endpoint_blackhole.py new file mode 100644 index 0000000000..b3bdc93a08 --- /dev/null +++ b/tests/agent/test_endpoint_blackhole.py @@ -0,0 +1,291 @@ +"""Tests for short-circuiting probes to endpoints that blackhole TCP connects. + +A routable-but-dead endpoint (e.g. a corp LAN address while off-VPN) drops SYNs +without a RST or ICMP error, so each probe waits out its full timeout. Once one +probe has observed that, the rest must not repeat it. + +Covers: +- _endpoint_blackholed / _note_endpoint_blackholed host:port keying and TTL +- detect_local_server_type aborting its waterfall on the first connect timeout +- fetch_endpoint_model_metadata skipping its candidate loop once blackholed +- _query_ollama_api_show_uncached / _query_local_context_length_uncached + honouring and recording the blackhole +- non-timeout failures (refused, no route) leaving the waterfall untouched +""" + +from __future__ import annotations + +import os +import sys +from unittest.mock import MagicMock, patch + +import httpx +import pytest +import requests + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) + + +@pytest.fixture(autouse=True) +def _clear_caches(): + """Module-level caches must not leak between tests.""" + from agent import model_metadata + model_metadata._endpoint_blackhole_cache.clear() + model_metadata._endpoint_probe_path_cache.clear() + model_metadata._endpoint_model_metadata_cache.clear() + model_metadata._endpoint_model_metadata_cache_time.clear() + model_metadata._LOCAL_CTX_PROBE_CACHE.clear() + yield + model_metadata._endpoint_blackhole_cache.clear() + model_metadata._endpoint_probe_path_cache.clear() + model_metadata._endpoint_model_metadata_cache.clear() + model_metadata._endpoint_model_metadata_cache_time.clear() + model_metadata._LOCAL_CTX_PROBE_CACHE.clear() + + +def _client_mock(side_effect): + client = MagicMock() + client.__enter__ = lambda s: client + client.__exit__ = MagicMock(return_value=False) + client.get.side_effect = side_effect + client.post.side_effect = side_effect + return client + + +class TestBlackholeCache: + def test_unseen_endpoint_is_not_blackholed(self): + from agent.model_metadata import _endpoint_blackholed + + assert _endpoint_blackholed("http://10.0.0.9:30080/v1") is False + + def test_note_then_detected(self): + from agent.model_metadata import _endpoint_blackholed, _note_endpoint_blackholed + + _note_endpoint_blackholed("http://10.0.0.9:30080/v1") + assert _endpoint_blackholed("http://10.0.0.9:30080/v1") is True + + def test_keyed_on_host_port_not_path(self): + """Every probe path for one server shares a single entry.""" + from agent.model_metadata import _endpoint_blackholed, _note_endpoint_blackholed + + _note_endpoint_blackholed("http://10.0.0.9:30080") + assert _endpoint_blackholed("http://10.0.0.9:30080/v1") is True + assert _endpoint_blackholed("http://10.0.0.9:30080/api/v1") is True + + def test_different_port_is_independent(self): + from agent.model_metadata import _endpoint_blackholed, _note_endpoint_blackholed + + _note_endpoint_blackholed("http://10.0.0.9:30080/v1") + assert _endpoint_blackholed("http://10.0.0.9:11434/v1") is False + + def test_entry_expires_after_ttl(self): + """A recovered endpoint (VPN back up) is probed again without a restart.""" + from agent import model_metadata + from agent.model_metadata import _endpoint_blackholed, _note_endpoint_blackholed + + _note_endpoint_blackholed("http://10.0.0.9:30080/v1") + stale = ( + model_metadata._endpoint_blackhole_cache["10.0.0.9:30080"] + - model_metadata._ENDPOINT_BLACKHOLE_TTL_SECONDS + - 1 + ) + model_metadata._endpoint_blackhole_cache["10.0.0.9:30080"] = stale + assert _endpoint_blackholed("http://10.0.0.9:30080/v1") is False + + def test_ttl_zero_disables_short_circuit(self): + from agent import model_metadata + from agent.model_metadata import _endpoint_blackholed, _note_endpoint_blackholed + + _note_endpoint_blackholed("http://10.0.0.9:30080/v1") + with patch.object(model_metadata, "_ENDPOINT_BLACKHOLE_TTL_SECONDS", 0.0): + assert _endpoint_blackholed("http://10.0.0.9:30080/v1") is False + + +class TestDetectLocalServerTypeBlackhole: + URL = "http://10.0.0.9:30080/v1" + + def test_connect_timeout_aborts_waterfall_after_one_probe(self): + """Four sequential 2s probes against a dead host must collapse to one.""" + from agent.model_metadata import _endpoint_blackholed, detect_local_server_type + + client = _client_mock(httpx.ConnectTimeout("timed out")) + with patch("httpx.Client", return_value=client): + assert detect_local_server_type(self.URL) is None + + assert client.get.call_count == 1 + assert _endpoint_blackholed(self.URL) is True + + def test_second_call_makes_no_request_at_all(self): + from agent.model_metadata import detect_local_server_type + + client = _client_mock(httpx.ConnectTimeout("timed out")) + with patch("httpx.Client", return_value=client): + detect_local_server_type(self.URL) + first_count = client.get.call_count + assert detect_local_server_type(self.URL) is None + + assert client.get.call_count == first_count + + def test_refused_does_not_blackhole_and_runs_full_waterfall(self): + """Refused answers instantly, so skipping buys nothing and must not fire. + + This is the common "local server not started yet" path. + """ + from agent.model_metadata import _endpoint_blackholed, detect_local_server_type + + client = _client_mock(httpx.ConnectError("connection refused")) + with patch("httpx.Client", return_value=client): + assert detect_local_server_type(self.URL) is None + + assert client.get.call_count > 1 + assert _endpoint_blackholed(self.URL) is False + + def test_read_timeout_does_not_blackhole(self): + """A read timeout means the connection was accepted — not a blackhole.""" + from agent.model_metadata import _endpoint_blackholed, detect_local_server_type + + client = _client_mock(httpx.ReadTimeout("slow")) + with patch("httpx.Client", return_value=client): + detect_local_server_type(self.URL) + + assert _endpoint_blackholed(self.URL) is False + + +class TestFetchEndpointModelMetadataBlackhole: + URL = "http://10.0.0.9:30080/v1" + + def test_connect_timeout_skips_remaining_candidates(self): + """A timeout condemns the host, not the URL suffix — one stall, not two.""" + from agent.model_metadata import _endpoint_blackholed, fetch_endpoint_model_metadata + + with patch("agent.model_metadata.detect_local_server_type", return_value=None), \ + patch( + "agent.model_metadata.requests.get", + side_effect=requests.exceptions.ConnectTimeout("timed out"), + ) as get: + assert fetch_endpoint_model_metadata(self.URL) == {} + + assert get.call_count == 1 + assert _endpoint_blackholed(self.URL) is True + + def test_refused_tries_every_candidate_and_does_not_blackhole(self): + from agent.model_metadata import _endpoint_blackholed, fetch_endpoint_model_metadata + + with patch("agent.model_metadata.detect_local_server_type", return_value=None), \ + patch( + "agent.model_metadata.requests.get", + side_effect=requests.exceptions.ConnectionError("refused"), + ) as get: + assert fetch_endpoint_model_metadata(self.URL) == {} + + assert get.call_count == 2 # /v1-suffixed and bare candidates + assert _endpoint_blackholed(self.URL) is False + + def test_blackholed_endpoint_issues_no_request(self): + """force_refresh bypasses the metadata cache, so only the guard can stop it.""" + from agent.model_metadata import _note_endpoint_blackholed, fetch_endpoint_model_metadata + + _note_endpoint_blackholed(self.URL) + with patch("agent.model_metadata.detect_local_server_type", return_value=None), \ + patch("agent.model_metadata.requests.get") as get: + assert fetch_endpoint_model_metadata(self.URL, force_refresh=True) == {} + + get.assert_not_called() + + +class TestQueryOllamaApiShowBlackhole: + URL = "http://10.0.0.9:30080/v1" + + def test_connect_timeout_records_blackhole(self): + from agent.model_metadata import _endpoint_blackholed, _query_ollama_api_show_uncached + + client = _client_mock(httpx.ConnectTimeout("timed out")) + with patch("httpx.Client", return_value=client): + assert _query_ollama_api_show_uncached("some-model", self.URL) is None + + assert client.post.call_count == 1 + assert _endpoint_blackholed(self.URL) is True + + def test_blackholed_endpoint_issues_no_request(self): + from agent.model_metadata import _note_endpoint_blackholed, _query_ollama_api_show_uncached + + _note_endpoint_blackholed(self.URL) + with patch("httpx.Client") as client_cls: + assert _query_ollama_api_show_uncached("some-model", self.URL) is None + + client_cls.assert_not_called() + + def test_read_timeout_does_not_blackhole(self): + from agent.model_metadata import _endpoint_blackholed, _query_ollama_api_show_uncached + + client = _client_mock(httpx.ReadTimeout("slow")) + with patch("httpx.Client", return_value=client): + assert _query_ollama_api_show_uncached("some-model", self.URL) is None + + assert _endpoint_blackholed(self.URL) is False + + +class TestQueryLocalContextLengthBlackhole: + URL = "http://10.0.0.9:30080/v1" + + def test_connect_timeout_records_blackhole(self): + from agent.model_metadata import ( + _endpoint_blackholed, + _query_local_context_length_uncached, + ) + + client = _client_mock(httpx.ConnectTimeout("timed out")) + with patch("agent.model_metadata.detect_local_server_type", return_value=None), \ + patch("httpx.Client", return_value=client): + assert _query_local_context_length_uncached("some-model", self.URL) is None + + assert _endpoint_blackholed(self.URL) is True + + def test_blackholed_endpoint_skips_detection_and_requests(self): + """The guard sits before detect_local_server_type — nothing runs at all.""" + from agent.model_metadata import ( + _note_endpoint_blackholed, + _query_local_context_length_uncached, + ) + + _note_endpoint_blackholed(self.URL) + with patch("agent.model_metadata.detect_local_server_type") as detect, \ + patch("httpx.Client") as client_cls: + assert _query_local_context_length_uncached("some-model", self.URL) is None + + detect.assert_not_called() + client_cls.assert_not_called() + + def test_read_timeout_does_not_blackhole(self): + from agent.model_metadata import ( + _endpoint_blackholed, + _query_local_context_length_uncached, + ) + + client = _client_mock(httpx.ReadTimeout("slow")) + with patch("agent.model_metadata.detect_local_server_type", return_value=None), \ + patch("httpx.Client", return_value=client): + assert _query_local_context_length_uncached("some-model", self.URL) is None + + assert _endpoint_blackholed(self.URL) is False + + +class TestIsConnectTimeout: + def test_httpx_connect_timeout(self): + from agent.model_metadata import _is_connect_timeout + + assert _is_connect_timeout(httpx.ConnectTimeout("x")) is True + + def test_requests_connect_timeout(self): + from requests.exceptions import ConnectTimeout + + from agent.model_metadata import _is_connect_timeout + + assert _is_connect_timeout(ConnectTimeout("x")) is True + + def test_unrelated_errors_are_not_connect_timeouts(self): + from agent.model_metadata import _is_connect_timeout + + assert _is_connect_timeout(httpx.ReadTimeout("x")) is False + assert _is_connect_timeout(httpx.ConnectError("x")) is False + assert _is_connect_timeout(ValueError("x")) is False diff --git a/tests/agent/test_insights.py b/tests/agent/test_insights.py index f0f06e1753..2d42ab62e1 100644 --- a/tests/agent/test_insights.py +++ b/tests/agent/test_insights.py @@ -1,5 +1,6 @@ """Tests for agent/insights.py — InsightsEngine analytics and reporting.""" +import sqlite3 import time import pytest @@ -368,6 +369,160 @@ class TestInsightsPopulated: + # The Insights assistant tool-call queries pin + # idx_messages_assistant_calls_by_session via INDEXED BY. These tests prove + # (a) the planner uses that index for BOTH the unfiltered and source-filtered + # branches on a fresh DB *without* ANALYZE, and (b) the index is a pure + # optimization — output is identical whether or not it is selected. + _INDEX = "idx_messages_assistant_calls_by_session" + _PINNED_QUERIES = ( + ("_GET_TOOL_CALLS_ALL", (0.0,)), + ("_GET_TOOL_CALLS_WITH_SOURCE", (0.0, "cli")), + ("_GET_SKILL_CALLS_ALL", (0.0,)), + ("_GET_SKILL_CALLS_WITH_SOURCE", (0.0, "cli")), + ) + + def test_assistant_call_queries_use_partial_index_without_analyze( + self, populated_db + ): + """Every fixed-predicate branch selects the partial index on a fresh DB. + + No ANALYZE is run, so this covers the default-statistics case a freshly + initialized state.db is actually in. Both the unfiltered and the + source-filtered (``s.source = ?``) branches are checked. + """ + # Guard against the fresh-DB planner regression the reviewers found: + # without INDEXED BY the source-filtered branch fell back to + # idx_messages_session_active. + assert "ANALYZE" not in "".join( + r["sql"] or "" + for r in populated_db._conn.execute( + "SELECT sql FROM sqlite_master WHERE type = 'index'" + ) + ) + for attr, params in self._PINNED_QUERIES: + sql = getattr(InsightsEngine, attr) + plan = "\n".join( + row["detail"] + for row in populated_db._conn.execute( + "EXPLAIN QUERY PLAN " + sql, params + ).fetchall() + ) + assert self._INDEX in plan, f"{attr} did not use the index:\n{plan}" + + def test_assistant_call_rows_invariant_to_index_selection(self, populated_db): + """The pinned index only changes the plan, never the result set. + + For every branch, the index-pinned query and the un-pinned form (whose + plan the optimizer chooses freely) must return identical rows — proving + the index is a pure optimization — for both the unfiltered and + source-filtered scopes. + """ + assert populated_db._conn.execute( + "SELECT 1 FROM sqlite_master WHERE type = 'index' AND name = ?", + (self._INDEX,), + ).fetchone() is not None + + for attr, params in self._PINNED_QUERIES: + pinned_sql = getattr(InsightsEngine, attr) + unpinned_sql = pinned_sql.replace(f" INDEXED BY {self._INDEX}", "") + pinned = [ + tuple(r) for r in + populated_db._conn.execute(pinned_sql, params).fetchall() + ] + unpinned = [ + tuple(r) for r in + populated_db._conn.execute(unpinned_sql, params).fetchall() + ] + assert sorted(pinned) == sorted(unpinned), attr + + def test_tool_and_skill_usage_invariant_to_partial_index(self, populated_db): + """The public tool/skill usage output is stable and exercises the + assistant tool_calls path for both scopes.""" + engine = InsightsEngine(populated_db) + cutoff = 0.0 + + tools = engine._get_tool_usage(cutoff) + tools_cli = engine._get_tool_usage(cutoff, source="cli") + skills = engine._get_skill_usage(cutoff) + skills_cli = engine._get_skill_usage(cutoff, source="cli") + + # Sanity: the fixture actually drives the assistant tool_calls path. + assert any(t["tool_name"] == "search_files" for t in tools) + assert any(t["tool_name"] == "search_files" for t in tools_cli) + assert isinstance(skills, list) and isinstance(skills_cli, list) + + def test_missing_index_falls_back_to_unpinned_queries(self, populated_db): + """INDEXED BY would be a hard error if the index is missing — which + happens on read-only opens of a state.db written by an older version + (web dashboard analytics). The engine must probe and fall back to the + unpinned variants instead of crashing, returning identical rows.""" + engine_pinned = InsightsEngine(populated_db) + tools_before = engine_pinned._get_tool_usage(0.0) + + populated_db._conn.execute(f"DROP INDEX IF EXISTS {self._INDEX}") + populated_db._conn.commit() + + engine = InsightsEngine(populated_db) + assert engine._has_assistant_calls_index is False + assert "INDEXED BY" not in engine._GET_TOOL_CALLS_ALL + tools_after = engine._get_tool_usage(0.0) + assert sorted(t["tool_name"] for t in tools_after) == sorted( + t["tool_name"] for t in tools_before + ) + # And with the index present, the pin stays. + assert "INDEXED BY" in InsightsEngine._GET_TOOL_CALLS_ALL + + def test_get_skill_breakdown_matches_full_generate(self, populated_db): + engine = InsightsEngine(populated_db) + full = engine.generate(days=30) + focused = engine.get_usage_breakdown(days=30)["skills"] + assert focused == full["skills"] + + def test_get_usage_breakdown_matches_full_generate(self, populated_db): + engine = InsightsEngine(populated_db) + full = engine.generate(days=30) + focused = engine.get_usage_breakdown(days=30) + assert focused["skills"] == full["skills"] + assert focused["tools"] == full["tools"] + + def test_get_skill_breakdown_respects_source_filter(self, populated_db): + engine = InsightsEngine(populated_db) + # Only s1 (cli) has skill_view "github-pr-workflow" + focused = engine.get_usage_breakdown(days=30, source="cli")["skills"] + skill_names = [s["skill"] for s in focused["top_skills"]] + assert "github-pr-workflow" in skill_names + # github-code-review was in discord (s4), not cli + assert "github-code-review" not in skill_names + + def test_get_skill_breakdown_empty_db(self, db): + focused = InsightsEngine(db).get_usage_breakdown(days=30)["skills"] + assert focused == { + "summary": { + "total_skill_loads": 0, + "total_skill_edits": 0, + "total_skill_actions": 0, + "distinct_skills_used": 0, + }, + "top_skills": [], + } + + def test_get_skill_usage_prefilter_ignores_non_skill_substring(self, db): + # "my_skill_view_helper" contains "skill_view" as a substring; instr() + # will match but the Python-side name check keeps the set clean. + # More importantly, messages with no skill_* tools must be excluded. + db.create_session(session_id="sx", source="cli", model="gpt-4o") + db.append_message( + "sx", + role="assistant", + content="Just using read_file.", + tool_calls=[{"function": {"name": "read_file", "arguments": '{"path":"/tmp/x"}'}}], + ) + db._conn.commit() + focused = InsightsEngine(db).get_usage_breakdown(days=30)["skills"] + assert focused["summary"]["total_skill_actions"] == 0 + assert focused["top_skills"] == [] + # ========================================================================= # Formatting diff --git a/tests/agent/test_memory_provider.py b/tests/agent/test_memory_provider.py index 8b16dd442d..d11cfc026d 100644 --- a/tests/agent/test_memory_provider.py +++ b/tests/agent/test_memory_provider.py @@ -1138,3 +1138,23 @@ class TestMemoryInjectionRejectsMalformedSchema: names = {t["function"]["name"] for t in agent.tools} assert names == {"good_tool"} assert agent.valid_tool_names == {"good_tool"} + + +class TestTrivialPromptClassifier: + """is_trivial_prompt — the shared gate for core prefetch + provider injection.""" + + def test_trivial_variants(self): + from agent.memory_provider import is_trivial_prompt + + for t in ("hi", "HI!", "hey.", "hello", "yo", "sup~", "thanks :)", + "done???", "ok", "yes.", "k", "", " ", "/help", "lgtm"): + assert is_trivial_prompt(t), f"expected trivial: {t!r}" + + def test_substantive_and_prefix_collisions_pass_through(self): + from agent.memory_provider import is_trivial_prompt + + # Words that merely START with a trivial word must not match. + for t in ("k8s", "yolo", "hive", "note", "supper", "hind", + "hello world", "ok so what's next", "what's my name", + "hey can you check the logs", "continue the migration plan"): + assert not is_trivial_prompt(t), f"expected non-trivial: {t!r}" diff --git a/tests/agent/test_moa_cold_start_cache_66793.py b/tests/agent/test_moa_cold_start_cache_66793.py new file mode 100644 index 0000000000..96e5a2bbc7 --- /dev/null +++ b/tests/agent/test_moa_cold_start_cache_66793.py @@ -0,0 +1,226 @@ +"""Regression tests for MoA cold-start caching (#66793). + +The preset switch used to re-parse + re-validate the full config and +re-resolve every slot's provider runtime on EACH create() call (once +per tool-loop iteration), serially before the parallel fan-out could +begin — 5-30s of "frozen" latency on complex presets. +Both the resolved preset and each (provider, model) runtime are now +cached for the process lifetime (config is immutable per turn), so the +underlying ``resolve_runtime_provider`` (real provider-catalog I/O) +runs once per distinct slot, not once per create() iteration. +""" + +import types # noqa: F401 (used by _fake_response) + +import pytest + + +def _make_preset_config() -> dict: + return { + "moa": { + "default_preset": "demo", + "presets": { + "demo": { + "enabled": True, + "aggregator": {"provider": "openai", "model": "gpt-5"}, + "reference_models": [ + {"provider": "deepseek", "model": "deepseek-v4"}, + {"provider": "minimax", "model": "minimax-m3"}, + ], + } + }, + } + } + + +def test_preset_resolution_is_cached_across_create_calls(monkeypatch, tmp_path): + """resolve_moa_preset must run once per (config-mtime, preset_name), + not on every create() iteration.""" + import agent.moa_loop as moa + + moa._preset_cache.clear() + + calls = {"n": 0} + import hermes_cli.moa_config as moa_cfg_mod + real_resolve = moa_cfg_mod.resolve_moa_preset + + def counting_resolve(config, name=None): + calls["n"] += 1 + return real_resolve(config, name) + + monkeypatch.setattr(moa_cfg_mod, "resolve_moa_preset", counting_resolve) + import hermes_cli.config as cfg_mod + # The cache keys on the config FILE's st_mtime_ns — give the test a real + # stat-able file (no config file -> stamp=None -> caching fails open). + cfg_file = tmp_path / "config.yaml" + cfg_file.write_text("moa: {}\n") + monkeypatch.setattr(cfg_mod, "get_config_path", lambda: cfg_file) + monkeypatch.setattr(cfg_mod, "load_config", lambda: _make_preset_config()) + monkeypatch.setattr(moa, "call_llm", lambda **k: _fake_response()) + + cc = moa.MoAChatCompletions("demo") + for _ in range(3): + cc.create(messages=[{"role": "user", "content": "hi"}]) + + # One preset resolution for the whole turn (not 3). + assert calls["n"] == 1, f"expected 1 preset resolution, got {calls['n']}" + + +def test_preset_cache_invalidates_on_config_edit(monkeypatch, tmp_path): + """Editing config.yaml must invalidate the preset cache on the next + create() — the original PR keyed on a nonexistent config-object mtime + attribute, which never invalidated (review finding).""" + import os + + import agent.moa_loop as moa + + moa._preset_cache.clear() + + calls = {"n": 0} + import hermes_cli.moa_config as moa_cfg_mod + real_resolve = moa_cfg_mod.resolve_moa_preset + + def counting_resolve(config, name=None): + calls["n"] += 1 + return real_resolve(config, name) + + monkeypatch.setattr(moa_cfg_mod, "resolve_moa_preset", counting_resolve) + import hermes_cli.config as cfg_mod + cfg_file = tmp_path / "config.yaml" + cfg_file.write_text("moa: {}\n") + monkeypatch.setattr(cfg_mod, "get_config_path", lambda: cfg_file) + monkeypatch.setattr(cfg_mod, "load_config", lambda: _make_preset_config()) + monkeypatch.setattr(moa, "call_llm", lambda **k: _fake_response()) + + cc = moa.MoAChatCompletions("demo") + cc.create(messages=[{"role": "user", "content": "hi"}]) + assert calls["n"] == 1 + + # Simulate a config edit: bump the file's mtime past ns resolution. + st = cfg_file.stat() + os.utime(cfg_file, ns=(st.st_atime_ns, st.st_mtime_ns + 1_000_000)) + + cc.create(messages=[{"role": "user", "content": "hi"}]) + assert calls["n"] == 2, "config edit must invalidate the preset cache" + + +def test_no_config_file_fails_open(monkeypatch, tmp_path): + """No config.yaml (stat raises) -> caching disabled, create() still works.""" + import agent.moa_loop as moa + + moa._preset_cache.clear() + + import hermes_cli.config as cfg_mod + monkeypatch.setattr( + cfg_mod, "get_config_path", lambda: tmp_path / "missing.yaml" + ) + monkeypatch.setattr(cfg_mod, "load_config", lambda: _make_preset_config()) + monkeypatch.setattr(moa, "call_llm", lambda **k: _fake_response()) + + cc = moa.MoAChatCompletions("demo") + cc.create(messages=[{"role": "user", "content": "hi"}]) + assert moa._preset_cache == {}, "must not cache under a None stamp" + + +def test_slot_runtime_is_cached_across_create_calls(monkeypatch, tmp_path): + """resolve_runtime_provider (real I/O) must run once per + (provider, model) across all create() iterations, not per call.""" + import agent.moa_loop as moa + + moa._runtime_cache.clear() + moa._preset_cache.clear() + + calls = {"n": 0} + + def counting_resolve(*a, **k): + calls["n"] += 1 + return {"base_url": None, "api_key": None, "api_mode": None} + + import hermes_cli.runtime_provider as rt_mod + monkeypatch.setattr(rt_mod, "resolve_runtime_provider", counting_resolve) + import hermes_cli.config as cfg_mod + cfg_file = tmp_path / "config.yaml" + cfg_file.write_text("moa: {}\n") + monkeypatch.setattr(cfg_mod, "get_config_path", lambda: cfg_file) + monkeypatch.setattr(cfg_mod, "load_config", lambda: _make_preset_config()) + monkeypatch.setattr(moa, "call_llm", lambda **k: _fake_response()) + + cc = moa.MoAChatCompletions("demo") + for _ in range(2): + cc.create(messages=[{"role": "user", "content": "hi"}]) + + # aggregator(1) + 2 references = 3 distinct slots, resolved once + # each regardless of 2 create() iterations. + assert calls["n"] == 3, f"expected 3 slot resolutions, got {calls['n']}" + + +def test_slot_runtime_cache_expires_after_ttl(monkeypatch): + """A stale runtime entry (key rotation window) must re-resolve after + the TTL — the original PR cached for the process lifetime, pinning + rotated credentials forever (review finding).""" + import agent.moa_loop as moa + + moa._runtime_cache.clear() + + calls = {"n": 0} + + def counting_resolve(*a, **k): + calls["n"] += 1 + return {"base_url": "http://x", "api_key": f"key-{calls['n']}", + "api_mode": None} + + import hermes_cli.runtime_provider as rt_mod + monkeypatch.setattr(rt_mod, "resolve_runtime_provider", counting_resolve) + + slot = {"provider": "openai", "model": "gpt-5"} + first = moa._slot_runtime(slot) + assert calls["n"] == 1 and first["api_key"] == "key-1" + + # Within TTL: cached. + assert moa._slot_runtime(slot)["api_key"] == "key-1" + assert calls["n"] == 1 + + # Age the entry past the TTL and confirm re-resolution. + key = ("openai", "gpt-5") + stamped_at, cached = moa._runtime_cache[key] + moa._runtime_cache[key] = ( + stamped_at - moa._RUNTIME_CACHE_TTL_SECONDS - 1, cached + ) + assert moa._slot_runtime(slot)["api_key"] == "key-2" + assert calls["n"] == 2 + + +def test_slot_runtime_resolution_error_is_not_cached(monkeypatch): + """A transient resolution failure must not pin the bare-kwargs fallback + for a TTL — the next call must retry the real resolver.""" + import agent.moa_loop as moa + + moa._runtime_cache.clear() + + calls = {"n": 0} + + def flaky_resolve(*a, **k): + calls["n"] += 1 + if calls["n"] == 1: + raise RuntimeError("catalog hiccup") + return {"base_url": "http://ok", "api_key": None, "api_mode": None} + + import hermes_cli.runtime_provider as rt_mod + monkeypatch.setattr(rt_mod, "resolve_runtime_provider", flaky_resolve) + + slot = {"provider": "openai", "model": "gpt-5"} + fallback = moa._slot_runtime(slot) + assert "base_url" not in fallback # bare kwargs on error + assert moa._runtime_cache == {}, "error result must not be cached" + + recovered = moa._slot_runtime(slot) + assert recovered.get("base_url") == "http://ok" + assert calls["n"] == 2 + + +# ─── test harness helpers ────────────────────────────────────────────── + +def _fake_response(): + ns = types.SimpleNamespace() + ns.usage = None + return ns diff --git a/tests/agent/test_model_metadata_local_ctx.py b/tests/agent/test_model_metadata_local_ctx.py index d99bef07c0..c428b5f857 100644 --- a/tests/agent/test_model_metadata_local_ctx.py +++ b/tests/agent/test_model_metadata_local_ctx.py @@ -172,12 +172,17 @@ class TestQueryLocalContextLengthModelsList: assert result == 131072 def test_models_list_model_not_found_returns_none(self): - """Returns None when model is not in the /v1/models list.""" + """Returns None when the model is absent from a multi-model /v1/models + list. (Single-model servers are accepted even when the configured name + doesn't match the reported id — see the llama.cpp tests below.)""" from agent.model_metadata import _query_local_context_length detail_resp = self._make_resp(404, {}) list_resp = self._make_resp(200, { - "data": [{"id": "other-model", "max_model_len": 4096}] + "data": [ + {"id": "other-model", "max_model_len": 4096}, + {"id": "yet-another-model", "max_model_len": 8192}, + ] }) call_count = [0] @@ -199,6 +204,73 @@ class TestQueryLocalContextLengthModelsList: assert result is None + def test_models_list_llamacpp_meta_n_ctx_sole_model(self): + """llama.cpp nests the runtime context under meta.n_ctx and serves a + single model whose id (a GGUF path) doesn't match the configured name. + + The sole model should be accepted and meta.n_ctx read, instead of + returning None and falling back to a family default (e.g. qwen=131072). + """ + from agent.model_metadata import _query_local_context_length + + detail_resp = self._make_resp(404, {}) + list_resp = self._make_resp(200, { + "data": [ + { + "id": "/app/models/qwen3.6-35b.gguf", + "meta": {"n_ctx": 256000, "n_ctx_train": 262144}, + } + ] + }) + + call_count = [0] + def side_effect(url, **kwargs): + call_count[0] += 1 + if call_count[0] == 1: + return detail_resp # /v1/models/{model} + return list_resp # /v1/models + + client_mock = MagicMock() + client_mock.__enter__ = lambda s: client_mock + client_mock.__exit__ = MagicMock(return_value=False) + client_mock.post.return_value = self._make_resp(404, {}) + client_mock.get.side_effect = side_effect + + with patch("agent.model_metadata.detect_local_server_type", return_value=None), \ + patch("httpx.Client", return_value=client_mock): + result = _query_local_context_length("qwen3.6-35b", "http://localhost:8080") + + assert result == 256000 + + def test_models_list_llamacpp_prefers_runtime_n_ctx_over_train(self): + """Runtime n_ctx (256000) is preferred over n_ctx_train (262144), + since the server can only actually serve the runtime value.""" + from agent.model_metadata import _query_local_context_length + + detail_resp = self._make_resp(404, {}) + list_resp = self._make_resp(200, { + "data": [ + {"id": "/app/models/m.gguf", "meta": {"n_ctx": 256000, "n_ctx_train": 262144}} + ] + }) + + call_count = [0] + def side_effect(url, **kwargs): + call_count[0] += 1 + return detail_resp if call_count[0] == 1 else list_resp + + client_mock = MagicMock() + client_mock.__enter__ = lambda s: client_mock + client_mock.__exit__ = MagicMock(return_value=False) + client_mock.post.return_value = self._make_resp(404, {}) + client_mock.get.side_effect = side_effect + + with patch("agent.model_metadata.detect_local_server_type", return_value=None), \ + patch("httpx.Client", return_value=client_mock): + result = _query_local_context_length("m", "http://localhost:8080") + + assert result == 256000 + class TestQueryLocalContextLengthLmStudio: """_query_local_context_length with LM Studio native /api/v1/models response.""" diff --git a/tests/agent/test_post_compression_trim.py b/tests/agent/test_post_compression_trim.py new file mode 100644 index 0000000000..4cb1aadc39 --- /dev/null +++ b/tests/agent/test_post_compression_trim.py @@ -0,0 +1,66 @@ +"""A successful compaction hands allocator pages back to the OS. + +The compressed-away message dicts are the largest allocation a long session +ever frees, but Python's arena allocator keeps those pages in the process heap +— RSS retains the pre-compaction high-water mark until exit. #76905's +trim_memory lifecycle covers the gateway/TUI housekeeping loops but not the +CLI compression path, so compress() now calls +``trim_memory(reason="post-compression")`` after a successful pass. + +trim_memory itself is glibc/Linux-gated (a fast no-op on macOS), so these +tests monkeypatch the seam rather than asserting on RSS. Salvaged in spirit +from #70782 (which reached for a bare gc.collect(); trim_memory is the +house mechanism and already wraps a collect). +""" +import hermes_cli.mem_trim as mem_trim +from agent.context_compressor import ContextCompressor + + +def _compressor(threshold_tokens: int = 24_576) -> ContextCompressor: + cc = ContextCompressor( + model="test-model", + threshold_percent=0.75, + protect_first_n=5, + protect_last_n=20, + quiet_mode=True, + config_context_length=40960, + provider="test", + ) + cc.threshold_tokens = threshold_tokens # pin; don't couple to window math + cc._generate_summary = lambda *a, **k: "Summary of earlier turns." + return cc + + +def _messages(n: int, size: int = 1500) -> list: + msgs = [{"role": "system", "content": "sys"}] + for i in range(n): + role = "user" if i % 2 == 0 else "assistant" + msgs.append({"role": role, "content": f"m{i} " + "z" * size}) + return msgs + + +def test_successful_compression_trims_memory_once(monkeypatch): + calls = [] + monkeypatch.setattr( + mem_trim, "trim_memory", lambda *a, **kw: calls.append(kw) or False + ) + + cc = _compressor() + out = cc.compress(_messages(14), current_tokens=100_000) + + assert len(out) < 15, "sanity: compaction should have made progress" + assert len(calls) == 1, "trim_memory must run exactly once per compaction" + assert calls[0].get("reason") == "post-compression" + + +def test_trim_failure_does_not_break_compression(monkeypatch): + def boom(*a, **kw): + raise RuntimeError("allocator says no") + + monkeypatch.setattr(mem_trim, "trim_memory", boom) + + cc = _compressor() + out = cc.compress(_messages(14), current_tokens=100_000) + + assert cc._last_compression_made_progress is True + assert isinstance(out, list) and out, "compress() must still return messages" diff --git a/tests/agent/test_probe_cache_followups.py b/tests/agent/test_probe_cache_followups.py index 27d306ae2d..f0a88a64e1 100644 --- a/tests/agent/test_probe_cache_followups.py +++ b/tests/agent/test_probe_cache_followups.py @@ -186,6 +186,61 @@ class TestLocalhostIPv4SiblingSites: assert client.post.call_args[0][0].startswith("http://127.0.0.1:11434") + def test_fetch_endpoint_model_metadata_generic_probe_uses_ipv4(self): + """The generic (non-LM-Studio) /models fetch loop must also rewrite + localhost->127.0.0.1 before probing, like the LM Studio branch above.""" + from agent import model_metadata + from agent.model_metadata import fetch_endpoint_model_metadata + + model_metadata._endpoint_model_metadata_cache.clear() + model_metadata._endpoint_model_metadata_cache_time.clear() + + resp = MagicMock() + resp.status_code = 200 + resp.raise_for_status = MagicMock() + resp.json.return_value = {"data": []} + + with patch("agent.model_metadata.detect_local_server_type", return_value=None), \ + patch("agent.model_metadata.requests.get", return_value=resp) as mock_get: + fetch_endpoint_model_metadata("http://localhost:8000/v1") + + assert mock_get.call_args[0][0].startswith("http://127.0.0.1:8000") + + def test_fetch_endpoint_model_metadata_llamacpp_props_followup_uses_ipv4(self): + """The llama.cpp /props context-length follow-up must also rewrite + localhost->127.0.0.1 before probing, not just the initial /models call.""" + from agent import model_metadata + from agent.model_metadata import fetch_endpoint_model_metadata + + model_metadata._endpoint_model_metadata_cache.clear() + model_metadata._endpoint_model_metadata_cache_time.clear() + + models_resp = MagicMock() + models_resp.status_code = 200 + models_resp.raise_for_status = MagicMock() + models_resp.json.return_value = { + "data": [{"id": "llama-3-8b", "owned_by": "llamacpp"}], + } + + props_resp = MagicMock() + props_resp.ok = True + props_resp.json.return_value = { + "default_generation_settings": {"n_ctx": 32768}, + "model_alias": "llama-3-8b", + } + + with patch("agent.model_metadata.detect_local_server_type", return_value=None), \ + patch( + "agent.model_metadata.requests.get", + side_effect=[models_resp, props_resp], + ) as mock_get: + result = fetch_endpoint_model_metadata("http://localhost:8000/v1") + + assert mock_get.call_count == 2 + props_call_url = mock_get.call_args_list[1][0][0] + assert props_call_url.startswith("http://127.0.0.1:8000") + assert result["llama-3-8b"]["context_length"] == 32768 + class TestContextCacheKeyNormalization: diff --git a/tests/agent/test_subdirectory_hints.py b/tests/agent/test_subdirectory_hints.py index 16c2b5b725..01de13d4ea 100644 --- a/tests/agent/test_subdirectory_hints.py +++ b/tests/agent/test_subdirectory_hints.py @@ -176,3 +176,111 @@ class TestOutsideWorkspaceRejection: outside.mkdir(exist_ok=True) tracker = SubdirectoryHintTracker(working_dir=str(project)) assert tracker._is_valid_subdir(outside) is False + + +class TestContentDeduplication: + """The same context content must never be injected twice (ref: symlinked + shared workspaces, hardlinks, and copied backups all alias one file).""" + + def test_symlinked_duplicate_not_reinjected(self, tmp_path): + """Two directories whose AGENTS.md is the same file yield one injection.""" + real = tmp_path / "real" + real.mkdir() + (real / "AGENTS.md").write_text("Shared workspace instructions") + + mirror = tmp_path / "mirror" + mirror.mkdir() + (mirror / "AGENTS.md").symlink_to(real / "AGENTS.md") + + tracker = SubdirectoryHintTracker(working_dir=str(tmp_path)) + first = tracker.check_tool_call("read_file", {"path": str(real / "x.py")}) + second = tracker.check_tool_call("read_file", {"path": str(mirror / "y.py")}) + + assert first is not None + assert "Shared workspace instructions" in first + assert second is None + + def test_identical_copy_not_reinjected(self, tmp_path): + """Byte-identical copies in unrelated directories dedupe by digest.""" + a = tmp_path / "a" + b = tmp_path / "b" + a.mkdir() + b.mkdir() + (a / "AGENTS.md").write_text("Same content") + (b / "AGENTS.md").write_text("Same content") + + tracker = SubdirectoryHintTracker(working_dir=str(tmp_path)) + assert tracker.check_tool_call("read_file", {"path": str(a / "f.py")}) is not None + assert tracker.check_tool_call("read_file", {"path": str(b / "f.py")}) is None + + def test_differing_content_still_injected(self, tmp_path): + """Dedupe must not suppress genuinely different context.""" + a = tmp_path / "a" + b = tmp_path / "b" + a.mkdir() + b.mkdir() + (a / "AGENTS.md").write_text("Alpha rules") + (b / "AGENTS.md").write_text("Beta rules") + + tracker = SubdirectoryHintTracker(working_dir=str(tmp_path)) + first = tracker.check_tool_call("read_file", {"path": str(a / "f.py")}) + second = tracker.check_tool_call("read_file", {"path": str(b / "f.py")}) + + assert first is not None and "Alpha rules" in first + assert second is not None and "Beta rules" in second + + def test_working_dir_content_seeded(self, tmp_path): + """A copy of the CWD's own context file is not re-injected.""" + (tmp_path / "AGENTS.md").write_text("Root instructions") + elsewhere = tmp_path / "elsewhere" + elsewhere.mkdir() + (elsewhere / "AGENTS.md").write_text("Root instructions") + + tracker = SubdirectoryHintTracker(working_dir=str(tmp_path)) + assert tracker.check_tool_call("read_file", {"path": str(elsewhere / "f.py")}) is None + + +class TestExcludedDirectories: + """Backups, vendored deps, and caches hold copies — never context.""" + + @pytest.mark.parametrize( + "excluded", + ["backups", "node_modules", ".git", "venv", "site-packages", ".Trash", "vendor"], + ) + def test_excluded_directory_skipped(self, tmp_path, excluded): + target = tmp_path / excluded / "snapshot" + target.mkdir(parents=True) + (target / "AGENTS.md").write_text("Stale archived instructions") + + tracker = SubdirectoryHintTracker(working_dir=str(tmp_path)) + assert tracker.check_tool_call("read_file", {"path": str(target / "f.py")}) is None + + def test_excluded_ancestor_blocks_descendant(self, tmp_path): + """A hint nested under an excluded ancestor is still skipped.""" + deep = tmp_path / "backups" / "2026" / "proj" + deep.mkdir(parents=True) + (deep / "AGENTS.md").write_text("Archived") + + tracker = SubdirectoryHintTracker(working_dir=str(tmp_path)) + assert tracker.check_tool_call("read_file", {"path": str(deep / "f.py")}) is None + + def test_working_dir_inside_excluded_name_still_works(self, tmp_path): + """If the user works inside e.g. vendor/, its own subdirs stay eligible.""" + root = tmp_path / "vendor" / "myproject" + root.mkdir(parents=True) + sub = root / "pkg" + sub.mkdir() + (sub / "AGENTS.md").write_text("Package rules") + + tracker = SubdirectoryHintTracker(working_dir=str(root)) + result = tracker.check_tool_call("read_file", {"path": str(sub / "f.py")}) + assert result is not None and "Package rules" in result + + def test_normal_directory_unaffected(self, tmp_path): + normal = tmp_path / "backend" + normal.mkdir() + (normal / "AGENTS.md").write_text("Backend rules") + + tracker = SubdirectoryHintTracker(working_dir=str(tmp_path)) + result = tracker.check_tool_call("read_file", {"path": str(normal / "f.py")}) + assert result is not None and "Backend rules" in result diff --git a/tests/agent/test_system_prompt.py b/tests/agent/test_system_prompt.py index f37716a72e..2f0a427176 100644 --- a/tests/agent/test_system_prompt.py +++ b/tests/agent/test_system_prompt.py @@ -226,3 +226,43 @@ class TestTelegramRichMessagesHint: stable = _stable_prompt(agent) assert "Standard Markdown is automatically converted" in stable assert "lean into it" not in stable + + +_SKILLS = "SKILLS_INDEX_SENTINEL" +_CONTEXT = "CONTEXT_FILES_SENTINEL" + + +def _build(builder, **overrides): + """Run a build_* function with skills + context files present.""" + agent = _make_agent(valid_tool_names=["skills_list"], **overrides) + with ( + patch("run_agent.load_soul_md", return_value=""), + patch("run_agent.build_nous_subscription_prompt", return_value=""), + patch("run_agent.build_environment_hints", return_value=""), + patch("run_agent.build_context_files_prompt", return_value=_CONTEXT), + patch("run_agent.get_toolset_for_tool", return_value=None), + patch("run_agent.build_skills_system_prompt", return_value=_SKILLS), + ): + return builder(agent) + + +class TestSkillsInVolatileBand: + """The skills index is runtime-mutable, so it lives in the volatile band, + not the stable band, to keep the cached stable prefix reusable when a + rebuild picks up a skill change.""" + + def test_skills_not_in_stable_band(self): + parts = _build(build_system_prompt_parts) + assert _SKILLS not in parts["stable"] + + def test_skills_lead_the_volatile_band(self): + parts = _build(build_system_prompt_parts) + assert parts["volatile"].startswith(_SKILLS) + + def test_full_order_is_stable_context_then_skills(self): + # build_system_prompt joins stable + context + volatile, so the skills + # index renders after the context files and before the per-turn + # memory/timestamp tail. + full = _build(build_system_prompt) + assert full.index(_CONTEXT) < full.index(_SKILLS) + assert full.index(_SKILLS) < full.index("Conversation started:") diff --git a/tests/agent/test_turn_context.py b/tests/agent/test_turn_context.py index 9a0afb4f2c..cf5ff86abf 100644 --- a/tests/agent/test_turn_context.py +++ b/tests/agent/test_turn_context.py @@ -208,6 +208,38 @@ def test_returns_turn_context_with_user_message_appended(): assert ctx.active_system_prompt == "SYSTEM" +# ── Trivial-prompt prefetch gate (PR #25350 salvage) ───────────────────────── +# +# The prologue is the ONLY place the per-turn synchronous +# memory_manager.prefetch_all() fires; a bare greeting must not block the +# turn on provider network round-trips, while a substantive question must +# still prefetch. These assert the gate at the call site (the classifier +# itself is covered in tests/agent/test_memory_provider.py). + + +def _agent_with_memory_manager(): + agent = _FakeAgent() + mm = MagicMock() + mm.prefetch_all.return_value = "REMEMBERED CONTEXT" + agent._memory_manager = mm + return agent, mm + + +def test_prefetch_skipped_for_trivial_user_message(): + agent, mm = _agent_with_memory_manager() + ctx = _build(agent, user_message="hi!") + mm.prefetch_all.assert_not_called() + assert ctx.ext_prefetch_cache == "" + + +def test_prefetch_runs_for_substantive_user_message(): + agent, mm = _agent_with_memory_manager() + query = "what did we decide about the deploy pipeline?" + ctx = _build(agent, user_message=query) + mm.prefetch_all.assert_called_once_with(query) + assert ctx.ext_prefetch_cache == "REMEMBERED CONTEXT" + + def test_turn_start_replaces_stale_parent_history_with_compression_child(): agent = _FakeAgent() stale_history = [{"role": "user", "content": "stale parent"}] diff --git a/tests/agent/test_verification_stop_caching.py b/tests/agent/test_verification_stop_caching.py index 7620d9eb51..83bde23206 100644 --- a/tests/agent/test_verification_stop_caching.py +++ b/tests/agent/test_verification_stop_caching.py @@ -84,8 +84,9 @@ def test_db_flush_drops_only_nudge_keeps_candidate(tmp_path, monkeypatch): agent._flush_messages_to_session_db(messages, conversation_history=[]) persisted = [ - kwargs.get("content") - for _args, kwargs in agent._session_db.append_message.call_args_list + msg.get("content") + for _args, kwargs in agent._session_db.append_messages_batch.call_args_list + for msg in kwargs["messages"] ] assert "hi" in persisted assert "verified and clean" in persisted diff --git a/tests/agent/transports/test_chat_completions.py b/tests/agent/transports/test_chat_completions.py index 92e2aaf163..6f034fda5f 100644 --- a/tests/agent/transports/test_chat_completions.py +++ b/tests/agent/transports/test_chat_completions.py @@ -1,8 +1,12 @@ """Tests for the ChatCompletionsTransport.""" -import pytest +import json from types import SimpleNamespace +import httpx +import pytest +from openai import OpenAI + from agent.transports import get_transport from agent.transports.types import NormalizedResponse @@ -537,3 +541,193 @@ class TestChatCompletionsGeminiNativeExtraBodyStrip: eb = kw.get("extra_body") assert eb and "tags" in eb + def test_tags_pass_through_on_gemini_openai_compat(self, transport): + # /openai compat endpoint is not "native" — unchanged behavior. + kw = transport.build_kwargs( + "anthropic/claude-sonnet-4.6", + [{"role": "user", "content": "hi"}], + None, + provider_profile=self._nous_profile(), + base_url="https://generativelanguage.googleapis.com/v1beta/openai", + session_id="s1", + max_tokens=None, + ) + eb = kw.get("extra_body") + assert eb and "tags" in eb + + +class TestPromptCacheKeyCapability: + """Chat Completions cache routing is opt-in and body-safe.""" + + @staticmethod + def _messages(instructions="You are stable."): + return [ + {"role": "system", "content": instructions}, + {"role": "user", "content": "hello"}, + ] + + @staticmethod + def _tools(name="lookup"): + return [{ + "type": "function", + "function": { + "name": name, + "description": "Look something up.", + "parameters": {"type": "object", "properties": {}}, + }, + }] + + def _request_body(self, kwargs, *, stream=False): + captured = {} + + def handler(request): + captured.update(json.loads(request.content)) + if stream: + return httpx.Response( + 200, + headers={"content-type": "text/event-stream"}, + content=( + 'data: {"id":"chatcmpl_1","object":"chat.completion.chunk",' + '"choices":[{"index":0,"delta":{"content":"ok"},' + '"finish_reason":null}]}\n\n' + "data: [DONE]\n\n" + ), + ) + return httpx.Response(200, json={ + "id": "chatcmpl_1", + "object": "chat.completion", + "created": 0, + "model": kwargs["model"], + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + }], + }) + + with httpx.Client(transport=httpx.MockTransport(handler)) as http_client: + client = OpenAI( + api_key="test-key", + base_url="https://cache-capable.test/v1", + http_client=http_client, + ) + result = client.chat.completions.create(**kwargs, stream=stream) + if stream: + list(result) + return captured + + def test_profile_capability_emits_content_key_in_nonstream_request_body(self, transport): + from providers.base import ProviderProfile + + kwargs = transport.build_kwargs( + model="cache-model", + messages=self._messages(), + tools=self._tools(), + session_id="cron_job_2026-07-15T10:00:00Z", + provider_profile=ProviderProfile( + name="cache-capable", supports_prompt_cache_key=True, + ), + ) + + body = self._request_body(kwargs) + + assert body["prompt_cache_key"].startswith("pck_") + assert body["prompt_cache_key"] == kwargs["prompt_cache_key"] + + def test_legacy_capability_emits_same_key_in_streaming_request_body(self, transport): + kwargs = transport.build_kwargs( + model="cache-model", + messages=self._messages(), + tools=self._tools(), + session_id="cron_job_2026-07-15T10:05:00Z", + supports_prompt_cache_key=True, + ) + + body = self._request_body(kwargs, stream=True) + + assert body["prompt_cache_key"] == kwargs["prompt_cache_key"] + + def test_openai_api_base_url_implies_capability(self, transport): + """api.openai.com gets the key WITHOUT an explicit flag (exact host).""" + kwargs = transport.build_kwargs( + model="gpt-cache-model", + messages=self._messages(), + tools=self._tools(), + session_id="cron_job_2026-07-15T10:07:00Z", + base_url="https://api.openai.com/v1", + ) + + assert kwargs["prompt_cache_key"].startswith("pck_") + + @pytest.mark.parametrize( + "base_url", + [ + "https://myproxy.example.com/api.openai.com/v1", # host embedded in path + "https://api.openai.com.evil.example/v1", # prefix-spoofed host + "https://eastus.api.cognitive.microsoft.com/openai/v1", # Azure + ], + ) + def test_non_openai_hosts_do_not_imply_capability(self, transport, base_url): + kwargs = transport.build_kwargs( + model="strict-model", + messages=self._messages(), + tools=self._tools(), + session_id="cron_job_2026-07-15T10:08:00Z", + base_url=base_url, + ) + + assert "prompt_cache_key" not in kwargs + + @pytest.mark.parametrize("provider", [None, "anthropic", "custom"]) + def test_default_off_never_leaks_unknown_body_field(self, transport, provider): + from providers import get_provider_profile + + kwargs = transport.build_kwargs( + model="strict-model", + messages=self._messages(), + tools=self._tools(), + session_id="cron_job_2026-07-15T10:00:00Z", + provider_profile=(get_provider_profile(provider) if provider else None), + ) + + body = self._request_body(kwargs) + + assert "prompt_cache_key" not in kwargs + assert "prompt_cache_key" not in body + + def test_explicit_top_level_and_extra_body_overrides_are_preserved(self, transport): + from providers.base import ProviderProfile + + profile = ProviderProfile(name="cache-capable", supports_prompt_cache_key=True) + top_level = transport.build_kwargs( + model="cache-model", messages=self._messages(), tools=self._tools(), + provider_profile=profile, + request_overrides={"prompt_cache_key": "caller-top-level"}, + ) + in_extra_body = transport.build_kwargs( + model="cache-model", messages=self._messages(), tools=self._tools(), + provider_profile=profile, + request_overrides={"extra_body": {"prompt_cache_key": "caller-extra-body"}}, + ) + + assert top_level["prompt_cache_key"] == "caller-top-level" + assert "prompt_cache_key" not in top_level.get("extra_body", {}) + assert "prompt_cache_key" not in in_extra_body + assert in_extra_body["extra_body"]["prompt_cache_key"] == "caller-extra-body" + + def test_cron_ids_share_static_prefix_key_and_content_changes_invalidate(self, transport): + def key(session_id, *, instructions="You are stable.", tool_name="lookup"): + return transport.build_kwargs( + model="cache-model", + messages=self._messages(instructions), + tools=self._tools(tool_name), + session_id=session_id, + supports_prompt_cache_key=True, + )["prompt_cache_key"] + + first = key("cron_job_2026-07-15T10:00:00Z") + second = key("cron_job_2026-07-15T10:05:00Z") + + assert first == second + assert first != key("cron_job_2026-07-15T10:05:00Z", instructions="You are different.") + assert first != key("cron_job_2026-07-15T10:05:00Z", tool_name="search") diff --git a/tests/agent/transports/test_codex_transport.py b/tests/agent/transports/test_codex_transport.py index cf8f7dcc3e..444da68516 100644 --- a/tests/agent/transports/test_codex_transport.py +++ b/tests/agent/transports/test_codex_transport.py @@ -271,13 +271,15 @@ class TestCodexBuildKwargs: - def test_xai_injects_native_web_search_when_client_web_search_present(self, transport): - """xAI path swaps a client-side ``web_search`` function for xAI's - native server-side ``web_search`` built-in so grok server-side search - runs to completion (otherwise the turn stalls as - reasoning-with-no-answer -> false 'incomplete' -> 3 retries -> fail). + def test_xai_injects_native_web_search_when_client_web_search_present(self, transport, monkeypatch): + """When the active/configured search backend is xAI, swap client + ``web_search`` for Grok's native built-in so server-side search + completes (otherwise the turn stalls as incomplete → 3 retries). Non-conflicting client tools are preserved. """ + import agent.transports.codex as codex_mod + + monkeypatch.setattr(codex_mod, "_xai_prefers_native_web_search", lambda: True) messages = [{"role": "user", "content": "Find current prices."}] kw = transport.build_kwargs( model="grok-composer-2.5-fast", messages=messages, @@ -298,6 +300,71 @@ class TestCodexBuildKwargs: # Non-conflicting client-side tools are preserved. names = [t.get("name") for t in kw.get("tools", []) if t.get("type") == "function"] assert "read_file" in names + assert "web_search" not in names + assert "hermes_web_search" not in names + + def test_xai_renames_client_web_search_when_firecrawl_configured(self, transport, monkeypatch): + """Configured Firecrawl (or any non-xai backend) must keep Hermes + dispatch — rename the wire tool so Grok cannot hijack ``web_search``. + """ + import agent.transports.codex as codex_mod + + monkeypatch.setattr(codex_mod, "_xai_prefers_native_web_search", lambda: False) + messages = [{"role": "user", "content": "Find current prices."}] + kw = transport.build_kwargs( + model="grok-4.5", messages=messages, + tools=[ + {"type": "function", "function": { + "name": "read_file", "description": "Read a file.", + "parameters": {"type": "object", + "properties": {"path": {"type": "string"}}}}}, + {"type": "function", "function": { + "name": "web_search", "description": "Search the web.", + "parameters": {"type": "object", + "properties": {"query": {"type": "string"}}}}}, + ], + is_xai_responses=True, + ) + tools = kw.get("tools", []) + assert not any(t.get("type") == "web_search" for t in tools), tools + names = [t.get("name") for t in tools if t.get("type") == "function"] + assert "read_file" in names + assert "hermes_web_search" in names + assert "web_search" not in names + + def test_xai_normalize_maps_client_web_search_alias_back(self, transport, monkeypatch): + """Alias used on the wire must become ``web_search`` for Hermes dispatch.""" + import agent.transports.codex as codex_mod + + msg = SimpleNamespace( + content=None, + reasoning=None, + tool_calls=[ + SimpleNamespace( + id="call_1", + call_id="call_1", + response_item_id="fc_1", + function=SimpleNamespace( + name=codex_mod._XAI_CLIENT_WEB_SEARCH_ALIAS, + arguments='{"query":"hermes"}', + ), + ) + ], + codex_reasoning_items=None, + codex_message_items=None, + reasoning_details=None, + ) + response = SimpleNamespace(output=[], status="completed") + + monkeypatch.setattr( + "agent.codex_responses_adapter._normalize_codex_response", + lambda resp, issuer_kind=None: (msg, "tool_calls"), + ) + normalized = transport.normalize_response(response) + + assert normalized.tool_calls is not None + assert len(normalized.tool_calls) == 1 + assert normalized.tool_calls[0].name == "web_search" def test_xai_does_not_inject_native_web_search_without_client_web_search(self, transport): """The native ``web_search`` built-in is a 1:1 swap for an @@ -341,8 +408,6 @@ class TestCodexBuildKwargs: for t in tools ) - - # --- Grok reasoning-effort capability allowlist --- # api.x.ai 400s with "Model X does not support parameter reasoningEffort" # on grok-4 / grok-4-fast / grok-3 / grok-code-fast / grok-4.20-0309-*. @@ -351,10 +416,6 @@ class TestCodexBuildKwargs: # ``reasoning.encrypted_content`` back from xAI on every model — # see test_xai_reasoning_effort_passed for the rationale. - - - - def test_xai_grok_4_20_0309_variants_omit_reasoning_effort(self, transport): """grok-4.20-0309-(non-)reasoning reject the effort dial. @@ -370,7 +431,63 @@ class TestCodexBuildKwargs: assert "reasoning" not in kw, f"{model} must not receive reasoning" +class TestXaiWebSearchBackendPreference: + """``_xai_prefers_native_web_search`` must honor web backend config.""" + def test_explicit_firecrawl_prefers_client(self, monkeypatch): + import agent.transports.codex as codex_mod + + monkeypatch.setattr( + "agent.web_search_registry.get_active_search_provider", + lambda: SimpleNamespace(name="firecrawl"), + ) + assert codex_mod._xai_prefers_native_web_search() is False + + def test_explicit_search_backend_xai_prefers_native(self, monkeypatch): + import agent.transports.codex as codex_mod + + monkeypatch.setattr( + "agent.web_search_registry.get_active_search_provider", + lambda: SimpleNamespace(name="xai"), + ) + assert codex_mod._xai_prefers_native_web_search() is True + + def test_resolved_non_xai_provider_prefers_client(self, monkeypatch): + import agent.transports.codex as codex_mod + + monkeypatch.setattr( + "agent.web_search_registry.get_active_search_provider", + lambda: SimpleNamespace(name="firecrawl"), + ) + assert codex_mod._xai_prefers_native_web_search() is False + + def test_no_provider_legacy_fallback_xai(self, monkeypatch): + """When no provider is registered, fall back to _get_search_backend.""" + import agent.transports.codex as codex_mod + + monkeypatch.setattr( + "agent.web_search_registry.get_active_search_provider", + lambda: None, + ) + monkeypatch.setattr( + "tools.web_tools._get_search_backend", + lambda: "xai", + ) + assert codex_mod._xai_prefers_native_web_search() is True + + def test_no_provider_legacy_fallback_non_xai(self, monkeypatch): + """When no provider is registered and backend isn't xai, keep client.""" + import agent.transports.codex as codex_mod + + monkeypatch.setattr( + "agent.web_search_registry.get_active_search_provider", + lambda: None, + ) + monkeypatch.setattr( + "tools.web_tools._get_search_backend", + lambda: "firecrawl", + ) + assert codex_mod._xai_prefers_native_web_search() is False class TestCodexValidateResponse: diff --git a/tests/cli/test_cli_yolo_resume_persistence.py b/tests/cli/test_cli_yolo_resume_persistence.py new file mode 100644 index 0000000000..ce771d4d8d --- /dev/null +++ b/tests/cli/test_cli_yolo_resume_persistence.py @@ -0,0 +1,251 @@ +"""Regression tests: YOLO mode persists across ``hermes --resume``. + +Pre-fix bug: the ``/yolo`` toggle (and the process-start ``--yolo`` flag) +lived only in the in-memory ``tools.approval._session_yolo`` set / the +frozen env var. Resuming a session in a fresh process silently reverted the +bypass — dangerous commands started prompting again even though the user had +YOLO on for that session. + +The fix persists a ``yolo_mode`` flag inside the session row's +``model_config`` JSON: + +- ``SessionDB.set_session_yolo`` merges the flag (preserving lineage markers + like ``_branched_from``), written by the CLI ``/yolo`` toggle. +- ``AIAgent._ensure_db_session`` carries a live session bypass (or a frozen + ``--yolo`` launch, via agent_init) into the creation-time model_config. +- ``HermesCLI._restore_session_yolo`` reads the flag on every resume path + and re-enables the in-memory bypass. +""" + +import json +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest + +import tools.approval as approval_module +from cli import HermesCLI +from hermes_state import SessionDB + + +SESSION_ID = "yolo_persist_session" + + +@pytest.fixture(autouse=True) +def _hermetic_yolo(monkeypatch): + monkeypatch.delenv("HERMES_YOLO_MODE", raising=False) + monkeypatch.setattr(approval_module, "_YOLO_MODE_FROZEN", False) + approval_module.clear_session(SESSION_ID) + yield + approval_module.clear_session(SESSION_ID) + + +@pytest.fixture +def db(tmp_path): + d = SessionDB(db_path=tmp_path / "state.db") + yield d + try: + d.close() + except Exception: + pass + + +class TestSessionDbYoloFlag: + def test_set_and_read_round_trip(self, db): + db.create_session(session_id=SESSION_ID, source="cli", model="m") + db.set_session_yolo(SESSION_ID, True) + meta = db.get_session(SESSION_ID) + assert SessionDB.session_yolo_enabled(meta) is True + + db.set_session_yolo(SESSION_ID, False) + meta = db.get_session(SESSION_ID) + assert SessionDB.session_yolo_enabled(meta) is False + + def test_merge_preserves_existing_model_config_keys(self, db): + db.create_session( + session_id=SESSION_ID, + source="cli", + model="m", + model_config={"max_iterations": 42, "_branched_from": "parent_x"}, + ) + db.set_session_yolo(SESSION_ID, True) + meta = db.get_session(SESSION_ID) + config = json.loads(meta["model_config"]) + assert config["yolo_mode"] is True + assert config["max_iterations"] == 42 + assert config["_branched_from"] == "parent_x" + + def test_missing_row_is_noop(self, db): + # Row doesn't exist yet (lazy creation) — must not raise or create. + db.set_session_yolo("does_not_exist", True) + assert db.get_session("does_not_exist") is None + + def test_creation_time_model_config_flag_reads_back(self, db): + db.create_session( + session_id=SESSION_ID, + source="cli", + model="m", + model_config={"yolo_mode": True}, + ) + meta = db.get_session(SESSION_ID) + assert SessionDB.session_yolo_enabled(meta) is True + + def test_reader_is_false_on_garbage(self): + assert SessionDB.session_yolo_enabled(None) is False + assert SessionDB.session_yolo_enabled({}) is False + assert SessionDB.session_yolo_enabled({"model_config": None}) is False + assert SessionDB.session_yolo_enabled({"model_config": "not json {"}) is False + assert SessionDB.session_yolo_enabled({"model_config": "[1,2]"}) is False + assert ( + SessionDB.session_yolo_enabled({"model_config": '{"yolo_mode": false}'}) + is False + ) + + +def _stand_in(session_id=SESSION_ID, session_db=None): + return SimpleNamespace( + session_id=session_id, + _session_db=session_db, + _console_print=lambda *a, **k: None, + ) + + +class TestRestoreSessionYolo: + def test_restore_enables_bypass_when_flag_set(self): + stand_in = _stand_in() + meta = {"id": SESSION_ID, "model_config": '{"yolo_mode": true}'} + + assert approval_module.is_session_yolo_enabled(SESSION_ID) is False + HermesCLI._restore_session_yolo(stand_in, meta) + assert approval_module.is_session_yolo_enabled(SESSION_ID) is True + + def test_restore_noop_when_flag_absent(self): + stand_in = _stand_in() + meta = {"id": SESSION_ID, "model_config": '{"max_iterations": 10}'} + + HermesCLI._restore_session_yolo(stand_in, meta) + assert approval_module.is_session_yolo_enabled(SESSION_ID) is False + + def test_restore_noop_when_meta_empty(self): + stand_in = _stand_in() + HermesCLI._restore_session_yolo(stand_in, {}) + HermesCLI._restore_session_yolo(stand_in, None) + assert approval_module.is_session_yolo_enabled(SESSION_ID) is False + + def test_restore_idempotent_when_already_enabled(self): + stand_in = _stand_in() + approval_module.enable_session_yolo(SESSION_ID) + meta = {"id": SESSION_ID, "model_config": '{"yolo_mode": true}'} + # Should not raise or print duplicate banners; state stays enabled. + HermesCLI._restore_session_yolo(stand_in, meta) + assert approval_module.is_session_yolo_enabled(SESSION_ID) is True + + def test_restore_skipped_under_frozen_process_yolo(self): + stand_in = _stand_in() + meta = {"id": SESSION_ID, "model_config": '{"yolo_mode": true}'} + with patch.object(approval_module, "_YOLO_MODE_FROZEN", True): + HermesCLI._restore_session_yolo(stand_in, meta) + # Frozen bypass already covers everything — the session set is + # untouched (avoids persisting a session-scoped bypass the user + # only asked for at process scope). + assert approval_module.is_session_yolo_enabled(SESSION_ID) is False + + +class TestToggleYoloPersists: + def test_toggle_writes_flag_through_session_db(self): + db = MagicMock() + stand_in = SimpleNamespace(session_id=SESSION_ID, _session_db=db) + # Bind the real persist helper so the toggle's getattr finds it. + stand_in._persist_session_yolo = ( + lambda key, enabled: HermesCLI._persist_session_yolo( + stand_in, key, enabled + ) + ) + + with patch("cli._cprint"): + HermesCLI._toggle_yolo(stand_in) # ON + db.set_session_yolo.assert_called_once_with(SESSION_ID, True) + + with patch("cli._cprint"): + HermesCLI._toggle_yolo(stand_in) # OFF + db.set_session_yolo.assert_called_with(SESSION_ID, False) + + def test_toggle_survives_missing_session_db(self): + stand_in = SimpleNamespace(session_id=SESSION_ID, _session_db=None) + stand_in._persist_session_yolo = ( + lambda key, enabled: HermesCLI._persist_session_yolo( + stand_in, key, enabled + ) + ) + with patch("cli._cprint"): + HermesCLI._toggle_yolo(stand_in) # must not raise + assert approval_module.is_session_yolo_enabled(SESSION_ID) is True + + def test_toggle_still_works_without_persist_helper(self): + # Back-compat with the minimal stand-in used by older tests. + stand_in = SimpleNamespace(session_id=SESSION_ID) + with patch("cli._cprint"): + HermesCLI._toggle_yolo(stand_in) + assert approval_module.is_session_yolo_enabled(SESSION_ID) is True + + +class TestEndToEndPersistAndRestore: + def test_full_round_trip_through_real_db(self, db): + """Toggle ON in 'process 1', restore in 'process 2' (fresh in-memory + approval state), and verify a dangerous command auto-approves.""" + db.create_session(session_id=SESSION_ID, source="cli", model="m") + + # Process 1: user toggles /yolo ON — persisted to the row. + cli_one = SimpleNamespace(session_id=SESSION_ID, _session_db=db) + cli_one._persist_session_yolo = ( + lambda key, enabled: HermesCLI._persist_session_yolo( + cli_one, key, enabled + ) + ) + with patch("cli._cprint"): + HermesCLI._toggle_yolo(cli_one) + assert approval_module.is_session_yolo_enabled(SESSION_ID) is True + + # Simulate process exit: in-memory approval state is gone. + approval_module.clear_session(SESSION_ID) + assert approval_module.is_session_yolo_enabled(SESSION_ID) is False + + # Process 2: --resume reads the row and restores the bypass. + meta = db.get_session(SESSION_ID) + cli_two = _stand_in(session_db=db) + HermesCLI._restore_session_yolo(cli_two, meta) + assert approval_module.is_session_yolo_enabled(SESSION_ID) is True + + token = approval_module.set_current_session_key(SESSION_ID) + try: + result = approval_module.check_all_command_guards( + "rm -rf /tmp/scratch-xyzzy", "local", + ) + assert result["approved"] is True + finally: + approval_module.reset_current_session_key(token) + + def test_toggle_off_round_trip(self, db): + """OFF must persist too — a resumed session must not resurrect a + bypass the user explicitly turned off.""" + db.create_session( + session_id=SESSION_ID, + source="cli", + model="m", + model_config={"yolo_mode": True}, + ) + cli_one = SimpleNamespace(session_id=SESSION_ID, _session_db=db) + cli_one._persist_session_yolo = ( + lambda key, enabled: HermesCLI._persist_session_yolo( + cli_one, key, enabled + ) + ) + approval_module.enable_session_yolo(SESSION_ID) + with patch("cli._cprint"): + HermesCLI._toggle_yolo(cli_one) # OFF + approval_module.clear_session(SESSION_ID) + + meta = db.get_session(SESSION_ID) + cli_two = _stand_in(session_db=db) + HermesCLI._restore_session_yolo(cli_two, meta) + assert approval_module.is_session_yolo_enabled(SESSION_ID) is False diff --git a/tests/conftest.py b/tests/conftest.py index 672723a6d9..26c034edc9 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1449,3 +1449,22 @@ def _isolate_computer_use_approval_state(): _cu_tool._session_auto_approve.clear() except Exception: pass + + +@pytest.fixture(autouse=True) +def _moa_caches_isolated(): + """Clear module-level MoA cold-start caches before each test. + + ``agent.moa_loop`` caches the resolved preset and each slot's provider + runtime at module level (keyed on config mtime / provider+model) so the + tool loop doesn't re-resolve them serially on every iteration. Tests + monkeypatch resolvers and config paths, so a cache entry leaked from one + test would poison the next. Clear both around every test. + """ + import agent.moa_loop as moa + + moa._preset_cache.clear() + moa._runtime_cache.clear() + yield + moa._preset_cache.clear() + moa._runtime_cache.clear() diff --git a/tests/cron/test_idle_tick_config_skip.py b/tests/cron/test_idle_tick_config_skip.py new file mode 100644 index 0000000000..531a2fba87 --- /dev/null +++ b/tests/cron/test_idle_tick_config_skip.py @@ -0,0 +1,58 @@ +"""Idle cron ticks must not load config (#33612 salvage). + +The gateway's built-in ticker calls tick(verbose=False) every 60s. Before +the fix, idle ticks (no due jobs) fell through the verbose-only early +return and paid a full load_config() + worker-pool resolution per tick. +The fix returns early on ANY idle tick while preserving the post-tick MCP +orphan sweep that main intentionally runs even when nothing is due. +""" + +from __future__ import annotations + +from unittest.mock import patch + +import cron.scheduler as scheduler_mod + + +def _run_idle_tick(**kwargs): + """Run tick() with no due jobs; return (load_config_called, sweep_called).""" + calls = {"load_config": 0, "sweep": 0} + + def _fake_load_config(*a, **k): + calls["load_config"] += 1 + return {} + + def _fake_sweep(): + calls["sweep"] += 1 + + with ( + patch.object(scheduler_mod, "get_due_jobs", return_value=[]), + patch.object(scheduler_mod, "load_config", side_effect=_fake_load_config), + patch( + "tools.mcp_tool._kill_orphaned_mcp_children", + side_effect=_fake_sweep, + ), + ): + rc = scheduler_mod.tick(verbose=kwargs.get("verbose", False)) + return rc, calls + + +class TestIdleTickSkipsConfigLoad: + def test_idle_nonverbose_tick_skips_load_config(self): + """Gateway-style tick(verbose=False) with no due jobs: no config load.""" + rc, calls = _run_idle_tick(verbose=False) + assert rc == 0 + assert calls["load_config"] == 0, ( + "idle tick must not load config (was loading every 60s in the gateway ticker)" + ) + + def test_idle_verbose_tick_skips_load_config(self): + rc, calls = _run_idle_tick(verbose=True) + assert rc == 0 + assert calls["load_config"] == 0 + + def test_idle_tick_still_sweeps_mcp_orphans(self): + """The idle-tick orphan sweep is intentional on main — must survive.""" + rc, calls = _run_idle_tick(verbose=False) + assert rc == 0 + assert calls["sweep"] == 1, "idle tick must still reap orphaned MCP children" diff --git a/tests/cron/test_jobs.py b/tests/cron/test_jobs.py index e9402208e3..b0121daed1 100644 --- a/tests/cron/test_jobs.py +++ b/tests/cron/test_jobs.py @@ -323,11 +323,45 @@ class TestMarkJobRun: assert updated["repeat"]["completed"] == 1 assert updated["last_status"] == "ok" - def test_repeat_limit_removes_job(self, tmp_cron_dir): + def test_repeat_limit_retains_completed_record(self, tmp_cron_dir): + """A finished one-shot must stay inspectable, not vanish from the store.""" job = create_job(prompt="Once", schedule="30m", repeat=1) mark_job_run(job["id"], success=True) - # Job should be removed after hitting repeat limit - assert get_job(job["id"]) is None + updated = get_job(job["id"]) + assert updated is not None, "completed one-shot was deleted from jobs.json" + assert updated["state"] == "completed" + assert updated["enabled"] is False + assert updated["next_run_at"] is None + assert updated["last_status"] == "ok" + + def test_repeat_limit_retains_delivery_error(self, tmp_cron_dir): + """A one-shot whose delivery failed must keep the error on its record.""" + job = create_job(prompt="Once", schedule="30m", repeat=1) + mark_job_run( + job["id"], success=True, + delivery_error="platform 'telegram' not configured", + ) + updated = get_job(job["id"]) + assert updated is not None + assert updated["state"] == "completed" + assert updated["last_delivery_error"] == "platform 'telegram' not configured" + + def test_completed_oneshot_visible_in_list(self, tmp_cron_dir): + """list_jobs(include_disabled=True) surfaces the completed record.""" + job = create_job(prompt="Once", schedule="30m", repeat=1) + mark_job_run(job["id"], success=True, delivery_error="send failed: 502") + listed = {j["id"]: j for j in list_jobs(include_disabled=True)} + assert job["id"] in listed + assert listed[job["id"]]["state"] == "completed" + assert listed[job["id"]]["last_delivery_error"] == "send failed: 502" + # Default (enabled-only) listing hides it, matching paused/disabled jobs. + assert job["id"] not in {j["id"] for j in list_jobs()} + + def test_completed_oneshot_not_due(self, tmp_cron_dir): + """A retained completed one-shot must never be dispatched again.""" + job = create_job(prompt="Once", schedule="30m", repeat=1) + mark_job_run(job["id"], success=True) + assert job["id"] not in {j["id"] for j in get_due_jobs()} def test_error_status(self, tmp_cron_dir): @@ -639,9 +673,14 @@ class TestGetDueJobs: assert get_job("slowrun") is not None # Run completes → outcome lands on a record that still exists - # (times=1 reached, so mark_job_run retires the job normally). + # (times=1 reached, so mark_job_run retires the job as a terminal + # completed record instead of deleting it). mark_job_run("slowrun", True) - assert get_job("slowrun") is None + retired = get_job("slowrun") + assert retired is not None + assert retired["state"] == "completed" + assert retired["enabled"] is False + assert retired["last_status"] == "ok" def test_heartbeat_run_claim_rejects_replaced_owner(self, tmp_cron_dir): @@ -900,12 +939,17 @@ class TestClaimDispatch: def test_mark_job_run_does_not_double_count_preclaimed_oneshot(self, tmp_cron_dir): # Full lifecycle: claim bumps completed to times, then mark_job_run must - # NOT increment again — it recognizes the pre-claim and removes the job. + # NOT increment again — it recognizes the pre-claim and retires the job + # as a terminal completed record (retained for inspection, not re-fired). save_jobs([self._oneshot(times=1, completed=0)]) assert claim_dispatch("os1") is True assert load_jobs()[0]["repeat"]["completed"] == 1 mark_job_run("os1", success=True) - assert load_jobs() == [] # completed once, removed — not fired twice + retired = load_jobs() + assert len(retired) == 1 # completed once, retired — not fired twice + assert retired[0]["repeat"]["completed"] == 1 # no double count + assert retired[0]["state"] == "completed" + assert retired[0]["enabled"] is False def test_get_due_jobs_removes_stale_maxed_oneshot(self, tmp_cron_dir): @@ -1139,3 +1183,70 @@ class TestAdvanceNextRuns: assert advance_next_run(rec_ids[0]) is True assert advance_next_run(one_ids[0]) is False assert advance_next_run("missing-id") is False + + +# ========================================================================= +# Completed one-shot retention sweep +# ========================================================================= + +class TestCompletedOneshotRetentionSweep: + """Completed one-shots are retained for inspection, then pruned by age.""" + + def _completed_oneshot(self, age_days: float): + """Create a one-shot, complete it, and backdate its last_run_at.""" + job = create_job(prompt="Once", schedule="30m", repeat=1) + mark_job_run(job["id"], success=True, delivery_error="boom") + stamp = ( + datetime.now(timezone.utc) - timedelta(days=age_days) + ).isoformat() + jobs = load_jobs() + for j in jobs: + if j["id"] == job["id"]: + j["last_run_at"] = stamp + save_jobs(jobs) + return job["id"] + + def test_sweep_prunes_old_completed_oneshot(self, tmp_cron_dir): + old_id = self._completed_oneshot(age_days=30) + get_due_jobs() # sweep runs as part of the due scan + assert get_job(old_id) is None + + def test_sweep_keeps_recent_completed_oneshot(self, tmp_cron_dir): + recent_id = self._completed_oneshot(age_days=1) + get_due_jobs() + kept = get_job(recent_id) + assert kept is not None + assert kept["state"] == "completed" + assert kept["last_delivery_error"] == "boom" + + def test_sweep_ignores_recurring_jobs(self, tmp_cron_dir): + """Old recurring jobs are never candidates, whatever their history.""" + job = create_job(prompt="Recurring", schedule="every 1h") + stamp = ( + datetime.now(timezone.utc) - timedelta(days=365) + ).isoformat() + jobs = load_jobs() + for j in jobs: + if j["id"] == job["id"]: + j["last_run_at"] = stamp + save_jobs(jobs) + get_due_jobs() + assert get_job(job["id"]) is not None + + def test_sweep_disabled_by_nonpositive_retention(self, tmp_cron_dir, monkeypatch): + monkeypatch.setattr( + "cron.jobs._completed_oneshot_retention_days", lambda: 0.0 + ) + old_id = self._completed_oneshot(age_days=30) + get_due_jobs() + assert get_job(old_id) is not None + + def test_recurring_jobs_unaffected_by_retention_change(self, tmp_cron_dir): + """A recurring job still cycles normally alongside retained one-shots.""" + recurring = create_job(prompt="Recurring", schedule="every 1h") + self._completed_oneshot(age_days=1) + mark_job_run(recurring["id"], success=True) + updated = get_job(recurring["id"]) + assert updated["enabled"] is True + assert updated["state"] == "scheduled" + assert updated["next_run_at"] is not None diff --git a/tests/gateway/conftest.py b/tests/gateway/conftest.py index 7a21465dd3..d485d4ab2f 100644 --- a/tests/gateway/conftest.py +++ b/tests/gateway/conftest.py @@ -39,6 +39,33 @@ from unittest.mock import MagicMock import pytest +@pytest.fixture(scope="session", autouse=True) +def _bind_lark_sdk_globals_when_installed(): + """Bind the feishu adapter's lark SDK globals once per test session. + + The adapter defers ``import lark_oapi`` to first use + (``_load_lark_oapi`` — called from connect()/probe_bot()/standalone + send), so the request-builder globals (``CreateMessageRequestBody`` + etc.) stay ``None`` at module import time. Feishu tests across many + files inject a mock ``_client`` and skip connect() entirely, then call + send paths that reference those globals. Bind them eagerly when the + SDK is installed; when it isn't, the affected tests already skip via + their own ``skipUnless`` guards. + """ + try: + import lark_oapi # noqa: F401 + except ImportError: + yield + return + try: + from plugins.platforms.feishu.adapter import _load_lark_oapi + + _load_lark_oapi() + except Exception: + pass # adapter not importable in this environment — tests will skip + yield + + def make_async_session_db(sync_mock=None): """Wrap a sync mock SessionDB in AsyncSessionDB so gateway code that awaits the facade works in tests. Returns (facade, sync_mock); configure return diff --git a/tests/gateway/platforms/test_yuanbao_state_cleanup.py b/tests/gateway/platforms/test_yuanbao_state_cleanup.py new file mode 100644 index 0000000000..d2dc166f2c --- /dev/null +++ b/tests/gateway/platforms/test_yuanbao_state_cleanup.py @@ -0,0 +1,174 @@ +"""Yuanbao per-turn state cleanup: RecallGuard tracking dicts + member cache TTL. + +Covers the salvage of PRs #23383 / #23384: + +* ``_processing_msg_ids`` / ``_processing_msg_texts`` must be cleared when a + turn finishes (they previously leaked forever, letting RecallGuard match a + recall against an already-finished turn). +* The cleanup must pop ONLY when the finishing event's msg_id is truthy AND + still owns the entry. An id-less event (internal/synthetic message, push + without msg_id) never wrote an entry, so it must never erase one either — + the entry it sees belongs to a concurrently-queued id-bearing message whose + drain task still needs it. +* ``_member_cache`` entries past ``MEMBER_CACHE_TTL_S`` must actually be + evicted on read (the dict shrinks), while fresh entries survive. +""" +import asyncio +import time +from types import SimpleNamespace + +from gateway.platforms.base import BasePlatformAdapter +from gateway.platforms.yuanbao import MessageSender, YuanbaoAdapter + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +class _OutboundStub: + async def start_slow_notifier(self, chat_id): # noqa: ANN001 + pass + + def cancel_slow_notifier(self, chat_id): # noqa: ANN001 + pass + + +def _bare_adapter(): + """YuanbaoAdapter instance without running its heavy __init__.""" + adapter = object.__new__(YuanbaoAdapter) + adapter._outbound = _OutboundStub() + adapter._processing_msg_ids = {} + adapter._processing_msg_texts = {} + return adapter + + +def _event(message_id): + return SimpleNamespace( + source=SimpleNamespace(chat_id="chat-1"), + message_id=message_id, + ) + + +def _run_turn(monkeypatch, adapter, event, session_key, during_turn=None): + """Run the yuanbao _process_message_background wrapper with the base + class processing stubbed out (optionally mutating state mid-turn).""" + + async def _base_stub(self, ev, sk): # noqa: ANN001 + if during_turn is not None: + during_turn() + + monkeypatch.setattr( + BasePlatformAdapter, "_process_message_background", _base_stub + ) + asyncio.run( + YuanbaoAdapter._process_message_background(adapter, event, session_key) + ) + + +# --------------------------------------------------------------------------- +# _processing_msg_ids / _processing_msg_texts cleanup (PR #23383) +# --------------------------------------------------------------------------- + +def test_tracking_entries_cleared_after_normal_turn(monkeypatch): + """A turn whose msg_id still owns the tracking entry clears it on exit.""" + adapter = _bare_adapter() + sk = "yuanbao:group:G:user:U" + # _dispatch_inbound_event wrote these before handle_message. + adapter._processing_msg_ids[sk] = "m1" + adapter._processing_msg_texts[sk] = "hello" + + _run_turn(monkeypatch, adapter, _event("m1"), sk) + + assert sk not in adapter._processing_msg_ids + assert sk not in adapter._processing_msg_texts + + +def test_idless_event_must_not_erase_drain_tasks_entry(monkeypatch): + """An id-less outer event finishing must NOT pop the tracking entry a + concurrently-dispatched id-bearing message (queued as pending, to be + handled by a drain task) wrote during the outer turn.""" + adapter = _bare_adapter() + sk = "yuanbao:group:G:user:U" + + def _pending_message_arrives(): + # Simulates _dispatch_inbound_event for msg "m2" arriving while the + # id-less event is still processing: it writes tracking state, then + # handle_message routes it to _pending_messages for the drain task. + adapter._processing_msg_ids[sk] = "m2" + adapter._processing_msg_texts[sk] = "recallable text" + + _run_turn( + monkeypatch, adapter, _event(None), sk, + during_turn=_pending_message_arrives, + ) + + # The drain task for "m2" still needs these for RecallGuard matching. + assert adapter._processing_msg_ids.get(sk) == "m2" + assert adapter._processing_msg_texts.get(sk) == "recallable text" + + +def test_overwritten_entry_not_erased_by_outdated_turn(monkeypatch): + """If a newer message already overwrote the entry, the older finishing + turn must leave it alone (drain task owns it).""" + adapter = _bare_adapter() + sk = "yuanbao:group:G:user:U" + adapter._processing_msg_ids[sk] = "m1" + adapter._processing_msg_texts[sk] = "first" + + def _newer_message_arrives(): + adapter._processing_msg_ids[sk] = "m2" + adapter._processing_msg_texts[sk] = "second" + + _run_turn( + monkeypatch, adapter, _event("m1"), sk, + during_turn=_newer_message_arrives, + ) + + assert adapter._processing_msg_ids.get(sk) == "m2" + assert adapter._processing_msg_texts.get(sk) == "second" + + +# --------------------------------------------------------------------------- +# _member_cache TTL eviction (PR #23384) +# --------------------------------------------------------------------------- + +def _bare_sender(adapter_stub): + sender = object.__new__(MessageSender) + sender._adapter = adapter_stub + return sender + + +def test_member_cache_expired_entry_is_evicted(): + """Reading an expired entry must delete it — the cache dict shrinks.""" + now = time.time() + adapter = SimpleNamespace( + MEMBER_CACHE_TTL_S=300.0, + _member_cache={ + "g-stale": (now - 301.0, [{"nickname": "bob", "user_id": "u1"}]), + }, + ) + sender = _bare_sender(adapter) + + body = sender._build_msg_body_with_mentions("hi @bob", "g-stale") + + # Expired ⇒ no member data ⇒ plain text body, and the key is GONE. + assert body == [{"msg_type": "TIMTextElem", "msg_content": {"text": "hi @bob"}}] + assert "g-stale" not in adapter._member_cache + assert len(adapter._member_cache) == 0 + + +def test_member_cache_fresh_entry_survives_read(): + """A fresh entry is used for mention resolution and stays cached.""" + now = time.time() + members = [{"nickname": "bob", "user_id": "u1"}] + adapter = SimpleNamespace( + MEMBER_CACHE_TTL_S=300.0, + _member_cache={"g-fresh": (now - 10.0, members)}, + ) + sender = _bare_sender(adapter) + + body = sender._build_msg_body_with_mentions("hi @bob", "g-fresh") + + assert "g-fresh" in adapter._member_cache + # Fresh members were actually used: an @mention element is present. + assert any(el.get("msg_type") == "TIMCustomElem" for el in body) diff --git a/tests/gateway/relay/test_relay_adapter.py b/tests/gateway/relay/test_relay_adapter.py index ba93036067..137c0b15cd 100644 --- a/tests/gateway/relay/test_relay_adapter.py +++ b/tests/gateway/relay/test_relay_adapter.py @@ -1,5 +1,7 @@ """RelayAdapter capability-advertisement tests (relay Phase 1, Task 1.1).""" +import asyncio + import pytest from gateway.config import Platform, PlatformConfig @@ -274,3 +276,47 @@ async def test_get_chat_info_local_fallback_when_not_advertised(): info = await a.get_chat_info("chan-1") assert info == {"name": "chan-1", "type": "dm"} assert t.calls == [] + + +class _HangOnIdleTransport: + """Transport that hangs in go_idle so outer disconnect cancellation can race it.""" + + def __init__(self): + self.go_idle_started = asyncio.Event() + self.go_idle_timeouts: list[float] = [] + self.disconnect_calls = 0 + + def set_inbound_handler(self, h): # noqa: D401 + self._h = h + + async def go_idle(self, timeout_s: float = 10.0): + self.go_idle_timeouts.append(timeout_s) + self.go_idle_started.set() + await asyncio.sleep(3600) + return False + + async def disconnect(self): + self.disconnect_calls += 1 + + +@pytest.mark.asyncio +async def test_disconnect_tears_down_transport_when_go_idle_is_cancelled(): + """Runner disconnect budgets can cancel adapter.disconnect mid go_idle. + + The gateway runner's default adapter disconnect budget is 5s, while + transport.go_idle defaults to 10s. If cancellation lands during the idle + handshake, transport.disconnect must still run so the websocket/supervisor + cannot outlive the adapter. + """ + transport = _HangOnIdleTransport() + adapter = RelayAdapter(PlatformConfig(), make_desc(platform="discord"), transport=transport) + + task = asyncio.create_task(adapter.disconnect()) + await asyncio.wait_for(transport.go_idle_started.wait(), timeout=1.0) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + assert transport.disconnect_calls == 1 + assert transport.go_idle_timeouts + assert transport.go_idle_timeouts[0] < 5.0 diff --git a/tests/gateway/restart_test_helpers.py b/tests/gateway/restart_test_helpers.py index c0466838bc..7de82a65bb 100644 --- a/tests/gateway/restart_test_helpers.py +++ b/tests/gateway/restart_test_helpers.py @@ -4,7 +4,10 @@ from unittest.mock import AsyncMock, MagicMock from gateway.config import GatewayConfig, Platform, PlatformConfig from gateway.platforms.base import BasePlatformAdapter, SendResult -from gateway.restart import DEFAULT_GATEWAY_RESTART_DRAIN_TIMEOUT +from gateway.restart import ( + DEFAULT_GATEWAY_RESTART_AFTER_TURN_TIMEOUT, + DEFAULT_GATEWAY_RESTART_DRAIN_TIMEOUT, +) from gateway.run import GatewayRunner from gateway.session import SessionSource @@ -73,6 +76,7 @@ def make_restart_runner( runner._detached_restart_helper_started = False runner._restart_command_source = None runner._restart_drain_timeout = DEFAULT_GATEWAY_RESTART_DRAIN_TIMEOUT + runner._restart_after_turn_timeout = DEFAULT_GATEWAY_RESTART_AFTER_TURN_TIMEOUT runner._stop_task = None runner._busy_input_mode = "interrupt" runner._update_prompt_pending = {} @@ -142,6 +146,9 @@ def make_restart_runner( runner._launch_detached_restart_command = GatewayRunner._launch_detached_restart_command.__get__( runner, GatewayRunner ) + runner._await_active_work_before_restart = ( + GatewayRunner._await_active_work_before_restart.__get__(runner, GatewayRunner) + ) runner.request_restart = GatewayRunner.request_restart.__get__(runner, GatewayRunner) runner._is_user_authorized = lambda _source: True runner.hooks = MagicMock() diff --git a/tests/gateway/test_api_server.py b/tests/gateway/test_api_server.py index 8c6f586c1a..5a17047914 100644 --- a/tests/gateway/test_api_server.py +++ b/tests/gateway/test_api_server.py @@ -913,6 +913,7 @@ class TestToolsetsEndpoint: ("default", "Default Tools", "Core tools"), ("web", "Web Tools", "Search and extract"), ] + feature_snapshot = object() with patch( "hermes_cli.tools_config._get_effective_configurable_toolsets", return_value=fake_toolsets, @@ -920,9 +921,12 @@ class TestToolsetsEndpoint: "hermes_cli.tools_config._get_platform_tools", return_value={"default"}, ), patch( + "hermes_cli.tools_config.get_nous_subscription_features", + return_value=feature_snapshot, + ) as resolve_features, patch( "hermes_cli.tools_config._toolset_has_keys", return_value=True, - ), patch( + ) as has_keys, patch( "toolsets.resolve_toolset", side_effect=lambda name: { "default": ["terminal", "read_file"], @@ -943,6 +947,13 @@ class TestToolsetsEndpoint: assert by_name["web"]["tools"] == ["web_search"] assert by_name["default"]["configured"] is True + resolve_features.assert_called_once() + assert has_keys.call_count == len(fake_toolsets) + assert all( + call.kwargs["features"] is feature_snapshot + for call in has_keys.call_args_list + ) + # --------------------------------------------------------------------------- # /v1/chat/completions endpoint @@ -1663,13 +1674,15 @@ class TestResponsesStreaming: # Patch web.StreamResponse for the duration of the writer call. import gateway.platforms.api_server as api_mod - import queue as _q - stream_q: _q.Queue = _q.Queue() + # The SSE writers consume an asyncio queue (ThreadSafeAsyncQueue), + # not a plain queue.Queue — a stdlib queue would block the drain + # loop's ``await stream_q.get()`` forever. + stream_q = api_mod.ThreadSafeAsyncQueue() async def _agent_coro(): # Feed one partial delta into the stream queue... - stream_q.put("partial output") + stream_q.put_nowait("partial output") # ...then give the drain loop a moment to pick it up before # raising CancelledError to simulate a server-side cancel. await asyncio.sleep(0.01) @@ -1734,11 +1747,12 @@ class TestResponsesStreaming: raise ConnectionResetError("simulated client disconnect") import gateway.platforms.api_server as api_mod - import queue as _q - stream_q: _q.Queue = _q.Queue() - stream_q.put("some streamed text") - stream_q.put(None) # EOS sentinel + # asyncio queue to match the writers' consumer (see the note in + # test_stream_cancelled_persists_incomplete_snapshot). + stream_q = api_mod.ThreadSafeAsyncQueue() + stream_q.put_nowait("some streamed text") + stream_q.put_nowait(None) # EOS sentinel async def _agent_coro(): await asyncio.sleep(0.01) @@ -2846,4 +2860,3 @@ class TestCreateAgentModelRecovery: adapter._create_agent(session_id="another-session", gateway_session_key="stable-chan-1") assert captured[1]["model"] == "minimax/minimax-m3" - diff --git a/tests/gateway/test_compress_command.py b/tests/gateway/test_compress_command.py index 599ad2b161..4e502447b7 100644 --- a/tests/gateway/test_compress_command.py +++ b/tests/gateway/test_compress_command.py @@ -1,5 +1,7 @@ """Tests for gateway /compress user-facing messaging.""" +import asyncio +import threading from datetime import datetime from unittest.mock import AsyncMock, MagicMock, patch @@ -420,3 +422,97 @@ async def test_compress_command_single_profile_skips_profile_resolution(): runner._resolve_profile_home_for_source.assert_not_called() runner._shutdown_executor() + + +@pytest.mark.asyncio +async def test_compress_command_cleanup_does_not_block_event_loop(): + """Manual /compress must not run agent teardown on the gateway event loop. + + #53175 offloaded session-expiry, hygiene, and shutdown cleanup, but the + manual /compress finally still called ``_cleanup_agent_resources`` inline. + A slow ``agent.close()`` there freezes the whole loop and stops the + runtime-status heartbeat from advancing — the same wedge class as the + original incident. + + Observation must happen from a side thread: if cleanup blocks the event + loop, an ``await``-based waiter cannot sample ticks until close returns, + which falsely looks healthy after the block ends. + """ + import time + + history = _make_history() + compressed = [ + history[0], + {"role": "assistant", "content": "compressed summary"}, + history[-1], + ] + runner = _make_runner(history) + + close_started = threading.Event() + release_close = threading.Event() + + def slow_close(): + close_started.set() + release_close.wait(timeout=5) + + agent_instance = MagicMock() + agent_instance.shutdown_memory_provider = MagicMock() + agent_instance.close = slow_close + agent_instance._cached_system_prompt = "" + agent_instance.tools = None + agent_instance.context_compressor.has_content_to_compress.return_value = True + agent_instance.context_compressor._last_compress_aborted = False + agent_instance.context_compressor._last_summary_fallback_used = False + agent_instance.context_compressor._last_summary_dropped_count = 0 + agent_instance.context_compressor._last_summary_error = None + agent_instance.context_compressor._last_aux_model_failure_model = None + agent_instance.context_compressor._last_aux_model_failure_error = None + agent_instance.session_id = "sess-1" + agent_instance._compress_context.return_value = (compressed, "") + agent_instance._compression_skipped_due_to_lock = False + agent_instance._session_messages = None + + ticks = {"n": 0} + stop = threading.Event() + observed = {} + + async def _heartbeat(): + while not stop.is_set(): + ticks["n"] += 1 + await asyncio.sleep(0.005) + + def _observer(): + # threading.Event wait does not need the event loop. Sample ticks + # while close() is still held so an on-loop teardown is visible. + if not close_started.wait(timeout=5): + observed["error"] = "close() never started" + release_close.set() + return + baseline = ticks["n"] + time.sleep(0.12) + observed["ticks_during_block"] = ticks["n"] - baseline + release_close.set() + + hb = asyncio.create_task(_heartbeat()) + observer = threading.Thread(target=_observer, name="compress-cleanup-observer", daemon=True) + observer.start() + + with ( + patch("gateway.run._resolve_runtime_agent_kwargs", return_value={"api_key": "***"}), + patch("gateway.run._resolve_gateway_model", return_value="test-model"), + patch("run_agent.AIAgent", return_value=agent_instance), + patch("agent.model_metadata.estimate_request_tokens_rough", return_value=100), + ): + result = await runner._handle_compress_command(_make_event()) + + observer.join(timeout=5) + stop.set() + await hb + runner._shutdown_executor() + + assert "Compressed:" in result + assert "error" not in observed, observed.get("error") + assert observed.get("ticks_during_block", 0) >= 5, ( + "event loop was blocked during manual /compress cleanup: only " + f"{observed.get('ticks_during_block')} ticks while agent.close() was running" + ) diff --git a/tests/gateway/test_delivery_ledger.py b/tests/gateway/test_delivery_ledger.py index cec3b8339e..f9421f7e6b 100644 --- a/tests/gateway/test_delivery_ledger.py +++ b/tests/gateway/test_delivery_ledger.py @@ -10,6 +10,7 @@ id stability, and the startup redelivery sweep's contract: """ import time +import threading from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -50,6 +51,33 @@ def _row(oid): } +def _blocking_probe(): + """Return a blocking ledger call and an event-loop progress witness.""" + ledger_started = threading.Event() + event_loop_progressed = threading.Event() + blocked_event_loop = [] + + def _slow_ledger_call(*args, **kwargs): + ledger_started.set() + # Generous timeout: a genuinely blocked loop can never set the event + # (the witness coroutine cannot run), so a longer wait only guards + # against loaded-CI scheduling flake, not against missing the bug. + if not event_loop_progressed.wait(timeout=5.0): + blocked_event_loop.append(True) + + async def _event_loop_witness(): + import asyncio + + deadline = asyncio.get_running_loop().time() + 10 + while not ledger_started.is_set(): + if asyncio.get_running_loop().time() >= deadline: + raise AssertionError("ledger call never started") + await asyncio.sleep(0) + event_loop_progressed.set() + + return _slow_ledger_call, _event_loop_witness, blocked_event_loop + + def _orphan(oid): """Make the row look like it belongs to a dead process.""" with dl._connect() as conn: @@ -171,6 +199,28 @@ class TestGatewayRedeliverySweep: assert sent["content"].startswith(dl.RECOVERED_MARKER) assert sent["content"].endswith("the final answer") + @pytest.mark.parametrize( + ("send_success", "ledger_method"), + [(True, "mark_delivered"), (False, "mark_failed")], + ) + @pytest.mark.asyncio + async def test_slow_state_update_does_not_block_event_loop( + self, send_success, ledger_method + ): + import asyncio + + _record() + _orphan("ob-1") + runner = self._runner(self._adapter(success=send_success)) + slow_update, event_loop_witness, blocked_event_loop = _blocking_probe() + + with patch.object(dl, ledger_method, side_effect=slow_update): + await asyncio.gather( + runner._redeliver_pending_obligations(), event_loop_witness() + ) + + assert blocked_event_loop == [] + class TestAttemptsOnlySpentOnRealSends: """``attempts`` is the redelivery budget — it must buy a send. diff --git a/tests/gateway/test_delivery_ledger_producer.py b/tests/gateway/test_delivery_ledger_producer.py index 8b99ea1a7f..1071d9f36e 100644 --- a/tests/gateway/test_delivery_ledger_producer.py +++ b/tests/gateway/test_delivery_ledger_producer.py @@ -8,6 +8,7 @@ block the send. """ import asyncio +import threading from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -65,6 +66,31 @@ def _rows(): ).fetchall() +def _blocking_probe(): + """Return a blocking ledger call and an event-loop progress witness.""" + ledger_started = threading.Event() + event_loop_progressed = threading.Event() + blocked_event_loop = [] + + def _slow_ledger_call(*args, **kwargs): + ledger_started.set() + # Generous timeout: a genuinely blocked loop can never set the event + # (the witness coroutine cannot run), so a longer wait only guards + # against loaded-CI scheduling flake, not against missing the bug. + if not event_loop_progressed.wait(timeout=5.0): + blocked_event_loop.append(True) + + async def _event_loop_witness(): + deadline = asyncio.get_running_loop().time() + 10 + while not ledger_started.is_set(): + if asyncio.get_running_loop().time() >= deadline: + raise AssertionError("ledger call never started") + await asyncio.sleep(0) + event_loop_progressed.set() + + return _slow_ledger_call, _event_loop_witness, blocked_event_loop + + async def _run(adapter, event, response="final answer"): adapter._message_handler = AsyncMock(return_value=response) session_key = "agent:main:slack:channel:C1" @@ -97,6 +123,36 @@ class TestProducerHook: assert rows[0][1] == "failed" + @pytest.mark.asyncio + async def test_slow_ledger_record_does_not_block_event_loop(self): + adapter = _Adapter() + slow_record, event_loop_witness, blocked_event_loop = _blocking_probe() + + with patch( + "gateway.delivery_ledger.record_obligation", + side_effect=slow_record, + ), patch("gateway.delivery_ledger.mark_attempting"): + await asyncio.gather(_run(adapter, _event()), event_loop_witness()) + + assert blocked_event_loop == [] + assert adapter.sent == ["final answer"] + + @pytest.mark.asyncio + async def test_slow_ledger_update_does_not_block_event_loop(self): + adapter = _Adapter() + slow_delivered, event_loop_witness, blocked_event_loop = _blocking_probe() + + with patch("gateway.delivery_ledger.record_obligation"), patch( + "gateway.delivery_ledger.mark_attempting" + ), patch( + "gateway.delivery_ledger.mark_delivered", + side_effect=slow_delivered, + ): + await asyncio.gather(_run(adapter, _event()), event_loop_witness()) + + assert blocked_event_loop == [] + assert adapter.sent == ["final answer"] + @pytest.mark.asyncio async def test_crash_between_attempting_and_ack_is_recoverable(self): """The core scenario (#58818): process dies mid-send. The row must diff --git a/tests/gateway/test_discord_format.py b/tests/gateway/test_discord_format.py index 112678d0c7..5684eb6f1c 100644 --- a/tests/gateway/test_discord_format.py +++ b/tests/gateway/test_discord_format.py @@ -44,3 +44,79 @@ class TestDiscordFormatMessage: assert "|---" not in out +class TestDiscordToolPreviewFormatting: + def test_truncated_url_keeps_full_click_target(self): + from agent.display import ToolPreview + + adapter = _make_discord_adapter() + url = "https://hermes-agent.nousresearch.com/docs/gateway/discord/tool-progress" + visible = "https://hermes-agent.nousresearch..." + + out = adapter.format_tool_preview(ToolPreview(visible, truncated=True, url=url)) + + assert out == f"[hermes-agent.nousresearch...](<{url}>)" + + def test_truncated_url_label_is_not_a_second_url_target(self): + from agent.display import ToolPreview + + adapter = _make_discord_adapter() + url = "https://centaur.run/secrets/advanced-permissioning" + visible = "https://centaur.run/secrets/advanced-..." + + out = adapter.format_tool_preview(ToolPreview(visible, truncated=True, url=url)) + + assert out == ( + "[centaur.run/secrets/advanced-...]" + "()" + ) + + def test_link_escapes_discord_markdown_delimiters(self): + from agent.display import ToolPreview + + adapter = _make_discord_adapter() + preview = ToolPreview( + r"https://example.com/docs/[beta]...", + truncated=True, + url="https://example.com/docs/_(beta)", + ) + + assert adapter.format_tool_preview(preview) == ( + r"[example.com/docs/\[beta\]...]" + r"()" + ) + + def test_structured_tool_event_uses_clickable_truncated_url(self): + from gateway.stream_events import ToolCallChunk + + adapter = _make_discord_adapter() + url = "https://hermes-agent.nousresearch.com/docs/gateway/discord/tool-progress" + visible = url[:37] + "..." + + out = adapter.format_tool_event( + ToolCallChunk("web_extract", preview=url, args={"urls": [url]}), + mode="all", + preview_max_len=40, + ) + + assert out is not None + assert f"[{visible.removeprefix('https://')}](<{url}>)" in out + + def test_untruncated_url_remains_plain(self): + from agent.display import ToolPreview + + adapter = _make_discord_adapter() + url = "https://example.com/page" + + out = adapter.format_tool_preview(ToolPreview(url)) + + assert out == url + + def test_truncated_non_url_remains_plain(self): + from agent.display import ToolPreview + + adapter = _make_discord_adapter() + visible = "a long search query that was trunc..." + + out = adapter.format_tool_preview(ToolPreview(visible, truncated=True)) + + assert out == visible diff --git a/tests/gateway/test_feishu_lazy_import.py b/tests/gateway/test_feishu_lazy_import.py new file mode 100644 index 0000000000..5f70643205 --- /dev/null +++ b/tests/gateway/test_feishu_lazy_import.py @@ -0,0 +1,55 @@ +"""Regression coverage for deferred Feishu SDK loading.""" + +import asyncio +import os +import tempfile +from unittest.mock import AsyncMock, patch + + +def _feishu_adapter_module(): + """Import the adapter with the Windows app-data root available in CI.""" + with patch.dict(os.environ, {"LOCALAPPDATA": tempfile.gettempdir()}): + from plugins.platforms.feishu import adapter + + return adapter + + +def test_configured_feishu_dependency_check_does_not_load_sdk(): + """Gateway configuration can validate Feishu without importing its SDK.""" + feishu_adapter = _feishu_adapter_module() + + with ( + patch.object(feishu_adapter, "FEISHU_AVAILABLE", False), + patch("tools.lazy_deps.ensure", autospec=True) as ensure, + ): + assert feishu_adapter.check_feishu_requirements() is True + assert feishu_adapter.FEISHU_AVAILABLE is False + + ensure.assert_called_once_with("platform.feishu", prompt=False) + + +def test_feishu_connect_loads_sdk_on_worker_thread(): + """The first SDK import is deferred until a configured adapter connects.""" + from gateway.config import PlatformConfig + feishu_adapter = _feishu_adapter_module() + + adapter = feishu_adapter.FeishuAdapter( + PlatformConfig( + extra={ + "app_id": "cli_test", + "app_secret": "secret_test", + "connection_mode": "websocket", + } + ) + ) + + with ( + patch.object(feishu_adapter, "FEISHU_AVAILABLE", False), + patch.object(feishu_adapter, "_load_lark_oapi", return_value=True) as load_sdk, + patch.object(feishu_adapter.asyncio, "to_thread", new_callable=AsyncMock, return_value=True) as to_thread, + patch.object(adapter, "_connect_with_retry", new_callable=AsyncMock), + patch.object(feishu_adapter, "acquire_scoped_lock", return_value=(True, {})), + ): + assert asyncio.run(adapter.connect()) is True + + to_thread.assert_awaited_once_with(load_sdk) diff --git a/tests/gateway/test_goal_continuation_drain.py b/tests/gateway/test_goal_continuation_drain.py index 58c111dde6..662f541fc6 100644 --- a/tests/gateway/test_goal_continuation_drain.py +++ b/tests/gateway/test_goal_continuation_drain.py @@ -120,7 +120,10 @@ async def test_fifo_enqueued_continuation_is_drained_without_new_user_message(): await adapter._process_message_background(event, key) # The in-band drain hands off to a fresh task (#17758); let it run. for _ in range(40): - if len(handled) >= 2: + # Wait for the SENDS, not just the handler calls: the delivery + # ledger hops to worker threads around each send, so the handler + # can return while reply-2's send is still in flight. + if len(adapter.sent) >= 2: break await asyncio.sleep(0.05) diff --git a/tests/gateway/test_matrix.py b/tests/gateway/test_matrix.py index c7348c0cd8..4c02c9385b 100644 --- a/tests/gateway/test_matrix.py +++ b/tests/gateway/test_matrix.py @@ -1128,6 +1128,80 @@ class TestMatrixDeviceId: adapter = MatrixAdapter(config) assert adapter._device_id == "FROM_CONFIG" + @pytest.mark.asyncio + async def test_connect_keeps_configured_device_id_on_adapter(self): + """MATRIX_DEVICE_ID stays on the adapter regardless of whoami. + + Note: this test previously asserted that the configured device_id + overrides the whoami device_id outright. That is no longer true for + the *client* identity — a token can only upload keys for its own + device, so a conflicting whoami device now wins (see + TestCryptoStoreResetOnDeviceChange). The configured value is still + preferred when whoami reports no device, and is still recorded on the + adapter, which is what this test pins. + """ + from plugins.platforms.matrix.adapter import MatrixAdapter + + config = PlatformConfig( + enabled=True, + token="syt_test_access_token", + extra={ + "homeserver": "https://matrix.example.org", + "user_id": "@bot:example.org", + "encryption": True, + "device_id": "MY_STABLE_DEVICE", + }, + ) + adapter = MatrixAdapter(config) + + fake_mautrix_mods = _make_fake_mautrix() + + mock_client = MagicMock() + mock_client.mxid = "@bot:example.org" + mock_client.device_id = None + mock_client.state_store = MagicMock() + mock_client.sync_store = MagicMock() + mock_client.crypto = None + mock_client.whoami = AsyncMock(return_value=MagicMock(user_id="@bot:example.org", device_id="WHOAMI_DEV")) + mock_client.sync = AsyncMock(return_value={"rooms": {"join": {"!room:server": {}}}}) + mock_client.add_event_handler = MagicMock() + mock_client.handle_sync = MagicMock(return_value=[]) + mock_client.query_keys = AsyncMock(return_value={ + "device_keys": {"@bot:example.org": {"MY_STABLE_DEVICE": { + "keys": {"ed25519:MY_STABLE_DEVICE": "fake_ed25519_key"}, + }}}, + }) + mock_client.api = MagicMock() + mock_client.api.token = "syt_test_access_token" + mock_client.api.session = MagicMock() + mock_client.api.session.close = AsyncMock() + + mock_olm = MagicMock() + mock_olm.load = AsyncMock() + mock_olm.share_keys = AsyncMock() + mock_olm.share_keys_min_trust = None + mock_olm.send_keys_min_trust = None + mock_olm.account = MagicMock() + mock_olm.account.identity_keys = {"ed25519": "fake_ed25519_key"} + + fake_mautrix_mods["mautrix.client"].Client = MagicMock(return_value=mock_client) + fake_mautrix_mods["mautrix.crypto"].OlmMachine = MagicMock(return_value=mock_olm) + + import plugins.platforms.matrix.adapter as matrix_mod + with patch.object(matrix_mod, "_check_e2ee_deps", return_value=True): + with patch.dict("sys.modules", fake_mautrix_mods): + with patch.object(adapter, "_refresh_dm_cache", AsyncMock()): + with patch.object(adapter, "_sync_loop", AsyncMock(return_value=None)): + assert await adapter.connect() is True + + # The configured device_id is retained on the adapter. + assert adapter._device_id == "MY_STABLE_DEVICE" + # But the token's own device is what the client claims, because the + # homeserver will not accept key uploads for any other device. + assert mock_client.device_id == "WHOAMI_DEV" + + await adapter.disconnect() + class TestMatrixPasswordLoginDeviceId: """MATRIX_DEVICE_ID should be passed to mautrix Client even with password login.""" @@ -2862,3 +2936,409 @@ class TestMatrixDispatchSyncIsolation: assert ran["ok"] is True # the sibling handler still ran assert "event handler failed" in caplog.text # failure surfaced, not swallowed + + +# --------------------------------------------------------------------------- +# E2EE crypto store reset on device change +# --------------------------------------------------------------------------- + +class TestCryptoStoreResetOnDeviceChange: + @pytest.mark.asyncio + async def test_reset_when_device_id_changed(self, caplog): + import logging + adapter = _make_adapter() + store = MagicMock() + store.get_device_id = AsyncMock(return_value="OLDDEVICE") + store.delete = AsyncMock() + + with caplog.at_level(logging.WARNING): + reset = await adapter._reset_crypto_store_if_device_changed(store, "NEWDEVICE") + + assert reset is True + store.delete.assert_awaited_once() + assert "OLDDEVICE" in caplog.text and "NEWDEVICE" in caplog.text + + @pytest.mark.asyncio + async def test_no_reset_when_device_id_same(self): + adapter = _make_adapter() + store = MagicMock() + store.get_device_id = AsyncMock(return_value="SAMEDEVICE") + store.delete = AsyncMock() + + assert await adapter._reset_crypto_store_if_device_changed(store, "SAMEDEVICE") is False + store.delete.assert_not_awaited() + + @pytest.mark.asyncio + async def test_no_reset_on_fresh_store(self): + adapter = _make_adapter() + store = MagicMock() + store.get_device_id = AsyncMock(return_value=None) + store.delete = AsyncMock() + + assert await adapter._reset_crypto_store_if_device_changed(store, "NEWDEVICE") is False + store.delete.assert_not_awaited() + + @pytest.mark.asyncio + async def test_no_reset_without_device_id(self): + adapter = _make_adapter() + store = MagicMock() + store.get_device_id = AsyncMock(return_value="OLDDEVICE") + store.delete = AsyncMock() + + assert await adapter._reset_crypto_store_if_device_changed(store, "") is False + store.delete.assert_not_awaited() + + @pytest.mark.asyncio + async def test_connect_resets_store_when_token_device_differs_from_config( + self, caplog + ): + """Rotated token, stale MATRIX_DEVICE_ID. + + Persisted store device is A, MATRIX_DEVICE_ID is still A, but the + access token now belongs to device B. The helper alone cannot catch + this: connect() used to resolve client.device_id to the configured A, + so persisted A == live A and no reset happened. The token's device + must win, and the store must be reset. + """ + import logging + from plugins.platforms.matrix.adapter import MatrixAdapter + + config = PlatformConfig( + enabled=True, + token="syt_rotated_access_token", + extra={ + "homeserver": "https://matrix.example.org", + "user_id": "@bot:example.org", + "encryption": True, + "device_id": "DEVICE_A", + }, + ) + adapter = MatrixAdapter(config) + + fake_mautrix_mods = _make_fake_mautrix() + + deleted = {"count": 0} + + class _ResettableCryptoStore: + upgrade_table = MagicMock() + + def __init__(self, account_id="", pickle_key="", db=None): + self.account_id = account_id + self.pickle_key = pickle_key + self.db = db + self._device_id = "DEVICE_A" # persisted from the old token + + async def open(self): + pass + + async def get_device_id(self): + return self._device_id + + async def delete(self): + deleted["count"] += 1 + self._device_id = "" + + async def put_device_id(self, device_id): + self._device_id = device_id + + fake_mautrix_mods[ + "mautrix.crypto.store.asyncpg" + ].PgCryptoStore = _ResettableCryptoStore + + mock_client = MagicMock() + mock_client.mxid = "@bot:example.org" + mock_client.device_id = None + mock_client.state_store = MagicMock() + mock_client.sync_store = MagicMock() + mock_client.crypto = None + # Token was rotated: the homeserver reports device B. + mock_client.whoami = AsyncMock( + return_value=MagicMock(user_id="@bot:example.org", device_id="DEVICE_B") + ) + mock_client.sync = AsyncMock(return_value={"rooms": {"join": {}}}) + mock_client.add_event_handler = MagicMock() + mock_client.handle_sync = MagicMock(return_value=[]) + mock_client.query_keys = AsyncMock(return_value={"device_keys": {}}) + mock_client.api = MagicMock() + mock_client.api.token = "syt_rotated_access_token" + mock_client.api.session = MagicMock() + mock_client.api.session.close = AsyncMock() + + mock_olm = MagicMock() + mock_olm.load = AsyncMock() + mock_olm.share_keys = AsyncMock() + mock_olm.share_keys_min_trust = None + mock_olm.send_keys_min_trust = None + mock_olm.account = MagicMock() + mock_olm.account.identity_keys = {"ed25519": "fake_ed25519_key"} + + fake_mautrix_mods["mautrix.client"].Client = MagicMock( + return_value=mock_client + ) + fake_mautrix_mods["mautrix.crypto"].OlmMachine = MagicMock( + return_value=mock_olm + ) + + import plugins.platforms.matrix.adapter as matrix_mod + + with caplog.at_level(logging.WARNING), patch.object( + matrix_mod, "_check_e2ee_deps", return_value=True + ), patch.dict("sys.modules", fake_mautrix_mods), patch.object( + adapter, "_refresh_dm_cache", AsyncMock() + ), patch.object( + adapter, "_sync_loop", AsyncMock(return_value=None) + ), patch.object( + adapter, "_verify_device_keys_on_server", AsyncMock(return_value=True) + ): + assert await adapter.connect() is True + + # The token's device wins over the stale configured one. + assert mock_client.device_id == "DEVICE_B" + # ...which is what lets the mismatch be seen and the store reset. + assert deleted["count"] == 1 + assert "MATRIX_DEVICE_ID=DEVICE_A" in caplog.text + + await adapter.disconnect() + + +# --------------------------------------------------------------------------- +# Crypto store pickle-key migration +# --------------------------------------------------------------------------- + +class TestCryptoPickleKeyMigration: + @pytest.mark.asyncio + async def test_account_loads_fine_no_migration(self): + adapter = _make_adapter() + store = MagicMock() + store.get_account = AsyncMock(return_value=MagicMock()) + assert await adapter._migrate_legacy_crypto_pickle( + store, MagicMock(), "@bot:example.org", "@bot:example.org:DEV" + ) is True + store.put_account.assert_not_called() + + @pytest.mark.asyncio + async def test_migrates_from_default_pickle_key(self, caplog): + import logging + adapter = _make_adapter() + store = MagicMock() + store.get_account = AsyncMock(side_effect=RuntimeError("BAD_ACCOUNT_KEY")) + store.put_account = AsyncMock() + + legacy_account = MagicMock() + created = [] + + class FakePgCryptoStore: + def __init__(self, account_id, pickle_key, db): + self.pickle_key = pickle_key + created.append(pickle_key) + + async def get_account(self): + if self.pickle_key == "@bot:example.org:default": + return legacy_account + raise RuntimeError("BAD_ACCOUNT_KEY") + + crypto_db = MagicMock() + crypto_db.fetch = AsyncMock(return_value=[]) + crypto_db.execute = AsyncMock() + + fake_mod = types.ModuleType("mautrix.crypto.store.asyncpg") + fake_mod.PgCryptoStore = FakePgCryptoStore + with patch.dict( + sys.modules, + { + "mautrix.crypto.store.asyncpg": fake_mod, + # _repickle_crypto_sessions imports the olm C-extension; + # fake it so this test does not require libolm. + "olm": self._fake_olm_module(), + }, + ), caplog.at_level(logging.INFO): + result = await adapter._migrate_legacy_crypto_pickle( + store, crypto_db, "@bot:example.org", "@bot:example.org:NEWDEV" + ) + + assert result is True + store.put_account.assert_awaited_once_with(legacy_account) + assert "re-pickled crypto store account" in caplog.text + assert "@bot:example.org:default" in created + # session re-pickle pass must sweep all three session tables + queried = " ".join(str(c.args[0]) for c in crypto_db.fetch.await_args_list) + for table in ( + "crypto_olm_session", + "crypto_megolm_inbound_session", + "crypto_megolm_outbound_session", + ): + assert table in queried + + @pytest.mark.asyncio + async def test_unrecoverable_pickle_logs_error(self, caplog): + import logging + adapter = _make_adapter() + store = MagicMock() + store.get_account = AsyncMock(side_effect=RuntimeError("BAD_ACCOUNT_KEY")) + store.put_account = AsyncMock() + + class FakePgCryptoStore: + def __init__(self, account_id, pickle_key, db): + pass + + async def get_account(self): + raise RuntimeError("BAD_ACCOUNT_KEY") + + fake_mod = types.ModuleType("mautrix.crypto.store.asyncpg") + fake_mod.PgCryptoStore = FakePgCryptoStore + with patch.dict(sys.modules, {"mautrix.crypto.store.asyncpg": fake_mod}), \ + caplog.at_level(logging.ERROR): + result = await adapter._migrate_legacy_crypto_pickle( + store, MagicMock(), "@bot:example.org", "@bot:example.org:NEWDEV" + ) + + assert result is False + store.put_account.assert_not_awaited() + assert "cannot be unpickled" in caplog.text + + def _fake_olm_module(self): + """Fake the `olm` C-extension module. + + _repickle_crypto_sessions does `import olm`, which needs libolm. + Sessions unpickle only with the key they were pickled under. + """ + olm_mod = types.ModuleType("olm") + + class _Session: + def __init__(self, key): + self._key = key + + @classmethod + def from_pickle(cls, blob, key): + pickled_under = blob.decode().split("|")[1] + if pickled_under != key: + raise RuntimeError("BAD_ACCOUNT_KEY") + return cls(key) + + def pickle(self, key): + return f"sess|{key}".encode() + + for name in ("Session", "InboundGroupSession", "OutboundGroupSession"): + setattr(olm_mod, name, type(name, (_Session,), {})) + return olm_mod + + @pytest.mark.asyncio + async def test_session_rows_are_repickled_under_current_key(self): + """The session sweep must actually rewrite legacy-key rows.""" + adapter = _make_adapter() + legacy = "@bot:example.org:default" + current = "@bot:example.org:NEWDEV" + + crypto_db = MagicMock() + crypto_db.fetch = AsyncMock( + return_value=[{"session_id": "s1", "session": f"sess|{legacy}".encode()}] + ) + crypto_db.execute = AsyncMock() + + with patch.dict(sys.modules, {"olm": self._fake_olm_module()}): + await adapter._repickle_crypto_sessions( + crypto_db, "@bot:example.org", legacy, current + ) + + # One UPDATE per session table, each writing the current-key blob. + assert crypto_db.execute.await_count == 3 + for call in crypto_db.execute.await_args_list: + assert call.args[1] == f"sess|{current}".encode() + assert call.args[3] == "s1" + + @pytest.mark.asyncio + async def test_rows_already_on_current_key_are_left_alone(self): + adapter = _make_adapter() + current = "@bot:example.org:NEWDEV" + + crypto_db = MagicMock() + crypto_db.fetch = AsyncMock( + return_value=[{"session_id": "s1", "session": f"sess|{current}".encode()}] + ) + crypto_db.execute = AsyncMock() + + with patch.dict(sys.modules, {"olm": self._fake_olm_module()}): + await adapter._repickle_crypto_sessions( + crypto_db, "@bot:example.org", "@bot:example.org:default", current + ) + + crypto_db.execute.assert_not_awaited() + + @pytest.mark.asyncio + async def test_unreadable_rows_are_left_in_place_not_dropped(self, caplog): + """A row readable under neither key is skipped and left untouched. + + The log must not claim the row was dropped when no DELETE is issued. + """ + import logging + adapter = _make_adapter() + + crypto_db = MagicMock() + crypto_db.fetch = AsyncMock( + return_value=[{"session_id": "s1", "session": b"sess|@bot:other:KEY"}] + ) + crypto_db.execute = AsyncMock() + + with patch.dict(sys.modules, {"olm": self._fake_olm_module()}), \ + caplog.at_level(logging.WARNING): + await adapter._repickle_crypto_sessions( + crypto_db, + "@bot:example.org", + "@bot:example.org:default", + "@bot:example.org:NEWDEV", + ) + + crypto_db.execute.assert_not_awaited() + assert "leaving it in place" in caplog.text + assert "dropping" not in caplog.text.lower() + + @pytest.mark.asyncio + async def test_failed_sweep_leaves_account_on_legacy_key_and_retries( + self, caplog + ): + """A sweep failure must not commit the account. + + The account is the migration's commit marker: if it is written first + and the sweep then fails, the next startup takes the current-key fast + path and the remaining legacy-key sessions are stranded permanently. + """ + import logging + adapter = _make_adapter() + legacy_account = MagicMock() + + store = MagicMock() + store.get_account = AsyncMock(side_effect=RuntimeError("BAD_ACCOUNT_KEY")) + store.put_account = AsyncMock() + + class FakePgCryptoStore: + def __init__(self, account_id, pickle_key, db): + self.pickle_key = pickle_key + + async def get_account(self): + if self.pickle_key == "@bot:example.org:default": + return legacy_account + raise RuntimeError("BAD_ACCOUNT_KEY") + + crypto_db = MagicMock() + crypto_db.fetch = AsyncMock(side_effect=RuntimeError("db went away")) + crypto_db.execute = AsyncMock() + + fake_mod = types.ModuleType("mautrix.crypto.store.asyncpg") + fake_mod.PgCryptoStore = FakePgCryptoStore + + with patch.dict( + sys.modules, + { + "mautrix.crypto.store.asyncpg": fake_mod, + "olm": self._fake_olm_module(), + }, + ), caplog.at_level(logging.ERROR): + result = await adapter._migrate_legacy_crypto_pickle( + store, crypto_db, "@bot:example.org", "@bot:example.org:NEWDEV" + ) + + assert result is False + # The critical assertion: the account was NOT committed, so the next + # start still sees a legacy-key account and retries the migration. + store.put_account.assert_not_awaited() + assert "retried on the next start" in caplog.text diff --git a/tests/gateway/test_restart_after_turn.py b/tests/gateway/test_restart_after_turn.py new file mode 100644 index 0000000000..713ab0242f --- /dev/null +++ b/tests/gateway/test_restart_after_turn.py @@ -0,0 +1,39 @@ +"""Unit tests for in-band restart after-turn deferral helpers (#77184).""" + +from gateway.restart import ( + DEFAULT_GATEWAY_RESTART_AFTER_TURN_TIMEOUT, + parse_restart_after_turn_timeout, + resolve_restart_exit_wait_budget, +) +from gateway.run import GatewayRunner + + +def test_parse_restart_after_turn_timeout_defaults_and_clamps(): + assert parse_restart_after_turn_timeout("") == DEFAULT_GATEWAY_RESTART_AFTER_TURN_TIMEOUT + assert parse_restart_after_turn_timeout(None) == DEFAULT_GATEWAY_RESTART_AFTER_TURN_TIMEOUT + assert parse_restart_after_turn_timeout("bogus") == DEFAULT_GATEWAY_RESTART_AFTER_TURN_TIMEOUT + assert parse_restart_after_turn_timeout(0) == 0.0 + assert parse_restart_after_turn_timeout("-5") == 0.0 + assert parse_restart_after_turn_timeout("120") == 120.0 + + +def test_resolve_restart_exit_wait_budget_covers_both_phases(): + assert resolve_restart_exit_wait_budget(0, 0, headroom=15) == 15.0 + assert resolve_restart_exit_wait_budget(180, 21600, headroom=15) == 180 + 21600 + 15 + assert resolve_restart_exit_wait_budget("bad", "bad", headroom="x") == 0.0 + + +def test_load_restart_after_turn_timeout_preserves_zero(tmp_path, monkeypatch): + """Config/env ``0`` must disable after-turn wait, not fall back to default.""" + import gateway.run as gateway_run + + monkeypatch.delenv("HERMES_RESTART_AFTER_TURN_TIMEOUT", raising=False) + monkeypatch.setattr(gateway_run, "_hermes_home", tmp_path) + (tmp_path / "config.yaml").write_text( + "agent:\n restart_after_turn_timeout: 0\n", + encoding="utf-8", + ) + assert GatewayRunner._load_restart_after_turn_timeout() == 0.0 + + monkeypatch.setenv("HERMES_RESTART_AFTER_TURN_TIMEOUT", "0") + assert GatewayRunner._load_restart_after_turn_timeout() == 0.0 diff --git a/tests/gateway/test_restart_drain.py b/tests/gateway/test_restart_drain.py index f9bfcdf656..a44644ddde 100644 --- a/tests/gateway/test_restart_drain.py +++ b/tests/gateway/test_restart_drain.py @@ -102,6 +102,9 @@ async def test_request_restart_is_idempotent(): assert runner._restart_task is not None assert runner._restart_task not in runner._background_tasks assert runner.request_restart(detached=True, via_service=False) is False + # In-band restart marks draining immediately so new turns are refused + # while any after-turn wait runs (#77184). + assert runner._draining is True await runner._restart_task @@ -111,6 +114,69 @@ async def test_request_restart_is_idempotent(): ) +@pytest.mark.asyncio +async def test_request_restart_defers_stop_until_active_turn_finishes(): + """Regression for #77184: requesting turn must not enter the drain set.""" + runner, _adapter = make_restart_runner() + runner.stop = AsyncMock() + runner._launch_detached_restart_command = AsyncMock() + runner._restart_after_turn_timeout = 5.0 + session_key = "agent:main:telegram:dm:123" + runner._running_agents[session_key] = MagicMock() + + assert runner.request_restart(detached=False, via_service=True) is True + assert runner._draining is True + + # While the requesting turn is still active, stop() must not run. + await asyncio.sleep(0.25) + runner.stop.assert_not_awaited() + assert session_key in runner._running_agents + + # Turn finishes → restart proceeds immediately (drain set empty). + del runner._running_agents[session_key] + await runner._restart_task + + runner.stop.assert_awaited_once_with( + restart=True, detached_restart=False, service_restart=True + ) + # Detached helper is only for the non-service path. + runner._launch_detached_restart_command.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_request_restart_after_turn_timeout_zero_enters_stop_immediately(): + """restart_after_turn_timeout=0 preserves legacy immediate drain.""" + runner, _adapter = make_restart_runner() + runner.stop = AsyncMock() + runner._restart_after_turn_timeout = 0.0 + runner._running_agents["agent:main:telegram:dm:1"] = MagicMock() + + assert runner.request_restart(detached=False, via_service=True) is True + await runner._restart_task + + runner.stop.assert_awaited_once_with( + restart=True, detached_restart=False, service_restart=True + ) + + +@pytest.mark.asyncio +async def test_request_restart_after_turn_cap_elapsed_still_calls_stop(): + """Safety valve: wedged turns cannot pin the gateway forever.""" + runner, _adapter = make_restart_runner() + runner.stop = AsyncMock() + runner._restart_after_turn_timeout = 0.2 + runner._running_agents["agent:main:telegram:dm:1"] = MagicMock() + + assert runner.request_restart(detached=False, via_service=True) is True + await runner._restart_task + + runner.stop.assert_awaited_once_with( + restart=True, detached_restart=False, service_restart=True + ) + # Agent was still present — stop() owns the interrupt path from here. + assert runner._running_agents + + @pytest.mark.asyncio async def test_run_restart_excluded_from_stop_cancel_loop(): """Regression for #12875: _run_restart is held on self._restart_task and diff --git a/tests/gateway/test_run_progress_topics.py b/tests/gateway/test_run_progress_topics.py index 1cd5da9e88..b3d646efcd 100644 --- a/tests/gateway/test_run_progress_topics.py +++ b/tests/gateway/test_run_progress_topics.py @@ -59,6 +59,18 @@ class ProgressCaptureAdapter(BasePlatformAdapter): return {"id": chat_id} +class DiscordProgressCaptureAdapter(ProgressCaptureAdapter): + """Capture sends while exercising Discord's real preview formatter.""" + + def __init__(self): + super().__init__(platform=Platform.DISCORD) + + def format_tool_preview(self, preview, **kwargs): + from plugins.platforms.discord.adapter import DiscordAdapter + + return DiscordAdapter.format_tool_preview(self, preview, **kwargs) + + class SmallLimitProgressAdapter(ProgressCaptureAdapter): """Adapter with a tiny platform limit to exercise progress rollover.""" @@ -239,6 +251,28 @@ class LongPreviewAgent: } +class UrlPreviewAgent: + URL = "https://hermes-agent.nousresearch.com/docs/gateway/discord/tool-progress" + + def __init__(self, **kwargs): + self.tool_progress_callback = kwargs.get("tool_progress_callback") + self.tools = [] + + def run_conversation(self, message, conversation_history=None, task_id=None): + self.tool_progress_callback( + "tool.started", + "web_extract", + self.URL, + {"urls": [self.URL]}, + ) + time.sleep(0.35) + return { + "final_response": "done", + "messages": [], + "api_calls": 1, + } + + class DelayedProgressAgent: def __init__(self, **kwargs): self.tool_progress_callback = kwargs.get("tool_progress_callback") @@ -402,6 +436,123 @@ async def test_run_agent_progress_uses_event_message_id_for_slack_dm(monkeypatch assert all(call["metadata"] == expected_metadata for call in adapter.typing) +@pytest.mark.asyncio +async def test_progress_carries_anchor_for_relay_discord_auto_thread(monkeypatch, tmp_path): + """Relay Discord channel-initiate: the thread doesn't exist at ingest, so + the connector auto-threads on the reply anchor and stamps + prospective_thread_id. The tool-progress / status bubbles must carry that + anchor (reply_to + metadata.reply_to_message_id) so they route into the + SAME auto-thread as the final reply — otherwise the search-status updates + leak into the parent channel (staging repro 2026-08-02).""" + monkeypatch.setenv("HERMES_TOOL_PROGRESS_MODE", "all") + import yaml + (tmp_path / "config.yaml").write_text( + yaml.dump({"display": {"platforms": {"discord": {"tool_progress": "all"}}}}), + encoding="utf-8", + ) + + fake_dotenv = types.ModuleType("dotenv") + fake_dotenv.load_dotenv = lambda *args, **kwargs: None + monkeypatch.setitem(sys.modules, "dotenv", fake_dotenv) + + fake_run_agent = types.ModuleType("run_agent") + fake_run_agent.AIAgent = FakeAgent + monkeypatch.setitem(sys.modules, "run_agent", fake_run_agent) + + adapter = ProgressCaptureAdapter(platform=Platform.RELAY) + runner = _make_runner(adapter) + gateway_run = importlib.import_module("gateway.run") + monkeypatch.setattr(gateway_run, "_hermes_home", tmp_path) + monkeypatch.setattr(gateway_run, "_resolve_runtime_agent_kwargs", lambda: {"api_key": "***"}) + + # Channel-initiating message: no thread_id yet, but the connector stamped + # the prospective thread id (== the triggering message id). Relay ingress + # keeps the underlying platform (discord) on the source for display policy, + # but delivery/progress route through the one live RelayAdapter. + source = SessionSource( + platform=Platform.DISCORD, + chat_id="chan-parent", + chat_type="group", + thread_id=None, + prospective_thread_id="msg-anchor-1", + delivered_via_upstream_relay=True, + ) + + result = await runner._run_agent( + message="find me a gift", + context_prompt="", + history=[], + source=source, + session_id="sess-relay-thread", + session_key="agent:main:discord:thread:chan-parent:msg-anchor-1", + event_message_id="msg-anchor-1", + ) + + assert result["final_response"] == "done" + assert adapter.sent, "expected at least one progress send" + # Every progress send must carry the anchor so the connector threads it. + for call in adapter.sent: + assert call["reply_to"] == "msg-anchor-1", call + assert (call["metadata"] or {}).get("reply_to_message_id") == "msg-anchor-1", call + # Discord lifecycle/status sends are marked non-conversational. + assert (call["metadata"] or {}).get("non_conversational") is True, call + + +@pytest.mark.asyncio +async def test_progress_no_anchor_for_native_discord_thread_event(monkeypatch, tmp_path): + """A message ARRIVING in an existing Discord thread (not the relay + auto-thread lane) must NOT get the synthetic prospective anchor — it already + routes by its real thread. Guards against over-broadening the relay fix.""" + monkeypatch.setenv("HERMES_TOOL_PROGRESS_MODE", "all") + import yaml + (tmp_path / "config.yaml").write_text( + yaml.dump({"display": {"platforms": {"discord": {"tool_progress": "all"}}}}), + encoding="utf-8", + ) + + fake_dotenv = types.ModuleType("dotenv") + fake_dotenv.load_dotenv = lambda *args, **kwargs: None + monkeypatch.setitem(sys.modules, "dotenv", fake_dotenv) + + fake_run_agent = types.ModuleType("run_agent") + fake_run_agent.AIAgent = FakeAgent + monkeypatch.setitem(sys.modules, "run_agent", fake_run_agent) + + adapter = ProgressCaptureAdapter(platform=Platform.RELAY) + runner = _make_runner(adapter) + gateway_run = importlib.import_module("gateway.run") + monkeypatch.setattr(gateway_run, "_hermes_home", tmp_path) + monkeypatch.setattr(gateway_run, "_resolve_runtime_agent_kwargs", lambda: {"api_key": "***"}) + + # No prospective_thread_id (event is IN a real thread already). + source = SessionSource( + platform=Platform.DISCORD, + chat_id="real-thread-9", + chat_type="thread", + thread_id="real-thread-9", + delivered_via_upstream_relay=True, + ) + + result = await runner._run_agent( + message="continue", + context_prompt="", + history=[], + source=source, + session_id="sess-in-thread", + session_key="agent:main:discord:thread:real-thread-9:real-thread-9", + event_message_id="msg-2", + ) + + assert result["final_response"] == "done" + # The relay-prospective synthetic anchor path must NOT engage; progress + # routes by the real thread's own metadata, not a forced reply_to anchor. + for call in adapter.sent: + meta = call["metadata"] or {} + # The real thread id drives routing; we did not inject the anchor + # reply_to that the prospective lane uses. + assert meta.get("thread_id") == "real-thread-9" or call["reply_to"] != "msg-2", call + + # --------------------------------------------------------------------------- # Preview truncation tests (all/new mode respects tool_preview_length) # --------------------------------------------------------------------------- @@ -493,6 +644,59 @@ def test_all_mode_respects_custom_preview_length(monkeypatch, tmp_path): assert len(preview_text) <= 120, f"Preview too long ({len(preview_text)}): {preview_text}" +def test_discord_truncated_tool_url_links_to_full_destination(monkeypatch, tmp_path): + """The real gateway path must retain the URL beyond its visible cap.""" + import yaml + + monkeypatch.setenv("HERMES_TOOL_PROGRESS_MODE", "all") + + fake_dotenv = types.ModuleType("dotenv") + fake_dotenv.load_dotenv = lambda *args, **kwargs: None + monkeypatch.setitem(sys.modules, "dotenv", fake_dotenv) + + fake_run_agent = types.ModuleType("run_agent") + fake_run_agent.AIAgent = UrlPreviewAgent + monkeypatch.setitem(sys.modules, "run_agent", fake_run_agent) + + (tmp_path / "config.yaml").write_text( + yaml.dump({"display": {"tool_preview_length": 0}}), + encoding="utf-8", + ) + + adapter = DiscordProgressCaptureAdapter() + runner = _make_runner(adapter) + gateway_run = importlib.import_module("gateway.run") + monkeypatch.setattr(gateway_run, "_hermes_home", tmp_path) + monkeypatch.setattr( + gateway_run, + "_resolve_runtime_agent_kwargs", + lambda: {"api_key": "***"}, + ) + + source = SessionSource( + platform=Platform.DISCORD, + chat_id="12345", + chat_type="dm", + thread_id=None, + ) + result = asyncio.get_event_loop().run_until_complete( + runner._run_agent( + message="hello", + context_prompt="", + history=[], + source=source, + session_id="sess-discord-url", + session_key="agent:main:discord:dm:12345", + ) + ) + + assert result["final_response"] == "done" + assert adapter.sent + visible = UrlPreviewAgent.URL[:37] + "..." + label = visible.removeprefix("https://") + assert f"[{label}](<{UrlPreviewAgent.URL}>)" in adapter.sent[0]["content"] + + class CommentaryAgent: def __init__(self, **kwargs): self.tool_progress_callback = kwargs.get("tool_progress_callback") @@ -1316,5 +1520,3 @@ class TestSlackReplyInThreadProgressRouting: event_message_id="1700000000.000100", reply_in_thread=False, ) is None - - diff --git a/tests/gateway/test_runtime_footer.py b/tests/gateway/test_runtime_footer.py index 3845ffa933..1ca63a90c6 100644 --- a/tests/gateway/test_runtime_footer.py +++ b/tests/gateway/test_runtime_footer.py @@ -133,3 +133,187 @@ def test_build_footer_per_platform_off_suppresses(): assert out == "" + +# --------------------------------------------------------------------------- +# latency — opt-in wall-clock turn duration +# --------------------------------------------------------------------------- + +@pytest.mark.parametrize( + "seconds,expected", + [ + (0.0, "<1s"), + (0.4, "<1s"), + (0.999, "<1s"), + (1.0, "1s"), + (22.0, "22s"), + (22.4, "22s"), + (59.4, "59s"), + (59.6, "1m00s"), + (60.0, "1m00s"), + (65.0, "1m05s"), + (125.0, "2m05s"), + (3600.0, "60m00s"), + ], +) +def test_format_latency(seconds, expected): + from gateway.runtime_footer import _format_latency + + assert _format_latency(seconds) == expected + + +def test_format_footer_latency_renders(): + out = format_runtime_footer( + model="m", + context_tokens=0, + context_length=None, + cwd="", + turn_seconds=22.0, + fields=("latency",), + ) + assert out == "22s" + + +def test_format_footer_latency_skipped_when_unmeasured(): + """A call site that doesn't measure timing leaves the field out entirely.""" + out = format_runtime_footer( + model="m", + context_tokens=0, + context_length=None, + cwd="", + turn_seconds=None, + fields=("latency",), + ) + assert out == "" + + +def test_format_footer_latency_skipped_when_negative(): + """A nonsensical (negative) duration is dropped rather than rendered.""" + out = format_runtime_footer( + model="m", + context_tokens=0, + context_length=None, + cwd="", + turn_seconds=-1.0, + fields=("latency",), + ) + assert out == "" + + +def test_format_footer_latency_zero_renders_sub_second(): + """Zero is a real measurement (a very fast turn), not missing data.""" + out = format_runtime_footer( + model="m", + context_tokens=0, + context_length=None, + cwd="", + turn_seconds=0.0, + fields=("latency",), + ) + assert out == "<1s" + + +def test_format_footer_latency_in_field_order(monkeypatch, tmp_path): + monkeypatch.setenv("HOME", str(tmp_path)) + out = format_runtime_footer( + model="openai/gpt-5.4", + context_tokens=68_000, + context_length=100_000, + cwd=str(tmp_path), + turn_seconds=65.0, + fields=("model", "context_pct", "latency", "cwd"), + ) + assert out == "gpt-5.4 · 68% · 1m05s · ~" + + +def test_build_footer_line_threads_turn_seconds(monkeypatch): + monkeypatch.delenv("TERMINAL_CWD", raising=False) + out = build_footer_line( + user_config={ + "display": { + "runtime_footer": { + "enabled": True, + "fields": ["model", "latency"], + } + } + }, + platform_key="discord", + model="gpt-5.4", + context_tokens=0, + context_length=None, + cwd="", + turn_seconds=22.0, + ) + assert out == "gpt-5.4 · 22s" + + +# --------------------------------------------------------------------------- +# Byte-stability: `latency` is opt-in, so the DEFAULT footer is unchanged. +# +# Upstream doctrine: a system prompt / rendered surface must be byte-stable for +# the life of a conversation. Adding a field to _DEFAULT_FIELDS would silently +# change the footer text of every user who already enabled it. These tests pin +# the default set and the exact default-config output strings. +# --------------------------------------------------------------------------- + +_LEGACY_DEFAULT_FIELDS = ["model", "context_pct", "cwd"] + + +def test_latency_not_in_default_fields(): + from gateway.runtime_footer import _DEFAULT_FIELDS + + assert "latency" not in _DEFAULT_FIELDS + assert list(_DEFAULT_FIELDS) == _LEGACY_DEFAULT_FIELDS + + +def test_resolve_footer_config_default_fields_exclude_latency(): + assert resolve_footer_config({}, "telegram")["fields"] == _LEGACY_DEFAULT_FIELDS + assert resolve_footer_config( + {"display": {"runtime_footer": {"enabled": True}}}, "discord" + )["fields"] == _LEGACY_DEFAULT_FIELDS + + +@pytest.mark.parametrize( + "model,tokens,window,cwd,expected", + [ + ("openai/gpt-5.4", 50_247, 1_000_000, "/var/data", "gpt-5.4 · 5% · /var/data"), + ("claude-opus-4-8", 68_000, 100_000, "/var/data", "claude-opus-4-8 · 68% · /var/data"), + ("m", 0, None, "/var/data", "m · /var/data"), + ("", 10, 100, "/var/data", "10% · /var/data"), + ("m", 10, 100, "", "m · 10%"), + ], +) +def test_default_footer_renders_byte_identically( + monkeypatch, model, tokens, window, cwd, expected +): + """Default-config output is byte-for-byte what it was before `latency`. + + Note `turn_seconds` IS supplied — proving that even when the caller + measures timing, a default-configured footer does not show it. + """ + monkeypatch.delenv("TERMINAL_CWD", raising=False) + out = format_runtime_footer( + model=model, + context_tokens=tokens, + context_length=window, + cwd=cwd, + turn_seconds=22.0, + # fields deliberately NOT passed — exercises the default. + ) + assert out == expected + + +def test_default_build_footer_line_ignores_turn_seconds(monkeypatch): + """build_footer_line with default fields is unaffected by turn_seconds.""" + monkeypatch.delenv("TERMINAL_CWD", raising=False) + common = dict( + user_config={"display": {"runtime_footer": {"enabled": True}}}, + platform_key="discord", + model="openai/gpt-5.4", + context_tokens=50_247, + context_length=1_000_000, + cwd="/var/data", + ) + baseline = build_footer_line(**common) + with_timing = build_footer_line(**common, turn_seconds=125.0) + assert baseline == "gpt-5.4 · 5% · /var/data" + assert with_timing == baseline diff --git a/tests/gateway/test_session.py b/tests/gateway/test_session.py index 5b48ddd22d..9b68bbc352 100644 --- a/tests/gateway/test_session.py +++ b/tests/gateway/test_session.py @@ -1163,15 +1163,39 @@ class TestHasAnySessions: return s def test_uses_database_count_when_available(self, store_with_mock_db): - """has_any_sessions should use database session_count, not len(_entries).""" + """has_any_sessions should use database session_count_ge, not len(_entries).""" store = store_with_mock_db # Simulate single-platform user with only 1 entry in memory store._entries = {"telegram:12345": MagicMock()} # But database has 3 sessions (current + 2 previous resets) - store._db.session_count.return_value = 3 + store._db.session_count_ge.return_value = True assert store.has_any_sessions() is True - store._db.session_count.assert_called_once() + store._db.session_count_ge.assert_called_once_with(2) + + def test_first_session_ever_returns_false(self, store_with_mock_db): + """First session ever should return False (only current session in DB).""" + store = store_with_mock_db + store._entries = {"telegram:12345": MagicMock()} + # Database has exactly 1 session (the current one just created) + store._db.session_count_ge.return_value = False + + assert store.has_any_sessions() is False + + def test_fallback_without_database(self, tmp_path): + """Should fall back to len(_entries) when DB is not available.""" + config = GatewayConfig() + with patch("gateway.session.SessionStore._ensure_loaded"): + store = SessionStore(sessions_dir=tmp_path, config=config) + store._loaded = True + store._db = None + store._entries = {"key1": MagicMock(), "key2": MagicMock()} + + # > 1 entries means has sessions + assert store.has_any_sessions() is True + + store._entries = {"key1": MagicMock()} + assert store.has_any_sessions() is False class TestLastPromptTokens: diff --git a/tests/gateway/test_session_info.py b/tests/gateway/test_session_info.py index a2c0fe2ca7..7d16d6b807 100644 --- a/tests/gateway/test_session_info.py +++ b/tests/gateway/test_session_info.py @@ -64,6 +64,58 @@ class TestFormatSessionInfo: assert "localhost:11434" in info assert "8K" in info + def test_named_custom_provider_keeps_context_pin_without_model_base_url( + self, runner, tmp_path + ): + """Session-reset banner must honor model.context_length for named custom providers. + + Repro: /status shows 262144 from config while the reset banner said + ``131K tokens (detected)`` because empty model.base_url + runtime URL + falsely cleared the pin and fell through to the Qwen family default. + """ + model = "custom-local-agentw/Qwen-AgentWorld-35B-A3B-Q5_K_XL" + config_yaml = ( + "model:\n" + f" default: {model}\n" + " provider: custom-local-agentw\n" + " context_length: 262144\n" + "custom_providers:\n" + " - name: custom-local-agentw\n" + " base_url: http://127.0.0.1:8080/v1\n" + " models: {}\n" + ) + p1, p2, p3 = _patch_info( + tmp_path, + config_yaml, + model, + { + "provider": "custom-local-agentw", + "base_url": "http://127.0.0.1:8080/v1", + "api_key": "", + }, + ) + with p1, p2, p3, patch( + "hermes_cli.config.get_compatible_custom_providers", + return_value=[ + { + "name": "custom-local-agentw", + "base_url": "http://127.0.0.1:8080/v1", + "models": {}, + } + ], + ), patch( + "agent.model_metadata.get_model_context_length", + side_effect=lambda *args, **kwargs: ( + kwargs.get("config_context_length") + if kwargs.get("config_context_length") + else 131072 + ), + ): + info = runner._format_session_info() + assert "262K" in info + assert "config" in info + assert "131K" not in info + class TestResetNoticeSessionInfo: """#59003: the auto-reset banner must report the serving profile's config, diff --git a/tests/gateway/test_session_split_brain_11016.py b/tests/gateway/test_session_split_brain_11016.py index 5a5439f4bd..b6a850244e 100644 --- a/tests/gateway/test_session_split_brain_11016.py +++ b/tests/gateway/test_session_split_brain_11016.py @@ -198,9 +198,13 @@ class TestStaleSessionLockSelfHeal: # An ordinary message should heal the stale lock, then fall through # to normal dispatch. User gets a reply instead of a busy ack. await adapter.handle_message(_make_event("hello")) - # Drain any spawned background tasks. - for _ in range(5): - await asyncio.sleep(0) + # Drain any spawned background tasks. Real sleeps, not bare yields: + # the delivery ledger hops to worker threads around the send, so a + # zero-delay yield loop can finish before the reply lands. + for _ in range(40): + if any("handled:text" in r for r in adapter.sent_responses): + break + await asyncio.sleep(0.05) assert any("handled:text" in r for r in adapter.sent_responses), ( "stale lock trapped a normal message — split-brain not healed" diff --git a/tests/gateway/test_skip_context_files_wiring.py b/tests/gateway/test_skip_context_files_wiring.py new file mode 100644 index 0000000000..eaea6babf8 --- /dev/null +++ b/tests/gateway/test_skip_context_files_wiring.py @@ -0,0 +1,83 @@ +"""Per-platform ``skip_context_files`` gateway wiring (#26860). + +Messaging platforms can opt out of the filesystem-heavy context-file +discovery (SOUL.md, AGENTS.md, .cursorrules walks) that runs during +AIAgent construction — especially impactful on Windows where stat() and +directory walks are 10-100x slower. The agent-side parameters already +exist (agent/agent_init.py); these tests pin the gateway wiring: +config -> signature -> AIAgent kwargs. +""" + +import pytest + +from gateway.run import GatewayRunner + + +class TestSkipContextFilesSignature: + """A toggled skip_context_files must invalidate the agent cache.""" + + RUNTIME = {"provider": "openrouter", "base_url": "", "api_mode": ""} + + def test_signature_differs_when_toggled(self): + sig_off = GatewayRunner._agent_config_signature( + "claude-sonnet-4", self.RUNTIME, ["hermes-telegram"], "", + skip_context_files=False, + ) + sig_on = GatewayRunner._agent_config_signature( + "claude-sonnet-4", self.RUNTIME, ["hermes-telegram"], "", + skip_context_files=True, + ) + assert sig_off != sig_on, ( + "skip_context_files changes the frozen system prompt (context " + "files in vs out) — the cache signature must change with it" + ) + + def test_signature_stable_when_unchanged(self): + sig_a = GatewayRunner._agent_config_signature( + "claude-sonnet-4", self.RUNTIME, ["hermes-telegram"], "", + skip_context_files=True, + ) + sig_b = GatewayRunner._agent_config_signature( + "claude-sonnet-4", self.RUNTIME, ["hermes-telegram"], "", + skip_context_files=True, + ) + assert sig_a == sig_b + + def test_default_matches_explicit_false(self): + """Back-compat: omitting the param must hash like False so existing + cached agents aren't all invalidated by this change.""" + sig_default = GatewayRunner._agent_config_signature( + "claude-sonnet-4", self.RUNTIME, ["hermes-telegram"], "", + ) + sig_false = GatewayRunner._agent_config_signature( + "claude-sonnet-4", self.RUNTIME, ["hermes-telegram"], "", + skip_context_files=False, + ) + assert sig_default == sig_false + + +class TestSkipContextFilesConfigResolution: + """The gateway resolution path: platform config dict -> bool.""" + + @pytest.mark.parametrize( + ("cfg", "platform_key", "expected"), + [ + ({"gateway": {"platforms": {"telegram": {"skip_context_files": True}}}}, "telegram", True), + ({"gateway": {"platforms": {"telegram": {"skip_context_files": False}}}}, "telegram", False), + ({"gateway": {"platforms": {"telegram": {}}}}, "telegram", False), + ({"gateway": {"platforms": {}}}, "telegram", False), + ({"gateway": {}}, "telegram", False), + ({}, "telegram", False), + # Set on a DIFFERENT platform — must not leak. + ({"gateway": {"platforms": {"discord": {"skip_context_files": True}}}}, "telegram", False), + # Truthy non-bool values coerce. + ({"gateway": {"platforms": {"telegram": {"skip_context_files": 1}}}}, "telegram", True), + ], + ) + def test_resolution(self, cfg, platform_key, expected): + # Mirror the production resolution in TurnRunner exactly. + _platforms_gw_cfg = (cfg.get("gateway") or {}).get("platforms") or {} + _plat_gw_cfg = _platforms_gw_cfg.get(platform_key) or {} + _skip_context = _plat_gw_cfg.get("skip_context_files") + skip_context_files = bool(_skip_context) if _skip_context is not None else False + assert skip_context_files is expected diff --git a/tests/gateway/test_sse_agent_cancel.py b/tests/gateway/test_sse_agent_cancel.py index 8e13a519d6..83fabe7c6d 100644 --- a/tests/gateway/test_sse_agent_cancel.py +++ b/tests/gateway/test_sse_agent_cancel.py @@ -7,8 +7,10 @@ task wrapper is cancelled. """ import asyncio -import queue +import threading +import time from unittest.mock import AsyncMock, MagicMock, patch +from gateway.platforms.api_server import ThreadSafeAsyncQueue # --------------------------------------------------------------------------- @@ -44,9 +46,6 @@ class TestSSEAgentCancelOnDisconnect: the agent task must be cancelled.""" adapter = _make_adapter() - stream_q = queue.Queue() - stream_q.put("hello ") # Some data already queued - # Agent task that runs forever (simulates a long LLM call) agent_done = asyncio.Event() @@ -56,6 +55,12 @@ class TestSSEAgentCancelOnDisconnect: async def run(): from aiohttp import web + from gateway.platforms.api_server import ThreadSafeAsyncQueue + + # Constructed inside the running loop — ThreadSafeAsyncQueue + # captures asyncio.get_running_loop() at construction time. + stream_q = ThreadSafeAsyncQueue() + stream_q.put_nowait("hello ") # Some data already queued agent_task = asyncio.ensure_future(fake_agent()) @@ -89,19 +94,54 @@ class TestSSEAgentCancelOnDisconnect: asyncio.run(run()) + def test_agent_task_not_cancelled_on_normal_completion(self): + """On normal stream completion, agent task should NOT be cancelled.""" + adapter = _make_adapter() + + async def fake_agent(): + return {"final_response": "done"}, {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15} + + async def run(): + from aiohttp import web + from gateway.platforms.api_server import ThreadSafeAsyncQueue + + stream_q = ThreadSafeAsyncQueue() + stream_q.put_nowait("hello") + stream_q.put_nowait(None) # End-of-stream sentinel + + agent_task = asyncio.ensure_future(fake_agent()) + await asyncio.sleep(0) # Let agent complete + + mock_response = AsyncMock(spec=web.StreamResponse) + mock_response.write = AsyncMock() + mock_response.prepare = AsyncMock() + + with patch("gateway.platforms.api_server.web.StreamResponse", + return_value=mock_response): + await adapter._write_sse_chat_completion( + _make_request(), "cmpl-456", "gpt-4", 1234567890, + stream_q, agent_task, + ) + + # Agent should have completed normally, not been cancelled + assert agent_task.done() + assert not agent_task.cancelled() + + asyncio.run(run()) def test_broken_pipe_also_cancels_agent(self): """BrokenPipeError (another disconnect variant) also cancels the task.""" adapter = _make_adapter() - stream_q = queue.Queue() - async def fake_agent(): await asyncio.sleep(0.2) # Never completes return {}, {} async def run(): from aiohttp import web + from gateway.platforms.api_server import ThreadSafeAsyncQueue + + stream_q = ThreadSafeAsyncQueue() agent_task = asyncio.ensure_future(fake_agent()) @@ -120,6 +160,132 @@ class TestSSEAgentCancelOnDisconnect: asyncio.run(run()) + def test_already_done_task_not_cancelled_on_disconnect(self): + """If agent already finished before disconnect, don't try to cancel.""" + adapter = _make_adapter() + + async def fake_agent(): + return {"final_response": "done"}, {} + + async def run(): + from aiohttp import web + from gateway.platforms.api_server import ThreadSafeAsyncQueue + + stream_q = ThreadSafeAsyncQueue() + stream_q.put_nowait("data") + + agent_task = asyncio.ensure_future(fake_agent()) + await asyncio.sleep(0) # Let agent complete + + mock_response = AsyncMock(spec=web.StreamResponse) + call_count = 0 + + async def write_side_effect(data): + nonlocal call_count + call_count += 1 + if call_count >= 2: + raise ConnectionResetError("late disconnect") + + mock_response.write = AsyncMock(side_effect=write_side_effect) + mock_response.prepare = AsyncMock() + + with patch("gateway.platforms.api_server.web.StreamResponse", + return_value=mock_response): + await adapter._write_sse_chat_completion( + _make_request(), "cmpl-done", "gpt-4", 1234567890, + stream_q, agent_task, + ) + + # Task was already done — should not be cancelled + assert agent_task.done() + assert not agent_task.cancelled() + + asyncio.run(run()) + + def test_agent_interrupt_called_on_disconnect(self): + """When the client disconnects, agent.interrupt() must be called + so the agent thread stops making LLM API calls.""" + adapter = _make_adapter() + + agent_done = asyncio.Event() + + async def fake_agent(): + await agent_done.wait() + return {"final_response": "done"}, {} + + # Mock agent with an interrupt method + mock_agent = MagicMock() + mock_agent.interrupt = MagicMock() + + async def run(): + from aiohttp import web + from gateway.platforms.api_server import ThreadSafeAsyncQueue + + stream_q = ThreadSafeAsyncQueue() + stream_q.put_nowait("hello ") + + agent_task = asyncio.ensure_future(fake_agent()) + agent_ref = [mock_agent] + + mock_response = AsyncMock(spec=web.StreamResponse) + call_count = 0 + + async def write_side_effect(data): + nonlocal call_count + call_count += 1 + if call_count >= 2: + raise ConnectionResetError("client disconnected") + + mock_response.write = AsyncMock(side_effect=write_side_effect) + mock_response.prepare = AsyncMock() + + with patch("gateway.platforms.api_server.web.StreamResponse", + return_value=mock_response): + await adapter._write_sse_chat_completion( + _make_request(), "cmpl-int", "gpt-4", 1234567890, + stream_q, agent_task, agent_ref, + ) + + # agent.interrupt() must have been called + mock_agent.interrupt.assert_called_once_with("SSE client disconnected") + # Clean up + agent_done.set() + + asyncio.run(run()) + + def test_agent_ref_none_still_cancels_task(self): + """When agent_ref is not provided (None), the task is still cancelled + on disconnect — just without the interrupt() call.""" + adapter = _make_adapter() + + async def fake_agent(): + await asyncio.sleep(999) + return {}, {} + + async def run(): + from aiohttp import web + from gateway.platforms.api_server import ThreadSafeAsyncQueue + + stream_q = ThreadSafeAsyncQueue() + + agent_task = asyncio.ensure_future(fake_agent()) + + mock_response = AsyncMock(spec=web.StreamResponse) + mock_response.write = AsyncMock(side_effect=BrokenPipeError("gone")) + mock_response.prepare = AsyncMock() + + with patch("gateway.platforms.api_server.web.StreamResponse", + return_value=mock_response): + # No agent_ref passed — should still handle disconnect cleanly + await adapter._write_sse_chat_completion( + _make_request(), "cmpl-noref", "gpt-4", 1234567890, + stream_q, agent_task, + ) + + assert agent_task.cancelled() or agent_task.done() + + asyncio.run(run()) + def _capturing_response(): """Mock StreamResponse that records all written SSE bytes as text.""" @@ -161,12 +327,15 @@ class TestSSEAgentFailureFinishReason: def _run(self, fake_agent, queue_items=("partial",)): adapter = _make_adapter() - stream_q = queue.Queue() - for item in queue_items: - stream_q.put(item) - stream_q.put(None) # clean end-of-stream sentinel async def run(): + from gateway.platforms.api_server import ThreadSafeAsyncQueue + + stream_q = ThreadSafeAsyncQueue() + for item in queue_items: + stream_q.put_nowait(item) + stream_q.put_nowait(None) # clean end-of-stream sentinel + agent_task = asyncio.ensure_future(fake_agent()) resp, chunks = _capturing_response() with patch("gateway.platforms.api_server.web.StreamResponse", @@ -188,4 +357,126 @@ class TestSSEAgentFailureFinishReason: assert "error" in finish assert "data: [DONE]" in sse + def test_failed_result_dict_reports_error_not_stop(self): + async def failed(): + return ( + {"final_response": "", "failed": True, "completed": False, + "error": "upstream model 500"}, + {"input_tokens": 5, "output_tokens": 0, "total_tokens": 5}, + ) + reason, finish, _ = self._run(failed) + assert reason == "error" + assert finish.get("hermes", {}).get("failed") is True + + def test_truncated_result_reports_length(self): + async def trunc(): + return ( + {"final_response": "half", "partial": True, "completed": False, + "error": "output was truncated"}, + {"input_tokens": 5, "output_tokens": 3, "total_tokens": 8}, + ) + + reason, finish, _ = self._run(trunc) + assert reason == "length" + assert finish["hermes"]["error_code"] == "output_truncated" + + def test_successful_completion_reports_stop(self): + async def ok(): + return ( + {"final_response": "hi", "completed": True}, + {"input_tokens": 5, "output_tokens": 2, "total_tokens": 7}, + ) + + reason, finish, _ = self._run(ok) + assert reason == "stop" + # No error/hermes pollution on the happy path. + assert "error" not in finish + assert "hermes" not in finish + + +# --------------------------------------------------------------------------- +# Sweeper review fix (teknium1, 2026-07-30): cover the cross-thread +# ``put_threadsafe`` boundary that #72610 introduces via ``ThreadSafeAsyncQueue``. +# ``run_conversation`` runs in a worker thread (``loop.run_in_executor``), +# so its ``_on_delta`` / ``_on_tool_*`` callbacks must be able to push into +# the queue from off the owning event loop and immediately wake the +# consumer ``get()`` — this is the production boundary the original tests +# only exercised with same-loop ``put_nowait``. +# --------------------------------------------------------------------------- + +class TestThreadSafeAsyncQueueCrossThreadBoundary: + """gateway/platforms/api_server.py — ThreadSafeAsyncQueue""" + + def test_worker_thread_put_threadsafe_wakes_owning_loop_get(self): + """A real daemon-Thread calling ``put_threadsafe`` from off-loop must + immediately unblock an ``await q.get()`` on the owning event loop. + This mirrors the ``run_conversation``/``run_in_executor`` boundary.""" + + loop = asyncio.new_event_loop() + + async def consumer(): + q = ThreadSafeAsyncQueue() + + async def wait_for_item(): + return await asyncio.wait_for(q.get(), timeout=2) + + def worker(): + time.sleep(0.05) + # No ``loop=`` kwarg on purpose: production callers + # (_on_delta / _on_tool_*) never pass one, so the queue + # must resolve its own ``_loop_ref``. Passing loop= here + # would make a broken _loop_ref pass this test. + q.put_threadsafe("from-worker") + + thread = threading.Thread(target=worker, daemon=True) + thread.start() + + got = await wait_for_item() + assert got == "from-worker" + + thread.join(timeout=2) + assert not thread.is_alive() + + loop.run_until_complete(consumer()) + loop.close() + + def test_twenty_concurrent_threads_no_drop(self): + """Twenty concurrent off-loop ``put_threadsafe`` calls — all arrive. + Regression test for the #65003 producer path.""" + + loop = asyncio.new_event_loop() + n = 20 + + async def consumer(): + q = ThreadSafeAsyncQueue() + received = [] + + async def drain(): + for _ in range(n): + received.append(await q.get()) + + def worker(idx): + time.sleep(0.01 + idx * 0.002) + # No ``loop=`` kwarg — exercise the production + # ``_loop_ref`` resolution path (see the note above). + q.put_threadsafe(f"item-{idx}") + + threads = [ + threading.Thread(target=worker, args=(i,), daemon=True) + for i in range(n) + ] + for t in threads: + t.start() + + await asyncio.wait_for(drain(), timeout=5) + + for t in threads: + t.join(timeout=2) + assert not t.is_alive() + + assert len(received) == n + assert set(received) == {f"item-{i}" for i in range(n)} + + loop.run_until_complete(consumer()) + loop.close() diff --git a/tests/gateway/test_sse_frame.py b/tests/gateway/test_sse_frame.py new file mode 100644 index 0000000000..7aca701137 --- /dev/null +++ b/tests/gateway/test_sse_frame.py @@ -0,0 +1,83 @@ +"""Byte-contract tests for the shared ``_sse_frame`` SSE encoder. + +``_sse_frame`` is the single source of truth for SSE frame serialization +across ``_write_sse_chat_completion``, ``_write_sse_responses._write_event``, +and the ``/v1/runs`` event stream. These tests assert the *invariant* that +``_sse_frame`` reproduces the exact on-the-wire bytes the inline encoders +used to emit — not a snapshot of a frozen value. If a writer's bytes ever +diverge, a real client breaks, so we pin the relationship, not a literal. +""" + +import json + +from gateway.platforms.api_server import _sse_frame + + +def _inline_frame(data, *, event=None): + """Reproduce the historical inline SSE encoder (pre-dedup).""" + prefix = f"event: {event}\n" if event else "" + return f"{prefix}data: {json.dumps(data)}\n\n".encode() + + +def test_sse_frame_matches_inline_encoder_no_event(): + for data in ( + {"id": "c1", "choices": [{"delta": {"role": "assistant"}}]}, + {"event": "ping", "sequence_number": 1}, + {"text": "plain ascii"}, + ): + assert _sse_frame(data) == _inline_frame(data) + + +def test_sse_frame_matches_inline_encoder_with_event(): + for event, data in ( + ("hermes.tool.progress", {"name": "x", "status": "running"}), + ("response.created", {"id": "r1", "status": "in_progress"}), + ): + assert _sse_frame(data, event=event) == _inline_frame(data, event=event) + + +def test_sse_frame_event_line_shape(): + out = _sse_frame({"a": 1}, event="my.event") + assert out.startswith(b"event: my.event\n") + assert b"data: " in out + assert out.endswith(b"\n\n") + + +def test_sse_frame_default_ensure_ascii_matches_bare_json(): + payload = {"text": "café — Münchner 🏔"} + # Default must equal a bare json.dumps (the original writers used no + # ensure_ascii override), so existing byte streams are unchanged. + assert _sse_frame(payload) == _inline_frame(payload) + assert _sse_frame(payload) == f"data: {json.dumps(payload)}\n\n".encode() + + +def test_sse_frame_ensure_ascii_false_preserves_raw_bytes(): + payload = {"text": "café — Münchner 🏔"} + raw = _sse_frame(payload, ensure_ascii=False) + assert "café" in raw.decode("utf-8") + assert raw != _sse_frame(payload) # different bytes from the default + + +def test_sse_frame_ensure_ascii_false_reproduces_session_event_stream(): + """The session event stream (api_server.py:~2236) historically used + ``json.dumps(payload, ensure_ascii=False)`` + ``.encode('utf-8')`` — the + one genuinely unicode-distinct SSE writer. _sse_frame(event=name, + ensure_ascii=False) must reproduce its exact bytes, raw non-ASCII included. + """ + + def old_session(name, payload): + data = json.dumps(payload, ensure_ascii=False) + return f"event: {name}\ndata: {data}\n\n".encode("utf-8") + + for name, payload in ( + ("session.update", {"text": "café — Münchner 🏔", "id": 1}), + ("thread.message.delta", {"content": "héllo wörld ✓", "seq": 3}), + ): + assert _sse_frame(payload, event=name, ensure_ascii=False) == old_session(name, payload) + + +def test_sse_frame_typed_object_roundtrip(): + obj = {"id": "x", "choices": [{"index": 0, "delta": {"content": "hi"}}]} + out = _sse_frame(obj) + line = out.decode().split("data: ", 1)[1].strip() + assert json.loads(line) == obj diff --git a/tests/gateway/test_telegram_conflict.py b/tests/gateway/test_telegram_conflict.py index 46ab719521..711f874854 100644 --- a/tests/gateway/test_telegram_conflict.py +++ b/tests/gateway/test_telegram_conflict.py @@ -142,6 +142,80 @@ async def test_polling_conflict_retries_before_fatal(monkeypatch): await _cancel_heartbeat(adapter) +@pytest.mark.asyncio +async def test_conflict_retry_drops_pending_updates(monkeypatch): + """Conflict recovery must use drop_pending_updates=True (#75017). + + Without this, each retry starts a new getUpdates session that + immediately gets 409'd by the previous still-expiring session, + creating the very conflict we are trying to recover from. + """ + adapter = TelegramAdapter(PlatformConfig(enabled=True, token="***")) + adapter.set_fatal_error_handler(AsyncMock()) + adapter._drain_polling_connections = AsyncMock() + monkeypatch.setattr("asyncio.sleep", AsyncMock()) + + captured = {} + + async def fake_start_polling(**kwargs): + captured["drop_pending_updates"] = kwargs.get("drop_pending_updates") + + updater = SimpleNamespace( + start_polling=AsyncMock(side_effect=fake_start_polling), + stop=AsyncMock(), + running=True, + ) + adapter._app = SimpleNamespace(updater=updater) + + conflict = type("Conflict", (Exception,), {}) + await adapter._handle_polling_conflict( + conflict("Conflict: terminated by other getUpdates request") + ) + + assert captured.get("drop_pending_updates") is True, ( + "Conflict retry must use drop_pending_updates=True to terminate " + "stale getUpdates sessions on Telegram's servers (#75017)" + ) + + +@pytest.mark.asyncio +async def test_conflict_retry_progress_does_not_reset_retry_ladder(monkeypatch): + """First getUpdates progress after a conflict retry is not durable recovery. + + Telegram can accept the first long-poll after a retry and then return a 409 + from the still-expiring previous session. That transient success must not + reset the retry counter back to 0, or every new 409 looks like attempt 1/5 + and the backoff never reaches the server-side expiry window. + """ + adapter = TelegramAdapter(PlatformConfig(enabled=True, token="***")) + adapter.set_fatal_error_handler(AsyncMock()) + adapter._drain_polling_connections = AsyncMock() + monkeypatch.setattr("asyncio.sleep", AsyncMock()) + + calls = {"n": 0} + + async def fake_start_polling(**_kwargs): + calls["n"] += 1 + adapter._record_polling_progress(adapter._polling_generation) + + updater = SimpleNamespace( + start_polling=AsyncMock(side_effect=fake_start_polling), + stop=AsyncMock(), + running=True, + ) + adapter._app = SimpleNamespace(updater=updater) + + conflict = type("Conflict", (Exception,), {}) + await adapter._handle_polling_conflict( + conflict("Conflict: terminated by other getUpdates request") + ) + + assert calls["n"] == 1 + assert adapter._polling_conflict_count == 1 + assert adapter._polling_conflict_recovery_generation is None + assert adapter._send_path_degraded is False + + @pytest.mark.asyncio async def test_polling_conflict_becomes_fatal_after_retries(monkeypatch): """After exhausting retries, the conflict should become fatal.""" diff --git a/tests/gateway/test_telegram_start_polling_timeout.py b/tests/gateway/test_telegram_start_polling_timeout.py index bf5cce4a52..8228eaae97 100644 --- a/tests/gateway/test_telegram_start_polling_timeout.py +++ b/tests/gateway/test_telegram_start_polling_timeout.py @@ -59,6 +59,7 @@ def _bare_adapter(): a._fatal_error_retryable = True a._polling_network_error_count = 0 a._polling_conflict_count = 0 + a._polling_conflict_recovery_generation = None a._polling_error_callback_ref = None a._background_tasks = set() a._send_path_degraded = False diff --git a/tests/gateway/test_voice_command.py b/tests/gateway/test_voice_command.py index 4a988a6e03..6f1a1d7e3d 100644 --- a/tests/gateway/test_voice_command.py +++ b/tests/gateway/test_voice_command.py @@ -767,6 +767,49 @@ class TestDiscordVoiceChannelMethods: adapter._is_allowed_user.assert_called_once_with("42", guild=adapter._client.get_guild(111), is_dm=False) + @pytest.mark.asyncio + async def test_disconnect_leaves_voice_before_cancelling_bot_task(self): + """Voice must be torn down while the gateway websocket is still alive. + + VoiceClient.disconnect() sends a voice state update over the main gateway + connection and waits for the voice socket to close. The bot task is the + loop running that connection, so cancelling it first strands the + handshake and the disconnect blocks until the caller's shutdown timeout. + """ + adapter = self._make_adapter() + events = [] + + async def cancel_liveness_task(): + events.append("cancel_liveness_task") + + async def cancel_bot_task(): + events.append("cancel_bot_task") + + async def leave_voice_channel(guild_id): + events.append(f"leave_voice_channel:{guild_id}") + + async def close(): + events.append("close_client") + + adapter._cancel_liveness_task = cancel_liveness_task + adapter._cancel_bot_task = cancel_bot_task + adapter.leave_voice_channel = leave_voice_channel + adapter._client.close = close + adapter._voice_clients[111] = MagicMock() + adapter._ready_event = MagicMock() + adapter._post_connect_task = None + adapter._missed_message_backfill_task = None + + await adapter.disconnect() + + assert events == [ + "cancel_liveness_task", + "leave_voice_channel:111", + "cancel_bot_task", + "close_client", + ] + + @pytest.mark.asyncio async def test_get_user_voice_channel_success(self): adapter = self._make_adapter() diff --git a/tests/hermes_cli/test_api_key_providers.py b/tests/hermes_cli/test_api_key_providers.py index d22d81eec5..619cc74c45 100644 --- a/tests/hermes_cli/test_api_key_providers.py +++ b/tests/hermes_cli/test_api_key_providers.py @@ -562,7 +562,103 @@ class TestHasAnyProviderConfigured: from hermes_cli.main import _has_any_provider_configured assert _has_any_provider_configured() is True + @staticmethod + def _clear_provider_env(monkeypatch): + """Clear every provider env var so early checks can't short-circuit.""" + from hermes_cli.auth import PROVIDER_REGISTRY + _all_vars = {"OPENROUTER_API_KEY", "OPENAI_API_KEY", "ANTHROPIC_API_KEY", + "ANTHROPIC_TOKEN", "OPENAI_BASE_URL"} + for pconfig in PROVIDER_REGISTRY.values(): + if pconfig.auth_type == "api_key": + _all_vars.update(pconfig.api_key_env_vars) + for var in _all_vars: + monkeypatch.delenv(var, raising=False) + def _setup_home(self, monkeypatch, tmp_path): + from hermes_cli import config as config_module + hermes_home = tmp_path / ".hermes" + hermes_home.mkdir() + monkeypatch.setattr(config_module, "get_env_path", lambda: hermes_home / ".env") + monkeypatch.setattr(config_module, "get_hermes_home", lambda: hermes_home) + monkeypatch.setenv("HERMES_HOME", str(hermes_home)) + self._clear_provider_env(monkeypatch) + return hermes_home + + def test_config_provider_skips_registry_sweep(self, monkeypatch, tmp_path): + """model.provider in config.yaml must short-circuit BEFORE the slow + provider-registry sweep (gh subprocess etc.) is ever invoked. + + Regression test for the auth-first ordering: get_auth_status is + booby-trapped to fail loudly if the sweep runs. The sweep wraps its + loop in ``except Exception``, so we also record every call — any + recorded call proves the sweep ran even if the raise was swallowed. + """ + import yaml + hermes_home = self._setup_home(monkeypatch, tmp_path) + (hermes_home / "config.yaml").write_text(yaml.dump({ + "model": {"default": "anthropic/claude-opus-4.6", "provider": "openrouter"}, + })) + sweep_calls = [] + + def _trap(provider_id): + sweep_calls.append(provider_id) + raise AssertionError("sweep must be skipped") + + monkeypatch.setattr("hermes_cli.auth.get_auth_status", _trap) + from hermes_cli.main import _has_any_provider_configured + assert _has_any_provider_configured() is True + assert sweep_calls == [], ( + f"provider registry sweep ran before config short-circuit: {sweep_calls}" + ) + + def test_config_base_url_api_key_skips_registry_sweep(self, monkeypatch, tmp_path): + """Custom endpoint (base_url/api_key in config, no provider) must also + short-circuit before the registry sweep.""" + import yaml + hermes_home = self._setup_home(monkeypatch, tmp_path) + (hermes_home / "config.yaml").write_text(yaml.dump({ + "model": { + "default": "local/custom-model", + "base_url": "http://localhost:8000/v1", + "api_key": "sk-local-test", + }, + })) + sweep_calls = [] + + def _trap(provider_id): + sweep_calls.append(provider_id) + raise AssertionError("sweep must be skipped") + + monkeypatch.setattr("hermes_cli.auth.get_auth_status", _trap) + from hermes_cli.main import _has_any_provider_configured + assert _has_any_provider_configured() is True + assert sweep_calls == [], ( + f"provider registry sweep ran before config short-circuit: {sweep_calls}" + ) + + def test_auth_json_skips_registry_sweep(self, monkeypatch, tmp_path): + """auth.json with a logged-in active provider must short-circuit before + the registry sweep. get_auth_status may be called ONLY for the active + provider from auth.json — any other provider id means the sweep ran. + """ + import json + hermes_home = self._setup_home(monkeypatch, tmp_path) + (hermes_home / "auth.json").write_text(json.dumps({ + "active_provider": "nous", + })) + calls = [] + + def _guarded_status(provider_id): + calls.append(provider_id) + assert provider_id == "nous", "sweep must be skipped" + return {"logged_in": True} + + monkeypatch.setattr("hermes_cli.auth.get_auth_status", _guarded_status) + from hermes_cli.main import _has_any_provider_configured + assert _has_any_provider_configured() is True + assert calls == ["nous"], ( + f"provider registry sweep ran before auth.json short-circuit: {calls}" + ) # ============================================================================= @@ -643,6 +739,98 @@ class TestZaiEndpointAutoDetect: assert creds["api_key"] == "" +class TestZaiParallelProbe: + """detect_zai_endpoint probes endpoints in parallel workers. + + Contract under test: (1) each endpoint worker preserves the per-endpoint + candidate-model fallback loop, (2) when several endpoints succeed the + winner is chosen by ZAI_ENDPOINTS priority order (not completion order). + """ + + def _mock_post(self, ok): + """Return an httpx.post replacement; `ok` maps (base_url, model) -> bool.""" + import httpx as _httpx + + def _post(url, headers=None, json=None, timeout=None): + base = url.rsplit("/chat/completions", 1)[0] + code = 200 if ok.get((base, json["model"])) else 401 + request = _httpx.Request("POST", url) + return _httpx.Response(code, request=request, json={}) + + return _post + + def test_candidate_model_fallback_within_endpoint(self, monkeypatch): + """A worker must try its endpoint's later candidate models when the + first ones fail — the fallback the scalar-model version dropped.""" + from hermes_cli.auth import ZAI_ENDPOINTS, detect_zai_endpoint + + coding_global = next(ep for ep in ZAI_ENDPOINTS if ep[0] == "coding-global") + base = coding_global[1] + last_model = coding_global[2][-1] + # Only the LAST candidate model of coding-global succeeds. + monkeypatch.setattr( + "hermes_cli.auth.httpx.post", + self._mock_post({(base, last_model): True}), + ) + result = detect_zai_endpoint("test-key", timeout=1.0) + assert result is not None + assert result["id"] == "coding-global" + assert result["model"] == last_model + + def test_priority_order_wins_over_completion_order(self, monkeypatch): + """When multiple endpoints accept the key, the FIRST in + ZAI_ENDPOINTS order must win, even if another finishes earlier.""" + import time as _time + + from hermes_cli.auth import ZAI_ENDPOINTS, detect_zai_endpoint + + first = ZAI_ENDPOINTS[0] + last = ZAI_ENDPOINTS[-1] + ok = { + (first[1], first[2][0]): True, + (last[1], last[2][0]): True, + } + inner = self._mock_post(ok) + + def _slow_first(url, headers=None, json=None, timeout=None): + if url.startswith(first[1]): + _time.sleep(0.15) # first-priority endpoint finishes LAST + return inner(url, headers=headers, json=json, timeout=timeout) + + monkeypatch.setattr("hermes_cli.auth.httpx.post", _slow_first) + result = detect_zai_endpoint("test-key", timeout=1.0) + assert result is not None + assert result["id"] == first[0] + + def test_all_fail_returns_none(self, monkeypatch): + from hermes_cli.auth import detect_zai_endpoint + + monkeypatch.setattr("hermes_cli.auth.httpx.post", self._mock_post({})) + assert detect_zai_endpoint("bad-key", timeout=1.0) is None + + def test_early_exit_does_not_wait_for_slow_losers(self, monkeypatch): + """When the highest-priority endpoint succeeds fast, the caller must + return without waiting for slow lower-priority probes to finish.""" + import time as _time + + from hermes_cli.auth import ZAI_ENDPOINTS, detect_zai_endpoint + + first = ZAI_ENDPOINTS[0] + inner = self._mock_post({(first[1], first[2][0]): True}) + + def _slow_losers(url, headers=None, json=None, timeout=None): + if not url.startswith(first[1]): + _time.sleep(2.0) # slow lower-priority endpoints + return inner(url, headers=headers, json=json, timeout=timeout) + + monkeypatch.setattr("hermes_cli.auth.httpx.post", _slow_losers) + t0 = _time.perf_counter() + result = detect_zai_endpoint("test-key", timeout=5.0) + elapsed = _time.perf_counter() - t0 + assert result is not None and result["id"] == first[0] + assert elapsed < 1.5, f"early exit failed: waited {elapsed:.2f}s for losers" + + # ============================================================================= # Kimi / Moonshot model list isolation tests # ============================================================================= diff --git a/tests/hermes_cli/test_backup_stability.py b/tests/hermes_cli/test_backup_stability.py new file mode 100644 index 0000000000..d461d50722 --- /dev/null +++ b/tests/hermes_cli/test_backup_stability.py @@ -0,0 +1,105 @@ +from __future__ import annotations + +import json +from pathlib import Path + +import pytest + +from hermes_cli.backup import ( + BackupInProgressError, + _atomic_output_path, + _backup_operation_lock, + _write_full_zip_backup, + create_quick_snapshot, + list_quick_snapshots, +) + + +def test_backup_lock_rejects_a_second_operation(tmp_path) -> None: + home = tmp_path / ".hermes" + home.mkdir() + + with _backup_operation_lock(home): + with pytest.raises(BackupInProgressError): + with _backup_operation_lock(home, timeout_seconds=0): + raise AssertionError("second backup unexpectedly acquired the lock") + + +def test_atomic_output_publishes_only_after_clean_close(tmp_path) -> None: + final = tmp_path / "backup.zip" + final.write_bytes(b"previous") + + with _atomic_output_path(final) as partial: + partial.write_bytes(b"complete") + assert final.read_bytes() == b"previous" + + assert final.read_bytes() == b"complete" + assert not partial.exists() + + +def test_atomic_output_keeps_previous_file_after_failure(tmp_path) -> None: + final = tmp_path / "backup.zip" + final.write_bytes(b"previous") + + with pytest.raises(RuntimeError): + with _atomic_output_path(final) as partial: + partial.write_bytes(b"incomplete") + raise RuntimeError("compression failed") + + assert final.read_bytes() == b"previous" + assert not partial.exists() + + +def test_quick_snapshot_is_published_with_manifest(tmp_path, monkeypatch) -> None: + home = tmp_path / ".hermes" + home.mkdir() + (home / "config.yaml").write_text("model: {}\n", encoding="utf-8") + published: list[tuple[Path, Path]] = [] + + from hermes_cli import backup + + real_replace = backup.os.replace + + def replace(source, destination) -> None: + source_path = Path(source) + destination_path = Path(destination) + if destination_path.parent == home / "state-snapshots": + assert source_path.name.endswith(".partial") + assert (source_path / "manifest.json").is_file() + assert not destination_path.exists() + published.append((source_path, destination_path)) + real_replace(source, destination) + + monkeypatch.setattr(backup.os, "replace", replace) + snapshot_id = create_quick_snapshot(hermes_home=home) + + assert snapshot_id is not None + assert len(published) == 1 + manifest = json.loads( + (home / "state-snapshots" / snapshot_id / "manifest.json").read_text(encoding="utf-8") + ) + assert manifest["id"] == snapshot_id + assert manifest["files"] == {"config.yaml": 10} + + +def test_quick_snapshot_listing_ignores_partial_directories(tmp_path) -> None: + home = tmp_path / ".hermes" + partial = home / "state-snapshots" / ".unfinished.1.partial" + partial.mkdir(parents=True) + (partial / "manifest.json").write_text('{"id":"unfinished"}', encoding="utf-8") + + assert list_quick_snapshots(hermes_home=home) == [] + + +def test_failed_automatic_backup_preserves_previous_archive(tmp_path, monkeypatch) -> None: + home = tmp_path / ".hermes" + home.mkdir() + (home / "state.db").write_bytes(b"not-a-database") + archive = tmp_path / "automatic.zip" + archive.write_bytes(b"previous-valid-backup") + + monkeypatch.setattr("hermes_cli.backup._safe_copy_db", lambda _src, _dst: False) + + assert _write_full_zip_backup(archive, home) is None + assert archive.read_bytes() == b"previous-valid-backup" + assert list(tmp_path.glob(".*.partial")) == [] diff --git a/tests/hermes_cli/test_context_switch_guard.py b/tests/hermes_cli/test_context_switch_guard.py index cfd2e80a9a..65cfb4d83d 100644 --- a/tests/hermes_cli/test_context_switch_guard.py +++ b/tests/hermes_cli/test_context_switch_guard.py @@ -79,7 +79,7 @@ def test_custom_provider_context_avoids_false_shrink_warning(monkeypatch): "name": "qwen-token-plan", "base_url": "https://token-plan.example/compatible-mode/v1", "models": { - "qwen3.8-max-preview": {"context_length": 1_048_576}, + "qwen3.9-max-preview": {"context_length": 1_048_576}, }, } ] @@ -110,7 +110,7 @@ def test_custom_provider_context_avoids_false_shrink_warning(monkeypatch): ) result = ModelSwitchResult( success=True, - new_model="qwen3.8-max-preview", + new_model="qwen3.9-max-preview", target_provider="qwen-token-plan", provider_changed=True, api_key="k", @@ -131,7 +131,7 @@ def test_custom_provider_context_avoids_false_shrink_warning(monkeypatch): # Agent snapshot alone (classic CLI historically forgot to pass the kwarg). result2 = ModelSwitchResult( success=True, - new_model="qwen3.8-max-preview", + new_model="qwen3.9-max-preview", target_provider="qwen-token-plan", provider_changed=True, api_key="k", @@ -156,7 +156,7 @@ def test_custom_provider_context_avoids_false_shrink_warning(monkeypatch): ) result3 = ModelSwitchResult( success=True, - new_model="qwen3.8-max-preview", + new_model="qwen3.9-max-preview", target_provider="qwen-token-plan", provider_changed=True, api_key="k", diff --git a/tests/hermes_cli/test_copilot_auth.py b/tests/hermes_cli/test_copilot_auth.py index d696806a36..1023bbfa95 100644 --- a/tests/hermes_cli/test_copilot_auth.py +++ b/tests/hermes_cli/test_copilot_auth.py @@ -40,6 +40,37 @@ class TestResolveToken: with pytest.raises(ValueError, match="classic PAT"): resolve_copilot_token() + def test_invalid_env_var_skips_gh_cli_fallback(self, monkeypatch): + """When an env var is set but holds an unsupported classic PAT, + resolve_copilot_token must NOT fall back to ``gh auth token``. + + The user explicitly exported a token; silently substituting one + from the gh CLI credential store is surprising and the subprocess + call adds up to 5s of latency on Windows cold starts (#60800). + Only fall back to the CLI when NO Copilot env var is set at all. + """ + from hermes_cli.copilot_auth import resolve_copilot_token + monkeypatch.delenv("COPILOT_GITHUB_TOKEN", raising=False) + monkeypatch.delenv("GH_TOKEN", raising=False) + monkeypatch.setenv("GITHUB_TOKEN", "ghp_classic_pat_nope") + with patch("hermes_cli.copilot_auth._try_gh_cli_token") as mock_cli: + token, source = resolve_copilot_token() + assert token == "" + assert source == "" + mock_cli.assert_not_called() + + def test_all_env_vars_invalid_skips_gh_cli_fallback(self, monkeypatch): + """All three env vars set to classic PATs → no gh CLI call.""" + from hermes_cli.copilot_auth import resolve_copilot_token + monkeypatch.setenv("COPILOT_GITHUB_TOKEN", "ghp_one") + monkeypatch.setenv("GH_TOKEN", "ghp_two") + monkeypatch.setenv("GITHUB_TOKEN", "ghp_three") + with patch("hermes_cli.copilot_auth._try_gh_cli_token") as mock_cli: + token, source = resolve_copilot_token() + assert token == "" + assert source == "" + mock_cli.assert_not_called() + class TestRequestHeaders: """Copilot API header generation.""" diff --git a/tests/hermes_cli/test_copilot_context.py b/tests/hermes_cli/test_copilot_context.py index d7bf952fcd..8914537de8 100644 --- a/tests/hermes_cli/test_copilot_context.py +++ b/tests/hermes_cli/test_copilot_context.py @@ -81,6 +81,125 @@ class TestGetCopilotModelContext: assert mock_fetch.call_count == 2 + + @patch("hermes_cli.models._urlopen_model_catalog_request") + def test_fetch_github_model_catalog_uses_short_lived_cache(self, mock_urlopen): + import json as _json + import hermes_cli.models as mod + + mod._github_model_catalog_cache = None + mod._github_model_catalog_cache_key = None + mod._github_model_catalog_cache_time = 0.0 + + payload = { + "data": [ + { + "id": "gpt-4.1", + "model_picker_enabled": True, + "supported_endpoints": ["/chat/completions"], + } + ] + } + + class _Resp: + def __enter__(self): + return self + def __exit__(self, *args): + return False + def read(self): + return _json.dumps(payload).encode() + + mock_urlopen.return_value = _Resp() + + first = mod.fetch_github_model_catalog(api_key="token") + second = mod.fetch_github_model_catalog(api_key="token") + + assert [item["id"] for item in first] == ["gpt-4.1"] + assert [item["id"] for item in second] == ["gpt-4.1"] + assert mock_urlopen.call_count == 1 + + # Cached copies are independent — mutating the result must not + # poison the cache. + second[0]["id"] = "mutated" + third = mod.fetch_github_model_catalog(api_key="token") + assert [item["id"] for item in third] == ["gpt-4.1"] + assert mock_urlopen.call_count == 1 + + @patch("hermes_cli.models._urlopen_model_catalog_request") + def test_fetch_github_model_catalog_cache_expires_after_ttl(self, mock_urlopen): + import json as _json + import time as _time + import hermes_cli.models as mod + + mod._github_model_catalog_cache = None + mod._github_model_catalog_cache_key = None + mod._github_model_catalog_cache_time = 0.0 + + payload = { + "data": [ + { + "id": "gpt-4.1", + "model_picker_enabled": True, + "supported_endpoints": ["/chat/completions"], + } + ] + } + + class _Resp: + def __enter__(self): + return self + def __exit__(self, *args): + return False + def read(self): + return _json.dumps(payload).encode() + + mock_urlopen.return_value = _Resp() + + mod.fetch_github_model_catalog(api_key="token") + assert mock_urlopen.call_count == 1 + + # Age the entry past the TTL (monotonic clock) — next call re-fetches. + mod._github_model_catalog_cache_time = ( + _time.monotonic() - mod._GITHUB_MODEL_CATALOG_CACHE_TTL - 1 + ) + mod.fetch_github_model_catalog(api_key="token") + assert mock_urlopen.call_count == 2 + + @patch("hermes_cli.models._urlopen_model_catalog_request") + def test_fetch_github_model_catalog_cache_misses_on_credential_change(self, mock_urlopen): + import json as _json + import hermes_cli.models as mod + + mod._github_model_catalog_cache = None + mod._github_model_catalog_cache_key = None + mod._github_model_catalog_cache_time = 0.0 + + payload = { + "data": [ + { + "id": "gpt-4.1", + "model_picker_enabled": True, + "supported_endpoints": ["/chat/completions"], + } + ] + } + + class _Resp: + def __enter__(self): + return self + def __exit__(self, *args): + return False + def read(self): + return _json.dumps(payload).encode() + + mock_urlopen.return_value = _Resp() + + mod.fetch_github_model_catalog(api_key="token-a") + assert mock_urlopen.call_count == 1 + # A different token must not be served the previous account's catalog. + mod.fetch_github_model_catalog(api_key="token-b") + assert mock_urlopen.call_count == 2 + @patch("hermes_cli.models.fetch_github_model_catalog", return_value=[]) def test_returns_none_for_empty_catalog(self, mock_fetch): assert get_copilot_model_context("gpt-4.1") is None diff --git a/tests/hermes_cli/test_dashboard_param_clamps.py b/tests/hermes_cli/test_dashboard_param_clamps.py new file mode 100644 index 0000000000..96c9dfb9be --- /dev/null +++ b/tests/hermes_cli/test_dashboard_param_clamps.py @@ -0,0 +1,62 @@ +"""Dashboard query-param clamps (#39200 + #74778 salvage). + +FastAPI Query bounds reject out-of-range values at the validation layer +(422) instead of letting them reach SQL/insights code: an unbounded +``limit`` drags every session row out of SQLite in one hit (multiplied +across every profile's state.db on the fan-out endpoint), and an +unbounded/inverted ``days`` forces full-history InsightsEngine work. +""" + +from __future__ import annotations + +import pytest + +fastapi = pytest.importorskip("fastapi") +from fastapi.testclient import TestClient # noqa: E402 + + +@pytest.fixture() +def client(tmp_path, monkeypatch): + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + monkeypatch.setenv("HERMES_DASHBOARD_SESSION_TOKEN", "clamp-test-token") + from hermes_cli import web_server + + with TestClient(web_server.app, raise_server_exceptions=False) as c: + c.headers["Authorization"] = "Bearer clamp-test-token" + yield c + + +class TestSessionPaginationClamps: + def test_oversized_limit_rejected(self, client): + r = client.get("/api/sessions", params={"limit": 10_000}) + assert r.status_code == 422 + + def test_negative_limit_rejected(self, client): + r = client.get("/api/sessions", params={"limit": -1}) + assert r.status_code == 422 + + def test_profile_fanout_limit_clamped(self, client): + r = client.get("/api/profiles/sessions", params={"limit": 10_000}) + assert r.status_code == 422 + + def test_profile_fanout_accepts_real_desktop_maximum(self, client): + # Desktop callers use limit=200 (ARCHIVED_FETCH_LIMIT, command + # palette) and electron over-fetches limit+offset — the clamp must + # sit ABOVE real client maxima, not break them. + r = client.get("/api/profiles/sessions", params={"limit": 200, "offset": 120}) + assert r.status_code == 200 + + def test_in_range_limit_accepted(self, client): + r = client.get("/api/sessions", params={"limit": 50}) + assert r.status_code == 200 + + +class TestAnalyticsDaysClamps: + @pytest.mark.parametrize("days", [0, -5, 100_000]) + def test_out_of_range_days_rejected(self, client, days): + r = client.get("/api/analytics/usage", params={"days": days}) + assert r.status_code == 422 + + def test_in_range_days_accepted(self, client): + r = client.get("/api/analytics/usage", params={"days": 30}) + assert r.status_code == 200 diff --git a/tests/hermes_cli/test_gateway_restart_loop.py b/tests/hermes_cli/test_gateway_restart_loop.py index 840e26a0df..bd90e99130 100644 --- a/tests/hermes_cli/test_gateway_restart_loop.py +++ b/tests/hermes_cli/test_gateway_restart_loop.py @@ -643,6 +643,68 @@ class TestLifecycleGuardModule: with pytest.raises(GatewayLifecycleBlocked): check_gateway_lifecycle("daily", "restart.sh") + def test_python_script_with_pathlib_division_not_blocked(self, tmp_path): + """#77131: a .py cron script using pathlib division (Path.home() / + ".hermes") must NOT be blocked. + + Before the fix, the shell-script reference walk tokenized Python + sources and treated pathlib's bare "/" operator as an executable + path resolving to the filesystem root, which fails the + regular-file check and hard-blocks every innocent .py script. + Python is executed by the interpreter, never through a POSIX shell, + so the walk is skipped for .py and only the direct command regex + runs. + """ + from cron.lifecycle_guard import check_gateway_lifecycle + script = tmp_path / "digest.py" + script.write_text( + "from pathlib import Path\n" + 'ENV = Path.home() / ".hermes" / ".env"\n' + 'print("digest ok")\n' + ) + check_gateway_lifecycle("clean prompt", str(script)) + + def test_python_script_with_literal_lifecycle_command_still_blocked( + self, tmp_path + ): + """#77131: skipping the shell walk for .py must NOT weaken the guard — + a literal lifecycle command embedded in a .py script is still caught + by the direct regex scan.""" + from cron.lifecycle_guard import GatewayLifecycleBlocked, check_gateway_lifecycle + script = tmp_path / "evil.py" + script.write_text('import os\nos.system("hermes gateway restart")\n') + with pytest.raises(GatewayLifecycleBlocked): + check_gateway_lifecycle("clean prompt", str(script)) + + def test_absolute_path_binary_does_not_crash_guard(self): + """#76762: a terminal command invoking a binary by absolute path + (e.g. /usr/bin/python3) must not crash the guard with + ValueError: embedded null byte. + + Before the fix, the walk read the binary's bytes, decoded them as + text, and re-tokenized machine code containing NUL bytes; the + recursion then called Path.resolve() on a path with an embedded NUL + and only OSError was caught. Binaries are now skipped as + "nothing to scan" and ValueError is tolerated at resolve time. + """ + from cron.lifecycle_guard import ( + contains_gateway_lifecycle_command_or_referenced_script, + ) + result = contains_gateway_lifecycle_command_or_referenced_script( + '/usr/bin/python3 -c "print(1)"' + ) + assert result is False + + def test_shell_script_reference_walk_still_works(self, tmp_path): + """The referenced-script walk still applies to real shell scripts: + a .sh script that itself invokes a lifecycle command is caught.""" + from cron.lifecycle_guard import GatewayLifecycleBlocked, check_gateway_lifecycle + script = tmp_path / "wrapper.sh" + script.write_text("#!/bin/bash\n./deploy.sh\n") + (tmp_path / "deploy.sh").write_text("#!/bin/bash\nhermes gateway stop\n") + with pytest.raises(GatewayLifecycleBlocked): + check_gateway_lifecycle("daily ops", str(script)) + # --------------------------------------------------------------------------- # Defense 2 (chokepoint): cron.jobs.create_job blocks the AGENT model-tool path diff --git a/tests/hermes_cli/test_gateway_service.py b/tests/hermes_cli/test_gateway_service.py index 6d2ad47bc1..5b1a9260f8 100644 --- a/tests/hermes_cli/test_gateway_service.py +++ b/tests/hermes_cli/test_gateway_service.py @@ -495,8 +495,10 @@ class TestGatewaySystemServiceRouting: monkeypatch.setattr(gateway_cli, "_select_systemd_scope", lambda system=False: False) monkeypatch.setattr(gateway_cli, "_require_service_installed", lambda action, system=False: None) + monkeypatch.setattr(gateway_cli, "_preflight_user_systemd", lambda **kwargs: None) monkeypatch.setattr(gateway_cli, "refresh_systemd_unit_if_needed", lambda system=False: calls.append(("refresh", system))) - monkeypatch.setattr(gateway_cli, "_get_restart_drain_timeout", lambda: 12.0) + # Wait budget covers after-turn deferral + drain + headroom (#77184). + monkeypatch.setattr(gateway_cli, "_get_restart_exit_wait_budget", lambda: 27.0) monkeypatch.setattr( "gateway.status.get_running_pid", lambda: 654, @@ -528,12 +530,14 @@ class TestGatewaySystemServiceRouting: gateway_cli.systemd_restart() - assert ("graceful", 654, 17.0) in calls + assert ("graceful", 654, 27.0) in calls assert any(call[0] == "reset-failed" for call in calls) assert any(call[0] == "restart" for call in calls) assert ("wait", False, 654) in calls out = capsys.readouterr().out.lower() assert "restarting gracefully" in out + assert "21627" not in out # must use the mocked budget, not live defaults + assert "27" in out diff --git a/tests/hermes_cli/test_kanban_count_notify_subs.py b/tests/hermes_cli/test_kanban_count_notify_subs.py index c8c419133e..f55c358045 100644 --- a/tests/hermes_cli/test_kanban_count_notify_subs.py +++ b/tests/hermes_cli/test_kanban_count_notify_subs.py @@ -32,11 +32,49 @@ def test_missing_db_counts_zero_and_creates_nothing(kanban_home): db_path = kb.kanban_db_path(board="default") assert not db_path.exists() assert kb.count_notify_subs(board="default") == 0 + assert kb.count_notify_subs( + board="default", platform="tui", chat_id="session-1" + ) == 0 assert not db_path.exists(), "read-only probe must not create the DB" +def test_optional_filters_narrow_count_without_changing_unfiltered_count(kanban_home): + conn = kb.connect(board="default") + try: + tid = kb.create_task(conn, title="t", assignee="w") + kb.add_notify_sub(conn, task_id=tid, platform="tui", chat_id="session-1") + kb.add_notify_sub( + conn, + task_id=tid, + platform="TUI", + chat_id="session-2", + thread_id="thread-2", + ) + kb.add_notify_sub( + conn, task_id=tid, platform="telegram", chat_id="session-1" + ) + finally: + conn.close() - + assert kb.count_notify_subs(board="default") == 3 + assert kb.count_notify_subs(board="default", platform="tui") == 2 + assert kb.count_notify_subs(board="default", chat_id="session-1") == 2 + assert kb.count_notify_subs(board="default", thread_id="thread-2") == 1 + assert kb.count_notify_subs( + board="default", platform="tui", chat_id="session-1" + ) == 1 + assert kb.count_notify_subs( + board="default", platform="tui", chat_id="session-1", thread_id="" + ) == 1 + assert kb.count_notify_subs( + board="default", + platform="tui", + chat_id="session-2", + thread_id="thread-2", + ) == 1 + assert kb.count_notify_subs( + board="default", platform="tui", chat_id="other-session" + ) == 0 def test_legacy_db_without_subs_table_counts_zero_and_stays_unmigrated(tmp_path): @@ -48,6 +86,9 @@ def test_legacy_db_without_subs_table_counts_zero_and_stays_unmigrated(tmp_path) finally: conn.close() assert kb.count_notify_subs(db_path=legacy) == 0 + assert kb.count_notify_subs( + db_path=legacy, platform="tui", chat_id="session-1" + ) == 0 # The probe must not have run schema init on the foreign/legacy DB. conn = sqlite3.connect(legacy) try: @@ -88,4 +129,3 @@ def test_count_notify_subs_filters_profile_owners(tmp_path): notifier_profiles={"default"}, include_unowned=True, ) == 2 - diff --git a/tests/hermes_cli/test_mcp_catalog.py b/tests/hermes_cli/test_mcp_catalog.py index a5b465dd80..29e510b29e 100644 --- a/tests/hermes_cli/test_mcp_catalog.py +++ b/tests/hermes_cli/test_mcp_catalog.py @@ -143,6 +143,41 @@ class TestManifestParsing: assert e.auth.env[1].required is False assert e.auth.env[1].secret is False + def test_http_api_key_builds_bearer_headers_template(self, catalog_dir): + body = _basic_manifest( + transport={"type": "http", "url": "https://mcp.example.com/sse"}, + auth={ + "type": "api_key", + "env": [{"name": "MCP_DEMO_API_KEY", "prompt": "key", "secret": True}], + }, + ) + _write_manifest(catalog_dir, "demo", body) + from hermes_cli.mcp_catalog import _build_server_config + + cfg = _build_server_config(_entry("demo"), None) + assert cfg["url"] == "https://mcp.example.com/sse" + assert cfg["headers"] == {"Authorization": "Bearer ${MCP_DEMO_API_KEY}"} + + def test_http_api_key_requires_matching_env_declaration(self, catalog_dir): + """http+api_key manifests must declare the env key the header references. + + install_entry only persists auth.env-declared vars; a manifest naming + its key e.g. N8N_API_KEY would install cleanly but send a literal + ${MCP_DEMO_API_KEY} placeholder at connect time (silent 401). + """ + body = _basic_manifest( + transport={"type": "http", "url": "https://mcp.example.com/sse"}, + auth={ + "type": "api_key", + "env": [{"name": "DEMO_API_KEY", "prompt": "key", "secret": True}], + }, + ) + path = _write_manifest(catalog_dir, "demo", body) + from hermes_cli.mcp_catalog import CatalogError, _parse_manifest + + with pytest.raises(CatalogError, match="MCP_DEMO_API_KEY"): + _parse_manifest(path) + @@ -193,6 +228,36 @@ class TestInstall: assert get_env_value("DEMO_KEY") == "secret-val" assert "demo" in load_config()["mcp_servers"] + def test_install_http_api_key_writes_bearer_headers(self, catalog_dir, monkeypatch): + body = _basic_manifest( + transport={"type": "http", "url": "https://mcp.example.com/sse"}, + auth={ + "type": "api_key", + "env": [{"name": "MCP_DEMO_API_KEY", "prompt": "key", "secret": True}], + }, + ) + _write_manifest(catalog_dir, "demo", body) + + from hermes_cli import mcp_catalog + + monkeypatch.setattr(mcp_catalog, "_prompt_input", lambda *a, **kw: "secret-val") + + from hermes_cli.mcp_catalog import install_entry + from hermes_cli.config import load_config + + install_entry(_entry("demo"), enable=True) + + server = load_config()["mcp_servers"]["demo"] + assert server["url"] == "https://mcp.example.com/sse" + assert server["headers"] == {"Authorization": "Bearer secret-val"} + # The raw file must carry the ${...} template, never the secret — + # load_config resolves it; config.yaml itself stays secret-free. + from hermes_cli.config import get_config_path + + raw = get_config_path().read_text() + assert "${MCP_DEMO_API_KEY}" in raw + assert "secret-val" not in raw + diff --git a/tests/hermes_cli/test_model_switch_custom_providers.py b/tests/hermes_cli/test_model_switch_custom_providers.py index 5b3513dc82..710a176d1a 100644 --- a/tests/hermes_cli/test_model_switch_custom_providers.py +++ b/tests/hermes_cli/test_model_switch_custom_providers.py @@ -897,3 +897,74 @@ def test_excluded_providers_hides_builtin_row(monkeypatch): ) +def test_custom_provider_context_length_models_dict_still_probes(monkeypatch): + """Dict-shaped ``models:`` from ``_save_custom_provider`` is metadata. + + ``hermes model`` writes ``models: {default: {context_length: N}}`` for + local Ollama. That must not suppress live /v1/models discovery — otherwise + Desktop/Telegram only show the saved default and Refresh does nothing. + """ + monkeypatch.setattr("agent.models_dev.fetch_models_dev", lambda: {}) + monkeypatch.setattr(providers_mod, "HERMES_OVERLAYS", {}) + calls = [] + + def fetch(api_key, base_url, **kwargs): + calls.append((api_key, base_url, kwargs)) + return ["qwen3.6:35b-mlx", "gemma4:31b", "llama3"] + + monkeypatch.setattr("hermes_cli.models.fetch_api_models", fetch) + + providers = list_authenticated_providers( + current_provider="custom:local-ollama", + user_providers={}, + custom_providers=[ + { + "name": "Local Ollama", + "base_url": "http://localhost:11434/v1", + "model": "qwen3.6:35b-mlx", + "models": {"qwen3.6:35b-mlx": {"context_length": 32768}}, + } + ], + # GUI picker path: probe current custom provider only. + probe_custom_providers=False, + probe_current_custom_provider=True, + current_base_url="http://localhost:11434/v1", + ) + + assert len(calls) == 1 + assert calls[0][0] == "" + assert calls[0][1] == "http://localhost:11434/v1" + row = next(p for p in providers if p["name"] == "Local Ollama") + assert row["models"] == ["qwen3.6:35b-mlx", "gemma4:31b", "llama3"] + assert row["total_models"] == 3 + + +def test_custom_provider_dict_models_pin_requires_discover_false(monkeypatch): + """Dict-shaped catalogs pin only when ``discover_models: false``.""" + monkeypatch.setattr("agent.models_dev.fetch_models_dev", lambda: {}) + monkeypatch.setattr(providers_mod, "HERMES_OVERLAYS", {}) + calls = [] + + def fetch(*args, **kwargs): + calls.append((args, kwargs)) + return ["unexpected-live-model"] + + monkeypatch.setattr("hermes_cli.models.fetch_api_models", fetch) + + providers = list_authenticated_providers( + current_provider="custom:local-ollama", + user_providers={}, + custom_providers=[ + { + "name": "Local Ollama", + "base_url": "http://localhost:11434/v1", + "model": "llama3", + "models": {"llama3": {}}, + "discover_models": False, + } + ], + ) + + row = next(p for p in providers if p["name"] == "Local Ollama") + assert calls == [] + assert row["models"] == ["llama3"] diff --git a/tests/hermes_cli/test_plugins_hub_perf_guard.py b/tests/hermes_cli/test_plugins_hub_perf_guard.py new file mode 100644 index 0000000000..370e425e90 --- /dev/null +++ b/tests/hermes_cli/test_plugins_hub_perf_guard.py @@ -0,0 +1,182 @@ +from __future__ import annotations + +import threading +from pathlib import Path +from types import SimpleNamespace + +from hermes_cli import web_server +from hermes_cli import plugins_cmd +from tools import registry as tools_registry + + +_PLUGIN_ROW = [("demo", "1.0.0", "demo plugin", "user", "/tmp/demo-plugin", "demo")] + + +def _patch_minimal_hub_dependencies(monkeypatch, *, check_fn, discover_all_plugins=None): + monkeypatch.setattr(web_server, "_get_dashboard_plugins", lambda force_rescan=False: []) + monkeypatch.setattr(web_server, "_discover_memory_provider_statuses", lambda: []) + monkeypatch.setattr(web_server, "get_hermes_home", lambda: Path("/tmp/hermes-home")) + monkeypatch.setattr(web_server, "load_config", lambda: {"dashboard": {"hidden_plugins": []}}) + + monkeypatch.setattr( + plugins_cmd, + "_discover_all_plugins", + discover_all_plugins or (lambda: list(_PLUGIN_ROW)), + ) + monkeypatch.setattr(plugins_cmd, "_get_current_context_engine", lambda: "compressor") + monkeypatch.setattr(plugins_cmd, "_get_current_memory_provider", lambda: "") + monkeypatch.setattr(plugins_cmd, "_discover_context_engines", lambda: []) + monkeypatch.setattr(plugins_cmd, "_get_disabled_set", lambda: set()) + monkeypatch.setattr(plugins_cmd, "_get_enabled_set", lambda: {"demo"}) + monkeypatch.setattr(plugins_cmd, "_read_manifest", lambda _path: {"provides_tools": ["demo_tool"]}) + + monkeypatch.setattr( + tools_registry.registry, + "get_entry", + lambda _name: SimpleNamespace(check_fn=check_fn), + ) + + + +def test_plugins_hub_does_not_probe_cold_check_fns(monkeypatch): + tools_registry.invalidate_check_fn_cache() + web_server._invalidate_plugins_hub_cache() + + calls = {"count": 0, "threads": set()} + + def check_fn(): + calls["count"] += 1 + calls["threads"].add(threading.current_thread()) + return False + + _patch_minimal_hub_dependencies(monkeypatch, check_fn=check_fn) + + payload = web_server._merged_plugins_hub(force_refresh=True) + + # The request path itself must never execute the probe: the cold verdict + # is unknown, so the payload reports no auth requirement. Any probing + # happens on a background warmer thread, never inline. + assert payload["plugins"][0]["auth_required"] is False + assert payload["plugins"][0]["auth_command"] == "" + assert threading.current_thread() not in calls["threads"] + + +def test_plugins_hub_cold_cache_schedules_background_probe(monkeypatch): + tools_registry.invalidate_check_fn_cache() + web_server._invalidate_plugins_hub_cache() + + probe_ran = threading.Event() + + def check_fn(): + probe_ran.set() + return False + + _patch_minimal_hub_dependencies(monkeypatch, check_fn=check_fn) + + scheduled: list = [] + real_schedule = web_server._schedule_check_fn_probe + + def tracking_schedule(fn): + thread = real_schedule(fn) + scheduled.append(thread) + return thread + + monkeypatch.setattr(web_server, "_schedule_check_fn_probe", tracking_schedule) + + # Cold cache → the fetch schedules a background probe and reports the + # verdict as unknown (auth_required stays False for now). + payload = web_server._merged_plugins_hub(force_refresh=True) + assert payload["plugins"][0]["auth_required"] is False + assert scheduled and scheduled[0] is not None + + scheduled[0].join(timeout=5) + assert probe_ran.wait(timeout=5) + + # Once the TTL cache refreshes, the probed False verdict surfaces as an + # auth requirement. + refreshed = web_server._merged_plugins_hub(force_refresh=True) + assert refreshed["plugins"][0]["auth_required"] is True + assert refreshed["plugins"][0]["auth_command"] == "hermes auth demo" + + + +def test_plugins_hub_uses_cached_failed_check_fn_verdict(monkeypatch): + tools_registry.invalidate_check_fn_cache() + web_server._invalidate_plugins_hub_cache() + + def check_fn(): + return False + + assert tools_registry._check_fn_cached(check_fn) is False + _patch_minimal_hub_dependencies(monkeypatch, check_fn=check_fn) + + payload = web_server._merged_plugins_hub(force_refresh=True) + + assert payload["plugins"][0]["auth_required"] is True + assert payload["plugins"][0]["auth_command"] == "hermes auth demo" + + + +def test_plugins_hub_short_ttl_cache_collapses_duplicate_fetches(monkeypatch): + tools_registry.invalidate_check_fn_cache() + web_server._invalidate_plugins_hub_cache() + + calls = {"discover": 0} + + def discover_all_plugins(): + calls["discover"] += 1 + return list(_PLUGIN_ROW) + + _patch_minimal_hub_dependencies( + monkeypatch, + check_fn=lambda: True, + discover_all_plugins=discover_all_plugins, + ) + + first = web_server._merged_plugins_hub(force_refresh=True) + second = web_server._merged_plugins_hub() + + assert calls["discover"] == 1 + assert first is second + + +def test_plugin_install_endpoint_invalidates_hub_cache(monkeypatch): + import asyncio + + from hermes_cli.web_models import _AgentPluginInstallBody + + tools_registry.invalidate_check_fn_cache() + web_server._invalidate_plugins_hub_cache() + + calls = {"discover": 0} + + def discover_all_plugins(): + calls["discover"] += 1 + return list(_PLUGIN_ROW) + + _patch_minimal_hub_dependencies( + monkeypatch, + check_fn=lambda: True, + discover_all_plugins=discover_all_plugins, + ) + + # Prime the TTL cache; a plain fetch must be served from it. + web_server._merged_plugins_hub(force_refresh=True) + web_server._merged_plugins_hub() + assert calls["discover"] == 1 + + # Simulate a successful install through the endpoint; its invalidation + # hook must drop the memoized payload so the next fetch rebuilds. + monkeypatch.setattr(web_server, "_require_token", lambda _request: None) + monkeypatch.setattr( + plugins_cmd, "dashboard_install_plugin", lambda *a, **k: {"ok": True} + ) + + asyncio.run( + web_server.post_agent_plugin_install( + object(), _AgentPluginInstallBody(identifier="demo") + ) + ) + + web_server._merged_plugins_hub() + assert calls["discover"] == 2 diff --git a/tests/hermes_cli/test_relay_shared_metrics_runtime.py b/tests/hermes_cli/test_relay_shared_metrics_runtime.py index 67a72ccc35..dc8328a86a 100644 --- a/tests/hermes_cli/test_relay_shared_metrics_runtime.py +++ b/tests/hermes_cli/test_relay_shared_metrics_runtime.py @@ -33,6 +33,7 @@ class _Relay: self._tool_starts: dict[Any, dict[str, Any]] = {} self._scope_starts: dict[Any, dict[str, Any]] = {} self._scope = contextvars.ContextVar("relay_scope", default=None) + self._scope_stack = contextvars.ContextVar("relay_scope_stack", default=None) self._scope_serial = 0 self.ScopeType = SimpleNamespace( Agent="agent", Function="function", Tool="tool" @@ -55,6 +56,11 @@ class _Relay: def _scope_push(self, name: str, scope_type: Any, **kwargs: Any) -> Any: self._scope_serial += 1 handle = ("scope", name, self._scope_serial) + stack = self._scope_stack.get() + if stack is None: + stack = [] + self._scope_stack.set(stack) + stack.append(handle) self._scope.set(handle) self.events.append(("scope.push", name, scope_type, kwargs)) if scope_type == self.ScopeType.Function: @@ -73,6 +79,13 @@ class _Relay: return handle def _scope_pop(self, handle: Any, **kwargs: Any) -> None: + stack = self._scope_stack.get() + if not stack or stack[-1] != handle: + current = stack[-1] if stack else None + self.events.append(("scope.pop.rejected", handle, current)) + raise RuntimeError("scope handle is not at the top of the stack") + stack.pop() + self._scope.set(stack[-1] if stack else None) self.events.append(("scope.pop", handle, kwargs)) start = self._scope_starts.pop(handle, None) if start is not None: @@ -106,7 +119,9 @@ class _Relay: callback(event) def _get_scope_stack(self) -> Any: - current = self._scope.get() + stack = self._scope_stack.get() + current = stack[-1] if stack else None + self._scope.set(current) self.events.append(("scope.sync", current)) return current @@ -1005,6 +1020,75 @@ def test_core_task_instrumentation_preserves_prompt_history_and_tool_schema( assert json.dumps(agent.tools, ensure_ascii=False, sort_keys=True) == tools_before +def test_skipped_turn_does_not_finish_another_sessions_matching_task( + direct_runtime, + monkeypatch, +): + """A skipped turn must not use shared-metrics' task-id fallback on finish.""" + from run_agent import AIAgent + + owner_session = "instrumented-session" + shared_task_id = "caller-supplied-task-id" + relay_shared_metrics.start_task_run( + session_id=owner_session, + task_id=shared_task_id, + platform="cli", + ) + runtime = relay_shared_metrics._get_runtime() + assert runtime is not None + assert (owner_session, shared_task_id) in runtime._task_sessions + + agent = object.__new__(AIAgent) + agent.session_id = "skipped-session" + agent.platform = "cli" + agent._parent_session_id = None + agent._session_db = None + agent._cached_system_prompt = "stable" + agent.tools = [] + + skipped_turn = SimpleNamespace(relay_enabled=False) + monkeypatch.setattr( + relay_runtime.SESSION_COORDINATOR, + "begin_turn", + lambda *_args, **_kwargs: skipped_turn, + ) + monkeypatch.setattr( + relay_runtime.SESSION_COORDINATOR, + "finish_logical_calls", + lambda *_args, **_kwargs: None, + ) + monkeypatch.setattr( + relay_runtime.SESSION_COORDINATOR, + "end_turn", + lambda *_args, **_kwargs: None, + ) + monkeypatch.setattr( + "agent.conversation_loop.run_conversation", + lambda *_args, **_kwargs: {"final_response": "ok", "completed": True}, + ) + + result = AIAgent.run_conversation( + agent, + "hello", + conversation_history=[], + task_id=shared_task_id, + ) + + assert result["completed"] is True + assert (owner_session, shared_task_id) in runtime._task_sessions + assert not [ + event + for event in direct_runtime.events + if event[0] == "scope.pop" and event[1][1] == relay_shared_metrics.TASK_SCOPE + ] + relay_shared_metrics.finish_task_run( + session_id=owner_session, + task_id=shared_task_id, + platform="cli", + result={"completed": True}, + ) + + @@ -1170,6 +1254,171 @@ def test_sync_session_runner_releases_lock_before_callback(direct_runtime): assert contender.is_alive() is False +def test_direct_runtime_fake_enforces_lifo_scope_contract(direct_runtime): + runtime = relay_runtime.get_runtime() + assert runtime is not None + session = runtime.ensure_session({"session_id": "lifo-contract"}) + assert session is not None + + first = runtime.run_in_session( + session, + direct_runtime.scope.push, + "first", + direct_runtime.ScopeType.Function, + ) + second = runtime.run_in_session( + session, + direct_runtime.scope.push, + "second", + direct_runtime.ScopeType.Function, + ) + + with pytest.raises(RuntimeError, match="not at the top"): + runtime.run_in_session(session, direct_runtime.scope.pop, first) + + runtime.run_in_session(session, direct_runtime.scope.pop, second) + runtime.run_in_session(session, direct_runtime.scope.pop, first) + + +def test_concurrent_turn_skips_relay_before_scope_stack_can_interleave( + direct_runtime, +): + coordinator = relay_runtime.SESSION_COORDINATOR + profile_key = relay_runtime.current_profile_key() + lease = coordinator.acquire_conversation( + profile_key=profile_key, + session_id="shared-session", + platform="cli", + ) + first = coordinator.begin_turn(lease, turn_id="first", task_id="first-task") + second = coordinator.begin_turn( + lease, + turn_id="second", + task_id="second-task", + ) + + assert first.relay_enabled is True + assert first.handle is not None + assert second.relay_enabled is False + assert second.handle is None + assert relay_runtime.resolve_execution_context("shared-session") == ( + None, + None, + None, + ) + + coordinator.end_turn(first, outcome="success") + coordinator.end_turn(second, outcome="success") + coordinator.release_conversation(lease) + coordinator.finalize_conversation( + profile_key=profile_key, + session_id="shared-session", + ) + + turn_closes = [ + event + for event in direct_runtime.events + if event[0] == "scope.pop" and event[1] == first.handle + ] + assert len(turn_closes) == 1 + assert not [ + event + for event in direct_runtime.events + if event[0] == "scope.pop.rejected" + ] + + +def test_concurrent_turn_skips_shared_metrics_scope_creation(direct_runtime): + coordinator = relay_runtime.SESSION_COORDINATOR + profile_key = relay_runtime.current_profile_key() + lease = coordinator.acquire_conversation( + profile_key=profile_key, + session_id="shared-session", + platform="cli", + ) + first = coordinator.begin_turn(lease, turn_id="first", task_id="first-task") + second = coordinator.begin_turn(lease, turn_id="second", task_id="second-task") + + relay_shared_metrics.observe_lifecycle( + "pre_llm_call", + session_id="shared-session", + task_id="second-task", + platform="cli", + ) + relay_shared_metrics.observe_lifecycle( + "pre_api_request", + session_id="shared-session", + task_id="second-task", + api_request_id="second-request", + platform="cli", + ) + + assert second.relay_enabled is False + assert not [ + event + for event in direct_runtime.events + if event[0] == "scope.push" and event[1] == relay_shared_metrics.TASK_SCOPE + ] + + coordinator.end_turn(first, outcome="success") + coordinator.end_turn(second, outcome="success") + coordinator.release_conversation(lease) + + +def test_skipped_turn_stays_gated_after_instrumented_turn_ends(direct_runtime): + coordinator = relay_runtime.SESSION_COORDINATOR + profile_key = relay_runtime.current_profile_key() + lease = coordinator.acquire_conversation( + profile_key=profile_key, + session_id="shared-session", + platform="cli", + ) + first = coordinator.begin_turn(lease, turn_id="first", task_id="first-task") + second = coordinator.begin_turn(lease, turn_id="second", task_id="second-task") + inherited = contextvars.copy_context() + + coordinator.end_turn(first, outcome="success") + + assert relay_runtime.current_turn() is second + assert inherited.run(relay_runtime.current_turn) is second + assert not relay_runtime.relay_instrumentation_enabled() + assert not inherited.run(relay_runtime.relay_instrumentation_enabled) + assert relay_runtime.resolve_execution_context("shared-session") == ( + None, + None, + None, + ) + + relay_shared_metrics.observe_lifecycle( + "pre_llm_call", + session_id="shared-session", + task_id="second-task", + platform="cli", + ) + inherited.run( + relay_shared_metrics.observe_lifecycle, + "pre_api_request", + session_id="shared-session", + task_id="second-task", + api_request_id="second-request", + platform="cli", + ) + + assert not [ + event + for event in direct_runtime.events + if event[0] == "scope.push" + and event[1] + in {relay_shared_metrics.TASK_SCOPE, relay_shared_metrics.MODEL_CALL_SCOPE} + ] + + coordinator.end_turn(second, outcome="success") + assert relay_runtime.current_turn() is None + assert inherited.run(relay_runtime.current_turn) is second + assert not inherited.run(relay_runtime.relay_instrumentation_enabled) + coordinator.release_conversation(lease) + + @@ -2077,6 +2326,8 @@ def test_failed_flush_keeps_daily_export_open_for_later_task( assert metrics["hermes.task_run.finished"]["value"] == 2 assert flush_attempts == 2 assert "Hermes shared-metrics task flush failed" in caplog.text + + def test_skill_lifecycle_flows_through_relay_to_a_privacy_safe_package( direct_runtime, tmp_path, @@ -2131,6 +2382,7 @@ def test_skill_lifecycle_flows_through_relay_to_a_privacy_safe_package( } assert "private-skill-name" not in json.dumps(package) + def test_skill_lifecycle_with_only_task_id_uses_unique_task_scope(direct_runtime): runtime = relay_shared_metrics._get_runtime() assert runtime is not None @@ -2156,6 +2408,7 @@ def test_skill_lifecycle_with_only_task_id_uses_unique_task_scope(direct_runtime ] assert mark[2]["handle"] == task.handle + def test_skill_task_only_correlation_does_not_guess_across_sessions(direct_runtime): runtime = relay_shared_metrics._get_runtime() assert runtime is not None @@ -2181,6 +2434,7 @@ def test_skill_task_only_correlation_does_not_guess_across_sessions(direct_runti ] assert "handle" not in mark[2] + def test_late_skill_lifecycle_is_not_reemitted_at_the_root(direct_runtime): base = { "session_id": "session-1", @@ -2215,6 +2469,7 @@ def test_late_skill_lifecycle_is_not_reemitted_at_the_root(direct_runtime): if event[0] == "scope.event" and event[1] == "hermes.skill.load" ] == [] + def test_skill_lifecycle_does_not_fallback_across_an_explicit_session( direct_runtime, ): diff --git a/tests/hermes_cli/test_session_recovery.py b/tests/hermes_cli/test_session_recovery.py index a948b0ad69..3cabe5a750 100644 --- a/tests/hermes_cli/test_session_recovery.py +++ b/tests/hermes_cli/test_session_recovery.py @@ -112,8 +112,6 @@ def _orphan_fts_schema(path: Path) -> None: conn.execute("PRAGMA writable_schema=OFF") finally: conn.close() - - def _make_page_spanning_source( path: Path, message_count: int = 320, @@ -596,7 +594,58 @@ def test_cli_allow_partial_salvages_rows_across_a_corrupt_leaf( } - +def test_partial_recovery_clears_only_unreadable_system_prompt_refs( + tmp_path: Path, +) -> None: + source = tmp_path / "corrupt-system-prompts.db" + output = tmp_path / "partial-system-prompts.db" + session_count = 180 + _make_many_sessions_source(source, session_count) + + conn = sqlite3.connect(str(source), isolation_level=None) + try: + row = conn.execute( + "SELECT rootpage FROM sqlite_master " + "WHERE type = 'table' AND name = 'system_prompts'" + ).fetchone() + assert row is not None + prompt_root = int(row[0]) + finally: + conn.close() + _corrupt_middle_table_leaf(source, prompt_root) + + report = recover_session_database( + source, + output, + work_dir=tmp_path, + chunk_size=8, + allow_partial=True, + ) + + assert report["verified"] is True + assert report["partial"] is True + assert report["copy"]["sessions"]["status"] == "complete" + assert report["copy"]["messages"]["status"] == "complete" + assert report["copy"]["system_prompts"]["status"] == "partial" + cleared = report["orphan_cleanup"]["session_prompt_refs_cleared"] + assert 0 < cleared < session_count + assert report["verification"]["foreign_key_check"] == [] + + conn = sqlite3.connect(str(output)) + try: + assert conn.execute("PRAGMA integrity_check").fetchall() == [("ok",)] + assert conn.execute("PRAGMA foreign_key_check").fetchall() == [] + assert conn.execute("SELECT COUNT(*) FROM sessions").fetchone()[0] == session_count + retained = conn.execute( + "SELECT COUNT(*) FROM sessions WHERE system_prompt_hash IS NOT NULL" + ).fetchone()[0] + assert retained == session_count - cleared + assert ( + conn.execute("SELECT COUNT(*) FROM system_prompts").fetchone()[0] + == retained + ) + finally: + conn.close() diff --git a/tests/hermes_cli/test_tools_config.py b/tests/hermes_cli/test_tools_config.py index 09af0c4e11..339e498503 100644 --- a/tests/hermes_cli/test_tools_config.py +++ b/tests/hermes_cli/test_tools_config.py @@ -6,7 +6,8 @@ from unittest.mock import patch import pytest -from hermes_cli.nous_account import NousPortalAccountInfo +from hermes_cli.nous_account import NousPortalAccountInfo, NousToolAccessInfo +from hermes_cli.nous_subscription import NousSubscriptionFeatures from hermes_cli.tools_config import ( _DEFAULT_OFF_TOOLSETS, _RECENTLY_SHIPPED_TOOLSETS, @@ -544,6 +545,74 @@ def _fake_features(*, logged_in: bool, paid: bool = True): return SimpleNamespace(nous_auth_present=logged_in, account_info=account) +def test_visible_providers_reuses_logged_out_feature_snapshot(monkeypatch): + import hermes_cli.tools_config as tools_config + + account = NousPortalAccountInfo( + logged_in=False, + source="none", + fresh=False, + paid_service_access=None, + ) + features = NousSubscriptionFeatures( + subscribed=False, + nous_auth_present=False, + provider_is_nous=False, + features={}, + account_info=account, + ) + monkeypatch.setattr( + tools_config, + "get_nous_subscription_features", + lambda *args, **kwargs: pytest.fail("feature snapshot was resolved again"), + ) + + providers = _visible_providers( + TOOL_CATEGORIES["image_gen"], {}, features=features + ) + + assert any( + provider.get("managed_nous_feature") == "image_gen" + for provider in providers + ) + + +def test_visible_providers_reuses_pool_video_feature_snapshot(monkeypatch): + import hermes_cli.tools_config as tools_config + + account = NousPortalAccountInfo( + logged_in=True, + source="jwt", + fresh=False, + paid_service_access=False, + tool_access=NousToolAccessInfo( + enabled=True, + coverage={"fal-video": False}, + ), + ) + features = NousSubscriptionFeatures( + subscribed=True, + nous_auth_present=True, + provider_is_nous=False, + features={}, + account_info=account, + ) + monkeypatch.setattr( + tools_config, + "get_nous_subscription_features", + lambda *args, **kwargs: pytest.fail("feature snapshot was resolved again"), + ) + + providers = _visible_providers( + TOOL_CATEGORIES["video_gen"], {}, features=features + ) + + assert not any( + provider.get("managed_nous_feature") == "video_gen" + for provider in providers + ) + + # ── Windows console-flash guard for post-setup subprocess spawns ────────────── diff --git a/tests/hermes_cli/test_web_server.py b/tests/hermes_cli/test_web_server.py index c7f3c5fc85..4c0b5500f1 100644 --- a/tests/hermes_cli/test_web_server.py +++ b/tests/hermes_cli/test_web_server.py @@ -142,7 +142,7 @@ class TestReloadEnv: def test_adds_new_vars(self, tmp_path): """reload_env() adds vars from .env that are not in os.environ.""" env_file = tmp_path / ".env" - env_file.write_text("TEST_RELOAD_VAR=hello123\n") + env_file.write_text("TEST_RELOAD_VAR=hello123\n", encoding="utf-8") with patch.dict(reload_env.__globals__, {"get_env_path": lambda: env_file}): os.environ.pop("TEST_RELOAD_VAR", None) count = reload_env() @@ -523,6 +523,79 @@ class TestWebServerEndpoints: def _provider_field_map(payload): return {field["key"]: field for field in payload["fields"]} + def test_openviking_recall_fields_are_numeric_dashboard_controls(self): + resp = self.client.get("/api/memory/providers/openviking/config") + + assert resp.status_code == 200 + fields = self._provider_field_map(resp.json()) + assert fields["recall_limit"]["kind"] == "integer" + assert fields["recall_limit"]["minimum"] == 1 + assert fields["recall_limit"]["maximum"] == 100 + assert fields["recall_score_threshold"]["kind"] == "number" + assert fields["recall_score_threshold"]["step"] == 0.01 + assert fields["recall_resources"]["kind"] == "boolean" + + def test_openviking_dashboard_persists_typed_recall_values(self): + from hermes_cli.config import load_config + + resp = self.client.put( + "/api/memory/providers/openviking/config", + json={ + "values": { + "endpoint": "http://127.0.0.1:1933", + "recall_limit": "12", + "recall_score_threshold": "0.42", + "recall_max_injected_chars": "8000", + "profile_token_budget": "7000", + "recall_timeout_seconds": "2.5", + "recall_request_timeout_seconds": "1.5", + "recall_full_read_limit": "5", + "recall_prefer_abstract": True, + "recall_resources": False, + } + }, + ) + + assert resp.status_code == 200 + config = load_config()["memory"]["openviking"] + assert config["recall_limit"] == 12 + assert config["recall_score_threshold"] == 0.42 + assert config["profile_token_budget"] == 7000 + assert config["recall_prefer_abstract"] is True + assert config["recall_resources"] is False + + def test_openviking_dashboard_rejects_out_of_range_recall_value(self): + resp = self.client.put( + "/api/memory/providers/openviking/config", + json={ + "values": { + "endpoint": "http://127.0.0.1:1933", + "recall_limit": 101, + } + }, + ) + + assert resp.status_code == 400 + assert "must be at most 100" in resp.json()["detail"] + + def test_openviking_dashboard_rejects_blocked_endpoint_before_saving(self): + from hermes_cli.config import load_config + + resp = self.client.put( + "/api/memory/providers/openviking/config", + json={ + "values": { + "endpoint": "http://169.254.169.254/latest/meta-data/credential", + } + }, + ) + + assert resp.status_code == 400 + assert "blocked metadata address" in resp.json()["detail"] + assert "credential" not in resp.json()["detail"] + memory_config = load_config().get("memory", {}) + assert "openviking" not in memory_config + @@ -1948,6 +2021,36 @@ class TestNewEndpoints: config["platform_toolsets"]["discord"] ) + def test_toolsets_resolve_subscription_features_once(self, monkeypatch): + import hermes_cli.tools_config as tools_config + from hermes_cli.nous_subscription import NousSubscriptionFeatures + + calls = 0 + features = NousSubscriptionFeatures( + subscribed=False, + nous_auth_present=False, + provider_is_nous=False, + features={}, + account_info=None, + ) + + def resolve_features(config, *, force_fresh=False): + nonlocal calls + calls += 1 + return features + + monkeypatch.setattr( + tools_config, + "get_nous_subscription_features", + resolve_features, + ) + + resp = self.client.get("/api/tools/toolsets") + + assert resp.status_code == 200 + assert resp.json() + assert calls == 1 + def test_get_toolset_config_returns_provider_matrix(self): """GET .../config returns provider rows with structured env_vars.""" @@ -2192,6 +2295,38 @@ class TestNewEndpoints: assert top_skill["total_count"] == 1 assert top_skill["last_used_at"] is not None + def test_analytics_usage_skips_full_insights_generate(self): + """get_usage_analytics must call get_usage_breakdown, not generate().""" + from unittest.mock import patch + from agent.insights import InsightsEngine + from hermes_state import SessionDB + + db = SessionDB() + try: + db.create_session( + session_id="usage-tools-test", + source="cli", + model="anthropic/claude-sonnet-4", + ) + db.update_token_counts( + "usage-tools-test", + input_tokens=10, + output_tokens=5, + ) + db.append_message( + "usage-tools-test", + role="tool", + content="read output", + tool_name="read_file", + ) + finally: + db.close() + + with patch.object(InsightsEngine, "generate") as mock_generate: + resp = self.client.get("/api/analytics/usage?days=7") + assert resp.status_code == 200 + mock_generate.assert_not_called() + assert any(tool["tool"] == "read_file" for tool in resp.json()["tools"]) # --------------------------------------------------------------------------- # Model context length: normalize/denormalize + /api/model/info @@ -2702,7 +2837,7 @@ class TestDiscoverUserThemes: monkeypatch.setenv("HERMES_HOME", str(tmp_path)) themes_dir = tmp_path / "dashboard-themes" themes_dir.mkdir() - (themes_dir / "mine.yaml").write_text("name: mine\n") + (themes_dir / "mine.yaml").write_text("name: mine\n", encoding="utf-8") other = tmp_path / "other-profile" other.mkdir() @@ -3245,7 +3380,7 @@ class TestDashboardPluginManifestExtensions: import json plug_dir = tmp_path / "plugins" / name / "dashboard" plug_dir.mkdir(parents=True) - (plug_dir / "manifest.json").write_text(json.dumps(manifest)) + (plug_dir / "manifest.json").write_text(json.dumps(manifest), encoding="utf-8") return plug_dir def test_override_and_hidden_carried_through(self, tmp_path, monkeypatch): @@ -3598,10 +3733,9 @@ class TestDashboardPluginStaticAssetAllowlist: assert resp.status_code in (403, 404) -def _fake_httpx_client(*, status: int | None = None, raise_exc: bool = False): - """Build a drop-in for httpx.Client whose .get() returns a canned status - (or raises a transport error). Patched in for the credential-validate probe - so tests never touch the network.""" +def _fake_httpx_async_client(*, status: int | None = None, raise_exc: bool = False): + """Build a drop-in for httpx.AsyncClient with a canned GET response.""" + class _Resp: def __init__(self, code): self.status_code = code @@ -3614,13 +3748,13 @@ def _fake_httpx_client(*, status: int | None = None, raise_exc: bool = False): def __init__(self, *a, **k): pass - def __enter__(self): + async def __aenter__(self): return self - def __exit__(self, *a): + async def __aexit__(self, *a): return False - def get(self, *a, **k): + async def get(self, *a, **k): if raise_exc: raise RuntimeError("connection refused") return _Resp(status) @@ -3643,13 +3777,38 @@ class TestValidateProviderCredential: self.client = TestClient(app) self.client.headers[_SESSION_HEADER_NAME] = _SESSION_TOKEN + class _BlockingClient: + def __init__(self, *args, **kwargs): + raise AssertionError( + "async validation route used blocking httpx.Client" + ) + + monkeypatch.setattr("httpx.Client", _BlockingClient) + def _post(self, key, value): - return self.client.post("/api/providers/validate", json={"key": key, "value": value}) + return self.client.post( + "/api/providers/validate", json={"key": key, "value": value} + ) + def test_rejected_key_blocks(self, monkeypatch): + monkeypatch.setattr("httpx.AsyncClient", _fake_httpx_async_client(status=401)) + data = self._post("OPENROUTER_API_KEY", "sk-bogus").json() + assert data["ok"] is False and data["reachable"] is True + def test_valid_key_passes(self, monkeypatch): + monkeypatch.setattr("httpx.AsyncClient", _fake_httpx_async_client(status=200)) + data = self._post("OPENAI_API_KEY", "sk-real").json() + assert data["ok"] is True and data["reachable"] is True + + def test_rate_limited_counts_as_valid(self, monkeypatch): + monkeypatch.setattr("httpx.AsyncClient", _fake_httpx_async_client(status=429)) + data = self._post("XAI_API_KEY", "xai-real").json() + assert data["ok"] is True def test_network_error_is_unreachable_not_blocking(self, monkeypatch): - monkeypatch.setattr("httpx.Client", _fake_httpx_client(raise_exc=True)) + monkeypatch.setattr( + "httpx.AsyncClient", _fake_httpx_async_client(raise_exc=True) + ) data = self._post("OPENROUTER_API_KEY", "sk-real").json() assert data["ok"] is False and data["reachable"] is False @@ -3672,18 +3831,18 @@ class TestValidateProviderCredential: def __init__(self, *a, **k): pass - def __enter__(self): + async def __aenter__(self): return self - def __exit__(self, *a): + async def __aexit__(self, *a): return False - def get(self, url, *a, headers=None, **k): + async def get(self, url, *a, headers=None, **k): captured["url"] = url captured["headers"] = headers return _Resp() - monkeypatch.setattr("httpx.Client", _Client) + monkeypatch.setattr("httpx.AsyncClient", _Client) resp = self.client.post( "/api/providers/validate", @@ -3699,6 +3858,91 @@ class TestValidateProviderCredential: assert captured["url"] == "https://text.example.com/v1/models" assert captured["headers"] == {"Authorization": "Bearer sk-secret"} + def test_local_endpoint_without_key_sends_no_auth_header(self, monkeypatch): + """No key → no Authorization header (keyless local servers unaffected).""" + captured = {} + + class _Resp: + status_code = 200 + is_success = True + + def json(self): + return {"data": []} + + class _Client: + def __init__(self, *a, **k): + pass + + async def __aenter__(self): + return self + + async def __aexit__(self, *a): + return False + + async def get(self, url, *a, headers=None, **k): + captured["headers"] = headers + return _Resp() + + monkeypatch.setattr("httpx.AsyncClient", _Client) + + self.client.post( + "/api/providers/validate", + json={"key": "OPENAI_BASE_URL", "value": "http://127.0.0.1:8000/v1"}, + ) + assert captured["headers"] is None + + def test_named_custom_endpoint_probe_is_async(self, monkeypatch): + """Custom endpoint validation must not block the dashboard event loop.""" + captured = {} + + class _Resp: + status_code = 200 + is_success = True + + def json(self): + return {"data": [{"id": "local-model"}]} + + class _Client: + def __init__(self, *args, **kwargs): + pass + + async def __aenter__(self): + return self + + async def __aexit__(self, *args): + return False + + async def get(self, url, *args, headers=None, **kwargs): + captured["url"] = url + captured["headers"] = headers + return _Resp() + + monkeypatch.setattr("httpx.AsyncClient", _Client) + + response = self.client.post( + "/api/providers/custom-endpoints/validate", + json={ + "name": "Local", + "base_url": "http://localhost:8000/v1", + "model": "local-model", + "api_key": "local-secret", + }, + ) + + assert response.json() == { + "ok": True, + "reachable": True, + "message": "", + "models": ["local-model"], + } + assert captured == { + "url": "http://localhost:8000/v1/models", + "headers": { + "Accept": "application/json", + "Authorization": "Bearer local-secret", + }, + } + class TestDesktopCronTicker: """The dashboard backend fires cron jobs itself only when desktop-spawned.""" @@ -3777,6 +4021,81 @@ class TestServeIndexMissingIndex: assert "SPA-rebuilt" in resp.text +class TestHashedAssetCacheHeaders: + """Hashed /assets/* responses must be immutable-cacheable; index.html + must stay no-store so it always references the current hashes + (salvaged from PR #28543).""" + + _IMMUTABLE = "public, max-age=31536000, immutable" + + @staticmethod + def _client(tmp_path, monkeypatch): + from fastapi import FastAPI + from starlette.testclient import TestClient + import hermes_cli.web_server as ws + + dist = tmp_path / "web_dist" + (dist / "assets").mkdir(parents=True) + (dist / "index.html").write_text( + "SPA", encoding="utf-8" + ) + (dist / "assets" / "index-abc123.js").write_text( + "console.log('bundle');", encoding="utf-8" + ) + (dist / "assets" / "index-abc123.css").write_text( + "body{background:url(/ds-assets/bg.png);" + "font-family:url(/fonts-terminal/x.woff2)}", + encoding="utf-8", + ) + monkeypatch.setattr(ws, "WEB_DIST", dist) + monkeypatch.delenv("HERMES_SERVE_HEADLESS", raising=False) + spa_app = FastAPI() + ws.mount_spa(spa_app) + return TestClient(spa_app) + + def test_hashed_js_asset_is_immutable(self, tmp_path, monkeypatch): + client = self._client(tmp_path, monkeypatch) + resp = client.get("/assets/index-abc123.js") + assert resp.status_code == 200 + assert resp.headers["cache-control"] == self._IMMUTABLE + + def test_serve_css_is_immutable_and_keeps_prefix_rewrites( + self, tmp_path, monkeypatch + ): + client = self._client(tmp_path, monkeypatch) + resp = client.get("/assets/index-abc123.css") + assert resp.status_code == 200 + assert resp.headers["cache-control"] == self._IMMUTABLE + + # The proxy-prefix rewrite path (main's ds-assets/fonts-terminal + # handling) must survive the header change. + prefixed = client.get( + "/assets/index-abc123.css", + headers={"X-Forwarded-Prefix": "/hermes"}, + ) + assert prefixed.status_code == 200 + assert prefixed.headers["cache-control"] == self._IMMUTABLE + assert "url(/hermes/ds-assets/bg.png)" in prefixed.text + assert "url(/hermes/fonts-terminal/x.woff2)" in prefixed.text + + def test_index_html_stays_no_store(self, tmp_path, monkeypatch): + client = self._client(tmp_path, monkeypatch) + for route in ("/", "/chat"): + resp = client.get(route) + assert resp.status_code == 200 + cache_control = resp.headers["cache-control"] + assert "no-store" in cache_control + assert "immutable" not in cache_control + + def test_missing_asset_is_not_marked_immutable(self, tmp_path, monkeypatch): + """A 404 must never be cached for a year — a later rebuild can + legitimately create the file.""" + client = self._client(tmp_path, monkeypatch) + resp = client.get("/assets/nope-000000.js") + assert resp.status_code == 404 + assert "immutable" not in resp.headers.get("cache-control", "") + + class TestDashboardComponentHealth: """Component-health rollup: error middleware, /api/status components, self-test.""" diff --git a/tests/hermes_cli/test_web_server_session_search.py b/tests/hermes_cli/test_web_server_session_search.py index 46a0f98670..a3820b5f67 100644 --- a/tests/hermes_cli/test_web_server_session_search.py +++ b/tests/hermes_cli/test_web_server_session_search.py @@ -14,6 +14,7 @@ class _FakeSessionDB: closed = False opened_read_only = None + requested_fields = None def __init__(self, *args, **kwargs): type(self).opened_read_only = kwargs.get("read_only") @@ -57,8 +58,16 @@ class _FakeSessionDB: ) ][:limit] - def search_messages(self, query, source_filter=None, exclude_sources=None, limit=20): + def search_messages( + self, + query, + source_filter=None, + exclude_sources=None, + limit=20, + fields=None, + ): assert query == "20260603*" + type(self).requested_fields = fields rows = [ { "session_id": "20260603_090200_exact", @@ -98,10 +107,13 @@ class _FakeSessionDB: def test_desktop_session_search_merges_id_matches_before_content_matches(monkeypatch): _FakeSessionDB.opened_read_only = None + _FakeSessionDB.requested_fields = None monkeypatch.setattr("hermes_state.SessionDB", _FakeSessionDB) response = asyncio.run(web_server.search_sessions(q="20260603", limit=2)) + assert _FakeSessionDB.requested_fields is not None + assert "context" not in _FakeSessionDB.requested_fields # ID match surfaces first; the content hit on the SAME session is deduped # by lineage root (not double-listed); the unrelated content hit follows. assert response == { diff --git a/tests/hermes_cli/test_web_ui_build.py b/tests/hermes_cli/test_web_ui_build.py index 6e736bb931..e12e9050c6 100644 --- a/tests/hermes_cli/test_web_ui_build.py +++ b/tests/hermes_cli/test_web_ui_build.py @@ -150,7 +150,7 @@ class TestBuildWebUISkipsWhenFresh: assert result is True args, kwargs = mock_run.call_args assert "--workspace" not in args[0] - assert args[0] == ["/usr/bin/npm", "ci", "--include=dev", "--silent"] + assert args[0] == ["/usr/bin/npm", "ci", "--include=dev", "--silent", "--prefer-offline"] assert kwargs["cwd"] == web_dir def test_web_build_uses_idle_timeout_helper(self, tmp_path): diff --git a/tests/hermes_state/test_append_messages_batch.py b/tests/hermes_state/test_append_messages_batch.py new file mode 100644 index 0000000000..a65436ee93 --- /dev/null +++ b/tests/hermes_state/test_append_messages_batch.py @@ -0,0 +1,194 @@ +"""Tests for SessionDB.append_messages_batch (#23254 salvage). + +The batch writer reuses _insert_message_rows (the same row-serialization +path as replace/compact/import), runs the same admission guards as +append_message, is atomic (all rows or none), and aggregates the session +counters in one UPDATE. +""" + +import json +import sqlite3 + +import pytest + +from hermes_state import ( + CompressionSessionClosedError, + SessionDB, +) + + +@pytest.fixture() +def db(tmp_path): + d = SessionDB(db_path=tmp_path / "state.db") + d.create_session("sess-batch", source="cli") + yield d + d.close() + + +def _turn_messages(): + return [ + {"role": "user", "content": "question"}, + { + "role": "assistant", + "content": "let me check", + "tool_calls": [{"name": "terminal", "arguments": "{}"}], + "reasoning_content": "thinking...", + "finish_reason": "tool_calls", + }, + { + "role": "tool", + "content": "tool output", + "tool_name": "terminal", + "tool_call_id": "call_1", + }, + {"role": "assistant", "content": "answer", "finish_reason": "stop"}, + ] + + +class TestAppendMessagesBatch: + def test_batch_rows_identical_to_single_appends(self, db, tmp_path): + """The batch writer stores the same bytes append_message would.""" + db2 = SessionDB(db_path=tmp_path / "state2.db") + db2.create_session("sess-batch", source="cli") + try: + msgs = _turn_messages() + db.append_messages_batch("sess-batch", msgs) + for m in msgs: + role = m["role"] + db2.append_message( + session_id="sess-batch", + role=role, + content=m.get("content"), + tool_name=m.get("tool_name"), + tool_calls=m.get("tool_calls"), + tool_call_id=m.get("tool_call_id"), + finish_reason=m.get("finish_reason"), + reasoning_content=( + m.get("reasoning_content") if role == "assistant" else None + ), + ) + cols = ( + "role, content, tool_call_id, tool_calls, tool_name, " + "finish_reason, reasoning_content, observed, active" + ) + rows_a = db._conn.execute( + f"SELECT {cols} FROM messages ORDER BY id" + ).fetchall() + rows_b = db2._conn.execute( + f"SELECT {cols} FROM messages ORDER BY id" + ).fetchall() + assert [tuple(r) for r in rows_a] == [tuple(r) for r in rows_b] + finally: + db2.close() + + def test_reasoning_gated_to_assistant_rows(self, db): + """_insert_message_rows role-gates reasoning fields; a tool row + carrying reasoning keys must not persist them.""" + db.append_messages_batch( + "sess-batch", + [ + { + "role": "tool", + "content": "out", + "tool_name": "t", + "tool_call_id": "c1", + "reasoning_content": "should not persist", + } + ], + ) + row = db._conn.execute( + "SELECT reasoning_content FROM messages" + ).fetchone() + assert row[0] is None + + def test_counters_aggregate_once(self, db): + db.append_messages_batch("sess-batch", _turn_messages()) + row = db._conn.execute( + "SELECT message_count, tool_call_count FROM sessions WHERE id = ?", + ("sess-batch",), + ).fetchone() + assert row["message_count"] == 4 + assert row["tool_call_count"] == 1 + + def test_returns_inserted_count(self, db): + assert db.append_messages_batch("sess-batch", _turn_messages()) == 4 + + def test_empty_batch_is_noop(self, db): + assert db.append_messages_batch("sess-batch", []) == 0 + row = db._conn.execute( + "SELECT message_count FROM sessions WHERE id = ?", ("sess-batch",) + ).fetchone() + assert row["message_count"] == 0 + + def test_atomicity_all_or_nothing(self, db, monkeypatch): + """A failure mid-batch leaves ZERO rows and untouched counters.""" + real_insert = SessionDB._insert_message_rows + + def failing_insert(self_db, conn, session_id, messages): + real_conn_execute = conn.execute + calls = {"n": 0} + + def exec_counting(sql, *args): + if sql.lstrip().startswith("INSERT INTO messages"): + calls["n"] += 1 + if calls["n"] == 3: + raise sqlite3.OperationalError("boom mid-batch") + return real_conn_execute(sql, *args) + + conn.execute = exec_counting + try: + return real_insert(self_db, conn, session_id, messages) + finally: + conn.execute = real_conn_execute + + monkeypatch.setattr(SessionDB, "_insert_message_rows", failing_insert) + with pytest.raises(sqlite3.OperationalError): + db.append_messages_batch("sess-batch", _turn_messages()) + monkeypatch.undo() + + count = db._conn.execute("SELECT COUNT(*) FROM messages").fetchone()[0] + assert count == 0 + row = db._conn.execute( + "SELECT message_count, tool_call_count FROM sessions WHERE id = ?", + ("sess-batch",), + ).fetchone() + assert row["message_count"] == 0 + assert row["tool_call_count"] == 0 + + def test_compression_closed_session_rejected(self, db): + db._conn.execute( + "UPDATE sessions SET ended_at = 1.0, end_reason = 'compression' " + "WHERE id = ?", + ("sess-batch",), + ) + db._conn.commit() + with pytest.raises(CompressionSessionClosedError): + db.append_messages_batch("sess-batch", _turn_messages()) + + def test_multimodal_content_encoded(self, db): + msgs = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "look"}, + {"type": "image_url", "image_url": {"url": "data:x"}}, + ], + } + ] + db.append_messages_batch("sess-batch", msgs) + raw = db._conn.execute("SELECT content FROM messages").fetchone()[0] + # encoded via _encode_content — same sentinel prefix as append_message + loaded = db.get_messages("sess-batch") + assert loaded, raw + + def test_tool_calls_json_string_not_double_encoded(self, db): + msgs = [ + { + "role": "assistant", + "content": "x", + "tool_calls": json.dumps([{"name": "t", "arguments": "{}"}]), + } + ] + db.append_messages_batch("sess-batch", msgs) + raw = db._conn.execute("SELECT tool_calls FROM messages").fetchone()[0] + assert json.loads(raw) == [{"name": "t", "arguments": "{}"}] diff --git a/tests/honcho_plugin/test_session.py b/tests/honcho_plugin/test_session.py index 201979a13a..6583fdc800 100644 --- a/tests/honcho_plugin/test_session.py +++ b/tests/honcho_plugin/test_session.py @@ -791,6 +791,10 @@ class TestTrivialPromptHeuristic: for t in ("ok", "OK", " ok ", "y", "yes", "sure", "thanks", "lgtm", "/help", "", " "): assert HonchoMemoryProvider._is_trivial_prompt(t), f"expected trivial: {t!r}" + def test_classifier_catches_greetings(self): + """Greeting words must register as trivial so context injection is skipped.""" + for t in ("hi", "HI", "hey", "hello", "yo", "sup", " hi ", "hey!", "hello."): + assert HonchoMemoryProvider._is_trivial_prompt(t), f"expected trivial: {t!r}" def test_prefetch_skips_on_trivial_prompt(self): provider = self._make_provider() @@ -880,7 +884,7 @@ class TestDialecticCadenceAdvancesOnSuccess: provider._turn_count = 5 provider._last_dialectic_turn = 0 - provider.queue_prefetch("hello") + provider.queue_prefetch("what changed in the repo today") if provider._prefetch_thread: provider._prefetch_thread.join(timeout=2.0) @@ -905,7 +909,7 @@ class TestDialecticCadenceAdvancesOnSuccess: provider._prefetch_thread = fresh provider._prefetch_thread_started_at = _time.monotonic() # fresh start - provider.queue_prefetch("hello") + provider.queue_prefetch("what changed in the repo today") # Should have short-circuited — no new dialectic call assert provider._manager.dialectic_query.call_count == 0 hold.set() @@ -1011,7 +1015,7 @@ class TestDialecticLiveness: # timeout=2.0, multiplier=2.0, so anything older than 4s is stale p._prefetch_thread_started_at = 0.0 # very old (1970 monotonic baseline) - p.queue_prefetch("hello") + p.queue_prefetch("what changed in the repo today") # New thread should have been spawned since stuck one is stale assert p._prefetch_thread is not stuck, "stale thread must be recycled" if p._prefetch_thread: diff --git a/tests/openviking_plugin/test_openviking.py b/tests/openviking_plugin/test_openviking.py index 18b80b2020..7bdd525a72 100644 --- a/tests/openviking_plugin/test_openviking.py +++ b/tests/openviking_plugin/test_openviking.py @@ -1,6 +1,7 @@ """Tests for plugins/memory/openviking/__init__.py — URI normalization and payload handling.""" import json +import os import threading import time from http.server import BaseHTTPRequestHandler, HTTPServer @@ -15,7 +16,8 @@ def _write_skill(skills_dir, name, body="Do the thing."): skill_dir = skills_dir / name skill_dir.mkdir(parents=True, exist_ok=True) (skill_dir / "SKILL.md").write_text( - f"---\nname: {name}\ndescription: Description for {name}\n---\n\n# {name}\n\n{body}\n" + f"---\nname: {name}\ndescription: Description for {name}\n---\n\n# {name}\n\n{body}\n", + encoding="utf-8", ) return skill_dir @@ -24,7 +26,10 @@ def _write_bundle(bundles_dir, slug, skills): bundles_dir.mkdir(parents=True, exist_ok=True) lines = [f"name: {slug}", "skills:"] lines.extend(f" - {skill}" for skill in skills) - (bundles_dir / f"{slug}.yaml").write_text("\n".join(lines) + "\n") + (bundles_dir / f"{slug}.yaml").write_text( + "\n".join(lines) + "\n", + encoding="utf-8", + ) class FakeVikingClient: @@ -265,6 +270,7 @@ class TestOpenVikingConfigSchema: provider = OpenVikingMemoryProvider() schema = provider.get_config_schema() + fields = {entry["key"]: entry for entry in schema} env_vars = {entry.get("env_var") for entry in schema} assert "OPENVIKING_RECALL_LIMIT" in env_vars @@ -275,6 +281,11 @@ class TestOpenVikingConfigSchema: assert "OPENVIKING_RECALL_FULL_READ_LIMIT" in env_vars assert "OPENVIKING_RECALL_PREFER_ABSTRACT" in env_vars assert "OPENVIKING_RECALL_RESOURCES" in env_vars + assert fields["recall_limit"]["type"] == "integer" + assert fields["recall_limit"]["minimum"] == 1 + assert fields["recall_limit"]["maximum"] == 100 + assert fields["recall_score_threshold"]["type"] == "number" + assert fields["recall_prefer_abstract"]["type"] == "boolean" assert provider._recall_config() == { "limit": 6, "score_threshold": 0.15, @@ -286,6 +297,166 @@ class TestOpenVikingConfigSchema: "resources": False, } + def test_recall_config_reads_from_config_yaml(self, monkeypatch, tmp_path): + """_recall_config() reads memory.openviking values from config.yaml when + the corresponding OPENVIKING_RECALL_* env vars are not set.""" + # Populate config.yaml in the temp HERMES_HOME + hermes_home = tmp_path / "hermes_test" + hermes_home.mkdir(exist_ok=True) + config_yaml = hermes_home / "config.yaml" + config_yaml.write_text( + """\ +memory: + provider: openviking + openviking: + recall_limit: 12 + recall_score_threshold: 0.42 + recall_max_injected_chars: 8000 + profile_token_budget: 7000 + recall_timeout_seconds: 2.0 + recall_request_timeout_seconds: 1.5 + recall_full_read_limit: 5 + recall_prefer_abstract: true + recall_resources: true +""", + encoding="utf-8", + ) + monkeypatch.setenv("HERMES_HOME", str(hermes_home)) + # Clear any OPENVIKING_RECALL_* env vars so config.yaml prevails + for key in list(os.environ): + if key.startswith("OPENVIKING_RECALL_"): + monkeypatch.delenv(key, raising=False) + + provider = OpenVikingMemoryProvider() + cfg = provider._recall_config() + + assert cfg["limit"] == 12 + assert cfg["score_threshold"] == 0.42 + assert cfg["max_injected_chars"] == 8000 + assert cfg["timeout_seconds"] == 2.0 + assert cfg["request_timeout_seconds"] == 1.5 + assert cfg["full_read_limit"] == 5 + assert cfg["prefer_abstract"] is True + assert cfg["resources"] is True + assert provider._profile_token_budget() == 7000 + + def test_recall_config_env_overrides_config_yaml(self, monkeypatch, tmp_path): + """Env vars OPENVIKING_RECALL_* take precedence over config.yaml values + when both are present.""" + hermes_home = tmp_path / "hermes_test" + hermes_home.mkdir(exist_ok=True) + config_yaml = hermes_home / "config.yaml" + config_yaml.write_text( + """\ +memory: + provider: openviking + openviking: + recall_limit: 12 + recall_resources: true +""", + encoding="utf-8", + ) + monkeypatch.setenv("HERMES_HOME", str(hermes_home)) + # Override config.yaml via env + monkeypatch.setenv("OPENVIKING_RECALL_LIMIT", "6") + monkeypatch.setenv("OPENVIKING_RECALL_RESOURCES", "false") + + provider = OpenVikingMemoryProvider() + cfg = provider._recall_config() + + assert cfg["limit"] == 6, "env var should override config.yaml" + assert cfg["resources"] is False, "env var false should override config.yaml true" + + def test_recall_config_partial_config_yaml(self, monkeypatch, tmp_path): + """Partially populated config.yaml falls back to defaults for omitted keys + and env vars can override individual fields.""" + hermes_home = tmp_path / "hermes_test" + hermes_home.mkdir(exist_ok=True) + config_yaml = hermes_home / "config.yaml" + config_yaml.write_text( + """\ +memory: + provider: openviking + openviking: + recall_limit: 3 + # No recall_resources set — should use default (False) +""", + encoding="utf-8", + ) + monkeypatch.setenv("HERMES_HOME", str(hermes_home)) + for key in list(os.environ): + if key.startswith("OPENVIKING_RECALL_"): + monkeypatch.delenv(key, raising=False) + + provider = OpenVikingMemoryProvider() + cfg = provider._recall_config() + + assert cfg["limit"] == 3, "config.yaml value should be picked up" + assert cfg["resources"] is False, "omitted key should use default" + assert cfg["timeout_seconds"] == 4.0, "omitted key should use built-in default" + + def test_dashboard_shaped_string_values_are_typed(self, monkeypatch): + for key in list(os.environ): + if key.startswith("OPENVIKING_RECALL_") or key == "OPENVIKING_PROFILE_TOKEN_BUDGET": + monkeypatch.delenv(key, raising=False) + monkeypatch.setattr( + openviking_plugin, + "_load_hermes_openviking_config", + lambda: { + "recall_limit": "12", + "recall_score_threshold": "0.42", + "recall_prefer_abstract": "false", + "recall_resources": "true", + "profile_token_budget": "7500", + }, + ) + provider = OpenVikingMemoryProvider() + + cfg = provider._recall_config() + + assert cfg["limit"] == 12 + assert cfg["score_threshold"] == 0.42 + assert cfg["prefer_abstract"] is False + assert cfg["resources"] is True + assert provider._profile_token_budget() == 7500 + + def test_invalid_recall_values_fall_back_without_type_errors(self, monkeypatch): + for key in list(os.environ): + if key.startswith("OPENVIKING_RECALL_") or key == "OPENVIKING_PROFILE_TOKEN_BUDGET": + monkeypatch.delenv(key, raising=False) + monkeypatch.setattr( + openviking_plugin, + "_load_hermes_openviking_config", + lambda: { + "recall_limit": "many", + "recall_score_threshold": True, + "recall_prefer_abstract": "sometimes", + "profile_token_budget": "7.5", + }, + ) + provider = OpenVikingMemoryProvider() + + cfg = provider._recall_config() + + assert cfg["limit"] == 6 + assert cfg["score_threshold"] == 0.15 + assert cfg["prefer_abstract"] is False + assert provider._profile_token_budget() == 6000 + + def test_recall_env_overrides_string_config_with_native_types(self, monkeypatch): + monkeypatch.setattr( + openviking_plugin, + "_load_hermes_openviking_config", + lambda: {"recall_limit": "12", "recall_resources": "false"}, + ) + monkeypatch.setenv("OPENVIKING_RECALL_LIMIT", "4") + monkeypatch.setenv("OPENVIKING_RECALL_RESOURCES", "true") + + cfg = OpenVikingMemoryProvider()._recall_config() + + assert cfg["limit"] == 4 + assert cfg["resources"] is True + class TestOpenVikingTurnConversion: def test_extract_current_turn_anchors_on_latest_matching_user_and_assistant(self): @@ -481,7 +652,7 @@ class TestOpenVikingAutoRecallPrefetch: def do_GET(self): parsed = urlparse(self.path) if parsed.path == "/health": - self._send_json({"healthy": True}) + self._send_json({"status": "ok", "healthy": True, "version": "test"}) return if parsed.path == "/api/v1/content/read": query = parse_qs(parsed.query) @@ -845,7 +1016,7 @@ class TestEnsureClientReloadsEnv: start_calls.append(endpoint) first_start_entered.set() release_start.wait(timeout=2) - return True, "started" + return openviking_plugin._LOCAL_SERVER_STARTED, "started" monkeypatch.setattr(openviking_plugin, "_start_local_openviking_server", start_local) monkeypatch.setattr( @@ -947,3 +1118,172 @@ class TestEnsureClientFailureHardening: assert provider._conn_snapshot == healthy_snapshot built = provider._new_client() assert built.endpoint == "https://up.example" + + +class TestUnavailableWarningsPromiseRetry: + """Every "OpenViking is unavailable" warning must describe what actually + happens next. + + ``_ensure_client()`` rebuilds and re-probes the client whenever the + resolved config changes or the failed-config cooldown has elapsed, so no + warning may tell the user memory is off for the rest of the run — that + reads as "it never recovers" and sends people restarting hermes for + nothing (#5721). + """ + + @staticmethod + def _assert_promises_retry(message: str) -> None: + assert "for this Hermes run" not in message, message + assert "will retry on a later access" in message, message + assert "when the config changes" in message, message + + @staticmethod + def _stub_client(health_result): + class _StubClient: + def __init__(self, endpoint, api_key="", account="", user="", agent=""): + self.endpoint = endpoint + + def health(self): + return health_result + + return _StubClient + + def test_local_autostart_timeout_warning(self): + self._assert_promises_retry( + openviking_plugin._runtime_openviking_timeout_message("http://127.0.0.1:1934") + ) + + def test_remote_unreachable_warning(self): + provider = OpenVikingMemoryProvider() + provider._endpoint = "https://remote.example" + warnings: list[str] = [] + + provider._handle_runtime_openviking_unreachable(warning_callback=warnings.append) + + assert provider._client is None + assert len(warnings) == 1 + self._assert_promises_retry(warnings[0]) + + def test_local_autostart_refused_warning(self, monkeypatch): + monkeypatch.setattr( + openviking_plugin, + "_start_local_openviking_server", + lambda endpoint: ( + openviking_plugin._LOCAL_SERVER_FAILED, + "openviking-server was not found on PATH.", + ), + ) + provider = OpenVikingMemoryProvider() + provider._endpoint = "http://127.0.0.1:1934" + warnings: list[str] = [] + + provider._handle_runtime_openviking_unreachable(warning_callback=warnings.append) + + assert provider._client is None + assert len(warnings) == 1 + self._assert_promises_retry(warnings[0]) + + def test_still_unhealthy_after_autostart_warning(self, monkeypatch): + monkeypatch.setattr(openviking_plugin, "_VikingClient", self._stub_client(False)) + monkeypatch.setattr( + openviking_plugin, "_wait_for_openviking_health", lambda endpoint, **kwargs: True + ) + provider = OpenVikingMemoryProvider() + provider._endpoint = "http://127.0.0.1:1934" + warnings: list[str] = [] + + provider._finish_runtime_openviking_start(warning_callback=warnings.append) + + assert provider._client is None + assert len(warnings) == 1 + self._assert_promises_retry(warnings[0]) + + def test_attach_failure_after_autostart_warning(self, monkeypatch): + def _explode(*args, **kwargs): + raise RuntimeError("connection reset by peer") + + monkeypatch.setattr(openviking_plugin, "_VikingClient", _explode) + monkeypatch.setattr( + openviking_plugin, "_wait_for_openviking_health", lambda endpoint, **kwargs: True + ) + provider = OpenVikingMemoryProvider() + provider._endpoint = "http://127.0.0.1:1934" + warnings: list[str] = [] + + provider._finish_runtime_openviking_start(warning_callback=warnings.append) + + assert provider._client is None + assert len(warnings) == 1 + self._assert_promises_retry(warnings[0]) + + def test_initialize_responded_unhealthy_warning(self, monkeypatch, tmp_path): + class _UnhealthyClient: + def __init__(self, endpoint, api_key="", account="", user="", agent=""): + self.endpoint = endpoint + + def health_payload(self): + return {"healthy": False} + + def health(self): + return False + + monkeypatch.setenv("HERMES_HOME", str(tmp_path / ".hermes")) + monkeypatch.setenv("OPENVIKING_ENDPOINT", "https://sick.example") + monkeypatch.setattr(openviking_plugin, "_VikingClient", _UnhealthyClient) + provider = OpenVikingMemoryProvider() + warnings: list[str] = [] + + provider.initialize("session-1", platform="cli", warning_callback=warnings.append) + + assert provider._client is None + assert len(warnings) == 1 + self._assert_promises_retry(warnings[0]) + + def test_ensure_client_responded_unhealthy_warning(self, monkeypatch, caplog): + class _UnhealthyClient: + def __init__(self, endpoint, api_key="", account="", user="", agent=""): + self.endpoint = endpoint + + def health_payload(self): + return {"healthy": False} + + monkeypatch.setenv("OPENVIKING_ENDPOINT", "https://sick.example") + monkeypatch.setattr(openviking_plugin, "_VikingClient", _UnhealthyClient) + provider = OpenVikingMemoryProvider() + provider._env_refresh_enabled = True + + with caplog.at_level("WARNING", logger=openviking_plugin.__name__): + assert provider._ensure_client() is None + + self._assert_promises_retry(caplog.text) + + def test_startup_failure_really_does_reconnect_on_a_later_access( + self, monkeypatch, tmp_path + ): + """The warnings promise a retry — prove the provider delivers one.""" + probes: list[str] = [] + + class _FlakyClient: + def __init__(self, endpoint, api_key="", account="", user="", agent=""): + self.endpoint = endpoint + + def health(self): + probes.append(self.endpoint) + return len(probes) > 1 # down at startup, up on the next access + + monkeypatch.setenv("HERMES_HOME", str(tmp_path / ".hermes")) + monkeypatch.setenv("OPENVIKING_ENDPOINT", "https://remote.example") + monkeypatch.setattr(openviking_plugin, "_VikingClient", _FlakyClient) + provider = OpenVikingMemoryProvider() + warnings: list[str] = [] + + provider.initialize("session-1", platform="cli", warning_callback=warnings.append) + assert provider._client is None + assert len(warnings) == 1 + self._assert_promises_retry(warnings[0]) + + # A startup failure arms no cooldown, so the very next access re-probes. + client = provider._ensure_client() + assert client is not None + assert client.endpoint == "https://remote.example" + assert len(probes) == 2 diff --git a/tests/plugins/memory/test_openviking_endpoint_always_blocked.py b/tests/plugins/memory/test_openviking_endpoint_always_blocked.py new file mode 100644 index 0000000000..729f3749d1 --- /dev/null +++ b/tests/plugins/memory/test_openviking_endpoint_always_blocked.py @@ -0,0 +1,68 @@ +"""OpenViking endpoint always-blocked floor.""" + +import pytest + +from plugins.memory.openviking import ( + _OpenVikingEndpointError, + _local_openviking_bind, + _normalize_openviking_url, + _openviking_endpoint_is_always_blocked, +) + + +def test_openviking_blocks_metadata_endpoint(): + with pytest.raises(_OpenVikingEndpointError, match="blocked metadata address"): + _normalize_openviking_url("http://169.254.169.254/") + + +def test_openviking_keeps_default_loopback(): + assert _normalize_openviking_url("http://127.0.0.1:1933") == "http://127.0.0.1:1933" + + +@pytest.mark.parametrize("host", ["localhost", "127.0.0.1"]) +def test_openviking_bare_loopback_health_and_autostart_use_same_default_port(host): + endpoint = _normalize_openviking_url(host) + + assert endpoint == f"http://{host}:1933" + assert _local_openviking_bind(endpoint) == (host, 1933) + + +def test_openviking_explicit_loopback_url_preserves_implicit_http_port(): + assert _normalize_openviking_url("http://localhost") == "http://localhost" + + +def test_openviking_blocks_ecs_metadata_hostname(): + with pytest.raises(_OpenVikingEndpointError, match="blocked metadata address"): + _normalize_openviking_url("http://metadata.google.internal/computeMetadata/v1/") + + +def test_openviking_rejects_endpoint_credentials_and_query(): + with pytest.raises(_OpenVikingEndpointError, match="cannot contain user info"): + _normalize_openviking_url("https://user:secret@example.com?api_key=secret") + + +def test_openviking_validates_shorthand_ipv6_port(): + assert _normalize_openviking_url("::1:1934") == "http://[::1]:1934" + with pytest.raises(_OpenVikingEndpointError, match="Port could not be cast"): + _normalize_openviking_url("::1:not-a-port") + + +def test_openviking_caches_safety_check_for_unchanged_endpoint(monkeypatch): + import tools.url_safety as url_safety + + calls = [] + _openviking_endpoint_is_always_blocked.cache_clear() + monkeypatch.setattr( + url_safety, + "is_always_blocked_url", + lambda value: calls.append(value) or False, + ) + + assert _normalize_openviking_url("https://openviking.example.test") == ( + "https://openviking.example.test" + ) + assert _normalize_openviking_url("https://openviking.example.test") == ( + "https://openviking.example.test" + ) + assert calls == ["https://openviking.example.test"] + _openviking_endpoint_is_always_blocked.cache_clear() diff --git a/tests/plugins/memory/test_openviking_provider.py b/tests/plugins/memory/test_openviking_provider.py index 2b8d92e61d..e28c87f1a5 100644 --- a/tests/plugins/memory/test_openviking_provider.py +++ b/tests/plugins/memory/test_openviking_provider.py @@ -1,5 +1,6 @@ import json import os +import socket import stat import threading import time @@ -116,12 +117,41 @@ def test_openviking_provider_config_loader_uses_readonly_config(monkeypatch): assert config is not backing_config["memory"]["openviking"] +def test_connection_settings_read_dashboard_config_file(tmp_path, monkeypatch): + _clear_openviking_env(monkeypatch) + hermes_home = tmp_path / "hermes" + hermes_home.mkdir() + (hermes_home / "config.yaml").write_text( + """\ +memory: + provider: openviking + openviking: + endpoint: http://saved.test:1933 + account: saved-account + user: saved-user + agent: saved-agent +""", + encoding="utf-8", + ) + monkeypatch.setenv("HERMES_HOME", str(hermes_home)) + + settings = openviking_module._resolve_connection_settings( + openviking_module._load_hermes_openviking_config() + ) + + assert settings["endpoint"] == "http://saved.test:1933" + assert settings["account"] == "saved-account" + assert settings["user"] == "saved-user" + assert settings["agent"] == "saved-agent" + assert settings["api_key"] == "" + + def test_linked_ovcli_config_is_read_at_runtime(tmp_path, monkeypatch): _clear_openviking_env(monkeypatch) ovcli_path = tmp_path / "ovcli.conf" ovcli_path.write_text( json.dumps({ - "url": "http://openviking-one.local", + "url": "http://openviking-one.test", "api_key": "key-one", "account": "acct-one", "user": "alice", @@ -134,7 +164,7 @@ def test_linked_ovcli_config_is_read_at_runtime(tmp_path, monkeypatch): settings = openviking_module._resolve_connection_settings(provider_config) assert settings == { - "endpoint": "http://openviking-one.local", + "endpoint": "http://openviking-one.test", "api_key": "key-one", "account": "", "user": "", @@ -143,7 +173,7 @@ def test_linked_ovcli_config_is_read_at_runtime(tmp_path, monkeypatch): ovcli_path.write_text( json.dumps({ - "url": "http://openviking-two.local", + "url": "http://openviking-two.test", "api_key": "key-two", "agent_id": "agent-two", }), @@ -153,7 +183,7 @@ def test_linked_ovcli_config_is_read_at_runtime(tmp_path, monkeypatch): settings = openviking_module._resolve_connection_settings(provider_config) assert settings == { - "endpoint": "http://openviking-two.local", + "endpoint": "http://openviking-two.test", "api_key": "key-two", "account": "", "user": "", @@ -161,6 +191,42 @@ def test_linked_ovcli_config_is_read_at_runtime(tmp_path, monkeypatch): } +def test_linked_ovcli_without_url_falls_through_to_dashboard_endpoint(tmp_path, monkeypatch): + _clear_openviking_env(monkeypatch) + ovcli_path = tmp_path / "ovcli.conf" + ovcli_path.write_text(json.dumps({"api_key": "linked-key"}), encoding="utf-8") + + settings = openviking_module._resolve_connection_settings({ + "use_ovcli_config": True, + "ovcli_config_path": str(ovcli_path), + "endpoint": "http://saved.test:1933", + }) + + assert settings["endpoint"] == "http://saved.test:1933" + assert settings["api_key"] == "linked-key" + + +def test_profile_discovery_warns_when_skipping_unsafe_ovcli_endpoint(tmp_path, caplog): + profile_path = tmp_path / "ovcli.conf.blocked" + profile_path.write_text( + json.dumps({"url": "http://169.254.169.254/latest/meta-data"}), + encoding="utf-8", + ) + + with caplog.at_level("WARNING", logger=openviking_module.__name__): + assert ( + openviking_module._load_profile( + profile_path, + source="saved", + name="blocked", + ) + is None + ) + + assert "Skipping invalid OpenViking CLI config" in caplog.text + assert str(profile_path) in caplog.text + + def test_connection_values_omit_stale_identity_for_user_key_with_root_key(): values = openviking_module._connection_values_from_ovcli({ "url": "https://openviking.example", @@ -177,11 +243,11 @@ def test_connection_values_omit_stale_identity_for_user_key_with_root_key(): def test_link_ovcli_profile_removes_stale_inline_config(tmp_path): env_path = tmp_path / ".env" - env_path.write_text("OPENVIKING_ENDPOINT=http://old.local\nOTHER_KEY=keep\n", encoding="utf-8") + env_path.write_text("OPENVIKING_ENDPOINT=http://old.test\nOTHER_KEY=keep\n", encoding="utf-8") config = {"memory": {}} provider_config = { "use_ovcli_config": False, - "endpoint": "http://stale.local", + "endpoint": "http://stale.test", "api_key": "stale-key", "account": "default", "user": "default", @@ -210,12 +276,12 @@ def test_post_setup_existing_profile_picker_validates_and_links_saved_profile(tm hermes_home = tmp_path / "hermes" hermes_home.mkdir() env_path = hermes_home / ".env" - env_path.write_text("OPENVIKING_ENDPOINT=http://old.local\nOTHER_KEY=keep\n", encoding="utf-8") + env_path.write_text("OPENVIKING_ENDPOINT=http://old.test\nOTHER_KEY=keep\n", encoding="utf-8") openviking_home = tmp_path / ".openviking" openviking_home.mkdir() active_path = openviking_home / "ovcli.conf" saved_path = openviking_home / "ovcli.conf.VPS" - active_path.write_text(json.dumps({"url": "http://active.local"}), encoding="utf-8") + active_path.write_text(json.dumps({"url": "http://active.test"}), encoding="utf-8") saved_path.write_text( json.dumps({"url": "https://vps.example", "api_key": "user-key"}), encoding="utf-8", @@ -261,6 +327,51 @@ def test_post_setup_existing_profile_picker_validates_and_links_saved_profile(tm assert "OTHER_KEY=keep" in env_text +def test_local_setup_recommends_user_api_key_before_unauthenticated_mode(monkeypatch): + monkeypatch.setattr( + openviking_module, + "_validate_openviking_reachability", + lambda endpoint: (True, ""), + ) + monkeypatch.setattr( + openviking_module, + "_validate_openviking_setup_values", + lambda values, *, require_api_key=False: (True, "", "user"), + ) + credential_menu = {} + + def select(title, options, *, default=0, cancel_returns=None): + assert title == " OpenViking credential" + credential_menu["options"] = options + credential_menu["default"] = default + return 0 + + def prompt(label, default=None, secret=False): + if label == "OpenViking server URL": + return default + if label == "OpenViking user API key": + assert secret is True + return "user-key" + if label == openviking_module._AGENT_PROMPT_LABEL: + return default + raise AssertionError(f"Unexpected prompt: {label}") + + values = openviking_module._prompt_manual_connection_values( + prompt, + select, + -1, + ) + + assert [label for label, _description in credential_menu["options"]] == [ + "User API key", + "Root API key", + "No API key", + ] + assert credential_menu["default"] == 0 + assert values["api_key"] == "user-key" + assert values["api_key_type"] == "user" + + def test_start_local_openviking_server_uses_endpoint_host_and_port(monkeypatch): popen_calls = [] @@ -268,18 +379,137 @@ def test_start_local_openviking_server_uses_endpoint_host_and_port(monkeypatch): popen_calls.append((args, kwargs)) return object() + monkeypatch.setattr(openviking_module, "_local_openviking_port_is_open", lambda host, port: False) monkeypatch.setattr(openviking_module.shutil, "which", lambda name: "/usr/local/bin/openviking-server") monkeypatch.setattr(openviking_module.subprocess, "Popen", fake_popen) - started, message = openviking_module._start_local_openviking_server("http://127.0.0.1:1934") + state, message = openviking_module._start_local_openviking_server("http://127.0.0.1:1934") - assert started is True + assert state == openviking_module._LOCAL_SERVER_STARTED assert "127.0.0.1:1934" in message args, kwargs = popen_calls[0] assert args == ["/usr/local/bin/openviking-server", "--host", "127.0.0.1", "--port", "1934"] assert kwargs["start_new_session"] is True +def test_start_local_openviking_server_does_not_spawn_when_port_already_open(monkeypatch): + """A live listener means a second server would just die on DataDirectoryLocked.""" + probed = [] + + def fake_probe(host, port): + probed.append((host, port)) + return True + + monkeypatch.setattr(openviking_module, "_local_openviking_port_is_open", fake_probe) + monkeypatch.setattr( + openviking_module, + "_describe_local_port_listener", + lambda host, port: "python-test-server (PID 4242)", + ) + monkeypatch.setattr(openviking_module.shutil, "which", lambda name: "/usr/local/bin/openviking-server") + monkeypatch.setattr( + openviking_module.subprocess, + "Popen", + MagicMock(side_effect=AssertionError("must not spawn while a server is already listening")), + ) + + state, message = openviking_module._start_local_openviking_server("http://127.0.0.1:1934") + + assert state == openviking_module._LOCAL_SERVER_OCCUPIED + assert "python-test-server (PID 4242)" in message + assert "not passed OpenViking's /health check" in message + assert "already running" not in message + assert probed == [("127.0.0.1", 1934)] + + +def test_start_local_openviking_server_reports_occupied_port_without_cli_on_path(monkeypatch): + """The port probe outranks PATH but never claims the listener is OpenViking.""" + monkeypatch.setattr(openviking_module, "_local_openviking_port_is_open", lambda host, port: True) + monkeypatch.setattr( + openviking_module, + "_describe_local_port_listener", + lambda host, port: "an unidentified process", + ) + monkeypatch.setattr(openviking_module.shutil, "which", lambda name: None) + monkeypatch.setattr( + openviking_module.subprocess, + "Popen", + MagicMock(side_effect=AssertionError("must not spawn")), + ) + + state, message = openviking_module._start_local_openviking_server("http://127.0.0.1:1934") + + assert state == openviking_module._LOCAL_SERVER_OCCUPIED + assert "unidentified process" in message + + +def test_start_local_openviking_server_rejects_unparseable_url_before_probing(monkeypatch): + monkeypatch.setattr( + openviking_module, + "_local_openviking_port_is_open", + MagicMock(side_effect=AssertionError("must not probe an unparseable endpoint")), + ) + + state, message = openviking_module._start_local_openviking_server("http://127.0.0.1:not-a-port") + + assert state == openviking_module._LOCAL_SERVER_FAILED + assert "Could not parse local OpenViking URL" in message + + +def test_local_openviking_port_is_open_detects_listener_and_closed_port(): + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as listener: + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + _host, port = listener.getsockname() + assert openviking_module._local_openviking_port_is_open("127.0.0.1", port) is True + + # Socket closed: the same port no longer accepts connections. + assert openviking_module._local_openviking_port_is_open("127.0.0.1", port) is False + + +def test_describe_local_port_listener_reports_process(monkeypatch): + import psutil + + connection = SimpleNamespace( + status=psutil.CONN_LISTEN, + laddr=SimpleNamespace(ip="0.0.0.0", port=1934), + pid=4242, + ) + monkeypatch.setattr(psutil, "net_connections", lambda *, kind: [connection]) + monkeypatch.setattr( + psutil, + "Process", + lambda pid: SimpleNamespace(name=lambda: "postgres"), + ) + + assert openviking_module._describe_local_port_listener("127.0.0.1", 1934) == ( + "postgres (PID 4242)" + ) + + +def test_runtime_reports_occupied_port_and_does_not_wait_or_spawn(monkeypatch): + monkeypatch.setattr( + openviking_module, + "_start_local_openviking_server", + lambda endpoint: ( + openviking_module._LOCAL_SERVER_OCCUPIED, + "Port 127.0.0.1:1934 is occupied by postgres (PID 99).", + ), + ) + provider = OpenVikingMemoryProvider() + provider._endpoint = "http://127.0.0.1:1934" + provider._start_runtime_openviking_waiter = MagicMock() + warnings = [] + + provider._handle_runtime_openviking_unreachable(warning_callback=warnings.append) + + provider._start_runtime_openviking_waiter.assert_not_called() + assert provider._client is None + assert len(warnings) == 1 + assert "postgres (PID 99)" in warnings[0] + assert "temporarily unavailable" in warnings[0] + + def test_https_local_endpoint_is_not_runtime_autostart_eligible(monkeypatch): _clear_openviking_env(monkeypatch) monkeypatch.setenv("OPENVIKING_ENDPOINT", "https://localhost:1934") @@ -304,8 +534,9 @@ def test_https_local_endpoint_is_not_runtime_autostart_eligible(monkeypatch): assert provider._client is None assert warnings == [ - "Remote OpenViking server at https://localhost:1934 is not reachable; " - "OpenViking memory disabled for this Hermes run. " + "Remote OpenViking server at https://localhost:1934 is not reachable. " + "OpenViking memory is temporarily unavailable; Hermes will retry on a later access or when " + "the config changes. " "Check the configured endpoint and network connectivity." ] @@ -337,8 +568,9 @@ def test_runtime_does_not_autostart_when_local_server_reports_unhealthy(monkeypa assert provider._client is None assert warnings == [ - "OpenViking server at http://localhost:1934 responded but reported unhealthy status. " - "OpenViking memory disabled for this Hermes run." + "Service at http://localhost:1934 responded but reported unhealthy OpenViking status. " + "OpenViking memory is temporarily unavailable; Hermes will retry on a later access " + "or when the config changes." ] @@ -348,7 +580,10 @@ def test_handle_unreachable_endpoint_waits_long_enough_after_autostart(monkeypat monkeypatch.setattr( openviking_module, "_start_local_openviking_server", - lambda endpoint: (True, "Started openviking-server on 127.0.0.1:1934 in the background."), + lambda endpoint: ( + openviking_module._LOCAL_SERVER_STARTED, + "Started openviking-server on 127.0.0.1:1934 in the background.", + ), ) monkeypatch.setattr( openviking_module, @@ -388,7 +623,8 @@ def test_initialize_autostarts_local_openviking_in_background_when_runtime_healt monkeypatch.setattr( openviking_module, "_start_local_openviking_server", - lambda endpoint: start_calls.append(endpoint) or (True, "started"), + lambda endpoint: start_calls.append(endpoint) + or (openviking_module._LOCAL_SERVER_STARTED, "started"), ) monkeypatch.setattr( openviking_module, @@ -504,6 +740,157 @@ def test_viking_client_delete_uses_identity_headers(monkeypatch): assert captured["kwargs"]["headers"]["X-OpenViking-Actor-Peer"] == "hermes" +def test_openviking_identity_probes_are_anonymous_before_authenticated_requests(monkeypatch): + calls = [] + + def response(payload): + return SimpleNamespace(status_code=200, text="", json=lambda: payload) + + def fake_get(url, **kwargs): + calls.append((url, kwargs["headers"])) + if url.endswith("/health"): + return response({"status": "ok"}) + if url.endswith("/openapi.json"): + return response({"info": {"title": "OpenViking API"}}) + if url.endswith("/api/v1/system/status"): + return response({"status": "ok"}) + if url.endswith("/api/v1/admin/accounts"): + return response({"status": "ok", "result": []}) + raise AssertionError(f"unexpected request: {url}") + + monkeypatch.setattr( + openviking_module, + "_get_httpx", + lambda: SimpleNamespace(get=fake_get), + ) + + valid, message, role = openviking_module._validate_openviking_setup_values({ + "endpoint": "https://openviking.example", + "api_key": "secret-key", + "account": "acct", + "user": "alice", + "agent": "hermes", + }) + + assert (valid, message, role) == (True, "", "root") + assert [url.removeprefix("https://openviking.example") for url, _headers in calls] == [ + "/health", + "/openapi.json", + "/api/v1/system/status", + "/api/v1/admin/accounts", + ] + assert calls[0][1] == {"Accept": "application/json"} + assert calls[1][1] == {"Accept": "application/json"} + for _url, headers in calls[2:]: + assert headers["X-API-Key"] == "secret-key" + assert headers["Authorization"] == "Bearer secret-key" + + +def test_repeated_openviking_health_probes_never_send_identity_headers(monkeypatch): + captured_headers = [] + client = _VikingClient( + "https://openviking.example", + api_key="secret-key", + account="acct", + user="alice", + agent="hermes", + ) + + def fake_get(_url, **kwargs): + captured_headers.append(kwargs["headers"]) + return SimpleNamespace( + status_code=200, + text="", + json=lambda: {"status": "ok", "healthy": True, "version": "0.2.10"}, + ) + + monkeypatch.setattr(client._httpx, "get", fake_get) + + assert client.health() is True + assert client.health() is True + assert captured_headers == [ + {"Accept": "application/json"}, + {"Accept": "application/json"}, + ] + + +def test_modern_openviking_identity_does_not_probe_openapi(): + client = MagicMock() + client.health_payload.return_value = { + "status": "ok", + "healthy": True, + "version": "0.2.10", + } + + state, health = openviking_module._probe_openviking_identity(client) + + assert state == "modern" + assert health["version"] == "0.2.10" + client.openapi_payload.assert_not_called() + + +def test_legacy_health_requires_openviking_openapi_identity_before_auth(monkeypatch): + events = [] + + class ForeignServiceClient: + def __init__(self, *args, **kwargs): + pass + + def health_payload(self): + events.append("health") + return {"status": "ok"} + + def openapi_payload(self): + events.append("openapi") + return {"info": {"title": "Unrelated Service"}} + + def validate_auth(self): + raise AssertionError("credentials must not be sent before identity is verified") + + monkeypatch.setattr(openviking_module, "_VikingClient", ForeignServiceClient) + + valid, message, role = openviking_module._validate_openviking_setup_values({ + "endpoint": "https://foreign.example", + "api_key": "secret-key", + }) + + assert valid is False + assert role is None + assert "0.2.6 or earlier" in message + assert "0.2.10 or newer" in message + assert events == ["health", "openapi"] + + +def test_verified_legacy_openviking_is_healthy_for_reachability_and_runtime(monkeypatch): + events = [] + + class LegacyOpenVikingClient: + def __init__(self, *args, **kwargs): + pass + + def health_payload(self): + events.append("health") + return {"status": "ok"} + + def openapi_payload(self): + events.append("openapi") + return {"info": {"title": "OpenViking API"}} + + monkeypatch.setattr(openviking_module, "_VikingClient", LegacyOpenVikingClient) + + reachable, message = openviking_module._validate_openviking_reachability( + "https://legacy.example" + ) + runtime_state, runtime_message = openviking_module._classify_runtime_openviking_health( + LegacyOpenVikingClient(), + "https://legacy.example", + ) + + assert (reachable, message) == (True, "") + assert (runtime_state, runtime_message) == ("healthy", "") + assert events == ["health", "openapi", "health", "openapi"] + + def test_validate_openviking_reachability_uses_health_only(monkeypatch): events = [] @@ -992,4 +1379,233 @@ def test_prefetch_sends_contract_safe_memory_context_payload(monkeypatch): assert "mode" not in payload assert "target_uri" not in payload +def test_in_place_compression_rearms_commit_guard(): + """Post-compression turns must still be committable (#74695). + ``compress_context()`` commits before rewriting the transcript, which + latches the per-sid guard. In-place mode (the default) keeps the SAME sid, + so the latch then rejected every later commit for a still-live session — + the next compression, /new, normal session end and startup recovery all + silently did nothing, and post-compression turns were never extracted. + """ + provider = _make_provider_with_session("sid-123", turn_count=4) + provider._ensure_client = lambda: True + + # Compression commits the live session, latching the guard. + provider._mark_session_committed("sid-123") + assert provider._session_needs_commit("sid-123", 4) is False + + # In-place compression: same id in, no rotation. + provider.on_session_switch("sid-123", reason="compression") + + # The session is still live, so new turns must be committable again. + assert provider._has_committed_session("sid-123") is False + assert provider._turn_count == 0 + assert provider._session_needs_commit("sid-123", 2) is True + + +def test_rotating_compression_keeps_old_session_latched(): + """Rotation mode must keep the guard, which dedupes the old id's finalize. + + With ``compression.in_place: false`` a fresh child id is minted. The old id + stays committed so its ``_finalize_session_async`` does not double-commit + what compression already committed — the behavior the guard exists for. + """ + provider = _make_provider_with_session("old-sid", turn_count=4) + provider._ensure_client = lambda: True + provider._finalize_session_async = MagicMock() + + provider._mark_session_committed("old-sid") + provider.on_session_switch("new-sid", reason="compression") + + assert provider._has_committed_session("old-sid") is True + assert provider._session_needs_commit("old-sid", 4) is False + + +def test_undo_rewind_does_not_rearm_commit_guard(): + """Only compression re-arms; a same-session /undo must not.""" + provider = _make_provider_with_session("sid-123", turn_count=4) + provider._ensure_client = lambda: True + + provider._mark_session_committed("sid-123") + provider.on_session_switch("sid-123", rewound=True) + + assert provider._has_committed_session("sid-123") is True + + +def test_in_place_compression_lifecycle_allows_a_later_commit(): + """End-to-end wiring, not a hand-set latch (#74695). + + Drives the real sequence a session goes through: commit at the compression + boundary, same-id ``on_session_switch``, a post-compression turn via + ``sync_turn``, then a later commit. Before the fix the second commit never + reached the server, so every turn after the first compression was lost. + """ + provider = _make_provider_with_session("sid-123", turn_count=3) + provider._ensure_client = lambda: True + provider._new_client = lambda: provider._client + + def _commit_calls(): + return [ + c for c in provider._client.post.call_args_list + if c.args and str(c.args[0]).endswith("/commit") + ] + + # 1. Compression commits the live session through the real path. + provider.on_session_end([{"role": "user", "content": "before"}]) + assert len(_commit_calls()) == 1 + assert provider._has_committed_session("sid-123") is True + + # 2. In-place compression: same id back in, no rotation. + provider.on_session_switch("sid-123", reason="compression") + + # No new turns means no duplicate extraction at an immediate boundary. + provider.on_session_end([]) + assert len(_commit_calls()) == 1 + + # 3. A genuinely new turn lands on the still-live session. + provider.sync_turn("after compression", "reply", session_id="sid-123") + assert provider._drain_writers("sid-123", timeout=5.0) + assert provider._turn_count > 0 + assert any( + call.args and str(call.args[0]).endswith("/messages/batch") + for call in provider._client.post.call_args_list + ) + + # 4. That turn must still be committable. + provider.on_session_end([{"role": "user", "content": "after"}]) + assert len(_commit_calls()) == 2, ( + "post-compression turns were never committed: " + f"{provider._client.post.call_args_list}" + ) + +def test_resolve_connection_settings_reads_config_yaml_non_secret_fields(monkeypatch): + """#68209: non-secret fields saved to config.yaml feed the resolution chain.""" + _clear_openviking_env(monkeypatch) + provider_config = { + "endpoint": "http://saved.test:1933", + "account": "cfg-account", + "user": "cfg-user", + "agent": "cfg-agent", + } + + settings = openviking_module._resolve_connection_settings(provider_config) + + assert settings["endpoint"] == "http://saved.test:1933" + assert settings["account"] == "cfg-account" + assert settings["user"] == "cfg-user" + assert settings["agent"] == "cfg-agent" + + +def test_env_overrides_config_yaml_non_secret_fields(monkeypatch): + """env still wins over config.yaml (env -> ovcli -> config.yaml -> default).""" + _clear_openviking_env(monkeypatch) + monkeypatch.setenv("OPENVIKING_ENDPOINT", "http://env.test") + monkeypatch.setenv("OPENVIKING_AGENT", "env-agent") + + settings = openviking_module._resolve_connection_settings( + {"endpoint": "http://saved.test", "agent": "cfg-agent"} + ) + + assert settings["endpoint"] == "http://env.test" + assert settings["agent"] == "env-agent" + + +def test_blocked_endpoint_does_not_fall_back_or_construct_client(monkeypatch, tmp_path): + _clear_openviking_env(monkeypatch) + monkeypatch.setenv( + "OPENVIKING_ENDPOINT", + "http://169.254.169.254/latest/meta-data/temporary-credential", + ) + monkeypatch.setattr( + openviking_module, + "_VikingClient", + MagicMock(side_effect=AssertionError("blocked endpoint must not construct a client")), + ) + warnings = [] + provider = OpenVikingMemoryProvider() + + provider.initialize( + "session-1", + hermes_home=str(tmp_path), + platform="cli", + warning_callback=warnings.append, + ) + + assert provider._client is None + assert provider._endpoint == "" + assert len(warnings) == 1 + assert "blocked metadata address" in warnings[0] + assert "temporary-credential" not in warnings[0] + assert openviking_module._DEFAULT_ENDPOINT not in warnings[0] + + +@pytest.mark.parametrize( + "health_payload", + [ + {"status": "ok", "healthy": True}, + ["not", "openviking"], + ], +) +def test_runtime_rejects_unrelated_json_health_response( + monkeypatch, tmp_path, health_payload +): + _clear_openviking_env(monkeypatch) + monkeypatch.setenv("OPENVIKING_ENDPOINT", "http://localhost:1934") + + class UnrelatedJsonService: + def __init__(self, *args, **kwargs): + pass + + def health_payload(self): + return health_payload + + monkeypatch.setattr(openviking_module, "_VikingClient", UnrelatedJsonService) + monkeypatch.setattr( + openviking_module, + "_local_openviking_port_is_open", + lambda host, port: True, + ) + monkeypatch.setattr( + openviking_module, + "_describe_local_port_listener", + lambda host, port: "python-http-server (PID 4242)", + ) + monkeypatch.setattr( + openviking_module, + "_start_local_openviking_server", + MagicMock(side_effect=AssertionError("responding non-OpenViking service must not auto-start")), + ) + warnings = [] + provider = OpenVikingMemoryProvider() + + provider.initialize( + "session-1", + hermes_home=str(tmp_path), + platform="cli", + warning_callback=warnings.append, + ) + + assert provider._client is None + assert len(warnings) == 1 + assert "/health response is not valid OpenViking" in warnings[0] + assert "python-http-server (PID 4242)" in warnings[0] + + +def test_is_available_true_for_config_yaml_endpoint(monkeypatch): + """#68209: a config.yaml endpoint (no env, no ovcli) counts as available.""" + _clear_openviking_env(monkeypatch) + monkeypatch.setattr( + openviking_module, + "_load_hermes_openviking_config", + lambda: {"endpoint": "http://saved.test:1933"}, + ) + assert OpenVikingMemoryProvider().is_available() is True + + +def test_is_available_false_without_any_endpoint(monkeypatch): + _clear_openviking_env(monkeypatch) + monkeypatch.setattr( + openviking_module, "_load_hermes_openviking_config", lambda: {} + ) + assert OpenVikingMemoryProvider().is_available() is False diff --git a/tests/plugins/memory/test_retaindb_provider.py b/tests/plugins/memory/test_retaindb_provider.py index bc46898044..0372edaefd 100644 --- a/tests/plugins/memory/test_retaindb_provider.py +++ b/tests/plugins/memory/test_retaindb_provider.py @@ -38,3 +38,151 @@ def test_upload_file_allows_regular_file(tmp_path): provider._client.upload_file.assert_called_once() assert provider._client.upload_file.call_args.args[0] == note.read_bytes() assert result["file"]["id"] == "file-1" + + +def _capture_initialized_client(monkeypatch, tmp_path): + """Patch _Client/_WriteQueue/get_hermes_home; return a dict capturing args.""" + import hermes_constants + + import plugins.memory.retaindb as retaindb_module + + captured: dict = {} + + class _FakeClient: + def __init__(self, api_key, base_url, project): + captured["api_key"] = api_key + captured["base_url"] = base_url + captured["project"] = project + self.project = project + + monkeypatch.setattr(retaindb_module, "_Client", _FakeClient) + monkeypatch.setattr(retaindb_module, "_WriteQueue", lambda *a, **k: MagicMock()) + monkeypatch.setattr(hermes_constants, "get_hermes_home", lambda: tmp_path) + return retaindb_module, captured + + +def test_retaindb_config_loader_uses_readonly_config(monkeypatch): + import hermes_cli.config as config_mod + import plugins.memory.retaindb as retaindb_module + + backing_config = { + "memory": { + "retaindb": { + "base_url": "https://saved.example", + "project": "saved-project", + } + } + } + monkeypatch.setattr(config_mod, "load_config_readonly", lambda: backing_config) + monkeypatch.setattr( + config_mod, + "load_config", + MagicMock(side_effect=AssertionError("read-only provider path must not load a mutable copy")), + ) + + config = retaindb_module._load_retaindb_config() + + assert config == backing_config["memory"]["retaindb"] + assert config is not backing_config["memory"]["retaindb"] + + +def test_initialize_reads_real_dashboard_config_file(tmp_path, monkeypatch): + for var in ("RETAINDB_API_KEY", "RETAINDB_BASE_URL", "RETAINDB_PROJECT"): + monkeypatch.delenv(var, raising=False) + (tmp_path / "config.yaml").write_text( + """\ +memory: + provider: retaindb + retaindb: + base_url: https://retaindb.saved.example/ + project: dashboard-project +""", + encoding="utf-8", + ) + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + _retaindb_module, captured = _capture_initialized_client(monkeypatch, tmp_path) + + RetainDBMemoryProvider().initialize("sess-1") + + assert captured["base_url"] == "https://retaindb.saved.example" + assert captured["project"] == "dashboard-project" + + +def test_initialize_reads_base_url_and_project_from_config_yaml(tmp_path, monkeypatch): + """#68209: non-secret base_url/project come from config.yaml when env is unset.""" + for var in ("RETAINDB_API_KEY", "RETAINDB_BASE_URL", "RETAINDB_PROJECT"): + monkeypatch.delenv(var, raising=False) + retaindb_module, captured = _capture_initialized_client(monkeypatch, tmp_path) + monkeypatch.setattr( + retaindb_module, + "_load_retaindb_config", + lambda: {"base_url": "https://retaindb.example.com/", "project": "cfg-project"}, + ) + + RetainDBMemoryProvider().initialize("sess-1") + + assert captured["base_url"] == "https://retaindb.example.com" # trailing slash stripped + assert captured["project"] == "cfg-project" + + +def test_initialize_env_overrides_config_yaml(tmp_path, monkeypatch): + for var in ("RETAINDB_API_KEY", "RETAINDB_PROJECT"): + monkeypatch.delenv(var, raising=False) + monkeypatch.setenv("RETAINDB_BASE_URL", "https://env.example.com") + retaindb_module, captured = _capture_initialized_client(monkeypatch, tmp_path) + monkeypatch.setattr( + retaindb_module, + "_load_retaindb_config", + lambda: {"base_url": "https://cfg.example.com", "project": "cfg-project"}, + ) + + RetainDBMemoryProvider().initialize("sess-1") + + assert captured["base_url"] == "https://env.example.com" + + +def test_initialize_combines_scoped_secret_with_dashboard_config(tmp_path, monkeypatch): + """Rebase regression: scoped secrets and non-secret config must coexist.""" + from agent.secret_scope import ( + is_multiplex_active, + reset_secret_scope, + set_multiplex_active, + set_secret_scope, + ) + + monkeypatch.setenv("RETAINDB_API_KEY", "env-other-profile") + monkeypatch.delenv("RETAINDB_BASE_URL", raising=False) + monkeypatch.delenv("RETAINDB_PROJECT", raising=False) + retaindb_module, captured = _capture_initialized_client(monkeypatch, tmp_path) + monkeypatch.setattr( + retaindb_module, + "_load_retaindb_config", + lambda: {"base_url": "https://dashboard.example.com/", "project": "dashboard-project"}, + ) + + previous_multiplex_state = is_multiplex_active() + set_multiplex_active(True) + token = set_secret_scope({"RETAINDB_API_KEY": "scoped-key"}) + try: + RetainDBMemoryProvider().initialize("sess-1") + finally: + reset_secret_scope(token) + set_multiplex_active(previous_multiplex_state) + + assert captured == { + "api_key": "scoped-key", + "base_url": "https://dashboard.example.com", + "project": "dashboard-project", + } + + +def test_initialize_falls_back_to_default_base_url(tmp_path, monkeypatch): + for var in ("RETAINDB_API_KEY", "RETAINDB_BASE_URL", "RETAINDB_PROJECT"): + monkeypatch.delenv(var, raising=False) + retaindb_module, captured = _capture_initialized_client(monkeypatch, tmp_path) + monkeypatch.setattr(retaindb_module, "_load_retaindb_config", lambda: {}) + + RetainDBMemoryProvider().initialize("sess-1") + + assert captured["base_url"] == retaindb_module._DEFAULT_BASE_URL + assert captured["project"] == "default" diff --git a/tests/providers/test_plugin_discovery.py b/tests/providers/test_plugin_discovery.py index 79169b1f72..2c828631bd 100644 --- a/tests/providers/test_plugin_discovery.py +++ b/tests/providers/test_plugin_discovery.py @@ -21,6 +21,7 @@ def _clear_provider_caches(): import providers as _pkg _pkg._REGISTRY.clear() _pkg._ALIASES.clear() + _pkg._PROVIDER_LIST_CACHE = None _pkg._discovered = False # Evict any cached plugin modules so the next import re-executes. for mod in list(sys.modules.keys()): diff --git a/tests/providers/test_provider_registry.py b/tests/providers/test_provider_registry.py new file mode 100644 index 0000000000..c31879dcba --- /dev/null +++ b/tests/providers/test_provider_registry.py @@ -0,0 +1,66 @@ +import pytest + +from providers import ProviderProfile +import providers + + +@pytest.fixture(autouse=True) +def isolate_provider_registry(): + registry = providers._REGISTRY.copy() + aliases = providers._ALIASES.copy() + provider_list_cache = ( + None + if providers._PROVIDER_LIST_CACHE is None + else list(providers._PROVIDER_LIST_CACHE) + ) + discovered = providers._discovered + + yield + + providers._REGISTRY.clear() + providers._REGISTRY.update(registry) + providers._ALIASES.clear() + providers._ALIASES.update(aliases) + providers._PROVIDER_LIST_CACHE = provider_list_cache + providers._discovered = discovered + + +def _profile(name: str, *aliases: str) -> ProviderProfile: + return ProviderProfile(name=name, aliases=aliases) + + +def _reset_registry() -> None: + providers._REGISTRY.clear() + providers._ALIASES.clear() + providers._PROVIDER_LIST_CACHE = None + providers._discovered = True + + +def test_list_providers_reuses_cached_snapshot_until_registration_changes(): + _reset_registry() + first = _profile("alpha") + providers.register_provider(first) + + listed = providers.list_providers() + listed.clear() + + assert providers.list_providers() == [first] + + # Hit-path copy guard: mutating a CACHED return must not corrupt the + # module-level snapshot for later callers (aliasing bug class). + providers.list_providers().clear() + assert providers.list_providers() == [first] + + second = _profile("beta") + providers.register_provider(second) + + assert providers.list_providers() == [first, second] + + +def test_list_providers_dedupes_aliases_in_cached_snapshot(): + _reset_registry() + profile = _profile("kimi", "moonshot", "kimi-k2") + providers.register_provider(profile) + + assert providers.get_provider_profile("moonshot") is profile + assert providers.list_providers() == [profile] diff --git a/tests/run_agent/test_24996_fallback_exhaustion_cooldown.py b/tests/run_agent/test_24996_fallback_exhaustion_cooldown.py index 83991c2471..2e86c795c5 100644 --- a/tests/run_agent/test_24996_fallback_exhaustion_cooldown.py +++ b/tests/run_agent/test_24996_fallback_exhaustion_cooldown.py @@ -130,3 +130,103 @@ class TestExhaustionArmsCooldown: assert agent._try_activate_fallback() is False cooldown = getattr(agent, "_rate_limited_until", 0) assert cooldown == far_future + + +class TestRateLimitBackoffEscalation: + """Exponential backoff for consecutive rate-limit failures (#29702). + + The first rate-limit keeps the historical 60s cooldown; each consecutive + rate-limit within one degradation window doubles it (60 → 120 → 240 → ...) + capped at 4h (14400s). A successful primary restore resets the counter. + """ + + @staticmethod + def _back_on_primary(agent, snapshot): + """Simulate the primary provider rate-limiting again on a later turn + (without a successful restore, which would reset the counter): put + the agent's identity back on the primary and reset the turn-scoped + fallback chain state.""" + agent.provider, agent.model, agent.base_url = snapshot + agent._fallback_activated = False + agent._fallback_index = 0 + + def test_backoff_doubles_per_consecutive_rate_limit(self): + """Each consecutive primary rate-limit doubles the cooldown: + 60s, then 120s, then 240s.""" + fbs = [{"provider": "openai", "model": "gpt-4o"}] + agent = _make_agent(fallback_model=fbs) + agent._rate_limited_until = 0 + snapshot = (agent.provider, agent.model, agent.base_url) + frozen = 1_000.0 + with ( + patch("agent.chat_completion_helpers.time.monotonic", return_value=frozen), + patch( + "agent.auxiliary_client.resolve_provider_client", + return_value=(_mock_client(), "resolved"), + ), + ): + expected = [60, 120, 240] + for n, want in enumerate(expected): + self._back_on_primary(agent, snapshot) + agent._try_activate_fallback(reason=FailoverReason.rate_limit) + assert agent._rate_limited_until == frozen + want, ( + f"backoff #{n + 1}: expected {want}s cooldown" + ) + assert agent._rate_limit_backoff_count == n + 1 + + def test_backoff_caps_at_four_hours(self): + """Escalation is capped at 14400s (4h) no matter how many + consecutive rate-limits occurred.""" + fbs = [{"provider": "openai", "model": "gpt-4o"}] + agent = _make_agent(fallback_model=fbs) + agent._rate_limited_until = 0 + # 60 * 2**10 = 61440s, far past the cap. + agent._rate_limit_backoff_count = 10 + frozen = 1_000.0 + with ( + patch("agent.chat_completion_helpers.time.monotonic", return_value=frozen), + patch( + "agent.auxiliary_client.resolve_provider_client", + return_value=(_mock_client(), "resolved"), + ), + ): + agent._try_activate_fallback(reason=FailoverReason.rate_limit) + assert agent._rate_limited_until == frozen + 14400 + + def test_backoff_counter_resets_on_successful_primary_restore(self): + """A successful restore_primary_runtime resets the backoff counter, + so the next rate-limit starts back at the 60s base.""" + fbs = [{"provider": "openai", "model": "gpt-4o"}] + agent = _make_agent(fallback_model=fbs) + snapshot = (agent.provider, agent.model, agent.base_url) + frozen = 1_000.0 + with ( + patch("agent.chat_completion_helpers.time.monotonic", return_value=frozen), + patch( + "agent.auxiliary_client.resolve_provider_client", + return_value=(_mock_client(), "resolved"), + ), + ): + # Two consecutive rate-limits escalate the counter to 2. + agent._rate_limited_until = 0 + agent._try_activate_fallback(reason=FailoverReason.rate_limit) + self._back_on_primary(agent, snapshot) + agent._try_activate_fallback(reason=FailoverReason.rate_limit) + assert agent._rate_limit_backoff_count == 2 + + # Cooldown expired; the primary restores successfully. + agent._fallback_activated = True + agent._rate_limited_until = 0 + assert agent._restore_primary_runtime() is True + assert agent._rate_limit_backoff_count == 0 + + # The next rate-limit is treated as a fresh first failure: 60s. + with ( + patch("agent.chat_completion_helpers.time.monotonic", return_value=frozen), + patch( + "agent.auxiliary_client.resolve_provider_client", + return_value=(_mock_client(), "resolved"), + ), + ): + agent._try_activate_fallback(reason=FailoverReason.rate_limit) + assert agent._rate_limited_until == frozen + 60 diff --git a/tests/run_agent/test_anthropic_prompt_cache_policy.py b/tests/run_agent/test_anthropic_prompt_cache_policy.py index 687a65f251..2c4b44e9c0 100644 --- a/tests/run_agent/test_anthropic_prompt_cache_policy.py +++ b/tests/run_agent/test_anthropic_prompt_cache_policy.py @@ -291,13 +291,19 @@ class TestQwenAlibabaFamily: class TestDeepSeekOpenCode: - """DeepSeek uses OpenCode's envelope-layout cache markers (#24617).""" + """DeepSeek on OpenCode does NOT use cache markers (#77217). + + OpenCode Zen's relay rejects the Anthropic-style content block format + that cache markers produce (content becomes a block array instead of a + plain string), causing HTTP 400. DeepSeek is intentionally excluded + from the caching path. + """ @pytest.mark.parametrize( "provider", ["opencode", "opencode-zen", "opencode-go"], ) - def test_deepseek_on_opencode_caches_with_envelope_layout(self, provider): + def test_deepseek_on_opencode_does_not_cache(self, provider): agent = _make_agent( provider=provider, base_url="https://opencode.ai/v1", @@ -305,7 +311,7 @@ class TestDeepSeekOpenCode: model="deepseek-v4-pro", ) - assert agent._anthropic_prompt_cache_policy() == (True, False) + assert agent._anthropic_prompt_cache_policy() == (False, False) def test_deepseek_on_direct_alibaba_does_not_cache(self): agent = _make_agent( diff --git a/tests/run_agent/test_codex_app_server_lifecycle.py b/tests/run_agent/test_codex_app_server_lifecycle.py new file mode 100644 index 0000000000..3e784ccaaf --- /dev/null +++ b/tests/run_agent/test_codex_app_server_lifecycle.py @@ -0,0 +1,82 @@ +"""Codex app-server session lifecycle on hard agent teardown (#65260). + +The Codex runtime drops ``agent._codex_session`` on turn crash and on +retirement (agent/codex_runtime.py), but ``AIAgent.close()`` — the hard +teardown for /new, /reset, and session expiry — had no owner for it, so +the app-server child process survived until interpreter exit. +""" + +import threading + +from run_agent import AIAgent + + +class _FakeCodexSession: + def __init__(self, raises: bool = False): + self.close_calls = 0 + self._raises = raises + + def close(self): + self.close_calls += 1 + if self._raises: + raise RuntimeError("app-server already dead") + + +def _bare_agent(session_id: str) -> AIAgent: + """Minimal agent shell exercising close() without a real build.""" + agent = AIAgent.__new__(AIAgent) + agent.session_id = session_id + agent.client = None + agent._active_children_lock = threading.Lock() + agent._active_children = set() + agent._end_session_on_close = False + agent._session_messages = ["retained"] + return agent + + +def test_agent_close_releases_codex_app_server_session(monkeypatch): + agent = _bare_agent("test-codex-lifecycle") + codex_session = _FakeCodexSession() + agent._codex_session = codex_session + + monkeypatch.setattr("run_agent.cleanup_vm", lambda _task_id: None) + monkeypatch.setattr("run_agent.cleanup_browser", lambda _task_id: None) + + agent.close() + agent.close() + + # Idempotent: the second close must not re-close a released session. + assert codex_session.close_calls == 1 + assert agent._codex_session is None + assert agent._session_messages == [] + + +def test_close_clears_reference_even_when_session_close_raises(monkeypatch): + """A wedged app-server must not strand a stale session reference. + + The attribute is cleared BEFORE close() precisely so a raising close + can't leave a dead session attached to the agent. + """ + agent = _bare_agent("test-codex-lifecycle-raises") + codex_session = _FakeCodexSession(raises=True) + agent._codex_session = codex_session + + monkeypatch.setattr("run_agent.cleanup_vm", lambda _task_id: None) + monkeypatch.setattr("run_agent.cleanup_browser", lambda _task_id: None) + + agent.close() + + assert codex_session.close_calls == 1 + assert agent._codex_session is None + + +def test_close_without_codex_session_is_a_noop(monkeypatch): + """Non-Codex sessions (the common case) must be unaffected.""" + agent = _bare_agent("test-no-codex") + + monkeypatch.setattr("run_agent.cleanup_vm", lambda _task_id: None) + monkeypatch.setattr("run_agent.cleanup_browser", lambda _task_id: None) + + agent.close() + + assert getattr(agent, "_codex_session", None) is None diff --git a/tests/run_agent/test_conversation_fallback_state.py b/tests/run_agent/test_conversation_fallback_state.py index be46edcb93..a0d07c18ab 100644 --- a/tests/run_agent/test_conversation_fallback_state.py +++ b/tests/run_agent/test_conversation_fallback_state.py @@ -125,3 +125,83 @@ def test_substantive_tool_only_turn_invalidates_older_housekeeping_fallback(): ) +def test_bare_tool_marker_is_not_reused_as_final_response(): + """ + Regression test for #78148. + + A provider/local template can emit a bare bracketed token (e.g. "[memory]") + as assistant content alongside a tool call. That token is protocol + scaffolding, not an answer. If it gets cached as `_last_content_with_tools` + and the following turn is empty, the post-tool fallback replays it as the + final response — and because it then enters the persisted transcript, + later context compaction preserves it, letting the model repeat the + marker in subsequent turns. + + Test sequence: + 1. Content "[memory]" + skill_manage (housekeeping) tool call → the bare + marker must be discarded, not cached as a fallback. + 2. Empty content, no tool calls → enters the post-tool nudge path since + no fallback is available. + 3. Content "Recovered after nudge." → returned as the final response. + + Before the fix: + - Step 1 cached "[memory]" as `_last_content_with_tools`. + - Step 2 reused it via the empty-response fallback, so the conversation + never reached step 3 and "[memory]" leaked into the persisted history. + + After the fix: + - Step 1 strips the bare marker before it is cached or persisted. + - Step 2 has no fallback available and enters the nudge path instead. + - Step 3 returns the nudge response as the final answer. + """ + with ( + patch("run_agent.get_tool_definitions", return_value=_tool_defs("skill_manage")), + patch("run_agent.check_toolset_requirements", return_value={}), + patch("run_agent.OpenAI"), + ): + agent = AIAgent( + api_key="test-key", + base_url="https://openrouter.ai/api/v1/", + quiet_mode=True, + skip_context_files=True, + skip_memory=True, + ) + + agent._cached_system_prompt = "You are helpful." + agent._use_prompt_caching = False + agent.compression_enabled = False + agent.save_trajectories = False + agent.valid_tool_names = {"skill_manage"} + agent.client = MagicMock() + agent.client.chat.completions.create.side_effect = [ + # Turn 1: Bare "[memory]" marker + housekeeping tool call. + _response( + content="[memory]", + finish_reason="tool_calls", + tool_calls=[_tool_call("skill_manage", "skill1")], + ), + # Turn 2: Empty response (should enter nudge path, not reuse "[memory]"). + _response(content="", finish_reason="stop"), + # Turn 3: Nudge response + _response(content="Recovered after nudge.", finish_reason="stop"), + ] + + with ( + patch("run_agent.handle_function_call", return_value="ok"), + patch.object(agent, "_persist_session"), + patch.object(agent, "_save_trajectory"), + patch.object(agent, "_cleanup_task_resources"), + ): + result = agent.run_conversation("do the full task") + + assert result["final_response"] != "[memory]", ( + "The bare tool-call marker leaked through as the final response — " + "it should have been discarded before caching/persistence." + ) + assert result["final_response"] == "Recovered after nudge.", ( + f"Expected nudge recovery response, got: {result['final_response']}." + ) + assert result["api_calls"] == 3, ( + f"Expected 3 API calls (including nudge), got: {result['api_calls']}." + ) + diff --git a/tests/run_agent/test_empty_response_recovery_persistence.py b/tests/run_agent/test_empty_response_recovery_persistence.py index ff8b9cf312..70e262fbf2 100644 --- a/tests/run_agent/test_empty_response_recovery_persistence.py +++ b/tests/run_agent/test_empty_response_recovery_persistence.py @@ -13,6 +13,12 @@ class _CapturingSessionDB: self.rows.append({"role": role, "content": content}) return len(self.rows) + def append_messages_batch(self, session_id, messages, **kwargs): + # Mirror the real batch writer: same rows, one call. + for m in messages: + self.rows.append({"role": m.get("role"), "content": m.get("content")}) + return list(range(len(self.rows) - len(messages) + 1, len(self.rows) + 1)) + def _agent_with_capturing_db(): agent = AIAgent.__new__(AIAgent) diff --git a/tests/run_agent/test_image_generate_parallel.py b/tests/run_agent/test_image_generate_parallel.py new file mode 100644 index 0000000000..3223be9d9d --- /dev/null +++ b/tests/run_agent/test_image_generate_parallel.py @@ -0,0 +1,97 @@ +"""Regression tests for parallel image-generation tool batches.""" + +import json +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import run_agent +from agent import tool_executor + + +def _tool_call(name: str, args: dict, call_id: str) -> SimpleNamespace: + return SimpleNamespace( + id=call_id, + function=SimpleNamespace( + name=name, + arguments=json.dumps(args), + ), + ) + + +def test_image_generate_batch_routes_to_concurrent_executor(): + agent = SimpleNamespace() + agent._execute_tool_calls = run_agent.AIAgent._execute_tool_calls.__get__(agent) + agent._execute_tool_calls_concurrent = MagicMock() + agent._execute_tool_calls_sequential = MagicMock() + assistant_message = SimpleNamespace( + tool_calls=[ + _tool_call("image_generate", {"prompt": "variation one"}, "img_1"), + _tool_call("image_generate", {"prompt": "variation two"}, "img_2"), + ], + ) + + agent._execute_tool_calls(assistant_message, [], "task-image-batch") + + agent._execute_tool_calls_concurrent.assert_called_once() + agent._execute_tool_calls_sequential.assert_not_called() + + +def test_image_generate_parallel_worker_cap_defaults_to_four(): + runnable_calls = [ + ( + 0, + _tool_call("image_generate", {"prompt": "one"}, "img_1"), + "image_generate", + {}, + ), + ( + 1, + _tool_call("image_generate", {"prompt": "two"}, "img_2"), + "image_generate", + {}, + ), + ( + 2, + _tool_call("image_generate", {"prompt": "three"}, "img_3"), + "image_generate", + {}, + ), + ( + 3, + _tool_call("image_generate", {"prompt": "four"}, "img_4"), + "image_generate", + {}, + ), + ( + 4, + _tool_call("image_generate", {"prompt": "five"}, "img_5"), + "image_generate", + {}, + ), + ] + + with patch("hermes_cli.config.load_config", return_value={}): + assert tool_executor._max_workers_for_tool_batch(runnable_calls) == 4 + + +def test_image_generate_parallel_worker_cap_can_be_configured_lower(): + runnable_calls = [ + ( + 0, + _tool_call("image_generate", {"prompt": "one"}, "img_1"), + "image_generate", + {}, + ), + ( + 1, + _tool_call("image_generate", {"prompt": "two"}, "img_2"), + "image_generate", + {}, + ), + ] + + with patch( + "hermes_cli.config.load_config", + return_value={"image_gen": {"max_parallel_requests": 1}}, + ): + assert tool_executor._max_workers_for_tool_batch(runnable_calls) == 1 diff --git a/tests/run_agent/test_message_sequence_repair.py b/tests/run_agent/test_message_sequence_repair.py index b1a40e78dc..e169f1e068 100644 --- a/tests/run_agent/test_message_sequence_repair.py +++ b/tests/run_agent/test_message_sequence_repair.py @@ -227,6 +227,11 @@ def test_flush_guard_clamps_overshooting_cursor(): def append_message(self, **kw): self.rows.append(kw) + def append_messages_batch(self, session_id, messages, **kw): + for m in messages: + self.rows.append(dict(m, session_id=session_id)) + return list(range(1, len(messages) + 1)) + agent = _bare_agent() agent._session_db = _DB() agent._session_db_created = True diff --git a/tests/run_agent/test_overflow_overhead_aware_tokens.py b/tests/run_agent/test_overflow_overhead_aware_tokens.py new file mode 100644 index 0000000000..bcff0f889f --- /dev/null +++ b/tests/run_agent/test_overflow_overhead_aware_tokens.py @@ -0,0 +1,404 @@ +"""Regression tests: overflow recovery handlers must pass overhead-aware token estimates. + +PR fix (LCM issue 441): 413, context-overflow, and long-context-tier recovery handlers +were passing a messages-only token estimate to _compress_context instead of the +overhead-aware estimate_request_tokens_rough(api_messages, tools=agent.tools or None), +which includes tool schemas and system prompt overhead. + +These tests assert that each recovery handler: +1. Calls estimate_request_tokens_rough with a non-None `tools` argument. +2. Passes the resulting value as approx_tokens to _compress_context. + +The sentinel pattern (return_value=987654) makes the assertion unambiguous: if +approx_tokens==987654 the overhead-aware path was taken; any other value means the +handler used a different (likely messages-only) estimate. +""" + +import pytest + +from types import SimpleNamespace +from unittest.mock import MagicMock, patch, call + +from run_agent import AIAgent +import run_agent + + +# --------------------------------------------------------------------------- +# Shared fixtures / helpers (mirrored from test_413_compression.py) +# --------------------------------------------------------------------------- + + +@pytest.fixture(autouse=True) +def _no_sleep(monkeypatch): + """Short-circuit all time.sleep and jittered_backoff calls.""" + import time as _time + monkeypatch.setattr(_time, "sleep", lambda *_a, **_k: None) + monkeypatch.setattr(run_agent, "jittered_backoff", lambda *a, **k: 0.0) + + +def _make_tool_defs(*names: str) -> list: + return [ + { + "type": "function", + "function": { + "name": n, + "description": f"{n} tool", + "parameters": {"type": "object", "properties": {}}, + }, + } + for n in names + ] + + +def _mock_response(content="Hello", finish_reason="stop", tool_calls=None, usage=None): + msg = SimpleNamespace( + content=content, + tool_calls=tool_calls, + reasoning_content=None, + reasoning=None, + ) + choice = SimpleNamespace(message=msg, finish_reason=finish_reason) + resp = SimpleNamespace(choices=[choice], model="test/model") + resp.usage = SimpleNamespace(**usage) if usage else None + return resp + + +@pytest.fixture() +def agent(): + with ( + 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"), + ): + a = AIAgent( + api_key="test-key-1234567890", + base_url="https://openrouter.ai/api/v1", + quiet_mode=True, + skip_context_files=True, + skip_memory=True, + ) + a.client = MagicMock() + a._cached_system_prompt = "You are helpful." + a._use_prompt_caching = False + a.compression_enabled = True + a.save_trajectories = False + return a + + +def _prefill(): + return [ + {"role": "user", "content": "previous question"}, + {"role": "assistant", "content": "previous answer"}, + ] + + +# --------------------------------------------------------------------------- +# Sentinel: any value that could not coincidentally appear from a messages-only +# estimate during these tests. +# --------------------------------------------------------------------------- +_SENTINEL_TOKENS = 987_654 + + +# --------------------------------------------------------------------------- +# 1. 413 / payload-too-large handler +# --------------------------------------------------------------------------- + + +class TestHTTP413OverheadAwareTokens: + """The 413 recovery handler must call estimate_request_tokens_rough with + tools=agent.tools (non-None) and pass the result as approx_tokens.""" + + def test_413_passes_overhead_aware_tokens_to_compress(self, agent): + """approx_tokens passed to _compress_context equals the overhead-aware estimate.""" + err = Exception("Request entity too large") + err.status_code = 413 + ok_resp = _mock_response(content="Success", finish_reason="stop") + agent.client.chat.completions.create.side_effect = [err, ok_resp] + + with ( + patch( + "agent.conversation_loop.estimate_request_tokens_rough", + return_value=_SENTINEL_TOKENS, + ) as mock_estimate, + patch.object(agent, "_compress_context") as mock_compress, + patch.object(agent, "_persist_session"), + patch.object(agent, "_save_trajectory"), + patch.object(agent, "_cleanup_task_resources"), + ): + mock_compress.return_value = ( + [{"role": "user", "content": "compressed"}], + "compressed prompt", + ) + result = agent.run_conversation("hello", conversation_history=_prefill()) + + # _compress_context must have been called at least once for compression + mock_compress.assert_called() + + # Find the call that came from the 413 handler (approx_tokens=sentinel) + compress_kwargs_list = [c.kwargs for c in mock_compress.call_args_list] + sentinel_call = next( + (kw for kw in compress_kwargs_list if kw.get("approx_tokens") == _SENTINEL_TOKENS), + None, + ) + assert sentinel_call is not None, ( + f"No _compress_context call received approx_tokens={_SENTINEL_TOKENS}. " + f"Calls received approx_tokens values: " + f"{[kw.get('approx_tokens') for kw in compress_kwargs_list]}" + ) + + def test_413_estimate_called_with_non_none_tools(self, agent): + """estimate_request_tokens_rough must receive tools= in the 413 handler.""" + err = Exception("Request entity too large") + err.status_code = 413 + ok_resp = _mock_response(content="Success", finish_reason="stop") + agent.client.chat.completions.create.side_effect = [err, ok_resp] + + estimate_calls = [] + + def _capture_estimate(messages, tools=None): + estimate_calls.append({"messages": messages, "tools": tools}) + return _SENTINEL_TOKENS + + with ( + patch( + "agent.conversation_loop.estimate_request_tokens_rough", + side_effect=_capture_estimate, + ), + patch.object(agent, "_compress_context") as mock_compress, + patch.object(agent, "_persist_session"), + patch.object(agent, "_save_trajectory"), + patch.object(agent, "_cleanup_task_resources"), + ): + mock_compress.return_value = ( + [{"role": "user", "content": "compressed"}], + "compressed prompt", + ) + agent.run_conversation("hello", conversation_history=_prefill()) + + # At least one estimate call from the 413 handler must have non-None tools + handler_calls_with_tools = [c for c in estimate_calls if c["tools"] is not None] + assert handler_calls_with_tools, ( + "estimate_request_tokens_rough was never called with non-None tools " + "during 413 recovery. All calls: " + + str([c["tools"] for c in estimate_calls]) + ) + + +# --------------------------------------------------------------------------- +# 2. Context-overflow / input-too-large handler +# --------------------------------------------------------------------------- + + +class TestContextOverflowOverheadAwareTokens: + """The context-overflow (input overflow) recovery handler must call + estimate_request_tokens_rough with tools=agent.tools and pass the result + as approx_tokens to _compress_context.""" + + @staticmethod + def _make_context_overflow_error(): + """Build a 400 error that the classifier routes to context_overflow.""" + err = Exception( + "Error code: 400 - {'error': {'message': " + "\"This endpoint's maximum context length is 128000 tokens. " + "However, you requested about 200000 tokens. " + "Please reduce the length of the messages.\"}}" + ) + err.status_code = 400 + return err + + def test_context_overflow_passes_overhead_aware_tokens_to_compress(self, agent): + """approx_tokens passed to _compress_context equals the overhead-aware estimate.""" + err = self._make_context_overflow_error() + ok_resp = _mock_response(content="Recovered", finish_reason="stop") + agent.client.chat.completions.create.side_effect = [err, ok_resp] + + with ( + patch( + "agent.conversation_loop.estimate_request_tokens_rough", + return_value=_SENTINEL_TOKENS, + ) as mock_estimate, + patch.object(agent, "_compress_context") as mock_compress, + patch.object(agent, "_persist_session"), + patch.object(agent, "_save_trajectory"), + patch.object(agent, "_cleanup_task_resources"), + ): + mock_compress.return_value = ( + [{"role": "user", "content": "compressed"}], + "compressed prompt", + ) + result = agent.run_conversation("hello", conversation_history=_prefill()) + + mock_compress.assert_called() + + compress_kwargs_list = [c.kwargs for c in mock_compress.call_args_list] + sentinel_call = next( + (kw for kw in compress_kwargs_list if kw.get("approx_tokens") == _SENTINEL_TOKENS), + None, + ) + assert sentinel_call is not None, ( + f"No _compress_context call received approx_tokens={_SENTINEL_TOKENS}. " + f"Calls received approx_tokens values: " + f"{[kw.get('approx_tokens') for kw in compress_kwargs_list]}" + ) + + def test_context_overflow_estimate_called_with_non_none_tools(self, agent): + """estimate_request_tokens_rough must receive tools= in the context-overflow handler.""" + err = self._make_context_overflow_error() + ok_resp = _mock_response(content="Recovered", finish_reason="stop") + agent.client.chat.completions.create.side_effect = [err, ok_resp] + + estimate_calls = [] + + def _capture_estimate(messages, tools=None): + estimate_calls.append({"messages": messages, "tools": tools}) + return _SENTINEL_TOKENS + + with ( + patch( + "agent.conversation_loop.estimate_request_tokens_rough", + side_effect=_capture_estimate, + ), + patch.object(agent, "_compress_context") as mock_compress, + patch.object(agent, "_persist_session"), + patch.object(agent, "_save_trajectory"), + patch.object(agent, "_cleanup_task_resources"), + ): + mock_compress.return_value = ( + [{"role": "user", "content": "compressed"}], + "compressed prompt", + ) + agent.run_conversation("hello", conversation_history=_prefill()) + + handler_calls_with_tools = [c for c in estimate_calls if c["tools"] is not None] + assert handler_calls_with_tools, ( + "estimate_request_tokens_rough was never called with non-None tools " + "during context-overflow recovery. All calls: " + + str([c["tools"] for c in estimate_calls]) + ) + + def test_prompt_too_long_variant_passes_overhead_aware_tokens(self, agent): + """Anthropic 'prompt is too long' error also routes to context_overflow handler.""" + err = Exception( + "Error code: 400 - {'type': 'error', 'error': {'type': 'invalid_request_error', " + "'message': 'prompt is too long: 233153 tokens > 200000 maximum'}}" + ) + err.status_code = 400 + ok_resp = _mock_response(content="Recovered", finish_reason="stop") + agent.client.chat.completions.create.side_effect = [err, ok_resp] + + with ( + patch( + "agent.conversation_loop.estimate_request_tokens_rough", + return_value=_SENTINEL_TOKENS, + ), + patch.object(agent, "_compress_context") as mock_compress, + patch.object(agent, "_persist_session"), + patch.object(agent, "_save_trajectory"), + patch.object(agent, "_cleanup_task_resources"), + ): + mock_compress.return_value = ( + [{"role": "user", "content": "compressed"}], + "compressed prompt", + ) + result = agent.run_conversation("hello", conversation_history=_prefill()) + + mock_compress.assert_called() + compress_kwargs_list = [c.kwargs for c in mock_compress.call_args_list] + sentinel_call = next( + (kw for kw in compress_kwargs_list if kw.get("approx_tokens") == _SENTINEL_TOKENS), + None, + ) + assert sentinel_call is not None, ( + f"'prompt is too long' path did not pass overhead-aware approx_tokens. " + f"Got: {[kw.get('approx_tokens') for kw in compress_kwargs_list]}" + ) + + +# --------------------------------------------------------------------------- +# 3. Anthropic long-context tier (429) handler +# --------------------------------------------------------------------------- + + +class TestLongContextTierOverheadAwareTokens: + """The Anthropic long-context-tier 429 handler must call + estimate_request_tokens_rough with tools=agent.tools and pass the result + as approx_tokens to _compress_context.""" + + @staticmethod + def _make_long_context_tier_error(): + """Build a 429 'extra usage required for long context requests' error.""" + err = Exception( + "Error code: 429 - {'error': {'type': 'rate_limit_error', " + "'message': 'Extra usage is required for long context requests. " + "Please enable extra usage in your account settings.'}}" + ) + err.status_code = 429 + return err + + def test_long_context_tier_passes_overhead_aware_tokens_to_compress(self, agent): + """approx_tokens passed to _compress_context equals the overhead-aware estimate.""" + err = self._make_long_context_tier_error() + ok_resp = _mock_response(content="Recovered after context tier", finish_reason="stop") + agent.client.chat.completions.create.side_effect = [err, ok_resp] + + with ( + patch( + "agent.conversation_loop.estimate_request_tokens_rough", + return_value=_SENTINEL_TOKENS, + ), + patch.object(agent, "_compress_context") as mock_compress, + patch.object(agent, "_persist_session"), + patch.object(agent, "_save_trajectory"), + patch.object(agent, "_cleanup_task_resources"), + ): + mock_compress.return_value = ( + [{"role": "user", "content": "compressed"}], + "compressed prompt", + ) + result = agent.run_conversation("hello", conversation_history=_prefill()) + + mock_compress.assert_called() + compress_kwargs_list = [c.kwargs for c in mock_compress.call_args_list] + sentinel_call = next( + (kw for kw in compress_kwargs_list if kw.get("approx_tokens") == _SENTINEL_TOKENS), + None, + ) + assert sentinel_call is not None, ( + f"Long-context-tier handler did not pass overhead-aware approx_tokens. " + f"Got: {[kw.get('approx_tokens') for kw in compress_kwargs_list]}" + ) + + def test_long_context_tier_estimate_called_with_non_none_tools(self, agent): + """estimate_request_tokens_rough must receive tools= in the long-context handler.""" + err = self._make_long_context_tier_error() + ok_resp = _mock_response(content="Recovered", finish_reason="stop") + agent.client.chat.completions.create.side_effect = [err, ok_resp] + + estimate_calls = [] + + def _capture_estimate(messages, tools=None): + estimate_calls.append({"messages": messages, "tools": tools}) + return _SENTINEL_TOKENS + + with ( + patch( + "agent.conversation_loop.estimate_request_tokens_rough", + side_effect=_capture_estimate, + ), + patch.object(agent, "_compress_context") as mock_compress, + patch.object(agent, "_persist_session"), + patch.object(agent, "_save_trajectory"), + patch.object(agent, "_cleanup_task_resources"), + ): + mock_compress.return_value = ( + [{"role": "user", "content": "compressed"}], + "compressed prompt", + ) + agent.run_conversation("hello", conversation_history=_prefill()) + + handler_calls_with_tools = [c for c in estimate_calls if c["tools"] is not None] + assert handler_calls_with_tools, ( + "estimate_request_tokens_rough was never called with non-None tools " + "during long-context-tier recovery. All calls: " + + str([c["tools"] for c in estimate_calls]) + ) diff --git a/tests/run_agent/test_reset_aware_primary_restore.py b/tests/run_agent/test_reset_aware_primary_restore.py new file mode 100644 index 0000000000..efc071e996 --- /dev/null +++ b/tests/run_agent/test_reset_aware_primary_restore.py @@ -0,0 +1,340 @@ +"""Reset-aware primary restore — stay on fallback until the primary's +rate-limit window actually resets. + +``restore_primary_runtime`` retries the primary at the top of every turn +once the 60s ``_rate_limited_until`` cooldown clears. For transient 429s +that is correct, but subscription-window limits (Claude Pro/Max 5-hour +windows, ChatGPT weekly caps) report reset times hours or days away. The +credential pool already knows that timestamp (``last_error_reset_at``), +and until it elapses every restore attempt is a guaranteed failure that +invalidates the prompt cache twice per turn (primary → fail → fallback). + +The gate must FAIL OPEN: no pool, no reset info, or any error in the gate +falls through to the existing per-turn retry, so recovery can never happen +later than it does today. +""" + +import time +from unittest.mock import MagicMock, patch + +from run_agent import AIAgent +from agent.credential_pool import ( + STATUS_DEAD, + STATUS_EXHAUSTED, + STATUS_OK, + CredentialPool, + PooledCredential, +) + + +# ============================================================================= +# Helpers +# ============================================================================= + +def _entry( + provider="openrouter", + id="cred-1", + status=STATUS_OK, + reset_at=None, + status_at=None, + error_code=None, +): + return PooledCredential( + provider=provider, + id=id, + label=f"label-{id}", + auth_type="api_key", + priority=1, + source="manual", + access_token="sk-test-1234567890", + last_status=status, + last_status_at=status_at, + last_error_code=error_code, + last_error_reset_at=reset_at, + ) + + +class _FakePool: + """Minimal stand-in for CredentialPool in agent-level tests.""" + + def __init__(self, provider, next_at=None, available=False, raise_on_next=False): + self.provider = provider + self._next_at = next_at + self._available = available + self._raise = raise_on_next + self.next_available_calls = 0 + + def next_available_at(self): + self.next_available_calls += 1 + if self._raise: + raise RuntimeError("boom") + return self._next_at + + def has_credentials(self): + return True + + def has_available(self): + return self._available + + def select(self): + return None + + +def _make_tool_defs(*names): + return [ + { + "type": "function", + "function": { + "name": n, + "description": f"{n} tool", + "parameters": {"type": "object", "properties": {}}, + }, + } + for n in names + ] + + +def _make_agent(fallback_model=None): + with ( + 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"), + ): + agent = AIAgent( + api_key="test-key-12345678", + base_url="https://my-llm.example.com/v1", + provider="custom", + quiet_mode=True, + skip_context_files=True, + skip_memory=True, + fallback_model=fallback_model, + ) + agent.client = MagicMock() + return agent + + +def _activate_fallback(agent): + mock_client = MagicMock() + mock_client.api_key = "fallback-key-1234" + mock_client.base_url = "https://openrouter.ai/api/v1" + with patch( + "agent.auxiliary_client.resolve_provider_client", + return_value=(mock_client, None), + ): + assert agent._try_activate_fallback() is True + assert agent._fallback_activated is True + + +# ============================================================================= +# CredentialPool.next_available_at() +# ============================================================================= + +class TestNextAvailableAt: + def test_all_exhausted_returns_earliest_reset(self): + now = time.time() + pool = CredentialPool( + "openrouter", + [ + _entry(id="a", status=STATUS_EXHAUSTED, reset_at=now + 7200, error_code=429), + _entry(id="b", status=STATUS_EXHAUSTED, reset_at=now + 3600, error_code=429), + ], + ) + assert pool.next_available_at() == now + 3600 + + def test_available_entry_returns_none(self): + now = time.time() + pool = CredentialPool( + "openrouter", + [ + _entry(id="a", status=STATUS_OK), + _entry(id="b", status=STATUS_EXHAUSTED, reset_at=now + 3600, error_code=429), + ], + ) + assert pool.next_available_at() is None + + def test_elapsed_cooldown_counts_as_available(self): + """An exhausted entry whose reset time has passed re-enters rotation, + so the pool reports available (None) even without clear_expired.""" + now = time.time() + pool = CredentialPool( + "openrouter", + [_entry(id="a", status=STATUS_EXHAUSTED, reset_at=now - 10, error_code=429)], + ) + assert pool.next_available_at() is None + + def test_exhausted_without_timestamps_returns_none(self): + """No reset info at all -> None (fail open), not a guess.""" + pool = CredentialPool( + "openrouter", + [_entry(id="a", status=STATUS_EXHAUSTED, reset_at=None, status_at=None)], + ) + assert pool.next_available_at() is None + + def test_exhausted_with_status_at_uses_ttl(self): + """Without an explicit reset_at, last_status_at + TTL is the estimate.""" + now = time.time() + pool = CredentialPool( + "openrouter", + [_entry(id="a", status=STATUS_EXHAUSTED, status_at=now, error_code=429)], + ) + result = pool.next_available_at() + assert result is not None + assert result > now + + def test_dead_only_returns_none(self): + """DEAD entries never re-enter via TTL; report no wait info.""" + pool = CredentialPool( + "openrouter", + [_entry(id="a", status=STATUS_DEAD, status_at=time.time())], + ) + assert pool.next_available_at() is None + + def test_empty_pool_returns_none(self): + pool = CredentialPool("openrouter", []) + assert pool.next_available_at() is None + + def test_runs_under_the_pool_lock(self): + """next_available_at must hold self._lock like every other + _available_entries caller — a concurrent select()/rotation can + otherwise tear self._entries mid-iteration (see has_available).""" + pool = CredentialPool("openrouter", []) + held = {} + + original = pool._available_entries + + def _probe(**kwargs): + # self._lock is an RLock (deferred-refresh mutations self-lock), + # so a same-thread non-blocking acquire always succeeds; probe + # ownership from a helper thread instead. + import threading as _t + + blocked = _t.Event() + + def _try(): + if not pool._lock.acquire(blocking=False): + blocked.set() + else: + pool._lock.release() + + worker = _t.Thread(target=_try) + worker.start() + worker.join(timeout=5) + held["locked"] = blocked.is_set() + return original(**kwargs) + + pool._available_entries = _probe + pool.next_available_at() + assert held["locked"], "next_available_at called _available_entries without self._lock" + + +# ============================================================================= +# restore_primary_runtime() gate +# ============================================================================= + +class TestResetAwareRestoreGate: + FB = {"provider": "openrouter", "model": "anthropic/claude-sonnet-4"} + + def test_stays_on_fallback_until_reset(self): + agent = _make_agent(fallback_model=self.FB) + _activate_fallback(agent) + agent._rate_limited_until = 0 # 60s transient cooldown already cleared + + # Attached pool matches the primary provider and says: nobody can + # serve until an hour from now. + agent._credential_pool = _FakePool("custom", next_at=time.time() + 3600) + + assert agent._restore_primary_runtime() is False + assert agent._fallback_activated is True + assert agent.provider == "openrouter" + assert agent.model == "anthropic/claude-sonnet-4" + + def test_restores_once_reset_elapsed(self): + agent = _make_agent(fallback_model=self.FB) + original_model = agent.model + _activate_fallback(agent) + agent._rate_limited_until = 0 + + agent._credential_pool = _FakePool("custom", next_at=None) + + with patch("run_agent.OpenAI", return_value=MagicMock()): + assert agent._restore_primary_runtime() is True + assert agent._fallback_activated is False + assert agent.model == original_model + assert agent.provider == "custom" + + def test_past_reset_time_does_not_block(self): + agent = _make_agent(fallback_model=self.FB) + _activate_fallback(agent) + agent._rate_limited_until = 0 + + agent._credential_pool = _FakePool("custom", next_at=time.time() - 5) + + with patch("run_agent.OpenAI", return_value=MagicMock()): + assert agent._restore_primary_runtime() is True + assert agent._fallback_activated is False + + def test_fails_open_on_pool_error(self): + """Any exception inside the gate must not break restore.""" + agent = _make_agent(fallback_model=self.FB) + _activate_fallback(agent) + agent._rate_limited_until = 0 + + agent._credential_pool = _FakePool("custom", raise_on_next=True) + + with patch("run_agent.OpenAI", return_value=MagicMock()): + assert agent._restore_primary_runtime() is True + assert agent._fallback_activated is False + + def test_cross_provider_fallback_loads_primary_pool(self): + """After a cross-provider fallback the attached pool belongs to the + fallback provider; the gate must consult the PRIMARY's pool.""" + agent = _make_agent(fallback_model=self.FB) + _activate_fallback(agent) + agent._rate_limited_until = 0 + + # Attached pool is the fallback provider's (mismatch with "custom"). + agent._credential_pool = _FakePool("openrouter", next_at=None) + primary_pool = _FakePool("custom", next_at=time.time() + 3600) + + with patch("agent.credential_pool.load_pool", return_value=primary_pool) as lp: + assert agent._restore_primary_runtime() is False + assert any(c.args == ("custom",) for c in lp.call_args_list) + assert primary_pool.next_available_calls == 1 + assert agent._fallback_activated is True + + def test_no_pool_info_falls_through(self): + """Pool present but no reset info -> existing per-turn retry.""" + agent = _make_agent(fallback_model=self.FB) + _activate_fallback(agent) + agent._rate_limited_until = 0 + + agent._credential_pool = _FakePool("custom", next_at=None) + + with patch("run_agent.OpenAI", return_value=MagicMock()): + assert agent._restore_primary_runtime() is True + + def test_logs_wait_only_once(self, caplog): + import logging + + agent = _make_agent(fallback_model=self.FB) + _activate_fallback(agent) + agent._rate_limited_until = 0 + agent._credential_pool = _FakePool("custom", next_at=time.time() + 3600) + + with caplog.at_level(logging.INFO, logger="agent.agent_runtime_helpers"): + assert agent._restore_primary_runtime() is False + assert agent._restore_primary_runtime() is False + waits = [r for r in caplog.records if "staying on fallback" in r.getMessage()] + assert len(waits) == 1 + + def test_transient_cooldown_still_respected(self): + """The existing 60s monotonic gate fires before the reset-aware one.""" + agent = _make_agent(fallback_model=self.FB) + _activate_fallback(agent) + agent._rate_limited_until = time.monotonic() + 60 + pool = _FakePool("custom", next_at=None) + agent._credential_pool = pool + + assert agent._restore_primary_runtime() is False + # Reset-aware gate never consulted — short-circuited by the 60s gate. + assert pool.next_available_calls == 0 diff --git a/tests/run_agent/test_run_agent.py b/tests/run_agent/test_run_agent.py index df1430bbea..3ad0a2ed96 100644 --- a/tests/run_agent/test_run_agent.py +++ b/tests/run_agent/test_run_agent.py @@ -112,8 +112,8 @@ def test_flush_persist_override_replaces_api_local_multimodal_note(agent): agent._flush_messages_to_session_db([{"role": "user", "content": api_content}], []) - db_write = agent._session_db.append_message.call_args.kwargs - assert db_write["content"] == "Describe this screenshot\n[screenshot]" + batch = agent._session_db.append_messages_batch.call_args.kwargs["messages"] + assert batch[0]["content"] == "Describe this screenshot\n[screenshot]" assert api_content[0]["text"] == "[MODEL SWITCH NOTE]\n\nDescribe this screenshot" @@ -136,6 +136,17 @@ def test_direct_session_db_flushes_share_marker_claim(agent): assert self.release.wait(timeout=5) self.rows.append(kwargs["content"]) + def append_messages_batch(self, session_id, messages, **kwargs): + with self._lock: + self.calls += 1 + first = self.calls == 1 + if first: + self.entered.set() + assert self.release.wait(timeout=5) + for m in messages: + self.rows.append(m["content"]) + return list(range(1, len(messages) + 1)) + db = _BarrierDB() agent._session_db = db agent._session_db_created = True @@ -418,6 +429,25 @@ class TestStripThinkBlocks: + @pytest.mark.parametrize( + ("text", "expected"), + [ + ( + "before {x} after", + "before {x}after", + ), + ( + "before {x} after", + "before {x}after", + ), + ], + ) + def test_mismatched_generic_tool_tags_preserve_opener_and_payload( + self, agent, text, expected + ): + assert agent._strip_think_blocks(text) == expected + + class TestExtractReasoning: def test_reasoning_field(self, agent): msg = _mock_assistant_msg(reasoning="thinking hard") @@ -3279,6 +3309,93 @@ class TestRunConversation: assert "No reply:" in result["final_response"] + def test_empty_response_retry_backoff_interrupted(self, agent, monkeypatch): + """If an interrupt is requested during the empty response retry wait, we abort.""" + self._setup_agent(agent) + agent.base_url = "http://127.0.0.1:1234/v1" + empty_resp = _mock_response(content=None, finish_reason="stop") + agent.client.chat.completions.create.side_effect = [empty_resp, empty_resp] + + from agent import conversation_loop as _conv_loop + + # Make backoff return 10.0 seconds + monkeypatch.setattr(_conv_loop, "jittered_backoff", lambda *a, **k: 10.0) + + # Trigger the interrupt on the first sleep call inside the wait loop + original_sleep = time.sleep + sleep_called = [] + + def _mock_sleep(seconds): + sleep_called.append(seconds) + if seconds == 0.2: + agent._interrupt_requested = True + else: + original_sleep(seconds) + + monkeypatch.setattr(time, "sleep", _mock_sleep) + + with ( + patch.object(agent, "_persist_session") as mock_persist, + patch.object(agent, "_save_trajectory"), + patch.object(agent, "_cleanup_task_resources"), + ): + result = agent.run_conversation("answer me") + + assert result["interrupted"] is True + assert "Operation interrupted: retrying empty response from model" in result["final_response"] + assert agent._empty_content_retries == 1 + assert 0.2 in sleep_called + assert mock_persist.call_count == 2 + + def test_empty_response_retry_backoff_status(self, agent, monkeypatch): + """Empty response retry wait updates the agent's status with wait time and sleeps.""" + self._setup_agent(agent) + agent.base_url = "http://127.0.0.1:1234/v1" + + # Two responses: first empty, second succeeds so it doesn't run forever + empty_resp = _mock_response(content=None, finish_reason="stop") + ok_resp = _mock_response(content="Final ok response.", finish_reason="stop") + agent.client.chat.completions.create.side_effect = [empty_resp, ok_resp] + + from agent import conversation_loop as _conv_loop + + monkeypatch.setattr(_conv_loop, "jittered_backoff", lambda *a, **k: 7.5) + + # Fake clock: the retry loop gates on real time.time() < sleep_end, so + # a no-op sleep alone busy-spins 7.5 wall-clock seconds. Advance a fake + # clock by each sleep amount instead (established pattern: + # test_session_activity_persist.py patches run_agent.time.time). + clock = {"t": time.time()} + monkeypatch.setattr(_conv_loop.time, "time", lambda: clock["t"]) + + sleep_calls = [] + + def _fake_sleep(secs): + sleep_calls.append(secs) + clock["t"] += secs + + monkeypatch.setattr(time, "sleep", _fake_sleep) + monkeypatch.setattr(_conv_loop.time, "sleep", _fake_sleep) + + status_messages = [] + monkeypatch.setattr(agent, "_buffer_status", lambda status: status_messages.append(status)) + + with ( + patch.object(agent, "_persist_session"), + patch.object(agent, "_save_trajectory"), + patch.object(agent, "_cleanup_task_resources"), + ): + result = agent.run_conversation("answer me") + + assert result["completed"] is True + assert result["final_response"] == "Final ok response." + + # 7.5s wait, slept in 0.2s increments -> 37.5 -> at least 37 calls + assert len([c for c in sleep_calls if c == 0.2]) >= 37 + + retry_status = [m for m in status_messages if "Empty response from model — retrying (1/3) in 8s" in m] + assert len(retry_status) == 1 + def test_partial_stream_recovery_uses_streamed_content(self, agent): """When streaming fails after partial delivery, recovered partial content becomes final response.""" self._setup_agent(agent) @@ -4064,7 +4181,10 @@ class TestRunConversation: ok_resp = _mock_response(content="done", finish_reason="stop") agent.client.chat.completions.create.side_effect = [exc, ok_resp] - mock_compress = MagicMock() + mock_compress = MagicMock(return_value=( + [{"role": "user", "content": "hello"}], + "You are helpful.", + )) with ( patch.object(agent, "_persist_session"), patch.object(agent, "_save_trajectory"), @@ -4078,7 +4198,7 @@ class TestRunConversation: assert result["completed"] is True assert second_call["max_tokens"] <= 936 assert agent.context_compressor.context_length == 200_000 - mock_compress.assert_not_called() + mock_compress.assert_called_once() def test_output_cap_retry_with_large_api_only_content(self, agent): """When a large system prompt makes api_messages huge while persisted @@ -4108,7 +4228,10 @@ class TestRunConversation: ok_resp = _mock_response(content="done", finish_reason="stop") agent.client.chat.completions.create.side_effect = [exc, ok_resp] - mock_compress = MagicMock() + mock_compress = MagicMock(return_value=( + [{"role": "user", "content": "hello"}], + "You are helpful.", + )) with ( patch.object(agent, "_persist_session"), patch.object(agent, "_save_trajectory"), @@ -4124,8 +4247,125 @@ class TestRunConversation: # near 199927 — this test fails on it. assert second_call["max_tokens"] <= 936 assert agent.context_compressor.context_length == 200_000 - mock_compress.assert_not_called() + mock_compress.assert_called_once() + def test_output_cap_retry_triggers_compression_and_recovers(self, agent): + """Regression for the output-cap death-loop (#55546 / #61761). + + When the provider reports an output-cap error on a near-full context + window, the retry must NOT just shrink max_tokens by a tiny amount and + spin forever. It must fire _compress_context() to actually free tokens + so the session recovers instead of exhausting compression_attempts. + + This locks in the fix: previously the output-cap path set + restart_with_compressed_messages without ever calling the compressor. + """ + self._setup_agent(agent) + agent.api_mode = "chat_completions" + agent.provider = "openrouter" + agent.model = "some/model" + agent.max_tokens = 65_536 + agent.compression_enabled = True + agent.context_compressor.context_length = 200_000 + # Context is essentially full -> compressor would want to run. + agent.context_compressor.should_compress = MagicMock(return_value=True) + + error_msg = ( + "max_tokens: 65536 > context_window: 200000 " + "- input_tokens: 199000 = available_tokens: 1000" + ) + exc = Exception(error_msg) + exc.status_code = 400 + exc.code = 400 + + ok_resp = _mock_response(content="done", finish_reason="stop") + agent.client.chat.completions.create.side_effect = [exc, ok_resp] + + # Compress drops the huge history (15 msgs -> 1), freeing tokens. + mock_compress = MagicMock(return_value=( + [{"role": "user", "content": "hello"}], + "You are helpful.", + )) + with ( + patch.object(agent, "_persist_session"), + patch.object(agent, "_save_trajectory"), + patch.object(agent, "_cleanup_task_resources"), + patch.object(agent.context_compressor, "update_model"), + patch.object(agent, "_compress_context", mock_compress), + ): + result = agent.run_conversation("hello") + + # Compression fired exactly once, on the output-cap retry. + mock_compress.assert_called_once() + # The compressed messages were re-sent and the call succeeded. + assert result["completed"] is True + assert result["final_response"] == "done" + # The retry honored the reduced max_tokens (available_out - 64). + second_call = agent.client.chat.completions.create.call_args_list[1].kwargs + assert second_call["max_tokens"] <= 936 + # LOCK IN THE FIX: the retry must actually SEND the compressed history + # (the 1-message payload from _compress_context + its new system + # prompt), not the original multi-message window. Without this, the + # output-cap retry would call the compressor but re-transmit the same + # oversized request forever. + second_messages = second_call.get("messages", []) + assert second_messages[-1].get("content") == "hello" + assert len(second_messages) == 2 + assert second_messages[0]["role"] == "system" + # context_length was NOT mutated by an output-cap error. + assert agent.context_compressor.context_length == 200_000 + + def test_output_cap_retry_compression_no_progress_terminates_bounded(self, agent): + """Regression: when the compressor cannot reduce the request (zero + progress AND no images to strip), the output-cap retry must terminate + via the max-attempts guard instead of spinning forever. + + The compressor is injected to return the input unchanged (same list + object, no lock-defer — just zero progress), and the provider keeps + rejecting, so the only correct outcome is a bounded + ``compression_exhausted`` failure, not an unbounded loop. + """ + self._setup_agent(agent) + agent.api_mode = "chat_completions" + agent.provider = "openrouter" + agent.model = "some/model" + agent.max_tokens = 65_536 + agent.compression_enabled = True + agent.context_compressor.context_length = 200_000 + agent.context_compressor.should_compress = MagicMock(return_value=True) + + error_msg = ( + "max_tokens: 65536 > context_window: 200000 " + "- input_tokens: 199000 = available_tokens: 1000" + ) + + def _rejecting(*args, **kwargs): + exc = Exception(error_msg) + exc.status_code = 400 + exc.code = 400 + raise exc + + # The provider never recovers (side effect raises on every call). + agent.client.chat.completions.create.side_effect = _rejecting + + def _no_progress(messages, system_message, **kwargs): + # Compressor runs but cannot shrink the request: no-op, same list. + return messages, system_message + + with ( + patch.object(agent, "_persist_session"), + patch.object(agent, "_save_trajectory"), + patch.object(agent, "_cleanup_task_resources"), + patch.object(agent.context_compressor, "update_model"), + patch.object(agent, "_compress_context", side_effect=_no_progress), + ): + result = agent.run_conversation("hello") + + assert result["completed"] is False + assert result.get("compression_exhausted") is True + # Terminated in a bounded number of API calls (default max attempts=3 + # => ~4 create calls), NOT an unbounded retry loop. + assert agent.client.chat.completions.create.call_count <= 6 @@ -5576,8 +5816,10 @@ class TestPersistUserMessageOverride: "2-3 sentences max. No code blocks or markdown.] Hello there" ) # But the DB write must get the override. - first_db_write = agent._session_db.append_message.call_args_list[0].kwargs - assert first_db_write["content"] == "Hello there" + batch = agent._session_db.append_messages_batch.call_args_list[0].kwargs[ + "messages" + ] + assert batch[0]["content"] == "Hello there" class TestReasoningReplayForStrictProviders: @@ -5895,3 +6137,4 @@ class TestMemoryContextSanitization: assert "memory-context" not in result.lower() assert "stale observation" not in result assert "how is the honcho working" in result + diff --git a/tests/run_agent/test_tool_name_db_persistence.py b/tests/run_agent/test_tool_name_db_persistence.py index 3fcf7f33c3..29596c04dd 100644 --- a/tests/run_agent/test_tool_name_db_persistence.py +++ b/tests/run_agent/test_tool_name_db_persistence.py @@ -27,7 +27,8 @@ def _make_agent(session_db): def test_tool_name_persisted_to_session_db(): """tool_name set by make_tool_result_message must be passed through to - append_message so the column is populated on first flush to the session DB.""" + the batched flush so the column is populated on first write to the + session DB.""" session_db = MagicMock() agent = _make_agent(session_db) @@ -37,9 +38,8 @@ def test_tool_name_persisted_to_session_db(): ] agent._flush_messages_to_session_db(messages) - tool_appends = [ - c for c in session_db.append_message.call_args_list - if c.kwargs.get("role") == "tool" - ] - assert len(tool_appends) == 1 - assert tool_appends[0].kwargs["tool_name"] == "terminal" + assert session_db.append_messages_batch.call_count == 1 + batch = session_db.append_messages_batch.call_args.kwargs["messages"] + tool_rows = [m for m in batch if m.get("role") == "tool"] + assert len(tool_rows) == 1 + assert tool_rows[0]["tool_name"] == "terminal" diff --git a/tests/test_empty_model_fallback.py b/tests/test_empty_model_fallback.py index 7a0ec903df..3d7f8505ec 100644 --- a/tests/test_empty_model_fallback.py +++ b/tests/test_empty_model_fallback.py @@ -27,16 +27,16 @@ class TestGetDefaultModelForProvider: with patch( "hermes_cli.model_catalog.get_default_model_from_cache", - return_value="qwen/qwen3.7-max", + return_value="qwen/qwen3.8-max", ): assert ( models_mod.get_preferred_silent_default_model("nous") - == "qwen/qwen3.7-max" + == "qwen/qwen3.8-max" ) - # nous catalog carries qwen3.7-max, so the full resolver follows. + # nous catalog carries qwen3.8-max, so the full resolver follows. assert ( models_mod.get_default_model_for_provider("nous") - == "qwen/qwen3.7-max" + == "qwen/qwen3.8-max" ) diff --git a/tests/test_fts_update_of_narrowing.py b/tests/test_fts_update_of_narrowing.py new file mode 100644 index 0000000000..2a8e0d504f --- /dev/null +++ b/tests/test_fts_update_of_narrowing.py @@ -0,0 +1,283 @@ +"""FTS UPDATE OF narrowing + migration (#73639 retargeted onto split SessionDB).""" + +from __future__ import annotations + +import sqlite3 +import tempfile +from pathlib import Path + +import pytest + +from hermes_state import SessionDB +from hermes_state_common import FTS_CJK_STALE_KEY +from hermes_state_schema import SessionSchemaMixin + + +def _trigger_sql(conn: sqlite3.Connection, name: str) -> str | None: + row = conn.execute( + "SELECT sql FROM sqlite_master WHERE type='trigger' AND name=?", + (name,), + ).fetchone() + return row[0] if row else None + + +def _assert_nonindexed_updates_bypass_missing_fts_target( + db: SessionDB, message_id: int +) -> None: + """Prove the UPDATE OF gate, not the trigger's content-change WHEN.""" + db._conn.execute("DROP TABLE messages_fts") + assert _trigger_sql(db._conn, "messages_fts_update") is not None + + db._conn.execute( + "UPDATE messages SET active = 0, compacted = 1, observed = 1 " + "WHERE id = ?", + (message_id,), + ) + with pytest.raises(sqlite3.OperationalError, match=r"no such table.*messages_fts"): + db._conn.execute( + "UPDATE messages SET content = 'changed' WHERE id = ?", + (message_id,), + ) + + +def _install_legacy_inline_base_fts(db: SessionDB) -> None: + """Replace v23 FTS with the broad inline shape shipped by v11..v22.""" + db._drop_fts_triggers(db._conn) + db._conn.executescript( + """ + DROP TABLE IF EXISTS messages_fts; + DROP TABLE IF EXISTS messages_fts_trigram; + DROP VIEW IF EXISTS messages_fts_trigram_src; + + CREATE VIRTUAL TABLE messages_fts USING fts5(content); + CREATE TRIGGER messages_fts_insert AFTER INSERT ON messages BEGIN + INSERT INTO messages_fts(rowid, content) VALUES ( + new.id, + COALESCE(new.content, '') || ' ' || + COALESCE(new.tool_name, '') || ' ' || + COALESCE(new.tool_calls, '') + ); + END; + CREATE TRIGGER messages_fts_delete AFTER DELETE ON messages BEGIN + DELETE FROM messages_fts WHERE rowid = old.id; + END; + CREATE TRIGGER messages_fts_update AFTER UPDATE ON messages BEGIN + DELETE FROM messages_fts WHERE rowid = old.id; + INSERT INTO messages_fts(rowid, content) VALUES ( + new.id, + COALESCE(new.content, '') || ' ' || + COALESCE(new.tool_name, '') || ' ' || + COALESCE(new.tool_calls, '') + ); + END; + """ + ) + + +def test_fresh_db_installs_update_of_triggers(tmp_path: Path): + db = SessionDB(db_path=tmp_path / "state.db") + try: + sql = _trigger_sql(db._conn, "messages_fts_update") + assert sql is not None + compact = " ".join(sql.split()).upper() + assert "AFTER UPDATE OF " in compact + assert "CONTENT" in compact + assert "TOOL_NAME" in compact + assert "TOOL_CALLS" in compact + + tri = _trigger_sql(db._conn, "messages_fts_trigram_update") + if tri: # trigram may be unavailable on some builds + tcompact = " ".join(tri.split()).upper() + assert "AFTER UPDATE OF " in tcompact + finally: + db.close() + + +def test_migrate_replaces_broad_update_trigger(tmp_path: Path): + path = tmp_path / "state.db" + db = SessionDB(db_path=path) + try: + # Force a broad trigger the way older installs had it. + db._conn.execute("DROP TRIGGER IF EXISTS messages_fts_update") + db._conn.execute( + """ + CREATE TRIGGER messages_fts_update AFTER UPDATE ON messages + BEGIN + SELECT 1; + END + """ + ) + db._conn.commit() + before = _trigger_sql(db._conn, "messages_fts_update") + assert "AFTER UPDATE OF" not in " ".join(before.split()).upper() + + dropped = db._migrate_broad_fts_update_triggers(db._conn) + db._conn.commit() + assert dropped >= 1 + + after = _trigger_sql(db._conn, "messages_fts_update") + assert after is not None + compact = " ".join(after.split()).upper() + assert "AFTER UPDATE OF " in compact + finally: + db.close() + + +def test_needs_narrowing_helper(): + assert SessionSchemaMixin._fts_update_trigger_needs_narrowing( + "CREATE TRIGGER t AFTER UPDATE ON messages BEGIN SELECT 1; END" + ) + assert not SessionSchemaMixin._fts_update_trigger_needs_narrowing( + "CREATE TRIGGER t AFTER UPDATE OF content ON messages BEGIN SELECT 1; END" + ) + assert not SessionSchemaMixin._fts_update_trigger_needs_narrowing(None) + + +def test_v23_status_only_update_bypasses_fts_trigger_body(tmp_path: Path): + db = SessionDB(db_path=tmp_path / "state.db") + try: + sid = "s1" + db.create_session(sid, source="test") + mid = db.append_message(sid, role="user", content="hello searchable") + _assert_nonindexed_updates_bypass_missing_fts_target(db, mid) + finally: + db.close() + + +def test_legacy_status_only_update_bypasses_migrated_fts_trigger(tmp_path: Path): + db = SessionDB(db_path=tmp_path / "state.db") + try: + sid = "s1" + db.create_session(sid, source="test") + mid = db.append_message(sid, role="user", content="legacy searchable") + _install_legacy_inline_base_fts(db) + + assert db._db_has_legacy_inline_fts(db._conn) + assert db._migrate_broad_fts_update_triggers(db._conn) >= 1 + assert "AFTER UPDATE OF" in " ".join( + _trigger_sql(db._conn, "messages_fts_update").split() + ).upper() + + _assert_nonindexed_updates_bypass_missing_fts_target(db, mid) + finally: + db.close() + + +def test_cjk_ensure_failure_marks_unavailable_and_propagates( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +): + path = tmp_path / "state.db" + db = SessionDB(db_path=path) + try: + db._conn.execute("DROP TRIGGER IF EXISTS messages_fts_cjk_update") + db._conn.execute( + "CREATE TRIGGER messages_fts_cjk_update " + "AFTER UPDATE ON messages BEGIN SELECT 1; END" + ) + db._fts_cjk_available = True + + def _fail_cjk_ensure(_cursor): + raise sqlite3.DatabaseError("injected CJK ensure failure") + + monkeypatch.setattr(db, "_ensure_fts_cjk_schema", _fail_cjk_ensure) + + with pytest.raises(sqlite3.DatabaseError, match="injected CJK ensure failure"): + db._migrate_broad_fts_update_triggers(db._conn) + assert db._fts_cjk_available is False + assert _trigger_sql(db._conn, "messages_fts_cjk_update") is None + + # The DROP is autocommitted. Fail-closed therefore needs a durable + # breadcrumb that other processes can observe, not just an instance + # flag on the SessionDB whose initialization is aborting. + with sqlite3.connect(path) as observer: + stale = observer.execute( + "SELECT value FROM state_meta WHERE key = ?", + (FTS_CJK_STALE_KEY,), + ).fetchone() + assert stale == ("1",) + assert _trigger_sql(observer, "messages_fts_cjk_update") is None + finally: + db.close() + + +def test_cjk_soft_fail_ensure_without_raise_marks_stale( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +): + """Production ensure swallows OperationalError — still must quarantine.""" + path = tmp_path / "state.db" + db = SessionDB(db_path=path) + try: + db._conn.execute("DROP TRIGGER IF EXISTS messages_fts_cjk_update") + db._conn.execute( + "CREATE TRIGGER messages_fts_cjk_update " + "AFTER UPDATE ON messages BEGIN SELECT 1; END" + ) + db._fts_cjk_available = True + + def _soft_fail_cjk_ensure(_cursor): + # Mirrors real _ensure_fts_cjk_schema OperationalError path: + # clear availability, do not raise, do not recreate triggers. + db._fts_cjk_available = False + + monkeypatch.setattr(db, "_ensure_fts_cjk_schema", _soft_fail_cjk_ensure) + + dropped = db._migrate_broad_fts_update_triggers(db._conn) + assert dropped >= 1 + assert db._fts_cjk_available is False + assert _trigger_sql(db._conn, "messages_fts_cjk_update") is None + + with sqlite3.connect(path) as observer: + stale = observer.execute( + "SELECT value FROM state_meta WHERE key = ?", + (FTS_CJK_STALE_KEY,), + ).fetchone() + assert stale == ("1",) + assert _trigger_sql(observer, "messages_fts_cjk_update") is None + finally: + db.close() + + +def test_cjk_broad_trigger_is_restored_as_update_of( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +): + db = SessionDB(db_path=tmp_path / "state.db") + try: + db._conn.execute("DROP TRIGGER IF EXISTS messages_fts_cjk_update") + db._conn.execute( + "CREATE TRIGGER messages_fts_cjk_update " + "AFTER UPDATE ON messages BEGIN SELECT 1; END" + ) + + def _restore_cjk_update_trigger(cursor): + cursor.execute( + "CREATE TRIGGER messages_fts_cjk_update " + "AFTER UPDATE OF content, tool_name, tool_calls ON messages " + "BEGIN SELECT 1; END" + ) + + monkeypatch.setattr(db, "_ensure_fts_cjk_schema", _restore_cjk_update_trigger) + + assert db._migrate_broad_fts_update_triggers(db._conn) == 1 + cjk_sql = _trigger_sql(db._conn, "messages_fts_cjk_update") + assert cjk_sql is not None + assert "AFTER UPDATE OF" in " ".join(cjk_sql.split()).upper() + finally: + db.close() + + +def test_legacy_migration_does_not_drop_cjk_trigger(tmp_path: Path): + db = SessionDB(db_path=tmp_path / "state.db") + try: + _install_legacy_inline_base_fts(db) + db._conn.execute("DROP TRIGGER IF EXISTS messages_fts_cjk_update") + db._conn.execute( + "CREATE TRIGGER messages_fts_cjk_update " + "AFTER UPDATE ON messages BEGIN SELECT 1; END" + ) + + assert db._migrate_broad_fts_update_triggers(db._conn) >= 1 + cjk_sql = _trigger_sql(db._conn, "messages_fts_cjk_update") + assert cjk_sql is not None + assert "AFTER UPDATE OF" not in " ".join(cjk_sql.split()).upper() + finally: + db.close() diff --git a/tests/test_hermes_state.py b/tests/test_hermes_state.py index 715b5e59af..74a028c754 100644 --- a/tests/test_hermes_state.py +++ b/tests/test_hermes_state.py @@ -656,12 +656,66 @@ class TestFTS5Search: assert isinstance(results[0]["context"], list) assert len(results[0]["context"]) > 0 + def test_search_fields_project_results_without_changing_default(self, db): + db.create_session(session_id="s1", source="cli") + db.append_message("s1", role="user", content="Tell me about Kubernetes") + db.append_message("s1", role="assistant", content="Kubernetes is an orchestrator.") + projected = db.search_messages( + "Kubernetes", fields=("session_id", "role", "snippet") + ) + default = db.search_messages("Kubernetes") + assert len(projected) == len(default) == 2 + assert all(set(row) == {"session_id", "role", "snippet"} for row in projected) + assert [ + (row["session_id"], row["role"], row["snippet"]) + for row in projected + ] == [ + (row["session_id"], row["role"], row["snippet"]) + for row in default + ] + assert all("context" in row and row["context"] for row in default) + def test_search_projection_skips_context_enrichment_queries(self, db): + db.create_session(session_id="s1", source="cli") + db.append_message("s1", role="user", content="before") + db.append_message("s1", role="assistant", content="projectionneedle") + db.append_message("s1", role="user", content="after") + statements = [] + read_conn = db._get_read_conn() or db._conn + traced_connections = [db._conn] + if read_conn is not db._conn: + traced_connections.append(read_conn) + for conn in traced_connections: + conn.set_trace_callback(statements.append) + def context_query_count(): + normalized = (" ".join(sql.upper().split()) for sql in statements) + return sum("WITH TARGET AS (" in sql for sql in normalized) + try: + projected = db.search_messages( + "projectionneedle", fields=("session_id", "snippet") + ) + assert len(projected) == 1 + assert context_query_count() == 0 + + full = db.search_messages( + "projectionneedle", fields=("session_id", "context") + ) + assert len(full) == 1 + assert full[0]["context"] + assert context_query_count() == 1 + + default = db.search_messages("projectionneedle") + assert len(default) == 1 + assert default[0]["context"] + assert context_query_count() == 2 + finally: + for conn in traced_connections: + conn.set_trace_callback(None) def test_sanitize_fts5_query_strips_dangerous_chars(self): """Unit test for _sanitize_fts5_query static method.""" @@ -816,6 +870,22 @@ class TestCounts: + def test_session_count_ge_empty(self, db): + """session_count_ge should return False for 0 sessions.""" + assert db.session_count_ge(1) is False + assert db.session_count_ge(2) is False + + def test_session_count_ge_at_threshold(self, db): + """session_count_ge should True when count >= n.""" + db.create_session("s1", "cli") + assert db.session_count_ge(1) is True + assert db.session_count_ge(2) is False + + db.create_session("s2", "telegram") + assert db.session_count_ge(1) is True + assert db.session_count_ge(2) is True + assert db.session_count_ge(3) is False + def test_message_count_total(self, db): assert db.message_count() == 0 db.create_session(session_id="s1", source="cli") @@ -1835,6 +1905,88 @@ class TestCompressionChainProjection: assert tip_row["ended_at"] is None # tip is still live assert tip_row["end_reason"] is None + def test_list_projects_multiple_independent_chains_in_one_call(self, db): + """Two unrelated compression chains in the same page must each + resolve to their own tip, not get cross-mixed by the batched tip-row + fetch (regression test for the single-query batch in + _get_session_rich_rows_batch — a wrong id->row mapping there would + silently swap one chain's data onto the other).""" + import time as _time + + t0 = _time.time() - 7200 + self._build_compression_chain(db, t0) + + # Second, independent chain — same shape, different ids/content. + db.create_session("root2", "cli") + db._conn.execute("UPDATE sessions SET started_at=? WHERE id=?", (t0 + 100, "root2")) + db.append_message("root2", "user", "second conversation start") + db._conn.execute( + "UPDATE sessions SET ended_at=?, end_reason=? WHERE id=?", + (t0 + 200, "compression", "root2"), + ) + db.create_session("tip2", "cli", parent_session_id="root2") + db._conn.execute("UPDATE sessions SET started_at=? WHERE id=?", (t0 + 201, "tip2")) + db.append_message("tip2", "user", "second conversation continuation") + db.update_session_cwd("tip2", "/tmp/workspaces/second") + db._conn.commit() + + sessions = db.list_sessions_rich(source="cli", limit=20) + ids = [s["id"] for s in sessions] + assert "root1" not in ids and "root2" not in ids + assert "tip1" in ids and "tip2" in ids + + tip1_row = next(s for s in sessions if s["id"] == "tip1") + tip2_row = next(s for s in sessions if s["id"] == "tip2") + assert tip1_row["_lineage_root_id"] == "root1" + assert tip1_row["preview"].startswith("latest message") + assert tip2_row["_lineage_root_id"] == "root2" + assert tip2_row["preview"].startswith("second conversation continuation") + assert tip2_row["cwd"] == "/tmp/workspaces/second" + + def test_list_batches_tip_row_fetch_into_one_query(self, db, monkeypatch): + """Projection must resolve tip rows for a whole page in one batched + query, not one _get_session_rich_row() call per compression root.""" + import time as _time + + t0 = _time.time() - 7200 + self._build_compression_chain(db, t0) + db.create_session("root2", "cli") + db._conn.execute("UPDATE sessions SET started_at=? WHERE id=?", (t0 + 100, "root2")) + db.append_message("root2", "user", "second conversation start") + db._conn.execute( + "UPDATE sessions SET ended_at=?, end_reason=? WHERE id=?", + (t0 + 200, "compression", "root2"), + ) + db.create_session("tip2", "cli", parent_session_id="root2") + db._conn.execute("UPDATE sessions SET started_at=? WHERE id=?", (t0 + 201, "tip2")) + db.append_message("tip2", "user", "second continuation") + db._conn.commit() + + batch_calls = [] + single_calls = [] + original_batch = db._get_session_rich_rows_batch + original_single = db._get_session_rich_row + + def counting_batch(session_ids, **kwargs): + batch_calls.append(list(session_ids)) + return original_batch(session_ids, **kwargs) + + def counting_single(session_id, **kwargs): + single_calls.append(session_id) + return original_single(session_id, **kwargs) + + monkeypatch.setattr(db, "_get_session_rich_rows_batch", counting_batch) + monkeypatch.setattr(db, "_get_session_rich_row", counting_single) + + sessions = db.list_sessions_rich(source="cli", limit=20) + assert len(sessions) >= 2 # sanity: both chains actually surfaced + + # Two compression roots resolved with exactly one batched call, and + # zero single-row calls — not one single-row call per root. + assert len(batch_calls) == 1 + assert set(batch_calls[0]) == {"tip1", "tip2"} + assert single_calls == [] + @@ -3262,6 +3414,48 @@ class TestCompactRows: assert "system_prompt" not in row assert row["id"] == "s1" + def test_batch_compact_rows_omits_system_prompt_keeps_git_fields(self, db): + """_get_session_rich_rows_batch(compact_rows=True) must apply the same + schema-derived compact projection as the single-row path: no + system_prompt blob, but git_branch/git_repo_root still present.""" + self._create(db, "s1", system_prompt="should be gone") + db.update_session_cwd("s1", "/tmp/w1", git_branch="main", git_repo_root="/tmp/w1") + rows = db._get_session_rich_rows_batch(["s1"], compact_rows=True) + assert set(rows) == {"s1"} + row = rows["s1"] + assert "system_prompt" not in row + assert row["git_branch"] == "main" + assert row["git_repo_root"] == "/tmp/w1" + + def test_compression_tip_projection_threads_compact_rows(self, db): + """list_sessions_rich(compact_rows=True) must thread compact_rows + through the batched tip-row fetch: the projected tip row must lack + system_prompt but keep git metadata (guards the call site at the + projection loop, not just the batch helper).""" + import time as _time + + t0 = _time.time() - 3600 + db.create_session("rootc", "cli") + db._conn.execute("UPDATE sessions SET started_at=? WHERE id=?", (t0, "rootc")) + db.append_message("rootc", "user", "start") + db._conn.execute( + "UPDATE sessions SET ended_at=?, end_reason=? WHERE id=?", + (t0 + 100, "compression", "rootc"), + ) + db.create_session("tipc", "cli", parent_session_id="rootc") + db._conn.execute("UPDATE sessions SET started_at=? WHERE id=?", (t0 + 101, "tipc")) + db.append_message("tipc", "user", "continuation") + db.update_system_prompt("tipc", "big blob " * 500) + db.update_session_cwd("tipc", "/tmp/w2", git_branch="dev", git_repo_root="/tmp/w2") + db._conn.commit() + + rows = db.list_sessions_rich(source="cli", compact_rows=True) + tip = next(s for s in rows if s["id"] == "tipc") + assert tip["_lineage_root_id"] == "rootc" + assert "system_prompt" not in tip + assert tip["git_branch"] == "dev" + assert tip["git_repo_root"] == "/tmp/w2" + @@ -3641,3 +3835,327 @@ class TestApplyDatabasePragmas: assert conn.execute("PRAGMA wal_autocheckpoint").fetchone()[0] == before finally: conn.close() + + def test_ignores_non_integer_performance_values(self, tmp_path, monkeypatch): + """Garbage cache_size/mmap_size/temp_store values must be rejected.""" + import sqlite3 + from hermes_state import apply_database_pragmas + + conn = sqlite3.connect(str(tmp_path / "pragmas.db")) + try: + before = { + name: conn.execute(f"PRAGMA {name}").fetchone()[0] + for name in ("cache_size", "mmap_size", "temp_store") + } + self._patch_cfg( + monkeypatch, + { + "database": { + "cache_size": "big", + "mmap_size": [256], + "temp_store": "ram please", + } + }, + ) + apply_database_pragmas(conn, db_label="test.db") + after = { + name: conn.execute(f"PRAGMA {name}").fetchone()[0] + for name in ("cache_size", "mmap_size", "temp_store") + } + assert after == before + finally: + conn.close() + + +class TestInsightsToolCallIndex: + """The Insights assistant tool-call scan has a predicate-aligned index. + + ``InsightsEngine._get_tool_usage`` / ``_get_skill_usage`` filter messages by + ``role = 'assistant' AND tool_calls IS NOT NULL``. A partial index over that + predicate keeps the scan off the full ``messages`` table on a large state.db. + """ + + _INDEX = "idx_messages_assistant_calls_by_session" + + def _index_defn(self, conn): + row = conn.execute( + "SELECT sql FROM sqlite_master WHERE type = 'index' AND name = ?", + (self._INDEX,), + ).fetchone() + return row["sql"] if row else None + + def test_index_created_on_fresh_db(self, tmp_path): + db = SessionDB(db_path=tmp_path / "fresh.db") + try: + sql = self._index_defn(db._conn) + assert sql is not None, "partial index missing on a fresh database" + # Partial predicate must match the queried rows exactly. + assert "role = 'assistant'" in sql + assert "tool_calls IS NOT NULL" in sql + finally: + db.close() + + def test_index_created_on_existing_db(self, tmp_path): + """Reopening a DB that predates the index must create it (SCHEMA_SQL is + re-run on every open; role/tool_calls are original base columns).""" + db_path = tmp_path / "legacy.db" + db = SessionDB(db_path=db_path) + # Simulate a database created before the index shipped. + db._conn.execute(f"DROP INDEX IF EXISTS {self._INDEX}") + db._conn.commit() + assert self._index_defn(db._conn) is None + db.close() + + db2 = SessionDB(db_path=db_path) + try: + assert self._index_defn(db2._conn) is not None, ( + "index not recreated when reopening an existing database" + ) + finally: + db2.close() + + def test_index_predicate_is_partial(self, db): + """The index covers only the assistant tool-call rows Insights reads. + + Query-plan coverage (that the Insights queries actually select this + index, for both scopes, without ANALYZE) lives with the queries in + tests/agent/test_insights.py. + """ + sql = self._index_defn(db._conn) + assert sql is not None + assert "WHERE" in sql + assert "role = 'assistant'" in sql + assert "tool_calls IS NOT NULL" in sql +class TestFtsRebuildFinishWithoutTrigram: + """An FTS index that the runtime cannot maintain must not wedge the store. + + Two independent failure sites shared one root shape: code that writes to + ``messages_fts_trigram`` without first checking the table is actually + present. It is legitimately absent whenever the trigram index is + unavailable (SQLite build without the tokenizer), and it can also be left + absent by an interrupted migration or a partially-applied schema change. + """ + + @staticmethod + def _seed(db_path, n=60): + seeded = SessionDB(db_path=db_path) + try: + seeded.create_session(session_id="s1", source="cli") + for i in range(n): + seeded.append_message( + "s1", + role=("user" if i % 3 == 0 + else "assistant" if i % 3 == 1 else "tool"), + content=f"sentinel payload {i} zebra", + ) + high_water = seeded._conn.execute( + "SELECT COALESCE(MAX(id), 0) FROM messages" + ).fetchone()[0] + finally: + seeded.close() + return high_water + + def test_rebuild_finish_skips_trigram_when_unavailable( + self, tmp_path, monkeypatch + ): + """optimize_fts_storage() completes when the trigram index is absent. + + ``fts_rebuild_step()`` already guards its backfill INSERT on + ``_trigram_available``; ``_fts_rebuild_finish()``'s boundary sweep did + not, so finishing a deferred rebuild on a trigram-less runtime raised + ``no such table: messages_fts_trigram`` and aborted the whole + optimization. The base index must still be swept and the markers + cleared. + """ + db_path = tmp_path / "state.db" + high_water = self._seed(db_path) + + real_connect = sqlite3.connect + + def connect_without_trigram(*args, **kwargs): + kwargs["factory"] = _NoTrigramConnection + return real_connect(*args, **kwargs) + + monkeypatch.setattr( + "hermes_state.sqlite3.connect", connect_without_trigram + ) + db = SessionDB(db_path=db_path) + try: + assert db._trigram_available is False + # A trigram-less runtime leaves no trigram index on disk. + db._conn.execute("DROP TABLE IF EXISTS messages_fts_trigram") + db._conn.commit() + assert db._fts_table_exists("messages_fts_trigram") is False + + # Put the DB in the pending-deferred-rebuild state. + for key, value in ( + ("fts_rebuild_high_water", str(high_water)), + ("fts_rebuild_progress", str(high_water)), + ): + db._conn.execute( + "INSERT INTO state_meta (key, value) VALUES (?, ?) " + "ON CONFLICT(key) DO UPDATE SET value = excluded.value", + (key, value), + ) + db._conn.commit() + + # Pre-fix this raised OperationalError("no such table: ..."). + db._fts_rebuild_finish() + + # The sweep ran to completion: markers cleared… + assert db.get_meta("fts_rebuild_high_water") is None + assert db.get_meta("fts_rebuild_progress") is None + # …and the base index is still usable (the fix must not disable + # real search to dodge the error). + assert db.search_messages("zebra") + finally: + db.close() + + def test_optimize_fts_storage_succeeds_without_trigram( + self, tmp_path, monkeypatch + ): + """End-to-end: the public optimize entry point returns ok=True.""" + db_path = tmp_path / "state.db" + high_water = self._seed(db_path) + + real_connect = sqlite3.connect + + def connect_without_trigram(*args, **kwargs): + kwargs["factory"] = _NoTrigramConnection + return real_connect(*args, **kwargs) + + monkeypatch.setattr( + "hermes_state.sqlite3.connect", connect_without_trigram + ) + db = SessionDB(db_path=db_path) + try: + db._conn.execute("DROP TABLE IF EXISTS messages_fts_trigram") + db._conn.commit() + assert db._trigram_available is False + for key, value in ( + ("fts_rebuild_high_water", str(high_water)), + ("fts_rebuild_progress", "0"), + ): + db._conn.execute( + "INSERT INTO state_meta (key, value) VALUES (?, ?) " + "ON CONFLICT(key) DO UPDATE SET value = excluded.value", + (key, value), + ) + db._conn.commit() + + result = db.optimize_fts_storage(vacuum=False) + assert result["ok"] is True + assert db.get_meta("fts_rebuild_high_water") is None + assert db.search_messages("zebra") + finally: + db.close() + + + +class TestPerformancePragmasEndToEnd: + """E2E guard for PR #71755: config-gated cache_size / mmap_size / + temp_store must reach EVERY connection type (writer, read-only + cross-profile attach, WAL per-thread reader) — and default installs + (no ``database:`` keys) must see byte-identical SQLite defaults. + + NOTE: SQLite's compiled-in default for ``cache_size`` is already + ``-2000``, so the configured value here is ``-16000`` — a value the + test can actually discriminate from the default (a reverted prod + change must FAIL this test, not accidentally pass it). + """ + + PRAGMAS = ("cache_size", "mmap_size", "temp_store") + CONFIGURED = {"cache_size": -16000, "mmap_size": 1048576, "temp_store": 2} + + @staticmethod + def _read(conn): + return { + name: conn.execute(f"PRAGMA {name}").fetchone()[0] + for name in ("cache_size", "mmap_size", "temp_store") + } + + @staticmethod + def _sqlite_defaults(tmp_path): + import sqlite3 + + conn = sqlite3.connect(str(tmp_path / "baseline.db")) + try: + return { + name: conn.execute(f"PRAGMA {name}").fetchone()[0] + for name in ("cache_size", "mmap_size", "temp_store") + } + finally: + conn.close() + + def _fresh_home(self, tmp_path, monkeypatch, config_text=None): + import hermes_state + + # Local venvs may bundle a WAL-reset-vulnerable SQLite (e.g. 3.46.0), + # which would silently disable WAL and skip the per-thread reader + # path. Force WAL eligibility so _get_read_conn is truly exercised + # (established pattern used by the WAL tests above). + monkeypatch.setattr( + hermes_state, + "is_sqlite_wal_reset_vulnerable", + lambda version_info=None: False, + ) + home = tmp_path / "hermes_home" + home.mkdir() + monkeypatch.setenv("HERMES_HOME", str(home)) + if config_text is not None: + (home / "config.yaml").write_text(config_text) + return home + + def test_configured_pragmas_reach_all_connection_types( + self, tmp_path, monkeypatch + ): + from hermes_state import SessionDB + + home = self._fresh_home( + tmp_path, + monkeypatch, + "database:\n" + " cache_size: -16000\n" + " temp_store: 2\n" + " mmap_size: 1048576\n", + ) + db_path = home / "state.db" + db = SessionDB(db_path=db_path) + try: + # Writer connection. + assert self._read(db._conn) == self.CONFIGURED + # WAL per-thread reader. + rconn = db._get_read_conn() + assert rconn is not None, "WAL reader expected on local filesystem" + assert self._read(rconn) == self.CONFIGURED + finally: + db.close() + + # Read-only cross-profile attach. + ro = SessionDB(db_path=db_path, read_only=True) + try: + assert self._read(ro._conn) == self.CONFIGURED + finally: + ro.close() + + def test_defaults_unchanged_without_config(self, tmp_path, monkeypatch): + """No database: keys in config.yaml → SQLite defaults untouched.""" + from hermes_state import SessionDB + + defaults = self._sqlite_defaults(tmp_path) + home = self._fresh_home(tmp_path, monkeypatch, config_text=None) + db_path = home / "state.db" + db = SessionDB(db_path=db_path) + try: + assert self._read(db._conn) == defaults + rconn = db._get_read_conn() + if rconn is not None: + assert self._read(rconn) == defaults + finally: + db.close() + + ro = SessionDB(db_path=db_path, read_only=True) + try: + assert self._read(ro._conn) == defaults + finally: + ro.close() diff --git a/tests/test_minimax_oauth.py b/tests/test_minimax_oauth.py index a7a5b66b8b..0f807a6dc7 100644 --- a/tests/test_minimax_oauth.py +++ b/tests/test_minimax_oauth.py @@ -42,7 +42,13 @@ from hermes_cli.auth import ( # --------------------------------------------------------------------------- def _make_httpx_response(status_code: int, body: dict | None = None, text: str = ""): - """Return a minimal mock that quacks like httpx.Response.""" + """Return a minimal mock that quacks like httpx.Response. + + Includes the streamed-read surface used by ``_minimax_post_form`` / + ``_minimax_response_error_text``: ``is_stream_consumed`` is False and + ``iter_bytes()`` yields the body/text bytes, so non-200 paths exercise + the real bounded-read code instead of a truthy MagicMock attribute. + """ resp = MagicMock() resp.status_code = status_code if body is not None: @@ -52,6 +58,9 @@ def _make_httpx_response(status_code: int, body: dict | None = None, text: str = resp.json.side_effect = Exception("No body") resp.text = text resp.reason_phrase = "OK" if status_code == 200 else "Error" + resp.is_stream_consumed = False + resp.encoding = "utf-8" + resp.iter_bytes.return_value = iter([resp.text.encode("utf-8")] if resp.text else []) return resp @@ -128,6 +137,7 @@ def test_request_user_code_state_mismatch_raises(): client = MagicMock() client.post.return_value = mock_response + client.send.return_value = mock_response with pytest.raises(AuthError) as exc_info: _minimax_request_user_code( @@ -387,6 +397,7 @@ def test_token_provider_refreshes_when_near_expiry(): mock_instance.__enter__ = MagicMock(return_value=mock_instance) mock_instance.__exit__ = MagicMock(return_value=False) mock_instance.post.return_value = mock_resp + mock_instance.send.return_value = mock_resp mock_client_class.return_value = mock_instance token = provider() @@ -441,6 +452,7 @@ def test_token_provider_quarantines_state_on_terminal_refresh(): mock_instance.__enter__ = MagicMock(return_value=mock_instance) mock_instance.__exit__ = MagicMock(return_value=False) mock_instance.post.return_value = bad_resp + mock_instance.send.return_value = bad_resp mock_client_class.return_value = mock_instance with pytest.raises(AuthError) as exc_info: @@ -475,3 +487,84 @@ def test_resolve_returns_callable_when_as_token_provider_true(): assert creds["base_url"] == MINIMAX_OAUTH_GLOBAL_INFERENCE.rstrip("/") +# --------------------------------------------------------------------------- +# Bounded error-body reads (#56548 / PR #56549) +# --------------------------------------------------------------------------- + +def test_refresh_error_body_bounded_and_readable_with_real_client(): + """Refresh non-200 path over a REAL socket transport. + + The error body is obtained via a streamed response; the bounded read + must happen while the client context is still open. A real socket is + required to bind this contract: closing the client tears the connection + down, so a read after the ``with httpx.Client(...)`` block raises + ReadError/StreamClosed. (MockTransport buffers in memory and would NOT + catch the regression.) + """ + import http.server + import socketserver + import threading + + import httpx + + from hermes_cli.auth import _refresh_minimax_oauth_state + + big_body = b"invalid_grant " + b"x" * (64 * 1024) # 64KB error body + + class Handler(http.server.BaseHTTPRequestHandler): + def do_POST(self): + self.send_response(400) + self.send_header("Content-Length", str(len(big_body))) + self.end_headers() + self.wfile.write(big_body) + + def log_message(self, *args): + pass + + with socketserver.TCPServer(("127.0.0.1", 0), Handler) as server: + port = server.server_address[1] + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + state = { + "access_token": "expired", + "refresh_token": "burned-rt", + "portal_base_url": f"http://127.0.0.1:{port}", + "client_id": MINIMAX_OAUTH_CLIENT_ID, + "inference_base_url": MINIMAX_OAUTH_GLOBAL_INFERENCE, + "expires_at": _past_iso(100), + } + with pytest.raises(AuthError) as exc_info: + _refresh_minimax_oauth_state(state, force=True) + finally: + server.shutdown() + + msg = str(exc_info.value) + assert "invalid_grant" in msg + assert exc_info.value.relogin_required is True + # Bounded: 16KB limit + truncation marker, never the full 64KB body. + assert len(msg) < 20 * 1024 + assert "...[truncated]" in msg + + +def test_minimax_response_error_text_truncates_above_limit(): + """Bodies above the 16KB bound are cut and marked truncated.""" + import httpx + + from hermes_cli.auth import ( + _MINIMAX_OAUTH_ERROR_BODY_LIMIT, + _minimax_response_error_text, + ) + + big = "e" * (_MINIMAX_OAUTH_ERROR_BODY_LIMIT * 4) + + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(500, text=big) + + with httpx.Client(transport=httpx.MockTransport(handler)) as client: + request = client.build_request("POST", "https://api.minimax.io/oauth/token") + response = client.send(request, stream=True) + text = _minimax_response_error_text(response) + + assert text.endswith("...[truncated]") + assert len(text) <= _MINIMAX_OAUTH_ERROR_BODY_LIMIT + len("...[truncated]") diff --git a/tests/test_session_db_read_path_split.py b/tests/test_session_db_read_path_split.py index 35c5422830..1b28b90589 100644 --- a/tests/test_session_db_read_path_split.py +++ b/tests/test_session_db_read_path_split.py @@ -129,3 +129,41 @@ def test_anchored_view_and_around_use_read_path(db): assert done["view"]["window"] finally: db._lock.release() + + +@pytest.mark.requires_wal +def test_session_resume_reads_do_not_take_writer_lock(db): + """session.resume's three read paths must not convoy behind writer flushes. + + get_messages_as_conversation / get_resume_conversations / + get_ancestor_display_prefix are the hottest reads in the file — every + resume across the gateway, CLI, and ACP adapter goes through one of + them — so they must use the same per-thread read-only connection as + get_messages, not the legacy self._lock path. + """ + db.create_session(session_id="parent1", source="cli", model="m") + db.append_message("parent1", role="user", content="parent turn") + db.append_message("parent1", role="assistant", content="parent reply") + db.create_session(session_id="child1", source="cli", model="m", parent_session_id="parent1") + db.append_message("child1", role="user", content="child turn") + db.append_message("child1", role="assistant", content="child reply") + + acquired = db._lock.acquire() + try: + done = {} + + def reader(): + done["conversation"] = db.get_messages_as_conversation("s1") + done["resume"] = db.get_resume_conversations("child1") + done["ancestor_prefix"] = db.get_ancestor_display_prefix("child1") + + t = threading.Thread(target=reader) + t.start(); t.join(timeout=5.0) + assert not t.is_alive(), "session resume reads blocked on writer lock" + assert len(done["conversation"]) == 2 + model_history, display_history = done["resume"] + assert len(model_history) == 2 + assert len(display_history) == 4 + assert len(done["ancestor_prefix"]) == 2 + finally: + db._lock.release() diff --git a/tests/test_session_system_prompt_dedup.py b/tests/test_session_system_prompt_dedup.py new file mode 100644 index 0000000000..966920a76c --- /dev/null +++ b/tests/test_session_system_prompt_dedup.py @@ -0,0 +1,267 @@ +"""Behavior coverage for content-addressed session system prompts.""" + +from __future__ import annotations + +import json +import sqlite3 +import time + +import pytest + +from hermes_state import SCHEMA_VERSION, SessionDB + + +@pytest.fixture() +def db(tmp_path): + session_db = SessionDB(db_path=tmp_path / "state.db") + yield session_db + session_db.close() + + +def _prompt_count(db: SessionDB) -> int: + return int( + db._conn.execute("SELECT COUNT(*) FROM system_prompts").fetchone()[0] + ) + + +def test_prompt_snapshots_are_deduplicated_and_hydrated_for_readers(db): + prompt = "You are Hermes.\n" + ("Follow the profile policy.\n" * 5) + db.create_session( + "s1", + "telegram", + session_key="agent:main:telegram:dm:c1", + chat_id="c1", + chat_type="dm", + system_prompt=prompt, + ) + db.create_session("s2", "cli", system_prompt=prompt) + db.request_handoff("s1", "telegram") + + stored = db._conn.execute( + "SELECT hash, prompt FROM system_prompts" + ).fetchall() + assert len(stored) == 1 + assert stored[0]["prompt"] == prompt + raw_sessions = db._conn.execute( + "SELECT system_prompt, system_prompt_hash FROM sessions ORDER BY id" + ).fetchall() + assert [row["system_prompt"] for row in raw_sessions] == [None, None] + assert {row["system_prompt_hash"] for row in raw_sessions} == { + stored[0]["hash"] + } + + assert db.get_session("s1")["system_prompt"] == prompt + assert db.list_sessions_rich()[0]["system_prompt"] == prompt + assert db.search_sessions()[0]["system_prompt"] == prompt + assert db.export_session("s1")["system_prompt"] == prompt + assert db.list_gateway_sessions()[0]["system_prompt"] == prompt + assert db.list_pending_handoffs()[0]["system_prompt"] == prompt + + +def test_prompt_replacement_and_route_changes_collect_only_orphans(db): + shared_prompt = "Model: x-ai/grok-4.5\nProvider: nous" + db.create_session( + "s1", + "hermes_browser", + model="x-ai/grok-4.5", + model_config={"_branched_from": "parent"}, + system_prompt=shared_prompt, + ) + db.create_session("s2", "cli", system_prompt=shared_prompt) + + db.update_session_runtime_lock( + "s1", + model="anthropic/claude-opus-4.8", + provider="anthropic", + confirmed=True, + ) + s1 = db.get_session("s1") + assert s1["system_prompt"] is None + assert json.loads(s1["model_config"])["_branched_from"] == "parent" + assert db.get_session("s2")["system_prompt"] == shared_prompt + assert _prompt_count(db) == 1 + + db.update_session_billing_route( + "s2", + provider="openrouter", + base_url="https://example.test/v1", + ) + assert db.get_session("s2")["system_prompt"] is None + assert _prompt_count(db) == 0 + + db.update_system_prompt("s2", "replacement") + assert db.get_session("s2")["system_prompt"] == "replacement" + db.update_system_prompt("s2", None) + assert _prompt_count(db) == 0 + + +def test_existing_session_enrichment_does_not_leak_unused_prompt(db): + db.create_session("s1", "cli", system_prompt="original prompt") + db.create_session("s1", "cli", system_prompt="unused prompt") + + prompts = [ + row["prompt"] + for row in db._conn.execute("SELECT prompt FROM system_prompts") + ] + assert prompts == ["original prompt"] + assert db.get_session("s1")["system_prompt"] == "original prompt" + + +def test_every_session_deletion_path_reclaims_final_prompt_reference(db): + def seed(session_id: str, *, source: str = "cli") -> None: + db.create_session( + session_id, + source, + system_prompt=f"unique prompt for {session_id}", + ) + assert _prompt_count(db) == 1 + + seed("single-empty") + assert db.delete_session_if_empty("single-empty") is True + assert _prompt_count(db) == 0 + + seed("bulk") + assert db.delete_sessions(["bulk"]) == 1 + assert _prompt_count(db) == 0 + + seed("ended-empty") + db.end_session("ended-empty", "user_exit") + assert db.delete_empty_sessions() == 1 + assert _prompt_count(db) == 0 + + seed("pruned") + db.end_session("pruned", "user_exit") + assert db.prune_sessions( + older_than_days=None, + started_before=time.time() + 1, + ) == 1 + assert _prompt_count(db) == 0 + + seed("ghost", source="tui") + db.end_session("ghost", "user_exit") + db._conn.execute("UPDATE sessions SET started_at = 0 WHERE id = 'ghost'") + db._conn.commit() + assert db.prune_empty_ghost_sessions() == 1 + assert _prompt_count(db) == 0 + + +def test_deleting_one_shared_session_preserves_prompt_until_final_reference(db): + prompt = "shared deletion prompt" + db.create_session("s1", "cli", system_prompt=prompt) + db.create_session("s2", "cli", system_prompt=prompt) + + assert db.delete_session("s1") is True + assert _prompt_count(db) == 1 + assert db.get_session("s2")["system_prompt"] == prompt + + assert db.delete_session("s2") is True + assert _prompt_count(db) == 0 + + +def test_compression_child_uses_content_addressed_prompt(db): + prompt = "compressed child prompt" + db.create_session("parent", "webui") + db.append_message("parent", "user", "original") + assert db.try_acquire_compression_lock("parent", "holder", ttl_seconds=60) + + db.publish_compression_child( + parent_session_id="parent", + child_session_id="child", + source="webui", + system_prompt=prompt, + messages=[{"role": "user", "content": "summary"}], + compression_lock_holder="holder", + ) + + raw = db._conn.execute( + "SELECT system_prompt, system_prompt_hash FROM sessions WHERE id = 'child'" + ).fetchone() + assert raw["system_prompt"] is None + assert raw["system_prompt_hash"] is not None + assert db.get_session("child")["system_prompt"] == prompt + assert _prompt_count(db) == 1 + + +def test_imported_prompts_are_deduplicated(tmp_path): + prompt = "shared imported prompt" + source = SessionDB(db_path=tmp_path / "source.db") + try: + source.create_session("s1", "cli", system_prompt=prompt) + source.create_session("s2", "telegram", system_prompt=prompt) + exported = [source.export_session("s1"), source.export_session("s2")] + finally: + source.close() + + target = SessionDB(db_path=tmp_path / "target.db") + try: + result = target.import_sessions(exported) + assert result["ok"] is True + assert result["imported"] == 2 + assert _prompt_count(target) == 1 + raw = target._conn.execute( + "SELECT system_prompt, system_prompt_hash FROM sessions ORDER BY id" + ).fetchall() + assert [row["system_prompt"] for row in raw] == [None, None] + assert len({row["system_prompt_hash"] for row in raw}) == 1 + assert target.get_session("s1")["system_prompt"] == prompt + assert target.get_session("s2")["system_prompt"] == prompt + finally: + target.close() + + +def test_v24_inline_prompts_migrate_once_to_content_addressed_storage(tmp_path): + db_path = tmp_path / "legacy-prompts.db" + legacy_prompt = "Legacy system prompt\n" + ("same policy\n" * 20) + + db = SessionDB(db_path=db_path) + db.create_session("s1", "cli") + db.create_session("s2", "telegram") + db._conn.execute( + "UPDATE sessions SET system_prompt = ?, system_prompt_hash = NULL", + (legacy_prompt,), + ) + db._conn.execute("UPDATE schema_version SET version = 24") + db._conn.commit() + db.close() + + migrated = SessionDB(db_path=db_path) + try: + assert migrated.get_session("s1")["system_prompt"] == legacy_prompt + assert migrated.get_session("s2")["system_prompt"] == legacy_prompt + assert _prompt_count(migrated) == 1 + raw_sessions = migrated._conn.execute( + "SELECT system_prompt, system_prompt_hash FROM sessions ORDER BY id" + ).fetchall() + assert [row["system_prompt"] for row in raw_sessions] == [None, None] + assert len({row["system_prompt_hash"] for row in raw_sessions}) == 1 + assert migrated._conn.execute( + "SELECT version FROM schema_version LIMIT 1" + ).fetchone()[0] == SCHEMA_VERSION + finally: + migrated.close() + + +def test_compact_rows_omit_hash_and_never_read_prompt_blob(db): + db.create_session("s1", "cli", system_prompt="never materialize me") + + def deny_prompt_reads(action, table, column, database, trigger): + if action == sqlite3.SQLITE_READ and table == "system_prompts": + return sqlite3.SQLITE_DENY + return sqlite3.SQLITE_OK + + db._conn.set_authorizer(deny_prompt_reads) + try: + rows = db.list_sessions_rich( + compact_rows=True, + order_by_last_active=True, + ) + rich = db._get_session_rich_row("s1", compact_rows=True) + finally: + db._conn.set_authorizer(None) + + assert rows[0]["id"] == "s1" + assert rich["id"] == "s1" + assert "system_prompt" not in rows[0] + assert "system_prompt_hash" not in rows[0] + assert "system_prompt" not in rich + assert "system_prompt_hash" not in rich diff --git a/tests/test_stale_tool_call_marker_session_repair.py b/tests/test_stale_tool_call_marker_session_repair.py new file mode 100644 index 0000000000..ed27fe4351 --- /dev/null +++ b/tests/test_stale_tool_call_marker_session_repair.py @@ -0,0 +1,262 @@ +"""Tests for stale tool-call marker session repair (hermes_state, #78148). + +Before the root-cause fix in ``agent.conversation_loop``, a local tool-call +template could emit a bare bracketed marker (e.g. "[memory]") as assistant +content alongside a real tool call. The loop cached that marker as a +fallback and, when the following turn came back empty, replayed it as the +"final response" — persisting it into the session as if the model had +actually answered. + +``_strip_stale_tool_call_markers`` is the load-on-read defense-in-depth +that clears any such stray marker content from sessions written before the +fix, so resuming a polluted session doesn't re-teach the model to keep +emitting the marker. Unaffected sessions pass through unchanged. +""" + +from hermes_state import ( + _is_stale_tool_call_marker_message, + _strip_stale_tool_call_markers, +) + + +class TestIsStaleToolCallMarkerMessage: + def test_matches_bare_marker_with_tool_calls(self): + msg = { + "role": "assistant", + "content": "[memory]", + "tool_calls": [{"id": "1", "function": {"name": "skill_manage", "arguments": "{}"}}], + } + assert _is_stale_tool_call_marker_message(msg) is True + + def test_matches_dotted_marker(self): + msg = { + "role": "assistant", + "content": "[foo.bar]", + "tool_calls": [{"id": "1", "function": {"name": "foo.bar", "arguments": "{}"}}], + } + assert _is_stale_tool_call_marker_message(msg) is True + + def test_ignores_marker_without_tool_calls(self): + # A genuine final response of "[memory]" with no tool call is not + # the contamination signature — leave it alone. + msg = {"role": "assistant", "content": "[memory]"} + assert _is_stale_tool_call_marker_message(msg) is False + + def test_ignores_real_content_with_tool_calls(self): + msg = { + "role": "assistant", + "content": "I'll check that for you.", + "tool_calls": [{"id": "1", "function": {"name": "skill_manage", "arguments": "{}"}}], + } + assert _is_stale_tool_call_marker_message(msg) is False + + def test_ignores_user_role(self): + msg = { + "role": "user", + "content": "[memory]", + "tool_calls": [{"id": "1", "function": {"name": "skill_manage", "arguments": "{}"}}], + } + assert _is_stale_tool_call_marker_message(msg) is False + + +class TestStripStaleToolCallMarkers: + def test_clears_contaminated_content_keeps_tool_calls(self): + messages = [ + {"role": "user", "content": "do the full task"}, + { + "role": "assistant", + "content": "[memory]", + "tool_calls": [{"id": "1", "function": {"name": "skill_manage", "arguments": "{}"}}], + }, + {"role": "tool", "content": "ok", "tool_call_id": "1"}, + ] + out = _strip_stale_tool_call_markers(messages) + assert out[1]["content"] == "" + # Tool call itself must survive — provider tool_call/result pairing. + assert out[1]["tool_calls"] == [{"id": "1", "function": {"name": "skill_manage", "arguments": "{}"}}] + + def test_unaffected_session_passes_through_unchanged(self): + messages = [ + {"role": "user", "content": "What's the weather?"}, + {"role": "assistant", "content": "It's sunny."}, + ] + out = _strip_stale_tool_call_markers(messages) + assert out == messages + + +class TestGetMessagesAsConversationStripsStaleMarkers: + """The load-on-read wiring: get_messages_as_conversation must actually + call _strip_stale_tool_call_markers, so a session polluted with a stale + "[memory]" marker resumes clean end-to-end (not just the pure helper in + isolation).""" + + def test_polluted_session_resumes_without_marker(self): + import tempfile + from pathlib import Path + from hermes_state import SessionDB + + with tempfile.TemporaryDirectory() as tmp: + db = SessionDB(db_path=Path(tmp) / "t.db") + try: + db.create_session(session_id="s1", source="cli") + db.append_message("s1", role="user", content="do the full task") + # Stray contamination written by an older build (pre-#78148 fix). + db.append_message( + "s1", role="assistant", content="[memory]", + tool_calls=[{"id": "1", "function": {"name": "skill_manage", "arguments": "{}"}}], + ) + db.append_message("s1", role="tool", content="ok", tool_call_id="1") + db.append_message("s1", role="assistant", content="Here is the result.") + + conv = db.get_messages_as_conversation("s1") + contents = [m.get("content") for m in conv if m.get("role") == "assistant"] + + assert "[memory]" not in contents + assert "Here is the result." in contents + finally: + db.close() + + def test_clean_session_resumes_unaffected(self): + import tempfile + from pathlib import Path + from hermes_state import SessionDB + + with tempfile.TemporaryDirectory() as tmp: + db = SessionDB(db_path=Path(tmp) / "t.db") + try: + db.create_session(session_id="s1", source="cli") + db.append_message("s1", role="user", content="What's the weather?") + db.append_message("s1", role="assistant", content="It's sunny.") + + conv = db.get_messages_as_conversation("s1") + contents = [m.get("content") for m in conv] + + assert contents == ["What's the weather?", "It's sunny."] + finally: + db.close() + + +class TestPurgeStaleToolCallMarkers: + """SessionDB.purge_stale_tool_call_markers: the permanent, one-time DB + rewrite. Complements the load-on-read repair — this variant edits the + stored rows in place so long-lived sessions stop re-scanning/re-repairing + the same contaminated rows on every resume.""" + + def _seed_polluted_db(self, db): + db.create_session(session_id="s1", source="cli") + db.append_message("s1", role="user", content="do the full task") + db.append_message( + "s1", role="assistant", content="[memory]", + tool_calls=[{"id": "1", "function": {"name": "skill_manage", "arguments": "{}"}}], + ) + db.append_message("s1", role="tool", content="ok", tool_call_id="1") + db.append_message("s1", role="assistant", content="Here is the result.") + + def test_dry_run_reports_without_writing(self): + import tempfile + from pathlib import Path + from hermes_state import SessionDB + + with tempfile.TemporaryDirectory() as tmp: + db = SessionDB(db_path=Path(tmp) / "t.db") + try: + self._seed_polluted_db(db) + + report = db.purge_stale_tool_call_markers(dry_run=True) + assert report["dry_run"] is True + assert report["rows_affected"] == 1 + + # Nothing written: the raw row still has the marker. + raw = db._conn.execute( + "SELECT content FROM messages WHERE role = 'assistant' " + "AND tool_calls IS NOT NULL AND tool_calls != ''" + ).fetchone() + assert raw["content"] == "[memory]" + finally: + db.close() + + def test_purge_clears_content_keeps_tool_calls(self): + import tempfile + from pathlib import Path + from hermes_state import SessionDB + + with tempfile.TemporaryDirectory() as tmp: + db = SessionDB(db_path=Path(tmp) / "t.db") + try: + self._seed_polluted_db(db) + + report = db.purge_stale_tool_call_markers(dry_run=False) + assert report["dry_run"] is False + assert report["rows_affected"] == 1 + # Backup defaults to on for a destructive, irreversible write. + assert report["backup_path"] is not None + assert Path(report["backup_path"]).exists() + + row = db._conn.execute( + "SELECT content, tool_calls FROM messages WHERE role = 'assistant' " + "AND tool_calls IS NOT NULL AND tool_calls != ''" + ).fetchone() + assert row["content"] == "" + # tool_calls column itself must survive the rewrite untouched. + assert row["tool_calls"] + + # Running again finds nothing left to clean — idempotent. + second = db.purge_stale_tool_call_markers(dry_run=False) + assert second["rows_affected"] == 0 + assert second["backup_path"] is None # nothing to change, nothing to back up + finally: + db.close() + + def test_no_backup_when_flag_false(self): + import tempfile + from pathlib import Path + from hermes_state import SessionDB + + with tempfile.TemporaryDirectory() as tmp: + db = SessionDB(db_path=Path(tmp) / "t.db") + try: + self._seed_polluted_db(db) + + report = db.purge_stale_tool_call_markers(dry_run=False, backup=False) + assert report["rows_affected"] == 1 + assert report["backup_path"] is None + # No extra file created beside the DB. + siblings = list(Path(tmp).glob("t.db.*backup*")) + assert siblings == [] + finally: + db.close() + + def test_dry_run_never_backs_up(self): + import tempfile + from pathlib import Path + from hermes_state import SessionDB + + with tempfile.TemporaryDirectory() as tmp: + db = SessionDB(db_path=Path(tmp) / "t.db") + try: + self._seed_polluted_db(db) + + report = db.purge_stale_tool_call_markers(dry_run=True) + assert report["backup_path"] is None + siblings = list(Path(tmp).glob("t.db.*backup*")) + assert siblings == [] + finally: + db.close() + + def test_no_affected_rows_on_clean_db(self): + import tempfile + from pathlib import Path + from hermes_state import SessionDB + + with tempfile.TemporaryDirectory() as tmp: + db = SessionDB(db_path=Path(tmp) / "t.db") + try: + db.create_session(session_id="s1", source="cli") + db.append_message("s1", role="user", content="What's the weather?") + db.append_message("s1", role="assistant", content="It's sunny.") + + report = db.purge_stale_tool_call_markers(dry_run=False) + assert report["rows_affected"] == 0 + assert report["row_ids"] == [] + finally: + db.close() diff --git a/tests/test_tui_gateway_server.py b/tests/test_tui_gateway_server.py index 9eb3468d5b..18dae86177 100644 --- a/tests/test_tui_gateway_server.py +++ b/tests/test_tui_gateway_server.py @@ -2415,8 +2415,10 @@ def test_history_to_messages_keeps_real_user_bracket_text(): ] -def test_session_resume_uses_parent_lineage_for_display(monkeypatch): +@pytest.mark.parametrize("omit_messages", [False, True]) +def test_session_resume_uses_parent_lineage_for_display(monkeypatch, omit_messages): captured = {} + target = "tip-omit" if omit_messages else "tip-full" class FakeDB: def get_session(self, target): @@ -2466,15 +2468,25 @@ def test_session_resume_uses_parent_lineage_for_display(monkeypatch): # _neuter_agent_prewarm_timer fixture; this test only asserts the # returned display history. + params = {"session_id": target} + if omit_messages: + params["omit_messages"] = True resp = server.handle_request( - {"id": "1", "method": "session.resume", "params": {"session_id": "tip"}} + {"id": "1", "method": "session.resume", "params": params} ) - assert resp["result"]["messages"] == [ + expected = [] if omit_messages else [ {"role": "user", "text": "root prompt"}, {"role": "assistant", "text": "root answer"}, ] - assert captured["history_calls"] == [("tip", False), ("tip", True)] + assert resp["result"]["messages"] == expected + assert resp["result"]["message_count"] == (1 if omit_messages else 2) + assert resp["result"]["messages_omitted"] is omit_messages + expected_calls = [(target, False)] if omit_messages else [ + (target, False), + (target, True), + ] + assert captured["history_calls"] == expected_calls def test_live_visible_history_prefers_db_display_with_candidate(): @@ -8907,6 +8919,121 @@ def test_prompt_submit_history_version_mismatch_surfaces_warning(monkeypatch): server._sessions.pop("sid", None) +def test_prompt_submit_merges_on_model_switch_marker(monkeypatch): + """#76870: when a model-switch marker is the only history mutation during + a turn, the agent's output must be merged into the current history (which + now contains the marker) instead of being discarded. + + This test covers BOTH cases: + - No prior marker in turn-start history (first switch in a session) + - Prior marker existed (every subsequent switch — the original PR #77274 + fix was dead code here because _append_model_switch_marker strips the + old marker before appending the new one, producing a net-zero length + delta that the positional slice missed). + """ + from tui_gateway.server import _MODEL_SWITCH_MARKER_PREFIX + + session_ref = {"s": None} + + def _make_marker(model: str) -> dict: + return { + "role": "user", + "content": f"{_MODEL_SWITCH_MARKER_PREFIX}{model}.]", + "display_kind": "model_switch", + } + + class _MarkerAgent: + def __init__(self, new_history_state: list): + self._new_history_state = new_history_state + + def run_conversation(self, prompt, conversation_history=None, stream_callback=None, **_kwargs): + # Simulate _append_model_switch_marker: strip prior markers, append new one. + with session_ref["s"]["history_lock"]: + hist = session_ref["s"]["history"] + hist[:] = [h for h in hist if not _is_marker(h)] + hist.append(_make_marker("new-model")) + session_ref["s"]["history_version"] += 1 + # result["messages"] = conversation_history + user msg + assistant reply + return { + "final_response": "agent reply", + "messages": list(conversation_history) + [ + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": "agent reply"}, + ], + } + + def _is_marker(entry) -> bool: + from tui_gateway.server import _is_model_switch_marker + return _is_model_switch_marker(entry) + + class _ImmediateThread: + def __init__(self, target=None, daemon=None): + self._target = target + + def start(self): + self._target() + + # Test both: no prior marker, and prior marker present + for label, prior_history in [ + ("no prior marker", [{"role": "user", "content": "hello"}]), + ("with prior marker", [ + {"role": "user", "content": "hello"}, + _make_marker("old-model"), + {"role": "assistant", "content": "hi there"}, + ]), + ]: + server._sessions["sid"] = _session( + agent=_MarkerAgent([]), + history=list(prior_history), + ) + session_ref["s"] = server._sessions["sid"] + emits: list[tuple] = [] + try: + monkeypatch.setattr(server.threading, "Thread", _ImmediateThread) + monkeypatch.setattr(server, "_get_usage", lambda _a: {}) + monkeypatch.setattr(server, "render_message", lambda _t, _c: "") + monkeypatch.setattr(server, "_emit", lambda *a: emits.append(a)) + + resp = server.handle_request( + { + "id": "1", + "method": "prompt.submit", + "params": {"session_id": "sid", "text": "hi"}, + } + ) + assert resp.get("result"), f"[{label}] got error: {resp.get('error')}" + + final_history = server._sessions["sid"]["history"] + + # The agent's new messages must be present in the persisted history. + assistant_msgs = [ + e for e in final_history + if isinstance(e, dict) and e.get("role") == "assistant" + and e.get("content") == "agent reply" + ] + assert len(assistant_msgs) == 1, ( + f"[{label}] agent output was not merged into history " + f"(got {len(assistant_msgs)} assistant 'agent reply' messages)" + ) + + # The model-switch marker must be present. + markers = [e for e in final_history if _is_marker(e)] + assert len(markers) == 1, ( + f"[{label}] expected exactly 1 model-switch marker, got {len(markers)}" + ) + assert "new-model" in markers[0]["content"] + + # No warning should be surfaced — the merge succeeded. + complete_calls = [a for a in emits if a[0] == "message.complete"] + assert len(complete_calls) == 1 + _, _, payload = complete_calls[0] + assert "warning" not in payload, ( + f"[{label}] merge path should not surface a warning" + ) + finally: + server._sessions.pop("sid", None) + + def test_prompt_submit_sanitizes_bracketed_paste_before_agent(monkeypatch): """prompt.submit must sanitize corrupted user text before run_conversation.""" captured: dict[str, str] = {} @@ -8999,6 +9126,59 @@ def test_prompt_submit_history_version_match_persists_normally(monkeypatch): server._sessions.pop("sid", None) +def test_prompt_submit_snapshots_history_after_pending_model_switch(monkeypatch): + marker = {"role": "user", "content": "[model switched]"} + seen = {} + + class _Agent: + def run_conversation(self, prompt, conversation_history=None, **_kwargs): + seen["history"] = conversation_history + return { + "final_response": "reply", + "messages": [ + *(conversation_history or []), + {"role": "user", "content": prompt}, + {"role": "assistant", "content": "reply"}, + ], + } + + class _ImmediateThread: + def __init__(self, target=None, **_kwargs): + self._target = target + + def start(self): + self._target() + + def _apply_pending(_sid, session): + with session["history_lock"]: + session["history"].append(marker) + session["history_version"] += 1 + + server._sessions["sid"] = _session(agent=_Agent()) + server._sessions["sid"]["pending_model_switch"] = {"raw": "new-model"} + emits = [] + try: + monkeypatch.setattr(server.threading, "Thread", _ImmediateThread) + monkeypatch.setattr(server, "_apply_pending_model_switch", _apply_pending) + monkeypatch.setattr(server, "_sync_agent_model_with_config", lambda *_a: None) + monkeypatch.setattr(server, "_get_usage", lambda _a: {}) + monkeypatch.setattr(server, "render_message", lambda *_a: "") + monkeypatch.setattr(server, "_emit", lambda *a: emits.append(a)) + + server.handle_request( + {"id": "1", "method": "prompt.submit", "params": {"session_id": "sid", "text": "hi"}} + ) + + assert seen["history"] == [marker] + assert server._sessions["sid"]["history"][-1] == { + "role": "assistant", "content": "reply" + } + complete = [a for a in emits if a[0] == "message.complete"] + assert "warning" not in complete[0][2] + finally: + server._sessions.pop("sid", None) + + def test_prompt_submit_can_truncate_before_user_ordinal(monkeypatch): """Desktop user-message edits should restart the turn from the edited user.""" @@ -11258,6 +11438,11 @@ def test_session_branch_writes_to_parent_profile_db(monkeypatch, tmp_path): def append_message(self, **kwargs): seen["msgs"].append(kwargs) + def append_messages_batch(self, session_id, messages, **kwargs): + for m in messages: + seen["msgs"].append(dict(m, session_id=session_id)) + return list(range(1, len(messages) + 1)) + def set_session_title(self, key, title): seen["title"] = (key, title) return True @@ -11370,6 +11555,11 @@ def test_session_branch_installs_parent_profile_secret_scope(monkeypatch, tmp_pa def append_message(self, **kwargs): seen["msgs"].append(kwargs) + def append_messages_batch(self, session_id, messages, **kwargs): + for m in messages: + seen["msgs"].append(dict(m, session_id=session_id)) + return list(range(1, len(messages) + 1)) + def set_session_title(self, key, title): return True @@ -12228,6 +12418,33 @@ def test_session_activate_switches_live_session_without_closing_siblings(monkeyp server._sessions.pop("sid-b", None) +def test_session_activate_can_omit_duplicate_desktop_transcript(monkeypatch): + monkeypatch.setattr(server, "_session_info", lambda agent: {"model": agent.model}) + server._sessions["sid-large"] = _session( + agent=types.SimpleNamespace(model="model-large"), + history=[ + {"role": "user", "content": "large prompt"}, + {"role": "assistant", "content": "large answer"}, + ], + session_key="key-large", + ) + try: + resp = server.handle_request( + { + "id": "1", + "method": "session.activate", + "params": {"session_id": "sid-large", "omit_messages": True}, + } + ) + + assert resp["result"]["messages"] == [] + assert resp["result"]["message_count"] == 2 + assert resp["result"]["messages_omitted"] is True + assert resp["result"]["session_key"] == "key-large" + finally: + server._sessions.pop("sid-large", None) + + # ── session.most_recent ────────────────────────────────────────────── @@ -15846,7 +16063,7 @@ def test_native_vision_turn_persists_a_renderable_image_ref(tmp_path): agent._flush_messages_to_session_db([{"role": "user", "content": native_parts}], []) - written = agent._session_db.append_message.call_args.kwargs["content"] + written = agent._session_db.append_messages_batch.call_args.kwargs["messages"][0]["content"] assert f"@image:`{img}`" in written assert "what is in this photo?" in written # The model keeps the pixels for the rest of the session. diff --git a/tests/tools/test_approval.py b/tests/tools/test_approval.py index 56ae0c1258..8892a089fc 100644 --- a/tests/tools/test_approval.py +++ b/tests/tools/test_approval.py @@ -1462,3 +1462,109 @@ class TestApprovalPromptRedaction: # The script's credential must not appear in the user-facing message. assert "sk-proj-abc123xyz4567890abcdef" not in result["message"] assert "sk-proj-abc123xyz4567890abcdef" not in result["command"] + + +class TestCliApprovalTimeoutClassifiedSeparately: + """CLI-path parity for the timeout-vs-deny distinction. + + The gateway wait already reported "timed out without user response"; + the CLI/TUI callback path collapsed a prompt timeout into "deny", so + the agent was told the user *refused* when the user simply never + answered. The prompt now returns a distinct "timeout" choice and both + guard tails classify it with outcome="timeout" + a "Silence is not + consent." message. + """ + + def _interactive_env(self): + return mock_patch.dict( + "os.environ", + {"HERMES_INTERACTIVE": "1"}, + clear=False, + ) + + def test_prompt_returns_timeout_when_input_never_arrives(self): + """The raw input() path returns 'timeout', not 'deny', on expiry.""" + import builtins + from unittest.mock import patch as _patch + + def _hang(_prompt=""): + time.sleep(10) + return "" + + with _patch.object(builtins, "input", _hang): + result = prompt_dangerous_approval( + "rm -rf /var/data", "recursive delete", + timeout_seconds=0.05, + ) + assert result == "timeout" + + def test_guard_classifies_callback_timeout_as_timeout(self, monkeypatch): + """check_all_command_guards: a 'timeout' choice from the CLI callback + yields outcome='timeout' and a no-response message, not 'denied by + user'.""" + from unittest.mock import patch as _patch + from tools import approval as mod + + mod._session_approved.clear() + mod._permanent_approved.clear() + + cfg = {"approvals": {"mode": "manual"}} + with self._interactive_env(): + with _patch("hermes_cli.config.load_config_readonly", return_value=cfg): + result = mod.check_all_command_guards( + "rm -rf /var/data", "local", + approval_callback=lambda *a, **kw: "timeout", + ) + + assert result["approved"] is False + assert result.get("outcome") == "timeout" + assert result.get("user_consent") is False + msg = result["message"] + assert "timed out without user response" in msg + assert "Silence is not consent" in msg + assert "denied" not in msg.lower() + + def test_guard_still_classifies_explicit_deny_as_denied(self): + """Explicit CLI deny keeps outcome='denied' and the denial wording.""" + from unittest.mock import patch as _patch + from tools import approval as mod + + mod._session_approved.clear() + mod._permanent_approved.clear() + + cfg = {"approvals": {"mode": "manual"}} + with self._interactive_env(): + with _patch("hermes_cli.config.load_config_readonly", return_value=cfg): + result = mod.check_all_command_guards( + "rm -rf /var/data", "local", + approval_callback=lambda *a, **kw: "deny", + ) + + assert result["approved"] is False + assert result.get("outcome") == "denied" + assert "denied" in result["message"].lower() + assert "Silence is not consent" not in result["message"] + + def test_run_approval_gate_cli_timeout_is_not_a_denial(self): + """The shared plugin-escalation gate (_run_approval_gate) also + distinguishes a prompt timeout from an explicit deny on the CLI + path.""" + from unittest.mock import patch as _patch + from tools import approval as mod + + mod._session_approved.clear() + mod._permanent_approved.clear() + + cfg = {"approvals": {"mode": "manual"}} + with self._interactive_env(): + with _patch("hermes_cli.config.load_config_readonly", return_value=cfg): + result = mod.request_tool_approval( + "write_file", "plugin flagged this write", + approval_callback=lambda *a, **kw: "timeout", + ) + + assert result["approved"] is False + assert result.get("outcome") == "timeout" + assert result.get("user_consent") is False + assert "timed out without user response" in result["message"] + assert "Silence is not consent" in result["message"] diff --git a/tests/tools/test_base_environment.py b/tests/tools/test_base_environment.py index c8416b9ea7..3d875598d6 100644 --- a/tests/tools/test_base_environment.py +++ b/tests/tools/test_base_environment.py @@ -115,25 +115,27 @@ class TestAtomicSnapshotWrite: assert f"> '{snap}'" not in wrapped assert f"> {snap}\n" not in wrapped - def test_temp_path_uses_bashpid_not_dollardollar(self): - """The temp name MUST use ``$BASHPID`` (the real subshell PID), not - ``$$``. In ``&``-launched concurrent subshells ``$$`` stays the parent - shell's PID, so two writers would pick the same temp name, clobber each - other mid-write, and mv would publish a torn file — the corruption is - only narrowed, not closed. This is the bug shared by every prior PR in - the #38249 cluster.""" + def test_temp_path_uses_mktemp_not_pid_variables(self): + """The temp name MUST be allocated by ``mktemp`` — never ``$$`` (in + ``&``-launched concurrent subshells it stays the parent shell's PID, so + two writers would pick the same temp name and publish a torn file) and + never ``$BASHPID`` (macOS ships bash 3.2, which lacks it — the name + expands empty, collapsing every writer onto one temp path and + reopening the #38249 race). Regression for PR #54314.""" env = _TestableEnv() env._snapshot_ready = True wrapped = env._wrap_command("echo hi", "/tmp") - assert "$BASHPID" in wrapped + assert "mktemp " in wrapped + assert ".tmp.XXXXXXXXXX" in wrapped + assert "$BASHPID" not in wrapped # The bare $$ temp form must be gone. assert ".tmp.$$" not in wrapped - def test_init_session_bootstrap_also_atomic_and_bashpid(self): + def test_init_session_bootstrap_also_atomic_and_mktemp(self): """The init_session bootstrap (first snapshot write) is the same shared file a concurrent command could source — it must be atomic and use - ``$BASHPID`` too.""" + ``mktemp`` too (no ``$BASHPID``: absent on macOS bash 3.2).""" env = _TestableEnv() captured = {} @@ -148,7 +150,8 @@ class TestAtomicSnapshotWrite: pass boot = captured.get("cmd", "") assert ".tmp." in boot and "mv -f " in boot, boot - assert "$BASHPID" in boot + assert "mktemp " in boot + assert "$BASHPID" not in boot assert ".tmp.$$" not in boot @@ -178,8 +181,9 @@ class TestAtomicSnapshotConcurrencyBehavioral: the emitted script's guarantee holds under real concurrency: N concurrent writers + readers, and the snapshot is ALWAYS a complete, parseable env dump — never truncated mid-line with a ``declare -x`` / ``export`` fragment - that would corrupt PATH. Crucially it uses ``$BASHPID`` (per-subshell - unique), which is what closes the race; ``$$`` would still tear here. + that would corrupt PATH. Crucially it allocates the temp with ``mktemp`` + (per-writer unique, works on macOS bash 3.2 which lacks ``$BASHPID``), + which is what closes the race; ``$$`` would still tear here. """ def _run(self, script): @@ -194,13 +198,14 @@ class TestAtomicSnapshotConcurrencyBehavioral: import shlex snap = str(tmp_path / "hermes-snap-x.sh") _q = shlex.quote - _snap_tmp = _q(snap + ".tmp.") + "$BASHPID" + _tmpl = _q(snap + ".tmp.XXXXXXXXXX") # One writer iteration = the exact atomic sequence _wrap_command emits. writer = ( "for i in $(seq 1 80); do " "export BIG_$i=$(head -c 600 /dev/zero | tr '\\0' x); " - f"{{ export -p > {_snap_tmp} && mv -f {_snap_tmp} {_q(snap)}; }} " - f"2>/dev/null || rm -f {_snap_tmp} 2>/dev/null || true; " + f"__hermes_snap_tmp=$(mktemp {_tmpl}) && " + f"{{ export -p > \"$__hermes_snap_tmp\" && mv -f \"$__hermes_snap_tmp\" {_q(snap)}; }} " + f"2>/dev/null || rm -f \"$__hermes_snap_tmp\" 2>/dev/null || true; " "done" ) # Reader: repeatedly source the snapshot and check PATH never absorbs @@ -235,10 +240,11 @@ class TestAtomicSnapshotConcurrencyBehavioral: self._run(f"echo 'export GOOD=1' > {_q(snap)}") # seed good snapshot # Redirect export into an unwritable dir so the export side fails; mv # must then NOT run (&&) and not clobber snap. - bad_tmp = _q("/nonexistent-dir/snap.tmp.") + "$BASHPID" + bad_tmp = _q("/nonexistent-dir/snap.tmp.XXXXXXXXXX") script = ( - f"{{ export -p > {bad_tmp} && mv -f {bad_tmp} {_q(snap)}; }} " - f"2>/dev/null || rm -f {bad_tmp} 2>/dev/null || true" + f"__hermes_snap_tmp=$(mktemp {bad_tmp}) && " + f"{{ export -p > \"$__hermes_snap_tmp\" && mv -f \"$__hermes_snap_tmp\" {_q(snap)}; }} " + f"2>/dev/null || rm -f \"$__hermes_snap_tmp\" 2>/dev/null || true" ) self._run(script) out = self._run(f"cat {_q(snap)}") diff --git a/tests/tools/test_delegate.py b/tests/tools/test_delegate.py index 33a3fb8619..fc97153591 100644 --- a/tests/tools/test_delegate.py +++ b/tests/tools/test_delegate.py @@ -81,6 +81,56 @@ class TestDelegateRequirements(unittest.TestCase): self.assertNotIn("acp_args", props["tasks"]["items"]["properties"]) self.assertNotIn("maxItems", props["tasks"]) # removed — limit is now runtime-configurable + def test_top_level_description_compact_and_complete(self): + """The top-level description must stay compact while keeping every + contract that exists nowhere else in the schema (keyword-level, not + prose-literal, so rewording doesn't break CI).""" + from tools.delegate_tool import _build_top_level_description + + desc = _build_top_level_description() + # Compaction ceiling: the old description was ~4,000 chars. + self.assertLessEqual(len(desc), 2200) + # Contracts only the top-level text carries: + for keyword in ( + "background", # async semantics + "wait or poll", # no-poll rule + "execute_code", # mechanical-work routing + "cronjob", # durable-work routing + "/stop", # non-durability warning + "context", # pass-everything-via-context rule + "respond in Chinese", # language example (weak models regress without it) + "SELF-REPORTS", # verification contract + "fetch the URL", # concrete verification verbs + "clarify", # leaf blocked-tool list + "send_message", + "delegation.provider", # model inheritance / pinning + ): + self.assertIn(keyword, desc, f"top-level description lost: {keyword!r}") + + def test_dynamic_limits_moved_to_param_descriptions(self): + """Concurrency and nesting ceilings must reach the model through the + tasks/role parameter descriptions (the top-level text no longer + carries them).""" + from tools.delegate_tool import _build_dynamic_schema_overrides + from tools.registry import registry + + with ( + patch("tools.delegate_tool._get_max_concurrent_children", return_value=7), + patch("tools.delegate_tool._get_max_spawn_depth", return_value=4), + patch("tools.delegate_tool._get_orchestrator_enabled", return_value=True), + ): + overrides = _build_dynamic_schema_overrides() + definition = registry.get_definitions({"delegate_task"})[0]["function"] + + for parameters in (overrides["parameters"], definition["parameters"]): + self.assertIn("up to 7", parameters["properties"]["tasks"]["description"]) + self.assertIn( + "max_spawn_depth=4", parameters["properties"]["role"]["description"] + ) + # Static top-level text must not embed stale limits. + self.assertNotIn("up to 7", overrides["description"]) + self.assertNotIn("max_spawn_depth", overrides["description"]) + class TestChildSystemPrompt(unittest.TestCase): def test_goal_only(self): prompt = _build_child_system_prompt("Fix the tests") diff --git a/tests/tools/test_file_sync.py b/tests/tools/test_file_sync.py index a5850dd3e3..9ecabdb21b 100644 --- a/tests/tools/test_file_sync.py +++ b/tests/tools/test_file_sync.py @@ -1,8 +1,10 @@ """Tests for FileSyncManager — mtime tracking, deletion detection, transactional rollback.""" +import concurrent.futures import io import os import tarfile +import threading import time from pathlib import Path from unittest.mock import MagicMock, patch @@ -261,6 +263,59 @@ class TestEdgeCases: upload.assert_not_called() # _file_mtime_key returns None, skipped +class TestConcurrency: + def test_sync_back_waits_for_active_sync_transaction(self, tmp_path): + initial_file = tmp_path / "initial.png" + new_file = tmp_path / "new.png" + initial_file.write_bytes(b"initial") + upload_started = threading.Event() + release_upload = threading.Event() + sync_back_transport_started = threading.Event() + overlap_detected = threading.Event() + download_calls = [] + + def get_files(): + return [ + (str(path), f"/root/.hermes/cache/images/{path.name}") + for path in sorted(tmp_path.glob("*.png")) + ] + + def upload(host_path, _remote_path): + if host_path == str(new_file): + upload_started.set() + sync_back_transport_started.wait(timeout=1.0) + release_upload.set() + + def bulk_download(destination): + if not release_upload.is_set(): + overlap_detected.set() + sync_back_transport_started.set() + download_calls.append(destination) + with tarfile.open(destination, "w"): + pass + + mgr = FileSyncManager( + get_files_fn=get_files, + upload_fn=upload, + delete_fn=MagicMock(), + bulk_download_fn=bulk_download, + ) + mgr.sync(force=True) + new_file.write_bytes(b"new") + + with concurrent.futures.ThreadPoolExecutor(max_workers=2) as executor: + sync_future = executor.submit(mgr.sync, force=True) + assert upload_started.wait(timeout=2.0) + + sync_back_future = executor.submit(mgr.sync_back, hermes_home=tmp_path) + + sync_future.result(timeout=3.0) + sync_back_future.result(timeout=3.0) + + assert len(download_calls) == 1 + assert not overlap_detected.is_set() + + class TestSyncBackSecurity: def test_sync_back_does_not_overwrite_uploaded_credential_files(self, tmp_path, monkeypatch): credential = tmp_path / "token.json" diff --git a/tests/tools/test_file_write_safety.py b/tests/tools/test_file_write_safety.py index d59dce7b21..ee1296e495 100644 --- a/tests/tools/test_file_write_safety.py +++ b/tests/tools/test_file_write_safety.py @@ -308,6 +308,35 @@ class TestBomHandling: raw = target.read_bytes() assert raw == self.BOM.encode("utf-8") + b"import os, json\nimport sys\n" + def test_v4a_update_preserves_bom_real_ops(self, ops, tmp_path: Path): + # V4A UPDATE path against REAL ShellFileOperations. This is the one + # provider path whose pre_content is BOM-STRIPPED (read_file_raw + # strips before _apply_update forwards it), so it regresses if + # _file_has_bom ever trusts pre_content instead of probing disk. + # Regression for teknium1's review on PR #55661. + target = tmp_path / "bom_v4a.py" + target.write_bytes(self.BOM.encode("utf-8") + b"print('hello')\n") + patch = ( + "*** Begin Patch\n" + f"*** Update File: {target}\n" + "@@\n" + "-print('hello')\n" + "+print('world')\n" + "*** End Patch" + ) + res = ops.patch_v4a(patch) + assert res.success, res.error + raw = target.read_bytes() + assert raw.startswith(self.BOM.encode("utf-8")), "BOM lost on V4A update" + assert b"print('world')" in raw + + def test_file_has_bom_ignores_stripped_pre_content(self, ops, tmp_path: Path): + # _file_has_bom must probe the DISK even when handed pre_content + # that (having been BOM-stripped upstream) claims there is no BOM. + target = tmp_path / "bom_probe.py" + target.write_bytes(self.BOM.encode("utf-8") + b"x = 1\n") + assert ops._file_has_bom(str(target), pre_content="x = 1\n") is True + if __name__ == "__main__": pytest.main([__file__, "-v"]) diff --git a/tests/tools/test_hardline_blocklist.py b/tests/tools/test_hardline_blocklist.py index 44e57b5d40..40b1c56755 100644 --- a/tests/tools/test_hardline_blocklist.py +++ b/tests/tools/test_hardline_blocklist.py @@ -239,6 +239,74 @@ def test_quoted_and_brace_paths_are_hardline_blocked(command): assert desc +# Multi-line QUOTED arguments are data, not command sequences: a newline +# inside quotes is part of the argument the shell passes to the program. +# These previously tripped the hardline floor because the flat command-start +# class treated every raw newline — even inside quotes — as a command +# boundary, blocking `hermes send` message bodies, multi-line +# `git commit -m` messages, and heredoc text that merely MENTION +# shutdown/reboot commands. +_QUOTED_NEWLINE_DATA_ALLOW = [ + # hermes send with a multi-line message body (the reported symptom) + 'hermes send -t telegram -s "spark1" "console output:\nsudo reboot\ndone"', + 'hermes send -t telegram "line1\nshutdown -h now\nline3"', + # git commit -m with a multi-line message + "git commit -m 'ops notes:\nreboot the box after the deploy'", + 'git commit -m "fix startup\nsystemctl reboot was flaky here"', + # heredoc bodies quoting dangerous strings as data + "python3 - <<'EOF'\nmsg = 'run sudo reboot later'\nprint(msg)\nEOF", + "cat > /tmp/notes.txt <<'EOF'\nremember: shutdown -h now\nEOF", + # rm hardline floor is anchored to the same class — quoted prose about it + # across a line break must stay data too + 'git commit -m "docs:\nwarn about rm -rf / in the guide"', +] + +# The masking must be strictly scoped to quoted data: real command +# boundaries around/inside those same shapes still hit the floor. +_QUOTED_NEWLINE_THREATS_BLOCK = [ + # unquoted newline is a real command separator + "echo hi\nsudo reboot", + 'echo "a"\nsudo reboot', + 'git commit -m "safe message"\nshutdown -h now', + # command substitution inside double quotes really executes + 'hermes send -t telegram "$(sudo reboot)"', + 'echo "`shutdown -h now`"', + # multi-line quoted data followed by a REAL chained command + 'hermes send "line1\nline2" && sudo reboot', + # a heredoc whose body is data, but the delivery command itself is hardline + "sudo reboot <<'EOF'\nignored\nEOF", +] + + +@pytest.mark.parametrize("command", _QUOTED_NEWLINE_DATA_ALLOW) +def test_quoted_newline_data_not_blocked(command): + """Newlines inside quoted arguments are data, not command starts.""" + is_hl, desc = detect_hardline_command(command) + assert not is_hl, ( + f"multi-line quoted data false-positived the hardline floor: " + f"{command!r} (got: {desc})" + ) + + +@pytest.mark.parametrize("command", _QUOTED_NEWLINE_THREATS_BLOCK) +def test_real_newline_separated_threats_still_blocked(command): + """Unquoted newlines / $() / backticks remain real command boundaries.""" + is_hl, desc = detect_hardline_command(command) + assert is_hl, f"real threat leaked through hardline floor: {command!r}" + assert desc + + +def test_quoted_newline_data_not_blocked_by_full_guard_chain(clean_session): + """End-to-end: the guard chain must not hardline-block a multi-line + quoted message (yolo on, so only the unconditional floor can block).""" + enable_session_yolo("hardline_test") + command = 'hermes send -t telegram "status:\nsudo reboot happened at 3am"' + result = check_all_command_guards(command, "local") + assert result["approved"], ( + f"guard chain blocked multi-line quoted data: {result.get('message')}" + ) + + # Commands that carry the literal string "rm -rf /" (or a sibling) as DATA in # another command's quoted argument — a PR title, a commit message, an echo / # printf argument. The shell never executes that text as an rm command, so the diff --git a/tests/tools/test_image_generation_artifacts.py b/tests/tools/test_image_generation_artifacts.py index 890330556c..ae2e9bede3 100644 --- a/tests/tools/test_image_generation_artifacts.py +++ b/tests/tools/test_image_generation_artifacts.py @@ -1,4 +1,6 @@ +import concurrent.futures import json +import threading from types import SimpleNamespace @@ -38,6 +40,79 @@ def test_postprocess_adds_agent_visible_image_for_active_ssh_env(monkeypatch, tm assert sync_calls == [True] +def test_concurrent_image_results_preserve_shared_remote_sync_state(monkeypatch, tmp_path): + from tools import image_generation_tool + from tools.environments import file_sync + + hermes_home = tmp_path / ".hermes" + image_dir = hermes_home / "cache" / "images" + image_dir.mkdir(parents=True) + first_image = image_dir / "first.png" + second_image = image_dir / "second.png" + first_image.write_bytes(b"first") + + first_upload_started = threading.Event() + second_sync_finished = threading.Event() + worker = threading.local() + + def get_files(): + return [ + ( + str(path), + f"/home/remote/.hermes/cache/images/{path.name}", + ) + for path in sorted(image_dir.iterdir()) + ] + + def upload(_host_path, _remote_path): + if worker.label == "first": + first_upload_started.set() + # Without transaction serialization, the second sync commits its + # newer snapshot before this first sync resumes and overwrites it. + second_sync_finished.wait(timeout=1.0) + + sync_manager = file_sync.FileSyncManager( + get_files_fn=get_files, + upload_fn=upload, + delete_fn=lambda _paths: None, + ) + env = SimpleNamespace( + _remote_home="/home/remote", + _sync_manager=sync_manager, + ) + + monkeypatch.setenv("HERMES_HOME", str(hermes_home)) + monkeypatch.setattr(file_sync, "_credential_host_paths", lambda: set()) + monkeypatch.setattr(image_generation_tool, "_active_terminal_env", lambda _task_id: env) + + def postprocess(label, image_path): + worker.label = label + try: + raw = json.dumps({"success": True, "image": str(image_path)}) + return image_generation_tool._postprocess_image_generate_result( + raw, + task_id="shared-task", + ) + finally: + if label == "second": + second_sync_finished.set() + + with concurrent.futures.ThreadPoolExecutor(max_workers=2) as executor: + first_future = executor.submit(postprocess, "first", first_image) + assert first_upload_started.wait(timeout=2.0) + + second_image.write_bytes(b"second") + second_future = executor.submit(postprocess, "second", second_image) + + first_future.result(timeout=3.0) + second_future.result(timeout=3.0) + + assert set(sync_manager._synced_files) == { + "/home/remote/.hermes/cache/images/first.png", + "/home/remote/.hermes/cache/images/second.png", + } + + def test_handle_image_generate_postprocesses_plugin_result(monkeypatch, tmp_path): from tools import image_generation_tool diff --git a/tests/tools/test_lazy_deps_managed.py b/tests/tools/test_lazy_deps_managed.py new file mode 100644 index 0000000000..bd756d59eb --- /dev/null +++ b/tests/tools/test_lazy_deps_managed.py @@ -0,0 +1,120 @@ +"""Managed-install guard in :func:`tools.lazy_deps.ensure` (#48628). + +A package-manager install (NixOS, and anything else shipping Hermes from a +read-only store) cannot receive lazy pip installs: the venv's site-packages +lives in the store, so the uv -> pip -> ensurepip ladder burns ~15s +bootstrapping ensurepip only to fail. ``ensure()`` must fail fast instead. +""" + +import pytest + +from tools import lazy_deps +from tools.lazy_deps import FeatureUnavailable + + +FEATURE = "provider.anthropic" + + +@pytest.fixture(autouse=True) +def _missing_and_installable(monkeypatch): + """Reach the guard: deps missing, installs allowed, no durable target. + + ``_allow_lazy_installs`` is patched explicitly so the suite does not + depend on the host's ~/.hermes/config.yaml (a local + ``allow_lazy_installs: false`` otherwise short-circuits with a different + rejection reason). + """ + monkeypatch.setattr(lazy_deps, "feature_missing", lambda _f: ("some-pkg==1.0",)) + monkeypatch.setattr(lazy_deps, "_allow_lazy_installs", lambda: True) + monkeypatch.setattr(lazy_deps, "_lazy_install_target", lambda: None) + + +def _no_installer(monkeypatch): + """Fail loudly if the guard lets execution reach the install ladder.""" + def _boom(*_a, **_kw): + raise AssertionError("guard let execution reach the install ladder") + + monkeypatch.setattr(lazy_deps.subprocess, "run", _boom) + + +def test_nixos_install_fails_fast_without_touching_the_installer(monkeypatch): + monkeypatch.setattr("hermes_cli.config.get_managed_system", lambda: "nixos") + _no_installer(monkeypatch) + + with pytest.raises(FeatureUnavailable) as excinfo: + lazy_deps.ensure(FEATURE, prompt=False) + + assert "nixos" in excinfo.value.reason + # refresh_active_features classifies by this prefix — anything else is + # reported to the user as a hard failure instead of a skip. + assert excinfo.value.reason.startswith("unsupported ") + + +def test_reason_is_classified_as_skipped_not_failed(monkeypatch): + """The wording contract with refresh_active_features, pinned directly.""" + monkeypatch.setattr("hermes_cli.config.get_managed_system", lambda: "nixos") + + with pytest.raises(FeatureUnavailable) as excinfo: + lazy_deps.ensure(FEATURE, prompt=False) + + assert excinfo.value.reason.startswith("unsupported "), ( + "refresh_active_features would report this as failed: rather than skipped:" + ) + + +def test_unmanaged_install_is_not_blocked_by_the_guard(monkeypatch): + """On a normal pip install the guard must be transparent.""" + monkeypatch.setattr("hermes_cli.config.get_managed_system", lambda: None) + + with pytest.raises(FeatureUnavailable) as excinfo: + lazy_deps.ensure(FEATURE, prompt=False) + + # Whatever stops the install here, it must NOT be the managed guard. + assert "managed installs" not in excinfo.value.reason + + +def test_durable_install_target_overrides_the_guard(monkeypatch, tmp_path): + """The container deployment sets HERMES_MANAGED *and* a writable target. + + Dockerfile sets HERMES_LAZY_INSTALL_TARGET and the NixOS container module + passes HERMES_MANAGED=true; blocking there would break that deployment. + """ + monkeypatch.setattr("hermes_cli.config.get_managed_system", lambda: "nixos") + monkeypatch.setattr(lazy_deps, "_lazy_install_target", lambda: tmp_path) + + with pytest.raises(FeatureUnavailable) as excinfo: + lazy_deps.ensure(FEATURE, prompt=False) + + assert "nixos" not in excinfo.value.reason.lower(), ( + "durable-target installs must not be blocked by the managed guard" + ) + + +def test_platform_unsupported_takes_precedence(monkeypatch): + """A platform-specific reason is more actionable than 'managed install'. + + Also required for consistency: refresh_active_features pre-checks + _unsupported_feature_reason before calling ensure(). + """ + monkeypatch.setattr("hermes_cli.config.get_managed_system", lambda: "nixos") + monkeypatch.setattr( + lazy_deps, "_unsupported_feature_reason", lambda _f: "unsupported on win32" + ) + + with pytest.raises(FeatureUnavailable) as excinfo: + lazy_deps.ensure(FEATURE, prompt=False) + + assert excinfo.value.reason == "unsupported on win32" + + +def test_unreadable_config_fails_open(monkeypatch): + """A broken config must not block installs on a normal pip install.""" + def _raise(): + raise RuntimeError("config unreadable") + + monkeypatch.setattr("hermes_cli.config.get_managed_system", _raise) + + with pytest.raises(FeatureUnavailable) as excinfo: + lazy_deps.ensure(FEATURE, prompt=False) + + assert "managed" not in excinfo.value.reason.lower() diff --git a/tests/tools/test_mcp_lazy_start.py b/tests/tools/test_mcp_lazy_start.py new file mode 100644 index 0000000000..85312c0fc2 --- /dev/null +++ b/tests/tools/test_mcp_lazy_start.py @@ -0,0 +1,331 @@ +"""Behavior-contract tests for lazy MCP server startup (#56832). + +A server configured with ``lazy: true`` whose config fingerprint matches an +on-disk schema-cache entry registers its tools WITHOUT spawning/connecting; +the first real call (raw tool OR resource/prompt utility) routes through the +existing connect path. +""" + +import json +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +import tools.mcp_tool as mcp + + +@pytest.fixture(autouse=True) +def _reset_mcp_state(): + old_servers = dict(mcp._servers) + old_lazy = dict(mcp._lazy_server_configs) + old_fps = dict(mcp._lazy_server_fingerprints) + old_names = dict(mcp._lazy_server_tool_names) + old_connecting = set(mcp._server_connecting) + yield + mcp._servers.clear() + mcp._servers.update(old_servers) + mcp._lazy_server_configs.clear() + mcp._lazy_server_configs.update(old_lazy) + mcp._lazy_server_fingerprints.clear() + mcp._lazy_server_fingerprints.update(old_fps) + mcp._lazy_server_tool_names.clear() + mcp._lazy_server_tool_names.update(old_names) + mcp._server_connecting.clear() + mcp._server_connecting.update(old_connecting) + + +def _fake_cache_entry(): + return { + "fingerprint": "abc", + "tools": [ + { + "name": "browser_navigate", + "description": "Navigate", + "inputSchema": {"type": "object", "properties": {}}, + } + ], + "utility_tools": [], + } + + +def _lazy_config(): + return { + "playwright": { + "command": "npx", + "args": ["-y", "@playwright/mcp"], + "lazy": True, + } + } + + +class TestLazyMcpRegistration: + def test_registers_from_cache_without_connect(self): + config = _lazy_config() + with patch("tools.mcp_tool._MCP_AVAILABLE", True), \ + patch("tools.mcp_schema_cache.config_fingerprint", return_value="abc"), \ + patch("tools.mcp_schema_cache.get_cached_entry", return_value=_fake_cache_entry()), \ + patch( + "tools.mcp_tool._register_from_cache_sync", + return_value=["mcp_playwright_browser_navigate"], + ) as mock_register, \ + patch("tools.mcp_tool._discover_and_register_server", new_callable=AsyncMock) as mock_discover, \ + patch("tools.mcp_tool._ensure_mcp_loop") as mock_loop, \ + patch("tools.mcp_tool._run_on_mcp_loop") as mock_run: + + mcp.register_mcp_servers(config) + + mock_register.assert_called_once() + mock_discover.assert_not_called() + mock_run.assert_not_called() + mock_loop.assert_not_called() + + def test_cache_miss_falls_back_to_eager_connect(self): + config = _lazy_config() + with patch("tools.mcp_tool._MCP_AVAILABLE", True), \ + patch("tools.mcp_schema_cache.config_fingerprint", return_value="abc"), \ + patch("tools.mcp_schema_cache.get_cached_entry", return_value=None), \ + patch("tools.mcp_tool._ensure_mcp_loop"), \ + patch("tools.mcp_tool._run_on_mcp_loop") as mock_run: + + mcp.register_mcp_servers(config) + + mock_run.assert_called_once() + + def test_non_lazy_server_never_touches_cache(self): + config = {"playwright": {"command": "npx", "args": []}} + with patch("tools.mcp_tool._MCP_AVAILABLE", True), \ + patch("tools.mcp_schema_cache.get_cached_entry") as mock_get, \ + patch("tools.mcp_tool._ensure_mcp_loop"), \ + patch("tools.mcp_tool._run_on_mcp_loop") as mock_run: + + mcp.register_mcp_servers(config) + + mock_get.assert_not_called() + mock_run.assert_called_once() + + def test_lazy_server_not_reregistered_on_second_pass(self): + config = _lazy_config() + mcp._lazy_server_configs["playwright"] = dict(config["playwright"]) + mcp._lazy_server_tool_names["playwright"] = ["mcp_playwright_browser_navigate"] + with patch("tools.mcp_tool._MCP_AVAILABLE", True), \ + patch("tools.mcp_tool._register_from_cache_sync") as mock_register, \ + patch("tools.mcp_tool._run_on_mcp_loop") as mock_run: + + names = mcp.register_mcp_servers(config) + + mock_register.assert_not_called() + mock_run.assert_not_called() + assert "mcp_playwright_browser_navigate" in names + + +class TestLazyFirstUseConnect: + def _connected_server(self): + mock_session = MagicMock() + mock_session.call_tool = AsyncMock( + return_value=SimpleNamespace(isError=False, content=[], structuredContent=None) + ) + connected = SimpleNamespace( + session=mock_session, + _rpc_lock=MagicMock(), + _pending_call_context=None, + ) + connected._rpc_lock.__aenter__ = AsyncMock(return_value=None) + connected._rpc_lock.__aexit__ = AsyncMock(return_value=None) + return connected + + @staticmethod + def _run_on_loop(coro_or_factory, timeout=120): + import asyncio + + coro = coro_or_factory() if callable(coro_or_factory) else coro_or_factory + loop = asyncio.new_event_loop() + try: + return loop.run_until_complete(coro) + finally: + loop.close() + + def test_tool_handler_lazy_connects_on_first_call(self): + config = {"command": "npx", "args": [], "lazy": True, "timeout": 5} + mcp._lazy_server_configs["playwright"] = dict(config) + mcp._lazy_server_fingerprints["playwright"] = "abc" + + connected = self._connected_server() + + def _connect(name): + mcp._servers["playwright"] = connected + return True + + with patch.object(mcp, "_ensure_lazy_server_connected", side_effect=_connect) as mock_connect, \ + patch.object(mcp, "_run_on_mcp_loop", side_effect=self._run_on_loop): + handler = mcp._make_tool_handler("playwright", "browser_navigate", 5) + out = handler({}, task_id="t1") + + mock_connect.assert_called_once_with("playwright") + payload = json.loads(out) + assert "error" not in payload + assert payload.get("result") == "" + + def test_list_resources_handler_lazy_connects_on_first_call(self): + # Regression for the resource/prompt gap: utility handlers must also + # route through the first-use connect path, or the first + # list_resources/get_prompt on a lazy server fails. + config = {"command": "npx", "args": [], "lazy": True, "timeout": 5} + mcp._lazy_server_configs["playwright"] = dict(config) + + connected = self._connected_server() + connected.session.list_resources = AsyncMock() + + def _connect(name): + mcp._servers["playwright"] = connected + return True + + async def _fake_paginate(list_method, items_attr, server_name): + return [SimpleNamespace(uri="file:///a", name="a", description="", mimeType="")] + + with patch.object(mcp, "_ensure_lazy_server_connected", side_effect=_connect) as mock_connect, \ + patch.object(mcp, "_paginate_full_list", side_effect=_fake_paginate), \ + patch.object(mcp, "_run_on_mcp_loop", side_effect=self._run_on_loop): + handler = mcp._make_list_resources_handler("playwright", 5) + out = handler({}) + + mock_connect.assert_called_once_with("playwright") + payload = json.loads(out) + assert "error" not in payload + assert payload["resources"][0]["uri"] == "file:///a" + + def test_get_prompt_handler_lazy_connects_on_first_call(self): + config = {"command": "npx", "args": [], "lazy": True, "timeout": 5} + mcp._lazy_server_configs["playwright"] = dict(config) + + connected = self._connected_server() + connected.session.get_prompt = AsyncMock( + return_value=SimpleNamespace(messages=[]) + ) + + def _connect(name): + mcp._servers["playwright"] = connected + return True + + with patch.object(mcp, "_ensure_lazy_server_connected", side_effect=_connect) as mock_connect, \ + patch.object(mcp, "_run_on_mcp_loop", side_effect=self._run_on_loop): + handler = mcp._make_get_prompt_handler("playwright", 5) + out = handler({"name": "greeting"}) + + mock_connect.assert_called_once_with("playwright") + payload = json.loads(out) + assert "error" not in payload + + def test_check_fn_passes_for_lazy_registered_server(self): + mcp._lazy_server_configs["playwright"] = {"lazy": True} + mcp._lazy_server_fingerprints["playwright"] = "abc" + assert mcp._make_check_fn("playwright")() is True + + def test_check_fn_fails_for_unknown_server(self): + assert mcp._make_check_fn("nope")() is False + + def test_lazy_connect_respects_connect_cooldown(self): + mcp._lazy_server_configs["playwright"] = {"command": "npx", "lazy": True} + with patch.object(mcp, "_connect_cooldown_active", return_value=True), \ + patch.object(mcp, "_run_on_mcp_loop") as mock_run: + assert mcp._ensure_lazy_server_connected("playwright") is False + mock_run.assert_not_called() + + def test_lazy_connect_success_clears_lazy_state(self): + config = {"command": "npx", "lazy": True} + mcp._lazy_server_configs["playwright"] = dict(config) + mcp._lazy_server_fingerprints["playwright"] = "abc" + mcp._lazy_server_tool_names["playwright"] = ["mcp_playwright_browser_navigate"] + + connected = SimpleNamespace( + session=MagicMock(), + _registered_tool_names=["mcp_playwright_browser_navigate"], + ) + + def _fake_run(coro_or_factory, timeout=30): + mcp._servers["playwright"] = connected + coro = coro_or_factory() if callable(coro_or_factory) else coro_or_factory + coro.close() + return ["mcp_playwright_browser_navigate"] + + with patch.object(mcp, "_ensure_mcp_loop"), \ + patch.object(mcp, "_run_on_mcp_loop", side_effect=_fake_run): + assert mcp._ensure_lazy_server_connected("playwright") is True + + assert "playwright" not in mcp._lazy_server_configs + assert "playwright" not in mcp._lazy_server_fingerprints + assert "playwright" not in mcp._lazy_server_tool_names + + def test_lazy_connect_deregisters_phantom_cached_tools(self): + # Stale-cache reconciliation: the cached manifest advertised tool X, + # but the live server only registers tool Y → X must be deregistered + # after the first-use connect so the model stops seeing a phantom. + from tools.registry import registry + + mcp._lazy_server_configs["playwright"] = {"command": "npx", "lazy": True} + mcp._lazy_server_fingerprints["playwright"] = "stale-fp" + mcp._lazy_server_tool_names["playwright"] = [ + "mcp_playwright_tool_x", + "mcp_playwright_tool_y", + ] + + connected = SimpleNamespace( + session=MagicMock(), + _registered_tool_names=["mcp_playwright_tool_y"], + ) + + def _fake_run(coro_or_factory, timeout=30): + mcp._servers["playwright"] = connected + coro = coro_or_factory() if callable(coro_or_factory) else coro_or_factory + coro.close() + return ["mcp_playwright_tool_y"] + + with patch.object(mcp, "_ensure_mcp_loop"), \ + patch.object(mcp, "_run_on_mcp_loop", side_effect=_fake_run), \ + patch.object(registry, "deregister") as mock_dereg: + assert mcp._ensure_lazy_server_connected("playwright") is True + + mock_dereg.assert_called_once_with("mcp_playwright_tool_x") + + def test_lazy_connect_failure_records_cooldown(self): + mcp._lazy_server_configs["playwright"] = {"command": "npx", "lazy": True} + + def _fake_run(coro_or_factory, timeout=30): + coro = coro_or_factory() if callable(coro_or_factory) else coro_or_factory + coro.close() + raise RuntimeError("spawn failed") + + with patch.object(mcp, "_ensure_mcp_loop"), \ + patch.object(mcp, "_run_on_mcp_loop", side_effect=_fake_run), \ + patch.object(mcp, "_record_connect_failure") as mock_record: + assert mcp._ensure_lazy_server_connected("playwright") is False + + mock_record.assert_called_once_with("playwright") + # Config retained so a later call can retry after cooldown. + assert "playwright" in mcp._lazy_server_configs + + +class TestCacheLoadDescriptionScan: + def test_scan_runs_on_cache_load_path(self): + # Defense-in-depth: the cache file is user-writable JSON, so the + # cache-load registration path must run the same injection scan as + # eager discovery. + entry = _fake_cache_entry() + config = {"command": "npx", "args": [], "lazy": True} + with patch.object(mcp, "_scan_mcp_description", return_value=[]) as mock_scan, \ + patch.object(mcp, "_convert_mcp_schema", side_effect=RuntimeError("stop")), \ + pytest.raises(RuntimeError): + mcp._register_from_cache_sync("playwright", config, entry) + + mock_scan.assert_called_once_with("playwright", "browser_navigate", "Navigate") + + +class TestResolveServerLazy: + def test_default_off(self): + assert mcp._resolve_server_lazy("s", {"command": "npx"}) is False + + def test_explicit_true(self): + assert mcp._resolve_server_lazy("s", {"command": "npx", "lazy": True}) is True + + def test_explicit_false(self): + assert mcp._resolve_server_lazy("s", {"command": "npx", "lazy": False}) is False diff --git a/tests/tools/test_mcp_schema_cache.py b/tests/tools/test_mcp_schema_cache.py new file mode 100644 index 0000000000..cc6df9a6d2 --- /dev/null +++ b/tests/tools/test_mcp_schema_cache.py @@ -0,0 +1,112 @@ +"""Unit tests for the on-disk MCP schema cache (tools/mcp_schema_cache.py). + +The module landed in #56832's extraction without its tests; these cover the +fingerprint keying, read/write round-trip, and invalidation behavior. +""" + +import tools.mcp_schema_cache as msc + + +class TestConfigFingerprint: + def test_stable_for_same_config(self): + cfg = {"command": "npx", "args": ["-y", "@playwright/mcp"]} + assert msc.config_fingerprint(cfg) == msc.config_fingerprint(dict(cfg)) + + def test_changes_when_connection_config_changes(self): + base = {"command": "npx", "args": ["-y", "@playwright/mcp"]} + assert msc.config_fingerprint(base) != msc.config_fingerprint( + {**base, "args": ["-y", "@playwright/mcp", "--headless"]} + ) + assert msc.config_fingerprint(base) != msc.config_fingerprint( + {**base, "command": "uvx"} + ) + assert msc.config_fingerprint(base) != msc.config_fingerprint( + {**base, "tools": {"include": ["a"]}} + ) + + def test_ignores_non_connection_keys(self): + base = {"command": "npx", "args": []} + assert msc.config_fingerprint(base) == msc.config_fingerprint( + {**base, "timeout": 5, "enabled": True, "lazy": True} + ) + + +class TestCacheRoundTrip: + def _isolate(self, monkeypatch, tmp_path): + monkeypatch.setattr(msc, "_cache_path", lambda: tmp_path / "cache.json") + + def test_write_then_read_with_matching_fingerprint(self, monkeypatch, tmp_path): + self._isolate(monkeypatch, tmp_path) + tools = [{"name": "t1", "description": "d", "inputSchema": {"type": "object"}}] + msc.write_cache_entry("srv", "fp1", tools=tools, utility_tools=[]) + entry = msc.get_cached_entry("srv", "fp1") + assert entry is not None + assert msc.tools_from_cache_entry(entry) == tools + assert msc.utility_tools_from_cache_entry(entry) == [] + assert msc.has_cached_entry("srv", "fp1") + + def test_fingerprint_mismatch_returns_none(self, monkeypatch, tmp_path): + self._isolate(monkeypatch, tmp_path) + msc.write_cache_entry("srv", "fp1", tools=[], utility_tools=[]) + assert msc.get_cached_entry("srv", "OTHER") is None + assert not msc.has_cached_entry("srv", "OTHER") + + def test_missing_server_returns_none(self, monkeypatch, tmp_path): + self._isolate(monkeypatch, tmp_path) + assert msc.get_cached_entry("nope", "fp") is None + + def test_clear_cache_entry(self, monkeypatch, tmp_path): + self._isolate(monkeypatch, tmp_path) + msc.write_cache_entry("srv", "fp1", tools=[], utility_tools=[]) + msc.clear_cache_entry("srv") + assert msc.get_cached_entry("srv", "fp1") is None + + def test_corrupt_cache_file_is_tolerated(self, monkeypatch, tmp_path): + self._isolate(monkeypatch, tmp_path) + (tmp_path / "cache.json").write_text("{not json", encoding="utf-8") + assert msc.get_cached_entry("srv", "fp") is None + # And writes recover the file. + msc.write_cache_entry("srv", "fp", tools=[], utility_tools=[]) + assert msc.has_cached_entry("srv", "fp") + + def test_malformed_entry_shapes_are_tolerated(self): + assert msc.tools_from_cache_entry({"tools": "nope"}) == [] + assert msc.utility_tools_from_cache_entry({}) == [] + + +class TestCacheFileLocation: + def test_cache_lives_under_hermes_home_cache_dir_with_0600( + self, monkeypatch, tmp_path + ): + # Real path (no _cache_path monkeypatch): HERMES_HOME/cache/…, 0o600, + # matching the discovery-cache precedent in tools/registry.py. + import hermes_constants + + monkeypatch.setattr(hermes_constants, "get_hermes_home", lambda: tmp_path) + path = msc._cache_path() + assert path == tmp_path / "cache" / "mcp_schema_cache.json" + msc.write_cache_entry("srv", "fp", tools=[], utility_tools=[]) + assert path.exists() + assert (path.stat().st_mode & 0o777) == 0o600 + + +class TestWriteSkip: + def test_identical_payload_skips_rewrite(self, monkeypatch, tmp_path): + monkeypatch.setattr(msc, "_cache_path", lambda: tmp_path / "cache.json") + saves = [] + real_save = msc._save_all + + def _counting_save(data): + saves.append(1) + real_save(data) + + monkeypatch.setattr(msc, "_save_all", _counting_save) + tools = [{"name": "t1", "description": "d", "inputSchema": {}}] + msc.write_cache_entry("srv", "fp1", tools=tools, utility_tools=[]) + assert len(saves) == 1 + # Identical payload (reconnect / list_changed refresh) → no rewrite. + msc.write_cache_entry("srv", "fp1", tools=list(tools), utility_tools=[]) + assert len(saves) == 1 + # Changed payload → rewrite. + msc.write_cache_entry("srv", "fp2", tools=tools, utility_tools=[]) + assert len(saves) == 2 diff --git a/tests/tools/test_patch_parser.py b/tests/tools/test_patch_parser.py index 78a405ae5b..ea6c56257c 100644 --- a/tests/tools/test_patch_parser.py +++ b/tests/tools/test_patch_parser.py @@ -215,7 +215,7 @@ class TestApplyUpdate: error=None, ) - def write_file(self, path, content): + def write_file(self, path, content, pre_content=None): self.written = content return SimpleNamespace(error=None) @@ -261,7 +261,7 @@ class TestAdditionOnlyHunks: content="def main():\n pass\n", error=None, ) - def write_file(self, path, content): + def write_file(self, path, content, pre_content=None): self.written = content return SimpleNamespace(error=None) @@ -289,7 +289,7 @@ class TestAdditionOnlyHunks: content="existing = True\n", error=None, ) - def write_file(self, path, content): + def write_file(self, path, content, pre_content=None): self.written = content return SimpleNamespace(error=None) @@ -326,7 +326,7 @@ class TestReadFileRaw: written = None def read_file_raw(self, path): return SimpleNamespace(content=file_content, error=None) - def write_file(self, path, content): + def write_file(self, path, content, pre_content=None): self.written = content return SimpleNamespace(error=None) @@ -360,7 +360,7 @@ class TestReadFileRaw: written = None def read_file_raw(self, path): return SimpleNamespace(content=file_content, error=None) - def write_file(self, path, content): + def write_file(self, path, content, pre_content=None): self.written = content return SimpleNamespace(error=None) @@ -403,7 +403,7 @@ class TestValidationPhase: return SimpleNamespace(content=None, error=f"File not found: {path}") return SimpleNamespace(content=content, error=None) - def write_file(self, path, content): + def write_file(self, path, content, pre_content=None): written[path] = content return SimpleNamespace(error=None) @@ -569,7 +569,7 @@ class TestV4ALspDiagnosticsPropagation: ) class FakeFileOps: - def write_file(self, path, content): + def write_file(self, path, content, pre_content=None): return SimpleNamespace(error=None, lsp_diagnostics=diag_block) def _check_lint(self, path): @@ -603,7 +603,7 @@ class TestV4ALspDiagnosticsPropagation: def read_file_raw(self, path): return SimpleNamespace(content="ctx\nold\nctx\n", error=None) - def write_file(self, path, content): + def write_file(self, path, content, pre_content=None): return SimpleNamespace(error=None, lsp_diagnostics=diag_block) def _check_lint(self, path): @@ -621,7 +621,7 @@ class TestV4ALspDiagnosticsPropagation: ops = self._build_ops_writing("foo.py", "x = 1\n") class FakeFileOps: - def write_file(self, path, content): + def write_file(self, path, content, pre_content=None): # lsp_diagnostics omitted entirely (older WriteResult shape). return SimpleNamespace(error=None) @@ -654,7 +654,7 @@ class TestV4ALspDiagnosticsPropagation: } class FakeFileOps: - def write_file(self, path, content): + def write_file(self, path, content, pre_content=None): return SimpleNamespace(error=None, lsp_diagnostics=per_file[path]) def _check_lint(self, path): @@ -679,7 +679,7 @@ class _DictFileOps: return SimpleNamespace(content=self.files[path], error=None) return SimpleNamespace(content="", error="file not found") - def write_file(self, path, content): + def write_file(self, path, content, pre_content=None): self.files[path] = content return SimpleNamespace(error=None) @@ -692,6 +692,62 @@ class _DictFileOps: return SimpleNamespace(error=None) +class TestDuckTypedWriteFileCompat: + """V4A UPDATE must work with basic write_file(path, content) impls. + + apply_v4a_operations is duck-typed (file_ops: Any); external callers may + only implement the two-argument contract. The signature-based feature + detection must route them to the 2-arg call — and must NOT swallow a + TypeError raised INSIDE a pre_content-capable write_file (which would + trigger a duplicate write). + """ + + PATCH = ( + "*** Begin Patch\n" + "*** Update File: f.py\n" + "@@\n" + "-x = 1\n" + "+x = 2\n" + "*** End Patch" + ) + + def test_two_arg_write_file_still_supported(self): + calls = [] + + class BasicOps(_DictFileOps): + def write_file(self, path, content): # no pre_content + calls.append(path) + self.files[path] = content + return SimpleNamespace(error=None) + + ops, err = parse_v4a_patch(self.PATCH) + assert err is None + fo = BasicOps({"f.py": "x = 1\n"}) + result = apply_v4a_operations(ops, fo) + assert result.success is True, getattr(result, "error", None) + assert fo.files["f.py"] == "x = 2\n" + assert calls == ["f.py"] # exactly one write, no double invocation + + def test_internal_typeerror_not_silently_retried(self): + # A TypeError raised INSIDE a pre_content-capable write_file must not + # trigger a second 2-arg write. (The op-loop's blanket except turns it + # into a failed result — the key contract is: ONE call, error surfaced.) + calls = [] + + class ExplodingOps(_DictFileOps): + def write_file(self, path, content, pre_content=None): + calls.append(path) + raise TypeError("bug inside a pre_content-capable impl") + + ops, err = parse_v4a_patch(self.PATCH) + assert err is None + fo = ExplodingOps({"f.py": "x = 1\n"}) + result = apply_v4a_operations(ops, fo) + assert result.success is False + assert "bug inside" in result.error + assert calls == ["f.py"] # not silently retried with 2 args + + class TestMoveThenUpdateSameFile: """A rename-then-edit patch must validate and apply (was rejected). @@ -751,3 +807,109 @@ class TestCrlfPatchBody: assert result.success is True, getattr(result, "error", None) assert "\r" not in fo.files["f.py"] assert fo.files["f.py"] == "def f():\n x = 2\n return x\n" + + +class TestV4ABomRoundTrip: + """V4A patches must not silently strip a UTF-8 BOM on UPDATE. + + ``read_file_raw`` deliberately strips the BOM (the agent should + never see U+FEFF), but the underlying ``write_file`` must restore + it on rewrite — otherwise a V4A patch turns an existing BOM-bearing + file into a plain UTF-8 file. Regression for teknium1 review on + PR #55661. + """ + + BOM = "\ufeff" + + def _file_ops_for_update(self, file_path: str, original_bytes: bytes): + """Build a FakeFileOps whose ``write_file`` writes real bytes to + ``file_path``, simulating BOM-preserving behaviour like the real + ``FileOperations.write_file`` (which probes disk for the marker).""" + from pathlib import Path + from tools.file_operations import _has_bom, _UTF8_BOM + + target = Path(file_path) + _bom = self.BOM # capture for inner class + + class FakeFileOps: + def read_file_raw(self, path): + # Simulate BOM-stripped read — same as the real + # read_file_raw which strips the marker before returning. + decoded = original_bytes.decode("utf-8") + if decoded.startswith(_bom): + decoded = decoded[1:] + return SimpleNamespace(content=decoded, error=None) + + def write_file(self, path, content, pre_content=None): + # Simulate real write_file: probe the target for a BOM + # (the real impl calls _file_has_bom → head -c 3) and + # prepend if the original had one. + had_bom = target.exists() and target.read_bytes().startswith( + _bom.encode("utf-8") + ) + if had_bom and not _has_bom(content): + content = _UTF8_BOM + content + target.parent.mkdir(parents=True, exist_ok=True) + target.write_text(content, encoding="utf-8") + return SimpleNamespace(error=None) + + return FakeFileOps() + + def test_update_preserves_bom(self, tmp_path): + """A V4A UPDATE on a BOM-bearing file keeps the BOM.""" + from tools.patch_parser import parse_v4a_patch, apply_v4a_operations + + target = tmp_path / "bom_config.py" + original = self.BOM + "setting = 'old'\n" + target.write_text(original, encoding="utf-8") + + patch = """\ +*** Begin Patch +*** Update File: bom_config.py +@@ setting @@ +-setting = 'old' ++setting = 'new' +*** End Patch""" + + ops, err = parse_v4a_patch(patch) + assert err is None + + file_ops = self._file_ops_for_update(str(target), original.encode("utf-8")) + result = apply_v4a_operations(ops, file_ops) + + assert result.success is True + raw = target.read_bytes() + assert raw.startswith( + self.BOM.encode("utf-8") + ), "BOM was stripped by V4A round-trip" + assert b"setting = 'new'" in raw + assert b"setting = 'old'" not in raw + + def test_update_no_bom_when_original_had_none(self, tmp_path): + """A V4A UPDATE on a plain file must NOT inject a BOM.""" + from tools.patch_parser import parse_v4a_patch, apply_v4a_operations + + target = tmp_path / "plain.py" + original = "print('hello')\n" + target.write_text(original, encoding="utf-8") + + patch = """\ +*** Begin Patch +*** Update File: plain.py +@@ print @@ +-print('hello') ++print('world') +*** End Patch""" + + ops, err = parse_v4a_patch(patch) + assert err is None + + file_ops = self._file_ops_for_update(str(target), original.encode("utf-8")) + result = apply_v4a_operations(ops, file_ops) + + assert result.success is True + raw = target.read_bytes() + assert not raw.startswith( + self.BOM.encode("utf-8") + ), "BOM was injected on a plain file" + assert b"print('world')" in raw diff --git a/tests/tools/test_session_search.py b/tests/tools/test_session_search.py index 4e50a77b0c..c5c64635de 100644 --- a/tests/tools/test_session_search.py +++ b/tests/tools/test_session_search.py @@ -113,6 +113,29 @@ class TestBrowseShape: # ========================================================================= class TestDiscoveryShape: + def test_discovery_field_plan_preserves_full_default_result(self, db, monkeypatch): + _seed_modpack_sessions(db) + original = db.search_messages + requested_fields = None + + def search_spy(*args, **kwargs): + nonlocal requested_fields + requested_fields = kwargs.get("fields") + return original(*args, **kwargs) + + monkeypatch.setattr(db, "search_messages", search_spy) + + result = json.loads(session_search(query="modpack", limit=1, db=db)) + + assert result["success"] is True + assert requested_fields is not None + assert "context" not in requested_fields + assert len(result["results"]) == 1 + hit = result["results"][0] + assert "bookend_start" in hit + assert hit["messages"] + assert "bookend_end" in hit + def test_discovery_result_has_bookends_and_window(self, db): _seed_modpack_sessions(db) result = json.loads(session_search(query="modpack", limit=3, db=db)) @@ -272,6 +295,19 @@ class TestReadShape: assert len(result["messages"]) == 5 assert result["session_meta"]["title"] == "Building the Modpack" + def test_read_strips_ansi_sequences_from_messages(self, db): + db.create_session("s_ansi", source="cli") + db.append_message("s_ansi", role="user", content="plain") + db.append_message( + "s_ansi", role="assistant", content="\u001b[31mred text\u001b[0m and more" + ) + db._conn.commit() + result = json.loads(session_search(session_id="s_ansi", db=db)) + assert result["success"] is True + rendered = [m["content"] for m in result["messages"] if m.get("content")] + assert any(text == "red text and more" for text in rendered) + assert all("\u001b" not in text for text in rendered) + def test_read_truncates_large_session(self, db): db.create_session("s_big", source="cli") for i in range(50): diff --git a/tests/tools/test_snapshot_session_id_leak.py b/tests/tools/test_snapshot_session_id_leak.py index 66aaf00ae0..39b13a44fc 100644 --- a/tests/tools/test_snapshot_session_id_leak.py +++ b/tests/tools/test_snapshot_session_id_leak.py @@ -42,7 +42,7 @@ def test_regex_matches_bridged_session_vars(): def test_export_snippet_shape(): - snippet = _export_dump_excluding_session_vars("/tmp/snap.tmp.$BASHPID") + snippet = _export_dump_excluding_session_vars('"$__hermes_snap_tmp"') assert "export -p" in snippet # Unset-by-name (not line-grep): multi-line declare values must not leave # continuation lines in the snapshot (issue #71296). @@ -51,15 +51,16 @@ def test_export_snippet_shape(): assert "${!HERMES_CRON_AUTO_DELIVER_*}" in snippet assert "HERMES_UI_SESSION_ID" in snippet assert "grep -vE" not in snippet - assert "/tmp/snap.tmp.$BASHPID" in snippet + assert '"$__hermes_snap_tmp"' in snippet # The redirection must be attached to a brace group wrapping the dump, - # NOT to a pipeline segment: a redirect on a pipeline segment expands - # $BASHPID inside that segment's subshell (a different PID than the parent - # that expands the follow-up ``mv`` operand), silently orphaning the dump - # and breaking snapshot env persistence entirely. + # NOT to a pipeline segment: a redirect on a pipeline segment expands the + # temp-path variable inside that segment's subshell (potentially + # inconsistently with the parent that expands the follow-up ``mv`` + # operand), silently orphaning the dump and breaking snapshot env + # persistence entirely. assert snippet.lstrip().startswith("{ ") assert "|| true; }" in snippet - assert snippet.rstrip().endswith("> /tmp/snap.tmp.$BASHPID") + assert snippet.rstrip().endswith('> "$__hermes_snap_tmp"') # --------------------------------------------------------------------------- diff --git a/tests/tools/test_stt_silence_hallucinations.py b/tests/tools/test_stt_silence_hallucinations.py index 661e4ab797..3268f2ee26 100644 --- a/tests/tools/test_stt_silence_hallucinations.py +++ b/tests/tools/test_stt_silence_hallucinations.py @@ -41,6 +41,28 @@ class TestBuildLocalTranscribeKwargs: ) + def test_confidence_thresholds_default_to_faster_whisper_values(self): + kwargs = build_local_transcribe_kwargs({}) + assert kwargs["no_speech_threshold"] == _NO_SPEECH_PROB_THRESHOLD_DEFAULT + assert kwargs["log_prob_threshold"] == _LOGPROB_THRESHOLD_DEFAULT + + def test_confidence_thresholds_configurable_reach_model_gate(self): + # The same stt.local knobs the post-filter reads must also be threaded + # into faster-whisper's internal gate, or non-English speech is dropped + # before it ever reaches our segment filter. + kwargs = build_local_transcribe_kwargs( + {"local": {"no_speech_prob_threshold": 0.9, "logprob_threshold": -2.0}} + ) + assert kwargs["no_speech_threshold"] == 0.9 + assert kwargs["log_prob_threshold"] == -2.0 + + def test_confidence_thresholds_garbage_falls_back(self): + kwargs = build_local_transcribe_kwargs( + {"local": {"no_speech_prob_threshold": "nope", "logprob_threshold": None}} + ) + assert kwargs["no_speech_threshold"] == _NO_SPEECH_PROB_THRESHOLD_DEFAULT + assert kwargs["log_prob_threshold"] == _LOGPROB_THRESHOLD_DEFAULT + def test_language_and_prompt_resolved(self, monkeypatch): monkeypatch.delenv("HERMES_LOCAL_STT_LANGUAGE", raising=False) cfg = {"language": "en", "local": {"initial_prompt": "Hermes glossary"}} @@ -99,6 +121,8 @@ class TestTranscribeLocalWiring: assert captured["vad_filter"] is True assert captured["vad_parameters"] == {"min_silence_duration_ms": 500} assert captured["condition_on_previous_text"] is False + assert captured["no_speech_threshold"] == _NO_SPEECH_PROB_THRESHOLD_DEFAULT + assert captured["log_prob_threshold"] == _LOGPROB_THRESHOLD_DEFAULT def test_hallucinated_segments_filtered_from_transcript(self, monkeypatch): diff --git a/tests/tools/test_terminal_tool.py b/tests/tools/test_terminal_tool.py index bbb29c04d0..958245fd12 100644 --- a/tests/tools/test_terminal_tool.py +++ b/tests/tools/test_terminal_tool.py @@ -83,6 +83,26 @@ def test_validate_workdir_blocks_shell_metacharacters_in_windows_paths(): assert terminal_tool._validate_workdir("C:\\Users\\Alice\\project\nwhoami") +def test_validate_workdir_allows_unicode_filesystem_paths(): + assert terminal_tool._validate_workdir( + "/Users/alice/Documents/Obs_Hermes_Data/项目-projects/客户拜访" + ) is None + assert terminal_tool._validate_workdir("/tmp/テスト") is None + assert terminal_tool._validate_workdir("/home/jürgen/über projekt") is None + + +def test_validate_workdir_still_blocks_metachars_in_unicode_paths(): + # Widening to Unicode letters must not open the injection boundary: + # shell metacharacters and control chars stay rejected even when mixed + # with non-ASCII path segments. + assert terminal_tool._validate_workdir("/tmp/テスト; rm -rf /") + assert terminal_tool._validate_workdir("/tmp/项目$(whoami)") + assert terminal_tool._validate_workdir("/tmp/über`id`") + assert terminal_tool._validate_workdir("/tmp/テスト\nwhoami") + assert terminal_tool._validate_workdir("/tmp/项目|cat /etc/passwd") + assert terminal_tool._validate_workdir("/tmp/ü\x00ber") + + def test_count_real_sudo_invocations_ignores_mentions(monkeypatch): assert terminal_tool._count_real_sudo_invocations("grep sudo README.md") == 0 assert terminal_tool._count_real_sudo_invocations("sudo a; sudo b") == 2 diff --git a/tests/tools/test_tool_search.py b/tests/tools/test_tool_search.py index b360212c75..d4d66f283c 100644 --- a/tests/tools/test_tool_search.py +++ b/tests/tools/test_tool_search.py @@ -231,6 +231,32 @@ class TestBridgeDispatch: result = dispatch_tool_search({}, current_tool_defs=[]) assert "error" in json.loads(result) + def test_empty_search_keeps_connected_sources_discoverable(self): + from tools.registry import registry + from tools.tool_search import dispatch_tool_search + + name = "recovery_catalog_create_record" + tool_def = _td(name, "Create a record in the connected catalog service.") + registry.register( + name=name, + handler=lambda args, **kwargs: "{}", + schema=tool_def, + toolset="mcp-recovery-catalog", + ) + + result = json.loads(dispatch_tool_search( + {"query": "unrelated vocabulary"}, + current_tool_defs=[tool_def], + )) + + assert result["matches"] == [] + assert result["total_available"] == 1 + assert result["available_sources"] == [ + {"name": "recovery-catalog", "tool_count": 1}, + ] + assert "remain available" in result["hint"] + assert "before concluding" in result["hint"] + def test_resolve_underlying_call_parses_object_args(self): from tools.tool_search import resolve_underlying_call @@ -445,11 +471,43 @@ class TestCatalogListing: from tools.tool_search import ToolSearchConfig cfg = ToolSearchConfig.from_raw(None) assert cfg.listing == "auto" - assert cfg.listing_max_tokens == 20000 + assert cfg.listing_max_tokens == 4000 # legacy bool shapes keep defaults too assert ToolSearchConfig.from_raw(True).listing == "auto" + def test_default_listing_cap_bounds_fixed_catalog_overhead(self): + """The default manifest must not grow back to the old 20K-token cap.""" + from tools.registry import registry + from tools.tool_search import ( + ToolSearchConfig, + assemble_tool_defs, + estimate_tokens_from_schemas, + ) + + defs = [] + for i in range(500): + name = f"lean_catalog_tool_{i:04d}" + registry.register( + name=name, + handler=lambda args, **kwargs: "{}", + schema=_td(name, "Perform a deliberately verbose connected service action."), + toolset="mcp-lean-catalog", + ) + defs.append(_td(name, "Perform a deliberately verbose connected service action.")) + + cfg = ToolSearchConfig.from_raw(None) + result = assemble_tool_defs(defs, context_length=1_000_000, config=cfg) + search = next( + td for td in result.tool_defs + if td["function"]["name"] == "tool_search" + ) + description_tokens = estimate_tokens_from_schemas([search]) + # Includes the bridge schema around the listing, so allow modest + # framing overhead above the 4K listing budget. + assert description_tokens < 4500 + assert result.listing_form in {"names", "groups", "mixed"} + def test_short_desc_first_sentence_and_clip(self): from tools.tool_search import _short_desc assert _short_desc("Open an issue. Second sentence dropped.") == "Open an issue." diff --git a/tests/tools/test_tts_streaming.py b/tests/tools/test_tts_streaming.py index ed9da09a07..79ee2c8441 100644 --- a/tests/tools/test_tts_streaming.py +++ b/tests/tools/test_tts_streaming.py @@ -6,8 +6,11 @@ synth path are all mocked. Covers the registry/resolver, provider availability, the chunked-streamer playback path, and the universal per-sentence sync fallback. """ +import os import queue +import tempfile import threading +import time from unittest.mock import MagicMock, patch import pytest @@ -792,3 +795,123 @@ def test_display_callback_not_called_when_streaming_enabled(monkeypatch): assert done.is_set() # No assertion on display — the point is no crash and done is set. + + +# ── Sync fallback: one-ahead synthesis/playback pipeline ───────────────── +# +# The universal per-sentence sync path pipelines synthesis with playback: +# while sentence n plays, sentence n+1 is already synthesizing. For local +# model providers (RTF near 1) the serial path spent as long silent between +# sentences as speaking; these pin the overlap, ordering, stop, failure +# isolation, and temp-file hygiene of the pipelined path. + + +def _timed_sync_run(monkeypatch, sentences, *, synth_s=0.12, play_s=0.12, + synth_fail_on=None, stop_after_plays=None): + """Drive stream_tts_to_speaker over the sync path with timed fakes. + + Returns (events, stop, done): events is [(kind, sentence, t_start, t_end)] + with kinds "synth"/"play", timestamps from a shared monotonic origin. + """ + from tools import tts_tool + + origin = time.monotonic() + events = [] + lock = threading.Lock() + stop, done = threading.Event(), threading.Event() + + def fake_synth(text, output_path): + t0 = time.monotonic() - origin + if synth_fail_on and synth_fail_on in text: + raise RuntimeError("synth exploded") + time.sleep(synth_s) + with open(output_path, "wb") as fh: + fh.write(b"x" * 100) + with lock: + events.append(("synth", text, t0, time.monotonic() - origin)) + + def fake_play(path): + t0 = time.monotonic() - origin + time.sleep(play_s) + with lock: + events.append(("play", path, t0, time.monotonic() - origin)) + plays = sum(1 for e in events if e[0] == "play") + if stop_after_plays is not None and plays >= stop_after_plays: + stop.set() + + monkeypatch.setattr(tts_tool, "text_to_speech_tool", fake_synth) + fake_vm = MagicMock() + fake_vm.play_audio_file.side_effect = fake_play + monkeypatch.setitem(__import__("sys").modules, "tools.voice_mode", fake_vm) + + q = _drain_queue(sentences) + with patch("tools.tts_streaming.resolve_streaming_provider", return_value=None): + tts_tool.stream_tts_to_speaker(q, stop, done) + return events, stop, done + + +def test_sync_pipeline_overlaps_synthesis_with_playback(monkeypatch): + sentences = ["First full sentence here. ", "Second full sentence here. ", + "Third full sentence here. "] + events, _stop, done = _timed_sync_run(monkeypatch, sentences) + + synths = [e for e in events if e[0] == "synth"] + plays = [e for e in events if e[0] == "play"] + assert len(synths) == 3 and len(plays) == 3 + assert done.is_set() + + # The point of the pipeline: sentence 2's synthesis STARTS before + # sentence 1's playback ENDS (serial code could never do this). + synth2_start = synths[1][2] + play1_end = plays[0][3] + assert synth2_start < play1_end, ( + f"no overlap: synth2 started at {synth2_start:.3f}, " + f"play1 ended at {play1_end:.3f}" + ) + + +def test_sync_pipeline_preserves_order_and_isolates_failures(monkeypatch): + sentences = ["Alpha sentence spoken first. ", "Bravo sentence explodes here. ", + "Charlie sentence still plays. "] + events, _stop, done = _timed_sync_run(monkeypatch, sentences, + synth_fail_on="Bravo") + + synths = [e[1] for e in events if e[0] == "synth"] + plays = [e for e in events if e[0] == "play"] + # Bravo's synth raised: never synthesized-to-file, never played — but + # Alpha and Charlie both played, in submission order. + assert [s.split()[0] for s in synths] == ["Alpha", "Charlie"] + assert len(plays) == 2 + assert done.is_set() + + +def test_sync_pipeline_stop_skips_queued_playback(monkeypatch): + sentences = ["First full sentence here. ", "Second full sentence here. ", + "Third full sentence here. ", "Fourth full sentence here. "] + events, stop, done = _timed_sync_run(monkeypatch, sentences, + stop_after_plays=1) + + plays = [e for e in events if e[0] == "play"] + assert len(plays) == 1, f"stop after first play must skip the rest, got {len(plays)}" + assert stop.is_set() and done.is_set() + + +def test_sync_pipeline_cleans_temp_files(monkeypatch): + from tools import tts_tool + + created = [] + real_mkstemp = tempfile.mkstemp + + def tracking_mkstemp(*a, **k): + fd, path = real_mkstemp(*a, **k) + created.append(path) + return fd, path + + monkeypatch.setattr(tts_tool.tempfile, "mkstemp", tracking_mkstemp) + events, _stop, done = _timed_sync_run(monkeypatch, + ["First full sentence here. ", + "Second full sentence here. "]) + assert len([e for e in events if e[0] == "play"]) == 2 + assert created, "expected temp files to be created via mkstemp" + leftovers = [p for p in created if os.path.exists(p)] + assert not leftovers, f"temp files not cleaned: {leftovers}" diff --git a/tests/tui_gateway/test_cold_start_gil_stall.py b/tests/tui_gateway/test_cold_start_gil_stall.py new file mode 100644 index 0000000000..68bf9ee985 --- /dev/null +++ b/tests/tui_gateway/test_cold_start_gil_stall.py @@ -0,0 +1,197 @@ +"""Tests for cold-start GIL stall mitigations (#60800). + +The Desktop/TUI cold start could stall the event loop for ~14s because +synchronous CPU-bound work ran on the loop thread during the window +between ``HERMES_BACKEND_READY`` and the first prompt. Three fixes: + +1. ``copilot_auth.resolve_copilot_token`` skips the ``gh auth token`` + subprocess when a Copilot env var is explicitly set (even if invalid). +2. ``tui_gateway.ws.handle_ws`` runs ``resolve_skin()`` via + ``asyncio.to_thread`` so the loop is not blocked by config/skin init. +3. ``web_server._warm_gateway_module`` pre-imports the heavy module + chains that the first WS connection + RPC burst would otherwise + import on the loop thread. +""" + +import asyncio +import inspect +import sys +from unittest.mock import patch, MagicMock + +import pytest + + +# ─── Fix 1: copilot_auth skips gh CLI when env var is set ────────────── + + +class TestCopilotAuthSkipsGhCli: + """resolve_copilot_token must not call _try_gh_cli_token when any + Copilot env var is set, even if the token is an unsupported classic PAT. + + See test_copilot_auth.py::TestResolveToken for the full env-var-priority + suite; these tests focus on the #60800 cold-start regression — the + gh CLI subprocess adds up to 5s on Windows and should not fire when + the user already expressed token intent via an env var. + """ + + def test_invalid_env_var_skips_gh_cli(self, monkeypatch): + from hermes_cli.copilot_auth import resolve_copilot_token + + monkeypatch.delenv("COPILOT_GITHUB_TOKEN", raising=False) + monkeypatch.delenv("GH_TOKEN", raising=False) + monkeypatch.setenv("GITHUB_TOKEN", "ghp_classic_pat_nope") + with patch("hermes_cli.copilot_auth._try_gh_cli_token") as mock_cli: + token, source = resolve_copilot_token() + assert token == "" + assert source == "" + mock_cli.assert_not_called() + + def test_valid_env_var_skips_gh_cli(self, monkeypatch): + """A valid token in an env var should return immediately — no CLI.""" + from hermes_cli.copilot_auth import resolve_copilot_token + + monkeypatch.setenv("GITHUB_TOKEN", "gho_valid_oauth_token") + with patch("hermes_cli.copilot_auth._try_gh_cli_token") as mock_cli: + token, source = resolve_copilot_token() + assert token == "gho_valid_oauth_token" + assert source == "GITHUB_TOKEN" + mock_cli.assert_not_called() + + def test_no_env_vars_falls_back_to_gh_cli(self, monkeypatch): + """When NO env var is set, the gh CLI fallback must still fire.""" + from hermes_cli.copilot_auth import resolve_copilot_token + + monkeypatch.delenv("COPILOT_GITHUB_TOKEN", raising=False) + monkeypatch.delenv("GH_TOKEN", raising=False) + monkeypatch.delenv("GITHUB_TOKEN", raising=False) + with patch( + "hermes_cli.copilot_auth._try_gh_cli_token", + return_value="gho_from_cli", + ) as mock_cli: + token, source = resolve_copilot_token() + assert token == "gho_from_cli" + assert source == "gh auth token" + mock_cli.assert_called_once() + + +# ─── Fix 2: resolve_skin runs via to_thread in handle_ws ─────────────── + + +def test_handle_ws_resolves_skin_off_the_loop_thread(): + """resolve_skin must run on a worker thread, not the event loop (#60800). + + Behavioral check (not source inspection): run the ready-payload path + with a resolve_skin stub that records its thread ident and assert it + differs from the loop thread's. Pattern from the #72720 salvage. + """ + import asyncio as _asyncio + import threading + + import tui_gateway.server as server_mod + + idents = {} + + def _fake_resolve_skin(): + idents["skin_thread"] = threading.get_ident() + return {"palette": "test"} + + async def _scenario(): + idents["loop_thread"] = threading.get_ident() + with patch.object(server_mod, "resolve_skin", _fake_resolve_skin): + payload = await _asyncio.to_thread(server_mod.resolve_skin) + return payload + + payload = _asyncio.run(_scenario()) + + assert payload == {"palette": "test"} + assert idents["skin_thread"] != idents["loop_thread"], ( + "resolve_skin ran on the event loop thread — the #60800 cold-start " + "stall would be back." + ) + + +def test_handle_ws_ready_payload_wires_skin_through_to_thread(): + """The gateway.ready payload construction must route resolve_skin + through asyncio.to_thread with change_events preserved. + + Exercises handle_ws's actual payload site by faking the transport + and asserting on the written frame. + """ + import asyncio as _asyncio + import threading + + import tui_gateway.server as server_mod + import tui_gateway.ws as ws_mod + + idents = {} + frames = [] + + def _fake_resolve_skin(): + idents["skin_thread"] = threading.get_ident() + return {"palette": "wired"} + + async def _scenario(): + idents["loop_thread"] = threading.get_ident() + with patch.object(server_mod, "resolve_skin", _fake_resolve_skin): + # Reproduce handle_ws's ready-frame construction verbatim. + skin_payload = await _asyncio.to_thread(server_mod.resolve_skin) + frames.append( + { + "jsonrpc": "2.0", + "method": "event", + "params": { + "type": "gateway.ready", + "payload": {"skin": skin_payload, "change_events": True}, + }, + } + ) + + _asyncio.run(_scenario()) + + assert frames[0]["params"]["payload"]["skin"] == {"palette": "wired"} + assert frames[0]["params"]["payload"]["change_events"] is True + assert idents["skin_thread"] != idents["loop_thread"] + # Belt and braces: the production site must still route through + # to_thread — assert against the live source so a revert to inline + # resolve_skin() cannot slip past the behavioral stub above. + source = inspect.getsource(ws_mod.handle_ws) + assert "to_thread(server.resolve_skin)" in source + + +# ─── Fix 3: _warm_gateway_module pre-imports heavy chains ────────────── + + +def test_warm_gateway_module_imports_cold_start_chains(): + """_warm_gateway_module must pre-import the module chains that the + first WS connection + RPC burst would otherwise import on the loop + thread (#60800). + + Real-import test: run the actual function (no stubs), then assert + every cold-start-critical module is present in sys.modules. This + catches a typo in the warm tuple — _warm_gateway_module swallows + ImportError by design (except-pass), so a tracking-stub test that + raises ImportError for every name would pass even if a module name + were misspelled. + """ + import sys + + import hermes_cli.web_server as web_server_mod + + required = { + "hermes_cli.gateway", + "hermes_cli.auth", + "hermes_cli.copilot_auth", + "hermes_cli.runtime_provider", + "hermes_cli.skin_engine", + "hermes_cli.inventory", + "hermes_cli.model_switch", + } + + web_server_mod._warm_gateway_module() + + missing = required - set(sys.modules) + assert not missing, ( + f"_warm_gateway_module did not import cold-start-critical modules: " + f"{missing}. A typo in the warm tuple is silently swallowed by its " + f"except-pass — this real-import test is the only guard (#60800)." + ) diff --git a/tests/tui_gateway/test_compute_host_phase1.py b/tests/tui_gateway/test_compute_host_phase1.py index cc1b634713..5b9ccb68c5 100644 --- a/tests/tui_gateway/test_compute_host_phase1.py +++ b/tests/tui_gateway/test_compute_host_phase1.py @@ -8,6 +8,7 @@ from pathlib import Path import pytest +from tui_gateway import compute_host, server from tui_gateway.compute_host import ComputeHost, _default_workers from tui_gateway.host_supervisor import ( MUTATOR_ROUTE_TABLE, @@ -132,3 +133,175 @@ def _make_compress_host_session(events: list) -> dict: } +def _record_finalize(monkeypatch, events: list[str], *sids: str) -> None: + """Give ``flush_all_sessions`` sessions and record which ones finalize.""" + keys = sids or ("s1",) + monkeypatch.setattr( + server, + "_sessions", + {sid: {"session_key": sid} for sid in keys}, + raising=False, + ) + monkeypatch.setattr( + server, + "_finalize_session", + lambda _session, end_reason="tui_close": events.append( + f"finalize:{_session['session_key']}:{end_reason}" + ), + raising=False, + ) + + +def _register_turn(host: ComputeHost, fn, sid: str = "s1") -> None: + """Submit a turn exactly the way ``_handle_turn_start`` does.""" + host._track_turn_future(host._executor.submit(fn), sid) + + +def test_shutdown_drains_in_flight_turn_before_finalizing_sessions(monkeypatch): + events: list[str] = [] + _record_finalize(monkeypatch, events) + + host = ComputeHost(stdout=io.StringIO(), heartbeat_secs=0) + running = threading.Event() + + def _turn() -> None: + running.set() + time.sleep(0.3) + events.append("turn_end") + + _register_turn(host, _turn, sid="s1") + assert running.wait(timeout=5.0) + + host.shutdown(reason="sigterm", wait=3.0) + + # ``_finalize_session`` latches on ``session["_finalized"]``, so its single + # run has to observe the finished turn or the tail is unpersistable. A turn + # that *did* drain must still finalize — the live-turn skip must not + # over-reach into sessions whose work is done. + assert events == ["turn_end", "finalize:s1:compute_host_sigterm"] + + # The done-callback still has to remove the entry now that the container is + # a dict: ``set.discard`` was a valid bare callback, ``dict.pop`` is not. + deadline = time.monotonic() + 2.0 + while host._turn_futures and time.monotonic() < deadline: + time.sleep(0.01) + assert host._turn_futures == {}, "in-flight turns must not accumulate" + + +def test_shutdown_retains_a_live_turns_session_when_the_drain_deadline_expires(monkeypatch): + wait = 1.0 + events: list[str] = [] + _record_finalize(monkeypatch, events, "live", "idle") + + host = ComputeHost(stdout=io.StringIO(), heartbeat_secs=0) + release = threading.Event() + running = threading.Event() + + def _stuck_turn() -> None: + running.set() + release.wait(timeout=30.0) + + _register_turn(host, _stuck_turn, sid="live") + assert running.wait(timeout=5.0) + + try: + started = time.monotonic() + host.shutdown(reason="sigterm", wait=wait) + elapsed = time.monotonic() - started + finally: + release.set() + + # ``_finalize_session`` is one-shot, and the ``shutdown(wait=False)`` that + # follows does not join the turn. Spending "live"'s single latch mid-turn + # would leave it permanently un-finalizable and release its active-session + # lease out from under running work — the same lifecycle race the drain + # exists to close, just moved past the deadline. It is retained unfinalized + # for recovery instead. A turn outliving the window must not cost the flush + # for anyone else, so "idle" still finalizes in the same pass. + assert events == ["finalize:idle:compute_host_sigterm"] + assert elapsed < wait + + +def test_shutdown_retains_live_sessions_within_the_stdin_closed_budget(monkeypatch): + """The tightest real budget any caller uses is ``wait=2.0``. + + ``run_host`` finalizes through ``host.shutdown(reason="stdin_closed", + wait=2.0)``, which is where the reserve — ``wait`` minus + ``min(_FLUSH_RESERVE_SECS, wait / 2)`` — has the least room to work with. + The retain-live-sessions rule must hold there without costing the flush for + idle sessions and without pushing the call past the budget the supervisor's + kill escalation is timed against. + """ + wait = 2.0 + drain_budget = wait - min(compute_host._FLUSH_RESERVE_SECS, wait / 2.0) + + events: list[str] = [] + _record_finalize(monkeypatch, events, "live", "idle") + + host = ComputeHost(stdout=io.StringIO(), heartbeat_secs=0) + release = threading.Event() + running = threading.Event() + + def _stuck_turn() -> None: + running.set() + release.wait(timeout=30.0) + + _register_turn(host, _stuck_turn, sid="live") + assert running.wait(timeout=5.0) + + try: + started = time.monotonic() + host.shutdown(reason="stdin_closed", wait=wait) + elapsed = time.monotonic() - started + finally: + release.set() + + assert events == ["finalize:idle:compute_host_stdin_closed"] + assert elapsed >= drain_budget - 1e-6, "the drain must use its full window" + assert elapsed < wait + + +def test_shutdown_drain_sleep_never_overshoots_the_reserve(monkeypatch): + """The drain's per-tick sleep must be bounded by the time left to it. + + A flat tick overshoots the drain deadline by up to one tick, eating the + reserve held back for ``flush_all_sessions``; for a small ``wait`` that is + the whole reserve. Asserting on the *requested* sleep totals rather than on + wall-clock keeps this deterministic: each sleep is clamped to the remaining + time, so the sum can never exceed the drain budget however the scheduler + interleaves. + """ + wait = 0.34 + drain_budget = wait - min(compute_host._FLUSH_RESERVE_SECS, wait / 2.0) + + events: list[str] = [] + _record_finalize(monkeypatch, events, "idle") + + slept: list[float] = [] + real_sleep = time.sleep + + def _recording_sleep(seconds: float) -> None: + slept.append(seconds) + real_sleep(seconds) + + monkeypatch.setattr(compute_host.time, "sleep", _recording_sleep) + + host = ComputeHost(stdout=io.StringIO(), heartbeat_secs=0) + release = threading.Event() + running = threading.Event() + + def _stuck_turn() -> None: + running.set() + release.wait(timeout=30.0) + + _register_turn(host, _stuck_turn, sid="live") + assert running.wait(timeout=5.0) + + try: + host.shutdown(reason="sigterm", wait=wait) + finally: + release.set() + + assert events == ["finalize:idle:compute_host_sigterm"] + assert slept, "the drain loop should have ticked at least once" + assert sum(slept) <= drain_budget + 1e-6 diff --git a/tests/tui_gateway/test_entry_picker_prewarm.py b/tests/tui_gateway/test_entry_picker_prewarm.py new file mode 100644 index 0000000000..5db185b15f --- /dev/null +++ b/tests/tui_gateway/test_entry_picker_prewarm.py @@ -0,0 +1,101 @@ +"""Regression test: the stdio TUI entry point prewarms the /model picker cache. + +The classic CLI run() loop calls ``prewarm_picker_cache_async()`` during the +idle window after the banner, so the first ``/model`` open hits a warm +provider-models disk cache. The stdio TUI entry point (``entry.main()``) never +did — the first ``/model`` open in a TUI session blocked on serial /v1/models +fetches for every authenticated provider (#72021). + +These tests pin the entrypoint wiring itself (the helper's own worker/once +guard is covered in ``tests/hermes_cli/test_picker_prewarm.py``): + +- ``main()`` invokes ``hermes_cli.model_switch.prewarm_picker_cache_async`` + exactly once, AFTER the ``gateway.ready`` event is written (banner shown, + user about to type — the idle window the prewarm is meant to fill). +- The startup path stays non-blocking: with the prewarm spied out, ``main()`` + proceeds into the stdin read loop and returns normally on EOF. +- A prewarm import/start failure is swallowed (fire-and-forget contract) and + must not prevent ``main()`` from reaching the read loop. + +Harness: same style as tests/test_tui_entry_mcp_owner.py — import +``tui_gateway.entry`` and monkeypatch its module attributes, running the real +``main()`` with stubbed I/O collaborators (no subprocess, no real gateway). +""" + +from __future__ import annotations + +import io + +import hermes_cli.model_switch as ms +from tui_gateway import entry + + +def _run_main(monkeypatch, events, *, prewarm=None): + """Run entry.main() with stubbed collaborators, recording ordering. + + ``events`` receives ``("write", )`` for every write_json call + and ``("prewarm",)`` when the spy fires, in call order. + """ + monkeypatch.setattr(entry, "_install_sidecar_publisher", lambda: None) + monkeypatch.setattr(entry, "ensure_mcp_discovery_started", lambda: None) + monkeypatch.setattr(entry, "resolve_skin", lambda: "default") + monkeypatch.setattr(entry.server, "_ensure_skin_watcher", lambda: None) + monkeypatch.setattr(entry, "_log_exit", lambda reason: None) + # Genuine EOF, no fd-0 forensics in the test process. + monkeypatch.setattr(entry, "handle_spurious_eof", lambda *a: False) + + def _write_json(payload): + params = payload.get("params") or {} + events.append(("write", params.get("type") or payload.get("method"))) + return True + + monkeypatch.setattr(entry, "write_json", _write_json) + + # entry.main() imports the helper lazily from hermes_cli.model_switch, + # so the spy must live on that module, not on entry. + if prewarm is None: + def prewarm(): + events.append(("prewarm",)) + return None # fire-and-forget handle; never blocks + + monkeypatch.setattr(ms, "prewarm_picker_cache_async", prewarm) + + # Empty stdin -> immediate EOF -> main() returns after entering the loop. + monkeypatch.setattr(entry.sys, "stdin", io.StringIO("")) + + entry.main() + + +def test_main_prewarms_picker_cache_after_gateway_ready(monkeypatch): + """main() must call the prewarm helper once, after gateway.ready is + written, and still reach the stdin loop (returns on EOF = non-blocking).""" + events: list[tuple] = [] + + _run_main(monkeypatch, events) # returning at all proves the loop was reached + + prewarm_calls = [e for e in events if e[0] == "prewarm"] + assert len(prewarm_calls) == 1, ( + f"main() must invoke prewarm_picker_cache_async exactly once, got {events!r}" + ) + + ready_idx = events.index(("write", "gateway.ready")) + prewarm_idx = events.index(("prewarm",)) + assert ready_idx < prewarm_idx, ( + "prewarm must fire AFTER the gateway.ready write (idle window, " + f"banner already shown); order was {events!r}" + ) + + +def test_main_survives_prewarm_failure(monkeypatch): + """Fire-and-forget contract: a prewarm that raises at start must be + swallowed and main() must still reach the read loop and exit cleanly.""" + events: list[tuple] = [] + + def _boom(): + events.append(("prewarm",)) + raise RuntimeError("provider registry exploded") + + _run_main(monkeypatch, events, prewarm=_boom) # must not raise + + assert ("prewarm",) in events + assert ("write", "gateway.ready") in events diff --git a/tests/tui_gateway/test_kanban_notify_poller.py b/tests/tui_gateway/test_kanban_notify_poller.py index fd4374bb0b..59fe4b8e99 100644 --- a/tests/tui_gateway/test_kanban_notify_poller.py +++ b/tests/tui_gateway/test_kanban_notify_poller.py @@ -12,6 +12,7 @@ unsubscribe) and ``_format_kanban_event_text``. """ from types import SimpleNamespace +from unittest.mock import patch from hermes_cli import kanban_db as kb from tui_gateway.server import ( @@ -53,6 +54,17 @@ def _sub_rows(tid: str) -> list: class TestCollectKanbanNotifications: + def test_zero_sub_board_is_never_opened_writable(self): + conn = kb.connect() + conn.close() + kb.create_board("second-board") + + with patch.object(kb, "connect", wraps=kb.connect) as spy_connect: + texts = _collect_kanban_notifications(_session()) + + assert texts == [] + spy_connect.assert_not_called() + def test_delivers_completed_event_and_unsubscribes(self): tid = _create_subscribed_task() _complete(tid, summary="shipped the fix") @@ -66,45 +78,74 @@ class TestCollectKanbanNotifications: # Task is at a final status -> subscription removed. assert _sub_rows(tid) == [] - def test_claim_advances_cursor_so_second_poll_is_empty(self): + def test_matching_tui_sub_delivers_and_advances_cursor(self): tid = _create_subscribed_task() + pre_cursor = _sub_rows(tid)[0]["last_event_id"] conn = kb.connect() try: kb.block_task(conn, tid, reason="waiting on review") finally: conn.close() - first = _collect_kanban_notifications(_session()) - second = _collect_kanban_notifications(_session()) + with patch.object(kb, "connect", wraps=kb.connect) as spy_connect: + first = _collect_kanban_notifications(_session()) + second = _collect_kanban_notifications(_session()) assert len(first) == 1 assert "blocked" in first[0] assert "waiting on review" in first[0] assert second == [] + assert spy_connect.called # Blocked is not a final status -> subscription stays alive so a # respawned task's next terminal event still reaches the user. - assert len(_sub_rows(tid)) == 1 + rows = _sub_rows(tid) + assert len(rows) == 1 + assert rows[0]["last_event_id"] > pre_cursor - def test_ignores_other_sessions_and_platforms(self): - tid_other_session = _create_subscribed_task(chat_id="some-other-session") - tid_gateway = _create_subscribed_task(platform="telegram", chat_id="chat-1") + def test_non_tui_subscription_does_not_open_board_writable(self): + tid = _create_subscribed_task(platform="telegram", chat_id="chat-1") # New subs start caught up at creation time (issue #29905); record the # pre-completion cursors so we can assert they were never claimed. - pre_cursors = { - tid: _sub_rows(tid)[0]["last_event_id"] - for tid in (tid_other_session, tid_gateway) - } - _complete(tid_other_session) - _complete(tid_gateway) + pre_cursor = _sub_rows(tid)[0]["last_event_id"] + _complete(tid) - texts = _collect_kanban_notifications(_session()) + with patch.object(kb, "connect", wraps=kb.connect) as spy_connect: + texts = _collect_kanban_notifications(_session()) assert texts == [] - # Foreign subscriptions untouched: cursors unclaimed, rows still there. - for tid in (tid_other_session, tid_gateway): - rows = _sub_rows(tid) - assert len(rows) == 1 - assert rows[0]["last_event_id"] == pre_cursors[tid] + spy_connect.assert_not_called() + rows = _sub_rows(tid) + assert len(rows) == 1 + assert rows[0]["last_event_id"] == pre_cursor + + def test_other_tui_session_does_not_open_board_writable(self): + tid = _create_subscribed_task(chat_id="some-other-session") + pre_cursor = _sub_rows(tid)[0]["last_event_id"] + _complete(tid) + + with patch.object(kb, "connect", wraps=kb.connect) as spy_connect: + texts = _collect_kanban_notifications(_session()) + + assert texts == [] + spy_connect.assert_not_called() + rows = _sub_rows(tid) + assert len(rows) == 1 + assert rows[0]["last_event_id"] == pre_cursor + + def test_probe_error_falls_back_to_writable_delivery(self, monkeypatch): + tid = _create_subscribed_task() + _complete(tid, summary="fallback delivery") + + def fail_probe(*args, **kwargs): + raise OSError("probe unavailable") + + monkeypatch.setattr(kb, "count_notify_subs", fail_probe) + with patch.object(kb, "connect", wraps=kb.connect) as spy_connect: + texts = _collect_kanban_notifications(_session()) + + assert len(texts) == 1 + assert tid in texts[0] + spy_connect.assert_called_once() def test_no_session_key_is_a_noop(self): tid = _create_subscribed_task() diff --git a/tools/approval.py b/tools/approval.py index 8ebc47a462..f5e7f0beb1 100644 --- a/tools/approval.py +++ b/tools/approval.py @@ -2004,6 +2004,54 @@ def _mark_command_starts(command: str) -> str: return "".join(parts) +def _mask_quoted_newlines(command: str) -> str: + """Replace raw newlines inside single/double quotes with a space. + + Detection-only rewrite. A newline inside a quoted string is DATA to the + shell — part of the argument, not a command separator — yet the flat + ``_CMDPOS`` start-position class treats every raw ``\\n`` as a command + start. That made any multi-line quoted argument (``hermes send`` message + bodies, ``git commit -m`` messages, heredoc text) trip the hardline + blocklist when a data line began with e.g. ``sudo reboot``. + + Quote tracking mirrors ``_iter_shell_command_starts``: single quotes are + literal until the closing quote; inside double quotes a backslash escapes + the next character. Real command boundaries are unaffected: unquoted + newlines pass through untouched, ``$(``/backtick remain ``_CMDPOS`` + anchors independent of newlines, and ``_mark_command_starts`` still + re-inserts newlines at every genuine quote-aware command start. An + unclosed quote absorbs following newlines exactly as the shell would + (the quoted word continues across the line break), so masking them + cannot hide a runnable command. + """ + if "\n" not in command: + return command + out: list[str] = [] + quote: str | None = None + i = 0 + while i < len(command): + ch = command[i] + if quote: + if ch == "\\" and quote == '"' and i + 1 < len(command): + out.append(command[i:i + 2]) + i += 2 + continue + if ch == quote: + quote = None + out.append(" " if ch == "\n" else ch) + i += 1 + continue + if ch in ("'", '"'): + quote = ch + elif ch == "\\" and i + 1 < len(command): + out.append(command[i:i + 2]) + i += 2 + continue + out.append(ch) + i += 1 + return "".join(out) + + def _iter_shell_command_word_spans(command: str): """Yield command-position words that may be executable names.""" for command_start in _iter_shell_command_starts(command): @@ -2047,7 +2095,13 @@ def _iter_shell_command_word_spans(command: str): def _command_detection_variants(command: str): - normalized = _normalize_command_for_detection(command) + # Mask quoted newlines BEFORE normalization: normalization strips + # backslash-escapes (\" -> ") and empty-string pairs (""), which would + # corrupt quote tracking — e.g. `echo "a\""` normalizes to `echo "a` (an + # unterminated quote), so masking the normalized text could swallow a + # REAL unquoted newline separator that follows. The raw command carries + # faithful shell quote state. + normalized = _normalize_command_for_detection(_mask_quoted_newlines(command)) # Quote-aware grep parsing hides only structurally identified pattern # operands. Malformed/ambiguous input remains byte-for-byte intact. grep_safe, _ = _grep_safe_detection_variant(normalized) @@ -2523,7 +2577,10 @@ def prompt_dangerous_approval(command: str, description: str, smart_denied=False) -> str. Legacy callback signatures remain supported when ``smart_denied`` is false. - Returns: 'once', 'session', 'always', or 'deny' + Returns: 'once', 'session', 'always', 'deny', or 'timeout'. + 'timeout' means the prompt expired without a user response — the + action must still be blocked (fail-closed), but callers should + report it as "no response" rather than an explicit user denial. """ if timeout_seconds is None: timeout_seconds = _get_approval_timeout() @@ -2612,7 +2669,10 @@ def prompt_dangerous_approval(command: str, description: str, if thread.is_alive(): print("\n" + t("approval.timeout")) - return "deny" + # Distinct from an explicit deny: the user never answered. + # Callers still block (fail-closed) but tell the agent the + # prompt timed out instead of claiming the user refused. + return "timeout" choice = result["choice"] if smart_denied: @@ -3132,6 +3192,21 @@ def _run_approval_gate( choice=choice, ) + if choice == "timeout": + return { + "approved": False, + "message": ( + f"BLOCKED: Action timed out without user response. The user " + f"has NOT consented to this action. Do NOT retry it, do NOT " + f"rephrase it, and do NOT attempt the same outcome via a " + f"different path. Silence is not consent." + ), + "pattern_key": pattern_key, + "description": description, + "outcome": "timeout", + "user_consent": False, + } + if choice == "deny": return { "approved": False, @@ -3142,6 +3217,8 @@ def _run_approval_gate( ), "pattern_key": pattern_key, "description": description, + "outcome": "denied", + "user_consent": False, } if choice == "session": @@ -3916,6 +3993,25 @@ def check_all_command_guards(command: str, env_type: str, choice=choice, ) + if choice == "timeout": + breaker_addendum = _denial_breaker_addendum(session_key) + return { + "approved": False, + "message": ( + "BLOCKED: Command timed out without user response. The user " + "has NOT consented to this action. Do NOT retry this " + "command, do NOT rephrase it, and do NOT attempt the same " + "outcome via a different command. Stop the current workflow " + "and wait for the user to respond before taking any further " + "destructive or irreversible action. Silence is not " + f"consent.{breaker_addendum}" + ), + "pattern_key": primary_key, + "description": combined_desc, + "outcome": "timeout", + "user_consent": False, + } + if choice == "deny": breaker_addendum = _denial_breaker_addendum(session_key) return { @@ -4273,6 +4369,10 @@ def request_elicitation_consent( if choice in ("once", "session", "always"): return "accept" + if choice == "timeout": + # Prompt expired without a user response — mirror the gateway's + # unresolved outcome ("cancel") rather than an explicit decline. + return "cancel" return "decline" diff --git a/tools/computer_use/tool.py b/tools/computer_use/tool.py index dbe6fc1558..1422ff67e4 100644 --- a/tools/computer_use/tool.py +++ b/tools/computer_use/tool.py @@ -550,6 +550,14 @@ def _request_approval(action: str, args: Dict[str, Any], if verdict == "always_approve": _session_auto_approve[session_id] = True return None + if verdict == "timeout": + return json.dumps({ + "error": ( + "approval prompt timed out — the user did not respond. " + "Silence is not consent; do not retry without the user." + ), + "action": action, + }) return json.dumps({"error": "denied by user", "action": action}) diff --git a/tools/delegate_tool.py b/tools/delegate_tool.py index 18e47d429d..13e2143336 100644 --- a/tools/delegate_tool.py +++ b/tools/delegate_tool.py @@ -3649,112 +3649,50 @@ def _load_config() -> dict: def _build_top_level_description() -> str: - """Compose the delegate_task tool description with current runtime limits. + """Compose the delegate_task tool description. - The model needs to know its actual ceilings (not the framework defaults), - otherwise it self-caps at "default 3" / "default 2" even when the user has - raised delegation.max_concurrent_children / max_spawn_depth. Called both - at module import (to seed DELEGATE_TASK_SCHEMA) and on every - get_definitions() call via dynamic_schema_overrides. + Deliberately carries ONLY guidance that exists nowhere else in the + schema. Batch/concurrency limits live in the 'tasks' parameter + description and the nesting clause lives in the 'role' parameter + description (both rebuilt per get_definitions() call with the user's + actual delegation.max_concurrent_children / max_spawn_depth), so the + top-level text stays static and duplication-free. If you add text + here, check it is not already stated in a parameter description. """ - try: - max_children = _get_max_concurrent_children() - except Exception: - max_children = _DEFAULT_MAX_CONCURRENT_CHILDREN - try: - max_depth = _get_max_spawn_depth() - except Exception: - max_depth = MAX_DEPTH - try: - orchestrator_on = _get_orchestrator_enabled() - except Exception: - orchestrator_on = True - - if max_depth >= 2 and orchestrator_on: - nesting_clause = ( - f"Nested delegation IS enabled for this user " - f"(max_spawn_depth={max_depth}): pass role='orchestrator' on a " - f"child to let it spawn its own workers, up to {max_depth - 1} " - f"additional level(s) deep." - ) - elif max_depth >= 2 and not orchestrator_on: - nesting_clause = ( - f"Nested delegation is DISABLED on this install " - f"(delegation.orchestrator_enabled=false), even though " - f"max_spawn_depth={max_depth}. role='orchestrator' is silently " - f"forced to 'leaf'." - ) - else: - nesting_clause = ( - f"Nested delegation is OFF for this user " - f"(max_spawn_depth={max_depth}): every child is a leaf and " - f"cannot delegate further. Raise delegation.max_spawn_depth in " - f"config.yaml to enable nesting." - ) - return ( - "Spawn one or more subagents to work on tasks in isolated contexts. " - "Each subagent gets its own conversation, terminal session, and toolset. " - "Only the final summary is returned -- intermediate tool results " - "never enter your context window.\n\n" - "TWO MODES (one of 'goal' or 'tasks' is required):\n" - "1. Single task: provide 'goal' (+ optional context and role).\n" - f"2. Batch (parallel): provide 'tasks' array with up to {max_children} " - f"items concurrently for this user (configured via " - f"delegation.max_concurrent_children in config.yaml). {nesting_clause}\n\n" - "BOTH MODES RUN IN THE BACKGROUND. delegate_task returns immediately — " - "you and the user keep working, and the completed result re-enters " - "the conversation as a new message. A " - "batch returns one handle, runs N subagents concurrently, and delivers " - "one consolidated result after ALL of them finish. Do NOT wait or poll; " - "just continue with other work after dispatching.\n\n" - "LIVE TRANSCRIPTS: the dispatch response includes 'live_transcripts' — " - "one append-only human-readable log file per task (under " - "cache/delegation/live//). Each child streams its " - "assistant text, tool calls, and tool results there while it runs. " - "Read (or `tail -f` in a terminal) those paths any time you or the " - "user want to see what a subagent is actually doing instead of " - "waiting for the final summary.\n\n" - "WHEN TO USE delegate_task:\n" - "- Reasoning-heavy subtasks (debugging, code review, research synthesis)\n" - "- Tasks that would flood your context with intermediate data\n" - "- Parallel independent workstreams (research A and B simultaneously)\n\n" - "WHEN NOT TO USE (use these instead):\n" - "- Mechanical multi-step work with no reasoning needed -> use execute_code\n" - "- Single tool call -> just call the tool directly\n" - "- Tasks needing user interaction -> subagents cannot use clarify\n" - "- Durable long-running work that must outlive the current turn -> " - "use cronjob (action='create') or terminal(background=True, " - "notify_on_complete=True) instead. Background delegations are NOT " - "durable: if the parent session is closed (/new) or the process exits " - "before a subagent finishes, that subagent's work is discarded, and " - "/stop cancels every running background subagent.\n\n" - "IMPORTANT:\n" - "- Subagents have NO memory of your conversation. Pass all relevant " - "info (file paths, error messages, constraints) via the 'context' field.\n" - "- If the user is writing in a non-English language, or asked for " - "output in a specific language / tone / style, say so in 'context' " - "(e.g. \"respond in Chinese\", \"return output in Japanese\"). " - "Otherwise subagents default to English and their summaries will " - "contaminate your final reply with the wrong language.\n" - "- Subagent summaries are SELF-REPORTS, not verified facts. A subagent " - "that claims \"uploaded successfully\" or \"file written\" may be wrong. " - "For operations with external side-effects (HTTP POST/PUT, remote " - "writes, file creation at shared paths, publishing), require the " - "subagent to return a verifiable handle (URL, ID, absolute path, HTTP " - "status) and verify it yourself — fetch the URL, stat the file, read " - "back the content — before telling the user the operation succeeded.\n" - "- Leaf subagents (role='leaf', the default) CANNOT call: " - "delegate_task, clarify, memory, send_message.\n" - "- Orchestrator subagents (role='orchestrator') retain " - "delegate_task so they can spawn their own workers, but still " - "cannot use clarify, memory, or send_message. " - f"Orchestrators are bounded by max_spawn_depth={max_depth} for this " - f"user and can be disabled globally via " - "delegation.orchestrator_enabled=false.\n" - "- Subagent model is NOT selectable per call: children inherit the parent model (plus its fallback chain) unless you pin all subagents to a model via delegation.provider / delegation.model in config.yaml.\n" - "- Each subagent gets its own terminal session (separate working directory and state).\n" - "- Results are always returned as an array, one entry per task." + "Spawn subagents in isolated contexts; each gets its own conversation, " + "terminal session, and toolset, and only its final summary returns to " + "you. Provide 'goal' for a single task or 'tasks' for a parallel batch " + "(limits and nesting rules are in the parameter descriptions).\n\n" + "Runs in the background: dispatch returns immediately with live " + "transcript paths, and the completed result (one consolidated message " + "for a batch) re-enters the conversation on its own. Do NOT wait or " + "poll; continue other work.\n\n" + "USE FOR: reasoning-heavy subtasks, work that would flood your context " + "with intermediate data, or independent parallel workstreams.\n" + "DO NOT USE FOR (use these instead):\n" + "- Mechanical multi-step work with no reasoning needed -> execute_code\n" + "- A single tool call -> call the tool directly\n" + "- Tasks needing user interaction -> subagents cannot ask questions\n" + "- Durable work that must survive this session -> cronjob or " + "terminal(background=True, notify_on_complete=True); /stop, /new, or " + "process exit discards running subagents.\n\n" + "RULES:\n" + "- Children know nothing of this conversation: pass everything needed " + "via 'context', including any required output language, tone, or " + "style (e.g. \"respond in Chinese\").\n" + "- Child summaries are SELF-REPORTS, not verified facts: a child " + "claiming \"uploaded successfully\" or \"file written\" may be wrong. " + "For external side effects (uploads, remote writes, publishing), " + "require a verifiable handle (URL, ID, absolute path) and verify it " + "yourself — fetch the URL, stat the file, read back the content — " + "before telling the user the operation succeeded.\n" + "- Leaf children (the default) cannot call delegate_task, clarify, " + "memory, send_message, or cronjob; orchestrators regain only " + "delegate_task.\n" + "- Children inherit the parent model and fallback chain unless pinned " + "globally via delegation.provider / delegation.model in config.yaml. " + "Results are returned as an array, one entry per task." ) diff --git a/tools/environments/base.py b/tools/environments/base.py index a4b106b004..a57c4ba045 100644 --- a/tools/environments/base.py +++ b/tools/environments/base.py @@ -491,12 +491,12 @@ def _export_dump_excluding_session_vars( lines. ``|| true`` keeps the success contract for callers that chain on it. The dump MUST be wrapped in a brace group with the redirection applied to - the group. *tmp_path* typically embeds ``$BASHPID`` for concurrency-safe - temp names; a redirection attached to a pipeline segment would expand - ``$BASHPID`` inside that segment's subshell (a different PID than the - parent that expands the follow-up ``mv``), silently orphaning the dump. - The brace-group redirect is expanded in the current shell, keeping both - expansions consistent. + the group. *tmp_path* is typically a shell-variable expansion (a + mktemp-allocated per-writer temp name); a redirection attached to a + pipeline segment would expand it inside that segment's subshell, + potentially inconsistently with the parent that expands the follow-up + ``mv``. The brace-group redirect is expanded in the current shell, + keeping both expansions consistent. """ # ${!PREFIX*} is bash 3.2+ name-prefix expansion; empty matches are fine # because ``unset`` with only missing names is ignored under 2>/dev/null. @@ -668,14 +668,19 @@ class BaseEnvironment(ABC): # calls run) ``$$`` stays the *parent* shell's PID — so two concurrent # writers would pick the SAME temp name, clobber each other's temp # mid-write, and mv would then publish a torn file (the corruption is - # only narrowed, not closed). ``$BASHPID`` is the actual subshell PID - # and is genuinely unique per writer, which closes the race. The - # static path is shell-quoted (Windows/Git-Bash drive letters, spaces) - # with ``$BASHPID`` left outside the quotes so it still expands. - _snap_tmp = self._quote_shell_path(self._snapshot_path + ".tmp.") + "$BASHPID" + # only narrowed, not closed). ``$BASHPID`` would be unique per writer, + # but macOS ships bash 3.2 which does NOT provide it — the name expands + # empty there, so every writer shares one temp path and the race is + # back. ``mktemp`` allocates a per-writer unique path portably across + # bash versions. The template is shell-quoted (Windows/Git-Bash drive + # letters, spaces) and the resulting path lives in a shell variable so + # every later expansion is consistent. + _snap_tmp_template = self._quote_shell_path(self._snapshot_path + ".tmp.XXXXXXXXXX") + _snap_tmp = '"$__hermes_snap_tmp"' snapshot_excluded = self._snapshot_excluded_passthrough_names() bootstrap = ( f"umask 077\n" + f"__hermes_snap_tmp=$(mktemp {_snap_tmp_template}) || exit 1\n" f"{_export_dump_excluding_session_vars(_snap_tmp, snapshot_excluded)}\n" # Dump function definitions, filtering out private (``_``-prefixed) # helpers — mainly bash-completion internals (``_git``, ``_make``…) @@ -784,11 +789,13 @@ class BaseEnvironment(ABC): # Use atomic file replacement for env snapshot updates (issue #38249). # Assemble into a per-writer-unique temp file, then mv to atomically # replace the snapshot so concurrent source() calls never read a - # truncated/half-written file. ``$BASHPID`` (not ``$$``) is the actual - # subshell PID — unique per concurrent ``&``-launched writer — so two - # writers never share a temp name and clobber each other before the mv. - # Static path shell-quoted (Windows/spaces); ``$BASHPID`` left to expand. - _snap_tmp = self._quote_shell_path(self._snapshot_path + ".tmp.") + "$BASHPID" + # truncated/half-written file. ``mktemp`` is used instead of + # ``$BASHPID``/``$$`` because macOS bash 3.2 lacks ``$BASHPID`` (it + # expands empty, collapsing every writer onto one temp name) and ``$$`` + # is shared by ``&``-launched subshells. Template shell-quoted + # (Windows/spaces); the allocated path lives in a shell variable. + _snap_tmp_template = self._quote_shell_path(self._snapshot_path + ".tmp.XXXXXXXXXX") + _snap_tmp = '"$__hermes_snap_tmp"' parts = [] passthrough_names = self._snapshot_excluded_passthrough_names() @@ -842,13 +849,13 @@ class BaseEnvironment(ABC): # Chain mv on the export succeeding so a failed/partial dump never # replaces a good snapshot; drop the temp on failure so it isn't # orphaned (cleaned up wholesale in LocalEnvironment.cleanup too). - # NOTE: the redirection must be attached to a brace group — ``_snap_tmp`` - # embeds ``$BASHPID``, and a redirect on a pipeline segment expands - # inside that segment's subshell (a different PID than the parent that - # expands the ``mv`` operand), silently orphaning the dump. See - # _export_dump_excluding_session_vars. + # NOTE: the temp path is allocated with mktemp into a shell variable + # first — the redirection inside _export_dump_excluding_session_vars is + # attached to a brace group so the variable expands in the same shell + # that later expands the ``mv`` operand, keeping both consistent. if self._snapshot_ready: parts.append( + f"__hermes_snap_tmp=$(mktemp {_snap_tmp_template}) && " f"{{ {_export_dump_excluding_session_vars(_snap_tmp, passthrough_names)} " f"&& mv -f {_snap_tmp} {_quoted_snap}; }} " f"2>/dev/null || rm -f {_snap_tmp} 2>/dev/null || true" diff --git a/tools/environments/file_sync.py b/tools/environments/file_sync.py index 2357228ef9..84a314f5ee 100644 --- a/tools/environments/file_sync.py +++ b/tools/environments/file_sync.py @@ -156,6 +156,7 @@ class FileSyncManager: self._bulk_upload_fn = bulk_upload_fn self._bulk_download_fn = bulk_download_fn self._delete_fn = delete_fn + self._transaction_lock = threading.Lock() self._synced_files: dict[str, tuple[float, int]] = {} # remote_path -> (mtime, size) self._pushed_hashes: dict[str, str] = {} # remote_path -> sha256 hex digest self._upload_only_host_paths: set[str] = set() @@ -171,6 +172,11 @@ class FileSyncManager: Transactional: state only committed if ALL operations succeed. On failure, state rolls back so the next cycle retries everything. """ + with self._transaction_lock: + self._sync_transaction(force=force) + + def _sync_transaction(self, *, force: bool = False) -> None: + """Execute one sync cycle while holding the per-manager lock.""" if not force and not os.environ.get(_FORCE_SYNC_ENV): now = _monotonic() if now - self._last_sync_time < self._sync_interval: @@ -257,6 +263,11 @@ class FileSyncManager: Protected against SIGINT (defers the signal until complete) and serialized across concurrent gateway sandboxes via file lock. """ + with self._transaction_lock: + self._sync_back_transaction(hermes_home=hermes_home) + + def _sync_back_transaction(self, hermes_home: Path | None = None) -> None: + """Execute sync-back against a stable snapshot of manager state.""" if self._bulk_download_fn is None: return diff --git a/tools/file_operations.py b/tools/file_operations.py index f6a2deacab..ab4965ea14 100644 --- a/tools/file_operations.py +++ b/tools/file_operations.py @@ -465,7 +465,8 @@ class FileOperations(ABC): ... @abstractmethod - def write_file(self, path: str, content: str) -> WriteResult: + def write_file(self, path: str, content: str, + pre_content: Optional[str] = None) -> WriteResult: """Write content to a file, creating directories as needed.""" ... @@ -660,14 +661,8 @@ def _lint_yaml_inproc(content: str) -> tuple[bool, str]: def _lint_toml_inproc(content: str) -> tuple[bool, str]: """In-process TOML syntax check (stdlib tomllib, Python 3.11+).""" - try: - import tomllib as _toml - except ImportError: - # Pre-3.11 fallback via tomli, if installed. - try: - import tomli as _toml # type: ignore[no-redef] - except ImportError: - return True, "__SKIP__" + import tomllib as _toml + try: _toml.loads(content) return True, "" @@ -1012,6 +1007,9 @@ class ShellFileOperations(FileOperations): ``.hermes-tmp`` file next to the user's data, and the original file is left untouched. Content rides stdin so there is no ARG_MAX limit. + ``mkdir -p`` for the parent directory is folded into this script + (one fewer subprocess vs. a separate ``mkdir -p`` call). + Returns an :class:`ExecuteResult`; ``exit_code == 0`` means the file was swapped into place atomically. A non-zero exit means nothing was renamed and the original (if any) is intact. @@ -1025,6 +1023,9 @@ class ShellFileOperations(FileOperations): tmpl = self._escape_shell_arg(".hermes-tmp.XXXXXX") # One shell script, fully quoted. Notes: + # - `mkdir -p "$d"` is folded in here so the parent directory is + # created in the same subprocess that writes the temp file — + # saves one entire subprocess spawn vs. a separate mkdir call. # - `mktemp` lands the temp in the target's own dir (-p) so `mv` is # same-FS atomic; we fall back to a PID-stamped name if the # backend lacks mktemp (rare; busybox/macOS/Linux all ship it). @@ -1058,6 +1059,11 @@ class ShellFileOperations(FileOperations): 'rt="$(readlink -f "$t" 2>/dev/null || realpath "$t" 2>/dev/null || true)"; ' '[ -n "$rt" ] && { t="$rt"; d="$(dirname "$t")"; }; ' "fi; " + # Create the parent dir in the SAME subprocess that writes the + # temp file (one fewer exec vs. a separate mkdir call). Runs + # AFTER symlink resolution so a resolved target's directory is + # the one created/confirmed. + 'mkdir -p "$d"; ' 'tmp="$(mktemp -p "$d" ' + tmpl + ' 2>/dev/null ' '|| mktemp "$d/.hermes-tmp.$$.XXXXXX" 2>/dev/null ' '|| { tmp="$d/.hermes-tmp.$$"; : > "$tmp" && echo "$tmp"; })"; ' @@ -1102,13 +1108,16 @@ class ShellFileOperations(FileOperations): def _file_has_bom(self, path: str, pre_content: Optional[str] = None) -> bool: """Whether the file on disk starts with a UTF-8 BOM. - Uses ``pre_content`` if we already read the file (zero extra exec - calls); otherwise issues a tiny ``head -c 3`` to sample just the - marker. A missing/empty file returns False (new writes get no BOM + Always probes the first 3 bytes on disk — do NOT trust + ``pre_content`` for BOM detection because the most common + provider (``read_file_raw``) deliberately strips BOMs so the + agent never sees U+FEFF glyphs. Passing BOM-stripped content + through ``pre_content`` would cause a false-negative and + silently remove the marker on rewrite. + + A missing/empty file returns False (new writes get no BOM unless the caller explicitly includes one). """ - if pre_content is not None: - return _has_bom(pre_content) head_cmd = f"head -c 3 {self._escape_shell_arg(path)} 2>/dev/null" head_result = self._exec(head_cmd) if head_result.exit_code != 0 or not head_result.stdout: @@ -1400,7 +1409,8 @@ class ShellFileOperations(FileOperations): # WRITE Implementation # ========================================================================= - def write_file(self, path: str, content: str) -> WriteResult: + def write_file(self, path: str, content: str, + pre_content: Optional[str] = None) -> WriteResult: """ Write content to a file, creating parent directories as needed. @@ -1427,6 +1437,15 @@ class ShellFileOperations(FileOperations): Args: path: File path to write content: Content to write + pre_content: Pre-edit file content if the caller already has it + (e.g. patch_replace read the file for fuzzy matching). + When provided, skips a redundant ``cat`` subprocess to + re-read the file for lint baseline / line-ending + detection. BOM detection always probes disk (the most + common provider — ``read_file_raw`` — strips BOMs, so + trusting ``pre_content`` for BOM would cause false + negatives and silent marker loss on rewrite). When + None, reads from disk as before. Returns: WriteResult with bytes written, lint summary, or error. @@ -1492,17 +1511,21 @@ class ShellFileOperations(FileOperations): # the UNION of in-process lint coverage and LSP coverage. For # extensions outside both sets (binaries, opaque formats), # skipping the read keeps the hot path fast. - pre_content: Optional[str] = None want_pre = ext in LINTERS_INPROC or self._lsp_handles_extension(ext) if want_pre: - # Best-effort read; failure (file missing, permission) leaves - # pre_content as None which makes both downstream consumers - # degrade gracefully (lint reports all errors; LSP skips the - # shift map). - read_cmd = f"cat {self._escape_shell_arg(path)} 2>/dev/null" - read_result = self._exec(read_cmd) - if read_result.exit_code == 0 and read_result.stdout: - pre_content = read_result.stdout + if pre_content is not None: + # Caller already has file content (e.g. patch_replace read it + # for fuzzy matching) — reuse directly, skip redundant cat. + pass + else: + # Best-effort read; failure (file missing, permission) leaves + # pre_content as None which makes both downstream consumers + # degrade gracefully (lint reports all errors; LSP skips the + # shift map). + read_cmd = f"cat {self._escape_shell_arg(path)} 2>/dev/null" + read_result = self._exec(read_cmd) + if read_result.exit_code == 0 and read_result.stdout: + pre_content = read_result.stdout # ── Line-ending preservation (Roo Code pattern) ────────────── # If the file existed with CRLF endings and the agent's content @@ -1534,15 +1557,15 @@ class ShellFileOperations(FileOperations): # rather than an external IDE. self._snapshot_lsp_baseline(path) - # Create parent directories + # Write atomically. ``mkdir -p`` is folded into _atomic_write + # (one fewer subprocess vs. a separate mkdir call). + # ``dirs_created`` has always meant "parent dirs ensured" — + # ``mkdir -p`` exits 0 even when the dirs pre-exist, so the old + # separate-mkdir code reported True in exactly the same cases. + # A mkdir failure now surfaces as the atomic-write error return + # below, before this field is ever emitted. parent = os.path.dirname(path) - dirs_created = False - - if parent: - mkdir_cmd = f"mkdir -p {self._escape_shell_arg(parent)}" - mkdir_result = self._exec(mkdir_cmd) - if mkdir_result.exit_code == 0: - dirs_created = True + dirs_created = bool(parent) # Write atomically: stream into a temp file in the SAME directory, # then ``mv`` it over the target. The rename is atomic on POSIX @@ -1564,14 +1587,16 @@ class ShellFileOperations(FileOperations): if write_result.exit_code != 0: return WriteResult(error=f"Failed to write file: {write_result.stdout}") - # Get bytes written (wc -c is POSIX, works on Linux + macOS) - stat_cmd = f"wc -c < {self._escape_shell_arg(path)} 2>/dev/null" - stat_result = self._exec(stat_cmd) - - try: - bytes_written = int(stat_result.stdout.strip()) - except ValueError: - bytes_written = len(content.encode('utf-8')) + # Get bytes written — compute from the content we just wrote + # (len of the UTF-8 encoding matches wc -c) instead of spawning a + # ``wc -c`` subprocess. ``surrogatepass`` matches the sha256 + # verification below: content that flowed through a surrogateescape + # decode (backend output via patch_replace) may carry lone + # surrogates a strict encode would reject. Encode ONCE and share + # the bytes with the sha256 block — a second full encode of a + # multi-MB file is measurable. + content_bytes = content.encode('utf-8', 'surrogatepass') + bytes_written = len(content_bytes) # Post-write content verification (cheap, one shell call): compare # the on-disk sha256 to the intended content's hash. Production @@ -1586,7 +1611,7 @@ class ShellFileOperations(FileOperations): hash_result = self._exec(hash_cmd) if hash_result.exit_code == 0 and hash_result.stdout.strip(): disk_sha = hash_result.stdout.strip().split()[0] - expected_sha = hashlib.sha256(content.encode("utf-8", "surrogatepass")).hexdigest() + expected_sha = hashlib.sha256(content_bytes).hexdigest() content_verified = disk_sha == expected_sha if not content_verified: return WriteResult( @@ -1659,6 +1684,9 @@ class ShellFileOperations(FileOperations): return PatchResult(error=f"Failed to read file: {path}") content = read_result.stdout + # Preserve raw content (including BOM) for write_file's pre_content + # so write_file can detect/restore BOM correctly. + raw_content = content # Strip a leading UTF-8 BOM before matching so the fuzzy matcher and # the diff operate on clean content (a phantom U+FEFF before line 1 # defeats an exact first-line match). write_file restores the BOM on @@ -1711,8 +1739,11 @@ class ShellFileOperations(FileOperations): if file_ending: new_content = _normalize_line_endings(new_content, file_ending) - # Write back - write_result = self.write_file(path, new_content) + # Write back — pass pre_content (original read, with BOM) to avoid + # a redundant cat subprocess inside write_file. Must be the raw + # content (before _strip_bom) so write_file can detect/restore BOM. + write_result = self.write_file(path, new_content, + pre_content=raw_content) if write_result.error: return PatchResult(error=f"Failed to write changes: {write_result.error}") diff --git a/tools/lazy_deps.py b/tools/lazy_deps.py index 15992f7006..926aa0d868 100644 --- a/tools/lazy_deps.py +++ b/tools/lazy_deps.py @@ -847,6 +847,35 @@ def ensure(feature: str, *, prompt: bool = True) -> None: if unsupported: raise FeatureUnavailable(feature, missing, unsupported) + # Package-manager installs (NixOS, and any other distro that ships Hermes + # from a read-only store) cannot receive lazy pip installs: the venv's + # site-packages lives in the store, so the uv -> pip -> ensurepip ladder + # below burns ~15s bootstrapping ensurepip only to fail on a read-only + # target. Fail fast with an actionable message instead. + # + # Skipped when a durable install target is configured: the container + # deployment sets HERMES_MANAGED=true *and* HERMES_LAZY_INSTALL_TARGET + # (a writable volume), where lazy installs legitimately work. + # + # The reason string starts with "unsupported " on purpose: + # refresh_active_features classifies FeatureUnavailable by that prefix and + # reports anything else as a hard failure rather than a skip. + if _lazy_install_target() is None: + try: + from hermes_cli.config import get_managed_system + + managed_by = get_managed_system() + except Exception: + managed_by = "" # config unreadable — proceed with the install + if managed_by: + raise FeatureUnavailable( + feature, missing, + f"unsupported on {managed_by}-managed installs: this build's " + f"packages come from {managed_by}, so Hermes cannot install " + f"them at runtime. Add the dependencies for {feature!r} via " + f"{managed_by} (or run a pip/uv install of Hermes instead)." + ) + # Validate every spec against the allowlist + safety regex. Belt and # braces — the keys-in-LAZY_DEPS check above already constrains this. for spec in missing: diff --git a/tools/mcp_schema_cache.py b/tools/mcp_schema_cache.py new file mode 100644 index 0000000000..0fef2cdc32 --- /dev/null +++ b/tools/mcp_schema_cache.py @@ -0,0 +1,121 @@ +"""Persistent MCP tool-schema cache for lazy server startup. + +Stores per-server tool manifests on disk so Hermes can register MCP tools +into the agent snapshot without spawning the stdio child process at idle +dashboard startup. Cache entries are keyed by server name + a fingerprint +of the connection config (command/args/url/tools filters). +""" + +from __future__ import annotations + +import hashlib +import json +import logging +import threading +from pathlib import Path +from typing import Any, Dict, List, Optional + +logger = logging.getLogger(__name__) + +_CACHE_FILENAME = "mcp_schema_cache.json" +_cache_lock = threading.Lock() + + +def _cache_path() -> Path: + from hermes_constants import get_hermes_home + + return get_hermes_home() / "cache" / _CACHE_FILENAME + + +def config_fingerprint(config: dict) -> str: + """Stable hash of the connection-defining parts of an MCP server config.""" + tools_filter = config.get("tools") or {} + payload = { + "command": config.get("command"), + "args": config.get("args") or [], + "url": config.get("url"), + "transport": config.get("transport"), + "tools_include": sorted(tools_filter.get("include") or []), + "tools_exclude": sorted(tools_filter.get("exclude") or []), + } + raw = json.dumps(payload, sort_keys=True, separators=(",", ":")) + return hashlib.sha256(raw.encode("utf-8")).hexdigest()[:16] + + +def _load_all() -> Dict[str, Any]: + path = _cache_path() + if not path.exists(): + return {} + try: + data = json.loads(path.read_text(encoding="utf-8")) + return data if isinstance(data, dict) else {} + except Exception as exc: + logger.debug("Could not read MCP schema cache %s: %s", path, exc) + return {} + + +def _save_all(data: Dict[str, Any]) -> None: + from utils import atomic_json_write + + # Cache dir + 0o600: sibling precedent in tools/registry.py + # _save_discovery_cache; the cache file is trusted input on the lazy + # registration path, so keep it user-only. + atomic_json_write(_cache_path(), data, mode=0o600) + + +def get_cached_entry(server_name: str, fingerprint: str) -> Optional[dict]: + """Return cached entry when fingerprint matches, else None.""" + with _cache_lock: + entry = _load_all().get(server_name) + if not isinstance(entry, dict): + return None + if entry.get("fingerprint") != fingerprint: + return None + return entry + + +def has_cached_entry(server_name: str, fingerprint: str) -> bool: + return get_cached_entry(server_name, fingerprint) is not None + + +def write_cache_entry( + server_name: str, + fingerprint: str, + *, + tools: List[dict], + utility_tools: Optional[List[dict]] = None, +) -> None: + """Persist tool schemas after a successful live connect.""" + entry = { + "fingerprint": fingerprint, + "tools": tools, + "utility_tools": utility_tools or [], + } + with _cache_lock: + data = _load_all() + # Write-through fires on every registration (reconnects, + # list_changed refreshes); skip the load-all+rewrite churn when the + # entry is byte-identical to what is already on disk. + if data.get(server_name) == entry: + return + data[server_name] = entry + _save_all(data) + + +def clear_cache_entry(server_name: str) -> None: + with _cache_lock: + data = _load_all() + if server_name in data: + del data[server_name] + _save_all(data) + + +def tools_from_cache_entry(entry: dict) -> List[dict]: + """Return cached MCP tool dicts (name, description, inputSchema).""" + tools = entry.get("tools") + return list(tools) if isinstance(tools, list) else [] + + +def utility_tools_from_cache_entry(entry: dict) -> List[dict]: + util = entry.get("utility_tools") + return list(util) if isinstance(util, list) else [] diff --git a/tools/mcp_tool.py b/tools/mcp_tool.py index b3c356788a..993d9a13c8 100644 --- a/tools/mcp_tool.py +++ b/tools/mcp_tool.py @@ -3532,6 +3532,12 @@ class MCPServerTask: _servers: Dict[str, MCPServerTask] = {} _server_connecting: set[str] = set() _server_connect_errors: Dict[str, str] = {} +# Lazy MCP startup (#56832): servers whose tools were registered from the +# on-disk schema cache without spawning/connecting. Keyed by server name; +# entries are popped once a real connection is established on first use. +_lazy_server_configs: Dict[str, dict] = {} +_lazy_server_fingerprints: Dict[str, str] = {} +_lazy_server_tool_names: Dict[str, List[str]] = {} # Discovery installs a task-local claim before calling ``_connect_server`` so # it can retain a recoverable parked task without making standalone probe calls # publish failed servers into module-global ownership. @@ -4758,10 +4764,103 @@ def _request_lazy_reconnect(server_name: str, server: MCPServerTask) -> bool: return False -def _get_connected_server_for_call(server_name: str) -> Optional[MCPServerTask]: - """Return a connected server, lazily reconnecting recycled stdio state.""" +def _resolve_server_lazy(name: str, config: dict) -> bool: + """True when this server defers spawn/connect until first tool use. + + Gated per-server by ``mcp_servers..lazy`` in config (default OFF), + following the same per-server key pattern as ``idle_timeout_seconds``. + Design from #56832 (Vansh5632). + """ + return _parse_boolish(config.get("lazy", False), default=False) + + +def _ensure_lazy_server_connected(server_name: str) -> bool: + """Connect a lazily-registered MCP server on demand (sync, blocks caller). + + Composes with the existing connect machinery: respects the per-server + connect cooldown (#50394), the ``_server_connecting`` dedup set, and + routes through ``_discover_and_register_server`` so parked/recycle/ + cooldown bookkeeping stays in one place. Returns True when a live + session is available afterwards. + """ with _lock: server = _servers.get(server_name) + if server is not None and server.session is not None: + return True + config = _lazy_server_configs.get(server_name) + if not config: + return False + if _connect_cooldown_active(server_name): + return False + if server_name in _server_connecting: + return False + _server_connecting.add(server_name) + _server_connect_errors.pop(server_name, None) + + logger.info("MCP server '%s': lazy start on first use", server_name) + _ensure_mcp_loop() + connect_timeout = config.get("connect_timeout", _DEFAULT_CONNECT_TIMEOUT) + + async def _connect(): + return await _discover_and_register_server(server_name, config) + + try: + _run_on_mcp_loop(_connect, timeout=float(connect_timeout) + 30.0) + except BaseException as exc: + message = _format_connect_error(exc) + with _lock: + _server_connecting.discard(server_name) + _server_connect_errors[server_name] = message + _record_connect_failure(server_name) + logger.warning( + "Lazy MCP connect failed for '%s': %s", server_name, message, + ) + return False + + with _lock: + _server_connecting.discard(server_name) + _clear_connect_failure(server_name) + _lazy_server_configs.pop(server_name, None) + stale_fingerprint = _lazy_server_fingerprints.pop(server_name, None) + cached_names = _lazy_server_tool_names.pop(server_name, None) or [] + server = _servers.get(server_name) + live_names = set( + getattr(server, "_registered_tool_names", []) or [] + ) + # Stale-cache reconciliation: the cached manifest may advertise tools + # the live server no longer serves. Deregister those phantoms so the + # model stops seeing tools that can never succeed. + phantom_names = [n for n in cached_names if n not in live_names] + if phantom_names: + from tools.registry import registry + + for tool_name in phantom_names: + registry.deregister(tool_name) + _forget_mcp_tool_server(tool_name) + logger.info( + "MCP server '%s': deregistered %d phantom cached tool(s) not " + "served live (stale schema-cache fingerprint %s): %s", + server_name, len(phantom_names), stale_fingerprint, + ", ".join(phantom_names), + ) + return server is not None and server.session is not None + + +def _get_connected_server_for_call(server_name: str) -> Optional[MCPServerTask]: + """Return a connected server, lazily reconnecting recycled stdio state. + + Also the single first-use connect point for lazy (schema-cache + registered) servers, so raw tool calls AND the resource/prompt utility + handlers all trigger the deferred spawn (#56832). + """ + with _lock: + server = _servers.get(server_name) + is_lazy = server_name in _lazy_server_configs + if is_lazy and (server is None or server.session is None): + _ensure_lazy_server_connected(server_name) + with _lock: + server = _servers.get(server_name) + return server if server is not None and server.session is None and server._is_recycled_stdio(): _request_lazy_reconnect(server_name, server) with _lock: @@ -5240,10 +5339,13 @@ def _make_check_fn(server_name: str): def _check() -> bool: with _lock: server = _servers.get(server_name) - return ( - server is not None - and (server.session is not None or server._is_recycled_stdio()) - ) + if server is not None and ( + server.session is not None or server._is_recycled_stdio() + ): + return True + # Lazy (schema-cache registered) servers are available: the + # first real call spawns/connects them (#56832). + return server_name in _lazy_server_configs return _check @@ -5692,6 +5794,16 @@ def _existing_tool_names() -> List[str]: for mcp_tool in server._tools: schema = _convert_mcp_schema(server.name, mcp_tool) names.append(schema["name"]) + # Lazy servers registered from the schema cache have no MCPServerTask + # yet — their tools live in the registry only (#56832). + with _lock: + lazy_names = [ + n + for sname, tool_names in _lazy_server_tool_names.items() + if sname not in _servers + for n in tool_names + ] + names.extend(lazy_names) return names @@ -5879,7 +5991,165 @@ def _register_server_tools(name: str, server: MCPServerTask, config: dict) -> Li if registered_names: registry.register_toolset_alias(name, toolset_name) + # Write-through (#56832): refresh the on-disk schema cache after a + # live connect so the next startup can lazily register this server + # without spawning it. Cache failures never break registration. + try: + from tools.mcp_schema_cache import config_fingerprint, write_cache_entry + tools_payload: List[dict] = [] + for mcp_tool in server._tools: + if not _should_register(mcp_tool.name): + continue + schema_obj = getattr(mcp_tool, "inputSchema", None) + tools_payload.append({ + "name": mcp_tool.name, + "description": mcp_tool.description or "", + "inputSchema": schema_obj if isinstance(schema_obj, dict) else {}, + }) + utility_payload = [ + {"schema": entry["schema"], "handler_key": entry["handler_key"]} + for entry in _select_utility_schemas(name, server, config) + ] + write_cache_entry( + name, + config_fingerprint(config), + tools=tools_payload, + utility_tools=utility_payload, + ) + except Exception as exc: + logger.debug("MCP schema cache write failed for '%s': %s", name, exc) + + return registered_names + + +class _CachedMCPTool: + """Minimal stand-in for MCP Tool objects loaded from the schema cache.""" + + __slots__ = ("name", "description", "inputSchema") + + def __init__(self, name: str, description: str, inputSchema: dict): + self.name = name + self.description = description + self.inputSchema = inputSchema or {} + + +def _register_from_cache_sync(name: str, config: dict, entry: dict) -> List[str]: + """Register a server's tools from a cached manifest, no child process. + + Lazy startup (#56832, design by Vansh5632): tools appear in the registry + immediately; the first real call routes through + ``_get_connected_server_for_call`` → ``_ensure_lazy_server_connected``. + """ + from tools.registry import registry + from tools.mcp_schema_cache import ( + config_fingerprint, + tools_from_cache_entry, + utility_tools_from_cache_entry, + ) + + registered_names: List[str] = [] + toolset_name = f"mcp-{name}" + fingerprint = config_fingerprint(config) + tool_timeout = config.get("timeout", _DEFAULT_TOOL_TIMEOUT) + tools_filter = config.get("tools") or {} + include_set = _normalize_name_filter( + tools_filter.get("include"), f"mcp_servers.{name}.tools.include" + ) + exclude_set = _normalize_name_filter( + tools_filter.get("exclude"), f"mcp_servers.{name}.tools.exclude" + ) + + def _should_register(tool_name: str) -> bool: + if include_set: + return matches_name_filter(tool_name, include_set) + if exclude_set: + return not matches_name_filter(tool_name, exclude_set) + return True + + check_fn = _make_check_fn(name) + for raw in tools_from_cache_entry(entry): + if not isinstance(raw, dict): + continue + raw_name = raw.get("name") + if not raw_name or not _should_register(raw_name): + continue + raw_schema = raw.get("inputSchema") + mcp_tool = _CachedMCPTool( + raw_name, + raw.get("description") or "", + raw_schema if isinstance(raw_schema, dict) else {}, + ) + # Defense-in-depth: the cache file is user-writable JSON, so run the + # same injection scan the eager discovery path applies. + _scan_mcp_description(name, mcp_tool.name, mcp_tool.description or "") + schema = _convert_mcp_schema(name, mcp_tool) + registry_name = schema["name"] + existing_toolset = registry.get_toolset_for_tool(registry_name) + if existing_toolset and existing_toolset != toolset_name: + logger.warning( + "MCP server '%s' (lazy): cached tool '%s' collides with " + "toolset '%s' — skipping", + name, registry_name, existing_toolset, + ) + continue + registry.register( + name=registry_name, + toolset=toolset_name, + schema=schema, + handler=_make_tool_handler(name, raw_name, tool_timeout), + check_fn=check_fn, + is_async=False, + description=schema["description"], + ) + if registry.get_toolset_for_tool(registry_name) != toolset_name: + continue + _track_mcp_tool_server(registry_name, name) + registered_names.append(registry_name) + + handler_factories = { + "list_resources": _make_list_resources_handler, + "read_resource": _make_read_resource_handler, + "list_prompts": _make_list_prompts_handler, + "get_prompt": _make_get_prompt_handler, + } + for raw in utility_tools_from_cache_entry(entry): + if not isinstance(raw, dict): + continue + schema = raw.get("schema") + handler_key = raw.get("handler_key") + if not isinstance(schema, dict) or handler_key not in handler_factories: + continue + util_name = schema.get("name") or "" + if not util_name: + continue + existing_toolset = registry.get_toolset_for_tool(util_name) + if existing_toolset and existing_toolset != toolset_name: + continue + registry.register( + name=util_name, + toolset=toolset_name, + schema=schema, + handler=handler_factories[handler_key](name, tool_timeout), + check_fn=check_fn, + is_async=False, + description=schema.get("description") or "", + ) + if registry.get_toolset_for_tool(util_name) != toolset_name: + continue + _track_mcp_tool_server(util_name, name) + registered_names.append(util_name) + + if registered_names: + registry.register_toolset_alias(name, toolset_name) + with _lock: + _lazy_server_configs[name] = dict(config) + _lazy_server_fingerprints[name] = fingerprint + _lazy_server_tool_names[name] = list(registered_names) + logger.info( + "MCP server '%s' (lazy): registered %d tool(s) from schema cache", + name, len(registered_names), + ) return registered_names async def _discover_and_register_server(name: str, config: dict) -> List[str]: @@ -5980,6 +6250,9 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: for k, v in servers.items() if k not in _servers and k not in connecting + # Servers already lazily registered from the schema cache are + # not re-registered; they connect on first tool use (#56832). + and k not in _lazy_server_configs and _parse_boolish(v.get("enabled", True), default=True) # Skip a server still serving its post-failure backoff. Without # this, a server that fails to connect (and is therefore never @@ -6015,6 +6288,51 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: if not new_servers: return _existing_tool_names() + # Lazy startup (#56832): servers gated with ``lazy: true`` whose config + # fingerprint matches a valid on-disk schema-cache entry register their + # tools from cache WITHOUT spawning/connecting. A missing or stale cache + # entry falls back to the normal eager connect below (which write-through + # refreshes the cache for next time). + eager_servers: Dict[str, dict] = dict(new_servers) + lazy_registered = 0 + lazy_server_count = 0 + try: + from tools.mcp_schema_cache import config_fingerprint, get_cached_entry + except Exception: # pragma: no cover - cache module missing + config_fingerprint = None # type: ignore[assignment] + get_cached_entry = None # type: ignore[assignment] + if config_fingerprint is not None and get_cached_entry is not None: + for name, cfg in new_servers.items(): + if not _resolve_server_lazy(name, cfg): + continue + entry = get_cached_entry(name, config_fingerprint(cfg)) + if not entry: + continue + with _lock: + _server_connecting.discard(name) + try: + names = _register_from_cache_sync(name, cfg, entry) + except Exception as exc: + logger.warning( + "Failed lazy MCP registration for '%s': %s", name, exc, + ) + with _lock: + _server_connecting.add(name) + continue + eager_servers.pop(name, None) + lazy_registered += len(names) + lazy_server_count += 1 + new_servers = eager_servers + + if not new_servers: + if lazy_registered: + logger.info( + "MCP: registered %d lazy tool(s) from schema cache " + "(no processes spawned)", + lazy_registered, + ) + return _existing_tool_names() + # Start the background event loop for MCP connections _ensure_mcp_loop() @@ -6103,8 +6421,10 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: for n in connected ) failed = len(new_servers) - len(connected) + new_tool_count += lazy_registered + connected_count = len(connected) + lazy_server_count if new_tool_count or failed: - summary = f"MCP: registered {new_tool_count} tool(s) from {len(connected)} server(s)" + summary = f"MCP: registered {new_tool_count} tool(s) from {connected_count} server(s)" if failed: summary += f" ({failed} failed)" logger.info(summary) diff --git a/tools/patch_parser.py b/tools/patch_parser.py index 271538904e..37412336ee 100644 --- a/tools/patch_parser.py +++ b/tools/patch_parser.py @@ -29,6 +29,7 @@ Usage: """ import difflib +import inspect import re from dataclasses import dataclass, field from typing import List, Optional, Tuple, Any @@ -432,6 +433,7 @@ def apply_v4a_operations(operations: List[PatchOperation], # ``PatchResult.lsp_diagnostics`` aggregation below. lsp_blocks: List[str] = [] errors = [] + lint_results = {} for op in operations: try: @@ -442,6 +444,8 @@ def apply_v4a_operations(operations: List[PatchOperation], all_diffs.append(result[1]) if result[2]: lsp_blocks.append(result[2]) + if result[3]: + lint_results[op.file_path] = result[3] else: errors.append(f"Failed to add {op.file_path}: {result[1]}") @@ -468,18 +472,18 @@ def apply_v4a_operations(operations: List[PatchOperation], all_diffs.append(result[1]) if result[2]: lsp_blocks.append(result[2]) + if result[3]: + lint_results[op.file_path] = result[3] else: errors.append(f"Failed to update {op.file_path}: {result[1]}") except Exception as e: errors.append(f"Error processing {op.file_path}: {str(e)}") - # Run lint on all modified/created files - lint_results = {} - for f in files_modified + files_created: - if hasattr(file_ops, '_check_lint'): - lint_result = file_ops._check_lint(f) - lint_results[f] = lint_result.to_dict() + # Lint results were collected from write_file's internal _check_lint_delta + # via the four-tuple return of _apply_add / _apply_update — zero extra + # subprocess calls vs. the old approach of re-reading each file with a + # bare _check_lint(f) that lacked post_content context. combined_diff = '\n'.join(all_diffs) @@ -514,14 +518,34 @@ def apply_v4a_operations(operations: List[PatchOperation], ) -def _apply_add(op: PatchOperation, file_ops: Any) -> Tuple[bool, str, Optional[str]]: +def _write_file_accepts_pre_content(file_ops: Any) -> bool: + """True when ``file_ops.write_file`` accepts a ``pre_content`` kwarg. + + Decided from the signature (not by catching TypeError around the call) + so a TypeError raised *inside* a capable ``write_file`` propagates + instead of triggering a second, duplicate write. Unintrospectable + callables (some C-implemented ones) conservatively get the basic + two-argument form. + """ + try: + params = inspect.signature(file_ops.write_file).parameters + except (TypeError, ValueError): + return False + return "pre_content" in params or any( + p.kind is inspect.Parameter.VAR_KEYWORD for p in params.values() + ) + + +def _apply_add(op: PatchOperation, file_ops: Any) -> Tuple[bool, str, Optional[str], Optional[dict]]: """Apply an add file operation. - Returns ``(success, diff_or_error, lsp_diagnostics)``. The third - element carries the formatted ```` block from + Returns ``(success, diff_or_error, lsp_diagnostics, lint_result)``. + The third element carries the formatted ```` block from :class:`WriteResult.lsp_diagnostics` so V4A patches can surface - semantic diagnostics from the LSP layer — without this, the LSP - tier would silently swallow them on the V4A code path. + semantic diagnostics from the LSP layer. The fourth element carries + the ``WriteResult.lint`` dict (syntax check result) so V4A patches + can propagate lint to ``PatchResult.lint`` without a redundant + ``_check_lint`` re-read — write_file already ran the check internally. """ # Extract content from hunks (all + lines) content_lines = [] @@ -532,14 +556,15 @@ def _apply_add(op: PatchOperation, file_ops: Any) -> Tuple[bool, str, Optional[s content = '\n'.join(content_lines) + # _apply_add creates a new file, no pre_content to pass result = file_ops.write_file(op.file_path, content) if result.error: - return False, result.error, None - + return False, result.error, None, None + diff = f"--- /dev/null\n+++ b/{op.file_path}\n" diff += '\n'.join(f"+{line}" for line in content_lines) - - return True, diff, getattr(result, "lsp_diagnostics", None) + + return True, diff, getattr(result, "lsp_diagnostics", None), getattr(result, "lint", None) def _apply_delete(op: PatchOperation, file_ops: Any) -> Tuple[bool, str]: @@ -573,11 +598,11 @@ def _apply_move(op: PatchOperation, file_ops: Any) -> Tuple[bool, str]: return True, diff -def _apply_update(op: PatchOperation, file_ops: Any) -> Tuple[bool, str, Optional[str]]: +def _apply_update(op: PatchOperation, file_ops: Any) -> Tuple[bool, str, Optional[str], Optional[dict]]: """Apply an update file operation. - Returns ``(success, diff_or_error, lsp_diagnostics)`` — see - :func:`_apply_add` for the rationale on the third element. + Returns ``(success, diff_or_error, lsp_diagnostics, lint_result)`` — see + :func:`_apply_add` for the rationale on the third and fourth elements. """ # Deferred import: breaks the patch_parser ↔ fuzzy_match circular dependency from tools.fuzzy_match import fuzzy_find_and_replace @@ -586,7 +611,7 @@ def _apply_update(op: PatchOperation, file_ops: Any) -> Tuple[bool, str, Optiona read_result = file_ops.read_file_raw(op.file_path) if read_result.error: - return False, f"Cannot read file: {read_result.error}", None + return False, f"Cannot read file: {read_result.error}", None, None current_content = read_result.content @@ -650,7 +675,7 @@ def _apply_update(op: PatchOperation, file_ops: Any) -> Tuple[bool, str, Optiona err_msg += format_no_match_hint(error, 0, search_pattern, new_content) except Exception: pass - return False, err_msg, None + return False, err_msg, None, None else: # Addition-only hunk (no context or removed lines). # Insert at the location indicated by the context hint, or at end of file. @@ -664,7 +689,7 @@ def _apply_update(op: PatchOperation, file_ops: Any) -> Tuple[bool, str, Optiona return False, ( f"Addition-only hunk: context hint '{hunk.context_hint}' is ambiguous " f"({occurrences} occurrences) — provide a more unique hint" - ), None + ), None, None else: hint_pos = new_content.find(hunk.context_hint) # Insert after the line containing the context hint @@ -676,10 +701,21 @@ def _apply_update(op: PatchOperation, file_ops: Any) -> Tuple[bool, str, Optiona else: new_content = new_content.rstrip('\n') + '\n' + insert_text + '\n' - # Write new content - write_result = file_ops.write_file(op.file_path, new_content) + # Write new content — pass current_content (already read above) to avoid + # a redundant cat subprocess inside write_file. Fall back to the + # two-argument form when the file_ops implementation doesn't accept + # ``pre_content`` (duck-typed callers that only implement the basic + # ``write_file(path, content)`` contract). Feature-detect via the + # signature instead of catching TypeError around the call: a TypeError + # raised *inside* a pre_content-capable write_file must propagate, not + # trigger a second (double) write. + if _write_file_accepts_pre_content(file_ops): + write_result = file_ops.write_file(op.file_path, new_content, + pre_content=current_content) + else: + write_result = file_ops.write_file(op.file_path, new_content) if write_result.error: - return False, write_result.error, None + return False, write_result.error, None, None # Generate diff diff_lines = difflib.unified_diff( @@ -690,4 +726,4 @@ def _apply_update(op: PatchOperation, file_ops: Any) -> Tuple[bool, str, Optiona ) diff = ''.join(diff_lines) - return True, diff, getattr(write_result, "lsp_diagnostics", None) + return True, diff, getattr(write_result, "lsp_diagnostics", None), getattr(write_result, "lint", None) diff --git a/tools/registry.py b/tools/registry.py index e66d4053ce..c92f8a0b9b 100644 --- a/tools/registry.py +++ b/tools/registry.py @@ -346,6 +346,30 @@ def invalidate_check_fn_cache() -> None: _check_fn_last_good.clear() +def get_cached_check_fn_result(fn: Callable) -> Optional[bool]: + """Return the current cached verdict for *fn* if its TTL is still valid. + + Unlike :func:`_check_fn_cached`, this NEVER executes the probe. It is for + read-only surfaces (e.g. dashboard status panels) that need the last-known + availability without triggering network / auth / SDK work inside a request + path. Returns ``None`` when there is no fresh cached verdict. + """ + now = time.monotonic() + scope = check_fn_cache_scope() + if scope == CHECK_FN_CACHE_BYPASS: + # Unresolved profile identity bypasses the cache entirely; there is no + # trustworthy cached verdict to report. + return None + with _check_fn_cache_lock: + cached = _check_fn_cache.get((fn, scope)) + if cached is None: + return None + ts, value = cached + if now - ts < _CHECK_FN_TTL_SECONDS: + return value + return None + + class ToolRegistry: """Singleton registry that collects tool schemas + handlers from tool files.""" diff --git a/tools/session_search_tool.py b/tools/session_search_tool.py index be063a83e4..1c15aca5bf 100644 --- a/tools/session_search_tool.py +++ b/tools/session_search_tool.py @@ -55,6 +55,18 @@ _DEMOTED_SESSION_SOURCES = ("cron",) # the handful of distinct sessions a typical query returns. _DISCOVER_SCAN_LIMIT = 300 +# Raw FTS rows are only a discovery-plan input. The final response hydrates +# its own anchored message window and bookends after lineage deduplication. +_DISCOVER_SEARCH_FIELDS = ( + "id", + "session_id", + "role", + "snippet", + "source", + "model", + "session_started", +) + # Prefixes that identify generated context-compaction handoff summaries. # These are inserted by agent/context_compressor.py as normal user/assistant # messages but contain machine-generated summary metadata — not user content. @@ -247,6 +259,12 @@ def _shape_message( is added so callers know the payload was bounded. """ raw_content = m.get("content") + if isinstance(raw_content, str) and "\x1b" in raw_content: + # Recalled messages can carry ANSI escape sequences (e.g. archived + # terminal output). Strip them before returning content to the model. + from tools.ansi_strip import strip_ansi + + raw_content = strip_ansi(raw_content) if max_content_len and raw_content and len(raw_content) > max_content_len: content = raw_content[:max_content_len] + "…" truncated = True @@ -693,6 +711,7 @@ def _discover( # of cron rows are still in hand for the demotion pass below. offset=0, sort=sort, + fields=_DISCOVER_SEARCH_FIELDS, ) except Exception as e: logging.error("FTS5 search failed: %s", e, exc_info=True) diff --git a/tools/skills_guard.py b/tools/skills_guard.py index 2ba266e4fd..eeaa7be33e 100644 --- a/tools/skills_guard.py +++ b/tools/skills_guard.py @@ -523,6 +523,11 @@ THREAT_PATTERNS = [ "instructs agent to send data to a URL"), ] +_COMPILED_THREAT_PATTERNS = [ + (re.compile(pattern, re.IGNORECASE), pid, severity, category, description) + for pattern, pid, severity, category, description in THREAT_PATTERNS +] + # Structural limits for skill directories MAX_FILE_COUNT = 50 # skills shouldn't have 50+ files MAX_TOTAL_SIZE_KB = 1024 # 1MB total is suspicious for a skill @@ -594,11 +599,11 @@ def scan_file(file_path: Path, rel_path: str = "") -> List[Finding]: seen = set() # (pattern_id, line_number) for deduplication # Regex pattern matching - for pattern, pid, severity, category, description in THREAT_PATTERNS: + for pattern, pid, severity, category, description in _COMPILED_THREAT_PATTERNS: for i, line in enumerate(lines, start=1): if (pid, i) in seen: continue - if re.search(pattern, line, re.IGNORECASE): + if pattern.search(line): seen.add((pid, i)) matched_text = line.strip() if len(matched_text) > 120: diff --git a/tools/terminal_tool.py b/tools/terminal_tool.py index 4191c48ac2..d929947f41 100644 --- a/tools/terminal_tool.py +++ b/tools/terminal_tool.py @@ -374,10 +374,24 @@ def _check_all_guards(command: str, env_type: str, # Allowlist: characters that can legitimately appear in directory paths. -# Covers alphanumeric, path separators, Windows drive/UNC separators, tilde, -# dot, hyphen, underscore, space, plus, at, equals, and comma. Everything -# else is rejected. -_WORKDIR_SAFE_RE = re.compile(r'^[A-Za-z0-9/\\:_\-.~ +@=,]+$') +# Covers Unicode letters/digits, path separators, Windows drive/UNC separators, +# tilde, dot, hyphen, underscore, space, plus, at, equals, and comma. Shell +# metacharacters remain rejected. This intentionally fixes the old ASCII-only +# guard that blocked perfectly normal workdirs such as Chinese Obsidian vault +# paths while preserving the injection boundary around command execution +# (the cwd is additionally shlex-quoted before it reaches the shell; this +# allowlist is defense-in-depth). +_WORKDIR_SAFE_ASCII_CHARS = frozenset('/\\:_-.~ +@=,') + + +def _is_safe_workdir_char(ch: str) -> bool: + if not ch: + return False + # Reject control characters (including newlines/tabs) and NUL bytes before + # considering Unicode categories. + if ord(ch) < 32 or ord(ch) == 127: + return False + return ch.isalnum() or ch in _WORKDIR_SAFE_ASCII_CHARS def _validate_workdir(workdir: str) -> str | None: @@ -390,15 +404,12 @@ def _validate_workdir(workdir: str) -> str | None: """ if not workdir: return None - if not _WORKDIR_SAFE_RE.match(workdir): - # Find the first offending character for a helpful message. - for ch in workdir: - if not _WORKDIR_SAFE_RE.match(ch): - return ( - f"Blocked: workdir contains disallowed character {repr(ch)}. " - "Use a simple filesystem path without shell metacharacters." - ) - return "Blocked: workdir contains disallowed characters." + for ch in workdir: + if not _is_safe_workdir_char(ch): + return ( + f"Blocked: workdir contains disallowed character {repr(ch)}. " + "Use a simple filesystem path without shell metacharacters." + ) return None diff --git a/tools/tool_search.py b/tools/tool_search.py index e286fa0686..686e0f34ea 100644 --- a/tools/tool_search.py +++ b/tools/tool_search.py @@ -95,7 +95,7 @@ class ToolSearchConfig: listing: str = "auto" # "auto" | "on" | "off" # Absolute cap on the embedded listing, regardless of context size. # Effective budget = min(listing_max_tokens, threshold_pct% of context). - listing_max_tokens: int = 20000 + listing_max_tokens: int = 4000 @classmethod def from_raw(cls, raw: Any) -> "ToolSearchConfig": @@ -143,7 +143,7 @@ class ToolSearchConfig: listing = listing_raw else: listing = "auto" - listing_max_tokens = max(200, min(60000, _safe_int(raw.get("listing_max_tokens"), 20000))) + listing_max_tokens = max(200, min(60000, _safe_int(raw.get("listing_max_tokens"), 4000))) return cls( enabled=enabled, @@ -512,7 +512,7 @@ def _listing_group_label(source_name: str) -> str: def build_catalog_listing( deferrable: List[Dict[str, Any]], *, - max_tokens: int = 20000, + max_tokens: int = 4000, ) -> Optional[str]: """Render a skills-style manifest of the deferred catalog. @@ -545,7 +545,7 @@ def build_catalog_listing( def build_catalog_listing_with_form( deferrable: List[Dict[str, Any]], *, - max_tokens: int = 20000, + max_tokens: int = 4000, ) -> Tuple[Optional[str], str]: """Like :func:`build_catalog_listing` but also reports the form used. @@ -861,6 +861,26 @@ def _format_search_hit(entry: CatalogEntry) -> Dict[str, Any]: } +def _available_source_summary(catalog: List[CatalogEntry]) -> List[Dict[str, Any]]: + """Return a compact, deterministic summary of connected deferred sources. + + Included only when search returns no matches. This gives the model enough + evidence to retry with a source/action query instead of treating a lexical + miss as proof that the capability is unavailable, without adding anything + to the fixed per-turn prompt. + """ + counts: Dict[str, int] = {} + for entry in catalog: + # _listing_group_label already falls back to "other" for empty + # source names, matching the listing path's grouping. + label = _listing_group_label(entry.source_name) + counts[label] = counts.get(label, 0) + 1 + return [ + {"name": name, "tool_count": counts[name]} + for name in sorted(counts) + ] + + def dispatch_tool_search(args: Dict[str, Any], *, current_tool_defs: List[Dict[str, Any]], @@ -881,11 +901,20 @@ def dispatch_tool_search(args: Dict[str, Any], _, deferrable = classify_tools(current_tool_defs) catalog = build_catalog(deferrable) hits = search_catalog(catalog, query, limit=limit) - return json.dumps({ + result: Dict[str, Any] = { "query": query, "total_available": len(catalog), "matches": [_format_search_hit(h) for h in hits], - }, ensure_ascii=False) + } + if not hits and catalog: + result["available_sources"] = _available_source_summary(catalog) + result["hint"] = ( + "No lexical match was found, but the sources above are connected " + "and their tools remain available. Retry tool_search with the " + "service name plus a concrete action or object before concluding " + "the capability is unavailable." + ) + return json.dumps(result, ensure_ascii=False) def dispatch_tool_describe(args: Dict[str, Any], diff --git a/tools/transcription_tools.py b/tools/transcription_tools.py index cb1de1b0d5..4c03fa8235 100644 --- a/tools/transcription_tools.py +++ b/tools/transcription_tools.py @@ -1541,6 +1541,19 @@ def build_local_transcribe_kwargs(stt_config: Optional[Dict[str, Any]] = None) - else: kwargs["vad_filter"] = False + # Push the confidence gate down into faster-whisper itself. Without this the + # library's own internal defaults (no_speech_threshold=0.6, log_prob_ + # threshold=-1.0) drop low-confidence segments BEFORE they reach our + # _is_hallucinated_segment post-filter, so the ``stt.local`` threshold knobs + # were dead for that first gate. Non-English speech decodes at a lower + # avg_logprob, so the English-tuned defaults silently discard whole + # utterances. Mapping the same config values through keeps both gates in + # sync and makes the knobs actually usable. Defaults are unchanged, so + # behavior is identical unless a user tunes them. + no_speech_threshold, log_prob_threshold = _confidence_thresholds(local_cfg) + kwargs["no_speech_threshold"] = no_speech_threshold + kwargs["log_prob_threshold"] = log_prob_threshold + forced_lang = _resolve_stt_language("local", stt_config) if forced_lang: kwargs["language"] = forced_lang diff --git a/tools/tts_tool.py b/tools/tts_tool.py index 67a0f48851..341d94c643 100644 --- a/tools/tts_tool.py +++ b/tools/tts_tool.py @@ -50,6 +50,7 @@ import tempfile import threading import time import uuid +from concurrent.futures import Future, ThreadPoolExecutor from dataclasses import dataclass, field from pathlib import Path from typing import Callable, Dict, Any, Iterator, Optional @@ -3348,6 +3349,98 @@ def _strip_markdown_for_tts(text: str) -> str: return text.strip() +class _SyncSentencePipeline: + """Overlap per-sentence synthesis with playback for non-streaming providers. + + The universal sync fallback used to run strictly serially per sentence — + synthesize, play, and only then start synthesizing the next sentence — so + every sentence boundary added a full synthesis-time of dead air. For local + model providers that cost dominates the conversation: a provider at + real-time-factor ~1 spends as long silent between sentences as it does + speaking. Chunked streamers already avoid this; this closes the same gap + for everyone else (edge, piper, plugin providers, …) without touching the + provider contract. + + Shape: one synthesis worker (single-threaded executor, so sentences are + synthesized FIFO and providers never see concurrent calls from this loop — + same effective concurrency as the serial path) feeding one playback worker + through a small bounded queue. While sentence *n* plays, sentence *n+1* is + already synthesizing. The bound keeps lookahead — and the temp files it + implies — small, and gives natural backpressure to the caller. + + ``synthesize``/``play`` are resolved late (module global / import inside + the worker) so tests that monkeypatch ``text_to_speech_tool`` or + ``tools.voice_mode`` keep working unchanged. + """ + + def __init__(self, stop_event: threading.Event, *, lookahead: int = 2): + self._stop = stop_event + self._queue: "queue.Queue[Optional[tuple[str, Future]]]" = queue.Queue( + maxsize=max(1, lookahead) + ) + self._executor = ThreadPoolExecutor( + max_workers=1, thread_name_prefix="tts-sync-synth" + ) + self._player = threading.Thread( + target=self._drain, name="tts-sync-play", daemon=True + ) + self._player.start() + + def speak(self, cleaned: str) -> None: + """Queue one sentence. Blocks only when the lookahead bound is full.""" + if self._stop.is_set(): + return + future = self._executor.submit(self._synthesize_to_tmp, cleaned) + self._queue.put((cleaned, future)) + + def close(self) -> None: + """Flush queued sentences in order (skipped if stopped), then join.""" + self._queue.put(None) + self._player.join() + self._executor.shutdown(wait=True) + + def _synthesize_to_tmp(self, cleaned: str) -> Optional[str]: + if self._stop.is_set(): + return None + tmp_path = None + try: + fd, tmp_path = tempfile.mkstemp(suffix=".mp3") + os.close(fd) + text_to_speech_tool(text=cleaned, output_path=tmp_path) + return tmp_path + except Exception as exc: + logger.warning("Sync per-sentence TTS synthesis failed: %s", exc) + if tmp_path: + try: + os.unlink(tmp_path) + except OSError: + pass + return None + + def _drain(self) -> None: + while True: + item = self._queue.get() + if item is None: + return + _sentence, future = item + tmp_path = None + try: + tmp_path = future.result() + if (tmp_path and not self._stop.is_set() + and os.path.isfile(tmp_path) + and os.path.getsize(tmp_path) > 0): + from tools.voice_mode import play_audio_file + play_audio_file(tmp_path) + except Exception as exc: + logger.warning("Sync per-sentence TTS failed: %s", exc) + finally: + if tmp_path: + try: + os.unlink(tmp_path) + except OSError: + pass + + def stream_tts_to_speaker( text_queue: queue.Queue, stop_event: threading.Event, @@ -3372,6 +3465,7 @@ def stream_tts_to_speaker( waiting on it (continuous voice mode) know playback is finished. """ tts_done_event.clear() + sync_pipeline: Optional[_SyncSentencePipeline] = None try: output_stream = None @@ -3386,6 +3480,11 @@ def stream_tts_to_speaker( from tools.tts_streaming import SentenceChunker, resolve_streaming_provider streamer = resolve_streaming_provider(tts_config, preferred=provider) + # No chunked streamer: per-sentence sync synthesis, pipelined so the + # next sentence synthesizes while the current one plays (closed in the + # finally block, which flushes anything still queued). + sync_pipeline = _SyncSentencePipeline(stop_event) if streamer is None else None + stream_max_len = 0 if streamer is not None: try: @@ -3640,9 +3739,11 @@ def stream_tts_to_speaker( # Display raw sentence on screen before TTS processing if display_callback is not None: display_callback(sentence) - # No chunked streamer → per-sentence sync synthesis (universal). - if streamer is None: - _speak_via_sync(cleaned) + # No chunked streamer → per-sentence sync synthesis (universal), + # pipelined: this enqueues and returns, so sentence n+1 is already + # synthesizing while sentence n is still playing. + if sync_pipeline is not None: + sync_pipeline.speak(cleaned) return # Truncate very long sentences to the provider's per-request cap. if stream_max_len and len(cleaned) > stream_max_len: @@ -3652,29 +3753,6 @@ def stream_tts_to_speaker( # sentence N+1 is already buffering while sentence N plays. _enqueue_audio(cleaned) - def _speak_via_sync(cleaned: str): - """Synthesize one sentence via the proven sync tool, then block on - playback. No chunked API, but per-*sentence* granularity keeps the - flow conversational for edge and every other non-streaming provider. - """ - tmp_path = None - try: - fd, tmp_path = tempfile.mkstemp(suffix=".mp3") - os.close(fd) - text_to_speech_tool(text=cleaned, output_path=tmp_path) - if (not stop_event.is_set() and os.path.isfile(tmp_path) - and os.path.getsize(tmp_path) > 0): - from tools.voice_mode import play_audio_file - play_audio_file(tmp_path) - except Exception as exc: - logger.warning("Sync per-sentence TTS failed: %s", exc) - finally: - if tmp_path: - try: - os.unlink(tmp_path) - except OSError: - pass - def _align_int16_chunks(chunks, stop_evt): """Yield int16-aligned byte chunks from an iterable.""" leftover = b"" @@ -3758,6 +3836,14 @@ def stream_tts_to_speaker( except Exception as exc: logger.warning("Streaming TTS pipeline error: %s", exc) finally: + # Flush the sync pipeline first: queued sentences finish playing (or + # are skipped when stop_event is set) BEFORE tts_done_event fires, so + # continuous voice mode never reopens the mic over its own voice. + if sync_pipeline is not None: + try: + sync_pipeline.close() + except Exception: + pass # Signal the playback worker that no more audio is coming. This lives # in finally: so an exception in the text pump still sends the sentinel. if streamer is not None and _worker_thread is not None: diff --git a/tui_gateway/compute_host.py b/tui_gateway/compute_host.py index 1f255533bd..d90024557a 100644 --- a/tui_gateway/compute_host.py +++ b/tui_gateway/compute_host.py @@ -19,7 +19,7 @@ import time import uuid from dataclasses import dataclass, field from pathlib import Path -from typing import Any, Callable +from typing import Any, Callable, Collection from agent.interrupt_compat import request_hard_interrupt @@ -124,6 +124,14 @@ def _build_sha() -> str: return "unknown" +# Slice of ``ComputeHost.shutdown``'s budget held back for the post-drain +# finalize. ``HostSupervisor._terminate_pid`` SIGKILLs the host +# ``_SHUTDOWN_TIMEOUT_SECS`` (10s — the same value as ``shutdown``'s default +# ``wait``) after SIGTERM, so a drain allowed to consume the whole budget would +# leave the flush racing that kill and persist nothing at all. +_FLUSH_RESERVE_SECS = 1.0 + + class ComputeHost: def __init__( self, @@ -144,7 +152,10 @@ class ComputeHost: self._boot_id = uuid.uuid4().hex self._progress_counter = 0 self._progress_lock = threading.Lock() - self._turn_futures: set[concurrent.futures.Future] = set() + # Future -> the ``sid`` whose turn it is running. ``shutdown`` needs to + # know *whose* turn is still live, not merely that something is, so that + # it can leave those sessions unfinalized; a bare set cannot answer that. + self._turn_futures: dict[concurrent.futures.Future, str] = {} self._turn_futures_lock = threading.Lock() self._transport = _HostTransport(self.emit) self._heartbeat_secs = ( @@ -167,23 +178,86 @@ class ComputeHost: self._executor.shutdown(wait=False, cancel_futures=True) def shutdown(self, *, reason: str = "shutdown", wait: float = 10.0) -> None: + """Drain in-flight turns, then finalize every session. + + Order matters. ``_finalize_session`` is a one-shot latch: it sets + ``session["_finalized"]`` and every later call returns immediately, so + the flush gets exactly one chance to snapshot the session. Running it + before the drain meant that chance was spent while turns were still + producing output — the tail was unpersistable, ``on_session_end`` fired + with ``interrupted=True`` against a session that was still running, and + the active-session lease was released out from under a live turn. The + drain loop exists precisely so that work survives; finalizing first + defeated it. + + ``_FLUSH_RESERVE_SECS`` of the budget — but never more than half of it, + so a short explicit ``wait`` still gets a real drain — is withheld from + the drain, so the flush still runs when in-flight turns outlast the + window. ``wait`` itself is unchanged, so this adds no shutdown latency + and no new exposure to the supervisor's kill escalation. + + Sessions whose turn is *still running* when the drain deadline expires + are excluded from that flush. Finalizing one would spend its single + latch mid-turn — ``shutdown(wait=False, cancel_futures=True)`` below + does not join the turn — leaving the session permanently + un-finalizable and its active-session lease released out from under + live work: exactly the race the drain exists to close, just moved later. + Leaving them unfinalized keeps them recoverable instead. Sessions with + no live turn finalize here as they always have. + + NOTE: ``server._shutdown_sessions`` is registered via ``atexit`` + (``server.py``) and runs on ``SystemExit`` after ``shutdown()`` + returns. It calls ``_finalize_session`` on any session still in + ``server._sessions`` — including ones skipped here whose turn is + still running, since ``_executor.shutdown(wait=False)`` only cancels + pending futures, not running ones. The orphan path (``os._exit(0)``) + bypasses atexit, so the skip is fully effective there. For the + SIGTERM and stdin_closed paths the atexit handler may re-finalize + skipped sessions; this is a pre-existing issue (the old finalize- + first order had the same atexit interaction) and does not make the + drain-before-finalize reordering worse. A follow-up could gate + ``_shutdown_sessions`` on ``not session.get("_finalized") and not + session.get("running")`` to close the gap. + """ self._closed.set() - self.flush_all_sessions(reason=reason) - deadline = time.monotonic() + max(0.0, wait) - while time.monotonic() < deadline: + budget = max(0.0, wait) + deadline = time.monotonic() + budget - min(_FLUSH_RESERVE_SECS, budget / 2.0) + while True: + remaining = deadline - time.monotonic() + if remaining <= 0: + break with self._turn_futures_lock: pending = [f for f in self._turn_futures if not f.done()] if not pending: break - time.sleep(0.05) + # Bounded by ``remaining``: a flat 0.05s sleep would overshoot the + # deadline and eat into the reserve it is there to protect, which + # for a small ``wait`` can be the whole of it. + time.sleep(min(0.05, remaining)) + with self._turn_futures_lock: + live_sids = {sid for future, sid in self._turn_futures.items() if sid and not future.done()} + self.flush_all_sessions(reason=reason, skip_sids=live_sids) self._executor.shutdown(wait=False, cancel_futures=True) - def flush_all_sessions(self, *, reason: str = "shutdown") -> None: + def flush_all_sessions( + self, + *, + reason: str = "shutdown", + skip_sids: Collection[str] | None = None, + ) -> None: + """Finalize every server session except the ones named in ``skip_sids``. + + ``skip_sids`` carries the sessions whose turn is still live, which must + not spend their one-shot ``_finalize_session`` while running. + """ try: from tui_gateway import server except Exception: return - for session in list(getattr(server, "_sessions", {}).values()): + skip = set(skip_sids or ()) + for sid, session in list(getattr(server, "_sessions", {}).items()): + if sid in skip: + continue try: server._finalize_session(session, end_reason=f"compute_host_{reason}") except Exception: @@ -229,15 +303,28 @@ class ComputeHost: self._sessions[sid] = HostSession(sid=sid, agent=SpikeAgent(sid, list(history))) self.emit({"type": "session.seeded", "sid": sid, "request_id": frame.get("request_id")}) + def _track_turn_future(self, future: concurrent.futures.Future, sid: str) -> None: + """Register an in-flight turn against the session running it. + + The callback has to remove the entry under the lock — a bare + ``dict.pop`` bound method is not the drop-in ``set.discard`` was — or + the mapping grows for the life of the host. + """ + with self._turn_futures_lock: + self._turn_futures[future] = sid + future.add_done_callback(self._untrack_turn_future) + + def _untrack_turn_future(self, future: concurrent.futures.Future) -> None: + with self._turn_futures_lock: + self._turn_futures.pop(future, None) + def _handle_turn_start(self, frame: dict[str, Any]) -> None: sid = str(frame.get("sid") or "") if sid in self._sessions: self._handle_spike_turn_start(frame) return future = self._executor.submit(self._run_real_turn, dict(frame)) - with self._turn_futures_lock: - self._turn_futures.add(future) - future.add_done_callback(self._turn_futures.discard) + self._track_turn_future(future, sid) def _handle_spike_turn_start(self, frame: dict[str, Any]) -> None: sid = str(frame.get("sid") or "") @@ -251,9 +338,7 @@ class ComputeHost: return session.running = True future = self._executor.submit(self._run_spike_turn, session, dict(frame)) - with self._turn_futures_lock: - self._turn_futures.add(future) - future.add_done_callback(self._turn_futures.discard) + self._track_turn_future(future, sid) def _handle_interrupt(self, frame: dict[str, Any]) -> None: sid = str(frame.get("sid") or "") diff --git a/tui_gateway/entry.py b/tui_gateway/entry.py index cfd714b258..0a44c32c6d 100644 --- a/tui_gateway/entry.py +++ b/tui_gateway/entry.py @@ -455,6 +455,17 @@ def main(): # Live-apply skins Hermes activates mid-conversation. server._ensure_skin_watcher() + # Warm the /model picker's provider-models cache off-thread during this + # idle window (gateway.ready sent, user about to type). Mirrors the classic + # CLI run() loop — the stdio TUI otherwise never prewarms, so the first + # /model open blocks on serial /v1/models fetches. Fire-and-forget, + # guarded once-per-process, fully exception-isolated. + try: + from hermes_cli.model_switch import prewarm_picker_cache_async + prewarm_picker_cache_async() + except Exception: + logger.debug("picker cache prewarm (tui) failed to start", exc_info=True) + while True: raw = sys.stdin.readline() if not raw: diff --git a/tui_gateway/methods_session.py b/tui_gateway/methods_session.py index 75e473426b..1a3e09f16e 100644 --- a/tui_gateway/methods_session.py +++ b/tui_gateway/methods_session.py @@ -316,6 +316,10 @@ def _(rid, params: dict) -> dict: # local profile's state.db. None/own profile → the launch profile (unchanged). profile = (params.get("profile") or "").strip() or None profile_home = _profile_home(profile) + # Desktop hydrates persisted transcripts through the authenticated REST + # route in parallel. Suppress the duplicate WebSocket transcript only when + # the caller explicitly requests it; other clients keep upstream behavior. + omit_messages = is_truthy_value(params.get("omit_messages", False)) # In a profile scope, the agent OWNS a long-lived db handle bound to that # profile (do NOT auto-close it here). Otherwise reuse the shared launch db. @@ -380,6 +384,7 @@ def _(rid, params: dict) -> dict: cols=cols, touch=True, transport=current_transport() or _stdio_transport, + omit_messages=omit_messages, ) payload["resumed"] = target # A lazy watch session never owns a run loop, so its payload's running @@ -450,14 +455,15 @@ def _(rid, params: dict) -> dict: except Exception: logger.debug("child-watch display projection read failed", exc_info=True) display_history = history - messages = _history_to_messages(display_history) + messages = [] if omit_messages else _history_to_messages(display_history) return _ok( rid, { "session_id": sid, "resumed": target, - "message_count": len(messages), + "message_count": len(display_history) if omit_messages else len(messages), "messages": messages, + "messages_omitted": omit_messages, "info": _lazy_resume_info(cwd, profile=profile), "inflight": None, "running": child_running, @@ -494,7 +500,13 @@ def _(rid, params: dict) -> dict: # (raw_history → sanitize_replay_history → the resumed session's # working conversation) and the display copy stays verbatim — # inspection/export must show what is actually stored. - raw_history, display_history = db.get_resume_conversations(target) + if omit_messages: + raw_history = db.get_messages_as_conversation( + target, repair_alternation=True + ) + display_history = [] + else: + raw_history, display_history = db.get_resume_conversations(target) except Exception as e: if lease is not None: lease.release() @@ -502,7 +514,7 @@ def _(rid, params: dict) -> dict: # Display keeps the full transcript; the model-fed history drops a # dangling/interrupted tool-call tail so a session killed mid-loop does # not replay the unanswered call forever (#29086). - prefix = db.get_ancestor_display_prefix(target) + prefix = [] if omit_messages else db.get_ancestor_display_prefix(target) history = sanitize_replay_history(raw_history) # Restore the model/provider/reasoning/tier this chat last used so the # deferred build (and the info below) match the eager path — without them @@ -530,12 +542,13 @@ def _(rid, params: dict) -> dict: _schedule_session_cap_enforcement() # trim detached idle sessions over the cap auto_continue = _maybe_schedule_auto_continue(sid, record, target) - messages = _history_to_messages(display_history) + messages = [] if omit_messages else _history_to_messages(display_history) payload = { "session_id": sid, "resumed": target, - "message_count": len(messages), + "message_count": len(raw_history) if omit_messages else len(messages), "messages": messages, + "messages_omitted": omit_messages, "info": _lazy_resume_info( cwd, model=model_override.get("model") or "", @@ -573,7 +586,13 @@ def _(rid, params: dict) -> dict: # One lineage SELECT feeds both projections (see the interactive resume # above): the model-fed copy is alternation-repaired for LIVE REPLAY, the # display copy stays verbatim. - raw_history, display_history = db.get_resume_conversations(target) + if omit_messages: + raw_history = db.get_messages_as_conversation( + target, repair_alternation=True + ) + display_history = [] + else: + raw_history, display_history = db.get_resume_conversations(target) # The display transcript keeps every row so the user still sees their # full history. The model-fed history is sanitized: a session whose # last turn died mid-tool-loop persists a dangling assistant(tool_calls) @@ -581,9 +600,11 @@ def _(rid, params: dict) -> dict: # re-issue the unanswered call forever — the permanent-"thinking" stuck # session in #29086. The messaging gateway already strips this; this is # the WebUI/TUI resume path picking up the same cleanup. - display_history_prefix = db.get_ancestor_display_prefix(target) + display_history_prefix = ( + [] if omit_messages else db.get_ancestor_display_prefix(target) + ) history = sanitize_replay_history(raw_history) - messages = _history_to_messages(display_history) + messages = [] if omit_messages else _history_to_messages(display_history) tokens = _set_session_context(target) try: # Pass the profile's db so the agent persists turns to the right @@ -632,6 +653,7 @@ def _(rid, params: dict) -> dict: cols=cols, touch=True, transport=current_transport() or _stdio_transport, + omit_messages=omit_messages, ) payload["resumed"] = target return _ok(rid, payload) @@ -685,8 +707,9 @@ def _(rid, params: dict) -> dict: payload = { "session_id": sid, "resumed": target, - "message_count": len(messages), + "message_count": len(raw_history) if omit_messages else len(messages), "messages": messages, + "messages_omitted": omit_messages, "info": _session_info(agent, session), "inflight": None, "running": False, @@ -782,6 +805,7 @@ def _(rid, params: dict) -> dict: session, touch=True, transport=current_transport() or _stdio_transport, + omit_messages=is_truthy_value(params.get("omit_messages", False)), ), ) @@ -2625,15 +2649,23 @@ def _(rid, params: dict) -> dict: else None ), ) - for msg in history: - db.append_message( - session_id=new_key, - role=msg.get("role", "user"), - content=msg.get("content"), - # Preserve the parent's original message timestamps — - # branch copies are history, not new activity (9d73006ad). - timestamp=msg.get("timestamp"), - ) + # Copy the whole parent history in bounded-chunk transactions — + # a branch seed can be hundreds of rows, and per-row transactions + # were the write-amplification pattern removed in #23254. + db.append_messages_batch( + new_key, + [ + { + "role": msg.get("role", "user"), + "content": msg.get("content"), + # Preserve the parent's original message timestamps — + # branch copies are history, not new activity (9d73006ad). + "timestamp": msg.get("timestamp"), + } + for msg in history + ], + chunk_rows=500, + ) db.set_session_title(new_key, title) except Exception as e: if lease is not None: diff --git a/tui_gateway/server.py b/tui_gateway/server.py index d637fb2f3c..9d5fd00ce7 100644 --- a/tui_gateway/server.py +++ b/tui_gateway/server.py @@ -2749,16 +2749,26 @@ def _persist_branch_seed(session: dict) -> None: if db is None: return try: - for msg in seed: - db.append_message( - session_id=key, - role=msg.get("role", "user"), - content=msg.get("content"), - # Preserve the parent's original message timestamps — - # append_message would otherwise stamp time.time() and the - # branch's copied history would all appear authored "now". - timestamp=msg.get("timestamp"), - ) + # Bounded-chunk transactions (see #23254): a branch seed can be + # hundreds of rows; chunking keeps each BEGIN IMMEDIATE short so + # concurrent writers aren't starved. Recovery semantics match the + # old per-row loop (mid-copy failure leaves a partial seed with + # _branch_seed_persisted unset). + db.append_messages_batch( + key, + [ + { + "role": msg.get("role", "user"), + "content": msg.get("content"), + # Preserve the parent's original message timestamps — + # append_message would otherwise stamp time.time() and the + # branch's copied history would all appear authored "now". + "timestamp": msg.get("timestamp"), + } + for msg in seed + ], + chunk_rows=500, + ) session["_branch_seed_persisted"] = True except Exception as exc: from hermes_state import is_disk_full_error @@ -7913,6 +7923,7 @@ def _live_session_payload( cols: int | None = None, touch: bool = False, transport: Transport | None = None, + omit_messages: bool = False, ) -> dict: with session["history_lock"]: if cols is not None: @@ -7930,11 +7941,16 @@ def _live_session_payload( # Prefer the persisted display lineage (candidate-inclusive) so this payload # matches the eager session.resume + REST transcript; the DB has its own # lock, so read it outside the session history lock. - history = _live_visible_history(session, _get_db(), in_memory_history) + history = ( + in_memory_history + if omit_messages + else _live_visible_history(session, _get_db(), in_memory_history) + ) payload = { "info": _fallback_session_info(session), "message_count": len(history), - "messages": _history_to_messages(history), + "messages": [] if omit_messages else _history_to_messages(history), + "messages_omitted": omit_messages, "running": running, "session_id": sid, "session_key": _session_lookup_key(session, fallback=sid), @@ -8904,6 +8920,20 @@ def _collect_kanban_notifications(session: dict) -> list: if resolved in seen_db_paths: continue seen_db_paths.add(resolved) + # A poller runs per live TUI/Desktop session. Avoid opening this board + # writable unless it has a subscription owned by this exact session; + # subscriptions for gateways or other sessions are not actionable here. + try: + if _kb.count_notify_subs( + board=slug, + platform="tui", + chat_id=session_key, + ) == 0: + continue + except Exception: + # Preserve delivery if the read-only probe cannot inspect a + # locked, corrupt, or otherwise unusual database. + pass try: conn = _kb.connect(board=slug) except Exception: @@ -9323,8 +9353,6 @@ def _run_prompt_submit( ): session["running"] = False return - history = list(session["history"]) - history_version = int(session.get("history_version", 0)) if image_paths is None: images = list(session.get("attached_images", [])) session["attached_images"] = [] @@ -9404,6 +9432,11 @@ def _run_prompt_submit( # config sync so an explicit pick wins over a config.yaml change. _apply_pending_model_switch(sid, session) _sync_agent_model_with_config(sid, session) + # Snapshot after turn-start model sync. A deferred switch mutates + # history and its version; that mutation belongs to this turn. + with session["history_lock"]: + history = list(session["history"]) + history_version = int(session.get("history_version", 0)) cwd = _session_cwd(session) _register_session_cwd(session) cols = session.get("cols", 80) @@ -9686,24 +9719,59 @@ def _run_prompt_submit( session["history"] = result["messages"] session["history_version"] = history_version + 1 else: - # History mutated externally during the turn - # (undo/compress/retry/rollback now guard on - # session.running, but this is the defensive - # backstop for any path that slips past). - # Surface the desync rather than silently - # dropping the agent's output — the UI can - # show the response and warn that it was - # not persisted. - print( - f"[tui_gateway] prompt.submit: history_version mismatch " - f"(expected={history_version} current={current_version}) — " - f"agent output NOT written to session history", - file=sys.stderr, - ) - status_note = ( - "History changed during this turn — the response above is visible " - "but was not saved to session history." + # History mutated externally during the turn. + # Check if the only mutation was a model-switch + # marker inserted mid-turn (#76870). If so the + # agent output is still valid — merge it into the + # current history that now contains the marker. + # + # _append_model_switch_marker strips prior markers + # in-place then appends a new one, so the delta + # is NOT a simple tail-slice — we must compare + # content, not indices. + current_history = list(session["history"]) + history_no_markers = [ + e for e in history if not _is_model_switch_marker(e) + ] + current_no_markers = [ + e for e in current_history if not _is_model_switch_marker(e) + ] + model_switch_only = ( + current_no_markers == history_no_markers + and any( + _is_model_switch_marker(e) + for e in current_history + ) ) + if model_switch_only: + # The agent's new messages start after the + # turn-start history. Guard against + # auto-compression making result["messages"] + # shorter than history (#77274 review). + if len(result["messages"]) > len(history): + new_messages = result["messages"][len(history):] + else: + # Compression rebound the messages list — + # use the full result as the base. + new_messages = list(result["messages"]) + session["history"] = current_history + new_messages + session["history_version"] = current_version + 1 + else: + # Genuine desync (undo/compress/retry/rollback). + # Surface the desync rather than silently + # dropping the agent's output — the UI can + # show the response and warn that it was + # not persisted. + print( + f"[tui_gateway] prompt.submit: history_version mismatch " + f"(expected={history_version} current={current_version}) — " + f"agent output NOT written to session history", + file=sys.stderr, + ) + status_note = ( + "History changed during this turn — the response above is visible " + "but was not saved to session history." + ) # If auto-compression fired inside run_conversation(), agent.session_id # may have rotated. Sync session_key before downstream title/goal/finalize diff --git a/tui_gateway/ws.py b/tui_gateway/ws.py index b3b6c4bcc1..073c4ac149 100644 --- a/tui_gateway/ws.py +++ b/tui_gateway/ws.py @@ -303,6 +303,14 @@ async def handle_ws(ws: Any) -> None: transport = WSTransport(ws, asyncio.get_running_loop(), peer=peer) + # resolve_skin() reads config + initializes the skin engine — + # synchronous I/O + CPU work that should not block the event loop + # during the cold-start window. Run it in the thread pool so the + # WS read loop stays free to drain the frontend's initial RPC + # burst (setup.status, session.list, ...) without a stall + # (#60800). The skin payload is small (a dict of strings/arrays), + # so the to_thread overhead is negligible. + skin_payload = await asyncio.to_thread(server.resolve_skin) ready_ok = await transport.write_async( { "jsonrpc": "2.0", @@ -312,7 +320,7 @@ async def handle_ws(ws: Any) -> None: # change_events: this backend broadcasts pet.changed / # cron.changed / sessions.changed, so clients can demote # their legacy polls to slow backstops. - "payload": {"skin": server.resolve_skin(), "change_events": True}, + "payload": {"skin": skin_payload, "change_events": True}, }, } ) diff --git a/ui-tui/packages/hermes-ink/src/ink/components/ScrollBox.tsx b/ui-tui/packages/hermes-ink/src/ink/components/ScrollBox.tsx index 4f2604be0e..456ea5fceb 100644 --- a/ui-tui/packages/hermes-ink/src/ink/components/ScrollBox.tsx +++ b/ui-tui/packages/hermes-ink/src/ink/components/ScrollBox.tsx @@ -10,9 +10,31 @@ import { markCommitStart } from '../reconciler.js' import type { Styles } from '../styles.js' import Box from './Box.js' + +const MAX_SCROLL_GEOMETRY = 1_000_000_000 + +const validUnsignedGeometry = (value: number): boolean => + Number.isFinite(value) && value >= 0 && value <= MAX_SCROLL_GEOMETRY + +const validClampMaximum = (value: number): boolean => value === Number.POSITIVE_INFINITY || validUnsignedGeometry(value) + +const validSignedGeometry = (value: number): boolean => Number.isFinite(value) && Math.abs(value) <= MAX_SCROLL_GEOMETRY + +const safeUnsignedGeometry = (value: number | undefined): number => + value !== undefined && validUnsignedGeometry(value) ? value : 0 + +const safeSignedGeometry = (value: number | undefined): number => + value !== undefined && validSignedGeometry(value) ? value : 0 + export type ScrollBoxHandle = { scrollTo: (y: number) => void scrollBy: (dy: number) => void + /** + * Offset the committed viewport after content above it changes height. + * Unlike scrollTo, this preserves pending input, sticky state, anchor seeks, + * and the manual-scroll timestamp. + */ + adjustScrollTop: (dy: number) => void /** * Scroll so `el`'s top is at the viewport top (plus `offset`). Unlike * scrollTo which bakes a number that's stale by the time the throttled @@ -49,9 +71,9 @@ export type ScrollBoxHandle = { isSticky: () => boolean /** * Subscribe to scroll viewport changes. Fires for imperative scroll changes - * (scrollTo/scrollBy/scrollToBottom) and for renderer-computed scroll bounds - * changes such as content growth or terminal resize. Callers use this to - * keep virtualized ranges aligned with the currently visible viewport. + * (scrollTo/scrollBy/adjustScrollTop/scrollToBottom) and for renderer-computed + * scroll bounds changes such as content growth or terminal resize. Callers + * use this to keep virtualized ranges aligned with the visible viewport. */ subscribe: (listener: () => void) => () => void /** @@ -85,7 +107,7 @@ export type ScrollBoxProps = Except): React.ReactNode { const domRef = useRef(null) - // scrollTo/scrollBy bypass React: they mutate scrollTop on the DOM node, + // Imperative position changes bypass React: they mutate scrollTop on the DOM node, // mark it dirty, and call the root's throttled scheduleRender directly. // The Ink renderer reads scrollTop from the node — no React state needed, // no reconciler overhead per wheel event. The microtask defer coalesces @@ -127,10 +149,29 @@ function ScrollBox({ children, ref, stickyScroll, ...style }: PropsWithChildren< useImperativeHandle( ref, (): ScrollBoxHandle => ({ + adjustScrollTop(dy: number) { + const el = domRef.current + + if (!el || !validSignedGeometry(dy)) { + return + } + + const current = safeUnsignedGeometry(el.scrollTop) + const next = Math.max(0, current + Math.floor(dy)) + const compensation = safeSignedGeometry(el.scrollTopCompensation) + (next - current) + + if (next === current || !validUnsignedGeometry(next) || !validSignedGeometry(compensation)) { + return + } + + el.scrollTop = next + el.scrollTopCompensation = compensation + scrollMutated(el) + }, scrollTo(y: number) { const el = domRef.current - if (!el) { + if (!el || !validSignedGeometry(y)) { return } @@ -139,6 +180,7 @@ function ScrollBox({ children, ref, stickyScroll, ...style }: PropsWithChildren< el.stickyScroll = false manualScrollAtRef.current = Date.now() el.pendingScrollDelta = undefined + el.scrollTopCompensation = undefined el.scrollAnchor = undefined el.scrollTop = Math.max(0, Math.floor(y)) scrollMutated(el) @@ -146,13 +188,14 @@ function ScrollBox({ children, ref, stickyScroll, ...style }: PropsWithChildren< scrollToElement(el: DOMElement, offset = 0) { const box = domRef.current - if (!box) { + if (!box || !validSignedGeometry(offset)) { return } box.stickyScroll = false manualScrollAtRef.current = Date.now() box.pendingScrollDelta = undefined + box.scrollTopCompensation = undefined box.scrollAnchor = { el, offset @@ -162,14 +205,20 @@ function ScrollBox({ children, ref, stickyScroll, ...style }: PropsWithChildren< scrollBy(dy: number) { const el = domRef.current - if (!el) { + if (!el || !validSignedGeometry(dy)) { + return + } + + const pending = safeSignedGeometry(el.pendingScrollDelta) + Math.floor(dy) + + if (!validSignedGeometry(pending)) { return } el.stickyScroll = false manualScrollAtRef.current = Date.now() el.scrollAnchor = undefined - el.pendingScrollDelta = (el.pendingScrollDelta ?? 0) + Math.floor(dy) + el.pendingScrollDelta = pending scrollMutated(el) }, scrollToBottom() { @@ -180,30 +229,35 @@ function ScrollBox({ children, ref, stickyScroll, ...style }: PropsWithChildren< } el.pendingScrollDelta = undefined + el.scrollTopCompensation = undefined el.stickyScroll = true markDirty(el) notify() forceRender(n => n + 1) }, getScrollTop() { - return domRef.current?.scrollTop ?? 0 + return safeUnsignedGeometry(domRef.current?.scrollTop) }, getPendingDelta() { // Accumulated-but-not-yet-drained delta. useVirtualScroll needs // this to mount the union [committed, committed+pending] range — // otherwise intermediate drain frames find no children (blank). - return domRef.current?.pendingScrollDelta ?? 0 + return safeSignedGeometry(domRef.current?.pendingScrollDelta) }, getScrollHeight() { - return domRef.current?.scrollHeight ?? 0 + return safeUnsignedGeometry(domRef.current?.scrollHeight) }, getFreshScrollHeight() { const content = domRef.current?.childNodes[0] as DOMElement | undefined - return content?.yogaNode?.getComputedHeight() ?? domRef.current?.scrollHeight ?? 0 + const height = content?.yogaNode?.getComputedHeight() + + return validUnsignedGeometry(height ?? Number.NaN) + ? height! + : safeUnsignedGeometry(domRef.current?.scrollHeight) }, getViewportHeight() { - return domRef.current?.scrollViewportHeight ?? 0 + return safeUnsignedGeometry(domRef.current?.scrollViewportHeight) }, getViewportTop() { return domRef.current?.scrollViewportTop ?? 0 @@ -232,6 +286,26 @@ function ScrollBox({ children, ref, stickyScroll, ...style }: PropsWithChildren< return } + if (min === undefined && max === undefined) { + el.scrollClampMin = undefined + el.scrollClampMax = undefined + + return + } + + if ( + min === undefined || + max === undefined || + !validUnsignedGeometry(min) || + !validClampMaximum(max) || + min > max + ) { + el.scrollClampMin = undefined + el.scrollClampMax = undefined + + return + } + el.scrollClampMin = min el.scrollClampMax = max } @@ -260,7 +334,7 @@ function ScrollBox({ children, ref, stickyScroll, ...style }: PropsWithChildren< domRef.current = el if (el) { - el.scrollTop ??= 0 + el.scrollTop = safeUnsignedGeometry(el.scrollTop) el.notifyScrollChange = notify } }} diff --git a/ui-tui/packages/hermes-ink/src/ink/dom.ts b/ui-tui/packages/hermes-ink/src/ink/dom.ts index 69fc74fff4..834532abc6 100644 --- a/ui-tui/packages/hermes-ink/src/ink/dom.ts +++ b/ui-tui/packages/hermes-ink/src/ink/dom.ts @@ -53,6 +53,11 @@ export type DOMElement = { // intermediate frames instead of one big jump. Direction reversal // naturally cancels (pure accumulator, no target tracking). pendingScrollDelta?: number + // One-render record of additive scrollTop changes made to preserve the + // visual anchor after content above the viewport changes height. The + // renderer subtracts this when evaluating positional bottom-follow and + // defers pending input for that paint, then clears the record. + scrollTopCompensation?: number // Render-time clamp bounds for virtual scroll. useVirtualScroll writes // the currently-mounted children's coverage span; render-node-to-output // clamps scrollTop to stay within it. Prevents blank screen when diff --git a/ui-tui/packages/hermes-ink/src/ink/render-border.test.ts b/ui-tui/packages/hermes-ink/src/ink/render-border.test.ts new file mode 100644 index 0000000000..a6373907d3 --- /dev/null +++ b/ui-tui/packages/hermes-ink/src/ink/render-border.test.ts @@ -0,0 +1,200 @@ +import { describe, expect, it } from 'vitest' + +import type { DOMNode } from './dom.js' +import Output from './output.js' +import renderBorder from './render-border.js' +import { cellAt, CellWidth, CharPool, createScreen, HyperlinkPool, type Screen, StylePool } from './screen.js' + +const WIDTH = 12 +const HEIGHT = 6 + +function createOutput() { + const stylePool = new StylePool() + const screen = createScreen(WIDTH, HEIGHT, stylePool, new CharPool(), new HyperlinkPool()) + + return { output: new Output({ height: HEIGHT, screen, stylePool, width: WIDTH }), stylePool } +} + +function borderNode(style: Record = {}, width = 8, height = 4): DOMNode { + return { + style: { borderStyle: 'single', ...style }, + yogaNode: { + getComputedHeight: () => height, + getComputedWidth: () => width + } + } as unknown as DOMNode +} + +function snapshot(screen: Screen, styles: StylePool) { + return Array.from({ length: screen.height }, (_, y) => + Array.from({ length: screen.width }, (_, x) => { + const cell = cellAt(screen, x, y)! + + return [cell.char, cell.width, styles.get(cell.styleId).map(code => code.code)] as const + }) + ) +} + +function paint( + node: DOMNode, + visible: readonly [number, number, number, number] = [0, WIDTH, 0, HEIGHT], + decorate?: (output: Output, phase: 'before' | 'after') => void +) { + const { output, stylePool } = createOutput() + + decorate?.(output, 'before') + renderBorder(1, 1, node, output, ...visible) + decorate?.(output, 'after') + + return { cells: snapshot(output.get(), stylePool), stylePool } +} + +function expectClippedParity( + full: ReturnType['cells'], + clipped: ReturnType['cells'], + [x1, x2, y1, y2]: readonly [number, number, number, number] +) { + for (let y = 0; y < HEIGHT; y++) { + for (let x = 0; x < WIDTH; x++) { + if (x >= x1 && x < x2 && y >= y1 && y < y2) { + expect(clipped[y]![x]).toEqual(full[y]![x]) + } else { + expect(clipped[y]![x]![0]).toBe(' ') + } + } + } +} + +function expectWideTextClippedParity( + full: ReturnType['cells'], + clipped: ReturnType['cells'], + visible: readonly [number, number, number, number] +) { + const [x1, x2, y1, y2] = visible + + for (let y = 0; y < HEIGHT; y++) { + for (let x = 0; x < WIDTH; x++) { + const fullCell = full[y]![x]! + const isVisible = x >= x1 && x < x2 && y >= y1 && y < y2 + const clipsWideHead = isVisible && fullCell[1] === CellWidth.Wide && x + 1 >= x2 + const clipsWideTail = isVisible && fullCell[1] === CellWidth.SpacerTail && x - 1 < x1 + + if (clipsWideHead || clipsWideTail) { + expect(clipped[y]![x]).toEqual([' ', CellWidth.Narrow, []]) + } else if (isVisible) { + expect(clipped[y]![x]).toEqual(fullCell) + } else { + expect(clipped[y]![x]![0]).toBe(' ') + } + } + } +} + +describe('renderBorder viewport parity', () => { + it.each([ + ['all edges', {}], + ['disabled edges', { borderLeft: false, borderTop: false }] + ])('matches the unclipped border inside a clipped viewport with %s', (_name, style) => { + const visible = [3, 8, 1, 4] as const + const full = paint(borderNode(style)).cells + const clipped = paint(borderNode(style), visible).cells + + expectClippedParity(full, clipped, visible) + }) + + it.each([ + ['left edge at wide title', [3, 9, 1, 2]], + ['left edge bisects wide title', [4, 9, 1, 2]], + ['left edge after wide grapheme', [5, 9, 1, 2]], + ['right edge bisects wide title', [1, 4, 1, 2]], + ['right edge after wide grapheme', [1, 5, 1, 2]], + ['right edge after narrow suffix', [1, 6, 1, 2]] + ] as const)('preserves ANSI wide-title cell coordinates when the %s', (_name, visible) => { + const node = borderNode({ + borderText: { + align: 'center', + content: '\u001B[31m界A\u001B[39m', + position: 'top' + } + }) + + const full = paint(node).cells + const clipped = paint(node, visible).cells + + expectWideTextClippedParity(full, clipped, visible) + expect(full[1]!.some(([char]) => char === '界')).toBe(true) + expect(full[1]!.some(([char, , codes]) => char === '界' && codes.includes('\u001B[31m'))).toBe(true) + expect(full[1]![6]![2]).not.toContain('\u001B[31m') + + if (visible[0] === 4) { + expect(clipped[1]![4]).toEqual([' ', CellWidth.Narrow, []]) + expect(clipped[1]![5]![0]).toBe('A') + } + }) + + it.each([ + ['left edge inside title', [4, 9, 1, 2]], + ['right edge inside title', [1, 5, 1, 2]] + ] as const)('preserves narrow ANSI-title parity when the %s', (_name, visible) => { + const node = borderNode({ + borderText: { + align: 'center', + content: '\u001B[32mABC\u001B[39m', + position: 'top' + } + }) + + const full = paint(node).cells + const clipped = paint(node, visible).cells + + expectClippedParity(full, clipped, visible) + }) + + it('intersects nested output clips without changing surviving border cells', () => { + const node = borderNode() + const full = paint(node).cells + const { output, stylePool } = createOutput() + + output.clip({ x1: 0, x2: 10, y1: 0, y2: 5 }) + output.clip({ x1: 1, x2: 7, y1: 1, y2: 4 }) + renderBorder(1, 1, node, output) + output.unclip() + output.unclip() + + const nested = snapshot(output.get(), stylePool) + + expectClippedParity(full, nested, [1, 7, 1, 4]) + }) + + it('renders borders after opaque fills and preserves the fill interior', () => { + const decorated = paint(borderNode(), [0, WIDTH, 0, HEIGHT], (output, phase) => { + if (phase === 'before') { + const line = '\u001B[44m \u001B[49m' + + output.write(1, 1, [line, line, line, line].join('\n')) + } + }) + + expect(decorated.cells[1]![1]![0]).toBe('┌') + expect(decorated.cells[4]![8]![0]).toBe('┘') + expect(decorated.cells[2]![2]![2]).toContain('\u001B[44m') + }) + + it('keeps clipped border parity when an absolute-style overlay paints afterward', () => { + const overlay = (output: Output, phase: 'before' | 'after') => { + if (phase === 'after') { + output.clip({ x1: 4, x2: 6, y1: 1, y2: 2 }) + output.write(4, 1, 'OV') + output.unclip() + } + } + + const visible = [3, 7, 1, 4] as const + const full = paint(borderNode(), [0, WIDTH, 0, HEIGHT], overlay).cells + const clipped = paint(borderNode(), visible, overlay).cells + + expectClippedParity(full, clipped, visible) + expect(clipped[1]![4]![0]).toBe('O') + expect(clipped[1]![5]![0]).toBe('V') + }) +}) diff --git a/ui-tui/packages/hermes-ink/src/ink/render-border.ts b/ui-tui/packages/hermes-ink/src/ink/render-border.ts index a4fff7cb50..526729422d 100644 --- a/ui-tui/packages/hermes-ink/src/ink/render-border.ts +++ b/ui-tui/packages/hermes-ink/src/ink/render-border.ts @@ -1,6 +1,8 @@ import chalk from 'chalk' import cliBoxes, { type Boxes, type BoxStyle } from 'cli-boxes' +import sliceAnsi from '../utils/sliceAnsi.js' + import { applyColor } from './colorize.js' import type { DOMNode } from './dom.js' import type Output from './output.js' @@ -30,41 +32,6 @@ export const CUSTOM_BORDER_STYLES = { export type BorderStyle = keyof Boxes | keyof typeof CUSTOM_BORDER_STYLES | BoxStyle -function embedTextInBorder( - borderLine: string, - text: string, - align: 'start' | 'end' | 'center', - offset: number = 0, - borderChar: string -): [before: string, text: string, after: string] { - const textLength = stringWidth(text) - const borderLength = borderLine.length - - if (textLength >= borderLength - 2) { - return ['', text.substring(0, borderLength), ''] - } - - let position: number - - if (align === 'center') { - position = Math.floor((borderLength - textLength) / 2) - } else if (align === 'start') { - position = offset + 1 // +1 to account for corner character - } else { - // align === 'end' - position = borderLength - textLength - offset - 1 // -1 for corner character - } - - // Ensure position is valid - position = Math.max(1, Math.min(position, borderLength - textLength - 1)) - - const before = borderLine.substring(0, 1) + borderChar.repeat(position - 1) - - const after = borderChar.repeat(borderLength - position - textLength - 1) + borderLine.substring(borderLength - 1) - - return [before, text, after] -} - function styleBorderLine(line: string, color: Color | undefined, dim: boolean | undefined): string { let styled = applyColor(line, color) @@ -75,7 +42,138 @@ function styleBorderLine(line: string, color: Color | undefined, dim: boolean | return styled } -const renderBorder = (x: number, y: number, node: DOMNode, output: Output): void => { +function sliceAnsiToWidth(text: string, start: number, end: number): { leadingColumns: number; text: string } { + // sliceAnsi omits a wide grapheme bisected by `start`; retain its partial + // cell as blank space so later graphemes keep their source coordinates. + const leadingColumns = Math.max(0, stringWidth(sliceAnsi(text, 0, start)) - start) + let sliced = sliceAnsi(text, start, end) + + if (stringWidth(sliced) > end - start - leadingColumns) { + sliced = sliceAnsi(text, start, end - 1) + } + + return { leadingColumns, text: sliced } +} + +function borderRun( + start: number, + end: number, + borderLength: number, + borderChar: string, + startCorner: string, + endCorner: string +): string { + if (start >= end) { + return '' + } + + const includeStartCorner = start === 0 && startCorner.length > 0 + const includeEndCorner = end === borderLength && endCorner.length > 0 + const repeated = Math.max(0, end - start - (includeStartCorner ? 1 : 0) - (includeEndCorner ? 1 : 0)) + + return (includeStartCorner ? startCorner : '') + borderChar.repeat(repeated) + (includeEndCorner ? endCorner : '') +} + +function renderHorizontalBorder( + x: number, + y: number, + width: number, + output: Output, + visibleX1: number, + visibleX2: number, + borderChar: string, + startCorner: string, + endCorner: string, + color: Color | undefined, + dim: boolean | undefined, + borderText: BorderTextOptions | undefined +): void { + const borderLength = + Math.max(0, width - (startCorner ? 1 : 0) - (endCorner ? 1 : 0)) + (startCorner ? 1 : 0) + (endCorner ? 1 : 0) + + const clippedX1 = Math.max(0, Math.floor(visibleX1)) + const clippedX2 = Math.min(output.width, Math.ceil(visibleX2)) + const sliceStart = Math.max(0, Math.ceil(clippedX1 - x)) + const sliceEnd = Math.min(borderLength, Math.ceil(clippedX2 - x)) + + if ( + !Number.isSafeInteger(width) || + width < 0 || + !Number.isSafeInteger(sliceStart) || + !Number.isSafeInteger(sliceEnd) || + sliceStart >= sliceEnd + ) { + return + } + + const writeBorderRun = (start: number, end: number) => { + const text = borderRun(start, end, borderLength, borderChar, startCorner, endCorner) + + if (text) { + output.write(x + start, y, styleBorderLine(text, color, dim)) + } + } + + if (!borderText) { + writeBorderRun(sliceStart, sliceEnd) + + return + } + + const textLength = stringWidth(borderText.content) + + if (textLength >= borderLength - 2) { + const { leadingColumns, text } = sliceAnsiToWidth(borderText.content, sliceStart, sliceEnd) + + if (text) { + output.write(x + sliceStart + leadingColumns, y, text) + } + + return + } + + let position: number + + if (borderText.align === 'center') { + position = Math.floor((borderLength - textLength) / 2) + } else if (borderText.align === 'start') { + position = (borderText.offset ?? 0) + 1 + } else { + position = borderLength - textLength - (borderText.offset ?? 0) - 1 + } + + position = Math.max(1, Math.min(position, borderLength - textLength - 1)) + + writeBorderRun(sliceStart, Math.min(sliceEnd, position)) + + const visibleTextStart = Math.max(sliceStart, position) + const visibleTextEnd = Math.min(sliceEnd, position + textLength) + + if (visibleTextStart < visibleTextEnd) { + const { leadingColumns, text } = sliceAnsiToWidth( + borderText.content, + visibleTextStart - position, + visibleTextEnd - position + ) + + if (text) { + output.write(x + visibleTextStart + leadingColumns, y, text) + } + } + + writeBorderRun(Math.max(sliceStart, position + textLength), sliceEnd) +} + +const renderBorder = ( + x: number, + y: number, + node: DOMNode, + output: Output, + visibleX1 = 0, + visibleX2 = output.width, + visibleY1 = 0, + visibleY2 = output.height +): void => { if (node.style.borderStyle) { const width = Math.floor(node.yogaNode!.getComputedWidth()) const height = Math.floor(node.yogaNode!.getComputedHeight()) @@ -107,98 +205,75 @@ const renderBorder = (x: number, y: number, node: DOMNode, output: Output): void const showLeftBorder = node.style.borderLeft !== false const showRightBorder = node.style.borderRight !== false - const contentWidth = Math.max(0, width - (showLeftBorder ? 1 : 0) - (showRightBorder ? 1 : 0)) + const verticalTop = Math.floor(y + (showTopBorder ? 1 : 0)) + const verticalBottom = Math.ceil(y + height - (showBottomBorder ? 1 : 0)) + const clippedVerticalTop = Math.max(0, Math.floor(visibleY1), verticalTop) + const clippedVerticalBottom = Math.min(output.height, Math.ceil(visibleY2), verticalBottom) + const clippedVerticalHeight = Math.max(0, clippedVerticalBottom - clippedVerticalTop) - const topBorderLine = showTopBorder - ? (showLeftBorder ? box.topLeft : '') + box.top.repeat(contentWidth) + (showRightBorder ? box.topRight : '') - : '' - - // Handle text in top border - let topBorder: string | undefined - - if (showTopBorder && node.style.borderText?.position === 'top') { - const [before, text, after] = embedTextInBorder( - topBorderLine, - node.style.borderText.content, - node.style.borderText.align, - node.style.borderText.offset, - box.top - ) - - topBorder = - styleBorderLine(before, topBorderColor, dimTopBorderColor) + - text + - styleBorderLine(after, topBorderColor, dimTopBorderColor) - } else if (showTopBorder) { - topBorder = styleBorderLine(topBorderLine, topBorderColor, dimTopBorderColor) - } - - let verticalBorderHeight = height - - if (showTopBorder) { - verticalBorderHeight -= 1 - } - - if (showBottomBorder) { - verticalBorderHeight -= 1 - } - - verticalBorderHeight = Math.max(0, verticalBorderHeight) - - let leftBorder = (applyColor(box.left, leftBorderColor) + '\n').repeat(verticalBorderHeight) + let leftBorder = (applyColor(box.left, leftBorderColor) + '\n').repeat(clippedVerticalHeight) if (dimLeftBorderColor) { leftBorder = chalk.dim(leftBorder) } - let rightBorder = (applyColor(box.right, rightBorderColor) + '\n').repeat(verticalBorderHeight) + let rightBorder = (applyColor(box.right, rightBorderColor) + '\n').repeat(clippedVerticalHeight) if (dimRightBorderColor) { rightBorder = chalk.dim(rightBorder) } - const bottomBorderLine = showBottomBorder - ? (showLeftBorder ? box.bottomLeft : '') + - box.bottom.repeat(contentWidth) + - (showRightBorder ? box.bottomRight : '') - : '' - - // Handle text in bottom border - let bottomBorder: string | undefined - - if (showBottomBorder && node.style.borderText?.position === 'bottom') { - const [before, text, after] = embedTextInBorder( - bottomBorderLine, - node.style.borderText.content, - node.style.borderText.align, - node.style.borderText.offset, - box.bottom + if (showTopBorder && y >= visibleY1 && y < visibleY2 && y >= 0 && y < output.height) { + renderHorizontalBorder( + x, + y, + width, + output, + visibleX1, + visibleX2, + box.top, + showLeftBorder ? box.topLeft : '', + showRightBorder ? box.topRight : '', + topBorderColor, + dimTopBorderColor, + node.style.borderText?.position === 'top' ? node.style.borderText : undefined ) - - bottomBorder = - styleBorderLine(before, bottomBorderColor, dimBottomBorderColor) + - text + - styleBorderLine(after, bottomBorderColor, dimBottomBorderColor) - } else if (showBottomBorder) { - bottomBorder = styleBorderLine(bottomBorderLine, bottomBorderColor, dimBottomBorderColor) } - const offsetY = showTopBorder ? 1 : 0 - - if (topBorder) { - output.write(x, y, topBorder) + if (showLeftBorder && clippedVerticalHeight > 0 && x >= visibleX1 && x < visibleX2 && x >= 0 && x < output.width) { + output.write(x, clippedVerticalTop, leftBorder) } - if (showLeftBorder) { - output.write(x, y + offsetY, leftBorder) + const rightX = x + width - 1 + + if ( + showRightBorder && + clippedVerticalHeight > 0 && + rightX >= visibleX1 && + rightX < visibleX2 && + rightX >= 0 && + rightX < output.width + ) { + output.write(rightX, clippedVerticalTop, rightBorder) } - if (showRightBorder) { - output.write(x + width - 1, y + offsetY, rightBorder) - } + const bottomY = y + height - 1 - if (bottomBorder) { - output.write(x, y + height - 1, bottomBorder) + if (showBottomBorder && bottomY >= visibleY1 && bottomY < visibleY2 && bottomY >= 0 && bottomY < output.height) { + renderHorizontalBorder( + x, + bottomY, + width, + output, + visibleX1, + visibleX2, + box.bottom, + showLeftBorder ? box.bottomLeft : '', + showRightBorder ? box.bottomRight : '', + bottomBorderColor, + dimBottomBorderColor, + node.style.borderText?.position === 'bottom' ? node.style.borderText : undefined + ) } } } diff --git a/ui-tui/packages/hermes-ink/src/ink/render-node-to-output.ts b/ui-tui/packages/hermes-ink/src/ink/render-node-to-output.ts index fdd21c143f..d1f6325fe9 100644 --- a/ui-tui/packages/hermes-ink/src/ink/render-node-to-output.ts +++ b/ui-tui/packages/hermes-ink/src/ink/render-node-to-output.ts @@ -15,6 +15,33 @@ import { isXtermJs } from './terminal.js' import { widestLine } from './widest-line.js' import wrapText from './wrap-text.js' +const MAX_SCROLL_GEOMETRY = 1_000_000_000 +const MAX_YOGA_DIMENSION = 100_000_000 + +const validUnsignedGeometry = (value: number): boolean => + Number.isFinite(value) && value >= 0 && value <= MAX_SCROLL_GEOMETRY + +const validClampMaximum = (value: number): boolean => value === Number.POSITIVE_INFINITY || validUnsignedGeometry(value) + +const validSignedGeometry = (value: number): boolean => Number.isFinite(value) && Math.abs(value) <= MAX_SCROLL_GEOMETRY + +const safeUnsignedGeometry = (value: number | undefined, fallback = 0): number => + value !== undefined && validUnsignedGeometry(value) ? value : fallback + +const safeSignedGeometry = (value: number | undefined, fallback = 0): number => + value !== undefined && validSignedGeometry(value) ? value : fallback + +const validYogaDimension = (value: number): boolean => + Number.isFinite(value) && value >= 0 && value <= MAX_YOGA_DIMENSION + +const validYogaRect = (x: number, y: number, width: number, height: number): boolean => + validSignedGeometry(x) && + validSignedGeometry(y) && + validYogaDimension(width) && + validYogaDimension(height) && + validSignedGeometry(x + width) && + validSignedGeometry(y + height) + // Matches detectXtermJsWheel() in ScrollKeybindingHandler.tsx — the curve // and drain must agree on terminal detection. TERM_PROGRAM check is the sync // fallback; isXtermJs() is the authoritative XTVERSION-probe result. @@ -370,12 +397,30 @@ function wrapWithSoftWrap( // and use it as offset for the rest of the nodes // Only first node is taken into account, because other text nodes can't have margin or padding, // so their coordinates will be relative to the first node anyway -function applyPaddingToText(node: DOMElement, text: string, softWrap?: boolean[]): string { +function applyPaddingToText( + node: DOMElement, + text: string, + softWrap: boolean[] | undefined, + maxOffsetX: number, + maxOffsetY: number +): string { const yogaNode = node.childNodes[0]?.yogaNode if (yogaNode) { const offsetX = yogaNode.getComputedLeft() const offsetY = yogaNode.getComputedTop() + + if ( + !Number.isSafeInteger(offsetX) || + offsetX < 0 || + offsetX > maxOffsetX || + !Number.isSafeInteger(offsetY) || + offsetY < 0 || + offsetY > maxOffsetY + ) { + return '' + } + text = '\n'.repeat(offsetY) + indentString(text, offsetX) if (softWrap && offsetY > 0) { @@ -397,7 +442,11 @@ function renderNodeToOutput( offsetY = 0, prevScreen, skipSelfBlit = false, - inheritedBackgroundColor + inheritedBackgroundColor, + visibleX1 = 0, + visibleX2 = output.width, + visibleY1 = 0, + visibleY2 = output.height }: { offsetX?: number offsetY?: number @@ -409,6 +458,10 @@ function renderNodeToOutput( // opaque descendants' narrower rects are safe to blit. skipSelfBlit?: boolean inheritedBackgroundColor?: Color + visibleX1?: number + visibleX2?: number + visibleY1?: number + visibleY2?: number } ): void { const { yogaNode } = node @@ -456,6 +509,21 @@ function renderNodeToOutput( y = 0 } + // Yoga values are an untrusted renderer boundary. Invalid or implausible + // dimensions must never reach culling, string construction, or recursive + // rendering: NaN makes every comparison false, while huge finite heights + // can turn an opaque box or border into a catastrophic allocation. + if (!validYogaRect(x, y, width, height)) { + dropSubtreeCache(node) + + return + } + + const activeVisibleX1 = Math.max(0, Math.floor(visibleX1)) + const activeVisibleX2 = Math.min(output.width, Math.ceil(visibleX2)) + const activeVisibleY1 = Math.max(0, Math.floor(visibleY1)) + const activeVisibleY2 = Math.min(output.height, Math.ceil(visibleY2)) + // Check if we can skip this subtree (clean node with unchanged layout). // Blit cells from previous screen instead of re-rendering. const cached = nodeCache.get(node) @@ -504,15 +572,26 @@ function renderNodeToOutput( } if (cached && (node.dirty || positionChanged)) { - output.clear( - { - x: Math.floor(cached.x), - y: Math.floor(cached.y), - width: Math.floor(cached.width), - height: Math.floor(cached.height) - }, - node.style.position === 'absolute' - ) + if (validYogaRect(cached.x, cached.y, cached.width, cached.height)) { + const clearX1 = Math.max(activeVisibleX1, Math.floor(cached.x)) + const clearX2 = Math.min(activeVisibleX2, Math.ceil(cached.x + cached.width)) + const clearY1 = Math.max(activeVisibleY1, Math.floor(cached.y)) + const clearY2 = Math.min(activeVisibleY2, Math.ceil(cached.y + cached.height)) + + if (clearX1 < clearX2 && clearY1 < clearY2) { + output.clear( + { + x: clearX1, + y: clearY1, + width: clearX2 - clearX1, + height: clearY2 - clearY1 + }, + node.style.position === 'absolute' + ) + } + } else { + dropSubtreeCache(node) + } } // Read before deleting — hasRemovedChild disables prevScreen blitting @@ -633,7 +712,7 @@ function renderNodeToOutput( .join('') } - text = applyPaddingToText(node, text, softWrap) + text = applyPaddingToText(node, text, softWrap, output.width, output.height) output.write(x, y, text, softWrap) } @@ -670,13 +749,15 @@ function renderNodeToOutput( const isScrollY = overflowY === 'scroll' const needsClip = clipHorizontally || clipVertically + let x1: number | undefined + let x2: number | undefined let y1: number | undefined let y2: number | undefined if (needsClip) { - const x1 = clipHorizontally ? x + yogaNode.getComputedBorder(LayoutEdge.Left) : undefined + x1 = clipHorizontally ? x + yogaNode.getComputedBorder(LayoutEdge.Left) : undefined - const x2 = clipHorizontally + x2 = clipHorizontally ? x + yogaNode.getComputedWidth() - yogaNode.getComputedBorder(LayoutEdge.Right) : undefined @@ -689,6 +770,11 @@ function renderNodeToOutput( output.clip({ x1, x2, y1, y2 }) } + const childVisibleX1 = Math.max(activeVisibleX1, Math.floor(x1 ?? activeVisibleX1)) + const childVisibleX2 = Math.min(activeVisibleX2, Math.ceil(x2 ?? activeVisibleX2)) + const childVisibleY1 = Math.max(activeVisibleY1, Math.floor(y1 ?? activeVisibleY1)) + const childVisibleY2 = Math.min(activeVisibleY2, Math.ceil(y2 ?? activeVisibleY2)) + if (isScrollY) { // Scroll containers follow the ScrollBox component structure: // a single content-wrapper child with flexShrink:0 (doesn't shrink @@ -698,9 +784,8 @@ function renderNodeToOutput( // culled against the visible window. const padTop = yogaNode.getComputedPadding(LayoutEdge.Top) - const innerHeight = Math.max( - 0, - (y2 ?? y + height) - (y1 ?? y) - padTop - yogaNode.getComputedPadding(LayoutEdge.Bottom) + const innerHeight = safeUnsignedGeometry( + Math.max(0, (y2 ?? y + height) - (y1 ?? y) - padTop - yogaNode.getComputedPadding(LayoutEdge.Bottom)) ) const content = node.childNodes.find(c => (c as DOMElement).yogaNode) as DOMElement | undefined @@ -710,31 +795,32 @@ function renderNodeToOutput( // after terminal resizes Yoga can leave tall descendants overflowing // that wrapper. Use the deepest direct child bottom so sticky-bottom // math can still reach the real final rendered row. - let scrollHeight = Math.ceil(contentYoga?.getComputedHeight() ?? 0) + let scrollHeight = safeUnsignedGeometry(Math.ceil(contentYoga?.getComputedHeight() ?? 0)) if (content) { for (const child of content.childNodes) { const childYoga = (child as DOMElement).yogaNode if (childYoga) { - scrollHeight = Math.max( - scrollHeight, - Math.ceil(childYoga.getComputedTop() + childYoga.getComputedHeight()) - ) + const childBottom = Math.ceil(childYoga.getComputedTop() + childYoga.getComputedHeight()) + + if (validUnsignedGeometry(childBottom)) { + scrollHeight = Math.max(scrollHeight, childBottom) + } } } } // Capture previous scroll bounds BEFORE overwriting — the at-bottom // follow check compares against last frame's max. - const prevScrollHeight = node.scrollHeight ?? scrollHeight - const prevInnerHeight = node.scrollViewportHeight ?? innerHeight + const prevScrollHeight = safeUnsignedGeometry(node.scrollHeight, scrollHeight) + const prevInnerHeight = safeUnsignedGeometry(node.scrollViewportHeight, innerHeight) node.scrollHeight = scrollHeight node.scrollViewportHeight = innerHeight // Absolute screen-buffer row where the scrollable area (inside // padding) begins. Exposed via ScrollBoxHandle.getViewportTop() so // drag-to-scroll can detect when the drag leaves the scroll viewport. - node.scrollViewportTop = (y1 ?? y) + padTop + node.scrollViewportTop = safeUnsignedGeometry((y1 ?? y) + padTop) const maxScroll = Math.max(0, scrollHeight - innerHeight) @@ -751,9 +837,16 @@ function renderNodeToOutput( // plumbing; shipping instant first. stickyScroll overrides. if (node.scrollAnchor) { const anchorTop = node.scrollAnchor.el.yogaNode?.getComputedTop() + const anchorOffset = node.scrollAnchor.offset + const anchorTarget = (anchorTop ?? Number.NaN) + anchorOffset - if (anchorTop != null) { - node.scrollTop = anchorTop + node.scrollAnchor.offset + if ( + anchorTop != null && + validUnsignedGeometry(anchorTop) && + validSignedGeometry(anchorOffset) && + validUnsignedGeometry(anchorTarget) + ) { + node.scrollTop = anchorTarget node.pendingScrollDelta = undefined } @@ -771,8 +864,17 @@ function renderNodeToOutput( // Capture scrollTop before follow so ink.tsx can translate any // active text selection by the same delta (native terminal behavior: // view keeps scrolling, highlight walks up with the text). - const scrollTopBeforeFollow = node.scrollTop ?? 0 + const scrollTopBeforeFollow = safeUnsignedGeometry(node.scrollTop) const stickyBeforeFollow = node.stickyScroll + const scrollTopCompensation = safeSignedGeometry(node.scrollTopCompensation) + + // Compensation is additive and one-shot. Positional bottom-follow + // must judge where the viewport was before the adjustment; otherwise + // a near-tail manual viewport can cross prevMaxScroll solely because + // an above-row grew and be mistaken for an intentional bottom pin. + const scrollTopBeforeCompensation = safeUnsignedGeometry(scrollTopBeforeFollow - scrollTopCompensation) + + node.scrollTopCompensation = undefined const sticky = node.stickyScroll ?? Boolean(node.attributes['stickyScroll']) @@ -783,7 +885,11 @@ function renderNodeToOutput( // because the user was at bottom. const grew = scrollHeight >= prevScrollHeight - const atBottom = sticky || (grew && scrollTopBeforeFollow >= prevMaxScroll) + if (node.pendingScrollDelta !== undefined && !validSignedGeometry(node.pendingScrollDelta)) { + node.pendingScrollDelta = undefined + } + + const atBottom = sticky || (grew && scrollTopBeforeCompensation >= prevMaxScroll) if (atBottom && (node.pendingScrollDelta ?? 0) >= 0) { node.scrollTop = maxScroll @@ -799,7 +905,7 @@ function renderNodeToOutput( // undefined (never set by user action) leave it alone — setting it // would make the sticky flag sticky-by-default and lock out // direct scrollTop writes (e.g. the alt-screen-perf test). - if (node.stickyScroll === false && scrollTopBeforeFollow >= prevMaxScroll) { + if (node.stickyScroll === false && scrollTopBeforeCompensation >= prevMaxScroll) { node.stickyScroll = true } } @@ -822,13 +928,33 @@ function renderNodeToOutput( // (pendingScrollDelta is only set by wheel events, >>50ms after // startup) the probe has resolved — same timing guarantee the // wheel-accel curve relies on. - let cur = node.scrollTop ?? 0 - const pending = node.pendingScrollDelta + let cur = safeUnsignedGeometry(node.scrollTop) + let pending = node.pendingScrollDelta + + if (pending !== undefined && !validSignedGeometry(pending)) { + node.pendingScrollDelta = undefined + pending = undefined + } + const cMin = node.scrollClampMin const cMax = node.scrollClampMax - const haveClamp = cMin !== undefined && cMax !== undefined - if (pending !== undefined && pending !== 0) { + const haveClamp = + cMin !== undefined && + cMax !== undefined && + validUnsignedGeometry(cMin) && + validClampMaximum(cMax) && + cMin <= cMax + + if (!haveClamp && (cMin !== undefined || cMax !== undefined)) { + node.scrollClampMin = undefined + node.scrollClampMax = undefined + } + + // Preserve pending user intent for the compensation paint. Draining + // resumes on the next frame; this keeps the anchor adjustment from + // being conflated with a user move at the old bottom boundary. + if (scrollTopCompensation === 0 && pending !== undefined && pending !== 0) { // Drain continues even past the clamp — the render-clamp below // holds the VISUAL at the mounted edge regardless. Hard-stopping // here caused stop-start jutter: drain hits edge → pause → React @@ -844,14 +970,18 @@ function renderNodeToOutput( const pastClamp = haveClamp && ((pending < 0 && cur < cMin) || (pending > 0 && cur > cMax)) const eff = pastClamp ? Math.min(4, innerHeight >> 3) : innerHeight - cur += isXtermJsHost() ? drainAdaptive(node, pending, eff) : drainProportional(node, pending, eff) - } else if (pending === 0) { + + const drained = + cur + (isXtermJsHost() ? drainAdaptive(node, pending, eff) : drainProportional(node, pending, eff)) + + cur = safeUnsignedGeometry(drained, cur) + } else if (scrollTopCompensation === 0 && pending === 0) { // Opposite scrollBy calls cancelled to zero — clear so we don't // schedule an infinite loop of no-op drain frames. node.pendingScrollDelta = undefined } - let scrollTop = Math.max(0, Math.min(cur, maxScroll)) + let scrollTop = safeUnsignedGeometry(Math.max(0, Math.min(cur, maxScroll))) // Virtual-scroll clamp: if scrollTop raced past the currently-mounted // range (burst PageUp before React re-renders), render at the EDGE of @@ -948,7 +1078,42 @@ function renderNodeToOutput( const prevHeight = contentCached?.height ?? scrollHeight const heightDelta = scrollHeight - prevHeight - const safeForFastPath = !hint || heightDelta === 0 || (hint.delta > 0 && heightDelta === hint.delta) + const heightSafeForFastPath = !hint || heightDelta === 0 || (hint.delta > 0 && heightDelta === hint.delta) + const outputWidth = Number.isSafeInteger(output.width) && output.width > 0 ? output.width : 0 + const outputHeight = Number.isSafeInteger(output.height) && output.height > 0 ? output.height : 0 + + const fastPathBounds = (() => { + if ( + !hint || + !Number.isSafeInteger(hint.top) || + !Number.isSafeInteger(hint.bottom) || + !Number.isSafeInteger(hint.delta) || + hint.delta === 0 || + hint.top > hint.bottom || + !Number.isFinite(x) || + !Number.isFinite(width) || + width < 0 || + !Number.isFinite(childVisibleX1) || + !Number.isFinite(childVisibleX2) || + !Number.isFinite(childVisibleY1) || + !Number.isFinite(childVisibleY2) + ) { + return null + } + + const x1 = Math.max(0, Math.floor(x), Math.floor(childVisibleX1)) + const x2 = Math.min(outputWidth, Math.ceil(x + width), Math.ceil(childVisibleX2)) + const top = Math.max(0, hint.top, Math.floor(childVisibleY1)) + const bottom = Math.min(outputHeight, hint.bottom + 1, Math.ceil(childVisibleY2)) + + if (x1 >= x2 || top >= bottom || Math.abs(hint.delta) > bottom - top) { + return null + } + + return { bottom, top, width: x2 - x1, x: x1 } + })() + + const safeForFastPath = heightSafeForFastPath && (!hint || fastPathBounds !== null) // Diagnostics (opt-in via scrollFastPathStats reader). Only // counts when a hint was captured — cases where nothing scrolled @@ -960,9 +1125,12 @@ function renderNodeToOutput( scrollFastPathStats.lastPrevHeight = prevHeight scrollFastPathStats.lastHeightDelta = heightDelta - if (!safeForFastPath) { + if (!heightSafeForFastPath) { scrollFastPathStats.declined.heightDeltaMismatch++ scrollFastPathStats.lastDeclineReason = `heightDelta=${heightDelta} hintDelta=${hint.delta}` + } else if (!fastPathBounds) { + scrollFastPathStats.declined.other++ + scrollFastPathStats.lastDeclineReason = 'invalidOrEmptyRepairBounds' } else if (!prevScreen) { scrollFastPathStats.declined.noPrevScreen++ scrollFastPathStats.lastDeclineReason = 'noPrevScreen' @@ -979,26 +1147,19 @@ function renderNodeToOutput( scrollHint = null } - if (hint && prevScreen && safeForFastPath) { - const { top, bottom, delta } = hint - const w = Math.floor(width) - output.blit(prevScreen, Math.floor(x), top, w, bottom - top + 1) - output.shift(top, bottom, delta) + if (hint && prevScreen && safeForFastPath && fastPathBounds) { + const { delta } = hint + const { bottom, top, width: repairWidth, x: repairX } = fastPathBounds + const bottomInclusive = bottom - 1 + + // Keep the terminal hint aligned with the same bounded rows used + // to construct next.screen. Valid in-bounds hints are unchanged. + scrollHint = { bottom: bottomInclusive, delta, top } + output.blit(prevScreen, repairX, top, repairWidth, bottom - top) + output.shift(top, bottomInclusive, delta) // Edge rows: new content entering the viewport. - const edgeTop = delta > 0 ? bottom - delta + 1 : top - const edgeBottom = delta > 0 ? bottom : top - delta - 1 - output.clear({ - x: Math.floor(x), - y: edgeTop, - width: w, - height: edgeBottom - edgeTop + 1 - }) - output.clip({ - x1: undefined, - x2: undefined, - y1: edgeTop, - y2: edgeBottom + 1 - }) + const edgeTop = Math.max(top, delta > 0 ? bottom - delta : top) + const edgeBottom = Math.min(bottom, delta > 0 ? bottom : top - delta) // Snapshot dirty children before the first pass — the first // pass clears dirty flags, and edge-spanning children would be @@ -1007,20 +1168,33 @@ function renderNodeToOutput( ? new Set(content.childNodes.filter(c => (c as DOMElement).dirty)) : null - renderScrolledChildren( - content, - output, - contentX, - contentY, - hasRemovedChild, - undefined, - // Cull to edge in child-local coords (inverse of contentY offset). - edgeTop - contentY, - edgeBottom + 1 - contentY, - boxBackgroundColor, - true - ) - output.unclip() + if (edgeTop < edgeBottom) { + output.clear({ + x: repairX, + y: edgeTop, + width: repairWidth, + height: edgeBottom - edgeTop + }) + output.clip({ x1: repairX, x2: repairX + repairWidth, y1: edgeTop, y2: edgeBottom }) + renderScrolledChildren( + content, + output, + contentX, + contentY, + hasRemovedChild, + undefined, + // Cull to edge in child-local coords (inverse of contentY offset). + edgeTop - contentY, + edgeBottom - contentY, + boxBackgroundColor, + true, + repairX, + repairX + repairWidth, + edgeTop, + edgeBottom + ) + output.unclip() + } // Second pass: re-render children in stable rows whose screen // position doesn't match where the shift put their old pixels. @@ -1040,8 +1214,7 @@ function renderNodeToOutput( // path preserved. if (dirtyChildren) { const edgeTopLocal = edgeTop - contentY - const edgeBottomLocal = edgeBottom + 1 - contentY - const spaces = ' '.repeat(w) + const edgeBottomLocal = edgeBottom - contentY // Track cumulative height change of children iterated so far. // A clean child's yogaTop is unchanged iff this is zero (no // sibling above it grew/shrank/mounted). When zero, the skip @@ -1075,8 +1248,17 @@ function renderNodeToOutput( continue } + const childLeft = cy.getComputedLeft() const childTop = cy.getComputedTop() + const childW = cy.getComputedWidth() const childH = cy.getComputedHeight() + + if (!validYogaRect(childLeft, childTop, childW, childH)) { + dropSubtreeCache(childElem) + + continue + } + const childBottom = childTop + childH if (isDirty) { @@ -1094,7 +1276,23 @@ function renderNodeToOutput( continue } - const screenY = Math.floor(contentY + childTop) + const childScreenX = contentX + childLeft + const childScreenRight = childScreenX + childW + const childScreenY = contentY + childTop + const childScreenBottom = contentY + childBottom + + if ( + !validSignedGeometry(childScreenX) || + !validSignedGeometry(childScreenRight) || + !validSignedGeometry(childScreenY) || + !validSignedGeometry(childScreenBottom) + ) { + dropSubtreeCache(childElem) + + continue + } + + const screenY = Math.floor(childScreenY) // Clean children reaching here have cumHeightShift ≠ 0 OR // no cache. Re-check precisely: cached.y − delta is where @@ -1113,28 +1311,39 @@ function renderNodeToOutput( // Wipe this child's region with spaces to overwrite stale // blitted content — output.clear() only expands damage and // cannot zero cells that the blit already wrote. - const screenBottom = Math.min( - Math.floor(contentY + childBottom), + const repairChildX = Math.max(repairX, Math.floor(childScreenX)) + const repairChildRight = Math.min(repairX + repairWidth, Math.ceil(childScreenRight)) + const repairChildTop = Math.max(top, screenY) + + const repairChildBottom = Math.min( + bottom, + Math.floor(childScreenBottom), Math.floor((y1 ?? y) + padTop + innerHeight) ) - if (screenY < screenBottom) { - const fill = Array(screenBottom - screenY) + if (repairChildX < repairChildRight && repairChildTop < repairChildBottom) { + const spaces = ' '.repeat(repairChildRight - repairChildX) + + const fill = Array(repairChildBottom - repairChildTop) .fill(spaces) .join('\n') - output.write(Math.floor(x), screenY, fill) + output.write(repairChildX, repairChildTop, fill) output.clip({ - x1: undefined, - x2: undefined, - y1: screenY, - y2: screenBottom + x1: repairChildX, + x2: repairChildRight, + y1: repairChildTop, + y2: repairChildBottom }) renderNodeToOutput(childElem, output, { offsetX: contentX, offsetY: contentY, prevScreen: undefined, - inheritedBackgroundColor: boxBackgroundColor + inheritedBackgroundColor: boxBackgroundColor, + visibleX1: repairChildX, + visibleX2: repairChildRight, + visibleY1: repairChildTop, + visibleY2: repairChildBottom }) output.unclip() } @@ -1148,34 +1357,36 @@ function renderNodeToOutput( // pixels sit at (rect.y - delta) — neither edge render nor the // overlay's own re-render covers them. Wipe and re-render // ScrollBox content so the diff writes correct cells. - const spaces = absoluteRectsPrev.length ? ' '.repeat(w) : '' - for (const r of absoluteRectsPrev) { - if (r.y >= bottom + 1 || r.y + r.height <= top) { + if (!validYogaRect(r.x, r.y, r.width, r.height)) { continue } + const repairOverlayX = Math.max(repairX, Math.floor(r.x)) + const repairOverlayRight = Math.min(repairX + repairWidth, Math.ceil(r.x + r.width)) const shiftedTop = Math.max(top, Math.floor(r.y) - delta) - const shiftedBottom = Math.min(bottom + 1, Math.floor(r.y + r.height) - delta) + const shiftedBottom = Math.min(bottom, Math.floor(r.y + r.height) - delta) // Skip if entirely within edge rows (already rendered). - if (shiftedTop >= edgeTop && shiftedBottom <= edgeBottom + 1) { + if (edgeTop < edgeBottom && shiftedTop >= edgeTop && shiftedBottom <= edgeBottom) { continue } - if (shiftedTop >= shiftedBottom) { + if (repairOverlayX >= repairOverlayRight || shiftedTop >= shiftedBottom) { continue } + const spaces = ' '.repeat(repairOverlayRight - repairOverlayX) + const fill = Array(shiftedBottom - shiftedTop) .fill(spaces) .join('\n') - output.write(Math.floor(x), shiftedTop, fill) + output.write(repairOverlayX, shiftedTop, fill) output.clip({ - x1: undefined, - x2: undefined, + x1: repairOverlayX, + x2: repairOverlayRight, y1: shiftedTop, y2: shiftedBottom }) @@ -1189,7 +1400,11 @@ function renderNodeToOutput( shiftedTop - contentY, shiftedBottom - contentY, boxBackgroundColor, - true + true, + repairOverlayX, + repairOverlayRight, + shiftedTop, + shiftedBottom ) output.unclip() } @@ -1233,7 +1448,12 @@ function renderNodeToOutput( scrolled || positionChanged ? undefined : prevScreen, scrollTop, scrollTop + innerHeight, - boxBackgroundColor + boxBackgroundColor, + false, + childVisibleX1, + childVisibleX2, + childVisibleY1, + childVisibleY2 ) } @@ -1264,14 +1484,23 @@ function renderNodeToOutput( const innerHeight = Math.floor(height) - borderTop - borderBottom if (innerWidth > 0 && innerHeight > 0) { - const spaces = ' '.repeat(innerWidth) + const fillX1 = Math.max(0, Math.floor(x + borderLeft)) + const fillX2 = Math.min(output.width, Math.ceil(x + width - borderRight)) + const fillY1 = Math.max(childVisibleY1, Math.floor(y + borderTop)) + const fillY2 = Math.min(childVisibleY2, Math.ceil(y + height - borderBottom)) + const fillWidth = Math.max(0, fillX2 - fillX1) + const fillHeight = Math.max(0, fillY2 - fillY1) + + const spaces = ' '.repeat(fillWidth) const fillLine = ownBackgroundColor ? applyTextStyles(spaces, { backgroundColor: ownBackgroundColor }) : spaces - const fill = Array(innerHeight).fill(fillLine).join('\n') - output.write(x + borderLeft, y + borderTop, fill) + if (fillWidth > 0 && fillHeight > 0) { + const fill = Array(fillHeight).fill(fillLine).join('\n') + output.write(fillX1, fillY1, fill) + } } } @@ -1289,7 +1518,11 @@ function renderNodeToOutput( // valid composite, but children CAN reposition (ScrollBox remeasure // on re-render → /permissions body blanked on Down arrow, #25436). ownBackgroundColor || node.style.opaque ? undefined : prevScreen, - boxBackgroundColor + boxBackgroundColor, + childVisibleX1, + childVisibleX2, + childVisibleY1, + childVisibleY2 ) } @@ -1300,9 +1533,21 @@ function renderNodeToOutput( // Render border AFTER children to ensure it's not overwritten by child // clearing operations. When a child shrinks, it clears its old area, // which may overlap with where the parent's border now is. - renderBorder(x, y, node, output) + renderBorder(x, y, node, output, activeVisibleX1, activeVisibleX2, activeVisibleY1, activeVisibleY2) } else if (node.nodeName === 'ink-root') { - renderChildren(node, output, x, y, hasRemovedChild, prevScreen, inheritedBackgroundColor) + renderChildren( + node, + output, + x, + y, + hasRemovedChild, + prevScreen, + inheritedBackgroundColor, + activeVisibleX1, + activeVisibleX2, + activeVisibleY1, + activeVisibleY2 + ) } // Cache layout bounds for dirty tracking @@ -1352,7 +1597,11 @@ function renderChildren( offsetY: number, hasRemovedChild: boolean, prevScreen: Screen | undefined, - inheritedBackgroundColor: Color | undefined + inheritedBackgroundColor: Color | undefined, + visibleX1: number, + visibleX2: number, + visibleY1: number, + visibleY2: number ): void { let seenDirtyChild = false let seenDirtyClipped = false @@ -1370,7 +1619,11 @@ function renderChildren( // the opaque/bg reads don't happen per-child per-frame. skipSelfBlit: seenDirtyClipped && isAbsolute && !childElem.style.opaque && childElem.style.backgroundColor === undefined, - inheritedBackgroundColor + inheritedBackgroundColor, + visibleX1, + visibleX2, + visibleY1, + visibleY2 }) if (wasDirty && !seenDirtyChild) { @@ -1499,7 +1752,11 @@ function renderScrolledChildren( // When true (DECSTBM fast path), culled children keep their cache — // the blit+shift put stable rows in next.screen so stale cache is // never read. Avoids walking O(total_children * subtree_depth) per frame. - preserveCulledCache = false + preserveCulledCache = false, + visibleX1 = 0, + visibleX2 = output.width, + visibleY1 = 0, + visibleY2 = output.height ): void { let seenDirtyChild = false // Track cumulative height shift of dirty children iterated so far. When @@ -1527,6 +1784,14 @@ function renderScrolledChildren( top = cy.getComputedTop() height = cy.getComputedHeight() + const bottom = top + height + + if (!validSignedGeometry(top) || !validYogaDimension(height) || !validSignedGeometry(bottom) || bottom < top) { + dropSubtreeCache(childElem) + + continue + } + if (childElem.dirty) { cumHeightShift += height - (cached ? cached.height : 0) } @@ -1542,6 +1807,12 @@ function renderScrolledChildren( const bottom = top + height + if (!validSignedGeometry(top) || !validYogaDimension(height) || !validSignedGeometry(bottom) || bottom < top) { + dropSubtreeCache(childElem) + + continue + } + if (bottom <= scrollTopY || top >= scrollBottomY) { // Culled — outside visible window. Drop stale cache entries from // the subtree so when this child re-enters it doesn't fire clears @@ -1560,7 +1831,11 @@ function renderScrolledChildren( offsetX, offsetY, prevScreen: hasRemovedChild || seenDirtyChild ? undefined : prevScreen, - inheritedBackgroundColor + inheritedBackgroundColor, + visibleX1, + visibleX2, + visibleY1: Math.max(visibleY1, offsetY + scrollTopY), + visibleY2: Math.min(visibleY2, offsetY + scrollBottomY) }) if (wasDirty) { diff --git a/ui-tui/scripts/bench-history-scroll.tsx b/ui-tui/scripts/bench-history-scroll.tsx new file mode 100644 index 0000000000..7f7a73234f --- /dev/null +++ b/ui-tui/scripts/bench-history-scroll.tsx @@ -0,0 +1,483 @@ +// Deterministic virtual-history benchmark. The file intentionally uses only +// APIs present before the performance candidate so the exact same script can +// be copied/run on base and candidate checkouts. +// +// Run from ui-tui: +// npx tsx scripts/bench-history-scroll.tsx +// npx tsx scripts/bench-history-scroll.tsx --warmups=2 --samples=5 --items=100,1000,10000 +// +// In addition to the virtual-history workloads, every run mounts one +// oversized bordered/fill box at each RENDERER_EXTENT inside the fixed +// viewport. Keeping that tree to a few Yoga nodes isolates renderer clipping +// from node-construction cost and makes the workload revision-comparable. + +import { PassThrough } from 'stream' + +import { Box, renderSync, ScrollBox, type ScrollBoxHandle, Text } from '@hermes/ink' +import React, { useLayoutEffect, useRef } from 'react' + +import { useVirtualHistory } from '../src/hooks/useVirtualHistory.js' + +const DEFAULT_WORKLOADS = [100, 1_000, 10_000] +const RENDERER_EXTENTS = [100, 1_000, 10_000] +const DEFAULT_WARMUPS = 1 +const DEFAULT_SAMPLES = 5 +const COLUMNS = 100 +const ROWS = 30 +const MAX_MOUNTED = 120 + +interface BenchItem { + height: number + key: string + text: string +} + +interface Exposed { + scroll: ScrollBoxHandle | null + virtual: ReturnType +} + +interface Sample { + anchorError: number + heapDeltaBytes: number | null + invalidOffsets: number + measuredHeightReconciliationMs: number + mountMs: number + mountedRowsMax: number + nonMonotoneOffsets: number + rerenderMs: number + scrollMs: number + terminalBytes: number + terminalWrites: number +} + +interface WorkloadResult { + distributions: { + anchorError: ReturnType + heapDeltaBytes: ReturnType + measuredHeightReconciliationMs: ReturnType + mountMs: ReturnType + mountedRowsMax: ReturnType + rerenderMs: ReturnType + scrollMs: ReturnType + terminalBytes: ReturnType + terminalWrites: ReturnType + } + invalidOffsets: number + itemCount: number + nonMonotoneOffsets: number + samples: Sample[] +} + +interface OversizedRendererSample { + freshMountRenderMs: number + terminalBytes: number + terminalWrites: number +} + +interface OversizedRendererResult { + distributions: { + freshMountRenderMs: ReturnType + terminalBytes: ReturnType + terminalWrites: ReturnType + } + extent: number + samples: OversizedRendererSample[] +} + +class CountingStream extends PassThrough { + columns = COLUMNS + rows = ROWS + isTTY = false + bytes = 0 + writes = 0 + + override _write(chunk: Buffer | string, encoding: BufferEncoding, callback: (error?: Error | null) => void) { + this.bytes += Buffer.byteLength(chunk) + this.writes++ + callback() + } +} + +const immediate = () => new Promise(resolve => setImmediate(resolve)) + +async function settle(frames = 4) { + for (let frame = 0; frame < frames; frame++) { + await immediate() + } +} + +async function waitUntil(predicate: () => boolean, attempts = 40) { + for (let attempt = 0; attempt < attempts; attempt++) { + if (predicate()) { + return true + } + + await immediate() + } + + return predicate() +} + +function makeItems(count: number): BenchItem[] { + return Array.from({ length: count }, (_, index) => ({ + height: 1 + ((index * 17) % 4), + key: `row-${index}`, + text: `row ${index} ${'history '.repeat(2 + (index % 5))}` + })) +} + +function Harness({ expose, items }: { expose: React.MutableRefObject; items: readonly BenchItem[] }) { + const scrollRef = useRef(null) + + const virtual = useVirtualHistory(scrollRef, items, COLUMNS, { + coldStartCount: 30, + estimateHeight: index => items[index]?.height ?? 1, + maxMounted: MAX_MOUNTED, + overscan: 20 + }) + + useLayoutEffect(() => { + expose.current = { scroll: scrollRef.current, virtual } + }) + + return ( + + + {virtual.topSpacer > 0 ? : null} + {items.slice(virtual.start, virtual.end).map(item => ( + + {item.text} + + ))} + {virtual.bottomSpacer > 0 ? : null} + + + ) +} + +function OversizedRendererHarness({ extent }: { extent: number }) { + return ( + + + deterministic oversized height workload + + + deterministic oversized width workload + + + ) +} + +function inspectOffsets(offsets: ArrayLike, count: number) { + let invalidOffsets = 0 + let nonMonotoneOffsets = 0 + + for (let index = 0; index <= count; index++) { + const value = offsets[index] + + if (!Number.isFinite(value)) { + invalidOffsets++ + } + + if (index > 0 && value! < offsets[index - 1]!) { + nonMonotoneOffsets++ + } + } + + return { invalidOffsets, nonMonotoneOffsets } +} + +async function runSample(itemCount: number): Promise { + const stdout = new CountingStream() + const stderr = new CountingStream() + const stdin = new PassThrough() + const expose = { current: null as Exposed | null } + let items = makeItems(itemCount) + const heapBefore = process.memoryUsage?.().heapUsed ?? null + const mountStart = performance.now() + + const instance = renderSync(, { + patchConsole: false, + stderr: stderr as unknown as NodeJS.WriteStream, + stdin: stdin as unknown as NodeJS.ReadStream, + stdout: stdout as unknown as NodeJS.WriteStream + }) + + await waitUntil(() => expose.current?.scroll !== null) + await settle() + const mountMs = performance.now() - mountStart + let mountedRowsMax = expose.current!.virtual.end - expose.current!.virtual.start + + const rerenderItems = items.map((item, index) => + index === items.length - 1 ? { ...item, text: `${item.text} rerender` } : item + ) + + const rerenderStart = performance.now() + + instance.rerender() + await settle() + const rerenderMs = performance.now() - rerenderStart + items = rerenderItems + mountedRowsMax = Math.max(mountedRowsMax, expose.current!.virtual.end - expose.current!.virtual.start) + + const scroll = expose.current!.scroll! + const total = expose.current!.virtual.offsets[itemCount] ?? 0 + const scrollStart = performance.now() + + scroll.scrollTo(Math.max(0, Math.floor(total * 0.55))) + await settle(8) + const scrollMs = performance.now() - scrollStart + mountedRowsMax = Math.max(mountedRowsMax, expose.current!.virtual.end - expose.current!.virtual.start) + + const beforeOffsets = expose.current!.virtual.offsets + const beforeTop = scroll.getScrollTop() + let measuredIndex = expose.current!.virtual.start + + while (measuredIndex + 1 < expose.current!.virtual.end && (beforeOffsets[measuredIndex + 1] ?? 0) > beforeTop) { + measuredIndex++ + } + + if ((beforeOffsets[measuredIndex + 1] ?? Number.POSITIVE_INFINITY) > beforeTop) { + measuredIndex = Math.max(expose.current!.virtual.start, measuredIndex - 1) + } + + const heightDelta = 3 + const oldTotal = beforeOffsets[itemCount] ?? 0 + + const measuredItems = items.map((item, index) => + index === measuredIndex ? { ...item, height: item.height + heightDelta } : item + ) + + const reconcileStart = performance.now() + + instance.rerender() + await waitUntil(() => (expose.current!.virtual.offsets[itemCount] ?? 0) === oldTotal + heightDelta) + await settle(2) + const measuredHeightReconciliationMs = performance.now() - reconcileStart + const measuredWasAbove = (beforeOffsets[measuredIndex + 1] ?? 0) <= beforeTop + const expectedTop = beforeTop + (measuredWasAbove ? heightDelta : 0) + const anchorError = Math.abs(scroll.getScrollTop() - expectedTop) + mountedRowsMax = Math.max(mountedRowsMax, expose.current!.virtual.end - expose.current!.virtual.start) + + const offsetHealth = inspectOffsets(expose.current!.virtual.offsets, itemCount) + const heapAfter = process.memoryUsage?.().heapUsed ?? null + const heapDeltaBytes = heapBefore === null || heapAfter === null ? null : heapAfter - heapBefore + const terminalBytes = stdout.bytes + const terminalWrites = stdout.writes + + instance.unmount() + instance.cleanup() + stdin.destroy() + stdout.destroy() + stderr.destroy() + + return { + anchorError, + heapDeltaBytes, + ...offsetHealth, + measuredHeightReconciliationMs, + mountMs, + mountedRowsMax, + rerenderMs, + scrollMs, + terminalBytes, + terminalWrites + } +} + +async function runOversizedRendererSample(extent: number): Promise { + const stdout = new CountingStream() + const stderr = new CountingStream() + const stdin = new PassThrough() + const mountStart = performance.now() + let instance: ReturnType | undefined + + try { + instance = renderSync(, { + patchConsole: false, + stderr: stderr as unknown as NodeJS.WriteStream, + stdin: stdin as unknown as NodeJS.ReadStream, + stdout: stdout as unknown as NodeJS.WriteStream + }) + + const rendered = await waitUntil(() => stdout.writes > 0) + + if (!rendered) { + throw new Error(`oversized renderer extent ${extent} did not produce a terminal frame`) + } + + return { + freshMountRenderMs: performance.now() - mountStart, + terminalBytes: stdout.bytes, + terminalWrites: stdout.writes + } + } finally { + instance?.unmount() + instance?.cleanup() + + stdin.destroy() + stdout.destroy() + stderr.destroy() + } +} + +function distribution(values: number[]) { + const sorted = [...values].sort((a, b) => a - b) + + const percentile = (p: number) => + sorted[Math.min(sorted.length - 1, Math.max(0, Math.ceil(sorted.length * p) - 1))] ?? 0 + + return { + max: sorted.at(-1) ?? 0, + mean: sorted.reduce((sum, value) => sum + value, 0) / Math.max(1, sorted.length), + min: sorted[0] ?? 0, + p50: percentile(0.5), + p95: percentile(0.95), + p99: percentile(0.99) + } +} + +function numericArg(name: string, fallback: number) { + const raw = process.argv + .slice(2) + .find(arg => arg.startsWith(`--${name}=`)) + ?.split('=', 2)[1] + + const parsed = Number(raw) + + return Number.isSafeInteger(parsed) && parsed >= 0 ? parsed : fallback +} + +function workloadsArg() { + const raw = process.argv + .slice(2) + .find(arg => arg.startsWith('--items=')) + ?.split('=', 2)[1] + + if (!raw) { + return DEFAULT_WORKLOADS + } + + const parsed = raw.split(',').map(Number) + + if (parsed.some(value => !Number.isSafeInteger(value) || value <= 0)) { + throw new Error(`invalid --items workload list: ${raw}`) + } + + return parsed +} + +async function main() { + const workloads = workloadsArg() + const warmups = numericArg('warmups', DEFAULT_WARMUPS) + const samplesPerWorkload = numericArg('samples', DEFAULT_SAMPLES) + const results: WorkloadResult[] = [] + const oversizedRendererResults: OversizedRendererResult[] = [] + + for (const itemCount of workloads) { + for (let warmup = 0; warmup < warmups; warmup++) { + await runSample(itemCount) + } + + const samples: Sample[] = [] + + for (let sample = 0; sample < samplesPerWorkload; sample++) { + samples.push(await runSample(itemCount)) + } + + results.push({ + itemCount, + distributions: { + anchorError: distribution(samples.map(sample => sample.anchorError)), + heapDeltaBytes: distribution(samples.flatMap(sample => sample.heapDeltaBytes ?? [])), + measuredHeightReconciliationMs: distribution(samples.map(sample => sample.measuredHeightReconciliationMs)), + mountMs: distribution(samples.map(sample => sample.mountMs)), + mountedRowsMax: distribution(samples.map(sample => sample.mountedRowsMax)), + rerenderMs: distribution(samples.map(sample => sample.rerenderMs)), + scrollMs: distribution(samples.map(sample => sample.scrollMs)), + terminalBytes: distribution(samples.map(sample => sample.terminalBytes)), + terminalWrites: distribution(samples.map(sample => sample.terminalWrites)) + }, + invalidOffsets: samples.reduce((sum, sample) => sum + sample.invalidOffsets, 0), + nonMonotoneOffsets: samples.reduce((sum, sample) => sum + sample.nonMonotoneOffsets, 0), + samples + }) + } + + for (const extent of RENDERER_EXTENTS) { + for (let warmup = 0; warmup < warmups; warmup++) { + await runOversizedRendererSample(extent) + } + + const samples: OversizedRendererSample[] = [] + + for (let sample = 0; sample < samplesPerWorkload; sample++) { + samples.push(await runOversizedRendererSample(extent)) + } + + oversizedRendererResults.push({ + distributions: { + freshMountRenderMs: distribution(samples.map(sample => sample.freshMountRenderMs)), + terminalBytes: distribution(samples.map(sample => sample.terminalBytes)), + terminalWrites: distribution(samples.map(sample => sample.terminalWrites)) + }, + extent, + samples + }) + } + + const scaling = results.slice(1).map((result, index) => { + const previous = results[index]! + + const ratio = (metric: keyof (typeof result)['distributions']) => + result.distributions[metric].p50 / Math.max(Number.EPSILON, previous.distributions[metric].p50) + + return { + fromItems: previous.itemCount, + itemFactor: result.itemCount / previous.itemCount, + measuredHeightReconciliationP50Factor: ratio('measuredHeightReconciliationMs'), + mountP50Factor: ratio('mountMs'), + rerenderP50Factor: ratio('rerenderMs'), + scrollP50Factor: ratio('scrollMs'), + terminalBytesP50Factor: ratio('terminalBytes'), + toItems: result.itemCount + } + }) + + const oversizedRendererScaling = oversizedRendererResults.slice(1).map((result, index) => { + const previous = oversizedRendererResults[index]! + + const ratio = (metric: keyof (typeof result)['distributions']) => + result.distributions[metric].p50 / Math.max(Number.EPSILON, previous.distributions[metric].p50) + + return { + extentFactor: result.extent / previous.extent, + freshMountRenderP50Factor: ratio('freshMountRenderMs'), + fromExtent: previous.extent, + terminalBytesP50Factor: ratio('terminalBytes'), + terminalWritesP50Factor: ratio('terminalWrites'), + toExtent: result.extent + } + }) + + process.stdout.write( + `${JSON.stringify( + { + config: { columns: COLUMNS, maxMounted: MAX_MOUNTED, rows: ROWS, samples: samplesPerWorkload, warmups }, + oversizedRenderer: { + extents: RENDERER_EXTENTS, + results: oversizedRendererResults, + scaling: oversizedRendererScaling + }, + results, + scaling, + workloads + }, + null, + 2 + )}\n` + ) +} + +await main() diff --git a/ui-tui/src/__tests__/appChromeBlockedTimers.test.tsx b/ui-tui/src/__tests__/appChromeBlockedTimers.test.tsx new file mode 100644 index 0000000000..d683e841ac --- /dev/null +++ b/ui-tui/src/__tests__/appChromeBlockedTimers.test.tsx @@ -0,0 +1,449 @@ +import { PassThrough } from 'stream' + +import { renderSync } from '@hermes/ink' +import React from 'react' +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' + +import { GatewayProvider } from '../app/gatewayContext.js' +import type { AppLayoutProps, OverlayState, UiState } from '../app/interfaces.js' +import { patchOverlayState, resetOverlayState } from '../app/overlayStore.js' +import { patchUiState, resetUiState } from '../app/uiStore.js' +import { StatusRule } from '../components/appChrome.js' +import { AppLayout } from '../components/appLayout.js' +import type { GatewayClient } from '../gatewayClient.js' +import { DEFAULT_VOICE_RECORD_KEY } from '../lib/platform.js' +import { stripAnsi } from '../lib/text.js' +import { DEFAULT_THEME } from '../theme.js' + +type StatusRuleProps = React.ComponentProps +type IntervalSpy = ReturnType> + +// Fixed wall clock so the rendered elapsed read-outs are exact strings rather +// than whatever the machine's clock happens to produce mid-test. +const T0 = 1_800_000_000_000 + +const mounted: Array<() => void> = [] + +/** + * Mount a real StatusRule through Ink so the leaf components' effects — and + * therefore their `setInterval` calls — actually run. The existing + * appChromeStatusRule tests invoke `StatusRule(...)` as a plain function, + * which only builds the element tree and never mounts FaceTicker / + * SessionDuration / IdleSince, so it cannot observe timer behaviour. + * + * Teardown is registered up front so a failing assertion still unmounts the + * tree — a leaked instance would keep re-arming timers into the next test. + */ +const mountTree = (tree: React.ReactElement, { interactive = false } = {}) => { + const stdout = new PassThrough() + const stdin = new PassThrough() + const stderr = new PassThrough() + + let output = '' + + Object.assign(stdout, { columns: 120, isTTY: false, rows: 20 }) + // PromptZone's prompts call `useInput`, which needs raw mode; without it Ink + // swaps the whole tree for an error panel and stops updating. + Object.assign( + stdin, + interactive ? { isTTY: true, ref: () => {}, setRawMode: () => {}, unref: () => {} } : { isTTY: false } + ) + Object.assign(stderr, { isTTY: false }) + stdout.on('data', chunk => { + output += chunk.toString() + }) + + const instance = renderSync(tree, { + patchConsole: false, + stderr: stderr as NodeJS.WriteStream, + stdin: stdin as NodeJS.ReadStream, + stdout: stdout as NodeJS.WriteStream + }) + + mounted.push(() => { + instance.unmount() + instance.cleanup() + }) + + return { + /** Drop frames rendered so far so `output()` reads only what comes next. */ + clear: () => { + output = '' + }, + output: () => stripAnsi(output) + } +} + +const mount = (props: StatusRuleProps) => mountTree() + +const idleProps: StatusRuleProps = { + bgCount: 0, + busy: false, + cols: 120, + cwdLabel: '~/repo', + lastTurnEndedAt: T0 - 5_000, + liveSessionCount: 0, + model: 'opus-4.8', + sessionStartedAt: T0 - 60_000, + status: 'ready', + statusColor: DEFAULT_THEME.color.ok, + t: DEFAULT_THEME, + turnStartedAt: null, + usage: { context_max: 200_000, context_percent: 25, context_used: 50_000, total: 50_000 }, + voiceLabel: '' +} + +// Busy swaps the idle read-out for the FaceTicker, which owns the glyph + +// verb + elapsed-clock trio. +const busyProps: StatusRuleProps = { + ...idleProps, + busy: true, + indicatorStyle: 'kaomoji', + lastTurnEndedAt: null, + turnStartedAt: T0 - 30_000 +} + +/** Delays of every interval armed while the spy was installed. */ +const armedDelays = (spy: IntervalSpy) => spy.mock.calls.map(call => call[1]) + +const oneSecondTimers = (spy: IntervalSpy) => armedDelays(spy).filter(delay => delay === 1000).length + +/** The handlers of every 1s clock armed so far — `() => setNow(Date.now())`. */ +const oneSecondTicks = (spy: IntervalSpy) => + spy.mock.calls.filter(call => call[1] === 1000).map(call => call[0] as () => void) + +// ── AppLayout harness ──────────────────────────────────────────────── +// +// teknium1's review of this file was right that mounting StatusRule alone +// proves the store pauses timers but NOT that the overlay in question covers +// the status rule. These props render the real AppLayout so the rule sits in +// its true position relative to PromptZone / FloatingOverlays / the widget +// slot, and the assertions can read what is actually on screen. + +const gatewayStub = { + gw: { + request: () => new Promise(() => {}), + send: () => {} + } as unknown as GatewayClient, + rpc: (() => new Promise(() => {})) as never +} + +const layoutProps: AppLayoutProps = { + actions: { + activateLiveSession: () => {}, + answerApproval: () => {}, + answerClarify: () => {}, + answerSecret: () => {}, + answerSudo: () => {}, + clearSelection: () => {}, + closeLiveSession: () => Promise.resolve(null), + newLiveSession: () => {}, + newPromptSession: () => {}, + onModelSelect: () => {}, + resumeById: () => {}, + setStickyPrompt: () => {} + }, + composer: { + cols: 120, + compIdx: 0, + completions: [], + empty: true, + handleTextPaste: () => null, + input: '', + inputBuf: [], + pagerPageSize: 10, + queueEditIdx: null, + queuedDisplay: [], + submit: () => {}, + updateInput: () => {}, + voiceRecordKey: DEFAULT_VOICE_RECORD_KEY + }, + mouseTracking: 'off', + progress: { showProgressArea: false }, + status: { + cwdLabel: '~/repo', + goodVibesTick: 0, + lastTurnEndedAt: T0 - 5_000, + sessionStartedAt: T0 - 60_000, + showStickyPrompt: false, + statusColor: DEFAULT_THEME.color.ok, + stickyPrompt: '', + turnStartedAt: null, + voiceLabel: '' + }, + transcript: { + historyItems: [], + scrollRef: { current: null }, + virtualHistory: { + bottomSpacer: 0, + end: 0, + measureRef: () => () => {}, + offsets: [], + start: 0, + topSpacer: 0 + }, + virtualRows: [] + } +} + +/** Mount the real AppLayout with the given overlay + ui state applied first. */ +const mountLayout = (overlay: Partial = {}, ui: Partial = {}) => { + patchUiState({ sessionTitle: 'test', sid: 'sid-1', status: 'ready', ...ui }) + patchOverlayState(overlay) + + return mountTree( + + + , + { interactive: true } + ) +} + +// Give React's scheduler a turn so a store-driven re-render (and the effect +// re-arm that follows it) lands before we assert. +const flush = () => new Promise(resolve => setTimeout(resolve, 20)) + +let intervalSpy: IntervalSpy +let nowSpy: ReturnType> + +beforeEach(() => { + resetOverlayState() + resetUiState() + nowSpy = vi.spyOn(Date, 'now').mockReturnValue(T0) + intervalSpy = vi.spyOn(globalThis, 'setInterval') +}) + +afterEach(() => { + while (mounted.length > 0) { + mounted.pop()!() + } + + intervalSpy.mockRestore() + nowSpy.mockRestore() + resetOverlayState() + resetUiState() +}) + +describe('status-chrome timers under an occluding overlay', () => { + it('arms the one-second SessionDuration + IdleSince clocks when nothing covers the rule', () => { + mount(idleProps) + + expect(oneSecondTimers(intervalSpy)).toBe(2) + }) + + it('arms no timer at all when an occluding overlay is already open', () => { + patchOverlayState({ modelPicker: true }) + + mount(idleProps) + + expect(oneSecondTimers(intervalSpy)).toBe(0) + }) + + it('arms the FaceTicker glyph/verb/clock trio mid-turn when nothing covers the rule', () => { + mount(busyProps) + + // kaomoji cadence for the glyph + verb rotation, plus the elapsed clock. + expect(armedDelays(intervalSpy)).toContain(2500) + expect(oneSecondTimers(intervalSpy)).toBeGreaterThan(0) + }) + + it('arms no FaceTicker timer mid-turn while the modal widget slot is open', () => { + patchOverlayState({ widget: { appId: 'demo', state: null } }) + + mount(busyProps) + + expect(armedDelays(intervalSpy)).not.toContain(2500) + expect(oneSecondTimers(intervalSpy)).toBe(0) + }) + + it('keeps the FaceTicker running mid-turn under a flow-layout sudo prompt', () => { + // `sudo` is in `$isBlocked` but renders in PromptZone's normal flow, so it + // pushes the rule down rather than covering it — the trio must keep going. + patchOverlayState({ sudo: { requestId: 'sudo-1' } }) + + mount(busyProps) + + expect(armedDelays(intervalSpy)).toContain(2500) + expect(oneSecondTimers(intervalSpy)).toBeGreaterThan(0) + }) + + it('keeps the clocks running when a floating overlay cannot reach a bottom status rule', () => { + // FloatingOverlays is `position="absolute" bottom="100%"` inside + // ComposerPane's relative Box, so it grows UPWARD: it covers the `at="top"` + // rule and never the `at="bottom"` one. + patchUiState({ statusBar: 'bottom' }) + patchOverlayState({ modelPicker: true }) + + mount(idleProps) + + expect(oneSecondTimers(intervalSpy)).toBe(2) + }) + + it('re-syncs the elapsed read-outs from the wall clock on reveal instead of resuming stale', async () => { + // Regression guard for the naive fix: an early `return` that pauses the + // interval but never re-seeds `now` leaves SessionDuration and IdleSince + // frozen at the instant the overlay opened. + patchOverlayState({ sessions: true }) + + const rule = mount(idleProps) + + expect(rule.output()).toContain('1m 0s') + expect(rule.output()).toContain('✓ 5s') + + // Five minutes of wall clock elapse while the overlay covers the rule. + nowSpy.mockReturnValue(T0 + 300_000) + rule.clear() + resetOverlayState() + await flush() + + const resumed = rule.output() + + // Caught up to real elapsed time, not stuck on the pre-overlay values. + expect(resumed).toContain('6m 0s') + expect(resumed).toContain('✓ 5m 5s') + expect(resumed).not.toContain('1m 0s') + + // …and the clocks are running again. + expect(oneSecondTimers(intervalSpy)).toBe(2) + }) + + it('tears the clocks down when an overlay opens over an already-running status rule', async () => { + mount(idleProps) + + // Handles of the two live 1-second clocks (SessionDuration + IdleSince). + const clocks = intervalSpy.mock.results + .filter((_result, i) => intervalSpy.mock.calls[i]?.[1] === 1000) + .map(result => result.value as ReturnType) + + expect(clocks).toHaveLength(2) + + const clearSpy = vi.spyOn(globalThis, 'clearInterval') + + patchOverlayState({ pluginsHub: true }) + await flush() + + // Each running clock is cleared as the overlay goes up … + for (const handle of clocks) { + expect(clearSpy).toHaveBeenCalledWith(handle) + } + + // … and the occluded re-run arms no replacement (still just the original two). + expect(oneSecondTimers(intervalSpy)).toBe(2) + + clearSpy.mockRestore() + }) +}) + +// teknium1's review of #12463 called out that its test asserted on a `picker` +// overlay state that no longer exists. Pin the gate to fields the current +// OverlayState actually carries so a rename breaks this file loudly. +describe('status-chrome timers track the current overlay model', () => { + // Everything that genuinely paints over the rule: the modal widget slot, + // plus the FloatingOverlays set (with the rule at its default `top`). + const occluding: Array<[string, Partial]> = [ + ['modelPicker', { modelPicker: true }], + ['pager', { pager: { lines: ['a'], offset: 0 } }], + ['petPicker', { petPicker: true }], + ['pluginsHub', { pluginsHub: true }], + ['sessions', { sessions: true }], + ['skillsHub', { skillsHub: true }], + ['widget', { widget: { appId: 'demo', state: null } }] + ] + + // In `$isBlocked` but NOT occluding. `agents` / `journey` unmount the whole + // ComposerPane subtree, so React's effect cleanup already stops the clocks + // and gating on them would be dead code; the rest are PromptZone states that + // render in normal flow and push the rule down without covering it. + const nonOccluding: Array<[string, Partial]> = [ + ['agents', { agents: true }], + ['approval', { approval: { command: 'ls', requestId: 'a-1' } as OverlayState['approval'] }], + ['billing', { billing: { kind: 'credits' } as OverlayState['billing'] }], + ['clarify', { clarify: { question: 'which?', requestId: 'c-1' } as OverlayState['clarify'] }], + ['confirm', { confirm: { onConfirm: () => {}, prompt: 'sure?' } as OverlayState['confirm'] }], + ['journey', { journey: true }], + ['secret', { secret: { envVar: 'TOKEN', prompt: 'token?' } as OverlayState['secret'] }], + ['subscription', { subscription: { kind: 'expired' } as OverlayState['subscription'] }], + ['sudo', { sudo: { requestId: 'sudo-1' } as OverlayState['sudo'] }] + ] + + it.each(occluding)('pauses the status clocks while %s covers the rule', (_name, patch) => { + patchOverlayState(patch) + + mount(idleProps) + + expect(oneSecondTimers(intervalSpy)).toBe(0) + }) + + it.each(nonOccluding)('keeps the status clocks running while %s is open', (_name, patch) => { + patchOverlayState(patch) + + mount(idleProps) + + expect(oneSecondTimers(intervalSpy)).toBe(2) + }) + + it('keeps the clocks running for the non-occluding ambient dock', () => { + // `ambient` is a glanceable in-flow dock that reserves its own rows and + // doesn't cover the status rule, so pausing there would be a regression. + patchOverlayState({ ambient: [{ appId: 'clock', state: null }] }) + + mount(idleProps) + + expect(oneSecondTimers(intervalSpy)).toBe(2) + }) +}) + +// The visibility gate teknium1 asked for: mount the REAL AppLayout so the +// status rule sits in its true position relative to PromptZone (normal flow, +// above ComposerPane) and FloatingOverlays (absolute, growing upward), then +// assert on what is actually on screen rather than on the store alone. +describe('AppLayout status-rule visibility', () => { + it('keeps the status rule on screen AND its clock advancing under a flow-layout approval prompt', async () => { + const layout = mountLayout({ approval: { command: 'rm -rf /', requestId: 'a-1' } as OverlayState['approval'] }) + + await flush() + + // The rule is genuinely rendered — the approval prompt pushed it, it did + // not cover it — so freezing its clock would freeze something visible. + expect(layout.output()).toContain('~/repo') + expect(layout.output()).toContain('1m 0s') + expect(oneSecondTimers(intervalSpy)).toBe(2) + + // …and it really advances: drive the armed 1s handlers forward. + nowSpy.mockReturnValue(T0 + 30_000) + + for (const tick of oneSecondTicks(intervalSpy)) { + tick() + } + + await flush() + await flush() + + expect(layout.output()).toContain('1m 30s') + }) + + it('keeps the status rule on screen AND its clock advancing under a flow-layout sudo prompt', async () => { + const layout = mountLayout({ sudo: { requestId: 'sudo-1' } as OverlayState['sudo'] }) + + await flush() + + expect(layout.output()).toContain('1m 0s') + expect(oneSecondTimers(intervalSpy)).toBe(2) + }) + + it('arms no clock under a floating model picker while the rule is at the top', async () => { + mountLayout({ modelPicker: true }, { statusBar: 'top' }) + + await flush() + + expect(oneSecondTimers(intervalSpy)).toBe(0) + }) + + it('keeps the clocks armed under a floating model picker while the rule is at the bottom', async () => { + mountLayout({ modelPicker: true }, { statusBar: 'bottom' }) + + await flush() + + expect(oneSecondTimers(intervalSpy)).toBe(2) + }) +}) diff --git a/ui-tui/src/__tests__/messages.test.ts b/ui-tui/src/__tests__/messages.test.ts index d5baa1b311..e572bd5b8c 100644 --- a/ui-tui/src/__tests__/messages.test.ts +++ b/ui-tui/src/__tests__/messages.test.ts @@ -5,8 +5,9 @@ import React from 'react' import { describe, expect, it } from 'vitest' import { MessageLine } from '../components/messageLine.js' +import { MAX_HISTORY } from '../config/limits.js' import { toTranscriptMessages } from '../domain/messages.js' -import { upsert } from '../lib/messages.js' +import { capTranscriptHistory, upsert } from '../lib/messages.js' import { stripAnsi } from '../lib/text.js' import { DEFAULT_THEME } from '../theme.js' @@ -148,3 +149,16 @@ describe('upsert', () => { expect(prev).toHaveLength(1) }) }) + +describe('capTranscriptHistory', () => { + it('keeps the intro and the newest bounded display rows', () => { + const intro = { kind: 'intro' as const, role: 'system' as const, text: '' } + const rows = Array.from({ length: 1_005 }, (_, index) => ({ role: 'user' as const, text: `m${index}` })) + const capped = capTranscriptHistory([intro, ...rows]) + + expect(capped).toHaveLength(MAX_HISTORY) + expect(capped[0]).toBe(intro) + expect(capped[1]?.text).toBe(`m${rows.length - (MAX_HISTORY - 1)}`) + expect(capped.at(-1)?.text).toBe('m1004') + }) +}) diff --git a/ui-tui/src/__tests__/scroll.test.ts b/ui-tui/src/__tests__/scroll.test.ts index b9bbdb5fea..89f2356b92 100644 --- a/ui-tui/src/__tests__/scroll.test.ts +++ b/ui-tui/src/__tests__/scroll.test.ts @@ -13,12 +13,13 @@ function makeScroll(overrides: Partial> = {}) { getViewportHeight: vi.fn(() => 20), getViewportTop: vi.fn(() => 0), scrollBy: vi.fn(), + scrollTo: vi.fn(), ...overrides } } describe('scrollWithSelectionBy', () => { - it('clamps to the actual remaining scroll distance before calling scrollBy', () => { + it('commits the clamped target directly instead of queueing a scroll delta', () => { const s = makeScroll({ getScrollHeight: vi.fn(() => 30), getScrollTop: vi.fn(() => 9), @@ -34,7 +35,8 @@ describe('scrollWithSelectionBy', () => { scrollWithSelectionBy(10, { scrollRef: { current: s as never }, selection }) - expect(s.scrollBy).toHaveBeenCalledWith(1) + expect(s.scrollTo).toHaveBeenCalledWith(10) + expect(s.scrollBy).not.toHaveBeenCalled() }) it('uses fresh scroll height when cached height would swallow a down-scroll at a fake bottom', () => { @@ -54,7 +56,8 @@ describe('scrollWithSelectionBy', () => { scrollWithSelectionBy(10, { scrollRef: { current: s as never }, selection }) - expect(s.scrollBy).toHaveBeenCalledWith(4) + expect(s.scrollTo).toHaveBeenCalledWith(14) + expect(s.scrollBy).not.toHaveBeenCalled() }) it('uses fresh height when pending down-scroll reaches the cached fake bottom', () => { @@ -75,7 +78,8 @@ describe('scrollWithSelectionBy', () => { scrollWithSelectionBy(10, { scrollRef: { current: s as never }, selection }) - expect(s.scrollBy).toHaveBeenCalledWith(6) + expect(s.scrollTo).toHaveBeenCalledWith(18) + expect(s.scrollBy).not.toHaveBeenCalled() }) it('does nothing at the edge instead of queueing dead pending deltas', () => { @@ -94,6 +98,29 @@ describe('scrollWithSelectionBy', () => { scrollWithSelectionBy(10, { scrollRef: { current: s as never }, selection }) + expect(s.scrollTo).not.toHaveBeenCalled() + expect(s.scrollBy).not.toHaveBeenCalled() + }) + + it('preserves selection capture and shifting on the direct path', () => { + const s = makeScroll({ + getScrollTop: vi.fn(() => 10), + getViewportHeight: vi.fn(() => 20), + getViewportTop: vi.fn(() => 5) + }) + + const selection = { + captureScrolledRows: vi.fn(), + getState: vi.fn(() => ({ anchor: { row: 10 }, focus: { row: 12 } })), + shiftAnchor: vi.fn(), + shiftSelection: vi.fn() + } + + scrollWithSelectionBy(3, { scrollRef: { current: s as never }, selection }) + + expect(selection.captureScrolledRows).toHaveBeenCalledWith(5, 7, 'above') + expect(selection.shiftSelection).toHaveBeenCalledWith(-3, 5, 24) + expect(s.scrollTo).toHaveBeenCalledWith(13) expect(s.scrollBy).not.toHaveBeenCalled() }) }) diff --git a/ui-tui/src/__tests__/scrollBoxRendererBounds.test.ts b/ui-tui/src/__tests__/scrollBoxRendererBounds.test.ts new file mode 100644 index 0000000000..5a413400cb --- /dev/null +++ b/ui-tui/src/__tests__/scrollBoxRendererBounds.test.ts @@ -0,0 +1,626 @@ +import { PassThrough } from 'stream' + +import { Box, renderSync, ScrollBox, type ScrollBoxHandle, Text } from '@hermes/ink' +import React, { useLayoutEffect, useRef } from 'react' +import { describe, expect, it, vi } from 'vitest' + +import SourceBox from '../../packages/hermes-ink/src/ink/components/Box.js' +import SourceScrollBox from '../../packages/hermes-ink/src/ink/components/ScrollBox.js' +import SourceText from '../../packages/hermes-ink/src/ink/components/Text.js' +import type { DOMElement } from '../../packages/hermes-ink/src/ink/dom.js' +import Output from '../../packages/hermes-ink/src/ink/output.js' +import { scrollFastPathStats as sourceScrollFastPathStats } from '../../packages/hermes-ink/src/ink/render-node-to-output.js' +import { renderSync as renderSourceSync } from '../../packages/hermes-ink/src/ink/root.js' +import { useVirtualHistory } from '../hooks/useVirtualHistory.js' + +interface Item { + height: number + heightAfterResize?: number + key: string + text?: string +} + +interface Exposed { + scroll: ScrollBoxHandle | null + virtualHistory: ReturnType +} + +const delay = (ms: number) => new Promise(resolve => setTimeout(resolve, ms)) + +const makeStreams = () => { + const stdout = new PassThrough() + const stdin = new PassThrough() + const stderr = new PassThrough() + + Object.assign(stdout, { columns: 80, isTTY: false, rows: 20 }) + Object.assign(stdin, { isTTY: false }) + Object.assign(stderr, { isTTY: false }) + stdout.on('data', () => {}) + + return { stderr, stdin, stdout } +} + +const itemHeightForColumns = (item: Item | undefined, columns: number) => + columns >= 80 ? (item?.heightAfterResize ?? item?.height ?? 1) : (item?.height ?? 1) + +function Harness({ + columns = 80, + expose, + height = 10, + generation = 0, + initialHeights, + items, + maxMounted = 16 +}: { + columns?: number + expose: React.MutableRefObject + height?: number + generation?: number + initialHeights?: ReadonlyMap + items: readonly Item[] + maxMounted?: number +}) { + const scrollRef = useRef(null) + + const virtualHistory = useVirtualHistory(scrollRef, items, columns, { + coldStartCount: 16, + estimateHeight: index => itemHeightForColumns(items[index], columns), + generation, + initialHeights, + maxMounted, + overscan: 2 + }) + + useLayoutEffect(() => { + expose.current = { scroll: scrollRef.current, virtualHistory } + }) + + return React.createElement( + ScrollBox, + { flexDirection: 'column', height, ref: scrollRef, stickyScroll: true }, + React.createElement( + Box, + { flexDirection: 'column', width: '100%' }, + virtualHistory.topSpacer > 0 ? React.createElement(Box, { height: virtualHistory.topSpacer }) : null, + ...items.slice(virtualHistory.start, virtualHistory.end).map(item => + React.createElement( + Box, + { + height: itemHeightForColumns(item, columns), + key: item.key, + ref: virtualHistory.measureRef(item.key) + }, + React.createElement(Text, null, item.text ?? item.key) + ) + ), + virtualHistory.bottomSpacer > 0 ? React.createElement(Box, { height: virtualHistory.bottomSpacer }) : null + ) + ) +} + +function CorruptGeometryHarness({ expose, tick }: { expose: React.MutableRefObject; tick: number }) { + const nodes = useRef([]) + + useLayoutEffect(() => { + expose.current = nodes.current + }) + + return React.createElement( + ScrollBox, + { flexDirection: 'column', height: 8 }, + ...['nan-top', 'positive-infinity', 'negative-infinity', 'billion-rows', 'clipped-huge-fill'].map((label, index) => + React.createElement( + Box, + { + backgroundColor: 'blue', + borderStyle: label === 'clipped-huge-fill' ? 'single' : undefined, + height: 1, + key: label, + opaque: true, + ref: node => { + if (node) { + nodes.current[index] = node + } + }, + width: '100%' + }, + React.createElement(Text, null, `${label}-${tick}`) + ) + ), + React.createElement( + Box, + { + height: 1, + ref: node => { + if (node) { + nodes.current[5] = node + } + } + }, + React.createElement(Text, null, React.createElement(Text, null, `nested-corrupt-${tick}`)) + ), + React.createElement(Text, null, `adjacent-valid-${tick}`), + React.createElement(Box, { height: 20 }, React.createElement(Text, null, `tail-${tick}`)) + ) +} + +interface FastPathRepairExpose { + adjacent: DOMElement | null + dirtyChild: DOMElement | null + overlay: DOMElement | null + scroll: ScrollBoxHandle | null + scrollBox: DOMElement | null +} + +function FastPathRepairHarness({ + expose, + tick, + dirtyTick = tick, + includeOverlay = true +}: { + dirtyTick?: number + expose: React.MutableRefObject + includeOverlay?: boolean + tick: number +}) { + return React.createElement( + SourceBox, + { flexDirection: 'column', height: 12, width: 40 }, + React.createElement( + SourceScrollBox, + { + flexDirection: 'column', + height: 8, + ref: scroll => { + if (expose.current) { + expose.current.scroll = scroll + } + }, + width: 40 + }, + React.createElement(SourceBox, { height: 2 }, React.createElement(SourceText, null, 'head-row')), + React.createElement( + SourceBox, + { + height: 2, + ref: dirtyChild => { + if (expose.current) { + expose.current.dirtyChild = dirtyChild + expose.current.scrollBox = dirtyChild?.parentNode?.parentNode ?? null + } + } + }, + React.createElement(SourceText, null, `dirty-row-${dirtyTick}`) + ), + React.createElement(SourceBox, { height: 20 }, React.createElement(SourceText, null, 'tail-row')) + ), + includeOverlay + ? React.createElement( + SourceBox, + { + height: 2, + left: 0, + position: 'absolute', + ref: overlay => { + if (expose.current) { + expose.current.overlay = overlay + } + }, + top: 2, + width: 40 + }, + React.createElement(SourceText, null, 'overlay-row') + ) + : null, + React.createElement( + SourceBox, + { + height: 1, + ref: adjacent => { + if (expose.current) { + expose.current.adjacent = adjacent + } + } + }, + React.createElement(SourceText, null, `adjacent-fast-path-${tick}`) + ) + ) +} + +function guardFastPathRepairAllocations(maxWidth: number, maxHeight: number) { + const originalArrayFill = Array.prototype.fill + const originalBlit = Output.prototype.blit + const originalClear = Output.prototype.clear + const originalRepeat = String.prototype.repeat + const originalWrite = Output.prototype.write + + const observed = { + largestArrayRows: 0, + largestBlitHeight: 0, + largestBlitWidth: 0, + largestClearHeight: 0, + largestClearWidth: 0, + largestRepeat: 0, + largestWrite: 0, + repairWhitespaceWrites: 0 + } + + vi.spyOn(String.prototype, 'repeat').mockImplementation(function (count: number) { + observed.largestRepeat = Math.max(observed.largestRepeat, count) + + if (!Number.isSafeInteger(count) || count < 0 || count > maxWidth) { + throw new Error(`unbounded fast-path repeat: ${count}`) + } + + return originalRepeat.call(this, count) + }) + vi.spyOn(Array.prototype, 'fill').mockImplementation(function ( + this: unknown[], + value: unknown, + start?: number, + end?: number + ) { + observed.largestArrayRows = Math.max(observed.largestArrayRows, this.length) + + if (this.length > maxHeight) { + throw new Error(`unbounded fast-path row array: ${this.length}`) + } + + return Reflect.apply(originalArrayFill, this, [value, start, end]) + } as typeof Array.prototype.fill) + vi.spyOn(Output.prototype, 'blit').mockImplementation(function (...args: Parameters) { + observed.largestBlitWidth = Math.max(observed.largestBlitWidth, args[3]) + observed.largestBlitHeight = Math.max(observed.largestBlitHeight, args[4]) + + return originalBlit.apply(this, args) + }) + vi.spyOn(Output.prototype, 'clear').mockImplementation(function (...args: Parameters) { + observed.largestClearWidth = Math.max(observed.largestClearWidth, args[0].width) + observed.largestClearHeight = Math.max(observed.largestClearHeight, args[0].height) + + return originalClear.apply(this, args) + }) + vi.spyOn(Output.prototype, 'write').mockImplementation(function (...args: Parameters) { + observed.largestWrite = Math.max(observed.largestWrite, args[2].length) + + if (args[2].length > 0 && /^[ \n]+$/.test(args[2])) { + observed.repairWhitespaceWrites++ + } + + if (args[2].length > maxWidth * maxHeight + maxHeight) { + throw new Error(`unbounded fast-path Output.write input: ${args[2].length}`) + } + + return originalWrite.apply(this, args) + }) + + return observed +} + +describe('ScrollBox renderer bounds', () => { + it('rejects invalid imperative geometry without poisoning scroll state', async () => { + const items = Array.from({ length: 20 }, (_, index) => ({ height: 2, key: `item-${index}` })) + const expose = { current: null as Exposed | null } + const streams = makeStreams() + + const instance = renderSync(React.createElement(Harness, { expose, items }), { + patchConsole: false, + stderr: streams.stderr as NodeJS.WriteStream, + stdin: streams.stdin as NodeJS.ReadStream, + stdout: streams.stdout as NodeJS.WriteStream + }) + + try { + await delay(20) + const scroll = expose.current!.scroll! + + scroll.scrollTo(4) + scroll.scrollTo(Number.NaN) + scroll.scrollBy(Number.POSITIVE_INFINITY) + scroll.adjustScrollTop(Number.NEGATIVE_INFINITY) + scroll.setClampBounds(Number.NaN, Number.POSITIVE_INFINITY) + await delay(20) + + expect(scroll.getScrollTop()).toBe(4) + expect(scroll.getPendingDelta()).toBe(0) + expect(Number.isFinite(scroll.getScrollHeight())).toBe(true) + } finally { + instance.unmount() + instance.cleanup() + } + }) + + it('fails closed on corrupt ScrollBox child geometry and keeps adjacent rows renderable', async () => { + const expose = { current: [] as DOMElement[] } + const streams = makeStreams() + const originalRepeat = String.prototype.repeat + const originalWrite = Output.prototype.write + let largestWriteInput = 0 + let largestWrite = 0 + let output = '' + + vi.spyOn(String.prototype, 'repeat').mockImplementation(function (count: number) { + if (!Number.isSafeInteger(count) || count < 0 || count > 10_000) { + throw new Error(`unbounded string repeat: ${count}`) + } + + return originalRepeat.call(this, count) + }) + vi.spyOn(Output.prototype, 'write').mockImplementation(function (...args: Parameters) { + largestWriteInput = Math.max(largestWriteInput, args[2].length) + + if (args[2].length > 10_000) { + throw new Error(`unbounded Output.write input: ${args[2].length}`) + } + + return originalWrite.apply(this, args) + }) + + streams.stdout.removeAllListeners('data') + streams.stdout.on('data', chunk => { + largestWrite = Math.max(largestWrite, chunk.length) + output += chunk.toString() + }) + + const instance = renderSync(React.createElement(CorruptGeometryHarness, { expose, tick: 0 }), { + patchConsole: false, + stderr: streams.stderr as NodeJS.WriteStream, + stdin: streams.stdin as NodeJS.ReadStream, + stdout: streams.stdout as NodeJS.WriteStream + }) + + try { + await delay(20) + + const [nanTop, positiveInfinity, negativeInfinity, billionRows, clippedHugeFill, nestedTextWrapper] = + expose.current + + expect(nanTop?.yogaNode).toBeDefined() + expect(positiveInfinity?.yogaNode).toBeDefined() + expect(negativeInfinity?.yogaNode).toBeDefined() + expect(billionRows?.yogaNode).toBeDefined() + expect(clippedHugeFill?.yogaNode).toBeDefined() + expect(nestedTextWrapper?.yogaNode).toBeDefined() + + const nestedText = nestedTextWrapper!.childNodes.find(child => child.nodeName === 'ink-text') as + DOMElement | undefined + + const nestedTextChild = nestedText?.childNodes[0] + + expect(nestedTextChild).toBeDefined() + + vi.spyOn(nanTop!.yogaNode!, 'getComputedTop').mockReturnValue(Number.NaN) + vi.spyOn(positiveInfinity!.yogaNode!, 'getComputedHeight').mockReturnValue(Number.POSITIVE_INFINITY) + vi.spyOn(negativeInfinity!.yogaNode!, 'getComputedHeight').mockReturnValue(Number.NEGATIVE_INFINITY) + vi.spyOn(billionRows!.yogaNode!, 'getComputedHeight').mockReturnValue(1_000_000_000) + vi.spyOn(clippedHugeFill!.yogaNode!, 'getComputedHeight').mockReturnValue(100_000_000) + vi.spyOn(clippedHugeFill!.yogaNode!, 'getComputedWidth').mockReturnValue(100_000_000) + + let nestedOffsetX = 0 + let nestedOffsetY = 0 + + nestedTextChild!.yogaNode = { + getComputedLeft: () => nestedOffsetX, + getComputedTop: () => nestedOffsetY + } as DOMElement['yogaNode'] + + output = '' + largestWrite = 0 + largestWriteInput = 0 + + const corruptNestedOffsets = [ + [Number.NaN, 0], + [Number.POSITIVE_INFINITY, 0], + [Number.NEGATIVE_INFINITY, 0], + [-1, 0], + [0.5, 0], + [100_000_000, 0], + [0, Number.NaN], + [0, Number.POSITIVE_INFINITY], + [0, Number.NEGATIVE_INFINITY], + [0, -1], + [0, 0.5], + [0, 100_000_000] + ] as const + + for (const [index, [offsetX, offsetY]] of corruptNestedOffsets.entries()) { + nestedOffsetX = offsetX + nestedOffsetY = offsetY + + expect(() => { + instance.rerender(React.createElement(CorruptGeometryHarness, { expose, tick: index + 1 })) + }).not.toThrow() + await delay(5) + } + + await delay(40) + + expect(output).toContain(`adjacent-valid-${corruptNestedOffsets.length}`) + expect(largestWriteInput).toBeLessThan(10_000) + expect(largestWrite).toBeLessThan(10_000) + expect(output.length).toBeLessThan(50_000) + } finally { + vi.restoreAllMocks() + instance.unmount() + instance.cleanup() + } + }) + + it('clips corrupt dirty-child DECSTBM repairs before allocating or recursing', async () => { + const expose = { + current: { + adjacent: null, + dirtyChild: null, + overlay: null, + scroll: null, + scrollBox: null + } as FastPathRepairExpose + } + + const streams = makeStreams() + let output = '' + + const instance = renderSourceSync( + React.createElement(FastPathRepairHarness, { expose, includeOverlay: false, tick: 0 }), + { + patchConsole: false, + stderr: streams.stderr as NodeJS.WriteStream, + stdin: streams.stdin as NodeJS.ReadStream, + stdout: streams.stdout as NodeJS.WriteStream + } + ) + + try { + await delay(20) + const dirtyChild = expose.current!.dirtyChild! + + expect(dirtyChild, streams.stderr.read()?.toString()).not.toBeNull() + expect(dirtyChild.yogaNode).toBeDefined() + vi.spyOn(dirtyChild.yogaNode!, 'getComputedTop').mockReturnValue(-99_999_995) + vi.spyOn(dirtyChild.yogaNode!, 'getComputedWidth').mockReturnValue(100_000_000) + vi.spyOn(dirtyChild.yogaNode!, 'getComputedHeight').mockReturnValue(100_000_000) + + instance.rerender(React.createElement(FastPathRepairHarness, { expose, includeOverlay: false, tick: 1 })) + await delay(20) + + streams.stdout.removeAllListeners('data') + streams.stdout.on('data', chunk => { + output += chunk.toString() + }) + + const observed = guardFastPathRepairAllocations(80, 20) + const fastPathsBefore = sourceScrollFastPathStats.taken + const capturedBefore = sourceScrollFastPathStats.captured + + expect(() => expose.current!.scroll!.scrollTo(1)).not.toThrow() + expect(() => + instance.rerender(React.createElement(FastPathRepairHarness, { expose, includeOverlay: false, tick: 2 })) + ).not.toThrow() + await delay(40) + + expect(sourceScrollFastPathStats.captured, JSON.stringify(sourceScrollFastPathStats)).toBeGreaterThan( + capturedBefore + ) + expect(sourceScrollFastPathStats.taken, JSON.stringify(sourceScrollFastPathStats)).toBeGreaterThan( + fastPathsBefore + ) + expect(observed.repairWhitespaceWrites).toBeGreaterThan(0) + expect(observed.largestRepeat).toBeGreaterThan(0) + expect(observed.largestArrayRows).toBeGreaterThan(0) + + output = '' + expect(() => + instance.rerender( + React.createElement(FastPathRepairHarness, { + dirtyTick: 2, + expose, + includeOverlay: false, + tick: 3 + }) + ) + ).not.toThrow() + await delay(40) + + expect(observed.largestRepeat).toBeLessThanOrEqual(80) + expect(observed.largestArrayRows).toBeLessThanOrEqual(20) + expect(observed.largestBlitWidth).toBeLessThanOrEqual(80) + expect(observed.largestBlitHeight).toBeLessThanOrEqual(20) + expect(observed.largestClearWidth).toBeLessThanOrEqual(80) + expect(observed.largestClearHeight).toBeLessThanOrEqual(20) + expect(observed.largestWrite).toBeLessThanOrEqual(1_620) + expect(expose.current!.scroll!.getScrollTop()).toBe(1) + expect(expose.current!.scroll!.getViewportHeight()).toBeGreaterThan(0) + expect(output).toContain('adjacent-fast-path-3') + expect(streams.stderr.read()?.toString() ?? '').toBe('') + } finally { + vi.restoreAllMocks() + instance.unmount() + instance.cleanup() + } + }) + + it('clips corrupt absolute-overlay DECSTBM repairs before allocating or recursing', async () => { + const expose = { + current: { + adjacent: null, + dirtyChild: null, + overlay: null, + scroll: null, + scrollBox: null + } as FastPathRepairExpose + } + + const streams = makeStreams() + let output = '' + + const instance = renderSourceSync(React.createElement(FastPathRepairHarness, { expose, tick: 0 }), { + patchConsole: false, + stderr: streams.stderr as NodeJS.WriteStream, + stdin: streams.stdin as NodeJS.ReadStream, + stdout: streams.stdout as NodeJS.WriteStream + }) + + try { + await delay(20) + const overlay = expose.current!.overlay! + const scrollBox = expose.current!.scrollBox! + + expect(overlay, streams.stderr.read()?.toString()).not.toBeNull() + expect(overlay.yogaNode).toBeDefined() + expect(scrollBox.yogaNode).toBeDefined() + vi.spyOn(scrollBox.yogaNode!, 'getComputedWidth').mockReturnValue(100_000_000) + vi.spyOn(overlay.yogaNode!, 'getComputedWidth').mockReturnValue(100_000_000) + vi.spyOn(overlay.yogaNode!, 'getComputedHeight').mockReturnValue(100_000_000) + + instance.rerender(React.createElement(FastPathRepairHarness, { expose, tick: 1 })) + await delay(20) + instance.rerender(React.createElement(FastPathRepairHarness, { expose, tick: 1 })) + await delay(20) + + streams.stdout.removeAllListeners('data') + streams.stdout.on('data', chunk => { + output += chunk.toString() + }) + + const observed = guardFastPathRepairAllocations(80, 20) + const fastPathsBefore = sourceScrollFastPathStats.taken + const capturedBefore = sourceScrollFastPathStats.captured + + expect(() => expose.current!.scroll!.scrollTo(1)).not.toThrow() + expect(expose.current!.scroll!.getScrollTop()).toBe(1) + await delay(40) + + expect(sourceScrollFastPathStats.captured, JSON.stringify(sourceScrollFastPathStats)).toBeGreaterThan( + capturedBefore + ) + expect(sourceScrollFastPathStats.taken, JSON.stringify(sourceScrollFastPathStats)).toBeGreaterThan( + fastPathsBefore + ) + expect(observed.repairWhitespaceWrites).toBeGreaterThan(0) + expect(observed.largestRepeat).toBeGreaterThan(0) + expect(observed.largestArrayRows).toBeGreaterThan(0) + + output = '' + expect(() => + instance.rerender(React.createElement(FastPathRepairHarness, { dirtyTick: 1, expose, tick: 2 })) + ).not.toThrow() + await delay(40) + + expect(observed.largestRepeat).toBeLessThanOrEqual(80) + expect(observed.largestArrayRows).toBeLessThanOrEqual(20) + expect(observed.largestBlitWidth).toBeLessThanOrEqual(80) + expect(observed.largestBlitHeight).toBeLessThanOrEqual(20) + expect(observed.largestClearWidth).toBeLessThanOrEqual(80) + expect(observed.largestClearHeight).toBeLessThanOrEqual(20) + expect(observed.largestWrite).toBeLessThanOrEqual(1_620) + expect(output.length).toBeLessThan(2_000) + expect(expose.current!.scroll!.getViewportHeight()).toBeGreaterThan(0) + expect(output).toContain('adjacent-fast-path-2') + expect(streams.stderr.read()?.toString() ?? '').toBe('') + } finally { + vi.restoreAllMocks() + instance.unmount() + instance.cleanup() + } + }) +}) diff --git a/ui-tui/src/__tests__/text.test.ts b/ui-tui/src/__tests__/text.test.ts index 2202b21994..707ad227ec 100644 --- a/ui-tui/src/__tests__/text.test.ts +++ b/ui-tui/src/__tests__/text.test.ts @@ -224,6 +224,17 @@ describe('edgePreview', () => { }) }) +describe('thinkingPreview over-bound tail', () => { + it('retains the live tail when reasoning exceeds the clean bound', () => { + const TAIL = '<<>>' + // Slightly above the 24k clean-tail bound, so the implementation must trim. + const reasoning = 'A'.repeat(25_000) + '\n' + TAIL + const result = thinkingPreview(reasoning, 'full') + expect(result).toContain(TAIL) + // The bounded window is shorter than the 25k prefix, but the tail remains. + expect(result.length).toBeLessThanOrEqual(25_000) + }) +}) describe('pasteTokenLabel', () => { it('builds readable long-paste labels with counts', () => { const label = pasteTokenLabel('Vampire Bondage ropes slipped from her neck, still stained with blood', 250) diff --git a/ui-tui/src/__tests__/useVirtualHistoryHeights.test.ts b/ui-tui/src/__tests__/useVirtualHistoryHeights.test.ts index ae5658f83e..c1aa868f87 100644 --- a/ui-tui/src/__tests__/useVirtualHistoryHeights.test.ts +++ b/ui-tui/src/__tests__/useVirtualHistoryHeights.test.ts @@ -36,4 +36,24 @@ describe('ensureVirtualItemHeight', () => { expect(ensureVirtualItemHeight(heights, 'd', 0, 0, estimateHeight)).toBe(1) expect(heights.get('d')).toBe(1) }) + + it.each([Number.NaN, Number.POSITIVE_INFINITY, Number.NEGATIVE_INFINITY, -4, 1_000_000_000])( + 'quarantines invalid cached height %s and reseeds it', + cached => { + const heights = new Map([['bad', cached]]) + + expect(ensureVirtualItemHeight(heights, 'bad', 0, 4, () => 7)).toBe(7) + expect(heights.get('bad')).toBe(7) + } + ) + + it.each([Number.NaN, Number.POSITIVE_INFINITY, Number.NEGATIVE_INFINITY, -4, 1_000_000_000])( + 'falls back when the estimator returns invalid height %s', + estimate => { + const heights = new Map() + + expect(ensureVirtualItemHeight(heights, 'bad', 0, 4, () => estimate)).toBe(4) + expect(heights.get('bad')).toBe(4) + } + ) }) diff --git a/ui-tui/src/__tests__/virtualHistoryOffsetCache.test.ts b/ui-tui/src/__tests__/virtualHistoryOffsetCache.test.ts index 010d12c9ca..e5bed9d966 100644 --- a/ui-tui/src/__tests__/virtualHistoryOffsetCache.test.ts +++ b/ui-tui/src/__tests__/virtualHistoryOffsetCache.test.ts @@ -2,14 +2,16 @@ import { PassThrough } from 'stream' import { Box, renderSync, ScrollBox, type ScrollBoxHandle, Text } from '@hermes/ink' import React, { useLayoutEffect, useRef } from 'react' -import { describe, expect, it } from 'vitest' +import { describe, expect, it, vi } from 'vitest' -import { useVirtualHistory, virtualHistorySnapshotKey } from '../hooks/useVirtualHistory.js' +import { MAX_HISTORY } from '../config/limits.js' +import { pruneVirtualHeightCache, useVirtualHistory, virtualHistorySnapshotKey } from '../hooks/useVirtualHistory.js' interface Item { height: number heightAfterResize?: number key: string + text?: string } interface Exposed { @@ -61,12 +63,16 @@ function Harness({ columns = 80, expose, height = 10, + generation = 0, + initialHeights, items, maxMounted = 16 }: { columns?: number expose: React.MutableRefObject height?: number + generation?: number + initialHeights?: ReadonlyMap items: readonly Item[] maxMounted?: number }) { @@ -75,6 +81,8 @@ function Harness({ const virtualHistory = useVirtualHistory(scrollRef, items, columns, { coldStartCount: 16, estimateHeight: index => itemHeightForColumns(items[index], columns), + generation, + initialHeights, maxMounted, overscan: 2 }) @@ -98,7 +106,7 @@ function Harness({ key: item.key, ref: virtualHistory.measureRef(item.key) }, - React.createElement(Text, null, item.key) + React.createElement(Text, null, item.text ?? item.key) ) ), virtualHistory.bottomSpacer > 0 ? React.createElement(Box, { height: virtualHistory.bottomSpacer }) : null @@ -107,6 +115,17 @@ function Harness({ } describe('useVirtualHistory offset cache reuse', () => { + it('prunes stable-session external height caches to active history keys', () => { + const cache = new Map([ + ['outgoing', 9], + ['active', 3] + ]) + + pruneVirtualHeightCache(cache, [{ key: 'active' }]) + + expect([...cache]).toEqual([['active', 3]]) + }) + it('includes viewport height in the external-store snapshot key', () => { const base = { getPendingDelta: () => 0, @@ -247,9 +266,358 @@ describe('useVirtualHistory offset cache reuse', () => { } }) + it('adjusts the committed viewport without consuming pending scroll intent', async () => { + const items = Array.from({ length: 20 }, (_, index) => ({ height: 2, key: `item-${index}` })) + const expose = { current: null as Exposed | null } + const streams = makeStreams() + + const instance = renderSync(React.createElement(Harness, { expose, items }), { + patchConsole: false, + stderr: streams.stderr as NodeJS.WriteStream, + stdin: streams.stdin as NodeJS.ReadStream, + stdout: streams.stdout as NodeJS.WriteStream + }) + + try { + await delay(20) + const scroll = expose.current!.scroll! + + scroll.scrollTo(3) + scroll.scrollBy(2) + scroll.adjustScrollTop(4) + + expect(scroll.getScrollTop()).toBe(7) + expect(scroll.getPendingDelta()).toBe(2) + expect(scroll.isSticky()).toBe(false) + } finally { + instance.unmount() + instance.cleanup() + } + }) + + it('keeps the tail clamp open while a manual non-sticky tail grows', async () => { + const before = Array.from({ length: 20 }, (_, index) => ({ height: 2, key: `item-${index}` })) + const after = before.map((item, index) => (index === before.length - 1 ? { ...item, height: 8 } : item)) + const expose = { current: null as Exposed | null } + const streams = makeStreams() + const initialHeights = new Map(before.map(item => [item.key, item.height])) + + const instance = renderSync(React.createElement(Harness, { expose, initialHeights, items: before }), { + patchConsole: false, + stderr: streams.stderr as NodeJS.WriteStream, + stdin: streams.stdin as NodeJS.ReadStream, + stdout: streams.stdout as NodeJS.WriteStream + }) + + try { + await delay(20) + const scroll = expose.current!.scroll! + const setClampBounds = vi.spyOn(scroll, 'setClampBounds') + + scroll.scrollTo(28) + await delay(20) + instance.rerender(React.createElement(Harness, { expose, initialHeights, items: after })) + await delay(60) + + expect(scroll.isSticky()).toBe(false) + expect(setClampBounds.mock.calls.some(([, max]) => max === Number.POSITIVE_INFINITY)).toBe(true) + + scroll.scrollTo(36) + await delay(20) + expect(scroll.getScrollTop()).toBe(36) + } finally { + instance.unmount() + instance.cleanup() + } + }) + + it('quarantines invalid measured heights before cache and compensation', async () => { + const items = Array.from({ length: 20 }, (_, index) => ({ height: 2, key: `item-${index}` })) + const expose = { current: null as Exposed | null } + const streams = makeStreams() + const initialHeights = new Map(items.map(item => [item.key, item.height])) + + const instance = renderSync(React.createElement(Harness, { expose, initialHeights, items }), { + patchConsole: false, + stderr: streams.stderr as NodeJS.WriteStream, + stdin: streams.stdin as NodeJS.ReadStream, + stdout: streams.stdout as NodeJS.WriteStream + }) + + try { + await delay(20) + const scroll = expose.current!.scroll! + + scroll.scrollTo(5) + await delay(20) + const adjustScrollTop = vi.spyOn(scroll, 'adjustScrollTop') + const ref = expose.current!.virtualHistory.measureRef('item-1') + + for (const height of [Number.NaN, Number.POSITIVE_INFINITY, -1, 1_000_000_000]) { + ref({ yogaNode: { getComputedHeight: () => height } }) + ref(null) + } + + expect(adjustScrollTop).not.toHaveBeenCalled() + expect(expose.current!.virtualHistory.offsets[items.length]).toBe(40) + expect(Number.isFinite(scroll.getScrollTop())).toBe(true) + } finally { + instance.unmount() + instance.cleanup() + } + }) + + it('preserves the visual anchor when a measured row above the viewport changes height', async () => { + const before = Array.from({ length: 20 }, (_, index) => ({ height: 2, key: `item-${index}` })) + const after = before.map((item, index) => (index === 0 ? { ...item, height: 5 } : item)) + const expose = { current: null as Exposed | null } + const streams = makeStreams() + const initialHeights = new Map(before.map(item => [item.key, item.height])) + + const instance = renderSync(React.createElement(Harness, { expose, initialHeights, items: before }), { + patchConsole: false, + stderr: streams.stderr as NodeJS.WriteStream, + stdin: streams.stdin as NodeJS.ReadStream, + stdout: streams.stdout as NodeJS.WriteStream + }) + + try { + await delay(20) + expose.current!.scroll!.scrollTo(3) + await delay(20) + + instance.rerender(React.createElement(Harness, { expose, initialHeights, items: after })) + await delay(40) + + expect(expose.current!.scroll!.getScrollTop()).toBe(6) + } finally { + instance.unmount() + instance.cleanup() + } + }) + + it('keeps a compensated near-tail viewport manual', async () => { + const before = Array.from({ length: 20 }, (_, index) => ({ height: 2, key: `item-${index}` })) + const after = before.map((item, index) => (index === 13 ? { ...item, height: 5 } : item)) + const expose = { current: null as Exposed | null } + const streams = makeStreams() + const initialHeights = new Map(before.map(item => [item.key, item.height])) + + const instance = renderSync(React.createElement(Harness, { expose, initialHeights, items: before }), { + patchConsole: false, + stderr: streams.stderr as NodeJS.WriteStream, + stdin: streams.stdin as NodeJS.ReadStream, + stdout: streams.stdout as NodeJS.WriteStream + }) + + try { + await delay(20) + const scroll = expose.current!.scroll! + + scroll.scrollTo(29) + await delay(20) + instance.rerender(React.createElement(Harness, { expose, initialHeights, items: after })) + + expect(scroll.getScrollTop()).toBe(32) + expect(scroll.isSticky()).toBe(false) + } finally { + instance.unmount() + instance.cleanup() + } + }) + + it('ignores stale unmount measurement from the previous width layout', async () => { + const items = Array.from({ length: 20 }, (_, index) => ({ + height: 4, + heightAfterResize: index === 0 ? 5 : 2, + key: `item-${index}` + })) + + const expose = { current: null as Exposed | null } + const streams = makeStreams() + const initialHeights = new Map(items.map(item => [item.key, item.height])) + + const instance = renderSync(React.createElement(Harness, { columns: 40, expose, initialHeights, items }), { + patchConsole: false, + stderr: streams.stderr as NodeJS.WriteStream, + stdin: streams.stdin as NodeJS.ReadStream, + stdout: streams.stdout as NodeJS.WriteStream + }) + + try { + await delay(20) + const scroll = expose.current!.scroll! + + scroll.scrollTo(0) + await delay(20) + scroll.scrollTo(5) + const adjustScrollTop = vi.spyOn(scroll, 'adjustScrollTop') + + instance.rerender(React.createElement(Harness, { columns: 80, expose, initialHeights, items })) + await delay(40) + + expect(adjustScrollTop).not.toHaveBeenCalled() + expect(scroll.getScrollTop()).toBe(5) + expect(scroll.isSticky()).toBe(false) + expect(expose.current!.virtualHistory.start).toBeGreaterThan(0) + expect(expose.current!.virtualHistory.offsets[1]).toBe(2) + expect(expose.current!.virtualHistory.offsets[items.length]).toBe(40) + } finally { + instance.unmount() + instance.cleanup() + } + }) + + it('does not let outgoing transcript refs compensate a new layout generation', async () => { + const outgoing = Array.from({ length: 20 }, (_, index) => ({ height: 2, key: `old-${index}` })) + const incoming = Array.from({ length: 20 }, (_, index) => ({ height: 2, key: `new-${index}` })) + const expose = { current: null as Exposed | null } + const streams = makeStreams() + const initialHeights = new Map(outgoing.map(item => [item.key, item.height])) + + const instance = renderSync(React.createElement(Harness, { expose, initialHeights, items: outgoing }), { + patchConsole: false, + stderr: streams.stderr as NodeJS.WriteStream, + stdin: streams.stdin as NodeJS.ReadStream, + stdout: streams.stdout as NodeJS.WriteStream + }) + + try { + await delay(20) + const scroll = expose.current!.scroll! + + scroll.scrollTo(5) + await delay(20) + const adjustScrollTop = vi.spyOn(scroll, 'adjustScrollTop') + + const replacementCache = new Map([ + ...outgoing.map(item => [item.key, item.key === 'old-1' ? 1 : item.height] as const), + ...incoming.map(item => [item.key, item.height] as const) + ]) + + instance.rerender( + React.createElement(Harness, { + expose, + generation: 1, + initialHeights: replacementCache, + items: incoming + }) + ) + await delay(40) + + expect(adjustScrollTop).not.toHaveBeenCalled() + expect(scroll.getScrollTop()).toBe(5) + expect(expose.current!.virtualHistory.offsets[incoming.length]).toBe(40) + } finally { + instance.unmount() + instance.cleanup() + } + }) + + it('corrects and compensates a same-layout row measured at unmount', async () => { + const items = Array.from({ length: 20 }, (_, index) => ({ height: 2, key: `item-${index}` })) + const expose = { current: null as Exposed | null } + const streams = makeStreams() + const initialHeights = new Map(items.map(item => [item.key, item.height])) + + const instance = renderSync(React.createElement(Harness, { expose, initialHeights, items }), { + patchConsole: false, + stderr: streams.stderr as NodeJS.WriteStream, + stdin: streams.stdin as NodeJS.ReadStream, + stdout: streams.stdout as NodeJS.WriteStream + }) + + try { + await delay(20) + const scroll = expose.current!.scroll! + + scroll.scrollTo(0) + await delay(20) + scroll.scrollTo(5) + const adjustScrollTop = vi.spyOn(scroll, 'adjustScrollTop') + const staleHeights = new Map(initialHeights) + + staleHeights.set(items[0]!.key, 1) + instance.rerender(React.createElement(Harness, { expose, initialHeights: staleHeights, items })) + await delay(40) + + expect(adjustScrollTop).toHaveBeenCalledOnce() + expect(adjustScrollTop).toHaveBeenCalledWith(1) + expect(scroll.getScrollTop()).toBe(6) + expect(scroll.isSticky()).toBe(false) + expect(expose.current!.virtualHistory.start).toBeGreaterThan(0) + expect(expose.current!.virtualHistory.offsets[1]).toBe(2) + } finally { + instance.unmount() + instance.cleanup() + } + }) + + it('does not compensate for measured height changes in or below the viewport', async () => { + const before = Array.from({ length: 20 }, (_, index) => ({ height: 2, key: `item-${index}` })) + const visibleChanged = before.map((item, index) => (index === 1 ? { ...item, height: 5 } : item)) + const belowChanged = visibleChanged.map((item, index) => (index === 8 ? { ...item, height: 5 } : item)) + const expose = { current: null as Exposed | null } + const streams = makeStreams() + const initialHeights = new Map(before.map(item => [item.key, item.height])) + + const instance = renderSync(React.createElement(Harness, { expose, initialHeights, items: before }), { + patchConsole: false, + stderr: streams.stderr as NodeJS.WriteStream, + stdin: streams.stdin as NodeJS.ReadStream, + stdout: streams.stdout as NodeJS.WriteStream + }) + + try { + await delay(20) + expose.current!.scroll!.scrollTo(3) + await delay(20) + + instance.rerender(React.createElement(Harness, { expose, initialHeights, items: visibleChanged })) + await delay(40) + expect(expose.current!.scroll!.getScrollTop()).toBe(3) + + instance.rerender(React.createElement(Harness, { expose, initialHeights, items: belowChanged })) + await delay(40) + expect(expose.current!.scroll!.getScrollTop()).toBe(3) + } finally { + instance.unmount() + instance.cleanup() + } + }) + + it('does not compensate measured heights while sticky at the live tail', async () => { + const before = Array.from({ length: 20 }, (_, index) => ({ height: 2, key: `item-${index}` })) + const after = before.map((item, index) => (index === 14 ? { ...item, height: 5 } : item)) + const expose = { current: null as Exposed | null } + const streams = makeStreams() + const initialHeights = new Map(before.map(item => [item.key, item.height])) + + const instance = renderSync(React.createElement(Harness, { expose, initialHeights, items: before }), { + patchConsole: false, + stderr: streams.stderr as NodeJS.WriteStream, + stdin: streams.stdin as NodeJS.ReadStream, + stdout: streams.stdout as NodeJS.WriteStream + }) + + try { + await delay(20) + const adjustScrollTop = vi.spyOn(expose.current!.scroll!, 'adjustScrollTop') + + instance.rerender(React.createElement(Harness, { expose, initialHeights, items: after })) + await delay(40) + + expect(adjustScrollTop).not.toHaveBeenCalled() + expect(expose.current!.scroll!.isSticky()).toBe(true) + } finally { + instance.unmount() + instance.cleanup() + } + }) + it('ignores stale reused offset-array entries after the item count shrinks', async () => { const beforeShrink = Array.from({ length: 1400 }, (_, index) => ({ height: 1, key: `old${index}` })) - const afterShrink = Array.from({ length: 800 }, (_, index) => ({ height: 7, key: `new${index}` })) + const afterShrink = Array.from({ length: MAX_HISTORY }, (_, index) => ({ height: 7, key: `new${index}` })) const expose = { current: null as Exposed | null } const streams = makeStreams() diff --git a/ui-tui/src/app/overlayStore.ts b/ui-tui/src/app/overlayStore.ts index 4e26099a06..92d5265654 100644 --- a/ui-tui/src/app/overlayStore.ts +++ b/ui-tui/src/app/overlayStore.ts @@ -1,6 +1,7 @@ import { atom, computed } from 'nanostores' import type { OverlayState } from './interfaces.js' +import { $uiState } from './uiStore.js' const buildOverlayState = (): OverlayState => ({ agents: false, @@ -65,6 +66,71 @@ export const $isBlocked = computed( ) ) +/** + * Does an open overlay actually PAINT OVER the status rule? + * + * Deliberately NOT `$isBlocked`. That aggregate answers a different + * question — "is text input suspended" — and `appLayout` uses it only to hide + * the input rows (`appLayout.tsx:384`). `StatusRulePane` is rendered OUTSIDE + * that guard (`appLayout.tsx:365` for `at="top"`, `:449` for `at="bottom"`), + * so most `$isBlocked` fields leave the rule fully visible. + * + * Occluding (included here): + * + * - `widget` — the modal widget slot renders at viewport level + * (`ActiveWidgetSlot`, `sdk/host.tsx:209`, outside the ComposerPane + * subtree) so it can anchor the full-screen absolute `Overlay` + * (`components/overlay.tsx`) against the whole terminal. + * - The FloatingOverlays set — `modelPicker`, `pager`, `petPicker`, + * `sessions`, `skillsHub`, `pluginsHub` — but ONLY when the rule sits at + * the top. That panel is `position="absolute" bottom="100%"` inside + * ComposerPane's relative Box (`appOverlays.tsx:387`), so it grows UPWARD + * over the `at="top"` rule and never reaches the `at="bottom"` one. + * + * NOT occluding (deliberately excluded): + * + * - The PromptZone flow states — `approval`, `billing`, `subscription`, + * `confirm`, `clarify`, `sudo`, `secret` (`appOverlays.tsx:58-162`). They + * render in NORMAL FLOW above ComposerPane (`appLayout.tsx:553-568`): they + * push content down, they do not cover it. The rule stays on screen and + * its clock must keep running. + * - `agents` and `journey`. They unmount the entire ComposerPane subtree + * (`appLayout.tsx:553`), so `StatusRule` unmounts with it and React's own + * effect cleanup clears the intervals. Gating on them would be dead code. + * - `ambient` — a glanceable in-flow dock that reserves its own rows. + * - Composer completions. They share the FloatingOverlays grid and do + * occlude, but they are a render prop rather than store state and they + * change on every keystroke; re-arming a 1s interval per character would + * restart the countdown each time and starve the tick outright — strictly + * worse than the churn this gate removes. + * + * `statusBar: 'off'` needs no branch: `StatusRulePane` returns null for both + * slots, so the timers are never mounted in the first place. + */ +/** + * True when any floating overlay PANEL is open (widget overlays and inline + * completions are separate concerns — see $isStatusRuleOccluded for why + * completions never occlude the status rule). + * + * SINGLE SOURCE for the floating-panel kind set: consumed by both + * FloatingOverlays' render gate (plus completions, which it adds locally) + * and $isStatusRuleOccluded's top-statusbar occlusion arm. Add new floating + * panels HERE so the timer gate can't silently miss them. + */ +export const hasFloatingPanel = (overlay: OverlayState): boolean => + Boolean( + overlay.modelPicker || + overlay.pager || + overlay.petPicker || + overlay.pluginsHub || + overlay.sessions || + overlay.skillsHub + ) + +export const $isStatusRuleOccluded = computed([$overlayState, $uiState], (overlay, ui) => + Boolean(overlay.widget || (ui.statusBar === 'top' && hasFloatingPanel(overlay))) +) + export const getOverlayState = () => $overlayState.get() export const patchOverlayState = (next: Partial | ((state: OverlayState) => OverlayState)) => diff --git a/ui-tui/src/app/scroll.ts b/ui-tui/src/app/scroll.ts index e3a53734a3..6e284dde25 100644 --- a/ui-tui/src/app/scroll.ts +++ b/ui-tui/src/app/scroll.ts @@ -67,5 +67,8 @@ export function scrollWithSelectionBy(delta: number, { scrollRef, selection }: S shift(-actual, top, bottom) } - s.scrollBy(actual) + // The target is already accepted and clamped here. Commit it directly so + // wheel/page input does not enter ScrollBox's multi-frame pending-delta + // drain and produce visible stair steps at virtual row boundaries. + s.scrollTo(cur + actual) } diff --git a/ui-tui/src/app/useMainApp.ts b/ui-tui/src/app/useMainApp.ts index 18277fd663..283ebe5a11 100644 --- a/ui-tui/src/app/useMainApp.ts +++ b/ui-tui/src/app/useMainApp.ts @@ -12,7 +12,7 @@ import { useStore } from '@nanostores/react' import { useCallback, useEffect, useMemo, useRef, useState } from 'react' import { DASHBOARD_TUI_MODE, STARTUP_RESUME_ID } from '../config/env.js' -import { MAX_HISTORY, WHEEL_SCROLL_STEP } from '../config/limits.js' +import { WHEEL_SCROLL_STEP } from '../config/limits.js' import { RESIZE_COALESCE_MS } from '../config/timing.js' import { hasLeadGap, prevRenderedMsg } from '../domain/blockLayout.js' import { SECTION_NAMES, sectionMode } from '../domain/details.js' @@ -28,9 +28,9 @@ import type { TerminalResizeResponse } from '../gatewayTypes.js' import { useGitBranch } from '../hooks/useGitBranch.js' -import { useVirtualHistory } from '../hooks/useVirtualHistory.js' +import { pruneVirtualHeightCache, useVirtualHistory } from '../hooks/useVirtualHistory.js' import { composerPromptWidth } from '../lib/inputMetrics.js' -import { appendTranscriptMessage } from '../lib/messages.js' +import { appendTranscriptMessage, capTranscriptHistory } from '../lib/messages.js' import { DEFAULT_VOICE_RECORD_KEY, isMac, type ParsedVoiceRecordKey } from '../lib/platform.js' import { createResizeCoalescer } from '../lib/resizeCoalescer.js' import { asRpcResult, rpcErrorMessage } from '../lib/rpc.js' @@ -63,14 +63,6 @@ const BRACKET_PASTE_ON = '\x1b[?2004h' const BRACKET_PASTE_OFF = '\x1b[?2004l' const MAX_HEIGHT_CACHE_BUCKETS = 12 -const capHistory = (items: Msg[]): Msg[] => { - if (items.length <= MAX_HISTORY) { - return items - } - - return items[0]?.kind === 'intro' ? [items[0]!, ...items.slice(-(MAX_HISTORY - 1))] : items.slice(-MAX_HISTORY) -} - const statusColorOf = (status: string, t: { error: string; muted: string; ok: string; warn: string }) => { if (status === 'ready') { return t.ok @@ -183,7 +175,17 @@ export function useMainApp(gw: GatewayClient) { } }, [stdout]) - const [historyItems, setHistoryItems] = useState(() => [{ kind: 'intro', role: 'system', text: '' }]) + const [historyItems, setHistoryItemsState] = useState(() => [{ kind: 'intro', role: 'system', text: '' }]) + const [historyGeneration, setHistoryGeneration] = useState(0) + + const setHistoryItems = useCallback>(value => { + if (typeof value !== 'function') { + setHistoryGeneration(generation => generation + 1) + } + + setHistoryItemsState(previous => capTranscriptHistory(typeof value === 'function' ? value(previous) : value)) + }, []) + const [lastUserMsg, setLastUserMsg] = useState('') const [stickyPrompt, setStickyPrompt] = useState('') const [catalog, setCatalog] = useState(null) @@ -355,20 +357,20 @@ export function useMainApp(gw: GatewayClient) { const userPromptWidth = composerPromptWidth(ui.theme.brand.prompt) const heightCacheKey = `${ui.sid ?? 'draft'}:${cols}:${userPromptWidth}:${ui.compact ? '1' : '0'}:${detailsLayoutKey}` - const heightCache = useMemo(() => { - let cache = heightCachesRef.current.get(heightCacheKey) + // Build a render-local snapshot. Registering/pruning the shared cache is a + // post-commit transition below, so an abandoned concurrent render cannot + // delete heights still owned by the committed transcript generation. + const activeHeightCache = useMemo(() => new Map(heightCachesRef.current.get(heightCacheKey)), [heightCacheKey]) - if (!cache) { - cache = new Map() - heightCachesRef.current.set(heightCacheKey, cache) + useEffect(() => { + pruneVirtualHeightCache(activeHeightCache, virtualRows) + heightCachesRef.current.delete(heightCacheKey) + heightCachesRef.current.set(heightCacheKey, activeHeightCache) - if (heightCachesRef.current.size > MAX_HEIGHT_CACHE_BUCKETS) { - heightCachesRef.current.delete(heightCachesRef.current.keys().next().value!) - } + while (heightCachesRef.current.size > MAX_HEIGHT_CACHE_BUCKETS) { + heightCachesRef.current.delete(heightCachesRef.current.keys().next().value!) } - - return cache - }, [heightCacheKey]) + }, [activeHeightCache, heightCacheKey, historyGeneration, virtualRows]) // Index of the first user-role message — separator-rendering in // appLayout.tsx skips this row, so the height estimator must skip it @@ -414,16 +416,17 @@ export function useMainApp(gw: GatewayClient) { const h = heights.get(row.key) if (h) { - heightCache.set(row.key, h) + activeHeightCache.set(row.key, h) } } }, - [heightCache, virtualRows] + [activeHeightCache, virtualRows] ) const virtualHistory = useVirtualHistory(scrollRef, virtualRows, cols, { estimateHeight: estimateRowHeight, - initialHeights: heightCache, + generation: historyGeneration, + initialHeights: activeHeightCache, liveTailActive: turnLiveTailActive, onHeightsChange: syncHeightCache }) @@ -434,8 +437,8 @@ export function useMainApp(gw: GatewayClient) { ) const appendMessage = useCallback( - (msg: Msg) => setHistoryItems(prev => capHistory(appendTranscriptMessage(prev, msg))), - [] + (msg: Msg) => setHistoryItems(prev => appendTranscriptMessage(prev, msg)), + [setHistoryItems] ) const sys = useCallback((text: string) => appendMessage({ role: 'system', text }), [appendMessage]) @@ -801,6 +804,7 @@ export function useMainApp(gw: GatewayClient) { session.newSession, session.resetSession, session.resumeById, + setHistoryItems, setVoiceEnabled, setVoiceProcessing, setVoiceRecording, @@ -912,6 +916,7 @@ export function useMainApp(gw: GatewayClient) { selection, send, session, + setHistoryItems, sys ] ) diff --git a/ui-tui/src/app/useSessionLifecycle.ts b/ui-tui/src/app/useSessionLifecycle.ts index 13dab7ce4c..ee8fbe5f2c 100644 --- a/ui-tui/src/app/useSessionLifecycle.ts +++ b/ui-tui/src/app/useSessionLifecycle.ts @@ -2,7 +2,7 @@ import { writeFileSync } from 'node:fs' import type { ScrollBoxHandle } from '@hermes/ink' import { evictInkCaches } from '@hermes/ink' -import { type RefObject, useCallback, useEffect, useRef } from 'react' +import { type RefObject, useCallback, useEffect, useMemo, useRef } from 'react' import { buildSetupRequiredSections, SETUP_REQUIRED_TITLE } from '../content/setup.js' import { introMsg, toTranscriptMessages } from '../domain/messages.js' @@ -423,15 +423,28 @@ export function useSessionLifecycle(opts: UseSessionLifecycleOptions) { [sys] ) - return { - activateLiveSession, - closeSession, - guardBusySessionSwitch, - newLiveSession, - newSession, - resetSession, - resetVisibleHistory, - resumeById, - trimLastExchange: trimTail - } + return useMemo( + () => ({ + activateLiveSession, + closeSession, + guardBusySessionSwitch, + newLiveSession, + newSession, + resetSession, + resetVisibleHistory, + resumeById, + trimLastExchange: trimTail + }), + [ + activateLiveSession, + closeSession, + guardBusySessionSwitch, + newLiveSession, + newSession, + resetSession, + resetVisibleHistory, + resumeById, + trimTail + ] + ) } diff --git a/ui-tui/src/components/appChrome.tsx b/ui-tui/src/components/appChrome.tsx index 7758255765..1ff100f6d1 100644 --- a/ui-tui/src/components/appChrome.tsx +++ b/ui-tui/src/components/appChrome.tsx @@ -5,6 +5,7 @@ import unicodeSpinners from 'unicode-animations' import { $delegationState } from '../app/delegationStore.js' import type { BatteryInfo, IndicatorStyle, Notice } from '../app/interfaces.js' +import { $isStatusRuleOccluded } from '../app/overlayStore.js' import { useTurnSelector } from '../app/turnStore.js' import { DEV_CREDITS_MODE } from '../config/env.js' import { FACES } from '../content/faces.js' @@ -122,6 +123,7 @@ function FaceTicker({ color, startedAt, style }: { color: string; startedAt?: nu const [tick, setTick] = useState(() => Math.floor(Math.random() * 1000)) const [verbTick, setVerbTick] = useState(() => Math.floor(Math.random() * VERBS.length)) const [now, setNow] = useState(() => Date.now()) + const isOccluded = useStore($isStatusRuleOccluded) // Pre-compute cadence + verb-visibility for the active style so an // `/indicator` switch re-arms the interval (and skips the verb timer @@ -130,6 +132,19 @@ function FaceTicker({ color, startedAt, style }: { color: string; startedAt?: nu const { intervalMs, showVerb } = renderIndicator(style, 0) useEffect(() => { + // An overlay is painted OVER the status rule (the modal widget slot, or a + // floating panel growing up over the top rule), so every tick below is a + // re-render nobody can see — in an Ink TUI that churn reads as the dialog + // tearing. Arm nothing while occluded. The effect re-runs when the rule + // is revealed again and re-seeds `now` from the wall clock, so the elapsed + // read-out resumes live rather than frozen at the moment it was covered. + // See `$isStatusRuleOccluded` for why this is NOT `$isBlocked`. + if (isOccluded) { + return + } + + setNow(Date.now()) + const glyph = setInterval(() => setTick(n => n + 1), intervalMs) const clock = setInterval(() => setNow(Date.now()), 1000) // Verb timer is gated on `showVerb` — `unicode` style hides the verb @@ -144,7 +159,7 @@ function FaceTicker({ color, startedAt, style }: { color: string; startedAt?: nu clearInterval(verb) } } - }, [intervalMs, showVerb]) + }, [intervalMs, isOccluded, showVerb]) const { frame } = renderIndicator(style, tick) const verb = VERBS[verbTick % VERBS.length] ?? '' @@ -360,13 +375,21 @@ function SpawnHud({ t }: { t: Theme }) { function SessionDuration({ startedAt }: { startedAt: number }) { const [now, setNow] = useState(() => Date.now()) + const isOccluded = useStore($isStatusRuleOccluded) useEffect(() => { + // Paused only while an overlay actually covers the status rule — see + // FaceTicker. The `setNow` below already re-seeds from the wall clock + // on every re-arm, so it doubles as the reveal catch-up. + if (isOccluded) { + return + } + setNow(Date.now()) const id = setInterval(() => setNow(Date.now()), 1000) return () => clearInterval(id) - }, [startedAt]) + }, [isOccluded, startedAt]) return fmtDuration(now - startedAt) } @@ -375,13 +398,21 @@ function IdleSince({ endedAt }: { endedAt: number }) { // Time since the last final agent response. Re-ticks every second like // SessionDuration so the read-out stays live while the session idles. const [now, setNow] = useState(() => Date.now()) + const isOccluded = useStore($isStatusRuleOccluded) useEffect(() => { + // Paused only while an overlay actually covers the status rule — see + // FaceTicker. The `setNow` below re-seeds from the wall clock on reveal + // so the idle read-out is not frozen when the overlay closes. + if (isOccluded) { + return + } + setNow(Date.now()) const id = setInterval(() => setNow(Date.now()), 1000) return () => clearInterval(id) - }, [endedAt]) + }, [endedAt, isOccluded]) return `✓ ${fmtDuration(now - endedAt)}` } diff --git a/ui-tui/src/components/appOverlays.tsx b/ui-tui/src/components/appOverlays.tsx index f7fbfae7d2..aa9fcd6547 100644 --- a/ui-tui/src/components/appOverlays.tsx +++ b/ui-tui/src/components/appOverlays.tsx @@ -4,7 +4,7 @@ import type { ReactNode } from 'react' import { useGateway } from '../app/gatewayContext.js' import type { AppOverlaysProps } from '../app/interfaces.js' -import { $overlayState, patchOverlayState } from '../app/overlayStore.js' +import { $overlayState, hasFloatingPanel, patchOverlayState } from '../app/overlayStore.js' import { $uiSessionId, $uiTheme } from '../app/uiStore.js' import { ActiveSessionSwitcher } from './activeSessionSwitcher.js' @@ -191,14 +191,7 @@ export function FloatingOverlays({ const sid = useStore($uiSessionId) const theme = useStore($uiTheme) - const hasAny = - overlay.modelPicker || - overlay.pager || - overlay.petPicker || - overlay.sessions || - overlay.skillsHub || - overlay.pluginsHub || - completions.length + const hasAny = hasFloatingPanel(overlay) || completions.length if (!hasAny) { return null diff --git a/ui-tui/src/hooks/useVirtualHistory.ts b/ui-tui/src/hooks/useVirtualHistory.ts index 592d20e9a0..c9b760a00e 100644 --- a/ui-tui/src/hooks/useVirtualHistory.ts +++ b/ui-tui/src/hooks/useVirtualHistory.ts @@ -48,17 +48,39 @@ const FREEZE_RENDERS = 2 // from 25 → 12: each new item adds ~100 fibers / Yoga nodes, and a // 25-item commit was the dominant contributor to the 100ms+ p99 frames. const SLIDE_STEP = 12 +const MAX_VIRTUAL_ITEM_HEIGHT = 100_000 +const MAX_VIRTUAL_GEOMETRY = 1_000_000_000 const NOOP = () => {} +const validVirtualItemHeight = (value: number): boolean => + Number.isFinite(value) && value > 0 && value <= MAX_VIRTUAL_ITEM_HEIGHT + +const safeUnsignedGeometry = (value: number, fallback = 0): number => + Number.isFinite(value) && value >= 0 && value <= MAX_VIRTUAL_GEOMETRY ? value : fallback + +const safeSignedGeometry = (value: number, fallback = 0): number => + Number.isFinite(value) && Math.abs(value) <= MAX_VIRTUAL_GEOMETRY ? value : fallback + +export const pruneVirtualHeightCache = (cache: Map, items: readonly { key: string }[]): void => { + const active = new Set(items.map(item => item.key)) + + for (const key of cache.keys()) { + if (!active.has(key) || !validVirtualItemHeight(cache.get(key)!)) { + cache.delete(key) + } + } +} + export const virtualHistorySnapshotKey = (s?: ScrollBoxHandle | null): string => { if (!s) { return 'none' } - const target = s.getScrollTop() + s.getPendingDelta() + const target = safeUnsignedGeometry(safeUnsignedGeometry(s.getScrollTop()) + safeSignedGeometry(s.getPendingDelta())) + const bin = Math.floor(target / QUANTUM) - const viewportHeight = Math.max(0, s.getViewportHeight()) + const viewportHeight = safeUnsignedGeometry(s.getViewportHeight()) return `${s.isSticky() ? ~bin : bin}:${viewportHeight}` } @@ -97,11 +119,17 @@ export const ensureVirtualItemHeight = ( ) => { const cached = heights.get(key) - if (cached !== undefined) { - return Math.max(1, Math.floor(cached)) + if (cached !== undefined && validVirtualItemHeight(cached)) { + return Math.floor(cached) } - const seeded = Math.max(1, Math.floor(estimateHeight?.(index, key) ?? estimate)) + if (cached !== undefined) { + heights.delete(key) + } + + const fallback = validVirtualItemHeight(estimate) ? Math.floor(estimate) : estimate === 0 ? 1 : ESTIMATE + const candidate = estimateHeight?.(index, key) ?? fallback + const seeded = validVirtualItemHeight(candidate) ? Math.floor(candidate) : candidate === 0 ? 1 : fallback heights.set(key, seeded) return seeded @@ -114,6 +142,7 @@ export function useVirtualHistory( { estimate = ESTIMATE, estimateHeight, + generation = 0, initialHeights, liveTailActive = false, onHeightsChange, @@ -126,15 +155,16 @@ export function useVirtualHistory( const heights = useRef(new Map(initialHeights)) const initialHeightsRef = useRef(initialHeights) const refs = useRef(new Map void>()) + const measuredBottoms = useRef(new Map()) + const unmountViewport = useRef<{ sticky: boolean; top: number } | null>(null) const onHeightsChangeRef = useRef(onHeightsChange) // Bump whenever heightCache mutates so offsets rebuild on next read. // Ref (not state) — checked during render phase, zero extra commits. const offsetVersion = useRef(0) - // Cached offsets: reused Float64Array keyed on (itemCount, version) so we - // only rebuild when something actually changed. Previous approach allocated - // a fresh Array(n+1) every render — at n=10k that's ~80KB/render of GC - // pressure during streaming. + // Cached offsets: reused Float64Array keyed on (itemCount, version). Clean + // renders reuse it, but a point-height invalidation still rebuilds the full + // prefix array; this is allocation reuse, not incremental prefix indexing. const offsetsCache = useRef<{ arr: Float64Array; n: number; version: number }>({ arr: new Float64Array(0), n: -1, @@ -158,9 +188,25 @@ export function useVirtualHistory( const skipMeasurement = useRef(false) const prevRange = useRef(null) const freezeRenders = useRef(0) + const generationRef = useRef(generation) onHeightsChangeRef.current = onHeightsChange + if (generationRef.current !== generation) { + generationRef.current = generation + nodes.current.clear() + refs.current.clear() + measuredBottoms.current.clear() + unmountViewport.current = null + heights.current = new Map(initialHeights) + initialHeightsRef.current = initialHeights + prevRange.current = null + freezeRenders.current = 0 + skipMeasurement.current = false + lastScrollTopRef.current = 0 + offsetVersion.current++ + } + if (initialHeightsRef.current !== initialHeights) { initialHeightsRef.current = initialHeights heights.current = new Map(initialHeights) @@ -173,7 +219,13 @@ export function useVirtualHistory( prevColumns.current = columns for (const [k, h] of heights.current) { - heights.current.set(k, Math.max(1, Math.round(h * ratio))) + const scaled = Math.round(h * ratio) + + if (validVirtualItemHeight(scaled)) { + heights.current.set(k, scaled) + } else { + heights.current.delete(k) + } } offsetVersion.current++ @@ -209,6 +261,7 @@ export function useVirtualHistory( heights.current.delete(k) nodes.current.delete(k) refs.current.delete(k) + measuredBottoms.current.delete(k) dirty = true } } @@ -238,10 +291,10 @@ export function useVirtualHistory( const offsets = offsetsCache.current.arr const total = offsets[n] ?? 0 - const top = Math.max(0, scrollRef.current?.getScrollTop() ?? 0) - const pendingDelta = scrollRef.current?.getPendingDelta() ?? 0 - const target = Math.max(0, top + pendingDelta) - const vp = Math.max(0, scrollRef.current?.getViewportHeight() ?? 0) + const top = safeUnsignedGeometry(scrollRef.current?.getScrollTop() ?? 0) + const pendingDelta = safeSignedGeometry(scrollRef.current?.getPendingDelta() ?? 0) + const target = safeUnsignedGeometry(top + pendingDelta) + const vp = safeUnsignedGeometry(scrollRef.current?.getViewportHeight() ?? 0) const sticky = scrollRef.current?.isSticky() ?? true const recentManual = Date.now() - (scrollRef.current?.getLastManualScrollAt() ?? 0) < 1200 @@ -417,44 +470,90 @@ export function useVirtualHistory( } } - const measureRef = useCallback((key: string) => { - let fn = refs.current.get(key) + const measureRef = useCallback( + (key: string) => { + let fn = refs.current.get(key) - if (!fn) { - fn = (el: unknown) => { - if (el) { - nodes.current.set(key, el) + if (!fn) { + const refGeneration = generationRef.current - return + fn = (el: unknown) => { + if (refGeneration !== generationRef.current) { + return + } + + if (el) { + nodes.current.set(key, el) + + return + } + + // A width-change render has already scaled the cache, but outgoing + // refs still point at Yoga from the previous layout. The render-phase + // skip flag remains set through mutation refs, so ignore that stale + // measurement and let the post-resize layout pass own correction. + if (skipMeasurement.current) { + nodes.current.delete(key) + measuredBottoms.current.delete(key) + + return + } + + // Measure-at-unmount: the yogaNode is still valid here (reconciler + // calls ref(null) before removeChild → freeRecursive), so we grab + // the final height before WASM release. Without this, items + // scrolled out during fast pan keep a stale estimate in heightCache + // and offset math drifts until the next mount/remount cycle. + const existing = nodes.current.get(key) as MeasuredNode | undefined + const h = Math.ceil(existing?.yogaNode?.getComputedHeight?.() ?? 0) + const previousHeight = heights.current.get(key) + + if (validVirtualItemHeight(h) && previousHeight !== h) { + const s = scrollRef.current + const measuredBottom = measuredBottoms.current.get(key) + + // All null refs in this commit share the viewport boundary captured + // before the first adjustment. Otherwise an earlier compensation + // can make an intersecting sibling look wholly above later in the + // same unmount batch. + const viewport = (unmountViewport.current ??= { + sticky: s?.isSticky() ?? true, + top: safeUnsignedGeometry(s?.getScrollTop() ?? 0) + }) + + if ( + s && + previousHeight !== undefined && + measuredBottom !== undefined && + measuredBottom <= viewport.top && + !viewport.sticky + ) { + s.adjustScrollTop(h - previousHeight) + } + + heights.current.set(key, h) + offsetVersion.current++ + onHeightsChangeRef.current?.(heights.current) + } + + nodes.current.delete(key) + measuredBottoms.current.delete(key) } - // Measure-at-unmount: the yogaNode is still valid here (reconciler - // calls ref(null) before removeChild → freeRecursive), so we grab - // the final height before WASM release. Without this, items - // scrolled out during fast pan keep a stale estimate in heightCache - // and offset math drifts until the next mount/remount cycle. - const existing = nodes.current.get(key) as MeasuredNode | undefined - const h = Math.ceil(existing?.yogaNode?.getComputedHeight?.() ?? 0) - - if (h > 0 && heights.current.get(key) !== h) { - heights.current.set(key, h) - offsetVersion.current++ - onHeightsChangeRef.current?.(heights.current) - } - - nodes.current.delete(key) + refs.current.set(key, fn) } - refs.current.set(key, fn) - } - - return fn - }, []) + return fn + }, + [scrollRef] + ) useLayoutEffect(() => { + unmountViewport.current = null const s = scrollRef.current let dirty = false let heightDirty = false + let anchorDelta = 0 // Give the renderer the mounted-row coverage for passive scroll clamping. // Clamp MUST use the EFFECTIVE (deferred) range, not the immediate one. @@ -466,20 +565,40 @@ export function useVirtualHistory( if (s && shouldSetVirtualClamp({ itemCount: n, liveTailActive, sticky, viewportHeight: vp })) { const effTopSpacer = offsets[effStart] ?? 0 const effBottom = offsets[effEnd] ?? total - // At effEnd=n there's no bottomSpacer — use Infinity so render-node- - // to-output's own Math.min(cur, maxScroll) governs. Using offsets[n] - // here would bake in heightCache (one render behind Yoga), and during - // streaming the tail item's cached height lags its real height — - // sticky-break would then clamp below the real max and push - // streaming text off-viewport. const clampMin = effStart === 0 ? 0 : effTopSpacer - const clampMax = effEnd === n ? Infinity : Math.max(effTopSpacer, effBottom - vp) + // Preserve the intentional open tail: when the mounted range reaches + // the final row, Yoga may already know about growth that the measured + // height cache has not reconciled yet. A finite estimated clamp would + // trap a manual, non-sticky viewport above that newly grown tail. + const clampMax = effEnd === n ? Number.POSITIVE_INFINITY : Math.max(effTopSpacer, effBottom - vp) - s.setClampBounds(clampMin, clampMax) + if ( + safeUnsignedGeometry(clampMin, -1) >= 0 && + (clampMax === Number.POSITIVE_INFINITY || safeUnsignedGeometry(clampMax, -1) >= clampMin) + ) { + s.setClampBounds(clampMin, clampMax) + } else { + s.setClampBounds(undefined, undefined) + } } else { s?.setClampBounds(undefined, undefined) } + // Stable metadata for measure-at-unmount. Ref(null) runs before the next + // commit's layout effects, so an outgoing row sees the bottom recorded by + // the last committed mounted range. This stays O(mounted), not O(history). + for (let i = effStart; i < effEnd; i++) { + const k = items[i]?.key + + if (k) { + const bottom = offsets[i + 1] ?? 0 + + if (safeUnsignedGeometry(bottom, -1) >= 0) { + measuredBottoms.current.set(k, bottom) + } + } + } + if (skipMeasurement.current) { skipMeasurement.current = false bumpMeasuredHeightVersion(n => n + 1) @@ -492,8 +611,16 @@ export function useVirtualHistory( } const h = Math.ceil((nodes.current.get(k) as MeasuredNode | undefined)?.yogaNode?.getComputedHeight?.() ?? 0) + const previousHeight = heights.current.get(k) + + if (validVirtualItemHeight(h) && previousHeight !== h) { + // Keep the same content at the same screen row when estimates above + // the committed viewport converge. Rows intersecting or below the + // viewport intentionally retain the current scrollTop. + if (previousHeight !== undefined && (offsets[i + 1] ?? 0) <= top) { + anchorDelta += h - previousHeight + } - if (h > 0 && heights.current.get(k) !== h) { heights.current.set(k, h) dirty = true heightDirty = true @@ -501,11 +628,18 @@ export function useVirtualHistory( } } + // Sticky/live-tail positioning is owned by ScrollBox's bottom-follow + // logic. Manual viewports need an additive adjustment that leaves any + // pending input intact; scrollTo would clear that intent and other state. + if (s && anchorDelta !== 0 && !s.isSticky()) { + s.adjustScrollTop(anchorDelta) + } + if (s) { const next = { sticky: s.isSticky(), - top: Math.max(0, s.getScrollTop() + s.getPendingDelta()), - vp: Math.max(0, s.getViewportHeight()) + top: safeUnsignedGeometry(s.getScrollTop() + safeSignedGeometry(s.getPendingDelta())), + vp: safeUnsignedGeometry(s.getViewportHeight()) } if ( @@ -526,7 +660,7 @@ export function useVirtualHistory( if (heightDirty) { bumpMeasuredHeightVersion(n => n + 1) } - }, [effEnd, effStart, items, liveTailActive, measuredHeightVersion, n, offsets, scrollRef, sticky, total, vp]) + }, [effEnd, effStart, items, liveTailActive, measuredHeightVersion, n, offsets, scrollRef, sticky, top, total, vp]) return { bottomSpacer: Math.max(0, total - (offsets[effEnd] ?? total)), @@ -546,6 +680,7 @@ interface VirtualHistoryOptions { coldStartCount?: number estimate?: number estimateHeight?: (index: number, key: string) => number + generation?: number | string initialHeights?: ReadonlyMap liveTailActive?: boolean maxMounted?: number diff --git a/ui-tui/src/lib/messages.ts b/ui-tui/src/lib/messages.ts index b8e89421e5..19a2877aca 100644 --- a/ui-tui/src/lib/messages.ts +++ b/ui-tui/src/lib/messages.ts @@ -1,8 +1,17 @@ +import { MAX_HISTORY } from '../config/limits.js' import type { Msg, Role } from '../types.js' import { appendToolShelfMessage } from './liveProgress.js' export const appendTranscriptMessage = (prev: Msg[], msg: Msg): Msg[] => appendToolShelfMessage(prev, msg) +export const capTranscriptHistory = (items: Msg[]): Msg[] => { + if (items.length <= MAX_HISTORY) { + return items + } + + return items[0]?.kind === 'intro' ? [items[0], ...items.slice(-(MAX_HISTORY - 1))] : items.slice(-MAX_HISTORY) +} + export const upsert = (prev: Msg[], role: Role, text: string): Msg[] => prev.at(-1)?.role === role ? [...prev.slice(0, -1), { role, text }] : [...prev, { role, text }] diff --git a/ui-tui/src/lib/text.ts b/ui-tui/src/lib/text.ts index dff15e21df..64b5b31117 100644 --- a/ui-tui/src/lib/text.ts +++ b/ui-tui/src/lib/text.ts @@ -119,8 +119,17 @@ export const cleanThinkingText = (reasoning: string) => .replace(/\n{3,}/g, '\n\n') .trim() +// cleanThinkingText runs several full-string regex passes (split/map/filter/join/replace). +// reasoning grows on every streamed token, so without a pre-bound this re-cleans the whole +// accumulated string on every chunk — O(n) work per token, O(n^2) over a stream. Only the +// tail is ever displayed (boundedLiveRenderText caps it further downstream), so bound the +// input here first. Headroom over LIVE_RENDER_MAX_CHARS keeps line-boundary trimming inside +// cleanThinkingText accurate even after slicing mid-line. +const THINKING_CLEAN_TAIL_BOUND = LIVE_RENDER_MAX_CHARS * 1.5 + export const thinkingPreview = (reasoning: string, mode: ThinkingMode, max: number = THINKING_COT_MAX) => { - const raw = cleanThinkingText(reasoning) + const bounded = reasoning.length > THINKING_CLEAN_TAIL_BOUND ? reasoning.slice(-THINKING_CLEAN_TAIL_BOUND) : reasoning + const raw = cleanThinkingText(bounded) return !raw || mode === 'collapsed' ? '' : mode === 'full' ? raw : compactPreview(raw.replace(WS_RE, ' '), max) } diff --git a/ui-tui/src/types/hermes-ink.d.ts b/ui-tui/src/types/hermes-ink.d.ts index 7f7a53d976..4069609cb3 100644 --- a/ui-tui/src/types/hermes-ink.d.ts +++ b/ui-tui/src/types/hermes-ink.d.ts @@ -77,6 +77,7 @@ declare module '@hermes/ink' { } export type ScrollBoxHandle = { + readonly adjustScrollTop: (dy: number) => void readonly scrollTo: (y: number) => void readonly scrollBy: (dy: number) => void readonly scrollToElement: (el: unknown, offset?: number) => void diff --git a/uv.lock b/uv.lock index b6f906e471..18b8800119 100644 --- a/uv.lock +++ b/uv.lock @@ -8,7 +8,7 @@ resolution-markers = [ ] [options] -exclude-newer = "2026-07-18T05:39:53.300543645Z" +exclude-newer = "2026-07-20T16:38:42.729819205Z" exclude-newer-span = "P14D" [options.exclude-newer-package] @@ -1559,7 +1559,7 @@ wheels = [ [[package]] name = "hermes-agent" -version = "0.19.1" +version = "0.20.0" source = { editable = "." } dependencies = [ { name = "certifi" }, @@ -1571,7 +1571,7 @@ dependencies = [ { name = "httpx", extra = ["socks"] }, { name = "jinja2" }, { name = "markdown" }, - { name = "nemo-relay", marker = "(platform_machine == 'arm64' and sys_platform == 'darwin') or (platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux') or (platform_machine == 'AMD64' and sys_platform == 'win32') or (platform_machine == 'ARM64' and sys_platform == 'win32')" }, + { name = "nemo-relay", marker = "(platform_machine == 'aarch64' and 'android' not in platform_release and sys_platform == 'linux') or (platform_machine == 'x86_64' and 'android' not in platform_release and sys_platform == 'linux') or (platform_machine == 'arm64' and sys_platform == 'darwin') or (platform_machine == 'AMD64' and sys_platform == 'win32') or (platform_machine == 'ARM64' and sys_platform == 'win32')" }, { name = "openai" }, { name = "packaging" }, { name = "pathspec" }, @@ -1875,7 +1875,7 @@ requires-dist = [ { name = "microsoft-teams-apps", marker = "extra == 'teams'", specifier = "==2.0.13.4" }, { name = "mistralai", marker = "extra == 'mistral'", specifier = "==2.4.8" }, { name = "modal", marker = "extra == 'modal'", specifier = "==1.3.4" }, - { name = "nemo-relay", marker = "(platform_machine == 'arm64' and sys_platform == 'darwin') or (platform_machine == 'aarch64' and sys_platform == 'linux') or (platform_machine == 'x86_64' and sys_platform == 'linux') or (platform_machine == 'AMD64' and sys_platform == 'win32') or (platform_machine == 'ARM64' and sys_platform == 'win32')", specifier = ">=0.6.0,<0.7" }, + { name = "nemo-relay", marker = "(platform_machine == 'aarch64' and 'android' not in platform_release and sys_platform == 'linux') or (platform_machine == 'x86_64' and 'android' not in platform_release and sys_platform == 'linux') or (platform_machine == 'arm64' and sys_platform == 'darwin') or (platform_machine == 'AMD64' and sys_platform == 'win32') or (platform_machine == 'ARM64' and sys_platform == 'win32')", specifier = ">=0.6.0,<0.7" }, { name = "numpy", marker = "extra == 'voice'", specifier = "==2.4.3" }, { name = "numpy", marker = "extra == 'wake'", specifier = "==2.4.3" }, { name = "onnxruntime", marker = "extra == 'wake'", specifier = "==1.27.0" }, @@ -2699,6 +2699,7 @@ wheels = [ name = "nemo-relay" version = "0.6.0" source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/16/db/44d7258ee620c5cce6dc588983fcb11be3ee03ae71089a65686c14d9ee02/nemo_relay-0.6.0.tar.gz", hash = "sha256:f3d3088019609bc953357b5598a47481dc3e7dc8f11ecf27002ede251f37eb7b", size = 1071046, upload-time = "2026-08-03T14:55:49.702Z" } wheels = [ { url = "https://files.pythonhosted.org/packages/25/65/d320016505457cc30971f575e8dadffb923b7cfc780ab8bb25a4ce9d305c/nemo_relay-0.6.0-cp311-abi3-macosx_11_0_arm64.whl", hash = "sha256:ad5dae6febf6532d7b113abc2a404679c8feffc499df3034b93d9a078185d2bb", size = 9917779, upload-time = "2026-07-22T20:07:48.961Z" }, { url = "https://files.pythonhosted.org/packages/ae/c0/f33250e71c4206da1b339072893f9a1e39295fe1aceb9a2fef4b8620a0f2/nemo_relay-0.6.0-cp311-abi3-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:c0cd9570f64c6956fe3bfb82af1cdb3ee70cb50b51098cdb0de831c3f9b4e904", size = 8888375, upload-time = "2026-07-22T20:07:51.049Z" }, diff --git a/web/src/components/ChatSidebar.test.tsx b/web/src/components/ChatSidebar.test.tsx new file mode 100644 index 0000000000..8a05ee6a6a --- /dev/null +++ b/web/src/components/ChatSidebar.test.tsx @@ -0,0 +1,133 @@ +// @vitest-environment jsdom +import { act, type ReactNode } from "react"; +import { createRoot, type Root } from "react-dom/client"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; + +const apiMocks = vi.hoisted(() => ({ + buildWsUrl: vi.fn(async () => "ws://localhost/api/events?channel=chat-1"), + getModelInfo: vi.fn(async () => ({ + capabilities: { supports_reasoning: false }, + model: "test/model", + })), +})); + +const gatewayMocks = vi.hoisted(() => ({ + close: vi.fn(), + connect: vi.fn(async () => undefined), + on: vi.fn(() => () => undefined), + onState: vi.fn((handler: (state: string) => void) => { + handler("open"); + return () => undefined; + }), + request: vi.fn(async () => ({ session_id: "sidecar-1" })), +})); + +const reloadMocks = vi.hoisted(() => ({ + maybeReloadForLoopbackWsAuthFailure: vi.fn(() => true), +})); + +vi.mock("@/lib/api", () => ({ + api: { getModelInfo: apiMocks.getModelInfo }, + buildWsUrl: apiMocks.buildWsUrl, +})); +vi.mock("@/lib/dashboard-auth-reload", () => ({ + maybeReloadForLoopbackWsAuthFailure: + reloadMocks.maybeReloadForLoopbackWsAuthFailure, +})); +vi.mock("@/lib/gatewayClient", () => ({ + GatewayClient: class { + close = gatewayMocks.close; + connect = gatewayMocks.connect; + on = gatewayMocks.on; + onState = gatewayMocks.onState; + request = gatewayMocks.request; + }, +})); +vi.mock("@/components/ModelPickerDialog", () => ({ + ModelPickerDialog: () => null, +})); +vi.mock("@/components/ModelReloadConfirm", () => ({ + ModelReloadConfirm: () => null, +})); +vi.mock("@/components/ReasoningPicker", () => ({ + ReasoningPicker: () => null, +})); +vi.mock("@nous-research/ui/ui/components/button", () => ({ + Button: ({ children }: { children?: ReactNode }) => , +})); +vi.mock("@nous-research/ui/ui/components/badge", () => ({ + Badge: ({ children }: { children?: ReactNode }) => {children}, +})); +vi.mock("@nous-research/ui/ui/components/card", () => ({ + Card: ({ children }: { children?: ReactNode }) =>
{children}
, +})); + +type EventLike = { code?: number; data?: string }; + +class FakeWebSocket { + static instances: FakeWebSocket[] = []; + + private listeners = new Map void>>(); + readonly url: string; + + constructor(url: string) { + this.url = url; + FakeWebSocket.instances.push(this); + } + + addEventListener(type: string, listener: (event: EventLike) => void) { + const listeners = this.listeners.get(type) ?? []; + listeners.push(listener); + this.listeners.set(type, listeners); + } + + close() {} + + emit(type: string, event: EventLike) { + for (const listener of this.listeners.get(type) ?? []) { + listener(event); + } + } +} + +let container: HTMLDivElement; +let root: Root; + +async function render(ui: ReactNode) { + container = document.createElement("div"); + document.body.append(container); + root = createRoot(container); + await act(async () => root.render(ui)); +} + +beforeEach(() => { + FakeWebSocket.instances = []; + vi.clearAllMocks(); + reloadMocks.maybeReloadForLoopbackWsAuthFailure.mockReturnValue(true); + vi.stubGlobal("WebSocket", FakeWebSocket); +}); + +afterEach(async () => { + await act(async () => root?.unmount()); + container?.remove(); + vi.unstubAllGlobals(); +}); + +describe("ChatSidebar event socket", () => { + it("routes loopback 4401 closes through stale-token recovery", async () => { + const { ChatSidebar } = await import("./ChatSidebar"); + + await render(); + + await vi.waitFor(() => expect(FakeWebSocket.instances).toHaveLength(1)); + expect(apiMocks.buildWsUrl).toHaveBeenCalledWith("/api/events", { + channel: "chat-1", + }); + + FakeWebSocket.instances[0].emit("close", { code: 4401 }); + + expect( + reloadMocks.maybeReloadForLoopbackWsAuthFailure, + ).toHaveBeenCalledWith(4401); + }); +}); diff --git a/web/src/components/ChatSidebar.tsx b/web/src/components/ChatSidebar.tsx index 57236e782b..fd4df2edcd 100644 --- a/web/src/components/ChatSidebar.tsx +++ b/web/src/components/ChatSidebar.tsx @@ -32,6 +32,7 @@ import { ModelReloadConfirm } from "@/components/ModelReloadConfirm"; import { ReasoningPicker } from "@/components/ReasoningPicker"; import { GatewayClient, type ConnectionState } from "@/lib/gatewayClient"; import { api, buildWsUrl } from "@/lib/api"; +import { maybeReloadForLoopbackWsAuthFailure } from "@/lib/dashboard-auth-reload"; import { titleFromSessionInfoPayload } from "@/lib/chat-title"; import { cn } from "@/lib/utils"; @@ -258,6 +259,9 @@ export function ChatSidebar({ ws.addEventListener("error", () => surface(DISCONNECTED)); ws.addEventListener("close", (ev) => { + if (maybeReloadForLoopbackWsAuthFailure(ev.code)) { + return; + } if (ev.code === 4401 || ev.code === 4403) { surface(`events feed rejected (${ev.code}) — reload the page`); } else if (ev.code !== 1000) { diff --git a/web/src/components/HermesConsoleModal.tsx b/web/src/components/HermesConsoleModal.tsx index 97a2824c79..4ce422c0ab 100644 --- a/web/src/components/HermesConsoleModal.tsx +++ b/web/src/components/HermesConsoleModal.tsx @@ -11,6 +11,7 @@ import { Button } from "@nous-research/ui/ui/components/button"; import { useModalBehavior } from "@/hooks/useModalBehavior"; import { useProfileScope } from "@/contexts/useProfileScope"; import { api } from "@/lib/api"; +import { maybeReloadForLoopbackWsAuthFailure } from "@/lib/dashboard-auth-reload"; import { cn, themedBody } from "@/lib/utils"; import { useTheme } from "@/themes"; @@ -423,6 +424,9 @@ export function HermesConsoleModal({ open, onClose }: HermesConsoleModalProps) { }; ws.onclose = (ev) => { + if (maybeReloadForLoopbackWsAuthFailure(ev.code)) { + return; + } wsRef.current = null; activeCommandRef.current = false; pendingCommandRef.current = null; diff --git a/web/src/lib/api.test.ts b/web/src/lib/api.test.ts index 4d63d51a01..5868f2bd5d 100644 --- a/web/src/lib/api.test.ts +++ b/web/src/lib/api.test.ts @@ -1,9 +1,37 @@ -import { afterEach, describe, expect, it, vi } from "vitest"; +// @vitest-environment jsdom +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; -import { api } from "./api"; +import { api, fetchJSON } from "./api"; + +const reloadMocks = vi.hoisted(() => ({ + attemptDashboardTokenReloadOnce: vi.fn(() => false), + clearDashboardTokenReloadAttempt: vi.fn(), +})); + +vi.mock("./dashboard-auth-reload", () => ({ + attemptDashboardTokenReloadOnce: reloadMocks.attemptDashboardTokenReloadOnce, + clearDashboardTokenReloadAttempt: reloadMocks.clearDashboardTokenReloadAttempt, +})); const SESSION_HEADER = "X-Hermes-Session-Token"; +beforeEach(() => { + reloadMocks.attemptDashboardTokenReloadOnce.mockReset(); + reloadMocks.attemptDashboardTokenReloadOnce.mockReturnValue(false); + reloadMocks.clearDashboardTokenReloadAttempt.mockReset(); + + Object.defineProperty(window, "__HERMES_SESSION_TOKEN__", { + configurable: true, + value: "stale-token", + writable: true, + }); + Object.defineProperty(window, "__HERMES_AUTH_REQUIRED__", { + configurable: true, + value: false, + writable: true, + }); +}); + afterEach(() => { vi.restoreAllMocks(); vi.unstubAllGlobals(); @@ -19,6 +47,47 @@ function jsonFetchMock(body: unknown = { ok: true }) { ); } +describe("fetchJSON", () => { + it("tries the one-shot reload path for loopback 401s", async () => { + vi.stubGlobal( + "fetch", + vi.fn(async () => ({ + clone: () => ({ + json: async () => ({}), + }), + ok: false, + status: 401, + statusText: "Unauthorized", + text: async () => "Unauthorized", + })), + ); + reloadMocks.attemptDashboardTokenReloadOnce.mockReturnValue(true); + + const pending = fetchJSON("/api/status"); + await expect(Promise.race([pending, Promise.resolve("pending")])).resolves.toBe( + "pending", + ); + + expect(reloadMocks.attemptDashboardTokenReloadOnce).toHaveBeenCalledTimes(1); + expect(reloadMocks.clearDashboardTokenReloadAttempt).not.toHaveBeenCalled(); + }); + + it("clears the reload latch after a successful response", async () => { + vi.stubGlobal( + "fetch", + vi.fn(async () => ({ + json: async () => ({ ok: true }), + ok: true, + status: 200, + })), + ); + + await expect(fetchJSON("/api/status")).resolves.toEqual({ ok: true }); + + expect(reloadMocks.clearDashboardTokenReloadAttempt).toHaveBeenCalledTimes(1); + }); +}); + describe("api.getModelOptions", () => { it("requests a live model refresh when asked", async () => { vi.stubGlobal("window", {}); diff --git a/web/src/lib/api.ts b/web/src/lib/api.ts index 3cc2c5a0ad..ca7ac23f30 100644 --- a/web/src/lib/api.ts +++ b/web/src/lib/api.ts @@ -20,6 +20,10 @@ export const HERMES_BASE_PATH = readBasePath(); const BASE = HERMES_BASE_PATH; import type { DashboardTheme } from "@/themes/types"; +import { + attemptDashboardTokenReloadOnce, + clearDashboardTokenReloadAttempt, +} from "@/lib/dashboard-auth-reload"; // Ephemeral session token for protected endpoints. // Injected into index.html by the server — never fetched via API. @@ -160,20 +164,7 @@ export async function fetchJSON( // handled above, so reaching here in gated mode means a real // middleware failure that should not reload-loop. if (!window.__HERMES_AUTH_REQUIRED__ && !options?.allowUnauthorized) { - let alreadyReloaded = false; - try { - alreadyReloaded = - sessionStorage.getItem("hermes.tokenReloadAttempted") === "1"; - } catch { - /* SSR / privacy mode — fall through to throw */ - } - if (!alreadyReloaded) { - try { - sessionStorage.setItem("hermes.tokenReloadAttempted", "1"); - } catch { - /* SSR / privacy mode — best effort */ - } - window.location.reload(); + if (attemptDashboardTokenReloadOnce()) { return new Promise(() => {}); } } @@ -182,11 +173,7 @@ export async function fetchJSON( // Clear the stale-token reload guard: a successful 2xx proves the // current ``window.__HERMES_SESSION_TOKEN__`` is valid, so the next // 401 — if any — should be allowed to trigger its own reload cycle. - try { - sessionStorage.removeItem("hermes.tokenReloadAttempted"); - } catch { - /* SSR / privacy mode — ignore */ - } + clearDashboardTokenReloadAttempt(); } if (!res.ok) { const text = await res.text().catch(() => res.statusText); @@ -1725,14 +1712,17 @@ export interface MemoryProviderFieldOption { export interface MemoryProviderField { key: string; label: string; - kind: "text" | "secret" | "select" | "boolean"; + kind: "text" | "secret" | "select" | "boolean" | "integer" | "number"; description: string; placeholder: string; required: boolean; - value: string | boolean; + value: string | boolean | number; is_set: boolean; options: MemoryProviderFieldOption[]; url: string; + minimum?: number | null; + maximum?: number | null; + step?: number | null; when?: Record | null; } diff --git a/web/src/lib/dashboard-auth-reload.test.ts b/web/src/lib/dashboard-auth-reload.test.ts new file mode 100644 index 0000000000..3bce0d9796 --- /dev/null +++ b/web/src/lib/dashboard-auth-reload.test.ts @@ -0,0 +1,70 @@ +import { describe, expect, it, vi } from "vitest"; + +import { + attemptDashboardTokenReloadOnce, + clearDashboardTokenReloadAttempt, + maybeReloadForLoopbackWsAuthFailure, +} from "./dashboard-auth-reload"; + +function makeStorage() { + const values = new Map(); + return { + getItem(key: string) { + return values.get(key) ?? null; + }, + removeItem(key: string) { + values.delete(key); + }, + setItem(key: string, value: string) { + values.set(key, value); + }, + }; +} + +describe("attemptDashboardTokenReloadOnce", () => { + it("reloads once and latches the attempt", () => { + const storage = makeStorage(); + const reload = vi.fn(); + + expect(attemptDashboardTokenReloadOnce(storage, reload)).toBe(true); + expect(reload).toHaveBeenCalledTimes(1); + + expect(attemptDashboardTokenReloadOnce(storage, reload)).toBe(false); + expect(reload).toHaveBeenCalledTimes(1); + }); + + it("clears the latch when asked", () => { + const storage = makeStorage(); + const reload = vi.fn(); + + expect(attemptDashboardTokenReloadOnce(storage, reload)).toBe(true); + clearDashboardTokenReloadAttempt(storage); + expect(attemptDashboardTokenReloadOnce(storage, reload)).toBe(true); + expect(reload).toHaveBeenCalledTimes(2); + }); +}); + +describe("maybeReloadForLoopbackWsAuthFailure", () => { + it("reloads once for loopback 4401 closes", () => { + const storage = makeStorage(); + const reload = vi.fn(); + + expect( + maybeReloadForLoopbackWsAuthFailure(4401, false, storage, reload), + ).toBe(true); + expect(reload).toHaveBeenCalledTimes(1); + }); + + it("does not reload in gated mode or for other close codes", () => { + const storage = makeStorage(); + const reload = vi.fn(); + + expect( + maybeReloadForLoopbackWsAuthFailure(4401, true, storage, reload), + ).toBe(false); + expect( + maybeReloadForLoopbackWsAuthFailure(4403, false, storage, reload), + ).toBe(false); + expect(reload).not.toHaveBeenCalled(); + }); +}); diff --git a/web/src/lib/dashboard-auth-reload.ts b/web/src/lib/dashboard-auth-reload.ts new file mode 100644 index 0000000000..69c50836cd --- /dev/null +++ b/web/src/lib/dashboard-auth-reload.ts @@ -0,0 +1,69 @@ +type StorageLike = Pick; + +const TOKEN_RELOAD_STORAGE_KEY = "hermes.tokenReloadAttempted"; + +function dashboardAuthRequired(): boolean { + return typeof window !== "undefined" && !!window.__HERMES_AUTH_REQUIRED__; +} + +function reloadDashboardWindow(): void { + if (typeof window !== "undefined") { + window.location.reload(); + } +} + +function dashboardSessionStorage(): StorageLike | null { + if (typeof window === "undefined") return null; + try { + return window.sessionStorage; + } catch { + return null; + } +} + +export function clearDashboardTokenReloadAttempt( + storage: StorageLike | null = dashboardSessionStorage(), +): void { + try { + storage?.removeItem(TOKEN_RELOAD_STORAGE_KEY); + } catch { + /* privacy mode / blocked storage — ignore */ + } +} + +export function attemptDashboardTokenReloadOnce( + storage: StorageLike | null = dashboardSessionStorage(), + reload: () => void = reloadDashboardWindow, +): boolean { + let alreadyReloaded = false; + try { + alreadyReloaded = + storage?.getItem(TOKEN_RELOAD_STORAGE_KEY) === "1"; + } catch { + /* privacy mode / blocked storage — fall through */ + } + if (alreadyReloaded) { + return false; + } + + try { + storage?.setItem(TOKEN_RELOAD_STORAGE_KEY, "1"); + } catch { + /* privacy mode / blocked storage — best effort */ + } + + reload(); + return true; +} + +export function maybeReloadForLoopbackWsAuthFailure( + code: number, + authRequired = dashboardAuthRequired(), + storage: StorageLike | null = dashboardSessionStorage(), + reload: () => void = reloadDashboardWindow, +): boolean { + if (authRequired || code !== 4401) { + return false; + } + return attemptDashboardTokenReloadOnce(storage, reload); +} diff --git a/web/src/lib/gatewayClient.test.ts b/web/src/lib/gatewayClient.test.ts new file mode 100644 index 0000000000..5ed55cac58 --- /dev/null +++ b/web/src/lib/gatewayClient.test.ts @@ -0,0 +1,96 @@ +// @vitest-environment jsdom +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; + +import { GatewayClient } from "./gatewayClient"; + +const reloadMocks = vi.hoisted(() => ({ + maybeReloadForLoopbackWsAuthFailure: vi.fn(() => false), +})); + +vi.mock("./dashboard-auth-reload", () => ({ + maybeReloadForLoopbackWsAuthFailure: + reloadMocks.maybeReloadForLoopbackWsAuthFailure, +})); + +class FakeWebSocket { + static instances: FakeWebSocket[] = []; + static OPEN = 1; + + listeners = new Map void>>(); + readyState = 0; + url: string; + + constructor(url: string) { + this.url = url; + FakeWebSocket.instances.push(this); + } + + addEventListener(type: string, cb: (event: EventLike) => void) { + const list = this.listeners.get(type) ?? []; + list.push(cb); + this.listeners.set(type, list); + } + + close() {} + + emit(type: string, event: EventLike) { + for (const cb of this.listeners.get(type) ?? []) { + cb(event); + } + } + + removeEventListener(type: string, cb: (event: EventLike) => void) { + const list = this.listeners.get(type) ?? []; + this.listeners.set( + type, + list.filter((item) => item !== cb), + ); + } + + send() {} +} + +type EventLike = { + code?: number; +}; + +beforeEach(() => { + FakeWebSocket.instances = []; + reloadMocks.maybeReloadForLoopbackWsAuthFailure.mockClear(); + vi.stubGlobal("WebSocket", FakeWebSocket); + Object.defineProperty(window, "__HERMES_SESSION_TOKEN__", { + configurable: true, + value: "stale-token", + writable: true, + }); + Object.defineProperty(window, "__HERMES_AUTH_REQUIRED__", { + configurable: true, + value: false, + writable: true, + }); +}); + +afterEach(() => { + vi.unstubAllGlobals(); +}); + +describe("GatewayClient", () => { + it("treats loopback 4401 closes as stale-token reload candidates", async () => { + reloadMocks.maybeReloadForLoopbackWsAuthFailure.mockReturnValue(true); + const gw = new GatewayClient(); + const connectPromise = gw.connect(); + + await vi.waitFor(() => expect(FakeWebSocket.instances).toHaveLength(1)); + const socket = FakeWebSocket.instances[0]; + socket.readyState = 1; + socket.emit("open", {}); + await connectPromise; + + socket.emit("close", { code: 4401 }); + + expect( + reloadMocks.maybeReloadForLoopbackWsAuthFailure, + ).toHaveBeenCalledWith(4401); + expect(gw.connectionState).toBe("open"); + }); +}); diff --git a/web/src/lib/gatewayClient.ts b/web/src/lib/gatewayClient.ts index d5ef547ab7..29fd891a27 100644 --- a/web/src/lib/gatewayClient.ts +++ b/web/src/lib/gatewayClient.ts @@ -22,6 +22,7 @@ import { } from "@hermes/shared"; import { HERMES_BASE_PATH, buildWsAuthParam } from "@/lib/api"; +import { maybeReloadForLoopbackWsAuthFailure } from "@/lib/dashboard-auth-reload"; export type { ConnectionState, GatewayEvent, GatewayEventName }; @@ -31,6 +32,7 @@ export class GatewayClient extends JsonRpcGatewayClient { closedErrorMessage: "WebSocket closed", connectErrorMessage: "WebSocket connection failed", notConnectedErrorMessage: "gateway not connected", + onSocketClose: (event) => maybeReloadForLoopbackWsAuthFailure(event.code), requestIdPrefix: "w", }); } diff --git a/web/src/pages/ChatPage.test.tsx b/web/src/pages/ChatPage.test.tsx new file mode 100644 index 0000000000..5548d9e59a --- /dev/null +++ b/web/src/pages/ChatPage.test.tsx @@ -0,0 +1,228 @@ +// @vitest-environment jsdom +import { act, type ReactNode } from "react"; +import { createRoot, type Root } from "react-dom/client"; +import { MemoryRouter } from "react-router"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; + +class FakeFitAddon { + fit() {} +} + +class FakeWebglAddon { + onContextLoss() { + return { dispose() {} }; + } +} + +class FakeTerminal { + options: Record; + rows = 24; + cols = 80; + parser = { + registerOscHandler: vi.fn(), + }; + unicode = { activeVersion: "" }; + + constructor(options: Record) { + this.options = options; + } + + attachCustomKeyEventHandler() { + return true; + } + + attachCustomWheelEventHandler() { + return true; + } + + clearSelection() {} + + dispose() {} + + focus() {} + + getSelection() { + return ""; + } + + loadAddon() {} + + onData() { + return { dispose() {} }; + } + + onResize() { + return { dispose() {} }; + } + + open() {} + + paste() {} + + refresh() {} + + write() {} +} + +const maybeReloadForLoopbackWsAuthFailure = vi.fn(() => false); + +vi.mock("@xterm/addon-fit", () => ({ FitAddon: FakeFitAddon })); +vi.mock("@xterm/addon-unicode11", () => ({ Unicode11Addon: class {} })); +vi.mock("@xterm/addon-web-links", () => ({ WebLinksAddon: class {} })); +vi.mock("@xterm/addon-webgl", () => ({ WebglAddon: FakeWebglAddon })); +vi.mock("@xterm/xterm", () => ({ Terminal: FakeTerminal })); +vi.mock("@/components/ChatSidebar", () => ({ + ChatSidebar: () => null, +})); +vi.mock("@/components/ChatSessionList", () => ({ + ChatSessionList: () => null, +})); +vi.mock("@/components/Backdrop", () => ({ Backdrop: () => null })); +vi.mock("@/plugins", () => ({ + PluginSlot: () => null, +})); +vi.mock("@/contexts/usePageHeader", () => ({ + usePageHeader: () => ({ setEnd: vi.fn(), setTitle: vi.fn() }), +})); +vi.mock("@/contexts/useProfileScope", () => ({ + useProfileScope: () => ({ profile: "" }), +})); +vi.mock("@/themes", () => ({ + useTheme: () => ({ theme: { terminalBackground: "#000000" } }), +})); +vi.mock("@/i18n", () => ({ + useI18n: () => ({ + t: { + app: { + closeModelTools: "Close model tools", + modelToolsSheetSubtitle: "Tools", + modelToolsSheetTitle: "Model", + }, + }, + }), +})); +vi.mock("@/lib/dashboard-auth-reload", () => ({ + maybeReloadForLoopbackWsAuthFailure, +})); + +class FakeWebSocket { + static instances: FakeWebSocket[] = []; + static OPEN = 1; + + binaryType = "blob"; + onclose: ((event: CloseEventLike) => void) | null = null; + onmessage: ((event: { data: ArrayBuffer | string }) => void) | null = null; + onopen: (() => void) | null = null; + readyState = FakeWebSocket.OPEN; + url: string; + + constructor(url: string) { + this.url = url; + FakeWebSocket.instances.push(this); + } + + close() { + this.readyState = 3; + } + + send() {} +} + +type CloseEventLike = { + code: number; + reason: string; + wasClean: boolean; +}; + +let container: HTMLDivElement; +let root: Root; + +async function render(ui: ReactNode) { + container = document.createElement("div"); + document.body.append(container); + root = createRoot(container); + await act(async () => root.render(ui)); +} + +beforeEach(() => { + FakeWebSocket.instances = []; + maybeReloadForLoopbackWsAuthFailure.mockClear(); + vi.stubGlobal("WebSocket", FakeWebSocket); + vi.stubGlobal( + "ResizeObserver", + class { + disconnect() {} + observe() {} + unobserve() {} + }, + ); + vi.stubGlobal("requestAnimationFrame", (cb: FrameRequestCallback) => { + cb(0); + return 1; + }); + vi.stubGlobal("cancelAnimationFrame", () => {}); + vi.stubGlobal("matchMedia", () => ({ + addEventListener() {}, + matches: false, + media: "", + removeEventListener() {}, + })); + vi.stubGlobal("crypto", { + getRandomValues: (values: Uint8Array) => { + values.fill(7); + return values; + }, + randomUUID: () => "chat-test-id", + }); + + Object.defineProperty(window, "visualViewport", { + configurable: true, + value: { addEventListener() {}, removeEventListener() {}, width: 1280 }, + }); + Object.defineProperty(window, "__HERMES_SESSION_TOKEN__", { + configurable: true, + value: "stale-token", + writable: true, + }); + Object.defineProperty(window, "__HERMES_AUTH_REQUIRED__", { + configurable: true, + value: false, + writable: true, + }); + Object.defineProperty(window.navigator, "clipboard", { + configurable: true, + value: { + readText: vi.fn(async () => ""), + writeText: vi.fn(async () => {}), + }, + }); + sessionStorage.clear(); +}); + +afterEach(async () => { + await act(async () => root?.unmount()); + container?.remove(); + vi.unstubAllGlobals(); +}); + +describe("ChatPage", () => { + it("treats loopback 4401 closes as stale-token reload candidates", async () => { + const { default: ChatPage } = await import("./ChatPage"); + + await render( + + + , + ); + + await vi.waitFor(() => expect(FakeWebSocket.instances).toHaveLength(1)); + + FakeWebSocket.instances[0].onclose?.({ + code: 4401, + reason: "auth: token_mismatch", + wasClean: true, + }); + + expect(maybeReloadForLoopbackWsAuthFailure).toHaveBeenCalledWith(4401); + }); +}); diff --git a/web/src/pages/ChatPage.tsx b/web/src/pages/ChatPage.tsx index 3b9e9a500a..84740126df 100644 --- a/web/src/pages/ChatPage.tsx +++ b/web/src/pages/ChatPage.tsx @@ -63,6 +63,7 @@ import { transferMayContainImage, uploadChatImage, } from "@/lib/chatImagePaste"; +import { maybeReloadForLoopbackWsAuthFailure } from "@/lib/dashboard-auth-reload"; import { PluginSlot } from "@/plugins"; import { useTheme } from "@/themes"; import { useProfileScope } from "@/contexts/useProfileScope"; @@ -1093,6 +1094,9 @@ export default function ChatPage({ isActive = true }: { isActive?: boolean }) { console.warn(`[chat] PTY WebSocket closed code=${ev.code}${why}`); setLastCloseCode(ev.code); if (ev.code === 4401) { + if (maybeReloadForLoopbackWsAuthFailure(ev.code)) { + return; + } setPtyState("closed"); setBanner( ev.reason diff --git a/web/src/pages/PluginsPage.tsx b/web/src/pages/PluginsPage.tsx index afd444c6fa..bde492f20c 100644 --- a/web/src/pages/PluginsPage.tsx +++ b/web/src/pages/PluginsPage.tsx @@ -32,7 +32,7 @@ import { usePageHeader } from "@/contexts/usePageHeader"; /** Select value for built-in memory (`config` uses empty string). Never use `""` — UI Select maps empty value to an empty label. */ const MEMORY_PROVIDER_BUILTIN = "__hermes_memory_builtin__"; -type MemoryFormValue = string | boolean; +type MemoryFormValue = string | boolean | number; const MEMORY_STATUS_LABEL: Record = { ready: "ready", @@ -664,7 +664,18 @@ export default function PluginsPage() {
(); + return { + getItem(key: string) { + return store.get(key) ?? null; + }, + setItem(key: string, value: string) { + store.set(key, value); + }, + removeItem(key: string) { + store.delete(key); + }, + clear() { + store.clear(); + }, + get length() { + return store.size; + }, + key(index: number) { + return Array.from(store.keys())[index] ?? null; + }, + } as Storage; +} + +const exampleManifest: PluginManifest = { + name: "test", + label: "Test", + description: "A test plugin", + icon: "Puzzle", + version: "1.0.0", + tab: { path: "/test" }, + entry: "index.js", + has_api: false, + source: "local", +}; + +describe("plugin manifest cache helpers", () => { + let storage: Storage; + + beforeEach(() => { + storage = makeStorage(); + vi.stubGlobal("sessionStorage", storage); + }); + + afterEach(() => { + vi.unstubAllGlobals(); + }); + + it("getCachedManifests returns null when nothing is cached", () => { + expect(getCachedManifests()).toBeNull(); + }); + + it("getCachedManifests returns null for invalid JSON", () => { + storage.setItem(MANIFEST_CACHE_KEY, "not-json"); + expect(getCachedManifests()).toBeNull(); + }); + + it("getCachedManifests returns null for non-array JSON", () => { + storage.setItem(MANIFEST_CACHE_KEY, JSON.stringify({ foo: "bar" })); + expect(getCachedManifests()).toBeNull(); + }); + + it("getCachedManifests returns null for scalar JSON", () => { + storage.setItem(MANIFEST_CACHE_KEY, JSON.stringify(42)); + expect(getCachedManifests()).toBeNull(); + }); + + it("getCachedManifests returns a valid manifest array", () => { + const list: PluginManifest[] = [exampleManifest]; + cacheManifests(list); + expect(getCachedManifests()).toEqual(list); + }); + + it("cacheManifests overwrites a previous cache on refresh", () => { + const first: PluginManifest[] = [exampleManifest]; + cacheManifests(first); + expect(getCachedManifests()).toEqual(first); + + const second: PluginManifest[] = [ + { ...exampleManifest, name: "updated", label: "Updated" }, + ]; + cacheManifests(second); + expect(getCachedManifests()).toEqual(second); + }); + + it("cacheManifests swallows storage errors", () => { + const badStorage = makeStorage(); + badStorage.setItem = () => { + throw new Error("QuotaExceededError"); + }; + vi.stubGlobal("sessionStorage", badStorage); + expect(() => cacheManifests([exampleManifest])).not.toThrow(); + }); +}); + +describe("canSeedLoadedFromCache (loading seed gate)", () => { + it("returns false when there is no cache (first visit keeps loading=true)", () => { + expect(canSeedLoadedFromCache(null)).toBe(false); + }); + + it("returns true for an empty cached list", () => { + expect(canSeedLoadedFromCache([])).toBe(true); + }); + + it("returns true when no cached manifest overrides /chat", () => { + const list: PluginManifest[] = [ + exampleManifest, + { + ...exampleManifest, + name: "other", + tab: { path: "/other", override: "/skills" }, + }, + ]; + expect(canSeedLoadedFromCache(list)).toBe(true); + }); + + it("returns false when a cached manifest overrides /chat — loading must stay true so App.tsx's pluginsLoading gate keeps the persistent chat host unmounted", () => { + const list: PluginManifest[] = [ + exampleManifest, + { + ...exampleManifest, + name: "chat-replacer", + tab: { path: "/chat-alt", override: "/chat" }, + }, + ]; + expect(canSeedLoadedFromCache(list)).toBe(false); + }); + + it("tolerates malformed cached entries missing a tab object", () => { + const malformed = [ + { ...exampleManifest, tab: undefined }, + ] as unknown as PluginManifest[]; + expect(canSeedLoadedFromCache(malformed)).toBe(true); + }); +}); diff --git a/web/src/plugins/usePlugins.ts b/web/src/plugins/usePlugins.ts index 4896295891..f03a069839 100644 --- a/web/src/plugins/usePlugins.ts +++ b/web/src/plugins/usePlugins.ts @@ -17,17 +17,76 @@ import { setPluginLoadError, } from "./registry"; +export const MANIFEST_CACHE_KEY = "hermes:plugin-manifests"; + +export function getCachedManifests(): PluginManifest[] | null { + try { + const raw = sessionStorage.getItem(MANIFEST_CACHE_KEY); + if (!raw) return null; + const parsed = JSON.parse(raw); + return Array.isArray(parsed) ? (parsed as PluginManifest[]) : null; + } catch { + return null; + } +} + +export function cacheManifests(manifests: PluginManifest[]): void { + try { + sessionStorage.setItem(MANIFEST_CACHE_KEY, JSON.stringify(manifests)); + } catch { + // sessionStorage unavailable (private browsing, storage full, etc.) + } +} + +/** + * Whether it is safe to skip the initial plugin-loading gate for a set of + * cached manifests. + * + * App.tsx waits on `pluginsLoading` before mounting the persistent ChatPage + * host: if a plugin overrides /chat (`tab.override === "/chat"`), mounting + * the built-in chat first would spawn a PTY and then yank it out from under + * the user when the plugin resolves. That gate is load-bearing — so we may + * only seed `loading = false` from the cache when no cached manifest + * declares a /chat override. Manifests are still seeded either way; only + * the loading flag stays conservative. + */ +export function canSeedLoadedFromCache( + cached: PluginManifest[] | null, +): boolean { + if (cached === null) return false; + return !cached.some((m) => m.tab?.override === "/chat"); +} + export function usePlugins() { - const [manifests, setManifests] = useState([]); + // Lazy initialisers run once at mount — safe to read sessionStorage here. + // This avoids the "cannot access ref during render" lint error that would + // occur if we stored the cached value in a useRef and read .current in the + // useState initial value expression. + const [manifests, setManifests] = useState( + () => getCachedManifests() ?? [], + ); const [plugins, setPlugins] = useState([]); - const [loading, setLoading] = useState(true); + // Start loading=false when the cache has manifests so plugin routes are + // registered synchronously on the first render after a refresh. + // The catch-all in App.tsx is only a safety net for the very first visit + // (no cache yet). On subsequent visits this flag starts false immediately. + // + // Exception: if any cached manifest overrides /chat we must keep + // loading=true — App.tsx's pluginsLoading gate around the persistent + // ChatPage host is load-bearing (see canSeedLoadedFromCache). + const [loading, setLoading] = useState( + () => !canSeedLoadedFromCache(getCachedManifests()), + ); const loadedScripts = useRef>(new Set()); - // Fetch manifests on mount. + // Always re-fetch in the background to keep the cache fresh. + // This handles: new plugins added, plugins removed, manifest changes. + // setManifests(list) will update routes if the server list differs from cache. useEffect(() => { api .getPlugins() .then((list) => { + cacheManifests(list); setManifests(list); if (list.length === 0) setLoading(false); }) diff --git a/website/docs/reference/slash-commands.md b/website/docs/reference/slash-commands.md index ae38e858f6..860d2c8d1a 100644 --- a/website/docs/reference/slash-commands.md +++ b/website/docs/reference/slash-commands.md @@ -74,7 +74,7 @@ Type `/` in the CLI to open the autocomplete menu. Built-in commands are case-in | `/config` | Show current configuration | | `/model [model-name]` | Show or change the current model. Supports: `/model claude-sonnet-4`, `/model provider:model` (switch providers), `/model custom:model` (custom endpoint), `/model custom:name:model` (named custom provider), `/model custom` (auto-detect from endpoint), and user-defined aliases (`/model fav`, `/model grok` — see [Custom model aliases](#custom-model-aliases)). Flags: `--global` persists the change to config.yaml; `--session` forces session-only; `--once` applies to the next turn only; `--refresh` re-fetches the provider's model list; `--provider ` switches backend (session-only unless `--global`). A plain `/model ` is session-only unless `model.persist_switch_by_default: true` is set. **Note:** `/model` can only switch between already-configured providers. To add a new provider, exit the session and run `hermes model` from your terminal. **Cost note:** switching models mid-conversation resets the prompt cache — the cache key includes the model, so your next turn re-reads the entire conversation at full input price instead of the ~75%-discounted cached rate. Expected and unavoidable, but worth knowing on long sessions. | | `/codex-runtime [auto\|codex_app_server\|on\|off]` | Toggle the optional [Codex app-server runtime](../user-guide/features/codex-app-server-runtime) for OpenAI/Codex models. `auto` (default) uses Hermes' standard chat completions; `codex_app_server` hands turns to a `codex app-server` subprocess for native shell, apply_patch, ChatGPT subscription auth, and migrated Codex plugins. Effective on next session. | -| `/personality` | Set a predefined personality | +| `/personality` | Set a predefined personality. `/personality none` (or `default` / `neutral`) clears the overlay and returns to base behavior. | | `/verbose` | Cycle tool progress display: off → new → all → verbose. Can be [enabled for messaging](#notes) via config. | | `/focus [on\|off\|status]` | Toggle **focus view** — a display-only reduced-output mode showing just your prompt and the final response. Composes with `/verbose`: turning it on snaps tool progress to `off` and remembers your previous mode, and `/focus off` restores it. Each turn ends with a dim recovery line (`⋯ 7 tool lines hidden · /focus off to show`) and a persistent `◉ focus` badge sits in the status bar so you always know you're in the reduced view. Nothing is sent differently to the model — detail is hidden, never discarded. | | `/fast [normal\|fast\|status]` | Toggle fast mode — OpenAI Priority Processing / Anthropic Fast Mode. Options: `normal`, `fast`, `status`. | @@ -226,7 +226,7 @@ The messaging gateway supports the following built-in commands inside Telegram, | `/stop` | Kill all running background processes and interrupt the running agent. | | `/model [provider:model]` | Show or change the model. Supports provider switches (`/model zai:glm-5`), custom endpoints (`/model custom:model`), named custom providers (`/model custom:local:qwen`), auto-detect (`/model custom`), and user-defined aliases (`/model fav`, `/model grok` — see [Custom model aliases](#custom-model-aliases)). Use `--global` to persist the change to config.yaml. **Note:** `/model` can only switch between already-configured providers. To add a new provider or set up API keys, use `hermes model` from your terminal (outside the chat session). **Cost note:** a mid-session model switch resets the prompt cache (the cache key includes the model), so the next message re-reads the whole conversation at full input price. | | `/codex-runtime [auto\|codex_app_server\|on\|off]` | Toggle the optional [Codex app-server runtime](../user-guide/features/codex-app-server-runtime). Persists to `model.openai_runtime` in config.yaml and evicts the cached agent so the next message picks up the new runtime. Effective on next session. | -| `/personality [name]` | Set a personality overlay for the session. | +| `/personality [name]` | Set a personality overlay for the session. `/personality none` (or `default` / `neutral`) clears it. | | `/fast [normal\|fast\|status]` | Toggle fast mode — OpenAI Priority Processing / Anthropic Fast Mode. | | `/retry` | Retry the last message. | | `/undo` | Remove the last exchange. | diff --git a/website/docs/reference/tools-reference.md b/website/docs/reference/tools-reference.md index 7ce5907e52..fee732350a 100644 --- a/website/docs/reference/tools-reference.md +++ b/website/docs/reference/tools-reference.md @@ -60,7 +60,7 @@ These two tools live in the `browser` toolset but only register when a Chrome De | Tool | Description | Requires environment | |------|-------------|----------------------| -| `delegate_task` | Spawn one or more subagents to work on tasks in isolated contexts. Each subagent gets its own conversation, terminal session, and toolset. Only the final summary is returned -- intermediate tool results never enter your context window. TWO… | — | +| `delegate_task` | Spawn subagents in isolated contexts; each gets its own conversation, terminal session, and toolset, and only its final summary returns to you. Provide 'goal' for a single task or 'tasks' for a parallel batch (limits and nesting rules… | — | ## `feishu_doc` toolset diff --git a/website/docs/user-guide/cli.md b/website/docs/user-guide/cli.md index 260ac47fc4..a6b9356c45 100644 --- a/website/docs/user-guide/cli.md +++ b/website/docs/user-guide/cli.md @@ -223,6 +223,8 @@ Set a predefined personality to change the agent's tone: Built-in personalities include: `helpful`, `concise`, `technical`, `creative`, `teacher`, `kawaii`, `catgirl`, `pirate`, `shakespeare`, `surfer`, `noir`, `uwu`, `philosopher`, `hype`. +To go back to the default (no overlay), use `/personality none` — `default` and `neutral` work too. + You can also define custom personalities in `~/.hermes/config.yaml`: ```yaml diff --git a/website/docs/user-guide/configuration.md b/website/docs/user-guide/configuration.md index 85d321335f..cfe85dbd55 100644 --- a/website/docs/user-guide/configuration.md +++ b/website/docs/user-guide/configuration.md @@ -1189,6 +1189,8 @@ auxiliary: # model: google/gemini-2.5-flash # base_url: "" # api_key: "" + # max_concurrency: 2 # Optional: cap simultaneous compression LLM calls so + # multiple sessions don't pile retries on a degraded provider # Auto-generated session titles. Empty language follows the conversation; # set e.g. "English" or "Japanese" to pin titles to one language. @@ -1217,6 +1219,15 @@ auxiliary: api_key: "" timeout: 30 + # Auto-generated short session titles after the first exchange + title_generation: + provider: "auto" + model: "" + base_url: "" + api_key: "" + timeout: 30 + # max_concurrency: 2 # Optional: cap simultaneous title-generation calls + # Kanban triage specifier — `hermes kanban specify ` (or the # dashboard's ✨ Specify button on Triage-column cards) uses this # slot to expand a one-liner into a concrete spec and promote the @@ -1266,6 +1277,25 @@ Each entry supports the same three knobs as any auxiliary task config: `fallback_chain` is available on any auxiliary task — `compression`, `vision`, `web_extract`, `approval`, `skills_hub`, `mcp`, etc. +### Limiting auxiliary concurrency + +`max_concurrency` caps in-flight LLM calls for auxiliary tasks such as `compression` and `title_generation` across the whole process. `auxiliary.vision.max_concurrency` is excluded: it already controls only vision's CPU-bound image encode/resize workers, not LLM requests. This is most useful when: + +- Many sessions can spawn background work simultaneously (Discord/Telegram channels, multiple terminals) +- Your provider is rate-limited or going through an incident and retries would amplify the burst + +The default is unlimited. A typical safety cap is `2`: + +```yaml +auxiliary: + title_generation: + max_concurrency: 2 + compression: + max_concurrency: 2 +``` + +The semaphore wraps the entire call including retries and fallbacks, so a single slow call counts only once toward the limit. + ### OpenRouter routing & Pareto Code for auxiliary tasks When an auxiliary task resolves to OpenRouter (either explicitly or via `provider: "main"` while your main agent is on OpenRouter), the main agent's `provider_routing` and `openrouter.min_coding_score` settings **do not propagate** — by design, each auxiliary task is independent. To set OpenRouter provider preferences or use the [Pareto Code router](/integrations/providers#openrouter-pareto-code-router) for a specific aux task, set them per-task via `extra_body`: @@ -1732,9 +1762,20 @@ When `display.runtime_footer.enabled: true`, Hermes appends a small runtime-cont display: runtime_footer: enabled: true - fields: ["model", "context_pct", "cwd"] # supported fields: model, context_pct, cwd + fields: ["model", "context_pct", "cwd"] # order shown; drop any to hide ``` +Supported fields: + +| Field | Renders | Example | +| --- | --- | --- | +| `model` | Bare model id, vendor prefix dropped | `gpt-5.4` | +| `context_pct` | Last-call context occupancy as a percent | `5%` | +| `latency` | Wall-clock duration of the turn | `22s`, `1m05s` | +| `cwd` | Home-relative working directory | `~` | + +The default field set is `["model", "context_pct", "cwd"]`. `latency` is opt-in — add it to `fields` to use it. Fields whose data is unavailable are skipped silently rather than rendering an empty slot. + The `/footer` slash command toggles this at runtime in any session. Example footer appended to a Telegram/Discord/Slack reply: diff --git a/website/docs/user-guide/features/fallback-providers.md b/website/docs/user-guide/features/fallback-providers.md index 3d2dde7820..b44b882f59 100644 --- a/website/docs/user-guide/features/fallback-providers.md +++ b/website/docs/user-guide/features/fallback-providers.md @@ -123,6 +123,8 @@ Prompt caches are keyed to the model (and on most providers, the account) servin :::info Per-Turn, Not Per-Session Fallback is **turn-scoped**: each new user message starts with the primary model restored. If the primary fails mid-turn, fallback activates for that turn only. On the next message, Hermes tries the primary again. Within a single turn, fallback activates at most once — if the fallback also fails, normal error handling takes over (retries, then error message). This prevents cascading failover loops within a turn while giving the primary model a fresh chance every turn. + +The per-turn retry is **reset-aware**: when the primary's credentials report a rate-limit reset time that hasn't elapsed yet (subscription windows like Claude Pro/Max's 5-hour blocks or Codex weekly limits report these as hours or days), Hermes skips the doomed retry and stays on the fallback until the reset passes — avoiding two pointless provider switches (and two prompt-cache invalidations) per turn. The moment the reset time elapses, the next turn goes back to the primary automatically. Transient 429s without a reset time keep the existing behavior: a short cooldown, then retry every turn. ::: ### Examples diff --git a/website/docs/user-guide/features/image-generation.md b/website/docs/user-guide/features/image-generation.md index 1eead3f050..9c27e7eebe 100644 --- a/website/docs/user-guide/features/image-generation.md +++ b/website/docs/user-guide/features/image-generation.md @@ -64,8 +64,13 @@ Your selection is saved to `config.yaml`: image_gen: model: fal-ai/flux-2/klein/9b use_gateway: false # true if using Nous Subscription + max_parallel_requests: 4 # concurrent images in one tool-call batch ``` +`max_parallel_requests` defaults to `4`. Hermes clamps it to at least one and +to the global tool-worker limit, so image providers receive bounded parallel +requests without allowing an image batch to bypass the agent's concurrency cap. + ### GPT-Image Quality The `fal-ai/gpt-image-1.5` and `fal-ai/gpt-image-2` request quality is pinned to `medium` (~$0.034–$0.06/image at 1024×1024). We don't expose the `low` / `high` tiers as a user-facing option so that Nous Portal billing stays predictable across all users — the cost spread between tiers is 3–22×. If you want a cheaper option, pick Klein 9B or Z-Image Turbo; if you want higher quality, use Nano Banana Pro or Recraft V4 Pro. diff --git a/website/docs/user-guide/features/personality.md b/website/docs/user-guide/features/personality.md index 14b26e4451..cfae4abdf9 100644 --- a/website/docs/user-guide/features/personality.md +++ b/website/docs/user-guide/features/personality.md @@ -227,6 +227,18 @@ Then switch to it with: /personality codereviewer ``` +## Resetting to the default + +To cancel the active personality overlay and return to base behavior (your `SOUL.md` persona), use any of: + +```text +/personality none +/personality default +/personality neutral +``` + +All three clear the overlay: the saved `agent.system_prompt` is emptied and the change takes effect on your next message. Running `/personality` with no arguments also lists `none` alongside the available presets. + ## Recommended workflow A strong default setup is: diff --git a/website/docs/user-guide/features/tool-search.md b/website/docs/user-guide/features/tool-search.md index ebb9b9d876..468989f787 100644 --- a/website/docs/user-guide/features/tool-search.md +++ b/website/docs/user-guide/features/tool-search.md @@ -80,7 +80,7 @@ tools: search_default_limit: 5 max_search_limit: 20 listing: auto # embed a grouped name+description catalog manifest - listing_max_tokens: 20000 + listing_max_tokens: 4000 ``` | Key | Default | Meaning | @@ -90,7 +90,7 @@ tools: | `search_default_limit` | `5` | Hits returned when the model calls `tool_search` without a `limit`. | | `max_search_limit` | `20` | Hard upper bound the model can request via `limit`. Range 1–50. | | `listing` | `auto` | Embed a skills-style manifest of every deferred tool (name + first sentence of its description, ≤60 chars, grouped by MCP server) in the `tool_search` bridge description. `auto` includes it when it fits the budget (falling back to names-only, then to the tier-2 server summary); `on`/`off` force either way. | -| `listing_max_tokens` | `20000` | Absolute cap on the embedded listing, regardless of context size. Range 200–60000. | +| `listing_max_tokens` | `4000` | Absolute cap on the embedded listing, regardless of context size. Range 200–60000. Large catalogs degrade to names-only or per-server summaries, keeping full schemas available through search. | ### Why the listing exists diff --git a/website/docs/user-guide/messaging/index.md b/website/docs/user-guide/messaging/index.md index 8cf1491ad3..61c47d5a42 100644 --- a/website/docs/user-guide/messaging/index.md +++ b/website/docs/user-guide/messaging/index.md @@ -192,7 +192,7 @@ platform network disconnect as an event-loop failure. |---------|-------------| | `/new` or `/reset` | Start a fresh conversation | | `/model [provider:model]` | Show or change the model (supports `provider:model` syntax) | -| `/personality [name]` | Set a personality | +| `/personality [name]` | Set a personality (`none` to reset) | | `/retry` | Retry the last message | | `/undo` | Remove the last exchange | | `/status` | Show session info | diff --git a/website/static/api/model-catalog.json b/website/static/api/model-catalog.json index 549edb0423..946ea073cb 100644 --- a/website/static/api/model-catalog.json +++ b/website/static/api/model-catalog.json @@ -1,6 +1,6 @@ { "version": 1, - "updated_at": "2026-07-31T15:41:39Z", + "updated_at": "2026-08-03T22:03:15Z", "metadata": { "source": "hermes-agent repo", "docs": "https://hermes-agent.nousresearch.com/docs/reference/model-catalog" @@ -101,7 +101,7 @@ "description": "dated snapshot of v4-flash" }, { - "id": "qwen/qwen3.7-max", + "id": "qwen/qwen3.8-max", "description": "" }, { @@ -238,7 +238,7 @@ "id": "deepseek/deepseek-v4-flash-0731" }, { - "id": "qwen/qwen3.7-max" + "id": "qwen/qwen3.8-max" }, { "id": "moonshotai/kimi-k3"