diff --git a/agent/agent_runtime_helpers.py b/agent/agent_runtime_helpers.py index 4c223bf607..3101db3e43 100644 --- a/agent/agent_runtime_helpers.py +++ b/agent/agent_runtime_helpers.py @@ -1588,7 +1588,11 @@ def create_openai_client(agent, client_kwargs: dict, *, reason: str, shared: boo # TCP keepalives so dead provider connections are detected (~60s) instead of hanging in # CLOSE-WAIT. Injected into the local copy only, so each client gets its own httpx.Client; # pinned by tests/run_agent/test_create_openai_client_reuse.py and - # test_sequential_chats_live.py. + # test_sequential_chats_live.py. What IS shared across those per-client wrappers is the + # connection pool: ``build_keepalive_http_client`` mounts a process-shared ``HTTPTransport`` + # behind a per-client view whose ``close()`` is a no-op for the pool, so a closed wrapper + # never takes a sibling's (or the successor's) connections with it + # (tests/agent/test_shared_http_transport.py). if "http_client" not in client_kwargs: keepalive_http = agent._build_keepalive_http_client(client_kwargs.get("base_url", ""), verify=httpx_verify) if keepalive_http is not None: @@ -2678,10 +2682,15 @@ def reapply_reasoning_echo_for_provider(agent, api_messages: list) -> int: return reapply_reasoning_echo(api_messages, agent._needs_thinking_reasoning_pad()) -def _iter_httpx_pool_objects(http_client: Any): - """Yield httpcore pool objects reachable from an httpx client, including mounted transports: +def _iter_httpx_pools_with_owner(http_client: Any): + """Yield ``(pool, owner)`` pairs reachable from an httpx client, including mounted transports: keepalive and proxy configs put live connections on ``client._mounts``, which a - ``_transport``-only walk misses.""" + ``_transport``-only walk misses. + + ``owner`` is ``None`` for a pool this client owns outright, or the ``_SharedTransport`` view + id when the pool is process-shared with other clients + (``process_bootstrap.build_keepalive_http_client``). Callers must then touch only the + in-flight requests stamped with that owner.""" seen_pools: set[int] = set() try: transports = [getattr(http_client, "_transport", None)] @@ -2696,11 +2705,18 @@ def _iter_httpx_pool_objects(http_client: Any): pool = transport if pool is not None and id(pool) not in seen_pools: seen_pools.add(id(pool)) - yield pool + owner = id(transport) if type(transport).__name__ == "_SharedTransport" else None + yield pool, owner except Exception: return +def _iter_httpx_pool_objects(http_client: Any): + """Yield httpcore pool objects reachable from an httpx client.""" + for pool, _owner in _iter_httpx_pools_with_owner(http_client): + yield pool + + def _connection_candidates(conn: Any): """Walk nested ``_connection`` wrappers (proxy tunnel → HTTP11/2).""" seen: set[int] = set() @@ -2734,20 +2750,30 @@ def _iter_pool_sockets(client: Any): try: # Some SDK wrappers *are* the httpx client; fall through so mount-aware discovery runs. http_client = getattr(client, "_client", None) - pools = list(_iter_httpx_pool_objects(client if http_client is None else http_client)) + pools = list(_iter_httpx_pools_with_owner(client if http_client is None else http_client)) except Exception: return + if not pools: + return + from agent.process_bootstrap import HERMES_TRANSPORT_OWNER_EXT seen: set[int] = set() - for pool in pools: + for pool, owner in pools: # ``is None``, not falsiness: an empty ``_connections`` must still let us walk in-flight ``_requests``. raw_conns = getattr(pool, "_connections", None) if raw_conns is None: raw_conns = getattr(pool, "_pool", None) - connections = list(raw_conns or []) - connections += [ - c for c in (getattr(r, "connection", None) for r in list(getattr(pool, "_requests", None) or [])) - if c is not None - ] + # A process-shared pool carries other clients' idle + in-flight connections: only this + # client's own in-flight requests (stamped by ``_SharedTransport.handle_request``) may be + # shut down. + connections = [] if owner is not None else list(raw_conns or []) + for pool_req in list(getattr(pool, "_requests", None) or []): + if owner is not None: + exts = getattr(getattr(pool_req, "request", None), "extensions", None) or {} + if exts.get(HERMES_TRANSPORT_OWNER_EXT) != owner: + continue + conn = getattr(pool_req, "connection", None) + if conn is not None: + connections.append(conn) for conn in connections: for candidate in _connection_candidates(conn): stream = getattr(candidate, "_network_stream", None) or getattr(candidate, "_stream", None) diff --git a/agent/auxiliary_client.py b/agent/auxiliary_client.py index 214bc812f2..1a2afc8262 100644 --- a/agent/auxiliary_client.py +++ b/agent/auxiliary_client.py @@ -1322,6 +1322,10 @@ class _CodexCompletionsAdapter: timeout = kwargs.get("timeout") if timeout is not None: resp_kwargs["timeout"] = timeout + # Per-request HTTP headers (OpenCode session affinity, Copilot x-initiator) map to real + # headers via the SDK kwarg — forward them. + if isinstance(kwargs.get("extra_headers"), dict) and kwargs["extra_headers"]: + resp_kwargs["extra_headers"] = dict(kwargs["extra_headers"]) # The Codex endpoint rejects max_output_tokens/temperature (400) — omit. extra_body = kwargs.get("extra_body") or {} if isinstance(extra_body, dict): @@ -1573,6 +1577,13 @@ class _AnthropicCompletionsAdapter: from agent.anthropic_adapter import _forbids_sampling_params if not _forbids_sampling_params(model): anthropic_kwargs["temperature"] = temperature + # Per-request HTTP headers (OpenCode session affinity) — the Anthropic SDK accepts + # ``extra_headers`` on messages.create/stream too. + if isinstance(kwargs.get("extra_headers"), dict) and kwargs["extra_headers"]: + anthropic_kwargs["extra_headers"] = { + **(anthropic_kwargs.get("extra_headers") or {}), + **kwargs["extra_headers"], + } # response_format: top-level gets the same translation as the extra_body form; when both # are present the extra_body form wins. Passthrough excludes ``reasoning``/``response_format`` # (already TRANSLATED to native fields — raw would 400 on strict gateways) and ``_`` Hermes plumbing. @@ -1839,15 +1850,21 @@ def _resolve_nous_pool_runtime_api(*, force_refresh: bool = False) -> Optional[t return api_key, base_url -def _resolve_nous_runtime_api(*, force_refresh: bool = False) -> Optional[tuple[str, str]]: - """Fresh Nous runtime credentials (pool first, then auth store + JWT refresh) — mirrors the main agent's 401 recovery.""" +def _resolve_nous_runtime_api( + *, force_refresh: bool = False, stale_access_token: Optional[str] = None +) -> Optional[tuple[str, str]]: + """Fresh Nous runtime credentials (pool first, then auth store + JWT refresh) — mirrors the main + agent's 401 recovery. ``stale_access_token`` is the bearer that just 401'd; with ``force_refresh`` + it lets the auth store adopt a sibling process's rotation instead of re-POSTing the shared grant.""" pooled = _resolve_nous_pool_runtime_api(force_refresh=force_refresh) if pooled is not None: return pooled try: from hermes_cli.auth import resolve_nous_runtime_credentials creds = resolve_nous_runtime_credentials( - timeout_seconds=env_float("HERMES_NOUS_TIMEOUT_SECONDS", 15), force_refresh=force_refresh, + timeout_seconds=env_float("HERMES_NOUS_TIMEOUT_SECONDS", 15), + force_refresh=force_refresh, + stale_access_token=stale_access_token or None, ) except Exception as exc: logger.debug("Auxiliary Nous runtime credential resolution failed: %s", exc) @@ -5040,7 +5057,7 @@ def _refresh_nous_auxiliary_client( the exact entry the stale one is served from. Keying on the resolved model or an empty task would leave the expired client immortal and every auxiliary call 401ing forever. """ - runtime = _resolve_nous_runtime_api(force_refresh=True) + runtime = _resolve_nous_runtime_api(force_refresh=True, stale_access_token=api_key) if runtime is None: return None, model fresh_key, fresh_base_url = runtime @@ -5803,7 +5820,10 @@ def _build_call_kwargs( or _endpoint_speaks_anthropic_messages(raw_base) or _is_anthropic_compat_endpoint(provider_norm, raw_base) ): kwargs["_reasoning_config"] = dict(reasoning_config) - return kwargs + # OpenCode relay session affinity — same key as the main turn so compression/title/vision + # calls stay on the conversation's warm backend. + from agent.opencode_affinity import merge_opencode_session_headers + return merge_opencode_session_headers(kwargs, provider, base_url, _runtime_main_value("session_id") or None) def _validate_llm_response( diff --git a/agent/chat_completion_helpers.py b/agent/chat_completion_helpers.py index a190cee4af..091b9052e7 100644 --- a/agent/chat_completion_helpers.py +++ b/agent/chat_completion_helpers.py @@ -1499,7 +1499,25 @@ def _build_chat_completions_kwargs(agent, api_messages, tools_for_api, reasoning def build_api_kwargs(agent, api_messages: list, tools_for_api: list | None = None) -> dict: - """Build the keyword arguments dict for the active API mode.""" + """Build the keyword arguments dict for the active API mode. + + Wraps the per-api_mode builder so the OpenCode ``x-opencode-session`` + affinity header rides on every OpenCode request regardless of transport + (chat_completions / codex_responses / anthropic_messages all route + OpenCode models). No-op for every other provider. + """ + from agent.opencode_affinity import merge_opencode_session_headers + + kwargs = _build_api_kwargs_for_mode(agent, api_messages, tools_for_api) + return merge_opencode_session_headers( + kwargs, + getattr(agent, "provider", None), + getattr(agent, "base_url", None), + getattr(agent, "session_id", None), + ) + + +def _build_api_kwargs_for_mode(agent, api_messages: list, tools_for_api: list | None = None) -> dict: # One-shot continuation override — consumed exactly once, on the FIRST # request this call builds (only one api_mode branch runs per invocation). reasoning_config = _reasoning_config_for_wire(agent) @@ -2079,6 +2097,9 @@ def _iteration_summary_api_messages(agent, messages: list) -> list: # Compression/resume can orphan a tool result whose parent tool_call was summarized away. api_messages = agent._sanitize_api_messages(api_messages) + # Same send-path vision eviction as the main loop (#89296). + from agent.context_compressor import evict_stale_outbound_tool_images + evict_stale_outbound_tool_images(api_messages) # Thinking-only assistant turns 400 on Anthropic-family providers; _thinking_prefill must # survive until here so the drop pass recognizes stubs after reasoning is stripped. api_messages = agent._drop_thinking_only_and_merge_users(api_messages) @@ -2541,6 +2562,15 @@ class _ToolCallAccumulator: self._notified: set = set() self._last_id_at_idx: dict = {} # raw_index -> last seen non-empty id self._active_slot_by_idx: dict = {} # raw_index -> current slot in acc + # Argument deltas are collected per slot and joined once in ``materialize`` — + # ``+=`` per chunk rebuilds the whole string every delta (quadratic on big args). + self._argument_parts: dict[int, list[str]] = {} + + def materialize(self) -> dict: + """Join buffered argument deltas into each entry's ``arguments``; idempotent. Returns ``acc``.""" + for idx, parts in self._argument_parts.items(): + self.acc[idx]["function"]["arguments"] = "".join(parts) + return self.acc def feed(self, tc_delta) -> Optional[str]: """Merge one delta; return the tool name the first time it is complete.""" @@ -2562,6 +2592,7 @@ class _ToolCallAccumulator: entry = self.acc.setdefault( idx, {"id": tc_id or "", "type": "function", "function": {"name": "", "arguments": ""}, "extra_content": None}, ) + parts = self._argument_parts.setdefault(idx, []) if tc_id: entry["id"] = tc_id tc_function = getattr(tc_delta, "function", None) @@ -2571,7 +2602,7 @@ class _ToolCallAccumulator: # NVIDIA NIM) resend the full name every chunk — += gives "read_fileread_file". entry["function"]["name"] = tc_function.name if getattr(tc_function, "arguments", None): - entry["function"]["arguments"] += tc_function.arguments + parts.append(tc_function.arguments) extra = getattr(tc_delta, "extra_content", None) if extra is None and hasattr(tc_delta, "model_extra"): extra = (tc_delta.model_extra if isinstance(tc_delta.model_extra, dict) else {}).get("extra_content") @@ -2832,6 +2863,7 @@ class _StreamingCall: return self._open_chat_stream({**next_api_kwargs, "stream": True, "timeout": timeout}) def _relay_final_response() -> dict[str, Any]: + tool_calls.materialize() message = {"role": role, "content": "".join(content_parts) or None, "reasoning_content": "".join(reasoning_parts) or None, "tool_calls": [tool_calls_acc[i] for i in sorted(tool_calls_acc)] or None} @@ -2920,6 +2952,7 @@ class _StreamingCall: # complete instead of silently discarding the action. self.result["partial_tool_names"].append(name) + tool_calls.materialize() self._close_managed_stream() if self._stream_attempt_was_cancelled(stream_attempt_id): raise _httpx.RemoteProtocolError(f"stream attempt {stream_attempt_id} was superseded") diff --git a/agent/client_lifecycle.py b/agent/client_lifecycle.py index 1833807832..c54b9c8ec2 100644 --- a/agent/client_lifecycle.py +++ b/agent/client_lifecycle.py @@ -490,7 +490,11 @@ class ClientLifecycleMixin: try: from hermes_cli.auth import resolve_nous_runtime_credentials timeout = env_float("HERMES_NOUS_TIMEOUT_SECONDS", 15) - creds = resolve_nous_runtime_credentials(timeout_seconds=timeout, force_refresh=force) + # Pass the bearer that just 401'd so a refresh already done by a sibling process is + # adopted instead of rotating the grant again. + creds = resolve_nous_runtime_credentials( + timeout_seconds=timeout, force_refresh=force, stale_access_token=self.api_key or None, + ) except Exception as exc: logger.debug("Nous credential refresh failed: %s", exc) return False diff --git a/agent/context_compressor.py b/agent/context_compressor.py index 65455d8c59..d33a834d6e 100644 --- a/agent/context_compressor.py +++ b/agent/context_compressor.py @@ -233,6 +233,21 @@ def _template_visible_role(message: Any) -> Optional[str]: return None if role == "tool" or (role == "assistant" and message.get("tool_calls")) else role +def _last_template_visible_role(messages: List[Dict[str, Any]]) -> Optional[str]: + """Last role a strict alternation template would count in *messages*. + + ``None`` when every row is template-exempt (tool flow only). + """ + return next( + ( + role + for role in (_template_visible_role(m) for m in reversed(messages)) + if role is not None + ), + None, + ) + + def _strip_persistence_markers(messages: List[Dict[str, Any]]) -> None: """Enforce the invariant: no assembled message carries a persistence marker. A leaked ``_db_persisted`` makes the child-session rotation flush skip the row, losing it from state.db. @@ -296,6 +311,22 @@ _SUMMARY_END_MARKER = "--- END OF CONTEXT SUMMARY — respond to the message bel _MERGED_PRIOR_CONTEXT_HEADER = "[PRIOR CONTEXT — for reference only; not a new message]" _MERGED_SUMMARY_DELIMITER = "[END OF PRIOR CONTEXT — COMPACTION SUMMARY BELOW]" +# Prefixes the copy of a still-running user task that compaction re-states after +# the handoff boundary (#100818). A cron run's only user turn is the job prompt +# in the protected head, so compaction leaves it BEFORE the summary — and +# SUMMARY_PREFIX tells the model to do nothing when no user message follows. +# Set on a compaction carrier when the in-flight task was merged onto it (the +# carrier ends the list, so a standalone user row would break alternation). +# conversation_compression._ensure_compressed_has_user_turn treats it as +# "intent present" so it does not insert a second copy of the same request. +_INFLIGHT_REPLAY_MERGED_KEY = "_inflight_replay_merged" + +_INFLIGHT_TASK_REPLAY_HEADER = ( + "[STILL IN PROGRESS — this is the active request, restated after the " + "compaction boundary because it was not finished yet. Continue it; do not " + "start over.]" +) + _SALVAGE_SUMMARY_MAX_CHARS = 8_000 _SALVAGE_KEEP_RECENT_TOOLS = 2 @@ -1140,6 +1171,23 @@ def _retire_stale_tool_result_images(result: List[Dict[str, Any]], keep_newest: return pruned +def evict_stale_outbound_tool_images( + api_messages: List[Dict[str, Any]], + keep_newest: int = _MAX_KEEP_TOOL_IMAGES, +) -> int: + """Drop stale screenshot/vision payloads from the per-call API copy. + + Compression's keep-newest pass only runs when prune/compress fires, and + the Anthropic adapter's screenshot eviction only sees nested + ``tool_result`` blocks. OpenAI-style ``image_url`` tool results + otherwise ride every subsequent request until a 413 forces the reactive + strip (#89286). Call this on the cloned ``api_messages`` list after + sanitization so older frames never leave the box (#89296). Do not pass + persisted history — the rewrite is send-path only. + """ + return _retire_stale_tool_result_images(api_messages, keep_newest=keep_newest) + + def _truncate_tool_call_args_json(args: str, head_chars: int = 200) -> str: """Shrink long string leaves in a tool-call arguments JSON blob, keeping it valid (providers 400 on malformed args).""" try: @@ -1716,7 +1764,7 @@ class ContextCompressor(MicroCompactionMixin, ContextEngine): self._reset_real_usage_pairing() self._last_compression_telemetry = self._active_compression_telemetry = None self._compression_telemetry_seed = None - self._proactive_prune_rearm_tokens = 0 + self._reset_proactive_prune_rearm() def bind_session_state(self, session_db: Any = None, session_id: str = "") -> None: """Bind the current session row so durable cooldowns can round-trip.""" @@ -1726,8 +1774,9 @@ class ContextCompressor(MicroCompactionMixin, ContextEngine): self._cooldown_persist_failed = False self._last_summary_error = None self._consecutive_timeout_failures = self._fallback_compression_streak = 0 - self._ineffective_compression_count = self._prellm_skip_count = self._proactive_prune_rearm_tokens = 0 + self._ineffective_compression_count = self._prellm_skip_count = 0 self._anti_thrash_recovery_deadline = self._structural_no_op_backoff_until = 0.0 + self._reset_proactive_prune_rearm() self.get_active_compression_failure_cooldown() self._load_fallback_compression_streak() self._load_ineffective_compression_count() @@ -2016,7 +2065,7 @@ class ContextCompressor(MicroCompactionMixin, ContextEngine): self._clear_compression_failure_cooldown() self._verify_compaction_cleared_threshold = self._last_compression_made_progress = False # Runway was computed against the previous model's trigger; clear the durable copy too. - self._proactive_prune_rearm_tokens = 0 + self._reset_proactive_prune_rearm() self._clear_durable_proactive_prune_rearm() # When the MINIMUM_CONTEXT_LENGTH floor binds on a small window, trigger near the top instead. @@ -2107,6 +2156,10 @@ class ContextCompressor(MicroCompactionMixin, ContextEngine): self.proactive_prune_min_reclaim_tokens = max(0, int(proactive_prune_min_reclaim_tokens or 0)) # A committed prune is a cache boundary: rearm only after the prompt regrows the reclaimed tokens. self._proactive_prune_rearm_tokens: int = 0 + # Dedup key for the over-threshold "reclamation no-oped" warning + # (#101889) so a tool loop riding above the threshold warns once per + # distinct reason + rearm snapshot instead of every iteration. + self._last_reclaim_block_warn: "tuple[str, int] | None" = None self.min_tail_user_messages = min_tail_user_messages self.summary_target_ratio = max(0.10, min(summary_target_ratio, 0.80)) self.quiet_mode = quiet_mode @@ -2513,35 +2566,115 @@ class ContextCompressor(MicroCompactionMixin, ContextEngine): ) return result, pruned + def _reset_proactive_prune_rearm(self) -> None: + """Fully rearm the proactive prune and let a future lockout warn again. + + Every path that zeroes the rearm mark (compaction, session + reset/end/rebind, model recalibration) is a reclamation or a fresh + start, so the over-threshold no-op dedup key must not survive it — + otherwise an identical lockout after a full compaction (rearm back + at 0) would be silent (#101889). + """ + self._proactive_prune_rearm_tokens = 0 + self._last_reclaim_block_warn = None + + def _billed_basis_over_threshold(self, current_tokens: "int | None") -> bool: + """Whether a provider-billed reading says the session is over threshold. + + ``current_tokens`` is the provider's ``prompt_tokens`` (or the + overhead-aware fallback estimate): it counts the system prompt and tool + schemas, which the message-only estimate behind + ``_proactive_prune_rearm_tokens`` does not. Used to stop schema + overhead from parking the prune rearm gate above a real request that is + already over ``threshold_tokens`` (#101889). + """ + return ( + current_tokens is not None + and self.threshold_tokens > 0 + and current_tokens >= self.threshold_tokens + ) + + def _warn_reclamation_no_op( + self, + reason: str, + current_tokens: "int | None", + before: "int | None" = None, + ) -> None: + """Warn when an over-threshold session's reclamation path no-ops. + + A session sitting above ``threshold_tokens`` with every reclamation + path declining is the failure mode from #101889: context keeps growing + until the provider's hard limit rejects the request, with nothing in + the log to explain it. Silent below the threshold (a declined prune + there is ordinary hysteresis, not a lockout). Deduped on + ``reason`` + the rearm snapshot so a busy tool loop logs once per + distinct state, not once per iteration; the key is cleared whenever + the session drops back under threshold or any reclamation resets the + rearm mark (prune commit, compaction, session reset/rebind, model + recalibration) so a later lockout warns again. + """ + # The explicit None check is redundant with the predicate; it narrows + # ``current_tokens`` for the type checker on the format below. + if current_tokens is None or not self._billed_basis_over_threshold( + current_tokens + ): + self._last_reclaim_block_warn = None + return + key = (reason, int(self._proactive_prune_rearm_tokens)) + if self._last_reclaim_block_warn == key: + return + self._last_reclaim_block_warn = key + logger.warning( + "Context is over the compression threshold (~%s of %s tokens) but " + "reclamation did not run: %s (message-token estimate %s, prune " + "rearm mark %s). The session may keep growing until the provider " + "rejects the request — /compact to compress history now.", + f"{int(current_tokens):,}", + f"{int(self.threshold_tokens):,}", + reason, + "n/a" if before is None else f"{int(before):,}", + f"{int(self._proactive_prune_rearm_tokens):,}", + ) + def prune_tool_results_only( self, messages: List[Dict[str, Any]], current_tokens: int | None = None, ) -> tuple[List[Dict[str, Any]], int]: """Deterministic, no-LLM tool-result prune gated on ``proactive_prune_tokens``. Protects the tail by message COUNT only. A commit breaks the prompt cache, so it requires ``proactive_prune_min_reclaim_tokens`` and a full regrowth runway; otherwise returns the INPUT - object as ``(messages, 0)``.""" - if ( - self.proactive_prune_tokens <= 0 - or (current_tokens is not None and current_tokens < self.proactive_prune_tokens) - or len(messages) <= self.protect_last_n + self._protect_head_size(messages) + 1 + object as ``(messages, 0)``. The rearm gate is measured on message bodies only, so it is + bypassed (never the reclaim gate) when a provider-billed ``current_tokens`` reading already + puts the request over ``threshold_tokens`` (#101889); every no-op taken while over threshold + is logged once per distinct reason.""" + if self.proactive_prune_tokens <= 0 or ( + current_tokens is not None and current_tokens < self.proactive_prune_tokens ): return messages, 0 + if len(messages) <= self.protect_last_n + self._protect_head_size(messages) + 1: + self._warn_reclamation_no_op("prune:tail_only", current_tokens) + return messages, 0 before = sum(_estimate_msg_budget_tokens(m) for m in messages) - if before < self._proactive_prune_rearm_tokens: + # Under-threshold runway skip is ordinary hysteresis (silent); above it the lockout is the bug. + if before < self._proactive_prune_rearm_tokens and not self._billed_basis_over_threshold(current_tokens): return messages, 0 # Capability gate first: a store without archive_and_compact makes every prune a no-op. session_db = getattr(self, "_session_db", None) session_id = getattr(self, "_session_id", "") if session_db and session_id and not callable(getattr(session_db, "archive_and_compact", None)): + self._warn_reclamation_no_op("prune:store_cannot_persist", current_tokens) return messages, 0 pruned_msgs, pruned_count = self._prune_old_tool_results( messages, protect_tail_count=self.protect_last_n, protect_tail_tokens=None, min_prune_chars=self.proactive_prune_min_result_chars, ) - # No-op contract: return the INPUT object so callers can gate on `result is not input`. + if not pruned_count: + # No-op contract: return the INPUT object so callers can gate on `result is not input`. + self._warn_reclamation_no_op("prune:nothing_eligible", current_tokens) + return messages, 0 # Prompt-cache hysteresis: commit only when the reclaim is meaningful. after = sum(_estimate_msg_budget_tokens(m) for m in pruned_msgs) reclaimed = max(0, before - after) - if not pruned_count or reclaimed < self.proactive_prune_min_reclaim_tokens: + if reclaimed < self.proactive_prune_min_reclaim_tokens: + self._warn_reclamation_no_op("prune:reclaim_below_minimum", current_tokens, before=before) return messages, 0 # Require a full trigger-sized regrowth before the next cache-breaking rewrite. runway = max(reclaimed, self.proactive_prune_tokens, self.proactive_prune_min_reclaim_tokens) @@ -2558,6 +2691,8 @@ class ContextCompressor(MicroCompactionMixin, ContextEngine): # Shared post-commit stamp site with the in-place commit and micro-compaction sync. stamp_db_persisted_markers(pruned_msgs) self._proactive_prune_rearm_tokens = next_rearm_tokens + # Reclamation just ran: let a future lockout warn again. + self._last_reclaim_block_warn = None return pruned_msgs, pruned_count def _compute_summary_budget(self, turns_to_summarize: List[Dict[str, Any]]) -> int: @@ -3602,6 +3737,164 @@ Write only the summary body. Do not include any preamble or prefix.""" return max(pair_end, head_end + 1) return adjusted + @classmethod + def _find_inflight_user_task( + cls, messages: List[Dict[str, Any]] + ) -> Optional[Dict[str, Any]]: + """Return the user turn that is still awaiting completion, or ``None``. + + Scans the WHOLE transcript, not just the compressible region: a cron + run's only user turn is the job prompt sitting in the protected head + (``protect_first_n`` keeps system + first user), which is exactly the + turn ``_find_last_user_message_idx`` cannot see (#100818). + + A turn is in-flight when the transcript does not already end with a + completed assistant reply — i.e. a text-bearing assistant message with + no pending ``tool_calls``. A trailing ``tool`` result or an assistant + message that still has ``tool_calls`` outstanding means the run was + interrupted mid-task and the instruction is still owed an answer. + + Handoff carriers and synthetic scaffolding rows are excluded via the + same filter pair as ``_find_last_user_message_idx``, so an idle session + whose only user-role row is an inherited summary yields ``None`` and is + never re-animated (#80622). + """ + from agent.conversation_compression import _is_real_user_message + + last_user_idx = -1 + for i in range(len(messages) - 1, -1, -1): + msg = messages[i] + # _is_real_user_message also rejects metadata-flagged scaffolding + # (_todo_snapshot_synthetic, recovery nudges, ...) that + # _is_actionable_user_turn cannot see. + if cls._is_actionable_user_turn(msg) and _is_real_user_message(msg): + last_user_idx = i + break + if isinstance(msg, dict) and msg.get(_INFLIGHT_REPLAY_MERGED_KEY): + # A previous cycle merged the live request onto this summary + # carrier; it is the only copy left, so it is still the task. + last_user_idx = i + break + if last_user_idx < 0: + return None + + for msg in reversed(messages[last_user_idx + 1:]): + if not isinstance(msg, dict) or msg.get("role") != "assistant": + # Trailing tool result (or anything else): still mid-task. + break + if msg.get("tool_calls"): + break + if _content_text_for_contains(msg.get("content")).strip(): + # Final answer already delivered — replaying the ask would + # hand the model finished work as a fresh instruction. + return None + # Empty assistant row (a bare reasoning/stub turn): keep looking. + return messages[last_user_idx] + + def _reappend_inflight_user_task( + self, + compressed: List[Dict[str, Any]], + inflight: Optional[Dict[str, Any]], + ) -> List[Dict[str, Any]]: + """Restate an unfinished user task after the compaction handoff. + + ``SUMMARY_PREFIX`` instructs the model to act only on a user message + that appears AFTER the summary, and to do nothing when none does. When + the single in-flight instruction lived in the protected head, the + assembled transcript orders it before the handoff and the run ends in a + ``[SILENT]`` no-op that the scheduler records as success (#100818). + + Re-append a copy of that turn after the surviving tail so the prefix's + "latest user message" pointer resolves to it again. If the transcript + already ends on a template-visible user row, appending a second one + would break user/assistant alternation, so the restatement is merged + onto the handoff carrier instead — after ``_SUMMARY_END_MARKER``, which + is the boundary the prefix's rule is written against. + """ + if inflight is None or not compressed: + return compressed + + carrier_idx = -1 + for idx in range(len(compressed) - 1, -1, -1): + if self._is_context_summary_message(compressed[idx]): + carrier_idx = idx + break + if carrier_idx < 0: + # No handoff was emitted — nothing reordered the instruction. + return compressed + + for msg in compressed[carrier_idx + 1:]: + if self._is_actionable_user_turn( + msg + ) and not self._is_synthetic_compression_user_turn(msg): + # A real request already follows the summary. + return compressed + + carrier = compressed[carrier_idx] + carrier_text = _content_text_for_contains(carrier.get("content")) + if _SUMMARY_END_MARKER not in carrier_text: + return compressed + if carrier_text.split(_SUMMARY_END_MARKER, 1)[1].strip(): + # The _force_user_leading layout keeps the live request on the + # carrier itself, after the marker. Already actionable. + return compressed + + task_text = _content_text_for_contains(inflight.get("content")).strip() + if _INFLIGHT_TASK_REPLAY_HEADER in task_text: + # Already a restatement from an earlier compaction (standalone row + # or merged onto a carrier): take the text after the header so a + # task that survives >1 cycle never stacks headers or drags the + # old summary along. + task_text = task_text.rsplit(_INFLIGHT_TASK_REPLAY_HEADER, 1)[1].strip() + if not task_text: + return compressed + + if not self.quiet_mode: + logger.info( + "Re-appending the in-flight user task after the compaction " + "handoff so it stays actionable (#100818)" + ) + + last_visible_role = _last_template_visible_role(compressed) + if inflight.get(_INFLIGHT_REPLAY_MERGED_KEY): + # Never copy a summary carrier (metadata would mark the replay + # synthetic): restate as a plain user row. + replay = {"role": "user", "content": task_text} + else: + replay = _fresh_compaction_message_copy(inflight) + replay.pop(_COMPACTION_TAIL_MARKER, None) + if isinstance(replay.get("content"), str): + # Plain text: rebuild from the header-stripped task text so a + # task surviving several compactions never stacks headers. + replay["content"] = _INFLIGHT_TASK_REPLAY_HEADER + "\n" + task_text + else: + # Multimodal parts: keep them, prepend the header text part. + replay["content"] = _append_text_to_content( + replay.get("content"), + _INFLIGHT_TASK_REPLAY_HEADER + "\n", + prepend=True, + ) + drop_stale_api_content(replay) + + if last_visible_role == "user": + # Alternation is judged on template-visible rows only (tool_calls / + # tool rows are exempt), so a user-pinned summary followed by a + # tool tail still "ends on user": a standalone user row would break + # the Mistral-style pre-flight check (#58753). Merge onto the + # carrier instead and flag it — the carrier's own metadata marks it + # synthetic, and without the flag _ensure_compressed_has_user_turn + # would insert a second copy of the same request. + carrier["content"] = _append_text_to_content( + carrier.get("content"), + "\n\n" + _INFLIGHT_TASK_REPLAY_HEADER + "\n" + task_text, + ) + carrier[_INFLIGHT_REPLAY_MERGED_KEY] = True + drop_stale_api_content(carrier) + return compressed + + compressed.append(replay) + return compressed + def _ensure_last_n_user_messages_in_tail( self, messages: List[Dict[str, Any]], cut_idx: int, head_end: int, n: int, ) -> int: @@ -3923,7 +4216,7 @@ Write only the summary body. Do not include any preamble or prefix.""" last_head_role: Optional[str] = "user" if compressed: # None = all-exempt head: the summary opens the visible sequence and must be "user". - last_head_role = next((r for r in map(_template_visible_role, reversed(compressed)) if r is not None), None) + last_head_role = _last_template_visible_role(compressed) first_tail_visible_idx, first_tail_role = next( ((idx, role) for idx, role in enumerate(map(_template_visible_role, tail_messages)) if role is not None), (None, None), @@ -3973,8 +4266,13 @@ Write only the summary body. Do not include any preamble or prefix.""" self, compressed: List[Dict[str, Any]], messages: List[Dict[str, Any]], n_messages: int, ) -> List[Dict[str, Any]]: """Post-assembly cleanup: orphan pairs, media, savings, markers, replay prune, mem trim.""" - self.compression_count += 1 + # Single-prompt cron shape: the only live instruction sits in the protected head, BEFORE the + # handoff, and SUMMARY_PREFIX reads that as "nothing to do" — restate it past the boundary + # (#100818). Sanitize FIRST: the trailing-in-flight exemption (#79278) walks back from the list + # end, and a replay user row there would strip a genuinely pending assistant(tool_calls). compressed = self._sanitize_tool_pairs(compressed) + compressed = self._reappend_inflight_user_task(compressed, self._find_inflight_user_task(messages)) + self.compression_count += 1 # Replace historical image payloads with placeholders; multi-MB base64 blobs otherwise # exceed body limits. compressed = _strip_historical_media(compressed) @@ -4010,7 +4308,7 @@ Write only the summary body. Do not include any preamble or prefix.""" # Batch marker holds MORE history than the rolling summary: reset micro state so it can't # supersede/defrag content it lacks; the next micro pass rehydrates from the batch marker. self._reset_micro_compact_cursor_state() - self._proactive_prune_rearm_tokens = 0 + self._reset_proactive_prune_rearm() return compressed def compress( diff --git a/agent/conversation_compression.py b/agent/conversation_compression.py index fc79d05b6a..c807ae85f0 100644 --- a/agent/conversation_compression.py +++ b/agent/conversation_compression.py @@ -59,6 +59,9 @@ _SPLIT_FAILURE_COOLDOWN_SECONDS = 60 # the phrase intact when rewording. Idle/preflight/retry lines lack it; is_compaction_progress_status covers those. COMPACTION_STATUS_MARKER = "Compacting context" COMPACTION_STATUS = f"🗜️ {COMPACTION_STATUS_MARKER} — summarizing earlier conversation so I can continue..." +# Periodic heartbeat re-emitted while a long compression is still running so remote transports with +# idle-turn watchdogs (#98371) see progress. Same marker as COMPACTION_STATUS so consumers classify it alike. +COMPACTION_HEARTBEAT_STATUS = f"🗜️ {COMPACTION_STATUS_MARKER} — still summarizing earlier conversation so I can continue..." COMPACTION_DONE_STATUS = "✓ Context compaction complete — continuing turn..." @@ -113,7 +116,8 @@ CONTEXT_OVERFLOW_BLOCKED_WARNING_TEMPLATE = ( # Formatted from the same constants the emission sites use, so noise-filter tests exercise the ACTUAL wording. ROUTINE_COMPRESSION_STATUS_SAMPLES = ( - COMPACTION_STATUS, COMPACTION_DONE_STATUS, PRE_API_COMPRESSION_STATUS_TEMPLATE.format(tokens=123456), + COMPACTION_STATUS, COMPACTION_HEARTBEAT_STATUS, COMPACTION_DONE_STATUS, + PRE_API_COMPRESSION_STATUS_TEMPLATE.format(tokens=123456), PREFLIGHT_COMPRESSION_STATUS_TEMPLATE.format(tokens=120000, threshold=100000), IDLE_COMPACTION_STATUS_TEMPLATE.format(idle_seconds=3600, tokens=120000), COMPRESSION_RETRY_TOO_LARGE_STATUS_TEMPLATE.format(tokens=250000, attempt=1, cap=3), @@ -1388,7 +1392,8 @@ class _CompressionActivityHeartbeat: """Refresh the agent inactivity tracker while compression blocks in an aux call.""" def __init__( - self, agent: Any, interval_seconds: float | None = None, commit_fence: Optional[CompressionCommitFence] = None + self, agent: Any, interval_seconds: float | None = None, *, emit_client_status: bool = False, + commit_fence: Optional[CompressionCommitFence] = None, ) -> None: self._agent = agent self._commit_fence = commit_fence @@ -1404,6 +1409,10 @@ class _CompressionActivityHeartbeat: except (TypeError, ValueError): interval_seconds = 60.0 self._interval_seconds = max(0.1, interval_seconds) + # Only a compression that opened a VISIBLE compaction phase (the + # routine start status was emitted) keeps it alive with heartbeats; + # quiet context engines emit neither (#98371 follow-up). + self._emit_client_status = emit_client_status self._stop = threading.Event() self._thread = threading.Thread(target=self._run, name="compression-activity-heartbeat", daemon=True) @@ -1449,11 +1458,40 @@ class _CompressionActivityHeartbeat: return touch(desc, provenance=ActivityProvenance.AGENT_COMPRESSION, force_persist=force_persist) + def _emit_progress_status(self) -> None: + """Re-publish the compacting status so remote transports see progress. + + Compression can stream for minutes with no deltas, tool events, or + status lines reaching remote transports. Idle-progress watchdogs on + those clients (e.g. the Android relay app's 180s turn watchdog) + treat the silence as a dead turn and fire ``session.interrupt`` — + killing a healthy compression mid-flight and rolling back its work, + which retriggers on the next prompt and loops forever on sessions + near the context ceiling (#98371). + + Routed through ``agent._emit_status`` like every other compaction + status: same "lifecycle" key (the TUI gateway re-tags it to + ``compacting``; Telegram edits one bubble per key), same chat-platform + filter, same CLI print path. + """ + if not self._emit_client_status: + return + emit = getattr(self._agent, "_emit_status", None) + if not callable(emit): + return + try: + emit(COMPACTION_HEARTBEAT_STATUS) + except Exception: + logger.debug( + "status emit error in compression heartbeat", exc_info=True + ) + def _run(self) -> None: while not self._stop.wait(self._interval_seconds): if self._should_suppress(): return self._touch("context compression in progress") + self._emit_progress_status() def _direct_messages_for_pre_compress_memory(messages: Any) -> list[dict[str, Any]]: @@ -1941,7 +1979,12 @@ def _ensure_compressed_has_user_turn(original_messages: list, compressed: list) """Preserve human intent, not merely a synthetic user-role placeholder.""" if any(_is_real_user_message(message) for message in compressed) or _compressed_has_busy_steer(compressed): return "already_present" - from agent.context_compressor import COMPRESSION_CONTINUATION_USER_CONTENT, _fresh_compaction_message_copy + from agent.context_compressor import ( + _INFLIGHT_REPLAY_MERGED_KEY, COMPRESSION_CONTINUATION_USER_CONTENT, _fresh_compaction_message_copy, + ) + if any(isinstance(message, dict) and message.get(_INFLIGHT_REPLAY_MERGED_KEY) for message in compressed): + # The in-flight request was restated onto the summary carrier (#100818); an anchor would duplicate it. + return "already_present" # One reversed scan over BOTH kinds: scanning steer then user would let an older # consumed steer outrank a newer real user request and replay it. for message in reversed(original_messages): @@ -2063,7 +2106,7 @@ class _CompactionLifecycle: def __init__(self, agent: Any, status_emitted: bool) -> None: self._agent = agent - self._status_emitted = status_emitted + self.status_emitted = status_emitted self._done_emitted = False self.commit_status = "aborted" @@ -2074,7 +2117,7 @@ class _CompactionLifecycle: # Suppressed start → no terminal edge. Non-compacting aborts (lock contender, # cancelled fence) opt in via force_terminal so clients can retire their phase. # Failure warnings go through _emit_warning and are never suppressed here. - if self._status_emitted and (self.commit_status == "committed" or force_terminal): + if self.status_emitted and (self.commit_status == "committed" or force_terminal): _emit_compaction_done(self._agent) @@ -2103,6 +2146,11 @@ class _CompressionLease: # cannot win between acquiring the lock and having a way to release it. self._lock_setup_entered = False + @property + def status_emitted(self) -> bool: + """True when the routine compaction start status was shown (heartbeats may follow it).""" + return self._lifecycle.status_emitted + def begin_lock_setup(self) -> bool: if self._commit_fence is None: return True @@ -2796,8 +2844,9 @@ def _warn_summary_or_aux_fallback(agent: Any) -> None: def _reset_read_dedup_caches(task_id: str, *, skills: bool = True) -> None: - """Clear the file-read (and skill_view) repeat-read dedup caches after a boundary. - Original read content was summarized away, so a re-read must return full content, not a "file unchanged" stub. + """Advance the file-read (and skill_view) repeat-read dedup to a fresh generation after a boundary. + The mtime map is kept: the first read of each unchanged key returns full content compaction may have + omitted; later reads return stubs, and stub-hit counters restart at the same boundary (#84857). """ with contextlib.suppress(Exception): from tools.file_tools import reset_file_dedup @@ -3144,7 +3193,9 @@ def _run_summary_phase( bypass_cooldown=bypass_cooldown, ) messages_before_compression = copy.deepcopy(messages) - _activity_heartbeat = _CompressionActivityHeartbeat(agent, commit_fence=commit_fence).start() + _activity_heartbeat = _CompressionActivityHeartbeat( + agent, commit_fence=commit_fence, emit_client_status=lease.status_emitted, + ).start() compressed = _run_summary_dispatch( agent, messages, compress_fn, compress_kwargs, commit_fence=commit_fence, attempt_generation=attempt.generation, hard_cancel_event=hard_cancel_event, @@ -3495,7 +3546,7 @@ def _compress_context_via_codex_app_server( logger.info("codex app-server compaction started: session=%s messages=%d tokens=~%s", _sid, len(messages), _tokens) with contextlib.suppress(Exception): agent._emit_status(COMPACTION_STATUS) - _activity_heartbeat = _CompressionActivityHeartbeat(agent).start() + _activity_heartbeat = _CompressionActivityHeartbeat(agent, emit_client_status=True).start() try: result = codex_session.compact_thread() except BaseException: @@ -3717,7 +3768,7 @@ def try_shrink_image_parts_in_messages(api_messages: list, *, max_dimension: int __all__ = [ - "COMPACTION_STATUS", "COMPACTION_DONE_STATUS", "COMPACTION_STATUS_MARKER", "is_compaction_progress_status", + "COMPACTION_STATUS", "COMPACTION_DONE_STATUS", "COMPACTION_HEARTBEAT_STATUS", "COMPACTION_STATUS_MARKER", "is_compaction_progress_status", "check_compression_model_feasibility", "replay_compression_warning", "compress_context", "try_shrink_image_parts_in_messages", ] diff --git a/agent/conversation_loop.py b/agent/conversation_loop.py index 39337f15ef..898fa3a52d 100644 --- a/agent/conversation_loop.py +++ b/agent/conversation_loop.py @@ -328,6 +328,36 @@ def _is_stale_copilot_credential_error(status_code: Optional[int], error_message )) +def _pressure_with_real_floor(compressor: Any, rough_tokens: int) -> int: + """Floor the ROUGH pre-API pressure estimate at the last REAL prompt size. + + Applied only on the fallback path -- when ``anchored_context_tokens`` has + no valid anchor (first request, transcript rewritten under the anchor, + provider never reported usage). A valid anchor is provider-exact and is + used as-is; in particular on MoA turns the anchor deliberately uses the + pre-fold aggregator usage while ``last_real_prompt_tokens`` holds the + folded figure, so flooring an anchored value would re-add fan-out tokens + the anchor exists to exclude. + + On the rough path, non-ASCII text (Cyrillic, Greek, Polish, ...) + under-counts by up to ~2x, so a session can sit at the provider's real + context ceiling while the rough figure stays under the compaction + threshold -- on silent-clip providers (ollama /v1) that is a truncation + death spiral the reactive overflow handler never sees (observed live: + real prompts 64,842->64,995 against a 55,705 threshold). The provider's + last reported prompt_tokens is authoritative; never let the rough figure + fall below it. Skipped for exactly one turn after a compaction, when + last_real_prompt_tokens still holds the stale pre-compression value + (#36718's awaiting_real_usage_after_compression window). + """ + last_real = int(getattr(compressor, "last_real_prompt_tokens", 0) or 0) + if last_real > rough_tokens and not getattr( + compressor, "awaiting_real_usage_after_compression", False + ): + return last_real + return rough_tokens + + def _ollama_context_limit_error(agent: Any, request_tokens: int) -> Optional[str]: """Return a user-facing error when Ollama is loaded with too little context.""" runtime_ctx = getattr(agent, "_ollama_num_ctx", None) diff --git a/agent/lsp/client.py b/agent/lsp/client.py index efa37dabe6..b8956cb37d 100644 --- a/agent/lsp/client.py +++ b/agent/lsp/client.py @@ -70,6 +70,11 @@ def file_uri(path: str) -> str: return "file://" + quote(abs_path, safe="/:") +def _folder(root: str) -> Dict[str, str]: + """Build an LSP ``WorkspaceFolder`` for ``root``.""" + return {"name": os.path.basename(root.rstrip(os.sep)) or root, "uri": file_uri(root)} + + def uri_to_path(uri: str) -> str: """Inverse of :func:`file_uri`.""" if not uri.startswith("file://"): @@ -125,6 +130,9 @@ class LSPClient: seed_diagnostics_on_first_push: bool = False) -> None: self.server_id = server_id self.workspace_root = workspace_root + # Roots this server serves. Single-root servers only ever hold ``workspace_root``; + # multi-root servers (pyright) grow this via ``add_workspace_folder`` instead of a second process. + self.workspace_folders: List[str] = [workspace_root] self._command = list(command) self._env = env self._cwd = cwd or workspace_root @@ -262,7 +270,17 @@ class LSPClient: await self._cleanup_process() def _workspace_folders(self) -> List[Dict[str, str]]: - return [{"name": "workspace", "uri": file_uri(self.workspace_root)}] + return [_folder(r) for r in self.workspace_folders] + + async def add_workspace_folder(self, root: str) -> None: + """Attach another root to a running multi-root server. Idempotent; the folder is recorded + before the notification is sent so concurrent callers for the same root only announce once.""" + if root in self.workspace_folders: + return + self.workspace_folders.append(root) + await self._send_notification( + "workspace/didChangeWorkspaceFolders", {"event": {"added": [_folder(root)], "removed": []}}, + ) async def _initialize(self) -> None: params = { diff --git a/agent/lsp/manager.py b/agent/lsp/manager.py index 9ff23c18ab..0d677649ef 100644 --- a/agent/lsp/manager.py +++ b/agent/lsp/manager.py @@ -2,7 +2,9 @@ :class:`LSPService` bridges the synchronous file_operations layer and the async :class:`agent.lsp.client.LSPClient`: one asyncio loop in a background thread, one lazily -spawned client per ``(server_id, workspace_root)``, a **broken-set** of pairs that failed +spawned client per ``(server_id, workspace_root)`` — servers flagged ``multi_root`` (pyright) get ONE +client per ``server_id`` and further roots (typically sibling git worktrees) are attached to the running +process via ``workspace/didChangeWorkspaceFolders`` — a **broken-set** of pairs that failed to spawn/initialize (never retried for the life of the service), and a **delta baseline** per file (``snapshot_baseline()`` runs BEFORE a write; the next ``get_diagnostics_sync()`` returns only diagnostics not in it). Off unless config enables it. @@ -30,6 +32,12 @@ _Key = Tuple[str, str] _Diags = List[Dict[str, Any]] +def _client_key(srv: ServerDef, root: str) -> _Key: + """Cache key for the client serving ``root``: multi-root servers share one process per + ``server_id``; everything else is keyed per resolved project root.""" + return (srv.server_id, "" if srv.multi_root else root) + + class _BackgroundLoop: """A daemon thread owning one asyncio loop; :meth:`run` blocks on a coroutine.""" @@ -274,9 +282,10 @@ class LSPService: return already_broken = key in self._broken self._broken.add(key) + ckey = _client_key(srv, key[1]) with self._state_lock: - client = self._clients.pop(key, None) - self._last_used.pop(key, None) + client = self._clients.pop(ckey, None) + self._last_used.pop(ckey, None) if client is not None: try: # Fire-and-forget shutdown — we're already on a slow path. @@ -301,8 +310,9 @@ class LSPService: """Return a snapshot of the service for ``hermes lsp status``.""" with self._state_lock: clients = [ - {"server_id": k[0], "workspace_root": k[1], "state": c.state, "running": c.is_running} - for k, c in self._clients.items() + {"server_id": c.server_id, "workspace_root": c.workspace_root, + "workspace_folders": list(c.workspace_folders), "state": c.state, "running": c.is_running} + for c in self._clients.values() ] broken = list(self._broken) return { @@ -349,7 +359,7 @@ class LSPService: if not (ws and gated and srv): return [] with self._state_lock: - client = self._clients.get((srv.server_id, ws)) + client = self._clients.get(_client_key(srv, ws)) return list(client.diagnostics_for(file_path, fresh_only=True)) if client else [] async def _get_or_spawn(self, file_path: str) -> Optional[LSPClient]: @@ -367,39 +377,47 @@ class LSPService: if root is None: eventlog.log_disabled(srv.server_id, file_path, "exclude marker hit (server gated off)") return None - key = (srv.server_id, root) - if key in self._broken: + if (srv.server_id, root) in self._broken: return None + key = _client_key(srv, root) with self._state_lock: client = self._clients.get(key) if client is not None and client.is_running: self._last_used[key] = time.time() - eventlog.log_active(*key) - return client + eventlog.log_active(srv.server_id, root) + return await self._attach_root(srv, client, root) spawning = self._spawning.get(key) owner = spawning is None if owner: spawning = self._spawning[key] = asyncio.get_running_loop().create_future() if not owner: try: - return await spawning + client = await spawning except Exception: # noqa: BLE001 return None + return await self._attach_root(srv, client, root) if client is not None else None try: client = await self._spawn_client(srv, root) if client is None: - self._broken.add(key) + self._broken.add((srv.server_id, root)) else: with self._state_lock: self._clients[key] = client self._last_used[key] = time.time() - eventlog.log_active(*key) + eventlog.log_active(srv.server_id, root) spawning.set_result(client) return client finally: with self._state_lock: self._spawning.pop(key, None) + @staticmethod + async def _attach_root(srv: ServerDef, client: LSPClient, root: str) -> LSPClient: + """Multi-root servers: announce ``root`` to the shared process instead of spawning another.""" + if srv.multi_root: + await client.add_workspace_folder(root) + return client + async def _spawn_client(self, srv: ServerDef, root: str) -> Optional[LSPClient]: """Resolve the binary and start a client; ``None`` (after logging) when either fails.""" ctx = ServerContext( @@ -425,10 +443,10 @@ class LSPService: def _touch(self, client: LSPClient) -> None: """Refresh last-used; guarded on membership so a client reaped mid-operation can't resurrect its entry.""" - key = (client.server_id, client.workspace_root) with self._state_lock: - if key in self._clients: - self._last_used[key] = time.time() + for key, c in self._clients.items(): + if c is client: + self._last_used[key] = time.time() async def _start_idle_reaper(self) -> None: self._idle_reaper_task = asyncio.create_task(self._idle_reaper_loop()) diff --git a/agent/lsp/servers.py b/agent/lsp/servers.py index c47016ca90..e148ded539 100644 --- a/agent/lsp/servers.py +++ b/agent/lsp/servers.py @@ -72,6 +72,9 @@ class ServerDef: build_spawn: _SpawnFn seed_first_push: bool = False description: str = "" + # Server handles ``workspace/didChangeWorkspaceFolders``: one process serves every project root + # (git worktrees included) as extra workspaceFolders instead of one process per root. + multi_root: bool = False def matches(self, file_path: str) -> bool: return _file_ext_or_basename(file_path) in self.extensions @@ -263,20 +266,21 @@ def _server(server_id: str, extensions: Tuple[str, ...], description: str, *, markers: Optional[Sequence[str]] = None, excludes: Sequence[str] = (), resolve_root: Optional[_RootFn] = None, build_spawn: Optional[_SpawnFn] = None, which: Sequence[str] = (), args: Sequence[str] = (), install_pkg: Optional[str] = None, - base_init: Optional[Dict[str, Any]] = None, seed: bool = False) -> ServerDef: + base_init: Optional[Dict[str, Any]] = None, seed: bool = False, + multi_root: bool = False) -> ServerDef: """Registry entry factory: defaults to marker-based root + single-binary spawn.""" return ServerDef( server_id, extensions, resolve_root or _markers_root(markers, excludes), build_spawn or _simple_spawn(server_id, which or (server_id,), args, install_pkg, base_init, seed), - seed_first_push=seed, description=description, + seed_first_push=seed, description=description, multi_root=multi_root, ) SERVERS: List[ServerDef] = [ _server("pyright", (".py", ".pyi"), "Python — Microsoft pyright", markers=["pyproject.toml", "setup.py", "setup.cfg", "requirements.txt", "Pipfile", "pyrightconfig.json"], - build_spawn=_spawn_pyright), + build_spawn=_spawn_pyright, multi_root=True), _server("typescript", (".ts", ".tsx", ".js", ".jsx", ".mjs", ".cjs", ".mts", ".cts"), "JavaScript/TypeScript — typescript-language-server", resolve_root=_root_typescript, which=("typescript-language-server",), args=("--stdio",), install_pkg="typescript-language-server", seed=True), diff --git a/agent/lsp/workspace.py b/agent/lsp/workspace.py index 5b220fcab7..799798cae6 100644 --- a/agent/lsp/workspace.py +++ b/agent/lsp/workspace.py @@ -120,7 +120,9 @@ def nearest_root( # Excludes are checked before markers at each level. if present(cur, excludes_list): return None - if present(cur, markers_list): + # A directory holding __init__.py is a Python package, never a project root (hermes_cli/setup.py + # matched the python marker list and gave every package dir its own pyright). + if not present(cur, ["__init__.py"]) and present(cur, markers_list): return str(cur) if ceiling_path is not None and cur == ceiling_path: return None diff --git a/agent/model_metadata.py b/agent/model_metadata.py index 10ada00b60..e094b27e1f 100644 --- a/agent/model_metadata.py +++ b/agent/model_metadata.py @@ -337,8 +337,10 @@ DEFAULT_CONTEXT_LENGTHS = { # https://api-docs.deepseek.com/zh-cn/quick_start/pricing "deepseek-v4-pro": 1_000_000, "deepseek-v4-flash": 1_000_000, "deepseek-chat": 1_000_000, "deepseek-reasoner": 1_000_000, "deepseek": 128000, - # Meta; Thinking Machines inkling (covers inkling-small and :free/:batch variants) - "llama": 131072, "inkling": 1_048_576, + # Meta; Muse Spark family (1.1/1.2/1.3, -contributor(-free), meta/ prefixed) is 1M per OpenRouter, + # models.dev and api.commandcode.ai /models — keep the "muse-spark" prefix (bare "muse" would match + # muse-image/muse-voice). Thinking Machines inkling (covers inkling-small and :free/:batch variants) + "llama": 131072, "muse-spark-1.3": 1_048_576, "muse-spark": 1_048_576, "inkling": 1_048_576, # Qwen — https://help.aliyun.com/zh/model-studio/developer-reference/ (3.8-max/flash # 1M verified on OpenRouter & Nous portal 2026-08; qwen3-max = 256K Coding Plan snapshot) "qwen3.8-max": 1_000_000, "qwen3.8-flash": 1_000_000, "qwen3.6-plus": 1048576, "qwen3.7-plus": 1048576, @@ -1269,6 +1271,7 @@ def _model_name_suggests_minimax_m3(model: str) -> bool: # shorter matching key and the 256K fallback — the threshold is inferred from them. _PRE_CATALOG_STALE_KEYS = frozenset({ "minimax-m3", # 1M; "minimax" catch-all persisted 204,800 + "muse-spark-1.3", "muse-spark", # 1M; pre-entry builds fell through to the 256K fallback "grok-4.3", "grok-4.6", # 1M / 500K; "grok-4" catch-all persisted 256,000 "grok-4-fast", "grok-4.20", # 2M; fell through to the 256K fallback "qwen3.6-plus", # 1M; "qwen" catch-all persisted 131,072 @@ -1775,8 +1778,9 @@ def _resolve_provider_aware_context_length(model: str, base_url: str, api_key: s if base_url and source == persist_on: save_context_length(model, base_url, ctx) return ctx - if effective_provider == "gmi" and base_url: - # GMI exposes authoritative context_length via /models, but it is not in models.dev yet. + if effective_provider in {"gmi", "commandcode", "commandcode-anthropic"} and base_url: + # GMI and CommandCode expose authoritative context_length via /models (e.g. muse-spark 1M) but are + # not in models.dev, and as known providers they skip step 2's probe — else they fell to 256K. ctx = _resolve_endpoint_context_length(model, base_url, api_key=api_key) if ctx is not None: return ctx @@ -1924,14 +1928,23 @@ def _is_cjk_token_dense_char(ch: str) -> bool: def estimate_tokens_rough(text: str) -> int: - """Rough token estimate: ceil(chars/4), CJK/Hangul/Kana codepoints ~1 token each. Ceiling keeps - short texts from estimating 0. Runs on every preflight walk, so the all-ASCII case stays O(1).""" + """Rough token estimate: CJK/Hangul/Kana codepoints ~1 token each; everything else ceil(UTF-8 bytes/4). + Ceiling keeps short texts from estimating 0. Runs on every preflight walk, so all-ASCII stays O(1). + + Byte-counting (not chars) is the corrective for non-CJK, non-ASCII text: Cyrillic/Greek/Arabic are 2 + bytes/char so count ~chars/2, matching real BPE cost (~2-3 chars/token) where chars/4 under-counted + ~2x and let sessions ride the provider ceiling below the compaction threshold. Calibrated vs + cl100k/o200k/Qwen2.5 (estimate/real): Russian 0.67->1.24, Arabic 0.53->0.96, Hindi 0.34->0.90, + Greek 0.37->0.68; accented Latin barely moves (French 1.02->1.03). errors="replace": lone surrogates + (routine in tool output; see message_sanitization) must not turn an estimate into a raise.""" if not text: return 0 text = str(text) - # ``str.isascii()`` is a flag check on CPython; non-ASCII without CJK (accents, Cyrillic, emoji) also gets chars/4. - dense = 0 if text.isascii() else len(text) - len(_CJK_DENSE_RE.sub("", text)) - return dense + ((len(text) - dense + 3) // 4) + if text.isascii(): # flag check on CPython; ASCII cannot contain token-dense CJK + return (len(text) + 3) // 4 + stripped = _CJK_DENSE_RE.sub("", text) + dense = len(text) - len(stripped) + return dense + ((len(stripped.encode("utf-8", "replace")) + 3) // 4) def estimate_messages_tokens_rough(messages: List[Dict[str, Any]], *, charge_stale_thinking: bool = True) -> int: diff --git a/agent/models_dev.py b/agent/models_dev.py index 5d4b93038a..1ae381ce38 100644 --- a/agent/models_dev.py +++ b/agent/models_dev.py @@ -114,7 +114,11 @@ PROVIDER_TO_MODELS_DEV: Dict[str, str] = { "minimax-oauth": "minimax", "minimax-cn": "minimax-cn", "deepseek": "deepseek", "alibaba": "alibaba", "qwen-oauth": "alibaba", "copilot": "github-copilot", "ai-gateway": "vercel", "opencode-zen": "opencode", - "opencode-go": "opencode-go", "kilocode": "kilo", "fireworks": "fireworks-ai", + "opencode-go": "opencode-go", + # opencode-free is Zen-hosted (hermes_cli/models.py) and models.dev's "opencode" catalog lists + # its *-contributor-free SKUs; without this alias every opencode-free lookup missed models.dev. + "opencode-free": "opencode", + "kilocode": "kilo", "fireworks": "fireworks-ai", "huggingface": "huggingface", "gemini": "google", "google": "google", "xai": "xai", "xai-oauth": "xai", # OAuth is a transport path for the same xAI catalog diff --git a/agent/opencode_affinity.py b/agent/opencode_affinity.py new file mode 100644 index 0000000000..4f0cc5522e --- /dev/null +++ b/agent/opencode_affinity.py @@ -0,0 +1,84 @@ +"""``x-opencode-session`` — OpenCode relay session-affinity header. + +OpenCode (opencode.ai Zen/Go/free relay) pins requests that share an +``x-opencode-session`` value to the same upstream backend, which is what +keeps its prompt cache warm across the turns of one conversation. The value +only has to be opaque and consistent per conversation, so it is derived the +same way as the other conversation-affinity hints Hermes already sends +(OpenRouter's sticky ``session_id``, xAI's ``x-grok-conv-id``): the +host-declared routing scope first, then the ambient conversation root, then +the physical session id — normalized through ``_cache_scope_from_session_id`` +so cron fires of one job share a scope. + +Every OpenCode request — main turn on any transport, auxiliary calls +(compression, titles, vision, MoA) — goes through :func:`opencode_session_headers` +so the header cannot drift per code path. +""" + +from __future__ import annotations + +from typing import Any, Optional + +OPENCODE_SESSION_HEADER = "x-opencode-session" + + +def is_opencode_target(provider: Optional[str], base_url: Optional[str]) -> bool: + """True when *provider* or *base_url* addresses the OpenCode relay. + + Matches the built-in opencode-zen/go/free providers, custom + ``opencode--*`` providers, and any base_url hosted on opencode.ai. + """ + try: + from hermes_cli.models import opencode_provider_family + + if opencode_provider_family(provider) is not None: + return True + except Exception: + pass + try: + from agent.anthropic_endpoints import _is_opencode_endpoint + + return _is_opencode_endpoint(str(base_url or "")) + except Exception: + return False + + +def opencode_session_headers( + provider: Optional[str], + base_url: Optional[str], + session_id: Optional[str] = None, +) -> dict[str, str]: + """Return ``{"x-opencode-session": }`` for OpenCode targets, else ``{}``.""" + if not is_opencode_target(provider, base_url): + return {} + try: + from agent.portal_tags import get_affinity_scope, get_conversation_context + from agent.transports.codex import _cache_scope_from_session_id + + key = _cache_scope_from_session_id( + get_affinity_scope() or get_conversation_context() or session_id + ) + except Exception: + key = str(session_id or "") + return {OPENCODE_SESSION_HEADER: key} if key else {} + + +def merge_opencode_session_headers( + kwargs: dict[str, Any], + provider: Optional[str], + base_url: Optional[str], + session_id: Optional[str] = None, +) -> dict[str, Any]: + """Merge the affinity header into ``kwargs["extra_headers"]`` (in place). + + Existing per-request headers win, so a caller-pinned value is preserved. + Non-OpenCode targets are left untouched. + """ + headers = opencode_session_headers(provider, base_url, session_id) + if headers: + existing = kwargs.get("extra_headers") + merged = dict(existing) if isinstance(existing, dict) else {} + for key, value in headers.items(): + merged.setdefault(key, value) + kwargs["extra_headers"] = merged + return kwargs diff --git a/agent/periodic_scheduler.py b/agent/periodic_scheduler.py new file mode 100644 index 0000000000..a1ba07d666 --- /dev/null +++ b/agent/periodic_scheduler.py @@ -0,0 +1,118 @@ +"""One process-wide timer thread for periodic maintenance callbacks. + +Replaces the per-child ``while not stop.wait(interval): body()`` daemon +threads (delegate heartbeat, durable turn-lease refresher, turn-liveness +watchdog). With ~130 in-process subagents those added 2-3 sleeping OS +threads per child; this module runs every periodic body on ONE daemon +thread ordered by a heap of due times. + +Semantics match the loop they replace: the first call happens ``interval`` +seconds after :func:`schedule`, and each following call ``interval`` seconds +after the previous body *returned* (drift-free wrt. body duration was never +a property of the old loops either). A body that returns ``False`` stops +itself; a body that raises is logged at debug and rescheduled — one bad +callback must never kill the shared thread. +""" + +from __future__ import annotations + +import heapq +import itertools +import logging +import threading +import time +from typing import Callable, Optional + +logger = logging.getLogger(__name__) + +_THREAD_NAME = "hermes-periodic-scheduler" + + +class ScheduledHandle: + """Cancel token for one scheduled periodic callback.""" + + __slots__ = ("_fn", "_interval", "_cancelled", "_scheduler") + + def __init__(self, scheduler: "PeriodicScheduler", fn: Callable[[], object], interval: float): + self._scheduler = scheduler + self._fn = fn + self._interval = interval + self._cancelled = False + + @property + def cancelled(self) -> bool: + return self._cancelled + + def cancel(self, wait: Optional[float] = None) -> None: + """Stop future runs. ``wait`` (seconds) additionally blocks until an + in-flight run of this callback finishes — the analogue of + ``thread.join(timeout=wait)`` on the old per-child thread.""" + self._scheduler._cancel(self, wait) + + +class PeriodicScheduler: + def __init__(self) -> None: + self._cond = threading.Condition() + self._heap: list = [] # (due, seq, handle) + self._seq = itertools.count() + self._thread: Optional[threading.Thread] = None + self._running: Optional[ScheduledHandle] = None + + def schedule(self, fn: Callable[[], object], interval: float) -> ScheduledHandle: + handle = ScheduledHandle(self, fn, float(interval)) + with self._cond: + heapq.heappush(self._heap, (time.monotonic() + handle._interval, next(self._seq), handle)) + if self._thread is None or not self._thread.is_alive(): + self._thread = threading.Thread(target=self._run, name=_THREAD_NAME, daemon=True) + self._thread.start() + self._cond.notify() + return handle + + def _cancel(self, handle: ScheduledHandle, wait: Optional[float]) -> None: + with self._cond: + handle._cancelled = True + self._cond.notify() + if wait and self._running is handle and threading.current_thread() is not self._thread: + self._cond.wait_for(lambda: self._running is not handle, timeout=wait) + + def _run(self) -> None: + while True: + with self._cond: + while True: + if not self._heap: + self._cond.wait() + continue + due, _, handle = self._heap[0] + if handle._cancelled: + heapq.heappop(self._heap) + continue + delay = due - time.monotonic() + if delay > 0: + self._cond.wait(delay) + continue + heapq.heappop(self._heap) + self._running = handle + break + stop = False + try: + stop = handle._fn() is False + except Exception: + logger.debug("periodic callback %r raised", handle._fn, exc_info=True) + with self._cond: + self._running = None + if stop: + handle._cancelled = True + elif not handle._cancelled: + heapq.heappush( + self._heap, + (time.monotonic() + handle._interval, next(self._seq), handle), + ) + self._cond.notify_all() + + +_DEFAULT = PeriodicScheduler() + + +def schedule(fn: Callable[[], object], interval: float) -> ScheduledHandle: + """Run ``fn()`` every ``interval`` seconds on the shared scheduler thread.""" + return _DEFAULT.schedule(fn, interval) diff --git a/agent/process_bootstrap.py b/agent/process_bootstrap.py index 3f0fae516f..4d640dbe30 100644 --- a/agent/process_bootstrap.py +++ b/agent/process_bootstrap.py @@ -13,6 +13,7 @@ import os import selectors import socket import sys +import threading import time import urllib.request from typing import Any, Optional @@ -23,6 +24,19 @@ from utils import base_url_hostname, normalize_proxy_url _OPENAI_CLS_CACHE = None _HAPPY_EYEBALLS_DELAY_SECONDS = 0.25 +# Process-wide pool of sync ``httpx.HTTPTransport`` objects shared by every +# keepalive client with the same (verify, proxy, happy-eyeballs) identity. +# Each delegated child AIAgent used to get its own transport = its own TLS +# pool, so a fan-out of N children held N separate socket sets to the same +# provider. Bounded: past the cap, callers get a private transport again. +_SHARED_TRANSPORTS: dict[tuple, Any] = {} +_SHARED_TRANSPORTS_LOCK = threading.Lock() +_SHARED_TRANSPORTS_MAX = 32 +# ``request.extensions`` key stamped by ``_SharedTransport.handle_request``; +# the socket-abort walker in agent_runtime_helpers uses it to find only the +# owning client's in-flight connections on a shared pool. +HERMES_TRANSPORT_OWNER_EXT = "hermes_transport_owner" + def _interleave_addrinfos(addrinfos: list[tuple]) -> list[tuple]: """Round-robin the resolved address families (deduped), preserving resolver order within each.""" @@ -283,6 +297,86 @@ def _get_proxy_for_base_url(base_url: Optional[str]) -> Optional[str]: return proxy +def _shared_transport_cls(): + """Lazily define the per-client transport view (httpx import is deferred).""" + global _SharedTransport + if _SharedTransport is not None: + return _SharedTransport + import httpx + + class _SharedTransportImpl(httpx.BaseTransport): + """Per-client view of a process-shared ``httpx.HTTPTransport``. + + ``httpx.Client.close()`` closes every mounted transport, and each OpenAI client still + owns its own ``httpx.Client`` (closing one client must never poison the next), so the + mounted object absorbs that close while the shared pool keeps serving other clients. + ``handle_request`` stamps the owning view into ``request.extensions`` so socket-abort + sweeps target only this client's in-flight connections on the shared pool. + """ + + __slots__ = ("_inner", "_closed") + + def __init__(self, inner: Any) -> None: + self._inner = inner + self._closed = False + + @property + def _pool(self) -> Any: # httpx-private; socket walkers and tests introspect it + return getattr(self._inner, "_pool", None) + + def handle_request(self, request: Any) -> Any: + if self._closed: + raise RuntimeError("Cannot send a request, as the client has been closed.") + request.extensions[HERMES_TRANSPORT_OWNER_EXT] = id(self) + return self._inner.handle_request(request) + + def close(self) -> None: + # Never closes the shared ``_inner``; idle connections are reaped by keepalive_expiry + # and the pool lives for the process (see ``close_shared_transports``). + self._closed = True + + _SharedTransportImpl.__name__ = _SharedTransportImpl.__qualname__ = "_SharedTransport" + _SharedTransport = _SharedTransportImpl + return _SharedTransport + + +_SharedTransport: Any = None + + +def _shared_transport_key(base_url: str, verify: Any, proxy: Optional[str]) -> tuple: + """Identity under which sync direct transports are pooled process-wide.""" + if verify is True or verify is False: + verify_key: Any = verify + elif isinstance(verify, str): + verify_key = ("path", verify) + else: + verify_key = ("id", id(verify)) # SSLContext / custom object: share by identity only + return (verify_key, proxy, _uses_codex_cloud_transport(base_url)) + + +def _get_shared_transport(key: tuple, build) -> Any: + with _SHARED_TRANSPORTS_LOCK: + transport = _SHARED_TRANSPORTS.get(key) + if transport is None: + transport = build() + if len(_SHARED_TRANSPORTS) < _SHARED_TRANSPORTS_MAX: + _SHARED_TRANSPORTS[key] = transport + return transport + + +def close_shared_transports() -> int: + """Really close every process-shared transport (test teardown / atexit).""" + with _SHARED_TRANSPORTS_LOCK: + transports = list(_SHARED_TRANSPORTS.values()) + _SHARED_TRANSPORTS.clear() + for transport in transports: + try: + transport.close() + except Exception: + pass + return len(transports) + + def build_keepalive_http_client(base_url: str = "", *, async_mode: bool = False, verify: Any = True) -> Optional[Any]: """httpx client for OpenAI SDK calls with env-only proxy policy (None on failure). @@ -291,6 +385,12 @@ def build_keepalive_http_client(base_url: str = "", *, async_mode: bool = False, reaps idle connections before reverse proxies' 30-60 s timeouts (a custom socket_options transport broke streaming and stripped TCP_NODELAY). ``verify`` goes on the client AND the mounts, since a mounted transport owns its SSL context. + + Every call returns a NEW ``httpx.Client`` (per-client close semantics), but sync clients + with the same (verify, proxy, happy-eyeballs) identity mount the SAME underlying + ``HTTPTransport`` through a ``_SharedTransport`` view, so N delegated children share one + connection pool + SSL context. Async clients are never shared: an httpcore async pool is + bound to the event loop that first used it. Proxy-backed clients keep httpx's own transport. """ try: import httpx @@ -301,12 +401,34 @@ def build_keepalive_http_client(base_url: str = "", *, async_mode: bool = False, client_cls = httpx.AsyncClient if async_mode else httpx.Client mounts = None if proxy is None: - mounts = {"http://": transport_cls(verify=verify), "https://": transport_cls(verify=verify)} - # Async transports race natively (anyio happy_eyeballs_delay=0.25). - if not async_mode and _uses_codex_cloud_transport(base_url): - for transport in mounts.values(): + happy_eyeballs = not async_mode and _uses_codex_cloud_transport(base_url) + # One pool serves every agent in the process, so its ceiling must cover a whole + # fan-out of concurrently streaming children. (Client-level ``limits`` never reach + # mounted transports — they used to run on httpx defaults, keepalive_expiry=5s.) + direct_limits = limits if async_mode else httpx.Limits( + max_keepalive_connections=50, max_connections=1000, keepalive_expiry=20.0, + ) + + def _build_direct(): + transport = transport_cls(verify=verify, limits=direct_limits) + # Async transports race natively (anyio happy_eyeballs_delay=0.25). + if happy_eyeballs: _enable_happy_eyeballs(transport) - return client_cls(limits=limits, timeout=timeout, proxy=proxy, mounts=mounts, verify=verify) + return transport + + if async_mode: + mounts = {"http://": _build_direct(), "https://": _build_direct()} + else: + key = _shared_transport_key(base_url, verify, proxy) + view_cls = _shared_transport_cls() + mounts = { + f"{scheme}://": view_cls(_get_shared_transport((scheme, *key), _build_direct)) + for scheme in ("http", "https") + } + # Default transport = the https view; otherwise httpx builds a third, never-used + # direct transport (pool + SSL context) per client. + return client_cls(limits=limits, timeout=timeout, transport=mounts["https://"], mounts=mounts) + return client_cls(limits=limits, timeout=timeout, proxy=proxy, mounts=mounts or None, verify=verify) except Exception: return None @@ -325,5 +447,6 @@ OpenAI = _OpenAIProxy() __all__ = [ "OpenAI", "_OpenAIProxy", "_load_openai_cls", "_SafeWriter", "_install_safe_stdio", "_get_proxy_from_env", - "_get_proxy_for_base_url", "build_keepalive_http_client", "enable_happy_eyeballs_on_client", + "_get_proxy_for_base_url", "build_keepalive_http_client", "close_shared_transports", + "enable_happy_eyeballs_on_client", ] diff --git a/agent/prompt_builder.py b/agent/prompt_builder.py index 0889204a47..e5fb6bbc99 100644 --- a/agent/prompt_builder.py +++ b/agent/prompt_builder.py @@ -8,6 +8,7 @@ import contextvars import json import logging import os +import queue import sys import threading from collections import OrderedDict @@ -31,6 +32,51 @@ from utils import atomic_json_write logger = logging.getLogger(__name__) +# Default read deadline for context files (SOUL.md, AGENTS.md, .cursorrules, +# ...); overridable via ``context_file_read_timeout`` in config.yaml. +# Intentionally short: network-backed filesystems (iCloud Drive, OneDrive, +# NFS) can fault-in an evicted file and block a cold read indefinitely, which +# stalls system-prompt assembly before the first turn. +_CONTEXT_FILE_READ_TIMEOUT_SECS = 5.0 + + +def _get_context_file_read_timeout() -> float: + """``context_file_read_timeout`` from config.yaml, else the 5s default.""" + val = _config_readonly("context_file_read_timeout").get("context_file_read_timeout") + if isinstance(val, (int, float)) and val > 0: + return float(val) + return _CONTEXT_FILE_READ_TIMEOUT_SECS + + +def _read_text_with_timeout(path: Path, timeout: Optional[float] = None) -> Optional[str]: + """``path.read_text()`` on a daemon thread so a slow file can't stall startup. + + Returns the text, or ``None`` after *timeout* seconds (logged at WARNING; + the orphaned reader thread finishes on its own). Read errors propagate to + the caller exactly as a direct ``read_text`` would, so existing + ``try/except`` handling at each site is unchanged. + """ + if timeout is None: + timeout = _get_context_file_read_timeout() + result: "queue.Queue[tuple[bool, object]]" = queue.Queue(maxsize=1) + + def _reader() -> None: + try: + result.put((True, path.read_text(encoding="utf-8"))) + except Exception as exc: # re-raised on the caller thread + result.put((False, exc)) + + threading.Thread(target=_reader, daemon=True, name=f"context-read:{path.name}").start() + try: + ok, value = result.get(timeout=timeout) + except queue.Empty: + logger.warning("Context file %s read timed out after %.1fs; skipping", path, timeout) + return None + if ok: + return value # type: ignore[return-value] + raise value # type: ignore[misc] + + def _scan_context_content(content: str, filename: str) -> str: """Scan a context file (AGENTS.md, .cursorrules, SOUL.md) for injection; matches are BLOCKED. @@ -245,14 +291,17 @@ TOOL_USE_ENFORCEMENT_GUIDANCE = ( "user. Responses that only describe intentions without acting are not acceptable." ) -TOOL_USE_ENFORCEMENT_MODELS = ("gpt", "codex", "gemini", "gemma", "grok", "glm", "qwen", "deepseek") +# "muse" = Meta Muse Spark: on defaults it answers in prose with 0 tool calls and the turn closes on +# finish_reason=stop (#96550). +TOOL_USE_ENFORCEMENT_MODELS = ("gpt", "codex", "gemini", "gemma", "grok", "glm", "qwen", "deepseek", "muse") # Models that receive OPENAI_MODEL_EXECUTION_GUIDANCE when agent.execution_guidance is "auto" (agentic-eval -# traces showed the same failure modes). Gemini/Gemma get GOOGLE_MODEL_OPERATIONAL_GUIDANCE instead; Claude -# does not exhibit these modes. Any model can opt in via config.yaml (`true` or a substring list). +# traces showed the same failure modes; Muse Spark stops after a chat-only turn on defaults). Gemini/Gemma get +# GOOGLE_MODEL_OPERATIONAL_GUIDANCE instead; Claude does not exhibit these modes. Any model can opt in via +# config.yaml (`true` or a substring list). EXECUTION_GUIDANCE_MODELS = ( "gpt", "codex", "grok", - "deepseek", "kimi", "qwen", "glm", "minimax", "mimo", "mistral", + "deepseek", "kimi", "qwen", "glm", "minimax", "mimo", "mistral", "muse", ) # Universal "finish the job" guidance (ALL models): don't stop after a stub, never @@ -734,17 +783,30 @@ _BACKEND_PROBE_CMD = ( def _run_backend_probe(env_type: str, terminal_tool) -> str: """Execute the probe command inside a freshly built backend; "" when it yields nothing.""" + from tools.terminal_tool_backends import _ssh_config_from_config + config = terminal_tool._get_env_config() # Mirrors tools/terminal_tool.py's live-command assembly (`_create_environment` is the factory). env = terminal_tool._create_environment( env_type=env_type, image=config.get(_BACKEND_IMAGE_KEYS[env_type], "") if env_type in _BACKEND_IMAGE_KEYS else "", cwd=config.get("cwd", ""), timeout=config.get("timeout", 180), - ssh_config=terminal_tool._ssh_config_from_config(config) if env_type == "ssh" else None, + ssh_config=_ssh_config_from_config(config) if env_type == "ssh" else None, container_config=({k: config.get(k, d) for k, d in _CONTAINER_CONFIG_DEFAULTS} if terminal_tool._is_container_backend(env_type) else None), task_id="prompt-backend-probe", host_cwd=config.get("host_cwd"), ) - result = env.execute(_BACKEND_PROBE_CMD, timeout=4) + try: + result = env.execute(_BACKEND_PROBE_CMD, timeout=4) + finally: + # One-shot `uname`; without teardown the backend leaves a second idle sandbox + # (task_id="prompt-backend-probe") running for the whole process next to the agent's own. + # ssh is left alone: no task-scoped sandbox, and its cleanup() closes a ControlMaster socket + # (keyed by user@host:port) shared with the agent's real environment; ControlPersist expires it. + if env_type != "ssh": + try: + terminal_tool._cleanup_env(env, force_remove=True) + except Exception: + logger.debug("Backend probe cleanup failed", exc_info=True) if result.get("returncode") != 0: logger.debug("Backend probe returned non-zero: %r", result) return "" @@ -781,6 +843,11 @@ def _probe_remote_backend(env_type: str) -> str | None: return formatted or None +def _clear_backend_probe_cache() -> None: + """Test helper — drop the backend probe cache so monkeypatched backends take effect.""" + _BACKEND_PROBE_CACHE.clear() + + def _local_host_hints() -> list[str]: """Host OS / home / cwd block for a local terminal backend (tools run on this host).""" import platform @@ -1295,7 +1362,7 @@ def load_soul_md(context_length: Optional[int] = None, home_override: "Path | No if not soul_path.exists(): return None try: - content = soul_path.read_text(encoding="utf-8").strip() + content = (_read_text_with_timeout(soul_path) or "").strip() if not content: return None return _truncate_content(_scan_context_content(content, "SOUL.md"), "SOUL.md", context_length=context_length, @@ -1310,7 +1377,7 @@ def _read_context_file(path: Path) -> str: if not path.exists(): return "" try: - return path.read_text(encoding="utf-8").strip() + return (_read_text_with_timeout(path) or "").strip() except Exception as e: logger.debug("Could not read %s: %s", path, e) return "" diff --git a/agent/relay_llm.py b/agent/relay_llm.py index df52bc269e..d75f5ae8ff 100644 --- a/agent/relay_llm.py +++ b/agent/relay_llm.py @@ -711,6 +711,12 @@ def _provider_request( headers = getattr(request, "headers", None) if isinstance(headers, dict): headers = {k: v for k, v in headers.items() if str(k).lower() not in _RELAY_INTERNAL_PROVIDER_HEADERS} + # Relay's managed-call trace header maps to ``extra_headers`` for known SDK adapters and custom + # requests that already use that container; other native transports take protocol kwargs directly + # and may reject an SDK-only argument. Non-trace middleware headers are preserved as before. + supports_extra_headers = _RELAY_PROTOCOL_BY_API_MODE.get(_api_mode(metadata)) is not None or "extra_headers" in original + if headers and not supports_extra_headers: + headers = {k: v for k, v in headers.items() if str(k).lower() != "traceparent"} if headers: final["extra_headers"] = {**dict(final.get("extra_headers") or {}), **headers} return final diff --git a/agent/relay_runtime.py b/agent/relay_runtime.py index 6b709b376e..9039e431c9 100644 --- a/agent/relay_runtime.py +++ b/agent/relay_runtime.py @@ -269,7 +269,8 @@ class _ProcessRelayPluginConfiguration: except Exception as exc: raise RuntimeError("Hermes Relay dynamic plugin activation failed") from exc if self._activation is None: - # Reached only after explicit opt-in; Relay owns any ambient layering. + # Reached only after explicit opt-in. Relay 0.8 no longer layers repository-local + # configuration onto this explicitly selected payload. _resolve_plugin_awaitable(relay.plugin.initialize(plugin_config)) return True diff --git a/agent/relay_tools.py b/agent/relay_tools.py index 54ccb74bd4..bba27a7db8 100644 --- a/agent/relay_tools.py +++ b/agent/relay_tools.py @@ -15,7 +15,7 @@ logger = logging.getLogger(__name__) def execute( tool_name: str, args: dict[str, Any], callback: Callable[[dict[str, Any]], Any], *, - session_id: str, metadata: dict[str, Any] | None = None, + session_id: str, tool_call_id: str | None = None, metadata: dict[str, Any] | None = None, ) -> tuple[Any, dict[str, Any]]: """Run one tool call through Relay and return its final arguments.""" runtime, session, parent = relay_runtime.resolve_execution_context(session_id) @@ -42,13 +42,13 @@ def execute( callback_error = exc raise raw_result.update(value=result, json=_jsonable(result)) - return raw_result["json"] + return runtime.relay.ToolExecutionResult(raw_result["json"]) try: managed = _run_awaitable( runtime.run_in_session_async( session, runtime.relay.tools.execute, tool_name, _jsonable(args), invoke, - handle=parent, metadata=_jsonable(metadata or {}), + handle=parent, metadata=_jsonable(metadata or {}), tool_call_id=tool_call_id or None, ) ) except BaseException as exc: @@ -61,9 +61,12 @@ def execute( ) return raw_result["value"], observed_args raise - if "value" in raw_result and _json_equal(managed, raw_result["json"]): + managed_result = managed.result + if "value" in raw_result and _json_equal(managed_result, raw_result["json"]): return raw_result["value"], observed_args - return (managed if isinstance(managed, str) else json.dumps(_jsonable(managed), ensure_ascii=False)), observed_args + if isinstance(managed_result, str): + return managed_result, observed_args + return json.dumps(_jsonable(managed_result), ensure_ascii=False), observed_args def _jsonable(value: Any) -> Any: diff --git a/agent/search_policy.py b/agent/search_policy.py new file mode 100644 index 0000000000..2972d5cae8 --- /dev/null +++ b/agent/search_policy.py @@ -0,0 +1,30 @@ +"""Shared directory pruning policy for broad recursive scans. + +These names identify version-control internals, dependency trees, generated +artifacts, caches, and backup copies that are not useful results for broad +agent-facing discovery. Ordinary search callers may still target an explicit +path; broad diagnostic probes should apply this policy to recursive walks. +""" + +from __future__ import annotations + + +# Keep this policy conservative and name-based so it works for local and remote +# shell backends alike. The same set is used by context discovery and search +# probes; adding a directory here protects every broad recursive consumer. +SEARCH_PRUNE_DIR_NAMES = frozenset({ + # Version-control internals. + ".git", ".hg", ".svn", + # Dependency and vendored trees. + "node_modules", "venv", ".venv", "site-packages", "dist-packages", + "vendor", "third_party", + # Generated/build output. + "build", "dist", "target", "out", "coverage", + ".next", ".turbo", ".parcel-cache", ".nuxt", ".svelte-kit", + # Python and package-manager caches. + "__pycache__", ".cache", ".Trash", ".tox", ".nox", ".mypy_cache", + ".pytest_cache", ".ruff_cache", ".npm", ".yarn", ".pnpm-store", + ".gradle", ".m2", ".nuget", + # Backup copies. + "backups", "backup", ".backups", +}) diff --git a/agent/ssl_verify.py b/agent/ssl_verify.py index 277fe5f953..d61186ca58 100644 --- a/agent/ssl_verify.py +++ b/agent/ssl_verify.py @@ -5,6 +5,7 @@ from __future__ import annotations import logging import os import ssl +import threading from pathlib import Path from typing import Any, Optional @@ -12,6 +13,24 @@ logger = logging.getLogger(__name__) _CA_BUNDLE_ENV_VARS = ("HERMES_CA_BUNDLE", "SSL_CERT_FILE", "REQUESTS_CA_BUNDLE", "CURL_CA_BUNDLE") _INSECURE_STRINGS = {"false", "0", "no", "off"} +_CA_CONTEXTS: dict[str, ssl.SSLContext] = {} +_CA_CONTEXTS_LOCK = threading.Lock() + + +def _context_for_ca_bundle(ca_path: str) -> ssl.SSLContext: + """One ``SSLContext`` per CA bundle path, process-wide. + + ``ssl.create_default_context(cafile=...)`` parses the whole bundle each call, and httpx + transport sharing keys on context identity — so a per-agent context cost one parsed bundle + AND one private connection pool per agent (and per delegated child). An ``SSLContext`` is + safe to share across connections. + """ + with _CA_CONTEXTS_LOCK: + ctx = _CA_CONTEXTS.get(ca_path) + if ctx is None: + ctx = ssl.create_default_context(cafile=ca_path) + _CA_CONTEXTS[ca_path] = ctx + return ctx def resolve_httpx_verify(*, ca_bundle: Optional[str] = None, ssl_verify: Any = None, base_url: str = "") -> bool | ssl.SSLContext: @@ -32,6 +51,6 @@ def resolve_httpx_verify(*, ca_bundle: Optional[str] = None, ssl_verify: Any = N if effective_ca: ca_path = str(Path(effective_ca).expanduser()) if os.path.isfile(ca_path): - return ssl.create_default_context(cafile=ca_path) + return _context_for_ca_bundle(ca_path) logger.warning("CA bundle path does not exist: %s — falling back to default certificates", effective_ca) return True diff --git a/agent/stream_delivery.py b/agent/stream_delivery.py index 4c2e02ec81..399e99a8ae 100644 --- a/agent/stream_delivery.py +++ b/agent/stream_delivery.py @@ -65,10 +65,25 @@ class StreamDeliveryMixin: deliver(ctx_scrubber.flush()) self._current_streamed_assistant_text = "" + @property + def _current_streamed_assistant_text(self) -> str: + """Visible assistant text streamed so far this turn. Backed by a list of pieces: ``+=`` on a string + attribute copies the whole reply on every delta (quadratic). Hot-path emptiness checks look at + ``_streamed_assistant_text_parts`` so they do not join per token.""" + parts = getattr(self, "_streamed_assistant_text_parts", None) + return "".join(parts) if parts else "" + + @_current_streamed_assistant_text.setter + def _current_streamed_assistant_text(self, value: str) -> None: + self._streamed_assistant_text_parts = [value] if value else [] + def _record_streamed_assistant_text(self, text: str) -> None: """Accumulate visible assistant text emitted through stream callbacks (superseded writers excluded).""" if isinstance(text, str) and text and not self._stream_writer_superseded(): - self._current_streamed_assistant_text = getattr(self, "_current_streamed_assistant_text", "") + text + parts = getattr(self, "_streamed_assistant_text_parts", None) + if parts is None: + parts = self._streamed_assistant_text_parts = [] + parts.append(text) @staticmethod def _normalize_interim_visible_text(text: str) -> str: @@ -263,7 +278,8 @@ class StreamDeliveryMixin: text = think_scrubber.feed(text) if think_scrubber is not None else self._strip_think_blocks(text) text = scrubber.feed(text) if scrubber is not None else sanitize_context(text) # Only strip leading newlines on the first delta — mid-stream "\n" is legitimate markdown. - if not prepended_break and not getattr(self, "_current_streamed_assistant_text", ""): + # Check the parts list, not the joined property (joining per token copies the whole reply). + if not prepended_break and not getattr(self, "_streamed_assistant_text_parts", None): text = text.lstrip("\n") if not text: return diff --git a/agent/subagent_lifecycle.py b/agent/subagent_lifecycle.py index a7f333137f..9bf556943d 100644 --- a/agent/subagent_lifecycle.py +++ b/agent/subagent_lifecycle.py @@ -14,6 +14,7 @@ import secrets import threading import time import contextlib +import weakref from contextlib import contextmanager from concurrent.futures import Future, TimeoutError from typing import Any, Callable, Mapping, Optional @@ -162,8 +163,19 @@ _ACTIVE_PARENT_AGENT: contextvars.ContextVar[Any] = contextvars.ContextVar("herm @contextmanager def bind_subagent_parent(parent_agent: Any): - """Bind the host-owned parent for the current agent turn.""" - token = _ACTIVE_PARENT_AGENT.set(parent_agent) + """Bind the host-owned parent for the current agent turn. + + Stored as a weakref: every asyncio Handle/Future scheduled from the turn + (LSP reader loops, kernel pipes, ...) snapshots the Context, and those + snapshots outlive the turn. A strong ref there pinned finished delegate + children — each of which binds itself here for its own turn — in the + parent process heap for the life of the background loop. + """ + try: + ref = weakref.ref(parent_agent) + except TypeError: + ref = lambda: parent_agent # noqa: E731 — non-weakrefable test doubles + token = _ACTIVE_PARENT_AGENT.set(ref) try: yield finally: @@ -172,7 +184,8 @@ def bind_subagent_parent(parent_agent: Any): def get_active_subagent_parent() -> Any: """Return the parent bound to this execution context, if any.""" - return _ACTIVE_PARENT_AGENT.get() + ref = _ACTIVE_PARENT_AGENT.get() + return ref() if ref is not None else None def _opt_str(value: Any) -> bool: diff --git a/agent/subdirectory_hints.py b/agent/subdirectory_hints.py index a981101ec3..b310783244 100644 --- a/agent/subdirectory_hints.py +++ b/agent/subdirectory_hints.py @@ -11,7 +11,8 @@ import shlex from pathlib import Path from typing import Dict, Any, Optional, Set -from agent.prompt_builder import _scan_context_content +from agent.prompt_builder import _read_text_with_timeout, _scan_context_content +from agent.search_policy import SEARCH_PRUNE_DIR_NAMES logger = logging.getLogger(__name__) @@ -22,13 +23,9 @@ _PATH_ARG_KEYS = {"path", "file_path", "workdir"} _COMMAND_TOOLS = {"terminal"} _MAX_ANCESTOR_WALK = 5 # ancestor levels walked per path — bounds deep-path scans -# Directories that hold *copies* of context files (backups, vendored deps, -# VCS internals, caches), never authoritative project context. -_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", -}) +# Shared with broad recursive search probes so context discovery and search never drift into +# different dependency/cache/build trees (those hold *copies* of context files, never authoritative ones). +_EXCLUDED_DIR_NAMES = SEARCH_PRUNE_DIR_NAMES def _digest(content: str) -> str: @@ -164,7 +161,7 @@ class SubdirectoryHintTracker: except OSError: continue try: - content = hint_path.read_text(encoding="utf-8").strip() + content = (_read_text_with_timeout(hint_path) or "").strip() if not content: continue digest = _digest(content) diff --git a/agent/tool_executor.py b/agent/tool_executor.py index a8f33ba90d..42e05826ef 100644 --- a/agent/tool_executor.py +++ b/agent/tool_executor.py @@ -701,6 +701,7 @@ def _run_agent_tool_execution_middleware( function_args, _hermes_pipeline, session_id=str(getattr(agent, "session_id", "") or ""), + tool_call_id=tool_call_id or None, metadata={ "task_id": effective_task_id or "", "turn_id": getattr(agent, "_current_turn_id", "") or "", diff --git a/agent/turn_context.py b/agent/turn_context.py index 082cd0ca76..6d266ddb9a 100644 --- a/agent/turn_context.py +++ b/agent/turn_context.py @@ -283,18 +283,26 @@ def _should_run_preflight_estimate( def _should_idle_compact( *, enabled: bool, idle_after_seconds: int, idle_gap_seconds: float, tokens: int, - floor_tokens: int, cooldown_active: bool, + floor_tokens: int, cooldown_active: bool, last_compaction_tokens: int = 0, ) -> bool: """Pure predicate: idle compaction fires after a wall-clock gap of ``idle_after_seconds`` (opt-in, <= 0 disables), independent of ``threshold_tokens``; - never at/below ``floor_tokens`` or during a compression-failure cooldown.""" - return bool( - enabled - and idle_after_seconds > 0 - and idle_gap_seconds >= idle_after_seconds - and not cooldown_active - and tokens > floor_tokens - ) + never at/below ``floor_tokens`` or during a compression-failure cooldown. + + ``floor_tokens`` (``threshold_tokens × summary_target_ratio``) is a theoretical target a + real pass routinely misses (system prompt, tool schemas and protected head/tail are + incompressible), so a session compacted to above it would re-summarise on every idle + resume without growing. ``last_compaction_tokens`` — what the previous pass actually + produced (``ContextCompressor.last_compression_rough_tokens``, same rough shape as + ``tokens``) — raises the floor to ``last + floor_tokens`` so the transcript must gain a + floor's worth of NEW content first. ``0`` (nothing compacted yet / counter reset) keeps + the original semantics exactly.""" + if not enabled or idle_after_seconds <= 0 or idle_gap_seconds < idle_after_seconds or cooldown_active: + return False + effective_floor = floor_tokens + if last_compaction_tokens > 0: + effective_floor = max(effective_floor, last_compaction_tokens + floor_tokens) + return tokens > effective_floor @dataclass diff --git a/agent/turn_context_compaction.py b/agent/turn_context_compaction.py index ca9c2d4aa8..14481f11c1 100644 --- a/agent/turn_context_compaction.py +++ b/agent/turn_context_compaction.py @@ -56,16 +56,26 @@ def _reset_retry_state_after_compaction(agent: Any) -> None: agent._mute_post_response = False -def _blocked_compress_reason(compressor: Any, tokens: int) -> Optional[str]: +def _blocked_compress_reason( + compressor: Any, tokens: int, attempts_spent: Optional[int] = None +) -> Optional[str]: """Why an over-threshold request is blocked (``None`` below threshold or when the - engine lacks ``should_compress_info`` / raises).""" + engine lacks ``should_compress_info`` / raises). + + ``attempts_spent``: when given and the engine says compression SHOULD run + (``(True, None)``) yet the caller skipped it, the per-turn attempt budget is + spent — name it ``attempts_exhausted:`` instead of dropping the + ``(True, None)`` on the floor (silent-lockout case, #101889).""" _info = getattr(compressor, "should_compress_info", None) if not callable(_info): return None try: - return _info(tokens)[1] + _should_now, _reason = _info(tokens) except Exception: return None + if attempts_spent is not None and _should_now and not _reason: + return f"attempts_exhausted:{attempts_spent}" + return _reason def _apply_grown_window(agent: Any, compressor: Any, grown: int) -> None: @@ -144,15 +154,22 @@ def _idle_compaction( _idle_cooldown = getattr( _compressor, "get_active_compression_failure_cooldown", lambda: None )() + # What the previous pass actually produced — the honest floor versus the theoretical + # ``_idle_floor``. Type pin: compressor doubles expose truthy non-ints here; only a real + # int may raise the floor, anything else falls back to 0 (original semantics). + _idle_last_compaction = getattr(_compressor, "last_compression_rough_tokens", 0) + if not isinstance(_idle_last_compaction, int) or isinstance(_idle_last_compaction, bool): + _idle_last_compaction = 0 if not _tc._should_idle_compact( enabled=agent.compression_enabled, idle_after_seconds=_idle_after, idle_gap_seconds=_idle_gap, tokens=_idle_tokens, floor_tokens=_idle_floor, - cooldown_active=bool(_idle_cooldown), + cooldown_active=bool(_idle_cooldown), last_compaction_tokens=_idle_last_compaction, ): return logger.info( - "Idle compaction: %ss idle >= %ss, ~%s tokens > %s floor (session %s)", + "Idle compaction: %ss idle >= %ss, ~%s tokens > %s floor (last compaction produced ~%s) (session %s)", int(_idle_gap), _idle_after, f"{_idle_tokens:,}", f"{_idle_floor:,}", + f"{_idle_last_compaction:,}" if _idle_last_compaction > 0 else "n/a", agent.session_id or "none", ) _idle_status = automatic_compaction_status_message( diff --git a/agent/turn_explainers.py b/agent/turn_explainers.py index dae098d0d9..945b7e7660 100644 --- a/agent/turn_explainers.py +++ b/agent/turn_explainers.py @@ -115,8 +115,14 @@ _PERSISTENCE_CAUSE_EXPLANATIONS: Dict[str, str] = { "have been lost on restart). Freeing disk space will " "not help. Recovery options:\n" "1. Run `hermes doctor --fix`\n" - "2. Salvage with: sqlite3 ~/.hermes/state.db \".recover\" " - "(then replace state.db)\n" + "2. Stop the gateway, then recover with:\n" + " hermes sessions recover --source {db_path} --inspect-only\n" + " (if it reports recoverable) hermes sessions recover " + "--source {db_path} --output recovered-state.db\n" + " — recovery snapshots the damaged file first; do NOT " + "run `sqlite3 ... \".recover\"` against the live " + "state.db, a vulnerable sqlite3 CLI can corrupt it " + "further\n" "3. Restore from a backup in ~/.hermes/backups/\n" "Then send your message again." ), @@ -287,4 +293,9 @@ class TurnExplainersMixin: body = _PERSISTENCE_CAUSE_EXPLANATIONS.get( persistence_cause or "unknown", _PERSISTENCE_DEFAULT_EXPLANATION ) + if persistence_cause == "corrupt": + # Copy-pasteable, so name the real store (profiles / HERMES_HOME do not live under ~/.hermes). + from hermes_state import _default_db_path + + body = body.replace("{db_path}", str(_default_db_path())) return _NO_REPLY + body if body else "" diff --git a/agent/turn_facade_lease.py b/agent/turn_facade_lease.py index 7f66863151..4987b1dd76 100644 --- a/agent/turn_facade_lease.py +++ b/agent/turn_facade_lease.py @@ -2,8 +2,9 @@ One process at a time may load -> run -> flush a session shared through state.db (Desktop, CLI resume, gateway, background delivery). ``admit_durable_turn_lease`` acquires the row lease (or -returns the early result the façade must hand back); ``DurableTurnLease`` owns the refresher -daemon thread, the turn-liveness watchdog wiring, and the lease-loss / stall interrupt plumbing. +returns the early result the façade must hand back); ``DurableTurnLease`` owns the periodic +refresher, the turn-liveness watchdog wiring, and the lease-loss / stall interrupt plumbing. Both +timers run on the shared scheduler thread (``agent/periodic_scheduler.py``), not per-turn threads. """ import logging import os @@ -20,7 +21,7 @@ LEASE_WAIT_SECONDS = 1800.0 class DurableTurnLease: - """An admitted session turn lease plus the threads that keep it alive and watch the turn. + """An admitted session turn lease plus the periodic timers that keep it alive and watch the turn. ``stop`` is shared by the refresher and the liveness watchdog; ``turn_active`` gates every interrupt so a late refresher miss can never hard-interrupt the NEXT turn. Both are read and @@ -37,18 +38,15 @@ class DurableTurnLease: self._lock = threading.Lock() self.turn_active = False self.interrupt_message: Optional[str] = None - self.refresh_thread: Optional[threading.Thread] = None - self.liveness_thread: Optional[threading.Thread] = None + self.watchdog = None # TurnLivenessWatchdog when configured + self.timer_handles: list = [] # periodic_scheduler handles, cancelled in join_threads def _current_session_id(self) -> str: return getattr(self.agent, "session_id", None) or self.session_id def build_threads(self) -> None: - """Create (not start) the refresher thread and, when configured, the liveness watchdog: - lease renewal is NOT evidence of progress; a silently stalled turn would renew forever.""" - self.refresh_thread = threading.Thread( - target=self.refresh_loop, name="session-turn-lease-refresh", daemon=True - ) + """Create (not schedule) the liveness watchdog when configured: lease renewal is NOT + evidence of progress; a silently stalled turn would renew forever.""" try: from hermes_cli.config import load_config_readonly @@ -59,13 +57,13 @@ class DurableTurnLease: timeout_s, poll_s = turn_liveness.resolve_turn_liveness_settings(liveness_config) if timeout_s is not None: - self.liveness_thread = turn_liveness.TurnLivenessWatchdog( + self.watchdog = turn_liveness.TurnLivenessWatchdog( self.agent, session_id=self._current_session_id(), timeout_s=timeout_s, poll_s=poll_s, stop_event=self.stop, activity_lock=self.agent._liveness_activity_lock(), is_turn_active=self.is_turn_active, commit_abort=self.commit_liveness_abort, deactivate_turn=self.stop_refresher, - ).make_thread() + ) def start(self) -> None: with self._lock: @@ -73,9 +71,11 @@ class DurableTurnLease: # Stamp the activity clock at turn entry: `_last_activity_ts` persists across turns, so # without this the watchdog would measure idle from the PREVIOUS turn and abort a fresh one. self.agent._touch_activity("starting new turn") - self.refresh_thread.start() - if self.liveness_thread is not None: - self.liveness_thread.start() + from agent.periodic_scheduler import schedule + + self.timer_handles.append(schedule(self.refresh_tick, self.refresh_interval)) + if self.watchdog is not None: + self.timer_handles.append(self.watchdog.schedule()) def stop_refresher(self) -> None: """Stop renewal and deactivate the turn. Also the watchdog's deactivate callback: a wedge the @@ -88,9 +88,10 @@ class DurableTurnLease: deactivate_after_liveness_abort = stop_refresher def join_threads(self, timeout: float = 1.0) -> None: - for thread in (self.refresh_thread, self.liveness_thread): - if thread is not None and thread.is_alive(): - thread.join(timeout=timeout) + """Cancel both timers; ``wait=timeout`` mirrors the old ``thread.join(timeout)`` so an + in-flight tick finishes before ``clear_interrupt`` runs.""" + for handle in self.timer_handles: + handle.cancel(wait=timeout) def release(self) -> None: """Release the row and drop the agent's holder attrs (only if they still name this lease).""" @@ -171,34 +172,36 @@ class DurableTurnLease: if agent._execution_thread_id is not None: _set_interrupt(False, agent._execution_thread_id) - def refresh_loop(self) -> None: - """Renew the lease every ``refresh_interval``; a miss or error interrupts the turn. + def refresh_tick(self): + """One periodic renewal (every ``refresh_interval`` on the shared scheduler); a miss or + error interrupts the turn. Returning False stops the timer. The holder-qualified UPDATE fences a late refresher from a successor lease. The façade's finally sets ``stop`` before releasing, so a holder-fenced miss observed after stop is not a loss.""" - while not self.stop.wait(self.refresh_interval): - try: - if self.db.refresh_session_turn_lease( - self._current_session_id(), self.holder, ttl_seconds=LEASE_TTL_SECONDS - ): - continue - if self.stop.is_set(): - return - logger.error( - "Lost session turn lease while turn is active: %s", self._current_session_id() - ) - self._interrupt_turn("Session turn lease lost; stopping to protect the transcript.") - except Exception: - if self.stop.is_set(): - return - logger.warning( - "Failed to refresh session turn lease: %s", self._current_session_id(), exc_info=True, - ) - self._interrupt_turn( - "Session turn lease could not be refreshed; stopping to protect the transcript." - ) - return + if self.stop.is_set(): + return False + try: + if self.db.refresh_session_turn_lease( + self._current_session_id(), self.holder, ttl_seconds=LEASE_TTL_SECONDS + ): + return None + if self.stop.is_set(): + return False + logger.error( + "Lost session turn lease while turn is active: %s", self._current_session_id() + ) + self._interrupt_turn("Session turn lease lost; stopping to protect the transcript.") + except Exception: + if self.stop.is_set(): + return False + logger.warning( + "Failed to refresh session turn lease: %s", self._current_session_id(), exc_info=True, + ) + self._interrupt_turn( + "Session turn lease could not be refreshed; stopping to protect the transcript." + ) + return False @dataclass diff --git a/agent/turn_liveness.py b/agent/turn_liveness.py index a1f5a552f8..a1eace5761 100644 --- a/agent/turn_liveness.py +++ b/agent/turn_liveness.py @@ -86,7 +86,8 @@ def resolve_turn_liveness_settings( class TurnLivenessWatchdog: - """Sampled-idle watchdog thread bound to one conversation turn. + """Sampled-idle watchdog bound to one conversation turn (polls on the + shared periodic scheduler thread). ``activity_lock`` must be the SAME lock ``AIAgent._touch_activity`` stamps the activity clock with; run_agent owns the lease state and callbacks. @@ -108,28 +109,32 @@ class TurnLivenessWatchdog: self._commit_abort = commit_abort self._deactivate_turn = deactivate_turn - def make_thread(self) -> threading.Thread: - """Build the (not yet started) watcher thread; started at turn entry, after the - turn-active flag and activity clock are stamped.""" - return threading.Thread(target=self._watch, name="turn-liveness-watchdog", daemon=True) + def schedule(self): + """Start polling on the shared periodic scheduler thread; returns the cancel handle. + Scheduled at turn entry, after the turn-active flag and activity clock are stamped.""" + from agent.periodic_scheduler import schedule - def _watch(self) -> None: - while not self._stop_event.wait(self._poll_s): - snapshot = self._sample() - if snapshot is None: - return # turn no longer active - if snapshot.idle_seconds < self._timeout_s: - continue - # Observational only: the commit below can still veto the abort if progress - # resumed; the definitive settlement is _surface_committed_abort. - self._surface_stall(snapshot) - message = f"Turn made no progress for {int(snapshot.idle_seconds)}s; aborting to release the session." - if not self._commit_abort(snapshot, message): - continue - # Stop renewing the lease so a wedge the interrupt cannot unwind expires via TTL. - self._deactivate_turn() - self._surface_committed_abort(snapshot) - return + return schedule(self._tick, self._poll_s) + + def _tick(self): + """One poll. Returns False when the watchdog is finished.""" + if self._stop_event.is_set(): + return False + snapshot = self._sample() + if snapshot is None: + return False # turn no longer active + if snapshot.idle_seconds < self._timeout_s: + return None + # Observational only: the commit below can still veto the abort if progress + # resumed; the definitive settlement is _surface_committed_abort. + self._surface_stall(snapshot) + message = f"Turn made no progress for {int(snapshot.idle_seconds)}s; aborting to release the session." + if not self._commit_abort(snapshot, message): + return None + # Stop renewing the lease so a wedge the interrupt cannot unwind expires via TTL. + self._deactivate_turn() + self._surface_committed_abort(snapshot) + return False def _sample(self) -> Optional[ActivitySnapshot]: with self._activity_lock: diff --git a/agent/turn_preflight.py b/agent/turn_preflight.py index a1096f788d..a7801786a2 100644 --- a/agent/turn_preflight.py +++ b/agent/turn_preflight.py @@ -315,8 +315,12 @@ def compress_after_tool_results( return _verdict(True) elif agent.compression_enabled: # Over threshold but compression blocked (cooldown/anti-thrash): deduped - # warning so context can't silently overflow. - _block_reason = _blocked_compress_reason(_compressor, _real_tokens) + # warning so context can't silently overflow. ``attempts_spent`` names the + # attempts_exhausted lockout when the engine says RUN but the per-turn + # budget is spent (#101889). + _block_reason = _blocked_compress_reason( + _compressor, _real_tokens, attempts_spent=compression_attempts + ) if _block_reason: agent._warn_context_overflow_blocked( _block_reason, _real_tokens, int(getattr(_compressor, "threshold_tokens", 0) or 0) diff --git a/agent/turn_request_assembly.py b/agent/turn_request_assembly.py index cb5490e368..09bf5a5728 100644 --- a/agent/turn_request_assembly.py +++ b/agent/turn_request_assembly.py @@ -113,7 +113,7 @@ def assemble_api_request( user merge and surrogate stripping, so the same row's bytes never vary across turns.""" from agent.conversation_loop import ( _apply_context_engine_selection, _canonicalize_api_tool_calls, _clone_message_for_send, - _midturn_request_pressure_tokens, estimate_messages_tokens_rough, + _midturn_request_pressure_tokens, _pressure_with_real_floor, estimate_messages_tokens_rough, ) api_messages, effective_system = build_api_messages( @@ -145,6 +145,13 @@ def assemble_api_request( # Runs unconditionally (not gated on context_compressor) so orphaned tool # results from session loading or manual message edits are always caught. api_messages = agent._sanitize_api_messages(api_messages) + # Send-path vision eviction (#89296): compression only strips stale screenshots + # when prune fires, and the Anthropic adapter's keep-window never sees + # OpenAI-style tool-result image_url parts. The per-call clone is rewritten in + # place; persisted history is untouched. + from agent.context_compressor import evict_stale_outbound_tool_images + + evict_stale_outbound_tool_images(api_messages) # One-time repeated-heal notice goes out via the status/warning callback, NEVER # appended to messages: the cached prompt prefix stays byte-identical. @@ -235,6 +242,13 @@ def assemble_api_request( _anchored_pressure = anchored_context_tokens(messages, getattr(agent, "_usage_anchor", None)) if _anchored_pressure is not None: request_pressure_tokens = _anchored_pressure + else: + # Rough fallback only: floor at the provider's last REAL prompt size (an anchored + # figure is provider-exact and is never floored — on MoA turns that would re-add + # the fan-out tokens the anchor excludes). + request_pressure_tokens = _pressure_with_real_floor( + agent.context_compressor, request_pressure_tokens + ) # Stash the rough estimate so update_from_response() can pair it with the real # count (should_defer_preflight_to_real_usage). getattr: test doubles lack it. _note_rough = getattr(agent.context_compressor, "note_request_rough_estimate", None) diff --git a/agent/usage_pricing.py b/agent/usage_pricing.py index d40b269ccf..a0cbb0137e 100644 --- a/agent/usage_pricing.py +++ b/agent/usage_pricing.py @@ -422,6 +422,10 @@ def get_pricing_entry( return _INCLUDED_ENTRY if route.provider == "openrouter": return _openrouter_pricing_entry(route) + + bundled_entry = _lookup_official_docs_pricing(route) + if bundled_entry: + return bundled_entry if route.base_url: entry = _pricing_entry_from_metadata( fetch_endpoint_model_metadata(route.base_url, api_key=api_key or ""), route.model, @@ -430,7 +434,7 @@ def get_pricing_entry( ) if entry: return entry - return _lookup_official_docs_pricing(route) + return None # Usage-field candidate paths per API shape: (input/prompt total, output, cache diff --git a/apps/desktop/electron/backend-claim.ts b/apps/desktop/electron/backend-claim.ts index f00a4d5bc9..16793d33ae 100644 --- a/apps/desktop/electron/backend-claim.ts +++ b/apps/desktop/electron/backend-claim.ts @@ -36,13 +36,26 @@ export function execText(command: string, args: string[], { timeout = 3000 } = { }) } +/** + * Probe budget for the ORPHAN-REAP path (matchesParent / matchesIdentity / + * stopOwnedBackend). The claim path keeps the full 30s headroom — a freshly + * spawned backend's marker is load-bearing and a slow probe must not kill a + * healthy child (#93608). Reap only needs to tell "same process" from "gone + * or reused" for OLD records, and the ownership file can accumulate dozens of + * them (one per profile per launch), so a 30s budget per record would let a + * cold PowerShell 5.1 stall boot for minutes (#87169). 5s is plenty for a + * warm probe; a timeout degrades to "unknown" and the record is preserved for + * the next launch instead of blocking boot. + */ +export const REAP_PROBE_TIMEOUT_MS = 5_000 + /** * Cross-platform process start marker: a value that changes when a PID is * reused, so `pid + marker` identifies one specific process incarnation. * Throws when the probe fails — callers decide what a failure means (see * `claimDecision` / `probeStartMarker`). */ -export async function processStartMarker(pid: number): Promise { +export async function processStartMarker(pid: number, timeoutMs: number = 30_000): Promise { // Cheap native dead-PID gate. Windows Get-Process / macOS `ps -p` exit 1 // on a missing PID (not ESRCH), so the identity matchers used to keep the // orphan and re-probe it every launch (#92875). ESRCH is the code those @@ -85,7 +98,9 @@ export async function processStartMarker(pid: number): Promise { ], // PowerShell 5.1 cold starts routinely exceed the default 3s execText // budget (2.4-8s observed in #87169); give the marker probe headroom. - { timeout: 30_000 } + // The claim path keeps this 30s budget; the orphan-reap path passes + // REAP_PROBE_TIMEOUT_MS so a slow probe cannot stall boot. + { timeout: timeoutMs } ) if (!/^\d+$/.test(ticks)) { diff --git a/apps/desktop/electron/backend-ownership.test.ts b/apps/desktop/electron/backend-ownership.test.ts index 7910478a67..77a6e46e2e 100644 --- a/apps/desktop/electron/backend-ownership.test.ts +++ b/apps/desktop/electron/backend-ownership.test.ts @@ -178,6 +178,51 @@ test('startup reap preserves failed stops for the next launch', async () => { assert.deepEqual(parseBackendOwnership(store.value()), [entry]) }) +test('startup reap stops at the deadline and preserves the unprocessed records', async () => { + const first = ownershipEntry({ pid: 60 }) + const second = ownershipEntry({ pid: 61 }) + const store = memoryStore(stored([first, second])) + const stop = vi.fn() + + const ownership = createOwnership(store, { + // Each probe is slow enough to blow a 1ms budget after the first entry. + matchesIdentity: async () => { + await new Promise(resolve => setTimeout(resolve, 10)) + + return false + }, + stop, + reapDeadlineMs: 1 + }) + + assert.deepEqual(await ownership.reapOrphans(), []) + // The first entry was processed (dropped); the second was preserved for the + // next launch instead of stalling boot on a slow identity probe. + assert.deepEqual(parseBackendOwnership(store.value()), [second]) +}) + +test('startup reap preserves would-be-reaped records when the budget runs out', async () => { + const first = ownershipEntry({ pid: 62 }) + const second = ownershipEntry({ pid: 63 }) + const store = memoryStore(stored([first, second])) + const stop = vi.fn() + + const ownership = createOwnership(store, { + matchesIdentity: async () => { + await new Promise(resolve => setTimeout(resolve, 10)) + + return true + }, + stop, + reapDeadlineMs: 1 + }) + + assert.deepEqual(await ownership.reapOrphans(), [62]) + // The second would have been reaped too, but the budget ran out — it is + // preserved so a later launch retries it. + assert.deepEqual(parseBackendOwnership(store.value()), [second]) +}) + test('startup reap never stops a backend whose parent Electron is still alive', async () => { const entry = { ...ownershipEntry({ pid: 54 }), parentPid: 100, parentStartMarker: 'os-start-parent' } const store = memoryStore(stored([entry])) diff --git a/apps/desktop/electron/backend-ownership.ts b/apps/desktop/electron/backend-ownership.ts index 682dbe3f5b..3ef1080f0e 100644 --- a/apps/desktop/electron/backend-ownership.ts +++ b/apps/desktop/electron/backend-ownership.ts @@ -28,8 +28,22 @@ export interface BackendOwnershipDeps { matchesParent: (entry: BackendOwnershipEntry) => Promise stop: (identity: BackendIdentity) => Promise | void store: BackendOwnershipStore + /** + * Overall time budget for one reap sweep. The ownership file legitimately + * accumulates one record per profile per launch, and each record can cost + * up to two identity probes (parent + backend) plus a stop — on Windows + * those shell out to PowerShell, whose 5.1 cold starts are slow (#87169). + * Without a bound, a large roster could stall boot for minutes while the + * renderer's 45s backend-boot budget expires and the user stares at the + * connecting screen. When the budget is exhausted the sweep preserves the + * unprocessed records for the next launch and returns what it reaped. + */ + reapDeadlineMs?: number } +/** Default budget for one reap sweep (see `reapDeadlineMs`). */ +export const REAP_ORPHANS_DEADLINE_MS = 5_000 + export interface BackendClaim extends BackendIdentity { command?: string parentPid?: number @@ -222,8 +236,21 @@ export function createBackendOwnership(deps: BackendOwnershipDeps) { const survivors: BackendOwnershipEntry[] = [] const reaped: number[] = [] + const deadline = Date.now() + (deps.reapDeadlineMs ?? REAP_ORPHANS_DEADLINE_MS) + + for (let i = 0; i < entries.length; i += 1) { + // Budget exhausted: preserve the unprocessed records so a later launch + // can retry them. A slow identity probe must never stall boot — the + // renderer's backend-boot budget is 45s and the spawn itself needs + // most of it. + if (Date.now() >= deadline) { + survivors.push(...entries.slice(i)) + + break + } + + const entry = entries[i] - for (const entry of entries) { // A backend whose Electron parent is still running is NOT an orphan: // reaping it would kill a live instance's session. This is what stops // a second launch from SIGTERMing the running instance's backend even diff --git a/apps/desktop/electron/main.ts b/apps/desktop/electron/main.ts index 8a6a53e65a..0e5ec73363 100644 --- a/apps/desktop/electron/main.ts +++ b/apps/desktop/electron/main.ts @@ -42,7 +42,8 @@ import { isPidOnlyStartMarker, pidOnlyStartMarker, probeStartMarker, - processStartMarker + processStartMarker, + REAP_PROBE_TIMEOUT_MS } from './backend-claim' import { dashboardFallbackArgs, sourceDeclaresServe } from './backend-command' import { createBackendConnectionState } from './backend-connection-state' @@ -280,6 +281,12 @@ import { undialedSshRouteSeeds } from './plugin-profile-routes' import { selectPoolEvictions } from './pool-eviction' +import { clampPoolLimits, parsePoolLimits, POOL_LIMITS_DEFAULTS } from './pool-limits' +import { + LocalBackendSpawnCoordinator, + type LocalBackendSpawnRequest, + releaseLocalBackendSlotAfterExit +} from './pool-spawn-coordinator' import { createPoolStopper } from './pool-stop' import { poolTouchKeys } from './pool-touch-scope' import { createKeepAwake } from './power-save' @@ -1408,8 +1415,97 @@ const profileDeletionGate = new ProfileDeletionGate() // Keep the pool light: cap concurrent profile backends (LRU eviction) and reap // idle ones. A user idles at exactly the primary backend; pool backends only // exist while a non-primary profile is actively being chatted through. -const POOL_MAX_BACKENDS = Math.max(1, Number(process.env.HERMES_DESKTOP_POOL_MAX) || 3) -const POOL_IDLE_MS = Math.max(60_000, Number(process.env.HERMES_DESKTOP_POOL_IDLE_MS) || 10 * 60_000) +// Pool sizing is a device preference (Settings → Advanced → pool rows), not a +// launch constant: mutable at runtime, persisted in userData, applied live. +// The legacy HERMES_DESKTOP_POOL_* env vars remain the initial-value fallback +// for scripted/headless setups; after launch the stored preference wins. +const POOL_LIMITS_PATH = path.join(app.getPath('userData'), 'pool-limits.json') + +function readPersistedPoolLimits() { + try { + const limits = parsePoolLimits(fs.readFileSync(POOL_LIMITS_PATH, 'utf8')) + rememberLog( + `[pool-limits] loaded from ${POOL_LIMITS_PATH}: maxBackends=${limits.maxBackends}, idleMs=${limits.idleMs}` + ) + + return limits + } catch { + // No persisted file yet — fall back to the legacy env vars so scripted + // setups keep working. Log which source won: a silently-ignored env var + // here costs a scripted-setup user a debugging session. + const fromEnv = clampPoolLimits({ + maxBackends: Number(process.env.HERMES_DESKTOP_POOL_MAX) || undefined, + idleMs: Number(process.env.HERMES_DESKTOP_POOL_IDLE_MS) || undefined + }) + + if (fromEnv.maxBackends !== POOL_LIMITS_DEFAULTS.maxBackends || fromEnv.idleMs !== POOL_LIMITS_DEFAULTS.idleMs) { + rememberLog( + `[pool-limits] no saved file; using env-var overrides: maxBackends=${fromEnv.maxBackends}, idleMs=${fromEnv.idleMs}` + ) + } else { + rememberLog('[pool-limits] no saved file and no env overrides; using defaults') + } + + return fromEnv + } +} + +function persistPoolLimits(limits) { + try { + fs.mkdirSync(path.dirname(POOL_LIMITS_PATH), { recursive: true }) + // Atomic write: write to a temp file in the same directory, then rename. + // A crash mid-write would otherwise leave truncated JSON and silently + // lose the user's saved sizing. + const tmpPath = `${POOL_LIMITS_PATH}.tmp` + fs.writeFileSync(tmpPath, JSON.stringify(limits, null, 2), 'utf8') + fs.renameSync(tmpPath, POOL_LIMITS_PATH) + } catch (error) { + rememberLog(`[pool-limits] write failed: ${error.message}`) + } +} + +// rememberLog() state. Declared here, before the top-level +// readPersistedPoolLimits() call below, because that call logs during module +// evaluation; declaring these later crashed launch with `undefined.push` in +// the packaged build (esbuild lowers the TDZ to undefined instead of throwing). +const hermesLog = [] +let desktopLogBuffer = '' +let desktopLogFlushTimer = null +let desktopLogFlushPromise = Promise.resolve() + +let poolLimits = readPersistedPoolLimits() +// Hard cap on local backends that are starting OR running (the LRU eviction +// above is soft — it spares keepalive-fresh entries). Follows the live +// preference: setPoolLimits() pushes a new max into the coordinator. +const localBackendSpawnCoordinator = new LocalBackendSpawnCoordinator(poolLimits.maxBackends) +// How long a spawn may wait for a free local slot. Must stay under the +// renderer's BACKEND_BOOT_WAIT_TIMEOUT_MS (45s, src/lib/with-timeout.ts) so +// the queued ticket fails before the renderer does and the user sees why. +const POOL_SLOT_WAIT_MS = 30_000 + +function poolMaxBackends() { + return poolLimits.maxBackends +} + +function poolIdleMs() { + return poolLimits.idleMs +} + +/** + * Apply new limits live: persist, then converge the running pool — evict + * LRU backends down to the new max, and let the (already running) idle + * reaper handle a shortened idle window on its next tick. Returns the + * limits actually in force (post-clamp). + */ +function setPoolLimits(raw) { + poolLimits = clampPoolLimits(raw) + persistPoolLimits(poolLimits) + localBackendSpawnCoordinator.setLimit(poolLimits.maxBackends) + evictLruPoolBackends(poolMaxBackends()) + startPoolIdleReaper() + + return { ...poolLimits } +} // A backend touched within this window has a live renderer socket (the keepalive // pings every 60s for every open profile). LRU eviction must spare these — a @@ -1429,7 +1525,7 @@ const POOL_IDLE_MS = Math.max(60_000, Number(process.env.HERMES_DESKTOP_POOL_IDL // re-allocating pooled gateway secondaries ~700×/day). // * 3× ping + 60s headroom = ~4 min, comfortable margin for two missed // pings + WSL2 IPC stall. The hard ceiling for the cap-eligible set is -// POOL_IDLE_MS above (default 10 min) — this constant only governs the +// pool idle window above (default 10 min) — this constant only governs the // "is this backend plausibly still alive" question for LRU eviction, // not when the idle reaper definitively tears a backend down. const POOL_KEEPALIVE_FRESH_MS = Math.max( @@ -1487,12 +1583,8 @@ let connectionRegistryCache = null let connectionRegistryCacheMtime = null let remoteHeaderRulesInstalled = false const remoteWsHeaderStore = createRemoteWsHeaderStore() -const hermesLog = [] const previewWatchers = new Map() let previewShortcutActive = false -let desktopLogBuffer = '' -let desktopLogFlushTimer = null -let desktopLogFlushPromise = Promise.resolve() let nativeThemeListenerInstalled = false let bootProgressState = { @@ -3401,7 +3493,7 @@ async function backendCommandForPid(pid) { } } -async function processIdentityMatches(identity) { +async function processIdentityMatches(identity, timeoutMs: number = 30_000) { // Degraded PID-only identity (#93608): the start-marker probe failed while // the child was verifiably alive, so only PID liveness can be checked here. // backendIdentityMatches layers the command-line check on top before @@ -3419,14 +3511,14 @@ async function processIdentityMatches(identity) { } try { - return (await processStartMarker(identity.pid)) === identity.startMarker + return (await processStartMarker(identity.pid, timeoutMs)) === identity.startMarker } catch (error) { return error?.code === 'ENOENT' || error?.code === 'ESRCH' ? false : undefined } } async function backendIdentityMatches(identity) { - const processMatches = await processIdentityMatches(identity) + const processMatches = await processIdentityMatches(identity, REAP_PROBE_TIMEOUT_MS) if (processMatches !== true) { return processMatches @@ -3447,17 +3539,26 @@ async function backendParentMatches(entry) { } try { - return (await processStartMarker(entry.parentPid)) === entry.parentStartMarker + return (await processStartMarker(entry.parentPid, REAP_PROBE_TIMEOUT_MS)) === entry.parentStartMarker } catch (error) { return error?.code === 'ENOENT' || error?.code === 'ESRCH' ? false : undefined } } async function stopOwnedBackend(identity) { - if ((await processIdentityMatches(identity)) !== true) { + const matches = await processIdentityMatches(identity, REAP_PROBE_TIMEOUT_MS) + + if (matches === false) { return } + if (matches !== true) { + // Identity probe failed (not confirmed gone): preserve the record so a + // later launch retries the stop instead of dropping it and leaking the + // backend. reapOrphans keeps the entry when stop() throws. + throw new Error(`Could not verify backend PID ${identity.pid} before stopping it.`) + } + if (IS_WINDOWS) { forceKillProcessTree(identity.pid) } else { @@ -3474,7 +3575,7 @@ async function stopOwnedBackend(identity) { const deadline = Date.now() + 1500 while (Date.now() < deadline) { - if ((await processIdentityMatches(identity)) !== true) { + if ((await processIdentityMatches(identity, REAP_PROBE_TIMEOUT_MS)) !== true) { return } @@ -3483,7 +3584,7 @@ async function stopOwnedBackend(identity) { // Revalidate immediately before escalation so PID reuse cannot target a // replacement process. - if ((await processIdentityMatches(identity)) === true) { + if ((await processIdentityMatches(identity, REAP_PROBE_TIMEOUT_MS)) === true) { try { process.kill(-identity.pid, 'SIGKILL') } catch { @@ -3493,7 +3594,7 @@ async function stopOwnedBackend(identity) { } await new Promise(resolve => setTimeout(resolve, 50)) - const remaining = await processIdentityMatches(identity) + const remaining = await processIdentityMatches(identity, REAP_PROBE_TIMEOUT_MS) if (remaining !== false) { throw new Error(`Backend PID ${identity.pid} did not stop cleanly.`) @@ -11251,7 +11352,7 @@ async function ensureBackend(profile) { return connection } - evictLruPoolBackends(POOL_MAX_BACKENDS - 1) + evictLruPoolBackends(poolMaxBackends() - 1) const entry = { process: null, @@ -11259,7 +11360,10 @@ async function ensureBackend(profile) { token: null, connectionPromise: null, lastActiveAt: Date.now(), - remoteBaseUrl: null + remoteBaseUrl: null, + releaseLocalBackendSlot: null, + localBackendSlotKey: null, + localBackendSpawnRequest: null } entry.connectionPromise = spawnPoolBackend(key, entry).catch(async error => { @@ -11270,12 +11374,7 @@ async function ensureBackend(profile) { `Hermes backend for profile "${key}" failed to start: ${error instanceof Error ? error.message : String(error)}` ) - if (backendPool.get(key) === entry) { - backendPool.delete(key) - } - - stopBackendChild(entry.process) - await waitForBackendExit(entry.process) + await teardownFailedLocalBackend(key, entry) throw error }) backendPool.set(key, entry) @@ -11419,7 +11518,7 @@ async function ensureRegistryBackend(connectionId, profile, managedUpdateCorrela return existingLocal.connectionPromise } - evictLruPoolBackends(POOL_MAX_BACKENDS - 1) + evictLruPoolBackends(poolMaxBackends() - 1) const localEntry = { process: null, @@ -11427,7 +11526,10 @@ async function ensureRegistryBackend(connectionId, profile, managedUpdateCorrela token: null, connectionPromise: null, lastActiveAt: Date.now(), - remoteBaseUrl: null + remoteBaseUrl: null, + releaseLocalBackendSlot: null, + localBackendSlotKey: null, + localBackendSpawnRequest: null } localEntry.connectionPromise = spawnPoolBackend(profileKey, localEntry, { @@ -11440,12 +11542,7 @@ async function ensureRegistryBackend(connectionId, profile, managedUpdateCorrela `Hermes backend for profile "${profileKey}" (forced-local) failed to start: ${error instanceof Error ? error.message : String(error)}` ) - if (backendPool.get(localRoute.poolKey) === localEntry) { - backendPool.delete(localRoute.poolKey) - } - - stopBackendChild(localEntry.process) - await waitForBackendExit(localEntry.process) + await teardownFailedLocalBackend(localRoute.poolKey, localEntry) throw error }) backendPool.set(localRoute.poolKey, localEntry) @@ -11492,7 +11589,7 @@ async function ensureRegistryBackend(connectionId, profile, managedUpdateCorrela ) } - evictLruPoolBackends(POOL_MAX_BACKENDS - 1) + evictLruPoolBackends(poolMaxBackends() - 1) const entry = { process: null, @@ -12147,7 +12244,7 @@ function evictLruPoolBackends(keep) { const evictions = selectPoolEvictions(backendPool.entries(), Math.max(0, keep), Date.now(), POOL_KEEPALIVE_FRESH_MS) for (const profile of evictions) { - rememberLog(`Evicting idle profile backend "${profile}" (LRU cap ${POOL_MAX_BACKENDS})`) + rememberLog(`Evicting idle profile backend "${profile}" (LRU cap ${poolMaxBackends()})`) stopPoolBackend(profile) } } @@ -12161,8 +12258,8 @@ function startPoolIdleReaper() { const now = Date.now() for (const [profile, entry] of [...backendPool.entries()]) { - if (now - (entry.lastActiveAt || 0) > POOL_IDLE_MS) { - rememberLog(`Reaping idle profile backend "${profile}" (idle > ${Math.round(POOL_IDLE_MS / 1000)}s)`) + if (now - (entry.lastActiveAt || 0) > poolIdleMs()) { + rememberLog(`Reaping idle profile backend "${profile}" (idle > ${Math.round(poolIdleMs() / 1000)}s)`) stopPoolBackend(profile) } } @@ -12178,6 +12275,67 @@ function startPoolIdleReaper() { } } +function releaseLocalBackendSlot(entry: any) { + if (!entry) { + return + } + + const release = entry.releaseLocalBackendSlot + const request = entry.localBackendSpawnRequest as LocalBackendSpawnRequest | null + entry.releaseLocalBackendSlot = null + entry.localBackendSlotKey = null + entry.localBackendSpawnRequest = null + + if (release) { + release() + } else { + request?.cancel() + } +} + +function assertPoolEntryStillOwned(poolKey: string, entry: any) { + if (backendPool.get(poolKey) !== entry) { + releaseLocalBackendSlot(entry) + throw new Error(`Profile backend start for "${poolKey}" was cancelled before spawn.`) + } +} + +const failedLocalBackendTeardowns = new WeakMap>() + +function teardownFailedLocalBackend(poolKey: string, entry: any): Promise { + const existing = failedLocalBackendTeardowns.get(entry) + + if (existing) { + return existing + } + + if (backendPool.get(poolKey) === entry) { + backendPool.delete(poolKey) + } + + const child = entry.process + + const teardown = releaseLocalBackendSlotAfterExit( + () => releaseLocalBackendSlot(entry), + async () => { + stopBackendChild(child) + await waitForBackendExit(child) + + if (child && child.exitCode === null && child.signalCode === null) { + throw new Error(`Profile backend for "${poolKey}" did not exit; keeping the local slot occupied.`) + } + + releaseBackendChild(child) + } + ) + + // Keep the settled promise in the WeakMap for the lifetime of this entry. + // Error + exit + outer catch may all request cleanup; none may run it twice. + failedLocalBackendTeardowns.set(entry, teardown) + + return teardown +} + // Spawn an additional dashboard backend pinned to a named profile. Mirrors the // local-spawn portion of startHermes() but without the boot-progress UI, // bootstrap, or remote handling (those belong to the primary backend only). @@ -12215,6 +12373,29 @@ async function spawnPoolBackend(profile, entry, opts: { forceLocal?: boolean; po } } + // Bound the slot wait BELOW the renderer's backend-boot budget (45s): once + // the renderer has given up on this spawn, a ticket still queued for the + // pool-idle window (10 min) would hold the pool key hostage and every + // later click on the profile would join that stale wait. Failing here + // surfaces the "all N slots busy" reason instead of a generic boot timeout. + const spawnRequest = localBackendSpawnCoordinator.request(poolKey, { timeoutMs: POOL_SLOT_WAIT_MS }) + entry.localBackendSlotKey = poolKey + entry.localBackendSpawnRequest = spawnRequest + + if (localBackendSpawnCoordinator.activeCount >= poolMaxBackends()) { + rememberLog( + `Profile backend "${profile}" waiting for a free local slot (${localBackendSpawnCoordinator.activeCount}/${poolMaxBackends()} busy, ${localBackendSpawnCoordinator.queuedCount} queued)` + ) + } + + entry.releaseLocalBackendSlot = await spawnRequest.acquired + + if (entry.localBackendSpawnRequest === spawnRequest) { + entry.localBackendSpawnRequest = null + } + + assertPoolEntryStillOwned(poolKey, entry) + const token = crypto.randomBytes(32).toString('base64url') // Same update mutual exclusion as the primary window's waitForLocalStart @@ -12258,12 +12439,12 @@ async function spawnPoolBackend(profile, entry, opts: { forceLocal?: boolean; po assertLocalProfileCanStart(profile, profileDeletionGate, key => directoryExists(path.join(HERMES_HOME, 'profiles', key)) ) - rememberLog(`Starting Hermes backend for profile "${profile}" via ${backend.label}`) const parentStartMarker = await desktopParentStartMarker() const backendNonce = crypto.randomBytes(16).toString('hex') const parentIdentityEnv = parentWatchdogEnv(process.pid, parentStartMarker, backendNonce) + assertPoolEntryStillOwned(poolKey, entry) const child = spawn( backend.command, @@ -12318,6 +12499,7 @@ async function spawnPoolBackend(profile, entry, opts: { forceLocal?: boolean; po // surface as an unhandled rejection before the Promise.race below attaches. portAnnouncement.catch(() => {}) await claimBackendChild(child, `${backend.command} ${backend.args.join(' ')}`, profile, backendNonce, outputTail) + assertPoolEntryStillOwned(poolKey, entry) child.stdout.on('data', rememberLog) child.stderr.on('data', rememberLog) @@ -12331,14 +12513,21 @@ async function spawnPoolBackend(profile, entry, opts: { forceLocal?: boolean; po child.once('error', error => { rememberLog(`Hermes backend for profile "${profile}" failed to start: ${error.message}`) - releaseBackendChild(child) - backendPool.delete(poolKey) + void teardownFailedLocalBackend(poolKey, entry).catch(cleanupError => { + rememberLog( + `Hermes backend for profile "${profile}" cleanup failed: ${cleanupError instanceof Error ? cleanupError.message : String(cleanupError)}` + ) + }) rejectStart?.(error) }) child.once('exit', (code, signal) => { rememberLog(`Hermes backend for profile "${profile}" exited (${signal || code})`) + releaseLocalBackendSlot(entry) releaseBackendChild(child) - backendPool.delete(poolKey) + + if (backendPool.get(poolKey) === entry) { + backendPool.delete(poolKey) + } if (!ready) { rejectStart?.( @@ -12406,16 +12595,20 @@ const poolStopper = createPoolStopper({ waitForExit: child => waitForBackendExit(child) }) -function stopPoolBackend(profile) { - return poolStopper.stop(profile) +async function stopPoolBackend(profile: string) { + const entry = backendPool.get(profile) + await poolStopper.stop(profile) + releaseLocalBackendSlot(entry) } async function teardownPoolBackendAndWait(profile) { - await Promise.all(localProfilePoolKeys(profile).map(key => poolStopper.stop(key))) + await Promise.all(localProfilePoolKeys(profile).map(key => stopPoolBackend(key))) } -function stopAllPoolBackends() { - return poolStopper.stopAll() +async function stopAllPoolBackends() { + const entries = [...backendPool.values()] + await poolStopper.stopAll() + entries.forEach(releaseLocalBackendSlot) } const backendShutdown = createBackendShutdownCoordinator(async () => { @@ -14610,6 +14803,18 @@ ipcMain.handle('hermes:backend:touch', async (_event, profile) => { return { ok: true } }) +// Pool sizing (Settings → Advanced): device-local, live-applied. Main is +// authoritative (it owns the pool and the persisted copy); the returned +// limits are what actually took effect post-clamp. +ipcMain.handle('hermes:pool-limits:get', async () => ({ ...poolLimits })) +ipcMain.handle('hermes:pool-limits:set', async (_event, raw) => { + const next = setPoolLimits({ + maxBackends: typeof raw?.maxBackends === 'number' ? raw.maxBackends : poolLimits.maxBackends, + idleMs: typeof raw?.idleMs === 'number' ? raw.idleMs : poolLimits.idleMs + }) + + return { ok: true, limits: next } +}) ipcMain.handle('hermes:gateway:ws-url', async (_event, profile) => { return gatewayWsUrlIpcResult(() => freshGatewayWsUrl(profile)) }) @@ -16817,7 +17022,9 @@ ipcMain.on('hermes:translucency:support', event => { // shortcut edit), and survives self-relaunches because collectRelaunchArgs // only strips internal flags. ipcMain.on('hermes:launch-flags', event => { - event.returnValue = { localModels: process.argv.includes('--local') } + event.returnValue = { + localModels: process.argv.includes('--local') || process.platform === 'win32' || process.platform === 'darwin' + } }) ipcMain.on('hermes:translucency', (_event, payload) => { diff --git a/apps/desktop/electron/pool-limits.test.ts b/apps/desktop/electron/pool-limits.test.ts new file mode 100644 index 0000000000..5aff45f088 --- /dev/null +++ b/apps/desktop/electron/pool-limits.test.ts @@ -0,0 +1,53 @@ +import { describe, expect, it } from 'vitest' + +import { + clampPoolLimits, + parsePoolLimits, + POOL_LIMITS_BOUNDS, + POOL_LIMITS_DEFAULTS, + POOL_LIMITS_MIN +} from './pool-limits' + +describe('parsePoolLimits', () => { + it('falls back to defaults for null/empty/corrupt input', () => { + expect(parsePoolLimits(null)).toEqual(POOL_LIMITS_DEFAULTS) + expect(parsePoolLimits(undefined)).toEqual(POOL_LIMITS_DEFAULTS) + expect(parsePoolLimits('')).toEqual(POOL_LIMITS_DEFAULTS) + expect(parsePoolLimits('not json {')).toEqual(POOL_LIMITS_DEFAULTS) + }) + + it('parses a valid persisted blob', () => { + expect(parsePoolLimits(JSON.stringify({ maxBackends: 11, idleMs: 7_200_000 }))).toEqual({ + maxBackends: 11, + idleMs: 7_200_000 + }) + }) + + it('fills missing keys from defaults', () => { + expect(parsePoolLimits(JSON.stringify({ maxBackends: 5 }))).toEqual({ ...POOL_LIMITS_DEFAULTS, maxBackends: 5 }) + expect(parsePoolLimits('{}')).toEqual(POOL_LIMITS_DEFAULTS) + }) + + it('ignores non-numeric junk instead of NaN-poisoning the pool', () => { + expect(parsePoolLimits(JSON.stringify({ maxBackends: 'lots', idleMs: null }))).toEqual(POOL_LIMITS_DEFAULTS) + }) +}) + +describe('clampPoolLimits', () => { + it('clamps below the floors', () => { + expect(clampPoolLimits({ maxBackends: 0 }).maxBackends).toBe(POOL_LIMITS_MIN.maxBackends) + expect(clampPoolLimits({ idleMs: 100 }).idleMs).toBe(POOL_LIMITS_MIN.idleMs) + }) + + it('clamps absurdly high backend counts', () => { + expect(clampPoolLimits({ maxBackends: 10_000 }).maxBackends).toBeLessThanOrEqual(64) + }) + + it('clamps idleMs to the shared ceiling (7 days)', () => { + expect(clampPoolLimits({ idleMs: 999_000_000 }).idleMs).toBe(POOL_LIMITS_BOUNDS.idleMsMax) + }) + + it('floors fractional values', () => { + expect(clampPoolLimits({ maxBackends: 2.9 }).maxBackends).toBe(2) + }) +}) diff --git a/apps/desktop/electron/pool-limits.ts b/apps/desktop/electron/pool-limits.ts new file mode 100644 index 0000000000..ee7e5b8a8a --- /dev/null +++ b/apps/desktop/electron/pool-limits.ts @@ -0,0 +1,81 @@ +/** + * Pool limits — how many bot backends may stay spawned, and how long an + * unused one survives. + * + * A device-local preference (each machine trades RAM against switching + * speed for itself), stored in userData like keep-awake. The main process + * is authoritative: it owns the pool AND the persisted copy, and applies a + * new max IMMEDIATELY by evicting least-recently-used idle backends — no + * app restart. The renderer mirrors the values for its UI and prewarm + * guard over IPC. + * + * Defaults preserve the historical hard-coded behavior (3 backends, 10min + * idle) so machines that never open Settings behave exactly as before. + */ + +export interface PoolLimits { + /** Max concurrently spawned non-primary profile backends. */ + maxBackends: number + /** Idle lifetime of an unused pool backend, in milliseconds. */ + idleMs: number +} + +export const POOL_LIMITS_DEFAULTS: PoolLimits = { + maxBackends: 3, + idleMs: 10 * 60_000 +} + +/** Hard floors — match the clamps the env-var path always applied. */ +export const POOL_LIMITS_MIN: PoolLimits = { + maxBackends: 1, + idleMs: 60_000 +} + +/** Shared bounds for both pool knobs — imported by the Settings UI so the + * advertised input ranges can never drift from what main actually clamps + * to. idleMs has no ceiling: a user who wants backends kept warm all week + * may have exactly that. */ +export const POOL_LIMITS_BOUNDS = { + maxBackendsMax: 64, + /** 7 days, matching the UI's suggestion ceiling. */ + idleMsMax: 7 * 24 * 60 * 60_000 +} as const + +const MAX_BACKENDS_CEILING = POOL_LIMITS_BOUNDS.maxBackendsMax +const IDLE_MS_CEILING = POOL_LIMITS_BOUNDS.idleMsMax + +/** Clamp a raw partial to the floors/ceilings; missing keys fall to defaults. */ +export function clampPoolLimits(raw: Partial): PoolLimits { + const maxBackends = Number.isFinite(raw.maxBackends) + ? Math.min(MAX_BACKENDS_CEILING, Math.max(POOL_LIMITS_MIN.maxBackends, Math.floor(Number(raw.maxBackends)))) + : POOL_LIMITS_DEFAULTS.maxBackends + + const idleMs = Number.isFinite(raw.idleMs) + ? Math.min(IDLE_MS_CEILING, Math.max(POOL_LIMITS_MIN.idleMs, Math.floor(Number(raw.idleMs)))) + : POOL_LIMITS_DEFAULTS.idleMs + + return { maxBackends, idleMs } +} + +function clampLimits(raw: Partial): PoolLimits { + return clampPoolLimits(raw) +} + +/** Parse + clamp a persisted JSON blob; anything unreadable falls back to + * defaults so a corrupted file can never wedge the pool. */ +export function parsePoolLimits(json: string | null | undefined): PoolLimits { + if (!json) { + return { ...POOL_LIMITS_DEFAULTS } + } + + try { + const parsed = JSON.parse(json) + + return clampLimits({ + maxBackends: typeof parsed?.maxBackends === 'number' ? parsed.maxBackends : undefined, + idleMs: typeof parsed?.idleMs === 'number' ? parsed.idleMs : undefined + }) + } catch { + return { ...POOL_LIMITS_DEFAULTS } + } +} diff --git a/apps/desktop/electron/pool-spawn-coordinator.test.ts b/apps/desktop/electron/pool-spawn-coordinator.test.ts new file mode 100644 index 0000000000..73b2291fe1 --- /dev/null +++ b/apps/desktop/electron/pool-spawn-coordinator.test.ts @@ -0,0 +1,344 @@ +import assert from 'node:assert/strict' +import { spawn } from 'node:child_process' +import fs from 'node:fs' +import path from 'node:path' +import { fileURLToPath } from 'node:url' + +import { test } from 'vitest' + +import { LocalBackendSpawnCoordinator, releaseLocalBackendSlotAfterExit } from './pool-spawn-coordinator' + +const deferred = () => { + let resolve!: () => void + + const promise = new Promise(done => { + resolve = done + }) + + return { promise, resolve } +} + +const flush = () => new Promise(resolve => setImmediate(resolve)) + +test('100 concurrent local requests never hold more than the configured slots', async () => { + const limit = 12 + const coordinator = new LocalBackendSpawnCoordinator(limit) + const gates = Array.from({ length: 100 }, deferred) + let active = 0 + let maxActive = 0 + + const tasks = gates.map(async (gate, index) => { + const release = await coordinator.acquire(`profile-${index}`) + active += 1 + maxActive = Math.max(maxActive, active) + + await gate.promise + + active -= 1 + release() + }) + + await flush() + assert.equal(active, limit) + assert.equal(coordinator.activeCount, limit) + assert.equal(coordinator.queuedCount, 100 - limit) + + for (let start = 0; start < gates.length; start += limit) { + for (const gate of gates.slice(start, start + limit)) { + gate.resolve() + } + + await flush() + } + + await Promise.all(tasks) + assert.equal(maxActive, limit) + assert.equal(coordinator.activeCount, 0) + assert.equal(coordinator.queuedCount, 0) +}) + +test('a queued start can be cancelled without waiting for an active backend', async () => { + const coordinator = new LocalBackendSpawnCoordinator(1) + const releaseFirst = await coordinator.acquire('first') + const queued = coordinator.request('cancelled') + + assert.equal(coordinator.queuedCount, 1) + assert.equal(queued.cancel(), true) + await assert.rejects(queued.acquired, /cancelled while queued/) + assert.equal(coordinator.activeCount, 1) + assert.equal(coordinator.queuedCount, 0) + + releaseFirst() + assert.equal(coordinator.activeCount, 0) +}) + +test('cancelling an old same-key request never rejects a newer waiter', async () => { + const coordinator = new LocalBackendSpawnCoordinator(1) + const blocker = coordinator.request('blocker') + const releaseBlocker = await blocker.acquired + const old = coordinator.request('same-profile') + + releaseBlocker() + const newer = coordinator.request('same-profile') + + assert.equal(old.cancel(), false, 'the old request was already granted') + assert.equal(coordinator.queuedCount, 1, 'the newer same-key waiter must remain queued') + + const releaseOld = await old.acquired + releaseOld() + const releaseNewer = await newer.acquired + releaseNewer() + + assert.equal(coordinator.activeCount, 0) + assert.equal(coordinator.queuedCount, 0) +}) + +test('a queued start times out with a clear error and frees its queue position', async () => { + const coordinator = new LocalBackendSpawnCoordinator(1) + const releaseFirst = await coordinator.acquire('first') + const queued = coordinator.request('timed-out', { timeoutMs: 10 }) + + await assert.rejects(queued.acquired, /timed out while waiting for a free slot/) + assert.equal(coordinator.activeCount, 1) + assert.equal(coordinator.queuedCount, 0) + + releaseFirst() + assert.equal(coordinator.activeCount, 0) +}) + +test('100 real child processes never exceed twelve simultaneous local slots', async () => { + const limit = 12 + const coordinator = new LocalBackendSpawnCoordinator(limit) + const livePids = new Set() + const seenPids = new Set() + let maxLive = 0 + + await Promise.all( + Array.from({ length: 100 }, async (_, index) => { + const release = await coordinator.acquire(`real-profile-${index}`) + + try { + const child = spawn(process.execPath, ['-e', 'setTimeout(() => {}, 40)'], { + stdio: 'ignore' + }) + + assert.ok(child.pid) + livePids.add(child.pid) + seenPids.add(child.pid) + maxLive = Math.max(maxLive, livePids.size) + + await new Promise((resolve, reject) => { + child.once('error', reject) + child.once('exit', code => { + if (code === 0) { + resolve() + } else { + reject(new Error(`child ${child.pid} exited with ${code}`)) + } + }) + }) + + livePids.delete(child.pid) + } finally { + release() + } + }) + ) + + assert.equal(seenPids.size, 100) + assert.equal(maxLive, limit) + assert.equal(livePids.size, 0) + assert.equal(coordinator.activeCount, 0) + assert.equal(coordinator.queuedCount, 0) +}) + +test('failed start keeps its slot until the child has actually exited', async () => { + const coordinator = new LocalBackendSpawnCoordinator(1) + const childExit = deferred() + const releaseFailed = await coordinator.acquire('failed') + let successorEntered = false + + const successor = coordinator.acquire('successor').then(release => { + successorEntered = true + + return release + }) + + const cleanup = releaseLocalBackendSlotAfterExit(releaseFailed, () => childExit.promise) + await flush() + + assert.equal(successorEntered, false) + assert.equal(coordinator.activeCount, 1) + assert.equal(coordinator.queuedCount, 1) + + childExit.resolve() + await cleanup + const releaseSuccessor = await successor + + assert.equal(successorEntered, true) + assert.equal(coordinator.activeCount, 1) + assert.equal(coordinator.queuedCount, 0) + + releaseSuccessor() + assert.equal(coordinator.activeCount, 0) +}) + +test('a rejected wait keeps the slot occupied', async () => { + const coordinator = new LocalBackendSpawnCoordinator(1) + const releaseFailed = await coordinator.acquire('failed') + let successorEntered = false + + const successor = coordinator.acquire('successor').then(release => { + successorEntered = true + + return release + }) + + const cleanup = releaseLocalBackendSlotAfterExit(releaseFailed, async () => { + throw new Error('exit unproven') + }) + + await assert.rejects(cleanup, /exit unproven/) + await flush() + + assert.equal(successorEntered, false) + assert.equal(coordinator.activeCount, 1) + assert.equal(coordinator.queuedCount, 1) + + releaseFailed() + const releaseSuccessor = await successor + assert.equal(successorEntered, true) + releaseSuccessor() + assert.equal(coordinator.activeCount, 0) + assert.equal(coordinator.queuedCount, 0) +}) + +test('an invalid timeout never enqueues a waiter', async () => { + const coordinator = new LocalBackendSpawnCoordinator(1) + const releaseFirst = await coordinator.acquire('first') + + assert.throws(() => coordinator.request('invalid', { timeoutMs: 0 }), /timeout must be a positive number/) + assert.throws(() => coordinator.request('invalid', { timeoutMs: Number.NaN }), /timeout must be a positive number/) + assert.throws(() => coordinator.request('invalid', { timeoutMs: -5 }), /timeout must be a positive number/) + + assert.equal(coordinator.activeCount, 1) + assert.equal(coordinator.queuedCount, 0) + + releaseFirst() + assert.equal(coordinator.activeCount, 0) +}) + +test('a failed or repeated cleanup releases exactly one slot', async () => { + const coordinator = new LocalBackendSpawnCoordinator(1) + const releaseFirst = await coordinator.acquire('first') + let secondEntered = false + + const second = coordinator.acquire('second').then(release => { + secondEntered = true + + return release + }) + + await flush() + assert.equal(secondEntered, false) + assert.equal(coordinator.activeCount, 1) + assert.equal(coordinator.queuedCount, 1) + + releaseFirst() + releaseFirst() + const releaseSecond = await second + + assert.equal(secondEntered, true) + assert.equal(coordinator.activeCount, 1) + assert.equal(coordinator.queuedCount, 0) + + releaseSecond() + assert.equal(coordinator.activeCount, 0) +}) + +test('raising the limit at runtime drains queued waiters into the new slots', async () => { + const coordinator = new LocalBackendSpawnCoordinator(1) + const first = await coordinator.acquire('a') + const queuedB = coordinator.request('b') + const queuedC = coordinator.request('c') + await flush() + assert.equal(coordinator.activeCount, 1) + assert.equal(coordinator.queuedCount, 2) + + coordinator.setLimit(2) + const releaseB = await queuedB.acquired + assert.equal(coordinator.activeCount, 2) + assert.equal(coordinator.queuedCount, 1) + + first() + const releaseC = await queuedC.acquired + assert.equal(coordinator.activeCount, 2) + releaseB() + releaseC() + assert.equal(coordinator.activeCount, 0) +}) + +test('lowering the limit never revokes granted slots; new requests queue until under cap', async () => { + const coordinator = new LocalBackendSpawnCoordinator(3) + const releases = await Promise.all(['a', 'b', 'c'].map(key => coordinator.acquire(key))) + coordinator.setLimit(1) + assert.equal(coordinator.activeCount, 3, 'granted slots stay granted') + + const queued = coordinator.request('d') + await flush() + assert.equal(coordinator.queuedCount, 1) + + releases[0]() + releases[1]() + await flush() + assert.equal(coordinator.queuedCount, 1, 'still over the new cap of 1') + + releases[2]() + const releaseD = await queued.acquired + assert.equal(coordinator.activeCount, 1) + releaseD() +}) + +test('setLimit rejects a non-positive or fractional cap', () => { + const coordinator = new LocalBackendSpawnCoordinator(2) + assert.throws(() => coordinator.setLimit(0), RangeError) + assert.throws(() => coordinator.setLimit(1.5), RangeError) + assert.equal(coordinator.limit, 2) +}) + +// ── main.ts wiring ────────────────────────────────────────────────────────── +// The coordinator is only as good as the timeout main.ts hands it. A queued +// ticket that outlives the renderer's backend-boot budget holds the pool key +// hostage: the renderer has already reported "backend didn't come up", and +// every later click on that profile joins the stale wait instead of failing +// fast with a reason. +{ + const here = path.dirname(fileURLToPath(import.meta.url)) + const mainSource = fs.readFileSync(path.join(here, 'main.ts'), 'utf8').replace(/\r\n/g, '\n') + + const withTimeoutSource = fs + .readFileSync(path.join(here, '..', 'src', 'lib', 'with-timeout.ts'), 'utf8') + .replace(/\r\n/g, '\n') + + test('main.ts bounds the slot wait below the renderer backend-boot budget', () => { + const slotWait = Number(/const POOL_SLOT_WAIT_MS = ([\d_]+)/.exec(mainSource)?.[1]?.replace(/_/g, '')) + + const bootBudget = Number( + /export const BACKEND_BOOT_WAIT_TIMEOUT_MS = ([\d_]+)/.exec(withTimeoutSource)?.[1]?.replace(/_/g, '') + ) + + assert.ok(Number.isFinite(slotWait) && slotWait > 0, 'POOL_SLOT_WAIT_MS must be a literal in main.ts') + assert.ok(Number.isFinite(bootBudget), 'BACKEND_BOOT_WAIT_TIMEOUT_MS must be a literal') + assert.ok(slotWait < bootBudget, `slot wait ${slotWait}ms must be below the boot budget ${bootBudget}ms`) + assert.match(mainSource, /localBackendSpawnCoordinator\.request\(poolKey, \{ timeoutMs: POOL_SLOT_WAIT_MS \}\)/) + assert.doesNotMatch(mainSource, /request\(poolKey, \{ timeoutMs: POOL_IDLE_MS \}\)/) + }) + + test('main.ts pushes the live pool max into the coordinator when the preference changes', () => { + // Pool sizing is a live device preference (#92581); the hard cap must + // follow it, otherwise raising the max in Settings would leave spawns + // queued behind the launch-time value. + assert.match(mainSource, /new LocalBackendSpawnCoordinator\(poolLimits\.maxBackends\)/) + assert.match(mainSource, /localBackendSpawnCoordinator\.setLimit\(poolLimits\.maxBackends\)/) + }) +} diff --git a/apps/desktop/electron/pool-spawn-coordinator.ts b/apps/desktop/electron/pool-spawn-coordinator.ts new file mode 100644 index 0000000000..8e565ed133 --- /dev/null +++ b/apps/desktop/electron/pool-spawn-coordinator.ts @@ -0,0 +1,153 @@ +export type ReleaseLocalBackendSlot = () => void + +export type LocalBackendSpawnRequest = { + acquired: Promise + cancel: () => boolean +} + +type Waiter = { + key: string + resolve: (release: ReleaseLocalBackendSlot) => void + reject: (error: Error) => void + timer: ReturnType | null +} + +export async function releaseLocalBackendSlotAfterExit( + release: ReleaseLocalBackendSlot, + waitForExit: () => Promise +): Promise { + await waitForExit() + release() +} + +/** + * Bounds the number of local profile backends that are starting or running. + * + * A lease is acquired immediately before local start work and is held until + * the child exits or the start fails. Remote descriptors never call request(). + */ +export class LocalBackendSpawnCoordinator { + #limit: number + #active = 0 + #queue: Waiter[] = [] + + constructor(limit: number) { + if (!Number.isInteger(limit) || limit < 1) { + throw new RangeError('Local backend spawn limit must be a positive integer.') + } + + this.#limit = limit + } + + get activeCount(): number { + return this.#active + } + + get limit(): number { + return this.#limit + } + + /** + * Adopt a new cap at runtime (the pool size is a live device preference). + * Raising it drains waiters into the newly freed slots immediately; lowering + * it never revokes a granted slot — the running backends simply stay over + * the cap until they exit, and LRU eviction (main.ts) converges the pool. + */ + setLimit(limit: number): void { + if (!Number.isInteger(limit) || limit < 1) { + throw new RangeError('Local backend spawn limit must be a positive integer.') + } + + this.#limit = limit + this.#drain() + } + + get queuedCount(): number { + return this.#queue.length + } + + request(key: string, options: { timeoutMs?: number } = {}): LocalBackendSpawnRequest { + if (options.timeoutMs !== undefined && (!Number.isFinite(options.timeoutMs) || options.timeoutMs < 1)) { + throw new RangeError('Local backend spawn timeout must be a positive number.') + } + + if (this.#active < this.#limit) { + return { + acquired: Promise.resolve(this.#grant()), + cancel: () => false + } + } + + let waiter!: Waiter + + const acquired = new Promise((resolve, reject) => { + waiter = { key, resolve, reject, timer: null } + this.#queue.push(waiter) + + if (options.timeoutMs !== undefined) { + waiter.timer = setTimeout(() => { + this.#rejectWaiter( + waiter, + new Error(`Local backend start for "${key}" timed out while waiting for a free slot.`) + ) + }, options.timeoutMs) + waiter.timer.unref?.() + } + }) + + return { + acquired, + cancel: () => + this.#rejectWaiter(waiter, new Error(`Local backend start for "${key}" was cancelled while queued.`)) + } + } + + acquire(key: string): Promise { + return this.request(key).acquired + } + + #rejectWaiter(waiter: Waiter, error: Error): boolean { + const index = this.#queue.indexOf(waiter) + + if (index === -1) { + return false + } + + this.#queue.splice(index, 1) + this.#clearTimer(waiter) + waiter.reject(error) + + return true + } + + #clearTimer(waiter: Waiter): void { + if (waiter.timer) { + clearTimeout(waiter.timer) + waiter.timer = null + } + } + + #grant(): ReleaseLocalBackendSlot { + this.#active += 1 + let released = false + + return () => { + if (released) { + return + } + + released = true + this.#active -= 1 + this.#drain() + } + } + + /** Hand free slots to queued waiters while under the (possibly lowered) cap. */ + #drain(): void { + while (this.#active < this.#limit && this.#queue.length > 0) { + const next = this.#queue.shift()! + this.#clearTimer(next) + next.resolve(this.#grant()) + } + } +} diff --git a/apps/desktop/electron/preload.ts b/apps/desktop/electron/preload.ts index 8176ebfc07..fd9668752b 100644 --- a/apps/desktop/electron/preload.ts +++ b/apps/desktop/electron/preload.ts @@ -24,6 +24,8 @@ contextBridge.exposeInMainWorld('hermesDesktop', { getProfileRoutes: profiles => ipcRenderer.invoke('hermes:plugin-profile-routes', profiles), revalidateConnection: () => ipcRenderer.invoke('hermes:connection:revalidate'), touchBackend: profile => ipcRenderer.invoke('hermes:backend:touch', profile), + getPoolLimits: () => ipcRenderer.invoke('hermes:pool-limits:get'), + setPoolLimits: limits => ipcRenderer.invoke('hermes:pool-limits:set', limits), getGatewayWsUrl: profile => ipcRenderer.invoke('hermes:gateway:ws-url', profile), // Registry-scoped fresh WS URL: { connectionId, profile } → result shape of // getGatewayWsUrl, minted against that connection's backend. diff --git a/apps/desktop/src/app/chat/index.tsx b/apps/desktop/src/app/chat/index.tsx index 14651b7f72..1b95bc769a 100644 --- a/apps/desktop/src/app/chat/index.tsx +++ b/apps/desktop/src/app/chat/index.tsx @@ -78,7 +78,7 @@ import { mergeOlderTranscriptPage, transcriptBackfillAvailable } from './transcript-backfill' -import { advanceTranscriptWindow, type TranscriptWindowState } from './transcript-window' +import { advanceSessionTranscriptWindow, type SessionWindowMemo } from './transcript-window' interface ChatViewProps extends Omit, 'onSubmit'> { gateway: HermesGateway | null @@ -247,22 +247,37 @@ function ChatRuntimeBoundary({ const [windowPages, setWindowPages] = useState(1) const [windowSessionKey, setWindowSessionKey] = useState(runtimeId) - // Sticky-cut continuity across flushes (advanceTranscriptWindow). A ref, not - // state: it is derived from `messages` and must never trigger a render. - const windowStateRef = useRef(null) + // Per-session sticky-cut continuity (advanceSessionTranscriptWindow). A ref, + // not state: it is derived from `messages` and must never trigger a render. + // Keyed by runtime id so a warm switch back to a session whose transcript + // is unchanged reuses the previous windowed slice BY REFERENCE — no window + // re-index, no runtime-repository rebuild, no per-row re-parse/re-highlight + // (#95595). Bounded internally (oldest session evicted). + const windowStateRef = useRef(new Map()) + // The memo below intentionally skips `runtimeId` in its deps (a switch + // always changes the messages array too, which re-runs it), so the current + // value must come from a ref rather than the stale render closure. + const runtimeIdRef = useRef(runtimeId) + runtimeIdRef.current = runtimeId // Reset the window on session swap during RENDER, so a large expand from the - // previous chat can't leak into the next one's first paint (#55191). + // previous chat can't leak into the next one's first paint (#55191). The + // per-session map above keeps each session's own cut; only the page count + // resets on a switch. if (windowSessionKey !== runtimeId) { setWindowSessionKey(runtimeId) setWindowPages(1) - windowStateRef.current = null } const { messages: windowedMessages, windowed } = useMemo(() => { - const next = advanceTranscriptWindow(windowStateRef.current, messages, windowPages) - - windowStateRef.current = next + const next = advanceSessionTranscriptWindow( + windowStateRef.current, + // Draft state has no runtime id yet; a single shared slot is fine there + // (mirrors the old single-slot behaviour for the no-runtime case). + runtimeIdRef.current ?? '', + messages, + windowPages + ) return next.window }, [messages, windowPages]) @@ -296,7 +311,7 @@ function ChatRuntimeBoundary({ // something older to show. Fire-and-forget: the prepend lands through the // session-state write path and re-renders this boundary. if ( - !windowStateRef.current?.window.windowed && + !windowStateRef.current.get(runtimeIdRef.current ?? '')?.state.window.windowed && runtimeId && storedId && transcriptBackfillAvailable(storedId, tailProfile) diff --git a/apps/desktop/src/app/chat/right-rail/preview-file.tsx b/apps/desktop/src/app/chat/right-rail/preview-file.tsx index f3723f19d7..a16d64f280 100644 --- a/apps/desktop/src/app/chat/right-rail/preview-file.tsx +++ b/apps/desktop/src/app/chat/right-rail/preview-file.tsx @@ -347,19 +347,7 @@ function MarkdownCode({ className, children, ...props }: ComponentProps<'code'>) const code = String(children).replace(/\n$/, '') - const highlighted = ( - - {code} - - ) + const highlighted = // ```mermaid / ```svg fences route to the shared lazy renderers (same // registry the chat transcript uses); everything else stays on Shiki. @@ -661,17 +649,7 @@ export function SourceView({ filePath, language, text }: { filePath?: string; la })}
- - {chunk.text} - +
))} diff --git a/apps/desktop/src/app/chat/transcript-window.test.ts b/apps/desktop/src/app/chat/transcript-window.test.ts index 600056e0b7..963b914a79 100644 --- a/apps/desktop/src/app/chat/transcript-window.test.ts +++ b/apps/desktop/src/app/chat/transcript-window.test.ts @@ -4,8 +4,10 @@ import type { ChatMessage } from '@/lib/chat-messages' import { RENDER_WEIGHT_CHARS } from '@/lib/render-weight' import { + advanceSessionTranscriptWindow, advanceTranscriptWindow, alignToBranchGroup, + MAX_SESSION_WINDOWS, selectTranscriptWindow, TRANSCRIPT_WINDOW_BUDGET, TRANSCRIPT_WINDOW_MIN_MESSAGES, @@ -214,6 +216,87 @@ describe('advanceTranscriptWindow', () => { }) }) +describe('advanceSessionTranscriptWindow', () => { + const heavyChars = RENDER_WEIGHT_CHARS * 40 + + it('matches a fresh walk on first visit', () => { + const memos = new Map() + const messages = transcript(400, heavyChars) + + const state = advanceSessionTranscriptWindow(memos, 'session-a', messages) + + expect(state.window).toEqual(selectTranscriptWindow(messages)) + expect(state.anchorId).toBe(state.window.messages[0].id) + }) + + it('returns the SAME windowed slice by reference on a warm re-visit with an unchanged transcript', () => { + const memos = new Map() + const sessionA = transcript(400, heavyChars) + const sessionB = transcript(300, heavyChars).map(m => ({ ...m, id: `b-${m.id}` })) + + // Visit B, then A, then B again — the exact warm-switch shape of #95595. + const firstB = advanceSessionTranscriptWindow(memos, 'session-b', sessionB) + advanceSessionTranscriptWindow(memos, 'session-a', sessionA) + const secondB = advanceSessionTranscriptWindow(memos, 'session-b', sessionB) + + expect(secondB.window.windowed).toBe(true) + // THE perf guard: same transcript array => same windowed slice reference, + // so the runtime repository and every row keep their identity. + expect(secondB.window.messages).toBe(firstB.window.messages) + expect(secondB.anchorId).toBe(firstB.anchorId) + }) + + it('holds the sticky cut when a re-visited session grew while away', () => { + const memos = new Map() + const messages = transcript(400, heavyChars) + const state = advanceSessionTranscriptWindow(memos, 'session-a', messages) + + // While away, the session streamed a few light turns (within slack). + const grown = [...messages, ...transcript(10, 100)] + const next = advanceSessionTranscriptWindow(memos, 'session-a', grown) + + // Sticky cut survived the switch-away: same anchor, no fresh re-walk. + expect(next.anchorId).toBe(state.anchorId) + expect(next.window.messages[0].id).toBe(state.window.messages[0].id) + // The anchored slice now simply includes the 10 new light turns. + expect(next.window.messages.length).toBe(state.window.messages.length + 10) + }) + + it('falls back to a fresh walk when the anchor vanished while away', () => { + const memos = new Map() + const messages = transcript(400, heavyChars) + advanceSessionTranscriptWindow(memos, 'session-a', messages) + + // Compression rewrite: disjoint ids while the user was elsewhere. + const rewritten = transcript(300, heavyChars).map(m => ({ ...m, id: `compressed-${m.id}` })) + const next = advanceSessionTranscriptWindow(memos, 'session-a', rewritten) + + expect(next.window).toEqual(selectTranscriptWindow(rewritten)) + }) + + it('re-walks when pages change on re-entry', () => { + const memos = new Map() + const messages = transcript(400, heavyChars) + + const one = advanceSessionTranscriptWindow(memos, 'session-a', messages, 1) + const two = advanceSessionTranscriptWindow(memos, 'session-a', messages, 2) + + expect(two.window.messages.length).toBeGreaterThan(one.window.messages.length) + }) + + it('keeps sessions independent and evicts the oldest memo past the cap', () => { + const memos = new Map() + + for (let i = 0; i < MAX_SESSION_WINDOWS + 5; i++) { + advanceSessionTranscriptWindow(memos, `session-${i}`, transcript(400, heavyChars)) + } + + expect(memos.size).toBeLessThanOrEqual(MAX_SESSION_WINDOWS) + expect(memos.has('session-0')).toBe(false) + expect(memos.has(`session-${MAX_SESSION_WINDOWS + 4}`)).toBe(true) + }) +}) + describe('alignToBranchGroup', () => { const messages = [message('u-1', 10), message('a-1', 10, 'g'), message('a-2', 10, 'g'), message('u-2', 10)] diff --git a/apps/desktop/src/app/chat/transcript-window.ts b/apps/desktop/src/app/chat/transcript-window.ts index b63e6d316c..6437331646 100644 --- a/apps/desktop/src/app/chat/transcript-window.ts +++ b/apps/desktop/src/app/chat/transcript-window.ts @@ -170,3 +170,65 @@ export function advanceTranscriptWindow( return { anchorId: window.windowed ? window.messages[0].id : null, pages, window } } + +/** How many sessions keep a sticky window before the oldest is evicted. */ +export const MAX_SESSION_WINDOWS = 12 + +/** + * A window state plus the exact message array it was computed from. + * The array identity is load-bearing: when a session is re-entered with the + * IDENTICAL transcript (the warm-switch path of #95595), the stored window — + * including the exact `window.messages` slice reference — is reused as-is. + * The reference reuse is what stops `useRuntimeMessageRepository` from + * rebuilding (and every row from re-rendering) on a warm switch. + */ +export interface SessionWindowMemo { + messages: readonly ChatMessage[] + state: TranscriptWindowState +} + +/** + * `advanceTranscriptWindow` with a STICKY cut that survives session switches. + * + * The previous single-slot state was nulled on every switch, so a warm + * re-entry always re-ran the weight walk and rebuilt the windowed slice — + * which re-indexed the whole windowed transcript (markdown re-parse + + * re-highlight per row) even though nothing had changed. This keeps one memo + * per session: + * + * - Re-entering a session with the same transcript array returns the cached + * windowed slice BY REFERENCE — the runtime repository and every message + * row stay mounted, so the switch is O(1). + * - Re-entering with a changed transcript keeps the sticky cut (anchor still + * present, tail within budget + slack) instead of re-walking from scratch. + * - The anchor vanishing (compression rewrite) or a pages change falls + * through to `advanceTranscriptWindow`'s existing fresh-walk behaviour. + * + * The map is bounded (oldest session evicted) so an unbounded session list + * cannot grow it without limit. + */ +export function advanceSessionTranscriptWindow( + memos: Map, + sessionKey: string, + messages: readonly ChatMessage[], + pages = 1 +): TranscriptWindowState { + const memo = memos.get(sessionKey) + + // Warm re-visit with the identical transcript and page count: reuse the + // cached state wholesale, preserving the windowed slice reference. + if (memo && memo.messages === messages && memo.state.pages === pages) { + return memo.state + } + + const state = advanceTranscriptWindow(memo?.state ?? null, messages, pages) + + memos.set(sessionKey, { messages, state }) + + if (memos.size > MAX_SESSION_WINDOWS) { + const oldest = memos.keys().next().value as string + memos.delete(oldest) + } + + return state +} 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 29dd2c7a5a..f5ee2ca66c 100644 --- a/apps/desktop/src/app/gateway/hooks/use-gateway-boot.ts +++ b/apps/desktop/src/app/gateway/hooks/use-gateway-boot.ts @@ -46,6 +46,7 @@ import { } from '@/store/gateway-switch' import { checkLocalRuntimeUpdate, watchLocalRuntimeJobs } from '@/store/local-runtime-jobs' import { notify, notifyError } from '@/store/notifications' +import { loadPoolLimits } from '@/store/pool-limits' import { $activeGatewayProfile, normalizeProfileKey, @@ -900,6 +901,10 @@ export function useGatewayBoot({ // this a socket dropped during sleep sits closed until the user clicks. window.addEventListener('focus', onFocus) + // Pool limits are main-process state; mirror them once for the Settings + // rows and prewarmProfileBackend's saturation guard. + void loadPoolLimits() + // Keep live pool backends alive while this window is open (the main process // can't observe the direct renderer↔backend WS). No-op for the primary. const keepaliveTimer = setInterval(() => { 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 4790f02765..cec1ebf521 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 @@ -157,6 +157,7 @@ import { isSessionGoneError, overlayConcurrentMessageChanges, patchSessionWorkspace, + preserveEquivalentTranscript, preserveLocalPendingTurnMessages, reconcileResumeMessages, removeRepresentedLocalLiveProjection, @@ -1372,26 +1373,37 @@ export function useSessionActions({ const activatedState = updateSessionState( cachedRuntimeId, - state => ({ - ...state, - messages: visibleActivatedMessages, - transcriptProvenance: - acceptedPersistedDisplayTranscript || hasValidProvenance - ? (expectedProvenance ?? undefined) - : undefined, - ...(pendingClarifyProjection - ? { - awaitingResponse: false, - sawAssistantPayload: true, - streamId: pendingClarifyProjection.streamId - } - : {}), - ...(clearedClarifyProjection - ? { - streamId: state.busy ? (clearedClarifyProjection.streamId ?? state.streamId) : null - } - : {}) - }), + state => { + // #95595: the reconcilers above always produce fresh + // message objects, so an unconditional publish replaces the + // warm-cached array with new-object equivalents and every + // visible row re-normalizes + remounts (markdown re-parse + + // shiki re-highlight per row, seconds of main-thread work). + // Keep the existing array when the content is unchanged — + // same guard the cold-resume path uses below. + const messages = preserveEquivalentTranscript(state.messages, visibleActivatedMessages) + + return { + ...state, + messages, + transcriptProvenance: + acceptedPersistedDisplayTranscript || hasValidProvenance + ? (expectedProvenance ?? undefined) + : undefined, + ...(pendingClarifyProjection + ? { + awaitingResponse: false, + sawAssistantPayload: true, + streamId: pendingClarifyProjection.streamId + } + : {}), + ...(clearedClarifyProjection + ? { + streamId: state.busy ? (clearedClarifyProjection.streamId ?? state.streamId) : null + } + : {}) + } + }, storedSessionId ) 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 d5f76b1423..19bff08385 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 @@ -26,6 +26,7 @@ import { goneSessionVerdict, isSessionGoneError, overlayConcurrentMessageChanges, + preserveEquivalentTranscript, preserveLocalPendingTurnMessages, reconcileResumeMessages, removeRepresentedLocalLiveProjection, @@ -1717,3 +1718,43 @@ describe('overlayConcurrentMessageChanges', () => { ]) }) }) + +describe('preserveEquivalentTranscript', () => { + it('keeps the current array BY REFERENCE when the replacement is content-equivalent', () => { + // The exact warm-resume shape of #95595: fresh objects, identical content. + const current = [msg('u-1', 'user', 'hello'), msg('a-1', 'assistant', 'const x = 1')] + const freshObjects = current.map(message => ({ ...message, parts: [...message.parts] })) + + const preserved = preserveEquivalentTranscript(current, freshObjects) + + expect(preserved).toBe(current) + expect(preserved[0]).toBe(current[0]) + }) + + it('keeps the current array when the arrays are the same reference', () => { + const current = [msg('u-1', 'user', 'hello')] + + expect(preserveEquivalentTranscript(current, current)).toBe(current) + }) + + it('accepts the replacement when anything changed', () => { + const current = [msg('u-1', 'user', 'hello')] + const next = [msg('u-1', 'user', 'hello'), msg('a-1', 'assistant', 'new turn')] + + expect(preserveEquivalentTranscript(current, next)).toBe(next) + }) + + it('rejects the replacement when a message diverges in content', () => { + const current = [msg('u-1', 'user', 'hello')] + const next = [msg('u-1', 'user', 'hello world')] + + expect(preserveEquivalentTranscript(current, next)).toBe(next) + }) + + it('rejects the replacement when metadata a row renders diverges', () => { + const current = [msg('u-1', 'user', 'hello')] + const next = [msg('u-1', 'user', 'hello', { pending: true })] + + expect(preserveEquivalentTranscript(current, next)).toBe(next) + }) +}) 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 113539132b..875b059ed2 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 @@ -298,6 +298,23 @@ export function chatMessageArraysEquivalent(a: ChatMessage[], b: ChatMessage[]): return a.length === b.length && a.every((message, index) => chatMessagesEquivalent(message, b[index])) } +/** + * Keep the CURRENT array when the replacement is content-equivalent. + * + * The resume reconcilers create fresh `ChatMessage` objects via + * `toChatMessages` even when nothing changed. Publishing those unconditionally + * replaces the `$messages`/session-slice array with a new reference of fresh + * objects — and because `useRuntimeMessageRepository` keys its normalization + * cache (and React keys its rows) by object identity, every message in the + * window re-normalizes and remounts: full markdown re-parse + shiki + * re-highlight per row, on the main thread, per warm session switch (#95595). + * Returning `current` when the content is equivalent keeps array AND object + * identity, so the warm switch is O(1) paint. + */ +export function preserveEquivalentTranscript(current: ChatMessage[], next: ChatMessage[]): ChatMessage[] { + return chatMessageArraysEquivalent(current, next) ? current : next +} + export function reconcileResumeMessages(nextMessages: ChatMessage[], previousMessages: ChatMessage[]): ChatMessage[] { if (!previousMessages.length) { return nextMessages diff --git a/apps/desktop/src/app/settings/config-settings.tsx b/apps/desktop/src/app/settings/config-settings.tsx index 40a30662da..73cb0d205b 100644 --- a/apps/desktop/src/app/settings/config-settings.tsx +++ b/apps/desktop/src/app/settings/config-settings.tsx @@ -45,6 +45,7 @@ import { import { MemoryConnect } from './memory/connect' import { ProviderConfigPanel } from './memory/provider-config-panel' import { ModelSettings, ModelSettingsSkeleton } from './model-settings' +import { PoolLimitsSetting } from './pool-limits-setting' import { EmptyState, ListRow, SettingsContent, SettingsSkeleton, ToggleRow } from './primitives' import { SettingsProfileScope } from './profile-scope' import { QuickEntrySettings } from './quick-entry-settings' @@ -405,6 +406,7 @@ function ConfigSettingsInner({ label={c.disableF12Title} onChange={setDisableF12} /> + )} diff --git a/apps/desktop/src/app/settings/pool-limits-setting.tsx b/apps/desktop/src/app/settings/pool-limits-setting.tsx new file mode 100644 index 0000000000..4db5ee8824 --- /dev/null +++ b/apps/desktop/src/app/settings/pool-limits-setting.tsx @@ -0,0 +1,113 @@ +import { useStore } from '@nanostores/react' +import { useEffect, useState } from 'react' + +import { ListRow } from '@/app/settings/primitives' +import { Input } from '@/components/ui/input' +import { $poolLimits, loadPoolLimits, savePoolLimits } from '@/store/pool-limits' + +// Bounds imported from main's clamp module so the advertised input ranges +// can never drift from what the pool actually enforces (review note on #92581). +import { POOL_LIMITS_BOUNDS } from '../../../electron/pool-limits' + +const MAX_BACKENDS_MAX = POOL_LIMITS_BOUNDS.maxBackendsMax +const IDLE_MS_MAX = POOL_LIMITS_BOUNDS.idleMsMax + +/** Settings → Advanced: warm-bot-backends count + backend idle timeout. + * Device-local (not profile-scoped): the pool is sized once per machine and + * changes apply live — main evicts/reaps to converge without a restart. */ +export function PoolLimitsSetting() { + const limits = useStore($poolLimits) + const [maxDraft, setMaxDraft] = useState(String(limits.maxBackends)) + const [idleDraft, setIdleDraft] = useState(String(limits.idleMs)) + + useEffect(() => { + void loadPoolLimits() + }, []) + + useEffect(() => { + setMaxDraft(String(limits.maxBackends)) + setIdleDraft(String(limits.idleMs)) + }, [limits]) + + const commitMax = () => { + const parsed = Number(maxDraft) + + if (!Number.isFinite(parsed) || parsed === limits.maxBackends) { + setMaxDraft(String(limits.maxBackends)) + + return + } + + void savePoolLimits({ maxBackends: parsed }) + .then(() => undefined) + .catch(() => setMaxDraft(String($poolLimits.get().maxBackends))) + } + + const commitIdle = () => { + const parsed = Number(idleDraft) + + if (!Number.isFinite(parsed) || parsed === limits.idleMs) { + setIdleDraft(String(limits.idleMs)) + + return + } + + void savePoolLimits({ idleMs: parsed }) + .then(() => undefined) + .catch(() => setIdleDraft(String($poolLimits.get().idleMs))) + } + + return ( + <> + + setMaxDraft(event.target.value)} + onKeyDown={event => { + if (event.key === 'Enter') { + event.currentTarget.blur() + } + }} + type="number" + value={maxDraft} + /> + + } + description="How many bot backends stay running for instant switching. Higher = faster switches, more memory (~60MB per backend). Applies immediately." + title="Warm Bot Backends" + /> + + setIdleDraft(event.target.value)} + onKeyDown={event => { + if (event.key === 'Enter') { + event.currentTarget.blur() + } + }} + type="number" + value={idleDraft} + /> + ms + + } + description="How long an unused bot backend stays warm before it is shut down. Raise this so bots you revisit every few minutes never pay a cold start." + title="Backend Idle Timeout" + /> + + ) +} diff --git a/apps/desktop/src/components/assistant-ui/thread/index.test.tsx b/apps/desktop/src/components/assistant-ui/thread/index.test.tsx new file mode 100644 index 0000000000..4bf2c82232 --- /dev/null +++ b/apps/desktop/src/components/assistant-ui/thread/index.test.tsx @@ -0,0 +1,82 @@ +import { render } from '@testing-library/react' +import { describe, expect, it, vi } from 'vitest' + +/** + * Issue #95595 proposed-fix #3: the `messageComponents` map handed to + * ThreadMessageList must keep its REFERENCE IDENTITY across a session switch. + * If it re-minted, React would unmount/remount every visible message — async + * re-rendered parts (shiki code blocks) collapse and re-expand, and the whole + * thread visibly jumps on every tab switch. + * + * The memo deps are deliberately only the boolean "definedness" gates (the + * callbacks themselves reach the composer through a ref), so a plain switch + * — sessionId changing, callbacks unchanged — must not change the map. + */ +let lastComponents: unknown + +vi.mock('@/components/assistant-ui/thread/list', () => ({ + ThreadMessageList: (props: { components: unknown }) => { + lastComponents = props.components + + return null + } +})) + +vi.mock('@/components/assistant-ui/thread/timeline', () => ({ + ThreadTimeline: () => null +})) + +vi.mock('@/components/assistant-ui/thread/status', () => ({ + BackgroundResumeNotice: () => null, + CenteredThreadSpinner: () => null +})) + +vi.mock('@/i18n', () => ({ + useI18n: () => ({ + t: { + assistant: { + thread: { + restoreBody: 'restore body', + restoreConfirm: 'Restore', + restoreTitle: 'Restore this turn?' + } + }, + common: { + cancel: 'Cancel', + confirm: 'Confirm', + done: 'Done', + loading: 'Loading' + } + } + }) +})) + +import { Thread } from './index' + +describe('Thread messageComponents identity across session switches', () => { + it('does not re-mint messageComponents when only the session changes', () => { + const { rerender } = render() + const first = lastComponents + + expect(first).toBeDefined() + + rerender() + + // THE guard: a switch must keep the component map reference, so the + // incoming transcript reconciles instead of remounting. + expect(lastComponents).toBe(first) + + rerender() + + expect(lastComponents).toBe(first) + }) + + it('keeps the map stable across a plain parent re-render', () => { + const { rerender } = render() + const first = lastComponents + + rerender() + + expect(lastComponents).toBe(first) + }) +}) diff --git a/apps/desktop/src/components/chat/shiki-block.test.tsx b/apps/desktop/src/components/chat/shiki-block.test.tsx new file mode 100644 index 0000000000..26273a73e7 --- /dev/null +++ b/apps/desktop/src/components/chat/shiki-block.test.tsx @@ -0,0 +1,122 @@ +import { cleanup, render, screen } from '@testing-library/react' +import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest' + +/** + * Perf regression guard for #95595: switching to a warm session remounts the + * incoming transcript, and every mounted code block used to re-tokenize from + * scratch on the main thread. The content-keyed cache must make a remount of + * an unchanged block a cache hit — ZERO highlighter calls. + * + * shiki itself is mocked (jsdom cannot run the oniguruma wasm engine); the + * mock counts `codeToHtml` invocations, which is the cost we are guarding. + */ +const { codeToHtml } = vi.hoisted(() => ({ + codeToHtml: vi.fn((code: string) => { + const escaped = String(code).replace(/&/g, '&').replace(//g, '>') + + return `
${escaped}
` + }) +})) + +vi.mock('shiki', () => ({ + bundledLanguages: { typescript: 'typescript-loader', text: 'text-loader' }, + getSingletonHighlighter: vi.fn(async () => ({ + codeToHtml: (code: string) => codeToHtml(code), + getLoadedLanguages: () => ['text', 'typescript'], + loadLanguage: vi.fn(async () => undefined) + })) +})) + +vi.mock('shiki/engine/oniguruma', () => ({ + createOnigurumaEngine: vi.fn(() => ({}) as never) +})) + +import CachedShikiBlock from '@/components/chat/shiki-block' +import { highlightCache } from '@/components/chat/shiki-highlight-cache' + +const TS_BLOCK = { language: 'typescript', code: 'const answer: number = 42\n' } + +async function waitForHighlighted(): Promise { + await screen.findByTestId('shiki-container', undefined, { timeout: 2_000 }) +} + +beforeEach(() => { + codeToHtml.mockClear() + highlightCache.clear() +}) + +afterEach(() => { + cleanup() +}) + +describe('CachedShikiBlock (warm-switch perf guard)', () => { + it('highlights on first mount and reuses the cached HTML on remount', async () => { + const { unmount } = render() + await waitForHighlighted() + + expect(codeToHtml).toHaveBeenCalledTimes(1) + expect(screen.getByTestId('shiki-container').innerHTML).toContain('const answer') + + // Warm session switch: the row unmounts and the SAME block mounts again. + unmount() + render() + await waitForHighlighted() + + // The guard: remounting an unchanged block must NOT re-tokenize. + expect(codeToHtml).toHaveBeenCalledTimes(1) + }) + + it('re-highlights a block whose code changed (cache miss)', async () => { + const { unmount } = render() + await waitForHighlighted() + + unmount() + render() + await waitForHighlighted() + + expect(codeToHtml).toHaveBeenCalledTimes(2) + }) + + it('keeps blocks independent: N blocks highlight N times across two mounts', async () => { + const { unmount } = render( + <> + + + + ) + + await screen.findAllByTestId('shiki-container', undefined, { timeout: 2_000 }) + + expect(codeToHtml).toHaveBeenCalledTimes(2) + + unmount() + render( + <> + + + + ) + await screen.findAllByTestId('shiki-container', undefined, { timeout: 2_000 }) + + // Two mounts of the same two blocks: exactly two tokenizations total. + expect(codeToHtml).toHaveBeenCalledTimes(2) + }) + + it('does not cache a failed highlight, so a retry can succeed', async () => { + codeToHtml.mockRejectedValueOnce(new Error('boom')) + + const { unmount } = render() + await waitForHighlighted() + + // The failure degrades to escaped plain text (and is NOT cached). + expect(screen.getByTestId('shiki-container').innerHTML).toContain('const answer') + expect(codeToHtml).toHaveBeenCalledTimes(1) + + unmount() + render() + await waitForHighlighted() + + // Second mount tries the highlighter again instead of serving stale HTML. + expect(codeToHtml).toHaveBeenCalledTimes(2) + }) +}) diff --git a/apps/desktop/src/components/chat/shiki-block.tsx b/apps/desktop/src/components/chat/shiki-block.tsx index 4279d21835..3704185caf 100644 --- a/apps/desktop/src/components/chat/shiki-block.tsx +++ b/apps/desktop/src/components/chat/shiki-block.tsx @@ -1,15 +1,173 @@ 'use client' /** - * The ONLY static importer of `react-shiki` (and through it the multi-MB - * shiki language/theme bundle). Every consumer reaches this module through + * The ONLY static importer of shiki (and through it the multi-MB shiki + * language/theme/wasm bundle). Every consumer reaches this module through * `React.lazy(() => import('./shiki-block'))` — see `LazyShiki` in * shiki-highlighter.tsx — so the shiki chunk stays entirely off the * cold-start path and loads on the first highlighted code block instead. * * Do NOT import this module statically from anything the entry graph * reaches, or the chunk moves back into boot. + * + * Unlike the previous pass-through of `react-shiki`'s component, this module + * is cache-aware: highlighted output is stored in a module-level LRU cache + * keyed by (theme scope, language, code), so a REMOUNT of an unchanged code + * block (the warm-session-switch path of #95595) paints the cached HTML + * synchronously and never re-tokenizes. Only cache misses run shiki, and + * misses are debounced so a streaming block settles before the heavy work + * starts. */ -import ShikiHighlighter from 'react-shiki' +import { useEffect, useMemo, useState } from 'react' +import { bundledLanguages, getSingletonHighlighter } from 'shiki' +import type { BundledLanguage, BundledTheme, Highlighter } from 'shiki' +import { createOnigurumaEngine } from 'shiki/engine/oniguruma' -export default ShikiHighlighter +import { SHIKI_HIGHLIGHT_SCOPE, SHIKI_THEME } from '@/components/chat/shiki-config' +import { highlightCache, highlightCacheKey } from '@/components/chat/shiki-highlight-cache' + +/** Same debounce react-shiki's `delay` used to throttle highlight work with. */ +const HIGHLIGHT_DELAY_MS = 120 + +// Stable identity for "no color replacements" so the memo/effect deps below +// never churn on renders that don't pass the prop. +const NO_COLOR_REPLACEMENTS: Record> = {} + +export interface CachedShikiBlockProps { + language: string + code: string + /** Theme override; defaults to the shared SHIKI_THEME. */ + theme?: { dark: string; light: string } + /** Color replacements; defaults to none (the chat passes its own). */ + colorReplacements?: Record> +} + +function isLoadableLanguage(language: string): boolean { + return language === 'text' || language in bundledLanguages +} + +/** + * Cache scope for one theme configuration. Part of the cache key: a block + * highlighted under a different theme is a different render. + */ +function highlightScope( + theme: { dark: string; light: string }, + colorReplacements: Record> +): string { + return `${SHIKI_HIGHLIGHT_SCOPE}:${theme.dark}:${theme.light}:${JSON.stringify(colorReplacements)}` +} + +let highlighterPromise: Promise | null = null +let loadedThemes = new Set([SHIKI_THEME.dark, SHIKI_THEME.light]) + +/** + * Lazily-created shiki singleton, mirroring react-shiki's full bundle: only + * the languages actually seen are loaded into it (the singleton is created + * with the first block's language, later ones are `loadLanguage`d on demand; + * override themes are `loadTheme`d the same way). Unknown languages are left + * unloaded and fall through to shiki's plain-text handling, as before. + */ +async function highlightToHtml( + language: string, + code: string, + theme: { dark: string; light: string }, + colorReplacements: Record> +): Promise { + if (!highlighterPromise) { + highlighterPromise = getSingletonHighlighter({ + // Only bundled languages are ever passed here (isLoadableLanguage + // guards both call sites), so the cast is safe. + langs: isLoadableLanguage(language) ? [language as BundledLanguage] : [], + themes: [SHIKI_THEME.dark, SHIKI_THEME.light], + engine: createOnigurumaEngine(import('shiki/wasm')) + }) + } + + const highlighter = await highlighterPromise + + if (isLoadableLanguage(language) && !highlighter.getLoadedLanguages().includes(language)) { + await highlighter.loadLanguage(language as BundledLanguage) + } + + const missingThemes = [theme.dark, theme.light].filter(name => !loadedThemes.has(name)) + + if (missingThemes.length > 0) { + await highlighter.loadTheme(...(missingThemes as BundledTheme[])) + missingThemes.forEach(name => loadedThemes.add(name)) + } + + return highlighter.codeToHtml(code, { + lang: language, + themes: { dark: theme.dark, light: theme.light }, + defaultColor: 'light-dark()', + colorReplacements + }) +} + +function escapeHtml(text: string): string { + return text.replace(/&/g, '&').replace(//g, '>').replace(/"/g, '"') +} + +/** Never let a highlight failure blank a block — degrade to escaped plain text. */ +function plainTextHtml(code: string): string { + return `
${escapeHtml(code)}
` +} + +export default function CachedShikiBlock({ language, code, theme, colorReplacements }: CachedShikiBlockProps) { + const themeConfig = theme ?? SHIKI_THEME + const replacements = colorReplacements ?? NO_COLOR_REPLACEMENTS + + const cacheKey = useMemo( + () => highlightCacheKey(highlightScope(themeConfig, replacements), language, code), + [language, code, replacements, themeConfig] + ) + + const [html, setHtml] = useState(() => highlightCache.get(cacheKey) ?? null) + + useEffect(() => { + let cancelled = false + + // Cache hit — no highlighter work at all. This is the warm-switch path: + // the previous visit already rendered this block, so paint it again. + const cached = highlightCache.get(cacheKey) + + if (cached !== undefined) { + setHtml(cached) + + return + } + + const timer = window.setTimeout(() => { + highlightToHtml(language, code, themeConfig, replacements) + .then(result => { + if (cancelled) { + return + } + + highlightCache.set(cacheKey, result) + setHtml(result) + }) + .catch(error => { + if (cancelled) { + return + } + + console.error('shiki highlight failed; rendering plain code', error) + setHtml(plainTextHtml(code)) + }) + }, HIGHLIGHT_DELAY_MS) + + return () => { + cancelled = true + window.clearTimeout(timer) + } + }, [cacheKey, code, language, replacements, themeConfig]) + + if (html === null) { + // Nothing to paint yet (miss, debounce pending). Matches react-shiki's + // own empty render while the highlight is in flight. + return null + } + + return
+} diff --git a/apps/desktop/src/components/chat/shiki-config.ts b/apps/desktop/src/components/chat/shiki-config.ts new file mode 100644 index 0000000000..4ba81aaf06 --- /dev/null +++ b/apps/desktop/src/components/chat/shiki-config.ts @@ -0,0 +1,33 @@ +// Shiki theme/color constants shared by the chat code-block renderer +// (shiki-highlighter.tsx) and the lazy shiki chunk (shiki-block.tsx). Kept in +// their own dependency-free module so the lazy chunk can import them without +// pulling the main chat module (or react-shiki) into the shiki bundle. + +// `github-dark-dimmed` is GitHub's lower-contrast dark palette — the vivid +// `github-dark-default` tokens read harsh at our small code size. Shared by the +// inline diff renderer too (see diff-lines.tsx) so code + diffs match. +export const SHIKI_THEME = { dark: 'github-dark-dimmed', light: 'github-light-default' } as const + +/** + * `github-light-default` colors comments `#6e7781` (~4.2:1 against the code + * card background) — borderline unreadable at our 11px code size, and worst of + * all for shell snippets where a single `#` turns the rest of the line into one + * long comment span. Remap light-mode comments to GitHub's darker muted gray + * (`#57606a`, ~6.4:1). Dark mode (`#8b949e`, ~6.1:1) already reads fine, so we + * leave it untouched. Keyed per theme name so the bump only applies in light. + */ +export const SHIKI_COLOR_REPLACEMENTS: Record> = { + 'github-light-default': { '#6e7781': '#57606a' } +} + +/** + * Cache-key scope for the content-addressed highlight cache. Bumping this + * invalidates every cached highlight at once — bump it whenever the rendering + * options (themes, color replacements) change, because keys are NOT allowed to + * silently produce a different DOM than the one they were computed with. + */ +export const SHIKI_HIGHLIGHT_SCOPE = `hermes-shiki-v1:${JSON.stringify({ + dark: SHIKI_THEME.dark, + light: SHIKI_THEME.light, + colorReplacements: SHIKI_COLOR_REPLACEMENTS +})}` diff --git a/apps/desktop/src/components/chat/shiki-highlight-cache.test.ts b/apps/desktop/src/components/chat/shiki-highlight-cache.test.ts new file mode 100644 index 0000000000..77dde0d941 --- /dev/null +++ b/apps/desktop/src/components/chat/shiki-highlight-cache.test.ts @@ -0,0 +1,95 @@ +import { describe, expect, it } from 'vitest' + +import { + HIGHLIGHT_CACHE_MAX_CHARS, + HIGHLIGHT_CACHE_MAX_ENTRIES, + HighlightCache, + highlightCacheKey +} from '@/components/chat/shiki-highlight-cache' + +describe('highlightCacheKey', () => { + it('separates scope, language and code so distinct blocks never collide', () => { + const a = highlightCacheKey('scope-1', 'ts', 'const x = 1') + const b = highlightCacheKey('scope-1', 'ts', 'const x = 2') + + expect(a).not.toBe(b) + expect(highlightCacheKey('scope-1', 'js', 'const x = 1')).not.toBe(a) + expect(highlightCacheKey('scope-2', 'ts', 'const x = 1')).not.toBe(a) + }) +}) + +describe('HighlightCache', () => { + it('round-trips an entry and refreshes recency on get', () => { + const cache = new HighlightCache(3, 1000) + + cache.set('a', 'A') + cache.set('b', 'B') + cache.set('c', 'C') + // Touch the oldest entry so it becomes the newest. + expect(cache.get('a')).toBe('A') + cache.set('d', 'D') + + // 'b' is now the least recently used and must be evicted first. + expect(cache.has('b')).toBe(false) + expect(cache.get('a')).toBe('A') + expect(cache.get('c')).toBe('C') + expect(cache.get('d')).toBe('D') + }) + + it('evicts oldest-first past the entry cap', () => { + const cache = new HighlightCache(2, 1_000_000) + + cache.set('a', 'A') + cache.set('b', 'B') + cache.set('c', 'C') + + expect(cache.size).toBe(2) + expect(cache.has('a')).toBe(false) + expect(cache.has('b')).toBe(true) + expect(cache.has('c')).toBe(true) + }) + + it('evicts oldest-first past the total-char cap', () => { + const cache = new HighlightCache(HIGHLIGHT_CACHE_MAX_ENTRIES, 10) + + cache.set('a', '12345') + cache.set('b', '123456') + + expect(cache.size).toBe(1) + expect(cache.has('a')).toBe(false) + expect(cache.has('b')).toBe(true) + expect(cache.totalChars).toBeLessThanOrEqual(10) + }) + + it('replaces an existing key in place without double-counting chars', () => { + const cache = new HighlightCache(2, 100) + + cache.set('a', '12345') + cache.set('a', '1234567890') + + expect(cache.size).toBe(1) + expect(cache.totalChars).toBe(10) + }) + + it('stays within both caps for a large burst of unique blocks', () => { + const cache = new HighlightCache(HIGHLIGHT_CACHE_MAX_ENTRIES, HIGHLIGHT_CACHE_MAX_CHARS) + + for (let i = 0; i < 2_000; i++) { + cache.set(`block-${i}`, `html-${i}`.repeat(100)) + } + + expect(cache.size).toBeLessThanOrEqual(HIGHLIGHT_CACHE_MAX_ENTRIES) + expect(cache.totalChars).toBeLessThanOrEqual(HIGHLIGHT_CACHE_MAX_CHARS) + }) + + it('clear drops everything', () => { + const cache = new HighlightCache() + + cache.set('a', 'A') + cache.clear() + + expect(cache.size).toBe(0) + expect(cache.totalChars).toBe(0) + expect(cache.get('a')).toBeUndefined() + }) +}) diff --git a/apps/desktop/src/components/chat/shiki-highlight-cache.ts b/apps/desktop/src/components/chat/shiki-highlight-cache.ts new file mode 100644 index 0000000000..3b0b7fe93b --- /dev/null +++ b/apps/desktop/src/components/chat/shiki-highlight-cache.ts @@ -0,0 +1,90 @@ +// ── Content-addressed syntax-highlight cache (#95595) ──────────────────────── +// Switching to a warm session remounts the incoming transcript, and every +// mounted code block used to be re-tokenized from scratch by shiki — N blocks +// × full tokenization on the renderer main thread, on every switch, even +// though the code had not changed. The fix is a module-level LRU cache keyed +// by (scope, language, code) holding the final highlighted HTML, so a remount +// of an unchanged block paints the cached markup synchronously and never +// touches the highlighter. +// +// Bounds: shiki's tokenized HTML is ~5-10x the source size, so an unbounded +// cache would leak renderer memory over a long session list. Cap both the +// entry count and the total cached characters; evict oldest-first. +// +// This module is intentionally dependency-free (no React, no shiki) so the +// cache logic can be unit-tested in isolation. + +export const HIGHLIGHT_CACHE_MAX_ENTRIES = 512 +export const HIGHLIGHT_CACHE_MAX_CHARS = 6 * 1024 * 1024 + +/** Unique key for one highlighted block: scope + language + exact code. */ +export function highlightCacheKey(scope: string, language: string, code: string): string { + return `${scope}\u0000${language}\u0000${code}` +} + +/** + * Bounded LRU map of highlight cache keys to rendered HTML. `get` refreshes + * recency (Map insertion order is used as the LRU clock); `set` evicts the + * oldest entries until both caps hold. + */ +export class HighlightCache { + private readonly entries = new Map() + private chars = 0 + + constructor( + private readonly maxEntries: number = HIGHLIGHT_CACHE_MAX_ENTRIES, + private readonly maxChars: number = HIGHLIGHT_CACHE_MAX_CHARS + ) {} + + get size(): number { + return this.entries.size + } + + get totalChars(): number { + return this.chars + } + + has(key: string): boolean { + return this.entries.has(key) + } + + get(key: string): string | undefined { + const value = this.entries.get(key) + + if (value !== undefined) { + // Refresh recency: re-inserting moves the entry to the newest position. + this.entries.delete(key) + this.entries.set(key, value) + } + + return value + } + + set(key: string, html: string): void { + if (this.entries.has(key)) { + this.chars -= this.entries.get(key)!.length + this.entries.delete(key) + } + + this.entries.set(key, html) + this.chars += html.length + this.evict() + } + + clear(): void { + this.entries.clear() + this.chars = 0 + } + + private evict(): void { + while ((this.entries.size > this.maxEntries || this.chars > this.maxChars) && this.entries.size > 0) { + const oldestKey = this.entries.keys().next().value as string + const oldest = this.entries.get(oldestKey)! + this.entries.delete(oldestKey) + this.chars -= oldest.length + } + } +} + +/** The renderer-wide highlight cache. Lives for the lifetime of the module. */ +export const highlightCache = new HighlightCache() diff --git a/apps/desktop/src/components/chat/shiki-highlighter.tsx b/apps/desktop/src/components/chat/shiki-highlighter.tsx index 6823bba7cf..dcb27bd9a1 100644 --- a/apps/desktop/src/components/chat/shiki-highlighter.tsx +++ b/apps/desktop/src/components/chat/shiki-highlighter.tsx @@ -1,15 +1,20 @@ 'use client' import type { SyntaxHighlighterProps } from '@assistant-ui/react-streamdown' -import { type ComponentProps, type FC, lazy, Suspense, useMemo } from 'react' -import type ShikiHighlighter from 'react-shiki' +import { type FC, lazy, Suspense, useMemo } from 'react' import { CodeCard, CodeCardBody } from '@/components/chat/code-card' import { ExpandableBlock } from '@/components/chat/expandable-block' +// Theme constants live in shiki-config (dependency-free) so the lazy shiki +// chunk can import them without pulling this module into the shiki bundle. +import { SHIKI_COLOR_REPLACEMENTS } from '@/components/chat/shiki-config' import { CopyButton } from '@/components/ui/copy-button' import { useI18n } from '@/i18n' import { isLikelyProseCodeBlock } from '@/lib/markdown-code' +import type { CachedShikiBlockProps } from './shiki-block' +export { SHIKI_COLOR_REPLACEMENTS, SHIKI_THEME } from '@/components/chat/shiki-config' + /** * Streamdown's code adapter renders header + body as inline siblings, so we * own the wrapping `` here and neutralize the upstream @@ -17,46 +22,34 @@ import { isLikelyProseCodeBlock } from '@/lib/markdown-code' * background-only — no header row, no language label — so a fence reads as a * tinted slab of the reply; copy is a hover-reveal control in the corner. * - * `react-shiki` full bundle so all `bundledLanguages` work; theme switches - * follow the document `color-scheme` via `defaultColor="light-dark()"`. + * The heavy lifting lives in the lazy `shiki-block` chunk (full bundle so all + * `bundledLanguages` work; theme switches follow the document `color-scheme` + * via `defaultColor="light-dark()"`), and its output is cached by content so + * warm-session switches never re-tokenize unchanged blocks (#95595). */ interface HermesSyntaxHighlighterProps extends SyntaxHighlighterProps { defer?: boolean } -// `github-dark-dimmed` is GitHub's lower-contrast dark palette — the vivid -// `github-dark-default` tokens read harsh at our small code size. Shared by the -// inline diff renderer too (see diff-lines.tsx) so code + diffs match. -export const SHIKI_THEME = { dark: 'github-dark-dimmed', light: 'github-light-default' } as const - -/** - * `github-light-default` colors comments `#6e7781` (~4.2:1 against the code - * card background) — borderline unreadable at our 11px code size, and worst of - * all for shell snippets where a single `#` turns the rest of the line into one - * long comment span. Remap light-mode comments to GitHub's darker muted gray - * (`#57606a`, ~6.4:1). Dark mode (`#8b949e`, ~6.1:1) already reads fine, so we - * leave it untouched. Keyed per theme name so the bump only applies in light. - */ -const SHIKI_COLOR_REPLACEMENTS: Record> = { - 'github-light-default': { '#6e7781': '#57606a' } -} - const MAX_HIGHLIGHT_CHARS = 150_000 const MAX_HIGHLIGHT_LINES = 3_000 const CHUNK_LINES = 200 const EST_LINE_PX = 16 -// react-shiki (and through it the multi-MB shiki grammar/theme bundle) is the +// shiki (and through it the multi-MB grammar/theme/wasm bundle) is the // heaviest dependency in the renderer. `shiki-block.tsx` is its only static // importer, so this lazy() is the single seam that keeps shiki out of the // entry chunk — it loads on the first highlighted code block, not at boot. +// The lazy module is cache-aware (#95595): unchanged blocks paint from a +// content-keyed cache instead of re-tokenizing on every mount. const ShikiBlock = lazy(() => import('./shiki-block')) -/** Drop-in ShikiHighlighter that suspends on first use and renders the code - * as plain preformatted text until the shiki chunk arrives. */ -export const LazyShiki: FC> = props => ( - }> - +/** Suspends on first use and renders the code as plain preformatted text + * until the shiki chunk arrives. Highlighted output is cached by + * (theme, language, code), so revisits never re-tokenize (#95595). */ +export const LazyShiki: FC = ({ language, code, theme, colorReplacements }) => ( + }> + ) @@ -160,18 +153,7 @@ export const SyntaxHighlighter: FC = ({ {plain ? ( ) : ( - - {trimmed} - + )} diff --git a/apps/desktop/src/components/notifications.tsx b/apps/desktop/src/components/notifications.tsx index 788873c531..e8a8dda8d6 100644 --- a/apps/desktop/src/components/notifications.tsx +++ b/apps/desktop/src/components/notifications.tsx @@ -236,9 +236,9 @@ function NotificationItem({ notification }: { notification: AppNotification }) { notification.action?.onClick() dismissNotification(notification.id) }} - size="xs" + size="sm" type="button" - variant="textStrong" + variant="default" > {notification.action.label} diff --git a/apps/desktop/src/components/pane-shell/tree/store.ts b/apps/desktop/src/components/pane-shell/tree/store.ts index a451e570e3..07e9ed09f7 100644 --- a/apps/desktop/src/components/pane-shell/tree/store.ts +++ b/apps/desktop/src/components/pane-shell/tree/store.ts @@ -1040,6 +1040,13 @@ export function revealTreePane(paneId: string) { // Reveal beats a Close: un-dismiss and let adoption put the pane back. if ($dismissedPanes.get().has(paneId)) { setDismissed(paneId, false) + } + + // A layout replacement can omit a still-registered pane without dismissing + // it. Reconcile that saved contribution before claiming to reveal it. + const currentTree = $layoutTree.get() + + if (currentTree && !findGroupOfPane(currentTree, paneId)) { adoptContributedPanes() } @@ -1064,8 +1071,8 @@ export function revealTreePane(paneId: string) { if (hiddenNow.has(paneId)) { setTreePaneHidden(paneId, false) - - return + // Reactive unhide preserves a visible sibling. Explicit reveal must also + // front this pane and restore its group below. } const tree = $layoutTree.get() diff --git a/apps/desktop/src/global.d.ts b/apps/desktop/src/global.d.ts index 8067308df7..0362b8ff12 100644 --- a/apps/desktop/src/global.d.ts +++ b/apps/desktop/src/global.d.ts @@ -1,6 +1,8 @@ import type { GatewayWsUrlResult } from '@hermes/shared' import type { TranslucencyState } from '@hermes/shared/translucency' +import type { PoolLimits } from '../electron/pool-limits' + import type { WakeIndicatorState } from './lib/wake-indicator' import type { PetOverlayBounds, @@ -46,6 +48,14 @@ declare global { // Keepalive: mark a pool profile backend as recently used so the idle // reaper spares it while its chat is active. touchBackend: (profile?: string | null) => Promise<{ ok: boolean }> + // Pool sizing (Settings → Advanced): device-local, live-applied by the + // main process. get resolves the limits currently in force; set applies + // (and persists) new ones, evicting/reaping to converge immediately. + getPoolLimits: () => Promise + setPoolLimits: (limits: { maxBackends?: number; idleMs?: number }) => Promise<{ + ok: boolean + limits: PoolLimits + }> getGatewayWsUrl: (profile?: null | string) => Promise // Open (or focus) a standalone OS window for a single chat session so // the user can work with multiple chats side by side. Returns ok:false diff --git a/apps/desktop/src/plugins/hermes-bots/bot-row-opens-canonical-chat.test.ts b/apps/desktop/src/plugins/hermes-bots/bot-row-opens-canonical-chat.test.ts index ca87583ff5..a50c880280 100644 --- a/apps/desktop/src/plugins/hermes-bots/bot-row-opens-canonical-chat.test.ts +++ b/apps/desktop/src/plugins/hermes-bots/bot-row-opens-canonical-chat.test.ts @@ -71,6 +71,24 @@ describe('a row click lands on the canonical chat, never a remembered side tab', }) }) + it('fronting an already-open Bot Chat refreshes its transcript in place', async () => { + // The front is presentation-only: the pane keeps whatever transcript it + // last painted, which can predate rows the bot wrote while the user was + // elsewhere (a cron delivery, a teammate's message_agent, another bot's + // turn). Fronting must force a registry open so forceResume re-pulls the + // latest rows instead of leaving a stale snapshot until the next turn + // (#99393 class; #95600 only covered the not-yet-open path). + host.focusOpenWorkspaceSession = vi.fn((_key: string, _probe: unknown, only?: readonly string[]) => + only?.includes('bot-chat-tip') ? 'bot-chat-tip' : null + ) as never + $selectedStoredSessionId.set('bot-chat-tip') + + await expect(openRosterBot(canonicalBot)).resolves.toBe(true) + + expect(openBotCanonicalChat).toHaveBeenCalledWith(canonicalBot, expect.any(Function)) + $selectedStoredSessionId.set(null) + }) + it('resolves the registry when only a side thread is open', async () => { // The shell would happily front 'side-thread' — the allowlist excludes it. host.focusOpenWorkspaceSession = vi.fn((_key: string, _probe: unknown, only?: readonly string[]) => diff --git a/apps/desktop/src/plugins/hermes-bots/plugin.mentions.test.ts b/apps/desktop/src/plugins/hermes-bots/plugin.mentions.test.ts index 1844349c38..de02c50025 100644 --- a/apps/desktop/src/plugins/hermes-bots/plugin.mentions.test.ts +++ b/apps/desktop/src/plugins/hermes-bots/plugin.mentions.test.ts @@ -253,14 +253,37 @@ describe('the mention middleware', () => { expect(hostMock.requestProfile).not.toHaveBeenCalled() }) - it('hands the agent the connection-qualified message_agent target', async () => { + it('hands the agent a relay-resolvable canonical target', async () => { const { handler } = await contributions() const result = await handler({ text: 'ping @default-vera' }) - expect(result.text).toMatch(/message_agent target: "default-vera@vera"/) + // The UI alias ('default-vera') is not a relay identity: resolve_remote_target() + // accepts only a roster row's handle/profile, optionally @connection-qualified. + // The annotation must carry the canonical profile@connection form (#97678). + expect(result.text).toMatch(/message_agent target: "default@vera"/) + expect(result.text).not.toMatch(/message_agent target: "default-vera/) expect(result.text).toMatch(/on Vera/) }) + it('annotates the resolvable handle for a local row whose UI alias differs', async () => { + // The reporter's shape (#97678 / Discord video): the LOCAL twin carries + // the 'default-this-device' alias when the remote gateway is active. + // The local resolver only knows bare profile names / 'hermes'. + const { handler } = await contributions({ + focused: 'ops', + profiles: [ + { connectionId: 'local', connectionKind: 'local', handle: 'default-this-device', name: 'default' }, + { name: 'ops' } + ] + }) + + const result = await handler({ text: 'ping @default-this-device' }) + + expect(result.text).toMatch(/@default-this-device = agent profile "default"/) + expect(result.text).toMatch(/message_agent target: "hermes"/) + expect(result.text).not.toMatch(/message_agent target: "default-this-device/) + }) + it('passes a draft with no mention straight through', async () => { const { handler } = await contributions() const draft = { text: 'no tags here' } diff --git a/apps/desktop/src/plugins/hermes-bots/plugin.tsx b/apps/desktop/src/plugins/hermes-bots/plugin.tsx index 02f52826a8..ffdcaeb8fa 100644 --- a/apps/desktop/src/plugins/hermes-bots/plugin.tsx +++ b/apps/desktop/src/plugins/hermes-bots/plugin.tsx @@ -734,11 +734,23 @@ export default { botRosterMeta(bot, $botMeta.get())?.title || bot.ui_meta?.['hermes-bots']?.title || bot.title || '' ).trim() - const target = bot.remoteSource && bot.connectionId ? `${handle}@${bot.connectionId}` : handle + // message_agent only resolves canonical identities: the relay + // matches a roster row's handle/profile (± @connection-id), the + // local path a bare profile name or 'hermes'. botHandle() prefers + // the row's source-qualified UI alias ('default-vera'), which + // neither resolver accepts — annotate the canonical form instead. + const target = + bot.remoteSource && bot.connectionId ? `${bot.name}@${bot.connectionId}` : botHandle(bot.name) + // Local rows get the same annotation whenever their UI alias + // ('default-this-device') differs from the resolvable handle — + // otherwise the agent has only the alias to go on and the local + // path rejects it the same way (#97678). const where = bot.remoteSource ? ` — on ${bot.connectionLabel || bot.connectionId} (message_agent target: "${target}")` - : '' + : handle !== target + ? ` (message_agent target: "${target}")` + : '' return `@${handle} = agent profile "${bot.name}"${title ? ` ("${title}")` : ''}${where}` }) diff --git a/apps/desktop/src/plugins/hermes-bots/roster-actions.ts b/apps/desktop/src/plugins/hermes-bots/roster-actions.ts index 29f2ce81d2..fb7cce2397 100644 --- a/apps/desktop/src/plugins/hermes-bots/roster-actions.ts +++ b/apps/desktop/src/plugins/hermes-bots/roster-actions.ts @@ -246,6 +246,14 @@ export async function openRosterBot(bot: RosterRow): Promise { // the roster-activity refresh treat it exactly like a registry open. $openBotChat.set({ key, openedRegistryId: fronted.registryId, openedSessionId: fronted.storedSessionId }) + // Fronting is presentation-only: the pane keeps whatever transcript it + // last painted, which can predate rows the bot wrote while the user was + // elsewhere (another bot's turn, a cron delivery, a teammate's + // message_agent). Force a registry open so forceResume re-pulls the + // latest transcript instead of leaving a stale snapshot until the next + // user turn (#99393 class; #95600 only covered the not-yet-open path). + refreshOpenBotChat(bot) + return true } diff --git a/apps/desktop/src/store/gateway.ts b/apps/desktop/src/store/gateway.ts index b07971d13f..5bf36abaea 100644 --- a/apps/desktop/src/store/gateway.ts +++ b/apps/desktop/src/store/gateway.ts @@ -1584,6 +1584,24 @@ export function reconnectSecondaryGateways({ forceOpenSockets = false }: { force } } +// How many non-primary backends currently hold an open socket. Hover-intent +// prewarming consults this before spawning: a speculative spawn that pushes +// the pool past its cap causes the Electron main to LRU-evict a warm backend +// — often one the user is about to click — turning the prewarm into churn +// (the #91545 evict/respawn cascade). The active gateway's backend is +// primary-routed and never counts toward the pool cap. +export function openSecondaryCount(): number { + let count = 0 + + for (const entry of g.secondaries.values()) { + if (isOpen(entry.gateway)) { + count += 1 + } + } + + return count +} + // Keep the idle reaper from killing a backend we still need: ping every live // secondary. The active one is pinged separately (touchActiveGatewayBackend). export function touchSecondaryGateways(): void { diff --git a/apps/desktop/src/store/pool-limits.ts b/apps/desktop/src/store/pool-limits.ts new file mode 100644 index 0000000000..71bd9e4ade --- /dev/null +++ b/apps/desktop/src/store/pool-limits.ts @@ -0,0 +1,65 @@ +/** + * Pool limits — how many bot backends may stay spawned, and how long an + * unused one survives before it is shut down. + * + * A device-local preference (each machine trades RAM against switching + * speed for itself). The MAIN process is authoritative: it owns the pool + * and the persisted copy, and applies a new max immediately by evicting + * least-recently-used idle backends — no restart. This store mirrors the + * live values for the Settings rows and feeds prewarmProfileBackend's + * saturation guard. + */ + +import { atom } from 'nanostores' + +export interface PoolLimits { + /** Max concurrently spawned non-primary profile backends. */ + maxBackends: number + /** Idle lifetime of an unused pool backend, in milliseconds. */ + idleMs: number +} + +export const POOL_LIMITS_DEFAULTS: PoolLimits = { + maxBackends: 3, + idleMs: 10 * 60_000 +} + +export const $poolLimits = atom({ ...POOL_LIMITS_DEFAULTS }) + +/** Seed from main's authoritative state once at startup; no-op without the + * bridge (web/older builds just keep the defaults for the UI). */ +export async function loadPoolLimits(): Promise { + try { + const limits = await window.hermesDesktop?.getPoolLimits?.() + + if (limits) { + $poolLimits.set(limits) + } + } catch { + // Keep defaults — Settings rows still render and can retry on save. + } +} + +/** Push new limits to main; adopt the post-clamp values it reports. */ +export async function savePoolLimits(next: { maxBackends?: number; idleMs?: number }): Promise { + const current = $poolLimits.get() + + const optimistic: PoolLimits = { + maxBackends: next.maxBackends ?? current.maxBackends, + idleMs: next.idleMs ?? current.idleMs + } + + // Optimistic paint, then honest reconciliation with the clamped result. + $poolLimits.set(optimistic) + + try { + const result = await window.hermesDesktop?.setPoolLimits?.(next) + + if (result?.limits) { + $poolLimits.set(result.limits) + } + } catch { + $poolLimits.set(current) + throw new Error('Applying pool limits failed') + } +} diff --git a/apps/desktop/src/store/profile.test.ts b/apps/desktop/src/store/profile.test.ts index e6825de85a..33c9578427 100644 --- a/apps/desktop/src/store/profile.test.ts +++ b/apps/desktop/src/store/profile.test.ts @@ -9,10 +9,24 @@ import type { ProfileInfo } from '@/types/hermes' const ensureGatewayForProfile = vi.fn(async () => undefined) const ensureGatewayForAgent = vi.fn(async () => undefined) const openGatewayForProfile = vi.fn(async (_profile: string) => undefined) +const openSecondaryCount = vi.fn(() => 0) const $gateway = atom({ id: 'live-socket', connectionState: 'open' }) const resetStarmapGraph = vi.fn() -vi.mock('@/store/gateway', () => ({ $gateway, ensureGatewayForAgent, ensureGatewayForProfile, openGatewayForProfile })) +vi.mock('@/store/gateway', () => ({ + $gateway, + ensureGatewayForAgent, + ensureGatewayForProfile, + openGatewayForProfile, + openSecondaryCount +})) +// The pool-limits atom is profile.ts's live saturation signal — keep the real +// one so tests can move the cap via the store, but stub its IPC bridge. +vi.mock('@/store/pool-limits', async () => { + const { atom } = await import('nanostores') + + return { $poolLimits: atom({ idleMs: 600_000, maxBackends: 3 }) } +}) vi.mock('@/hermes', () => ({ getProfiles: vi.fn(async () => ({ profiles: [] })), setApiRequestProfile: vi.fn() @@ -29,6 +43,8 @@ const { refreshProfiles } = await import('./profile') +const { $poolLimits } = await import('@/store/pool-limits') + const { $connection } = await import('./session') const { invalidateProfileScopedQueries } = await import('@/lib/query-client') const { getProfiles } = await import('@/hermes') @@ -55,6 +71,7 @@ beforeEach(() => { getConnection.mockReset() ensureGatewayForProfile.mockClear() openGatewayForProfile.mockClear() + openSecondaryCount.mockReturnValue(0) $gateway.set({ id: 'live-socket', connectionState: 'open' }) $activeGatewayProfile.set('default') $connection.set(localConn()) @@ -169,6 +186,43 @@ describe('prewarmProfileBackend (hover-intent pool spawn)', () => { expect(() => prewarmProfileBackend('warm-failing')).not.toThrow() }) + + it('skips pre-warm when the pool is saturated (#91545 evict/respawn cascade)', () => { + // Every pool slot occupied: a speculative spawn would LRU-evict a warm + // backend — often the one the user is about to click. Default limit 3, + // 3 open secondaries → the next spawn would exceed the cap. + openSecondaryCount.mockReturnValue(3) + + prewarmProfileBackend('warm-saturated') + + expect(openGatewayForProfile).not.toHaveBeenCalled() + }) + + it('pre-warms while pool slots are free', () => { + openSecondaryCount.mockReturnValue(1) + + prewarmProfileBackend('warm-slot-free') + + expect(openGatewayForProfile).toHaveBeenCalledWith('warm-slot-free') + }) + + it('follows the live pool-limit atom, not a hard-coded cap', () => { + // User raises Warm Bot Backends to 8 in Settings: prewarming must keep + // working well past the old default of 3. + openSecondaryCount.mockReturnValue(5) + $poolLimits.set({ idleMs: 600_000, maxBackends: 8 }) + + prewarmProfileBackend('warm-raised-cap') + + expect(openGatewayForProfile).toHaveBeenCalledWith('warm-raised-cap') + + // And lowering the cap re-engages the guard at the new boundary. + $poolLimits.set({ idleMs: 600_000, maxBackends: 2 }) + + prewarmProfileBackend('warm-lowered-cap') + + expect(openGatewayForProfile).not.toHaveBeenCalledWith('warm-lowered-cap') + }) }) describe('refreshProfiles shared rail list (#49289)', () => { diff --git a/apps/desktop/src/store/profile.ts b/apps/desktop/src/store/profile.ts index 9c9fdaf29a..cf903e120e 100644 --- a/apps/desktop/src/store/profile.ts +++ b/apps/desktop/src/store/profile.ts @@ -21,9 +21,11 @@ import { ensureGatewayForAgent, ensureGatewayForProfile, openGatewayForAgent, - openGatewayForProfile + openGatewayForProfile, + openSecondaryCount } from '@/store/gateway' import { notifyError } from '@/store/notifications' +import { $poolLimits } from '@/store/pool-limits' import { notifyRemoteOverrideAuthFailure } from '@/store/profile-remote-override' import { clearComposerSelectionOwner, setComposerSelectionOwner, setConnection } from '@/store/session' import type { SessionOwnerRoute } from '@/store/session-request-router' @@ -423,6 +425,17 @@ export function prewarmProfileBackend(name: string): void { return } + // Prewarm/cap harmony (#91545): the pool caps spawned backends at the + // configured max, and a spawn over the cap LRU-evicts the warmest idle + // backend. A hover sweep across the rail therefore evicted backends for + // profiles the user was about to click — prewarming caused the exact churn + // it exists to prevent. Skip speculative spawns once every pool slot is + // occupied by an open socket; the real click still spawns on demand, it + // just doesn't get a head start. + if (openSecondaryCount() + 1 > $poolLimits.get().maxBackends) { + return + } + prewarmedAt.set(key, now) openGatewayForProfile(key).catch(() => undefined) } diff --git a/apps/desktop/src/store/session-pane-focus.test.ts b/apps/desktop/src/store/session-pane-focus.test.ts new file mode 100644 index 0000000000..0c0b2688b1 --- /dev/null +++ b/apps/desktop/src/store/session-pane-focus.test.ts @@ -0,0 +1,82 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest' + +async function setup() { + const tree = await import('@/components/pane-shell/tree/store') + const model = await import('@/components/pane-shell/tree/model') + const { registry } = await import('@/contrib/registry') + const session = await import('@/store/session') + const states = await import('@/store/session-states') + const { paneMirror } = await import('@/app/chat/pane-mirror') + + registry.register({ + area: 'panes', + data: { placement: 'main', uncloseable: true }, + id: 'workspace', + render: () => null, + title: 'Chat' + }) + tree.declareDefaultTree(model.group(['workspace'], { active: 'workspace', id: 'main' })) + tree.watchContributedPanes() + paneMirror({ + source: states.$sessionTiles, + key: tile => tile.storedSessionId, + prefix: 'session-tile', + dir: () => 'center', + minWidth: '20rem', + title: id => id, + render: () => null, + close: states.closeSessionTile + })() + session.$selectedStoredSessionId.set('previous-chat') + + const scope = { + ownerRoute: { connectionId: 'remote-a', mode: 'remote' as const, profile: 'writer' }, + workspaceMode: 'bots' as const, + workspaceOwnerKey: 'remote-a::writer', + workspaceTabTitle: 'Bot Chat' + } + + states.openSessionTile('canonical-chat', 'center', 'workspace', undefined, scope) + + return { model, scope, session, states, tree } +} + +describe('focusing a saved Bot Chat requires a visible pane', () => { + let ctx: Awaited> + const paneId = 'session-tile:canonical-chat' + + beforeEach(async () => { + window.localStorage.clear() + vi.resetModules() + ctx = await setup() + }) + + it('re-adopts a saved tab after a profile overlay replaces the layout', async () => { + const { applyDesktopOverlay } = await import('@/store/profile-share') + const { model, scope, states, tree } = ctx + const saved = states.$sessionTiles.get() + applyDesktopOverlay('imported-profile', { + version: 1, + layoutTree: model.group(['workspace'], { active: 'workspace', id: 'imported-main' }) + }) + expect(model.findGroupOfPane(tree.$layoutTree.get()!, paneId)).toBeNull() + + expect(states.focusWorkspaceOwnerSessionTile(scope.workspaceOwnerKey, undefined, ['canonical-chat'])).toBe( + 'canonical-chat' + ) + expect(tree.isPaneVisible(paneId)).toBe(true) + expect(tree.$activeTreeGroup.get()).toBe('imported-main') + expect(states.$sessionTiles.get()).toEqual(saved) + expect(states.sessionTileOwnerRoute('canonical-chat')).toEqual(scope.ownerRoute) + }) + + it('reports a miss through both helpers if the layout cannot place the saved tab', () => { + const { scope, session, states, tree } = ctx + tree.$layoutTree.set(null) + + expect(states.focusOpenSession('canonical-chat', scope)).toBeNull() + expect(states.focusWorkspaceOwnerSessionTile(scope.workspaceOwnerKey, undefined, ['canonical-chat'])).toBeNull() + expect(session.$selectedStoredSessionId.get()).toBe('previous-chat') + expect(states.$sessionTiles.get().map(tile => tile.storedSessionId)).toEqual(['canonical-chat']) + }) +}) diff --git a/apps/desktop/src/store/session-states.test.ts b/apps/desktop/src/store/session-states.test.ts index a5f887cbb0..1926948f27 100644 --- a/apps/desktop/src/store/session-states.test.ts +++ b/apps/desktop/src/store/session-states.test.ts @@ -320,6 +320,7 @@ describe('SessionTile workspace scope', () => { $selectedStoredSessionId.set('bot-chat') openSessionTile('bot-chat', 'center', undefined, undefined, scope) + $layoutTree.set(group(['workspace', tilePane('bot-chat')], { active: 'workspace', id: 'main' })) expect($sessionTiles.get()).toEqual([ expect.objectContaining({ @@ -337,6 +338,7 @@ describe('SessionTile workspace scope', () => { // new tip must front that tile, not open the same chat twice. setSessions([{ _lineage_ids: ['seg-1', 'seg-2', 'seg-3'], _lineage_root_id: 'seg-1', id: 'seg-3' } as never]) openSessionTile('seg-2') + $layoutTree.set(group(['workspace', tilePane('seg-2')], { active: 'workspace', id: 'main' })) expect(focusOpenSession('seg-3')).toBe('tile') expect($sessionTiles.get().map(t => t.storedSessionId)).toEqual(['seg-2']) @@ -487,6 +489,7 @@ describe('focusWorkspaceOwnerSessionTile', () => { openSessionTile('thread', 'center', 'workspace', undefined, botA) rememberActivePane(workspaceScopeKey('bots', 'bot:a'), tilePane('closed-bot-chat')) $sessionTiles.set($sessionTiles.get().filter(t => t.storedSessionId !== 'closed-bot-chat')) + $layoutTree.set(group(['workspace', tilePane('thread')], { active: 'workspace', id: 'main' })) expect(focusWorkspaceOwnerSessionTile('bot:a')).toBe('thread') }) @@ -533,6 +536,7 @@ describe('focusWorkspaceOwnerSessionTile', () => { it('a throwing probe keeps the tile — reconciliation must not break the click', () => { openSessionTile('bot-chat', 'center', 'workspace', undefined, botA) + $layoutTree.set(group(['workspace', tilePane('bot-chat')], { active: 'workspace', id: 'main' })) expect( focusWorkspaceOwnerSessionTile('bot:a', () => { @@ -542,8 +546,9 @@ describe('focusWorkspaceOwnerSessionTile', () => { expect($sessionTiles.get().map(t => t.storedSessionId)).toEqual(['bot-chat']) }) - it('no probe keeps the old behavior byte for byte', () => { + it('fronts a visible tile without a probe', () => { openSessionTile('bot-chat', 'center', 'workspace', undefined, botA) + $layoutTree.set(group(['workspace', tilePane('bot-chat')], { active: 'workspace', id: 'main' })) expect(focusWorkspaceOwnerSessionTile('bot:a')).toBe('bot-chat') expect($sessionTiles.get().map(t => t.storedSessionId)).toEqual(['bot-chat']) diff --git a/apps/desktop/src/store/session-states.ts b/apps/desktop/src/store/session-states.ts index fd9cb8cf60..afd217b79b 100644 --- a/apps/desktop/src/store/session-states.ts +++ b/apps/desktop/src/store/session-states.ts @@ -25,6 +25,7 @@ import { $activeTreeGroup, $layoutTree, focusedSessionTabAnchor, + isPaneVisible, moveTreePane, noteActiveTreeGroup, revealTreePane @@ -1551,10 +1552,12 @@ export function focusOpenSession( const tree = $layoutTree.get() const group = tree ? findGroupOfPane(tree, paneId) : null - if (group) { - noteActiveTreeGroup(group.id) + if (!group || !isPaneVisible(paneId)) { + return null } + noteActiveTreeGroup(group.id) + return 'tile' } @@ -1634,9 +1637,9 @@ export function focusWorkspaceOwnerSessionTile( const paneId = resolveRememberedActivePane(workspaceScopeKey('bots', workspaceOwnerKey), paneIds) ?? paneIds[0] const storedSessionId = paneId.slice(TILE_PANE_PREFIX.length) - focusOpenSession(storedSessionId, { workspaceMode: 'bots', workspaceOwnerKey }) - - return storedSessionId + return focusOpenSession(storedSessionId, { workspaceMode: 'bots', workspaceOwnerKey }) === 'tile' + ? storedSessionId + : null } /** Does a sidebar click still need to navigate after `focusOpenSession`? A miss diff --git a/contributors/emails/16833782+Fatmylin@users.noreply.github.com b/contributors/emails/16833782+Fatmylin@users.noreply.github.com new file mode 100644 index 0000000000..f608e6e066 --- /dev/null +++ b/contributors/emails/16833782+Fatmylin@users.noreply.github.com @@ -0,0 +1,2 @@ +Fatmylin +# PR #95964 salvage diff --git a/contributors/emails/210261288+Christopher-Schulze@users.noreply.github.com b/contributors/emails/210261288+Christopher-Schulze@users.noreply.github.com new file mode 100644 index 0000000000..e40f1c1a02 --- /dev/null +++ b/contributors/emails/210261288+Christopher-Schulze@users.noreply.github.com @@ -0,0 +1,2 @@ +Christopher-Schulze +# PR #85806 salvage diff --git a/contributors/emails/296402666+ciabata-git@users.noreply.github.com b/contributors/emails/296402666+ciabata-git@users.noreply.github.com new file mode 100644 index 0000000000..317632f3c5 --- /dev/null +++ b/contributors/emails/296402666+ciabata-git@users.noreply.github.com @@ -0,0 +1,2 @@ +ciabata-git +# PR #96011 author email preserved by PR #97330 diff --git a/contributors/emails/298902573+pierrenode@users.noreply.github.com b/contributors/emails/298902573+pierrenode@users.noreply.github.com new file mode 100644 index 0000000000..ef6b9888a4 --- /dev/null +++ b/contributors/emails/298902573+pierrenode@users.noreply.github.com @@ -0,0 +1,2 @@ +pierrenode +# PR #84168 salvage diff --git a/contributors/emails/JLHunzicker@gmail.com b/contributors/emails/JLHunzicker@gmail.com new file mode 100644 index 0000000000..865bba7f71 --- /dev/null +++ b/contributors/emails/JLHunzicker@gmail.com @@ -0,0 +1,2 @@ +JackHunzicker +# PR #98371 salvage diff --git a/contributors/emails/agent@dynamicagency.com b/contributors/emails/agent@dynamicagency.com new file mode 100644 index 0000000000..3696dc87bf --- /dev/null +++ b/contributors/emails/agent@dynamicagency.com @@ -0,0 +1,2 @@ +agentdynamic +# PR #92413 salvage diff --git a/contributors/emails/andrewwikel@gmail.com b/contributors/emails/andrewwikel@gmail.com new file mode 100644 index 0000000000..dd9535b42e --- /dev/null +++ b/contributors/emails/andrewwikel@gmail.com @@ -0,0 +1,2 @@ +slash1andy +# PR #88217 salvage diff --git a/contributors/emails/benjaminperry6@yahoo.fr b/contributors/emails/benjaminperry6@yahoo.fr new file mode 100644 index 0000000000..16b94399dd --- /dev/null +++ b/contributors/emails/benjaminperry6@yahoo.fr @@ -0,0 +1,2 @@ +benperry6 +# PR #97330 author email diff --git a/contributors/emails/csreyes92@gmail.com b/contributors/emails/csreyes92@gmail.com new file mode 100644 index 0000000000..6d1d2cf38e --- /dev/null +++ b/contributors/emails/csreyes92@gmail.com @@ -0,0 +1 @@ +csreyes diff --git a/contributors/emails/emanuele.cornaggia@gmail.com b/contributors/emails/emanuele.cornaggia@gmail.com new file mode 100644 index 0000000000..830fe71ba2 --- /dev/null +++ b/contributors/emails/emanuele.cornaggia@gmail.com @@ -0,0 +1,2 @@ +EmanueleCornaggia +# PR #101778 salvage diff --git a/contributors/emails/michel.alexander@gmail.com b/contributors/emails/michel.alexander@gmail.com new file mode 100644 index 0000000000..d990d4f5e5 --- /dev/null +++ b/contributors/emails/michel.alexander@gmail.com @@ -0,0 +1,2 @@ +SelfParody +# PR #19000 salvage diff --git a/contributors/emails/sosxradar@gmail.com b/contributors/emails/sosxradar@gmail.com new file mode 100644 index 0000000000..02d4ffb5d0 --- /dev/null +++ b/contributors/emails/sosxradar@gmail.com @@ -0,0 +1 @@ +GTHell diff --git a/contributors/emails/talmoredder@gmail.com b/contributors/emails/talmoredder@gmail.com new file mode 100644 index 0000000000..0e8b281e2c --- /dev/null +++ b/contributors/emails/talmoredder@gmail.com @@ -0,0 +1,2 @@ +EdderTalmor +# PR #10110 salvage diff --git a/contributors/emails/yj2761@nyu.edu b/contributors/emails/yj2761@nyu.edu new file mode 100644 index 0000000000..a73c993b32 --- /dev/null +++ b/contributors/emails/yj2761@nyu.edu @@ -0,0 +1,2 @@ +yaojiejia +# PR #95160 salvage diff --git a/cron/delivery_queue.py b/cron/delivery_queue.py new file mode 100644 index 0000000000..635a295b2f --- /dev/null +++ b/cron/delivery_queue.py @@ -0,0 +1,379 @@ +"""Profile-local durable handoff for cron delivery through live gateway adapters. + +A restart-safe cron worker executes outside the gateway cgroup. It cannot own +relay/E2EE adapter objects, so it queues the final send here. A gateway claims +each row at most once. If that gateway dies after claiming, the outcome is +marked unknown and never retried: losing a delivery is safer than duplicating a +possibly-completed send. +""" + +from __future__ import annotations + +import json +import logging +import os +import sqlite3 +import threading +import time +import uuid +from contextlib import contextmanager +from pathlib import Path +from typing import Any, Callable, Iterator, Optional + +from agent.redact import redact_sensitive_text +from cron.executions import _owner_is_live, _process_start_time +from hermes_cli.sqlite_util import add_column_if_missing +from hermes_constants import get_hermes_home +from hermes_time import now as _hermes_now + +logger = logging.getLogger(__name__) + +DELIVERY_DB: Optional[Path] = None +_PROCESS_ID = uuid.uuid4().hex +_lock = threading.RLock() +_ACTIVE_DELIVERIES: set[str] = set() +_TERMINAL = ("delivered", "failed", "unknown") +MAX_TERMINAL_DELIVERIES = 1000 +DEFAULT_DELIVERY_WAIT_TIMEOUT_SECONDS = 300.0 + + +def _prune_terminal_unlocked(conn: sqlite3.Connection) -> None: + """Redact terminal payloads and retain only bounded outcome metadata.""" + conn.execute( + """UPDATE deliveries SET job_json='{}', content='' + WHERE status IN ('delivered','failed','unknown') + AND (job_json != '{}' OR content != '')""" + ) + keep = max(0, int(MAX_TERMINAL_DELIVERIES)) + terminal_count = int( + conn.execute( + "SELECT COUNT(*) FROM deliveries " + "WHERE status IN ('delivered','failed','unknown')" + ).fetchone()[0] + ) + excess = terminal_count - keep + if excess > 0: + conn.execute( + """INSERT OR IGNORE INTO delivery_tombstones + (execution_id, terminal_status, finished_at) + SELECT execution_id, status, finished_at FROM deliveries + WHERE status IN ('delivered','failed','unknown') + ORDER BY finished_at, created_at, execution_id + LIMIT ?""", + (excess,), + ) + conn.execute( + """DELETE FROM deliveries WHERE execution_id IN ( + SELECT execution_id FROM deliveries + WHERE status IN ('delivered','failed','unknown') + ORDER BY finished_at, created_at, execution_id + LIMIT ? + )""", + (excess,), + ) + + +def _path() -> Path: + return DELIVERY_DB or (get_hermes_home().resolve() / "cron" / "deliveries.db") + + +@contextmanager +def _transaction() -> Iterator[sqlite3.Connection]: + with _lock: + path = _path() + path.parent.mkdir(parents=True, exist_ok=True) + conn = sqlite3.connect(path, timeout=5) + try: + path.chmod(0o600) + except OSError: + pass + conn.row_factory = sqlite3.Row + try: + from hermes_state import apply_wal_with_fallback + + conn.execute("PRAGMA busy_timeout=5000") + apply_wal_with_fallback(conn, db_label="cron/deliveries.db") + conn.execute("PRAGMA synchronous=FULL") + conn.execute( + """CREATE TABLE IF NOT EXISTS deliveries ( + execution_id TEXT PRIMARY KEY, + job_json TEXT NOT NULL, + content TEXT NOT NULL, + for_failure INTEGER NOT NULL DEFAULT 0, + status TEXT NOT NULL CHECK(status IN + ('pending','delivering','delivered','failed','unknown')), + owner_process_id TEXT, + owner_pid INTEGER, + owner_started_at INTEGER, + created_at TEXT NOT NULL, + finished_at TEXT, + error TEXT + )""" + ) + conn.execute( + """CREATE TABLE IF NOT EXISTS delivery_tombstones ( + execution_id TEXT PRIMARY KEY, + terminal_status TEXT NOT NULL CHECK(terminal_status IN + ('delivered','failed','unknown')), + finished_at TEXT + )""" + ) + add_column_if_missing( + conn, "deliveries", "for_failure", + "for_failure INTEGER NOT NULL DEFAULT 0", + ) + # Pruning is done explicitly by the paths that create terminal + # rows (_finish / recover_abandoned / _terminalize_wait_timeout); + # read-only polls must not pay for a full-table UPDATE + COUNT. + with conn: + yield conn + finally: + conn.close() + + +def enqueue( + execution_id: str, + job: dict, + content: str, + *, + for_failure: bool = False, +) -> dict: + """Persist one idempotent delivery request before the worker waits.""" + with _transaction() as conn: + tombstone = conn.execute( + "SELECT terminal_status, finished_at FROM delivery_tombstones " + "WHERE execution_id=?", + (str(execution_id),), + ).fetchone() + if tombstone is not None: + return { + "execution_id": str(execution_id), + "status": tombstone["terminal_status"], + "finished_at": tombstone["finished_at"], + } + conn.execute( + """INSERT OR IGNORE INTO deliveries + (execution_id, job_json, content, for_failure, status, created_at) + VALUES (?, ?, ?, ?, 'pending', ?)""", + ( + str(execution_id), + json.dumps(job, ensure_ascii=False, sort_keys=True), + str(content), + int(bool(for_failure)), + _hermes_now().isoformat(), + ), + ) + row = conn.execute( + "SELECT * FROM deliveries WHERE execution_id=?", (str(execution_id),) + ).fetchone() + return dict(row) + + +def get_status(execution_id: str) -> Optional[dict]: + with _transaction() as conn: + row = conn.execute( + "SELECT * FROM deliveries WHERE execution_id=?", (str(execution_id),) + ).fetchone() + if row is not None: + return dict(row) + tombstone = conn.execute( + "SELECT execution_id, terminal_status, finished_at " + "FROM delivery_tombstones WHERE execution_id=?", + (str(execution_id),), + ).fetchone() + if tombstone is None: + return None + return { + "execution_id": tombstone["execution_id"], + "status": tombstone["terminal_status"], + "finished_at": tombstone["finished_at"], + "error": None, + } + + +def claim_next() -> Optional[dict]: + """Atomically claim one pending send before touching the transport.""" + pid = os.getpid() + started = _process_start_time(pid) + with _transaction() as conn: + row = conn.execute( + "SELECT execution_id FROM deliveries WHERE status='pending' " + "ORDER BY created_at, execution_id LIMIT 1" + ).fetchone() + if row is None: + return None + cur = conn.execute( + """UPDATE deliveries SET status='delivering', owner_process_id=?, + owner_pid=?, owner_started_at=? + WHERE execution_id=? AND status='pending'""", + (_PROCESS_ID, pid, started, row["execution_id"]), + ) + if cur.rowcount != 1: + return None + claimed = conn.execute( + "SELECT * FROM deliveries WHERE execution_id=?", (row["execution_id"],) + ).fetchone() + _ACTIVE_DELIVERIES.add(row["execution_id"]) + result = dict(claimed) + result["job"] = json.loads(result.pop("job_json")) + return result + + +def _finish(execution_id: str, *, error: Optional[str]) -> bool: + status = "failed" if error else "delivered" + safe_error = ( + redact_sensitive_text(str(error), force=True, redact_url_credentials=True) + if error + else None + ) + with _transaction() as conn: + cur = conn.execute( + """UPDATE deliveries SET status=?, finished_at=?, error=? + WHERE execution_id=? AND status='delivering' + AND owner_process_id=? AND owner_pid=?""", + ( + status, + _hermes_now().isoformat(), + safe_error, + execution_id, + _PROCESS_ID, + os.getpid(), + ), + ) + _prune_terminal_unlocked(conn) + return cur.rowcount == 1 + + +def recover_abandoned() -> int: + """Fence dead delivery owners as unknown; never replay uncertain sends.""" + changed = 0 + with _transaction() as conn: + rows = conn.execute( + "SELECT execution_id, owner_process_id, owner_pid, owner_started_at " + "FROM deliveries WHERE status='delivering'" + ).fetchall() + for row in rows: + same_process = row["owner_process_id"] == _PROCESS_ID + if same_process: + with _lock: + if row["execution_id"] in _ACTIVE_DELIVERIES: + continue + elif _owner_is_live(int(row["owner_pid"]), row["owner_started_at"]): + continue + error = ( + "Gateway finished delivery but could not persist its outcome; " + "send was not retried." + if same_process + else "Gateway exited during delivery; send outcome is unknown and was not retried." + ) + cur = conn.execute( + """UPDATE deliveries SET status='unknown', finished_at=?, error=? + WHERE execution_id=? AND status='delivering'""", + ( + _hermes_now().isoformat(), + error, + row["execution_id"], + ), + ) + changed += cur.rowcount + _prune_terminal_unlocked(conn) + return changed + + +def drain( + send: Callable[[dict, str, bool], Optional[str]], *, limit: int = 20 +) -> int: + """Deliver pending rows through *send*, terminalizing every claimed row.""" + recover_abandoned() + processed = 0 + for _ in range(max(0, limit)): + row = claim_next() + if row is None: + break + with _lock: + _ACTIVE_DELIVERIES.add(row["execution_id"]) + try: + try: + error = send( + row["job"], row["content"], bool(row["for_failure"]) + ) + except BaseException as exc: + error = f"{type(exc).__name__}: {exc}" + _finish(row["execution_id"], error=error) + finally: + with _lock: + _ACTIVE_DELIVERIES.discard(row["execution_id"]) + processed += 1 + return processed + + +def _terminalize_wait_timeout(execution_id: str) -> str: + """Fence a delivery whose worker can no longer wait for confirmation. + + A row still ``pending`` was provably never attempted, so it is left queued + for whichever gateway comes up next (a restart that includes an update can + easily exceed the worker's wait budget). That is a deferral, not a + failure: report success so the job is not recorded ``delivery_failed`` for + a message the drain will still send. Only a row caught mid-send is + uncertain and gets fenced ``unknown``. + """ + now = _hermes_now().isoformat() + uncertain_error = ( + "timed out while gateway delivery was in progress; outcome is unknown and " + "was not retried" + ) + with _transaction() as conn: + row = conn.execute( + "SELECT status FROM deliveries WHERE execution_id=?", + (str(execution_id),), + ).fetchone() + if row is not None and row["status"] == "pending": + logger.warning( + "Cron delivery %s: no live gateway within the wait budget; " + "left queued for the next gateway", + execution_id, + ) + return "" + conn.execute( + """UPDATE deliveries SET status='unknown', finished_at=?, error=? + WHERE execution_id=? AND status='delivering'""", + (now, uncertain_error, str(execution_id)), + ) + row = conn.execute( + "SELECT status, error FROM deliveries WHERE execution_id=?", + (str(execution_id),), + ).fetchone() + _prune_terminal_unlocked(conn) + if row is None: + return "timed out waiting for live gateway delivery" + if row["status"] == "delivered": + return "" + return str(row["error"] or f"delivery {row['status']}") + + +def enqueue_and_wait( + execution_id: str, + job: dict, + content: str, + *, + for_failure: bool = False, + timeout: Optional[float] = None, +) -> Optional[str]: + """Queue delivery and wait for a gateway's terminal at-most-once outcome.""" + queued = enqueue(execution_id, job, content, for_failure=for_failure) + if queued["status"] in _TERMINAL: + return None if queued["status"] == "delivered" else str( + queued.get("error") or f"delivery {queued['status']}" + ) + wait_timeout = ( + DEFAULT_DELIVERY_WAIT_TIMEOUT_SECONDS if timeout is None else max(0.0, timeout) + ) + deadline = time.monotonic() + wait_timeout + while time.monotonic() < deadline: + row = get_status(execution_id) + if row and row["status"] in _TERMINAL: + return None if row["status"] == "delivered" else str( + row.get("error") or f"delivery {row['status']}" + ) + time.sleep(1.0) + return _terminalize_wait_timeout(execution_id) or None diff --git a/cron/executions.py b/cron/executions.py index 8c9fa24f0f..e7577155e0 100644 --- a/cron/executions.py +++ b/cron/executions.py @@ -10,6 +10,7 @@ from __future__ import annotations import os import sqlite3 import threading +import time import uuid from contextlib import contextmanager from pathlib import Path @@ -23,6 +24,7 @@ from hermes_time import now as _hermes_now # home. EXECUTIONS_FILE: Optional[Path] = None MAX_TERMINAL_EXECUTIONS = 1000 +HANDOFF_ADOPTION_GRACE_SECONDS = 30.0 _TERMINAL_STATES = ("completed", "failed", "unknown") _lock = threading.RLock() _PROCESS_ID = uuid.uuid4().hex @@ -88,12 +90,23 @@ def _initialize_schema(conn: sqlite3.Connection) -> None: process_started_at INTEGER, status TEXT NOT NULL CHECK(status IN ('claimed','running','completed','failed','unknown')), + handoff_pending INTEGER NOT NULL DEFAULT 0, + handoff_started_at REAL, claimed_at TEXT NOT NULL, started_at TEXT, finished_at TEXT, error TEXT )""" ) + from hermes_cli.sqlite_util import add_column_if_missing + + add_column_if_missing( + conn, "executions", "handoff_pending", + "handoff_pending INTEGER NOT NULL DEFAULT 0", + ) + add_column_if_missing( + conn, "executions", "handoff_started_at", "handoff_started_at REAL" + ) conn.execute( "CREATE INDEX IF NOT EXISTS idx_executions_job_claimed " "ON executions(job_id, claimed_at DESC, id DESC)" @@ -153,7 +166,7 @@ def _prune_unlocked(conn: sqlite3.Connection) -> None: """DELETE FROM executions WHERE id IN ( SELECT id FROM executions WHERE status IN ('completed','failed','unknown') - ORDER BY claimed_at DESC, id DESC LIMIT -1 OFFSET ? + ORDER BY finished_at DESC, claimed_at DESC, id DESC LIMIT -1 OFFSET ? )""", (max(0, int(MAX_TERMINAL_EXECUTIONS)),), ) @@ -178,14 +191,60 @@ def create_execution(job_id: str, *, source: str) -> Dict[str, Any]: return record # type: ignore[return-value] +def mark_execution_handoff_pending(execution_id: str) -> Optional[Dict[str, Any]]: + """Fence restart recovery while an external worker is adopting a claim.""" + with _transaction() as conn: + cur = conn.execute( + """UPDATE executions + SET handoff_pending=1, handoff_started_at=? + WHERE id=? AND status='claimed' + AND process_id=? AND pid=?""", + (time.time(), execution_id, _PROCESS_ID, os.getpid()), + ) + if cur.rowcount != 1: + return None + record = _fetch(conn, execution_id) + _emit_execution_state(record) + return record + + +def adopt_claimed_execution(execution_id: str) -> Optional[Dict[str, Any]]: + """Atomically transfer and start an attempt in its worker process. + + The dispatching gateway creates the row before spawning a restart-safe + worker. Adoption is the single ``claimed`` → ``running`` gate: only the + winner may acknowledge ownership or run side effects. + """ + pid = os.getpid() + process_started_at = _process_start_time(pid) + now = _hermes_now().isoformat() + with _transaction() as conn: + cur = conn.execute( + """UPDATE executions + SET process_id=?, pid=?, process_started_at=?, + status='running', started_at=?, handoff_pending=0, + handoff_started_at=NULL + WHERE id=? AND status='claimed' AND handoff_pending=1""", + (_PROCESS_ID, pid, process_started_at, now, execution_id), + ) + if cur.rowcount != 1: + return None + record = _fetch(conn, execution_id) + _emit_execution_state(record) + return record + + def mark_execution_running(execution_id: str) -> Optional[Dict[str, Any]]: """Transition one claimed attempt to running exactly once.""" now = _hermes_now().isoformat() with _transaction() as conn: cur = conn.execute( - """UPDATE executions SET status='running', started_at=? - WHERE id=? AND status='claimed'""", - (now, execution_id), + """UPDATE executions + SET status='running', started_at=?, handoff_pending=0, + handoff_started_at=NULL + WHERE id=? AND status='claimed' AND handoff_pending=0 + AND process_id=? AND pid=?""", + (now, execution_id, _PROCESS_ID, os.getpid()), ) if cur.rowcount != 1: return None @@ -204,9 +263,12 @@ def finish_execution( detail = None if success else (str(error) if error else "unknown failure") with _transaction() as conn: cur = conn.execute( - """UPDATE executions SET status=?, finished_at=?, error=? - WHERE id=? AND status IN ('claimed','running')""", - (status, now, detail, execution_id), + """UPDATE executions + SET status=?, finished_at=?, error=?, handoff_pending=0, + handoff_started_at=NULL + WHERE id=? AND status IN ('claimed','running') + AND process_id=? AND pid=?""", + (status, now, detail, execution_id, _PROCESS_ID, os.getpid()), ) if cur.rowcount != 1: return None @@ -223,7 +285,9 @@ def recover_interrupted_executions() -> int: recovered: List[Dict[str, Any]] = [] with _transaction() as conn: rows = conn.execute( - """SELECT id, process_id, pid, process_started_at FROM executions + """SELECT id, status, process_id, pid, process_started_at, + handoff_pending, handoff_started_at + FROM executions WHERE status IN ('claimed','running')""" ).fetchall() for row in rows: @@ -231,13 +295,26 @@ def recover_interrupted_executions() -> int: continue if _owner_is_live(int(row["pid"]), row["process_started_at"]): continue + handoff_started_at = row["handoff_started_at"] + if ( + row["handoff_pending"] + and handoff_started_at is not None + and time.time() - float(handoff_started_at) + < HANDOFF_ADOPTION_GRACE_SECONDS + ): + continue cur = conn.execute( - """UPDATE executions SET status='unknown', finished_at=?, error=? - WHERE id=? AND status IN ('claimed','running')""", + """UPDATE executions + SET status='unknown', finished_at=?, error=?, + handoff_pending=0, handoff_started_at=NULL + WHERE id=? AND status=? AND process_id=? AND pid=? + AND handoff_pending=? + AND handoff_started_at IS ?""", (now, "Scheduler restarted after this execution's owner exited before a durable " "terminal state; whether side effects ran is unknown.", - row["id"]), + row["id"], row["status"], row["process_id"], row["pid"], + row["handoff_pending"], row["handoff_started_at"]), ) changed += cur.rowcount if cur.rowcount: @@ -274,6 +351,16 @@ def list_executions( return [dict(row) for row in rows] +def get_execution(execution_id: str) -> Optional[Dict[str, Any]]: + """Return one exact execution attempt, or ``None`` when it is absent.""" + with _transaction() as conn: + row = conn.execute( + "SELECT * FROM executions WHERE id=?", + (str(execution_id),), + ).fetchone() + return dict(row) if row is not None else None + + def latest_execution(job_id: str) -> Optional[Dict[str, Any]]: rows = list_executions(job_id=job_id, limit=1) return rows[0] if rows else None diff --git a/cron/lifecycle_guard.py b/cron/lifecycle_guard.py index a27e24e15a..d68d1b8bb6 100644 --- a/cron/lifecycle_guard.py +++ b/cron/lifecycle_guard.py @@ -249,6 +249,95 @@ def contains_gateway_lifecycle_command(text: str) -> bool: return _contains_launchctl_gateway_lifecycle(normalized) +# Whole-walk work limits. The per-file cap and depth bound above limit one read, not the walk: a +# command can reference arbitrarily many scripts, and the pure-Python shlex pass (quadratic on a +# giant token) once held the GIL for minutes. These caps bound one whole walk and are charged +# BEFORE any text reaches shlex. Exhaustion fails closed (an unscanned script could hide a +# lifecycle command) and is logged at WARNING so an operator can tell it from a real block. Sizes +# sit well above any legitimate wrapper graph; remote reads are a backend roundtrip each, so they +# get a far tighter cap. +_MAX_LIFECYCLE_SCAN_BYTES = _MAX_REFERENCED_SCRIPT_BYTES # 1 MiB across the walk +_MAX_LIFECYCLE_SCAN_LINES = 16384 +_MAX_LIFECYCLE_SCAN_LINE_BYTES = 64 * 1024 +_MAX_LIFECYCLE_SCAN_PATHS = 1024 +_MAX_LIFECYCLE_SCAN_REMOTE_READS = 64 + + +class _LifecycleScanBudget: + """Shared work budget for one complete referenced-script walk.""" + + __slots__ = ("bytes_remaining", "lines_remaining", "paths_remaining", "remote_reads_remaining") + + def __init__(self) -> None: + # Read the module constants at construction so tests/operators can lower them at runtime. + self.bytes_remaining = _MAX_LIFECYCLE_SCAN_BYTES + self.lines_remaining = _MAX_LIFECYCLE_SCAN_LINES + self.paths_remaining = _MAX_LIFECYCLE_SCAN_PATHS + self.remote_reads_remaining = _MAX_LIFECYCLE_SCAN_REMOTE_READS + + def charge_text(self, text: str) -> bool: + """Charge *text* before tokenization; False when it does not fit.""" + # UTF-8 is >= one byte per code point, so the char count is a free lower bound. + if len(text) > self.bytes_remaining: + return False + encoded = len(text.encode("utf-8", errors="replace")) + if encoded > self.bytes_remaining: + return False + lines = text.count("\n") + 1 + if lines > self.lines_remaining: + return False + # One huge token is the quadratic shlex case; bound the longest physical line (chars, a + # lower bound on bytes — tight enough for a DoS bound without a per-line encode). + longest = max((len(line) for line in text.split("\n")), default=0) + if longest > _MAX_LIFECYCLE_SCAN_LINE_BYTES: + return False + self.bytes_remaining -= encoded + self.lines_remaining -= lines + return True + + def charge_path(self) -> bool: + """Charge one unique referenced path before any local/remote read.""" + if self.paths_remaining <= 0: + return False + self.paths_remaining -= 1 + return True + + def charge_remote_read(self) -> bool: + """Charge one remote-backend read (a network roundtrip each).""" + if self.remote_reads_remaining <= 0: + return False + self.remote_reads_remaining -= 1 + return True + + +def _capped_read_limit(max_bytes: Optional[int]) -> int: + """Per-read byte cap: never above the per-file cap, never negative. One definition so local and + remote reads cannot diverge.""" + if max_bytes is None: + return _MAX_REFERENCED_SCRIPT_BYTES + return min(_MAX_REFERENCED_SCRIPT_BYTES, max(0, int(max_bytes))) + + +def lifecycle_scan_root_within_budget(text: str) -> bool: + """Whether *text* may safely enter an optional tokenizer pass (``tools/terminal_tool.py`` gates + its launchctl pre-scan on this). A FRESH budget, independent of the full guard's walk: the + pre-scan may pass while the walk later exhausts, still fail-closed — only the friendlier + launchctl diagnostic is lost. ``False`` is not a verdict: callers must still run the full guard.""" + try: + return _LifecycleScanBudget().charge_text(text) + except Exception: + return False + + +def _budget_exhausted(what: str, depth: int) -> bool: + logger.warning( + "lifecycle guard scan budget exhausted (%s at depth %d); " + "failing closed — see _MAX_LIFECYCLE_SCAN_* in cron/lifecycle_guard.py", + what, depth, + ) + return True + + # --- shell tokenization ----------------------------------------------------------------------- def _split_logical_lines(text: str) -> list[str]: @@ -622,13 +711,17 @@ def _has_binary_magic(data: bytes) -> bool: return data.startswith(_BINARY_MAGICS) -def _read_referenced_script(path: Path) -> tuple[Optional[str], bool]: +def _read_referenced_script( + path: Path, *, max_bytes: Optional[int] = None +) -> tuple[Optional[str], bool]: """Return ``(text, unsafe)`` using bounded, regular-file-only reads. Shared choke point for every local script read, so the cloud-placeholder refusal lives here: a FileProvider path is never opened — not even to check hydration — because an evicted placeholder's ``open()`` can hang preflight. Lexical check: direct paths; resolved: symlinks. + ``max_bytes`` lowers the per-file cap to what the calling walk can still afford. """ + byte_limit = _capped_read_limit(max_bytes) if _on_cloud_path(path): return None, True flags = os.O_RDONLY | getattr(os, "O_NONBLOCK", 0) @@ -649,9 +742,13 @@ def _read_referenced_script(path: Path) -> tuple[Optional[str], bool]: data = os.read(descriptor, _BINARY_SNIFF_BYTES) if _has_binary_magic(data): return None, False + # A regular file whose size already exceeds the cap fails closed without reading it (the + # walk budget can be far below 1 MiB). + if metadata.st_size > byte_limit: + return None, True # Read the remainder (bounded); loop because os.read may return short. - while len(data) <= _MAX_REFERENCED_SCRIPT_BYTES: - chunk = os.read(descriptor, _MAX_REFERENCED_SCRIPT_BYTES + 1 - len(data)) + while len(data) <= byte_limit: + chunk = os.read(descriptor, byte_limit + 1 - len(data)) if not chunk: break data += chunk @@ -663,21 +760,26 @@ def _read_referenced_script(path: Path) -> tuple[Optional[str], bool]: return None, False # Size check BEFORE NUL stripping: stripping shrinks the buffer and would let an oversized file # slip under the threshold past this fail-closed branch. - if len(data) > _MAX_REFERENCED_SCRIPT_BYTES: + if len(data) > byte_limit: return None, True if b"\x00" in data: data = data.replace(b"\x00", b"") return data.decode("utf-8", errors="replace"), False -def _sanitize_remote_script_text(text: Optional[str]) -> tuple[Optional[str], bool]: +def _sanitize_remote_script_text( + text: Optional[str], *, max_bytes: Optional[int] = None +) -> tuple[Optional[str], bool]: """Apply the local-read contract to text from an untrusted ``read_remote_script`` callback: NUL means binary (nothing to scan, checked first); oversized fails closed. Size compares re-encoded *bytes* (matching the ``head -c`` wire bound): a >1 MiB multibyte file truncated at the byte cap decodes to fewer chars, and a char count would scan instead of failing.""" if not text or "\x00" in text: return None, False - if len(text.encode("utf-8", errors="replace")) > _MAX_REFERENCED_SCRIPT_BYTES: + byte_limit = _capped_read_limit(max_bytes) + if len(text) > byte_limit: + return None, True # chars <= bytes: over the cap without encoding + if len(text.encode("utf-8", errors="replace")) > byte_limit: return None, True return text, False @@ -698,9 +800,12 @@ def _read_script_for_scanning(script_path: str) -> str: # --- recursive walk --------------------------------------------------------------------------- def _contains_unsafe_gateway_action( - command: str, *, cwd: Optional[str], depth: int, visited: set[Path], + command: str, *, cwd: Optional[str], depth: int, visited: set[Path], budget: _LifecycleScanBudget, read_remote_script: Optional[_ReadRemoteScriptFn] = None, ) -> bool: + # Charge BEFORE _direct_lifecycle_scan: every scan in it tokenizes with shlex. + if not budget.charge_text(command): + return _budget_exhausted("text", depth) if _direct_lifecycle_scan(command): return True if depth >= _MAX_REFERENCED_SCRIPT_DEPTH: @@ -708,7 +813,8 @@ def _contains_unsafe_gateway_action( def recurse(text: str, cwd: Optional[str]) -> bool: return _contains_unsafe_gateway_action( - text, cwd=cwd, depth=depth + 1, visited=visited, read_remote_script=read_remote_script + text, cwd=cwd, depth=depth + 1, visited=visited, budget=budget, + read_remote_script=read_remote_script, ) for payload in _iter_shell_command_payloads(command): @@ -722,14 +828,22 @@ def _contains_unsafe_gateway_action( resolved = _resolve_lenient(script_path) if resolved in visited: continue + if not budget.charge_path(): + return _budget_exhausted("paths", depth) visited.add(resolved) - script_text, unsafe = _read_referenced_script(script_path) + # Never read more than the walk can still afford to tokenize; a file larger than the + # remainder fails closed exactly like an oversized one. + script_text, unsafe = _read_referenced_script(script_path, max_bytes=budget.bytes_remaining) if unsafe: return True if script_text is None and read_remote_script is not None: # Local path missing; the remote backend's output crosses the same trust boundary as a # local read — sanitize identically (binary skip + size fail-closed). - script_text, unsafe = _sanitize_remote_script_text(read_remote_script(str(script_path))) + if not budget.charge_remote_read(): + return _budget_exhausted("remote reads", depth) + script_text, unsafe = _sanitize_remote_script_text( + read_remote_script(str(script_path)), max_bytes=budget.bytes_remaining + ) if unsafe: return True if not script_text: @@ -752,7 +866,8 @@ def contains_gateway_lifecycle_command_or_referenced_script( """ try: return _contains_unsafe_gateway_action( - command, cwd=cwd, depth=0, visited=set(), read_remote_script=read_remote_script + command, cwd=cwd, depth=0, visited=set(), budget=_LifecycleScanBudget(), + read_remote_script=read_remote_script, ) except Exception: logger.warning( @@ -795,7 +910,11 @@ def check_gateway_lifecycle(prompt: Optional[str], script: Optional[str] = None) # Python runs via the interpreter, never a POSIX shell, and the shell reference walk is a # false-positive generator on Python sources (pathlib "/" resolves to the filesystem root). # The regex still scans the full text; non-regular/oversized files fail closed (sentinel). - unsafe = _lifecycle_command_scan_with_data_exemption(combined) + # The data-exemption masker tokenizes with shlex, so it is charged against the walk budget. + if not _LifecycleScanBudget().charge_text(combined): + unsafe = _budget_exhausted("text", 0) + else: + unsafe = _lifecycle_command_scan_with_data_exemption(combined) else: unsafe = contains_gateway_lifecycle_command_or_referenced_script( combined, cwd=_resolve_script_directory(script) if script else None diff --git a/cron/scheduler.py b/cron/scheduler.py index 2be5f564bc..0829788f97 100644 --- a/cron/scheduler.py +++ b/cron/scheduler.py @@ -434,7 +434,9 @@ from cron.jobs import ( _ensure_cron_dir, advance_next_runs, claim_dispatch, claim_job_for_fire, fire_claim_fence, clear_run_claim, get_due_jobs, heartbeat_fire_claim, heartbeat_run_claim, mark_job_run, save_job_output, use_cron_store) -from cron.executions import create_execution, finish_execution, mark_execution_running +from cron.executions import ( + _TERMINAL_STATES, create_execution, finish_execution, get_execution, + mark_execution_handoff_pending, mark_execution_running, recover_interrupted_executions) # Response marker that suppresses delivery (output is still saved locally for audit). SILENT_MARKER = "[SILENT]" @@ -453,6 +455,10 @@ _parallel_pool: Optional[concurrent.futures.ThreadPoolExecutor] = None _parallel_pool_max_workers: Optional[int] = None _running_job_ids: set = set() _running_fire_owners: dict[str, dict[object, tuple[Optional[str], Path]]] = {} +# Parent gateway threads synchronously waiting on restart-safe scope workers. +# Shutdown must not misclassify these as ownerless in-process runs: the tool +# process sweep cannot reach the worker's transient scope. +_restart_safe_waiter_job_ids: set[str] = set() _running_lock = threading.Lock() # Per in-flight id: time.time() claim instant + the future owning its release (``_FUTURE_PENDING`` @@ -781,9 +787,11 @@ def mark_running_jobs_interrupted( ``run_one_job`` sees them. """ with _running_lock: + restart_safe_waiters = set(_restart_safe_waiter_job_ids) active_fires = [ (token, job_id, owner, profile_home) for job_id, executions in _running_fire_owners.items() + if job_id not in restart_safe_waiters for token, (owner, profile_home) in executions.items() ] if only_owners is not None: @@ -792,7 +800,9 @@ def mark_running_jobs_interrupted( if only_owners is None: active_fires.extend( (None, job_id, None, _get_hermes_home()) - for job_id in _running_job_ids - registered_ids + for job_id in ( + _running_job_ids - registered_ids - restart_safe_waiters + ) ) _interrupted_job_ids.update( token if token is not None else job_id @@ -987,6 +997,26 @@ def _reclaim_fds_best_effort() -> None: apply_nofile_soft_limit(None) +def drain_delivery_queue(adapters, loop) -> int: + """Send queued worker results through this gateway's live adapters.""" + from cron.delivery_queue import _path, drain + + # Only restart-safe workers create the queue file. Every gateway (macOS, + # Windows, launchd, Docker) runs this housekeeping tick, so skip the sqlite + # open/create entirely until a worker has actually queued something. + if not _path().exists(): + return 0 + return drain( + lambda queued_job, queued_content, queued_for_failure: _deliver_result( + queued_job, + queued_content, + adapters=adapters, + loop=loop, + for_failure=queued_for_failure, + ) + ) + + _DEFAULT_SCRIPT_TIMEOUT = 3600 # seconds (1 hour) # Backward-compatible module override used by tests and emergency monkeypatches. _SCRIPT_TIMEOUT = _DEFAULT_SCRIPT_TIMEOUT @@ -2264,6 +2294,34 @@ def run_one_job( claim (callers use the store CAS) but keeps it alive. True if processed (a job failure is recorded via ``mark_job_run``), False only if processing raised. ``cancel_event``: optional transport-level cancel (dashboard drain).""" + # Every gateway path (built-in scheduler, external providers, and direct + # API fires) crosses this seam. Ensure the detached worker has a durable + # attempt to adopt before any launch can occur. + if not job.get("execution_id"): + execution = create_execution(job["id"], source="direct") + job["execution_id"] = execution["id"] + + execution_id = str(job["execution_id"]) + external_owner = os.environ.get("_HERMES_CRON_EXTERNAL_WORKER") == execution_id + if not external_owner: + try: + if _launch_external_cron_worker(job): + return True + except Exception as handoff_error: + error = f"Restart-safe cron worker dispatch failed: {handoff_error}" + logger.error("Job '%s': %s", job["id"], error) + claim = job.get("fire_claim") + owner = str(claim.get("by") or "") if isinstance(claim, dict) else "" + try: + mark_job_run( + job["id"], + False, + error, + **({"expected_fire_owner": owner} if owner else {}), + ) + finally: + finish_execution(execution_id, success=False, error=error) + return True if extra_prompt is None: # Gateway-forwarded manual run stamps its prompt on the job via trigger_job; the fire that # consumes the manual occurrence picks it up here. Single-fire: mark_job_run clears it. @@ -2613,7 +2671,13 @@ def _run_one_job_body( error="Dispatch claim rejected; execution was not started.") return True # not an error — already handled/removed - mark_execution_running(execution_id) + # Claimed durably before dispatch; becomes running only right before the actual run. + # Detached workers transition to running while adopting; in-process paths must win the + # claimed->running CAS here before any user script or agent side effect may begin. + external_owner = os.environ.get("_HERMES_CRON_EXTERNAL_WORKER") == execution_id + if not external_owner and mark_execution_running(execution_id) is None: + logger.warning("Cron job %s lost execution ownership before start; skipping", job["id"]) + return True # get_secret() fails closed outside a scope; the ticker thread has none. Delivery adapters # resolve credentials, so the scope must span delivery too (reset in the outer finally). @@ -2733,6 +2797,313 @@ def _run_one_job_body( reset_terminal_scope(_terminal_scope_token) +def _wait_for_external_cron_worker_body( + process: subprocess.Popen, + *, + execution_id: str, +) -> bool: + """Preserve ``run_one_job``'s synchronous contract after handoff. + + The worker owns the durable execution and survives this gateway process. + The caller nevertheless waits while it remains alive so manual/background + callers do not release their in-process guard or report stale job state. + A gateway replacement may kill this waiter; it does not kill the scoped + worker or change its ledger ownership. + """ + def _is_terminal() -> bool: + current = get_execution(execution_id) + return bool(current and current.get("status") in _TERMINAL_STATES) + + # The worker commits its terminal row before its process exits, so exit is + # the correct wakeup. Each ledger read opens a connection and re-runs + # schema init; polling it at 50ms for an hours-long agent run is ~72k + # opens/hour of pure contention with the worker's own writes. + while True: + try: + returncode = process.wait(timeout=1.0) + except subprocess.TimeoutExpired: + if _is_terminal(): + return True + continue + # The worker can commit its terminal row and exit between the first + # read and wait(). Re-read the exact attempt before declaring that + # it died without terminalizing. + if _is_terminal(): + return True + # If the adopted worker died without terminalizing, its owner is + # now provably gone. Recover to ``unknown`` rather than routing the + # exception through the pre-handoff dispatch-failure path, which + # would falsely assert that no side effect could have happened. + recover_interrupted_executions() + if _is_terminal(): + return True + raise RuntimeError( + "cron external worker exited before durable recovery could " + f"terminalize its execution state (exit {returncode})" + ) + + +def _wait_for_external_cron_worker( + process: subprocess.Popen, + *, + execution_id: str, + job_id: Optional[str] = None, + handoff_files: tuple[Path, ...] = (), +) -> bool: + try: + return _wait_for_external_cron_worker_body( + process, execution_id=execution_id + ) + finally: + if job_id is not None: + with _running_lock: + _restart_safe_waiter_job_ids.discard(job_id) + # The execution is terminal or its worker is dead: nobody will read a + # payload or acknowledgement left behind by a late/unread handoff. + for stale in handoff_files: + try: + stale.unlink(missing_ok=True) + except OSError: + pass + + +def _launch_external_cron_worker(job: dict) -> bool: + """Launch *job* outside a managed gateway cgroup when required. + + Returns ``False`` when the caller is not a managed systemd gateway and the + existing in-process path should be used. In managed topology, failure to + establish the transient scope raises: falling back would recreate the + restart interruption this handoff exists to prevent. + """ + execution_id = str(job["execution_id"]) + job_id = str(job["id"]) + handoff_dir = _get_hermes_home() / "cron" / "external-workers" + payload_path = handoff_dir / f"{execution_id}.json" + ack_path = handoff_dir / f"{execution_id}.ready" + command = [ + sys.executable, + "-m", + "cron.scheduler", + "--external-worker-file", + str(payload_path), + "--ack-file", + str(ack_path), + ] + + from agent.secret_scope import is_multiplex_active + from tools.environments.local import build_subprocess_env + from tools.process_registry import restart_safe_gateway_child_argv + + multiplex_active = is_multiplex_active() + scoped_command = restart_safe_gateway_child_argv( + command, + unit_suffix=f"cron-{job_id}-exec-{execution_id}", + ) + if scoped_command == command: + return False + + if mark_execution_handoff_pending(execution_id) is None: + raise RuntimeError( + "cron execution claim changed before external worker handoff" + ) + + _ensure_cron_dir(handoff_dir) + try: + handoff_dir.chmod(0o700) + except OSError: + pass + fd = os.open(payload_path, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600) + try: + with os.fdopen(fd, "w", encoding="utf-8") as payload_file: + json.dump( + { + "job": job, + "profile_home": str(_get_hermes_home().resolve()), + "multiplex_active": multiplex_active, + }, + payload_file, + ) + payload_file.flush() + os.fsync(payload_file.fileno()) + except BaseException: + payload_path.unlink(missing_ok=True) + raise + + worker_env = build_subprocess_env( + scrub_secrets=multiplex_active, + inherit_profile_home=True, + extra={"HERMES_HOME": str(_get_hermes_home().resolve())}, + ) + try: + process = subprocess.Popen( + scoped_command, + cwd=str(Path(__file__).resolve().parent.parent), + env=worker_env, + stdin=subprocess.DEVNULL, + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + start_new_session=True, + creationflags=windows_hide_flags(), + ) + except BaseException: + payload_path.unlink(missing_ok=True) + raise + + with _running_lock: + _restart_safe_waiter_job_ids.add(job_id) + + deadline = time.monotonic() + 5.0 + while time.monotonic() < deadline: + if ack_path.exists(): + try: + acknowledgement = json.loads(ack_path.read_text(encoding="utf-8")) + except Exception: + logger.exception( + "Cron external worker %s published an unreadable acknowledgement; " + "treating handoff as ownership-uncertain", + execution_id, + ) + return _wait_for_external_cron_worker( + process, + execution_id=execution_id, + job_id=job_id, + handoff_files=(payload_path,), + ) + finally: + ack_path.unlink(missing_ok=True) + if ( + not isinstance(acknowledgement, dict) + or acknowledgement.get("execution_id") != execution_id + ): + logger.error( + "Cron external worker acknowledgement mismatch for %s; " + "treating handoff as ownership-uncertain", + execution_id, + ) + return _wait_for_external_cron_worker( + process, + execution_id=execution_id, + job_id=job_id, + handoff_files=(payload_path,), + ) + logger.info( + "Cron job '%s' handed to restart-safe worker pid=%s execution=%s", + job_id, + acknowledgement.get("pid"), + execution_id, + ) + return _wait_for_external_cron_worker( + process, + execution_id=execution_id, + job_id=job_id, + handoff_files=(payload_path,), + ) + returncode = process.poll() + if returncode is not None: + with _running_lock: + _restart_safe_waiter_job_ids.discard(job_id) + payload_path.unlink(missing_ok=True) + raise RuntimeError( + f"cron external worker exited before ownership acknowledgement " + f"(exit {returncode})" + ) + time.sleep(0.05) + + # The child may have adopted the durable row just before publishing its + # acknowledgement. Never fall back to in-process execution on an uncertain + # handoff: that could duplicate side effects. The execution owner/dead-owner + # recovery ledger remains the authority. + logger.warning( + "Cron external worker for job '%s' did not acknowledge within 5s; " + "leaving the durable execution claim untouched", + job_id, + ) + return _wait_for_external_cron_worker( + process, + execution_id=execution_id, + job_id=job_id, + handoff_files=(payload_path, ack_path), + ) + + +def _run_external_worker_payload(payload_path: Path, ack_path: Path) -> bool: + """Adopt and execute one gateway-dispatched cron payload. + + The execution row is created by the gateway before spawn, then transferred + here before the ready acknowledgement is published. No side effect runs + unless that durable ownership transfer succeeds. + """ + try: + payload = json.loads(payload_path.read_text(encoding="utf-8")) + job = payload["job"] + profile_home = Path(payload["profile_home"]).resolve() + execution_id = str(job["execution_id"]) + except Exception: + logger.exception("Cron external worker could not load payload %s", payload_path) + return False + finally: + try: + payload_path.unlink(missing_ok=True) + except OSError: + pass + + from agent.secret_scope import ( + build_profile_secret_scope, + is_multiplex_active, + reset_secret_scope, + set_multiplex_active, + set_secret_scope, + ) + from cron.executions import adopt_claimed_execution + from hermes_cli.env_loader import hydrate_profile_secret_sources + from hermes_constants import ( + reset_hermes_home_override, + set_hermes_home_override, + ) + + home_token = set_hermes_home_override(profile_home) + previous_multiplex = is_multiplex_active() + multiplex_active = bool(payload.get("multiplex_active", False)) + set_multiplex_active(multiplex_active) + hydrate_profile_secret_sources(profile_home) + secret_token = set_secret_scope(build_profile_secret_scope(profile_home)) + try: + with use_cron_store(profile_home): + if adopt_claimed_execution(execution_id) is None: + logger.error( + "Cron external worker refused execution %s: durable ownership " + "could not be established", + execution_id, + ) + return False + try: + ack_path.parent.mkdir(parents=True, exist_ok=True) + fd = os.open(ack_path, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600) + with os.fdopen(fd, "w", encoding="utf-8") as ack_file: + json.dump({"pid": os.getpid(), "execution_id": execution_id}, ack_file) + ack_file.flush() + os.fsync(ack_file.fileno()) + except Exception: + logger.exception( + "Cron external worker could not publish ready acknowledgement for %s", + execution_id, + ) + return False + old_external_execution = os.environ.get("_HERMES_CRON_EXTERNAL_WORKER") + os.environ["_HERMES_CRON_EXTERNAL_WORKER"] = execution_id + try: + return run_one_job(job, adapters=None, loop=None, verbose=False) + finally: + if old_external_execution is None: + os.environ.pop("_HERMES_CRON_EXTERNAL_WORKER", None) + else: + os.environ["_HERMES_CRON_EXTERNAL_WORKER"] = old_external_execution + finally: + reset_secret_scope(secret_token) + set_multiplex_active(previous_multiplex) + reset_hermes_home_override(home_token) + + def _notify_provider_jobs_changed() -> None: """Best-effort: tell the active scheduler provider the job set changed. Call AFTER a successful store mutation so an external provider can re-provision/cancel the one-shot; no-op for the @@ -3174,10 +3545,6 @@ def tick( _release_tick_lock(lock_fd) -if __name__ == "__main__": - tick(verbose=True) - - # --------------------------------------------------------------------------- # Split modules — re-exported so ``scheduler.`` keeps resolving (and stays the single # monkeypatch target). @@ -3221,3 +3588,27 @@ from cron.scheduler_preflight import ( # noqa: E402,F401 _is_transient_provider_resolve_error, _preflight_check_delivery, _preflight_check_provider_key, _preflight_check_skills, _preflight_job_config, _primary_profile_routes_for_current_home, ) + + +# `python -m cron.scheduler` entry: MUST stay below the split-module re-exports so the worker / +# tick paths see every name the module re-exports. +if __name__ == "__main__": + if "--external-worker-file" in sys.argv: + import argparse + + parser = argparse.ArgumentParser(add_help=False) + parser.add_argument("--external-worker-file", type=Path, required=True) + parser.add_argument("--ack-file", type=Path, required=True) + args = parser.parse_args() + # The gateway spawns this worker with stdout/stderr on DEVNULL; without + # a handler every adoption/ack failure below would be invisible. + try: + from hermes_logging import setup_logging + + setup_logging(hermes_home=_get_hermes_home(), mode="cron") + except Exception: + pass + raise SystemExit( + 0 if _run_external_worker_payload(args.external_worker_file, args.ack_file) else 1 + ) + tick(verbose=True) diff --git a/cron/scheduler_delivery.py b/cron/scheduler_delivery.py index e4ab56e21f..3b5da860a2 100644 --- a/cron/scheduler_delivery.py +++ b/cron/scheduler_delivery.py @@ -1539,6 +1539,19 @@ def _deliver_result( targets = _sched._resolve_delivery_targets(job, for_failure=for_failure) if not targets: return _unresolved_delivery_outcome(job, for_failure) + + # Restart-safe workers have no live gateway adapters: hand the send back through a durable + # queue so the current or replacement gateway performs it with relay/E2EE parity. The execution + # id is the idempotency key (the queue never retries an uncertain claimed send). Match on THIS + # job's own attempt: a worker's script may dispatch another job in-process (`hermes cron run`), + # and that nested delivery must not be keyed under the outer execution id. + external_execution = os.environ.get("_HERMES_CRON_EXTERNAL_WORKER", "") + if (external_execution and adapters is None + and external_execution == str(job.get("execution_id") or "")): + from cron.delivery_queue import enqueue_and_wait + + return enqueue_and_wait(external_execution, job, content, for_failure=for_failure) + from gateway.config import load_gateway_config # Wrap with header/footer unless cron.wrap_response: false. diff --git a/docs/observability/relay-shared-metrics.md b/docs/observability/relay-shared-metrics.md index a736bf6edd..14f98bbb90 100644 --- a/docs/observability/relay-shared-metrics.md +++ b/docs/observability/relay-shared-metrics.md @@ -18,18 +18,22 @@ as a no-op compatibility alias for existing installation commands. > longer activate exporters. Without the new variable, Hermes does not run > Relay plugin discovery, configuration layering, middleware, or exporters. -Hermes requires NeMo Relay 0.7.1 or later within the 0.7 release line. That -release establishes the lossless provider-codec contract used for Anthropic -Messages, OpenAI Chat Completions, and OpenAI Responses requests. +Hermes requires NeMo Relay 0.8.3 or later within the 0.8 release line. That +line provides the provider-codec and canonical tool-result contracts Hermes +uses for managed provider and tool calls. ## Runtime Dependency and Data Boundary Hermes installs the platform-specific `nemo-relay` native wheel from the -bounded `>=0.7.1,<0.8` dependency range. The published package is built from +bounded `>=0.8.3,<0.9` dependency range. The published package is built from the [NVIDIA NeMo Relay repository](https://github.com/NVIDIA/NeMo-Relay). Unsupported platforms use the explicit no-op runtime described above rather than downloading a different implementation. +Operator-supplied typed native plugins must be rebuilt for Relay 0.8. `grpc-v1` +workers must be regenerated and rebuilt when they use tool callbacks, tool +execution intercepts, or manual tool-end APIs. + When Relay managed execution is active, the provider request and response pass through that native module in the Hermes process so configured interceptors can operate on the real call. This is separate from the shared-metrics data @@ -59,12 +63,12 @@ opt-in. Set `HERMES_NEMO_RELAY_PLUGINS_TOML` to a selected `plugins.toml` to activate configured middleware, exporters, or dynamic plugins. When the variable is unset, Hermes does not invoke Relay's plugin initializer, so Relay does not perform plugin configuration discovery or layering. When it is set -and the selected file loads successfully, Relay performs its normal static -`plugins.toml` discovery and layers the selected static configuration over the -discovered configuration. Dynamic `[[plugins.dynamic]]` records are loaded -from the selected file only. If the selected file cannot be loaded, Hermes -reports the error and does not invoke Relay initialization or fall back to -ambient discovery. +and the selected file loads successfully, Relay discovers supported user and +system `plugins.toml` files and layers the selected static configuration over +them. Repository-local `.nemo-relay/plugins.toml` files are ignored. Dynamic +`[[plugins.dynamic]]` records are loaded from the selected file only. If the +selected file cannot be loaded, Hermes reports the error and does not invoke +Relay initialization or fall back to ambient discovery. ## Session-Span Segmentation for Continuous Sessions diff --git a/evals/fanout_resource_bench.py b/evals/fanout_resource_bench.py new file mode 100644 index 0000000000..fac2bc928f --- /dev/null +++ b/evals/fanout_resource_bench.py @@ -0,0 +1,232 @@ +#!/usr/bin/env python3 +"""Fan-out resource benchmark for hermes-agent. + +Spawns N in-process child AIAgents via the REAL delegate_task code path +(tools.delegate_tool.delegate_task) against a local fake OpenAI server, with +children editing python files across W distinct git worktrees so the LSP +(pyright) path is exercised for real. Measures, for the host process: + + threads, RSS MB, open fds, TCP ESTAB sockets, child processes (pyright, + kernels), state.db growth, wall time. + +Usage: + python evals/fanout_resource_bench.py --repo --children 24 --worktrees 6 --label before + +Prints one JSON line; append several and compare with --compare a.json b.json. +""" +from __future__ import annotations + +import argparse +import http.server +import json +import os +import shutil +import socket +import subprocess +import sys +import tempfile +import threading +import time + +# -------------------------------------------------------------------------- +# Fake OpenAI chat-completions server: each child does +# turn 1: call write_file on /hermes_cli/bench_.py +# turn 2: call execute_code print(1) +# turn 3: final text +# -------------------------------------------------------------------------- +_REPLY_KB = [0] + + +class _Fake(http.server.BaseHTTPRequestHandler): + def log_message(self, format, *args): # quiet + pass + + def do_POST(self): + n = int(self.headers.get("Content-Length", 0)) + body = json.loads(self.rfile.read(n) or b"{}") + msgs = body.get("messages", []) + goal = next((m["content"] for m in msgs if m.get("role") == "user"), "") + try: + plan = json.loads(goal) + except Exception: + plan = {} + n_tool = sum(1 for m in msgs if m.get("role") == "tool") + if n_tool == 0 and plan.get("file"): + tc = {"id": "c1", "type": "function", "function": {"name": "write_file", "arguments": json.dumps({"path": plan["file"], "content": "import os\nx: int = 'bad'\n"})}} + msg = {"role": "assistant", "content": None, "tool_calls": [tc]} + finish = "tool_calls" + elif n_tool == 1 and plan.get("file"): + tc = {"id": "c2", "type": "function", "function": {"name": "execute_code", "arguments": json.dumps({"code": "print(1)"})}} + msg = {"role": "assistant", "content": None, "tool_calls": [tc]} + finish = "tool_calls" + else: + msg = {"role": "assistant", "content": "done " + ("x" * (_REPLY_KB[0] * 1024))} + finish = "stop" + if body.get("stream") is True: + self.send_response(200) + self.send_header("Content-Type", "text/event-stream") + self.end_headers() + delta = {"role": "assistant", "content": msg.get("content") or ""} + if msg.get("tool_calls"): + tc = msg["tool_calls"][0] + delta["tool_calls"] = [{"index": 0, "id": tc["id"], "type": "function", "function": tc["function"]}] + for chunk in ( + {"id": "m", "object": "chat.completion.chunk", "choices": [{"index": 0, "delta": delta, "finish_reason": None}]}, + {"id": "m", "object": "chat.completion.chunk", "choices": [{"index": 0, "delta": {}, "finish_reason": finish}], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}}, + ): + self.wfile.write(f"data: {json.dumps(chunk)}\n\n".encode()) + self.wfile.write(b"data: [DONE]\n\n") + self.wfile.flush() + return + resp = {"id": "x", "object": "chat.completion", "created": 0, "model": body.get("model", "m"), + "choices": [{"index": 0, "message": msg, "finish_reason": finish}], + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}} + data = json.dumps(resp).encode() + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(data))) + self.end_headers() + self.wfile.write(data) + + +def _serve(): + srv = http.server.ThreadingHTTPServer(("127.0.0.1", 0), _Fake) + threading.Thread(target=srv.serve_forever, daemon=True).start() + return srv + + +def _count_live(cls_name: str) -> int: + import gc + return sum(1 for o in gc.get_objects() if type(o).__name__ == cls_name) + + +def _snap(pid: int, db_path: str) -> dict: + st = open(f"/proc/{pid}/status", encoding="utf-8").read() + g = lambda k: int(st.split(k + ":")[1].split()[0]) + tcp = subprocess.run(f"ss -tanp 2>/dev/null | grep -c 'pid={pid},'", shell=True, capture_output=True, text=True).stdout.strip() + kids = subprocess.run(["ps", "-o", "args=", "--ppid", str(pid)], capture_output=True, text=True).stdout + return { + "threads": g("Threads"), "rss_mb": g("VmRSS") // 1024, "fds": len(os.listdir(f"/proc/{pid}/fd")), + "tcp": int(tcp or 0), "pyright": kids.count("pyright"), "kernels": kids.count("hermes_kernel_runner"), + "db_mb": round(os.path.getsize(db_path) / 2**20, 1) if os.path.exists(db_path) else 0, + "httpx_clients": _count_live("Client"), "transports": _count_live("HTTPTransport"), "session_dbs": _count_live("SessionDB"), "live_agents": _count_live("AIAgent"), + } + + +def main() -> None: + ap = argparse.ArgumentParser() + ap.add_argument("--repo", required=True) + ap.add_argument("--children", type=int, default=24) + ap.add_argument("--worktrees", type=int, default=6) + ap.add_argument("--label", default="") + ap.add_argument("--out", default="") + ap.add_argument("--compare", nargs=2) + ap.add_argument("--reply-kb", type=int, default=0, help="pad each child's final reply to N KB (transcript-size realism)") + a = ap.parse_args() + if a.compare: + b, c = (json.load(open(p, encoding="utf-8")) for p in a.compare) + print(f"| metric | {b['label']} | {c['label']} | delta |\n|---|---|---|---|") + for k in ("threads", "rss_mb", "fds", "tcp", "pyright", "kernels", "db_mb", "httpx_clients", "transports", "session_dbs"): + bv, cv = b["peak"][k], c["peak"][k] + print(f"| {k} (peak) | {bv} | {cv} | {cv - bv:+} |") + for k in ("rss_mb", "live_agents", "db_mb", "threads"): + bv, cv = b["after"].get(k), c["after"].get(k) + if bv is not None and cv is not None: + print(f"| {k} (after, children done) | {bv} | {cv} | {cv - bv:+} |") + print(f"| wall_s | {b['wall_s']} | {c['wall_s']} | {c['wall_s'] - b['wall_s']:+.1f} |") + return + + _REPLY_KB[0] = a.reply_kb + home = tempfile.mkdtemp(prefix="hermes_bench_home_") + os.environ["HERMES_HOME"] = home + os.environ["TERMINAL_ENV"] = "local" + os.environ.pop("OPENROUTER_API_KEY", None) + sys.path.insert(0, a.repo) + os.chdir(a.repo) + pyright = shutil.which("pyright-langserver", path=os.path.expanduser("~/.hermes/lsp/bin") + os.pathsep + os.environ.get("PATH", "")) + with open(os.path.join(home, "config.yaml"), "w", encoding="utf-8") as f: + f.write("lsp:\n enabled: true\n wait_timeout: 5.0\n install_strategy: manual\n") + if pyright: + f.write(f" servers:\n pyright:\n command: [{json.dumps(pyright)}, \"--stdio\"]\n") + f.write("delegation:\n max_concurrent_children: 64\n subagent_auto_approve: true\n") + + # W git worktrees, each a real python project (pyproject + package) so pyright roots resolve. + wts = [] + base = tempfile.mkdtemp(prefix="hermes_bench_wt_") + for w in range(a.worktrees): + d = os.path.join(base, f"wt{w}") + os.makedirs(os.path.join(d, "hermes_cli")) + subprocess.run(["git", "init", "-q", d], check=True) + open(os.path.join(d, "pyproject.toml"), "w", encoding="utf-8").write("[project]\nname='b'\n") + open(os.path.join(d, "hermes_cli", "__init__.py"), "w", encoding="utf-8").write("") + wts.append(d) + + srv = _serve() + port = srv.server_address[1] + from run_agent import AIAgent + from tools import delegate_tool + + from hermes_state import SessionDB + db_path = os.path.join(home, "state.db") + from pathlib import Path + session_db = SessionDB(db_path=Path(db_path)) + parent = AIAgent(api_key="bench", base_url=f"http://127.0.0.1:{port}/v1", model="bench-model", + quiet_mode=True, skip_context_files=True, skip_memory=True, + enabled_toolsets=["delegation", "file", "code_execution"], + session_db=session_db, session_id="bench-root") + # Children reference parent_session_id; the parent row is normally created + # lazily on the parent's first turn, which this harness never runs. + parent._ensure_db_session() + pid = os.getpid() + before = _snap(pid, db_path) + peak = dict(before) + stop = threading.Event() + + def sampler(): + while not stop.wait(0.5): + s = _snap(pid, db_path) + for k, v in s.items(): + peak[k] = max(peak[k], v) + threading.Thread(target=sampler, daemon=True).start() + + tasks = [{"goal": json.dumps({"file": os.path.join(wts[i % len(wts)], "hermes_cli", f"bench_{i}.py")}), + "context": "bench"} for i in range(a.children)] + t0 = time.monotonic() + res = delegate_tool.delegate_task(tasks=tasks, parent_agent=parent, background=False) + wall = round(time.monotonic() - t0, 1) + if os.environ.get("BENCH_DEBUG"): + sys.__stderr__.write(str(res)[:3000] + "\n") + time.sleep(2.0) + stop.set() + import gc + gc.collect() + after = _snap(pid, db_path) + try: + parsed = json.loads(res) + items = parsed if isinstance(parsed, list) else parsed.get("results") or parsed.get("tasks") or [] + ok = sum(1 for r in items if str(r.get("status", "")) in ("completed", "success")) + except Exception: + ok = None + try: + session_db.checkpoint() if hasattr(session_db, "checkpoint") else None + except Exception: + pass + out = {"label": a.label, "children": a.children, "worktrees": a.worktrees, "ok": ok, "wall_s": wall, + "before": before, "peak": peak, "after": after} + sys.__stderr__.write("BENCH " + json.dumps(out) + "\n"); sys.__stderr__.flush() + if a.out: + open(a.out, "w", encoding="utf-8").write(json.dumps(out, indent=1)) + try: + from agent.lsp import shutdown_service + shutdown_service() + from tools.code_kernel import shutdown_all_kernels + shutdown_all_kernels() + except Exception: + pass + shutil.rmtree(base, ignore_errors=True) + os._exit(0) + + +if __name__ == "__main__": + main() diff --git a/gateway/platforms/base.py b/gateway/platforms/base.py index eab21d3b78..de4a3982cd 100644 --- a/gateway/platforms/base.py +++ b/gateway/platforms/base.py @@ -423,7 +423,7 @@ sys.path.insert(0, str(Path(__file__).resolve().parents[2])) from gateway.config import Platform, PlatformConfig from gateway.platforms.helpers import fence_state_after -from gateway.session import SessionSource, build_session_key +from gateway.session import SessionSource, TranscriptReadError, build_session_key from hermes_constants import get_default_hermes_root, get_hermes_dir, get_hermes_home if TYPE_CHECKING: @@ -599,6 +599,11 @@ def cache_image_from_bytes(data: bytes, ext: str = ".jpg") -> str: return _write_cache_file(get_image_cache_dir(), "img", ext, data) +async def cache_image_from_bytes_async(data: bytes, ext: str = ".jpg") -> str: + """Cache image bytes without blocking the caller's event loop.""" + return await asyncio.to_thread(cache_image_from_bytes, data, ext) + + async def _cache_media_from_url(url: str, ext: str, retries: int, *, media_type: str, accept: str, cache_fn, log_label: str) -> str: """Shared downloader behind ``cache_*_from_url``: SSRF-checked (pre-flight + per-redirect; @@ -616,7 +621,7 @@ async def _cache_media_from_url(url: str, ext: str, retries: int, *, media_type: async with client.stream("GET", url, headers=headers) as response: response.raise_for_status() content = await _read_httpx_body_with_limit(response, media_type=media_type) - return cache_fn(content, ext) + return await asyncio.to_thread(cache_fn, content, ext) except (httpx.TimeoutException, httpx.HTTPStatusError) as exc: if isinstance(exc, httpx.HTTPStatusError) and exc.response.status_code < 429: raise @@ -662,6 +667,11 @@ def cache_audio_from_bytes(data: bytes, ext: str = ".ogg") -> str: return _write_cache_file(get_audio_cache_dir(), "audio", sniff_audio_ext(data, ext), data) +async def cache_audio_from_bytes_async(data: bytes, ext: str = ".ogg") -> str: + """Cache audio bytes without blocking the caller's event loop.""" + return await asyncio.to_thread(cache_audio_from_bytes, data, ext) + + async def cache_audio_from_url(url: str, ext: str = ".ogg", retries: int = 2) -> str: """Download an audio URL into the audio cache; return the absolute path.""" return await _cache_media_from_url( @@ -685,6 +695,11 @@ def cache_video_from_bytes(data: bytes, ext: str = ".mp4") -> str: return _write_cache_file(get_video_cache_dir(), "video", ext, data) +async def cache_video_from_bytes_async(data: bytes, ext: str = ".mp4") -> str: + """Cache video bytes without blocking the caller's event loop.""" + return await asyncio.to_thread(cache_video_from_bytes, data, ext) + + # Document / screenshot cache utilities (same pattern; referenced by local path). DOCUMENT_CACHE_DIR = get_hermes_dir("cache/documents", "document_cache") SCREENSHOT_CACHE_DIR = get_hermes_dir("cache/screenshots", "browser_screenshots") @@ -1272,6 +1287,11 @@ def cache_document_from_bytes(data: bytes, filename: str) -> str: return str(filepath) +async def cache_document_from_bytes_async(data: bytes, filename: str) -> str: + """Cache document bytes without blocking the caller's event loop.""" + return await asyncio.to_thread(cache_document_from_bytes, data, filename) + + # Unified media caching: classify attachment bytes by ext/MIME, route to cache_*_from_bytes. @dataclass class CachedMedia: @@ -1333,6 +1353,23 @@ def cache_media_bytes(data: bytes, *, filename: str = "", mime_type: str = "", return CachedMedia(to_agent_visible_cache_path(path), out_mime, "document", display or fallback_name) +async def cache_media_bytes_async( + data: bytes, + *, + filename: str = "", + mime_type: str = "", + default_kind: Optional[str] = None, +) -> Optional[CachedMedia]: + """Classify and cache attachment bytes without blocking the event loop.""" + return await asyncio.to_thread( + cache_media_bytes, + data, + filename=filename, + mime_type=mime_type, + default_kind=default_kind, + ) + + class MessageType(Enum): """Types of incoming messages.""" TEXT = "text" @@ -1887,11 +1924,38 @@ class BasePlatformAdapter(ABC): def set_fatal_error_handler(self, handler: Callable[["BasePlatformAdapter"], Awaitable[None] | None]) -> None: self._fatal_error_handler = handler + #: Published when an adapter is installed and running but its receive + #: path is not yet confirmed (e.g. Telegram polling has not proven a + #: getUpdates round-trip). Same ``retrying`` platform_state the runner + #: uses for queued reconnects, so readers see "not delivering" (#101391). + DEGRADED_STATUS_MESSAGE = "connected but not yet confirmed active; recovering in background" + + @property + def send_path_degraded(self) -> bool: + """True while connect() succeeded but delivery is not confirmed. + + Adapters with a separately-proven receive path override this; the + default adapter is either connected or not. + """ + return False + def _mark_connected(self) -> None: self._running = True self._fatal_error_code = self._fatal_error_message = None self._fatal_error_retryable = True - self._write_runtime_status_safe("connected", platform_state="connected", error_code=None, error_message=None) + if self.send_path_degraded: + self._mark_degraded() + else: + self._write_runtime_status_safe("connected", platform_state="connected", error_code=None, error_message=None) + + def _mark_degraded(self) -> None: + """Publish ``retrying`` for a running adapter whose delivery path is unproven.""" + self._write_runtime_status_safe( + "connected_degraded", + platform_state="retrying", + error_code=None, + error_message=self.DEGRADED_STATUS_MESSAGE, + ) def _mark_disconnected(self) -> None: self._running = False @@ -2172,6 +2236,12 @@ class BasePlatformAdapter(ABC): peek = getattr(store, "peek_session_id", None) session_id = peek(session_key) if callable(peek) else None transcript = store.load_transcript(session_id or session_key) + except TranscriptReadError: + logger.warning( + "Transcript read failed for session %s; media dedup runs " + "with no history this turn (#100788)", session_key, + ) + return None except Exception: return None if not transcript: diff --git a/gateway/platforms/bluebubbles.py b/gateway/platforms/bluebubbles.py index fbf33f29de..8a0bead2ea 100644 --- a/gateway/platforms/bluebubbles.py +++ b/gateway/platforms/bluebubbles.py @@ -20,7 +20,7 @@ from gateway.config import Platform, PlatformConfig from gateway.platforms._shared import get_scoped_secret as _get_scoped_secret from gateway.platforms.base import ( BasePlatformAdapter, MessageEvent, MessageType, SendResult, - cache_image_from_bytes, cache_audio_from_bytes, cache_document_from_bytes) + cache_image_from_bytes_async, cache_audio_from_bytes_async, cache_document_from_bytes_async) from .media_cache import ext_for_mime from gateway.platforms.helpers import compile_mention_patterns, strip_markdown from utils import TRUTHY_STRINGS @@ -471,11 +471,11 @@ class BlueBubblesAdapter(BasePlatformAdapter): data = resp.content mime = (att_meta.get("mimeType") or "").lower() if mime.startswith("image/"): - return cache_image_from_bytes(data, _closed_ext(mime, _BLUEBUBBLES_IMAGE_EXT_OVERRIDES, ".jpg")) + return await cache_image_from_bytes_async(data, _closed_ext(mime, _BLUEBUBBLES_IMAGE_EXT_OVERRIDES, ".jpg")) if mime.startswith("audio/"): - return cache_audio_from_bytes(data, _closed_ext(mime, _BLUEBUBBLES_AUDIO_EXT_OVERRIDES, ".mp3")) + return await cache_audio_from_bytes_async(data, _closed_ext(mime, _BLUEBUBBLES_AUDIO_EXT_OVERRIDES, ".mp3")) # Videos, documents, and everything else - return cache_document_from_bytes(data, att_meta.get("transferName", "") or f"file_{uuid.uuid4().hex[:8]}") + return await cache_document_from_bytes_async(data, att_meta.get("transferName", "") or f"file_{uuid.uuid4().hex[:8]}") except Exception as exc: logger.warning("[bluebubbles] failed to download attachment %s: %s", _redact(att_guid), exc) return None diff --git a/gateway/platforms/qqbot/adapter.py b/gateway/platforms/qqbot/adapter.py index 5c10c86292..5148e25980 100644 --- a/gateway/platforms/qqbot/adapter.py +++ b/gateway/platforms/qqbot/adapter.py @@ -40,7 +40,7 @@ except ImportError: from gateway.config import Platform, PlatformConfig from gateway.platforms.base import ( gateway_trust_env, BasePlatformAdapter, MessageEvent, MessageType, SendResult, - _ssrf_redirect_guard, cache_document_from_bytes, cache_image_from_bytes) + _ssrf_redirect_guard, cache_document_from_bytes_async, cache_image_from_bytes_async) from gateway.platforms.helpers import strip_markdown from gateway.platforms.media_cache import ext_for_mime @@ -976,12 +976,12 @@ class QQAdapter(BasePlatformAdapter): if content_type.startswith("image/"): # Historical qqbot mapping: trust mimetypes' guess (never the shared table), fall back to .jpg. ext = ext_for_mime(content_type, use_defaults=False, use_mimetypes=True, fallback=".jpg") or ".jpg" - return cache_image_from_bytes(data, ext) + return await cache_image_from_bytes_async(data, ext) if content_type == "voice" or content_type.startswith("audio/"): # QQ voice is usually .amr/.silk — convert to .wav for STT engines. return await self._convert_audio_to_wav(data, url) filename = original_name or Path(urlparse(url).path).name or "qq_attachment" - return cache_document_from_bytes(data, filename) + return await cache_document_from_bytes_async(data, filename) @staticmethod def _is_voice_content_type(content_type: str, filename: str) -> bool: @@ -1245,16 +1245,16 @@ class QQAdapter(BasePlatformAdapter): try: if not await convert(src_path, wav_path): logger.warning("[%s] audio conversion failed for %s (format=%s)", self._log_tag, source_url[:60], ext) - return cache_document_from_bytes(audio_data, f"qq_voice{ext}") + return await cache_document_from_bytes_async(audio_data, f"qq_voice{ext}") except Exception: - return cache_document_from_bytes(audio_data, f"qq_voice{ext}") + return await cache_document_from_bytes_async(audio_data, f"qq_voice{ext}") finally: self._unlink_quiet(src_path) try: wav_data = Path(wav_path).read_bytes() os.unlink(wav_path) - return cache_document_from_bytes(wav_data, "qq_voice.wav") + return await cache_document_from_bytes_async(wav_data, "qq_voice.wav") except Exception as exc: logger.debug("[%s] Failed to read converted wav: %s", self._log_tag, exc) return None diff --git a/gateway/platforms/signal.py b/gateway/platforms/signal.py index c50650ab24..3883cd00eb 100644 --- a/gateway/platforms/signal.py +++ b/gateway/platforms/signal.py @@ -25,8 +25,8 @@ import httpx from gateway.config import Platform, PlatformConfig from gateway.platforms.base import ( - BasePlatformAdapter, MessageEvent, MessageType, ProcessingOutcome, SendResult, cache_image_from_bytes, - cache_audio_from_bytes, cache_document_from_bytes, cache_image_from_url, utf16_len) + BasePlatformAdapter, MessageEvent, MessageType, ProcessingOutcome, SendResult, cache_image_from_bytes_async, + cache_audio_from_bytes_async, cache_document_from_bytes_async, cache_image_from_url, utf16_len) from gateway.platforms.helpers import redact_phone from gateway.platforms.media_cache import mime_for_ext from tools.audio_container import CONTAINER_TO_EXT, sniff_container @@ -592,9 +592,9 @@ class SignalAdapter(BasePlatformAdapter): # to .m4a. Without ffmpeg the raw file is cached as-is (no downstream remux fallback). if ext == ".aac": raw_data, ext = (await asyncio.to_thread(_remux_aac_to_m4a, raw_data)) or (raw_data, ext) - cache = (cache_image_from_bytes if _is_image_ext(ext) - else cache_audio_from_bytes if _is_audio_ext(ext) else cache_document_from_bytes) - return cache(raw_data, ext), ext + cache = (cache_image_from_bytes_async if _is_image_ext(ext) + else cache_audio_from_bytes_async if _is_audio_ext(ext) else cache_document_from_bytes_async) + return await cache(raw_data, ext), ext async def _rpc(self, method: str, params: dict, rpc_id: str = None, *, log_failures: bool = True, raise_on_rate_limit: bool = False, timeout: float = 30.0) -> Any: diff --git a/gateway/platforms/weixin.py b/gateway/platforms/weixin.py index adc9d3efa9..3a70663617 100644 --- a/gateway/platforms/weixin.py +++ b/gateway/platforms/weixin.py @@ -8,7 +8,7 @@ import asyncio, base64, contextlib, hashlib, json, logging, mimetypes, os, re, s from datetime import datetime from functools import partial from pathlib import Path -from typing import Any, Callable, Dict, List, Optional, Tuple +from typing import Any, Awaitable, Callable, Dict, List, Optional, Tuple from urllib.parse import quote, urlparse logger = logging.getLogger(__name__) @@ -31,7 +31,7 @@ from gateway.config import Platform, PlatformConfig from gateway.platforms.helpers import MessageDeduplicator, greedy_pack_blocks from gateway.platforms.base import ( _IMAGE_EXTS, _VIDEO_EXTS, gateway_trust_env, BasePlatformAdapter, MessageEvent, MessageType, SendResult, - cache_audio_from_bytes, cache_document_from_bytes, cache_image_from_bytes, + cache_audio_from_bytes_async, cache_document_from_bytes_async, cache_image_from_bytes_async, ) from hermes_constants import get_hermes_home from utils import atomic_json_write @@ -175,6 +175,10 @@ class ContextTokenStore: def __init__(self, hermes_home: str): self._root = _account_dir(hermes_home) self._cache: Dict[str, str] = {} + # Serializes the offloaded flushes so two concurrent set() calls + # cannot land their writes out of order (last-writer-wins would drop + # the newer token from disk). + self._persist_lock = asyncio.Lock() @staticmethod def _key(account_id: str, user_id: str) -> str: @@ -197,10 +201,16 @@ class ContextTokenStore: def get(self, account_id: str, user_id: str) -> Optional[str]: return self._cache.get(self._key(account_id, user_id)) - def set(self, account_id: str, user_id: str, token: str) -> None: + async def set(self, account_id: str, user_id: str, token: str) -> None: self._cache[self._key(account_id, user_id)] = token - prefix = f"{account_id}:" - payload = {key[len(prefix):]: value for key, value in self._cache.items() if key.startswith(prefix)} + # atomic_json_write() fsyncs, so the flush is offloaded off the loop; the payload is snapshotted + # here (the worker never iterates ``_cache`` mid-mutation) and the lock keeps flushes in order. + async with self._persist_lock: + prefix = f"{account_id}:" + payload = {key[len(prefix):]: value for key, value in self._cache.items() if key.startswith(prefix)} + await asyncio.to_thread(self._persist, account_id, payload) + + def _persist(self, account_id: str, payload: Dict[str, str]) -> None: try: atomic_json_write(self._root / f"{account_id}.context-tokens.json", payload) except Exception as exc: @@ -649,11 +659,11 @@ _video_item, _voice_item = partial(_media_item, ITEM_VIDEO), partial(_media_item # Inbound media dispatch: item type -> (item key, download timeout, cache fn, mime or None (= guess from # file_name), log label). Cache fns are lambdas so monkeypatching the module names takes effect at call time. -_INBOUND_MEDIA: Dict[int, Tuple[str, float, Callable[[bytes, str], str], Optional[str], str]] = { - ITEM_IMAGE: ("image_item", 30.0, lambda data, _name: cache_image_from_bytes(data, ".jpg"), "image/jpeg", "image"), - ITEM_VIDEO: ("video_item", 120.0, lambda data, _name: cache_document_from_bytes(data, "video.mp4"), "video/mp4", "video"), - ITEM_FILE: ("file_item", 60.0, lambda data, name: cache_document_from_bytes(data, name), None, "file"), - ITEM_VOICE: ("voice_item", 60.0, lambda data, _name: cache_audio_from_bytes(data, ".silk"), "audio/silk", "voice"), +_INBOUND_MEDIA: Dict[int, Tuple[str, float, Callable[[bytes, str], Awaitable[str]], Optional[str], str]] = { + ITEM_IMAGE: ("image_item", 30.0, lambda data, _name: cache_image_from_bytes_async(data, ".jpg"), "image/jpeg", "image"), + ITEM_VIDEO: ("video_item", 120.0, lambda data, _name: cache_document_from_bytes_async(data, "video.mp4"), "video/mp4", "video"), + ITEM_FILE: ("file_item", 60.0, lambda data, name: cache_document_from_bytes_async(data, name), None, "file"), + ITEM_VOICE: ("voice_item", 60.0, lambda data, _name: cache_audio_from_bytes_async(data, ".silk"), "audio/silk", "voice"), } # Outbound local-file dispatch by extension: (extensions, sender method, path kwarg); default = send_document. @@ -870,7 +880,7 @@ class WeixinAdapter(BasePlatformAdapter): return context_token = str(message.get("context_token") or "").strip() if context_token: - self._token_store.set(self._account_id, sender_id, context_token) + await self._token_store.set(self._account_id, sender_id, context_token) if self._poll_session and self._token and not self._typing_cache.get(sender_id): asyncio.create_task(self._fetch_typing_ticket(self._poll_session, sender_id, context_token or None, "getConfig failed")) media_paths, media_types = [], [] # type: List[str], List[str] @@ -951,7 +961,7 @@ class WeixinAdapter(BasePlatformAdapter): data = await _download_and_decrypt_media( self._poll_session, cdn_base_url=self._cdn_base_url, encrypted_query_param=media.get("encrypt_query_param"), aes_key_b64=aes_key_b64, full_url=media.get("full_url"), timeout_seconds=timeout_seconds) - return cache_fn(data, filename), mime + return await cache_fn(data, filename), mime except Exception as exc: logger.warning("[%s] %s download failed: %s", self.name, label, exc) return None, mime diff --git a/gateway/platforms/yuanbao.py b/gateway/platforms/yuanbao.py index 82bb67a007..3e42248dcb 100644 --- a/gateway/platforms/yuanbao.py +++ b/gateway/platforms/yuanbao.py @@ -44,7 +44,7 @@ except ImportError: from gateway.config import Platform, PlatformConfig from gateway.platforms.base import ( BasePlatformAdapter, MessageEvent, MessageType, SendResult, - cache_document_from_bytes, cache_image_from_bytes, cache_video_from_bytes, + cache_document_from_bytes_async, cache_image_from_bytes_async, cache_video_from_bytes_async, ) from gateway.platforms import helpers as _mdchunk from gateway.platforms._shared import get_scoped_secret as _yb_secret @@ -62,7 +62,7 @@ from gateway.platforms.yuanbao_proto import ( encode_send_private_heartbeat, encode_send_group_heartbeat, encode_query_group_info, encode_get_group_member_list, next_seq_no, ) -from gateway.session import build_session_key +from gateway.session import TranscriptReadError, build_session_key logger = logging.getLogger(__name__) @@ -613,6 +613,14 @@ class RecallGuardMiddleware(InboundMiddleware): await asyncio.sleep(0.5) try: transcript = store.load_transcript(sid) + except TranscriptReadError as exc: + # No readable rows means nothing to redact; polling on + # would just re-log the same failure (#100788). + logger.warning( + "[%s] Recall redact: transcript unreadable for " + "session %s: %s", adapter.name, sid, exc, + ) + return except Exception: continue for entry in transcript: @@ -638,6 +646,11 @@ class RecallGuardMiddleware(InboundMiddleware): return try: transcript = store.load_transcript(sid) + except TranscriptReadError as exc: + # Not an empty transcript — the rows are unreadable, so recall has + # nothing to match against (#100788). + logger.warning("[%s] Recall: transcript unreadable: %s", adapter.name, exc) + return except Exception as exc: logger.warning("[%s] Recall: failed to load transcript: %s", adapter.name, exc) return @@ -1200,6 +1213,13 @@ class QuoteContextMiddleware(InboundMiddleware): if isinstance(_content, str) and "|ybres:" in _content: media_refs.extend(_iter_ybres_refs(_YB_RES_REF_RE.finditer(_content))) break + except TranscriptReadError as exc: + # Quote resolution degrades to "no refs" rather than pretending + # the quoted message was never seen (#100788). + logger.warning( + "[%s] quote transcript lookup: transcript unreadable: %s", + getattr(adapter, "name", "yuanbao"), exc, + ) except Exception as exc: logger.warning("[%s] quote transcript lookup failed: %s", getattr(adapter, "name", "yuanbao"), exc) return media_refs @@ -1422,7 +1442,7 @@ class MediaResolveMiddleware(InboundMiddleware): if kind == "image": ext = cls._guess_image_ext_from_url(fetch_url) try: - local_path = cache_image_from_bytes(file_bytes, ext=ext) + local_path = await cache_image_from_bytes_async(file_bytes, ext=ext) except ValueError as exc: logger.warning("[%s] inbound image cache rejected: %s err=%s", adapter.name, log_tag, exc) return None @@ -1431,12 +1451,12 @@ class MediaResolveMiddleware(InboundMiddleware): mime = content_type if content_type.startswith("image/") else "image/jpeg" elif kind == "video": # Yuanbao video resources carry no reliable extension; default to mp4. - local_path = cache_video_from_bytes(file_bytes) + local_path = await cache_video_from_bytes_async(file_bytes) mime = guess_mime_type(local_path) or (content_type if content_type.startswith("video/") else "video/mp4") else: # file file_name = file_name or os.path.basename(urllib.parse.urlparse(fetch_url).path) or "file" try: - local_path = cache_document_from_bytes(file_bytes, file_name) + local_path = await cache_document_from_bytes_async(file_bytes, file_name) except Exception as exc: logger.warning("[%s] inbound file cache failed: %s err=%s", adapter.name, log_tag, exc) return None @@ -1532,6 +1552,10 @@ class MediaResolveMiddleware(InboundMiddleware): return [], [] try: history = store.load_transcript(store.get_or_create_session(source).session_id) + except TranscriptReadError as exc: + # Hydrate nothing rather than silently acting as if the session had no observed media. + logger.warning("[%s] Observed-media hydration: transcript unreadable: %s", adapter.name, exc) + return [], [] except Exception as exc: logger.warning("[%s] Observed-media hydration setup failed: %s", adapter.name, exc) return [], [] diff --git a/gateway/run.py b/gateway/run.py index 26a64de9f3..0ef86b40dc 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -31,7 +31,7 @@ from typing import Callable, Dict, Optional, Any, List, Tuple, cast from agent.async_utils import safe_schedule_threadsafe from agent.conversation_compression import ( - COMPACTION_DONE_STATUS, COMPACTION_STATUS, COMPRESSION_RETRY_CONTEXT_REDUCED_STATUS_TEMPLATE, + COMPACTION_DONE_STATUS, COMPACTION_HEARTBEAT_STATUS, COMPACTION_STATUS, COMPRESSION_RETRY_CONTEXT_REDUCED_STATUS_TEMPLATE, COMPRESSION_RETRY_MESSAGES_STATUS_TEMPLATE, COMPRESSION_RETRY_TOKENS_STATUS_TEMPLATE, COMPRESSION_RETRY_TOO_LARGE_STATUS_TEMPLATE, IDLE_COMPACTION_STATUS_TEMPLATE, PRE_API_COMPRESSION_STATUS_TEMPLATE, PREFLIGHT_COMPRESSION_STATUS_TEMPLATE) @@ -71,7 +71,8 @@ _TELEGRAM_NOISY_STATUS_RE = re.compile( r"|auto-lowered\s+(?:this\s+)?session'?s?\s+threshold" r"|configured\s+auxiliary\s+compression\s+provider\s+.+\s+unavailable" r"|skipping\s+concurrent\s+compression" - r"|compacting\s+context\s+[—-]\s+summarizing\s+earlier\s+conversation" + rf"|{re.escape(COMPACTION_STATUS)}" + rf"|{re.escape(COMPACTION_HEARTBEAT_STATUS)}" r"|resumed\s+after\s+\d+s\s+idle\s+[—-]\s+compacting" r"|preflight\s+compression" r"|pre[- ]api\s+compression" @@ -286,7 +287,7 @@ _COMPRESSION_PROGRESS_STATUS_RE = re.compile( "|".join( _status_template_to_regex(_template) for _template in ( - COMPACTION_STATUS, COMPACTION_DONE_STATUS, PRE_API_COMPRESSION_STATUS_TEMPLATE, + COMPACTION_STATUS, COMPACTION_HEARTBEAT_STATUS, COMPACTION_DONE_STATUS, PRE_API_COMPRESSION_STATUS_TEMPLATE, PREFLIGHT_COMPRESSION_STATUS_TEMPLATE, IDLE_COMPACTION_STATUS_TEMPLATE, COMPRESSION_RETRY_TOO_LARGE_STATUS_TEMPLATE, COMPRESSION_RETRY_MESSAGES_STATUS_TEMPLATE, COMPRESSION_RETRY_TOKENS_STATUS_TEMPLATE, @@ -4162,11 +4163,43 @@ def _housekeeping_memory_trim() -> None: trim_memory(reason="messaging gateway housekeeping") -def _start_gateway_housekeeping(stop_event: threading.Event, adapters=None, loop=None, interval: int = 60, cron_provider=None): +def _drain_restart_safe_cron_deliveries(adapters, loop, runner=None) -> None: + """Drain each profile's worker queue through its matching live adapters. A credential-less satellite + profile (empty adapter map) drains through the primary's adapters routed by its own profile routes.""" + from cron import scheduler as cron_scheduler + + if runner is None: + if adapters is not None: + cron_scheduler.drain_delivery_queue(adapters, loop) + return + for profile_name, profile_home in _handoff_watch_scopes(runner): + if profile_name is None: + profile_adapters = adapters + else: + profile_adapters = getattr(runner, "_profile_adapters", {}).get(profile_name) + if profile_adapters is None: + continue + with _profile_runtime_scope(profile_home or get_hermes_home()): + if profile_name is not None and not profile_adapters and adapters: + routes = cron_scheduler._primary_profile_routes_for_current_home() + if routes: + profile_adapters = cron_scheduler.SharedRouteAdapters(adapters, routes) + cron_scheduler.drain_delivery_queue(profile_adapters, loop) + + +def _start_gateway_housekeeping( + stop_event: threading.Event, adapters=None, loop=None, interval: int = 60, cron_provider=None, runner=None, +): """Background thread for gateway-only periodic chores (NOT cron). Separate from the cron trigger so chores run under any ``CronScheduler`` provider (external scale-to-zero has no 60s loop). Cadences are ticks of ``interval``; inner gates own the real cadence.""" - chores: list[tuple[int, str, Any]] = [ + chores: list[tuple[int, str, Any]] = [] + if adapters is not None or runner is not None: + # Restart-safe cron workers run outside the gateway cgroup and queue their final send for + # whichever gateway is live; drained here (not the scheduler tick) so external providers get it too. + chores.append((1, "Cron durable delivery queue drain", + lambda: _drain_restart_safe_cron_deliveries(adapters, loop, runner))) + chores += [ (5, "Channel directory refresh", lambda: adapters and _housekeeping_channel_directory(adapters, loop)), (60, "Media cache cleanup", _housekeeping_media_caches), (60, "Paste sweep", _housekeeping_paste_sweep)] @@ -4694,7 +4727,7 @@ def _start_gateway_start_cron_and_housekeeping(runner): housekeeping_thread = threading.Thread( target=_start_gateway_housekeeping, args=(cron_stop,), kwargs={"adapters": runner.adapters, "loop": asyncio.get_running_loop(), - "cron_provider": cron_provider}, + "cron_provider": cron_provider, "runner": runner}, daemon=True, name="gateway-housekeeping") housekeeping_thread.start() return cron_stop, cron_provider, cron_thread, housekeeping_thread diff --git a/gateway/run_adapters.py b/gateway/run_adapters.py index 3ed61e4cbe..2e619d83d0 100644 --- a/gateway/run_adapters.py +++ b/gateway/run_adapters.py @@ -679,11 +679,19 @@ class GatewayAdapterLifecycleMixin: self._publish_primary_adapter(platform, adapter) self.delivery_router.adapters = self.adapters del self._failed_platforms[platform] + # connect() returning True does not mean the receive path is confirmed -- Telegram's degraded + # reconnect returns True so the gateway stays up while its own ladder retries. Stamping "connected" + # here would undo the adapter's accurate status. + _degraded = adapter.send_path_degraded self._update_platform_runtime_status( - platform.value, platform_state="connected", error_code=None, error_message=None, + platform.value, platform_state="retrying" if _degraded else "connected", error_code=None, + error_message=adapter.DEGRADED_STATUS_MESSAGE if _degraded else None, needs_attention=False, retrying_since=None, ) - logger.info("✓ %s reconnected successfully", platform.value) + if _degraded: + logger.info("⚠ %s reconnected in degraded mode (receive path not yet confirmed)", platform.value) + else: + logger.info("✓ %s reconnected successfully", platform.value) # Responses rejected while down are owned by this live process (startup recovery cannot claim them). with _log_suppressed( logging.DEBUG, "failed-obligation redelivery after %s reconnect failed", diff --git a/gateway/run_goals.py b/gateway/run_goals.py index 8d532e529a..4b451316e2 100644 --- a/gateway/run_goals.py +++ b/gateway/run_goals.py @@ -365,7 +365,12 @@ class GatewayGoalsMixin: return mgr = LoopManager(session_id=sid) - wakeup = mgr.fire_tick() if mgr.is_due(now) else None + if not mgr.is_due(now): + return + # fire_tick()/complete_tick() are writes (BEGIN IMMEDIATE) taking the SessionDB writer lock; a slow + # writer elsewhere holding it while the loop thread blocked froze the gateway until the watchdog + # fired. The context-preserving executor keeps the profile HERMES_HOME override under multiplex. + wakeup = await self._run_in_executor_with_context(mgr.fire_tick) if not wakeup: return try: @@ -378,7 +383,7 @@ class GatewayGoalsMixin: # Slash-command loops dispatch through the command path and never hit the post-turn # completion hook — complete the tick immediately (caps + scheduling). if wakeup.lstrip().startswith("/"): - mgr.complete_tick("") + await self._run_in_executor_with_context(mgr.complete_tick, "") except Exception as exc: logger.warning("loop wakeup injection failed for %s: %s", sid, exc) with suppress(Exception): @@ -398,8 +403,10 @@ class GatewayGoalsMixin: # Warm once per scan: the scan reads every persisted loop and a cold cache would # run the state.db init on the loop thread before the first read. await self._warm_goals_session_db("loop wakeup") + # Off-loop too: the read is lock-free under WAL but convoys on the writer lock without it. + active_loops = await self._run_in_executor_with_context(list_active_loops) now = time.time() - for sid, state in list_active_loops(): + for sid, state in active_loops: await self._loop_wakeup_fire_one(sid, state, now, warned_no_route) except Exception as exc: logger.debug("loop wakeup watcher error: %s", exc) diff --git a/gateway/run_notifications.py b/gateway/run_notifications.py index 494fc934a7..7928eb27a7 100644 --- a/gateway/run_notifications.py +++ b/gateway/run_notifications.py @@ -708,13 +708,24 @@ class GatewayNotificationsMixin: error = getattr(self, "_session_db_init_error", None) if not error: return - from hermes_state import classify_persistence_error, format_session_db_unavailable + from hermes_state import _default_db_path, classify_persistence_error, format_session_db_unavailable if classify_persistence_error(error) == "corrupt": + # Copy-pasteable, so name the real store (profiles / HERMES_HOME do not live under ~/.hermes). + db_path = _default_db_path() message = ( - "⚠️ Session database corruption detected. Messages may not be persisted. Recovery " - "options:\n1. Run `hermes doctor --fix`\n2. Salvage with: sqlite3 ~/.hermes/state.db " - "\".recover\" (then replace state.db)\n3. Restore from a backup in ~/.hermes/backups/\nRun " - "`hermes doctor` for sanitized diagnostics." + "⚠️ Session database corruption detected. Messages may not be " + "persisted. Recovery options:\n" + "1. Run `hermes doctor --fix`\n" + "2. Stop the gateway, then recover with:\n" + f" hermes sessions recover --source {db_path} " + "--inspect-only\n" + " (if it reports recoverable) hermes sessions recover " + f"--source {db_path} --output recovered-state.db\n" + " — recovery snapshots the damaged file first; do NOT run " + "`sqlite3 ... \".recover\"` against the live state.db, a " + "vulnerable sqlite3 CLI can corrupt it further\n" + "3. Restore from a backup in ~/.hermes/backups/\n" + "Run `hermes doctor` for sanitized diagnostics." ) else: message = ( diff --git a/gateway/run_startup.py b/gateway/run_startup.py index 3ff27838ee..0ee6bd31e7 100644 --- a/gateway/run_startup.py +++ b/gateway/run_startup.py @@ -972,10 +972,13 @@ class GatewayStartupMixin: if outcome == "ok": self._publish_primary_adapter(platform, adapter) connected_count += 1 + # connect() may return True on a degraded (unconfirmed) receive path; don't stamp "connected". + _degraded = adapter.send_path_degraded self._update_platform_runtime_status( - platform.value, platform_state="connected", error_code=None, error_message=None, + platform.value, platform_state="retrying" if _degraded else "connected", error_code=None, + error_message=adapter.DEGRADED_STATUS_MESSAGE if _degraded else None, ) - logger.info("\u2713 %s connected", platform.value) + logger.info("\u2713 %s connected%s", platform.value, " (degraded)" if _degraded else "") continue # outcome == "failed" logger.warning("\u2717 %s failed to connect", platform.value) diff --git a/gateway/slash_commands.py b/gateway/slash_commands.py index ceae7abc77..19801e1a11 100644 --- a/gateway/slash_commands.py +++ b/gateway/slash_commands.py @@ -22,13 +22,13 @@ from typing import Optional, Union from agent.i18n import t from gateway.config import HomeChannel, Platform, PlatformConfig, persist_home_channel from gateway.platforms.base import EphemeralReply, MessageEvent -from gateway.session import AsyncSessionStore +from gateway.session import AsyncSessionStore, TranscriptReadError from gateway.slash_commands_goals import GatewayGoalCommandsMixin from gateway.slash_commands_model import ( # noqa: F401 — _model_switch_skew_guard re-exported for tests GatewayModelCommandsMixin, _model_switch_skew_guard) from gateway.slash_commands_session import GatewaySessionCommandsMixin -from gateway.slash_commands_status import GatewayStatusCommandsMixin +from gateway.slash_commands_status import HISTORY_UNREADABLE, GatewayStatusCommandsMixin # noqa: F401 — re-exported from hermes_cli.config import atomic_config_write, cfg_get from utils import atomic_json_write, is_truthy_value @@ -768,7 +768,10 @@ class GatewaySlashCommandsMixin( return t("gateway.btw.usage") source = event.source session_entry = await self.async_session_store.get_or_create_session(source) - history = await self.async_session_store.load_transcript(session_entry.session_id) + try: + history = await self.async_session_store.load_transcript(session_entry.session_id) + except TranscriptReadError: + return HISTORY_UNREADABLE if not history: return t("gateway.btw.no_history") try: diff --git a/gateway/slash_commands_session.py b/gateway/slash_commands_session.py index 7ca4a9cee2..8e3c7764d4 100644 --- a/gateway/slash_commands_session.py +++ b/gateway/slash_commands_session.py @@ -18,7 +18,8 @@ from agent.i18n import t from agent.turn_context import extract_api_content_sidecar from gateway.config import Platform from gateway.platforms.base import EphemeralReply, MessageEvent, MessageType -from gateway.session import SessionSource, build_session_key, is_shared_multi_user_session +from gateway.session import SessionSource, TranscriptReadError, build_session_key, is_shared_multi_user_session +from gateway.slash_commands_status import HISTORY_UNREADABLE logger = logging.getLogger("gateway.run") # log-record parity with gateway/run.py @@ -378,7 +379,10 @@ class GatewaySessionCommandsMixin: source = event.source session_entry = await self.async_session_store.get_or_create_session(source) - history = await self.async_session_store.load_transcript(session_entry.session_id) + try: + history = await self.async_session_store.load_transcript(session_entry.session_id) + except TranscriptReadError: + return HISTORY_UNREADABLE last_user_idx = next((i for i in range(len(history) - 1, -1, -1) if user_originated_turn_view(history[i]) is not None), None) if last_user_idx is None: @@ -484,7 +488,10 @@ class GatewaySessionCommandsMixin: source = event.source session_entry = await self.async_session_store.get_or_create_session(source) - history = await self.async_session_store.load_transcript(session_entry.session_id) + try: + history = await self.async_session_store.load_transcript(session_entry.session_id) + except TranscriptReadError: + return HISTORY_UNREADABLE if not history or len(history) < 4: return t("gateway.compress.not_enough") # Flags are stripped before positional parsing so they coexist with the boundary-aware @@ -899,7 +906,11 @@ class GatewaySessionCommandsMixin: # provider cached _session_id at initialize() and would keep writing to the wrong session. self._evict_cached_agent(session_key) title = await self._session_db.get_session_title(target_id) or name - history = await self.async_session_store.load_transcript(target_id) + try: + history = await self.async_session_store.load_transcript(target_id) + except TranscriptReadError: + # The resume itself succeeded; only the count is missing — say so rather than "empty". + return t("gateway.resume.resumed_no_count", title=title) + "\n" + HISTORY_UNREADABLE msg_count = len([m for m in history if m.get("role") == "user"]) if history else 0 if source.platform == Platform.MATRIX and allow_cross_room: msg_part = f" ({msg_count} message{'s' if msg_count != 1 else ''})" if msg_count else "" @@ -993,7 +1004,10 @@ class GatewaySessionCommandsMixin: source = event.source session_key = self._session_key_for_source(source) current_entry = await self.async_session_store.get_or_create_session(source) - history = await self.async_session_store.load_transcript(current_entry.session_id) + try: + history = await self.async_session_store.load_transcript(current_entry.session_id) + except TranscriptReadError: + return HISTORY_UNREADABLE if not history: return t("gateway.branch.no_conversation") new_session_id = f"{_dt.now().strftime('%Y%m%d_%H%M%S')}_{_uuid.uuid4().hex[:6]}" diff --git a/gateway/slash_commands_status.py b/gateway/slash_commands_status.py index 5693fe6291..90915903cb 100644 --- a/gateway/slash_commands_status.py +++ b/gateway/slash_commands_status.py @@ -15,6 +15,7 @@ from agent.account_usage import fetch_account_usage, render_account_usage_lines from agent.i18n import t from gateway.config import Platform from gateway.platforms.base import MessageEvent +from gateway.session import TranscriptReadError # Log-record parity with gateway/run.py and the origin module. logger = logging.getLogger("gateway.run") @@ -66,6 +67,10 @@ async def _quiet(call, default=None): return default +HISTORY_UNREADABLE = ("⚠️ Conversation history is unreadable (state.db). " + "This is not a new conversation — earlier messages exist but cannot be loaded.") + + def _quiet_sync(call, default=None): """Sync twin of ``_quiet``.""" try: @@ -328,7 +333,10 @@ class GatewayStatusCommandsMixin: breakdown = await asyncio.to_thread(self._context_breakdown_block, agent, source, expanded) if agent else [] return "\n".join(lines + ([""] + breakdown if breakdown else [])) # Last resort: rough estimate from transcript - history = await self.async_session_store.load_transcript(session_entry.session_id) + try: + history = await self.async_session_store.load_transcript(session_entry.session_id) + except TranscriptReadError: + return HISTORY_UNREADABLE if not history: return t("gateway.context.no_data") approx, count = _transcript_estimate(history) @@ -440,7 +448,10 @@ class GatewayStatusCommandsMixin: Runs in a thread; returns [] and never raises.""" try: from agent.context_breakdown import compute_context_details, render_context_breakdown_lines - payload = self._session_context_breakdown(agent, source) + try: + payload = self._session_context_breakdown(agent, source) + except TranscriptReadError: + return [HISTORY_UNREADABLE] # a read failure is not an empty transcript if not (payload.get("categories") or []): return [] details = _quiet_sync(lambda: compute_context_details(agent), {"skills": [], "toolsets": []}) if expanded else None @@ -449,16 +460,25 @@ class GatewayStatusCommandsMixin: return [] def _session_context_breakdown(self, agent, source) -> dict: - """Per-category context estimate (chars/4) for *agent* over the session transcript (sync).""" + """Per-category context estimate (chars/4) for *agent* over the session transcript (sync). + Raises ``TranscriptReadError`` (unreadable rows must not pass as an empty transcript).""" from agent.context_breakdown import compute_session_context_breakdown store = self.session_store - history = _quiet_sync(lambda: store.load_transcript(store.get_or_create_session(source).session_id) or [], []) + try: + history = store.load_transcript(store.get_or_create_session(source).session_id) or [] + except TranscriptReadError: + raise + except Exception: + history = [] return compute_session_context_breakdown(agent, history) def _context_breakdown_lines(self, agent, source) -> list[str]: """/usage per-category context breakdown (chars/4 estimate). Returns [] and never raises.""" try: - payload = self._session_context_breakdown(agent, source) + try: + payload = self._session_context_breakdown(agent, source) + except TranscriptReadError: + return [HISTORY_UNREADABLE] categories = payload.get("categories") or [] if not categories: return [] @@ -542,7 +562,10 @@ class GatewayStatusCommandsMixin: # No agent at all -- rough count from session history session_entry = await self.async_session_store.get_or_create_session(source) - history = await self.async_session_store.load_transcript(session_entry.session_id) + try: + history = await self.async_session_store.load_transcript(session_entry.session_id) + except TranscriptReadError: + return HISTORY_UNREADABLE if history: approx, count = _transcript_estimate(history) return _with_account_blocks([ diff --git a/hermes_cli/active_sessions.py b/hermes_cli/active_sessions.py index 57c774ffa3..b61cabbf7a 100644 --- a/hermes_cli/active_sessions.py +++ b/hermes_cli/active_sessions.py @@ -19,7 +19,7 @@ from dataclasses import dataclass from pathlib import Path from typing import Any, Iterator, Optional -from hermes_constants import get_hermes_home +from hermes_constants import get_default_hermes_root, get_hermes_home logger = logging.getLogger(__name__) @@ -579,20 +579,37 @@ def transfer_active_session( return True -def release_orphaned_leases(live_lease_ids: set[str]) -> int: - """Drop this process's registry entries that no live session owns. +# A lease this process wrote in the last few seconds may not be in the caller's +# ``own_live_lease_ids`` yet: ``try_acquire_active_session`` writes the registry entry under +# the file lock and the server attaches the lease to its session record only after that +# returns. A concurrent finalize that snapshotted its live ids in between would otherwise +# read the brand-new lease as an orphan and drop it. Real orphans are minutes old. +_SELF_ORPHAN_GRACE_SECONDS = 30.0 - ``_prune_dead`` only reclaims leases of dead processes, so on a days-long server a - lease whose session skipped teardown is held until restart. The owning process is the - only authority on its own leases — exact, no heartbeat on the turn path, no threshold. - """ + +def _drop_self_orphans( + entries: list[dict[str, Any]], own_live_lease_ids: set[str] | None +) -> list[dict[str, Any]]: + """Drop this process's leases only when its caller can vouch for owners.""" + if own_live_lease_ids is None: + return entries pid = os.getpid() - state_path = _state_path() + cutoff = time.time() - _SELF_ORPHAN_GRACE_SECONDS + return [ + entry for entry in entries + if entry.get("pid") != pid + or str(entry.get("lease_id") or "") in own_live_lease_ids + or (_optional_float(entry.get("started_at")) or 0.0) > cutoff + ] + + +def _release_orphaned_leases_in_home(registry_home: Path, live_lease_ids: set[str]) -> int: + state_path = _state_path(registry_home) # No registry file yet means no leases have ever been written under this # home — don't take a lock (or create its file) on the idle-reaper tick. if not state_path.exists(): return 0 - with _FileLock(_lock_path()): + with _FileLock(_lock_path(registry_home)): loaded = _read_live_entries( state_path, track_liveness=False, warn="Active-session registry is unavailable; skipping orphaned-lease sweep", @@ -600,13 +617,35 @@ def release_orphaned_leases(live_lease_ids: set[str]) -> int: if loaded is None: return 0 entries = loaded[1] - kept = [ - entry for entry in entries - if entry.get("pid") != pid or str(entry.get("lease_id") or "") in live_lease_ids - ] + kept = _drop_self_orphans(entries, live_lease_ids) dropped = len(entries) - len(kept) if dropped: _write_entries(state_path, kept) + return dropped + + +def release_orphaned_leases(live_lease_ids: set[str]) -> int: + """Drop this process's registry entries that no live session owns. + + ``_prune_dead`` only reclaims leases of dead processes, so on a days-long server a + lease whose session skipped teardown is held until restart. The owning process is the + only authority on its own leases — exact, no heartbeat on the turn path, no threshold. + Sweeps the root home and every profile home (a multiplexed server leases across them). + """ + root = get_default_hermes_root() + homes = [root] + try: + homes.extend(p for p in (root / "profiles").iterdir() + if p.is_dir() and not p.name.startswith(".")) + except OSError: + pass + + dropped = 0 + for home in homes: + try: + dropped += _release_orphaned_leases_in_home(home, live_lease_ids) + except OSError as exc: + logger.debug("orphaned-lease sweep failed for %s: %s", home, exc) return dropped @@ -625,32 +664,39 @@ def active_session_registry_snapshot( @contextmanager def active_session_liveness_guard( - session_id: str, *, registry_home: str | Path | None = None + session_id: str, *, registry_home: str | Path | None = None, + own_live_lease_ids: set[str] | None = None, ) -> Iterator[bool]: """Hold the registry lock while reporting whether ``session_id`` is leased, so no new backend can acquire a lease between the check and the caller's ``end_session``.""" state_path, lock_path = _lease_paths(registry_home=registry_home) with _FileLock(lock_path): entries = _prune_dead(_read_entries(state_path, strict=True), strict=True) + entries = _drop_self_orphans(entries, own_live_lease_ids) _write_entries(state_path, entries) yield _holds_session(entries, session_id) @contextmanager def release_active_session_liveness_guard( - lease: ActiveSessionLease, session_id: str + lease: ActiveSessionLease, session_id: str, *, own_live_lease_ids: set[str] | None = None, ) -> Iterator[bool]: """Remove ``lease`` and hold its registry lock through a lifecycle write, making cleanup one atomic decision (release, check siblings, end the durable row).""" if not lease.enabled or lease.released: home = lease.state_path.parent.parent if lease.state_path is not None else None - with active_session_liveness_guard(session_id, registry_home=home) as active: + with active_session_liveness_guard( + session_id, registry_home=home, own_live_lease_ids=own_live_lease_ids, + ) as active: yield active return state_path, lock_path = _lease_paths(lease) with _FileLock(lock_path): entries = _prune_dead(_read_entries(state_path, strict=True), strict=True) - kept = _drop_lease(state_path, entries, lease.lease_id) + kept = [e for e in entries if str(e.get("lease_id") or "") != lease.lease_id] + kept = _drop_self_orphans(kept, own_live_lease_ids) + if len(kept) != len(entries): + _write_entries(state_path, kept) lease.released = True yield _holds_session(kept, session_id) diff --git a/hermes_cli/auth.py b/hermes_cli/auth.py index 01c9c6a66c..3507903bf4 100644 --- a/hermes_cli/auth.py +++ b/hermes_cli/auth.py @@ -501,6 +501,19 @@ def _same_path(left: Path, right: Path) -> bool: return left == right +def _is_same_auth_store(left: Path, right: Path) -> bool: + """True when two auth paths name ONE store rather than two copies. + ``_same_path`` resolves symlinks and ``..``; ``samefile`` adds hardlinks and bind-mounts + (same inode under two resolved names). Used by the forked-grant heal: a shared store has + no "other side" to consolidate.""" + if _same_path(left, right): + return True + try: + return left.samefile(right) + except OSError: + return False + + def _resolved_key(path: Path) -> str: """Canonical string for *path* (resolved when possible) used as a cache / lock-holder key.""" try: diff --git a/hermes_cli/auth_oauth_grants.py b/hermes_cli/auth_oauth_grants.py index 42c6843f87..707b270b68 100644 --- a/hermes_cli/auth_oauth_grants.py +++ b/hermes_cli/auth_oauth_grants.py @@ -376,6 +376,9 @@ class _HealPass: def heal_profile_singleton(self, profile_singleton: Optional[Path]) -> None: if profile_singleton is None or not profile_singleton.exists(): return + from hermes_cli.auth import _is_same_auth_store + if self.root_singleton is not None and _is_same_auth_store(profile_singleton, self.root_singleton): + return # an aliased singleton pair is one shared grant, not a fork: never self-compare/unlink p_single = _singleton_as_row(profile_singleton) root_has_grant = bool(self.r_oauth) or self.root_singleton_row is not None # Otherwise root has NO grant for this provider (or the file is not a grant): the @@ -444,7 +447,8 @@ class _HealPass: def _heal_forked_single_use_oauth_grants(provider_id: str) -> Optional[Dict[str, Any]]: from hermes_cli.auth import ( _auth_file_path, _auth_store_lock, _global_auth_file_path, _load_auth_store, - _oauth_heal_clean_marks, _oauth_heal_notices, _same_path, _save_auth_store) + _is_same_auth_store, _oauth_heal_clean_marks, _oauth_heal_notices, _same_path, + _save_auth_store) root_path = _global_auth_file_path() if root_path is None: return None # classic mode: nothing to consolidate into @@ -469,6 +473,14 @@ def _heal_forked_single_use_oauth_grants(provider_id: str) -> Optional[Dict[str, if fingerprint[1] is None and fingerprint[2] is None: _oauth_heal_clean_marks[provider_id] = fingerprint return None + if _is_same_auth_store(profile_path, root_path): + # The profile's auth.json IS the root store (symlink/hardlink alias — a deliberate way to + # share one grant). Both "sides" would read the same file, every OAuth row would match + # itself, and the strip would write through the alias and delete the shared credential. + # Nothing to consolidate; the mtime mark keeps this off the per-call hot path. + _oauth_heal_clean_marks[provider_id] = fingerprint + logger.debug("%s: forked-OAuth heal skipped, %s is the root store", provider_id, profile_path) + return None # Lock order: active (profile) store first, then the root source store — the same order # ``_provider_state_transaction`` uses. diff --git a/hermes_cli/backup.py b/hermes_cli/backup.py index fa911da453..113356fc1d 100644 --- a/hermes_cli/backup.py +++ b/hermes_cli/backup.py @@ -77,7 +77,8 @@ def _in_excluded_root_dir(rel_path: Path) -> bool: # SQLite sidecars are excluded because ``*.db`` is snapshotted via ``sqlite3.backup()``: # shipping the live WAL/SHM/journal alongside would pair a fresh snapshot with stale sidecar # state and produce a torn restore on next open. They are regenerated on first connection. -_EXCLUDED_SUFFIXES = (".pyc", ".pyo", ".db-wal", ".db-shm", ".db-journal") +_SQLITE_SIDECAR_SUFFIXES = (".db-wal", ".db-shm", ".db-journal") +_EXCLUDED_SUFFIXES = (".pyc", ".pyo", *_SQLITE_SIDECAR_SUFFIXES) # File names to skip (runtime state that's meaningless on another machine) _EXCLUDED_NAMES = {".backup.lock", "gateway.pid", "cron.pid"} @@ -430,6 +431,9 @@ def _safe_restore_db(src: Path, dst: Path) -> bool: Writing pages into the live file preserves its inode and WAL state, so other holders (gateway, dashboard, another CLI) see the restored data instead of stale pages from a replaced inode. + The fallback runs ONLY when no other process or in-process connection holds the file + (replacing the inode under a live holder is the #90950 split-brain); otherwise it fails closed + (``False``) and the caller reports the file as skipped. """ try: dst_conn = sqlite3.connect(str(dst)) @@ -740,6 +744,68 @@ def _extract_member_atomically( raise +def _count_session_rows(path: Path) -> Optional[Tuple[int, int]]: + """``(sessions, messages)`` in session database *path*; read-only, best effort. + + ``None`` means "unknown" (missing, not a Hermes session store, unreadable) — never "zero": + acting on an unreadable database would mask the very loss this count exists to surface. + Same contract as :func:`_count_cron_jobs`. + """ + if not path.is_file(): + return None + try: + conn = sqlite3.connect(f"file:{path}?mode=ro", uri=True) + except sqlite3.Error: + return None + try: + sessions = conn.execute("SELECT COUNT(*) FROM sessions").fetchone()[0] + messages = conn.execute("SELECT COUNT(*) FROM messages").fetchone()[0] + return int(sessions), int(messages) + except (sqlite3.Error, TypeError, ValueError): + return None + finally: + conn.close() + + +def _import_db_member( + zf: zipfile.ZipFile, member: str, target: Path, new_file_mode: Optional[int] = None) -> None: + """Publish a SQLite ``.db`` member onto *target* without replacing its inode. + + A rename-publish over a live database is the #65942 / #90950 corruption class: a gateway, + dashboard, or WebUI holding it open keeps serving the unlinked inode and writing sessions no + other process will see, and a sidecar WAL beside the new file describes the old database — + nothing fails, the sessions are simply gone (#100960). Route the member through the same + ``_safe_restore_db`` page copy ``/snapshot restore`` uses, so the live inode is preserved and + every open connection converges. A target that does not exist yet has no holders, so it takes + the ordinary atomic publish. Raises ``OSError`` when the database could not be replaced + safely, so the caller reports a skipped file instead of a silent success. + """ + if not target.exists(): + _extract_member_atomically(zf, member, target, new_file_mode) + return + # The database keeps its own mode/ownership: the bytes come from the archive, the file does not. + mode = _preserve_file_mode(target) + owner = _preserve_file_owner(target) + fd, tmp_name = tempfile.mkstemp(dir=str(target.parent), prefix=f".{target.name[:80]}.", suffix=".dbimport") + try: + with os.fdopen(fd, "wb") as dst: + with zf.open(member) as src: # stream: never hold a multi-GB state.db in memory + shutil.copyfileobj(src, dst) + dst.flush() + os.fsync(dst.fileno()) + if not _safe_restore_db(Path(tmp_name), target): + raise OSError( + "live-safe restore refused or failed; the existing database was " + "left untouched. Stop the gateway/dashboard processes holding it " + "open and re-run the import." + ) + _restore_file_owner(target, owner) + _restore_file_mode(target, mode) + finally: + with suppress(OSError): + os.unlink(tmp_name) + + def _confirm_import_overwrite(hermes_root: Path) -> bool: """Prompt before importing over an existing installation; True when import may proceed.""" if not any((hermes_root / m).exists() for m in ("config.yaml", ".env")): @@ -759,10 +825,15 @@ def _confirm_import_overwrite(hermes_root: Path) -> bool: def _import_members( zf: zipfile.ZipFile, members: List[str], prefix: str, hermes_root: Path, file_count: int -) -> tuple[int, int, list[str], list[str]]: - """Publish every member; return ``(restored, restored_external, errors, skipped_runtime)``.""" +) -> tuple[int, int, list[str], list[str], list[tuple[str, tuple[int, int], tuple[int, int]]]]: + """Publish every member; return ``(restored, restored_external, errors, skipped_runtime, db_shrunk)``. + + ``db_shrunk`` holds ``(rel, live_counts, imported_counts)`` for every session database the + import replaced with one holding fewer rows — allowed, but never silent (#100960). + """ errors: list[str] = [] skipped_runtime: list[str] = [] + db_shrunk: list[tuple[str, tuple[int, int], tuple[int, int]]] = [] restored = restored_external = 0 home_dir = Path.home().resolve() new_file_mode = _default_new_file_mode() # once: every member is published via mkstemp (0600) @@ -780,6 +851,13 @@ def _import_members( if rel and Path(rel).name in _IMPORT_SKIP_NAMES: # see ``_IMPORT_SKIP_NAMES`` skipped_runtime.append(rel) continue + # A ``.db`` member is page-restored into the live file; an archived WAL/SHM/journal + # describes a different database image and installed beside it (over a live sidecar) + # would replay a foreign WAL on next open. Current backups never ship these + # (_EXCLUDED_SUFFIXES); older or hand-built archives might. + if rel.endswith(_SQLITE_SIDECAR_SUFFIXES): + skipped_runtime.append(rel) + continue target = hermes_root / rel root = hermes_root.resolve() tighten = target.name in _SECRET_FILE_NAMES @@ -792,7 +870,15 @@ def _import_members( else: try: target.parent.mkdir(parents=True, exist_ok=True) - _extract_member_atomically(zf, member, target, new_file_mode) + if target.suffix == ".db": + # Count before the write: afterwards the dropped rows are gone. + before = _count_session_rows(target) + _import_db_member(zf, member, target, new_file_mode) + after = _count_session_rows(target) + if before and after and after[1] < before[1]: + db_shrunk.append((rel, before, after)) + else: + _extract_member_atomically(zf, member, target, new_file_mode) if tighten: try: os.chmod(target, 0o600) @@ -807,7 +893,7 @@ def _import_members( if restored % 500 == 0: print(f" {restored}/{file_count} files ...") - return restored, restored_external, errors, skipped_runtime + return restored, restored_external, errors, skipped_runtime, db_shrunk def run_import(args) -> None: @@ -838,7 +924,7 @@ def run_import(args) -> None: print(f"\nImporting {file_count} files ...") hermes_root.mkdir(parents=True, exist_ok=True) t0 = time.monotonic() - restored, restored_external, errors, skipped_runtime = _import_members( + restored, restored_external, errors, skipped_runtime, db_shrunk = _import_members( zf, members, prefix, hermes_root, file_count) elapsed = time.monotonic() - t0 print(f"\nImport complete: {restored} files restored in {elapsed:.1f}s\n Target: {display_hermes_home()}") @@ -847,6 +933,15 @@ def run_import(args) -> None: f"their original location(s) outside {display_hermes_home()}.") if errors: _print_capped(f"\n Warnings ({len(errors)} files skipped):", errors, " ") + if db_shrunk: + # The backup predates work that is now overwritten — say so (#100960: twelve sessions + # disappeared with nothing logged anywhere). + print("\n ⚠ Session data replaced by older backup contents:") + for rel, before, after in db_shrunk: + print(f" {rel}: {before[0]} session(s) / {before[1]} message(s)" + f" -> {after[0]} / {after[1]}") + print(" Anything recorded after the backup was taken is not in it. " + "Recover from a newer backup or snapshot: hermes snapshot list") if skipped_runtime: _print_capped(f"\n Preserved {len(skipped_runtime)} runtime state " f"file(s) (kept this machine's, not the backup's):", @@ -1144,7 +1239,10 @@ def restore_quick_snapshot(snapshot_id: str, hermes_home: Optional[Path] = None) if dst.suffix == ".db": # Through the backup API so live connections see the restored data instead of # stale pages from a replaced inode (#65942). - _safe_restore_db(src, dst) + if not _safe_restore_db(src, dst): + # Refused (live holder) or failed: destination untouched — a failure, not a restore. + logger.error("Failed to restore %s: live-safe restore refused", rel) + continue else: shutil.copy2(src, dst) restored += 1 diff --git a/hermes_cli/config.py b/hermes_cli/config.py index 694de674a1..2ce83382de 100644 --- a/hermes_cli/config.py +++ b/hermes_cli/config.py @@ -960,11 +960,14 @@ def _coerce_config_version(value: Any) -> int: return max(version, 0) -def check_config_version() -> Tuple[int, int]: +def check_config_version(*, raise_on_parse_error: bool = False) -> Tuple[int, int]: """Return ``(current_version, latest_version)`` from the raw on-disk config. Reads the raw file rather than ``load_config()``: the deep-merge would make a file lacking ``_config_version`` inherit the latest version, hiding that the schema was never migrated. - Invalid YAML gets a parse warning, not an automatic schema rewrite.""" + Invalid YAML gets a parse warning, not an automatic schema rewrite. Tolerant runtime status + callers keep the historical latest/latest fallback for malformed YAML; mutation and explicit + validation paths set ``raise_on_parse_error`` so a parse failure or a non-mapping root cannot + be mistaken for an up-to-date config.""" latest = _coerce_config_version(DEFAULT_CONFIG.get("_config_version", 1)) or 1 config_path = get_config_path() if not config_path.exists(): @@ -972,12 +975,25 @@ def check_config_version() -> Tuple[int, int]: try: with open(config_path, encoding="utf-8") as f: - config = fast_safe_load(f) or {} + config = fast_safe_load(f) except Exception as e: _warn_config_parse_failure(config_path, e) + if raise_on_parse_error: + raise InvalidUserConfigError( + f"Cannot inspect {config_path}: config.yaml is not valid YAML ({e})" + ) from e return latest, latest + if config is None: + config = {} # empty file / bare document: valid first-run state if not isinstance(config, dict): + # A list/scalar root parses fine but is just as unusable as broken YAML: save_config() + # would refuse it later, after .env was already rewritten. Strict callers see it up front. + if raise_on_parse_error: + raise InvalidUserConfigError( + f"Cannot inspect {config_path}: config.yaml top-level value must be " + f"a mapping, got {type(config).__name__}" + ) config = {} return _coerce_config_version(config.get("_config_version")), latest @@ -1245,6 +1261,10 @@ def migrate_config(interactive: bool = True, quiet: bool = False) -> Dict[str, A """Migrate config to latest version, prompting for new required fields.""" results = {"env_added": [], "config_added": [], "warnings": []} + # Validate config.yaml before any migration side effect: sanitize_env_file() rewrites .env, + # which must not happen when the migration will be refused for malformed YAML. + current_ver, latest_ver = check_config_version(raise_on_parse_error=True) + try: fixes = sanitize_env_file() if fixes and not quiet: @@ -1252,8 +1272,6 @@ def migrate_config(interactive: bool = True, quiet: bool = False) -> Dict[str, A except Exception: pass # best-effort; never block migration on sanitize failure - current_ver, latest_ver = check_config_version() - # Auto-migration support floor (v12): an EXPLICIT on-disk ``_config_version`` below the # floor is NOT migrated and NOT rewritten — surface a message and leave the file untouched # (deep-merge supplies defaults at read time). A config with NO version key is a fresh @@ -3497,7 +3515,7 @@ def _cmd_config_migrate(args): missing_env = get_missing_env_vars(required_only=False) missing_config = get_missing_config_fields() - current_ver, latest_ver = check_config_version() + current_ver, latest_ver = check_config_version(raise_on_parse_error=True) if not missing_env and not missing_config and current_ver >= latest_ver: print(color("✓ Configuration is up to date!", Colors.GREEN)) @@ -3536,7 +3554,7 @@ def _cmd_config_check(args): """Non-interactive report of what's missing.""" _print_banner("📋 Configuration Status") - current_ver, latest_ver = check_config_version() + current_ver, latest_ver = check_config_version(raise_on_parse_error=True) if current_ver >= latest_ver: print(f" Config version: {current_ver} ✓") else: diff --git a/hermes_cli/config_defaults.py b/hermes_cli/config_defaults.py index 5c51cfe577..b156d1a461 100644 --- a/hermes_cli/config_defaults.py +++ b/hermes_cli/config_defaults.py @@ -452,6 +452,9 @@ DEFAULT_CONFIG = { # .cursorrules) before head/tail truncation. null = scale with the model's context window (floor # 20K, ceiling 500K); a positive int pins a fixed cap. Separate from read_file limits. "context_file_max_chars": None, + # Seconds to wait for a single context file read before skipping it with a warning. Guards startup + # against network-backed filesystems (iCloud Drive, OneDrive, NFS) that can block a cold read. + "context_file_read_timeout": 5.0, # Max chars per read_file call; larger reads are rejected with offset+limit guidance. 100K chars # ≈ 25–35K tokens. "file_read_max_chars": 100_000, @@ -2105,6 +2108,7 @@ DEFAULT_CONFIG = { # cua-driver's upstream PostHog telemetry defaults ON; Hermes sets # CUA_DRIVER_RS_TELEMETRY_ENABLED=0 in every child env unless this is true. "cua_telemetry": False, + "native_wayland": False, # Cap driver screenshot longest edge (pixels) via set_config at session start; shrinks SOM # multimodal payloads. 0 disables. "max_image_dimension": 1456, diff --git a/hermes_cli/copilot_auth.py b/hermes_cli/copilot_auth.py index b178bdc835..d520b1310b 100644 --- a/hermes_cli/copilot_auth.py +++ b/hermes_cli/copilot_auth.py @@ -421,9 +421,18 @@ def exchange_copilot_token( so it is None. Cached in-process until close to expiry. Raises ``ValueError`` on failure. """ fp = _token_fingerprint(raw_token) - cached = _jwt_cache.get(fp) # fast path outside the lock + # Fast paths outside the lock: a valid in-process JWT needs no exchange, and a recent failure + # means queueing behind the in-flight holder (up to ~50 s) would only park an executor thread + # to learn the same answer. + cached = _jwt_cache.get(fp) if _cache_entry_fresh(cached): return cached + _fail_until = _exchange_failure_cache.get(fp, 0.0) + if time.time() < _fail_until: + raise ValueError("Copilot token exchange recently failed; skipping re-attempt " + f"for another {int(_fail_until - time.time())}s") + # Note: a waiter's own ``timeout`` is not honoured across the lock wait — by design of + # single-flight, it observes the holder's outcome instead. with _exchange_lock_for(fp): return _exchange_copilot_token_locked(raw_token, fp, timeout=timeout) diff --git a/hermes_cli/doctor_state.py b/hermes_cli/doctor_state.py index 360d36caa1..3550cb07d4 100644 --- a/hermes_cli/doctor_state.py +++ b/hermes_cli/doctor_state.py @@ -10,6 +10,7 @@ from hermes_cli.doctor_report import ( warn_on_error, ) from hermes_cli.sizefmt import format_bytes as _human_bytes +from hermes_state_common import FTS_STORAGE_VERSION def _honcho_is_configured_for_doctor() -> bool: @@ -82,12 +83,13 @@ def _render_state_db_stats(stats: dict, holders=None) -> list: lines.append(("warn", f"state.db FTS repair is blocked after {deferral.get('attempts') or '?'} deferral(s) " f"by PID(s) {deferral.get('holder_pids') or [] or 'unknown'}", "(stop the listed processes, then run 'hermes sessions optimize-storage' with the gateway stopped)")) - # Oversized DB: suggest auto_prune, plus the offline optimize-storage pass when the v23 FTS rebuild is - # pending OR the DB still carries the legacy inline trigram layout (fts_storage_version marker absent). + # Oversized DB: suggest auto_prune, plus the offline optimize-storage pass when the FTS rebuild is + # pending OR the DB predates the current trigram layout (fts_storage_version < FTS_STORAGE_VERSION). if logical is not None and logical > STATE_DB_SIZE_WARN_BYTES: detail = "consider enabling sessions.auto_prune in config.yaml to bound growth" - legacy_trigram = fts is not None and fts.get("messages_fts_trigram") and stats.get("fts_storage_version") is None - if stats.get("fts_rebuild_pending") or legacy_trigram: + stale_trigram = (fts is not None and fts.get("messages_fts_trigram") + and (stats.get("fts_storage_version") or 0) < FTS_STORAGE_VERSION) + if stats.get("fts_rebuild_pending") or stale_trigram: detail += "; run 'hermes sessions optimize-storage' offline (with the gateway stopped) to compact FTS storage" lines.append(("warn", f"state.db is large ({_human_bytes(logical)})", f"({detail})")) # WAL runaway is deliberately NOT warned here: _state_db_wal already warns above 50 MB and offers --fix. diff --git a/hermes_cli/inventory.py b/hermes_cli/inventory.py index 266732bf5b..7f3ecd613a 100644 --- a/hermes_cli/inventory.py +++ b/hermes_cli/inventory.py @@ -3,9 +3,14 @@ from __future__ import annotations +from contextvars import copy_context from dataclasses import dataclass, replace +from threading import Lock, Thread, current_thread from typing import Any, Optional +_pricing_prewarm_lock = Lock() +_pricing_prewarm_threads: dict[tuple[str, tuple[tuple[str, str], ...]], Thread] = {} + @dataclass(frozen=True) class ConfigContext: @@ -69,13 +74,15 @@ def _without_slug(rows: list[dict], slug: str) -> list[dict]: def build_models_payload( ctx: ConfigContext, *, explicit_only: bool = False, include_unconfigured: bool = False, picker_hints: bool = False, canonical_order: bool = False, pricing: bool = False, + pricing_cache_only: bool = False, capabilities: bool = False, featured: bool = False, force_fresh_nous_tier: bool = False, refresh: bool = False, probe_custom_providers: bool = True, probe_current_custom_provider: bool = False, for_picker: bool = False, max_models: int | None = None, ) -> dict: """Build the ``{providers, model, provider}`` shape every consumer needs. ``explicit_only`` keeps only providers the user explicitly configured — hides ambient/auto-seeded credentials from - desktop chat pickers.""" + desktop chat pickers. ``pricing_cache_only``: with ``pricing``, use only values already resident + in process caches (normal picker opens, while a background worker warms cold endpoints).""" from hermes_cli.model_switch import list_authenticated_providers rows = list_authenticated_providers( @@ -131,7 +138,7 @@ def build_models_payload( if canonical_order: rows = _reorder_canonical(rows) if pricing: - _apply_pricing(rows, force_fresh_nous_tier=force_fresh_nous_tier) + _apply_pricing(rows, force_fresh_nous_tier=force_fresh_nous_tier, cached_only=pricing_cache_only) if capabilities: _apply_capabilities(rows) if featured: @@ -174,11 +181,16 @@ def build_model_options_payload( """Shared API-server/dashboard/TUI payload. Normal open probes only the current custom provider so offline saved endpoints don't block the picker; explicit refresh probes all and busts the cache.""" refresh = bool(refresh) - return build_models_payload( + payload = build_models_payload( ctx, explicit_only=bool(explicit_only), include_unconfigured=bool(include_unconfigured), - picker_hints=True, canonical_order=True, pricing=True, capabilities=True, featured=True, + picker_hints=True, canonical_order=True, pricing=True, pricing_cache_only=not refresh, + capabilities=True, featured=True, refresh=refresh, probe_custom_providers=refresh, probe_current_custom_provider=not refresh, ) + if not refresh: + _prewarm_pricing_async(payload["providers"], current_provider=ctx.current_provider, + current_base_url=ctx.current_base_url) + return payload # ─── Public: auxiliary-task pickers ───────────────────────────────────── @@ -538,12 +550,14 @@ def _reorder_canonical(rows: list[dict]) -> list[dict]: return canon + extras -def _apply_pricing(rows: list[dict], *, force_fresh_nous_tier: bool = False) -> None: +def _apply_pricing(rows: list[dict], *, force_fresh_nous_tier: bool = False, cached_only: bool = False) -> None: """Set ``row["pricing"] = {model_id: {input, output, cache | None, free}}``; for Nous also - ``free_tier`` (account is free-tier) and ``unavailable_models`` (paid models a free user can't pick).""" + ``free_tier`` (account is free-tier) and ``unavailable_models`` (paid models a free user can't pick). + ``cached_only`` never hits the network: unknown Nous entitlement fails closed (``free_tier_pending``, + all models locked) and missing pricing is marked ``pricing_pending``.""" from hermes_cli.models import ( - _format_price_per_mtok, check_nous_free_tier, compute_sale_discount, get_pricing_for_provider, - partition_nous_models_by_tier, + _format_price_per_mtok, check_nous_free_tier, compute_sale_discount, get_cached_nous_free_tier, + get_pricing_for_provider, partition_nous_models_by_tier, ) nous_free_tier: Optional[bool] = None # resolved once (cached in models.py for the TTL window) @@ -554,10 +568,27 @@ def _apply_pricing(rows: list[dict], *, force_fresh_nous_tier: bool = False) -> if not models: continue try: - raw_pricing = get_pricing_for_provider(slug) or {} + pricing_kwargs = {"cached_only": True} if cached_only else {} + raw_pricing = get_pricing_for_provider(slug, **pricing_kwargs) or {} except Exception: raw_pricing = {} + cached_nous_tier: Optional[bool] = None + if slug == "nous" and cached_only: + cached_nous_tier = get_cached_nous_free_tier() + if cached_nous_tier is None: + # Entitlement unknown: stay nonblocking but fail closed until the prewarm has populated + # both caches, else a free account could briefly select paid models on first open. + row["free_tier_pending"] = True + row["unavailable_models"] = list(models) + if not row.get("warning"): # say why every model renders locked + row["warning"] = ("Checking Nous plan entitlement… models unlock on the " + "next picker open or refresh.") + continue if not raw_pricing: + if slug == "nous": + row["free_tier"] = bool(cached_nous_tier) + row["pricing_pending"] = True + row["unavailable_models"] = list(models) if cached_nous_tier else [] continue formatted: dict[str, dict] = {} @@ -592,7 +623,8 @@ def _apply_pricing(rows: list[dict], *, force_fresh_nous_tier: bool = False) -> if slug == "nous": try: if nous_free_tier is None: - nous_free_tier = check_nous_free_tier(force_fresh=force_fresh_nous_tier) + nous_free_tier = (cached_nous_tier if cached_only + else check_nous_free_tier(force_fresh=force_fresh_nous_tier)) row["free_tier"] = bool(nous_free_tier) row["unavailable_models"] = ( partition_nous_models_by_tier(list(models), raw_pricing, free_tier=True)[1] @@ -630,6 +662,42 @@ def _local_runtime_row(ctx: "ConfigContext") -> dict | None: return None +def _prewarm_pricing_async( + rows: list[dict], *, current_provider: str = "", current_base_url: str = "", +) -> Optional[Thread]: + """Warm picker pricing caches without delaying the current payload (one worker per + profile + endpoint scope; a live worker is reused).""" + from hermes_constants import hermes_home_key + from hermes_cli.models import pricing_cache_scope + + slugs = {str(row.get("slug") or "").lower() for row in rows if row.get("slug")} + endpoint_scope = tuple(sorted( + (slug, pricing_cache_scope(slug, current_provider=current_provider, current_base_url=current_base_url)) + for slug in slugs)) + prewarm_key = (hermes_home_key(), endpoint_scope) + + with _pricing_prewarm_lock: + current = _pricing_prewarm_threads.get(prewarm_key) + if current is not None and current.is_alive(): + return current + # The worker mutates only private copies; the pricing helpers populate shared process caches. + worker_rows = [{**row, "models": list(row.get("models") or [])} for row in rows] + + def _worker() -> None: + try: + _apply_pricing(worker_rows) + finally: + with _pricing_prewarm_lock: + if _pricing_prewarm_threads.get(prewarm_key) is current_thread(): + _pricing_prewarm_threads.pop(prewarm_key, None) + + thread = Thread(target=copy_context().run, args=(_worker,), + name="hermes-picker-pricing-prewarm", daemon=True) + _pricing_prewarm_threads[prewarm_key] = thread + thread.start() + return thread + + def _moa_provider_row(current_provider: str = "") -> dict | None: """The virtual ``moa`` row shared by the CLI inventory and gateway picker; ``None`` without presets.""" try: diff --git a/hermes_cli/kanban_db.py b/hermes_cli/kanban_db.py index 991e327b4e..3688790822 100644 --- a/hermes_cli/kanban_db.py +++ b/hermes_cli/kanban_db.py @@ -4183,6 +4183,7 @@ from hermes_cli.kanban_db_dispatch import ( # noqa: E402,F401 _record_worker_exit, _resolve_hermes_argv, _resolve_worker_cli_toolsets, + _restart_safe_worker_argv, _retag_legacy_worker_sessions, _set_worker_pid, _system_memory_sample, diff --git a/hermes_cli/kanban_db_dispatch.py b/hermes_cli/kanban_db_dispatch.py index c83b1c4188..bebd257421 100644 --- a/hermes_cli/kanban_db_dispatch.py +++ b/hermes_cli/kanban_db_dispatch.py @@ -2105,6 +2105,30 @@ def _open_worker_log(task: Task, board: Optional[str]): return open(log_path, "ab") +def _restart_safe_worker_argv(task: Task, command: list[str]) -> list[str]: + """Wrap a managed-gateway worker in the shared restart-safe scope.""" + from tools.process_registry import restart_safe_gateway_child_argv + + if task.current_run_id is None: + # Outside managed systemd this is harmless, but a managed dispatch must + # never mint an untraceable scope. Check topology through the shared + # helper first, using a placeholder suffix that cannot be launched. + scoped = restart_safe_gateway_child_argv( + command, unit_suffix=f"kanban-{task.id}-run-missing" + ) + if scoped is not command: + raise RuntimeError( + "cannot create restart-safe systemd scope for Kanban worker: " + "the claimed task has no current run id" + ) + return command + + return restart_safe_gateway_child_argv( + command, + unit_suffix=f"kanban-{task.id}-run-{task.current_run_id}", + ) + + def _default_spawn(task: Task, workspace: str, *, board: Optional[str] = None) -> Optional[int]: """Fire-and-forget ``hermes -p chat -q ...`` subprocess. @@ -2121,7 +2145,13 @@ def _default_spawn(task: Task, workspace: str, *, board: Optional[str] = None) - profile_arg = normalize_profile_name(task.assignee) - env = dict(os.environ) + from agent.secret_scope import is_multiplex_active + from tools.environments.local import build_subprocess_env + + env = build_subprocess_env( + scrub_secrets=is_multiplex_active(), + inherit_profile_home=True, + ) # The dispatcher is detached from every conversation; its worker must never # inherit routing mirrored by a previous gateway turn. from gateway.session_context import _VAR_MAP @@ -2183,6 +2213,10 @@ def _default_spawn(task: Task, workspace: str, *, board: Optional[str] = None) - env.pop("HERMES_TUI", None) cmd = _worker_argv(task, profile_arg, env.get("HERMES_HOME")) + # A worker spawned by a managed systemd gateway must leave the gateway's + # cgroup before startup; otherwise restarting the service kills the worker + # that is performing the handoff. + cmd = _restart_safe_worker_argv(task, cmd) log_f = _open_worker_log(task, board) try: proc = subprocess.Popen( # noqa: S603 -- argv is a fixed list built above diff --git a/hermes_cli/kanban_ops.py b/hermes_cli/kanban_ops.py index 7b2def47bd..ef65c3aeed 100644 --- a/hermes_cli/kanban_ops.py +++ b/hermes_cli/kanban_ops.py @@ -374,7 +374,9 @@ def _cmd_repair(args: argparse.Namespace) -> int: if report.backup_path: err(f" corrupt copy quarantined at: {report.backup_path}") err( - " Recover manually (e.g. `sqlite3 kanban.db \".recover\"` into a " - "fresh file) or move the file aside to start a new board." + " Recover manually (copy kanban.db aside FIRST, then run " + "`sqlite3 \".recover\"` into a fresh file — never against " + "the live path, a WAL-reset-vulnerable sqlite3 CLI can corrupt it " + "further) or move the file aside to start a new board." ) return 1 diff --git a/hermes_cli/main.py b/hermes_cli/main.py index 501072a24c..8dcea5753a 100644 --- a/hermes_cli/main.py +++ b/hermes_cli/main.py @@ -2827,7 +2827,19 @@ _AGENT_SUBCOMMANDS = { def _is_tui_chat_launch(args) -> bool: - return bool(getattr(args, "tui", False) or os.environ.get("HERMES_TUI") == "1") + if getattr(args, "tui", False) or os.environ.get("HERMES_TUI") == "1": + return True + # The chat path decides TUI-vs-classic via _resolve_use_tui (--cli/--tui + # flags, TTY gate, HERMES_TUI env, display.interface config). Bare + # `hermes`/`hermes chat` with a TUI display config was previously missed + # here, so the wrapper pre-warmed its own MCP discovery while the TUI + # gateway (spawned moments later) ran a second one — an idle stdio MCP + # server copy held dead for the whole session. Only chat commands can + # launch the TUI; other commands (mcp serve, gateway, acp, cron) keep + # their own discovery behavior untouched. + if getattr(args, "command", None) not in {None, "chat"}: + return False + return _resolve_use_tui(args) def _agent_subcommand_selected(args) -> bool: @@ -2877,6 +2889,17 @@ def _prepare_agent_startup(args) -> None: "plugin discovery failed at CLI startup", exc_info=True, ) + # -t/--toolsets narrows which configured MCP servers get spawned, on + # every discovery path (inline below, background thread, TUI/desktop + # deferred start). Built-in toolset names never match a server key, so + # `-t terminal` simply spawns nothing; `-t all` keeps the full set. + try: + from hermes_cli.mcp_startup import set_mcp_server_filter + + set_mcp_server_filter(getattr(args, "toolsets", None)) + except Exception: + logger.debug("MCP server filter setup failed", exc_info=True) + # TUI launches hand off to a startup path that backgrounds MCP discovery # with a bounded join; acp/gateway/cron do their own on the runtime path. _run_inline_mcp_discovery = not ( @@ -2898,9 +2921,14 @@ def _prepare_agent_startup(args) -> None: _run_inline_mcp_discovery = False if _run_inline_mcp_discovery: try: # synchronous for entrypoints without a later bounded startup path + from hermes_cli.mcp_startup import get_mcp_server_filter from tools.mcp_tool import discover_mcp_tools - discover_mcp_tools() + _mcp_filter = get_mcp_server_filter() + if _mcp_filter is None: + discover_mcp_tools() + else: + discover_mcp_tools(allowed_mcp_names=_mcp_filter) except Exception: logger.debug( "MCP tool discovery failed at CLI startup", diff --git a/hermes_cli/mcp_startup.py b/hermes_cli/mcp_startup.py index b87c82b87f..8a2c587e47 100644 --- a/hermes_cli/mcp_startup.py +++ b/hermes_cli/mcp_startup.py @@ -12,6 +12,37 @@ _mcp_discovery_lock = threading.Lock() _mcp_discovery_started = False _mcp_discovery_thread: Optional[threading.Thread] = None _mcp_discovery_deferred: Optional[threading.Timer] = None +# Process-wide MCP server-name allowlist derived from ``-t/--toolsets``. +# ``None`` = no filter (spawn every configured server). Set once at CLI +# startup by ``set_mcp_server_filter`` and honored by every discovery path +# in this module (inline, background, deferred), so a ``-t terminal`` +# oneshot never cold-starts MCP subprocesses it cannot use. +_mcp_server_filter: Optional[list[str]] = None + + +def set_mcp_server_filter(toolsets: object) -> Optional[list[str]]: + """Derive the MCP spawn allowlist from a ``-t/--toolsets`` value. + + Built-in toolset names in the list are harmless (they never match a + configured ``mcp_servers`` key). ``all``/``*`` or an empty/absent value + clears the filter. Returns the stored list for logging/tests. + """ + global _mcp_server_filter + names: list[str] = [] + if isinstance(toolsets, str): + names = [t.strip() for t in toolsets.split(",") if t.strip()] + elif isinstance(toolsets, (list, tuple, set)): + for item in toolsets: + names.extend(t.strip() for t in str(item).split(",") if t.strip()) + if not names or "all" in names or "*" in names: + _mcp_server_filter = None + else: + _mcp_server_filter = names + return _mcp_server_filter + + +def get_mcp_server_filter() -> Optional[list[str]]: + return _mcp_server_filter def _has_configured_mcp_servers() -> bool: @@ -124,7 +155,13 @@ def _discover_mcp_tools_without_interactive_oauth() -> None: with suppress_interactive_oauth(): from tools.mcp_tool import discover_mcp_tools - discover_mcp_tools() + # Only pass the kwarg when a filter is set: many tests (and any + # out-of-tree caller) stub discover_mcp_tools with a zero-arg + # callable, and the unfiltered call shape is unchanged. + if _mcp_server_filter is None: + discover_mcp_tools() + else: + discover_mcp_tools(allowed_mcp_names=_mcp_server_filter) def defer_background_mcp_discovery(*, logger, thread_name: str, delay: float) -> None: diff --git a/hermes_cli/model_data_policy_guard.py b/hermes_cli/model_data_policy_guard.py index 8735ff5a32..05c7bf28cd 100644 --- a/hermes_cli/model_data_policy_guard.py +++ b/hermes_cli/model_data_policy_guard.py @@ -32,20 +32,18 @@ def _is_meta_contributor(model_lower: str, provider_lower: str) -> bool: _META_CONTRIBUTOR_MESSAGE = ( "!!! CONTRIBUTOR TIER — TRAINS ON YOUR DATA !!!\n" "\n" - "muse-spark-1.2-contributor is Meta's contributor tier: heavily discounted\n" - "token pricing in exchange for permission to use your prompts and completions\n" - "to train future Meta models.\n" + "This is Meta's contributor tier. Selecting it permits Meta to use your\n" + "prompts and completions to train future Meta models.\n" "\n" - " Price per 1M tokens: input $0.10 | output $0.20 | cached input $0.002\n" - " (vs. standard muse-spark-1.2: input $1.25 | output $4.25 | cached $0.15)\n" + "See current pricing and rate limits for the Meta Model API here:\n" + " https://dev.meta.ai/docs/pricing-rate-limits/\n" "\n" "It lowers the barrier to entry for prototyping, testing integrations, and\n" "scaling experiments where training on your data is acceptable. Do NOT use it\n" "for confidential, proprietary, personal, or otherwise sensitive data. For the\n" - "same model at standard pricing with no training on your data, select the\n" - "standard variant, muse-spark-1.2.\n" + "same model with no training on your data, select the standard variant\n" + "(without the -contributor suffix).\n" "\n" - "Source: https://dev.meta.ai/docs/pricing-rate-limits/\n" "Confirm only if training on your prompts and completions is acceptable." ) diff --git a/hermes_cli/models.py b/hermes_cli/models.py index 189b25f390..5849ac3ca0 100644 --- a/hermes_cli/models.py +++ b/hermes_cli/models.py @@ -104,9 +104,12 @@ from hermes_cli.models_pricing import ( # noqa: F401 (re-exported; tests patch compute_sale_discount, fetch_ai_gateway_pricing, fetch_models_with_pricing, + _pricing_provider_cache_keys, + get_cached_nous_inference_base_url, get_pricing_for_provider, nous_policy_allowed_ids, peek_cached_pricing, + pricing_cache_scope, restrict_to_nous_policy) from hermes_cli.models_validate import validate_requested_model # noqa: F401 (re-exported) @@ -276,25 +279,45 @@ def union_with_portal_paid_recommendations( force_refresh=force_refresh, synthesize_free_pricing=False) -# Free-tier detection cache — short so an account upgrade shows within minutes. +# Free-tier detection cache, per profile — short so an account upgrade shows within minutes. _FREE_TIER_CACHE_TTL: int = 180 # seconds -_free_tier_cache: tuple[bool, float] | None = None # (result, timestamp) +_free_tier_cache: dict[str, tuple[bool, float]] = {} # profile key -> (result, timestamp) -def check_nous_free_tier(*, force_fresh: bool = False) -> bool: +def _pricing_profile_key() -> str: + """Stable profile identity for process-local pricing caches.""" + from hermes_constants import hermes_home_key + + return hermes_home_key() + + +def get_cached_nous_free_tier() -> Optional[bool]: + """This profile's live cached entitlement, or ``None`` if unknown/expired.""" + cached = _free_tier_cache.get(_pricing_profile_key()) + if cached is None or time.monotonic() - cached[1] >= _FREE_TIER_CACHE_TTL: + return None + return cached[0] + + +def check_nous_free_tier(*, force_fresh: bool = False, cached_only: bool = False) -> bool: """True only when the Nous Portal user is KNOWN to be free-tier (unknown/error → False so this - never blocks users). Cached ``_FREE_TIER_CACHE_TTL`` seconds so an upgrade shows within minutes.""" - global _free_tier_cache + never blocks users). Cached ``_FREE_TIER_CACHE_TTL`` seconds so an upgrade shows within minutes. + ``cached_only`` returns the live cached answer or the fail-open ``False`` without contacting Portal.""" now = time.monotonic() - if not force_fresh and _free_tier_cache is not None and now - _free_tier_cache[1] < _FREE_TIER_CACHE_TTL: - return _free_tier_cache[0] + profile_key = _pricing_profile_key() + if not force_fresh: + cached_result = get_cached_nous_free_tier() + if cached_result is not None: + return cached_result + if cached_only: + return False try: from hermes_cli.nous_account import get_nous_portal_account_info result = get_nous_portal_account_info(force_fresh=force_fresh).is_free_tier except Exception: result = False # default to paid on error — don't block users - _free_tier_cache = (result, now) + _free_tier_cache[profile_key] = (result, now) return result diff --git a/hermes_cli/models_catalog_static.py b/hermes_cli/models_catalog_static.py index 9e001d7709..83d52ef965 100644 --- a/hermes_cli/models_catalog_static.py +++ b/hermes_cli/models_catalog_static.py @@ -33,7 +33,8 @@ OPENROUTER_MODELS: list[tuple[str, str]] = [ "deepseek/deepseek-v4-pro-0813", "deepseek/deepseek-v4-flash", "deepseek/deepseek-v4-flash-0731", "qwen/qwen3.8-max", "qwen/qwen3.8-flash", "moonshotai/kimi-k3", "minimax/minimax-m3", "z-ai/glm-5.3", "z-ai/glm-5.3-flash", "z-ai/glm-5.2", "xiaomi/mimo-v2.5-pro", "tencent/hy4-preview", "tencent/hy3", - "stepfun/step-3.7-flash", "nvidia/nemotron-3-super-120b-a12b", "meta/muse-spark-1.2", "sakana/fugu-ultra", + "stepfun/step-3.7-flash", "nvidia/nemotron-3-super-120b-a12b", "meta/muse-spark-1.2", + "meta/muse-spark-1.2-contributor", "meta/muse-spark-1.3", "meta/muse-spark-1.3-contributor", "sakana/fugu-ultra", "openrouter/pareto-code", "thinkingmachines/inkling:free", "thinkingmachines/inkling-small:free", "minimax/minimax-m3:free", "z-ai/glm-5.2:free", "poolside/laguna-s-2.1:free", "poolside/laguna-xs-2.1:free", "nvidia/nemotron-3-super-120b-a12b:free", "nvidia/nemotron-3-ultra-550b-a55b:free", @@ -43,7 +44,8 @@ OPENROUTER_MODELS: list[tuple[str, str]] = [ # OpenRouter entries the Nous Portal does not carry (routing/fast variants, free tier). _OPENROUTER_ONLY = { - "anthropic/claude-opus-5-fast", "anthropic/claude-opus-4.8-fast", "meta/muse-spark-1.2", "openrouter/pareto-code", + "anthropic/claude-opus-5-fast", "anthropic/claude-opus-4.8-fast", "meta/muse-spark-1.2", + "meta/muse-spark-1.2-contributor", "meta/muse-spark-1.3", "meta/muse-spark-1.3-contributor", "openrouter/pareto-code", } @@ -221,7 +223,7 @@ _PROVIDER_MODELS: dict[str, list[str]] = { "glm-5.3", "glm-5.3-flash", "glm-5.2", "glm-5.1", "glm-5", "kimi-k2.7-code", "deepseek-v4-pro", "deepseek-v4-flash", "deepseek-v4-flash-free", "qwen3.6-plus", "qwen3.5-plus", "big-pickle", "mimo-v2.5-free", "hy3-free", "laguna-s-2.1-free", "nemotron-3-ultra-free", "nemotron-3.5-lightning-free", - "muse-spark-1.2-contributor-free", + "muse-spark-1.2-contributor-free", "muse-spark-1.3-contributor-free", ], # OpenCode keyless free tier — OFFLINE FLOOR only. provider_model_ids("opencode-free") # revalidates live against GET /zen/v1/models and filters to the anonymous tier, so this list @@ -230,6 +232,7 @@ _PROVIDER_MODELS: dict[str, list[str]] = { "opencode-free": [ "deepseek-v4-flash-free", "hy3-free", "mimo-v2.5-free", "laguna-s-2.1-free", "nemotron-3-ultra-free", "nemotron-3.5-lightning-free", "muse-spark-1.2-contributor-free", + "muse-spark-1.3-contributor-free", ], # Synced against opencode.ai/docs/go + live GET /zen/go/v1/models. "ox-alpha-free" is the # Go-subscription twin of Zen's keyless Ox Alpha (NOT keyless — the Go relay requires a Go key). @@ -238,7 +241,8 @@ _PROVIDER_MODELS: dict[str, list[str]] = { "glm-5.3-flash", "glm-5.2", "glm-5.1", "glm-5", "mimo-v2.5-pro", "mimo-v2.5", "mimo-v2-pro", "mimo-v2-omni", "minimax-m3", "minimax-m2.7", "minimax-m2.5", "deepseek-v4-pro", "deepseek-v4-flash", "qwen3.8-max", "qwen3.7-max", "qwen3.7-plus", "qwen3.6-plus", - "qwen3.5-plus", "hy3", "hy3-preview", "muse-spark-1.2-contributor", "ox-alpha-free", + "qwen3.5-plus", "hy3", "hy3-preview", "muse-spark-1.2-contributor", "muse-spark-1.3-contributor", + "ox-alpha-free", ], "kilocode": [ "anthropic/claude-opus-4.6", "anthropic/claude-sonnet-4.6", "openai/gpt-5.4", @@ -512,7 +516,7 @@ _BORROWED_MODEL_PROVIDERS: frozenset[str] = frozenset() # entries lead, curated-only append). Every OTHER provider keeps curated-first so a deliberately # surfaced newest model stays on top when the live API lags. Zen/Go re-expose dozens of vendors # and rotate them often, so their stale curated entries must not pollute the top. -_LIVE_FIRST_PICKER_PROVIDERS: frozenset[str] = frozenset({"opencode-zen", "opencode-go"}) +_LIVE_FIRST_PICKER_PROVIDERS: frozenset[str] = frozenset({"opencode-zen", "opencode-go", "meta-ai"}) # Models supporting OpenAI Priority Processing (service_tier="priority"; see diff --git a/hermes_cli/models_pricing.py b/hermes_cli/models_pricing.py index 2bebefb5db..24f8c113fa 100644 --- a/hermes_cli/models_pricing.py +++ b/hermes_cli/models_pricing.py @@ -19,6 +19,8 @@ from hermes_cli.models_reasoning_caps import _seed_reasoning_caps # Cache: maps model_id → {"prompt": str, "completion": str} per endpoint _pricing_cache: dict[str, dict[str, dict[str, str]]] = {} +# (profile key, provider) → endpoint cache key last fetched, so cached_only reads find the right entry. +_pricing_provider_cache_keys: dict[tuple[str, str], str] = {} # A failed fetch caches its empty result too, so an unreachable endpoint isn't re-dialed on every # call — but only until this deadline. Cached forever, one blip at startup would mean no live model @@ -370,28 +372,141 @@ def restrict_to_nous_policy( return kept +def _remember_provider_cache_key(provider: str, cache_key: str) -> None: + from hermes_cli.models import _pricing_profile_key, _pricing_provider_cache_keys + _pricing_provider_cache_keys[(_pricing_profile_key(), provider)] = cache_key + + def _fetch_openrouter_pricing(*, force_refresh: bool = False) -> dict[str, dict[str, Any]]: from hermes_cli.models import fetch_models_with_pricing + _remember_provider_cache_key("openrouter", _OPENROUTER_PRICING_BASE) return fetch_models_with_pricing( api_key=_resolve_openrouter_api_key(), - base_url="https://openrouter.ai/api", + base_url=_OPENROUTER_PRICING_BASE, force_refresh=force_refresh, ) +def _fetch_ai_gateway_pricing_for_provider(*, force_refresh: bool = False) -> dict[str, dict[str, Any]]: + from hermes_cli.models import fetch_ai_gateway_pricing + _remember_provider_cache_key("ai-gateway", _ai_gateway_pricing_scope()) + return fetch_ai_gateway_pricing(force_refresh=force_refresh) + + +def _fetch_novita_pricing_for_provider(*, force_refresh: bool = False) -> dict[str, dict[str, Any]]: + from hermes_cli.models import _fetch_novita_pricing + _remember_provider_cache_key("novita", _novita_pricing_scope()) + return _fetch_novita_pricing(force_refresh=force_refresh) + + +def _fetch_fireworks_pricing_for_provider(*, force_refresh: bool = False) -> dict[str, dict[str, Any]]: + _remember_provider_cache_key("fireworks", _FIREWORKS_PRICING_KEY) + return _fireworks_pricing_from_models_dev(force_refresh=force_refresh) + + def _fetch_nous_pricing_for_provider(*, force_refresh: bool = False) -> dict[str, dict[str, Any]]: from hermes_cli.models import _resolve_nous_pricing_credentials api_key, base_url = _resolve_nous_pricing_credentials() if not base_url: return {} + _remember_provider_cache_key("nous", base_url.rstrip("/")) return _fetch_nous_pricing(api_key, base_url, force_refresh=force_refresh) -def get_pricing_for_provider(provider: str, *, force_refresh: bool = False) -> dict[str, dict[str, str]]: +_OPENROUTER_PRICING_BASE = "https://openrouter.ai/api" +_FIREWORKS_PRICING_KEY = "models.dev/fireworks" + + +def _ai_gateway_pricing_scope() -> str: + from hermes_constants import AI_GATEWAY_BASE_URL + return AI_GATEWAY_BASE_URL.rstrip("/") + + +def _novita_pricing_scope() -> str: + return (os.getenv("NOVITA_BASE_URL", "").strip() or "https://api.novita.ai/openai/v1").rstrip("/") + + +def get_cached_nous_inference_base_url() -> str: + """The profile's persisted Nous endpoint (bare origin, no ``/v1``) without refreshing auth.""" + try: + from hermes_cli.auth import ( + _load_auth_store, _load_provider_state, _optional_base_url, _validate_nous_inference_url_from_network, + ) + + state = _load_provider_state(_load_auth_store(), "nous") or {} + url = _validate_nous_inference_url_from_network(_optional_base_url(state.get("inference_base_url"))) or "" + return url.rstrip("/").removesuffix("/v1") + except Exception: + return "" + + +# Static endpoint identity per provider; dynamic ones (deepinfra, nous) are resolved in pricing_cache_scope. +_STATIC_PRICING_SCOPES = { + "openrouter": lambda: _OPENROUTER_PRICING_BASE, + "ai-gateway": _ai_gateway_pricing_scope, + "novita": _novita_pricing_scope, + "fireworks": lambda: _FIREWORKS_PRICING_KEY, +} + + +def pricing_cache_scope(provider: str, *, current_provider: str = "", current_base_url: str = "") -> str: + """The current endpoint identity a provider's pricing cache is keyed on. Resolves local configuration + only, never fetches: picker prewarm single-flight uses it so an endpoint rotation can start a new + worker while the previous endpoint is still slow or unreachable.""" + from hermes_cli.models import ( + _deepinfra_catalog_url, _pricing_profile_key, _pricing_provider_cache_keys, normalize_provider, + ) + normalized = normalize_provider(provider) + static = _STATIC_PRICING_SCOPES.get(normalized) + if static: + return static() + if normalized == "deepinfra": + return _deepinfra_catalog_url()[0] + if normalized == "nous": + try: + from hermes_cli.auth import _nous_inference_env_override + + env_base = _nous_inference_env_override() + except Exception: + env_base = None + if env_base: + return env_base.rstrip("/").removesuffix("/v1") + if normalize_provider(current_provider) == "nous" and current_base_url: + return current_base_url.rstrip("/").removesuffix("/v1") + persisted_base = get_cached_nous_inference_base_url() + if persisted_base: + return persisted_base + return _pricing_provider_cache_keys.get((_pricing_profile_key(), normalized), _DEFAULT_NOUS_INFERENCE_BASE) + return "" + + +def _cached_only_pricing(normalized: str) -> dict[str, dict[str, str]]: + """Process-resident pricing for *normalized* without any provider I/O.""" + from hermes_cli.models import ( + _deepinfra_catalog_cache, _deepinfra_catalog_url, _fetch_deepinfra_pricing, _pricing_profile_key, + _pricing_provider_cache_keys, + ) + if normalized == "deepinfra": + cache_key, _url = _deepinfra_catalog_url() + return _fetch_deepinfra_pricing() if cache_key in _deepinfra_catalog_cache else {} + cache_key = _pricing_provider_cache_keys.get((_pricing_profile_key(), normalized)) + if cache_key is None and normalized in ("openrouter", "ai-gateway", "fireworks"): + cache_key = _STATIC_PRICING_SCOPES[normalized]() + return (_cached_catalog(cache_key) or {}) if cache_key else {} + + +def get_pricing_for_provider( + provider: str, *, force_refresh: bool = False, cached_only: bool = False +) -> dict[str, dict[str, str]]: """Return live pricing for providers that support it (openrouter, nous, ai-gateway, novita, - deepinfra, fireworks); ``{}`` for everything else.""" + deepinfra, fireworks); ``{}`` for everything else. ``cached_only`` never starts provider I/O: + normal picker opens use it so cold endpoints cannot hold the response path, while a background + prewarm fills the same caches for later opens.""" from hermes_cli.models import normalize_provider - fetcher = _PRICING_FETCHERS.get(normalize_provider(provider)) + normalized = normalize_provider(provider) + if cached_only: + return _cached_only_pricing(normalized) + fetcher = _PRICING_FETCHERS.get(normalized) return fetcher(force_refresh=force_refresh) if fetcher else {} @@ -482,9 +597,9 @@ def _fetch_deepinfra_pricing(timeout: float = 5.0, *, force_refresh: bool = Fals _PRICING_FETCHERS = { "openrouter": _fetch_openrouter_pricing, - "ai-gateway": fetch_ai_gateway_pricing, - "novita": _fetch_novita_pricing, + "ai-gateway": _fetch_ai_gateway_pricing_for_provider, + "novita": _fetch_novita_pricing_for_provider, "deepinfra": _fetch_deepinfra_pricing, - "fireworks": _fireworks_pricing_from_models_dev, + "fireworks": _fetch_fireworks_pricing_for_provider, "nous": _fetch_nous_pricing_for_provider, } diff --git a/hermes_cli/observability/relay_shared_metrics.py b/hermes_cli/observability/relay_shared_metrics.py index 76990929b1..129905743e 100644 --- a/hermes_cli/observability/relay_shared_metrics.py +++ b/hermes_cli/observability/relay_shared_metrics.py @@ -678,8 +678,11 @@ class _Runtime: ) self._guarded( "Hermes shared-metrics tool call close failed", - self._run_in_task, task, self.relay.tools.call_end, tool_call.handle, fields, - metadata=self._event_metadata(), + lambda: self._run_in_task( + task, self.relay.tools.call_end, tool_call.handle, + self.relay.ToolExecutionResult(fields), + metadata=self._event_metadata(), + ), ) def _end_pending_tool_calls( diff --git a/hermes_cli/proxy/adapters/nous_portal.py b/hermes_cli/proxy/adapters/nous_portal.py index 25015ed072..4660cd223b 100644 --- a/hermes_cli/proxy/adapters/nous_portal.py +++ b/hermes_cli/proxy/adapters/nous_portal.py @@ -59,19 +59,22 @@ class NousPortalAdapter(UpstreamAdapter): def get_retry_credential( self, *, failed_credential: UpstreamCredential, status_code: int ) -> Optional[UpstreamCredential]: - _ = failed_credential if status_code != 401: return None logger.info("proxy: Nous upstream rejected bearer; force-refreshing invoke JWT") - return self._get_credential(force_refresh=True) + return self._get_credential(force_refresh=True, stale_access_token=failed_credential.bearer) - def _get_credential(self, *, force_refresh: bool = False) -> UpstreamCredential: + def _get_credential( + self, *, force_refresh: bool = False, stale_access_token: Optional[str] = None + ) -> UpstreamCredential: with self._lock: state = self._read_state() if state is None: raise RuntimeError("Not logged into Nous Portal. Run `hermes auth add nous` first.") try: - refreshed = resolve_nous_runtime_credentials(force_refresh=force_refresh) + refreshed = resolve_nous_runtime_credentials( + force_refresh=force_refresh, stale_access_token=stale_access_token or None + ) except Exception as exc: if isinstance(exc, AuthError) and _is_terminal_nous_refresh_error(exc): _quarantine_nous_oauth_state(state, exc, reason="proxy_refresh_failure") diff --git a/hermes_cli/session_lost_and_found.py b/hermes_cli/session_lost_and_found.py index bea2036ff5..9800515e44 100644 --- a/hermes_cli/session_lost_and_found.py +++ b/hermes_cli/session_lost_and_found.py @@ -4,6 +4,7 @@ from __future__ import annotations +import logging import re import shutil import sqlite3 @@ -17,6 +18,8 @@ from hermes_cli.session_recovery import ( _placeholder_titles, _quoted_columns, _table_columns, ) +logger = logging.getLogger(__name__) + # Hermes session ids are timestamps (20260812_135332_ab12cd): the strongest sentinel for schema-less rows. SESSION_ID_PATTERN = re.compile(r"^\d{8}_\d{6}_") MESSAGE_ROLES = frozenset({"user", "assistant", "tool", "system"}) @@ -36,6 +39,12 @@ SESSION_MODEL_USAGE_NFIELD = 18 # Plausible unix-epoch window for started_at heuristics on legacy layouts. _EPOCH_LOW = 1_000_000_000.0 # 2001 _EPOCH_HIGH = 4_000_000_000.0 # 2096 + +# Title label/prefix of every session row this lane synthesises (legacy-layout rows and stubbed parents). +# The recovery verifier keys on the prefix to tell synthesised rows from positionally mapped ones. +_STUB_TITLE_LABEL = "best-effort recovered" +STUB_TITLE_PREFIX = f"[{_STUB_TITLE_LABEL}" + SQLITE3_CLI_GUIDANCE = ( "A last-resort page-level salvage is available when a `.recover`-capable `sqlite3` command-line shell is " "installed: its `.recover` command can rebuild rows into lost_and_found tables even when the table schemas are " @@ -45,16 +54,94 @@ SQLITE3_CLI_GUIDANCE = ( "re-run with --allow-partial." ) +# SQLite's WAL-reset bug (https://sqlite.org/wal.html#walresetbug) lets a +# fresh opener unlink a live WAL/SHM sidecar pair and split the database into +# two concurrent generations whose acknowledged writes can silently vanish. +# It is real in CLI builds up to 3.51.2; fixed in 3.51.3+ with backports +# 3.50.7 and 3.44.6 — the same version gate hermes_state applies to the +# embedded library (#69784). The system `sqlite3` CLI on Debian/Ubuntu is +# routinely in the vulnerable band (e.g. 3.45.1), and #100368's forensics +# caught exactly this shell converting a live Hermes state.db into two +# generations. A salvage shell must therefore be version-gated, not just +# capability-gated, before it is pointed at (a copy of) a Hermes database. +# +# The predicate lives in hermes_cli.sqlite_runtime (stdlib-only, shared with +# the installer/update gates) so the embedded runtime and the salvage shell +# can never disagree about which versions are safe. +from hermes_cli.sqlite_runtime import is_sqlite_wal_reset_vulnerable as _wal_reset_vulnerable # noqa: E502 + +_WAL_RESET_VULNERABLE_GUIDANCE = ( + "salvage against a Hermes database with the WAL-reset bug " + "(https://sqlite.org/wal.html#walresetbug, fixed in 3.51.3+ / backports " + "3.50.7 / 3.44.6; the vulnerable fresh-opener can unlink a live WAL/SHM " + "pair and split the database into two generations, losing acknowledged " + "writes — #100368). Install a fixed sqlite3 CLI (3.51.3+, e.g. `brew " + "install sqlite` or the precompiled sqlite-tools from sqlite.org)" +) + class LostAndFoundError(RuntimeError): """Raised when the CLI .recover pass cannot produce a usable database.""" +def _parse_sqlite3_cli_version(binary: str) -> Optional[tuple[int, int, int]]: + """Version of the sqlite3 CLI at *binary* via ``--version``, or None when it cannot run or be parsed.""" + try: + probe = subprocess.run([binary, "--version"], capture_output=True, timeout=30) + except (OSError, subprocess.SubprocessError): + return None + if probe.returncode != 0: + return None + match = re.search(rb"(\d+)\.(\d+)\.(\d+)", probe.stdout) + if match is None: + return None + return tuple(int(part) for part in match.groups()) + + +_last_cli_refusal: dict[str, Any] = {} + + +def find_sqlite3_cli_refusal() -> dict[str, Any]: + """Why the last :func:`find_sqlite3_cli` call in this process refused: ``{"reason": ...}`` with reason in + ``missing``, ``no_dbpage`` (shell cannot run ``.recover``), ``wal_reset_vulnerable``; empty if it succeeded.""" + return dict(_last_cli_refusal) + + def find_sqlite3_cli() -> Optional[str]: - """A ``.recover``-capable sqlite3 CLI path, or None. PATH presence is not enough: distro builds can - lack the ``sqlite_dbpage`` virtual table ``.recover`` needs, so probe on a scratch DB once.""" + """A salvage-safe ``.recover``-capable sqlite3 CLI path, or None. + + PATH presence is not enough, and neither is ``.recover`` support alone: (1) distro builds can lack the + ``sqlite_dbpage`` virtual table ``.recover`` needs — probed once on a scratch DB; (2) a capable CLI can still + carry the WAL-reset opener bug (fixed 3.51.3+ / backports 3.50.7 / 3.44.6). The salvage lane only runs it on a + snapshot copy, but refusing it keeps vulnerable shells out of the documented workflow. Refusals are recorded + for :func:`find_sqlite3_cli_refusal` so callers can say exactly what to install. + """ + global _last_cli_refusal + _last_cli_refusal = {} binary = shutil.which("sqlite3") - return binary if binary is not None and _cli_supports_recover(binary) else None + if binary is None: + _last_cli_refusal = {"reason": "missing"} + return None + if not _cli_supports_recover(binary): + _last_cli_refusal = {"reason": "no_dbpage", "binary": binary} + return None + version = _parse_sqlite3_cli_version(binary) + if version is not None and _wal_reset_vulnerable(version): + version_str = ".".join(str(part) for part in version) + logger.warning( + "sqlite3 CLI %s reports version %s, which still carries the " + "WAL-reset opener bug; refusing to use it for salvage", + binary, + version_str, + ) + _last_cli_refusal = { + "reason": "wal_reset_vulnerable", + "binary": binary, + "version": version_str, + "detail": f"reports version {version_str}, which has " + _WAL_RESET_VULNERABLE_GUIDANCE, + } + return None + return binary def _cli_supports_recover(binary: str) -> bool: @@ -280,7 +367,7 @@ def map_lost_and_found_rows(lf_conn: sqlite3.Connection, dest: sqlite3.Connectio row_values = ( cells[0], cells[1] if _looks_like_source(cells[1]) else "recovered", _heuristic_started_at(cells), - "[best-effort recovered] legacy session row (layout unknown)", + f"{STUB_TITLE_PREFIX}] legacy session row (layout unknown)", ) inserted = dest.execute( "INSERT OR IGNORE INTO sessions (id, source, started_at, title) VALUES (?, ?, ?, ?)", @@ -322,7 +409,7 @@ def stub_missing_parent_sessions(dest: sqlite3.Connection) -> dict[str, Any]: "EXISTS (SELECT 1 FROM sessions WHERE sessions.id = u.session_id)" ): orphan_ids.setdefault(str(session_id), {"started_at": 0.0, "message_count": 0}) - titles = _placeholder_titles(dest, "best-effort recovered") + titles = _placeholder_titles(dest, _STUB_TITLE_LABEL) for session_id, info in sorted(orphan_ids.items()): title = next(titles) dest.execute( diff --git a/hermes_cli/session_recovery.py b/hermes_cli/session_recovery.py index b2640d43b6..9a1f1d4e4f 100644 --- a/hermes_cli/session_recovery.py +++ b/hermes_cli/session_recovery.py @@ -899,6 +899,49 @@ def _finalize_derived_metadata(destination: sqlite3.Connection) -> dict[str, Any return result +def _lost_and_found_plausibility_errors( + conn: sqlite3.Connection, +) -> list[str]: + """Flag systematic timestamp mis-mapping in a salvaged database. + + Structural checks (integrity, FK, FTS, row counts) pass on mis-mapped + salvage because every row still inserts. Only semantics give it away: + the physical column order of a source upgraded via ALTER TABLE differs + from the destination template's declared order, so positional cell + mapping lands counters/strings where ``started_at``/``timestamp`` + belong — and the NOT NULL substitutes turn gaps into 0.0. When every + mapped row violates the epoch floor, the mapping was wrong. + + Stub rows written by ``stub_missing_parent_sessions`` legitimately carry + ``started_at = 0.0`` when no timestamped message survived, so they are + excluded from the denominator. + """ + from hermes_cli.session_lost_and_found import _EPOCH_LOW, STUB_TITLE_PREFIX + + errors: list[str] = [] + checks = ( + ("sessions", "started_at", f"WHERE COALESCE(title, '') NOT LIKE '{STUB_TITLE_PREFIX}%'"), + ("messages", "timestamp", ""), + ) + for table, column, mapped_filter in checks: + (total,) = conn.execute(f"SELECT COUNT(*) FROM {table} {mapped_filter}").fetchone() + if not total: + continue + (implausible,) = conn.execute( + f"SELECT COUNT(*) FROM {table} {mapped_filter} " + f"{'AND' if mapped_filter else 'WHERE'} ({column} IS NULL OR {column} < ?)", + (_EPOCH_LOW,), + ).fetchone() + if implausible == total: + errors.append( + f"{table}.{column} is implausible in all {total} salvaged row(s) " + "(NULL or before 2001-09): the source's physical column order " + "did not match the destination template, so cells were mapped " + "onto the wrong columns" + ) + return errors + + def _recover_via_lost_and_found( *, source: Path, snapshot_source: Path, snapshot_dir: Path, output: Path, inspection: dict[str, Any], disk_space: dict[str, Any], missing_required: list[str], @@ -907,12 +950,22 @@ def _recover_via_lost_and_found( (shell-only, not in Python's ``sqlite3``) rebuilds rows into a scratch lost_and_found database which is then heuristically mapped into a fresh current-schema database.""" from hermes_cli.session_lost_and_found import ( - SQLITE3_CLI_GUIDANCE, LostAndFoundError, find_sqlite3_cli, map_lost_and_found_rows, rebuild_fts_indexes, - run_cli_lost_and_found_recover, stub_missing_parent_sessions, + SQLITE3_CLI_GUIDANCE, LostAndFoundError, find_sqlite3_cli, find_sqlite3_cli_refusal, map_lost_and_found_rows, + rebuild_fts_indexes, run_cli_lost_and_found_recover, stub_missing_parent_sessions, ) missing = ", ".join(missing_required) sqlite3_bin = find_sqlite3_cli() if sqlite3_bin is None: + refusal = find_sqlite3_cli_refusal() + if refusal.get("reason") == "wal_reset_vulnerable": + raise SessionRecoverySourceError( + "Partial recovery requires a page-level salvage shell, but " + "the only sqlite3 CLI on PATH is not safe to use for it: it " + + refusal["detail"] + + ". The readable table schemas for: " + + ", ".join(missing_required) + + " are still required." + ) raise SessionRecoverySourceError( f"Partial recovery still requires readable table schemas for: {missing}. {SQLITE3_CLI_GUIDANCE}" ) @@ -956,6 +1009,16 @@ def _recover_via_lost_and_found( "pages and mapped heuristically. Review every count before trusting this output." ) verification.update(loss_detected=True, complete=False) + # Structural checks cannot see a positional mis-mapping: every row still inserts, so integrity/FK/FTS + # stay green. A systematic timestamp violation is the semantic tell — never report such a salvage as verified. + plausibility_conn = sqlite3.connect(str(output), isolation_level=None) + try: + plausibility_errors = _lost_and_found_plausibility_errors(plausibility_conn) + finally: + plausibility_conn.close() + if plausibility_errors: + verification["errors"].extend(plausibility_errors) + verification["healthy"] = False return _recovery_report( source, output, inspection, disk_space, verification, on_source_change="healthy", allow_partial=True, mode="lost_and_found_salvage", best_effort=True, unreadable_schemas=missing_required, diff --git a/hermes_cli/update_cmd_config.py b/hermes_cli/update_cmd_config.py index fae62176d0..f453e991e4 100644 --- a/hermes_cli/update_cmd_config.py +++ b/hermes_cli/update_cmd_config.py @@ -35,7 +35,7 @@ def _run_config_check_fresh() -> tuple: from hermes_cli.update_cmd import _reload_config_modules _reload_config_modules() from hermes_cli.config import check_config_version - return check_config_version() + return check_config_version(raise_on_parse_error=True) def _run_migrate_config_fresh(*, interactive: bool = False, quiet: bool = False) -> dict: diff --git a/hermes_cli/update_cmd_maint.py b/hermes_cli/update_cmd_maint.py index c000156484..b479268f11 100644 --- a/hermes_cli/update_cmd_maint.py +++ b/hermes_cli/update_cmd_maint.py @@ -140,7 +140,7 @@ def _print_curator_first_run_notice() -> None: def _print_fts_optimize_available_notice() -> None: - """Advertise the opt-in v23 FTS optimization when state.db is still on the legacy layout. + """Advertise the opt-in FTS storage rebuild when state.db still needs one. ``sessions.fts_optimize_notice``: ``advise`` (default), ``require`` (firmer), ``off``. """ @@ -168,13 +168,15 @@ def _print_fts_optimize_available_notice() -> None: if size_gb < 0.5: return db = None + needs_upgrade = False try: db = SessionDB(db_path=db_path, read_only=True) - # read_only opens skip schema init; probe the layout directly. + # read_only opens skip schema init; probe the stored layout directly. row = db._conn.execute( "SELECT sql FROM sqlite_master " "WHERE type = 'table' AND name = 'messages_fts'" ).fetchone() + needs_upgrade = bool(row) and getattr(db, "_db_needs_fts_storage_upgrade")(db._conn) # Interrupted optimize-storage: v23 table shape but backfill markers / trash # tables remain. Re-running resumes it, so offer the command again. interrupted = bool( @@ -197,9 +199,8 @@ def _print_fts_optimize_available_notice() -> None: if db is not None: with suppress(Exception): db.close() - sql = (row[0] if row else "") or "" - if not sql or ("tool_name" in sql and not interrupted): - return + if not needs_upgrade and not interrupted: + return # current layout already present (fresh/optimized) if interrupted: print() diff --git a/hermes_constants.py b/hermes_constants.py index bc37ef72b9..067c1c695d 100644 --- a/hermes_constants.py +++ b/hermes_constants.py @@ -89,13 +89,43 @@ def get_hermes_home() -> Path: return get_process_hermes_home() +# Resolved keys, keyed by the path string that was handed in. Path.resolve() +# is a filesystem call, and this function sits under every ToolRegistry +# lookup through current_scope_key(), so without this the registry pays a +# syscall per lookup. A process only ever sees a handful of home paths, so +# the dict stays tiny. Only paths that really exist are stored, see below. +_HOME_KEY_CACHE: dict[str, str] = {} + + def hermes_home_key(path: str | Path | None = None) -> str: """Stable registry key for a Hermes home/profile dir. ``strict=False`` so profiles whose directories don't exist yet still get a key. + + The resolved value is remembered per input path. A directory that does + not exist yet is resolved without touching the cache, because the answer + can change once it is created (e.g. part of the path turns out to be a symlink). """ candidate = Path(path) if path is not None else get_hermes_home() - return os.path.normcase(str(candidate.expanduser().resolve(strict=False))) + raw = str(candidate) + cached = _HOME_KEY_CACHE.get(raw) + if cached is not None: + return cached + expanded = candidate.expanduser() + try: + resolved = expanded.resolve(strict=True) + except OSError: + # Not on disk yet: lenient resolve, not stored, so the real answer is + # picked up once the directory appears. + return os.path.normcase(str(expanded.resolve(strict=False))) + key = os.path.normcase(str(resolved)) + _HOME_KEY_CACHE[raw] = key + return key + + +def reset_hermes_home_key_cache() -> None: + """Forget every remembered home key (for tests that move a home dir on disk).""" + _HOME_KEY_CACHE.clear() def get_process_hermes_home() -> Path: diff --git a/hermes_state.py b/hermes_state.py index e782dd6ec5..586d975bee 100644 --- a/hermes_state.py +++ b/hermes_state.py @@ -58,10 +58,10 @@ from hermes_state_fts import SessionFtsSetupMixin, load_fts5_cjk_extension # no from hermes_state_portability import SessionPortabilityMixin from hermes_state_telegram import SessionTelegramTopicsMixin from hermes_state_schema import SessionSchemaMixin +import hermes_state_holders as _state_holders from hermes_state_dbfile import ( # noqa: F401 (re-exported; tests patch hermes_state.) _canonical_sqlite_path, _concrete_state_db_holder_pids, _connect_tracked_db, - _is_inactive_orphan_desktop_holder, _looks_like_hermes, _read_proc_cmdline, - _read_sqlite_application_id, _stat_sqlite_sidecar_identity, _watched_sqlite_sidecar_paths, + _is_inactive_orphan_desktop_holder, _read_sqlite_application_id, _stat_sqlite_sidecar_identity, _watched_sqlite_sidecar_paths, collect_state_db_stats, count_db_holders, is_zeroed_state_db, iter_deleted_sqlite_sidecar_holders, quarantine_cross_process_lock, quarantine_zeroed_state_db, refuse_deleted_wal_generation, ) @@ -81,7 +81,8 @@ from hermes_state_repair import ( # noqa: F401 (re-exported; tests patch herme _persistent_repair_attempts_exhausted, _probe_journal_mode_for_repair, _prune_malformed_backups, _read_repair_ledger, _record_repair_outcome, _release_auto_maintenance_lock, _repair_backup_headroom_bytes, _repair_ledger_path, _repair_scratch_space_error, - _repair_snapshot_timeout_seconds, _repair_state_db_schema_locked, _run_repair_strategies, + _repair_snapshot_timeout_seconds, _repair_state_db_schema_locked, _restore_journal_mode_after_repair, + _exclusive_repair_db_guard, _run_repair_strategies, _try_acquire_auto_maintenance_lock, _unlink_db_triple, apply_durability_barriers, preflight_db_writability, repair_state_db_schema, ) @@ -337,6 +338,11 @@ def divert_session_transcript_jsonl(session_id: str, messages) -> "Optional[Path # Process-wide shared SessionDB registry: long-lived in-process callers share ONE writer # connection per resolved path via get_shared_session_db(); one-shots use SessionDB() + close(). +def _foreign_state_db_holders(db_path: Path) -> List[Tuple[int, str]]: + """Compatibility delegate to the state-holder authority.""" + return _state_holders.foreign_state_db_holders(db_path) + + from hermes_state_registry import ( # noqa: F401 (re-export) close_shared_session_dbs, get_shared_session_db, release_or_close, ) @@ -1059,53 +1065,8 @@ class SessionDB( return True def _foreign_state_db_holders(self) -> List[Tuple[int, str]]: - """Foreign processes holding this DB or its WAL sidecars: automatic FTS - repair must not run while another process is attached (a sidecar reset - under it splits the WAL inodes). A scan failure is reported as an unknown - holder — skipping optional maintenance beats assuming quiescence.""" - # Split-brain needs POSIX unlink semantics (Windows refuses to replace open sidecars); - # psutil.open_files() there can block for minutes. - if _IS_WINDOWS: - return [] - if psutil is None: - return [(-1, "open-file scan unavailable")] - db_path = os.path.abspath(os.fspath(self.db_path)) - watched = {_canonical_sqlite_path(db_path + suffix) for suffix in ("", "-wal", "-shm")} - holders: List[Tuple[int, str]] = [] - own_pid = os.getpid() - try: - if sys.platform.startswith("linux"): - # readlink /proc//fd directly; psutil.open_files() stats the literal path - # and silently drops "state.db-wal (deleted)" entries. - for pid in (int(p) for p in os.listdir("/proc") if p.isdigit()): - if pid == own_pid: - continue - try: - targets = list(_proc_fd_targets(pid)) - except OSError: - # Unreadable fd table (other user); flag only Hermes-looking holders via cmdline. - cmdline = _read_proc_cmdline(pid) - if cmdline is not None and _looks_like_hermes(cmdline): - holders.append((pid, f"uninspectable holder: {cmdline[:80]}")) - continue - holders.extend((pid, t) for t in targets if _canonical_sqlite_path(t) in watched) - else: - # macOS / BSD: psutil.open_files() (no "(deleted)" convention; AccessDenied -> empty). - for process in psutil.process_iter(["pid", "open_files"]): - pid = int(process.info["pid"]) - if pid == own_pid: - continue - for opened in process.info.get("open_files") or (): - path = getattr(opened, "path", "") - if path and _canonical_sqlite_path(path) in watched: - holders.append((pid, path)) - except Exception as exc: - logger.warning( - "Could not prove state.db has no foreign holders; " - "deferring automatic FTS maintenance: %s", exc, - ) - return holders or [(-1, f"open-file scan failed: {exc}")] - return holders + """Foreign processes holding this DB or its WAL sidecars (see hermes_state_holders).""" + return _foreign_state_db_holders(self.db_path) def _try_wal_checkpoint(self) -> None: """Best-effort PASSIVE WAL checkpoint; never raises. PASSIVE never blocks writers; diff --git a/hermes_state_common.py b/hermes_state_common.py index 376a79512f..6eb00fdf83 100644 --- a/hermes_state_common.py +++ b/hermes_state_common.py @@ -185,15 +185,42 @@ def _sql_session_last_active_by_id(session_id_expr: str) -> str: f"(SELECT started_at FROM sessions _act_s WHERE _act_s.id = {session_id_expr})") -SCHEMA_VERSION = 28 +SCHEMA_VERSION = 30 # Auto-maintenance VACUUMs only above this freelist fraction; below it a rewrite costs more I/O than it returns. AUTO_VACUUM_MIN_FREELIST_RATIO = 0.25 -# FTS layout, tracked INDEPENDENTLY of SCHEMA_VERSION (state_meta ``fts_storage_version``): it changes only -# when a DB is born fresh or via ``hermes sessions optimize-storage``. 0 (marker absent) = legacy inline -# index, still working; 1 = v23 external-content layout. -FTS_STORAGE_VERSION = 1 +# FTS storage-layout version, tracked INDEPENDENTLY of SCHEMA_VERSION in the +# state_meta key ``fts_storage_version``. The main schema version advances +# freely on open (so future migrations always land); the FTS *layout* only +# reaches the current version when a DB is either born fresh or explicitly +# optimized via ``hermes sessions optimize-storage``. A legacy DB sits at +# layout 0 (marker absent) with a working inline index until the user opts in. +# 1 = v23 external-content layout with a tool-row-excluded trigram +# 2 = trigram also excludes structured tool_calls JSON +FTS_STORAGE_VERSION = 2 + +# Tool results are often multi-megabyte machine payloads. Index a useful +# prefix for new tool rows instead of tokenizing the entire body while the +# canonical message write holds SQLite's single writer lock. The high-water +# marker lets upgraded databases retain the exact token stream already stored +# for historical rows, so external-content delete/update commands stay valid +# without an eager full-index rebuild. +FTS_TOOL_CONTENT_PREFIX_CHARS = 8_192 +FTS_TOOL_FULL_CONTENT_HIGH_WATER_KEY = "fts_tool_full_content_high_water" + + +def _fts_indexed_content_sql(alias: str) -> str: + return f"""CASE WHEN {alias}.role = 'tool' + AND {alias}.id > COALESCE((SELECT CAST(value AS INTEGER) + FROM state_meta + WHERE key = '{FTS_TOOL_FULL_CONTENT_HIGH_WATER_KEY}'), -1) + THEN substr(COALESCE({alias}.content, ''), 1, {FTS_TOOL_CONTENT_PREFIX_CHARS}) + ELSE {alias}.content END""" + + +_FTS_NEW_INDEXED_CONTENT_SQL = _fts_indexed_content_sql("new") +_FTS_OLD_INDEXED_CONTENT_SQL = _fts_indexed_content_sql("old") # Cap on user-controlled FTS5 query input before sanitizer processing. MAX_FTS5_QUERY_CHARS = 2_048 @@ -483,14 +510,34 @@ 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); +-- Recent-session browsing must never derive recency by scanning messages. +-- This expression is the durable, indexable approximation used to preselect +-- a small candidate set before compression-chain and preview hydration. +CREATE INDEX IF NOT EXISTS idx_sessions_effective_activity + ON sessions(COALESCE(last_activity_at, started_at) DESC, started_at DESC); """ -# Deferred FTS rebuild bookkeeping: while a background rebuild is pending, state_meta H = fts_rebuild_high_water -# (MAX(messages.id) when the old indexes were dropped) and P = fts_rebuild_progress (highest backfilled id) -# define the indexed rows, id <= P OR id > H (post-drop rows are indexed live by the insert triggers). Every -# trigger gates on that predicate: an external-content 'delete' for an unindexed row corrupts the index and -# skipping one for an indexed row leaves a stale entry. No rebuild pending => both absent, COALESCE tautology. -FTS_SQL = """ + +# ── Deferred FTS rebuild bookkeeping (schema v23) ── +# While a background index rebuild is pending, two state_meta keys define +# which message rows are currently IN the FTS indexes: +# +# fts_rebuild_high_water H — MAX(messages.id) at the moment the old +# indexes were dropped +# fts_rebuild_progress P — highest id the chunked backfill has indexed +# +# A row is indexed iff id <= P (backfilled) OR id > H (inserted after +# the drop; ids are AUTOINCREMENT so new rows are always > H and the insert +# triggers index them live). Rows in (P, H] are not yet indexed. +# +# Every trigger below gates on that same predicate: firing an FTS5 +# external-content 'delete' for a row that is NOT in the index corrupts the +# index, and skipping it for a row that IS indexed leaves a stale entry. +# When no rebuild is pending both keys are absent and COALESCE turns the +# predicate into a tautology (id > -1 OR id <= -1), i.e. normal operation. +# The two state_meta PK probes per write are negligible next to the FTS +# insert itself. +FTS_SQL = f""" CREATE VIRTUAL TABLE IF NOT EXISTS messages_fts USING fts5( content, tool_name, @@ -506,7 +553,12 @@ WHEN (new.id > COALESCE((SELECT CAST(value AS INTEGER) FROM state_meta WHERE key = 'fts_rebuild_progress'), -1)) BEGIN INSERT INTO messages_fts(rowid, content, tool_name, tool_calls) - VALUES (new.id, new.content, new.tool_name, new.tool_calls); + VALUES ( + new.id, + {_FTS_NEW_INDEXED_CONTENT_SQL}, + new.tool_name, + new.tool_calls + ); END; CREATE TRIGGER IF NOT EXISTS messages_fts_delete AFTER DELETE ON messages @@ -516,71 +568,19 @@ WHEN (old.id > COALESCE((SELECT CAST(value AS INTEGER) FROM state_meta WHERE key = 'fts_rebuild_progress'), -1)) BEGIN INSERT INTO messages_fts(messages_fts, rowid, content, tool_name, tool_calls) - VALUES ('delete', old.id, old.content, old.tool_name, old.tool_calls); + VALUES ( + 'delete', + old.id, + {_FTS_OLD_INDEXED_CONTENT_SQL}, + old.tool_name, + old.tool_calls + ); END; -- 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) - AND (old.id > COALESCE((SELECT CAST(value AS INTEGER) FROM state_meta - WHERE key = 'fts_rebuild_high_water'), -1) - OR old.id <= COALESCE((SELECT CAST(value AS INTEGER) FROM state_meta - WHERE key = 'fts_rebuild_progress'), -1)) -BEGIN - INSERT INTO messages_fts(messages_fts, rowid, content, tool_name, tool_calls) - VALUES ('delete', old.id, old.content, old.tool_name, old.tool_calls); - INSERT INTO messages_fts(rowid, content, tool_name, tool_calls) - VALUES (new.id, new.content, new.tool_name, new.tool_calls); -END; -""" - -# Trigram FTS5 table for CJK substring search (unicode61 splits CJK into single tokens). The trigram index is -# ~2.6x the text it covers and ``role='tool'`` rows are ~90% of message bytes, so it reads through the -# ``messages_fts_trigram_src`` view excluding tool rows; those stay searchable via ``messages_fts`` and -# ``search_messages`` routes CJK queries filtered on role='tool' to LIKE. -FTS_TRIGRAM_SQL = """ -CREATE VIEW IF NOT EXISTS messages_fts_trigram_src AS - SELECT id, role, content, tool_name, tool_calls - FROM messages - WHERE role <> 'tool'; - -CREATE VIRTUAL TABLE IF NOT EXISTS messages_fts_trigram USING fts5( - content, - tool_name, - tool_calls, - content='messages_fts_trigram_src', - content_rowid='id', - tokenize='trigram' -); - -CREATE TRIGGER IF NOT EXISTS messages_fts_trigram_insert AFTER INSERT ON messages -WHEN new.role <> 'tool' - AND (new.id > COALESCE((SELECT CAST(value AS INTEGER) FROM state_meta - WHERE key = 'fts_rebuild_high_water'), -1) - OR new.id <= COALESCE((SELECT CAST(value AS INTEGER) FROM state_meta - WHERE key = 'fts_rebuild_progress'), -1)) -BEGIN - INSERT INTO messages_fts_trigram(rowid, content, tool_name, tool_calls) - VALUES (new.id, new.content, new.tool_name, new.tool_calls); -END; - -CREATE TRIGGER IF NOT EXISTS messages_fts_trigram_delete AFTER DELETE ON messages -WHEN old.role <> 'tool' - AND (old.id > COALESCE((SELECT CAST(value AS INTEGER) FROM state_meta - WHERE key = 'fts_rebuild_high_water'), -1) - OR old.id <= COALESCE((SELECT CAST(value AS INTEGER) FROM state_meta - WHERE key = 'fts_rebuild_progress'), -1)) -BEGIN - INSERT INTO messages_fts_trigram(messages_fts_trigram, rowid, content, tool_name, tool_calls) - VALUES ('delete', old.id, old.content, old.tool_name, old.tool_calls); -END; - -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 @@ -591,12 +591,129 @@ WHEN (old.content IS NOT new.content OR old.id <= COALESCE((SELECT CAST(value AS INTEGER) FROM state_meta WHERE key = 'fts_rebuild_progress'), -1)) BEGIN - INSERT INTO messages_fts_trigram(messages_fts_trigram, rowid, content, tool_name, tool_calls) - SELECT 'delete', old.id, old.content, old.tool_name, old.tool_calls - WHERE old.role <> 'tool'; - INSERT INTO messages_fts_trigram(rowid, content, tool_name, tool_calls) - SELECT new.id, new.content, new.tool_name, new.tool_calls - WHERE new.role <> 'tool'; + INSERT INTO messages_fts(messages_fts, rowid, content, tool_name, tool_calls) + VALUES ( + 'delete', + old.id, + {_FTS_OLD_INDEXED_CONTENT_SQL}, + old.tool_name, + old.tool_calls + ); + INSERT INTO messages_fts(rowid, content, tool_name, tool_calls) + VALUES ( + new.id, + {_FTS_NEW_INDEXED_CONTENT_SQL}, + new.tool_name, + new.tool_calls + ); +END; +""" + + +# Trigram FTS5 table for CJK substring search. The default unicode61 +# tokenizer splits CJK characters into individual tokens, breaking phrase +# matching. The trigram tokenizer creates overlapping 3-byte sequences so +# substring queries work natively for any script (CJK, Thai, etc.). +# +# The trigram index is the most expensive index in state.db (~2.6x the size +# of the text it covers). Tool output (~90% of message bytes, machine noise) +# and cron transcripts are excluded: the index reads through +# ``messages_fts_trigram_src``, a view that skips both classes. They stay +# fully stored in ``messages`` and searchable via the standard +# ``messages_fts`` index; they just don't get trigram (CJK substring) +# treatment. ``search_messages`` routes explicit tool/cron CJK searches to +# LIKE for the same reason. Structured ``tool_calls`` JSON likewise stays +# searchable through ``messages_fts``; excluding it here avoids indexing +# repetitive JSON syntax as trigrams (FTS_STORAGE_VERSION 2). +# +# Delegate-child (subagent) transcripts are excluded the same way (v30): +# on a fan-out-heavy install they were ~70% of all message bytes and +# ``session_search`` hides ``source='subagent'`` sessions anyway. A child +# is recognised by its source OR by the ``_delegate_from`` creation marker +# (children spawned under a gateway turn inherit the gateway's source). +# Compression/branch continuations of interactive sessions also carry +# ``parent_session_id`` but NOT the marker, so they stay trigram-indexed. +FTS_TRIGRAM_EXCLUDED_SOURCES = ("cron", "subagent") + +# Predicate over a ``sessions`` row (unqualified column names) selecting +# sessions whose rows belong in the trigram index. Shared by the view, the +# sync triggers, and the deferred-backfill INSERT ... SELECTs so they can +# never disagree about the index boundary. +FTS_TRIGRAM_SESSION_SQL = ( + "source NOT IN (" + + ", ".join(f"'{src}'" for src in FTS_TRIGRAM_EXCLUDED_SOURCES) + + ") AND json_extract(COALESCE(model_config, '{}'), '$._delegate_from') IS NULL" +) + + +def fts_trigram_session_sql(alias: str) -> str: + """``FTS_TRIGRAM_SESSION_SQL`` with every column qualified by ``alias``.""" + return FTS_TRIGRAM_SESSION_SQL.replace("source ", f"{alias}.source ").replace( + "COALESCE(model_config", f"COALESCE({alias}.model_config" + ) + + +FTS_TRIGRAM_SQL = f""" +CREATE VIEW IF NOT EXISTS messages_fts_trigram_src AS + SELECT m.id, m.role, m.content, m.tool_name + FROM messages AS m + JOIN sessions AS s ON s.id = m.session_id + WHERE m.role <> 'tool' AND {fts_trigram_session_sql('s')}; + +CREATE VIRTUAL TABLE IF NOT EXISTS messages_fts_trigram USING fts5( + content, + tool_name, + content='messages_fts_trigram_src', + content_rowid='id', + tokenize='trigram' +); + +CREATE TRIGGER IF NOT EXISTS messages_fts_trigram_insert AFTER INSERT ON messages +WHEN new.role <> 'tool' + AND EXISTS (SELECT 1 FROM sessions + WHERE id = new.session_id AND {FTS_TRIGRAM_SESSION_SQL}) + AND (new.id > COALESCE((SELECT CAST(value AS INTEGER) FROM state_meta + WHERE key = 'fts_rebuild_high_water'), -1) + OR new.id <= COALESCE((SELECT CAST(value AS INTEGER) FROM state_meta + WHERE key = 'fts_rebuild_progress'), -1)) +BEGIN + INSERT INTO messages_fts_trigram(rowid, content, tool_name) + VALUES (new.id, new.content, new.tool_name); +END; + +CREATE TRIGGER IF NOT EXISTS messages_fts_trigram_delete AFTER DELETE ON messages +WHEN old.role <> 'tool' + AND EXISTS (SELECT 1 FROM sessions + WHERE id = old.session_id AND {FTS_TRIGRAM_SESSION_SQL}) + AND (old.id > COALESCE((SELECT CAST(value AS INTEGER) FROM state_meta + WHERE key = 'fts_rebuild_high_water'), -1) + OR old.id <= COALESCE((SELECT CAST(value AS INTEGER) FROM state_meta + WHERE key = 'fts_rebuild_progress'), -1)) +BEGIN + INSERT INTO messages_fts_trigram(messages_fts_trigram, rowid, content, tool_name) + VALUES ('delete', old.id, old.content, old.tool_name); +END; + +CREATE TRIGGER IF NOT EXISTS messages_fts_trigram_update +AFTER UPDATE OF content, tool_name, role ON messages +WHEN (old.content IS NOT new.content + OR old.tool_name IS NOT new.tool_name + OR old.role IS NOT new.role) + AND (old.id > COALESCE((SELECT CAST(value AS INTEGER) FROM state_meta + WHERE key = 'fts_rebuild_high_water'), -1) + OR old.id <= COALESCE((SELECT CAST(value AS INTEGER) FROM state_meta + WHERE key = 'fts_rebuild_progress'), -1)) +BEGIN + INSERT INTO messages_fts_trigram(messages_fts_trigram, rowid, content, tool_name) + SELECT 'delete', old.id, old.content, old.tool_name + WHERE old.role <> 'tool' + AND EXISTS (SELECT 1 FROM sessions + WHERE id = old.session_id AND {FTS_TRIGRAM_SESSION_SQL}); + INSERT INTO messages_fts_trigram(rowid, content, tool_name) + SELECT new.id, new.content, new.tool_name + WHERE new.role <> 'tool' + AND EXISTS (SELECT 1 FROM sessions + WHERE id = new.session_id AND {FTS_TRIGRAM_SESSION_SQL}); END; """ @@ -613,10 +730,19 @@ FTS_STALE_KEY = "fts_stale" # Durable diagnostic for stale FTS recovery blocked across process restarts. FTS_REBUILD_DEFERRAL_KEY = "fts_rebuild_deferral" -# Legacy (v22 / inline-content) FTS DDL: ONLY keeps a pre-v23 install's search working and its triggers -# repairable until `optimize_fts_storage()` migrates it (inline content || tool_name || tool_calls, trigram over -# every row). Never created fresh; the v23 DDL on a legacy DB would leave a mixed, broken state. -LEGACY_FTS_SQL = """ + +# ── Legacy (v22 / inline-content) FTS DDL ────────────────────────────── +# Used ONLY to keep an existing pre-v23 install's search working and its +# triggers repairable UNTIL the user opts into `hermes db optimize`. This is +# the exact inline shape v11..v22 shipped: each virtual table stores its own +# copy of ``content || tool_name || tool_calls`` and the trigram table indexes +# every row (including role='tool'). We never CREATE these on a fresh install — +# fresh installs are born on the v23 external-content schema above. These +# constants exist so a legacy DB is never accidentally handed the v23 DDL +# (which would create the external-content trigram source VIEW and leave the +# DB in a mixed, broken state). `optimize_fts_storage()` is what migrates a +# legacy DB to the v23 shape. +LEGACY_FTS_SQL = f""" CREATE VIRTUAL TABLE IF NOT EXISTS messages_fts USING fts5( content ); @@ -624,7 +750,8 @@ CREATE VIRTUAL TABLE IF NOT EXISTS messages_fts USING fts5( CREATE TRIGGER IF NOT EXISTS 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, '') + COALESCE({_FTS_NEW_INDEXED_CONTENT_SQL}, '') + || ' ' || COALESCE(new.tool_name, '') || ' ' || COALESCE(new.tool_calls, '') ); END; @@ -633,16 +760,18 @@ CREATE TRIGGER IF NOT EXISTS messages_fts_delete AFTER DELETE ON messages BEGIN END; CREATE TRIGGER IF NOT EXISTS messages_fts_update -AFTER UPDATE OF content, tool_name, tool_calls ON messages BEGIN +AFTER UPDATE OF content, tool_name, tool_calls, role 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, '') + COALESCE({_FTS_NEW_INDEXED_CONTENT_SQL}, '') + || ' ' || COALESCE(new.tool_name, '') || ' ' || COALESCE(new.tool_calls, '') ); END; """ -LEGACY_FTS_TRIGRAM_SQL = """ + +LEGACY_FTS_TRIGRAM_SQL = f""" CREATE VIRTUAL TABLE IF NOT EXISTS messages_fts_trigram USING fts5( content, tokenize='trigram' @@ -651,7 +780,8 @@ CREATE VIRTUAL TABLE IF NOT EXISTS messages_fts_trigram USING fts5( CREATE TRIGGER IF NOT EXISTS messages_fts_trigram_insert AFTER INSERT ON messages BEGIN INSERT INTO messages_fts_trigram(rowid, content) VALUES ( new.id, - COALESCE(new.content, '') || ' ' || COALESCE(new.tool_name, '') || ' ' || COALESCE(new.tool_calls, '') + COALESCE({_FTS_NEW_INDEXED_CONTENT_SQL}, '') + || ' ' || COALESCE(new.tool_name, '') || ' ' || COALESCE(new.tool_calls, '') ); END; @@ -660,11 +790,12 @@ CREATE TRIGGER IF NOT EXISTS messages_fts_trigram_delete AFTER DELETE ON message END; CREATE TRIGGER IF NOT EXISTS messages_fts_trigram_update -AFTER UPDATE OF content, tool_name, tool_calls ON messages BEGIN +AFTER UPDATE OF content, tool_name, tool_calls, role ON messages BEGIN DELETE FROM messages_fts_trigram WHERE rowid = old.id; INSERT INTO messages_fts_trigram(rowid, content) VALUES ( new.id, - COALESCE(new.content, '') || ' ' || COALESCE(new.tool_name, '') || ' ' || COALESCE(new.tool_calls, '') + COALESCE({_FTS_NEW_INDEXED_CONTENT_SQL}, '') + || ' ' || COALESCE(new.tool_name, '') || ' ' || COALESCE(new.tool_calls, '') ); END; """ diff --git a/hermes_state_dbfile.py b/hermes_state_dbfile.py index 1c2b10ca85..b0f208959f 100644 --- a/hermes_state_dbfile.py +++ b/hermes_state_dbfile.py @@ -378,23 +378,3 @@ def _concrete_state_db_holder_pids(db_path: Path, holders: List[Tuple[int, str]] watched = {canonical_db, canonical_db + "-wal", canonical_db + "-shm"} return list(dict.fromkeys( pid for pid, path in holders if pid > 0 and _canonical_sqlite_path(path) in watched)) - - -def _read_proc_cmdline(pid: int) -> Optional[str]: - """Space-joined /proc//cmdline (readable even when the fd table is not); None if unreadable.""" - try: - with open(f"/proc/{pid}/cmdline", "rb") as f: - raw = f.read() - return raw.replace(b"\x00", b" ").decode("utf-8", "replace").strip() if raw else None - except OSError: - return None - - -_HERMES_CMDLINE_MARKERS = ("hermes_cli.main", "hermes_cli/main", "hermes serve", - "hermes-agent", "hermes gateway", "hermes chat") - - -def _looks_like_hermes(cmdline: str) -> bool: - """Heuristic: is this a Hermes process? Decides whether an uninspectable process (fd - table unreadable, other user) counts as a potential state.db holder; daemons are not.""" - return any(marker in cmdline.lower() for marker in _HERMES_CMDLINE_MARKERS) diff --git a/hermes_state_fts.py b/hermes_state_fts.py index ce3c977b98..40b596d396 100644 --- a/hermes_state_fts.py +++ b/hermes_state_fts.py @@ -140,6 +140,19 @@ class SessionFtsSetupMixin: ).fetchone() return row is not None and "tool_name" not in (row[0] or "") + @staticmethod + def _db_has_trigram_tool_calls_projection(cursor: sqlite3.Cursor) -> bool: + """True when the trigram vtable still includes the tool_calls payload (FTS_STORAGE_VERSION 1).""" + row = cursor.execute( + "SELECT sql FROM sqlite_master WHERE type = 'table' AND name = 'messages_fts_trigram'" + ).fetchone() + return row is not None and "tool_calls" in (row[0] or "").lower() + + @classmethod + def _db_needs_fts_storage_upgrade(cls, cursor: sqlite3.Cursor) -> bool: + """True when the current FTS storage layout should be treated as stale (optimize-storage has work).""" + return cls._db_has_legacy_inline_fts(cursor) or cls._db_has_trigram_tool_calls_projection(cursor) + def _warn_trigram_unavailable(self, exc: sqlite3.OperationalError) -> None: """Log once that the trigram tokenizer is missing; base FTS5 stays enabled.""" if getattr(self, "_trigram_unavailable_warned", False): # attr is lazily created here diff --git a/hermes_state_holders.py b/hermes_state_holders.py new file mode 100644 index 0000000000..9f4176d40a --- /dev/null +++ b/hermes_state_holders.py @@ -0,0 +1,292 @@ +"""Process and descriptor authority for state.db structural maintenance. + +This module owns the proof that no foreign process still holds the active or +an unlinked SQLite DB/WAL/SHM generation. ``hermes_state`` supplies only the +SQLite connection factory needed by the final lock probe. +""" + +from __future__ import annotations + +import errno +import logging +import os +import sqlite3 +import sys +from pathlib import Path +from typing import Callable, List, Optional, Sequence, Set, Tuple + +try: # Hard dependency, but tolerate scaffold-phase imports before pip install. + import psutil +except ImportError: # pragma: no cover - stripped/scaffold installs only + psutil = None # type: ignore[assignment] + + +logger = logging.getLogger(__name__) + +_IS_WINDOWS = sys.platform == "win32" +_HERMES_EXECUTABLES = frozenset({"hermes", "hermes-agent", "hermes-acp"}) +_HERMES_PYTHON_MODULES = frozenset({"acp_adapter", "hermes_cli.main"}) +_HERMES_PYTHON_SCRIPTS = frozenset({"hermes_cli/main.py", "run_agent.py"}) +_PYTHON_SHORT_OPTIONS_WITH_OPERANDS = frozenset({"Q", "W", "X"}) +_PYTHON_LONG_OPTIONS_WITH_OPERANDS = frozenset( + {"--check-hash-based-pycs", "--jit"} +) + + +def _read_proc_argv(pid: int) -> Optional[List[str]]: + """Read /proc//cmdline without losing argv boundaries.""" + try: + with open(f"/proc/{pid}/cmdline", "rb") as handle: + raw = handle.read() + if not raw: + return None + argv = raw.decode("utf-8", "replace").split("\x00") + if argv[-1] == "": + argv.pop() + return argv or None + except OSError: + return None + + +def _looks_like_python_executable(program: str) -> bool: + name = os.path.basename(program).lower().removesuffix(".exe") + for prefix in ("python", "pypy"): + if name.startswith(prefix): + suffix = name[len(prefix) :] + return not suffix or all(char.isdigit() or char == "." for char in suffix) + return False + + +def _python_execution_target(argv: Sequence[str]) -> Optional[Tuple[str, str]]: + """Return the Python module or script selected by interpreter options.""" + index = 1 + while index < len(argv): + arg = argv[index] + if arg == "--": + index += 1 + return ("script", argv[index]) if index < len(argv) else None + if arg in _PYTHON_LONG_OPTIONS_WITH_OPERANDS: + index += 2 + continue + if arg.startswith("--check-hash-based-pycs=") or arg.startswith("--jit="): + index += 1 + continue + if arg.startswith("--"): + index += 1 + continue + if arg.startswith("-") and arg != "-": + options = arg[1:] + option_index = 0 + consumed_next = False + while option_index < len(options): + option = options[option_index] + attached = options[option_index + 1 :] + if option == "c": + return None + if option == "m": + if attached: + return "module", attached + index += 1 + return ("module", argv[index]) if index < len(argv) else None + if option in _PYTHON_SHORT_OPTIONS_WITH_OPERANDS: + consumed_next = not attached + break + option_index += 1 + index += 2 if consumed_next else 1 + continue + return "script", arg + return None + + +def _looks_like_hermes(argv: Sequence[str]) -> bool: + """Return whether argv identifies a supported Hermes execution target.""" + if not argv: + return False + program = os.path.basename(argv[0]).lower().removesuffix(".exe") + if program in _HERMES_EXECUTABLES: + return True + if not _looks_like_python_executable(program): + return False + target = _python_execution_target(argv) + if target is None: + return False + kind, value = target + normalized = value.lower().replace("\\", "/") + if kind == "module": + return normalized in _HERMES_PYTHON_MODULES + return any( + normalized == script or normalized.endswith(f"/{script}") + for script in _HERMES_PYTHON_SCRIPTS + ) + + +def canonical_sqlite_path(path: str) -> str: + """Normalize a /proc fd target, stripping the Linux `` (deleted)`` suffix.""" + return os.path.normcase(os.path.abspath(path.removesuffix(" (deleted)"))) + + +def foreign_state_db_holders(db_path: Path) -> List[Tuple[int, str]]: + """Return foreign holders of the DB or one of its WAL sidecars. + + A scan failure is represented as an unknown holder. Structural maintenance + must not assume quiescence when an old, unlinked SQLite generation may + still be open by another process. + """ + if _IS_WINDOWS: + return [] + + db_path_str = os.path.abspath(os.fspath(db_path)) + watched = { + canonical_sqlite_path(db_path_str), + canonical_sqlite_path(db_path_str + "-wal"), + canonical_sqlite_path(db_path_str + "-shm"), + } + holders: List[Tuple[int, str]] = [] + watched_ids: Set[Tuple[int, int]] = set() + db_dev: Optional[int] = None + for candidate in (db_path_str, db_path_str + "-wal", db_path_str + "-shm"): + try: + stat_result = os.stat(candidate) + except OSError as exc: + if exc.errno not in (errno.ENOENT, errno.ESRCH): + holders.append( + (-1, f"watched-file stat failed: {candidate}: {exc}") + ) + continue + watched_ids.add((stat_result.st_dev, stat_result.st_ino)) + if candidate == db_path_str: + db_dev = stat_result.st_dev + + if sys.platform.startswith("linux"): + try: + own_pid = os.getpid() + for pid_str in os.listdir("/proc"): + if not pid_str.isdigit(): + continue + pid = int(pid_str) + if pid == own_pid: + continue + fd_dir = f"/proc/{pid}/fd" + try: + fds = os.listdir(fd_dir) + except OSError: + argv = _read_proc_argv(pid) + if argv is not None and _looks_like_hermes(argv): + cmdline = " ".join(argv) + holders.append((pid, f"uninspectable holder: {cmdline[:80]}")) + continue + for fd in fds: + fd_path = f"{fd_dir}/{fd}" + try: + target = os.readlink(fd_path) + except OSError as exc: + if exc.errno in (errno.ENOENT, errno.ESRCH): + continue + argv = _read_proc_argv(pid) + if argv is not None and _looks_like_hermes(argv): + holders.append( + ( + pid, + f"uninspectable descriptor: {fd_path}: {exc}", + ) + ) + continue + target_is_watched = canonical_sqlite_path(target) in watched + try: + fd_stat = os.stat(fd_path) + except OSError as exc: + if exc.errno in (errno.ENOENT, errno.ESRCH): + continue + if target_is_watched: + holders.append( + (pid, f"uninspectable descriptor: {target}: {exc}") + ) + else: + argv = _read_proc_argv(pid) + if argv is not None and _looks_like_hermes(argv): + holders.append( + ( + pid, + "uninspectable descriptor: " + f"{target}: {exc}", + ) + ) + continue + if (fd_stat.st_dev, fd_stat.st_ino) in watched_ids or ( + target_is_watched + and target.endswith(" (deleted)") + and db_dev is not None + and fd_stat.st_dev == db_dev + ): + holders.append((pid, target)) + except Exception as exc: + logger.warning( + "Could not prove state.db has no foreign holders; " + "deferring structural maintenance: %s", + exc, + ) + holders.append((-1, f"open-file scan failed: {exc}")) + return holders + + if psutil is None: + return [(-1, "open-file scan unavailable")] + try: + for process in psutil.process_iter(["pid", "open_files"]): + info = process.info + pid = int(info["pid"]) + if pid == os.getpid(): + continue + for opened in info.get("open_files") or (): + path = getattr(opened, "path", "") + if path and canonical_sqlite_path(path) in watched: + holders.append((pid, path)) + except Exception as exc: + logger.warning( + "Could not prove state.db has no foreign holders; " + "deferring structural maintenance: %s", + exc, + ) + holders.append((-1, f"open-file scan failed: {exc}")) + return holders + + +def live_writer_holds_db( + db_path: Path, + *, + connect_repair_durable: Callable[..., sqlite3.Connection], +) -> bool: + """Return whether repair lacks proven exclusive ownership of ``db_path``.""" + foreign_holders = foreign_state_db_holders(db_path) + if any( + pid < 0 + or path.startswith("uninspectable holder:") + or path.startswith("uninspectable descriptor:") + or path.endswith(" (deleted)") + for pid, path in foreign_holders + ): + return True + + probe = None + try: + probe = connect_repair_durable(db_path, timeout=0.0) + probe.execute("PRAGMA locking_mode=EXCLUSIVE") + probe.execute("BEGIN IMMEDIATE") + probe.execute("ROLLBACK") + return False + except sqlite3.OperationalError as exc: + lowered = str(exc).lower() + return "locked" in lowered or "busy" in lowered + except sqlite3.DatabaseError: + return False + except Exception: + return False + finally: + if probe is not None: + try: + probe.execute("PRAGMA locking_mode=NORMAL") + except Exception: + pass + try: + probe.close() + except Exception: + pass diff --git a/hermes_state_portability.py b/hermes_state_portability.py index 1d23308a46..54b4da8980 100644 --- a/hermes_state_portability.py +++ b/hermes_state_portability.py @@ -98,9 +98,10 @@ class SessionPortabilityMixin: s["preview"] = _shape_preview(s.pop("_preview_raw", "")) return s - def _locked_rows(self, sql: str, params=()) -> list: - with self._lock: - return self._conn.execute(sql, params).fetchall() + def _read_rows(self, sql: str, params=()) -> list: + """Pure-read query via ``_read_ctx()`` (never the writer lock: turn persistence must not convoy).""" + with self._read_ctx() as conn: + return conn.execute(sql, params).fetchall() def distinct_session_cwds(self, include_archived: bool = False) -> List[Dict[str, Any]]: """Distinct non-empty session cwds with usage stats, for repo discovery. Aggregates @@ -109,7 +110,7 @@ class SessionPortabilityMixin: where = "cwd IS NOT NULL AND TRIM(cwd) != ''" if not include_archived: where += " AND archived = 0" - rows = self._locked_rows( + rows = self._read_rows( "SELECT cwd AS cwd, COUNT(*) AS sessions, MAX(COALESCE(ended_at, started_at, 0)) AS last_active " f"FROM sessions WHERE {where} GROUP BY cwd" ) @@ -130,7 +131,7 @@ class SessionPortabilityMixin: "\n ORDER BY s.started_at DESC, s.id DESC\n LIMIT ? OFFSET ?", prompt_select=f",\n {_PROMPT_RESOLVED_SQL}", ) - return [self._rich_row(row) for row in self._locked_rows(query, (prefix, prefix_hi, limit, offset))] + return [self._rich_row(row) for row in self._read_rows(query, (prefix, prefix_hi, limit, offset))] def _get_session_rich_row(self, session_id: str, compact_rows: bool = False) -> Optional[Dict[str, Any]]: """One session with the ``list_sessions_rich`` enriched columns, or None. @@ -160,13 +161,13 @@ class SessionPortabilityMixin: self._compact_session_cols() if compact_rows else "s.*", f"s.id IN ({','.join('?' for _ in ids)})", prompt_select=None if compact_rows else f", {_PROMPT_RESOLVED_SQL}", ) - return {s["id"]: s for s in map(self._rich_row, self._locked_rows(query, ids))} + return {s["id"]: s for s in map(self._rich_row, self._read_rows(query, ids))} def list_skill_scaffolded_sessions(self, limit: int = 200) -> List[Dict[str, Any]]: """Titled sessions whose first user turn was a ``/skill`` invocation (their titles describe the expanded skill body, not the request). Returns ``id``, ``title`` and the first-turn ``content`` so callers can re-derive what was typed. Newest first.""" - rows = self._locked_rows(""" + rows = self._read_rows(""" SELECT s.id, s.title, m.content FROM sessions s JOIN messages m ON m.id = ( diff --git a/hermes_state_repair.py b/hermes_state_repair.py index 29357f309c..ac0064c174 100644 --- a/hermes_state_repair.py +++ b/hermes_state_repair.py @@ -51,7 +51,8 @@ _FINGERPRINT_VOLATILE_HEADER_RANGES = ((24, 28), (92, 96)) _REPAIR_BACKUP_MIN_FREE_BYTES = 256 * 1024 * 1024 # 256 MiB absolute floor _REPAIR_BACKUP_FREE_FRACTION = 0.02 # plus 2% of the volume _FTS_TABLES = ("messages_fts", "messages_fts_trigram", "messages_fts_cjk") -_MANUAL_RECOVER_HINT = 'Free disk space, then retry (or recover manually with `sqlite3 {db_path} ".recover"`).' +_MANUAL_RECOVER_HINT = ("Free disk space, then retry (or recover manually with " + "`hermes sessions recover --source {db_path} --inspect-only` first).") def _sidecars(db_path: Path): @@ -372,7 +373,10 @@ def _persistent_repair_exhausted_error(db_path: Path) -> str: """The stable operator-facing diagnostic for an exhausted repair budget.""" return (f"automatic repair has already failed {_MAX_PERSISTENT_REPAIR_ATTEMPTS} times on this exact file — the " f"corruption is beyond the schema/FTS repair strategies (likely b-tree page damage). Manual recovery " - f"required: restore a backup, or salvage with `sqlite3 {db_path} \".recover\"`. " + f"required: restore a backup, or salvage with `hermes sessions recover --source {db_path} " + f"--inspect-only`, then (if it reports recoverable) `hermes sessions recover --source {db_path} " + f"--output recovered-state.db` (recovery snapshots the damaged file first, then runs the page-level " + f"`.recover` lane on the copy; do NOT point a raw `sqlite3` shell at the live database). " f"Delete {_repair_ledger_path(db_path).name} to force another automatic attempt.") @@ -458,7 +462,10 @@ def _backup_db_file(db_path: Path) -> "Tuple[Optional[Path], Optional[str]]": Dedupe: reuse the newest backup when byte-identical to the current recovery image (``_backup_content_identity`` — NOT mtime, NOT ``_db_fingerprint``); a repair loop once re-copied the same bytes on every restart. Staging names live OUTSIDE the ``.malformed-backup-`` prefix: inside it they count as a backup, sort NEWEST (prune kept - partials, deleted intact copies) and dedupe could return one with no real forensic copy on disk.""" + partials, deleted intact copies) and dedupe could return one with no real forensic copy on disk. + + Refusal reasons (``_backup_free_space_error`` / ``_MANUAL_RECOVER_HINT``) point operators at the safe lane, + `hermes sessions recover --source --inspect-only`, never at a raw sqlite3 shell on the live file.""" with contextlib.suppress(ImportError): # scaffold/embed installs without hermes_cli track no connections from hermes_cli.sqlite_safe_read import has_live_connection if has_live_connection(db_path): @@ -715,16 +722,11 @@ def _live_writer_holds_db(db_path: Path) -> bool: it with SQLITE_BUSY; neither statement parses the schema, so it works on malformed DBs. Fails **open** (False) on anything but a positive busy/locked signal: refusing to repair a DB nobody holds would strand the self-heal path. In ``journal_mode=DELETE`` a held reader takes only SHARED and this returns False; - repair is then serialised only by the cross-process repairer lock.""" - try: - probe = _open_exclusive(db_path, "BEGIN IMMEDIATE") - with contextlib.suppress(Exception): # a close() error is not evidence of a holder - _close_unpinned(probe) - return False - except sqlite3.OperationalError as exc: - return "locked" in str(exc).lower() or "busy" in str(exc).lower() - except Exception: # malformed/unreadable: no evidence of a live holder either way - return False + repair is then serialised only by the cross-process repairer lock. Before probing, the foreign-holder scan + (``hermes_state_holders``) fails closed on deleted-WAL-generation, uninspectable, or unknown holders.""" + import hermes_state_holders as _state_holders + from hermes_state import _connect_repair_durable as _connect # call-time lookup: tests patch hermes_state. + return _state_holders.live_writer_holds_db(db_path, connect_repair_durable=_connect) def _repair_skip(report: Dict[str, Any], verb: str, error: str, exc: Optional[BaseException] = None) -> Dict[str, Any]: @@ -789,10 +791,10 @@ def repair_state_db_schema(db_path: Path, *, backup: bool = True) -> Dict[str, A # Probe journal mode BEFORE surgery: a rebuilt file comes back in the default (delete) mode and nothing # else records the flip. Unprobeable (damaged file) -> database.journal_mode is the restore target. before_mode = _probe_journal_mode_for_repair(db_path) - result = _repair_state_db_schema_locked(db_path, backup=backup, report=report) + result = _repair_state_db_schema_locked(db_path, backup=backup, report=report, + journal_mode_before=before_mode) if result.get("repaired"): result["journal_mode_before"] = before_mode - _restore_journal_mode_after_repair(db_path, before_mode) # Environmental aborts (before a strategy mutates the snapshot) are retriable, not proof of exhaustion; # the private marker stays out of the public report. The ledger update stays under the cross-process # lock so two repairers cannot lose each other's updates; a queued loser must not record at all. @@ -813,9 +815,13 @@ def _probe_journal_mode_for_repair(db_path: Path) -> Optional[str]: return None -def _restore_journal_mode_after_repair(db_path: Path, before_mode: Optional[str]) -> None: +def _restore_journal_mode_after_repair(db_path: Path, before_mode: Optional[str], *, conn=None) -> None: """Re-apply the journal mode after schema surgery. + ``conn`` must be the exclusive repair guard connection when called from the repair path: opening a fresh + connection AFTER the guard released let a writer still holding the unlinked old ``-wal`` inode coexist with + a brand-new ``state.db-wal`` — two generations of one store. The reopen is the hazard, not the mode. + A rebuilt file comes back in the default (delete) mode; without this a corruption event silently moves a WAL store out of WAL (the open-time WAL-reset gate never sees a flip made inside repair). Routed through :func:`apply_wal_with_fallback`, not a direct pragma, so it inherits the WAL-reset gate (a vulnerable @@ -824,7 +830,10 @@ def _restore_journal_mode_after_repair(db_path: Path, before_mode: Optional[str] is ``database.journal_mode``. Best-effort: the repair already succeeded, so failures log at WARNING.""" from hermes_state import apply_wal_with_fallback try: - with _repair_conn(db_path) as conn: + if conn is None: + with _repair_conn(db_path) as owned: + after = apply_wal_with_fallback(owned, db_label=db_path.name) + else: after = apply_wal_with_fallback(conn, db_label=db_path.name) if before_mode and after != before_mode: logger.warning("state.db repair changed journal_mode %r -> %r (pre-surgery probe %r; restore resolved " @@ -835,7 +844,9 @@ def _restore_journal_mode_after_repair(db_path: Path, before_mode: Optional[str] "journal_mode on the next open", db_path, exc) -def _repair_state_db_schema_locked(db_path: Path, *, backup: bool, report: Dict[str, Any]) -> Dict[str, Any]: +def _repair_state_db_schema_locked( + db_path: Path, *, backup: bool, report: Dict[str, Any], journal_mode_before: Optional[str] = None, +) -> Dict[str, Any]: """Repair strategies for :func:`repair_state_db_schema`; caller holds the cross-process repair lock. Strategies run on a SCRATCH COPY, copied back through SQLite's transactional backup API only once proven to open @@ -844,8 +855,9 @@ def _repair_state_db_schema_locked(db_path: Path, *, backup: bool, report: Dict[ file from the schema SQLite can still parse — when the damage IS in the schema b-tree (the ``malformed database schema ()`` class) every table hanging off the unreadable part is silently dropped, the probe still reports malformed, and repair returned ``repaired=False`` having destroyed what it was asked to save.""" - from hermes_state import (_backup_db_file, _copy_database_snapshot, _db_opens_cleanly, - _repair_scratch_space_error, _run_repair_strategies, _unlink_db_triple) + from hermes_state import (_backup_db_file, _copy_database_snapshot, _db_opens_cleanly, _exclusive_repair_db_guard, + _repair_scratch_space_error, _restore_journal_mode_after_repair, _run_repair_strategies, + _unlink_db_triple) scratch = db_path.with_name(f"{db_path.name}.repair-scratch") if (cleanup_error := _unlink_db_triple(scratch)) is not None: return _repair_skip(report, "aborted", f"could not remove a stale repair snapshot before probing state.db: {cleanup_error}") @@ -894,6 +906,7 @@ def _repair_state_db_schema_locked(db_path: Path, *, backup: bool, report: Dict[ else: logger.warning("state.db repaired via '%s' and promoted transactionally: %s", report.get("strategy"), db_path) + _restore_journal_mode_after_repair(db_path, journal_mode_before, conn=live_guard) if not report.get("repaired"): # Logged HERE, not in the strategies: they see the scratch copy, and the message a human acts on # must name a path that still exists. diff --git a/hermes_state_schema.py b/hermes_state_schema.py index d4387cf239..f16862d03d 100644 --- a/hermes_state_schema.py +++ b/hermes_state_schema.py @@ -22,7 +22,8 @@ from hermes_startup_watchdog import report_startup_progress from utils import safe_json_loads from hermes_state_common import ( DEFERRED_INDEX_SQL, FTS_CJK_STALE_KEY, FTS_REBUILD_DEFERRAL_KEY, FTS_STALE_KEY, FTS_SQL, - FTS_STORAGE_VERSION, FTS_TRIGRAM_SQL, LEGACY_FTS_SQL, LEGACY_FTS_TRIGRAM_SQL, SCHEMA_SQL, + FTS_STORAGE_VERSION, FTS_TOOL_FULL_CONTENT_HIGH_WATER_KEY, FTS_TRIGRAM_SQL, LEGACY_FTS_SQL, + LEGACY_FTS_TRIGRAM_SQL, SCHEMA_SQL, SCHEMA_VERSION, _FTS_CJK_TRIGGERS, _FTS_TRIGGERS, _ephemeral_child_sql, fts_rebuild_admission, ) @@ -117,6 +118,9 @@ _TITLE_UNIQUE_INDEX_SQL = ( _STALE_KEY_UPSERT_SQL = ( "INSERT INTO state_meta (key, value) VALUES (?, '1') ON CONFLICT(key) DO UPDATE SET value = excluded.value" ) +_STATE_META_UPSERT_SQL = ( + "INSERT INTO state_meta (key, value) VALUES (?, ?) ON CONFLICT(key) DO UPDATE SET value = excluded.value" +) _CLEAR_REBUILD_MARKERS_SQL = "DELETE FROM state_meta WHERE key IN ('fts_rebuild_high_water', 'fts_rebuild_progress')" @@ -261,6 +265,93 @@ class SessionSchemaMixin: logger.info("Migrated %d broad FTS UPDATE trigger(s) to AFTER UPDATE OF (no rebuild required)", len(to_drop)) return len(to_drop) + @staticmethod + def _stamp_fts_tool_high_water(cursor: sqlite3.Cursor) -> None: + """Record MAX(messages.id) as the bounded-tool-content high-water mark: rows at or below it keep + their exact stored token stream; newer tool rows index only the prefix (see ``_fts_indexed_content_sql``).""" + high_water = cursor.execute("SELECT COALESCE(MAX(id), 0) FROM messages").fetchone()[0] + cursor.execute(_STATE_META_UPSERT_SQL, (FTS_TOOL_FULL_CONTENT_HIGH_WATER_KEY, str(high_water))) + + @staticmethod + def _execute_ddl_script_transactional(cursor: sqlite3.Cursor, ddl: str) -> None: + """Execute a DDL script without ``executescript``'s implicit commit.""" + statement = "" + for line in ddl.splitlines(): + statement += line + "\n" + if sqlite3.complete_statement(statement): + cursor.execute(statement) + statement = "" + if statement.strip(): + raise sqlite3.OperationalError("incomplete FTS DDL statement") + + def _migrate_bounded_tool_fts_triggers(self, cursor: sqlite3.Cursor, *, legacy: bool) -> None: + """Replace FTS triggers without rebuilding historical indexes. Existing rows keep their + full-content token stream; the durable high-water id makes new tool rows use the bounded + prefix in INSERT and the matching external-content delete/update. One savepoint, so no + concurrent writer lands in a trigger gap.""" + marker = cursor.execute( + "SELECT 1 FROM state_meta WHERE key = ? LIMIT 1", (FTS_TOOL_FULL_CONTENT_HIGH_WATER_KEY,), + ).fetchone() + if marker is not None: + return + trigram_present = self._sqlite_table_exists(cursor, "messages_fts_trigram") + names = _FTS_BASE_TRIGGERS + (_FTS_TRIGRAM_TRIGGERS if legacy and trigram_present else ()) + has_messages = cursor.execute("SELECT 1 FROM messages LIMIT 1").fetchone() is not None + self._fts_tool_prefix_migration_requires_rebuild = bool( + self._sqlite_table_exists(cursor, "messages_fts") and has_messages + and self._fts_triggers_missing(cursor, names) + ) + cursor.execute("SAVEPOINT bounded_tool_fts") + try: + self._stamp_fts_tool_high_water(cursor) + for name in names: + cursor.execute(f"DROP TRIGGER IF EXISTS {name}") + if legacy: + self._execute_ddl_script_transactional(cursor, LEGACY_FTS_SQL) + if trigram_present: + self._execute_ddl_script_transactional(cursor, LEGACY_FTS_TRIGRAM_SQL) + else: + self._execute_ddl_script_transactional(cursor, FTS_SQL) + cursor.execute("RELEASE SAVEPOINT bounded_tool_fts") + except BaseException: + cursor.execute("ROLLBACK TO SAVEPOINT bounded_tool_fts") + cursor.execute("RELEASE SAVEPOINT bounded_tool_fts") + raise + + @staticmethod + def _sqlite_table_exists(cursor: sqlite3.Cursor, name: str) -> bool: + return cursor.execute( + "SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = ?", (name,), + ).fetchone() is not None + + def _migrate_trigram_cron_exclusion(self, cursor: sqlite3.Cursor) -> bool: + """Install the source-filtered trigram view and purge historical rows (v29 cron exclusion, + v30 subagent exclusion: both only change the view/trigger predicate and rebuild from it). + Legacy inline indexes stay opt-in (their content is private to the vtable). A v1 external + layout whose vtable still declares ``tool_calls`` is left to ``optimize-storage``: swapping + the view underneath would make 'rebuild' read a column the view no longer has. Otherwise + the inverted index still holds excluded rows until FTS5 rebuilds from the new view, which + runs under the shared cross-process admission gate. Returns False to hold schema_version back.""" + if self._db_has_legacy_inline_fts(cursor) or self._db_has_trigram_tool_calls_projection(cursor): + return True + trigram_exists = self._fts_table_probe(cursor, "messages_fts_trigram") + if trigram_exists is not True: + # Absent: the normal ensure path creates/backfills it. None: this runtime cannot + # safely inspect an existing one, so leave the version behind for retry. + return trigram_exists is False + for name in _FTS_TRIGRAM_TRIGGERS: + cursor.execute(f"DROP TRIGGER IF EXISTS {name}") + cursor.execute("DROP VIEW IF EXISTS messages_fts_trigram_src") + if not self._ensure_fts_schema(cursor, "messages_fts_trigram", FTS_TRIGRAM_SQL): + return False + # Always rebuild while schema_version is behind, even if the view already has the new + # predicate: a process can die between replacing the view and rebuilding/stamping. + self._run_admitted_startup_rebuild( + cursor, + lambda: cursor.execute("INSERT INTO messages_fts_trigram(messages_fts_trigram) VALUES('rebuild')"), + ) + return True + def _quarantine_cjk_after_update_of_migration(self, cursor: sqlite3.Cursor) -> None: """Fail closed after dropping the CJK UPDATE trigger mid-migration: clear availability, persist ``fts_cjk_stale``, drop any residual trigger so a later open cannot @@ -281,6 +372,7 @@ class SessionSchemaMixin: markers are cleared or the worker would re-insert covered rows (duplicates). ``legacy`` (pre-v23 inline layout) has no external-content 'rebuild' source, so it DELETEs + reinserts the concatenated content the legacy triggers produced.""" + SessionSchemaMixin._stamp_fts_tool_high_water(cursor) tables = ("messages_fts", "messages_fts_trigram") if include_trigram else ("messages_fts",) for tbl in tables: if legacy: @@ -401,7 +493,15 @@ class SessionSchemaMixin: start a full rebuild, and a gateway opens state.db once for days, so "next open" never comes. Bounded doubling backoff, non-blocking admission, no new thread. True only when the index was rebuilt and sync triggers restored. Never raises.""" - if not self._fts_stale or self.read_only or self._conn is None: + if not self._fts_stale: + return False + if getattr(self, "_db_corrupt", False): + # Quarantined: never run FTS DDL/DML against a damaged image (mirrors _try_wal_checkpoint / + # close). Reset the backoff so a future un-quarantine starts from the default interval. + self._fts_stale_retry_after = 0.0 + self._fts_stale_retry_interval = 0.0 + return False + if self.read_only or self._conn is None: return False now = time.monotonic() if now < getattr(self, "_fts_stale_retry_after", 0.0): @@ -761,11 +861,17 @@ class SessionSchemaMixin: # install only gets a flag; `hermes sessions optimize-storage` performs it. The FTS # layout is tracked by the independent `fts_storage_version` marker, so # schema_version still advances for legacy-FTS users. - if current_version < 23 and fts5_available and self._db_has_legacy_inline_fts(cursor): + if current_version < 23 and fts5_available and self._db_needs_fts_storage_upgrade(cursor): self.set_meta("fts_optimize_available", "1", cursor=cursor) if current_version < 25: # v25: de-duplicate system prompt snapshots (old column stays a read fallback). self._dedupe_legacy_system_prompts(cursor) + fts_migrations_complete = True + if current_version < 30 and fts5_available: + # v29: cron sessions leave the trigram substring index (they stay in the word index); + # v30: delegate-child transcripts too (FTS_TRIGRAM_EXCLUDED_SOURCES + _delegate_from). + # Rebuild once so rows indexed by older view/trigger definitions do not linger. + fts_migrations_complete = self._migrate_trigram_cron_exclusion(cursor) # Stamp the FTS layout version (fresh/optimized DBs); a legacy DB keeps its absent/0 # marker until optimize-storage runs. An INTERRUPTED optimize (markers, trash, or an @@ -773,7 +879,7 @@ class SessionSchemaMixin: # source of truth for "fully optimized" and keeps the resume offer alive. if ( fts5_available - and not self._db_has_legacy_inline_fts(cursor) + and not self._db_needs_fts_storage_upgrade(cursor) and cursor.execute( "SELECT 1 FROM state_meta WHERE key = 'fts_rebuild_high_water' LIMIT 1" ).fetchone() is None @@ -785,7 +891,7 @@ class SessionSchemaMixin: # Advance schema_version — deliberately NOT gated on the FTS opt-in (that would block # every future migration for a user who never optimizes). FTS5 unavailable is the # one skip: claiming current would lie. - if current_version < SCHEMA_VERSION and fts5_available: + if current_version < SCHEMA_VERSION and fts_migrations_complete and fts5_available: cursor.execute("UPDATE schema_version SET version = ?", (SCHEMA_VERSION,)) def _migrate_v22_session_model_usage(self, cursor: sqlite3.Cursor) -> None: @@ -847,6 +953,8 @@ class SessionSchemaMixin: OPT-IN v23 boundary: a legacy v22 inline install keeps its inline schema + triggers (the v23 DDL would create the trigram source VIEW and leave a mixed state).""" legacy_fts = self._db_has_legacy_inline_fts(cursor) + if not self._fts_stale: + self._migrate_bounded_tool_fts_triggers(cursor, legacy=legacy_fts) if self._fts_stale: if self._recover_stale_fts(cursor, legacy=legacy_fts): # CJK was detached alongside the base indexes; its ensure path decides when it returns. @@ -857,7 +965,8 @@ class SessionSchemaMixin: base_sql, trigram_sql = _FTS_DDL[legacy_fts] # Measure BEFORE the DDL below runs (pre-repair state). Whether the trigram half is # creatable is only known AFTER _ensure_fts_schema, hence the halves combine at the `if`. - base_triggers_missing = self._fts_triggers_missing(cursor, _FTS_BASE_TRIGGERS) + base_triggers_missing = self._fts_triggers_missing(cursor, _FTS_BASE_TRIGGERS) or getattr( + self, "_fts_tool_prefix_migration_requires_rebuild", False) trigram_triggers_missing = self._fts_triggers_missing(cursor, _FTS_TRIGRAM_TRIGGERS) self._fts_enabled = self._ensure_fts_schema(cursor, "messages_fts", base_sql) if self._fts_enabled: diff --git a/hermes_state_search.py b/hermes_state_search.py index 478420f0f9..e6e769205c 100644 --- a/hermes_state_search.py +++ b/hermes_state_search.py @@ -14,9 +14,10 @@ from typing import Any, Callable, Collection, Dict, List, Optional, Tuple from agent.skill_commands import describe_skill_invocation from utils import env_float from hermes_state_common import ( - FTS_CJK_STALE_KEY, FTS_SQL, FTS_STALE_KEY, FTS_STORAGE_VERSION, FTS_TRIGRAM_SQL, + FTS_CJK_STALE_KEY, FTS_SQL, FTS_STALE_KEY, FTS_STORAGE_VERSION, FTS_TOOL_CONTENT_PREFIX_CHARS, + FTS_TOOL_FULL_CONTENT_HIGH_WATER_KEY, FTS_TRIGRAM_EXCLUDED_SOURCES, FTS_TRIGRAM_SQL, MAX_FTS5_QUERY_CHARS, SCHEMA_VERSION, _FTS_CJK_TRIGGERS, - escape_like as _escape_like, fts_rebuild_admission, + escape_like as _escape_like, fts_rebuild_admission, fts_trigram_session_sql, ) # Pre-split logger identity so log filtering/capture is unchanged. @@ -214,43 +215,64 @@ class SessionSearchMixin: return {"pending": True, "total": total, "indexed": progress, "percent": min(100, int(100 * progress / total))} # Re-index rows in an id window the index is missing. docsize has one row - # per indexed doc, so the anti-join is exact. + # per indexed doc, so the anti-join is exact. Params: (lo, hi) — the base sweep + # takes (hw, prefix_chars, lo, hi): tool rows past the high water index only a prefix. _BOUNDARY_SWEEP_SQL = ( "INSERT INTO {table}(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 <= ? {extra}" "AND NOT EXISTS (SELECT 1 FROM {table}_docsize d WHERE d.id = m.id)" ) + _BASE_BOUNDARY_SWEEP_SQL = ( + "INSERT INTO messages_fts(rowid, content, tool_name, tool_calls) " + "SELECT m.id, CASE WHEN m.role = 'tool' AND m.id > ? THEN substr(COALESCE(m.content, ''), 1, ?) " + "ELSE m.content END, m.tool_name, m.tool_calls FROM messages m WHERE m.id > ? AND m.id <= ? " + "AND NOT EXISTS (SELECT 1 FROM messages_fts_docsize d WHERE d.id = m.id)" + ) + # Trigram excludes tool rows and FTS_TRIGRAM_EXCLUDED_SOURCES sessions; no tool_calls column. + _TRIGRAM_BOUNDARY_SWEEP_SQL = ( + "INSERT INTO messages_fts_trigram(rowid, content, tool_name) " + "SELECT m.id, m.content, m.tool_name FROM messages m JOIN sessions s ON s.id = m.session_id " + f"WHERE m.id > ? AND m.id <= ? AND m.role <> 'tool' AND {fts_trigram_session_sql('s')} " + "AND NOT EXISTS (SELECT 1 FROM messages_fts_trigram_docsize d WHERE d.id = m.id)" + ) _CHUNK_INSERT_SQL = ( "INSERT INTO {table}(rowid, content, tool_name, tool_calls) " "SELECT id, content, tool_name, tool_calls FROM messages WHERE id > ? AND id <= ?{extra}" ) + _TRIGRAM_CHUNK_INSERT_SQL = ( + "INSERT INTO messages_fts_trigram(rowid, content, tool_name) " + "SELECT m.id, m.content, m.tool_name FROM messages m JOIN sessions s ON s.id = m.session_id " + f"WHERE m.id > ? AND m.id <= ? AND m.role <> 'tool' AND {fts_trigram_session_sql('s')}" + ) def _fts_rebuild_finish(self) -> None: """Finalize the deferred rebuild: boundary sweep + clear markers. The sweep is cheap insurance against a write that slipped between high_water capture and trigger activation. The trigram half is gated on ``_trigram_available``: without the tokenizer/table an unconditional INSERT raises and aborts the whole rebuild.""" - sweeps = [self._BOUNDARY_SWEEP_SQL.format(table="messages_fts", extra="")] + sweeps = [(self._BASE_BOUNDARY_SWEEP_SQL, True)] if self._trigram_available: - sweeps.append(self._BOUNDARY_SWEEP_SQL.format(table="messages_fts_trigram", extra="AND m.role <> 'tool' ")) + sweeps.append((self._TRIGRAM_BOUNDARY_SWEEP_SQL, False)) self._rebuild_finish("fts_rebuild", sweeps) logger.info("Deferred FTS rebuild complete — all messages indexed.") def _fts_cjk_rebuild_finish(self) -> None: """Boundary sweep + clear the cjk markers; index becomes servable.""" sweep = self._BOUNDARY_SWEEP_SQL.format(table="messages_fts_cjk", extra="AND m.role <> 'tool' ") - self._rebuild_finish("fts_cjk_rebuild", [sweep]) + self._rebuild_finish("fts_cjk_rebuild", [(sweep, False)]) self._fts_cjk_available = True logger.info("CJK FTS index backfill complete — serving CJK search.") - def _rebuild_finish(self, prefix: str, sweep_sqls: List[str]) -> None: - """Sweep a generous window around the high-water boundary, then clear the markers.""" + def _rebuild_finish(self, prefix: str, sweep_sqls: List[Tuple[str, bool]]) -> None: + """Sweep a generous window around the high-water boundary, then clear the markers. + ``(sql, bounded)``: a bounded sweep takes the (hw, prefix_chars) tool-content params first.""" def _do(conn): hw_row = _meta_row(conn, f"{prefix}_high_water") if hw_row is not None: hw = int(hw_row[0]) - for sql in sweep_sqls: - conn.execute(sql, (hw - 1000, hw + 1000)) + for sql, bounded in sweep_sqls: + params = (hw, FTS_TOOL_CONTENT_PREFIX_CHARS) if bounded else () + conn.execute(sql, (*params, hw - 1000, hw + 1000)) _delete_meta(conn, f"{prefix}_high_water", f"{prefix}_progress") self._execute_write(_do) @@ -262,9 +284,9 @@ class SessionSearchMixin: return False inserts = [self._CHUNK_INSERT_SQL.format(table="messages_fts", extra="")] if self._trigram_available: - inserts.append(self._CHUNK_INSERT_SQL.format(table="messages_fts_trigram", extra=" AND role <> 'tool'")) + inserts.append(self._TRIGRAM_CHUNK_INSERT_SQL) return self._rebuild_step("fts_rebuild", inserts, fail_msg="FTS rebuild chunk failed (will retry): %s", - finish=self._fts_rebuild_finish) + finish=self._fts_rebuild_finish, finish_when_empty=True) def fts_cjk_rebuild_step(self) -> bool: """Backfill one chunk of the CJK index. True while work remains.""" @@ -274,8 +296,10 @@ class SessionSearchMixin: return self._rebuild_step("fts_cjk_rebuild", [insert], finish=self._fts_cjk_rebuild_finish, fail_msg="CJK FTS rebuild chunk failed (will retry): %s") - def _rebuild_step(self, prefix: str, insert_sqls: List[str], *, fail_msg: str, finish) -> bool: - """Shared chunk engine for the base and CJK deferred backfills.""" + def _rebuild_step(self, prefix: str, insert_sqls: List[str], *, fail_msg: str, finish, + finish_when_empty: bool = False) -> bool: + """Shared chunk engine for the base and CJK deferred backfills. ``finish_when_empty`` + finalizes a high_water <= 0 marker (empty messages table) instead of leaving it pending.""" high_water_raw = self.get_meta(f"{prefix}_high_water") if high_water_raw is None: return False @@ -306,7 +330,9 @@ class SessionSearchMixin: return True # transient (lock contention) — caller retries if more is False: status = self._rebuild_status(prefix) - if status is not None and status["indexed"] >= status["total"]: + if (finish_when_empty and high_water <= 0) or ( + status is not None and status["indexed"] >= status["total"] + ): finish() return False return bool(more) @@ -316,8 +342,8 @@ class SessionSearchMixin: work remains. INTEGER single-column-key tables drain with a high-water marker so each chunk's scan is bounded (restarting the scan was O(n²)); compound-key tables keep the chunked ``LIMIT`` delete — they are small by construction.""" - with self._lock: - trash = [r[0] for r in self._conn.execute( + with self._read_ctx() as conn: + trash = [r[0] for r in conn.execute( "SELECT name FROM sqlite_master WHERE type = 'table' AND name LIKE ? ESCAPE '\\'", (self._FTS_TRASH_PREFIX.replace("_", "\\_") + "%",), ).fetchall()] @@ -409,16 +435,26 @@ class SessionSearchMixin: without an anti-join, so it needs a known-empty index. A missing docsize table counts as empty.""" if _meta_row(conn, "fts_rebuild_progress") is None: - try: - known_empty = int(conn.execute("SELECT COUNT(*) FROM messages_fts_docsize").fetchone()[0]) == 0 - except sqlite3.OperationalError: - known_empty = True - if not known_empty: - for tbl in ("messages_fts", "messages_fts_trigram"): - with contextlib.suppress(sqlite3.OperationalError): # table absent — already an empty surface - conn.execute(f"INSERT INTO {tbl}({tbl}) VALUES('delete-all')") + if not self._fts_index_known_empty(conn): + self._reset_fts_index_to_empty(conn) self.set_meta("fts_rebuild_progress", "0", cursor=conn) + @staticmethod + def _fts_index_known_empty(conn) -> bool: + """True when the base external-content index holds no rows (a missing table counts as empty).""" + try: + return int(conn.execute("SELECT COUNT(*) FROM messages_fts_docsize").fetchone()[0]) == 0 + except sqlite3.OperationalError: + return True + + @staticmethod + def _reset_fts_index_to_empty(conn) -> None: + """Truncate the v23 external-content tables via FTS5 ``'delete-all'`` (O(1); a plain DELETE is + O(rows) and corrupts the index when indexed rows diverged from ``messages``).""" + for tbl in ("messages_fts", "messages_fts_trigram"): + with contextlib.suppress(sqlite3.OperationalError): # table absent — already an empty surface + conn.execute(f"INSERT INTO {tbl}({tbl}) VALUES('delete-all')") + def _seed_fts_rebuild_markers(self, conn, *, force: bool = False) -> int: """Write ``fts_rebuild_high_water`` / ``fts_rebuild_progress`` for a full backfill; returns the high-water id. Without ``force`` an existing high_water only gets a missing @@ -426,10 +462,12 @@ class SessionSearchMixin: existing_hw = _meta_row(conn, "fts_rebuild_high_water") if existing_hw is not None and not force: self._reseed_missing_progress(conn) + self.set_meta(FTS_TOOL_FULL_CONTENT_HIGH_WATER_KEY, str(int(existing_hw[0])), cursor=conn) return int(existing_hw[0]) hw = conn.execute("SELECT COALESCE(MAX(id), 0) FROM messages").fetchone()[0] self.set_meta("fts_rebuild_high_water", str(hw), cursor=conn) self.set_meta("fts_rebuild_progress", "0", cursor=conn) + self.set_meta(FTS_TOOL_FULL_CONTENT_HIGH_WATER_KEY, str(hw), cursor=conn) return int(hw) def _repair_optimize_bookkeeping(self) -> None: @@ -449,15 +487,15 @@ class SessionSearchMixin: self._execute_write(_do) def fts_optimize_available(self) -> bool: - """True when `optimize_fts_storage()` has work: legacy inline FTS, an interrupted optimize + """True when `optimize_fts_storage()` has work: legacy inline FTS or a v23 trigram still + carrying ``tool_calls`` (``_db_needs_fts_storage_upgrade``), an interrupted optimize (markers/trash), a CJK backfill on this tokenizer-capable host, or an empty external index without markers. False when FTS5 is unavailable.""" if not self._fts_enabled or self.read_only: return False - with self._lock: - conn = self._conn + with self._read_ctx() as conn: return ( - self._db_has_legacy_inline_fts(conn) + self._db_needs_fts_storage_upgrade(conn) or _meta_row(conn, "fts_rebuild_high_water") is not None # interrupted optimize # CJK work is only offerable when THIS process can tokenize. or (self._fts_cjk_loaded and ( @@ -469,7 +507,7 @@ class SessionSearchMixin: ) def _demote_legacy_fts_to_trash(self) -> int: - """Demote the legacy inline FTS vtables and stage their shadow tables for chunked + """Demote upgrade-eligible FTS vtables and stage their shadow tables for chunked teardown; returns MAX(messages.id) as the rebuild high water. O(1) schema surgery — the heavy delete is deferred. Markers land in the same BEGIN IMMEDIATE, BEFORE the empty v23 schema is created (``executescript`` implicitly COMMITs), closing @@ -555,8 +593,9 @@ class SessionSearchMixin: def optimize_fts_storage( self, *, progress_cb: Optional[Callable[[Dict[str, Any]], None]] = None, vacuum: bool = True ) -> Dict[str, Any]: - """Migrate a legacy v22 inline-FTS DB to the v23 external-content schema, foreground and - to completion; re-running resumes. ``progress_cb`` receives {"phase", "percent", + """Repair an older FTS layout into the current v23 shape, foreground and to completion: + legacy-v22 inline -> external-content, or a v23 ``messages_fts_trigram`` that still stores + ``tool_calls``. Re-running resumes. ``progress_cb`` receives {"phase", "percent", "indexed", "total"}. A missing trigram tokenizer is not fatal (CJK falls back to LIKE).""" if not self._fts_enabled: return {"ok": False, "reason": "fts5_unavailable"} @@ -566,11 +605,11 @@ class SessionSearchMixin: # Heal bookkeeping BEFORE deciding whether to demote again. self._repair_optimize_bookkeeping() with self._lock: - legacy = self._db_has_legacy_inline_fts(self._conn) + needs_storage_upgrade = self._db_needs_fts_storage_upgrade(self._conn) pending = self.get_meta("fts_rebuild_high_water") is not None - if legacy and not pending: + if needs_storage_upgrade and not pending: self._demote_legacy_fts_to_trash() - elif pending and not legacy: + elif pending and not needs_storage_upgrade: # Resume mid-demote: the process may have died between the staged demote # commit and schema ensure. self._ensure_v23_fts_tables("failed to re-create v23 messages_fts on optimize-storage resume") @@ -607,10 +646,10 @@ class SessionSearchMixin: # Phase 2: tear down the demoted legacy shadow tables in chunks. _emit("teardown") _drive("teardown", self._fts_teardown_trash_step) - with self._lock: - still_pending = _meta_row(self._conn, "fts_rebuild_high_water") is not None - still_trash = self._has_fts_trash(self._conn) - empty_index = self._fts_external_index_empty_with_messages(self._conn) + with self._read_ctx() as conn: + still_pending = _meta_row(conn, "fts_rebuild_high_water") is not None + still_trash = self._has_fts_trash(conn) + empty_index = self._fts_external_index_empty_with_messages(conn) if still_pending or still_trash or empty_index: reason = "backfill_incomplete" if still_pending or empty_index else "teardown_incomplete" logger.warning("FTS storage optimization did not settle (%s): pending=%s trash=%s empty_index=%s", @@ -688,8 +727,8 @@ class SessionSearchMixin: with a DB pick that includes them.""" active_clause = "" if include_inactive else " AND active = 1" display_clause = " AND (display_kind IS NULL OR display_kind = '')" - with self._lock: - rows = self._conn.execute( + with self._read_ctx() as conn: + rows = conn.execute( "SELECT id, timestamp, content FROM messages WHERE session_id = ? AND role = 'user'" f"{active_clause}{display_clause} " "ORDER BY id DESC LIMIT ?", @@ -977,6 +1016,11 @@ class SessionSearchMixin: return [] filters = dict(include_inactive=include_inactive, source_filter=source_filter, exclude_sources=exclude_sources, role_filter=role_filter) + # New oversized tool results index only a bounded prefix; an explicit tool-role search is the + # opt-in full-body path and scans canonical rows via LIKE. + if role_filter and "tool" in role_filter: + matches = self._search_messages_like_fallback(query, limit=limit, offset=offset, sort=sort, **filters) + return self._finalize_search_matches(matches, result_fields=result_fields) self._refresh_fts_stale_state() if self._fts_stale: matches = self._search_messages_like_fallback(query, limit=limit, offset=offset, sort=sort, **filters) @@ -986,11 +1030,13 @@ class SessionSearchMixin: order_by_sql = _FTS_ORDER_BY.get(sort.strip().lower() if isinstance(sort, str) else None, "ORDER BY rank") route = dict(order_by_sql=order_by_sql, limit=limit, offset=offset, **filters) - # Tool rows are excluded from the trigram/cjk indexes (see FTS_TRIGRAM_SQL). - wants_tool_rows = bool(role_filter) and "tool" in role_filter + # Tool rows and FTS_TRIGRAM_EXCLUDED_SOURCES sessions are excluded from the trigram/cjk + # indexes (see FTS_TRIGRAM_SQL); an explicit filter for them must scan the base table. + wants_unindexed_rows = (bool(role_filter) and "tool" in role_filter) or ( + bool(source_filter) and any(src in FTS_TRIGRAM_EXCLUDED_SOURCES for src in source_filter)) is_cjk = self._contains_cjk(query) if is_cjk: - matches = self._search_cjk(query, wants_tool_rows, route) + matches = self._search_cjk(query, wants_unindexed_rows, route) else: sql, params = self._fts_match_sql("messages_fts", query, **route) try: @@ -1020,7 +1066,7 @@ class SessionSearchMixin: # substring-capable indexes: cjk first (exact ranked match), then trigram (>=3-char # tokens). Gated on a miss so hits keep their ranking ("cat" may then match # "concatenate"). Skipped for role='tool' (both indexes exclude tool rows). - if not matches and not is_cjk and not wants_tool_rows: + if not matches and not is_cjk and not (bool(role_filter) and "tool" in role_filter): fb_query = _quote_fts_tokens(query.strip('"').strip()) if self._fts_cjk_available: matches = self._match_rows("messages_fts_cjk", fb_query, **route) or matches @@ -1028,21 +1074,22 @@ class SessionSearchMixin: matches = self._match_rows("messages_fts_trigram", fb_query, **route) or matches return self._finalize_search_matches(matches, result_fields=result_fields) - def _search_cjk(self, query: str, wants_tool_rows: bool, route: Dict[str, Any]) -> List[Dict[str, Any]]: + def _search_cjk(self, query: str, wants_unindexed_rows: bool, route: Dict[str, Any]) -> List[Dict[str, Any]]: """CJK routing: the unicode61 table splits CJK into single characters (false positives, - missed phrases). cjk-bigram serves every shape except role='tool' queries and LONE + missed phrases). cjk-bigram serves every shape except queries wanting rows the + substring indexes exclude (role='tool', cron/subagent sources) and LONE 1-char CJK runs (bigrams only exist for runs >=2 — LIKE is broader); then trigram (>=3 CJK chars per token); then a LIKE substring scan with one clause per non-operator token so "广西 OR 桂林 OR 漓江" matches each term.""" raw_query = query.strip('"').strip() match_query = _quote_fts_tokens(raw_query) - if self._fts_cjk_available and not wants_tool_rows and not self._has_lone_cjk_run(raw_query): + if self._fts_cjk_available and not wants_unindexed_rows and not self._has_lone_cjk_run(raw_query): matches = self._match_rows( "messages_fts_cjk", match_query, fail_open="CJK-bigram", operational_debug="messages_fts_cjk query failed; falling back to trigram/LIKE", **route) if matches is not None: return matches - if self._trigram_route_ok(raw_query) and not wants_tool_rows: + if self._trigram_route_ok(raw_query) and not wants_unindexed_rows: matches = self._match_rows("messages_fts_trigram", match_query, fail_open="Trigram", **route) if matches is not None: return matches @@ -1136,6 +1183,11 @@ class SessionSearchMixin: "Deferred in-place FTS rebuild: another process holds the rebuild authority for this state.db.") return 0 with self._lock: + high_water = self._conn.execute("SELECT COALESCE(MAX(id), 0) FROM messages").fetchone()[0] + self._conn.execute( + "INSERT INTO state_meta (key, value) VALUES (?, ?) ON CONFLICT(key) DO UPDATE SET value = excluded.value", + (FTS_TOOL_FULL_CONTENT_HIGH_WATER_KEY, str(high_water)), + ) for tbl in self._present_fts_tables(): try: self._conn.execute(f"INSERT INTO {tbl}({tbl}) VALUES('rebuild')") diff --git a/hermes_state_sessions.py b/hermes_state_sessions.py index 174733b887..c5fb21a2ea 100644 --- a/hermes_state_sessions.py +++ b/hermes_state_sessions.py @@ -888,6 +888,191 @@ class SessionSessionsMixin: projected.append(merged) return projected + def list_recent_sessions_bounded( + self, + *, + limit: int = 20, + exclude_sources: List[str] = None, + timeout_seconds: float = 3.0, + candidate_limit: int = None, + lineage_limit: int = None, + ) -> List[Dict[str, Any]]: + """Latency-bounded recent-conversation browse (``session_search()``): preselect a small candidate set + from the indexed durable activity timestamp (fallback ``started_at``), resolve only those across + compression ancestry/chains, then hydrate activity/previews for that bounded set. Lineage traversal + uses ``UNION`` plus a total-row ceiling so a corrupt cycle or a deep/branching lineage cannot defeat + the bound; a lineage that hits the ceiling before a terminal root/tip is omitted, not expanded. + A cooperative SQLite progress deadline interrupts sustained work past ``timeout_seconds`` and raises + ``TimeoutError`` (cheap statements may finish between callbacks). Supports only the agent-tool + browse filters; rich callers keep using :meth:`list_sessions_rich`.""" + limit = max(1, int(limit)) + timeout_seconds = max(0.0, float(timeout_seconds)) + if candidate_limit is None: + candidate_limit = max(128, limit * 8) + candidate_limit = max(limit, min(int(candidate_limit), 2048)) + if lineage_limit is None: + lineage_limit = min(8192, candidate_limit * 8) + lineage_limit = max(candidate_limit, min(int(lineage_limit), 8192)) + + candidate_clauses = [ + "s.archived = 0", + "s.hidden = 0", + f"{_delegate_from_json('s.model_config')} IS NULL", + ] + candidate_params: List[Any] = [] + if exclude_sources: + placeholders = ",".join("?" for _ in exclude_sources) + candidate_clauses.append(f"s.source NOT IN ({placeholders})") + candidate_params.extend(exclude_sources) + candidate_where = " AND ".join(candidate_clauses) + + # A compression continuation is an implementation edge, unlike /new + # reset and /branch children which are independent user-visible + # conversations. The same predicate is used in both directions so a + # candidate tip maps to its logical root and the root maps back to the + # freshest live tip. + compression_parent_edge = f""" + parent.end_reason = 'compression' + AND child.parent_session_id = parent.id + AND json_extract( + COALESCE(child.model_config, '{{}}'), '$._branched_from' + ) IS NULL + AND {_delegate_from_json('child.model_config')} IS NULL + AND COALESCE(child.source, '') != 'tool' + """ + + query = f""" + WITH RECURSIVE + recent_candidates(id) AS ( + SELECT s.id + FROM sessions s + WHERE {candidate_where} + ORDER BY COALESCE(s.last_activity_at, s.started_at) DESC, + s.started_at DESC, s.id DESC + LIMIT ? + ), + ancestors(candidate_id, cur_id) AS ( + SELECT id, id FROM recent_candidates + UNION + SELECT a.candidate_id, parent.id + FROM ancestors a + JOIN sessions child ON child.id = a.cur_id + JOIN sessions parent ON {compression_parent_edge} + LIMIT ? + ), + candidate_roots(root_id) AS ( + SELECT DISTINCT a.cur_id + FROM ancestors a + JOIN sessions child ON child.id = a.cur_id + WHERE NOT EXISTS ( + SELECT 1 + FROM sessions parent + WHERE {compression_parent_edge} + ) + ), + chain(root_id, cur_id) AS ( + SELECT root_id, root_id FROM candidate_roots + UNION + SELECT c.root_id, child.id + FROM chain c + JOIN sessions parent ON parent.id = c.cur_id + JOIN sessions child ON {compression_parent_edge} + LIMIT ? + ), + chain_rows AS ( + SELECT + c.root_id, + c.cur_id, + {_sql_session_last_active_by_id('c.cur_id')} AS activity, + CASE WHEN EXISTS ( + SELECT 1 + FROM sessions parent + JOIN sessions child ON {compression_parent_edge} + WHERE parent.id = c.cur_id + ) THEN 0 ELSE 1 END AS is_tip + FROM chain c + ), + ranked_tips AS ( + SELECT root_id, cur_id, activity, + ROW_NUMBER() OVER ( + PARTITION BY root_id + ORDER BY activity DESC, cur_id DESC + ) AS rank_in_root + FROM chain_rows + WHERE is_tip = 1 + ) + SELECT + tip.id, + tip.source, + tip.model, + tip.title, + s.started_at AS started_at, + tip.ended_at, + tip.end_reason, + tip.message_count, + tip.tool_call_count, + rt.activity AS last_active, + COALESCE( + (SELECT {_PREVIEW_RAW_SELECT} + FROM messages m + WHERE m.session_id = tip.id + AND m.role = 'user' + AND m.content IS NOT NULL + AND {_PREVIEW_ELIGIBLE_SQL} + ORDER BY m.timestamp, m.id LIMIT 1), + '' + ) AS _preview_raw, + CASE WHEN s.id != tip.id THEN s.id ELSE NULL END + AS _lineage_root_id + FROM ranked_tips rt + JOIN sessions s ON s.id = rt.root_id + JOIN sessions tip ON tip.id = rt.cur_id + WHERE rt.rank_in_root = 1 + AND s.archived = 0 + AND s.hidden = 0 + AND {_LISTABLE_CHILD_SQL} + AND {_delegate_from_json('s.model_config')} IS NULL + ORDER BY rt.activity DESC, s.started_at DESC, tip.id DESC + LIMIT ? + """ + params = candidate_params + [ + candidate_limit, + lineage_limit, + lineage_limit, + limit, + ] + deadline = time.monotonic() + timeout_seconds + interrupted_by_deadline = False + + def _deadline_progress_handler() -> int: + nonlocal interrupted_by_deadline + if time.monotonic() >= deadline: + interrupted_by_deadline = True + return 1 + return 0 + + try: + with self._read_ctx() as conn: + conn.set_progress_handler(_deadline_progress_handler, 1000) + try: + rows = conn.execute(query, params).fetchall() + finally: + conn.set_progress_handler(None, 0) + except sqlite3.OperationalError as exc: + if interrupted_by_deadline and "interrupt" in str(exc).lower(): + raise TimeoutError( + f"recent-session browse exceeded {timeout_seconds:g}s deadline" + ) from exc + raise + + sessions = [] + for row in rows: + session = self._session_row_dict(row) + session["preview"] = _shape_preview(session.pop("_preview_raw", "")) + session["unread"] = self.session_unread(session) + sessions.append(session) + return sessions + @classmethod def _list_row(cls, row: sqlite3.Row) -> Dict[str, Any]: """Project a list_sessions_rich row: shape the preview, drop internal ordering columns.""" diff --git a/plugins/model-providers/meta-ai/__init__.py b/plugins/model-providers/meta-ai/__init__.py index fcef99fc25..439a9fa4bc 100644 --- a/plugins/model-providers/meta-ai/__init__.py +++ b/plugins/model-providers/meta-ai/__init__.py @@ -17,6 +17,27 @@ from providers.base import ProviderProfile class MetaAIProfile(ProviderProfile): """Meta Model API — top-level reasoning_effort, self-contained.""" + # Non-chat model prefixes excluded from the agent picker. The live + # /v1/models catalog includes image-generation and transcription models + # that are not suitable for agentic chat. + _NON_CHAT_PREFIXES = ("muse-image-", "muse-voice-") + + def fetch_models( + self, + *, + api_key: str | None = None, + base_url: str | None = None, + timeout: float = 8.0, + ) -> list[str] | None: + """Fetch and filter the live catalog, excluding non-chat models.""" + live = super().fetch_models(api_key=api_key, base_url=base_url, timeout=timeout) + if live is None: + return None + return [ + m for m in live + if not any(m.startswith(p) for p in self._NON_CHAT_PREFIXES) + ] + def build_api_kwargs_extras( self, *, reasoning_config: dict | None = None, supports_reasoning: bool = False, **context: Any ) -> tuple[dict[str, Any], dict[str, Any]]: @@ -43,10 +64,14 @@ meta_ai = MetaAIProfile( # Responses API engages Muse prompt caching (0 cached tokens on chat/completions vs # 93-99% hits on /v1/responses); the hook above still covers custom non-api.meta.ai base URLs. api_mode="codex_responses", - supports_vision=True, default_aux_model="muse-spark-1.2-contributor", + # Natively multimodal, but only on user turns: an image envelope inside a role:tool + # message 400s "content did not match any supported type". + supports_vision=True, supports_vision_tool_messages=False, + default_aux_model="muse-spark-1.2-contributor", # Muse spends completion budget on hidden reasoning first; low caps can finish with empty content. default_max_tokens=16384, - fallback_models=("muse-spark-1.2", "muse-spark-1.2-contributor"), + # Single safety-net entry, shown only when the live /v1/models fetch fails. + fallback_models=("muse-spark-1.2",), ) register_provider(meta_ai) diff --git a/plugins/platforms/buzz/adapter.py b/plugins/platforms/buzz/adapter.py index 8401d3c770..0f151330ff 100644 --- a/plugins/platforms/buzz/adapter.py +++ b/plugins/platforms/buzz/adapter.py @@ -74,7 +74,7 @@ def _scoped_platform_setting(env_name, extra, key): logger = logging.getLogger(__name__) from gateway.platforms.base import ( - BasePlatformAdapter, CachedMedia, SendResult, MessageEvent, MessageType, cache_media_bytes, + BasePlatformAdapter, CachedMedia, SendResult, MessageEvent, MessageType, cache_media_bytes_async, ) from gateway.config import Platform @@ -1396,7 +1396,7 @@ class BuzzAdapter(BasePlatformAdapter): logger.warning("Buzz: attachment %s does not match imeta", what) return None try: - return cache_media_bytes(bytes(data), filename=metadata["filename"], mime_type=metadata["mime_type"]) + return await cache_media_bytes_async(bytes(data), filename=metadata["filename"], mime_type=metadata["mime_type"]) except (OSError, ValueError) as exc: logger.warning("Buzz: attachment cache write failed: %s", exc) return None @@ -1637,7 +1637,7 @@ class BuzzAdapter(BasePlatformAdapter): media_urls: List[str] = [] media_types: List[str] = [] media_kinds: List[str] = [] - from gateway.platforms.base import cache_media_bytes, validate_inbound_media_size + from gateway.platforms.base import cache_media_bytes_async, validate_inbound_media_size for url in urls: path_match = _MEDIA_PATH_RE.fullmatch(urlsplit(url).path) if path_match is None: @@ -1652,7 +1652,9 @@ class BuzzAdapter(BasePlatformAdapter): continue validate_inbound_media_size(download_path.stat().st_size, media_type="Buzz media") mime_type = mimetypes.guess_type(download_path.name)[0] or "application/octet-stream" - cached = cache_media_bytes(download_path.read_bytes(), filename=download_path.name, mime_type=mime_type) + # Up to the inbound media cap (128 MiB) — read off the loop too. + data = await asyncio.to_thread(download_path.read_bytes) + cached = await cache_media_bytes_async(data, filename=download_path.name, mime_type=mime_type) except Exception as exc: logger.warning("Buzz: failed to localize inbound media %s (%s)", label, type(exc).__name__) continue diff --git a/plugins/platforms/discord/adapter.py b/plugins/platforms/discord/adapter.py index b41b2e9897..4a98793c5a 100644 --- a/plugins/platforms/discord/adapter.py +++ b/plugins/platforms/discord/adapter.py @@ -256,8 +256,8 @@ from gateway.platforms.helpers import ( from utils import atomic_json_write, env_float from gateway.platforms.base import ( BasePlatformAdapter, MessageEvent, MessageType, ProcessingOutcome, SendResult, - cache_image_from_url, cache_image_from_bytes, cache_audio_from_url, cache_audio_from_bytes, - cache_document_from_bytes, SUPPORTED_DOCUMENT_TYPES, _TEXT_INJECT_EXTENSIONS, + cache_image_from_url, cache_image_from_bytes_async, cache_audio_from_url, cache_audio_from_bytes_async, + cache_document_from_bytes_async, SUPPORTED_DOCUMENT_TYPES, _TEXT_INJECT_EXTENSIONS, _prefix_within_utf16_limit, utf16_len, validate_inbound_media_size, ) from tools.url_safety import is_safe_url @@ -429,6 +429,10 @@ class _DiscordNonConversationalMessageTracker: def __init__(self, max_tracked: int = _MAX_TRACKED): self._max_tracked = max_tracked self._ids: dict[str, None] = dict.fromkeys(self._load()) + # Serializes the offloaded flushes so two concurrent mark_many() calls + # cannot land their writes out of order (last-writer-wins would drop + # the newer ids from disk). + self._persist_lock = asyncio.Lock() def _state_path(self) -> _Path: from hermes_constants import get_hermes_home @@ -450,17 +454,21 @@ class _DiscordNonConversationalMessageTracker: logger.debug("[%s] Failed to load non-conversational Discord IDs", "Discord") return [] - def _save(self) -> None: + def _snapshot(self) -> list[str]: + """Trim in-memory state and return the ids to persist (loop-side).""" ids = list(self._ids) if len(ids) > self._max_tracked: ids = ids[-self._max_tracked:] self._ids = dict.fromkeys(ids) + return ids + + def _save(self, ids: list[str]) -> None: try: atomic_json_write(self._state_path(), ids, indent=None) except Exception: logger.debug("[%s] Failed to save non-conversational Discord IDs", "Discord", exc_info=True) - def mark_many(self, message_ids: List[str]) -> None: + async def mark_many(self, message_ids: List[str]) -> None: changed = False for message_id in message_ids: key = str(message_id or "").strip() @@ -468,7 +476,16 @@ class _DiscordNonConversationalMessageTracker: self._ids[key] = None changed = True if changed: - self._save() + # atomic_json_write() calls os.fsync(), which blocks until the + # write reaches stable storage. Both callers of mark_many() run + # on the event loop, so offload the flush the same way #83906 + # did for the other gateway persist paths. The snapshot (and the + # trim that reassigns ``_ids``) stays on the loop so the worker + # never touches the dict while another task mutates it; the lock + # keeps flushes in mutation order. + async with self._persist_lock: + ids = self._snapshot() + await asyncio.to_thread(self._save, ids) def __contains__(self, message_id: str) -> bool: return str(message_id or "") in self._ids @@ -2816,7 +2833,7 @@ class DiscordAdapter(BasePlatformAdapter): if message_ids: _target_id = thread_id or chat_id if nonconversational: - self._nonconversational_messages.mark_many(message_ids) + await self._nonconversational_messages.mark_many(message_ids) elif not _looks_like_nonconversational_history_message(content): self._last_self_message_id[_target_id] = message_ids[-1] result = SendResult( @@ -5536,7 +5553,7 @@ class DiscordAdapter(BasePlatformAdapter): return {"content": content, "embed": embed, "view": view}, view result = await self._send_prompt(chat_id, metadata, _build) if result.success and _metadata_marks_nonconversational(metadata): - self._nonconversational_messages.mark_many([result.message_id]) + await self._nonconversational_messages.mark_many([result.message_id]) return result async def send_model_picker( @@ -5666,7 +5683,7 @@ class DiscordAdapter(BasePlatformAdapter): raw_bytes = await self._read_attachment_bytes(att, media_type="image") if raw_bytes is not None: try: - return cache_image_from_bytes(raw_bytes, ext=ext) + return await cache_image_from_bytes_async(raw_bytes, ext=ext) except Exception as e: logger.debug( "[Discord] cache_image_from_bytes rejected att.read() data; falling back to URL: %s", @@ -5679,7 +5696,7 @@ class DiscordAdapter(BasePlatformAdapter): raw_bytes = await self._read_attachment_bytes(att, media_type="audio") if raw_bytes is not None: try: - return cache_audio_from_bytes(raw_bytes, ext=ext) + return await cache_audio_from_bytes_async(raw_bytes, ext=ext) except Exception as e: logger.debug("[Discord] cache_audio_from_bytes failed; falling back to URL: %s", e) return await cache_audio_from_url(att.url, ext=ext) @@ -5752,7 +5769,7 @@ class DiscordAdapter(BasePlatformAdapter): continue try: raw_bytes = await self._cache_discord_document(att, ext) - cached_path = cache_document_from_bytes(raw_bytes, att.filename or f"document{ext or '.bin'}") + cached_path = await cache_document_from_bytes_async(raw_bytes, att.filename or f"document{ext or '.bin'}") if in_allowlist: doc_mime = SUPPORTED_DOCUMENT_TYPES[ext] else: diff --git a/plugins/platforms/feishu/adapter.py b/plugins/platforms/feishu/adapter.py index 5cd53d594e..6160996418 100644 --- a/plugins/platforms/feishu/adapter.py +++ b/plugins/platforms/feishu/adapter.py @@ -83,8 +83,8 @@ FEISHU_WEBHOOK_AVAILABLE = aiohttp is not None from gateway.config import Platform, PlatformConfig from gateway.platforms.base import ( BasePlatformAdapter, MessageEvent, MessageType, ProcessingOutcome, SendResult, - SUPPORTED_DOCUMENT_TYPES, cache_document_from_bytes, cache_image_from_url, - cache_audio_from_bytes, cache_image_from_bytes, + SUPPORTED_DOCUMENT_TYPES, cache_document_from_bytes_async, cache_image_from_url, + cache_audio_from_bytes_async, cache_image_from_bytes_async, ) from gateway.status import acquire_scoped_lock, release_scoped_lock from hermes_constants import get_hermes_home @@ -1199,6 +1199,9 @@ class FeishuAdapter(BasePlatformAdapter): self._seen_message_order: List[str] = [] self._dedup_state_path = get_hermes_home() / "feishu_seen_message_ids.json" self._dedup_lock = threading.Lock() + # Serializes the offloaded dedup-state flushes so two concurrent + # inbound messages cannot land their writes out of order. + self._dedup_persist_lock = asyncio.Lock() self._sender_name_cache: Dict[str, tuple[str, float]] = {} # sender_id → (name, expire_at) self._webhook_rate_counts: Dict[str, tuple[int, float]] = {} # rate_key → (count, window_start) self._webhook_anomaly_counts: Dict[str, tuple[int, str, float]] = {} # ip → (count, last_status, first_seen) @@ -1949,7 +1952,7 @@ class FeishuAdapter(BasePlatformAdapter): logger.debug("[Feishu] Dropping malformed inbound event: missing message/sender") return message_id = getattr(message, "message_id", None) - if not message_id or self._is_duplicate(message_id): + if not message_id or await self._is_duplicate(message_id): logger.debug("[Feishu] Dropping duplicate/missing message_id: %s", message_id) return reason = self._admit(sender, message) @@ -2595,7 +2598,7 @@ class FeishuAdapter(BasePlatformAdapter): filename = self._derive_remote_filename( file_url, content_type=content_type_hdr, default_name=preferred_name, default_ext=default_ext, ) - return cache_document_from_bytes(body, filename), filename + return await cache_document_from_bytes_async(body, filename), filename @staticmethod def _guess_remote_extension(url: str, *, default: str) -> str: @@ -2943,7 +2946,7 @@ class FeishuAdapter(BasePlatformAdapter): content_type = self._get_response_header(response, "Content-Type") filename = getattr(response, "file_name", None) or f"{image_key}.jpg" ext = self._guess_extension(filename, content_type, ".jpg", allowed=_IMAGE_EXTENSIONS) - cached_path = cache_image_from_bytes(raw_bytes, ext=ext) + cached_path = await cache_image_from_bytes_async(raw_bytes, ext=ext) return cached_path, self._normalize_media_type(content_type, default=self._default_image_media_type(ext)) except Exception: logger.warning("[Feishu] Failed to cache image resource %s", image_key, exc_info=True) @@ -2979,20 +2982,20 @@ class FeishuAdapter(BasePlatformAdapter): if media_type.startswith("image/"): ext = self._guess_extension(filename, content_type, ".jpg", allowed=_IMAGE_EXTENSIONS) - kind, cached_path = "image", cache_image_from_bytes(raw_bytes, ext=ext) + kind, cached_path = "image", await cache_image_from_bytes_async(raw_bytes, ext=ext) media_type = media_type or self._default_image_media_type(ext) elif request_type == "audio" or media_type.startswith("audio/"): ext = self._guess_extension(filename, content_type, ".ogg", allowed=_AUDIO_EXTENSIONS) - kind, cached_path = "audio", cache_audio_from_bytes(raw_bytes, ext=ext) + kind, cached_path = "audio", await cache_audio_from_bytes_async(raw_bytes, ext=ext) media_type = media_type or f"audio/{ext.lstrip('.') or 'ogg'}" elif media_type.startswith("video/"): if not Path(filename).suffix: filename = f"{filename}.mp4" - kind, cached_path = "video", cache_document_from_bytes(raw_bytes, filename) + kind, cached_path = "video", await cache_document_from_bytes_async(raw_bytes, filename) else: if not Path(filename).suffix and media_type in _DOCUMENT_MIME_TO_EXT: filename = f"{filename}{_DOCUMENT_MIME_TO_EXT[media_type]}" - kind, cached_path = "document", cache_document_from_bytes(raw_bytes, filename) + kind, cached_path = "document", await cache_document_from_bytes_async(raw_bytes, filename) media_type = media_type or self._guess_document_media_type(filename) logger.info("[Feishu] Cached message %s resource at %s", kind, cached_path) return cached_path, media_type @@ -3387,14 +3390,15 @@ class FeishuAdapter(BasePlatformAdapter): def _persist_seen_message_ids(self) -> None: try: self._dedup_state_path.parent.mkdir(parents=True, exist_ok=True) - recent = self._seen_message_order[-self._dedup_cache_size:] - # Save as {msg_id: timestamp} so TTL filtering works across restarts. - payload = {"message_ids": {k: self._seen_message_ids[k] for k in recent if k in self._seen_message_ids}} + with self._dedup_lock: + recent = self._seen_message_order[-self._dedup_cache_size:] + # Save as {msg_id: timestamp} so TTL filtering works across restarts. + payload = {"message_ids": {k: self._seen_message_ids[k] for k in recent if k in self._seen_message_ids}} atomic_json_write(self._dedup_state_path, payload, indent=None) except OSError: logger.warning("[Feishu] Failed to persist dedup state to %s", self._dedup_state_path, exc_info=True) - def _is_duplicate(self, message_id: str) -> bool: + async def _is_duplicate(self, message_id: str) -> bool: now, ttl = time.time(), _FEISHU_DEDUP_TTL_SECONDS with self._dedup_lock: seen_at = self._seen_message_ids.get(message_id) @@ -3404,8 +3408,20 @@ class FeishuAdapter(BasePlatformAdapter): self._seen_message_order.append(message_id) while len(self._seen_message_order) > self._dedup_cache_size: self._seen_message_ids.pop(self._seen_message_order.pop(0), None) - self._persist_seen_message_ids() - return False + # atomic_json_write() fsyncs; this runs on the event loop for every inbound message, so + # offload the flush. The lock keeps flushes in mutation order (the snapshot inside the + # worker is taken under _dedup_lock, but the write itself is not). + async with self._dedup_persist_lock_or_create(): + await asyncio.to_thread(self._persist_seen_message_ids) + return False + + def _dedup_persist_lock_or_create(self) -> asyncio.Lock: + # Tests build bare adapters via object.__new__ and install dedup state + # by hand; create the lock lazily so those fixtures keep working. + lock = getattr(self, "_dedup_persist_lock", None) + if lock is None: + lock = self._dedup_persist_lock = asyncio.Lock() + return lock # --- Outbound payload construction and send pipeline --- def _build_outbound_payload(self, content: str, *, prefer_post: bool = False) -> tuple[str, str]: diff --git a/plugins/platforms/feishu/feishu_meeting_invite.py b/plugins/platforms/feishu/feishu_meeting_invite.py index 6a55280db7..58c0b2ef1e 100644 --- a/plugins/platforms/feishu/feishu_meeting_invite.py +++ b/plugins/platforms/feishu/feishu_meeting_invite.py @@ -126,7 +126,7 @@ async def handle_meeting_invited_event(adapter: Any, data: Any) -> None: return logger.warning("[Feishu-MeetingInvite] Dropping malformed meeting invite event") dedup_key = _dedup_key(payload) is_duplicate = getattr(adapter, "_is_duplicate", None) - if callable(is_duplicate) and is_duplicate(dedup_key): + if callable(is_duplicate) and await is_duplicate(dedup_key): return logger.debug("[Feishu-MeetingInvite] Dropping duplicate event: %s", dedup_key) inviter = payload.inviter if inviter is None or not inviter.open_id: diff --git a/plugins/platforms/google_chat/adapter.py b/plugins/platforms/google_chat/adapter.py index 60975ac293..4008c0b97b 100644 --- a/plugins/platforms/google_chat/adapter.py +++ b/plugins/platforms/google_chat/adapter.py @@ -125,7 +125,8 @@ Platform("google_chat") from gateway.platforms.helpers import MessageDeduplicator from gateway.platforms.base import ( gateway_trust_env, BasePlatformAdapter, MessageEvent, MessageType, ProcessingOutcome, SendResult, - cache_audio_from_bytes, cache_document_from_bytes, cache_image_from_bytes, cache_video_from_bytes, + cache_audio_from_bytes_async, cache_document_from_bytes_async, cache_image_from_bytes_async, + cache_video_from_bytes_async, ) # Pinned to the legacy module path so operator log filters keep matching. @@ -158,8 +159,8 @@ _REDACTIONS = ( ) _MIME_MESSAGE_TYPES = (("image/", MessageType.PHOTO), ("audio/", MessageType.AUDIO), ("video/", MessageType.VIDEO)) _MEDIA_CACHERS = ( - ("image/", cache_image_from_bytes, ".jpg"), ("audio/", cache_audio_from_bytes, ".ogg"), - ("video/", cache_video_from_bytes, ".mp4"), + ("image/", cache_image_from_bytes_async, ".jpg"), ("audio/", cache_audio_from_bytes_async, ".ogg"), + ("video/", cache_video_from_bytes_async, ".mp4"), ) @@ -975,8 +976,8 @@ class GoogleChatAdapter(BasePlatformAdapter): ext = "." + filename.rsplit(".", 1)[-1].lower() if "." in filename else "" for prefix, cache_fn, default_ext in _MEDIA_CACHERS: if mime.startswith(prefix): - return cache_fn(data, ext=ext or default_ext), mime - return cache_document_from_bytes(data, filename), mime + return await cache_fn(data, ext=ext or default_ext), mime + return await cache_document_from_bytes_async(data, filename), mime # -- outbound ------------------------------------------------------------ def _note_rate_limit(self, chat_id: str) -> int: diff --git a/plugins/platforms/line/adapter.py b/plugins/platforms/line/adapter.py index b9265dd65f..fdcbcf1aa7 100644 --- a/plugins/platforms/line/adapter.py +++ b/plugins/platforms/line/adapter.py @@ -36,7 +36,8 @@ from urllib.parse import quote as _urlquote from gateway.platforms._shared import get_scoped_secret as _get_scoped_secret from gateway.platforms.base import ( gateway_trust_env, BasePlatformAdapter, MessageEvent, MessageType, SendResult, - cache_audio_from_bytes, cache_document_from_bytes, cache_image_from_bytes, cache_video_from_bytes) + cache_audio_from_bytes_async, cache_document_from_bytes_async, cache_image_from_bytes_async, + cache_video_from_bytes_async) from gateway.config import Platform logger = logging.getLogger(__name__) @@ -340,7 +341,7 @@ _OUTBOUND_MEDIA = { # Inbound media kinds → cached file extension. _INBOUND_MEDIA_EXT = {"image": ".jpg", "audio": ".m4a", "video": ".mp4", "file": ".bin"} -_INBOUND_AV_CACHERS = {"audio": cache_audio_from_bytes, "video": cache_video_from_bytes} +_INBOUND_AV_CACHERS = {"audio": cache_audio_from_bytes_async, "video": cache_video_from_bytes_async} _LIFECYCLE_EVENTS = frozenset({"follow", "unfollow", "join", "leave"}) _ENV_SEED_KEYS = (("LINE_HOST", "host"), ("LINE_PUBLIC_URL", "public_url"), ("LINE_HOME_CHANNEL", "home_channel")) @@ -596,12 +597,12 @@ class LineAdapter(BasePlatformAdapter): ext = _INBOUND_MEDIA_EXT.get(msg_type, ".bin") try: if msg_type == "image": - return cache_image_from_bytes(data, ext=ext), "image/jpeg" + return await cache_image_from_bytes_async(data, ext=ext), "image/jpeg" if msg_type in _INBOUND_AV_CACHERS: - return _INBOUND_AV_CACHERS[msg_type](data, ext=ext), mimetypes.guess_type(f"{msg_type}{ext}")[0] or f"{msg_type}/mp4" + return await _INBOUND_AV_CACHERS[msg_type](data, ext=ext), mimetypes.guess_type(f"{msg_type}{ext}")[0] or f"{msg_type}/mp4" document_name = filename or f"line_file{ext}" mime = mimetypes.guess_type(document_name)[0] or "application/octet-stream" - return cache_document_from_bytes(data, document_name), mime + return await cache_document_from_bytes_async(data, document_name), mime except Exception as exc: logger.warning("LINE: failed to cache %s payload: %s", msg_type, exc) return None, "" diff --git a/plugins/platforms/matrix/adapter.py b/plugins/platforms/matrix/adapter.py index 64e30a131e..41bdd26374 100644 --- a/plugins/platforms/matrix/adapter.py +++ b/plugins/platforms/matrix/adapter.py @@ -2056,17 +2056,21 @@ class MatrixAdapter(BasePlatformAdapter): logger.warning("[Matrix] Encrypted media event missing decryption metadata for %s", event_id) return None file_bytes = decrypt_attachment(file_bytes, key_value, hash_value, iv_value) - from gateway.platforms.base import cache_audio_from_bytes, cache_document_from_bytes, cache_image_from_bytes + from gateway.platforms.base import ( + cache_audio_from_bytes_async, + cache_document_from_bytes_async, + cache_image_from_bytes_async, + ) if msg_type == MessageType.PHOTO: ext_map = {"image/jpeg": ".jpg", "image/png": ".png", "image/gif": ".gif", "image/webp": ".webp"} - cached_path = cache_image_from_bytes(file_bytes, ext=ext_map.get(media_type, ".jpg")) + cached_path = await cache_image_from_bytes_async(file_bytes, ext=ext_map.get(media_type, ".jpg")) logger.info("[Matrix] Cached user image at %s", cached_path) return cached_path if msg_type in {MessageType.AUDIO, MessageType.VOICE}: ext = Path(body or ("voice.ogg" if is_voice_message else "audio.ogg")).suffix or ".ogg" - return cache_audio_from_bytes(file_bytes, ext=ext) + return await cache_audio_from_bytes_async(file_bytes, ext=ext) filename = body or ("video.mp4" if msg_type == MessageType.VIDEO else "document") - return cache_document_from_bytes(file_bytes, filename) + return await cache_document_from_bytes_async(file_bytes, filename) async def _on_invite(self, event: Any) -> None: """Auto-join rooms when invited, recording DM rooms in m.direct.""" diff --git a/plugins/platforms/mattermost/adapter.py b/plugins/platforms/mattermost/adapter.py index 12005e4b2f..53c6d17aba 100644 --- a/plugins/platforms/mattermost/adapter.py +++ b/plugins/platforms/mattermost/adapter.py @@ -512,9 +512,13 @@ class MattermostAdapter(BasePlatformAdapter): async def _download_attachments(self, file_ids: List[str]) -> Tuple[List[str], List[str]]: """Download attachments now (URLs need auth headers downstream tools lack) → (paths, mime types).""" import aiohttp - from gateway.platforms.base import cache_audio_from_bytes, cache_document_from_bytes, cache_image_from_bytes + from gateway.platforms.base import ( + cache_audio_from_bytes_async, + cache_document_from_bytes_async, + cache_image_from_bytes_async, + ) media_urls, media_types = [], [] - cache_fns = {"image/": cache_image_from_bytes, "audio/": cache_audio_from_bytes} + cache_fns = {"image/": cache_image_from_bytes_async, "audio/": cache_audio_from_bytes_async} for fid in file_ids: try: file_info = await self._api_get(f"files/{fid}/info") @@ -529,9 +533,10 @@ class MattermostAdapter(BasePlatformAdapter): file_data = await resp.read() prefix = next((p for p in cache_fns if mime.startswith(p)), None) if prefix: - media_urls.append(cache_fns[prefix](file_data, Path(fname).suffix or _INBOUND_CACHE_EXT[prefix])) + media_urls.append( + await cache_fns[prefix](file_data, Path(fname).suffix or _INBOUND_CACHE_EXT[prefix])) else: - media_urls.append(cache_document_from_bytes(file_data, fname)) + media_urls.append(await cache_document_from_bytes_async(file_data, fname)) media_types.append(mime) except Exception as exc: logger.warning("Mattermost: error downloading file %s: %s", fid, exc) diff --git a/plugins/platforms/photon/adapter.py b/plugins/platforms/photon/adapter.py index baa5f0d65d..fcbec01b13 100644 --- a/plugins/platforms/photon/adapter.py +++ b/plugins/platforms/photon/adapter.py @@ -419,6 +419,7 @@ _CONTENT_NORMALIZERS: Dict[Any, Callable[[Dict[str, Any]], _Normalized]] = { "richlink": lambda c: (_format_richlink_content(c), MessageType.TEXT, [], []), "group": _normalize_group_content, } +_BINARY_CONTENT_TYPES = {"attachment", "voice", "group"} # may decode/cache media bytes → run off the event loop def _normalize_content(content: Dict[str, Any]) -> _Normalized: @@ -805,7 +806,11 @@ class PhotonAdapter(BasePlatformAdapter): return await self.handle_message(_event(choice)) return - text, mtype, media_urls, media_types = _normalize_content(content) + if ctype in _BINARY_CONTENT_TYPES: + # Base64 decode + media-cache write of possibly multi-MB payloads — keep it off the event loop. + text, mtype, media_urls, media_types = await asyncio.to_thread(_normalize_content, content) + else: + text, mtype, media_urls, media_types = _normalize_content(content) if chat_type == "group" and self.require_mention: if not self._message_matches_mention_patterns(text): logger.debug("[photon] ignoring group message (require_mention=true, no mention pattern matched)") diff --git a/plugins/platforms/slack/adapter.py b/plugins/platforms/slack/adapter.py index 8bf4630f62..8457a2a1a8 100644 --- a/plugins/platforms/slack/adapter.py +++ b/plugins/platforms/slack/adapter.py @@ -39,7 +39,7 @@ from gateway.platforms.base import ( gateway_trust_env, BasePlatformAdapter, MessageEvent, MessageType, ProcessingOutcome, SendResult, SUPPORTED_DOCUMENT_TYPES, SUPPORTED_VIDEO_TYPES, _TEXT_INJECT_EXTENSIONS, is_host_excluded_by_no_proxy, resolve_proxy_url, safe_url_for_log, _ssrf_redirect_guard, - cache_document_from_bytes, cache_video_from_bytes) + cache_document_from_bytes_async, cache_video_from_bytes_async) try: # sibling module; support both package and flat plugin-dir import from .block_kit import render_blocks, sanitize_blocks @@ -4222,7 +4222,7 @@ class SlackAdapter(BasePlatformAdapter): mime_to_ext = {v: k for k, v in SUPPORTED_VIDEO_TYPES.items()} ext = mime_to_ext.get(mimetype.split(";", 1)[0].lower(), ".mp4") raw_bytes = await self._download_slack_file_bytes(url, team_id=team_id) - cached_path = cache_video_from_bytes(raw_bytes, ext=ext) + cached_path = await cache_video_from_bytes_async(raw_bytes, ext=ext) logger.debug("[Slack] Cached user video: %s", cached_path) return cached_path, SUPPORTED_VIDEO_TYPES.get(ext, mimetype or "video/mp4"), "" return await self._cache_slack_document(f, url, mimetype, team_id) @@ -4243,7 +4243,7 @@ class SlackAdapter(BasePlatformAdapter): logger.warning("[Slack] Document too large or unknown size: %s", file_size) return None raw_bytes = await self._download_slack_file_bytes(url, team_id=team_id) - cached_path = cache_document_from_bytes( + cached_path = await cache_document_from_bytes_async( raw_bytes, original_filename or f"document{ext or '.bin'}") doc_mime = SUPPORTED_DOCUMENT_TYPES.get(ext, mimetype or "application/octet-stream") logger.debug("[Slack] Cached user document: %s (%s)", cached_path, doc_mime) @@ -5276,9 +5276,9 @@ class SlackAdapter(BasePlatformAdapter): async def _download_slack_file( self, url: str, ext: str, audio: bool = False, team_id: str = "") -> str: """Download a Slack image/audio file and cache it; returns the cached path.""" - from gateway.platforms.base import cache_audio_from_bytes, cache_image_from_bytes + from gateway.platforms.base import cache_audio_from_bytes_async, cache_image_from_bytes_async data = await self._download_slack_file_bytes(url, team_id=team_id, html_label="media") - return (cache_audio_from_bytes if audio else cache_image_from_bytes)(data, ext) + return await (cache_audio_from_bytes_async if audio else cache_image_from_bytes_async)(data, ext) # ── Channel mention gating ───────────────────────────────────────────── diff --git a/plugins/platforms/teams/adapter.py b/plugins/platforms/teams/adapter.py index 301d39cd81..1a9cf21123 100644 --- a/plugins/platforms/teams/adapter.py +++ b/plugins/platforms/teams/adapter.py @@ -51,7 +51,7 @@ HttpMethod = str # type: ignore[assignment,misc] from gateway.config import Platform, PlatformConfig from gateway.platforms.helpers import MessageDeduplicator from gateway.platforms.base import ( - gateway_trust_env, BasePlatformAdapter, MessageEvent, MessageType, SendResult, cache_image_from_url, cache_media_bytes, + gateway_trust_env, BasePlatformAdapter, MessageEvent, MessageType, SendResult, cache_image_from_url, cache_media_bytes_async, ) from gateway.platforms._shared import coerce_port, get_scoped_secret as _get_scoped_secret from plugins.platforms.teams.summary_writer import TeamsSummaryWriter # noqa: F401 — re-exported for teams_pipeline @@ -496,7 +496,7 @@ class TeamsAdapter(BasePlatformAdapter): filename = att_name or (f"document.{file_type}" if file_type else "document") try: data = await self._fetch_attachment_bytes(download_url) - cached = cache_media_bytes(data, filename=filename, mime_type="") + cached = await cache_media_bytes_async(data, filename=filename, mime_type="") if not cached: logger.warning("[teams] Unsupported document type for attachment '%s', skipping", filename) return None @@ -510,7 +510,7 @@ class TeamsAdapter(BasePlatformAdapter): # Connector URL needs the bot's bearer token; the generic cache helper sends none. data = await self._fetch_attachment_bytes(content_url) ext = content_type.split("/")[-1].split(";")[0] or "png" - cached = cache_media_bytes(data, filename=att_name or f"image.{ext}", mime_type=content_type) + cached = await cache_media_bytes_async(data, filename=att_name or f"image.{ext}", mime_type=content_type) if not cached: logger.warning( "[teams] Bot Framework attachment '%s' returned data that failed image validation, skipping", @@ -525,7 +525,7 @@ class TeamsAdapter(BasePlatformAdapter): if content_url: # direct-URL non-image attachment (video/audio/document) try: data = await self._fetch_attachment_bytes(content_url) - cached = cache_media_bytes(data, filename=att_name, mime_type=content_type) + cached = await cache_media_bytes_async(data, filename=att_name, mime_type=content_type) return (cached.path, cached.media_type, cached.kind) if cached else None except Exception as e: logger.warning("[teams] Failed to cache attachment '%s' (%s): %s", att_name or content_url, content_type, e) diff --git a/plugins/platforms/telegram/adapter.py b/plugins/platforms/telegram/adapter.py index 818d4a3e1e..7c2cbeb47f 100644 --- a/plugins/platforms/telegram/adapter.py +++ b/plugins/platforms/telegram/adapter.py @@ -178,7 +178,7 @@ from gateway.authz_mixin import _coerce_allow_set from gateway.config import Platform, PlatformConfig from gateway.platforms.base import ( BasePlatformAdapter, MessageEvent, MessageType, ProcessingOutcome, SendResult, - classify_send_error, cache_image_from_bytes, cache_audio_from_bytes, cache_video_from_bytes, + classify_send_error, cache_image_from_bytes_async, cache_audio_from_bytes_async, cache_video_from_bytes_async, resolve_proxy_url, SUPPORTED_VIDEO_TYPES, SUPPORTED_DOCUMENT_TYPES, SUPPORTED_IMAGE_DOCUMENT_TYPES, _TEXT_INJECT_EXTENSIONS, utf16_len, ) @@ -632,6 +632,13 @@ class TelegramAdapter(BasePlatformAdapter): # slow Bot API call (set_my_commands stall) can't blow the gateway connect timeout. self._post_connect_task: Optional[asyncio.Task] = None + @property + def send_path_degraded(self) -> bool: + # True from polling-generation start until the first getUpdates + # round-trip is proven (_record_polling_progress), and again at every + # polling-death site. getattr: tests build adapters via object.__new__(). + return bool(getattr(self, "_send_path_degraded", False)) + def _mark_connected(self) -> None: self._drop_delayed_deliveries = False super()._mark_connected() @@ -2001,6 +2008,13 @@ class TelegramAdapter(BasePlatformAdapter): self._polling_conflict_recovery_generation = None else: self._polling_conflict_count = 0 + # First proof getUpdates is flowing for this generation: flip a + # published "retrying" (degraded connect, reconnect stamp, or the + # mid-session recovery below) back to "connected" (#101391). + if self._send_path_degraded and getattr(self, "_running", False) and not self.has_fatal_error: + self._write_runtime_status_safe( + "connected", platform_state="connected", error_code=None, error_message=None, + ) self._send_path_degraded = False def _observe_polling_request_result(self, request, generation, result): @@ -2181,6 +2195,11 @@ class TelegramAdapter(BasePlatformAdapter): ) return self._send_path_degraded = True + # Polling died mid-session on an adapter that published "connected" + # at connect time. Without this, gateway_state.json keeps saying + # connected for as long as the recovery ladder runs (#101391: 11 h). + if getattr(self, "_running", False): + self._mark_degraded() logger.warning( "[%s] Telegram polling degraded (%s); gateway stays alive and will retry. Error: %s", self.name, reason, _redact_telegram_error_text(error), @@ -7061,7 +7080,7 @@ class TelegramAdapter(BasePlatformAdapter): ``"oversized"`` (skipped, ``cached`` is the raw file_size), ``"failed"`` (download error, logged), ``"unreadable"`` (cache rejected it) or ``"ok"``. """ - from gateway.platforms.base import cache_media_bytes + from gateway.platforms.base import cache_media_bytes_async source, filename, mime, kind = self._observed_media_source(msg) if source is None: return "none", None @@ -7078,7 +7097,7 @@ class TelegramAdapter(BasePlatformAdapter): data = bytes(await file_obj.download_as_bytearray()) if not filename: filename = os.path.basename(getattr(file_obj, "file_path", "") or "") - cached = cache_media_bytes(data, filename=filename, mime_type=mime, default_kind=kind) + cached = await cache_media_bytes_async(data, filename=filename, mime_type=mime, default_kind=kind) except Exception as exc: logger.warning("[Telegram] Failed to cache %s: %s", what, _redact_telegram_error_text(exc), exc_info=True) return "failed", None @@ -7550,10 +7569,10 @@ class TelegramAdapter(BasePlatformAdapter): if file_obj.file_path.lower().endswith(candidate): ext = candidate break - cached_path = cache_video_from_bytes(bytes(data), ext=ext) + cached_path = await cache_video_from_bytes_async(bytes(data), ext=ext) mime = SUPPORTED_VIDEO_TYPES.get(ext, "video/mp4") else: - cached_path = cache_audio_from_bytes(bytes(data), ext=ext) + cached_path = await cache_audio_from_bytes_async(bytes(data), ext=ext) event.media_urls = [cached_path] event.media_types = [mime] logger.info("[Telegram] Cached user %s at %s", kind, cached_path) @@ -7610,7 +7629,7 @@ class TelegramAdapter(BasePlatformAdapter): if file_obj.file_path.lower().endswith(candidate): ext = candidate break - cached_path = cache_image_from_bytes(bytes(image_bytes), ext=ext) + cached_path = await cache_image_from_bytes_async(bytes(image_bytes), ext=ext) event.media_urls = [cached_path] event.media_types = [f"image/{ext.lstrip('.')}"] logger.info("[Telegram] Cached user photo at %s", cached_path) @@ -7659,7 +7678,7 @@ class TelegramAdapter(BasePlatformAdapter): image_bytes = await file_obj.download_as_bytearray() image_ext = ext if ext in _TELEGRAM_IMAGE_EXTENSIONS else _TELEGRAM_IMAGE_MIME_TO_EXT.get(doc_mime, ".jpg") try: - cached_path = cache_image_from_bytes(bytes(image_bytes), ext=image_ext) + cached_path = await cache_image_from_bytes_async(bytes(image_bytes), ext=image_ext) except ValueError as e: logger.warning("[Telegram] Failed to cache image document: %s", _redact_telegram_error_text(e), exc_info=True) event.text = f"Image document '{original_filename or doc_mime or ext or 'unknown'}' could not be read as an image." @@ -7683,7 +7702,7 @@ class TelegramAdapter(BasePlatformAdapter): if ext in SUPPORTED_VIDEO_TYPES: file_obj = await doc.get_file() video_bytes = await file_obj.download_as_bytearray() - cached_path = cache_video_from_bytes(bytes(video_bytes), ext=ext) + cached_path = await cache_video_from_bytes_async(bytes(video_bytes), ext=ext) event.media_urls = [cached_path] event.media_types = [SUPPORTED_VIDEO_TYPES[ext]] event.message_type = MessageType.VIDEO @@ -7697,8 +7716,8 @@ class TelegramAdapter(BasePlatformAdapter): file_obj = await doc.get_file() doc_bytes = await file_obj.download_as_bytearray() raw_bytes = bytes(doc_bytes) - from gateway.platforms.base import cache_media_bytes - cached = cache_media_bytes( + from gateway.platforms.base import cache_media_bytes_async + cached = await cache_media_bytes_async( raw_bytes, filename=original_filename or f"document{ext or '.bin'}", mime_type=doc_mime ) if cached is None: @@ -7807,7 +7826,7 @@ class TelegramAdapter(BasePlatformAdapter): try: file_obj = await sticker.get_file() image_bytes = await file_obj.download_as_bytearray() - cached_path = cache_image_from_bytes(bytes(image_bytes), ext=".webp") + cached_path = await cache_image_from_bytes_async(bytes(image_bytes), ext=".webp") logger.info("[Telegram] Analyzing sticker at %s", cached_path) from tools.vision_tools import vision_analyze_tool result_json = await vision_analyze_tool(image_url=cached_path, user_prompt=STICKER_VISION_PROMPT) diff --git a/plugins/platforms/wecom/media.py b/plugins/platforms/wecom/media.py index f31b45c085..6081724763 100644 --- a/plugins/platforms/wecom/media.py +++ b/plugins/platforms/wecom/media.py @@ -17,7 +17,7 @@ from pathlib import Path from typing import Any, Dict, List, Optional, Tuple from urllib.parse import unquote, urlparse -from gateway.platforms.base import SendResult, cache_document_from_bytes, cache_image_from_bytes +from gateway.platforms.base import SendResult, cache_document_from_bytes_async, cache_image_from_bytes_async logger = logging.getLogger("plugins.platforms.wecom.adapter") @@ -102,9 +102,9 @@ class WeComMediaMixin: return None if kind == "image": ext = self._detect_image_ext(raw) - return self._cache_image(raw, ext, self._mime_for_ext(ext, fallback="image/jpeg"), "") + return await self._cache_image(raw, ext, self._mime_for_ext(ext, fallback="image/jpeg"), "") filename = str(media.get("filename") or media.get("name") or "wecom_file") - return cache_document_from_bytes(raw, filename), mimetypes.guess_type(filename)[0] or "application/octet-stream" + return await cache_document_from_bytes_async(raw, filename), mimetypes.guess_type(filename)[0] or "application/octet-stream" url = str(media.get("url") or "").strip() if not url: @@ -124,13 +124,13 @@ class WeComMediaMixin: content_type = str(headers.get("content-type") or "").split(";", 1)[0].strip() or "application/octet-stream" if kind == "image": ext = self._guess_extension(url, content_type, fallback=self._detect_image_ext(raw)) - return self._cache_image(raw, ext, content_type or self._mime_for_ext(ext, fallback="image/jpeg"), f" from {url}") + return await self._cache_image(raw, ext, content_type or self._mime_for_ext(ext, fallback="image/jpeg"), f" from {url}") filename = self._guess_filename(url, headers.get("content-disposition"), content_type) - return cache_document_from_bytes(raw, filename), content_type + return await cache_document_from_bytes_async(raw, filename), content_type - def _cache_image(self, raw: bytes, ext: str, mime: str, origin: str) -> Optional[Tuple[str, str]]: + async def _cache_image(self, raw: bytes, ext: str, mime: str, origin: str) -> Optional[Tuple[str, str]]: try: - return cache_image_from_bytes(raw, ext), mime + return await cache_image_from_bytes_async(raw, ext), mime except ValueError as exc: logger.warning("[%s] Rejected non-image bytes%s: %s", self.name, origin, exc) return None diff --git a/pyproject.toml b/pyproject.toml index fbd324439c..9f29c87976 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -161,13 +161,13 @@ dependencies = [ # Hence the ``sys_platform == 'win32'`` marker: the dep (and its portalocker # / pywin32 tree) ships only where it's actually used. "concurrent-log-handler==0.9.29; sys_platform == 'win32'", - # 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. 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). + # First-party lifecycle and shared-metrics runtime. Relay 0.8 is the supported + # native runtime and provider-codec baseline. Managed calls pass request and + # response data through this native module in-process; shared metrics installs + # no network exporter and consumes only its bounded projection. This marker + # must stay false anywhere no compatible native wheel tag can match; otherwise + # installing Python dependencies fails 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 @@ -175,7 +175,7 @@ dependencies = [ # 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.7.1,<0.8; (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')", + "nemo-relay>=0.8.3,<0.9; (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/run_agent.py b/run_agent.py index ad259fc81a..191801bd10 100644 --- a/run_agent.py +++ b/run_agent.py @@ -911,8 +911,12 @@ class AIAgent( _quietly(self._drop_shared_client, lambda c: self._close_openai_client(c, reason="agent_close", shared=True)) self._close_request_clients("agent_close") _quietly(self._close_codex_session) - # Free conversation history proactively: callers may still hold the closed agent. + # Free conversation history proactively: callers may still hold the closed agent. The DB-flush + # settled-prefix snapshot and the streamed-text accumulator are shadow copies of the same transcript; + # on a closed delegate child they were the only remaining owners, pinning its history in the parent heap. self._session_messages = [] + self._db_flush_scan_prefix = None + self._streamed_assistant_text_parts = [] _quietly(self._trim_process_memory) _quietly(self._finalize_owned_session_row) diff --git a/scripts/install.ps1 b/scripts/install.ps1 index 07930fd72d..6e64bc79c9 100644 --- a/scripts/install.ps1 +++ b/scripts/install.ps1 @@ -4733,11 +4733,15 @@ function Write-Completion { # or arrange to provide answers another way." $InstallStages = @( @{ Name = "uv"; Title = "Installing uv package manager"; Category = "prereqs"; NeedsUserInput = $false; Worker = "Stage-Uv" } - @{ Name = "python"; Title = "Verifying Python $PythonVersion"; Category = "prereqs"; NeedsUserInput = $false; Worker = "Stage-Python" } @{ Name = "git"; Title = "Installing Git"; Category = "prereqs"; NeedsUserInput = $false; Worker = "Stage-Git" } @{ Name = "node"; Title = "Detecting Node.js"; Category = "prereqs"; NeedsUserInput = $false; Worker = "Stage-Node" } @{ Name = "system-packages"; Title = "Installing ripgrep and ffmpeg"; Category = "prereqs"; NeedsUserInput = $false; Worker = "Stage-SystemPackages" } @{ Name = "repository"; Title = "Cloning Hermes repository"; Category = "install"; NeedsUserInput = $false; Worker = "Stage-Repository" } + # Managed Python lives under $InstallDir\.hermes-runtime, so the checkout + # must exist before this stage creates that directory. Otherwise the later + # repository stage treats the runtime-only directory as a broken checkout, + # parks it, and leaves Stage-Venv with no managed interpreter. + @{ Name = "python"; Title = "Verifying Python $PythonVersion"; Category = "prereqs"; NeedsUserInput = $false; Worker = "Stage-Python" } @{ Name = "venv"; Title = "Creating Python virtual environment"; Category = "install"; NeedsUserInput = $false; Worker = "Stage-Venv" } @{ Name = "dependencies"; Title = "Installing Python dependencies"; Category = "install"; NeedsUserInput = $false; Worker = "Stage-Dependencies" } @{ Name = "node-deps"; Title = "Installing Node.js dependencies"; Category = "install"; NeedsUserInput = $false; Worker = "Stage-NodeDeps" } diff --git a/tests/agent/lsp/_mock_lsp_server.py b/tests/agent/lsp/_mock_lsp_server.py index d7ce410151..57bfedbe19 100644 --- a/tests/agent/lsp/_mock_lsp_server.py +++ b/tests/agent/lsp/_mock_lsp_server.py @@ -103,6 +103,14 @@ def main(): if msg.get("method") == "workspace/didChangeWatchedFiles": continue + if msg.get("method") == "workspace/didChangeWorkspaceFolders": + # Multi-root tests observe attached folders through this log. + log_path = os.environ.get("MOCK_LSP_FOLDERS_LOG") + if log_path: + with open(log_path, "a", encoding="utf-8") as fh: + fh.write(json.dumps(msg.get("params")) + "\n") + continue + if msg.get("method") in {"textDocument/didOpen", "textDocument/didChange"}: params = msg.get("params") or {} td = params.get("textDocument") or {} diff --git a/tests/agent/lsp/test_multi_root.py b/tests/agent/lsp/test_multi_root.py new file mode 100644 index 0000000000..4f05aaac80 --- /dev/null +++ b/tests/agent/lsp/test_multi_root.py @@ -0,0 +1,125 @@ +"""Multi-root servers share ONE process across project roots. + +A profiled session with subagents editing across ~30 git worktrees ran +30-60 pyright processes. Pyright supports multi-root workspaces, so +the service keys such clients by ``server_id`` alone and attaches each +new root via ``workspace/didChangeWorkspaceFolders``. Single-root +servers keep the one-client-per-root behaviour. +""" +from __future__ import annotations + +import json +import sys +from pathlib import Path + +import pytest + +from agent.lsp.manager import LSPService +from agent.lsp.servers import SERVERS, ServerContext, ServerDef, SpawnSpec +from agent.lsp.workspace import clear_cache + +MOCK_SERVER = str(Path(__file__).parent / "_mock_lsp_server.py") + + +@pytest.fixture(autouse=True) +def _clear_workspace_cache(): + clear_cache() + yield + clear_cache() + + +def _make_repo(tmp_path: Path, name: str) -> Path: + repo = tmp_path / name + repo.mkdir() + (repo / ".git").mkdir() + (repo / "pyproject.toml").write_text("", encoding="utf-8") + (repo / "x.py").write_text("print('hi')\n", encoding="utf-8") + return repo + + +@pytest.fixture +def two_repos(tmp_path): + return _make_repo(tmp_path, "repo-a"), _make_repo(tmp_path, "repo-b") + + +@pytest.fixture +def mock_pyright(monkeypatch, tmp_path): + """Install the mock as ``pyright``; yield (spawn_count, folders_log, set_multi_root).""" + idx = next(i for i, s in enumerate(SERVERS) if s.server_id == "pyright") + original = SERVERS[idx] + spawns = {"value": 0} + folders_log = tmp_path / "folders.jsonl" + + def _spawn(root: str, ctx: ServerContext) -> SpawnSpec: + spawns["value"] += 1 + return SpawnSpec( + command=[sys.executable, MOCK_SERVER], + workspace_root=root, + cwd=root, + env={"MOCK_LSP_SCRIPT": "errors", "MOCK_LSP_FOLDERS_LOG": str(folders_log)}, + ) + + def _install(multi_root: bool) -> None: + SERVERS[idx] = ServerDef( + server_id="pyright", + extensions=original.extensions, + resolve_root=lambda fp, ws: ws, + build_spawn=_spawn, + multi_root=multi_root, + description="mock pyright", + ) + + yield spawns, folders_log, _install + SERVERS[idx] = original + + +def _service() -> LSPService: + return LSPService( + enabled=True, wait_mode="document", wait_timeout=3.0, install_strategy="manual" + ) + + +def test_multi_root_server_shares_one_client_across_roots(two_repos, mock_pyright, monkeypatch): + repo_a, repo_b = two_repos + spawns, folders_log, install = mock_pyright + install(multi_root=True) + svc = _service() + try: + monkeypatch.chdir(str(repo_a)) + diags_a = svc.get_diagnostics_sync(str(repo_a / "x.py")) + monkeypatch.chdir(str(repo_b)) + diags_b = svc.get_diagnostics_sync(str(repo_b / "x.py")) + + # Exactly one process; the second root arrived as a folder change. + assert spawns["value"] == 1 + assert len(svc._clients) == 1 + client = next(iter(svc._clients.values())) + assert client.workspace_folders == [str(repo_a), str(repo_b)] + events = [json.loads(line) for line in folders_log.read_text(encoding="utf-8").splitlines()] + assert [f["uri"] for e in events for f in e["event"]["added"]] == [ + Path(repo_b).as_uri() + ] + # Diagnostics still resolve per file in both folders. + assert len(diags_a) == 1 and len(diags_b) == 1 + status = svc.get_status()["clients"][0] + assert status["workspace_root"] == str(repo_a) + assert status["workspace_folders"] == [str(repo_a), str(repo_b)] + finally: + svc.shutdown() + + +def test_single_root_server_still_spawns_per_root(two_repos, mock_pyright, monkeypatch): + repo_a, repo_b = two_repos + spawns, folders_log, install = mock_pyright + install(multi_root=False) + svc = _service() + try: + monkeypatch.chdir(str(repo_a)) + svc.get_diagnostics_sync(str(repo_a / "x.py")) + monkeypatch.chdir(str(repo_b)) + svc.get_diagnostics_sync(str(repo_b / "x.py")) + assert spawns["value"] == 2 + assert set(svc._clients) == {("pyright", str(repo_a)), ("pyright", str(repo_b))} + assert not folders_log.exists() + finally: + svc.shutdown() diff --git a/tests/agent/lsp/test_workspace.py b/tests/agent/lsp/test_workspace.py index 8d96f902d6..7a0a6abcd7 100644 --- a/tests/agent/lsp/test_workspace.py +++ b/tests/agent/lsp/test_workspace.py @@ -49,6 +49,19 @@ def test_nearest_root_finds_first_marker(tmp_path: Path): assert found == str(root) +def test_nearest_root_skips_package_dirs(tmp_path: Path): + # hermes_cli/setup.py is a module inside a package, not a project + # marker; treating it as one spawned a second pyright per worktree. + root = tmp_path / "p" + pkg = root / "hermes_cli" + pkg.mkdir(parents=True) + (root / "pyproject.toml").write_text("") + (pkg / "__init__.py").write_text("") + (pkg / "setup.py").write_text("") + found = nearest_root(str(pkg / "main.py"), ["pyproject.toml", "setup.py"]) + assert found == str(root) + + diff --git a/tests/agent/test_auxiliary_client_nous_401_cache_key.py b/tests/agent/test_auxiliary_client_nous_401_cache_key.py index 194ab1c978..dcfde071b5 100644 --- a/tests/agent/test_auxiliary_client_nous_401_cache_key.py +++ b/tests/agent/test_auxiliary_client_nous_401_cache_key.py @@ -74,7 +74,7 @@ def test_call_llm_auto_provider_evicts_stale_client_end_to_end(monkeypatch): # The 401 refresh rebuilds a fresh client from refreshed runtime creds. monkeypatch.setattr( ac, "_resolve_nous_runtime_api", - lambda *, force_refresh=False: ("fresh-key", NOUS_BASE_URL), + lambda *, force_refresh=False, stale_access_token=None: ("fresh-key", NOUS_BASE_URL), ) monkeypatch.setattr( ac, "_create_openai_client", @@ -123,7 +123,7 @@ async def test_async_call_llm_auto_provider_evicts_stale_client_end_to_end(monke ) monkeypatch.setattr( ac, "_resolve_nous_runtime_api", - lambda *, force_refresh=False: ("fresh-key", NOUS_BASE_URL), + lambda *, force_refresh=False, stale_access_token=None: ("fresh-key", NOUS_BASE_URL), ) # Async refresh builds a sync client then wraps it; patch the wrap to `fresh`. monkeypatch.setattr( diff --git a/tests/agent/test_cjk_token_estimation.py b/tests/agent/test_cjk_token_estimation.py index 003d39d1eb..3c39c54cd0 100644 --- a/tests/agent/test_cjk_token_estimation.py +++ b/tests/agent/test_cjk_token_estimation.py @@ -49,15 +49,16 @@ def test_cjk_tail_does_not_expand_to_english_char_budget(): def _reference_per_char_estimate(text: str) -> int: - """The pre-perf-gate per-character reference implementation.""" + """Per-character reference: CJK ~1 token/char, everything else UTF-8 + bytes/4 (the byte width corrects Cyrillic/Greek/Arabic under-counting).""" dense = 0 - sparse = 0 + sparse_bytes = 0 for ch in text: if _is_cjk_token_dense_char(ch): dense += 1 else: - sparse += 1 - return dense + ((sparse + 3) // 4) + sparse_bytes += len(ch.encode("utf-8")) + return dense + ((sparse_bytes + 3) // 4) def test_perf_gated_estimator_matches_per_char_reference(): @@ -79,3 +80,40 @@ def test_perf_gated_estimator_matches_per_char_reference(): +def test_cyrillic_counts_by_utf8_bytes(): + # «русский текст» = 12 Cyrillic chars (2 bytes each) + 1 ASCII space: + # 25 bytes -> ceil(25/4) = 7 tokens; the old chars/4 rule said 4 — + # the ~2x under-count that let real prompts ride the context ceiling. + from agent.model_metadata import estimate_tokens_rough + assert estimate_tokens_rough("русский текст") == 7 + # Pure ASCII unchanged. + assert estimate_tokens_rough("a" * 400) == 100 + + +def test_accented_latin_is_not_inflated_by_byte_counting(): + # Byte-counting must not punish Western-European text: only the accented + # chars are 2 bytes, so the estimate moves by a few percent, not 2x. + from agent.model_metadata import estimate_tokens_rough + fr = "La compression du contexte permet aux longues sessions de rester dans la fenêtre du fournisseur sans perdre le fil de la tâche." + ascii_rule = (len(fr) + 3) // 4 + est = estimate_tokens_rough(fr) + assert ascii_rule <= est <= int(ascii_rule * 1.10), (ascii_rule, est) + + +def test_mixed_cyrillic_and_ascii_code_counts_ascii_at_one_byte(): + from agent.model_metadata import estimate_tokens_rough + code = "def compress(ctx):\n # Сжимаем контекст\n return summarize(ctx)\n" + ascii_part = "def compress(ctx):\n # \n return summarize(ctx)\n" + cyr = "Сжимаем контекст" + expected = (len(ascii_part.encode()) + len(cyr.encode()) + 3) // 4 + assert estimate_tokens_rough(code) == expected + # and strictly more than the old chars/4 rule for the same text + assert estimate_tokens_rough(code) > (len(code) + 3) // 4 + + +def test_lone_surrogates_do_not_raise(): + # main's estimator was total (len/regex never raise); byte-counting must + # stay total too — tool output routinely carries unpaired surrogates. + from agent.model_metadata import estimate_tokens_rough + assert estimate_tokens_rough("abc\ud800def") >= 2 + assert estimate_tokens_rough("漢字\udfff") >= 2 diff --git a/tests/agent/test_compression_concurrent_fork.py b/tests/agent/test_compression_concurrent_fork.py index a155b5ab65..831f1f9584 100644 --- a/tests/agent/test_compression_concurrent_fork.py +++ b/tests/agent/test_compression_concurrent_fork.py @@ -167,6 +167,105 @@ def test_compression_activity_heartbeat_touches_agent_during_long_compress(tmp_p assert db.get_compression_lock_holder(session_id) is None +def test_compression_activity_heartbeat_emits_client_status_events(tmp_path: Path) -> None: + """The heartbeat must re-emit the compacting status, not just DB touches. + + Remote transports (e.g. the Android relay app) run idle-progress turn + watchdogs that ``session.interrupt`` a turn after ~180s with no gateway + events. Compression is silent on the event stream, so without periodic + status heartbeats a long compression is killed mid-flight and retriggers + forever on sessions near the context ceiling. + """ + from agent.conversation_compression import ( + COMPACTION_HEARTBEAT_STATUS, + is_compaction_progress_status, + ) + + db = SessionDB(db_path=tmp_path / "state.db") + session_id = "HEARTBEAT_STATUS_TEST" + db.create_session(session_id, source="test") + + agent = _build_agent_with_db(db, session_id) + agent._compression_activity_heartbeat_interval = 0.1 + touch_calls: list[str] = [] + agent._touch_activity = lambda desc, **_kw: touch_calls.append(desc) + status_events: list[tuple[str, str]] = [] + setattr( + agent, + "status_callback", + lambda event, message: status_events.append((event, message)), + ) + + def _slow_compress(*_a, **_kw): + _wait_for_touch(touch_calls, "context compression in progress") + return [ + {"role": "user", "content": "[CONTEXT COMPACTION] summary"}, + {"role": "user", "content": "tail"}, + ] + + agent.context_compressor.compress.side_effect = _slow_compress + messages = [{"role": "user", "content": f"m{i}"} for i in range(20)] + + agent._compress_context(messages, "sys", approx_tokens=120_000) + + heartbeats = [e for e in status_events if e[1] == COMPACTION_HEARTBEAT_STATUS] + assert heartbeats, "no heartbeat status reached the client" + # Same "lifecycle" key as the other compaction statuses so the TUI gateway + # re-tags it to kind="compacting" and Telegram edits one bubble in place. + assert {event for event, _ in heartbeats} == {"lifecycle"} + assert is_compaction_progress_status(COMPACTION_HEARTBEAT_STATUS) + # Exactly one routine start line precedes the first heartbeat; the + # heartbeat no longer re-emits a start of its own (adapters without + # send_or_update_status would otherwise post two messages). + assert status_events[0][1] != COMPACTION_HEARTBEAT_STATUS + # Every heartbeat is a periodic tick: none may precede the first + # "in progress" DB touch, which is what start() would have produced. + first_tick_touch = touch_calls.index("context compression in progress") + assert first_tick_touch >= 1 # "started" touch came first + assert len(heartbeats) <= touch_calls.count("context compression in progress") + + +def test_compression_heartbeat_is_silent_for_quiet_context_engines(tmp_path: Path) -> None: + """A context engine that suppresses the routine start status opens no + visible compaction phase; the heartbeat must not open one either (there + would be no terminal edge to close it).""" + from agent.conversation_compression import COMPACTION_HEARTBEAT_STATUS + + db = SessionDB(db_path=tmp_path / "state.db") + session_id = "HEARTBEAT_QUIET_TEST" + db.create_session(session_id, source="test") + + agent = _build_agent_with_db(db, session_id) + agent._compression_activity_heartbeat_interval = 0.1 + touch_calls: list[str] = [] + agent._touch_activity = lambda desc, **_kw: touch_calls.append(desc) + status_events: list[tuple[str, str]] = [] + setattr( + agent, + "status_callback", + lambda event, message: status_events.append((event, message)), + ) + + def _slow_compress(*_a, **_kw): + _wait_for_touch(touch_calls, "context compression in progress") + return [ + {"role": "user", "content": "[CONTEXT COMPACTION] summary"}, + {"role": "user", "content": "tail"}, + ] + + agent.context_compressor.compress.side_effect = _slow_compress + messages = [{"role": "user", "content": f"m{i}"} for i in range(20)] + + with patch( + "agent.conversation_compression.automatic_compaction_status_message", + return_value="", + ): + agent._compress_context(messages, "sys", approx_tokens=120_000) + + assert all(m != COMPACTION_HEARTBEAT_STATUS for _, m in status_events) + assert "context compression in progress" in touch_calls # DB touches still ran + + def test_lock_contender_preserves_terminal_compaction_lifecycle(tmp_path: Path) -> None: """A lock loser still closes the structured compaction lifecycle. diff --git a/tests/agent/test_credential_pool_nous_refresh_stampede.py b/tests/agent/test_credential_pool_nous_refresh_stampede.py index fa328fa4a9..34a3128e69 100644 --- a/tests/agent/test_credential_pool_nous_refresh_stampede.py +++ b/tests/agent/test_credential_pool_nous_refresh_stampede.py @@ -99,3 +99,32 @@ def test_lock_timeout_during_nous_refresh_does_not_bench_entry(monkeypatch, capl assert result is entry assert pool._entries[0].last_status is None, "lock contention is not a credential failure" + + +def test_agent_401_refresh_passes_failed_bearer_as_stale_hint(monkeypatch): + """Every 401-recovery caller must hand the auth store the bearer that + failed — without it ``_already_rotated_by_peer`` can never fire and each + subagent rotates the shared grant again (the "no crash, N refreshes" + variant of the Sep 2 stampede). + """ + from run_agent import AIAgent + + agent = AIAgent.__new__(AIAgent) + agent.provider = "nous" + agent.api_mode = "chat_completions" + agent.api_key = "jwt-that-just-401d" + agent.base_url = "https://inference-api.nousresearch.com/v1" + agent._client_kwargs = {} + monkeypatch.setattr(agent, "_replace_primary_openai_client", lambda **k: True) + + seen = {} + + def _fake_resolve(**kwargs): + seen.update(kwargs) + return {"api_key": "fresh", "base_url": agent.base_url} + + monkeypatch.setattr(auth_mod, "resolve_nous_runtime_credentials", _fake_resolve) + + assert agent._try_refresh_nous_client_credentials(force=True) is True + assert seen["force_refresh"] is True + assert seen["stale_access_token"] == "jwt-that-just-401d" diff --git a/tests/agent/test_credential_pool_profile_oauth_fork.py b/tests/agent/test_credential_pool_profile_oauth_fork.py index 057db3021d..eec0c6c52e 100644 --- a/tests/agent/test_credential_pool_profile_oauth_fork.py +++ b/tests/agent/test_credential_pool_profile_oauth_fork.py @@ -450,3 +450,112 @@ def test_heal_is_a_noop_in_classic_mode(fleet): before = (fleet["root"] / "auth.json").read_text() assert heal_forked_single_use_oauth_grants("anthropic") is None assert (fleet["root"] / "auth.json").read_text() == before + + +# ── C. a SHARED root store is not a fork (#101356) ─────────────────────── + +def _seed_codex_grant(root): + """Give the root store an openai-codex pool row AND a providers block.""" + fresh = int((time.time() + 3600) * 1000) + store = json.loads((root / "auth.json").read_text()) + store["credential_pool"]["openai-codex"] = [{ + "id": "cdx001", "label": "codex", "auth_type": "oauth", "priority": 0, + "source": "manual:device_code", "access_token": "cdx-AT0", + "refresh_token": "cdx-RT0", "expires_at_ms": fresh, + }] + store["providers"]["openai-codex"] = { + "tokens": {"access_token": "cdx-AT0", "refresh_token": "cdx-RT0", "expires_at_ms": fresh}, + "last_refresh": fresh / 1000.0, + } + (root / "auth.json").write_text(json.dumps(store)) + + +def _shared_profile(fleet, name, *, link): + """Profile whose auth.json IS the root store (``link`` makes the alias).""" + pdir = _profile(fleet, name) + pdir.mkdir(parents=True, exist_ok=True) + alias = pdir / "auth.json" + if alias.is_symlink() or alias.exists(): + alias.unlink() + link(fleet["root"] / "auth.json", alias) + return pdir + + +def test_heal_skips_profile_auth_json_symlinked_to_the_root_store(fleet): + """#101356: `ln -s ~/.hermes/auth.json /auth.json` shares ONE store. + Both sides of the consolidation read the same file, so every row looks like + a fork of itself — healing would strip the shared grant through the link.""" + from hermes_cli.auth import consume_oauth_heal_notices, heal_forked_single_use_oauth_grants + + root = fleet["root"] + _seed_codex_grant(root) + before = (root / "auth.json").read_text() + + shared = _shared_profile(fleet, "shared", link=lambda target, alias: alias.symlink_to(target)) + fleet["use"](shared) + + assert heal_forked_single_use_oauth_grants("openai-codex") is None + assert (root / "auth.json").read_text() == before + assert (shared / "auth.json").is_symlink() + assert consume_oauth_heal_notices() == [] + store = json.loads((root / "auth.json").read_text()) + assert [r["id"] for r in store["credential_pool"]["openai-codex"]] == ["cdx001"] + assert store["providers"]["openai-codex"]["tokens"]["refresh_token"] == "cdx-RT0" + + +def test_heal_skips_profile_auth_json_hardlinked_to_the_root_store(fleet): + """Same class as the symlink: a hardlink resolves to a different name but + is the same inode, so it is still one store, not a forked copy.""" + from hermes_cli.auth import heal_forked_single_use_oauth_grants + + root = fleet["root"] + _seed_codex_grant(root) + before = (root / "auth.json").read_text() + + shared = _shared_profile(fleet, "twin", link=lambda target, alias: os.link(target, alias)) + fleet["use"](shared) + + assert heal_forked_single_use_oauth_grants("openai-codex") is None + assert (root / "auth.json").read_text() == before + assert (shared / "auth.json").samefile(root / "auth.json") + + +def test_heal_leaves_an_aliased_anthropic_singleton_alone(fleet): + """Separate auth.jsons but a profile `.anthropic_oauth.json` symlinked to + root's: one shared grant, not a fork. The heal must not self-compare it + or unlink the alias (#101356 sibling site).""" + from hermes_cli.auth import heal_forked_single_use_oauth_grants + + root = fleet["root"] + (root / ".anthropic_oauth.json").write_text(json.dumps({ + "accessToken": "AT-shared", "refreshToken": "RT-shared", + "expiresAt": int((time.time() + 3600) * 1000), + })) + kid = _profile(fleet, "kid") + kid.mkdir(parents=True, exist_ok=True) + (kid / "auth.json").write_text(json.dumps({"providers": {}, "credential_pool": {}})) + (kid / ".anthropic_oauth.json").symlink_to(root / ".anthropic_oauth.json") + before = (root / ".anthropic_oauth.json").read_text() + + fleet["use"](kid) + assert heal_forked_single_use_oauth_grants("anthropic") is None + assert (kid / ".anthropic_oauth.json").is_symlink() + assert (root / ".anthropic_oauth.json").read_text() == before + + +def test_heal_same_store_skip_is_memoized_off_the_hot_path(fleet, monkeypatch): + """The shared-store skip must record the clean mark so load_pool()'s + per-call heal does not re-stat/resolve both paths every model call.""" + from hermes_cli import auth as auth_mod + + root = fleet["root"] + _seed_codex_grant(root) + shared = _shared_profile(fleet, "shared", link=lambda target, alias: alias.symlink_to(target)) + fleet["use"](shared) + + assert auth_mod.heal_forked_single_use_oauth_grants("openai-codex") is None + assert "openai-codex" in auth_mod._oauth_heal_clean_marks + calls = [] + monkeypatch.setattr(auth_mod, "_is_same_auth_store", lambda *a: calls.append(a) or True) + assert auth_mod.heal_forked_single_use_oauth_grants("openai-codex") is None + assert calls == [], "same-store check ran again despite the clean mark" diff --git a/tests/agent/test_cron_inflight_prompt_reappend_100818.py b/tests/agent/test_cron_inflight_prompt_reappend_100818.py new file mode 100644 index 0000000000..ef1f8ccf24 --- /dev/null +++ b/tests/agent/test_cron_inflight_prompt_reappend_100818.py @@ -0,0 +1,333 @@ +"""Regression coverage for #100818: compaction of a single-prompt session +(the cron shape) must not leave the model with nothing to obey. + +A cron run is one user message — the job prompt — followed by nothing but +assistant/tool turns. When ContextCompressor fires mid-run, that prompt is +folded into the handoff summary and no user message survives *after* it. +SUMMARY_PREFIX then reads literally: + + If no user message appears AFTER this summary, do nothing. + +so the model correctly does nothing, the scheduler sees the ``[SILENT]`` +sentinel, and records ``last_status: ok`` — a silent failure. + +The fix re-appends the in-flight user task after the handoff so the prefix's +"latest user message" pointer resolves to the job prompt again. The +#80622 contract is unchanged: an idle session with no in-flight task must +still be left with nothing to act on. +""" + +from typing import Any, Dict, List +from unittest.mock import MagicMock, patch + +from agent.context_compressor import ( + _SUMMARY_END_MARKER, + SUMMARY_PREFIX, + ContextCompressor, +) + + +JOB_SENTINEL = "CRON_JOB_PROMPT_sentinel_brief_the_inbox_and_write_a_digest" + + +def _make_compressor() -> ContextCompressor: + with patch( + "agent.context_compressor.get_model_context_length", return_value=100_000 + ): + compressor = ContextCompressor( + model="test", + quiet_mode=True, + protect_first_n=2, + protect_last_n=2, + ) + compressor.tail_token_budget = 500 + return compressor + + +def _tool_pairs(count: int, start: int = 0) -> List[Dict[str, Any]]: + """``count`` assistant(tool_calls) + tool result pairs.""" + turns: List[Dict[str, Any]] = [] + for i in range(start, start + count): + turns.append( + { + "role": "assistant", + "content": f"step {i}", + "tool_calls": [ + {"id": f"c{i}", "function": {"name": "terminal", "arguments": "{}"}} + ], + } + ) + turns.append( + { + "role": "tool", + "tool_call_id": f"c{i}", + "content": ("tool output " * 200) + f" {i}", + } + ) + return turns + + +def _cron_transcript() -> List[Dict[str, Any]]: + """system + one user job prompt + many tool turns, NO trailing user.""" + return [ + { + "role": "system", + "content": "You are Hermes. Cron preamble: if nothing to report, " + "return [SILENT].", + }, + {"role": "user", "content": JOB_SENTINEL}, + *_tool_pairs(40), + ] + + +def _compress(messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + response = MagicMock() + response.choices = [MagicMock()] + response.choices[0].message.content = ( + "## Historical Task Snapshot\nUser asked: '" + JOB_SENTINEL + "'\n" + "## Summary\nRan a bunch of terminal steps." + ) + compressor = _make_compressor() + with patch("agent.context_compressor.call_llm", return_value=response): + return compressor.compress(messages, current_tokens=200_000, force=True) + + +def _handoff_idx(compressed: List[Dict[str, Any]]) -> int: + """Index of the handoff row (standalone summary or merged carrier).""" + for idx in range(len(compressed) - 1, -1, -1): + content = compressed[idx].get("content") + text = content if isinstance(content, str) else str(content) + if SUMMARY_PREFIX[:60] in text or _SUMMARY_END_MARKER in text: + return idx + return -1 + + +def _text(message: Dict[str, Any]) -> str: + content = message.get("content") + return content if isinstance(content, str) else str(content) + + +def _actionable_user_rows(rows: List[Dict[str, Any]]) -> List[Dict[str, Any]]: + return [ + m + for m in rows + if ContextCompressor._is_actionable_user_turn(m) + and not ContextCompressor._is_synthetic_compression_user_turn(m) + ] + + +def test_cron_job_prompt_survives_after_the_handoff(): + """The in-flight job prompt must be readable AFTER the summary boundary.""" + compressed = _compress(_cron_transcript()) + + idx = _handoff_idx(compressed) + assert idx >= 0, "expected a compaction handoff in the compressed transcript" + + after = compressed[idx + 1:] + # The re-append may also land inside the handoff carrier itself, after the + # end marker (the alternation-safe merge layout) — accept either shape. + carrier_tail = _text(compressed[idx]).split(_SUMMARY_END_MARKER)[-1] + + job_after_summary = any( + JOB_SENTINEL in _text(m) for m in _actionable_user_rows(after) + ) or (JOB_SENTINEL in carrier_tail) + assert job_after_summary, ( + "the in-flight cron job prompt must appear in a user message AFTER the " + "handoff summary — SUMMARY_PREFIX orders the model to do nothing " + "otherwise (#100818)" + ) + + +def test_model_is_not_left_without_a_user_message_after_the_handoff(): + """The 'no user message after this summary → do nothing' branch of + SUMMARY_PREFIX must not be what a mid-run cron compaction produces.""" + compressed = _compress(_cron_transcript()) + + idx = _handoff_idx(compressed) + after = compressed[idx + 1:] + has_user_after = bool(_actionable_user_rows(after)) or bool( + _text(compressed[idx]).split(_SUMMARY_END_MARKER)[-1].strip() + ) + assert has_user_after, ( + "compaction left no user message after the handoff; the model is " + "instructed to do nothing and the cron run fails silently" + ) + + +def test_role_alternation_and_head_are_preserved(): + """The re-append must not create two same-role rows in a row, and must + not disturb the cached head prefix.""" + messages = _cron_transcript() + compressed = _compress([dict(m) for m in messages]) + + assert compressed[0]["role"] == "system" + visible = [ + m.get("role") + for m in compressed + if not ( + m.get("role") == "tool" + or (m.get("role") == "assistant" and m.get("tool_calls")) + ) + ] + for previous, current in zip(visible, visible[1:]): + assert not (previous == current == "user"), ( + f"consecutive user rows in compressed transcript: {visible}" + ) + + +def test_idle_session_without_inflight_task_is_not_reanimated(): + """#80622 must hold: a session whose only user-role row is an inherited + handoff has no in-flight task, so compaction must not manufacture one.""" + messages: List[Dict[str, Any]] = [ + {"role": "system", "content": "You are Hermes."}, + { + "role": "user", + "content": ( + SUMMARY_PREFIX + + "\n## Historical Task Snapshot\nUser asked: 'a finished task'\n\n" + + _SUMMARY_END_MARKER + ), + }, + *_tool_pairs(40), + ] + compressed = _compress(messages) + + assert not _actionable_user_rows(compressed), ( + "no real user turn existed before compaction — none may be invented" + ) + + +def test_completed_exchange_is_not_replayed(): + """Only an in-flight task is re-appended. A turn that already produced a + final assistant reply must not be handed back to the model as a fresh + instruction.""" + messages = [ + *_cron_transcript(), + {"role": "assistant", "content": "Digest written. Nothing else to do."}, + ] + compressed = _compress(messages) + + idx = _handoff_idx(compressed) + after = compressed[idx + 1:] + assert not any( + JOB_SENTINEL in _text(m) for m in _actionable_user_rows(after) + ), "a completed exchange must not be re-appended as a new user instruction" + + +# --------------------------------------------------------------------------- +# Follow-up coverage (salvage of #101170) +# --------------------------------------------------------------------------- + + +def _pending_tail_transcript() -> List[Dict[str, Any]]: + """Cron shape whose LAST row is an assistant tool_calls turn still awaiting + its result — compaction fired inside the tool-execution window.""" + msgs = _cron_transcript() + msgs.append( + { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "pending", "function": {"name": "terminal", "arguments": "{}"}} + ], + } + ) + return msgs + + +def test_pending_trailing_tool_call_survives_the_replay(): + """The replay row must not make a genuinely in-flight tool_call look + orphaned: _sanitize_tool_pairs runs before the re-append, so the trailing + assistant(tool_calls) keeps its calls and the late tool result still pairs.""" + out = _compress(_pending_tail_transcript()) + pending = [ + m + for m in out + if m.get("role") == "assistant" + and any(tc.get("id") == "pending" for tc in (m.get("tool_calls") or [])) + ] + assert pending, "trailing in-flight tool_calls were stripped" + replay_idx = max( + i for i, m in enumerate(out) + if m.get("role") == "user" and JOB_SENTINEL in str(m.get("content")) + ) + assert replay_idx > out.index(pending[0]) + + +def _compress_with(protect_first_n: int, compression_count: int, messages): + response = MagicMock() + response.choices = [MagicMock()] + response.choices[0].message.content = SUMMARY_PREFIX + "\n## Summary\nran steps." + with patch( + "agent.context_compressor.get_model_context_length", return_value=100_000 + ): + compressor = ContextCompressor( + model="test", + quiet_mode=True, + protect_first_n=protect_first_n, + protect_last_n=2, + ) + compressor.tail_token_budget = 500 + compressor.compression_count = compression_count + with patch.object( + compressor, "_generate_summary", return_value=response.choices[0].message.content + ): + return compressor.compress(messages, current_tokens=90_000) + + +def _job_copies(messages) -> int: + return sum( + str(m.get("content")).count(JOB_SENTINEL) + for m in messages + if m.get("role") == "user" + ) + + +def test_merged_restatement_is_not_anchored_twice(): + """Later cycle: protect_first_n has decayed, the prompt is summarised away + and the summary is pinned to role=user followed by an exempt tool tail, so + the restatement is MERGED onto the carrier. The later + _ensure_compressed_has_user_turn pass must see intent as present instead + of inserting a second copy of the job prompt.""" + from agent.conversation_compression import _ensure_compressed_has_user_turn + + original = [{"role": "user", "content": JOB_SENTINEL}, *_tool_pairs(40)] + out = _compress_with(2, 1, original) + assert any(m.get("_inflight_replay_merged") for m in out), "expected merge layout" + assert _job_copies(out) == 1 + assert _ensure_compressed_has_user_turn(original, out) == "already_present" + assert _job_copies(out) == 1 + + +def test_restatement_survives_repeated_compactions_without_stacking(): + """A task alive across three compactions is restated exactly once per + output — one copy, one header, always after the summary.""" + from agent.context_compressor import _INFLIGHT_TASK_REPLAY_HEADER + + out = [{"role": "user", "content": JOB_SENTINEL}, *_tool_pairs(40)] + for cycle in (1, 2, 3): + extra = _tool_pairs(40, start=100 * cycle) if cycle > 1 else [] + out = _compress_with(2, cycle, out + extra) + users = [m for m in out if m.get("role") == "user"] + assert _job_copies(out) == 1, cycle + assert max( + str(m.get("content")).count(_INFLIGHT_TASK_REPLAY_HEADER) for m in users + ) == 1, cycle + last = str(users[-1].get("content")) + assert last.rfind(JOB_SENTINEL) > last.rfind(_SUMMARY_END_MARKER), cycle + + +def test_flagged_scaffolding_row_is_never_the_inflight_task(): + """A trailing user-role scaffolding row flagged synthetic (todo snapshot) + must not be mistaken for the live request and replayed as an instruction.""" + msgs = _cron_transcript() + msgs.append( + { + "role": "user", + "content": "[Your active task list was preserved across context compression]\n- x", + "_todo_snapshot_synthetic": True, + } + ) + found = ContextCompressor._find_inflight_user_task(msgs) + assert found is not None + assert JOB_SENTINEL in str(found.get("content")) diff --git a/tests/agent/test_idle_compaction.py b/tests/agent/test_idle_compaction.py index f4de6c819a..ba89099cfe 100644 --- a/tests/agent/test_idle_compaction.py +++ b/tests/agent/test_idle_compaction.py @@ -39,3 +39,66 @@ class TestShouldIdleCompact: def test_fires_just_above_floor(self): assert _decide(tokens=40_001, floor_tokens=40_000) is True + +class TestPostCompactionFloor: + """The floor also honours what the previous pass actually produced (#97239). + + ``floor_tokens`` is the theoretical target (threshold × target_ratio); a + real pass lands well above it because the system prompt, the tool schemas + and the protected head/tail are incompressible. Without this, an already + compacted session re-summarises itself on every idle resume forever. + """ + + def test_unrecorded_last_compaction_keeps_original_floor(self): + # 0 = nothing compacted yet (or state reset) — original semantics. + assert _decide(tokens=40_001, floor_tokens=40_000, + last_compaction_tokens=0) is True + + def test_skips_when_transcript_has_not_grown_since_last_compaction(self): + # Previous pass produced 44,000; the transcript is still ~that size. + assert _decide(tokens=44_100, floor_tokens=40_000, + last_compaction_tokens=44_000) is False + + def test_fires_once_a_full_floor_of_new_content_accumulated(self): + assert _decide(tokens=84_001, floor_tokens=40_000, + last_compaction_tokens=44_000) is True + + def test_does_not_fire_at_exactly_the_raised_floor(self): + assert _decide(tokens=84_000, floor_tokens=40_000, + last_compaction_tokens=44_000) is False + + def test_reported_session_stops_recompacting_itself(self): + """Exact numbers from issue #97239. + + The 17:01 pass reduced 64,105 -> 44,579 tokens; the 17:36 resume + re-fired on that same transcript because 44,579 > the 25,502 + theoretical floor, blocking the prompt for another 256 s. + """ + common = dict(idle_after_seconds=1, idle_gap_seconds=747.0, + floor_tokens=25_502) + # Before the fix the second resume fired: 44,579 > 25,502. + assert _decide(tokens=44_579, last_compaction_tokens=0, **common) is True + # With the previous pass's real output known, it sits the round out. + assert _decide(tokens=44_579, last_compaction_tokens=44_579, + **common) is False + + def test_an_effective_pass_still_raises_the_floor(self): + # 100K -> 10K is a good pass; another one is worth it only once about + # a floor's worth of new content has landed on top of the 10K. + assert _decide(tokens=30_000, floor_tokens=25_000, + last_compaction_tokens=10_000) is False + assert _decide(tokens=35_001, floor_tokens=25_000, + last_compaction_tokens=10_000) is True + + def test_other_gates_still_win_over_the_raised_floor(self): + # Growth alone must not defeat the cooldown / opt-out gates. + assert _decide(tokens=200_000, floor_tokens=40_000, + last_compaction_tokens=44_000, + cooldown_active=True) is False + assert _decide(tokens=200_000, floor_tokens=40_000, + last_compaction_tokens=44_000, + idle_after_seconds=0) is False + assert _decide(tokens=200_000, floor_tokens=40_000, + last_compaction_tokens=44_000, + idle_gap_seconds=0.5) is False + diff --git a/tests/agent/test_idle_compaction_lock_and_guards.py b/tests/agent/test_idle_compaction_lock_and_guards.py index 613488abe4..d1c30a024a 100644 --- a/tests/agent/test_idle_compaction_lock_and_guards.py +++ b/tests/agent/test_idle_compaction_lock_and_guards.py @@ -53,18 +53,20 @@ def _prep_idle_agent(db: SessionDB, session_id: str, *, idle_after: int = 60, return agent -def _run_prologue(agent, history, user_message="hello again"): +def _run_prologue(agent, history, user_message="hello again", + rough_tokens: int = 999_999): """Invoke ``build_turn_context`` the way ``conversation_loop`` does. The token-threshold preflight gate is pinned False so these tests exercise the IDLE trigger in isolation (the preflight path has its own - coverage in ``test_turn_context.py``). + coverage in ``test_turn_context.py``). ``rough_tokens`` pins the estimate + that the idle floor is compared against. """ with patch("agent.auxiliary_client.set_runtime_main", lambda *a, **k: None), \ patch("agent.turn_context._should_run_preflight_estimate", return_value=False), \ patch("agent.turn_context.estimate_request_tokens_rough", - return_value=999_999): + return_value=rough_tokens): return build_turn_context( agent=agent, user_message=user_message, @@ -145,6 +147,96 @@ def test_idle_compaction_defers_to_held_compression_lock(tmp_path: Path) -> None assert ctx.messages[ctx.current_turn_user_idx]["content"] == "hello again" +def _prep_recompaction_agent(db: SessionDB, sid: str): + """Idle-eligible agent with the #97239 threshold/floor numbers. + + threshold 127,510 x target_ratio 0.20 => a 25,502 theoretical floor, the + same one the reported session kept clearing while never actually + shrinking below ~44,579. + """ + agent = _prep_idle_agent(db, sid, idle_after=1, idle_gap=747.0) + agent.context_compressor.threshold_tokens = 127_510 + agent.context_compressor.summary_target_ratio = 0.20 + agent.context_compressor.emit_automatic_compaction_status = True + del agent.context_compressor.get_automatic_compaction_status_message + return agent + + +def _pin_compress_seam(agent): + """Stub the forwarder so these tests assert the idle DECISION only. + + Whether ``compress_context`` then rotates, locks or aborts is covered by + the tests above; here the question is purely whether the idle floor let + the turn through. Returning the input list is the documented "skipped" + shape, so the caller's re-baseline stays disarmed either way. + """ + seam = MagicMock(side_effect=lambda messages, *a, **k: (messages, "SYSTEM")) + agent._compress_context = seam + return seam + + +def test_idle_compaction_skips_a_transcript_that_has_not_grown(tmp_path: Path) -> None: + """The reported loop: re-compacting a session the last pass just produced. + + ``last_compression_rough_tokens`` records what the previous pass actually + emitted (44,579). The theoretical floor (25,502) is far below it, so the + old predicate re-fired a full multi-minute summary on every idle resume + even though the transcript had not grown at all (#97239). + """ + db = SessionDB(db_path=tmp_path / "state.db") + sid = "IDLE_RECOMPACT" + db.create_session(sid, source="cli") + agent = _prep_recompaction_agent(db, sid) + agent.context_compressor.last_compression_rough_tokens = 44_579 + seam = _pin_compress_seam(agent) + + ctx = _run_prologue(agent, _history(), rough_tokens=44_579) + + seam.assert_not_called() + agent.context_compressor.compress.assert_not_called() + assert agent.session_id == sid + assert len(ctx.messages) == len(_history()) + 1 + assert ctx.current_turn_user_idx == len(ctx.messages) - 1 + + +def test_idle_compaction_fires_again_once_the_transcript_grows(tmp_path: Path) -> None: + """The raised floor is a deferral, not an off switch.""" + db = SessionDB(db_path=tmp_path / "state.db") + sid = "IDLE_REGROWN" + db.create_session(sid, source="cli") + agent = _prep_recompaction_agent(db, sid) + agent.context_compressor.last_compression_rough_tokens = 44_579 + seam = _pin_compress_seam(agent) + + # 44,579 + 25,502 = 70,081 — one floor's worth of new content on top. + _run_prologue(agent, _history(), rough_tokens=70_082) + + seam.assert_called_once() + + +def test_idle_compaction_ignores_a_non_int_last_compaction_reading( + tmp_path: Path, +) -> None: + """Compressor doubles expose a Mock here — it must not raise the floor. + + An unset/derived attribute falls back to 0, which restores the original + ``tokens > floor_tokens`` semantics exactly. + """ + db = SessionDB(db_path=tmp_path / "state.db") + sid = "IDLE_MOCKREAD" + db.create_session(sid, source="cli") + agent = _prep_recompaction_agent(db, sid) + # Left as the MagicMock auto-attribute (a truthy non-int). + assert not isinstance( + agent.context_compressor.last_compression_rough_tokens, int + ) + seam = _pin_compress_seam(agent) + + _run_prologue(agent, _history(), rough_tokens=44_579) + + seam.assert_called_once() + + def test_idle_compaction_respects_anti_thrash_breaker(tmp_path: Path) -> None: """A tripped ineffective-compression breaker must block the idle trigger. diff --git a/tests/agent/test_model_metadata.py b/tests/agent/test_model_metadata.py index 3d0ccd4101..f2f954675b 100644 --- a/tests/agent/test_model_metadata.py +++ b/tests/agent/test_model_metadata.py @@ -303,6 +303,21 @@ class TestDefaultContextLengths: model, provider="kimi-coding", base_url=base_url ) == 1_048_576 + @pytest.mark.parametrize("model, provider, base_url", [ + ("muse-spark-1.3-contributor-free", "opencode-free", "https://opencode.ai/zen/v1"), + ("muse-spark-1.3-contributor", "opencode-go", "https://opencode.ai/zen/go/v1"), + ("muse-spark-1.3", "meta-ai", "https://api.meta.ai/v1"), + ("meta/muse-spark-1.3", "commandcode", "https://api.commandcode.ai/provider/v1"), + ]) + def test_muse_spark_resolves_1m_without_network(self, model, provider, base_url): + """Muse Spark is 1,048,576 on every host even when models.dev and the + live /models probe are unavailable (fresh HERMES_HOME, offline).""" + with patch("agent.model_metadata.get_cached_context_length", return_value=None), \ + patch("agent.model_metadata._query_ollama_api_show", return_value=None), \ + patch("agent.model_metadata.fetch_endpoint_model_metadata", return_value={}), \ + patch("agent.models_dev.fetch_models_dev", return_value={}): + assert get_model_context_length(model, provider=provider, base_url=base_url) == 1_048_576 + def test_empty_model_uses_fallback_context(self): assert get_model_context_length("") == DEFAULT_FALLBACK_CONTEXT assert get_model_context_length(None) == DEFAULT_FALLBACK_CONTEXT # type: ignore[arg-type] @@ -1631,6 +1646,18 @@ class TestGrok43StaleCacheGuard: assert ctx == 256_000, f"{slug} should stay 256000, got {ctx}" +class TestMuseSparkStaleCacheGuard: + """Muse Spark (1M window per OpenRouter live metadata) had no catalog + entry, so older builds persisted the 256K default fallback. The cache + guard must flag that stale value and keep correct/probed values.""" + + def test_stale_muse_spark_detected_by_generic_guard(self): + from agent.model_metadata import _stale_pre_catalog_cache_entry + for slug in ("muse-spark-1.3", "meta/muse-spark-1.3-contributor", "muse-spark-1.2-contributor"): + assert _stale_pre_catalog_cache_entry(slug, 256_000), slug + assert not _stale_pre_catalog_cache_entry(slug, 1_048_576), slug + + class TestGrok46StaleCacheGuard: """Pre-catalog builds resolved grok-4.6 via the generic 'grok-4' catch-all (256,000) and persisted it before the 500K catalog entry existed. diff --git a/tests/agent/test_models_dev_meta_mapping.py b/tests/agent/test_models_dev_meta_mapping.py index 01dc09256e..a4c110cf5a 100644 --- a/tests/agent/test_models_dev_meta_mapping.py +++ b/tests/agent/test_models_dev_meta_mapping.py @@ -1,4 +1,4 @@ -"""Meta Model API maps to the models.dev 'meta' provider id (context/pricing).""" +"""Muse Spark hosts map to the right models.dev provider id (context/pricing).""" from agent.models_dev import PROVIDER_TO_MODELS_DEV @@ -6,3 +6,8 @@ from agent.models_dev import PROVIDER_TO_MODELS_DEV def test_meta_ai_maps_to_meta(): assert PROVIDER_TO_MODELS_DEV.get("meta-ai") == "meta" assert PROVIDER_TO_MODELS_DEV.get("meta") == "meta" + + +def test_opencode_free_maps_to_zen_catalog(): + # The free tier is served by the Zen relay, whose models.dev id is "opencode". + assert PROVIDER_TO_MODELS_DEV.get("opencode-free") == "opencode" diff --git a/tests/agent/test_opencode_session_affinity.py b/tests/agent/test_opencode_session_affinity.py new file mode 100644 index 0000000000..b2b6cf2536 --- /dev/null +++ b/tests/agent/test_opencode_session_affinity.py @@ -0,0 +1,62 @@ +"""x-opencode-session rides on every OpenCode request, on every transport.""" + +from __future__ import annotations + +import pytest + +from agent import auxiliary_client as aux +from agent.chat_completion_helpers import build_api_kwargs +from run_agent import AIAgent + +_MSGS = [{"role": "user", "content": "hi"}] + + +def _agent(provider, model, base_url, api_mode=None): + agent = AIAgent( + api_key="test-key", + base_url=base_url, + model=model, + provider=provider, + quiet_mode=True, + skip_context_files=True, + skip_memory=True, + session_id="sess-affinity-1", + ) + if api_mode: + agent.api_mode = api_mode + agent._transport = None + agent._anthropic_base_url = base_url + return agent + + +@pytest.mark.parametrize( + "provider, model, base_url, api_mode", + [ + ("opencode-go", "glm-5", "https://opencode.ai/zen/go/v1", None), # chat_completions + ("opencode-go", "gpt-5.6-luna", "https://opencode.ai/zen/go/v1", None), # codex_responses + ("opencode-go", "minimax-m2.7", "https://opencode.ai/zen/go/v1", "anthropic_messages"), + ("opencode-free", "laguna-s-2.1-free", "https://opencode.ai/zen/v1", None), + ("custom", "glm-5", "https://opencode.ai/zen/go/v1", None), # URL-only detection + ], +) +def test_main_turn_sends_stable_session_header_on_every_transport(provider, model, base_url, api_mode): + agent = _agent(provider, model, base_url, api_mode) + first = build_api_kwargs(agent, _MSGS)["extra_headers"]["x-opencode-session"] + second = build_api_kwargs(agent, _MSGS)["extra_headers"]["x-opencode-session"] + assert first == second == "sess-affinity-1" + + other = _agent("openrouter", "anthropic/claude-sonnet-4.6", "https://openrouter.ai/api/v1") + assert "x-opencode-session" not in (build_api_kwargs(other, _MSGS).get("extra_headers") or {}) + + +def test_auxiliary_calls_share_the_main_turn_session_key(): + token = aux.set_runtime_main( + "opencode-go", "glm-5", base_url="https://opencode.ai/zen/go/v1", session_id="sess-affinity-1" + ) + try: + kwargs = aux._build_call_kwargs("opencode-go", "glm-5", _MSGS, base_url="https://opencode.ai/zen/go/v1") + assert kwargs["extra_headers"]["x-opencode-session"] == "sess-affinity-1" + other = aux._build_call_kwargs("openrouter", "x", _MSGS, base_url="https://openrouter.ai/api/v1") + assert "x-opencode-session" not in (other.get("extra_headers") or {}) + finally: + aux._RUNTIME_MAIN_CONTEXT.reset(token) diff --git a/tests/agent/test_outbound_stale_vision.py b/tests/agent/test_outbound_stale_vision.py new file mode 100644 index 0000000000..1cab4234a0 --- /dev/null +++ b/tests/agent/test_outbound_stale_vision.py @@ -0,0 +1,119 @@ +"""Send-path eviction of stale vision_analyze / screenshot tool payloads. + +Issue #89296: compression only retires older image-bearing tool results when +prune/compress fires, so OpenAI-style screenshots are re-serialized on every +later turn until a 413. ``evict_stale_outbound_tool_images`` is the +unconditional per-call chokepoint. +""" + +from __future__ import annotations + +from agent.agent_runtime_helpers import sanitize_api_messages +from agent.context_compressor import ( + _MAX_KEEP_TOOL_IMAGES, + _tool_content_has_images, + evict_stale_outbound_tool_images, +) + + +def _image_tool(i: int, *, blob: str = "A" * 80) -> list[dict]: + return [ + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": f"call_{i}", + "type": "function", + "function": { + "name": "vision_analyze", + "arguments": f'{{"image_url":"shot{i}.png"}}', + }, + } + ], + }, + { + "role": "tool", + "tool_call_id": f"call_{i}", + "content": [ + {"type": "text", "text": f"Image attached natively shot {i}"}, + { + "type": "image_url", + "image_url": {"url": f"data:image/png;base64,{blob}{i}"}, + }, + ], + }, + ] + + +def _history_with_screenshots(n: int) -> list[dict]: + msgs: list[dict] = [{"role": "user", "content": "look at these"}] + for i in range(n): + msgs.extend(_image_tool(i)) + msgs.append({"role": "user", "content": "compare them"}) + return msgs + + +def _image_bearing_tool_ids(messages: list[dict]) -> list[str]: + return [ + m["tool_call_id"] + for m in messages + if m.get("role") == "tool" and _tool_content_has_images(m.get("content")) + ] + + +class TestOutboundStaleVisionEviction: + def test_sanitize_alone_keeps_every_screenshot(self): + """The previous send chokepoint does not close #89296 by itself.""" + history = _history_with_screenshots(5) + sanitized = sanitize_api_messages(history) + assert _image_bearing_tool_ids(sanitized) == [f"call_{i}" for i in range(5)] + + def test_eviction_keeps_only_newest_window(self): + history = _history_with_screenshots(5) + outbound = sanitize_api_messages(history) + pruned = evict_stale_outbound_tool_images(outbound) + assert pruned == 5 - _MAX_KEEP_TOOL_IMAGES + kept = _image_bearing_tool_ids(outbound) + assert kept == [f"call_{i}" for i in range(5 - _MAX_KEEP_TOOL_IMAGES, 5)] + + oldest = next(m for m in outbound if m.get("tool_call_id") == "call_0") + assert isinstance(oldest["content"], list) + assert not _tool_content_has_images(oldest["content"]) + assert any( + isinstance(part, dict) + and part.get("type") == "text" + and "screenshot removed" in str(part.get("text", "")) + for part in oldest["content"] + ) + + def test_does_not_rewrite_persisted_history(self): + from agent.conversation_loop import _clone_message_for_send + + history = _history_with_screenshots(5) + outbound = [_clone_message_for_send(m) for m in history] + evict_stale_outbound_tool_images(outbound) + assert _image_bearing_tool_ids(history) == [f"call_{i}" for i in range(5)] + assert _image_bearing_tool_ids(outbound) == [ + f"call_{i}" for i in range(5 - _MAX_KEEP_TOOL_IMAGES, 5) + ] + + def test_user_uploads_are_not_evicted(self): + history = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "look"}, + { + "type": "image_url", + "image_url": {"url": "data:image/png;base64,USERUPLOAD"}, + }, + ], + } + ] + for i in range(_MAX_KEEP_TOOL_IMAGES + 2): + history.extend(_image_tool(i)) + outbound = sanitize_api_messages(history) + evict_stale_outbound_tool_images(outbound) + user = next(m for m in outbound if m.get("role") == "user") + assert user["content"][1]["image_url"]["url"].endswith("USERUPLOAD") diff --git a/tests/agent/test_periodic_scheduler.py b/tests/agent/test_periodic_scheduler.py new file mode 100644 index 0000000000..e1a2b3765a --- /dev/null +++ b/tests/agent/test_periodic_scheduler.py @@ -0,0 +1,92 @@ +"""agent/periodic_scheduler: one shared thread runs every periodic timer.""" + +import threading +import time + +from agent import periodic_scheduler +from agent.periodic_scheduler import PeriodicScheduler, schedule + + +def _wait_until(pred, timeout=3.0): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if pred(): + return True + time.sleep(0.005) + return pred() + + +def test_two_intervals_fire_proportionally_and_cancel_stops_one(): + sched = PeriodicScheduler() + fast, slow = [], [] + h_fast = sched.schedule(lambda: fast.append(time.monotonic()), 0.01) + h_slow = sched.schedule(lambda: slow.append(time.monotonic()), 0.05) + + assert _wait_until(lambda: len(slow) >= 3) + assert len(fast) > len(slow) # 5x interval ratio -> clearly more fast ticks + # Both ran on this scheduler's single thread, not on new threads. + before = threading.active_count() + sched.schedule(lambda: None, 0.01).cancel() + assert threading.active_count() == before + assert sched._thread is not None and sched._thread.is_alive() + + h_fast.cancel() + n_fast = len(fast) + time.sleep(0.1) + assert len(fast) == n_fast, "cancelled callback kept firing" + assert len(slow) > 3, "sibling callback stopped when another was cancelled" + h_slow.cancel() + + +def test_raising_callback_is_rescheduled_and_does_not_kill_sibling(): + sched = PeriodicScheduler() + boom, ok = [], [] + + def raises(): + boom.append(1) + raise RuntimeError("bad callback") + + h1 = sched.schedule(raises, 0.01) + h2 = sched.schedule(lambda: ok.append(1), 0.01) + assert _wait_until(lambda: len(boom) >= 3 and len(ok) >= 3) + h1.cancel() + h2.cancel() + + +def test_returning_false_stops_callback_and_cancel_wait_joins_inflight(): + sched = PeriodicScheduler() + calls = [] + sched.schedule(lambda: (calls.append(1), False)[1], 0.01) + assert _wait_until(lambda: len(calls) == 1) + time.sleep(0.05) + assert calls == [1] + + entered = threading.Event() + release = threading.Event() + + def blocking(): + entered.set() + release.wait(2.0) + + h = sched.schedule(blocking, 0.01) + assert entered.wait(2.0) + threading.Timer(0.05, release.set).start() + t0 = time.monotonic() + h.cancel(wait=2.0) # returns once the in-flight run finished + assert release.is_set() + assert time.monotonic() - t0 < 1.5 + + +def test_module_level_schedule_uses_shared_default(): + hits = [] + h = schedule(lambda: hits.append(1), 0.01) + assert _wait_until(lambda: hits) + h.cancel() + thread = periodic_scheduler._DEFAULT._thread + assert thread is not None and thread.name == "hermes-periodic-scheduler" + # Scheduling more timers on the shared default adds no OS threads. + before = threading.active_count() + handles = [schedule(lambda: None, 0.01) for _ in range(20)] + assert threading.active_count() == before + for handle in handles: + handle.cancel() diff --git a/tests/agent/test_proactive_prune_rearm_threshold.py b/tests/agent/test_proactive_prune_rearm_threshold.py new file mode 100644 index 0000000000..9d49eb5626 --- /dev/null +++ b/tests/agent/test_proactive_prune_rearm_threshold.py @@ -0,0 +1,227 @@ +"""Proactive-prune rearm must not lock out an over-threshold session (#101889). + +``_proactive_prune_rearm_tokens`` is armed from a message-bodies-only estimate, +but the provider bills the system prompt and tool schemas too. On a schema-heavy +session the message-only estimate can sit permanently just below the rearm mark +while the real request rides *above* ``threshold_tokens`` — the prune declines +every iteration, full compression never gets there, and nothing is logged. The +session then grows until the provider rejects the request. + +Pinned here as invariants (no frozen config literals): the gates are evaluated +against this compressor's own ``threshold_tokens`` / ``proactive_prune_tokens``. +""" + +from __future__ import annotations + +import logging +from typing import Any, Dict, List +from unittest.mock import patch + +from agent.context_compressor import ContextCompressor, _estimate_msg_budget_tokens + +LARGE_WINDOW = 1_000_000 + + +def _compressor(**kw: Any) -> ContextCompressor: + defaults = dict( + model="test", + quiet_mode=True, + threshold_percent=0.50, + protect_first_n=2, + protect_last_n=4, + proactive_prune_tokens=48_000, + proactive_prune_min_result_chars=8_000, + ) + defaults.update(kw) + with patch( + "agent.context_compressor.get_model_context_length", + return_value=LARGE_WINDOW, + ): + return ContextCompressor(**defaults) + + +def _history(n_pairs: int = 8, big: int = 9_000) -> List[Dict[str, Any]]: + msgs: List[Dict[str, Any]] = [{"role": "system", "content": "sys"}] + for i in range(n_pairs): + cid = f"call_{i}" + msgs.append({ + "role": "assistant", + "content": "", + "tool_calls": [{ + "id": cid, + "type": "function", + "function": {"name": "terminal", "arguments": '{"cmd":"ls"}'}, + }], + }) + msgs.append({ + "role": "tool", + "tool_call_id": cid, + "content": chr(65 + i) * big if i < 3 else "ok", + }) + return msgs + + +def _park_rearm_just_above_messages( + compressor: ContextCompressor, messages: List[Dict[str, Any]] +) -> int: + """Reproduce the reporter's state: message-only estimate stuck 913 tokens + below the rearm mark (schema overhead makes up the rest of the request).""" + before = sum(_estimate_msg_budget_tokens(m) for m in messages) + compressor._proactive_prune_rearm_tokens = before + 913 + assert before < compressor._proactive_prune_rearm_tokens + return before + + +def _over_threshold_warnings(caplog) -> list: + return [ + r for r in caplog.records + if r.levelno >= logging.WARNING + and "over the compression threshold" in r.getMessage() + ] + + +def test_billed_basis_over_threshold_defeats_message_only_rearm_lockout() -> None: + """Over ``threshold_tokens`` on the provider-billed basis, the rearm gate + must not short-circuit the prune on the message-only estimate alone.""" + c = _compressor() + msgs = _history() + _park_rearm_just_above_messages(c, msgs) + billed = c.threshold_tokens + 1 # provider says: over threshold, now + + scans: List[int] = [] + # Stand in for the real multi-pass scan: a NEW list whose old tool outputs + # are reclaimed, so the (untouched) reclaim gate can commit it. + reclaimed = [dict(m) for m in msgs] + for m in reclaimed[:-2]: + if m.get("role") == "tool": + m["content"] = "[pruned]" + + def _scan(*args: Any, **kwargs: Any) -> tuple[List[Dict[str, Any]], int]: + scans.append(1) + return reclaimed, 3 + + with patch.object(c, "_prune_old_tool_results", _scan): + result, pruned = c.prune_tool_results_only(msgs, current_tokens=billed) + + assert scans, "rearm gate short-circuited on the message-only estimate" + assert pruned == 3 + assert result is not msgs + + +def test_message_only_rearm_still_holds_below_threshold() -> None: + """Prompt-cache hysteresis is intact while the real request is under the + compression threshold — the rearm bypass is an overflow escape hatch only.""" + c = _compressor() + msgs = _history() + _park_rearm_just_above_messages(c, msgs) + under = c.threshold_tokens - 1 + assert under >= c.proactive_prune_tokens # above the prune trigger + + with patch.object( + c, + "_prune_old_tool_results", + side_effect=AssertionError("scan must not run below threshold"), + ): + result, pruned = c.prune_tool_results_only(msgs, current_tokens=under) + + assert result is msgs + assert pruned == 0 + + +def test_no_op_below_the_prune_trigger() -> None: + """Under ``proactive_prune_tokens`` nothing is reclaimed, rearm or not — + the bypass must not turn into over-pruning of small sessions.""" + c = _compressor() + msgs = _history() + c.on_session_reset() # fully rearmed; only the trigger gates + + with patch.object( + c, + "_prune_old_tool_results", + side_effect=AssertionError("scan must not run below the trigger"), + ): + result, pruned = c.prune_tool_results_only( + msgs, current_tokens=c.proactive_prune_tokens - 1 + ) + + assert result is msgs + assert pruned == 0 + + +def test_over_threshold_reclamation_no_op_warns_once(caplog) -> None: + """A session riding above the threshold with every reclamation path + declining must be distinguishable in the log — and must not spam the same + reason on every tool iteration.""" + # Reclaim floor above anything this transcript can free: the scan runs, + # finds candidates, and the commit gate rejects it — a silent no-op today. + c = _compressor(proactive_prune_min_reclaim_tokens=10_000_000) + msgs = _history() + billed = c.threshold_tokens + 5_000 + + with caplog.at_level(logging.WARNING, logger="agent.context_compressor"): + result, pruned = c.prune_tool_results_only(msgs, current_tokens=billed) + assert (result, pruned) == (msgs, 0) + + warnings = _over_threshold_warnings(caplog) + assert warnings, "over-threshold reclamation no-op was silent" + + # Same state on the next tool iteration: deduped, not re-logged. + with caplog.at_level(logging.WARNING, logger="agent.context_compressor"): + c.prune_tool_results_only(msgs, current_tokens=billed) + assert len(_over_threshold_warnings(caplog)) == len(warnings) + + +def test_under_threshold_no_op_is_not_warned(caplog) -> None: + """Ordinary hysteresis below the threshold stays quiet.""" + c = _compressor(proactive_prune_min_reclaim_tokens=10_000_000) + msgs = _history() + + with caplog.at_level(logging.WARNING, logger="agent.context_compressor"): + result, pruned = c.prune_tool_results_only( + msgs, current_tokens=c.threshold_tokens - 1 + ) + + assert (result, pruned) == (msgs, 0) + assert not _over_threshold_warnings(caplog) + + +def test_lockout_warns_again_after_rearm_reset(caplog) -> None: + """A full compaction (or session rebind / model recalibration) zeroes the + rearm mark. An identical lockout afterwards must warn again — the dedup key + must not outlive the reclamation that should have cleared it.""" + c = _compressor(proactive_prune_min_reclaim_tokens=10_000_000) + msgs = _history() + billed = c.threshold_tokens + 5_000 + + with caplog.at_level(logging.WARNING, logger="agent.context_compressor"): + c.prune_tool_results_only(msgs, current_tokens=billed) + assert len(_over_threshold_warnings(caplog)) == 1 + # Same state, same key (reason, rearm=0): deduped. + with caplog.at_level(logging.WARNING, logger="agent.context_compressor"): + c.prune_tool_results_only(msgs, current_tokens=billed) + assert len(_over_threshold_warnings(caplog)) == 1 + + # A public rearm boundary (same helper as compress(), on_session_end, + # bind_session_state and update_model): pins the wiring, not just the body. + c.on_session_reset() + assert c._proactive_prune_rearm_tokens == 0 + + with caplog.at_level(logging.WARNING, logger="agent.context_compressor"): + c.prune_tool_results_only(msgs, current_tokens=billed) + assert len(_over_threshold_warnings(caplog)) == 2, ( + "lockout after a rearm reset was deduped against the stale key" + ) + + +def test_dropping_under_threshold_clears_dedup_key(caplog) -> None: + """Back under threshold (e.g. compaction elsewhere shrank the request), the + key is released so the next over-threshold lockout is reported.""" + c = _compressor(proactive_prune_min_reclaim_tokens=10_000_000) + msgs = _history() + billed = c.threshold_tokens + 5_000 + + with caplog.at_level(logging.WARNING, logger="agent.context_compressor"): + c.prune_tool_results_only(msgs, current_tokens=billed) + c.prune_tool_results_only(msgs, current_tokens=c.threshold_tokens - 1) + c.prune_tool_results_only(msgs, current_tokens=billed) + assert len(_over_threshold_warnings(caplog)) == 2 diff --git a/tests/agent/test_proactive_tool_result_pruning.py b/tests/agent/test_proactive_tool_result_pruning.py index bbb01e4b16..2a963f7f01 100644 --- a/tests/agent/test_proactive_tool_result_pruning.py +++ b/tests/agent/test_proactive_tool_result_pruning.py @@ -130,7 +130,12 @@ def test_rearms_only_after_reclaimed_token_runway(): _tool_msg("call_9", "ok"), ] assert sum(map(_estimate_msg_budget_tokens, grown)) < rearm_tokens - blocked, n2 = c.prune_tool_results_only(grown, current_tokens=1_000_000) + # Below the full-compression threshold, where the runway is pure + # prompt-cache hysteresis. (Above it the runway is bypassed on the + # provider-billed reading instead — see + # tests/agent/test_proactive_prune_rearm_threshold.py, #101889.) + _under_threshold = c.threshold_tokens - 1 + blocked, n2 = c.prune_tool_results_only(grown, current_tokens=_under_threshold) assert n2 == 0 assert blocked is grown assert len(_tool_by_id(blocked, "call_6")["content"]) == 9000 @@ -139,7 +144,7 @@ def test_rearms_only_after_reclaimed_token_runway(): missing = rearm_tokens - sum(map(_estimate_msg_budget_tokens, grown)) regrown = grown + [{"role": "user", "content": "x" * (missing * 4)}] assert sum(map(_estimate_msg_budget_tokens, regrown)) >= rearm_tokens - rearmed, n3 = c.prune_tool_results_only(regrown, current_tokens=1_000_000) + rearmed, n3 = c.prune_tool_results_only(regrown, current_tokens=_under_threshold) assert n3 >= 2 assert rearmed is not regrown diff --git a/tests/agent/test_prompt_builder.py b/tests/agent/test_prompt_builder.py index 8dafc52438..910260fff5 100644 --- a/tests/agent/test_prompt_builder.py +++ b/tests/agent/test_prompt_builder.py @@ -5,6 +5,8 @@ import importlib import logging import os import sys +import time +from pathlib import Path import pytest @@ -820,6 +822,126 @@ class TestEnvironmentHints: assert "Linux 6.8.0" in line assert "root" in line + def test_probe_remote_backend_tears_down_its_sandbox(self, monkeypatch): + """THE BUG: the probe leaked a second, permanently idle sandbox. + + ``_probe_remote_backend`` spins up an environment with + ``task_id="prompt-backend-probe"`` purely to run one ``uname``. Container + backends default to ``container_persistent`` / + ``docker_persist_across_processes``, so that throwaway sandbox stayed up + for the whole process lifetime *next to* the agent's own ``default`` + sandbox — one wasted idle container per profile, forever. The probe owns + that environment, so it must tear it down. + """ + import agent.prompt_builder as _pb + + monkeypatch.setenv("TERMINAL_ENV", "docker") + _pb._clear_backend_probe_cache() + + cleaned = {} + + class _FakeEnv: + def execute(self, cmd, timeout=None): + return { + "returncode": 0, + "output": ( + "os=Linux\nkernel=6.8.0\nhome=/root\n" + "cwd=/workspace\nuser=root\n" + ), + } + + def cleanup(self, *, force_remove=False): + cleaned["force_remove"] = force_remove + + import tools.terminal_tool as _tt + monkeypatch.setattr(_tt, "_create_environment", lambda **kw: _FakeEnv()) + + assert _pb._probe_remote_backend("docker") is not None + # force_remove=True: persist mode would otherwise leave it running. + assert cleaned == {"force_remove": True} + + def test_probe_remote_backend_tears_down_sandbox_on_failure(self, monkeypatch): + """Teardown must also run when the probe command blows up — a flaky + backend would otherwise leak the container the probe just created.""" + import agent.prompt_builder as _pb + + monkeypatch.setenv("TERMINAL_ENV", "docker") + _pb._clear_backend_probe_cache() + + cleaned = [] + + class _ExplodingEnv: + def execute(self, cmd, timeout=None): + raise RuntimeError("backend went away") + + def cleanup(self, *, force_remove=False): + cleaned.append(force_remove) + + import tools.terminal_tool as _tt + monkeypatch.setattr(_tt, "_create_environment", lambda **kw: _ExplodingEnv()) + + assert _pb._probe_remote_backend("docker") is None + assert cleaned == [True] + + def test_probe_remote_backend_tolerates_kwargless_cleanup(self, monkeypatch): + """Backends that inherit the base ``cleanup(self)`` take no kwargs; the + probe must use the bare call instead of dying on TypeError.""" + import agent.prompt_builder as _pb + + monkeypatch.setenv("TERMINAL_ENV", "singularity") + _pb._clear_backend_probe_cache() + + calls = [] + + class _LegacyEnv: + def execute(self, cmd, timeout=None): + return { + "returncode": 0, + "output": ( + "os=Linux\nkernel=6.8.0\nhome=/home/u\n" + "cwd=/home/u\nuser=u\n" + ), + } + + def cleanup(self): + calls.append("bare") + + import tools.terminal_tool as _tt + monkeypatch.setattr(_tt, "_create_environment", lambda **kw: _LegacyEnv()) + + assert _pb._probe_remote_backend("singularity") is not None + assert calls == ["bare"] + + def test_probe_remote_backend_does_not_tear_down_ssh(self, monkeypatch): + """SSH has no task-scoped sandbox: its cleanup() closes a ControlMaster + socket shared with the agent's real environment, so the probe must + leave it alone (nothing leaks — ControlPersist expires the master).""" + import agent.prompt_builder as _pb + + monkeypatch.setenv("TERMINAL_ENV", "ssh") + _pb._clear_backend_probe_cache() + + calls = [] + + class _SharedSshEnv: + def execute(self, cmd, timeout=None): + return { + "returncode": 0, + "output": ( + "os=Linux\nkernel=6.8.0\nhome=/home/u\n" + "cwd=/home/u\nuser=u\n" + ), + } + + def cleanup(self): + calls.append("cleanup") + + import tools.terminal_tool as _tt + monkeypatch.setattr(_tt, "_create_environment", lambda **kw: _SharedSshEnv()) + + assert _pb._probe_remote_backend("ssh") is not None + assert calls == [] + def test_environment_hint_from_env_var_is_appended(self, monkeypatch): """HERMES_ENVIRONMENT_HINT lets an embedder describe the runtime env.""" @@ -986,6 +1108,13 @@ class TestExecutionGuidanceModels: for fam in ("deepseek", "kimi", "qwen", "glm", "minimax", "mimo", "mistral"): assert fam in EXECUTION_GUIDANCE_MODELS + def test_muse_spark_gets_both_guidance_blocks(self): + # Muse Spark closes the turn after a chat-only response on defaults + # (#96550) — it needs tool-use enforcement AND execution guidance. + from agent.prompt_builder import EXECUTION_GUIDANCE_MODELS + assert any(p in "meta/muse-spark-1.3-contributor" for p in TOOL_USE_ENFORCEMENT_MODELS) + assert any(p in "meta/muse-spark-1.3-contributor" for p in EXECUTION_GUIDANCE_MODELS) + def test_excludes_google_and_claude(self): # Gemini/Gemma get GOOGLE_MODEL_OPERATIONAL_GUIDANCE instead; # Claude doesn't exhibit the targeted failure modes. @@ -1020,3 +1149,40 @@ class TestParallelToolCallGuidance: # ========================================================================= + + +class TestContextFileReadTimeout: + def test_slow_hermes_md_is_skipped_and_agents_md_still_loads(self, tmp_path, monkeypatch, caplog): + (tmp_path / ".git").mkdir() + (tmp_path / ".hermes.md").write_text("Hermes project rules.") + (tmp_path / "AGENTS.md").write_text("Agent fallback rules.") + # Patch the module object build_context_files_prompt actually closes + # over: an earlier test re-imports agent.prompt_builder, so the + # sys.modules entry can be a different module object. + pb_mod = sys.modules[build_context_files_prompt.__module__] + monkeypatch.setattr(pb_mod, "_get_context_file_read_timeout", lambda: 0.05) + + original_read_text = Path.read_text + + def slow_read_text(self, *args, **kwargs): + if self.name == ".hermes.md": + time.sleep(0.6) + return original_read_text(self, *args, **kwargs) + + monkeypatch.setattr(Path, "read_text", slow_read_text) + + start = time.monotonic() + with caplog.at_level(logging.WARNING, logger=pb_mod.__name__): + result = build_context_files_prompt(cwd=str(tmp_path)) + elapsed = time.monotonic() - start + + assert elapsed < 0.4, f"context load blocked for {elapsed:.2f}s" + assert "Agent fallback rules" in result + assert "Hermes project rules" not in result + assert "timed out" in caplog.text.lower() + + def test_read_errors_still_propagate_to_caller(self, tmp_path): + from agent.prompt_builder import _read_text_with_timeout + + with pytest.raises(FileNotFoundError): + _read_text_with_timeout(tmp_path / "missing.md", timeout=1.0) diff --git a/tests/agent/test_relay_llm.py b/tests/agent/test_relay_llm.py index f091fb1222..a31553af88 100644 --- a/tests/agent/test_relay_llm.py +++ b/tests/agent/test_relay_llm.py @@ -259,6 +259,141 @@ def test_provider_request_overlays_interceptor_added_extra_body(): assert provider_request["extra_body"] == {"prompt_cache_retention": "24h"} +@pytest.mark.parametrize( + "api_mode", + ["chat_completions", "codex_responses", "anthropic_messages"], +) +def test_provider_request_maps_headers_for_supported_sdk_modes(api_mode): + original = {"model": "test-model"} + relay_request_body = relay_llm._relay_request_body( + original, + {"api_mode": api_mode}, + ) + + provider_request = relay_llm._provider_request( + original, + SimpleNamespace( + content=relay_request_body, + headers={ + "traceparent": ( + "00-11111111111111111111111111111111-" + "2222222222222222-01" + ) + }, + ), + relay_request_body=relay_request_body, + codec_baseline_body=dict(relay_request_body), + metadata={"api_mode": api_mode}, + ) + + assert provider_request["extra_headers"] == { + "traceparent": ( + "00-11111111111111111111111111111111-2222222222222222-01" + ) + } + + +def test_provider_request_preserves_custom_headers_for_native_transport(): + original = {"payload": "provider-native"} + + provider_request = relay_llm._provider_request( + original, + SimpleNamespace( + content=original, + headers={ + "traceparent": ( + "00-11111111111111111111111111111111-" + "2222222222222222-01" + ), + "x-custom-route": "private", + }, + ), + relay_request_body=original, + codec_baseline_body=dict(original), + metadata={"api_mode": "strict_native"}, + ) + + assert provider_request["extra_headers"] == { + "x-custom-route": "private" + } + + +def test_provider_request_traces_custom_transport_with_header_capability(): + original = { + "payload": "provider-native", + "extra_headers": {"authorization": "Bearer provider-token"}, + } + traceparent = ( + "00-11111111111111111111111111111111-2222222222222222-01" + ) + + provider_request = relay_llm._provider_request( + original, + SimpleNamespace( + content=original, + headers={"traceparent": traceparent}, + ), + relay_request_body=original, + codec_baseline_body=dict(original), + metadata={"api_mode": "custom"}, + ) + + assert provider_request["extra_headers"] == { + "authorization": "Bearer provider-token", + "traceparent": traceparent, + } + + +def test_managed_request_does_not_add_sdk_headers_to_strict_callback(relay_turn): + del relay_turn + observed = [] + + def strict_transport(*, payload): + observed.append(payload) + return {"content": payload} + + result = relay_llm.execute( + {"payload": "provider-native"}, + lambda request: strict_transport(**request), + session_id="session-1", + name="strict-native", + model_name="strict-model", + metadata={ + "api_mode": "bedrock_converse", + "api_request_id": "strict-native-request", + }, + ) + + assert observed == ["provider-native"] + assert result == {"content": "provider-native"} + + +def test_managed_stream_does_not_add_sdk_headers_to_strict_callback(relay_turn): + del relay_turn + observed = [] + chunks = [{"delta": "provider-native"}] + + def strict_transport(*, payload): + observed.append(payload) + return iter(chunks) + + stream = relay_llm.stream( + {"payload": "provider-native"}, + lambda request: strict_transport(**request), + session_id="session-1", + name="strict-native", + model_name="strict-model", + finalizer=lambda: {"content": "provider-native"}, + metadata={ + "api_mode": "bedrock_converse", + "api_request_id": "strict-native-stream", + }, + ) + + assert list(stream) == chunks + assert observed == ["provider-native"] + + def test_stream_uses_rewritten_request_and_post_intercept_chunks(relay_turn): relay, turn = relay_turn captured_requests = [] @@ -355,9 +490,15 @@ def test_stream_uses_rewritten_request_and_post_intercept_chunks(relay_turn): relay.intercepts.deregister_llm_request("hermes-test-request") assert captured_requests[0]["temperature"] == 0.25 - assert captured_requests[0]["extra_headers"] == { - "authorization": "Bearer provider-token" - } + headers = captured_requests[0]["extra_headers"] + assert headers["authorization"] == "Bearer provider-token" + version, trace_id, parent_id, flags = headers["traceparent"].split("-") + assert version == "00" + assert len(trace_id) == 32 + assert len(parent_id) == 16 + assert flags == "01" + int(trace_id, 16) + int(parent_id, 16) assert chunks[0].choices[0].delta.content == "HELLO" assert stream.output_modified is True assert turn.logical_llm_calls == {} diff --git a/tests/agent/test_relay_runtime_plugins.py b/tests/agent/test_relay_runtime_plugins.py index a906f94f29..214d80d24d 100644 --- a/tests/agent/test_relay_runtime_plugins.py +++ b/tests/agent/test_relay_runtime_plugins.py @@ -1049,7 +1049,7 @@ mode = "overwrite" assert not (atof_dir / "events.jsonl").exists() -def test_real_binding_layers_project_config_after_explicit_opt_in( +def test_real_binding_ignores_project_config_with_explicit_opt_in( tmp_path, monkeypatch, ): @@ -1061,7 +1061,8 @@ def test_real_binding_layers_project_config_after_explicit_opt_in( working_directory = project_root / "workspace" config_directory = project_root / ".nemo-relay" selected_directory = tmp_path / "selected-config" - atof_dir = tmp_path / "atof" + project_atof_dir = tmp_path / "project-atof" + selected_atof_dir = tmp_path / "selected-atof" working_directory.mkdir(parents=True) config_directory.mkdir() selected_directory.mkdir() @@ -1081,14 +1082,35 @@ enabled = true [[components.config.atof.sinks]] type = "file" -output_directory = "{atof_dir}" +output_directory = "{project_atof_dir}" filename = "events.jsonl" mode = "overwrite" """.strip(), encoding="utf-8", ) selected_config = selected_directory / "plugins.toml" - selected_config.write_text("version = 1", encoding="utf-8") + selected_config.write_text( + f""" +version = 1 + +[[components]] +kind = "observability" +enabled = true + +[components.config] +version = 4 + +[components.config.atof] +enabled = true + +[[components.config.atof.sinks]] +type = "file" +output_directory = "{selected_atof_dir}" +filename = "events.jsonl" +mode = "overwrite" +""".strip(), + encoding="utf-8", + ) xdg_config_home = tmp_path / "xdg" xdg_config_home.mkdir() monkeypatch.chdir(working_directory) @@ -1102,12 +1124,13 @@ mode = "overwrite" host = relay_runtime.RelayRuntime(relay=relay, profile_key="profile") try: assert host.managed_execution_enabled() - host.ensure_session({"session_id": "native-layered-plugins"}) + host.ensure_session({"session_id": "native-explicit-plugins"}) finally: host.shutdown() relay_runtime._reset_for_tests() - assert (atof_dir / "events.jsonl").is_file() + assert (selected_atof_dir / "events.jsonl").is_file() + assert not (project_atof_dir / "events.jsonl").exists() def test_real_binding_keeps_two_profile_trajectories_separate_in_shared_exporters( diff --git a/tests/agent/test_relay_tools.py b/tests/agent/test_relay_tools.py index 22dec2fe52..760aa7335b 100644 --- a/tests/agent/test_relay_tools.py +++ b/tests/agent/test_relay_tools.py @@ -78,7 +78,10 @@ def test_request_rewrite_reaches_authorized_callback_once(relay_turn): async def wrap_execution(_name, args, next_call): result = await next_call(args) - return relay.ToolExecutionInterceptOutcome({**result, "wrapped": True}) + return relay.ToolExecutionInterceptOutcome( + {**result.result, "wrapped": True}, + annotation={"audit": "annotation-canary"}, + ) relay.intercepts.register_tool_request( "hermes-test-tool-request", 1, False, rewrite_request @@ -102,10 +105,33 @@ def test_request_rewrite_reaches_authorized_callback_once(relay_turn): assert observed_args == {"path": "/approved/path"} assert isinstance(result, str) assert json.loads(result) == {"ok": True, "wrapped": True} + assert "annotation-canary" not in result +def test_tool_call_id_uses_canonical_relay_argument(relay_turn, monkeypatch): + relay = relay_turn + captured = {} + async def capture_execute(_name, args, callback, **kwargs): + captured.update(kwargs) + result = callback(args) + assert isinstance(result, relay.ToolExecutionResult) + return result + monkeypatch.setattr(relay.tools, "execute", capture_execute) + original_result = {"ok": True} + + result, observed_args = relay_tools.execute( + "write_file", + {"path": "/tmp/output"}, + lambda _args: original_result, + session_id="session-1", + tool_call_id="call-42", + ) + + assert result is original_result + assert observed_args == {"path": "/tmp/output"} + assert captured["tool_call_id"] == "call-42" def test_tool_error_is_preserved_from_relay_wrapper_suffix(relay_turn, monkeypatch): @@ -135,8 +161,3 @@ def test_tool_error_is_preserved_from_relay_wrapper_suffix(relay_turn, monkeypat ) assert caught.value is tool_error - - - - - diff --git a/tests/agent/test_shared_http_transport.py b/tests/agent/test_shared_http_transport.py new file mode 100644 index 0000000000..d0c541723a --- /dev/null +++ b/tests/agent/test_shared_http_transport.py @@ -0,0 +1,185 @@ +"""Keepalive httpx clients share one HTTPTransport per (verify, proxy) identity. + +Every AIAgent (and every delegated child) gets its own ``httpx.Client`` — the +#10933 contract that closing one client must never poison the next. What is +shared underneath is the connection pool + SSL context, so a fan-out of N +children no longer holds N TLS socket sets to the same provider. +""" + +import ssl +import threading +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer + +import certifi +import httpx +import pytest + +from agent import process_bootstrap +from agent.agent_runtime_helpers import _iter_pool_sockets, force_close_tcp_sockets +from agent.process_bootstrap import build_keepalive_http_client + + +@pytest.fixture +def no_proxy_env(monkeypatch): + for name in ( + "HTTPS_PROXY", "HTTP_PROXY", "ALL_PROXY", + "https_proxy", "http_proxy", "all_proxy", "NO_PROXY", "no_proxy", + ): + monkeypatch.delenv(name, raising=False) + process_bootstrap.close_shared_transports() + yield + process_bootstrap.close_shared_transports() + + +class _Handler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" # keep-alive so pooled connections persist + + def do_GET(self): # noqa: N802 + body = b"ok" + self.send_response(200) + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def log_message(self, *_args): + pass + + +@pytest.fixture +def local_server(): + server = ThreadingHTTPServer(("127.0.0.1", 0), _Handler) + server.daemon_threads = True + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + yield f"http://127.0.0.1:{server.server_address[1]}" + finally: + server.shutdown() + server.server_close() + + +def _inner(client, scheme="https://"): + mount = next(t for pat, t in client._mounts.items() if str(pat.pattern) == scheme) + return mount._inner + + +def test_same_identity_clients_share_transport_but_not_client(no_proxy_env): + a = build_keepalive_http_client("https://api.example.com/v1") + b = build_keepalive_http_client("https://api.example.com/v1") + assert isinstance(a, httpx.Client) and isinstance(b, httpx.Client) + assert a is not b + assert _inner(a) is _inner(b) + assert _inner(a, "http://") is _inner(b, "http://") + # The per-client view is distinct, so each client has its own close state. + assert a._mounts is not b._mounts + a.close() + b.close() + + +def test_closing_one_client_leaves_sibling_functional(no_proxy_env, local_server): + a = build_keepalive_http_client(local_server) + b = build_keepalive_http_client(local_server) + assert _inner(a, "http://") is _inner(b, "http://") + assert a.get(local_server + "/x").status_code == 200 + a.close() + assert a.is_closed + # #10933 shape: the shared pool must still serve the surviving client and + # any successor client built after the close. + assert b.get(local_server + "/y").status_code == 200 + c = build_keepalive_http_client(local_server) + assert _inner(c, "http://") is _inner(b, "http://") + assert c.get(local_server + "/z").status_code == 200 + with pytest.raises(RuntimeError): + a.get(local_server + "/closed") + b.close() + c.close() + + +def test_pool_survives_client_close(no_proxy_env, local_server): + a = build_keepalive_http_client(local_server) + a.get(local_server + "/warm") + pool = _inner(a, "http://")._pool + before = len(pool.connections) + assert before >= 1 + a.close() + assert len(pool.connections) == before, "client close must not drain the shared pool" + + +def test_different_verify_or_proxy_get_different_transports(no_proxy_env, monkeypatch): + default = build_keepalive_http_client("https://api.example.com/v1") + insecure = build_keepalive_http_client("https://api.example.com/v1", verify=False) + ctx = ssl.create_default_context(cafile=certifi.where()) + with_ctx = build_keepalive_http_client("https://api.example.com/v1", verify=ctx) + with_ctx2 = build_keepalive_http_client("https://api.example.com/v1", verify=ctx) + codex = build_keepalive_http_client("https://chatgpt.com/backend-api/codex") + assert _inner(default) is not _inner(insecure) + assert _inner(default) is not _inner(with_ctx) + assert _inner(with_ctx) is _inner(with_ctx2) + assert _inner(with_ctx)._pool._ssl_context is ctx + assert _inner(insecure)._pool._ssl_context.check_hostname is False + # Codex cloud gets the happy-eyeballs backend, so it can't share a pool. + assert _inner(codex) is not _inner(default) + assert isinstance( + _inner(codex)._pool._network_backend, process_bootstrap._HappyEyeballsSyncBackend + ) + for c in (default, insecure, with_ctx, with_ctx2, codex): + c.close() + + monkeypatch.setenv("HTTPS_PROXY", "http://127.0.0.1:3128") + proxied = build_keepalive_http_client("https://api.example.com/v1") + # Proxy clients keep httpx's own per-client proxy transport (unshared). + assert all( + type(t).__name__ != "_SharedTransport" for t in proxied._mounts.values() if t + ) + proxied.close() + + +def test_async_clients_are_not_shared(no_proxy_env): + a = build_keepalive_http_client("https://api.example.com/v1", async_mode=True) + b = build_keepalive_http_client("https://api.example.com/v1", async_mode=True) + assert isinstance(a, httpx.AsyncClient) + ta = [t for t in a._mounts.values() if t is not None] + tb = [t for t in b._mounts.values() if t is not None] + assert all(isinstance(t, httpx.AsyncHTTPTransport) for t in ta + tb) + assert not {id(t) for t in ta} & {id(t) for t in tb} + + +def test_force_close_only_touches_owning_clients_inflight_sockets(no_proxy_env, local_server): + """A stranger-thread abort on client A must not shut down client B's + idle/in-flight connections that live on the same shared pool.""" + a = build_keepalive_http_client(local_server) + b = build_keepalive_http_client(local_server) + b.get(local_server + "/warm") # idle keepalive connection on the shared pool + pool = _inner(a, "http://")._pool + assert pool.connections + # A has nothing in flight: nothing of A's may be touched. + assert list(_iter_pool_sockets(a)) == [] + assert force_close_tcp_sockets(a) == 0 + # B's idle connection is still healthy. + assert b.get(local_server + "/again").status_code == 200 + + # Now hold a B stream open and confirm A's abort still sees zero sockets + # while B's abort sees exactly its own. + with b.stream("GET", local_server + "/stream") as resp: + assert resp.status_code == 200 + assert list(_iter_pool_sockets(a)) == [] + assert len(list(_iter_pool_sockets(b))) == 1 + a.close() + b.close() + + +def test_shared_transport_cache_is_bounded(no_proxy_env, monkeypatch): + monkeypatch.setattr(process_bootstrap, "_SHARED_TRANSPORTS_MAX", 2) + clients = [ + build_keepalive_http_client("https://api.example.com/v1", verify=False), + build_keepalive_http_client("https://api.example.com/v1"), + ] + assert len(process_bootstrap._SHARED_TRANSPORTS) == 2 + ctx = ssl.create_default_context() + extra = build_keepalive_http_client("https://api.example.com/v1", verify=ctx) + assert len(process_bootstrap._SHARED_TRANSPORTS) == 2 + # Past the cap the caller still gets a working (private) transport. + assert _inner(extra)._pool._ssl_context is ctx + for c in clients + [extra]: + c.close() + assert process_bootstrap.close_shared_transports() == 2 diff --git a/tests/agent/test_subdirectory_hints.py b/tests/agent/test_subdirectory_hints.py index 85b89f647e..8972027adc 100644 --- a/tests/agent/test_subdirectory_hints.py +++ b/tests/agent/test_subdirectory_hints.py @@ -1,9 +1,12 @@ """Tests for progressive subdirectory hint discovery.""" +import time + import pytest from pathlib import Path from unittest.mock import patch +from agent.search_policy import SEARCH_PRUNE_DIR_NAMES from agent.subdirectory_hints import SubdirectoryHintTracker @@ -116,6 +119,40 @@ class TestSubdirectoryHintTracker: + def test_timeout_skips_slow_hint_files(self, project, monkeypatch, caplog): + """Slow hint reads time out instead of blocking the turn.""" + backend = project / "backend" + (backend / "AGENTS.md").write_text("Backend-specific instructions") + import sys + + from agent import subdirectory_hints as sh_mod + + # Patch the module object the hint tracker's helper closes over. + pb_mod = sys.modules[sh_mod._read_text_with_timeout.__module__] + monkeypatch.setattr(pb_mod, "_get_context_file_read_timeout", lambda: 0.05) + + original_read_text = Path.read_text + + def slow_read_text(self, *args, **kwargs): + if self.name.lower() == "agents.md" and self.parent == backend: + time.sleep(0.6) + return original_read_text(self, *args, **kwargs) + + monkeypatch.setattr(Path, "read_text", slow_read_text) + + tracker = SubdirectoryHintTracker(working_dir=str(project)) + start = time.monotonic() + with caplog.at_level("WARNING", logger="agent.prompt_builder"): + result = tracker.check_tool_call( + "read_file", {"path": str(project / "backend" / "src" / "main.py")} + ) + elapsed = time.monotonic() - start + + assert elapsed < 0.4, f"hint load blocked for {elapsed:.2f}s" + assert result is None + assert "timed out" in caplog.text.lower() + + class TestPermissionErrorHandling: """Regression tests for PermissionError in filesystem checks (ref #6214).""" @@ -245,7 +282,7 @@ class TestExcludedDirectories: @pytest.mark.parametrize( "excluded", - ["backups", "node_modules", ".git", "venv", "site-packages", ".Trash", "vendor"], + sorted(SEARCH_PRUNE_DIR_NAMES), ) def test_excluded_directory_skipped(self, tmp_path, excluded): target = tmp_path / excluded / "snapshot" diff --git a/tests/agent/test_turn_facade_lease.py b/tests/agent/test_turn_facade_lease.py index cb73f5818e..689a59e57c 100644 --- a/tests/agent/test_turn_facade_lease.py +++ b/tests/agent/test_turn_facade_lease.py @@ -83,7 +83,7 @@ def test_admission_sets_holder_attrs_and_release_clears_them(monkeypatch): assert agent._active_session_turn_lease_holder == lease.holder assert agent._active_session_turn_lease_ttl_seconds == LEASE_TTL_SECONDS assert lease.holder.startswith("pid=") and ":platform=cli" in lease.holder - assert lease.refresh_thread is not None and lease.liveness_thread is None + assert lease.watchdog is None and lease.timer_handles == [] assert lease.is_turn_active() is False lease.stop_refresher() diff --git a/tests/agent/test_usage_pricing.py b/tests/agent/test_usage_pricing.py index 1b6862aac6..4d4e8a1ea8 100644 --- a/tests/agent/test_usage_pricing.py +++ b/tests/agent/test_usage_pricing.py @@ -105,6 +105,48 @@ def test_deepseek_v4_pro_pricing_entry_exists(): assert float(entry.cache_read_cost_per_million) == 0.003625 +def test_bundled_pricing_skips_endpoint_metadata(monkeypatch): + """An exact bundled price must not block on the provider's /models API.""" + monkeypatch.setattr( + "agent.usage_pricing.fetch_endpoint_model_metadata", + lambda *_args, **_kwargs: (_ for _ in ()).throw( + AssertionError("endpoint metadata should not be fetched") + ), + ) + + entry = get_pricing_entry( + "deepseek-chat", + provider="deepseek", + base_url="https://api.deepseek.com/v1", + ) + + assert entry is not None + assert entry.source == "official_docs_snapshot" + + +def test_unknown_model_falls_back_to_endpoint_metadata(monkeypatch): + """Models absent from the bundled table still use endpoint pricing.""" + monkeypatch.setattr( + "agent.usage_pricing.fetch_endpoint_model_metadata", + lambda *_args, **_kwargs: { + "deepseek-future": { + "pricing": {"prompt": "0.000001", "completion": "0.000002"} + } + }, + ) + + entry = get_pricing_entry( + "deepseek-future", + provider="deepseek", + base_url="https://api.deepseek.com/v1", + ) + + assert entry is not None + assert entry.source == "provider_models_api" + assert entry.input_cost_per_million == Decimal("1") + assert entry.output_cost_per_million == Decimal("2") + + def test_deepseek_deprecated_aliases_price_as_v4_flash(): diff --git a/tests/computer_use/test_cua_wayland_env.py b/tests/computer_use/test_cua_wayland_env.py new file mode 100644 index 0000000000..df06e132d6 --- /dev/null +++ b/tests/computer_use/test_cua_wayland_env.py @@ -0,0 +1,22 @@ +from unittest.mock import patch + +from tools.computer_use import cua_backend + + +_VAR = "CUA_DRIVER_RS_ENABLE_WAYLAND" + + +def _child_env(base_env, native_wayland): + config = {"computer_use": {"native_wayland": native_wayland}} + with patch("hermes_cli.config.load_config", return_value=config), \ + patch.object(cua_backend.sys, "platform", "linux"): + return cua_backend.cua_driver_child_env(base_env) + + +def test_configured_native_wayland_reaches_linux_wayland_child(): + assert _child_env({"WAYLAND_DISPLAY": "wayland-1"}, True)[_VAR] == "1" + + +def test_native_wayland_not_injected_without_wayland_display_or_opt_in(): + assert _VAR not in _child_env({"DISPLAY": ":0"}, True) + assert _VAR not in _child_env({"WAYLAND_DISPLAY": "wayland-1"}, False) diff --git a/tests/cron/test_cron_failure_deliver.py b/tests/cron/test_cron_failure_deliver.py index 1695447ad0..88e6fee244 100644 --- a/tests/cron/test_cron_failure_deliver.py +++ b/tests/cron/test_cron_failure_deliver.py @@ -76,7 +76,7 @@ def run_env(monkeypatch, tmp_path): monkeypatch.setattr(s, "create_execution", lambda *_a, **_kw: {"id": "exec-t"}) monkeypatch.setattr(s, "claim_dispatch", lambda _job_id: True) - monkeypatch.setattr(s, "mark_execution_running", lambda _execution_id: None) + monkeypatch.setattr(s, "mark_execution_running", lambda _execution_id: {}) monkeypatch.setattr( s, "save_job_output", lambda jid, out: state["saved"].append(jid) or f"/tmp/{jid}.txt", diff --git a/tests/cron/test_delivery_queue.py b/tests/cron/test_delivery_queue.py new file mode 100644 index 0000000000..b4c56edfaf --- /dev/null +++ b/tests/cron/test_delivery_queue.py @@ -0,0 +1,212 @@ +"""Durable at-most-once delivery handoff for restart-safe cron workers.""" + +from __future__ import annotations + +import sqlite3 +from unittest.mock import Mock + +import pytest + + +def test_pending_delivery_is_claimed_and_sent_once(tmp_path, monkeypatch): + import cron.delivery_queue as queue + + monkeypatch.setattr(queue, "DELIVERY_DB", tmp_path / "deliveries.db") + queue.enqueue("exec-1", {"id": "job-1"}, "brief") + send = Mock(return_value=None) + + assert queue.drain(send) == 1 + assert queue.drain(send) == 0 + send.assert_called_once_with({"id": "job-1"}, "brief", False) + status = queue.get_status("exec-1") + assert status["status"] == "delivered" + assert status["job_json"] == "{}" + assert status["content"] == "" + + +def test_terminal_delivery_retention_is_bounded(tmp_path, monkeypatch): + import cron.delivery_queue as queue + + monkeypatch.setattr(queue, "DELIVERY_DB", tmp_path / "deliveries.db") + monkeypatch.setattr(queue, "MAX_TERMINAL_DELIVERIES", 2, raising=False) + for index in range(4): + execution_id = f"exec-{index}" + queue.enqueue(execution_id, {"id": f"job-{index}"}, f"brief-{index}") + assert queue.claim_next()["execution_id"] == execution_id + assert queue._finish(execution_id, error=None) + + # Pruning may discard verbose outcome rows, but never the durable + # idempotency tombstone for an execution that could be replayed later. + pruned = queue.get_status("exec-0") + assert pruned is not None + assert pruned["status"] == "delivered" + assert queue.get_status("exec-1")["status"] == "delivered" + assert queue.get_status("exec-2")["status"] == "delivered" + assert queue.get_status("exec-3")["status"] == "delivered" + + # Pruning may discard verbose outcome rows, but never the durable + # idempotency tombstone for an execution that could be replayed later. + queue.enqueue("exec-0", {"id": "job-replayed"}, "duplicate brief") + send = Mock(return_value=None) + assert queue.drain(send) == 0 + send.assert_not_called() + + +def test_failure_delivery_lane_survives_durable_handoff(tmp_path, monkeypatch): + import cron.delivery_queue as queue + + monkeypatch.setattr(queue, "DELIVERY_DB", tmp_path / "deliveries.db") + queue.enqueue( + "exec-failure", + {"id": "job-failure", "failure_deliver": "local"}, + "failed", + for_failure=True, + ) + send = Mock(return_value=None) + + assert queue.drain(send) == 1 + send.assert_called_once_with( + {"id": "job-failure", "failure_deliver": "local"}, + "failed", + True, + ) + + +def test_legacy_queue_schema_adds_failure_lane_before_enqueue(tmp_path, monkeypatch): + import cron.delivery_queue as queue + + db = tmp_path / "deliveries.db" + with sqlite3.connect(db) as conn: + conn.execute( + """CREATE TABLE deliveries ( + execution_id TEXT PRIMARY KEY, + job_json TEXT NOT NULL, + content TEXT NOT NULL, + status TEXT NOT NULL, + owner_process_id TEXT, + owner_pid INTEGER, + owner_started_at INTEGER, + created_at TEXT NOT NULL, + finished_at TEXT, + error TEXT + )""" + ) + monkeypatch.setattr(queue, "DELIVERY_DB", db) + + queue.enqueue( + "exec-migrated", + {"id": "job-migrated"}, + "failed", + for_failure=True, + ) + + assert queue.get_status("exec-migrated")["for_failure"] == 1 + + +def test_wait_timeout_marks_inflight_delivery_unknown_without_retry( + tmp_path, monkeypatch +): + import cron.delivery_queue as queue + + monkeypatch.setattr(queue, "DELIVERY_DB", tmp_path / "deliveries.db") + queue.enqueue("exec-inflight", {"id": "job-inflight"}, "result") + assert queue.claim_next() is not None + + error = queue.enqueue_and_wait( + "exec-inflight", {"id": "job-inflight"}, "result", timeout=0 + ) + + assert error is not None + assert "outcome is unknown" in error + status = queue.get_status("exec-inflight") + assert status is not None + assert status["status"] == "unknown" + send = Mock(return_value=None) + assert queue.drain(send) == 0 + send.assert_not_called() + + +def test_dead_delivery_owner_becomes_unknown_and_is_not_retried( + tmp_path, monkeypatch +): + import cron.delivery_queue as queue + + monkeypatch.setattr(queue, "DELIVERY_DB", tmp_path / "deliveries.db") + queue.enqueue("exec-1", {"id": "job-1"}, "brief") + assert queue.claim_next() is not None + monkeypatch.setattr(queue, "_PROCESS_ID", "replacement-gateway") + monkeypatch.setattr(queue, "_owner_is_live", lambda _pid, _started: False) + + assert queue.recover_abandoned() == 1 + send = Mock() + assert queue.drain(send) == 0 + send.assert_not_called() + assert queue.get_status("exec-1")["status"] == "unknown" + + +def test_delivery_failure_is_terminal_not_retried_and_redacted( + tmp_path, monkeypatch +): + import cron.delivery_queue as queue + + monkeypatch.setattr(queue, "DELIVERY_DB", tmp_path / "deliveries.db") + queue.enqueue("exec-1", {"id": "job-1"}, "brief") + send = Mock(return_value="request failed: https://example.test/?token=TOKEN123") + + assert queue.drain(send) == 1 + assert queue.drain(send) == 0 + assert send.call_count == 1 + status = queue.get_status("exec-1") + assert status is not None + assert status["status"] == "failed" + assert "TOKEN123" not in status["error"] + assert "token=***" in status["error"] + + +def test_wait_timeout_leaves_unclaimed_delivery_queued_for_next_gateway( + tmp_path, monkeypatch +): + """A row nobody claimed was never attempted: it is not uncertain, so a + gateway outage longer than the worker's wait budget must not lose it.""" + import cron.delivery_queue as queue + + monkeypatch.setattr(queue, "DELIVERY_DB", tmp_path / "deliveries.db") + job = {"id": "job-3", "deliver": "origin"} + + error = queue.enqueue_and_wait("exec-3", job, "result", timeout=0) + + # Deferred, not failed: the worker must not record delivery_failed. + assert error is None + status = queue.get_status("exec-3") + assert status is not None + assert status["status"] == "pending" + send = Mock(return_value=None) + assert queue.drain(send) == 1 + send.assert_called_once_with(job, "result", False) + assert queue.get_status("exec-3")["status"] == "delivered" + + +def test_same_gateway_recovers_terminalization_failure_without_resending( + tmp_path, monkeypatch +): + import cron.delivery_queue as queue + + monkeypatch.setattr(queue, "DELIVERY_DB", tmp_path / "deliveries.db") + queue.enqueue("exec-4", {"id": "job-4"}, "result") + send = Mock(return_value=None) + original_finish = queue._finish + monkeypatch.setattr( + queue, + "_finish", + Mock(side_effect=OSError("database temporarily unavailable")), + ) + + with pytest.raises(OSError, match="temporarily unavailable"): + queue.drain(send) + + monkeypatch.setattr(queue, "_finish", original_finish) + assert queue.drain(send) == 0 + send.assert_called_once() + status = queue.get_status("exec-4") + assert status["status"] == "unknown" + assert "not retried" in status["error"] diff --git a/tests/cron/test_execution_ledger.py b/tests/cron/test_execution_ledger.py index ffe39a5164..d68259ae4b 100644 --- a/tests/cron/test_execution_ledger.py +++ b/tests/cron/test_execution_ledger.py @@ -39,6 +39,102 @@ def test_execution_transitions_are_durable(monkeypatch, tmp_path): assert persisted == [completed] +def test_execution_can_be_loaded_by_exact_attempt_id(monkeypatch, tmp_path): + executions = _point_ledger(monkeypatch, tmp_path) + first = executions.create_execution("same-job", source="builtin") + second = executions.create_execution("same-job", source="builtin") + + assert executions.get_execution(first["id"]) == first + assert executions.get_execution(second["id"]) == second + assert executions.get_execution("missing") is None + + +def test_fresh_external_handoff_is_not_recovered_before_worker_adopts( + monkeypatch, tmp_path +): + executions = _point_ledger(monkeypatch, tmp_path) + record = executions.create_execution("handoff-job", source="builtin") + assert executions.mark_execution_handoff_pending(record["id"]) is not None + + monkeypatch.setattr(executions, "_PROCESS_ID", "replacement-gateway") + monkeypatch.setattr(executions, "_owner_is_live", lambda _pid, _started: False) + + assert executions.recover_interrupted_executions() == 0 + assert executions.get_execution(record["id"])["status"] == "claimed" + adopted = executions.adopt_claimed_execution(record["id"]) + assert adopted["status"] == "running" + assert adopted["handoff_pending"] == 0 + + +def test_stale_external_handoff_is_recovered_unknown(monkeypatch, tmp_path): + executions = _point_ledger(monkeypatch, tmp_path) + record = executions.create_execution("handoff-job", source="builtin") + pending = executions.mark_execution_handoff_pending(record["id"]) + + monkeypatch.setattr(executions, "_PROCESS_ID", "replacement-gateway") + monkeypatch.setattr(executions, "_owner_is_live", lambda _pid, _started: False) + monkeypatch.setattr( + executions.time, + "time", + lambda: pending["handoff_started_at"] + + executions.HANDOFF_ADOPTION_GRACE_SECONDS + + 1, + ) + + assert executions.recover_interrupted_executions() == 1 + recovered = executions.get_execution(record["id"]) + assert recovered["status"] == "unknown" + assert recovered["handoff_pending"] == 0 + + +def test_recovery_does_not_overwrite_concurrent_worker_adoption(monkeypatch, tmp_path): + executions = _point_ledger(monkeypatch, tmp_path) + record = executions.create_execution("adoption-race", source="builtin") + pending = executions.mark_execution_handoff_pending(record["id"]) + assert pending is not None + monkeypatch.setattr(executions, "_PROCESS_ID", "replacement-scheduler") + monkeypatch.setattr( + executions.time, + "time", + lambda: pending["handoff_started_at"] + + executions.HANDOFF_ADOPTION_GRACE_SECONDS + + 1, + ) + + def adopt_while_liveness_is_checked(_pid, _started_at): + monkeypatch.setattr(executions, "_PROCESS_ID", "external-worker") + monkeypatch.setattr(executions.os, "getpid", lambda: 4242) + monkeypatch.setattr(executions, "_process_start_time", lambda _pid: 9876) + assert executions.adopt_claimed_execution(record["id"]) is not None + return False + + monkeypatch.setattr(executions, "_owner_is_live", adopt_while_liveness_is_checked) + + assert executions.recover_interrupted_executions() == 0 + current = executions.get_execution(record["id"]) + assert current is not None + assert current["status"] == "running" + assert current["process_id"] == "external-worker" + assert current["pid"] == 4242 + + +def test_foreign_process_cannot_start_or_finish_execution(monkeypatch, tmp_path): + executions = _point_ledger(monkeypatch, tmp_path) + record = executions.create_execution("owner-fence", source="builtin") + original_process_id = executions._PROCESS_ID + original_pid = record["pid"] + + monkeypatch.setattr(executions, "_PROCESS_ID", "foreign-process") + monkeypatch.setattr(executions.os, "getpid", lambda: original_pid + 1) + assert executions.mark_execution_running(record["id"]) is None + assert executions.finish_execution(record["id"], success=True) is None + + monkeypatch.setattr(executions, "_PROCESS_ID", original_process_id) + monkeypatch.setattr(executions.os, "getpid", lambda: original_pid) + assert executions.mark_execution_running(record["id"]) is not None + assert executions.finish_execution(record["id"], success=True) is not None + + def test_execution_ledger_follows_the_current_profile_home(monkeypatch, tmp_path): import cron.executions as executions @@ -83,6 +179,24 @@ def test_retention_bounds_terminal_history_but_preserves_inflight(monkeypatch, t assert executions.latest_execution("live")["status"] == "running" +def test_recently_finished_long_running_execution_survives_retention( + monkeypatch, tmp_path +): + executions = _point_ledger(monkeypatch, tmp_path) + monkeypatch.setattr(executions, "MAX_TERMINAL_EXECUTIONS", 1) + long_running = executions.create_execution("long-running", source="builtin") + assert executions.mark_execution_running(long_running["id"]) is not None + newer = executions.create_execution("newer", source="builtin") + assert executions.finish_execution(newer["id"], success=True) is not None + + finished = executions.finish_execution(long_running["id"], success=True) + + assert finished is not None + assert finished["status"] == "completed" + assert executions.get_execution(long_running["id"])["status"] == "completed" + assert executions.get_execution(newer["id"]) is None + + def test_corrupt_store_fails_closed_without_overwrite(monkeypatch, tmp_path): executions = _point_ledger(monkeypatch, tmp_path) executions.EXECUTIONS_FILE.parent.mkdir(parents=True) @@ -221,7 +335,7 @@ def test_run_one_job_records_running_then_terminal(monkeypatch): monkeypatch.setattr( scheduler, "mark_execution_running", - lambda execution_id: events.append(("running", execution_id)), + lambda execution_id: events.append(("running", execution_id)) or {}, raising=False, ) monkeypatch.setattr( diff --git a/tests/cron/test_lifecycle_guard_budget.py b/tests/cron/test_lifecycle_guard_budget.py new file mode 100644 index 0000000000..181461744f --- /dev/null +++ b/tests/cron/test_lifecycle_guard_budget.py @@ -0,0 +1,241 @@ +"""Whole-walk work budget for the gateway lifecycle guard (#78398). + +The per-file byte cap and recursion depth bound one read, not the walk. These +tests pin the shared budget that bounds the whole referenced-script walk and +is charged *before* any text reaches ``shlex``. + +Budget constants are monkeypatched to tiny values so the tests are fast and +deterministic; ``_LifecycleScanBudget`` reads them at construction time. +""" + +from __future__ import annotations + +import pytest + +import cron.lifecycle_guard as lifecycle_guard + +guard = lifecycle_guard.contains_gateway_lifecycle_command_or_referenced_script + + +def _explode(*_args, **_kwargs): + raise AssertionError("over-budget text reached shlex") + + +# --- root command (depth 0) ----------------------------------------------- + + +def test_root_byte_limit_allows_exact_and_rejects_plus_one(monkeypatch): + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_BYTES", 8) + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_LINE_BYTES", 8) + + assert guard("x" * 8) is False + + monkeypatch.setattr(lifecycle_guard.shlex, "shlex", _explode) + assert guard("x" * 9) is True + + +def test_root_line_limit_allows_exact_and_rejects_plus_one(monkeypatch): + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_LINES", 2) + + assert guard("one\ntwo") is False + assert guard("one\ntwo\nthree") is True + + +def test_single_giant_line_rejected_before_shlex(monkeypatch): + """One enormous token is the quadratic shlex case.""" + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_LINE_BYTES", 8) + monkeypatch.setattr(lifecycle_guard.shlex, "shlex", _explode) + + assert guard("xxxxxxxxx\necho ok") is True + + +def test_root_budget_counts_utf8_bytes(monkeypatch): + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_BYTES", 4) + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_LINE_BYTES", 4) + + assert guard("éé") is False + assert guard("ééé") is True + + +def test_exhaustion_is_logged_at_warning(monkeypatch, caplog): + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_BYTES", 4) + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_LINE_BYTES", 4) + + with caplog.at_level("WARNING", logger=lifecycle_guard.logger.name): + assert guard("echo hello") is True + assert "budget exhausted" in caplog.text + + +def test_lifecycle_scan_root_within_budget_is_not_a_verdict(monkeypatch): + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_BYTES", 8) + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_LINE_BYTES", 8) + + assert lifecycle_guard.lifecycle_scan_root_within_budget("x" * 8) is True + assert lifecycle_guard.lifecycle_scan_root_within_budget("x" * 9) is False + + +# --- referenced-script walk ------------------------------------------------ + + +def test_unique_path_budget_bounds_reads_and_fails_closed(monkeypatch, tmp_path): + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_PATHS", 2) + for i in range(3): + (tmp_path / f"s{i}.sh").write_text("echo ok\n", encoding="utf-8") + + two = " && ".join(f"bash {tmp_path}/s{i}.sh" for i in range(2)) + three = " && ".join(f"bash {tmp_path}/s{i}.sh" for i in range(3)) + + assert guard(two) is False + assert guard(three) is True + + +def test_repeated_path_does_not_spend_unique_path_budget(monkeypatch, tmp_path): + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_PATHS", 1) + script = tmp_path / "s.sh" + script.write_text("echo ok\n", encoding="utf-8") + + assert guard(f"bash {script} && bash {script} && sh {script}") is False + + +def test_remote_read_budget_charged_before_remote_read(monkeypatch): + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_REMOTE_READS", 1) + reads: list[str] = [] + + def remote(path: str): + reads.append(path) + return "echo ok\n" + + assert ( + guard( + "bash /remote/a.sh && bash /remote/b.sh", + read_remote_script=remote, + ) + is True + ) + assert reads == ["/remote/a.sh"] + + +def test_cumulative_text_budget_bounds_recursive_scan(monkeypatch, tmp_path): + """Two scripts individually under the per-file cap exceed the walk cap. + + Relative references keep the root command short so the budget arithmetic + is about the scripts, not the tmp_path length.""" + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_BYTES", 48) + (tmp_path / "a.sh").write_text("echo " + "a" * 10 + "\n", encoding="utf-8") # 16 + (tmp_path / "b.sh").write_text("echo " + "b" * 10 + "\n", encoding="utf-8") # 16 + cwd = str(tmp_path) + + # 9 (root) + 16 fits in 48; 19 (root) + 16 + 16 does not → fail closed. + assert guard("bash a.sh", cwd=cwd) is False + assert guard("bash a.sh;bash b.sh", cwd=cwd) is True + + +def test_referenced_read_is_capped_at_remaining_budget(monkeypatch, tmp_path): + """A file bigger than what the walk can still afford is never read whole: + the read helper receives the remaining budget as its cap.""" + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_BYTES", 64) + (tmp_path / "big.sh").write_text("echo " + "x" * 200 + "\n", encoding="utf-8") + + caps: list = [] + original = lifecycle_guard._read_referenced_script + + def spy(path, *, max_bytes=None): + caps.append(max_bytes) + return original(path, max_bytes=max_bytes) + + monkeypatch.setattr(lifecycle_guard, "_read_referenced_script", spy) + + root = "bash big.sh" + assert guard(root, cwd=str(tmp_path)) is True + assert caps == [64 - len(root)] + + +def test_remote_script_sanitizer_honours_remaining_budget(): + text, unsafe = lifecycle_guard._sanitize_remote_script_text( + "echo ok\n", max_bytes=4 + ) + assert (text, unsafe) == (None, True) + text, unsafe = lifecycle_guard._sanitize_remote_script_text( + "echo ok\n", max_bytes=8 + ) + assert (text, unsafe) == ("echo ok\n", False) + + +def test_line_budget_fails_closed_before_tokenizing_every_line( + monkeypatch, tmp_path +): + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_LINES", 4) + script = tmp_path / "many.sh" + script.write_text("echo ok\n" * 10, encoding="utf-8") + + lexers = 0 + real_shlex = lifecycle_guard.shlex.shlex + + def counting(*args, **kwargs): + nonlocal lexers + lexers += 1 + return real_shlex(*args, **kwargs) + + monkeypatch.setattr(lifecycle_guard.shlex, "shlex", counting) + root = f"bash {script}" + assert guard(root) is True + # Only the one-line root was tokenized (a handful of lexers across the + # direct scans); the 10-line script never was. + assert 0 < lexers < 10 + + +# --- scheduler entry point -------------------------------------------------- + + +def test_check_gateway_lifecycle_shell_script_budget(monkeypatch, tmp_path): + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_BYTES", 8) + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_LINE_BYTES", 8) + script = tmp_path / "long-line.sh" + + script.write_text("x" * 7, encoding="utf-8") + lifecycle_guard.check_gateway_lifecycle("", str(script)) + + script.write_text("x" * 9, encoding="utf-8") + with pytest.raises(lifecycle_guard.GatewayLifecycleBlocked): + lifecycle_guard.check_gateway_lifecycle("", str(script)) + + +def test_check_gateway_lifecycle_python_path_charges_masker(monkeypatch, tmp_path): + """The .py branch's data-exemption masker tokenizes too, so it is budgeted + and fails closed before shlex on an over-budget line.""" + monkeypatch.setattr(lifecycle_guard, "_MAX_LIFECYCLE_SCAN_LINE_BYTES", 16) + + small = tmp_path / "small.py" + small.write_text("x = 1\n", encoding="utf-8") + lifecycle_guard.check_gateway_lifecycle("run report", str(small)) + + monkeypatch.setattr(lifecycle_guard.shlex, "shlex", _explode) + long_line = tmp_path / "long.py" + long_line.write_text("x = 1\n" + "y" * 40 + "\n", encoding="utf-8") + with pytest.raises(lifecycle_guard.GatewayLifecycleBlocked): + lifecycle_guard.check_gateway_lifecycle("run report", str(long_line)) + + +# --- no regression on realistic benign graphs ------------------------------ + + +def test_default_budget_admits_a_wide_benign_wrapper_graph(tmp_path): + """Issue #78398's shape: one wrapper invoking 200 small legitimate scripts + must still be allowed under the DEFAULT limits (an earlier fail-closed + attempt with a 64-path cap blocked exactly this).""" + children = [] + for i in range(200): + child = tmp_path / f"c{i}.sh" + child.write_text("echo step && ls -la /tmp\n" * 20, encoding="utf-8") + children.append(child) + hub = tmp_path / "hub.sh" + hub.write_text("".join(f"bash {c}\n" for c in children), encoding="utf-8") + + assert guard(f"bash {hub}") is False + + # ...and a lifecycle command hidden behind the 200 benign scripts is still + # found: the budget bounds work, it does not stop the walk early. + evil = tmp_path / "evil.sh" + evil.write_text("hermes gateway restart\n", encoding="utf-8") + hub.write_text(hub.read_text() + f"bash {evil}\n", encoding="utf-8") + assert guard(f"bash {hub}") is True diff --git a/tests/cron/test_parallel_pool.py b/tests/cron/test_parallel_pool.py index b159f265d7..9853dbc229 100644 --- a/tests/cron/test_parallel_pool.py +++ b/tests/cron/test_parallel_pool.py @@ -187,7 +187,7 @@ class TestRunningJobGuard: if job_id == "healthy-job" else None, ) - monkeypatch.setattr(sched, "mark_execution_running", lambda *_a, **_kw: None) + monkeypatch.setattr(sched, "mark_execution_running", lambda *_a, **_kw: {}) monkeypatch.setattr(sched, "heartbeat_fire_claim", lambda *_a, **_kw: True) n = sched.tick(verbose=False) diff --git a/tests/cron/test_restart_safe_worker.py b/tests/cron/test_restart_safe_worker.py new file mode 100644 index 0000000000..1be7e30654 --- /dev/null +++ b/tests/cron/test_restart_safe_worker.py @@ -0,0 +1,603 @@ +"""Restart-safe cron worker handoff and ownership contracts.""" + +from __future__ import annotations + +import asyncio +import json +import os +import signal +import subprocess +import sys +import threading +import time +from pathlib import Path +from unittest.mock import Mock + +import pytest + + +@pytest.fixture +def execution_ledger(tmp_path, monkeypatch): + import cron.executions as executions + + monkeypatch.setattr(executions, "EXECUTIONS_FILE", tmp_path / "executions.db") + return executions + + +def test_execution_owner_moves_to_external_worker_before_running( + execution_ledger, monkeypatch +): + record = execution_ledger.create_execution("job-1", source="builtin") + assert execution_ledger.mark_execution_handoff_pending(record["id"]) is not None + monkeypatch.setattr(execution_ledger.os, "getpid", lambda: 4242) + monkeypatch.setattr(execution_ledger, "_process_start_time", lambda pid: 9876) + + adopted = execution_ledger.adopt_claimed_execution(record["id"]) + + assert adopted is not None + assert adopted["pid"] == 4242 + assert adopted["process_started_at"] == 9876 + assert adopted["status"] == "running" + assert execution_ledger.adopt_claimed_execution(record["id"]) is None + assert execution_ledger.mark_execution_running(record["id"]) is None + + +def test_external_worker_cannot_adopt_execution_without_handoff_fence( + execution_ledger, monkeypatch +): + record = execution_ledger.create_execution("job-unfenced", source="builtin") + monkeypatch.setattr(execution_ledger.os, "getpid", lambda: 4242) + monkeypatch.setattr(execution_ledger, "_process_start_time", lambda _pid: 9876) + + assert execution_ledger.adopt_claimed_execution(record["id"]) is None + assert execution_ledger.get_execution(record["id"])["status"] == "claimed" + + +def test_genuine_external_worker_crash_is_recovered_unknown( + execution_ledger, monkeypatch +): + record = execution_ledger.create_execution("job-crash", source="builtin") + assert execution_ledger.mark_execution_handoff_pending(record["id"]) is not None + script = ( + "import os\n" + "from pathlib import Path\n" + "import cron.executions as executions\n" + f"executions.EXECUTIONS_FILE = Path({str(execution_ledger.EXECUTIONS_FILE)!r})\n" + f"assert executions.adopt_claimed_execution({record['id']!r}) is not None\n" + "os._exit(9)\n" + ) + + crashed = subprocess.run([sys.executable, "-c", script], check=False) + assert crashed.returncode == 9 + + monkeypatch.setattr(execution_ledger, "_PROCESS_ID", "replacement-scheduler") + assert execution_ledger.recover_interrupted_executions() == 1 + recovered = execution_ledger.latest_execution("job-crash") + assert recovered["status"] == "unknown" + assert "whether side effects ran is unknown" in recovered["error"] + + +@pytest.mark.linux_only +def test_restart_safe_gateway_child_fails_closed_without_scope(monkeypatch): + import tools.process_registry as process_registry + + monkeypatch.setattr(process_registry, "_is_supervised_gateway_process", lambda: True) + monkeypatch.setenv("INVOCATION_ID", "managed-service") + monkeypatch.setattr(process_registry, "_systemd_run_user_scope_available", lambda: False) + + with pytest.raises(RuntimeError, match="systemd-run --user --scope is unavailable"): + process_registry.restart_safe_gateway_child_argv( + ["python", "worker.py"], unit_suffix="cron-job-1" + ) + + +def test_restart_safe_gateway_child_is_unchanged_outside_managed_gateway(monkeypatch): + import tools.process_registry as process_registry + + command = ["python", "worker.py"] + monkeypatch.setattr(process_registry, "_is_supervised_gateway_process", lambda: False) + + assert process_registry.restart_safe_gateway_child_argv( + command, unit_suffix="cron-job-1" + ) is command + + +def test_restart_safe_gateway_child_never_probes_systemd_off_linux(monkeypatch): + import tools.process_registry as process_registry + + command = ["python", "worker.py"] + probe = Mock(side_effect=AssertionError("systemd probe ran off Linux")) + monkeypatch.setattr(process_registry, "_IS_LINUX", False) + monkeypatch.setattr(process_registry, "_is_supervised_gateway_process", lambda: True) + monkeypatch.setattr(process_registry, "_systemd_run_user_scope_available", probe) + monkeypatch.setenv("INVOCATION_ID", "managed-service") + + assert process_registry.restart_safe_gateway_child_argv( + command, unit_suffix="cron-job-1" + ) is command + probe.assert_not_called() + + +def test_external_worker_adopts_execution_and_runs_payload_once( + tmp_path, monkeypatch +): + import cron.scheduler as scheduler + + payload = tmp_path / "payload.json" + ack = tmp_path / "ready.json" + payload.write_text( + json.dumps({ + "job": {"id": "job-1", "execution_id": "exec-1"}, + "profile_home": str(tmp_path / "profile"), + }), + encoding="utf-8", + ) + from hermes_constants import get_hermes_home + + observed_homes = [] + adopted = Mock( + side_effect=lambda execution_id: ( + observed_homes.append(get_hermes_home().resolve()) + or {"id": execution_id, "status": "running"} + ) + ) + run = Mock( + side_effect=lambda *_args, **_kwargs: ( + observed_homes.append(get_hermes_home().resolve()) or True + ) + ) + monkeypatch.setattr("cron.executions.adopt_claimed_execution", adopted) + monkeypatch.setattr(scheduler, "run_one_job", run) + + assert scheduler._run_external_worker_payload(payload, ack) is True + + adopted.assert_called_once_with("exec-1") + run.assert_called_once() + assert run.call_args.args[0]["id"] == "job-1" + expected_home = (tmp_path / "profile").resolve() + assert observed_homes == [expected_home, expected_home] + assert ack.exists() + assert not payload.exists() + + +def test_external_worker_refuses_to_run_without_durable_ownership( + tmp_path, monkeypatch +): + import cron.scheduler as scheduler + + payload = tmp_path / "payload.json" + ack = tmp_path / "ready.json" + payload.write_text( + json.dumps({ + "job": {"id": "job-1", "execution_id": "exec-1"}, + "profile_home": str(tmp_path / "profile"), + }), + encoding="utf-8", + ) + monkeypatch.setattr("cron.executions.adopt_claimed_execution", lambda _id: None) + run = Mock() + monkeypatch.setattr(scheduler, "run_one_job", run) + + assert scheduler._run_external_worker_payload(payload, ack) is False + + run.assert_not_called() + assert not ack.exists() + + +def test_launch_external_worker_uses_restart_safe_scope_and_acknowledges( + tmp_path, monkeypatch +): + import cron.scheduler as scheduler + + job = {"id": "job-1", "execution_id": "exec-1", "prompt": "work"} + monkeypatch.setattr(scheduler, "_get_hermes_home", lambda: tmp_path) + wrapped_commands = [] + + def wrap(command, *, unit_suffix): + wrapped_commands.append((command, unit_suffix)) + return ["scope", "--", *command] + + monkeypatch.setattr( + "tools.process_registry.restart_safe_gateway_child_argv", wrap + ) + + class FakeProcess: + returncode = None + + def poll(self): + return self.returncode + + def wait(self, timeout=None): + if self.returncode is None: + raise subprocess.TimeoutExpired(cmd="worker", timeout=timeout) + return self.returncode + + spawned = [] + + payloads = [] + + def popen(command, **kwargs): + spawned.append((command, kwargs)) + payload_index = command.index("--external-worker-file") + 1 + payloads.append(json.loads(Path(command[payload_index]).read_text())) + ack_index = command.index("--ack-file") + 1 + Path(command[ack_index]).write_text( + json.dumps({"pid": 4321, "execution_id": "exec-1"}), + encoding="utf-8", + ) + return FakeProcess() + + handoff = Mock(return_value={"id": "exec-1", "handoff_pending": 1}) + monkeypatch.setattr(scheduler, "mark_execution_handoff_pending", handoff) + monkeypatch.setattr(scheduler.subprocess, "Popen", popen) + observed_statuses = iter( + [ + {"id": "exec-1", "status": "running"}, + {"id": "exec-1", "status": "completed"}, + ] + ) + get = Mock(side_effect=lambda _execution_id: next(observed_statuses)) + monkeypatch.setattr(scheduler, "get_execution", get) + monkeypatch.setenv("ANTHROPIC_API_KEY", "should-not-cross-profile") + from agent.secret_scope import set_multiplex_active + + set_multiplex_active(True) + try: + assert scheduler._launch_external_cron_worker(job) is True + finally: + set_multiplex_active(False) + assert wrapped_commands[0][1] == "cron-job-1-exec-exec-1" + assert spawned[0][0][0:2] == ["scope", "--"] + assert spawned[0][1]["start_new_session"] is True + assert "ANTHROPIC_API_KEY" not in spawned[0][1]["env"] + handoff.assert_called_once_with("exec-1") + assert get.call_count == 2 + assert payloads[0]["multiplex_active"] is True + # Once the attempt is terminal the parent reaps its own handoff artifacts. + assert not (tmp_path / "cron/external-workers/exec-1.json").exists() + + +def test_external_worker_exit_rechecks_exact_execution_before_failure(monkeypatch): + import cron.scheduler as scheduler + + statuses = iter( + [ + {"id": "exec-1", "status": "running"}, + {"id": "exec-1", "status": "completed"}, + ] + ) + get = Mock(side_effect=lambda _execution_id: next(statuses)) + monkeypatch.setattr(scheduler, "get_execution", get, raising=False) + process = Mock() + process.poll.return_value = 0 + process.wait.return_value = 0 + + assert scheduler._wait_for_external_cron_worker( + process, execution_id="exec-1" + ) is True + assert get.call_count == 2 + + +def test_external_worker_crash_recovers_uncertain_attempt(monkeypatch): + import cron.scheduler as scheduler + + statuses = iter( + [ + {"id": "exec-1", "status": "running"}, + {"id": "exec-1", "status": "unknown"}, + ] + ) + get = Mock(side_effect=lambda _execution_id: next(statuses)) + recover = Mock(return_value=1) + monkeypatch.setattr(scheduler, "get_execution", get) + monkeypatch.setattr( + scheduler, "recover_interrupted_executions", recover, raising=False + ) + process = Mock() + process.poll.return_value = 9 + process.wait.return_value = 9 + + assert scheduler._wait_for_external_cron_worker( + process, execution_id="exec-1" + ) is True + recover.assert_called_once_with() + assert get.call_count == 2 + + +def test_launch_external_worker_stays_in_process_outside_managed_gateway( + monkeypatch, +): + import cron.scheduler as scheduler + + command_calls = [] + + def unchanged(command, *, unit_suffix): + command_calls.append((command, unit_suffix)) + return command + + monkeypatch.setattr( + "tools.process_registry.restart_safe_gateway_child_argv", unchanged + ) + popen = Mock() + monkeypatch.setattr(scheduler.subprocess, "Popen", popen) + + assert scheduler._launch_external_cron_worker( + {"id": "job-1", "execution_id": "exec-1"} + ) is False + assert command_calls + popen.assert_not_called() + + +def test_shared_run_path_hands_gateway_fire_to_external_worker(monkeypatch): + import cron.scheduler as scheduler + + launch = Mock(return_value=True) + run = Mock(side_effect=AssertionError("agent ran inside gateway")) + monkeypatch.setattr(scheduler, "_launch_external_cron_worker", launch) + monkeypatch.setattr(scheduler, "run_job", run) + job = {"id": "job-1", "execution_id": "exec-1"} + + assert scheduler.run_one_job(job, adapters={"discord": object()}) is True + + launch.assert_called_once_with(job) + run.assert_not_called() + + +def test_shutdown_does_not_interrupt_restart_safe_waiter(): + import cron.scheduler as scheduler + + job_id = "external-waiter" + scheduler._running_job_ids.add(job_id) + scheduler._restart_safe_waiter_job_ids.add(job_id) + try: + assert scheduler.mark_running_jobs_interrupted("gateway restart") == [] + assert job_id not in scheduler._interrupted_job_ids + finally: + scheduler._restart_safe_waiter_job_ids.discard(job_id) + scheduler._running_job_ids.discard(job_id) + scheduler._interrupted_job_ids.discard(job_id) + + +def test_worker_delivery_queue_is_keyed_by_the_delivering_jobs_own_execution( + monkeypatch, tmp_path +): + """A nested in-process dispatch inside a worker (e.g. a script running + ``hermes cron run ``) must not queue under the OUTER execution id.""" + import cron.scheduler as scheduler + + queued = [] + monkeypatch.setattr( + "cron.delivery_queue.enqueue_and_wait", + lambda execution_id, job, content, **kw: ( + queued.append(execution_id) or "queued-marker" + ), + ) + monkeypatch.setattr( + scheduler, + "_resolve_delivery_targets", + lambda job, for_failure=False: [{"platform": "telegram", "chat_id": "123"}], + ) + + def _standalone(*_args, **_kwargs): + raise RuntimeError("standalone path reached") + + # First call the standalone (non-queue) path makes after the guard; the + # failure is reported as the delivery error string. + monkeypatch.setattr("gateway.config.load_gateway_config", _standalone) + monkeypatch.setenv("_HERMES_CRON_EXTERNAL_WORKER", "exec-outer") + + # Own attempt: routed through the durable queue. + assert scheduler._deliver_result( + {"id": "job-1", "execution_id": "exec-outer", "deliver": "telegram:123"}, + "done", + adapters=None, + loop=None, + ) == "queued-marker" + assert queued == ["exec-outer"] + + # A different job's attempt: must NOT be queued under exec-outer; it falls + # through to the standalone path. + error = scheduler._deliver_result( + {"id": "job-2", "execution_id": "exec-inner", "deliver": "telegram:123"}, + "done", + adapters=None, + loop=None, + ) + assert error == "failed to load gateway config: standalone path reached" + assert queued == ["exec-outer"] + + +def test_gateway_tool_run_without_adapter_objects_hands_off(monkeypatch): + import cron.scheduler as scheduler + + created = Mock(return_value={"id": "exec-tool"}) + launch = Mock(return_value=True) + run = Mock(side_effect=AssertionError("agent ran inside gateway")) + monkeypatch.setattr(scheduler, "create_execution", created) + monkeypatch.setattr(scheduler, "_launch_external_cron_worker", launch) + monkeypatch.setattr(scheduler, "run_job", run) + job = {"id": "tool-job"} + + assert scheduler.run_one_job(job, adapters=None) is True + + created.assert_called_once_with("tool-job", source="direct") + assert job["execution_id"] == "exec-tool" + launch.assert_called_once_with(job) + run.assert_not_called() + + +def test_shared_run_path_creates_execution_before_managed_handoff(monkeypatch): + import cron.scheduler as scheduler + + created = Mock(return_value={"id": "exec-new"}) + launch = Mock(return_value=True) + monkeypatch.setattr(scheduler, "create_execution", created) + monkeypatch.setattr(scheduler, "_launch_external_cron_worker", launch) + job = {"id": "manual-job"} + + assert scheduler.run_one_job(job, adapters={"discord": object()}) is True + + created.assert_called_once_with("manual-job", source="direct") + assert job["execution_id"] == "exec-new" + launch.assert_called_once_with(job) + + +def test_lost_execution_start_cas_prevents_side_effects(monkeypatch): + import cron.scheduler as scheduler + + run = Mock(side_effect=AssertionError("side effect ran without ownership")) + monkeypatch.setattr(scheduler, "claim_dispatch", lambda _job_id: True) + monkeypatch.setattr(scheduler, "mark_execution_running", lambda _execution_id: None) + monkeypatch.setattr(scheduler, "run_job", run) + + assert scheduler.run_one_job( + {"id": "job-1", "execution_id": "exec-1"}, adapters=None + ) is True + run.assert_not_called() + + +@pytest.mark.linux_only +@pytest.mark.live_system_guard_bypass +def test_managed_gateway_restart_preserves_active_worker_and_single_side_effect( + tmp_path, monkeypatch +): + import cron.delivery_queue as delivery_queue + import cron.executions as executions + import cron.scheduler as scheduler + from cron.jobs import create_job, use_cron_store + from gateway.config import Platform, PlatformConfig + from gateway.status import _pid_exists + from tools import process_registry + + if not process_registry._systemd_run_user_scope_available(): + pytest.skip("systemd-run --user --scope is unavailable on this host") + + home = tmp_path / "profile" + scripts_dir = home / "scripts" + scripts_dir.mkdir(parents=True) + started = tmp_path / "started" + release = tmp_path / "release" + side_effect = tmp_path / "side-effect" + probe = scripts_dir / "restart_probe.py" + probe.write_text( + "import pathlib, time\n" + f"started = pathlib.Path({str(started)!r})\n" + f"release = pathlib.Path({str(release)!r})\n" + f"side_effect = pathlib.Path({str(side_effect)!r})\n" + "started.write_text('started')\n" + "deadline = time.monotonic() + 15\n" + "while not release.exists() and time.monotonic() < deadline:\n" + " time.sleep(0.05)\n" + "if not release.exists():\n" + " raise SystemExit('release timeout')\n" + "with side_effect.open('a') as handle:\n" + " handle.write('once\\n')\n" + "print('completed')\n", + encoding="utf-8", + ) + monkeypatch.setenv("HERMES_HOME", str(home)) + with use_cron_store(home): + job = create_job( + prompt=None, + schedule="every 1h", + name="restart probe", + script=probe.name, + no_agent=True, + deliver="telegram:123", + ) + payload = tmp_path / "job.json" + launched = tmp_path / "launched.json" + payload.write_text(json.dumps(job), encoding="utf-8") + + sent = [] + adapter = Mock() + + async def send(_chat_id, content, metadata=None): + sent.append((content, metadata)) + return {"success": True, "message_id": "restart-delivery-1"} + + adapter.send = send + gateway_config = Mock() + gateway_config.platforms = { + Platform.TELEGRAM: PlatformConfig(enabled=True), + } + gateway_config.get_home_channel = lambda _platform: None + monkeypatch.setattr( + "gateway.config.load_gateway_config", lambda: gateway_config + ) + monkeypatch.setattr( + scheduler, "load_config", lambda: {"cron": {"wrap_response": False}} + ) + replacement_loop = asyncio.new_event_loop() + replacement_thread = threading.Thread( + target=replacement_loop.run_forever, + daemon=True, + ) + replacement_thread.start() + deadline = time.monotonic() + 2 + while not replacement_loop.is_running() and time.monotonic() < deadline: + time.sleep(0.01) + assert replacement_loop.is_running() + + harness = ( + "import json, os, pathlib, time\n" + f"os.environ['HERMES_HOME'] = {str(home)!r}\n" + "os.environ['INVOCATION_ID'] = 'restart-fixture'\n" + "from cron import scheduler\n" + "from tools import process_registry\n" + "process_registry._is_supervised_gateway_process = lambda: True\n" + f"job = json.loads(pathlib.Path({str(payload)!r}).read_text())\n" + "if not scheduler.run_one_job(job, adapters=None, loop=None):\n" + " raise SystemExit('worker was not isolated')\n" + f"pathlib.Path({str(launched)!r}).write_text('returned')\n" + ) + parent = subprocess.Popen([sys.executable, "-c", harness]) + worker_pid = None + try: + deadline = time.monotonic() + 10 + current = None + while time.monotonic() < deadline: + if parent.poll() is not None: + pytest.fail(f"gateway fixture exited early with {parent.returncode}") + current = executions.latest_execution(job["id"]) + if started.exists() and current and current.get("pid") != os.getpid(): + break + time.sleep(0.05) + assert started.exists() + assert current is not None + execution = current + worker_pid = int(current["pid"]) + assert not launched.exists(), "handoff returned before execution completed" + + # Replacing a managed gateway kills its old process tree. The active + # cron owner must remain in its transient scope and keep the same PID. + parent.terminate() + parent.wait(timeout=5) + assert _pid_exists(worker_pid) + + release.write_text("go", encoding="utf-8") + deadline = time.monotonic() + 10 + while time.monotonic() < deadline: + row = delivery_queue.get_status(execution["id"]) + if row and row["status"] == "pending": + scheduler.drain_delivery_queue( + {Platform.TELEGRAM: adapter}, replacement_loop + ) + current = executions.latest_execution(job["id"]) + if current and current["status"] == "completed": + break + time.sleep(0.05) + assert executions.latest_execution(job["id"])["status"] == "completed" + assert side_effect.read_text(encoding="utf-8").splitlines() == ["once"] + assert delivery_queue.get_status(execution["id"])["status"] == "delivered" + assert len(sent) == 1 + assert "completed" in sent[0][0] + finally: + replacement_loop.call_soon_threadsafe(replacement_loop.stop) + replacement_thread.join(timeout=2) + replacement_loop.close() + if parent.poll() is None: + parent.terminate() + parent.wait(timeout=5) + if worker_pid is not None and _pid_exists(worker_pid): + os.kill(worker_pid, signal.SIGKILL) diff --git a/tests/cron/test_run_one_job.py b/tests/cron/test_run_one_job.py index 93bcef00a5..9b241de1b6 100644 --- a/tests/cron/test_run_one_job.py +++ b/tests/cron/test_run_one_job.py @@ -88,7 +88,7 @@ def test_run_one_job_exception_delivers_failure_alert(monkeypatch): s, "create_execution", lambda *_a, **_kw: {"id": "exec-j3"} ) monkeypatch.setattr(s, "claim_dispatch", lambda _job_id: True) - monkeypatch.setattr(s, "mark_execution_running", lambda _execution_id: None) + monkeypatch.setattr(s, "mark_execution_running", lambda _execution_id: {}) monkeypatch.setattr( s, "run_job", @@ -141,7 +141,7 @@ def test_run_one_job_exception_records_failure_alert_delivery_error(monkeypatch) s, "create_execution", lambda *_a, **_kw: {"id": "exec-j4"} ) monkeypatch.setattr(s, "claim_dispatch", lambda _job_id: True) - monkeypatch.setattr(s, "mark_execution_running", lambda _execution_id: None) + monkeypatch.setattr(s, "mark_execution_running", lambda _execution_id: {}) monkeypatch.setattr( s, "run_job", @@ -165,7 +165,7 @@ def _patch_escaped_failure(monkeypatch, delivered, *, exec_id, err): """Make run_job raise, and capture what the escape handler delivers.""" monkeypatch.setattr(s, "create_execution", lambda *_a, **_kw: {"id": exec_id}) monkeypatch.setattr(s, "claim_dispatch", lambda _job_id: True) - monkeypatch.setattr(s, "mark_execution_running", lambda _execution_id: None) + monkeypatch.setattr(s, "mark_execution_running", lambda _execution_id: {}) monkeypatch.setattr( s, "run_job", @@ -246,7 +246,7 @@ def test_run_one_job_exception_after_delivery_does_not_redeliver(monkeypatch): s, "create_execution", lambda *_a, **_kw: {"id": "exec-j5"} ) monkeypatch.setattr(s, "claim_dispatch", lambda _job_id: True) - monkeypatch.setattr(s, "mark_execution_running", lambda _execution_id: None) + monkeypatch.setattr(s, "mark_execution_running", lambda _execution_id: {}) monkeypatch.setattr( s, "run_job", @@ -288,7 +288,7 @@ def test_run_one_job_keyboard_interrupt_skips_delivery_and_reraises(monkeypatch) s, "create_execution", lambda *_a, **_kw: {"id": "exec-j6"} ) monkeypatch.setattr(s, "claim_dispatch", lambda _job_id: True) - monkeypatch.setattr(s, "mark_execution_running", lambda _execution_id: None) + monkeypatch.setattr(s, "mark_execution_running", lambda _execution_id: {}) monkeypatch.setattr( s, "run_job", diff --git a/tests/cron/test_script_claim_heartbeat.py b/tests/cron/test_script_claim_heartbeat.py index effb3f8dd4..11fb4501b0 100644 --- a/tests/cron/test_script_claim_heartbeat.py +++ b/tests/cron/test_script_claim_heartbeat.py @@ -392,7 +392,7 @@ def test_lost_fire_claim_stops_stale_delivery(monkeypatch): monkeypatch.setattr(scheduler, "heartbeat_fire_claim", _heartbeat) monkeypatch.setattr(scheduler, "run_job", _run_job) monkeypatch.setattr(scheduler, "claim_dispatch", lambda job_id: True) - monkeypatch.setattr(scheduler, "mark_execution_running", lambda execution_id: None) + monkeypatch.setattr(scheduler, "mark_execution_running", lambda execution_id: {}) monkeypatch.setattr(scheduler, "finish_execution", lambda *args, **kwargs: None) save_output = MagicMock() deliver_result = MagicMock() @@ -564,7 +564,7 @@ def test_terminal_owner_cas_failure_marks_ledger_ownership_lost(monkeypatch): finish = MagicMock() monkeypatch.setattr(scheduler, "heartbeat_fire_claim", lambda *args, **kwargs: True) monkeypatch.setattr(scheduler, "claim_dispatch", lambda *_args, **_kwargs: True) - monkeypatch.setattr(scheduler, "mark_execution_running", lambda *_args: None) + monkeypatch.setattr(scheduler, "mark_execution_running", lambda *_args: {}) monkeypatch.setattr( scheduler, "run_job", diff --git a/tests/gateway/feishu_helpers.py b/tests/gateway/feishu_helpers.py index 97771daaa3..f9b7822a21 100644 --- a/tests/gateway/feishu_helpers.py +++ b/tests/gateway/feishu_helpers.py @@ -2,6 +2,7 @@ from __future__ import annotations +import asyncio import threading from types import SimpleNamespace from typing import Any, Optional @@ -59,6 +60,7 @@ def install_dedup_state(adapter: Any, seen: Optional[dict] = None) -> None: adapter._seen_message_order = list((seen or {}).keys()) adapter._dedup_cache_size = 100 adapter._dedup_lock = threading.Lock() + adapter._dedup_persist_lock = asyncio.Lock() adapter._dedup_state_path = None adapter._persist_seen_message_ids = lambda: None diff --git a/tests/gateway/test_async_media_cache.py b/tests/gateway/test_async_media_cache.py new file mode 100644 index 0000000000..1745d9fbe6 --- /dev/null +++ b/tests/gateway/test_async_media_cache.py @@ -0,0 +1,99 @@ +import asyncio +import threading +from pathlib import Path + +import pytest + +import gateway.platforms.base as base + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("async_name", "sync_name", "args"), + [ + ("cache_image_from_bytes_async", "cache_image_from_bytes", (b"data", ".png")), + ("cache_audio_from_bytes_async", "cache_audio_from_bytes", (b"data", ".ogg")), + ("cache_video_from_bytes_async", "cache_video_from_bytes", (b"data", ".mp4")), + ( + "cache_document_from_bytes_async", + "cache_document_from_bytes", + (b"data", "report.pdf"), + ), + ], +) +async def test_async_cache_wrappers_keep_event_loop_responsive( + monkeypatch, async_name, sync_name, args +): + loop_thread = threading.get_ident() + cache_started = threading.Event() + release_cache = threading.Event() + observed = {} + + def blocking_cache(*call_args): + observed["thread"] = threading.get_ident() + observed["args"] = call_args + cache_started.set() + observed["ticker_ran_during_cache"] = release_cache.wait(timeout=1) + return "cached" + + monkeypatch.setattr(base, sync_name, blocking_cache) + + async def ticker(): + while not cache_started.is_set(): + await asyncio.sleep(0) + release_cache.set() + + ticker_task = asyncio.create_task(ticker()) + result = await getattr(base, async_name)(*args) + await ticker_task + + assert result == "cached" + assert observed["args"] == args + assert observed["thread"] != loop_thread + assert observed["ticker_ran_during_cache"] is True + + +@pytest.mark.asyncio +async def test_async_cache_wrapper_propagates_validation_errors(monkeypatch): + def reject_image(data, ext): + raise ValueError("invalid image") + + monkeypatch.setattr(base, "cache_image_from_bytes", reject_image) + + with pytest.raises(ValueError, match="invalid image"): + await base.cache_image_from_bytes_async(b"not-an-image", ".png") + + +@pytest.mark.asyncio +async def test_cache_media_bytes_async_runs_off_loop_and_forwards_kwargs(monkeypatch): + loop_thread = threading.get_ident() + observed = {} + + def fake_cache_media_bytes(data, *, filename="", mime_type="", default_kind=None): + observed["thread"] = threading.get_ident() + observed["call"] = (data, filename, mime_type, default_kind) + return "cached-media" + + monkeypatch.setattr(base, "cache_media_bytes", fake_cache_media_bytes) + + result = await base.cache_media_bytes_async( + b"payload", filename="report.pdf", mime_type="application/pdf", default_kind="document" + ) + + assert result == "cached-media" + assert observed["call"] == (b"payload", "report.pdf", "application/pdf", "document") + assert observed["thread"] != loop_thread + + +@pytest.mark.asyncio +async def test_async_cache_wrapper_uses_active_profile_home(monkeypatch, tmp_path): + profile_home = tmp_path / "profile" + monkeypatch.setenv("HERMES_HOME", str(profile_home)) + + cached = await base.cache_image_from_bytes_async( + b"\x89PNG\r\n\x1a\nminimal", ".png" + ) + + cached_path = Path(cached) + assert cached_path.parent == profile_home / "cache" / "images" + assert cached_path.read_bytes() == b"\x89PNG\r\n\x1a\nminimal" diff --git a/tests/gateway/test_bluebubbles.py b/tests/gateway/test_bluebubbles.py index 695a26297d..7d95190cd0 100644 --- a/tests/gateway/test_bluebubbles.py +++ b/tests/gateway/test_bluebubbles.py @@ -259,13 +259,13 @@ class TestBlueBubblesAttachmentDownload: cached_path = None - def mock_cache_image(data, ext): + async def mock_cache_image(data, ext): nonlocal cached_path cached_path = f"/tmp/test_image{ext}" return cached_path monkeypatch.setattr( - "gateway.platforms.bluebubbles.cache_image_from_bytes", + "gateway.platforms.bluebubbles.cache_image_from_bytes_async", mock_cache_image, ) diff --git a/tests/gateway/test_buzz_adapter.py b/tests/gateway/test_buzz_adapter.py index 18587187f0..bee8af29e3 100644 --- a/tests/gateway/test_buzz_adapter.py +++ b/tests/gateway/test_buzz_adapter.py @@ -1134,8 +1134,8 @@ class TestInboundAttachments: ) monkeypatch.setattr( _buzz_mod, - "cache_media_bytes", - MagicMock(side_effect=OSError(36, "File name too long")), + "cache_media_bytes_async", + AsyncMock(side_effect=OSError(36, "File name too long")), ) adapter = _make_adapter() diff --git a/tests/gateway/test_compaction_heartbeat_gateway_filter.py b/tests/gateway/test_compaction_heartbeat_gateway_filter.py new file mode 100644 index 0000000000..1689157162 --- /dev/null +++ b/tests/gateway/test_compaction_heartbeat_gateway_filter.py @@ -0,0 +1,36 @@ +"""The compaction heartbeat (#98371) must be classified like COMPACTION_STATUS. + +Chat platforms suppress routine compression chatter unless +``compression.progress_notices`` is enabled; a heartbeat that slipped past +that gate would post a bubble per tick to Telegram/Discord. The TUI gateway, +by contrast, must keep receiving it so idle-turn watchdogs see progress. +""" + +from types import SimpleNamespace +from unittest.mock import patch + +from agent.conversation_compression import ( + COMPACTION_HEARTBEAT_STATUS, + COMPACTION_STATUS, + is_compaction_progress_status, +) +from gateway.run import _prepare_gateway_status_message + + +def _telegram(): + return SimpleNamespace(value="telegram") + + +def test_heartbeat_is_compaction_progress_for_tui_retagging(): + assert is_compaction_progress_status(COMPACTION_HEARTBEAT_STATUS) + + +def test_heartbeat_suppressed_on_chat_platforms_by_default(): + with patch("gateway.run._gateway_compression_progress_notices_enabled", return_value=False): + assert _prepare_gateway_status_message(_telegram(), "lifecycle", COMPACTION_STATUS) is None + assert _prepare_gateway_status_message(_telegram(), "lifecycle", COMPACTION_HEARTBEAT_STATUS) is None + + +def test_heartbeat_passes_when_progress_notices_enabled(): + with patch("gateway.run._gateway_compression_progress_notices_enabled", return_value=True): + assert _prepare_gateway_status_message(_telegram(), "lifecycle", COMPACTION_HEARTBEAT_STATUS) == COMPACTION_HEARTBEAT_STATUS diff --git a/tests/gateway/test_cron_delivery_housekeeping.py b/tests/gateway/test_cron_delivery_housekeeping.py new file mode 100644 index 0000000000..1228b04772 --- /dev/null +++ b/tests/gateway/test_cron_delivery_housekeeping.py @@ -0,0 +1,157 @@ +"""Gateway-independent draining of restart-safe cron deliveries.""" + +from contextlib import contextmanager +from types import SimpleNamespace + +import cron.scheduler as scheduler +import gateway.run as gateway_run + + +class _OneTickStopEvent: + def __init__(self): + self.waited = False + + def is_set(self): + return self.waited + + def wait(self, timeout=None): + self.waited = True + return True + + +def test_gateway_housekeeping_drains_cron_delivery_with_live_adapters(monkeypatch): + adapters = {"discord": object()} + loop = object() + calls = [] + monkeypatch.setattr( + scheduler, + "drain_delivery_queue", + lambda live_adapters, live_loop: calls.append((live_adapters, live_loop)), + raising=False, + ) + + gateway_run._start_gateway_housekeeping( + _OneTickStopEvent(), adapters=adapters, loop=loop, interval=0 + ) + + assert calls == [(adapters, loop)] + + +def test_gateway_housekeeping_drains_cron_delivery_without_connected_adapters(monkeypatch): + adapters = {} + loop = object() + calls = [] + monkeypatch.setattr( + scheduler, + "drain_delivery_queue", + lambda live_adapters, live_loop: calls.append((live_adapters, live_loop)), + raising=False, + ) + + gateway_run._start_gateway_housekeeping( + _OneTickStopEvent(), adapters=adapters, loop=loop, interval=0 + ) + + assert calls == [(adapters, loop)] + + +def test_multiplex_housekeeping_scopes_primary_and_drains_each_profile( + tmp_path, monkeypatch +): + root_adapters = {} + secondary_adapters = {} + runner = SimpleNamespace( + config=SimpleNamespace(multiplex_profiles=True), + adapters=root_adapters, + _profile_adapters={"secondary": secondary_adapters}, + ) + root_home = tmp_path / "root" + secondary_home = tmp_path / "secondary" + calls = [] + + monkeypatch.setattr(gateway_run, "get_hermes_home", lambda: root_home) + + monkeypatch.setattr( + gateway_run, + "_handoff_watch_scopes", + lambda _runner: [(None, None), ("secondary", secondary_home)], + ) + + @contextmanager + def fake_scope(home): + calls.append(("scope", home)) + yield + + monkeypatch.setattr(gateway_run, "_profile_runtime_scope", fake_scope) + monkeypatch.setattr( + scheduler, + "drain_delivery_queue", + lambda adapters, loop: calls.append(("drain", adapters)), + ) + + gateway_run._start_gateway_housekeeping( + _OneTickStopEvent(), + adapters=root_adapters, + loop=object(), + interval=0, + runner=runner, + ) + + assert calls == [ + ("scope", root_home), + ("drain", root_adapters), + ("scope", secondary_home), + ("drain", secondary_adapters), + ] + + +def test_multiplex_housekeeping_uses_primary_routes_for_credentialless_satellite( + tmp_path, monkeypatch +): + root_adapters = {"slack": object()} + secondary_home = tmp_path / "secondary" + runner = SimpleNamespace( + config=SimpleNamespace(multiplex_profiles=True), + adapters=root_adapters, + _profile_adapters={"secondary": {}}, + ) + calls = [] + routed = object() + + monkeypatch.setattr( + gateway_run, + "_handoff_watch_scopes", + lambda _runner: [(None, None), ("secondary", secondary_home)], + ) + + @contextmanager + def fake_scope(_home): + yield + + class FakeSharedRouteAdapters: + def __new__(cls, adapters, routes): + calls.append(("routed", adapters, routes)) + return routed + + monkeypatch.setattr(gateway_run, "_profile_runtime_scope", fake_scope) + monkeypatch.setattr(scheduler, "SharedRouteAdapters", FakeSharedRouteAdapters) + monkeypatch.setattr( + scheduler, + "_primary_profile_routes_for_current_home", + lambda: ["route-to-secondary"], + ) + monkeypatch.setattr( + scheduler, + "drain_delivery_queue", + lambda adapters, _loop: calls.append(("drain", adapters)), + ) + + gateway_run._drain_restart_safe_cron_deliveries( + root_adapters, object(), runner + ) + + assert calls == [ + ("drain", root_adapters), + ("routed", root_adapters, ["route-to-secondary"]), + ("drain", routed), + ] diff --git a/tests/gateway/test_discord_attachment_download.py b/tests/gateway/test_discord_attachment_download.py index a97632aa15..5b9575ee91 100644 --- a/tests/gateway/test_discord_attachment_download.py +++ b/tests/gateway/test_discord_attachment_download.py @@ -132,8 +132,8 @@ class TestCacheDiscordImage: att = _make_attachment_with_read(b"forbidden") with patch( - "plugins.platforms.discord.adapter.cache_image_from_bytes", - side_effect=ValueError("not a valid image"), + "plugins.platforms.discord.adapter.cache_image_from_bytes_async", + new=AsyncMock(side_effect=ValueError("not a valid image")), ), patch( "plugins.platforms.discord.adapter.cache_image_from_url", new_callable=AsyncMock, @@ -156,8 +156,8 @@ class TestCacheDiscordAudio: att = _make_attachment_with_read(_OGG_BYTES) with patch( - "plugins.platforms.discord.adapter.cache_audio_from_bytes", - return_value="/tmp/voice.ogg", + "plugins.platforms.discord.adapter.cache_audio_from_bytes_async", + new=AsyncMock(return_value="/tmp/voice.ogg"), ) as mock_bytes, patch( "plugins.platforms.discord.adapter.cache_audio_from_url", new_callable=AsyncMock, @@ -165,7 +165,7 @@ class TestCacheDiscordAudio: result = await adapter._cache_discord_audio(att, ".ogg") assert result == "/tmp/voice.ogg" - mock_bytes.assert_called_once_with(_OGG_BYTES, ext=".ogg") + mock_bytes.assert_awaited_once_with(_OGG_BYTES, ext=".ogg") mock_url.assert_not_called() @@ -215,8 +215,8 @@ class TestHandleMessageUsesAuthenticatedRead: adapter.handle_message = AsyncMock() with patch( - "plugins.platforms.discord.adapter.cache_image_from_bytes", - return_value="/tmp/img_from_read.png", + "plugins.platforms.discord.adapter.cache_image_from_bytes_async", + new=AsyncMock(return_value="/tmp/img_from_read.png"), ), patch( "plugins.platforms.discord.adapter.cache_image_from_url", new_callable=AsyncMock, diff --git a/tests/gateway/test_discord_free_response.py b/tests/gateway/test_discord_free_response.py index fdba6fea2a..fc58b982a3 100644 --- a/tests/gateway/test_discord_free_response.py +++ b/tests/gateway/test_discord_free_response.py @@ -2,7 +2,7 @@ from datetime import datetime, timezone from types import SimpleNamespace -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch import sys import pytest @@ -378,7 +378,7 @@ async def test_fetch_channel_context_skips_self_improvement_boundary_message(ada ], channel_id=123, ) - adapter._nonconversational_messages.mark_many(["9"]) + await adapter._nonconversational_messages.mark_many(["9"]) result = await adapter._fetch_channel_context(channel, before=make_message(channel=channel, content="trigger")) @@ -827,3 +827,58 @@ async def test_discord_reply_in_free_channel_triggers_backfill(adapter, monkeypa ) +class TestNonConversationalTrackerOffload: + """atomic_json_write() calls os.fsync(), which blocks until the write + reaches stable storage. mark_many() runs on the event loop from both + DiscordAdapter.send() and send_update_prompt(), so the persist step + must be offloaded to a thread — mirrors + test_directory_write_runs_off_event_loop_thread in + test_channel_directory.py for the same #83906 bug class. + """ + + @pytest.mark.asyncio + async def test_mark_many_persist_runs_off_event_loop_thread(self): + import threading + + tracker = discord_platform._DiscordNonConversationalMessageTracker() + loop_thread = threading.get_ident() + write_threads = [] + + def fake_write(path, data, *args, **kwargs): + write_threads.append(threading.get_ident()) + + with patch.object(discord_platform, "atomic_json_write", side_effect=fake_write): + await tracker.mark_many(["999"]) + + assert "999" in tracker + assert write_threads + assert all(tid != loop_thread for tid in write_threads) + + @pytest.mark.asyncio + async def test_concurrent_mark_many_persists_land_in_order(self): + """Two in-flight mark_many() calls (send() racing a history fetch) must + not let an older snapshot overwrite a newer one on disk.""" + import asyncio as _asyncio + import time + + tracker = discord_platform._DiscordNonConversationalMessageTracker() + tracker._ids = {} + writes = [] + calls = [0] + + def slow_first_write(path, data, *args, **kwargs): + idx = calls[0] + calls[0] += 1 + if idx == 0: + time.sleep(0.05) + writes.append(list(data)) + + with patch.object(discord_platform, "atomic_json_write", side_effect=slow_first_write): + first = _asyncio.create_task(tracker.mark_many(["1"])) + await _asyncio.sleep(0.005) + second = _asyncio.create_task(tracker.mark_many(["2"])) + await _asyncio.gather(first, second) + + assert sorted(writes[-1]) == ["1", "2"] + + diff --git a/tests/gateway/test_feishu.py b/tests/gateway/test_feishu.py index a923ea5f5d..2d1783e051 100644 --- a/tests/gateway/test_feishu.py +++ b/tests/gateway/test_feishu.py @@ -1105,8 +1105,8 @@ class TestAdapterBehavior(unittest.TestCase): side_effect=lambda **_kwargs: _FakeAsyncClient(), ): with patch( - "plugins.platforms.feishu.adapter.cache_document_from_bytes", - return_value="/tmp/cached-doc.bin", + "plugins.platforms.feishu.adapter.cache_document_from_bytes_async", + new=AsyncMock(return_value="/tmp/cached-doc.bin"), ): return await adapter._download_remote_document( "https://example.com/doc.bin", @@ -1186,9 +1186,9 @@ class TestAdapterBehavior(unittest.TestCase): with tempfile.TemporaryDirectory() as temp_home: with patch.dict(os.environ, {"HERMES_HOME": temp_home}, clear=False): first = FeishuAdapter(PlatformConfig()) - self.assertFalse(first._is_duplicate("om_same")) + self.assertFalse(asyncio.run(first._is_duplicate("om_same"))) second = FeishuAdapter(PlatformConfig()) - self.assertTrue(second._is_duplicate("om_same")) + self.assertTrue(asyncio.run(second._is_duplicate("om_same"))) @patch.dict(os.environ, {}, clear=True) @@ -1622,7 +1622,7 @@ class TestDedupTTL(unittest.TestCase): with patch.object(adapter, "_persist_seen_message_ids"): adapter._seen_message_ids = {"om_dup": time.time()} adapter._seen_message_order = ["om_dup"] - self.assertTrue(adapter._is_duplicate("om_dup")) + self.assertTrue(asyncio.run(adapter._is_duplicate("om_dup"))) @patch.dict(os.environ, {}, clear=True) @@ -1656,6 +1656,61 @@ class TestDedupTTL(unittest.TestCase): assert "om_bad_str" not in adapter._seen_message_ids assert "om_bad_null" not in adapter._seen_message_ids + @patch.dict(os.environ, {}, clear=True) + def test_persist_on_new_message_runs_off_event_loop_thread(self): + """atomic_json_write() calls os.fsync(), which blocks until the write + reaches stable storage. _is_duplicate() runs on the event loop for + every inbound message (_handle_message_event_data), so the persist + step must be offloaded to a thread — mirrors + test_directory_write_runs_off_event_loop_thread in + test_channel_directory.py for the same #83906 bug class.""" + import threading + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + loop_thread = threading.get_ident() + write_threads = [] + + def fake_write(path, data, *args, **kwargs): + write_threads.append(threading.get_ident()) + + with patch("plugins.platforms.feishu.adapter.atomic_json_write", side_effect=fake_write): + is_dup = asyncio.run(adapter._is_duplicate("om_new")) + + self.assertFalse(is_dup) + self.assertTrue(write_threads) + self.assertTrue(all(tid != loop_thread for tid in write_threads)) + + @patch.dict(os.environ, {}, clear=True) + def test_concurrent_dedup_persists_land_in_order(self): + """Two in-flight _is_duplicate() calls (two chats) must not let an + older seen-ids snapshot overwrite a newer one on disk.""" + from gateway.config import PlatformConfig + from plugins.platforms.feishu.adapter import FeishuAdapter + + adapter = FeishuAdapter(PlatformConfig()) + writes = [] + calls = [0] + + def slow_first_write(path, data, *args, **kwargs): + idx = calls[0] + calls[0] += 1 + if idx == 0: + time.sleep(0.05) + writes.append(sorted(data["message_ids"])) + + async def run(): + first = asyncio.create_task(adapter._is_duplicate("om_a")) + await asyncio.sleep(0.005) + second = asyncio.create_task(adapter._is_duplicate("om_b")) + await asyncio.gather(first, second) + + with patch("plugins.platforms.feishu.adapter.atomic_json_write", side_effect=slow_first_write): + asyncio.run(run()) + + self.assertEqual(writes[-1], ["om_a", "om_b"]) + class TestGroupMentionAtAll(unittest.TestCase): """Tests for @_all (Feishu @everyone) group mention routing.""" diff --git a/tests/gateway/test_feishu_meeting_invite.py b/tests/gateway/test_feishu_meeting_invite.py index 47ce7472d0..d8e4725a64 100644 --- a/tests/gateway/test_feishu_meeting_invite.py +++ b/tests/gateway/test_feishu_meeting_invite.py @@ -76,7 +76,7 @@ class _Adapter: self.dedup_keys = [] self.profile_requests = [] - def _is_duplicate(self, key): + async def _is_duplicate(self, key): self.dedup_keys.append(key) return self.duplicate @@ -166,6 +166,19 @@ class TestMeetingInviteHandler(unittest.TestCase): self.assertIn("You have been invited to join a meeting: 赵磊的视频会议", event.text) self.assertNotIn("{'open_id'", event.text) + def test_duplicate_event_is_dropped_without_routing(self): + """_is_duplicate() is async on the real FeishuAdapter (dedup persist + is offloaded off the event loop); the dedup check here must await + it — a missing await would leave an un-awaited coroutine, which is + always truthy, and drop every event as a false duplicate.""" + adapter = _Adapter(duplicate=True) + + self._run(handle_meeting_invited_event(adapter, _make_payload())) + + self.assertEqual(adapter.dedup_keys, ["vc_invite:evt_1"]) + self.assertEqual(adapter.events, []) + self.assertEqual(adapter.profile_requests, []) + class TestMeetingInviteSendRouting(unittest.TestCase): def _run(self, coro): diff --git a/tests/gateway/test_google_chat.py b/tests/gateway/test_google_chat.py index d3a05ea00c..c9f1f1584f 100644 --- a/tests/gateway/test_google_chat.py +++ b/tests/gateway/test_google_chat.py @@ -1336,9 +1336,9 @@ class TestAttachmentSSRFGuard: monkeypatch.setattr(asyncio, "to_thread", _fake_to_thread) from plugins.platforms.google_chat import adapter as gc_mod monkeypatch.setattr( - gc_mod, "cache_document_from_bytes", - lambda data, ext=None, filename=None: str(tmp_path / "out.pdf"), - raising=False, + gc_mod, + "cache_document_from_bytes_async", + AsyncMock(return_value=str(tmp_path / "out.pdf")), ) path, mime = await adapter._download_attachment(attachment) diff --git a/tests/gateway/test_line_plugin.py b/tests/gateway/test_line_plugin.py index e59bd8286e..5a6386a7be 100644 --- a/tests/gateway/test_line_plugin.py +++ b/tests/gateway/test_line_plugin.py @@ -205,10 +205,14 @@ class TestInboundMedia: return adapter.handle_message.await_args.args[0] def test_image_message_uses_photo_type_and_image_mime(self, adapter): - with patch.object(_line, "cache_image_from_bytes", return_value="/cache/image.jpg") as cache: + with patch.object( + _line, + "cache_image_from_bytes_async", + new=AsyncMock(return_value="/cache/image.jpg"), + ) as cache: asyncio.run(adapter._handle_message_event(self._event("image"))) - cache.assert_called_once_with(b"line-bytes", ext=".jpg") + cache.assert_awaited_once_with(b"line-bytes", ext=".jpg") event = self._captured_event(adapter) assert event.message_type is _line.MessageType.PHOTO assert event.media_urls == ["/cache/image.jpg"] diff --git a/tests/gateway/test_loop_command.py b/tests/gateway/test_loop_command.py index 7e8caee8fc..19c9c88520 100644 --- a/tests/gateway/test_loop_command.py +++ b/tests/gateway/test_loop_command.py @@ -1,7 +1,10 @@ """Gateway /loop command tests — dispatch, routing capture, mid-run guard.""" +import asyncio import logging +import threading import time +from types import SimpleNamespace from unittest.mock import AsyncMock, Mock import pytest @@ -255,3 +258,147 @@ async def test_post_turn_session_resolution_failure_is_logged(loop_env, caplog): ) assert "post-turn session resolution failed: store unavailable" in caplog.text + + +@pytest.mark.asyncio +async def test_loop_wakeup_watcher_keeps_event_loop_responsive_under_writer_lock(loop_env): + """The wakeup scan's SessionDB calls (list_active_loops / fire_tick / + complete_tick) must run off the loop thread. A slow writer holding the + SessionDB writer lock used to block the whole gateway event loop for the + duration of the hold (#92413).""" + runner = _make_runner() + runner._running = True + runner._running_agents = {} + runner.adapters = {} + + # Persist an active loop that is due now, routed to a platform with no + # adapter (so the scan exits after list_active_loops(), before fire_tick). + await GatewayRunner._handle_loop_command(runner, _make_event("/loop 5m poll CI")) + state = loops.load_loop("sid-gateway-loop") + state.next_due_at = time.time() - 1 + loops.save_loop("sid-gateway-loop", state) + + db = loops._get_session_db() + hold_s = 0.6 + released = threading.Event() + + def _hold_writer_lock(): + with db._lock: + time.sleep(hold_s) + released.set() + + # Force the read path onto the writer lock (non-WAL degradation) so the + # test binds the offload regardless of the host SQLite's WAL support. + db._wal_active = False + + # One scan only: patch asyncio.sleep inside the watcher to stop the loop + # after the first iteration. + orig_sleep = asyncio.sleep + calls = {"n": 0} + + async def _one_pass_sleep(delay): + calls["n"] += 1 + if calls["n"] >= 2: # first call is the 5s connect grace, second ends the scan + runner._running = False + return await orig_sleep(0) + + holder = threading.Thread(target=_hold_writer_lock) + with pytest.MonkeyPatch.context() as mp: + mp.setattr(asyncio, "sleep", _one_pass_sleep) + holder.start() + time.sleep(0.05) # ensure the lock is held before the scan starts + watcher = asyncio.ensure_future(GatewayRunner._loop_wakeup_watcher(runner, interval=0)) + + # Heartbeat coroutine: measures the longest gap between loop turns + # while the watcher is (supposedly) blocked in the executor. + gaps = [] + last = time.monotonic() + while not watcher.done(): + await orig_sleep(0.01) + now = time.monotonic() + gaps.append(now - last) + last = now + await watcher + holder.join() + + assert released.is_set() + # If the DB call ran on the loop thread, one heartbeat gap would be ~hold_s. + assert max(gaps) < hold_s / 2, f"event loop stalled for {max(gaps):.3f}s" + + +@pytest.mark.asyncio +async def test_loop_wakeup_watcher_runs_every_sessiondb_call_off_loop_thread(loop_env): + """Full wakeup path (slash-command loop, adapter present): list_active_loops, + fire_tick and complete_tick must each execute on an executor thread, never + on the event-loop thread (#92413).""" + runner = _make_runner() + runner._running = True + runner._running_agents = {} + + class _Adapter: + handled = [] + + async def handle_message(self, event): + self.handled.append(event.text) + + runner.adapters = {Platform.DISCORD: _Adapter()} + runner._build_process_event_source = lambda evt: SimpleNamespace( + platform=Platform.DISCORD, chat_id=evt["chat_id"], chat_type=evt["chat_type"], + thread_id=evt["thread_id"] or None, user_id=evt["user_id"], user_name=evt["user_name"], + ) + runner._session_key_for_source = lambda source: "agent:main:discord:channel:chat-loop" + + # A slash-command loop that is due now. + await GatewayRunner._handle_loop_command(runner, _make_event("/loop 5m /status")) + state = loops.load_loop("sid-gateway-loop") + state.next_due_at = time.time() - 1 + loops.save_loop("sid-gateway-loop", state) + + loop_thread = threading.current_thread() + on_loop_calls = [] + seen_calls = [] + + def _record(name): + # Positive count guards against a future import hoist in run.py that + # would bypass these wrappers and leave the off-loop assertion vacuous. + seen_calls.append(name) + if threading.current_thread() is loop_thread: + on_loop_calls.append(name) + + real_list = loops.list_active_loops + real_fire = loops.LoopManager.fire_tick + real_complete = loops.LoopManager.complete_tick + + def _list_active_loops(*a, **k): + _record("list_active_loops") + return real_list(*a, **k) + + def _fire_tick(self): + _record("fire_tick") + return real_fire(self) + + def _complete_tick(self, last_response): + _record("complete_tick") + return real_complete(self, last_response) + + orig_sleep = asyncio.sleep + calls = {"n": 0} + + async def _one_pass_sleep(delay): + calls["n"] += 1 + if calls["n"] >= 2: + runner._running = False + return await orig_sleep(0) + + with pytest.MonkeyPatch.context() as mp: + mp.setattr(loops, "list_active_loops", _list_active_loops) + mp.setattr(loops.LoopManager, "fire_tick", _fire_tick) + mp.setattr(loops.LoopManager, "complete_tick", _complete_tick) + mp.setattr(asyncio, "sleep", _one_pass_sleep) + await GatewayRunner._loop_wakeup_watcher(runner, interval=0) + + assert _Adapter.handled == ["/status"], _Adapter.handled + assert seen_calls == ["list_active_loops", "fire_tick", "complete_tick"], seen_calls + assert on_loop_calls == [], f"SessionDB calls ran on the event-loop thread: {on_loop_calls}" + # complete_tick ran (slash-command loops complete immediately). + assert loops.load_loop("sid-gateway-loop").ticks_fired == 1 diff --git a/tests/gateway/test_platform_reconnect.py b/tests/gateway/test_platform_reconnect.py index 75029072c8..e923fac809 100644 --- a/tests/gateway/test_platform_reconnect.py +++ b/tests/gateway/test_platform_reconnect.py @@ -199,6 +199,55 @@ class TestPlatformReconnectWatcher: ) assert Platform.TELEGRAM in runner.adapters + @pytest.mark.asyncio + @pytest.mark.parametrize("degraded", [False, True]) + async def test_reconnect_stamp_honours_adapter_send_path_degraded(self, degraded): + """connect() returning True is not proof the receive path is live: + Telegram's degraded reconnect returns True while its own ladder + retries. The watcher's status stamp must publish what the adapter + reports, not an unconditional "connected" (#101391).""" + runner = _make_runner() + runner._sync_voice_mode_state_to_adapter = MagicMock() + runner._update_platform_runtime_status = MagicMock() + runner._failed_platforms[Platform.TELEGRAM] = { + "config": PlatformConfig(enabled=True, token="test"), + "attempts": 1, + "next_retry": time.monotonic() - 1, + } + + class _DegradableAdapter(StubAdapter): + @property + def send_path_degraded(self) -> bool: + return degraded + + adapter = _DegradableAdapter(succeed=True) + real_sleep = asyncio.sleep + + with patch.object(runner, "_create_adapter", return_value=adapter): + with patch("gateway.run.build_channel_directory", create=True): + runner._running = True + call_count = 0 + + async def fake_sleep(n): + nonlocal call_count + call_count += 1 + if call_count > 1: + runner._running = False + await real_sleep(0) + + with patch("asyncio.sleep", side_effect=fake_sleep): + await runner._platform_reconnect_watcher() + + stamps = [ + c.kwargs for c in runner._update_platform_runtime_status.call_args_list + if c.args and c.args[0] == Platform.TELEGRAM.value + ] + assert stamps, "watcher never stamped telegram" + final = stamps[-1] + assert final["platform_state"] == ("retrying" if degraded else "connected") + assert final["error_message"] == (adapter.DEGRADED_STATUS_MESSAGE if degraded else None) + assert final["retrying_since"] is None + @pytest.mark.asyncio async def test_cold_connect_defaults_to_is_reconnect_false(self): """The cold-start connect path (_connect_adapter_with_timeout with no diff --git a/tests/gateway/test_signal.py b/tests/gateway/test_signal.py index 078d787d43..fb668dc34d 100644 --- a/tests/gateway/test_signal.py +++ b/tests/gateway/test_signal.py @@ -270,7 +270,10 @@ class TestSignalAttachmentFetch: adapter._rpc, captured = _stub_rpc({"data": b64_data}) - with patch("gateway.platforms.signal.cache_image_from_bytes", return_value="/tmp/test.png"): + with patch( + "gateway.platforms.signal.cache_image_from_bytes_async", + new=AsyncMock(return_value="/tmp/test.png"), + ): await adapter._fetch_attachment("attachment-123") call = captured[0] @@ -1327,7 +1330,10 @@ class TestSignalContentlessEnvelope: b64_data = base64.b64encode(png_data).decode() adapter._rpc, _ = _stub_rpc({"data": b64_data}) - with patch("gateway.platforms.signal.cache_image_from_bytes", return_value="/tmp/img.png"): + with patch( + "gateway.platforms.signal.cache_image_from_bytes_async", + new=AsyncMock(return_value="/tmp/img.png"), + ): await adapter._handle_envelope({ "envelope": { "sourceNumber": "+155****9999", diff --git a/tests/gateway/test_teams.py b/tests/gateway/test_teams.py index dc3b489e2f..fe7ff4c8a7 100644 --- a/tests/gateway/test_teams.py +++ b/tests/gateway/test_teams.py @@ -702,12 +702,12 @@ class TestTeamsBotFrameworkAttachments: adapter._fetch_attachment_bytes = AsyncMock(return_value=b"\x89PNG fake") adapter._get_botframework_token = AsyncMock(return_value="tok") - def fake_cache_media_bytes(data, **kwargs): + async def fake_cache_media_bytes(data, **kwargs): return SimpleNamespace( path="/tmp/img.png", media_type="image/png", kind="image" ) - with patch.object(_teams_mod, "cache_media_bytes", fake_cache_media_bytes): + with patch.object(_teams_mod, "cache_media_bytes_async", fake_cache_media_bytes): activity = self._make_activity([self._bf_image_attachment()]) await adapter._on_message(self._make_ctx(activity)) @@ -943,7 +943,10 @@ class TestTeamsBotFrameworkAttachments: adapter = self._make_adapter() adapter._fetch_attachment_bytes = AsyncMock(return_value=b"error page") - with patch.object(_teams_mod, "cache_media_bytes", lambda *a, **kw: None): + async def _no_media(*a, **kw): + return None + + with patch.object(_teams_mod, "cache_media_bytes_async", _no_media): with patch.object(_teams_mod.logger, "warning") as warn: activity = self._make_activity([self._bf_image_attachment()]) await adapter._on_message(self._make_ctx(activity)) diff --git a/tests/gateway/test_telegram_documents.py b/tests/gateway/test_telegram_documents.py index a82003ce9d..fd3b406a08 100644 --- a/tests/gateway/test_telegram_documents.py +++ b/tests/gateway/test_telegram_documents.py @@ -308,7 +308,10 @@ class TestMediaGroups: msg1 = _make_message(caption="two images", photo=[first_photo]) msg2 = _make_message(photo=[second_photo]) - with patch("plugins.platforms.telegram.adapter.cache_image_from_bytes", side_effect=["/tmp/burst-one.jpg", "/tmp/burst-two.jpg"]): + with patch( + "plugins.platforms.telegram.adapter.cache_image_from_bytes_async", + new=AsyncMock(side_effect=["/tmp/burst-one.jpg", "/tmp/burst-two.jpg"]), + ): await adapter._handle_media_message(_make_update(msg1), MagicMock()) await adapter._handle_media_message(_make_update(msg2), MagicMock()) assert adapter.handle_message.await_count == 0 diff --git a/tests/gateway/test_telegram_send_path_health.py b/tests/gateway/test_telegram_send_path_health.py index a16faa4ecd..91e0a01123 100644 --- a/tests/gateway/test_telegram_send_path_health.py +++ b/tests/gateway/test_telegram_send_path_health.py @@ -76,3 +76,110 @@ async def test_send_short_flood_still_retries_inline(monkeypatch): sleep.assert_awaited_once_with(2.0) +def test_mark_connected_publishes_connected_when_healthy(): + """A normal connect (never degraded) still publishes platform_state=connected.""" + adapter = _make_adapter() + adapter._send_path_degraded = False + + with patch.object(adapter, "_write_runtime_status_safe") as write_status: + adapter._mark_connected() + + write_status.assert_called_once() + _, kwargs = write_status.call_args + assert kwargs["platform_state"] == "connected" + + +def test_mark_connected_publishes_retrying_when_send_path_degraded(): + """connect() can return True while polling never proved a first getUpdates + round-trip (the degraded branch, or a reconnect where require_progress is + skipped). _mark_connected() must not publish "connected" for that case -- + it is indistinguishable from a healthy adapter to anything reading + gateway_state.json (#101391).""" + adapter = _make_adapter() + adapter._send_path_degraded = True + + with patch.object(adapter, "_write_runtime_status_safe") as write_status: + adapter._mark_connected() + + write_status.assert_called_once() + _, kwargs = write_status.call_args + assert kwargs["platform_state"] == "retrying" + + +def test_record_polling_progress_republishes_connected_after_degraded_connect(): + """Once getUpdates actually proves a round-trip after a degraded connect, + the previously-published "retrying" status must be corrected back to + "connected" -- otherwise it stays wedged until the next disconnect.""" + adapter = _make_adapter() + generation, _event = adapter._begin_polling_generation() + # Simulate connect() having already run and published the degraded state. + adapter._running = True + + with patch.object(adapter, "_write_runtime_status_safe") as write_status: + adapter._record_polling_progress(generation) + + write_status.assert_called_once() + _, kwargs = write_status.call_args + assert kwargs["platform_state"] == "connected" + assert adapter._send_path_degraded is False + + +def test_mid_session_polling_death_publishes_retrying_while_running(): + """#101391's measured incident: a HEALTHY connect published "connected", + then getUpdates silently died mid-session and nothing republished for 11h. + The recovery ladder's entry point must flip the file to "retrying".""" + adapter = _make_adapter() + adapter._running = True + adapter._send_path_degraded = False + adapter._polling_error_task = None + + class _Loop: + def create_task(self, coro): + coro.close() + return MagicMock(done=lambda: False) + + with patch.object(adapter, "_write_runtime_status_safe") as write_status, \ + patch("asyncio.get_running_loop", return_value=_Loop()): + adapter._schedule_polling_recovery(RuntimeError("boom"), reason="heartbeat probe") + + assert adapter._send_path_degraded is True + write_status.assert_called_once() + _, kwargs = write_status.call_args + assert kwargs["platform_state"] == "retrying" + assert kwargs["error_message"] == TelegramAdapter.DEGRADED_STATUS_MESSAGE + + +def test_polling_death_before_connect_does_not_publish(): + """Not yet running (cold connect still in progress): connect()'s own + _mark_connected publishes; the recovery path must not write early.""" + adapter = _make_adapter() + adapter._running = False + adapter._polling_error_task = None + + class _Loop: + def create_task(self, coro): + coro.close() + return MagicMock(done=lambda: False) + + with patch.object(adapter, "_write_runtime_status_safe") as write_status, \ + patch("asyncio.get_running_loop", return_value=_Loop()): + adapter._schedule_polling_recovery(RuntimeError("boom"), reason="polling bootstrap") + + write_status.assert_not_called() + + +@pytest.mark.parametrize("running, fatal", [(False, False), (True, True)]) +def test_record_polling_progress_does_not_flip_when_not_running_or_fatal(running, fatal): + """Cold connect: progress arrives while _running is still False -- the + connect path publishes, not the flip. Fatal: never overwrite "fatal".""" + adapter = _make_adapter() + generation, _event = adapter._begin_polling_generation() + adapter._running = running + if fatal: + adapter._fatal_error_message = "dead" + + with patch.object(adapter, "_write_runtime_status_safe") as write_status: + adapter._record_polling_progress(generation) + + write_status.assert_not_called() + assert adapter._send_path_degraded is False diff --git a/tests/gateway/test_transcript_read_failure_100788.py b/tests/gateway/test_transcript_read_failure_100788.py new file mode 100644 index 0000000000..b4dd6c3708 --- /dev/null +++ b/tests/gateway/test_transcript_read_failure_100788.py @@ -0,0 +1,103 @@ +"""A failed transcript read must not masquerade as an empty history (#100788). + +The gateway restore path (``_handle_message``) already fails closed on +current main (#100910). This file covers the surviving half of PR #100887: +the slash-command handlers, which used to let ``TranscriptReadError`` +propagate into the dispatch wrapper and reply with nothing at all. + +Incident shape: a malformed ``state.db`` made every +``SessionStore.load_transcript`` raise; the except-block swallowed it and +returned ``[]``. Restore then rebuilt the turn from "no history", so a +long-running chat silently restarted as a brand-new conversation and the +model happily answered as if nothing had ever been discussed. + +Two guarantees under test: + A. ``load_transcript`` raises ``TranscriptReadError`` on a read failure, + while a genuinely empty session still returns ``[]``. + B. Slash-command handlers that read the transcript reply with + ``HISTORY_UNREADABLE`` instead of raising into the dispatch wrapper + (which logs and sends nothing). + +Offline: SQLite on tmp_path only, no network. +""" + +import sqlite3 + +import pytest + +from gateway.config import GatewayConfig +from gateway.session import SessionStore, TranscriptReadError + + +@pytest.fixture +def store(tmp_path): + return SessionStore(sessions_dir=tmp_path / "gw", config=GatewayConfig()) + + +# -------------------------------------------------------------------------- +# A. read failure != empty transcript (landed on main via #100910; kept as +# the contract the slash-command handlers below rely on) +# -------------------------------------------------------------------------- + + +class TestLoadTranscriptReadFailure: + def test_read_failure_raises_instead_of_returning_empty(self, store, monkeypatch): + db = store._db + assert db is not None + db.create_session("s1", "telegram", session_key="telegram:1") + db.append_message("s1", "user", "the conversation we must not forget") + + boom = sqlite3.DatabaseError("database disk image is malformed") + + def _raise(*_args, **_kwargs): + raise boom + + monkeypatch.setattr(db, "get_messages_as_conversation", _raise) + + with pytest.raises(TranscriptReadError) as excinfo: + store.load_transcript("s1") + + assert excinfo.value.session_id == "s1" + assert excinfo.value.__cause__ is boom + + def test_genuinely_empty_session_still_returns_empty_list(self, store): + db = store._db + assert db is not None + db.create_session("s2", "telegram", session_key="telegram:2") + + assert store.load_transcript("s2") == [] + + def test_no_db_still_returns_empty_list(self, store): + # "No DB for this session" really is an empty transcript, not a + # failure — that path must keep its [] contract. + store._db = None + assert store.load_transcript("nope") == [] + + +# -------------------------------------------------------------------------- +# B. slash-command handlers surface the failure instead of dying silently. +# Before: the handler raised, base.py's dispatch wrapper logged +# "Command '/x' dispatch failed" and the user got NO reply at all. +# -------------------------------------------------------------------------- + + +class TestSlashCommandsOnUnreadableTranscript: + def test_history_unreadable_text_is_explicit(self): + from gateway.slash_commands import HISTORY_UNREADABLE + + assert "unreadable" in HISTORY_UNREADABLE + assert "not a new conversation" in HISTORY_UNREADABLE + + def test_every_transcript_reading_handler_catches_the_error(self): + """No `await ...load_transcript(` in the mixin may be left uncaught.""" + import inspect + import re + + from gateway import slash_commands as sc + + src = inspect.getsource(sc) + # Each awaited load_transcript must sit inside a try: whose handlers + # include TranscriptReadError within the following ~6 lines. + for m in re.finditer(r"await self\.async_session_store\.load_transcript\(", src): + window = src[m.end() : m.end() + 400] + assert "except TranscriptReadError" in window, src[m.start() - 200 : m.end() + 100] diff --git a/tests/gateway/test_weixin.py b/tests/gateway/test_weixin.py index dbe93c72b2..5329c39d28 100644 --- a/tests/gateway/test_weixin.py +++ b/tests/gateway/test_weixin.py @@ -167,6 +167,57 @@ class TestWeixinStatePersistence: assert json.loads(account_path.read_text(encoding="utf-8")) == original + @pytest.mark.asyncio + async def test_context_token_persist_runs_off_event_loop_thread(self, tmp_path): + """atomic_json_write() calls os.fsync(), which blocks until the write + reaches stable storage. ContextTokenStore.set() runs on the event + loop for every inbound message carrying a context_token + (_process_message), so the persist step must be offloaded to a + thread — mirrors test_directory_write_runs_off_event_loop_thread in + test_channel_directory.py for the same #83906 bug class.""" + import threading + + store = ContextTokenStore(str(tmp_path)) + loop_thread = threading.get_ident() + write_threads = [] + + def fake_write(path, data, *args, **kwargs): + write_threads.append(threading.get_ident()) + + with patch("gateway.platforms.weixin.atomic_json_write", side_effect=fake_write): + await store.set("acct-1", "user-1", "ctx-token-abc") + + assert store.get("acct-1", "user-1") == "ctx-token-abc" + assert write_threads + assert all(tid != loop_thread for tid in write_threads) + + @pytest.mark.asyncio + async def test_concurrent_context_token_persists_land_in_order(self, tmp_path): + """Two in-flight set() calls (two concurrent inbound messages) must not + let an older snapshot overwrite a newer one on disk. Without + serialization the first (slow) flush lands last and drops user-2.""" + import asyncio as _asyncio + import time + + store = ContextTokenStore(str(tmp_path)) + writes = [] + calls = [0] + + def slow_first_write(path, data, *args, **kwargs): + idx = calls[0] + calls[0] += 1 + if idx == 0: + time.sleep(0.05) + writes.append(dict(data)) + + with patch("gateway.platforms.weixin.atomic_json_write", side_effect=slow_first_write): + first = _asyncio.create_task(store.set("acct-1", "user-1", "t1")) + await _asyncio.sleep(0.005) + second = _asyncio.create_task(store.set("acct-1", "user-2", "t2")) + await _asyncio.gather(first, second) + + assert writes[-1] == {"user-1": "t1", "user-2": "t2"} + class TestWeixinQrLogin: @pytest.mark.asyncio @@ -685,8 +736,8 @@ class TestWeixinVoiceAlwaysDownloaded: adapter._poll_session = Mock() fake_audio_bytes = b"\\x00\\x01\\x02FAKE_SILK" - monkeypatch.setattr(weixin, "cache_audio_from_bytes", - lambda data, ext: str(tmp_path / f"voice.{ext.lstrip('.')}")) + monkeypatch.setattr(weixin, "cache_audio_from_bytes_async", + AsyncMock(side_effect=lambda data, ext: str(tmp_path / f"voice.{ext.lstrip('.')}"))) async def _fake_download(session, *, cdn_base_url, encrypted_query_param, aes_key_b64, full_url, timeout_seconds): @@ -739,8 +790,8 @@ class TestWeixinVoiceAlwaysDownloaded: adapter._cdn_base_url = "https://example.invalid" adapter._poll_session = Mock() - monkeypatch.setattr(weixin, "cache_audio_from_bytes", - lambda data, ext: str(tmp_path / f"voice.{ext.lstrip('.')}")) + monkeypatch.setattr(weixin, "cache_audio_from_bytes_async", + AsyncMock(side_effect=lambda data, ext: str(tmp_path / f"voice.{ext.lstrip('.')}"))) async def _fake_download(session, *, cdn_base_url, encrypted_query_param, aes_key_b64, full_url, timeout_seconds): @@ -803,8 +854,8 @@ class TestWeixinVoiceGatewayHandoff: adapter._token = None adapter._cdn_base_url = "https://example.invalid" - monkeypatch.setattr(weixin, "cache_audio_from_bytes", - lambda data, ext: str(tmp_path / f"voice.{ext.lstrip('.')}")) + monkeypatch.setattr(weixin, "cache_audio_from_bytes_async", + AsyncMock(side_effect=lambda data, ext: str(tmp_path / f"voice.{ext.lstrip('.')}"))) async def _fake_download(*a, **k): return b"\x00\x01FAKE_SILK" monkeypatch.setattr(weixin, "_download_and_decrypt_media", _fake_download) diff --git a/tests/hermes_cli/test_active_sessions.py b/tests/hermes_cli/test_active_sessions.py index 5741aed178..0fe223101f 100644 --- a/tests/hermes_cli/test_active_sessions.py +++ b/tests/hermes_cli/test_active_sessions.py @@ -12,6 +12,17 @@ import pytest from hermes_cli import active_sessions + +def _backdate_leases(*homes, age_seconds=600.0): + """Age every lease in the given registries past the self-orphan grace.""" + for home in homes: + state_path = active_sessions._state_path(home) + entries = active_sessions._read_entries(state_path) + for entry in entries: + entry["started_at"] = time.time() - age_seconds + active_sessions._write_entries(state_path, entries) + + def test_resolve_max_concurrent_sessions_values(caplog): assert active_sessions.resolve_max_concurrent_sessions({}) is None assert active_sessions.resolve_max_concurrent_sessions({"max_concurrent_sessions": None}) is None @@ -164,6 +175,7 @@ def test_release_orphaned_leases_reclaims_only_unowned_own_pid_entries(tmp_path, + [{"lease_id": "elsewhere", "session_id": "other", "surface": "cli", "pid": os.getpid() }], ) + _backdate_leases(tmp_path / ".hermes") assert active_sessions.release_orphaned_leases({kept.lease_id, "elsewhere"}) == 1 assert sorted( entry["session_id"] @@ -172,6 +184,47 @@ def test_release_orphaned_leases_reclaims_only_unowned_own_pid_entries(tmp_path, assert orphan is not None +def test_release_orphaned_leases_sweeps_profile_runtime_registries( + tmp_path, monkeypatch +): + root = tmp_path / "hermes" + profile = root / "profiles" / "worker" + profile.mkdir(parents=True) + monkeypatch.setenv("HERMES_HOME", str(root)) + + root_lease, root_error = active_sessions.try_acquire_active_session( + session_id="root-orphan", surface="desktop", config={}, registry_home=root + ) + profile_lease, profile_error = active_sessions.try_acquire_active_session( + session_id="profile-orphan", + surface="desktop", + config={}, + registry_home=profile, + ) + assert root_lease is not None and root_error is None + assert profile_lease is not None and profile_error is None + + # A lease written seconds ago is never an orphan: a sibling finalize that + # snapshotted its live ids before this acquire must not reap it (#101415). + assert active_sessions.release_orphaned_leases(set()) == 0 + _backdate_leases(root, profile) + assert active_sessions.release_orphaned_leases(set()) == 2 + assert active_sessions.active_session_registry_snapshot(root) == [] + assert active_sessions.active_session_registry_snapshot(profile) == [] + + +def test_drop_self_orphans_spares_foreign_and_vouched_leases(): + own = os.getpid() + entries = [ + {"lease_id": "orphan", "pid": own}, + {"lease_id": "live", "pid": own}, + {"lease_id": "foreign", "pid": own + 1}, + ] + + assert active_sessions._drop_self_orphans(entries, None) == entries + assert active_sessions._drop_self_orphans(entries, {"live"}) == entries[1:] + + def test_release_under_profile_home_override_targets_acquisition_registry( tmp_path, monkeypatch ): @@ -566,3 +619,30 @@ def test_release_wins_against_transfer_waiting_on_same_lease_lock( assert lease.released is True assert active_sessions.active_session_registry_snapshot() == [] + + +def test_liveness_guard_keeps_a_just_acquired_own_lease_it_cannot_vouch_for( + tmp_path, monkeypatch +): + """Race in #101415's fix: the finalizing session snapshots its live lease + ids, then a sibling session acquires a lease before the registry lock is + taken. That lease is absent from the snapshot but is not an orphan.""" + home = tmp_path / ".hermes" + monkeypatch.setenv("HERMES_HOME", str(home)) + fresh, error = active_sessions.try_acquire_active_session( + session_id="fresh", surface="desktop", config={}, registry_home=home + ) + assert fresh is not None and error is None + + with active_sessions.active_session_liveness_guard( + "fresh", registry_home=home, own_live_lease_ids=set() + ) as active: + assert active is True + assert [e["lease_id"] for e in active_sessions.active_session_registry_snapshot(home)] == [fresh.lease_id] + + _backdate_leases(home) + with active_sessions.active_session_liveness_guard( + "fresh", registry_home=home, own_live_lease_ids=set() + ) as active: + assert active is False + assert active_sessions.active_session_registry_snapshot(home) == [] diff --git a/tests/hermes_cli/test_backup.py b/tests/hermes_cli/test_backup.py index aa4c80056b..459c5ec8b1 100644 --- a/tests/hermes_cli/test_backup.py +++ b/tests/hermes_cli/test_backup.py @@ -1336,6 +1336,32 @@ class TestQuickSnapshot: assert "state.db" not in data.get("files", {}) assert "state.db" in data.get("failed_dbs", []) + def test_restore_refused_db_is_not_counted(self, hermes_home, monkeypatch): + """A refused live-safe restore (holder detected, backup leg failed) must + not be counted as a restored file — `hermes import` reports it, and + /snapshot restore must not claim success for that file either.""" + import hermes_cli.backup as backup_mod + from hermes_cli.backup import create_quick_snapshot, restore_quick_snapshot + + snap_id = create_quick_snapshot(hermes_home=hermes_home) + monkeypatch.setattr(backup_mod, "_safe_restore_db", lambda src, dst: False) + restored_log: list[str] = [] + real_info = backup_mod.logger.info + monkeypatch.setattr( + backup_mod.logger, "info", + lambda msg, *a, **kw: restored_log.append(msg % a if a else msg) or real_info(msg, *a, **kw), + ) + + restore_quick_snapshot(snap_id, hermes_home=hermes_home) + + manifest = json.loads( + (backup_mod._quick_snapshot_root(hermes_home) / snap_id / "manifest.json").read_text() + ) + non_db = [rel for rel in manifest.get("files", {}) if not rel.endswith(".db")] + summary = [line for line in restored_log if line.startswith("Restored ")] + assert summary, restored_log + assert summary[-1].startswith(f"Restored {len(non_db)} files"), summary[-1] + def test_restore_state_db_live_connection(self, hermes_home): """Restoring state.db must update data visible through a live connection. @@ -2225,3 +2251,175 @@ class TestImportHonorsHermesHomeOverride: backup_mod.run_import(args) assert calls and calls[0].get("context") == "import" + + +# --------------------------------------------------------------------------- +# Live session database import (issue #100960) +# --------------------------------------------------------------------------- + +def _write_session_db(path: Path, sessions: int, messages_per_session: int) -> None: + """Create a minimal Hermes-shaped session database at *path*.""" + conn = sqlite3.connect(str(path)) + try: + conn.execute( + "CREATE TABLE IF NOT EXISTS sessions " + "(session_id TEXT PRIMARY KEY, message_count INTEGER)" + ) + conn.execute( + "CREATE TABLE IF NOT EXISTS messages " + "(id INTEGER PRIMARY KEY, session_id TEXT, content TEXT)" + ) + for s in range(sessions): + sid = f"sess-{s}" + conn.execute( + "INSERT INTO sessions VALUES (?, ?)", (sid, messages_per_session) + ) + for m in range(messages_per_session): + conn.execute( + "INSERT INTO messages (session_id, content) VALUES (?, ?)", + (sid, f"{sid}-msg-{m}"), + ) + conn.commit() + finally: + conn.close() + + +class TestImportLiveSessionDatabase: + """`hermes import` must not swap the inode of a database Hermes holds open. + + Publishing state.db with a rename leaves any live gateway/dashboard/WebUI + connection reading and writing the unlinked inode, so its sessions vanish + from the database everyone else opens and nothing is logged (#100960). + """ + + def _zip_with_db(self, zip_path: Path, db_path: Path) -> None: + with zipfile.ZipFile(zip_path, "w") as zf: + zf.write(db_path, "state.db") + + def _prepare(self, tmp_path, monkeypatch, live=(3, 4), backup=(2, 2)): + home = tmp_path / ".hermes" + home.mkdir() + monkeypatch.setenv("HERMES_HOME", str(home)) + monkeypatch.setattr(Path, "home", lambda: tmp_path) + + live_db = home / "state.db" + _write_session_db(live_db, *live) + + staged = tmp_path / "backup-state.db" + _write_session_db(staged, *backup) + zip_path = tmp_path / "backup.zip" + self._zip_with_db(zip_path, staged) + return home, live_db, zip_path + + def test_live_holder_sees_imported_rows(self, tmp_path, monkeypatch): + """A connection open across the import converges on the imported data.""" + from hermes_cli.backup import run_import + + home, live_db, zip_path = self._prepare(tmp_path, monkeypatch) + + holder = sqlite3.connect(str(live_db)) + # Read first so the connection has cached pages of the pre-import file. + assert holder.execute("SELECT COUNT(*) FROM messages").fetchone()[0] == 12 + inode_before = os.stat(live_db).st_ino + + try: + run_import(Namespace(zipfile=str(zip_path), force=True)) + assert holder.execute("SELECT COUNT(*) FROM messages").fetchone()[0] == 4 + finally: + holder.close() + + assert os.stat(live_db).st_ino == inode_before + assert _count_rows(live_db) == (2, 4) + + def test_older_backup_reports_replaced_sessions(self, tmp_path, monkeypatch, capsys): + """Importing a backup that predates recorded work says what it dropped.""" + from hermes_cli.backup import run_import + + home, live_db, zip_path = self._prepare(tmp_path, monkeypatch) + run_import(Namespace(zipfile=str(zip_path), force=True)) + + out = capsys.readouterr().out + assert "Session data replaced by older backup contents" in out + assert "3 session(s) / 12 message(s) -> 2 / 4" in out + + def test_newer_backup_reports_nothing(self, tmp_path, monkeypatch, capsys): + """No warning when the import does not shrink the database.""" + from hermes_cli.backup import run_import + + home, live_db, zip_path = self._prepare( + tmp_path, monkeypatch, live=(1, 1), backup=(3, 4) + ) + run_import(Namespace(zipfile=str(zip_path), force=True)) + + out = capsys.readouterr().out + assert "Session data replaced by older backup contents" not in out + + def test_refused_restore_is_reported_and_leaves_db_intact( + self, tmp_path, monkeypatch, capsys + ): + """A refused live-safe restore is a warning, not a counted success.""" + import hermes_cli.backup as backup_mod + + home, live_db, zip_path = self._prepare(tmp_path, monkeypatch) + monkeypatch.setattr(backup_mod, "_safe_restore_db", lambda src, dst: False) + + backup_mod.run_import(Namespace(zipfile=str(zip_path), force=True)) + + out = capsys.readouterr().out + assert "files skipped" in out + assert "state.db" in out + # The pre-import database is still the one on disk. + assert _count_rows(live_db) == (3, 12) + + def test_sidecar_members_are_not_installed_beside_a_restored_db( + self, tmp_path, monkeypatch + ): + """A `state.db-wal` member from an old/hand-built archive must not be + os.replace'd next to the page-restored database: it describes a + different image and SQLite would replay it on the next open.""" + from hermes_cli.backup import run_import + + home, live_db, zip_path = self._prepare(tmp_path, monkeypatch) + with zipfile.ZipFile(zip_path, "a") as zf: + zf.writestr("state.db-wal", b"foreign-wal-from-archive") + zf.writestr("state.db-shm", b"foreign-shm") + zf.writestr("state.db-journal", b"foreign-journal") + + run_import(Namespace(zipfile=str(zip_path), force=True)) + + for suffix, payload in ( + ("-wal", b"foreign-wal-from-archive"), + ("-shm", b"foreign-shm"), + ("-journal", b"foreign-journal"), + ): + sidecar = live_db.with_name("state.db" + suffix) + assert not sidecar.exists() or sidecar.read_bytes() != payload, suffix + assert _count_rows(live_db) == (2, 4) + + def test_missing_target_takes_the_plain_publish(self, tmp_path, monkeypatch): + """A fresh install has no inode to preserve; the member still lands.""" + from hermes_cli.backup import run_import + + home = tmp_path / ".hermes" + home.mkdir() + monkeypatch.setenv("HERMES_HOME", str(home)) + monkeypatch.setattr(Path, "home", lambda: tmp_path) + + staged = tmp_path / "backup-state.db" + _write_session_db(staged, 2, 3) + zip_path = tmp_path / "backup.zip" + self._zip_with_db(zip_path, staged) + + run_import(Namespace(zipfile=str(zip_path), force=True)) + assert _count_rows(home / "state.db") == (2, 6) + + +def _count_rows(db_path: Path) -> tuple[int, int]: + conn = sqlite3.connect(str(db_path)) + try: + return ( + conn.execute("SELECT COUNT(*) FROM sessions").fetchone()[0], + conn.execute("SELECT COUNT(*) FROM messages").fetchone()[0], + ) + finally: + conn.close() diff --git a/tests/hermes_cli/test_config.py b/tests/hermes_cli/test_config.py index 0db5db5226..3ed68bba14 100644 --- a/tests/hermes_cli/test_config.py +++ b/tests/hermes_cli/test_config.py @@ -10,6 +10,7 @@ import yaml from hermes_cli.config import ( DEFAULT_CONFIG, + InvalidUserConfigError, check_config_version, get_hermes_home, ensure_hermes_home, @@ -738,7 +739,9 @@ class TestConfigMigrationSecretPrompts: saved = {} monkeypatch.setattr(cfg_mod, "sanitize_env_file", lambda: 0) - monkeypatch.setattr(cfg_mod, "check_config_version", lambda: (999, 999)) + monkeypatch.setattr( + cfg_mod, "check_config_version", lambda **_kwargs: (999, 999) + ) monkeypatch.setattr(cfg_mod, "get_missing_config_fields", lambda: []) monkeypatch.setattr(cfg_mod, "get_missing_skill_config_vars", lambda: []) monkeypatch.setattr( @@ -783,6 +786,48 @@ class TestConfigVersionDetection: assert load_config()["_config_version"] == DEFAULT_CONFIG["_config_version"] assert check_config_version() == (0, DEFAULT_CONFIG["_config_version"]) + _LATEST = DEFAULT_CONFIG["_config_version"] + # (bytes, strict match, tolerant return): tolerant malformed YAML keeps + # the historical latest/latest fallback; a parseable non-mapping root is + # reported as legacy (0). + _INVALID_CONFIG_CASES = [ + pytest.param( + b"model: [unterminated\n", "not valid YAML", (_LATEST, _LATEST), id="malformed-yaml" + ), + pytest.param(b"- just_a_list\n", "must be a mapping", (0, _LATEST), id="list-root"), + pytest.param(b"[]\n", "must be a mapping", (0, _LATEST), id="empty-list-root"), + ] + + @pytest.mark.parametrize("config_bytes, match, tolerant", _INVALID_CONFIG_CASES) + def test_strict_check_rejects_invalid_config( + self, tmp_path, config_bytes, match, tolerant + ): + config_path = tmp_path / "config.yaml" + config_path.write_bytes(config_bytes) + + with patch.dict(os.environ, {"HERMES_HOME": str(tmp_path)}): + with pytest.raises(InvalidUserConfigError, match=match): + check_config_version(raise_on_parse_error=True) + # Tolerant callers keep the historical non-raising behavior. + assert check_config_version() == tolerant + + @pytest.mark.parametrize("config_bytes, match, _tolerant", _INVALID_CONFIG_CASES) + def test_migration_rejects_invalid_config_before_sanitizing_env( + self, tmp_path, config_bytes, match, _tolerant + ): + config_path = tmp_path / "config.yaml" + config_path.write_bytes(config_bytes) + env_path = tmp_path / ".env" + env_bytes = b"OPENAI_API_KEY=test-without-final-newline" + env_path.write_bytes(env_bytes) + + with patch.dict(os.environ, {"HERMES_HOME": str(tmp_path)}): + with pytest.raises(InvalidUserConfigError, match=match): + migrate_config(interactive=False, quiet=True) + + assert config_path.read_bytes() == config_bytes + assert env_path.read_bytes() == env_bytes + class TestConfigSupportFloor: """Auto-migration support floor (v12). diff --git a/tests/hermes_cli/test_credential_pool_off_loop.py b/tests/hermes_cli/test_credential_pool_off_loop.py index 881d861b54..6ffe6a6137 100644 --- a/tests/hermes_cli/test_credential_pool_off_loop.py +++ b/tests/hermes_cli/test_credential_pool_off_loop.py @@ -154,6 +154,22 @@ class TestExchangeSingleFlight: assert len(errors) == 5 assert any("recently failed" in e for e in errors) + def test_negative_cache_short_circuits_before_taking_the_lock(self, monkeypatch): + """While one exchange holds the lock, a caller whose fingerprint is + already in the failure cache must raise immediately rather than park + an executor thread behind the holder.""" + fp = copilot_auth._token_fingerprint("ghu_" + "z" * 30) + copilot_auth._exchange_failure_cache[fp] = time.time() + 60 + lock = copilot_auth._exchange_lock_for(fp) + lock.acquire() # simulate an in-flight holder + try: + started = time.monotonic() + with pytest.raises(ValueError, match="recently failed"): + copilot_auth.exchange_copilot_token("ghu_" + "z" * 30) + assert time.monotonic() - started < 0.5 + finally: + lock.release() + # --------------------------------------------------------------------------- # web_server credential-pool handlers off the loop @@ -186,7 +202,7 @@ async def test_list_credential_pool_keeps_loop_responsive(monkeypatch): from hermes_cli import web_server def slow_read(*args, **kwargs): - time.sleep(0.2) + time.sleep(0.5) return {} monkeypatch.setattr(auth_mod, "read_credential_pool", slow_read) @@ -206,4 +222,6 @@ async def test_list_credential_pool_keeps_loop_responsive(monkeypatch): await web_server.list_credential_pool() stop.set() await t - assert max(gaps) < 0.1, f"event loop stalled for {max(gaps) * 1000:.0f} ms" + # 0.25 s threshold vs a 0.5 s blocking read: a regression (read on the + # loop) trips it by 2x, while runner-noise descheduling would need >200 ms. + assert max(gaps) < 0.25, f"event loop stalled for {max(gaps) * 1000:.0f} ms" diff --git a/tests/hermes_cli/test_fts_optimize_notice.py b/tests/hermes_cli/test_fts_optimize_notice.py new file mode 100644 index 0000000000..8ba415c0a3 --- /dev/null +++ b/tests/hermes_cli/test_fts_optimize_notice.py @@ -0,0 +1,52 @@ +"""Regression coverage for FTS storage upgrade discoverability.""" + +import sqlite3 +from types import SimpleNamespace + + +def test_update_notice_offers_v1_trigram_tool_calls_rebuild(tmp_path, monkeypatch, capsys): + """A deployed v1 trigram projection still receives the opt-in notice.""" + from hermes_cli import update_cmd + import hermes_constants + import hermes_state + + db_path = tmp_path / "state.db" + db_path.touch() + conn = sqlite3.connect(db_path) + conn.executescript( + """ + CREATE TABLE state_meta (key TEXT PRIMARY KEY, value TEXT); + CREATE TABLE messages_fts (content TEXT, tool_name TEXT, tool_calls TEXT); + CREATE TABLE messages_fts_trigram (content TEXT, tool_name TEXT, tool_calls TEXT); + """ + ) + + class FakeSessionDB: + def __init__(self, **_kwargs): + self._conn = conn + + def close(self): + pass + + _db_needs_fts_storage_upgrade = staticmethod( + hermes_state.SessionDB._db_needs_fts_storage_upgrade + ) + + monkeypatch.setattr(hermes_constants, "get_hermes_home", lambda: tmp_path) + monkeypatch.setattr(hermes_state, "SessionDB", FakeSessionDB) + # Report a large state.db without patching Path.stat globally: a + # 1-arg lambda on the class breaks pathlib.exists(follow_symlinks=...) + # for every caller in the process (pytest's own teardown included). + real_stat = update_cmd.Path.stat + + def _stat(path, *args, **kwargs): + if path.name == "state.db": + return SimpleNamespace(st_size=512 * 1024 ** 2) + return real_stat(path, *args, **kwargs) + + monkeypatch.setattr(update_cmd.Path, "stat", _stat) + + update_cmd._print_fts_optimize_available_notice() + + assert "hermes sessions optimize-storage" in capsys.readouterr().out + conn.close() diff --git a/tests/hermes_cli/test_gateway_restart_loop.py b/tests/hermes_cli/test_gateway_restart_loop.py index d310d7b4ca..58eb2bcd6f 100644 --- a/tests/hermes_cli/test_gateway_restart_loop.py +++ b/tests/hermes_cli/test_gateway_restart_loop.py @@ -587,6 +587,33 @@ class TestTerminalToolGatewayLifecycleGuard: assert result["exit_code"] == 1 assert "KeepAlive" in result["error"] + def test_oversized_root_skips_launchctl_prescan_and_fails_closed( + self, monkeypatch + ): + """#78398: an over-budget root must never reach shlex — not even via + the launchctl pre-scan that runs before the full guard.""" + import cron.lifecycle_guard as lifecycle_guard + import tools.terminal_tool as tt + + self._patch_env(monkeypatch, self._make_fake_env(), inside_gateway=True) + monkeypatch.setattr( + lifecycle_guard, "_MAX_LIFECYCLE_SCAN_BYTES", 8, raising=False + ) + monkeypatch.setattr( + lifecycle_guard, "_MAX_LIFECYCLE_SCAN_LINE_BYTES", 8, raising=False + ) + + def explode_if_tokenized(*args, **kwargs): + raise AssertionError("over-budget root reached shlex") + + monkeypatch.setattr(lifecycle_guard.shlex, "shlex", explode_if_tokenized) + + result = json.loads(tt.terminal_tool(command="x" * 9)) + + assert result["exit_code"] == 1 + assert "command or referenced script" in result["error"] + assert "KeepAlive" not in result["error"] + @pytest.mark.parametrize("command", [ # Neutral, non-hermes label: label-independent detection is the point # (#62891 second reproduction used `ai.hermes.svc-reload-tmp`). diff --git a/tests/hermes_cli/test_inventory_pricing.py b/tests/hermes_cli/test_inventory_pricing.py index 5fb7c39490..d626593f34 100644 --- a/tests/hermes_cli/test_inventory_pricing.py +++ b/tests/hermes_cli/test_inventory_pricing.py @@ -5,6 +5,9 @@ columns + Free/Pro badges and gate paid models on free Nous accounts, the same way the `hermes model` CLI picker does. """ +from threading import Event +from time import monotonic + import hermes_cli.inventory as inv import hermes_cli.models as models_mod @@ -101,3 +104,389 @@ def test_apply_pricing_omits_sale_when_original_not_cheaper(monkeypatch): assert "discount_percent" not in rows[0]["pricing"]["a/eq"] +def test_model_options_cold_pricing_fetch_runs_off_the_request_path(monkeypatch): + """A cold pricing endpoint must not delay the first picker payload.""" + fetch_started = Event() + release_fetch = Event() + + def fake_pricing(_slug, *, force_refresh=False, cached_only=False): + if cached_only: + return {} + fetch_started.set() + release_fetch.wait(timeout=5) + return {} + + row = { + "slug": "openrouter", + "name": "OpenRouter", + "models": ["vendor/model"], + "total_models": 1, + "is_current": True, + "is_user_defined": False, + "source": "built-in", + } + monkeypatch.setattr(models_mod, "get_pricing_for_provider", fake_pricing) + monkeypatch.setattr( + "hermes_cli.model_switch.list_authenticated_providers", + lambda **_kwargs: [row], + ) + monkeypatch.setattr(inv, "_moa_provider_row", lambda *_args, **_kwargs: None) + monkeypatch.setattr(inv, "_apply_capabilities", lambda _rows: None) + monkeypatch.setattr(inv, "_apply_featured", lambda _rows: None) + monkeypatch.setattr(inv, "_pricing_prewarm_threads", {}) + + try: + started_at = monotonic() + payload = inv.build_model_options_payload( + inv.ConfigContext( + current_provider="openrouter", + current_model="vendor/model", + current_base_url="", + user_providers={}, + custom_providers=[], + ) + ) + elapsed = monotonic() - started_at + assert payload["providers"][0]["slug"] == "openrouter" + assert "pricing" not in payload["providers"][0] + assert elapsed < 2.0, f"cold picker blocked for {elapsed:.2f}s" + assert fetch_started.wait(timeout=1), "pricing should prewarm in the background" + finally: + threads = list(inv._pricing_prewarm_threads.values()) + release_fetch.set() + for thread in threads: + thread.join(timeout=2) + + +def test_cold_nous_entitlement_keeps_models_unselectable(monkeypatch): + """A cold nonblocking response must not expose paid models fail-open.""" + monkeypatch.setattr( + models_mod, "get_pricing_for_provider", lambda *_args, **_kwargs: {} + ) + monkeypatch.setattr(models_mod, "get_cached_nous_free_tier", lambda: None) + rows = [{"slug": "nous", "models": ["free/model", "paid/model"]}] + + inv._apply_pricing(rows, cached_only=True) + + assert rows[0]["free_tier_pending"] is True + assert rows[0]["unavailable_models"] == ["free/model", "paid/model"] + # The whole list renders locked — the picker's per-provider warning + # surface must say why, without clobbering an existing auth warning. + assert "entitlement" in rows[0]["warning"] + + rows = [{"slug": "nous", "models": ["m"], "warning": "paste NOUS_API_KEY to activate"}] + inv._apply_pricing(rows, cached_only=True) + assert rows[0]["warning"] == "paste NOUS_API_KEY to activate" + + +def test_prewarm_preserves_context_and_runs_once_per_profile(tmp_path, monkeypatch): + """Concurrent multiplex profiles retain their own home and secret scope.""" + from agent.secret_scope import ( + current_secret_scope, + reset_secret_scope, + set_secret_scope, + ) + from hermes_constants import ( + hermes_home_key, + reset_hermes_home_override, + set_hermes_home_override, + ) + + monkeypatch.setattr(inv, "_pricing_prewarm_threads", {}) + release = Event() + started = {"a": Event(), "b": Event()} + observed = {} + + def capture_context(_rows): + scope = current_secret_scope() + label = scope["PROFILE_MARKER"] + observed[label] = (hermes_home_key(), dict(scope)) + started[label].set() + release.wait(timeout=5) + + monkeypatch.setattr(inv, "_apply_pricing", capture_context) + + threads = [] + try: + for label in ("a", "b"): + home = tmp_path / label + home_token = set_hermes_home_override(str(home)) + secret_token = set_secret_scope({"PROFILE_MARKER": label}) + try: + threads.append(inv._prewarm_pricing_async([{"models": []}])) + finally: + reset_secret_scope(secret_token) + reset_hermes_home_override(home_token) + + assert threads[0] is not threads[1] + assert started["a"].wait(timeout=1) + assert started["b"].wait(timeout=1) + assert observed["a"] == ( + hermes_home_key(tmp_path / "a"), + {"PROFILE_MARKER": "a"}, + ) + assert observed["b"] == ( + hermes_home_key(tmp_path / "b"), + {"PROFILE_MARKER": "b"}, + ) + finally: + release.set() + for thread in threads: + if thread is not None: + thread.join(timeout=2) + + +def test_prewarm_deduplicates_inflight_scope_and_cleans_up(monkeypatch): + """Rapid opens share one worker, then a completed scope can run again.""" + monkeypatch.setattr(inv, "_pricing_prewarm_threads", {}) + started = Event() + release = Event() + calls = [] + + def blocked_prewarm(_rows): + calls.append(None) + started.set() + release.wait(timeout=5) + + monkeypatch.setattr(inv, "_apply_pricing", blocked_prewarm) + rows = [{"slug": "openrouter", "models": ["vendor/model"]}] + + first = inv._prewarm_pricing_async(rows) + try: + assert started.wait(timeout=1) + second = inv._prewarm_pricing_async(rows) + assert second is first + assert len(calls) == 1 + finally: + release.set() + first.join(timeout=2) + + assert not first.is_alive() + assert inv._pricing_prewarm_threads == {} + + retry = inv._prewarm_pricing_async(rows) + retry.join(timeout=2) + assert retry is not first + assert len(calls) == 2 + assert inv._pricing_prewarm_threads == {} + + +def test_prewarm_endpoint_rotation_starts_a_new_worker(tmp_path, monkeypatch): + """A live endpoint-A worker must not suppress endpoint B for its profile.""" + from hermes_constants import ( + reset_hermes_home_override, + set_hermes_home_override, + ) + + endpoint_a = "https://endpoint-a.example" + endpoint_b = "https://endpoint-b.example" + active_endpoint = {"value": endpoint_a} + started = {endpoint_a: Event(), endpoint_b: Event()} + release_a = Event() + expected = { + endpoint_a: {"a/model": {"prompt": "1", "completion": "2"}}, + endpoint_b: {"b/model": {"prompt": "3", "completion": "4"}}, + } + monkeypatch.setattr(inv, "_pricing_prewarm_threads", {}) + monkeypatch.setattr(models_mod, "_pricing_cache", {}) + monkeypatch.setattr(models_mod, "_pricing_cache_retry_after", {}) + monkeypatch.setattr(models_mod, "_pricing_provider_cache_keys", {}) + monkeypatch.setattr( + models_mod, + "_resolve_nous_pricing_credentials", + lambda: ("", active_endpoint["value"]), + ) + + def fetch_pricing(*, base_url, **_kwargs): + started[base_url].set() + if base_url == endpoint_a: + release_a.wait(timeout=5) + return models_mod._cache_catalog(base_url, expected[base_url]) + + monkeypatch.setattr(models_mod, "fetch_models_with_pricing", fetch_pricing) + monkeypatch.setattr( + inv, + "_apply_pricing", + lambda _rows: models_mod.get_pricing_for_provider("nous"), + ) + + token = set_hermes_home_override(str(tmp_path / "profile")) + threads = [] + try: + threads.append( + inv._prewarm_pricing_async( + [{"slug": "nous", "models": ["a/model"]}], + current_provider="nous", + current_base_url=endpoint_a, + ) + ) + assert started[endpoint_a].wait(timeout=1) + + active_endpoint["value"] = endpoint_b + threads.append( + inv._prewarm_pricing_async( + [{"slug": "nous", "models": ["b/model"]}], + current_provider="nous", + current_base_url=endpoint_b, + ) + ) + + assert threads[0] is not threads[1] + assert started[endpoint_b].wait(timeout=1) + threads[1].join(timeout=2) + assert not threads[1].is_alive() + assert models_mod.get_pricing_for_provider( + "nous", cached_only=True + ) == expected[endpoint_b] + finally: + release_a.set() + for thread in threads: + if thread is not None: + thread.join(timeout=2) + reset_hermes_home_override(token) + + +def test_prewarm_nous_rotation_when_another_provider_is_current(tmp_path, monkeypatch): + """Nous endpoint identity must not depend on Nous being selected.""" + from hermes_constants import ( + reset_hermes_home_override, + set_hermes_home_override, + ) + + endpoint_a = "https://endpoint-a.example" + endpoint_b = "https://endpoint-b.example" + active_endpoint = {"value": endpoint_a} + started = {endpoint_a: Event(), endpoint_b: Event()} + release_a = Event() + expected = { + endpoint_a: {"a/model": {"prompt": "1", "completion": "2"}}, + endpoint_b: {"b/model": {"prompt": "3", "completion": "4"}}, + } + monkeypatch.setattr(inv, "_pricing_prewarm_threads", {}) + monkeypatch.setattr(models_mod, "_pricing_cache", {}) + monkeypatch.setattr(models_mod, "_pricing_cache_retry_after", {}) + monkeypatch.setattr(models_mod, "_pricing_provider_cache_keys", {}) + monkeypatch.setattr( + models_mod, + "get_cached_nous_inference_base_url", + lambda: active_endpoint["value"], + ) + monkeypatch.setattr( + models_mod, + "_resolve_nous_pricing_credentials", + lambda: ("", active_endpoint["value"]), + ) + + def fetch_pricing(*, base_url, **_kwargs): + started[base_url].set() + if base_url == endpoint_a: + release_a.wait(timeout=5) + return models_mod._cache_catalog(base_url, expected[base_url]) + + monkeypatch.setattr(models_mod, "fetch_models_with_pricing", fetch_pricing) + monkeypatch.setattr( + inv, + "_apply_pricing", + lambda _rows: models_mod.get_pricing_for_provider("nous"), + ) + + token = set_hermes_home_override(str(tmp_path / "profile")) + threads = [] + try: + threads.append( + inv._prewarm_pricing_async( + [{"slug": "nous", "models": ["a/model"]}], + current_provider="openrouter", + current_base_url="https://openrouter.ai/api/v1", + ) + ) + assert started[endpoint_a].wait(timeout=1) + + active_endpoint["value"] = endpoint_b + threads.append( + inv._prewarm_pricing_async( + [{"slug": "nous", "models": ["b/model"]}], + current_provider="openrouter", + current_base_url="https://openrouter.ai/api/v1", + ) + ) + + assert threads[0] is not threads[1] + assert started[endpoint_b].wait(timeout=1) + threads[1].join(timeout=2) + assert not threads[1].is_alive() + assert models_mod.get_pricing_for_provider( + "nous", cached_only=True + ) == expected[endpoint_b] + finally: + release_a.set() + for thread in threads: + if thread is not None: + thread.join(timeout=2) + reset_hermes_home_override(token) + + +def test_cached_only_pricing_returns_a_warm_value_without_fetching(monkeypatch): + """Cache-only picker reads preserve pricing once the prewarm completes.""" + cache_key = "https://openrouter.ai/api" + expected = {"vendor/model": {"prompt": "0.000001", "completion": "0.000002"}} + monkeypatch.setattr(models_mod, "_pricing_cache", {cache_key: expected}) + monkeypatch.setattr(models_mod, "_pricing_cache_retry_after", {}) + monkeypatch.setattr(models_mod, "_pricing_provider_cache_keys", {}) + monkeypatch.setattr( + models_mod, + "fetch_models_with_pricing", + lambda **_kwargs: (_ for _ in ()).throw(AssertionError("network fetch started")), + ) + + assert models_mod.get_pricing_for_provider( + "openrouter", cached_only=True + ) == expected + + +def test_cached_only_dynamic_pricing_is_profile_scoped(tmp_path, monkeypatch): + """Alternating profiles read the endpoint each profile warmed.""" + from hermes_constants import ( + reset_hermes_home_override, + set_hermes_home_override, + ) + + endpoint_a = "https://profile-a.example" + endpoint_b = "https://profile-b.example" + expected_a = {"a/model": {"prompt": "1", "completion": "2"}} + expected_b = {"b/model": {"prompt": "3", "completion": "4"}} + monkeypatch.setattr( + models_mod, + "_pricing_cache", + {endpoint_a: expected_a, endpoint_b: expected_b}, + ) + monkeypatch.setattr(models_mod, "_pricing_cache_retry_after", {}) + monkeypatch.setattr(models_mod, "_pricing_provider_cache_keys", {}) + active_endpoint = {"value": endpoint_a} + monkeypatch.setattr( + models_mod, + "_resolve_nous_pricing_credentials", + lambda: ("", active_endpoint["value"]), + ) + monkeypatch.setattr( + models_mod, + "fetch_models_with_pricing", + lambda **kwargs: models_mod._pricing_cache[kwargs["base_url"]], + ) + + def in_profile(home, endpoint, *, cached_only): + token = set_hermes_home_override(str(home)) + active_endpoint["value"] = endpoint + try: + return models_mod.get_pricing_for_provider( + "nous", cached_only=cached_only + ) + finally: + reset_hermes_home_override(token) + + assert in_profile(tmp_path / "a", endpoint_a, cached_only=False) == expected_a + assert in_profile(tmp_path / "b", endpoint_b, cached_only=False) == expected_b + assert in_profile(tmp_path / "a", endpoint_b, cached_only=True) == expected_a + assert in_profile(tmp_path / "b", endpoint_a, cached_only=True) == expected_b + + diff --git a/tests/hermes_cli/test_kanban_gateway_restart_handoff.py b/tests/hermes_cli/test_kanban_gateway_restart_handoff.py new file mode 100644 index 0000000000..d0efb6ed89 --- /dev/null +++ b/tests/hermes_cli/test_kanban_gateway_restart_handoff.py @@ -0,0 +1,178 @@ +"""Managed-gateway isolation for dispatcher-owned Kanban workers.""" + +from __future__ import annotations + +import json +import subprocess +import sys +import time +from pathlib import Path + +import pytest + +from hermes_cli import kanban_db as kb + + +@pytest.fixture +def worker_setup(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> tuple[Path, kb.Task]: + root = tmp_path / ".hermes" + profile = root / "profiles" / "coder" + profile.mkdir(parents=True) + root.joinpath("config.yaml").write_text("{}\n", encoding="utf-8") + profile.joinpath("config.yaml").write_text("{}\n", encoding="utf-8") + monkeypatch.setenv("HERMES_HOME", str(root)) + monkeypatch.setattr(Path, "home", lambda: tmp_path) + monkeypatch.setattr(kb, "_resolve_hermes_argv", lambda: ["hermes"]) + + workspace = tmp_path / "candidate-worktree" + workspace.mkdir() + task = kb.Task( + id="t_candidate_restart", + title="activate candidate", + body=None, + assignee="coder", + status="running", + priority=0, + created_by="test", + created_at=1, + started_at=1, + completed_at=None, + workspace_kind="worktree", + workspace_path=str(workspace), + claim_lock="host:dispatcher", + claim_expires=999, + tenant=None, + branch_name="wt/t_candidate_restart", + current_run_id=23, + ) + return workspace, task + + +@pytest.mark.linux_only +def test_managed_gateway_worker_is_spawned_in_restart_safe_scope( + worker_setup: tuple[Path, kb.Task], monkeypatch: pytest.MonkeyPatch +) -> None: + workspace, task = worker_setup + captured_cmd: list[str] = [] + captured_env: dict[str, str] = {} + captured_cwd: str | None = None + + class FakeProc: + pid = 4242 + + def fake_popen(cmd, **kwargs): + nonlocal captured_cwd + captured_cmd.extend(cmd) + captured_env.update(kwargs.get("env") or {}) + captured_cwd = kwargs.get("cwd") + return FakeProc() + + monkeypatch.setenv("INVOCATION_ID", "managed-gateway-test") + monkeypatch.setenv("ANTHROPIC_API_KEY", "must-not-cross-profile") + monkeypatch.setattr("agent.secret_scope.is_multiplex_active", lambda: True) + monkeypatch.setattr(subprocess, "Popen", fake_popen) + monkeypatch.setattr("tools.process_registry._is_supervised_gateway_process", lambda: True) + monkeypatch.setattr("tools.process_registry._systemd_run_user_scope_available", lambda: True) + monkeypatch.setattr("tools.process_registry._worker_memory_max_bytes", lambda: 536_870_912) + monkeypatch.setattr("shutil.which", lambda name: "/usr/bin/systemd-run") + + assert kb._default_spawn(task, str(workspace)) == 4242 + assert captured_cmd[:4] == ["/usr/bin/systemd-run", "--user", "--scope", "--quiet"] + unit_index = captured_cmd.index("--unit") + assert captured_cmd[unit_index + 1] == "hermes-worker-kanban-t_candidate_restart-run-23" + assert "MemoryMax=536870912" in captured_cmd + separator = captured_cmd.index("--") + assert captured_cmd[separator + 1 : separator + 4] == ["hermes", "-p", "coder"] + assert captured_cwd == str(workspace) + assert captured_env["HERMES_KANBAN_TASK"] == task.id + assert captured_env["HERMES_KANBAN_RUN_ID"] == "23" + assert "ANTHROPIC_API_KEY" not in captured_env + + +@pytest.mark.linux_only +def test_managed_gateway_worker_spawn_fails_closed_without_scope( + worker_setup: tuple[Path, kb.Task], monkeypatch: pytest.MonkeyPatch +) -> None: + workspace, task = worker_setup + popen_calls: list[list[str]] = [] + monkeypatch.setenv("INVOCATION_ID", "managed-gateway-test") + monkeypatch.setattr(subprocess, "Popen", lambda cmd, **kwargs: popen_calls.append(list(cmd))) + monkeypatch.setattr("tools.process_registry._is_supervised_gateway_process", lambda: True) + monkeypatch.setattr("tools.process_registry._systemd_run_user_scope_available", lambda: False) + + with pytest.raises(RuntimeError, match="restart-safe systemd scope"): + kb._default_spawn(task, str(workspace)) + assert popen_calls == [] + + +@pytest.mark.linux_only +def test_managed_gateway_scope_builder_fails_closed_if_binary_disappears( + worker_setup: tuple[Path, kb.Task], monkeypatch: pytest.MonkeyPatch +) -> None: + workspace, task = worker_setup + monkeypatch.setenv("INVOCATION_ID", "managed-gateway-test") + monkeypatch.setattr("tools.process_registry._is_supervised_gateway_process", lambda: True) + monkeypatch.setattr("tools.process_registry._systemd_run_user_scope_available", lambda: True) + monkeypatch.setattr("shutil.which", lambda _name: None) + monkeypatch.setattr(subprocess, "Popen", lambda *_args, **_kwargs: pytest.fail("unsafe direct spawn")) + + with pytest.raises(RuntimeError, match="restart-safe systemd scope"): + kb._default_spawn(task, str(workspace)) + + +def test_standalone_dispatcher_keeps_direct_worker_spawn( + worker_setup: tuple[Path, kb.Task], monkeypatch: pytest.MonkeyPatch +) -> None: + workspace, task = worker_setup + captured_cmd: list[str] = [] + + class FakeProc: + pid = 4243 + + monkeypatch.setattr(subprocess, "Popen", lambda cmd, **kwargs: captured_cmd.extend(cmd) or FakeProc()) + monkeypatch.setattr("tools.process_registry._is_supervised_gateway_process", lambda: False) + monkeypatch.setattr( + "tools.process_registry._systemd_run_user_scope_available", + lambda: pytest.fail("scope probe must not run outside managed gateway"), + ) + + assert kb._default_spawn(task, str(workspace)) == 4243 + assert captured_cmd[:3] == ["hermes", "-p", "coder"] + + +@pytest.mark.linux_only +def test_real_user_systemd_scope_preserves_worker_context( + worker_setup: tuple[Path, kb.Task], monkeypatch: pytest.MonkeyPatch +) -> None: + from tools import process_registry + + if not process_registry._systemd_run_user_scope_available(): + pytest.skip("systemd-run --user --scope is unavailable on this host") + + workspace, task = worker_setup + receipt = workspace / "worker-receipt.json" + script = ( + "import json, os, pathlib, sys, time; " + "pathlib.Path(sys.argv[1]).write_text(json.dumps({" + "'pid': os.getpid(), 'cwd': os.getcwd(), " + "'task': os.environ.get('HERMES_KANBAN_TASK'), " + "'run': os.environ.get('HERMES_KANBAN_RUN_ID'), " + "'cgroup': pathlib.Path('/proc/self/cgroup').read_text()})); time.sleep(0.5)" + ) + monkeypatch.setattr(kb, "_resolve_hermes_argv", lambda: [sys.executable, "-c", script, str(receipt)]) + monkeypatch.setenv("INVOCATION_ID", "managed-gateway-test") + monkeypatch.setattr(process_registry, "_is_supervised_gateway_process", lambda: True) + + pid = kb._default_spawn(task, str(workspace)) + deadline = time.monotonic() + 5 + while not receipt.exists() and time.monotonic() < deadline: + time.sleep(0.05) + + assert receipt.exists() + payload = json.loads(receipt.read_text(encoding="utf-8")) + assert payload["pid"] == pid + assert payload["cwd"] == str(workspace) + assert payload["task"] == task.id + assert payload["run"] == "23" + assert ".scope" in payload["cgroup"] + assert "hermes-gateway.service" not in payload["cgroup"] diff --git a/tests/hermes_cli/test_mcp_startup.py b/tests/hermes_cli/test_mcp_startup.py index 9c4b94a182..f9be0410ac 100644 --- a/tests/hermes_cli/test_mcp_startup.py +++ b/tests/hermes_cli/test_mcp_startup.py @@ -102,6 +102,96 @@ def test_prepare_agent_startup_backgrounds_blocking_mcp_for_chat(monkeypatch): stop.set() +def test_prepare_agent_startup_skips_discovery_when_chat_resolves_to_tui( + monkeypatch, +): + """Bare ``hermes`` / ``hermes chat`` on a TTY with ``display.interface: + tui`` resolves to the TUI via ``_resolve_use_tui``, but does NOT pass + ``--tui`` or ``HERMES_TUI``. Discovery must be skipped in the wrapper: + the TUI gateway owns it, and the wrapper would otherwise hold a dead + MCP server for the entire session (3 copies per TUI instance). + """ + calls = {"background": 0, "inline": 0} + + monkeypatch.setattr(main_mod, "_resolve_use_tui", lambda _args: True) + monkeypatch.setattr( + mcp_startup, + "start_background_mcp_discovery", + lambda **_kwargs: calls.__setitem__("background", calls["background"] + 1), + ) + monkeypatch.setitem( + sys.modules, + "hermes_cli.plugins", + types.SimpleNamespace(discover_plugins=lambda: None), + ) + monkeypatch.setitem( + sys.modules, + "hermes_cli.config", + types.SimpleNamespace( + read_raw_config=lambda: {"mcp_servers": {"demo": {"transport": "stdio"}}}, + load_config=lambda: {}, + ), + ) + monkeypatch.setitem( + sys.modules, + "agent.shell_hooks", + types.SimpleNamespace(register_from_config=lambda *_a, **_k: None), + ) + monkeypatch.setitem( + sys.modules, + "tools.mcp_tool", + types.SimpleNamespace( + discover_mcp_tools=lambda: calls.__setitem__("inline", calls["inline"] + 1), + ), + ) + + main_mod._prepare_agent_startup(_agent_args(command=None)) + + assert calls["background"] == 0 + assert calls["inline"] == 0 + assert mcp_startup._mcp_discovery_thread is None + + +def test_prepare_agent_startup_keeps_discovery_for_non_chat_commands( + monkeypatch, +): + """Non-chat commands never launch the TUI, so they must keep their own + MCP discovery even when the ambient display config resolves to TUI — + ``_is_tui_chat_launch`` must not consult ``_resolve_use_tui`` there.""" + calls = {"inline": 0} + + monkeypatch.setattr(main_mod, "_resolve_use_tui", lambda _args: True) + monkeypatch.setitem( + sys.modules, + "hermes_cli.plugins", + types.SimpleNamespace(discover_plugins=lambda: None), + ) + monkeypatch.setitem( + sys.modules, + "hermes_cli.config", + types.SimpleNamespace( + read_raw_config=lambda: {"mcp_servers": {"demo": {"transport": "stdio"}}}, + load_config=lambda: {}, + ), + ) + monkeypatch.setitem( + sys.modules, + "agent.shell_hooks", + types.SimpleNamespace(register_from_config=lambda *_a, **_k: None), + ) + monkeypatch.setitem( + sys.modules, + "tools.mcp_tool", + types.SimpleNamespace( + discover_mcp_tools=lambda: calls.__setitem__("inline", calls["inline"] + 1), + ), + ) + + main_mod._prepare_agent_startup(_agent_args(command="mcp", mcp_action="serve")) + + assert calls["inline"] == 1 + + def test_background_mcp_discovery_suppresses_interactive_oauth(monkeypatch): state = {"active": False, "during_discover": None} @@ -199,3 +289,108 @@ def _install_retry_stubs(monkeypatch, *, connected: bool, calls: dict): ) + + +# --- -t/--toolsets MCP spawn filter (#19000) -------------------------------- + + +@pytest.fixture +def _reset_mcp_server_filter(): + saved = mcp_startup._mcp_server_filter + try: + yield + finally: + mcp_startup._mcp_server_filter = saved + + +@pytest.mark.parametrize( + ("toolsets", "expected"), + [ + (None, None), + ("", None), + ("all", None), + (["*"], None), + ("terminal,web", ["terminal", "web"]), + (["terminal", "code-mcp,web"], ["terminal", "code-mcp", "web"]), + ], +) +def test_set_mcp_server_filter_normalizes(_reset_mcp_server_filter, toolsets, expected): + assert mcp_startup.set_mcp_server_filter(toolsets) == expected + assert mcp_startup.get_mcp_server_filter() == expected + + +def test_discover_mcp_tools_spawns_only_allowed_servers(monkeypatch): + """The filter must narrow the spawn set before any server is connected; + built-in toolset names in the list are ignored.""" + from tools import mcp_tool + + servers = { + "code-mcp": {"command": "true"}, + "docs-mcp": {"command": "true"}, + } + seen: dict[str, dict] = {} + + sdk_probes = {"n": 0} + + def _fake_ensure_sdk(): + sdk_probes["n"] += 1 + return True + + monkeypatch.setattr(mcp_tool, "_load_mcp_config", lambda: dict(servers)) + monkeypatch.setattr(mcp_tool, "_ensure_mcp_sdk", _fake_ensure_sdk) + monkeypatch.setattr(mcp_tool, "_try_acquire_mcp_discovery_lock", lambda: mcp_tool._LOCK_UNAVAILABLE) + monkeypatch.setattr(mcp_tool, "_release_mcp_discovery_lock", lambda *_a, **_k: None, raising=False) + + def _fake_register(cfgs): + seen.update(cfgs) + return [] + + monkeypatch.setattr(mcp_tool, "register_mcp_servers", _fake_register) + monkeypatch.setattr(mcp_tool, "_servers", {}) + monkeypatch.setattr(mcp_tool, "_server_connecting", set()) + + # Everything (no filter) — both would be registered. + mcp_tool.discover_mcp_tools() + assert set(seen) == {"code-mcp", "docs-mcp"} + + # `-t terminal,code-mcp` — only the matching server; "terminal" is a no-op. + seen.clear() + mcp_tool.discover_mcp_tools(allowed_mcp_names=["terminal", "code-mcp"]) + assert set(seen) == {"code-mcp"} + + # `-t terminal` — no MCP server in the filter: skip the whole MCP load, + # including the ~260ms `mcp` SDK import. + seen.clear() + sdk_probes["n"] = 0 + assert mcp_tool.discover_mcp_tools(allowed_mcp_names=["terminal"]) == [] + assert seen == {} + assert sdk_probes["n"] == 0 + + +def test_background_discovery_honors_server_filter(monkeypatch, _reset_mcp_server_filter): + calls: list = [] + monkeypatch.setitem( + sys.modules, + "tools.mcp_tool", + types.SimpleNamespace(discover_mcp_tools=lambda allowed_mcp_names=None: calls.append(allowed_mcp_names)), + ) + monkeypatch.setitem( + sys.modules, + "tools.mcp_oauth", + types.SimpleNamespace(suppress_interactive_oauth=nullcontext), + ) + mcp_startup.set_mcp_server_filter("terminal,code-mcp") + mcp_startup._discover_mcp_tools_without_interactive_oauth() + assert calls == [["terminal", "code-mcp"]] + + +def test_prepare_agent_startup_installs_server_filter(monkeypatch, _reset_mcp_server_filter): + monkeypatch.setitem( + sys.modules, + "hermes_cli.plugins", + types.SimpleNamespace(discover_plugins=lambda: None), + ) + monkeypatch.setattr(main_mod, "_should_background_mcp_startup", lambda args: False) + monkeypatch.setattr(main_mod, "_command_has_dedicated_mcp_startup", lambda args: True) + main_mod._prepare_agent_startup(_agent_args(toolsets="terminal,code-mcp")) + assert mcp_startup.get_mcp_server_filter() == ["terminal", "code-mcp"] diff --git a/tests/hermes_cli/test_model_data_policy_guard.py b/tests/hermes_cli/test_model_data_policy_guard.py index 493ccd4bb8..2546dba15d 100644 --- a/tests/hermes_cli/test_model_data_policy_guard.py +++ b/tests/hermes_cli/test_model_data_policy_guard.py @@ -11,9 +11,10 @@ def test_fires_on_meta_contributor(): assert isinstance(w, DataTrainingWarning) assert w.model == "muse-spark-1.2-contributor" assert "train" in w.message.lower() - assert "muse-spark-1.2" in w.message # points to the no-training alternative - # Aligns with Meta's own pricing doc language + figures. - assert "$0.10" in w.message and "$0.20" in w.message and "$0.002" in w.message + assert "-contributor" in w.message.lower() or "contributor" in w.message.lower() # mentions the tier + # Points to Meta's live pricing page instead of hardcoding prices. + assert "pricing and rate limits" in w.message.lower() + assert "$0.10" not in w.message and "$0.20" not in w.message and "$0.002" not in w.message assert "prompts and completions" in w.message.lower() assert "dev.meta.ai/docs/pricing-rate-limits" in w.message diff --git a/tests/hermes_cli/test_models.py b/tests/hermes_cli/test_models.py index c4535f503a..0cb4213db0 100644 --- a/tests/hermes_cli/test_models.py +++ b/tests/hermes_cli/test_models.py @@ -282,10 +282,10 @@ class TestCheckNousFreeTierCache: """Tests for the TTL cache on check_nous_free_tier().""" def setup_method(self): - _models_mod._free_tier_cache = None + _models_mod._free_tier_cache.clear() def teardown_method(self): - _models_mod._free_tier_cache = None + _models_mod._free_tier_cache.clear() @patch("hermes_cli.nous_account.get_nous_portal_account_info") def test_result_is_cached(self, mock_account): @@ -303,6 +303,42 @@ class TestCheckNousFreeTierCache: assert result2 is True assert mock_account.call_count == 1 + @patch("hermes_cli.nous_account.get_nous_portal_account_info") + def test_cache_only_cold_lookup_does_not_call_portal(self, mock_account): + assert check_nous_free_tier(cached_only=True) is False + mock_account.assert_not_called() + + @patch("hermes_cli.nous_account.get_nous_portal_account_info") + def test_entitlement_cache_is_profile_scoped(self, mock_account, tmp_path): + from hermes_constants import ( + hermes_home_key, + reset_hermes_home_override, + set_hermes_home_override, + ) + + def account_for_active_profile(*, force_fresh=False): + is_free = hermes_home_key() == hermes_home_key(tmp_path / "free") + return NousPortalAccountInfo( + logged_in=True, + source="jwt", + fresh=force_fresh, + paid_service_access=not is_free, + ) + + mock_account.side_effect = account_for_active_profile + + def check_in(home): + token = set_hermes_home_override(str(home)) + try: + return check_nous_free_tier() + finally: + reset_hermes_home_override(token) + + assert check_in(tmp_path / "free") is True + assert check_in(tmp_path / "paid") is False + assert check_in(tmp_path / "free") is True + assert mock_account.call_count == 2 + @patch("hermes_cli.nous_account.get_nous_portal_account_info") def test_force_fresh_bypasses_cache(self, mock_account): diff --git a/tests/hermes_cli/test_opencode_free_live_catalog.py b/tests/hermes_cli/test_opencode_free_live_catalog.py index 24fd7fc6a8..67ac93b64a 100644 --- a/tests/hermes_cli/test_opencode_free_live_catalog.py +++ b/tests/hermes_cli/test_opencode_free_live_catalog.py @@ -41,6 +41,7 @@ _LIVE_FREE_MODELS = [ "nemotron-3-ultra-free", "nemotron-3.5-lightning-free", "muse-spark-1.2-contributor-free", + "muse-spark-1.3-contributor-free", ] # The raw live /zen/v1/models dump also lists paid/subscription + KEYED-free IDs diff --git a/tests/hermes_cli/test_relay_shared_metrics_runtime.py b/tests/hermes_cli/test_relay_shared_metrics_runtime.py index 5520ec8612..292c43abbe 100644 --- a/tests/hermes_cli/test_relay_shared_metrics_runtime.py +++ b/tests/hermes_cli/test_relay_shared_metrics_runtime.py @@ -25,6 +25,12 @@ class _Request: self.content = content +class _ToolExecutionResult: + def __init__(self, result: Any, annotation: Any = None) -> None: + self.result = result + self.annotation = annotation + + class _Relay: def __init__(self) -> None: self.events: list[tuple[Any, ...]] = [] @@ -39,6 +45,7 @@ class _Relay: Agent="agent", Function="function", Tool="tool" ) self.LLMRequest = _Request + self.ToolExecutionResult = _ToolExecutionResult self.scope = SimpleNamespace( push=self._scope_push, pop=self._scope_pop, @@ -174,11 +181,13 @@ class _Relay: def _tool_call_end( self, handle: Any, - result: dict[str, Any], + result: _ToolExecutionResult, **kwargs: Any, ) -> None: + assert isinstance(result, _ToolExecutionResult) + payload = result.result start = self._tool_starts.pop(handle) - self.events.append(("tool.call_end", handle, result, kwargs)) + self.events.append(("tool.call_end", handle, payload, kwargs)) event = SimpleNamespace( kind="scope", category="tool", @@ -190,7 +199,7 @@ class _Relay: **kwargs["metadata"], "otel.status_code": "OK", }, - data=result, + data=payload, ) for callback in list(self._callbacks.values()): callback(event) diff --git a/tests/hermes_cli/test_session_recovery_lost_and_found.py b/tests/hermes_cli/test_session_recovery_lost_and_found.py index d93bd69b6f..5831bbd602 100644 --- a/tests/hermes_cli/test_session_recovery_lost_and_found.py +++ b/tests/hermes_cli/test_session_recovery_lost_and_found.py @@ -17,6 +17,7 @@ import pytest from hermes_state import SessionDB from hermes_cli import session_recovery from hermes_cli.session_lost_and_found import ( + STUB_TITLE_PREFIX, classify_lost_and_found_row, map_lost_and_found_rows, rebuild_fts_indexes, @@ -569,3 +570,292 @@ def test_fingerprint_error_enumerates_parent_cli_session( assert "CLI session" in message assert "fresh shell" in message assert "snapshot" in message + + +# ── issue #101409: mis-mapped salvage must not be reported verified ───────── + + +def _map_salvage_rows( + tmp_path: Path, + *, + blank_session_started_at: bool, + blank_message_timestamp: bool, +) -> sqlite3.Connection: + """Map synthetic lost_and_found cells into a fresh template DB. + + With either ``blank_*`` flag the cells mimic an upgraded source's + *physical* column order (#101409): whatever lands on the declared + ``started_at``/``timestamp`` position is not an epoch timestamp, so + the NOT NULL substitute turns it into 0.0 on every row. + """ + + schema_ref = tmp_path / "schema-ref.db" + SessionDB(db_path=schema_ref).close() + schema = sqlite3.connect(str(schema_ref)) + try: + sessions_columns = [ + str(row[1]) for row in schema.execute("PRAGMA table_info(sessions)") + ] + messages_columns = [ + str(row[1]) for row in schema.execute("PRAGMA table_info(messages)") + ] + finally: + schema.close() + current_width = len(sessions_columns) + + lf_path = tmp_path / "lost_and_found.db" + lf_conn = sqlite3.connect(str(lf_path), isolation_level=None) + try: + lf_cells = ", ".join(f"c{i}" for i in range(current_width)) + lf_conn.execute( + "CREATE TABLE lost_and_found (rootpgno INTEGER, pgno INTEGER, " + "nfield INTEGER, id INTEGER, " + lf_cells + ")" + ) + + def insert(nfield: int, rowid: int, values: list) -> None: + padded = list(values) + [None] * (current_width - len(values)) + placeholders = ", ".join("?" for _ in range(4 + current_width)) + lf_conn.execute( + "INSERT INTO lost_and_found VALUES (" + placeholders + ")", + [2, 5, nfield, rowid, *padded], + ) + + def session_row(session_id: str) -> list: + # title is UNIQUE (idx_sessions_title_unique) — keep it distinct + # per row so the probe isolates timestamp mis-mapping. + row = { + "id": session_id, + "source": "telegram", + "started_at": None + if blank_session_started_at + else 1_754_000_000.0, + "message_count": 2, + "title": f"mis-mapped probe {session_id}", + } + return [row.get(column) for column in sessions_columns] + + for index in range(3): + insert( + current_width, + index + 1, + session_row(f"20260101_01010{index}_aaa00{index}"), + ) + + for index in range(2): + message = { + "id": None, + "session_id": "20260101_010100_aaa000", + "role": "user", + "content": "payload", + "timestamp": None + if blank_message_timestamp + else 1_754_000_100.0 + index, + } + insert( + 23, + 100 + index, + [message.get(column) for column in messages_columns[:23]], + ) + finally: + lf_conn.close() + + output = tmp_path / "mapped.db" + SessionDB(db_path=output).close() + lf_conn = sqlite3.connect(str(lf_path), isolation_level=None) + dest = sqlite3.connect(str(output), isolation_level=None) + try: + dest.execute("PRAGMA foreign_keys=OFF") + map_lost_and_found_rows(lf_conn, dest) + finally: + lf_conn.close() + dest.close() + return sqlite3.connect(str(output), isolation_level=None) + + +def test_plausibility_gate_flags_positional_mis_mapping( + tmp_path: Path, +) -> None: + """A salvage whose timestamps all landed below the epoch floor was + mapped onto the wrong columns and must be flagged, not verified + (#101409).""" + + conn = _map_salvage_rows( + tmp_path, + blank_session_started_at=True, + blank_message_timestamp=False, + ) + try: + # The mapper happily inserted every row; structural checks pass. + assert conn.execute("SELECT COUNT(*) FROM sessions").fetchone()[0] == 3 + assert conn.execute( + "SELECT COUNT(*) FROM sessions WHERE started_at = 0.0" + ).fetchone()[0] == 3 + + errors = session_recovery._lost_and_found_plausibility_errors(conn) + assert len(errors) == 1 + assert "sessions.started_at" in errors[0] + finally: + conn.close() + + +def test_plausibility_gate_flags_mis_mapped_message_timestamps( + tmp_path: Path, +) -> None: + conn = _map_salvage_rows( + tmp_path, + blank_session_started_at=False, + blank_message_timestamp=True, + ) + try: + errors = session_recovery._lost_and_found_plausibility_errors(conn) + assert len(errors) == 1 + assert "messages.timestamp" in errors[0] + finally: + conn.close() + + +def test_plausibility_gate_passes_correctly_mapped_salvage( + tmp_path: Path, +) -> None: + """Well-mapped rows — and partially damaged ones (a torn cell on some + rows is expected salvage noise) — must not trip the gate: it fires + only on a *systematic* violation.""" + + conn = _map_salvage_rows( + tmp_path, + blank_session_started_at=False, + blank_message_timestamp=False, + ) + try: + # Damage one of three sessions the way a torn cell would. + conn.execute( + "UPDATE sessions SET started_at = 0.0 WHERE id = ?", + ("20260101_010101_aaa001",), + ) + conn.commit() + + assert session_recovery._lost_and_found_plausibility_errors(conn) == [] + finally: + conn.close() + + +def _rebuild_with_started_at_appended(conn: sqlite3.Connection) -> None: + """Give ``sessions`` the physical layout of an upgraded DB: ``started_at`` + lands at the END (as ``ALTER TABLE ADD COLUMN`` would place a column that + the current template declares mid-definition). Data is preserved.""" + info = list(conn.execute("PRAGMA table_info(sessions)")) + declared = [row[1] for row in info] + + def coldef(row): + _, name, ctype, notnull, dflt, pk = row + parts = [f'"{name}" {ctype}'] + if pk: + parts.append("PRIMARY KEY") + if notnull: + parts.append("NOT NULL") + if dflt is not None: + parts.append(f"DEFAULT {dflt}") + return " ".join(parts) + + reordered = [r for r in info if r[1] != "started_at"] + [r for r in info if r[1] == "started_at"] + cols = ", ".join(f'"{c}"' for c in declared) + conn.executescript("PRAGMA foreign_keys=OFF;") + conn.execute("CREATE TABLE sessions_new (" + ", ".join(coldef(r) for r in reordered) + ")") + conn.execute(f"INSERT INTO sessions_new({cols}) SELECT {cols} FROM sessions") + conn.executescript("DROP TABLE sessions; ALTER TABLE sessions_new RENAME TO sessions;") + + +@pytest.mark.skipif( + not HAVE_SQLITE3_CLI, + reason="sqlite3 CLI not on PATH; .recover is a shell-only feature", +) +def test_lost_and_found_lane_refuses_to_verify_a_physically_shifted_source( + tmp_path: Path, +) -> None: + """#101409 end to end: a source whose physical column order differs from + the template's declared order maps every cell onto the wrong column. The + output still passes integrity/FK/FTS, so only the plausibility gate can + stop the report from claiming ``verified``.""" + source = tmp_path / "upgraded.db" + output = tmp_path / "upgraded-recovered.db" + db = SessionDB(db_path=source) + try: + for n in range(3): + sid = f"20260812_1400{n:02d}_def{n:03x}" + db.create_session(sid, "cli", cwd=f"/tmp/shift-{n}") + db.set_session_title(sid, f"shift {n}") + for m in range(4): + db.append_message(sid, "user" if m % 2 == 0 else "assistant", f"payload {n} {m}") + finally: + db.close() + conn = sqlite3.connect(str(source), isolation_level=None) + try: + conn.execute("PRAGMA wal_checkpoint(TRUNCATE)") + conn.execute("PRAGMA journal_mode=DELETE") + _rebuild_with_started_at_appended(conn) + conn.execute("VACUUM") + physical = [r[1] for r in conn.execute("PRAGMA table_info(sessions)")] + assert physical[-1] == "started_at" + finally: + conn.close() + # The reporter's damage: page 1 (header + sqlite_master) overwritten, so + # ``.recover`` cannot name any table and every row lands in + # lost_and_found, to be mapped positionally onto the template. + with open(source, "r+b") as fh: + fh.write(b"\0" * _page_size(source.read_bytes())) + + report = recover_session_database(source, output, work_dir=tmp_path, allow_partial=True) + + assert report["mode"] == "lost_and_found_salvage" + # Mis-mapped rows that trip a NOT NULL / type constraint are stubbed, not + # mapped (the reporter saw 190 of 1,875) — at least one lands positionally. + assert report["lost_and_found"]["mapped"]["sessions"] >= 1 + assert report["verification"]["healthy"] is False + assert report["verified"] is False + assert any("sessions.started_at is implausible" in e for e in report["verification"]["errors"]) + out = sqlite3.connect(str(output)) + try: + # The mis-mapping the gate caught: every mapped (non-stub) session got + # the NOT NULL substitute where its real start time should be. + mapped = out.execute( + f"SELECT started_at FROM sessions WHERE COALESCE(title, '') NOT LIKE '{STUB_TITLE_PREFIX}%'" + ).fetchall() + assert mapped and all(row[0] == 0.0 for row in mapped) + finally: + out.close() + + +def test_plausibility_gate_ignores_stub_only_sessions(tmp_path: Path) -> None: + """Stub rows from ``stub_missing_parent_sessions`` legitimately carry + ``started_at = 0.0``; a salvage where only stubs survived is depleted, + not mis-mapped, and must not be flagged.""" + output = tmp_path / "stubs.db" + SessionDB(db_path=output).close() + conn = sqlite3.connect(str(output)) + try: + now = 1_750_000_000.0 + conn.execute( + "INSERT INTO sessions (id, source, started_at, title) VALUES (?, ?, ?, ?)", + ("20260812_140000_aaa000", "recovered", 0.0, "[best-effort recovered 1] session metadata was unreadable"), + ) + conn.execute( + "INSERT INTO messages (session_id, role, content, timestamp) VALUES (?, ?, ?, ?)", + ("20260812_140000_aaa000", "user", "hi", now), + ) + conn.commit() + assert session_recovery._lost_and_found_plausibility_errors(conn) == [] + # One genuinely mapped row with a real timestamp keeps it clean too... + conn.execute( + "INSERT INTO sessions (id, source, started_at, title) VALUES (?, ?, ?, ?)", + ("20260812_140001_aaa001", "cli", now, None), + ) + conn.commit() + assert session_recovery._lost_and_found_plausibility_errors(conn) == [] + # ...and a mapped row at 0.0 with a NULL title (the mis-mapped shape: + # blank titles) is still counted as mapped, not as a stub. + conn.execute("UPDATE sessions SET started_at = 0.0 WHERE id = '20260812_140001_aaa001'") + conn.commit() + errors = session_recovery._lost_and_found_plausibility_errors(conn) + assert len(errors) == 1 and "sessions.started_at" in errors[0] + finally: + conn.close() diff --git a/tests/hermes_cli/test_sqlite3_cli_salvage_gate.py b/tests/hermes_cli/test_sqlite3_cli_salvage_gate.py new file mode 100644 index 0000000000..ced5461eef --- /dev/null +++ b/tests/hermes_cli/test_sqlite3_cli_salvage_gate.py @@ -0,0 +1,384 @@ +"""#100368 regression: the corruption guidance must not direct a WAL-reset- +vulnerable sqlite3 CLI at a live Hermes database. + +Field forensics (issue #100368, maintainer round 2 + the isolated reproducer +in its comments): when a shell with SQLite's WAL-reset opener bug (fixed +3.51.3+ / backports 3.50.7 / 3.44.6 — Debian/Ubuntu system CLIs 3.45.1 / +3.46.1 are in the vulnerable band) opens a live state.db whose writer's DMS +lock has been cancelled, it unlinks the live -wal/-shm pair and splits the +store into two concurrent generations. Both generations report +``integrity_check ok`` while an old-generation acknowledged write is lost. + +Hermes' own corruption banners used to instruct exactly that command +(`sqlite3 ~/.hermes/state.db ".recover"`). The fix routes operators to +`hermes sessions recover --source ...`, whose lane snapshots the damaged +bundle before any shell touches it, and refuses a WAL-reset-vulnerable +sqlite3 CLI for the page-level salvage lane even on the snapshot. +""" + +from __future__ import annotations + +import argparse +import inspect +import sqlite3 +from pathlib import Path +from unittest.mock import patch + +import pytest + +from hermes_cli.session_lost_and_found import ( + _parse_sqlite3_cli_version, + _wal_reset_vulnerable, + find_sqlite3_cli, + find_sqlite3_cli_refusal, +) +from hermes_cli.sqlite_runtime import is_sqlite_wal_reset_vulnerable + + +LIVE_DB_SALVAGE_COMMAND = 'sqlite3 ~/.hermes/state.db ".recover"' + + +# --------------------------------------------------------------------------- +# The version gate itself +# --------------------------------------------------------------------------- + + +class TestWalResetVersionGate: + @pytest.mark.parametrize( + "version", + [(3, 45, 1), (3, 46, 1), (3, 44, 5), (3, 50, 4), (3, 51, 2), (3, 8, 0)], + ) + def test_vulnerable_versions(self, version): + assert _wal_reset_vulnerable(version) is True + + @pytest.mark.parametrize( + "version", + [ + (3, 44, 6), + (3, 44, 7), + (3, 50, 7), + (3, 50, 8), + (3, 51, 3), + (3, 51, 4), + (3, 52, 0), + (3, 53, 1), + (4, 0, 0), + ], + ) + def test_fixed_versions(self, version): + assert _wal_reset_vulnerable(version) is False + + def test_gate_mirrors_library_gate(self): + """The salvage gate must agree with the shared runtime gate so the + embedded library and the salvage shell can never disagree.""" + versions = [ + (3, 44, 5), + (3, 44, 6), + (3, 45, 1), + (3, 50, 4), + (3, 50, 7), + (3, 51, 2), + (3, 51, 3), + (3, 53, 1), + ] + for version in versions: + assert _wal_reset_vulnerable(version) == ( + is_sqlite_wal_reset_vulnerable(version) + ), f"salvage gate disagrees with the runtime gate at {version}" + + +# --------------------------------------------------------------------------- +# find_sqlite3_cli refuses unsafe shells and explains why +# --------------------------------------------------------------------------- + + +class TestFindSqlite3CliRefusal: + def test_missing_binary_refusal(self, monkeypatch): + monkeypatch.setattr( + "hermes_cli.session_lost_and_found.shutil.which", lambda _: None + ) + assert find_sqlite3_cli() is None + assert find_sqlite3_cli_refusal()["reason"] == "missing" + + def test_no_dbpage_refusal(self, monkeypatch): + monkeypatch.setattr( + "hermes_cli.session_lost_and_found.shutil.which", + lambda _: "/usr/bin/sqlite3", + ) + monkeypatch.setattr( + "hermes_cli.session_lost_and_found._cli_supports_recover", + lambda _: False, + ) + assert find_sqlite3_cli() is None + assert find_sqlite3_cli_refusal()["reason"] == "no_dbpage" + + def test_wal_reset_vulnerable_refusal(self, monkeypatch): + """A .recover-capable but WAL-reset-vulnerable CLI must be refused. + + This is the Debian/Ubuntu shape from the #100368 incident: the + system sqlite3 (3.45.1) has sqlite_dbpage, so the capability probe + passes, while the WAL-reset opener bug is still present. + """ + monkeypatch.setattr( + "hermes_cli.session_lost_and_found.shutil.which", + lambda _: "/usr/bin/sqlite3", + ) + monkeypatch.setattr( + "hermes_cli.session_lost_and_found._cli_supports_recover", + lambda _: True, + ) + monkeypatch.setattr( + "hermes_cli.session_lost_and_found._parse_sqlite3_cli_version", + lambda _: (3, 45, 1), + ) + assert find_sqlite3_cli() is None + refusal = find_sqlite3_cli_refusal() + assert refusal["reason"] == "wal_reset_vulnerable" + assert refusal["version"] == "3.45.1" + assert "WAL-reset" in refusal["detail"] + + def test_fixed_capable_cli_accepted(self, monkeypatch): + monkeypatch.setattr( + "hermes_cli.session_lost_and_found.shutil.which", + lambda _: "/usr/local/bin/sqlite3", + ) + monkeypatch.setattr( + "hermes_cli.session_lost_and_found._cli_supports_recover", + lambda _: True, + ) + monkeypatch.setattr( + "hermes_cli.session_lost_and_found._parse_sqlite3_cli_version", + lambda _: (3, 51, 3), + ) + assert find_sqlite3_cli() == "/usr/local/bin/sqlite3" + assert find_sqlite3_cli_refusal() == {} + + def test_unparsable_version_still_usable(self, monkeypatch): + """A CLI whose version line cannot be parsed is not refused on + version grounds alone (the salvage lane runs against a snapshot + copy, not the live file).""" + monkeypatch.setattr( + "hermes_cli.session_lost_and_found.shutil.which", + lambda _: "/usr/bin/sqlite3", + ) + monkeypatch.setattr( + "hermes_cli.session_lost_and_found._cli_supports_recover", + lambda _: True, + ) + monkeypatch.setattr( + "hermes_cli.session_lost_and_found._parse_sqlite3_cli_version", + lambda _: None, + ) + assert find_sqlite3_cli() == "/usr/bin/sqlite3" + + +class TestParseSqlite3CliVersion: + def test_parses_modern_output(self): + class Probe: + returncode = 0 + stdout = b"3.51.4 2026-XX-XX 12:34:56\n" + + with patch( + "hermes_cli.session_lost_and_found.subprocess.run", + return_value=Probe(), + ): + assert _parse_sqlite3_cli_version("x") == (3, 51, 4) + + def test_unexecutable_returns_none(self): + with patch( + "hermes_cli.session_lost_and_found.subprocess.run", + side_effect=OSError("no such file"), + ): + assert _parse_sqlite3_cli_version("x") is None + + +# --------------------------------------------------------------------------- +# The operator-facing guidance never names the live DB +# --------------------------------------------------------------------------- + + +class TestGuidanceNeverNamesLiveDb: + def test_gateway_corruption_banner(self): + """The gateway broadcast must route to the two-stage `sessions + recover` contract and must warn against pointing a raw sqlite3 + shell at the live file.""" + import gateway.run as gateway_run + + body = inspect.getsource( + gateway_run.GatewayRunner._send_session_db_warning_notifications + ) + assert LIVE_DB_SALVAGE_COMMAND not in body + assert "sessions recover --source" in body + assert "--inspect-only" in body + assert "--output" in body + assert "do NOT" in body + + def test_run_agent_corrupt_explanation(self): + from run_agent import AIAgent + + explanation = AIAgent._format_turn_completion_explanation( + "session_persistence_failed", "corrupt" + ) + assert LIVE_DB_SALVAGE_COMMAND not in explanation + assert "hermes sessions recover --source" in explanation + assert "--inspect-only" in explanation + assert "--output recovered-state.db" in explanation + assert ".recover" in explanation # the warning still names the hazard + + def test_repair_budget_error_names_safe_lane(self, tmp_path: Path): + import hermes_state_repair as hermes_state # helper lives in the split-out repair module + + message = hermes_state._persistent_repair_exhausted_error( + tmp_path / "state.db" + ) + assert "Manual recovery required" in message + assert "sessions recover --source" in message + assert "--inspect-only" in message + assert "--output recovered-state.db" in message + # The old shape embedded the live path straight into a raw sqlite3 + # command: `sqlite3 {db_path} ".recover"`. + assert ".recover\"`" not in message + assert "do NOT" in message + + def test_forensic_backup_refusals_name_safe_lane(self): + """The low-disk and stat-failure forensic backup refusal strings + must not embed a raw sqlite3 command against the live path.""" + import hermes_state + + body = inspect.getsource(hermes_state._backup_db_file) + assert ".recover\"`" not in body + assert "sessions recover --source" in body + assert "--inspect-only" in body + + def test_kanban_manual_recovery_warns_about_live_db(self): + import hermes_cli.kanban_ops as kanban # ``_cmd_repair`` lives here (split from hermes_cli.kanban) + + source = inspect.getsource(kanban) + assert '`sqlite3 kanban.db ".recover"`' not in source + assert "copy kanban.db aside FIRST" in source + + +# --------------------------------------------------------------------------- +# The emitted command satisfies the real CLI contract +# --------------------------------------------------------------------------- +# The reviewer's blocker on the first iteration of this fix: the banners +# printed `hermes sessions recover --source ` — which cmd_sessions +# rejects with exit 2 ("--output is required unless --inspect-only is +# used") before any snapshot is taken. These tests dispatch the EXACT argv +# shapes the banners emit through the real parser + cmd_sessions, so a +# guidance string can never again pass a source-substring test while the +# command it prints deterministically fails. + + +class TestEmittedCommandsSatisfyCliContract: + """Every `sessions recover` argv the guidance prints must be accepted + by the real CLI contract — the reviewer's blocker on the first + iteration of this fix was exactly this: the banners printed + `hermes sessions recover --source `, which cmd_sessions rejects + with exit 2 ("--output is required unless --inspect-only is used") + before any snapshot is taken. + + These tests dispatch the EXACT argv shapes the banners emit through + the real `cmd_sessions` (the same function `hermes` main() hands the + parsed namespace to), so a guidance string can never again pass a + source-substring test while the command it prints deterministically + fails. + """ + + @staticmethod + def _namespace(source: Path, **overrides) -> "argparse.Namespace": + """The namespace hermes main() produces for `sessions recover`. + + Mirrors the registrations in hermes_cli/main.py (sessions_recover + subparser): --source, --output, --inspect-only, --work-dir, + --chunk-size (default 1000), --allow-partial, --report. + """ + fields = dict( + sessions_action="recover", + source=source, + output=None, + inspect_only=False, + work_dir=None, + chunk_size=1000, + allow_partial=False, + report=None, + ) + fields.update(overrides) + return argparse.Namespace(**fields) + + def test_old_v1_shape_is_still_rejected(self, tmp_path): + """Guard the test's own premise: the bare `--source ` shape the + v1 banner printed (neither --inspect-only nor --output) is rejected + with rc 2 by the real dispatcher.""" + import hermes_cli.sessions_cmd as sc + + rc = sc.cmd_sessions(self._namespace(tmp_path / "state.db")) + assert rc == 2 + + def test_inspect_stage_dispatches_past_gate(self, tmp_path): + """`--inspect-only` (stage 1 of the emitted sequence) must pass + the contract gate and reach actual inspection work (rc 0/1, not + the gate's 2).""" + import hermes_cli.sessions_cmd as sc + + source = tmp_path / "state.db" + conn = sqlite3.connect(str(source)) + try: + conn.execute("CREATE TABLE t (x)") + conn.commit() + finally: + conn.close() + + rc = sc.cmd_sessions( + self._namespace(source, inspect_only=True) + ) + assert rc != 2, "--inspect-only shape must pass the contract gate" + + def test_output_stage_dispatches_past_gate(self, tmp_path): + """`--output recovered-state.db` (stage 2) must pass the contract + gate and reach actual recovery work (rc 0/1, not the gate's 2).""" + import hermes_cli.sessions_cmd as sc + + source = tmp_path / "state.db" + conn = sqlite3.connect(str(source)) + try: + conn.execute("CREATE TABLE t (x)") + conn.commit() + finally: + conn.close() + + rc = sc.cmd_sessions( + self._namespace(source, output=tmp_path / "recovered-state.db") + ) + assert rc != 2, "--output shape must pass the contract gate" + + def test_banner_strings_emit_only_contract_valid_argv(self, tmp_path): + """The exact argv shapes embedded in the guidance strings, when + parsed and dispatched, must never return the contract-gate 2. + + Extracts each `sessions recover` invocation printed by the + banners' code and runs its flag set through the real dispatcher. + """ + import hermes_cli.sessions_cmd as sc + + source = tmp_path / "state.db" + conn = sqlite3.connect(str(source)) + try: + conn.execute("CREATE TABLE t (x)") + conn.commit() + finally: + conn.close() + + # Every emitted flag-set from the five guidance sites. Stage 1 + # (inspect) and stage 2 (output) as printed by the banners: + emitted_shapes = [ + {"inspect_only": True}, # --inspect-only + {"output": tmp_path / "recovered-state.db"}, # --output + ] + for overrides in emitted_shapes: + rc = sc.cmd_sessions(self._namespace(source, **overrides)) + assert rc != 2, ( + f"emitted shape {overrides} must pass the cmd_sessions " + "contract gate — the banner is printing a command the CLI " + "rejects before doing anything" + ) diff --git a/tests/hermes_cli/test_update_autostash.py b/tests/hermes_cli/test_update_autostash.py index f7147dfa5c..b86cd898da 100644 --- a/tests/hermes_cli/test_update_autostash.py +++ b/tests/hermes_cli/test_update_autostash.py @@ -91,7 +91,7 @@ def _setup_update_mocks(monkeypatch, tmp_path): monkeypatch.setattr(hermes_main, "_restore_stashed_changes", lambda *a, **kw: True) monkeypatch.setattr(hermes_config, "get_missing_env_vars", lambda required_only=True: []) monkeypatch.setattr(hermes_config, "get_missing_config_fields", lambda: []) - monkeypatch.setattr(hermes_config, "check_config_version", lambda: (5, 5)) + monkeypatch.setattr(hermes_config, "check_config_version", lambda **_kwargs: (5, 5)) monkeypatch.setattr(hermes_config, "migrate_config", lambda **kw: {"env_added": [], "config_added": []}) monkeypatch.setattr(hermes_main, "_upgrade_pip_before_lazy_refresh", lambda *a, **kw: None) monkeypatch.setattr(hermes_main, "_refresh_active_lazy_features", lambda *a, **kw: True) diff --git a/tests/hermes_cli/test_web_server.py b/tests/hermes_cli/test_web_server.py index 776ada2cc9..808a77c4d8 100644 --- a/tests/hermes_cli/test_web_server.py +++ b/tests/hermes_cli/test_web_server.py @@ -440,6 +440,9 @@ class TestWebServerEndpoints: legacy = sqlite3.connect(str(db_path)) try: + # SQLite refuses DROP COLUMN while an index references the + # column; a pre-column legacy store has neither. + legacy.execute("DROP INDEX IF EXISTS idx_sessions_effective_activity") legacy.execute(f"ALTER TABLE sessions DROP COLUMN {missing_column}") legacy.commit() finally: @@ -486,6 +489,7 @@ class TestWebServerEndpoints: legacy = sqlite3.connect(str(db_path)) try: + legacy.execute("DROP INDEX IF EXISTS idx_sessions_effective_activity") legacy.execute("ALTER TABLE sessions DROP COLUMN last_activity_at") legacy.commit() finally: diff --git a/tests/hermes_state/test_bounded_recent_sessions.py b/tests/hermes_state/test_bounded_recent_sessions.py new file mode 100644 index 0000000000..6211ff7de0 --- /dev/null +++ b/tests/hermes_state/test_bounded_recent_sessions.py @@ -0,0 +1,240 @@ +"""Regression coverage for latency-bounded recent-session browsing.""" + +import sqlite3 +import time + +import pytest + +from hermes_state import SessionDB + + +@pytest.fixture +def db(tmp_path): + return SessionDB(tmp_path / "state.db") + + +def _set_activity(db, session_id, when): + db._conn.execute( + "UPDATE sessions SET last_activity_at = ? WHERE id = ?", + (when, session_id), + ) + db._conn.commit() + + +def test_bounded_recent_uses_effective_activity_index(db): + indexes = { + row[0] + for row in db._conn.execute( + "SELECT name FROM sqlite_master WHERE type = 'index'" + ).fetchall() + } + assert "idx_sessions_effective_activity" in indexes + + +def test_writable_startup_reconciles_legacy_activity_column_before_index(tmp_path): + """A pre-last_activity_at store must heal through the real startup path.""" + path = tmp_path / "legacy-state.db" + original = SessionDB(path) + original.close() + + conn = sqlite3.connect(path) + try: + conn.execute("DROP INDEX IF EXISTS idx_sessions_effective_activity") + conn.execute("ALTER TABLE sessions DROP COLUMN last_activity_at") + conn.commit() + finally: + conn.close() + + healed = SessionDB(path) + try: + columns = { + row[1] for row in healed._conn.execute("PRAGMA table_info(sessions)") + } + indexes = { + row[0] + for row in healed._conn.execute( + "SELECT name FROM sqlite_master WHERE type = 'index'" + ) + } + assert "last_activity_at" in columns + assert "idx_sessions_effective_activity" in indexes + assert healed.list_recent_sessions_bounded(limit=1) == [] + finally: + healed.close() + + +def test_bounded_recent_orders_by_durable_activity_and_shapes_preview(db): + now = time.time() + db.create_session("older", source="cli") + db.append_message("older", role="user", content="older preview") + db.create_session("newer", source="cli") + db.append_message("newer", role="user", content="newer preview") + _set_activity(db, "older", now - 20) + _set_activity(db, "newer", now - 10) + + rows = db.list_recent_sessions_bounded(limit=2) + + assert [row["id"] for row in rows] == ["newer", "older"] + assert rows[0]["preview"] == "newer preview" + + +def test_bounded_recent_maps_recent_compression_tip_to_logical_root(db): + now = time.time() + db.create_session("root", source="cli") + db.append_message("root", role="user", content="root preview") + db.end_session("root", "compression") + db.create_session("tip", source="cli", parent_session_id="root") + db.append_message("tip", role="user", content="tip preview") + _set_activity(db, "root", now - 1000) + _set_activity(db, "tip", now) + + rows = db.list_recent_sessions_bounded(limit=1) + + assert rows[0]["id"] == "tip" + assert rows[0]["_lineage_root_id"] == "root" + assert rows[0]["preview"] == "tip preview" + + +def test_bounded_recent_keeps_reset_child_user_visible(db): + now = time.time() + db.create_session("before-reset", source="cli", session_key="cli:one") + db.end_session("before-reset", "session_reset") + db.create_session( + "after-reset", + source="cli", + parent_session_id="before-reset", + session_key="cli:one", + ) + db.append_message("after-reset", role="user", content="fresh conversation") + _set_activity(db, "after-reset", now) + + rows = db.list_recent_sessions_bounded(limit=5) + + assert "after-reset" in [row["id"] for row in rows] + + +def test_bounded_recent_keeps_branch_separate_from_compression_parent(db): + now = time.time() + db.create_session("branch-parent", source="cli") + db.end_session("branch-parent", "compression") + db.create_session( + "branch-child", + source="cli", + parent_session_id="branch-parent", + model_config={"_branched_from": "branch-parent"}, + ) + db.append_message("branch-child", role="user", content="branch preview") + _set_activity(db, "branch-child", now) + + rows = db.list_recent_sessions_bounded(limit=5) + + branch = next(row for row in rows if row["id"] == "branch-child") + assert branch.get("_lineage_root_id") is None + + +def test_bounded_recent_excludes_delegated_children_and_sources(db): + now = time.time() + db.create_session("visible", source="cli") + db.append_message("visible", role="user", content="visible") + _set_activity(db, "visible", now - 1) + db.create_session( + "delegated", + source="cli", + model_config={"_delegate_from": "parent"}, + ) + db.append_message("delegated", role="user", content="hidden delegate") + _set_activity(db, "delegated", now) + db.create_session("hidden-source", source="cron") + db.append_message("hidden-source", role="user", content="hidden source") + _set_activity(db, "hidden-source", now + 1) + + rows = db.list_recent_sessions_bounded( + limit=5, + exclude_sources=["cron"], + ) + + assert [row["id"] for row in rows] == ["visible"] + + +def test_bounded_recent_omits_deep_lineage_when_traversal_cap_is_reached(db): + now = time.time() + parent = None + for i in range(40): + sid = f"deep-{i}" + db.create_session(sid, source="cli", parent_session_id=parent) + if parent is not None: + db.end_session(parent, "compression") + _set_activity(db, sid, now + i) + parent = sid + db.create_session("visible-deep-peer", source="cli") + _set_activity(db, "visible-deep-peer", now + 100) + + rows = db.list_recent_sessions_bounded( + limit=5, + candidate_limit=8, + lineage_limit=8, + ) + + assert [row["id"] for row in rows] == ["visible-deep-peer"] + + +def test_bounded_recent_omits_branching_lineage_at_total_row_cap(db): + now = time.time() + db.create_session("fanout-root", source="cli") + db.end_session("fanout-root", "compression") + for i in range(40): + sid = f"fanout-{i}" + db.create_session(sid, source="cli", parent_session_id="fanout-root") + _set_activity(db, sid, now + i) + db.create_session("visible-fanout-peer", source="cli") + _set_activity(db, "visible-fanout-peer", now + 100) + + rows = db.list_recent_sessions_bounded( + limit=5, + candidate_limit=8, + lineage_limit=8, + ) + + assert [row["id"] for row in rows] == ["visible-fanout-peer"] + + +def test_bounded_recent_cycle_is_deduplicated_and_omitted(db): + now = time.time() + db.create_session("cycle-a", source="cli") + db.create_session("cycle-b", source="cli", parent_session_id="cycle-a") + db.end_session("cycle-a", "compression") + db.end_session("cycle-b", "compression") + db._conn.execute( + "UPDATE sessions SET parent_session_id = ? WHERE id = ?", + ("cycle-b", "cycle-a"), + ) + _set_activity(db, "cycle-a", now) + _set_activity(db, "cycle-b", now + 1) + db.create_session("visible-cycle-peer", source="cli") + _set_activity(db, "visible-cycle-peer", now + 2) + + rows = db.list_recent_sessions_bounded( + limit=5, + candidate_limit=8, + lineage_limit=8, + ) + + assert [row["id"] for row in rows] == ["visible-cycle-peer"] + + +def test_bounded_recent_deadline_interrupts_sqlite(db): + for i in range(300): + sid = f"session-{i}" + db.create_session(sid, source="cli") + db.append_message(sid, role="user", content=f"message {i}") + + with pytest.raises(TimeoutError, match="recent-session browse exceeded"): + db.list_recent_sessions_bounded( + limit=20, + candidate_limit=300, + timeout_seconds=0.0, + ) + + # The progress handler is removed in finally: the same connection remains + # usable after cancellation instead of poisoning subsequent gateway reads. + assert db.get_session("session-0")["id"] == "session-0" \ No newline at end of file diff --git a/tests/providers/test_meta_ai_profile.py b/tests/providers/test_meta_ai_profile.py index 91ecdcdd6e..5a0ea074ce 100644 --- a/tests/providers/test_meta_ai_profile.py +++ b/tests/providers/test_meta_ai_profile.py @@ -8,6 +8,7 @@ bundled, so profiles resolve through normal registry discovery. import pytest from providers import get_provider_profile +from providers.base import ProviderProfile def _profile(): @@ -27,10 +28,29 @@ class TestMetaAIProfile: assert p.api_mode == "codex_responses" assert "MODEL_API_KEY" in p.env_vars assert p.supports_vision is True + # Images are accepted on user turns only; tool-result envelopes 400 (#101668). + assert p.supports_vision_tool_messages is False assert p.default_aux_model == "muse-spark-1.2-contributor" assert p.default_max_tokens == 16384 - assert "muse-spark-1.2-contributor" in p.fallback_models - assert "muse-spark-1.2" in p.fallback_models + assert p.fallback_models == ("muse-spark-1.2",) + + def test_live_catalog_filters_non_chat_models(self, monkeypatch): + p = _profile() + seen = [] + + def fake_fetch_models(_self, **_kwargs): + seen.append(True) + return [ + "muse-voice-transcribe-1.0", + "muse-spark-latest", + "muse-image-1.0-eval", + "muse-nova-test", + ] + + monkeypatch.setattr(ProviderProfile, "fetch_models", fake_fetch_models) + + assert p.fetch_models() == ["muse-spark-latest", "muse-nova-test"] + assert seen @pytest.mark.parametrize("alias", ["meta", "muse", "muse-spark", "model-api", "msl"]) def test_aliases_resolve(self, alias): diff --git a/tests/run_agent/test_corruption_recovery_guidance.py b/tests/run_agent/test_corruption_recovery_guidance.py index e4b5063ff7..155d326e20 100644 --- a/tests/run_agent/test_corruption_recovery_guidance.py +++ b/tests/run_agent/test_corruption_recovery_guidance.py @@ -31,6 +31,27 @@ def test_format_turn_completion_corrupt_includes_recovery_options(): assert "Freeing disk space will not help" in explanation +def test_format_turn_completion_corrupt_never_names_the_live_db(): + """The 'corrupt' cause must not direct a raw sqlite3 shell at the live DB. + + #100368 forensics: the system sqlite3 CLI on Debian/Ubuntu (3.45.1/ + 3.46.1, below the 3.51.x WAL-reset fix) unlinks the live WAL/SHM pair + when pointed at a live state.db, splitting the store into two + generations whose acknowledged writes vanish. The guidance that ships + in the corruption banner must be the snapshot-copying + `hermes sessions recover` lane. + """ + from run_agent import AIAgent + + explanation = AIAgent._format_turn_completion_explanation( + "session_persistence_failed", "corrupt" + ) + assert "sessions recover" in explanation + assert 'sqlite3 ~/.hermes/state.db ".recover"' not in explanation + # The replacement guidance names the safe command. + assert "hermes sessions recover --source" in explanation + + def test_format_turn_completion_disk_still_advises_space(): """The 'disk' cause still gives disk-space advice (unchanged).""" from run_agent import AIAgent diff --git a/tests/run_agent/test_infinite_compaction_loop.py b/tests/run_agent/test_infinite_compaction_loop.py index 79c6b734be..132c55a5d4 100644 --- a/tests/run_agent/test_infinite_compaction_loop.py +++ b/tests/run_agent/test_infinite_compaction_loop.py @@ -262,3 +262,52 @@ class TestCodexSparkShortSessionBoundary: f"This would cause the silent context wipe described in #48621." ) assert comp.has_content_to_compress(messages) is True + + +class TestPressureRealFloor: + """Regression: Cyrillic-heavy sessions under-count in the rough estimate, + letting real prompts ride the provider window (64,842→64,995 observed) + while the pre-API gate saw sub-threshold pressure.""" + + def _compressor(self, last_real, awaiting=False): + class _C: + last_real_prompt_tokens = last_real + awaiting_real_usage_after_compression = awaiting + return _C() + + def test_real_floor_lifts_undercounted_rough(self): + from agent.conversation_loop import _pressure_with_real_floor + assert _pressure_with_real_floor(self._compressor(64_842), 45_000) == 64_842 + + def test_rough_wins_when_larger(self): + from agent.conversation_loop import _pressure_with_real_floor + assert _pressure_with_real_floor(self._compressor(30_000), 45_000) == 45_000 + + def test_stale_real_ignored_right_after_compaction(self): + from agent.conversation_loop import _pressure_with_real_floor + compressor = self._compressor(64_842, awaiting=True) + assert _pressure_with_real_floor(compressor, 20_000) == 20_000 + + def test_zero_and_missing_real_are_safe(self): + from agent.conversation_loop import _pressure_with_real_floor + assert _pressure_with_real_floor(self._compressor(0), 10_000) == 10_000 + assert _pressure_with_real_floor(object(), 10_000) == 10_000 + + def test_anchored_pressure_is_never_floored(self): + """A valid usage anchor is provider-exact and wins as-is. + + On MoA turns the anchor deliberately uses the pre-fold aggregator + usage while ``last_real_prompt_tokens`` holds the folded figure; + flooring the anchored value would re-add the advisor fan-out tokens + the anchor exists to exclude. Pin the wiring shape: the floor is + applied only on the ``else`` (rough fallback) branch. + """ + import inspect + from agent import turn_request_assembly + + src = inspect.getsource(turn_request_assembly.assemble_api_request) + i = src.index("if _anchored_pressure is not None:") + window = src[i : i + 400] + assert "request_pressure_tokens = _anchored_pressure" in window + assert "else:" in window + assert window.index("else:") < window.index("_pressure_with_real_floor(") diff --git a/tests/run_agent/test_proactive_prune_loop_wiring.py b/tests/run_agent/test_proactive_prune_loop_wiring.py index 957799f466..259bbbb25b 100644 --- a/tests/run_agent/test_proactive_prune_loop_wiring.py +++ b/tests/run_agent/test_proactive_prune_loop_wiring.py @@ -189,6 +189,32 @@ class TestProactivePruneLoopWiring: assert tool_rows, "expected tool rows in the final transcript" assert all(m["content"] == marker for m in tool_rows) + def test_should_compress_true_but_skipped_is_warned(self, agent): + """``should_compress_info`` says RUN (``(True, None)``) yet this branch + was taken — the per-turn compression budget is spent. Over threshold + with no reclamation running must not be swallowed silently (#101889). + + Faithful to the real engine: ``should_compress()`` is + ``should_compress_info()[0]``, so the only way into this branch with + ``(True, None)`` is an exhausted per-turn budget.""" + agent.max_compression_attempts = 0 # budget already spent this turn + agent.context_compressor.should_compress.return_value = True + agent.context_compressor.should_compress_info.return_value = (True, None) + agent.context_compressor.prune_tool_results_only = ( + lambda messages, current_tokens=None: (messages, 0) + ) + warned = [] + with patch.object( + agent, + "_warn_context_overflow_blocked", + side_effect=lambda reason, tokens, threshold: warned.append(reason), + ): + result = _run_tool_loop(agent, n_tool_iterations=1) + + assert result["completed"] is True + assert warned, "over-threshold turn with no compaction ran silently" + assert all(r.startswith("attempts_exhausted") for r in warned) + def test_noop_input_object_commits_nothing(self, agent): """Engine returns the INPUT object with a (bogus) non-zero count — the caller's ``result is not input`` gate must refuse the commit.""" diff --git a/tests/run_agent/test_streamed_text_accumulation.py b/tests/run_agent/test_streamed_text_accumulation.py new file mode 100644 index 0000000000..12c903b22a --- /dev/null +++ b/tests/run_agent/test_streamed_text_accumulation.py @@ -0,0 +1,173 @@ +"""Tests for how a turn's streamed assistant text is built up. + +The text used to be grown with ``+=`` on an attribute. Python cannot grow a +string in place there, so every delta copied the whole thing again and a long +reply cost the square of its length in copying. The text is now held as a list +of pieces and joined when something reads it. + +These tests cover the behaviour callers depend on, plus a check on the stored +pieces that fails if the copying ever comes back. +""" +from unittest.mock import patch + +import pytest + + +def _make_agent(): + from run_agent import AIAgent + + agent = AIAgent( + api_key="test-key", + base_url="https://openrouter.ai/api/v1", + model="test/model", + quiet_mode=True, + skip_context_files=True, + skip_memory=True, + ) + agent.api_mode = "chat_completions" + agent._interrupt_requested = False + return agent + + +class TestStreamedTextValue: + """The value callers read must not change.""" + + def test_starts_empty(self): + agent = _make_agent() + assert agent._current_streamed_assistant_text == "" + + def test_deltas_join_in_order(self): + agent = _make_agent() + for piece in ["Hello", ", ", "world", "!"]: + agent._record_streamed_assistant_text(piece) + assert agent._current_streamed_assistant_text == "Hello, world!" + + def test_reading_twice_gives_the_same_answer(self): + agent = _make_agent() + agent._record_streamed_assistant_text("one ") + agent._record_streamed_assistant_text("two") + first = agent._current_streamed_assistant_text + second = agent._current_streamed_assistant_text + assert first == second == "one two" + + def test_reading_does_not_stop_later_deltas(self): + agent = _make_agent() + agent._record_streamed_assistant_text("before ") + assert agent._current_streamed_assistant_text == "before " + agent._record_streamed_assistant_text("after") + assert agent._current_streamed_assistant_text == "before after" + + def test_direct_assignment_still_works(self): + # Several call sites set this attribute straight, both to seed a value + # and to clear it between turns. + agent = _make_agent() + agent._record_streamed_assistant_text("thrown away") + agent._current_streamed_assistant_text = "set by hand" + assert agent._current_streamed_assistant_text == "set by hand" + agent._record_streamed_assistant_text(" plus more") + assert agent._current_streamed_assistant_text == "set by hand plus more" + + def test_clearing_resets_to_empty(self): + agent = _make_agent() + agent._record_streamed_assistant_text("left over") + agent._current_streamed_assistant_text = "" + assert agent._current_streamed_assistant_text == "" + agent._record_streamed_assistant_text("new turn") + assert agent._current_streamed_assistant_text == "new turn" + + def test_empty_and_non_string_deltas_are_ignored(self): + agent = _make_agent() + agent._record_streamed_assistant_text("keep") + agent._record_streamed_assistant_text("") + agent._record_streamed_assistant_text(None) # type: ignore[arg-type] + agent._record_streamed_assistant_text(12345) # type: ignore[arg-type] + assert agent._current_streamed_assistant_text == "keep" + + def test_superseded_writer_is_still_fenced_out(self): + # The single-writer guard (#65991) must keep working now that the + # text is stored as pieces. + agent = _make_agent() + agent._record_streamed_assistant_text("allowed") + with patch.object(agent, "_stream_writer_superseded", return_value=True): + agent._record_streamed_assistant_text("blocked") + assert agent._current_streamed_assistant_text == "allowed" + + +class TestStreamedTextCost: + """Adding a delta must not touch the text already collected. + + Checked by looking at the stored pieces rather than by timing, so the + test gives the same answer on a busy CI box as it does on a quiet one. + """ + + def test_each_delta_is_stored_as_its_own_piece(self): + agent = _make_agent() + for i in range(500): + agent._record_streamed_assistant_text(f"delta-{i} ") + # One piece per delta means nothing joined or copied the text that was + # already there. If a delta ever rebuilds the whole string again, this + # collapses to a single piece and the test fails. + assert len(agent._streamed_assistant_text_parts) == 500 + + def test_reading_the_text_does_not_collapse_the_pieces(self): + # Collapsing on read would drop any delta that lands between the join + # and the write back, so reading has to leave the pieces alone. + agent = _make_agent() + for i in range(10): + agent._record_streamed_assistant_text(str(i)) + assert agent._current_streamed_assistant_text == "0123456789" + assert len(agent._streamed_assistant_text_parts) == 10 + + def test_a_long_reply_is_assembled_correctly(self): + agent = _make_agent() + delta = "x" * 8 + for _ in range(20000): + agent._record_streamed_assistant_text(delta) + assert agent._current_streamed_assistant_text == delta * 20000 + assert len(agent._streamed_assistant_text_parts) == 20000 + + +def _agent_with_sink(): + agent = _make_agent() + delivered = [] + agent.stream_delta_callback = delivered.append + agent._stream_callback = None + return agent, delivered + + +class TestFireStreamDeltaEmptiness: + """_fire_stream_delta used to join the whole reply on every token just + to decide whether to strip leading newlines. That check now looks at + the parts list. + """ + + def test_first_delta_strips_leading_newlines(self): + agent, delivered = _agent_with_sink() + agent._fire_stream_delta("\n\nhello") + assert delivered == ["hello"] + assert agent._current_streamed_assistant_text == "hello" + + def test_later_delta_keeps_leading_newlines(self): + agent, delivered = _agent_with_sink() + agent._fire_stream_delta("hello") + agent._fire_stream_delta("\n\nworld") + assert delivered == ["hello", "\n\nworld"] + assert agent._current_streamed_assistant_text == "hello\n\nworld" + + def test_after_clear_the_next_delta_strips_again(self): + agent, delivered = _agent_with_sink() + agent._fire_stream_delta("hello") + agent._current_streamed_assistant_text = "" + agent._fire_stream_delta("\n\nagain") + assert delivered[-1] == "again" + assert agent._current_streamed_assistant_text == "again" + + def test_fire_path_stores_one_piece_per_delta(self): + agent, _delivered = _agent_with_sink() + for i in range(200): + agent._fire_stream_delta(f"d{i} ") + assert len(agent._streamed_assistant_text_parts) == 200 + + +if __name__ == "__main__": + raise SystemExit(pytest.main([__file__, "-q"])) diff --git a/tests/run_agent/test_streaming.py b/tests/run_agent/test_streaming.py index e5b0471e64..7a04ab0e68 100644 --- a/tests/run_agent/test_streaming.py +++ b/tests/run_agent/test_streaming.py @@ -305,8 +305,171 @@ class TestStreamingAccumulator: assert tc[0].function.name == "terminal" assert tc[0].function.arguments == '{"command": "ls"}' + @patch("run_agent.AIAgent._create_request_openai_client") + @patch("run_agent.AIAgent._close_request_openai_client") + def test_tool_argument_deltas_are_collected_without_concatenating_each_chunk( + self, mock_close, mock_create + ): + """Large tool arguments must not rebuild the accumulated string per delta.""" + from run_agent import AIAgent + class AppendOnlyChunk(str): + def __radd__(self, other): + raise AssertionError("tool argument delta was concatenated eagerly") + chunks = [ + _make_stream_chunk(tool_calls=[ + _make_tool_call_delta( + index=0, tc_id="call_123", name="write_file" + ) + ]), + _make_stream_chunk(tool_calls=[ + _make_tool_call_delta( + index=0, arguments=AppendOnlyChunk('{"path":"out.txt",') + ) + ]), + _make_stream_chunk(tool_calls=[ + _make_tool_call_delta( + index=0, arguments=AppendOnlyChunk('"content":"hello"}') + ) + ]), + _make_stream_chunk(finish_reason="tool_calls"), + ] + mock_client = MagicMock() + mock_client.chat.completions.create.return_value = iter(chunks) + mock_create.return_value = mock_client + agent = AIAgent( + api_key="test-key", + base_url="https://openrouter.ai/api/v1", + model="test/model", + quiet_mode=True, + skip_context_files=True, + skip_memory=True, + ) + agent.api_mode = "chat_completions" + agent._interrupt_requested = False + + response = agent._interruptible_streaming_api_call({}) + + tool_call = response.choices[0].message.tool_calls[0] + assert tool_call.function.arguments == ( + '{"path":"out.txt","content":"hello"}' + ) + + @patch("run_agent.AIAgent._create_request_openai_client") + @patch("run_agent.AIAgent._close_request_openai_client") + @patch("agent.relay_llm.stream") + def test_relay_finalizer_emits_joined_tool_arguments( + self, mock_relay_stream, mock_close, mock_create + ): + """Relay receives the public string shape, not buffered fragments.""" + from run_agent import AIAgent + + captured = {} + fake_stream = MagicMock() + fake_stream.final_response = None + fake_stream.__iter__.return_value = iter([ + _make_stream_chunk(tool_calls=[ + _make_tool_call_delta( + index=0, + tc_id="call_123", + name="search", + arguments='{"q":', + ) + ]), + _make_stream_chunk(tool_calls=[ + _make_tool_call_delta(index=0, arguments='"hello"}') + ]), + _make_stream_chunk(finish_reason="tool_calls"), + ]) + + def relay_stream_impl(*args, **kwargs): + captured["finalizer"] = kwargs["finalizer"] + return fake_stream + + mock_relay_stream.side_effect = relay_stream_impl + mock_client = MagicMock() + mock_client.chat.completions.create.return_value = iter([]) + mock_create.return_value = mock_client + agent = AIAgent( + api_key="test-key", + base_url="https://openrouter.ai/api/v1", + model="test/model", + quiet_mode=True, + skip_context_files=True, + skip_memory=True, + ) + agent.api_mode = "chat_completions" + agent._interrupt_requested = False + + agent._interruptible_streaming_api_call({}) + + payload = captured["finalizer"]() + tool_calls = payload["choices"][0]["message"]["tool_calls"] + assert len(tool_calls) == 1 + assert tool_calls[0]["function"] == { + "name": "search", + "arguments": '{"q":"hello"}', + } + + @patch("run_agent.AIAgent._create_request_openai_client") + @patch("run_agent.AIAgent._close_request_openai_client") + def test_tool_argument_assembly_is_chunk_boundary_invariant( + self, mock_close, mock_create + ): + """Argument bytes are identical across ASCII and Unicode fragment sizes.""" + import json + + from run_agent import AIAgent + + payload = json.dumps( + {"path": "/tmp/x", "content": "héllo wörld 日本語 " * 50}, + ensure_ascii=False, + ) + + def assemble(fragment_size): + fragments = [ + payload[i : i + fragment_size] + for i in range(0, len(payload), fragment_size) + ] + chunks = [ + _make_stream_chunk(tool_calls=[ + _make_tool_call_delta( + index=0, + tc_id="call_123", + name="write_file", + arguments=fragments[0], + ) + ]) + ] + chunks.extend( + _make_stream_chunk(tool_calls=[ + _make_tool_call_delta(index=0, arguments=fragment) + ]) + for fragment in fragments[1:] + ) + chunks.append(_make_stream_chunk(finish_reason="tool_calls")) + + mock_client = MagicMock() + mock_client.chat.completions.create.return_value = iter(chunks) + mock_create.return_value = mock_client + agent = AIAgent( + api_key="test-key", + base_url="https://openrouter.ai/api/v1", + model="test/model", + quiet_mode=True, + skip_context_files=True, + skip_memory=True, + ) + agent.api_mode = "chat_completions" + agent._interrupt_requested = False + + response = agent._interruptible_streaming_api_call({}) + return response.choices[0].message.tool_calls[0].function.arguments + + for fragment_size in (len(payload), 64, 7, 3, 1): + arguments = assemble(fragment_size) + assert arguments.encode("utf-8") == payload.encode("utf-8") # ── Test: Streaming Callbacks ──────────────────────────────────────────── diff --git a/tests/state/test_fts_rebuild_admission.py b/tests/state/test_fts_rebuild_admission.py index ac923c6a6c..cac3ef1690 100644 --- a/tests/state/test_fts_rebuild_admission.py +++ b/tests/state/test_fts_rebuild_admission.py @@ -622,3 +622,51 @@ class TestDeferredFtsRetryInProcess: assert ro.retry_deferred_fts_recovery() is False finally: ro.close() + + def test_retry_skips_quarantined_handle(self, tmp_path, fast_timeout): + """A structurally corrupt handle must never run a full FTS rebuild — + the housekeeping tick calls this unconditionally for the life of a + long-running gateway process, so a stale-FTS flag left set on a + now-corrupt handle must not retry the rebuild forever against the + damaged image (real DDL/DML the quarantine exists to prevent).""" + db_path = tmp_path / "state.db" + d = SessionDB(db_path=db_path) + if not d._fts_enabled: + d.close() + pytest.skip("FTS5 unavailable in this build") + d.create_session("s1", source="test") + d.append_message("s1", "user", "hello quarantine") + d.close() + self._mark_stale(db_path) + + # Force the open-time recovery to defer (foreign rebuild-lock + # holder) so _fts_stale is still True once the handle is open — + # mirrors test_retry_is_non_blocking_while_live_holder_and_backs_off. + with _rebuild_lock_held_by_other_process(db_path): + gw = SessionDB(db_path=db_path) + try: + assert gw._fts_stale is True + gw._db_corrupt = True + gw._db_corrupt_reason = "database disk image is malformed" + # A retry that is DUE (backoff already elapsed) on a handle that + # had been backing off before it tripped quarantine. Seeding the + # deadline in the past matters: a future deadline would make the + # unguarded code short-circuit on the backoff check and this test + # would pass without the quarantine guard ever being exercised. + gw._fts_stale_retry_after = time.monotonic() - 1.0 + gw._fts_stale_retry_interval = 900.0 + assert gw.retry_deferred_fts_recovery() is False + # Untouched: still marked stale, triggers still absent — no + # rebuild ran against the "damaged" handle. + assert gw._fts_stale is True + # The backoff bookkeeping is reset too, mirroring the success + # path's own reset — a doubled interval left behind a flag + # nothing currently clears would otherwise make the next real + # retry (if this handle is ever un-quarantined) start from a + # stale multi-minute backoff instead of the default. + assert gw._fts_stale_retry_after == 0.0 + assert gw._fts_stale_retry_interval == 0.0 + finally: + gw.close() + assert _meta_value(db_path, FTS_STALE_KEY) == "1" + assert _base_fts_triggers(db_path) == set() diff --git a/tests/state/test_fts_runtime_rebuild.py b/tests/state/test_fts_runtime_rebuild.py index ac030f4cb6..d6f0a6eeea 100644 --- a/tests/state/test_fts_runtime_rebuild.py +++ b/tests/state/test_fts_runtime_rebuild.py @@ -16,11 +16,11 @@ rebuild later, outside the failed live write/search operation. import json import os import sqlite3 -from types import SimpleNamespace import pytest import hermes_state +import hermes_state_holders import hermes_state_schema from hermes_state import ( FTS_REBUILD_DEFERRAL_KEY, @@ -140,49 +140,55 @@ class TestRuntimeFtsRebuild: } ) - def test_foreign_holder_detection_includes_deleted_wal( - self, db, tmp_path, monkeypatch - ): - db_path = tmp_path / "state.db" + @pytest.mark.parametrize( + "argv", + ( + ("journalctl", "-u", "hermes-agent.service"), + ("grep", "hermes-agent", "/var/log/syslog"), + ( + "/usr/sbin/tailscaled", + "be-child", + "ssh", + "--cmd=python -m hermes_cli.main gateway", + ), + ("tmux", "new-session", "/opt/hermes-agent/.venv/bin/hermes gateway"), + ("python3", "/opt/hermes-agent/tools/check_state.py"), + ("hermes-monitor", "gateway"), + ("hermesctl", "serve"), + ("python3", "worker.py", "hermes_cli.main"), + ("python3", "-m", "other.module", "hermes_cli.main"), + ("python3", "-c", "hermes_cli.main"), + ("python3", "-Icprint('hermes_cli.main')", "hermes_cli/main.py"), + ), + ) + def test_uninspectable_non_hermes_process_is_not_a_holder(self, argv): + assert not hermes_state_holders._looks_like_hermes(argv) - class FakePsutil: - @staticmethod - def process_iter(_attrs): - return iter( - ( - SimpleNamespace( - info={ - "pid": 111, - "open_files": [SimpleNamespace(path=str(db_path))], - } - ), - SimpleNamespace( - info={ - "pid": 222, - "open_files": [ - SimpleNamespace(path=f"{db_path}-wal (deleted)") - ], - } - ), - SimpleNamespace( - info={ - "pid": 333, - "open_files": [SimpleNamespace(path=str(tmp_path / "other.db"))], - } - ), - ) - ) - - monkeypatch.setattr(hermes_state, "psutil", FakePsutil) - monkeypatch.setattr(hermes_state, "_IS_WINDOWS", False) - monkeypatch.setattr(hermes_state.os, "getpid", lambda: 111) - # Force the macOS/psutil path even on Linux test runners - monkeypatch.setattr(hermes_state.sys, "platform", "darwin") - - assert db._foreign_state_db_holders() == [ - (222, f"{db_path}-wal (deleted)") - ] + @pytest.mark.parametrize( + "argv", + ( + ("/usr/local/bin/hermes", "gateway"), + ("/usr/local/bin/hermes-agent", "serve"), + ("/usr/local/bin/hermes-acp", "--stdio"), + ("/usr/bin/python3", "-m", "hermes_cli.main", "gateway"), + ("/usr/bin/python3", "-m", "acp_adapter"), + ("/usr/bin/python3", "-Im", "hermes_cli.main", "gateway"), + ("/usr/bin/python3", "-mhermes_cli.main", "gateway"), + ("/usr/bin/python3", "-W", "ignore", "-m", "hermes_cli.main"), + ("/usr/bin/python3", "-Xdev", "-m", "hermes_cli.main"), + ( + "/opt/hermes-agent/.venv/bin/python", + "/opt/hermes-agent/hermes_cli/main.py", + "gateway", + ), + ("python.exe", "--", "hermes_cli/main.py", "gateway"), + ("python3", "/opt/hermes-agent/run_agent.py", "--query", "hello"), + ), + ) + def test_uninspectable_hermes_process_remains_a_holder(self, argv): + assert hermes_state_holders._looks_like_hermes(argv) + @pytest.mark.linux_only def test_foreign_holder_detection_proc_readlink_deleted_wal( self, db, tmp_path, monkeypatch ): @@ -209,24 +215,93 @@ class TestRuntimeFtsRebuild: other.touch() os.symlink(str(other), str(proc_root / "333" / "fd" / "3")) - monkeypatch.setattr(hermes_state, "_IS_WINDOWS", False) - monkeypatch.setattr(hermes_state.os, "getpid", lambda: 111) - monkeypatch.setattr(hermes_state.sys, "platform", "linux") + monkeypatch.setattr(hermes_state_holders.os, "getpid", lambda: 111) real_listdir = os.listdir def _listdir(path): if isinstance(path, str): path = path.replace("/proc", str(proc_root)) return real_listdir(path) - monkeypatch.setattr(hermes_state.os, "listdir", _listdir) + monkeypatch.setattr(hermes_state_holders.os, "listdir", _listdir) real_readlink = os.readlink def _readlink(path): path = path.replace("/proc", str(proc_root)) return real_readlink(path) - monkeypatch.setattr(hermes_state.os, "readlink", _readlink) + monkeypatch.setattr(hermes_state_holders.os, "readlink", _readlink) + real_stat = os.stat + def _stat(path, *args, **kwargs): + path_s = str(path).replace("/proc", str(proc_root)) + if path_s.endswith("/222/fd/3"): + # A real /proc fd remains statable after unlink and retains + # the deleted sidecar's filesystem identity: same device as + # state.db, but an inode no live watched path can reach. + fields = list(real_stat(db_path)) + fields[1] += 1000 + return os.stat_result(fields) + return real_stat(path_s, *args, **kwargs) + monkeypatch.setattr(hermes_state_holders.os, "stat", _stat) - holders = db._foreign_state_db_holders() + holders = hermes_state_holders.foreign_state_db_holders(db_path) assert holders == [(222, db_path_wal + " (deleted)")] + @pytest.mark.linux_only + @pytest.mark.parametrize("different_device", (False, True)) + def test_foreign_holder_ignores_same_path_with_different_file_identity( + self, db, tmp_path, monkeypatch, different_device + ): + """A namespace peer's different state.db is not a holder of the host's. + + A peer process can appear in /proc with a string-identical path for a + different inode, either on the same filesystem or a different one. + Matching on path alone -- or on device alone -- would defer automatic + FTS maintenance forever while corruption compounds. + + Identity must come from (st_dev, st_ino), not the path text. + """ + db_path = tmp_path / "state.db" + + proc_root = tmp_path / "proc" + for pid in (111, 222): + (proc_root / str(pid) / "fd").mkdir(parents=True) + # PID 222 = container process holding ITS OWN state.db, which happens + # to have the identical absolute path inside its mount namespace. + guest_db = tmp_path / "guest_state.db" + guest_db.touch() + os.symlink(str(guest_db), str(proc_root / "222" / "fd" / "3")) + + monkeypatch.setattr(hermes_state_holders.os, "getpid", lambda: 111) + + real_listdir = os.listdir + def _listdir(path): + if isinstance(path, str): + path = path.replace("/proc", str(proc_root)) + return real_listdir(path) + monkeypatch.setattr(hermes_state_holders.os, "listdir", _listdir) + + # The guest fd reports the host's path (identical string), which is + # exactly what the kernel shows across mount namespaces. + def _readlink(path): + path = path.replace("/proc", str(proc_root)) + if path.endswith(f"{proc_root}/222/fd/3") or "222" in path: + return str(db_path) + return os.readlink(path) + monkeypatch.setattr(hermes_state_holders.os, "readlink", _readlink) + + # ...but stat()ing the descriptor resolves to the peer's own inode. + real_stat = os.stat + def _stat(path, *a, **kw): + path_s = str(path).replace("/proc", str(proc_root)) + st = real_stat(path_s, *a, **kw) + if different_device and path_s.endswith("/222/fd/3"): + fields = list(st) + # os.stat_result positional layout: st_dev is index 2. + fields[2] = st.st_dev + 1000 + return os.stat_result(fields) + return st + monkeypatch.setattr(hermes_state_holders.os, "stat", _stat) + + assert hermes_state_holders.foreign_state_db_holders(db_path) == [] + + @pytest.mark.linux_only def test_foreign_holder_uninspectable_process_cmdline_fallback( self, db, tmp_path, monkeypatch ): @@ -241,32 +316,34 @@ class TestRuntimeFtsRebuild: os.chmod(proc_root / "222" / "fd", 0o000) # PID 222's cmdline is world-readable and looks like Hermes cmdline_path = proc_root / "222" / "cmdline" - cmdline_path.write_bytes(b"python3\x00hermes_cli.main\x00chat\x00") + cmdline_path.write_bytes( + b"python3\x00-m\x00hermes_cli.main\x00chat\x00" + ) - monkeypatch.setattr(hermes_state, "_IS_WINDOWS", False) - monkeypatch.setattr(hermes_state.os, "getpid", lambda: 111) - monkeypatch.setattr(hermes_state.sys, "platform", "linux") + monkeypatch.setattr(hermes_state_holders.os, "getpid", lambda: 111) real_listdir = os.listdir def _listdir(path): if isinstance(path, str): + if path == "/proc/222/fd": + raise PermissionError(path) path = path.replace("/proc", str(proc_root)) return real_listdir(path) - monkeypatch.setattr(hermes_state.os, "listdir", _listdir) - # _read_proc_cmdline opens /proc//cmdline directly; redirect + monkeypatch.setattr(hermes_state_holders.os, "listdir", _listdir) + # _read_proc_argv opens /proc//cmdline directly; redirect # it to our fake proc tree. - def _fake_cmdline(pid): + def _fake_argv(pid): fake_path = str(proc_root / str(pid) / "cmdline") try: with open(fake_path, "rb") as f: raw = f.read() if not raw: return None - return raw.replace(b"\x00", b" ").decode("utf-8", "replace").strip() + return raw.decode("utf-8", "replace").rstrip("\x00").split("\x00") except OSError: return None - monkeypatch.setattr(hermes_state, "_read_proc_cmdline", _fake_cmdline) + monkeypatch.setattr(hermes_state_holders, "_read_proc_argv", _fake_argv) - holders = db._foreign_state_db_holders() + holders = hermes_state_holders.foreign_state_db_holders(db_path) # Should include PID 222 with the cmdline info assert len(holders) == 1 assert holders[0][0] == 222 diff --git a/tests/state/test_fts_trigram_cron_exclusion.py b/tests/state/test_fts_trigram_cron_exclusion.py new file mode 100644 index 0000000000..5e115cc67d --- /dev/null +++ b/tests/state/test_fts_trigram_cron_exclusion.py @@ -0,0 +1,254 @@ +"""Cron-source exclusion from the external-content trigram FTS index.""" + +from __future__ import annotations + +import sqlite3 + +import pytest + +from hermes_state import FTS_TRIGRAM_SQL, SCHEMA_VERSION, SessionDB + + +@pytest.fixture +def db(tmp_path): + session_db = SessionDB(db_path=tmp_path / "state.db") + if not session_db._trigram_available: + session_db.close() + pytest.skip("trigram tokenizer unavailable in this SQLite build") + yield session_db + session_db.close() + + +def _trigram_rowids(db: SessionDB) -> set[int]: + return { + row[0] + for row in db._conn.execute( + "SELECT id FROM messages_fts_trigram_docsize ORDER BY id" + ).fetchall() + } + + +def _install_pre_v27_trigram(db: SessionDB, *, with_tool_calls: bool = False) -> None: + """Recreate the pre-cron-exclusion external-content trigram boundary. + + ``with_tool_calls=True`` reproduces the FTS_STORAGE_VERSION 1 vtable + (``tool_calls`` projected) that installs upgraded before #88217 carry; + the default is the v2 column set with only the view/trigger predicates + behind, which is what the in-place v29 migration handles. + """ + cols = "content, tool_name" + (", tool_calls" if with_tool_calls else "") + vals = "new.content, new.tool_name" + (", new.tool_calls" if with_tool_calls else "") + db._conn.executescript( + f""" + DROP TRIGGER messages_fts_trigram_insert; + DROP TRIGGER messages_fts_trigram_delete; + DROP TRIGGER messages_fts_trigram_update; + DROP TABLE messages_fts_trigram; + DROP VIEW messages_fts_trigram_src; + CREATE VIEW messages_fts_trigram_src AS + SELECT id, role, content, tool_name, tool_calls + FROM messages WHERE role <> 'tool'; + CREATE VIRTUAL TABLE messages_fts_trigram USING fts5( + {cols}, + content='messages_fts_trigram_src', + content_rowid='id', + tokenize='trigram' + ); + CREATE TRIGGER messages_fts_trigram_insert AFTER INSERT ON messages + WHEN new.role <> 'tool' + BEGIN + INSERT INTO messages_fts_trigram(rowid, {cols}) + VALUES (new.id, {vals}); + END; + """ + ) + + +def test_fresh_trigram_indexes_conversations_but_not_cron(db: SessionDB): + db.create_session("cli", source="cli") + db.create_session("cron", source="cron") + cli_id = db.append_message("cli", role="user", content="交付状态正常") + cron_id = db.append_message("cron", role="user", content="定时任务状态正常") + + assert _trigram_rowids(db) == {cli_id} + assert cron_id not in _trigram_rowids(db) + assert db._conn.execute( + "SELECT id FROM messages_fts_docsize WHERE id = ?", (cron_id,) + ).fetchone() is not None + + +def test_cron_remains_searchable_via_standard_fts_and_explicit_cjk_fallback( + db: SessionDB, +): + db.create_session("cron", source="cron") + db.append_message( + "cron", role="assistant", content="quarterly archive 大别山项目 complete" + ) + + assert [row["session_id"] for row in db.search_messages("quarterly")] == [ + "cron" + ] + assert [ + row["session_id"] + for row in db.search_messages("大别山项目", source_filter=["cron"]) + ] == ["cron"] + + +def test_deferred_rebuild_does_not_reintroduce_cron(db: SessionDB): + db.create_session("cli", source="cli") + db.create_session("cron", source="cron") + cli_id = db.append_message("cli", role="assistant", content="交互会话内容") + db.append_message("cron", role="assistant", content="定时会话内容") + + with db._lock: + db._reset_fts_index_to_empty(db._conn) + db._seed_fts_rebuild_markers(db._conn, force=True) + db._conn.commit() + while db.fts_rebuild_step(): + pass + + assert _trigram_rowids(db) == {cli_id} + + +def test_existing_external_layout_rebuilds_trigram_on_upgrade(tmp_path): + db_path = tmp_path / "state.db" + old = SessionDB(db_path=db_path) + if not old._trigram_available: + old.close() + pytest.skip("trigram tokenizer unavailable in this SQLite build") + # The virtual table keeps referring to the view by name, so this recreates + # the exact old external-content boundary without reading source text. + _install_pre_v27_trigram(old) + old.create_session("cli", source="cli") + old.create_session("cron", source="cron") + cli_id = old.append_message("cli", role="user", content="交互迁移内容") + cron_id = old.append_message("cron", role="user", content="定时迁移内容") + assert _trigram_rowids(old) == {cli_id, cron_id} + old._conn.execute("UPDATE schema_version SET version = ?", (SCHEMA_VERSION - 1,)) + old._conn.commit() + old.close() + + migrated = SessionDB(db_path=db_path) + try: + assert _trigram_rowids(migrated) == {cli_id} + view_sql = migrated._conn.execute( + "SELECT sql FROM sqlite_master " + "WHERE type = 'view' AND name = 'messages_fts_trigram_src'" + ).fetchone()[0] + assert "sessions" in view_sql + assert "cron" in view_sql + migrated._conn.execute( + "INSERT INTO messages_fts_trigram(messages_fts_trigram) VALUES('integrity-check')" + ) + finally: + migrated.close() + + +def test_install_already_at_v28_still_gets_the_cron_exclusion_migration(tmp_path): + """The migration gate must fire for installs that were on main's v28. + + The original PR gated on ``current_version < 27``; main had meanwhile + reached SCHEMA_VERSION 28 via column-reconciliation bumps, so a v28 + database would have skipped the rebuild and kept cron rows in the trigram + index forever. Pin the gate against the version main actually shipped. + """ + db_path = tmp_path / "state.db" + old = SessionDB(db_path=db_path) + if not old._trigram_available: + old.close() + pytest.skip("trigram tokenizer unavailable in this SQLite build") + _install_pre_v27_trigram(old) + old.create_session("cli", source="cli") + old.create_session("cron", source="cron") + cli_id = old.append_message("cli", role="user", content="交互迁移内容") + cron_id = old.append_message("cron", role="user", content="定时迁移内容") + assert _trigram_rowids(old) == {cli_id, cron_id} + old._conn.execute("UPDATE schema_version SET version = 28") + old._conn.commit() + old.close() + + migrated = SessionDB(db_path=db_path) + try: + assert _trigram_rowids(migrated) == {cli_id}, ( + "a v28 database kept cron rows in the trigram index: the migration gate did not fire" + ) + finally: + migrated.close() + + +def test_v1_tool_calls_layout_is_left_for_optimize_storage(tmp_path): + """A FTS_STORAGE_VERSION 1 trigram vtable (``tool_calls`` projected) must + survive the v29 startup migration untouched and be finished by the opt-in + ``optimize_fts_storage`` path — not half-migrated into a view/vtable + column mismatch (which used to fail the rebuild with + ``no such column: T.tool_calls``).""" + db_path = tmp_path / "state.db" + old = SessionDB(db_path=db_path) + if not old._trigram_available: + old.close() + pytest.skip("trigram tokenizer unavailable in this SQLite build") + _install_pre_v27_trigram(old, with_tool_calls=True) + old.create_session("cli", source="cli") + old.create_session("cron", source="cron") + cli_id = old.append_message("cli", role="user", content="交互迁移内容") + cron_id = old.append_message("cron", role="user", content="定时迁移内容") + assert _trigram_rowids(old) == {cli_id, cron_id} + old._conn.execute("UPDATE schema_version SET version = 28") + old._conn.commit() + old.close() + + migrated = SessionDB(db_path=db_path) # must not raise + try: + # Startup left the v1 layout alone (cron row still there) … + assert _trigram_rowids(migrated) == {cli_id, cron_id} + assert migrated.fts_optimize_available() is True + # … and the opt-in path completes the transition: v2 columns, + # cron-filtered view, cron row purged. + migrated.optimize_fts_storage() + cols = [r[1] for r in migrated._conn.execute("PRAGMA table_info(messages_fts_trigram)")] + assert "tool_calls" not in cols + assert _trigram_rowids(migrated) == {cli_id} + finally: + migrated.close() + + +def test_partial_upgrade_view_does_not_skip_historical_rebuild(tmp_path): + db_path = tmp_path / "state.db" + old = SessionDB(db_path=db_path) + if not old._trigram_available: + old.close() + pytest.skip("trigram tokenizer unavailable in this SQLite build") + _install_pre_v27_trigram(old) + old.create_session("cron", source="cron") + cron_id = old.append_message("cron", role="assistant", content="迁移中断内容") + assert _trigram_rowids(old) == {cron_id} + + # Simulate a crash after new DDL landed but before the rebuild/schema stamp. + for name in ( + "messages_fts_trigram_insert", + "messages_fts_trigram_delete", + "messages_fts_trigram_update", + ): + old._conn.execute(f"DROP TRIGGER IF EXISTS {name}") + old._conn.execute("DROP VIEW messages_fts_trigram_src") + old._conn.executescript(FTS_TRIGRAM_SQL) + old._conn.execute("UPDATE schema_version SET version = ?", (SCHEMA_VERSION - 1,)) + old._conn.commit() + old.close() + + migrated = SessionDB(db_path=db_path) + try: + assert _trigram_rowids(migrated) == set() + finally: + migrated.close() + + +def test_delete_of_unindexed_cron_row_keeps_trigram_consistent(db: SessionDB): + db.create_session("cron", source="cron") + cron_id = db.append_message("cron", role="user", content="不会进入索引") + assert cron_id not in _trigram_rowids(db) + + db._conn.execute("DELETE FROM messages WHERE id = ?", (cron_id,)) + db._conn.execute( + "INSERT INTO messages_fts_trigram(messages_fts_trigram) VALUES('integrity-check')" + ) diff --git a/tests/state/test_fts_trigram_subagent_exclusion.py b/tests/state/test_fts_trigram_subagent_exclusion.py new file mode 100644 index 0000000000..00d1a63fa6 --- /dev/null +++ b/tests/state/test_fts_trigram_subagent_exclusion.py @@ -0,0 +1,161 @@ +"""Delegate-child (subagent) transcripts stay out of the trigram FTS index (v30). + +Mirrors ``test_fts_trigram_cron_exclusion.py``: children are canonical rows +in ``messages`` and stay searchable through the standard ``messages_fts`` +word index; only the trigram (CJK substring) shadow index skips them. +""" + +from __future__ import annotations + +import pytest + +from hermes_state import SCHEMA_VERSION, SessionDB +from hermes_state_common import FTS_TRIGRAM_EXCLUDED_SOURCES, fts_trigram_session_sql + + +@pytest.fixture +def db(tmp_path): + session_db = SessionDB(db_path=tmp_path / "state.db") + if not session_db._trigram_available: + session_db.close() + pytest.skip("trigram tokenizer unavailable in this SQLite build") + yield session_db + session_db.close() + + +def _trigram_rowids(db: SessionDB) -> set[int]: + return { + row[0] + for row in db._conn.execute("SELECT id FROM messages_fts_trigram_docsize").fetchall() + } + + +def _fts_rowids(db: SessionDB) -> set[int]: + return { + row[0] for row in db._conn.execute("SELECT id FROM messages_fts_docsize").fetchall() + } + + +def _seed(db: SessionDB) -> dict[str, int]: + db.create_session("root", source="cli") + # delegate_tool children: source='subagent' via platform, plus the + # _delegate_from creation marker. + db.create_session( + "kid", source="subagent", parent_session_id="root", + model_config={"_delegate_from": "root"}, + ) + # A child spawned under a gateway turn inherits the gateway's source but + # still carries the marker. + db.create_session( + "gw-kid", source="telegram", parent_session_id="root", + model_config={"_delegate_from": "root"}, + ) + # Compression continuation: parent_session_id but NO marker -> indexed. + db.create_session("cont", source="cli", parent_session_id="root") + return { + "root": db.append_message("root", role="user", content="交付状态正常 root-word"), + "kid": db.append_message("kid", role="assistant", content="子任务状态正常 kid-word"), + "gw-kid": db.append_message("gw-kid", role="assistant", content="网关子任务 gwkid-word"), + "cont": db.append_message("cont", role="assistant", content="继续会话内容 cont-word"), + } + + +def test_subagent_rows_skip_trigram_but_stay_in_standard_fts(db: SessionDB): + ids = _seed(db) + assert _trigram_rowids(db) == {ids["root"], ids["cont"]} + assert _fts_rowids(db) >= set(ids.values()) + + +def test_subagent_rows_remain_word_searchable(db: SessionDB): + _seed(db) + assert [r["session_id"] for r in db.search_messages("kid-word")] == ["kid"] + assert [r["session_id"] for r in db.search_messages("gwkid-word")] == ["gw-kid"] + # Explicit CJK search scoped to the excluded source falls back to LIKE. + assert [ + r["session_id"] + for r in db.search_messages("子任务状态", source_filter=["subagent"]) + ] == ["kid"] + # Top-level CJK substring search unaffected. + assert [r["session_id"] for r in db.search_messages("交付状态")] == ["root"] + + +def test_update_and_delete_of_unindexed_child_row_keep_trigram_consistent(db: SessionDB): + ids = _seed(db) + db._conn.execute( + "UPDATE messages SET content = ? WHERE id = ?", ("改写后的内容", ids["kid"]) + ) + db._conn.execute("DELETE FROM messages WHERE id = ?", (ids["kid"],)) + db._conn.execute( + "INSERT INTO messages_fts_trigram(messages_fts_trigram) VALUES('integrity-check')" + ) + assert _trigram_rowids(db) == {ids["root"], ids["cont"]} + + +def test_deferred_rebuild_does_not_reintroduce_children(db: SessionDB): + ids = _seed(db) + with db._lock: + db._reset_fts_index_to_empty(db._conn) + db._seed_fts_rebuild_markers(db._conn, force=True) + db._conn.commit() + while db.fts_rebuild_step(): + pass + assert _trigram_rowids(db) == {ids["root"], ids["cont"]} + assert _fts_rowids(db) >= set(ids.values()) + + +def test_full_rebuild_honours_exclusion(db: SessionDB): + ids = _seed(db) + db.rebuild_fts() + assert _trigram_rowids(db) == {ids["root"], ids["cont"]} + + +def test_v29_install_purges_child_rows_on_upgrade(tmp_path): + db_path = tmp_path / "state.db" + old = SessionDB(db_path=db_path) + if not old._trigram_available: + old.close() + pytest.skip("trigram tokenizer unavailable in this SQLite build") + # Recreate the v29 (cron-only) view/trigger boundary. + old._conn.executescript( + """ + DROP TRIGGER messages_fts_trigram_insert; + DROP TRIGGER messages_fts_trigram_delete; + DROP TRIGGER messages_fts_trigram_update; + DROP VIEW messages_fts_trigram_src; + CREATE VIEW messages_fts_trigram_src AS + SELECT m.id, m.role, m.content, m.tool_name + FROM messages AS m JOIN sessions AS s ON s.id = m.session_id + WHERE m.role <> 'tool' AND s.source <> 'cron'; + CREATE TRIGGER messages_fts_trigram_insert AFTER INSERT ON messages + WHEN new.role <> 'tool' + AND EXISTS (SELECT 1 FROM sessions WHERE id = new.session_id AND source <> 'cron') + BEGIN + INSERT INTO messages_fts_trigram(rowid, content, tool_name) + VALUES (new.id, new.content, new.tool_name); + END; + """ + ) + ids = _seed(old) + assert _trigram_rowids(old) == set(ids.values()) + old._conn.execute("UPDATE schema_version SET version = 29") + old._conn.commit() + old.close() + + migrated = SessionDB(db_path=db_path) + try: + assert _trigram_rowids(migrated) == {ids["root"], ids["cont"]} + assert migrated._conn.execute( + "SELECT version FROM schema_version" + ).fetchone()[0] == SCHEMA_VERSION + migrated._conn.execute( + "INSERT INTO messages_fts_trigram(messages_fts_trigram) VALUES('integrity-check')" + ) + finally: + migrated.close() + + +def test_predicate_constants_agree(): + assert "subagent" in FTS_TRIGRAM_EXCLUDED_SOURCES + assert "cron" in FTS_TRIGRAM_EXCLUDED_SOURCES + sql = fts_trigram_session_sql("s") + assert sql.startswith("s.source NOT IN (") and "s.model_config" in sql diff --git a/tests/state/test_no_locked_readers_gate.py b/tests/state/test_no_locked_readers_gate.py index e4db17d5db..0664669a9a 100644 --- a/tests/state/test_no_locked_readers_gate.py +++ b/tests/state/test_no_locked_readers_gate.py @@ -17,6 +17,13 @@ convoying on the writer lock. Methods that write under the lock are the lock's legitimate users and pass. New violations fail with the method name and the fix (route through ``_read_ctx()``). +``SessionDB`` itself is declared in ``hermes_state.py`` as +``class SessionDB(SessionSearchMixin, SessionSchemaMixin, +SessionPortabilityMixin)`` — its actual methods live across four files. +A gate that only opens ``hermes_state.py`` never sees a locked reader +declared in one of the three mixin files, so ``_ALL_STATE_SOURCES`` scans +each of them under their own class name. + Deliberately NOT flagged: - methods that INSERT/UPDATE/DELETE/REPLACE under the lock (writers); - read-modify-write methods (the read is ordered against its own write); @@ -32,7 +39,18 @@ from pathlib import Path import pytest -_STATE_PY = Path(__file__).resolve().parents[2] / "hermes_state.py" +_REPO_ROOT = Path(__file__).resolve().parents[2] +_STATE_PY = _REPO_ROOT / "hermes_state.py" + +# SessionDB's own class body lives in hermes_state.py; the rest of its +# methods come from these mixins (see module docstring). Each entry is +# (source file, class name to scan in that file). +_ALL_STATE_SOURCES: list[tuple[Path, str]] = [ + (_STATE_PY, "SessionDB"), + (_REPO_ROOT / "hermes_state_search.py", "SessionSearchMixin"), + (_REPO_ROOT / "hermes_state_schema.py", "SessionSchemaMixin"), + (_REPO_ROOT / "hermes_state_portability.py", "SessionPortabilityMixin"), +] _WRITE_RE = re.compile( r"^\s*(INSERT|UPDATE|DELETE|REPLACE|CREATE|DROP|ALTER|VACUUM|BEGIN|COMMIT|ANALYZE)\b", @@ -124,17 +142,19 @@ def _is_self_lock_with(item: ast.withitem) -> bool: ) -def _scan_locked_readers(state_py: "Path | None" = None) -> list[str]: +def _scan_locked_readers( + state_py: "Path | None" = None, class_name: str = "SessionDB" +) -> list[str]: target = state_py if state_py is not None else _STATE_PY tree = ast.parse(target.read_text(encoding="utf-8")) violations: list[str] = [] session_db = None for node in tree.body: - if isinstance(node, ast.ClassDef) and node.name == "SessionDB": + if isinstance(node, ast.ClassDef) and node.name == class_name: session_db = node break - assert session_db is not None, "SessionDB class not found" + assert session_db is not None, f"{class_name} class not found in {target}" for method in session_db.body: if not isinstance(method, (ast.FunctionDef, ast.AsyncFunctionDef)): @@ -197,9 +217,22 @@ def _scan_locked_readers(state_py: "Path | None" = None) -> list[str]: return violations +def _scan_all_state_sources() -> list[str]: + """Run ``_scan_locked_readers`` over every file that contributes methods + to ``SessionDB`` — the class body in ``hermes_state.py`` plus each mixin + it inherits from (see module docstring). Violations are prefixed with + their source filename since methods can share names across mixins. + """ + violations: list[str] = [] + for path, class_name in _ALL_STATE_SOURCES: + for v in _scan_locked_readers(path, class_name): + violations.append(f"{path.name}: {v}") + return violations + + class TestNoPureReadersUnderWriterLock: def test_no_locked_pure_readers(self): - violations = _scan_locked_readers() + violations = _scan_all_state_sources() assert violations == [], ( "Pure-read SessionDB methods holding the writer lock " "(Pattern C — every concurrent turn's persistence convoys " @@ -235,3 +268,24 @@ class TestNoPureReadersUnderWriterLock: assert flagged == { "guilty_reader", "guilty_alias_reader", "guilty_variable_sql" }, violations + + def test_scan_all_state_sources_visits_every_mixin_file(self, tmp_path): + """Sabotage self-check for the multi-file scope itself: a locked + reader planted in a MIXIN file (not hermes_state.py) must still be + caught. Guards against the gate's scope silently narrowing back to + one file — exactly how the real 2026-08 gap (9 locked readers across + three mixin files, invisible to the single-file scanner) happened. + """ + mixin_sabotage = ( + "class FakeMixin:\n" + " def guilty_mixin_reader(self):\n" + " with self._lock:\n" + " return self._conn.execute(\"SELECT 1\").fetchone()\n" + ) + p = tmp_path / "fake_mixin.py" + p.write_text(mixin_sabotage, encoding="utf-8") + + violations = [ + f"{p.name}: {v}" for v in _scan_locked_readers(p, "FakeMixin") + ] + assert any("guilty_mixin_reader" in v for v in violations), violations diff --git a/tests/state/test_state_db_holders.py b/tests/state/test_state_db_holders.py new file mode 100644 index 0000000000..4fb5c24f38 --- /dev/null +++ b/tests/state/test_state_db_holders.py @@ -0,0 +1,50 @@ +"""Behavioral tests for the state-holder and repair-admission authority.""" + +import os + +import pytest + +import hermes_state_holders + + +@pytest.mark.linux_only +def test_foreign_holder_accepts_same_inode_reached_through_an_alias( + tmp_path, monkeypatch +): + """Descriptor identity is authoritative even when /proc spells another path.""" + db_path = tmp_path / "state.db" + db_path.touch() + alias_path = tmp_path / "namespace-alias" / "state.db" + + proc_root = tmp_path / "proc" + for pid in (111, 222): + (proc_root / str(pid) / "fd").mkdir(parents=True) + os.symlink(db_path, proc_root / "222" / "fd" / "3") + + monkeypatch.setattr(hermes_state_holders.os, "getpid", lambda: 111) + real_listdir = os.listdir + + def _listdir(path): + if isinstance(path, str): + path = path.replace("/proc", str(proc_root)) + return real_listdir(path) + + monkeypatch.setattr(hermes_state_holders.os, "listdir", _listdir) + + def _readlink(path): + if path == "/proc/222/fd/3": + return str(alias_path) + return os.readlink(path.replace("/proc", str(proc_root))) + + monkeypatch.setattr(hermes_state_holders.os, "readlink", _readlink) + real_stat = os.stat + + def _stat(path, *args, **kwargs): + path = str(path).replace("/proc", str(proc_root)) + return real_stat(path, *args, **kwargs) + + monkeypatch.setattr(hermes_state_holders.os, "stat", _stat) + + assert hermes_state_holders.foreign_state_db_holders(db_path) == [ + (222, str(alias_path)) + ] diff --git a/tests/state/test_state_db_wal_unlink_race.py b/tests/state/test_state_db_wal_unlink_race.py new file mode 100644 index 0000000000..5f8bf5271c --- /dev/null +++ b/tests/state/test_state_db_wal_unlink_race.py @@ -0,0 +1,84 @@ +"""Regression coverage for WAL restoration during state.db repair (#101064). + +Journal-mode restoration used to open a NEW connection after the exclusive +repair guard had released the live database. In WAL mode a writer could still +hold the unlinked old WAL inode while that second connection created a fresh +``state.db-wal`` path — two generations of one store. The restore must run +through the guard connection, before the guard releases. +""" + +import sqlite3 + +import pytest + +import hermes_state +from hermes_state import repair_state_db_schema + + +def _make_db(path): + conn = sqlite3.connect(str(path), isolation_level=None) + conn.execute("CREATE TABLE sessions (name TEXT)") + conn.execute("INSERT INTO sessions VALUES ('seed')") + conn.close() + + +def test_wal_restoration_reuses_exclusive_repair_connection(tmp_path, monkeypatch): + """Unit contract: given the guard connection, no reopen happens.""" + db_path = tmp_path / "state.db" + conn = sqlite3.connect(db_path, isolation_level=None) + conn.execute("CREATE TABLE marker (value TEXT)") + + def fail_if_reopened(_path): + pytest.fail("WAL restoration reopened state.db outside the repair guard") + + monkeypatch.setattr(hermes_state, "_connect_repair_durable", fail_if_reopened) + + hermes_state._restore_journal_mode_after_repair(db_path, None, conn=conn) + # The mode itself is whatever apply_wal_with_fallback resolves on this + # runtime (WAL, or DELETE on WAL-reset-vulnerable SQLite builds); the + # contract under test is the connection reuse, asserted above. + assert conn.execute("PRAGMA journal_mode").fetchone()[0].lower() in ("wal", "delete") + conn.close() + + +def test_repair_never_reopens_after_the_guard_releases(tmp_path, monkeypatch): + """End to end through repair_state_db_schema: every connection the repair + opens is opened while the exclusive guard is still held, and none after.""" + db = tmp_path / "state.db" + _make_db(db) + monkeypatch.setattr(hermes_state, "_db_opens_cleanly", lambda path: "forced-unhealthy") + # The scratch-space pre-flight wants ~10GB headroom; irrelevant here. + monkeypatch.setattr(hermes_state, "_repair_scratch_space_error", lambda path: None) + + def fake_strategies(scratch_path, report): + report["repaired"] = True + report["strategy"] = "test_strategy" + return report + + monkeypatch.setattr(hermes_state, "_run_repair_strategies", fake_strategies) + + events: list[str] = [] + real_guard = hermes_state._exclusive_repair_db_guard + real_connect = hermes_state._connect_repair_durable + + from contextlib import contextmanager + + @contextmanager + def tracing_guard(path): + events.append("guard-enter") + with real_guard(path) as pair: + yield pair + events.append("guard-exit") + + def tracing_connect(path, *a, **kw): + events.append("connect") + return real_connect(path, *a, **kw) + + monkeypatch.setattr(hermes_state, "_exclusive_repair_db_guard", tracing_guard) + monkeypatch.setattr(hermes_state, "_connect_repair_durable", tracing_connect) + + report = repair_state_db_schema(db, backup=False) + assert report["repaired"] is True + assert "guard-exit" in events + after_release = events[events.index("guard-exit") + 1 :] + assert "connect" not in after_release, events diff --git a/tests/test_fts_tool_write_bounds.py b/tests/test_fts_tool_write_bounds.py new file mode 100644 index 0000000000..fb38bd60d4 --- /dev/null +++ b/tests/test_fts_tool_write_bounds.py @@ -0,0 +1,204 @@ +import sqlite3 + +import pytest + +from hermes_state import SessionDB +from hermes_state_common import ( + FTS_TOOL_CONTENT_PREFIX_CHARS, + FTS_TOOL_FULL_CONTENT_HIGH_WATER_KEY, + LEGACY_FTS_SQL, + _FTS_TRIGGERS, +) + + +def _long_message(prefix: str, tail: str) -> str: + padding = "padding " * (FTS_TOOL_CONTENT_PREFIX_CHARS // len("padding ") + 8) + return f"{prefix} {padding} {tail}" + + +@pytest.fixture +def db(tmp_path): + session_db = SessionDB(db_path=tmp_path / "state.db") + if not session_db._fts_enabled: + session_db.close() + pytest.skip("SQLite FTS5 unavailable") + session_db.create_session("session", source="cli") + try: + yield session_db + finally: + session_db.close() + + +def test_new_tool_rows_bound_fts_content_but_explicit_tool_search_is_complete(db): + tool_id = db.append_message( + "session", + role="tool", + content=_long_message("indexed-prefix-token", "tool-tail-token"), + tool_name="terminal", + ) + user_id = db.append_message( + "session", + role="user", + content=_long_message("user-prefix-token", "user-tail-token"), + ) + + assert [row["id"] for row in db.search_messages("indexed-prefix-token")] == [ + tool_id + ] + assert db.search_messages("tool-tail-token") == [] + assert [ + row["id"] + for row in db.search_messages("tool-tail-token", role_filter=["tool"]) + ] == [tool_id] + assert [row["id"] for row in db.search_messages("user-tail-token")] == [ + user_id + ] + + +def test_trigger_migration_preserves_historical_tool_tokens_without_rebuild(tmp_path): + path = tmp_path / "state.db" + first = SessionDB(db_path=path) + if not first._fts_enabled: + first.close() + pytest.skip("SQLite FTS5 unavailable") + first.create_session("session", source="cli") + + # Model the pre-migration trigger contract: every id through this artificial + # boundary receives full-content indexing. + first.set_meta(FTS_TOOL_FULL_CONTENT_HIGH_WATER_KEY, str(2**62)) + old_id = first.append_message( + "session", + role="tool", + content=_long_message("old-prefix-token", "old-tail-token"), + ) + assert [row["id"] for row in first.search_messages("old-tail-token")] == [old_id] + first._conn.execute( + "DELETE FROM state_meta WHERE key = ?", + (FTS_TOOL_FULL_CONTENT_HIGH_WATER_KEY,), + ) + first.close() + + migrated = SessionDB(db_path=path) + try: + assert int(migrated.get_meta(FTS_TOOL_FULL_CONTENT_HIGH_WATER_KEY)) == old_id + assert [ + row["id"] for row in migrated.search_messages("old-tail-token") + ] == [old_id] + + new_id = migrated.append_message( + "session", + role="tool", + content=_long_message("new-prefix-token", "new-tail-token"), + ) + assert migrated.search_messages("new-tail-token") == [] + assert [ + row["id"] + for row in migrated.search_messages( + "new-tail-token", role_filter=["tool"] + ) + ] == [new_id] + + # Historical rows still use their full old token stream for the FTS5 + # external-content delete command; redaction must remove the tail token. + migrated._execute_write( + lambda conn: conn.execute( + "UPDATE messages SET content = '' WHERE id = ?", (old_id,) + ) + ) + assert migrated.search_messages("old-tail-token") == [] + + # New bounded rows use the same prefix for delete as insert. A mismatch + # corrupts external-content FTS and makes this delete or later write fail. + migrated._execute_write( + lambda conn: conn.execute("DELETE FROM messages WHERE id = ?", (new_id,)) + ) + migrated.append_message("session", role="assistant", content="fts-still-healthy") + assert migrated.search_messages("fts-still-healthy") + finally: + migrated.close() + + +def test_full_rebuild_moves_boundary_before_future_tool_writes(db): + before_id = db.append_message( + "session", + role="tool", + content=_long_message("before-prefix-token", "before-tail-token"), + ) + assert db.search_messages("before-tail-token") == [] + + assert db.rebuild_fts() >= 1 + assert int(db.get_meta(FTS_TOOL_FULL_CONTENT_HIGH_WATER_KEY)) == before_id + assert [row["id"] for row in db.search_messages("before-tail-token")] == [ + before_id + ] + + after_id = db.append_message( + "session", + role="tool", + content=_long_message("after-prefix-token", "after-tail-token"), + ) + assert db.search_messages("after-tail-token") == [] + assert [ + row["id"] + for row in db.search_messages("after-tail-token", role_filter=["tool"]) + ] == [after_id] + + +def test_role_changes_switch_between_bounded_and_full_indexing(db): + message_id = db.append_message( + "session", + role="tool", + content=_long_message("role-prefix-token", "role-tail-token"), + ) + assert db.search_messages("role-tail-token") == [] + + db._execute_write( + lambda conn: conn.execute( + "UPDATE messages SET role = 'assistant' WHERE id = ?", (message_id,) + ) + ) + assert [row["id"] for row in db.search_messages("role-tail-token")] == [ + message_id + ] + + db._execute_write( + lambda conn: conn.execute( + "UPDATE messages SET role = 'tool' WHERE id = ?", (message_id,) + ) + ) + assert db.search_messages("role-tail-token") == [] + + +def test_legacy_inline_fts_also_bounds_new_tool_rows(tmp_path): + path = tmp_path / "legacy.db" + initial = SessionDB(db_path=path) + initial.create_session("session", source="cli") + for trigger in _FTS_TRIGGERS: + initial._conn.execute(f"DROP TRIGGER IF EXISTS {trigger}") + initial._conn.execute("DROP TABLE IF EXISTS messages_fts_trigram") + initial._conn.execute("DROP VIEW IF EXISTS messages_fts_trigram_src") + initial._conn.execute("DROP TABLE IF EXISTS messages_fts") + initial._conn.executescript(LEGACY_FTS_SQL) + initial._conn.execute( + "DELETE FROM state_meta WHERE key IN (?, 'fts_storage_version')", + (FTS_TOOL_FULL_CONTENT_HIGH_WATER_KEY,), + ) + initial.close() + + legacy = SessionDB(db_path=path) + try: + assert legacy._db_has_legacy_inline_fts(legacy._conn.cursor()) is True + message_id = legacy.append_message( + "session", + role="tool", + content=_long_message("legacy-prefix-token", "legacy-tail-token"), + ) + assert legacy.search_messages("legacy-tail-token") == [] + assert [ + row["id"] + for row in legacy.search_messages( + "legacy-tail-token", role_filter=["tool"] + ) + ] == [message_id] + finally: + legacy.close() diff --git a/tests/test_hermes_home_key_cache.py b/tests/test_hermes_home_key_cache.py new file mode 100644 index 0000000000..f421c188f8 --- /dev/null +++ b/tests/test_hermes_home_key_cache.py @@ -0,0 +1,164 @@ +"""Tests for the remembered results in `hermes_home_key`. + +`Path.resolve()` is a filesystem call. `hermes_home_key` sits under +`ToolRegistry.current_scope_key()`, which runs on every registry lookup, so +before the results were remembered the registry paid a syscall per lookup. + +The value this returns must not change, so most of these tests compare +against the plain uncached calculation. +""" + +from __future__ import annotations + +import os +from pathlib import Path + +import pytest + +import hermes_constants as hc + + +def _uncached(path=None) -> str: + """The calculation as it was before results were remembered.""" + candidate = Path(path) if path is not None else hc.get_hermes_home() + return os.path.normcase(str(candidate.expanduser().resolve(strict=False))) + + +@pytest.fixture(autouse=True) +def _clear_cache(): + hc.reset_hermes_home_key_cache() + yield + hc.reset_hermes_home_key_cache() + + +class TestSameAnswerAsBefore: + """Remembering a result must not change what comes back.""" + + @pytest.mark.parametrize("case", ["real", "missing", "trailing_sep", "dot_dot"]) + def test_matches_the_uncached_calculation(self, tmp_path, case): + (tmp_path / "real").mkdir() + target = { + "real": str(tmp_path / "real"), + "missing": str(tmp_path / "not_there"), + "trailing_sep": str(tmp_path / "real") + os.sep, + "dot_dot": str(tmp_path / "real" / ".." / "real"), + }[case] + assert hc.hermes_home_key(target) == _uncached(target) + + def test_matches_for_the_default_home(self): + assert hc.hermes_home_key() == _uncached() + + def test_matches_for_a_tilde_path(self): + assert hc.hermes_home_key("~") == _uncached("~") + + def test_accepts_a_path_object(self, tmp_path): + (tmp_path / "real").mkdir() + assert hc.hermes_home_key(tmp_path / "real") == _uncached(tmp_path / "real") + + def test_second_call_returns_the_same_string(self, tmp_path): + (tmp_path / "real").mkdir() + first = hc.hermes_home_key(str(tmp_path / "real")) + second = hc.hermes_home_key(str(tmp_path / "real")) + assert first == second == _uncached(str(tmp_path / "real")) + + +class TestWhatGetsRemembered: + def test_an_existing_path_is_remembered(self, tmp_path): + (tmp_path / "real").mkdir() + hc.hermes_home_key(str(tmp_path / "real")) + assert len(hc._HOME_KEY_CACHE) == 1 + + def test_a_missing_path_is_not_remembered(self, tmp_path): + # The answer can change once the directory is created, for example + # when part of the path turns out to be a link, so it must not stick. + hc.hermes_home_key(str(tmp_path / "not_there")) + assert hc._HOME_KEY_CACHE == {} + + def test_a_path_created_later_picks_up_the_real_answer(self, tmp_path): + later = tmp_path / "later" + before = hc.hermes_home_key(str(later)) + later.mkdir() + after = hc.hermes_home_key(str(later)) + assert after == _uncached(str(later)) + assert hc._HOME_KEY_CACHE == {str(later): after} + # On a plain directory both answers agree anyway. The point is that + # the first one was never stored. + assert before == after + + def test_different_paths_get_their_own_entries(self, tmp_path): + for name in ("a", "b", "c"): + (tmp_path / name).mkdir() + hc.hermes_home_key(str(tmp_path / name)) + assert len(hc._HOME_KEY_CACHE) == 3 + + def test_reset_clears_everything(self, tmp_path): + (tmp_path / "real").mkdir() + hc.hermes_home_key(str(tmp_path / "real")) + assert hc._HOME_KEY_CACHE + hc.reset_hermes_home_key_cache() + assert hc._HOME_KEY_CACHE == {} + + +class TestHomeChanges: + def test_pointing_hermes_home_somewhere_else_gives_a_new_key( + self, tmp_path, monkeypatch, + ): + # A different home is a different input path, so it lands on its own + # entry rather than reusing the first one. + first = tmp_path / "home_one" + second = tmp_path / "home_two" + first.mkdir() + second.mkdir() + + monkeypatch.setenv("HERMES_HOME", str(first)) + key_one = hc.hermes_home_key() + monkeypatch.setenv("HERMES_HOME", str(second)) + key_two = hc.hermes_home_key() + + assert key_one != key_two + assert key_one == _uncached(str(first)) + assert key_two == _uncached(str(second)) + + +class TestSymlinks: + def test_a_link_resolves_to_its_target(self, tmp_path): + target = tmp_path / "target" + target.mkdir() + link = tmp_path / "link" + try: + link.symlink_to(target, target_is_directory=True) + except (OSError, NotImplementedError): + pytest.skip("this platform or account cannot create symlinks") + assert hc.hermes_home_key(str(link)) == _uncached(str(link)) + assert hc.hermes_home_key(str(link)) == hc.hermes_home_key(str(target)) + + +class TestRegistryLookupsDoNotHitTheDisk: + def test_scope_key_resolves_the_path_once(self, monkeypatch): + # The reason this cache exists. ToolRegistry.current_scope_key() runs + # on every registry lookup, so it must not resolve the home path on + # the filesystem every time. + from tools.registry import registry + + calls = {"n": 0} + real_resolve = Path.resolve + + def counting_resolve(self, *a, **kw): + calls["n"] += 1 + return real_resolve(self, *a, **kw) + + monkeypatch.setattr(Path, "resolve", counting_resolve) + + registry.current_scope_key() + after_first = calls["n"] + for _ in range(50): + registry.current_scope_key() + + assert calls["n"] == after_first, ( + f"current_scope_key resolved the path on the filesystem " + f"{calls['n'] - after_first} extra times across 50 calls" + ) + + +if __name__ == "__main__": + raise SystemExit(pytest.main([__file__, "-q"])) diff --git a/tests/test_hermes_state.py b/tests/test_hermes_state.py index 0b52e1b977..3861d2cd90 100644 --- a/tests/test_hermes_state.py +++ b/tests/test_hermes_state.py @@ -11,7 +11,13 @@ import pytest import hermes_state from agent.session_activity import ActivityProvenance -from hermes_state import SCHEMA_SQL, SCHEMA_VERSION, SessionDB +from hermes_state import ( + FTS_SQL, + FTS_STORAGE_VERSION, + SCHEMA_SQL, + SCHEMA_VERSION, + SessionDB, +) class _NoFtsCursor(sqlite3.Cursor): @@ -3634,6 +3640,160 @@ class TestFTSExternalContentMigration: finally: db.close() + @pytest.mark.parametrize("with_message", [False, True]) + def test_v23_rebuild_from_trigram_tool_calls_projection( + self, tmp_path, with_message + ): + """v23 installs built with historical trigram projection should be + repaired via optimize-storage: trigram must drop tool_calls while + standard messages_fts keeps indexing them.""" + db_path = tmp_path / "v23-toolcalls.db" + + # Build an external-content DB that is already at schema version 23, + # but with the old tool_calls-inclusive trigram projection. + conn = sqlite3.connect(str(db_path)) + conn.executescript(SCHEMA_SQL) + conn.executescript(FTS_SQL) + conn.executescript( + """ + DROP TRIGGER IF EXISTS messages_fts_trigram_insert; + DROP TRIGGER IF EXISTS messages_fts_trigram_delete; + DROP TRIGGER IF EXISTS messages_fts_trigram_update; + DROP TABLE IF EXISTS messages_fts_trigram; + DROP VIEW IF EXISTS messages_fts_trigram_src; + + CREATE VIEW IF NOT EXISTS messages_fts_trigram_src AS + SELECT id, role, content, tool_name, tool_calls + FROM messages + WHERE role <> 'tool'; + + CREATE VIRTUAL TABLE messages_fts_trigram USING fts5( + content, + tool_name, + tool_calls, + content='messages_fts_trigram_src', + content_rowid='id', + tokenize='trigram' + ); + + CREATE TRIGGER messages_fts_trigram_insert AFTER INSERT ON messages + WHEN new.role <> 'tool' + BEGIN + INSERT INTO messages_fts_trigram(rowid, content, tool_name, tool_calls) + VALUES (new.id, new.content, new.tool_name, new.tool_calls); + END; + + CREATE TRIGGER messages_fts_trigram_delete AFTER DELETE ON messages + WHEN old.role <> 'tool' + BEGIN + INSERT INTO messages_fts_trigram(messages_fts_trigram, rowid, content, tool_name, tool_calls) + VALUES ('delete', old.id, old.content, old.tool_name, old.tool_calls); + END; + + CREATE TRIGGER 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 + OR old.role IS NOT new.role) + BEGIN + INSERT INTO messages_fts_trigram(messages_fts_trigram, rowid, content, tool_name, tool_calls) + SELECT 'delete', old.id, old.content, old.tool_name, old.tool_calls + WHERE old.role <> 'tool'; + INSERT INTO messages_fts_trigram(rowid, content, tool_name, tool_calls) + SELECT new.id, new.content, new.tool_name, new.tool_calls + WHERE new.role <> 'tool'; + END; + """ + ) + # Simulate the historical v23 projection shipped before this fix. + conn.execute( + "INSERT OR REPLACE INTO state_meta (key, value) VALUES ('fts_storage_version', '1')" + ) + conn.execute( + "INSERT OR REPLACE INTO state_meta (key, value) VALUES ('fts_optimize_available', '1')" + ) + conn.commit() + conn.close() + + if with_message: + conn = sqlite3.connect(str(db_path)) + conn.execute( + "INSERT INTO sessions (id, source, started_at) VALUES (?, ?, ?)", + ("s1", "cli", time.time()), + ) + conn.execute( + "INSERT INTO messages (session_id, timestamp, role, content, tool_name, tool_calls) " + "VALUES (?, ?, ?, ?, ?, ?)", + ( + "s1", + time.time(), + "assistant", + "部署完成 assistant content", + "legacyTool", + '{"name": "legacy", "arguments": "UNIQUE_TOOLCALL_TOKEN_43701"}', + ), + ) + conn.commit() + assert conn.execute( + "SELECT rowid FROM messages_fts_trigram WHERE messages_fts_trigram MATCH 'UNIQUE_TOOLCALL_TOKEN_43701'" + ).fetchall() + conn.close() + + db = SessionDB(db_path=db_path) + try: + assert db._conn is not None + assert db.fts_optimize_available() is True + assert db.get_meta("fts_storage_version") == "1" + + original_ensure = db._ensure_fts_schema + + def interrupt_after_demote(cursor, table_name, ddl): + if table_name == "messages_fts_trigram": + raise RuntimeError("injected trigram rebuild interruption") + return original_ensure(cursor, table_name, ddl) + + db._ensure_fts_schema = interrupt_after_demote + with pytest.raises(RuntimeError, match="injected trigram"): + db.optimize_fts_storage(vacuum=False) + + db.close() + db = SessionDB(db_path=db_path) + assert db._conn is not None + assert db.get_meta("fts_rebuild_high_water") is None + assert db.get_meta("fts_rebuild_progress") is None + assert db._has_fts_trash(db._conn) is True + assert db.fts_optimize_available() is True + assert db.get_meta("fts_storage_version") == "1" + if with_message: + assert db._conn.execute( + "SELECT 1 FROM messages_fts_trigram " + "WHERE messages_fts_trigram MATCH '部署完成' LIMIT 1" + ).fetchone() + + result = db.optimize_fts_storage(vacuum=False) + assert result["ok"] is True + + if with_message: + # messages_fts stays in the tool-calls search path. + assert len(db.search_messages("UNIQUE_TOOLCALL_TOKEN_43701")) == 1 + # New trigram schema excludes tool_calls from trigram projection. + assert not db._conn.execute( + "SELECT 1 FROM messages_fts_trigram WHERE messages_fts_trigram MATCH 'UNIQUE_TOOLCALL_TOKEN_43701' LIMIT 1" + ).fetchone() + assert db._conn.execute( + "SELECT 1 FROM messages_fts_trigram " + "WHERE messages_fts_trigram MATCH '部署完成' LIMIT 1" + ).fetchone() + trigger_sql = db._conn.execute( + "SELECT sql FROM sqlite_master " + "WHERE type = 'trigger' AND name = 'messages_fts_trigram_update'" + ).fetchone()[0] + assert "tool_calls" not in trigger_sql + assert db.get_meta("fts_storage_version") == str(FTS_STORAGE_VERSION) + finally: + db.close() + diff --git a/tests/test_install_ps1_managed_python_provenance.py b/tests/test_install_ps1_managed_python_provenance.py index 5c91721a32..7cccce360e 100644 --- a/tests/test_install_ps1_managed_python_provenance.py +++ b/tests/test_install_ps1_managed_python_provenance.py @@ -18,6 +18,46 @@ pytestmark = pytest.mark.windows_only _INSTALL_PS1 = Path(__file__).resolve().parents[1] / "scripts" / "install.ps1" +def test_fresh_install_manifest_orders_repo_before_checkout_scoped_python( + tmp_path: Path, +) -> None: + powershell = shutil.which("powershell") + if not powershell: + pytest.skip("Windows PowerShell is required") + + install_dir = tmp_path / "install" + run = subprocess.run( + [ + powershell, + "-NoProfile", + "-ExecutionPolicy", + "Bypass", + "-File", + str(_INSTALL_PS1), + "-Manifest", + "-HermesHome", + str(tmp_path / "hermes-home"), + "-InstallDir", + str(install_dir), + ], + cwd=tmp_path, + capture_output=True, + text=True, + check=False, + timeout=45, + ) + + assert run.returncode == 0, run.stdout + run.stderr + manifest = json.loads(run.stdout) + stages = [stage["name"] for stage in manifest["stages"]] + assert ( + stages.index("repository") + < stages.index("python") + < stages.index("venv") + ) + assert not install_dir.exists(), "manifest lookup must remain read-only" + + def _run_venv_stage( powershell: str, tmp_path: Path, diff --git a/tests/test_schema_read_probe.py b/tests/test_schema_read_probe.py index 9770a8f050..a1c0c83978 100644 --- a/tests/test_schema_read_probe.py +++ b/tests/test_schema_read_probe.py @@ -62,6 +62,9 @@ class TestSchemaReadProbeStatements: """ conn = _fresh_schema_conn() try: + # Indexes that depend on the column must be removed before SQLite + # can emulate a pre-column legacy store via DROP COLUMN. + conn.execute("DROP INDEX IF EXISTS idx_sessions_effective_activity") conn.execute("ALTER TABLE sessions DROP COLUMN last_activity_at") # The failure must come from the sessions probe naming the exact # column — not incidentally from some other statement — so a diff --git a/tests/test_state_db_repair_live_writer_guard.py b/tests/test_state_db_repair_live_writer_guard.py index 8288cb322b..161592e1cb 100644 --- a/tests/test_state_db_repair_live_writer_guard.py +++ b/tests/test_state_db_repair_live_writer_guard.py @@ -18,12 +18,18 @@ live-writer guard.) from __future__ import annotations +import errno +import select import sqlite3 +import subprocess +import sys import uuid from pathlib import Path import pytest +import hermes_state +import hermes_state_holders from hermes_state import ( SessionDB, repair_state_db_schema, @@ -58,18 +64,17 @@ def _make_wal_db(tmp_path: Path) -> Path: def test_repair_refuses_while_another_connection_holds_the_db(tmp_path): """Surgery under concurrent writers is what spread the corruption. - Gated on ``requires_wal``: ``_live_writer_holds_db`` detects an - out-of-process holder via ``PRAGMA locking_mode=EXCLUSIVE`` + a - ``BEGIN IMMEDIATE`` that a concurrent connection makes fail with - SQLITE_BUSY through the WAL index. On SQLite builds carrying the - WAL-reset bug (and on NFS/SMB) Hermes deliberately runs ``state.db`` in - ``journal_mode=DELETE``, where a held reader takes only a SHARED lock and - ``BEGIN IMMEDIATE`` can still acquire RESERVED — so the probe cannot see - the holder and the guard fails open. In DELETE mode repair is instead - serialised only by the cross-process repairer lock (see - ``_live_writer_holds_db``'s docstring). The conftest auto-skips this test - where WAL is unusable rather than assert a guarantee the runtime doesn't - make there. + Gated on ``requires_wal``: repair admission + (``hermes_state_holders.live_writer_holds_db``) first scans for foreign + holders — deleted sidecar generations, uninspectable Hermes processes — + and then probes SQLite with ``PRAGMA locking_mode=EXCLUSIVE`` + + ``BEGIN IMMEDIATE``, which a concurrent connection makes fail with + SQLITE_BUSY through the WAL index. This test exercises the probe leg: on + SQLite builds carrying the WAL-reset bug (and on NFS/SMB) Hermes runs + ``state.db`` in ``journal_mode=DELETE``, where a held reader takes only a + SHARED lock and ``BEGIN IMMEDIATE`` still acquires RESERVED, so the probe + alone cannot see the holder. The conftest auto-skips this test where WAL + is unusable rather than assert a guarantee the probe doesn't make there. """ db = _make_wal_db(tmp_path) @@ -84,6 +89,312 @@ def test_repair_refuses_while_another_connection_holds_the_db(tmp_path): assert "live writer" in (report["error"] or "").lower() +def test_repair_checks_foreign_holders_before_opening_sqlite(tmp_path, monkeypatch): + """A replacement pathname cannot expose locks on the deleted old inode.""" + db = _make_wal_db(tmp_path) + monkeypatch.setattr( + hermes_state_holders, + "foreign_state_db_holders", + lambda _path: [(4242, f"{db}-wal (deleted)")], + ) + + def _unexpected_probe(*_args, **_kwargs): + pytest.fail("repair opened SQLite before excluding foreign holders") + + monkeypatch.setattr(hermes_state, "_connect_repair_durable", _unexpected_probe) + + report = repair_state_db_schema(db, backup=False) + + assert report["repaired"] is False + assert "live writer" in (report["error"] or "").lower() + + +@pytest.mark.linux_only +def test_linux_holder_scan_does_not_require_psutil(tmp_path, monkeypatch): + """The Linux safety scan must not make psutil a repair dependency.""" + monkeypatch.setattr(hermes_state_holders, "psutil", None) + + holders = hermes_state_holders.foreign_state_db_holders( + tmp_path / "absent-state.db" + ) + + assert holders == [] + + +@pytest.mark.linux_only +def test_incomplete_holder_scan_keeps_unknown_sentinel(tmp_path, monkeypatch): + """A partial scan must not hide uncertainty behind an ordinary holder.""" + db = tmp_path / "state.db" + db.touch() + + def _listdir(path): + if path == "/proc": + return ["4242", "4343"] + if path == "/proc/4242/fd": + return ["7"] + if path == "/proc/4343/fd": + raise RuntimeError("scan interrupted") + raise AssertionError(f"unexpected scan path: {path}") + + monkeypatch.setattr(hermes_state_holders.os, "listdir", _listdir) + monkeypatch.setattr(hermes_state_holders.os, "readlink", lambda _path: str(db)) + real_stat = hermes_state_holders.os.stat + + def _stat(path, *args, **kwargs): + if path == "/proc/4242/fd/7": + return real_stat(db) + return real_stat(path, *args, **kwargs) + + monkeypatch.setattr(hermes_state_holders.os, "stat", _stat) + + holders = hermes_state_holders.foreign_state_db_holders(db) + + assert (4242, str(db)) in holders + assert any(pid < 0 and "scan interrupted" in path for pid, path in holders) + + +@pytest.mark.linux_only +def test_uninspectable_watched_descriptor_blocks_repair_before_sqlite( + tmp_path, monkeypatch +): + """A watched fd whose identity cannot be read is not proven safe.""" + db = _make_wal_db(tmp_path) + + def _listdir(path): + if path == "/proc": + return ["4242"] + if path == "/proc/4242/fd": + return ["7"] + raise AssertionError(f"unexpected scan path: {path}") + + real_stat = hermes_state_holders.os.stat + + def _stat(path, *args, **kwargs): + if path == "/proc/4242/fd/7": + raise PermissionError(errno.EACCES, "descriptor denied", path) + return real_stat(path, *args, **kwargs) + + monkeypatch.setattr(hermes_state_holders.os, "listdir", _listdir) + monkeypatch.setattr(hermes_state_holders.os, "readlink", lambda _path: str(db)) + monkeypatch.setattr(hermes_state_holders.os, "stat", _stat) + + def _unexpected_probe(*_args, **_kwargs): + pytest.fail("repair opened SQLite with unproven descriptor identity") + + monkeypatch.setattr(hermes_state, "_connect_repair_durable", _unexpected_probe) + + report = repair_state_db_schema(db, backup=False) + + assert report["repaired"] is False + assert "live writer" in (report["error"] or "").lower() + + +@pytest.mark.linux_only +@pytest.mark.parametrize( + ("argv", "should_block"), + ( + (["python3", "backup.py"], False), + (["python3", "-m", "hermes_cli.main", "gateway"], True), + ), +) +def test_uninspectable_unknown_descriptor_uses_hermes_identity_at_repair_boundary( + tmp_path, monkeypatch, argv, should_block +): + """An unknown fd target blocks only when argv identifies Hermes.""" + db = _make_wal_db(tmp_path) + + def _listdir(path): + if path == "/proc": + return ["4242"] + if path == "/proc/4242/fd": + return ["7"] + raise AssertionError(f"unexpected scan path: {path}") + + def _readlink(path): + raise PermissionError(errno.EACCES, "descriptor denied", path) + + monkeypatch.setattr(hermes_state_holders.os, "listdir", _listdir) + monkeypatch.setattr(hermes_state_holders.os, "readlink", _readlink) + monkeypatch.setattr( + hermes_state_holders, + "_read_proc_argv", + lambda _pid: argv, + ) + + if should_block: + + def _unexpected_probe(*_args, **_kwargs): + pytest.fail("repair opened SQLite with an unproven Hermes descriptor") + + monkeypatch.setattr( + hermes_state, + "_connect_repair_durable", + _unexpected_probe, + ) + report = repair_state_db_schema(db, backup=False) + + assert report["repaired"] is False + assert "live writer" in (report["error"] or "").lower() + else: + real_connect = hermes_state._connect_repair_durable + probe_reached = False + + def _record_probe(*args, **kwargs): + nonlocal probe_reached + probe_reached = True + return real_connect(*args, **kwargs) + + monkeypatch.setattr( + hermes_state, + "_connect_repair_durable", + _record_probe, + ) + report = repair_state_db_schema(db, backup=False) + + assert probe_reached is True + assert "live writer" not in (report["error"] or "").lower() + + +@pytest.mark.linux_only +def test_uninspectable_watched_identity_blocks_alias_before_sqlite( + tmp_path, monkeypatch +): + """A non-disappearance stat error cannot prove an aliased holder safe.""" + db = _make_wal_db(tmp_path) + alias = tmp_path / "namespace-alias" / "state.db" + + def _listdir(path): + if path == "/proc": + return ["4242"] + if path == "/proc/4242/fd": + return ["7"] + raise AssertionError(f"unexpected scan path: {path}") + + real_stat = hermes_state_holders.os.stat + + def _stat(path, *args, **kwargs): + if str(path) == str(db) and not args and not kwargs: + raise PermissionError(errno.EACCES, "watched identity denied", path) + if path == "/proc/4242/fd/7": + return real_stat(db) + return real_stat(path, *args, **kwargs) + + monkeypatch.setattr(hermes_state_holders.os, "listdir", _listdir) + monkeypatch.setattr( + hermes_state_holders.os, "readlink", lambda _path: str(alias) + ) + monkeypatch.setattr(hermes_state_holders.os, "stat", _stat) + + def _unexpected_probe(*_args, **_kwargs): + pytest.fail("repair opened SQLite with an unproven watched identity") + + monkeypatch.setattr(hermes_state, "_connect_repair_durable", _unexpected_probe) + + report = repair_state_db_schema(db, backup=False) + + assert report["repaired"] is False + assert "live writer" in (report["error"] or "").lower() + + +@pytest.mark.linux_only +def test_uninspectable_alias_descriptor_for_hermes_blocks_before_sqlite( + tmp_path, monkeypatch +): + """Hermes cannot make an aliased fd safe when its identity is unreadable.""" + db = _make_wal_db(tmp_path) + alias = tmp_path / "namespace-alias" / "state.db" + + def _listdir(path): + if path == "/proc": + return ["4242"] + if path == "/proc/4242/fd": + return ["7"] + raise AssertionError(f"unexpected scan path: {path}") + + real_stat = hermes_state_holders.os.stat + + def _stat(path, *args, **kwargs): + if path == "/proc/4242/fd/7": + raise PermissionError(errno.EACCES, "descriptor denied", path) + return real_stat(path, *args, **kwargs) + + monkeypatch.setattr(hermes_state_holders.os, "listdir", _listdir) + monkeypatch.setattr( + hermes_state_holders.os, "readlink", lambda _path: str(alias) + ) + monkeypatch.setattr(hermes_state_holders.os, "stat", _stat) + monkeypatch.setattr( + hermes_state_holders, + "_read_proc_argv", + lambda _pid: ["python3", "-m", "hermes_cli.main", "gateway"], + ) + + def _unexpected_probe(*_args, **_kwargs): + pytest.fail("repair opened SQLite with an unproven Hermes alias fd") + + monkeypatch.setattr(hermes_state, "_connect_repair_durable", _unexpected_probe) + + report = repair_state_db_schema(db, backup=False) + + assert report["repaired"] is False + assert "live writer" in (report["error"] or "").lower() + + +@pytest.mark.requires_wal +@pytest.mark.linux_only +def test_repair_refuses_while_foreign_process_holds_deleted_wal(tmp_path): + """Reproduce the inode split that a pathname lock probe cannot observe.""" + db = _make_wal_db(tmp_path) + holder_code = """ +import sqlite3 +import sys + +conn = sqlite3.connect(sys.argv[1]) +conn.execute("PRAGMA journal_mode=WAL") +conn.execute("BEGIN IMMEDIATE") +print("ready", flush=True) +sys.stdin.read(1) +conn.rollback() +conn.close() +""" + holder = subprocess.Popen( + [sys.executable, "-c", holder_code, str(db)], + stdin=subprocess.PIPE, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + ) + try: + assert holder.stdout is not None + readable, _, _ = select.select([holder.stdout], [], [], 10) + assert readable, "holder subprocess did not signal readiness" + assert holder.stdout.readline().strip() == "ready" + deleted = [] + for suffix in ("-wal", "-shm"): + sidecar = Path(f"{db}{suffix}") + if sidecar.exists(): + sidecar.unlink() + deleted.append(sidecar) + assert deleted + + report = repair_state_db_schema(db, backup=False) + + assert report["repaired"] is False + assert "live writer" in (report["error"] or "").lower() + finally: + if holder.poll() is None and holder.stdin is not None: + try: + holder.stdin.write("x") + holder.stdin.close() + except (BrokenPipeError, ValueError): + pass + try: + holder.wait(timeout=10) + except subprocess.TimeoutExpired: + holder.kill() + holder.wait(timeout=10) + + def test_repair_proceeds_once_the_database_is_quiescent(tmp_path): """The guard must not deadlock repair on an exclusively-held file.""" db = _make_wal_db(tmp_path) diff --git a/tests/test_state_db_repair_loop_cap.py b/tests/test_state_db_repair_loop_cap.py index 7293d7a641..0eafbcb7a3 100644 --- a/tests/test_state_db_repair_loop_cap.py +++ b/tests/test_state_db_repair_loop_cap.py @@ -82,6 +82,13 @@ class TestPersistentAttemptCap: assert report["repaired"] is False assert "Manual recovery required" in report["error"] assert ".recover" in report["error"] + assert "sessions recover" in report["error"] + assert "sqlite3" in report["error"] + # The terminal error must not direct a raw sqlite3 shell at the + # live database: a WAL-reset-vulnerable CLI (Debian/Ubuntu 3.45.x/ + # 3.46.x, pre-#100368 forensics) unlinks the live WAL/SHM pair + # and splits the store into two generations. + assert 'sqlite3 state.db ".recover"' not in report["error"] assert len(_existing_malformed_backups(db)) == backups_before def test_changed_file_resets_the_budget(self, tmp_path): diff --git a/tests/test_state_db_stats.py b/tests/test_state_db_stats.py index f64ce4821e..8e51cc73ac 100644 --- a/tests/test_state_db_stats.py +++ b/tests/test_state_db_stats.py @@ -236,6 +236,20 @@ def test_render_large_db_legacy_trigram_suggests_optimize(): assert "optimize-storage" in blob +def test_render_large_db_v1_trigram_suggests_optimize(): + from hermes_cli.doctor import STATE_DB_SIZE_WARN_BYTES, _render_state_db_stats + + lines = _render_state_db_stats( + _base_stats( + logical_size_bytes=STATE_DB_SIZE_WARN_BYTES + 1, + fts_storage_version=1, + ), + holders=None, + ) + blob = " ".join(" ".join(str(p) for p in line) for line in lines) + assert "optimize-storage" in blob + + def test_render_does_not_duplicate_legacy_wal_warning(): """A large WAL must NOT warn here: doctor's pre-existing WAL check (50 MB threshold, with a --fix checkpoint) already covers it, and a diff --git a/tests/tools/file_ops_fakes.py b/tests/tools/file_ops_fakes.py new file mode 100644 index 0000000000..4cbb2d992b --- /dev/null +++ b/tests/tools/file_ops_fakes.py @@ -0,0 +1,61 @@ +"""Fakes for ``ShellFileOperations``' compound shell probes. + +``read_file`` and ``write_file`` ask the shell everything in ONE command +whose stdout is split on a per-call random sentinel line. Test doubles that +script ``env.execute`` / ``_exec`` need to answer that command with exactly +the stream the shell would produce; these helpers build it. Match the +sentinel out of the command first (it is random), then compose: + + m = READ_SENTINEL_RE.search(command) + if m: + return {"output": compound_read_output(m.group(0), size=5, sample=b"hello", + content="hello\\n", total_lines=1), + "returncode": 0} +""" + +import base64 +import re +from typing import Optional + +READ_SENTINEL_RE = re.compile(r"__HERMES_RF_[0-9a-f]{32}__") +WRITE_SENTINEL_RE = re.compile(r"__HERMES_WF_[0-9a-f]{32}__") + + +def compound_read_output( + sentinel: str, + *, + size: int, + sample: Optional[bytes], + content: str, + total_lines: int, + trailing_newline: bool = True, + sample_rc: int = 0, + read_rc: int = 0, +) -> str: + """Stdout of ``_read_probe_cmd`` for a regular file. + + ``content`` is the ``sed | cut`` page exactly as the shell prints it: + every line newline-terminated (``cut`` always adds one), or ``""`` for a + page past EOF. ``sample`` is the raw first-1000-bytes slice (``None`` + emits an empty base64 segment, e.g. a shell without ``base64``). + """ + sample_seg = base64.b64encode(sample).decode() + "\n" if sample else "" + return ( + f"{size}\n{sentinel}\n" + f"{sample_seg}{sentinel}\n" + f"{content}{sentinel}\n" + f"{total_lines}\n{sentinel}\n" + f"{1 if trailing_newline else 0}\n{sentinel}\n" + f"{sample_rc} {read_rc}\n" + ) + + +def compound_write_probe_output(sentinel: str, *, head3: bytes, body: str) -> str: + """Stdout of ``_write_probe_cmd`` for an existing file. + + ``head3`` is the first three bytes on disk (BOM detection); ``body`` is + the second segment: the whole file when pre-content was wanted, else + the 4 KB line-ending sample. + """ + head_seg = base64.b64encode(head3).decode() + "\n" if head3 else "" + return f"{head_seg}{sentinel}\n{body}" diff --git a/tests/tools/test_bot_relay.py b/tests/tools/test_bot_relay.py index 7b9fb528f1..5ec1322578 100644 --- a/tests/tools/test_bot_relay.py +++ b/tests/tools/test_bot_relay.py @@ -84,6 +84,9 @@ def test_resolve_remote_target_forms(root): assert bot_relay.resolve_remote_target("default", roster)["connection_id"] == "cloud-1" # exact connection-qualified form assert bot_relay.resolve_remote_target("hermes@cloud-1", roster)["profile"] == "default" + # profile@connection — the form Desktop's mention middleware annotates + # for remote bots (#97678); the UI alias form must not be required + assert bot_relay.resolve_remote_target("default@cloud-1", roster)["profile"] == "default" assert bot_relay.resolve_remote_target("hermes@nope", roster) is None assert bot_relay.resolve_remote_target("ghost", roster) is None diff --git a/tests/tools/test_browser_lightpanda_serve.py b/tests/tools/test_browser_lightpanda_serve.py index 73200bac14..db607ec77b 100644 --- a/tests/tools/test_browser_lightpanda_serve.py +++ b/tests/tools/test_browser_lightpanda_serve.py @@ -5,12 +5,17 @@ import json import os import stat import subprocess +from pathlib import Path from unittest.mock import patch import pytest import tools.browser_lightpanda as lp +# The autouse _isolate fixture swaps _binary_supports_http_cache for a lambda; +# the probe test needs the real (lru_cache-wrapped) function back. +_real_probe = lp._binary_supports_http_cache + class FakeProc: def __init__(self, pid=4242, exit_code=None): @@ -42,6 +47,8 @@ def _isolate(tmp_path, monkeypatch): # Never touch the developer's real ~/.local/bin/lightpanda. monkeypatch.setattr(lp, "_home_candidates", lambda: []) monkeypatch.setattr(lp, "_safe_start_time", lambda pid: 111) + lp._binary_supports_http_cache.cache_clear() + monkeypatch.setattr(lp, "_binary_supports_http_cache", lambda binary: True) with lp._servers_lock: lp._servers.clear() yield state @@ -112,6 +119,7 @@ class TestLaunch: assert err is None assert calls["argv"] == [ "/opt/lightpanda", "serve", "--host", "127.0.0.1", "--port", "43111", + "--http-cache-dir", str(_isolate / "http-cache"), ] kw = calls["kwargs"] assert kw["stdin"] is subprocess.DEVNULL @@ -129,6 +137,57 @@ class TestLaunch: assert record["owner_pid"] == os.getpid() assert record["start_time"] == 111 + def test_http_cache_dir_is_shared_across_sessions(self, monkeypatch, _isolate): + _, _, first = self._launch(monkeypatch) + with lp._servers_lock: + lp._servers.clear() + _, _, second = self._launch(monkeypatch) + cache = str(_isolate / "http-cache") + assert first["argv"][first["argv"].index("--http-cache-dir") + 1] == cache + assert second["argv"][second["argv"].index("--http-cache-dir") + 1] == cache + assert Path(cache).is_dir() + assert not list(Path(cache).glob("*.json")) # never confused with a session record + + def test_no_http_cache_flag_on_old_binary(self, monkeypatch, _isolate): + monkeypatch.setattr(lp, "_binary_supports_http_cache", lambda binary: False) + _, err, calls = self._launch(monkeypatch) + assert err is None + assert "--http-cache-dir" not in calls["argv"] + assert calls["argv"][-1] == "43111" + + def test_http_cache_probe_caches_and_detects_flag(self, monkeypatch, tmp_path): + monkeypatch.setattr(lp, "_binary_supports_http_cache", _real_probe) + _real_probe.cache_clear() + exe = _exe(tmp_path / "lightpanda") + runs = [] + + def fake_run(argv, **kwargs): + runs.append(argv) + return subprocess.CompletedProcess( + argv, returncode=0, + stdout="--http-cache-dir " if len(runs) == 1 else "", + ) + + monkeypatch.setattr(lp.subprocess, "run", fake_run) + assert lp._binary_supports_http_cache(str(exe)) is True + assert lp._binary_supports_http_cache(str(exe)) is True # cached, single probe + assert len(runs) == 1 + + lp._binary_supports_http_cache.cache_clear() + + def fake_run_old(argv, **kwargs): + return subprocess.CompletedProcess(argv, returncode=0, stdout="no such flag") + + monkeypatch.setattr(lp.subprocess, "run", fake_run_old) + assert lp._binary_supports_http_cache(str(exe)) is False + + def fake_run_hangs(argv, **kwargs): + raise subprocess.TimeoutExpired(cmd=argv, timeout=3.0) + + lp._binary_supports_http_cache.cache_clear() + monkeypatch.setattr(lp.subprocess, "run", fake_run_hangs) + assert lp._binary_supports_http_cache(str(exe)) is False + def test_block_private_networks_flag(self, monkeypatch): _, err, calls = self._launch(monkeypatch, block_private_networks=True) assert err is None diff --git a/tests/tools/test_code_execution.py b/tests/tools/test_code_execution.py index c67ce0247c..2ee6492a44 100644 --- a/tests/tools/test_code_execution.py +++ b/tests/tools/test_code_execution.py @@ -475,10 +475,11 @@ class TestStubSchemaDrift(unittest.TestCase): compile(src, "hermes_tools.py", "exec") # Verify specific parameter signatures are in the source - # search_files must accept context, offset, output_mode + # search_files must accept its pagination, output, and ordering controls self.assertIn("context", src) self.assertIn("offset", src) self.assertIn("output_mode", src) + self.assertIn("order", src) # patch must accept mode and patch params self.assertIn("mode", src) diff --git a/tests/tools/test_code_kernel.py b/tests/tools/test_code_kernel.py index d7df35720b..3741ffae65 100644 --- a/tests/tools/test_code_kernel.py +++ b/tests/tools/test_code_kernel.py @@ -18,9 +18,15 @@ tests patch ``_load_config`` directly, mirroring test_code_execution_modes. import json import os +import shutil +import subprocess import sys +import tempfile +import textwrap +import time import unittest from contextlib import contextmanager +from pathlib import Path from unittest.mock import patch import pytest @@ -97,6 +103,51 @@ class TestSessionStatePersistence(unittest.TestCase): class TestKernelLifecycle(unittest.TestCase): + def test_kernel_exits_when_its_backend_parent_dies(self): + """A kernel must not outlive the host that spawned it, even when the + host dies without cleanup (SIGKILL/OOM/crash). Windows: inherited + SYNCHRONIZE handle; POSIX: inherited death pipe. Both are proven the + same way — kill the host mid-cell, the kernel is gone within seconds.""" + import psutil + + repo_root = str(Path(__file__).resolve().parents[2]) + host_src = textwrap.dedent(f""" + import json, os, sys, time + os.environ["HERMES_HOME"] = sys.argv[1] + sys.path.insert(0, {repo_root!r}) + from tools.code_kernel import SessionKernel, _spawn + k = SessionKernel(("parent-death",)) + _spawn(k, task_id="parent-death", child_python=sys.executable, + child_cwd="", sandbox_tools=frozenset(), max_tool_calls=1) + cell = json.dumps({{"id": "x", "code": "import os, time\\n" + "assert 'HERMES_KERNEL_PARENT_PROCESS_HANDLE' not in os.environ\\n" + "assert 'HERMES_KERNEL_PARENT_DEATH_FD' not in os.environ\\n" + "time.sleep(300)"}}) + "\\n" + k.proc.stdin.write(cell.encode()); k.proc.stdin.flush() + print(k.proc.pid, flush=True) + time.sleep(600) + """) + with tempfile.TemporaryDirectory() as home: + host = subprocess.Popen( + [sys.executable, "-c", host_src, home], + stdout=subprocess.PIPE, text=True, + creationflags=getattr(subprocess, "CREATE_NO_WINDOW", 0), + ) + try: + kernel = psutil.Process(int(host.stdout.readline())) + time.sleep(0.5) + self.assertTrue(kernel.is_running(), "kernel never came up") + host.kill() + host.wait(timeout=10) + try: + kernel.wait(timeout=10) + except psutil.TimeoutExpired: + kernel.kill() + self.fail("session kernel survived its backend parent") + finally: + if host.poll() is None: + host.kill() + def test_timeout_kills_the_kernel_and_reports_state_loss(self): with _kernel_config(timeout=1): slow = _run("import time\ntime.sleep(30)") @@ -283,6 +334,31 @@ class TestKernelOwnershipAndLifecycle(unittest.TestCase): stale.proc.wait(timeout=10) self.assertFalse(stale.alive()) + def test_parallel_cells_share_one_kernel_process(self): + """Parallel cells for one owner race the first spawn. Each racer + used to see proc=None as 'dead', replace the registry entry, and + orphan the winner's process — 110 live kernels under a 4-capped + process (Sep 2026). Every kernel process must stay registry-owned.""" + import subprocess + import threading + + results = [] + with _kernel_config(): + def _cell(): + results.append(self._run_as("conv-a", "import time; time.sleep(0.3)", task_id="t")) + threads = [threading.Thread(target=_cell) for _ in range(6)] + for t in threads: + t.start() + for t in threads: + t.join() + self.assertEqual([r["status"] for r in results], ["success"] * 6) + self.assertEqual(len(_KERNELS), 1) + live = subprocess.run( + ["pgrep", "-fc", "-P", str(os.getpid()), "hermes_kernel_runner"], + capture_output=True, text=True, + ).stdout.strip() + self.assertEqual(live, "1") + class TestPerCellRpcAuthority(unittest.TestCase): """Interpreter state persists across cells; RPC authority must not.""" diff --git a/tests/tools/test_code_kernel_remote.py b/tests/tools/test_code_kernel_remote.py index 9b2a1343d2..bf8470dcc2 100644 --- a/tests/tools/test_code_kernel_remote.py +++ b/tests/tools/test_code_kernel_remote.py @@ -200,6 +200,81 @@ class TestOwnershipIsolation(RemoteKernelBase): self.assertEqual(remaining_owner, "owner-b") +class TestIdleReapAndCapEviction(RemoteKernelBase): + """Unlike local session kernels, remote kernels had no idle-reap or + process-wide cap: _REMOTE_KERNELS grew one entry per distinct + (owner, env_type, task_env_id) that was never revisited, for the life + of the gateway process.""" + + def test_idle_expired_kernel_is_reaped_on_next_call(self): + env = ScriptedEnv(_spawn_ok_handlers([_cell(), _cell()])) + execute_in_remote_kernel( + "print(1)", env=env, env_type="ssh", task_env_id="stale", + sandbox_tools=frozenset(), timeout=10, max_tool_calls=5, + reset=False, idle_exit=1800, + ) + self.assertEqual(len(_REMOTE_KERNELS), 1) + # Backdate the kernel's last_used past the idle window — simulates + # a key that is never revisited again. + for kernel in _REMOTE_KERNELS.values(): + kernel.last_used -= 2000 + # A call for a DIFFERENT key must reap the stale entry on entry, + # without ever touching or reviving it. + execute_in_remote_kernel( + "print(1)", env=env, env_type="ssh", task_env_id="fresh", + sandbox_tools=frozenset(), timeout=10, max_tool_calls=5, + reset=False, idle_exit=1800, + ) + owners = {key[0] for key in _REMOTE_KERNELS} + self.assertNotIn("stale", owners) + self.assertIn("fresh", owners) + + def test_over_cap_evicts_least_recently_used(self): + with patch("tools.code_kernel._lifecycle_limits", return_value=(2, 1800)): + env = ScriptedEnv(_spawn_ok_handlers([_cell() for _ in range(10)])) + for i in range(3): + execute_in_remote_kernel( + "print(1)", env=env, env_type="ssh", task_env_id=f"owner-{i}", + sandbox_tools=frozenset(), timeout=10, max_tool_calls=5, + reset=False, idle_exit=1800, + ) + self.assertEqual(len(_REMOTE_KERNELS), 2) + owners = {key[0] for key in _REMOTE_KERNELS} + self.assertNotIn("owner-0", owners) + self.assertIn("owner-1", owners) + self.assertIn("owner-2", owners) + + def test_eviction_skips_kernels_with_a_running_cell(self): + """Cap eviction must never kill a kernel mid-cell (the local-kernel + race from hermes-agent#101861): a busy kernel stays put and a + settled one goes instead, even if the busy one is older.""" + import threading + + gate = threading.Event() + + def slow_cat(command): + gate.wait(10) + return {"output": json.dumps(_cell()), "returncode": 0} + + busy_env = ScriptedEnv([ + ("nohup", lambda c: {"output": "PID:4242\n", "returncode": 0}), + ("kill -0", lambda c: {"output": "ALIVE\n", "returncode": 0}), + ("cat ", slow_cat), + ]) + with patch("tools.code_kernel._lifecycle_limits", return_value=(1, 1800)): + worker = threading.Thread(target=_run, args=(busy_env,), kwargs={"task": "busy"}) + worker.start() + while not any(k.attached for k in _REMOTE_KERNELS.values()): + pass + env = ScriptedEnv(_spawn_ok_handlers([_cell()])) + _run(env, task="settled") + owners = {key[0] for key in _REMOTE_KERNELS} + self.assertIn("busy", owners) + gate.set() + worker.join(10) + self.assertFalse(any("kill 4242" in c for c in busy_env.commands)) + + class TestDispatchIntegration(unittest.TestCase): """_execute_remote prefers the kernel and falls open to per-call.""" diff --git a/tests/tools/test_cronjob_run_background.py b/tests/tools/test_cronjob_run_background.py index a35d4c4a8f..c120273c03 100644 --- a/tests/tools/test_cronjob_run_background.py +++ b/tests/tools/test_cronjob_run_background.py @@ -237,6 +237,28 @@ class TestInFlightDedupe: assert seen_during_run["registered"] is True assert "job-bg-09" not in sched.get_running_job_ids() # released after + def test_run_claimed_job_reports_exact_unknown_execution_not_stale_success(self): + from tools.cronjob_tools import _run_claimed_job + + def probe_run(job, **_kwargs): + job["execution_id"] = "exec-unknown" + return True + + with patch("cron.scheduler.run_one_job", side_effect=probe_run), \ + patch("cron.executions.get_execution", return_value={ + "id": "exec-unknown", + "status": "unknown", + "error": "worker owner exited", + }), \ + patch("tools.cronjob_tools.get_job", return_value={ + "last_status": "ok", + "last_error": None, + }): + res = _run_claimed_job(_job("job-bg-unknown")) + + assert res["success"] is False + assert res["error"] == "worker owner exited" + def test_background_dispatch_reports_running_job_immediately(self): """The dispatch path pre-checks the running set so a mid-run job reports in the tool response, not as a delayed completion event.""" diff --git a/tests/tools/test_delegate_child_transcript_release.py b/tests/tools/test_delegate_child_transcript_release.py new file mode 100644 index 0000000000..de4e274b5a --- /dev/null +++ b/tests/tools/test_delegate_child_transcript_release.py @@ -0,0 +1,145 @@ +"""Finished delegate children must not pin their transcripts in the parent heap. + +Profiled parent (1,320 children over 13h) reached 2.6 GB RSS: every closed +child AIAgent stayed reachable, and each still owned a shallow copy of its +full message list. Two retainers were proven with ``gc.get_referrers``: + +1. ``AIAgent.close()`` cleared ``_session_messages`` but not the + ``_db_flush_scan_prefix`` snapshot (``messages[:]``) or the streamed-text + accumulator, so every message dict stayed alive through the agent. +2. ``bind_subagent_parent`` stored the agent (each child binds ITSELF for + its own turn) strongly in a ContextVar; asyncio Handles/Futures scheduled + during the turn snapshot that Context and live as long as the background + LSP / kernel loops do, so the child object itself was never collected. +""" + +from __future__ import annotations + +import gc +import json +import weakref +from unittest.mock import MagicMock + +from agent.subagent_lifecycle import ( + _ACTIVE_PARENT_AGENT, + bind_subagent_parent, + get_active_subagent_parent, +) +from run_agent import AIAgent + + +def _bare_agent() -> AIAgent: + agent = AIAgent.__new__(AIAgent) + agent._active_children = [] + import threading + + agent._active_children_lock = threading.Lock() + agent._session_db = None + agent.session_id = "child-x" + return agent + + +def test_close_releases_transcript_shadow_copies(): + agent = _bare_agent() + + class Payload(str): # weakref-able stand-in for a message content string + pass + + payload = Payload("x" * 50_000) + big = {"role": "tool", "content": payload} + agent._session_messages = [big] + agent._db_flush_scan_prefix = agent._session_messages[:] + agent._streamed_assistant_text_parts = ["y" * 10_000] + probe = weakref.ref(payload) + + agent.close() + + assert agent._session_messages == [] + assert agent._db_flush_scan_prefix is None + assert agent._streamed_assistant_text_parts == [] + del big, payload + gc.collect() + assert probe() is None, "closed agent still owns its message dicts" + + +def test_bind_subagent_parent_does_not_pin_agent(): + agent = _bare_agent() + probe = weakref.ref(agent) + snapshots = [] + with bind_subagent_parent(agent): + assert get_active_subagent_parent() is agent + import contextvars + + # An asyncio Handle scheduled inside the turn keeps this snapshot. + snapshots.append(contextvars.copy_context()) + assert get_active_subagent_parent() is None + assert snapshots[0][_ACTIVE_PARENT_AGENT] is not agent + del agent + gc.collect() + assert probe() is None, "Context snapshot still pins the agent" + + +def test_bind_subagent_parent_accepts_non_weakrefable_doubles(): + class Slots: + __slots__ = () + + double = Slots() + with bind_subagent_parent(double): + assert get_active_subagent_parent() is double + + +def _fake_child(messages): + child = MagicMock() + child._credential_pool = None + child._delegate_role = "leaf" + child.session_estimated_cost_usd = 0.0123 + child.session_cost_status = "estimated" + child.session_id = "child-sess" + child.run_conversation.return_value = { + "final_response": "the summary", + "completed": True, + "interrupted": False, + "api_calls": 3, + "messages": messages, + } + return child + + +def test_run_single_child_result_json_unchanged_by_transcript_release(): + """Pin: the parent-visible result entry is byte-identical whether or not + the child released its transcript at close() (the entry never carried + ``messages``; only summary/tool_trace/tokens/cost derive from them).""" + from tests.tools.test_delegate import _make_mock_parent + from tools.delegate_tool import _run_single_child + + messages = [ + {"role": "user", "content": "goal"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "c1", + "type": "function", + "function": {"name": "read_file", "arguments": '{"path": "a.py"}'}, + } + ], + }, + {"role": "tool", "tool_call_id": "c1", "content": "z" * 5000}, + {"role": "assistant", "content": "the summary"}, + ] + results = [] + for _ in range(2): + child = _fake_child([dict(m) for m in messages]) + entry = _run_single_child( + task_index=0, goal="goal", child=child, parent_agent=_make_mock_parent() + ) + child.close.assert_called_once() + entry.pop("duration_seconds", None) + results.append(json.dumps(entry, sort_keys=True, default=str)) + assert results[0] == results[1] + parsed = json.loads(results[0]) + assert parsed["summary"] == "the summary" + assert parsed["tool_trace"][0]["tool"] == "read_file" + assert parsed["cost_usd"] == 0.0123 + assert "messages" not in parsed diff --git a/tests/tools/test_file_operations.py b/tests/tools/test_file_operations.py index b26ca1ba62..b1c939dcc9 100644 --- a/tests/tools/test_file_operations.py +++ b/tests/tools/test_file_operations.py @@ -1,12 +1,13 @@ """Tests for tools/file_operations.py — deny list, result dataclasses, helpers.""" import os -import re import pytest import subprocess from pathlib import Path from unittest.mock import MagicMock +from tests.tools.file_ops_fakes import READ_SENTINEL_RE, compound_read_output +from tools.environments.local import _find_bash, _msys_to_windows_path, LocalEnvironment from tools.file_operations import ( _is_write_denied, ReadResult, @@ -154,10 +155,16 @@ class TestSearchResult: assert d["matches"][0]["path"] == "a.py" - def test_truncated_flag(self): + def test_truncated_flag_marks_total_as_lower_bound(self): r = SearchResult(total_count=100, truncated=True) d = r.to_dict() assert d["truncated"] is True + assert d["total_count_is_lower_bound"] is True + + def test_untruncated_total_omits_lower_bound_flag(self): + r = SearchResult(total_count=100) + d = r.to_dict() + assert "total_count_is_lower_bound" not in d class TestSearchResultDensify: @@ -254,16 +261,29 @@ def make_real_subprocess_env(cwd: str, include_stderr: bool = False) -> MagicMoc env.cwd = cwd def execute(command, **kwargs): + stdin_data = kwargs.get("stdin_data") + is_windows = os.name == "nt" + if is_windows: + # Match LocalEnvironment: commands are POSIX scripts executed by + # Git Bash, and stdin bytes must bypass Windows newline rewriting. + command = [_find_bash(), "-c", command] completed = subprocess.run( command, - shell=True, - text=True, + shell=not is_windows, + text=not is_windows, capture_output=True, - input=kwargs.get("stdin_data"), + input=(stdin_data.encode("utf-8", "surrogateescape") + if is_windows and stdin_data is not None else stdin_data), + ) + output = ( + completed.stdout.decode("utf-8", "replace") + if is_windows else completed.stdout ) - output = completed.stdout if include_stderr: - output += completed.stderr + output += ( + completed.stderr.decode("utf-8", "replace") + if is_windows else completed.stderr + ) return { "output": output, "returncode": completed.returncode, @@ -302,19 +322,14 @@ class TestShellFileOpsHelpers: def side_effect(command, **kwargs): commands.append(command) - # The size probe gates `wc -c` behind `[ -f ]` so a FIFO or device - # cannot block the read; it still reports a plain byte count. - if command.startswith("if [ -f ") or command.startswith("wc -c"): - return {"output": "5\n", "returncode": 0} - if command.startswith("head -c") and "| base64" in command: - import base64 as b64 - return {"output": b64.b64encode(b"hello").decode(), "returncode": 0} - if command.startswith("head -c"): - return {"output": "hello", "returncode": 0} - if command.startswith("sed -n"): - return {"output": "hello\n", "returncode": 0} - if command.startswith("wc -l"): - return {"output": "1\n", "returncode": 0} + m = READ_SENTINEL_RE.search(command) + if m: + return { + "output": compound_read_output( + m.group(0), size=5, sample=b"hello", content="hello\n", total_lines=1 + ), + "returncode": 0, + } return {"output": "", "returncode": 0} mock_env.execute.side_effect = side_effect @@ -322,16 +337,22 @@ class TestShellFileOpsHelpers: result = ops.read_file(r"C:\Users\alice\notes.txt") assert result.error is None - assert commands[0] == ( + # One compound probe carries every stage; each embeds the MSYS path. + # The size probe gates `wc -c` behind `[ -f ]` so a FIFO or device + # cannot block the read; it still reports a plain byte count. + assert len(commands) == 1 + probe = commands[0] + assert probe.startswith( "if [ -f '/c/Users/alice/notes.txt' ]; " "then wc -c < '/c/Users/alice/notes.txt' 2>/dev/null; " + ) + assert "head -c 1000 '/c/Users/alice/notes.txt' 2>/dev/null | base64" in probe + assert "sed -n '1,2000p' '/c/Users/alice/notes.txt' 2>/dev/null | cut -b1-8001" in probe + assert "wc -l < '/c/Users/alice/notes.txt'" in probe + assert ( "elif [ -e '/c/Users/alice/notes.txt' ]; " "then echo __hermes_not_regular__; " - "else exit 1; fi" - ) - assert commands[1] == "head -c 1000 '/c/Users/alice/notes.txt' 2>/dev/null | base64" - assert commands[2] == "sed -n '1,2000p' '/c/Users/alice/notes.txt' | cut -b1-8001" - assert commands[3] == "wc -l < '/c/Users/alice/notes.txt'" + ) in probe def test_is_likely_binary_by_extension(self, file_ops): assert file_ops._is_likely_binary("photo.png") is True @@ -354,14 +375,15 @@ class TestShellFileOpsHelpers: ) def side_effect(command, **kwargs): - if command.startswith("if [ -f ") or command.startswith("wc -c"): - return {"output": "12\n", "returncode": 0} - if command.startswith("head -c"): - return {"output": "print('ok')\n", "returncode": 0} - if command.startswith("sed -n"): - return {"output": leaked, "returncode": 0} - if command.startswith("wc -l"): - return {"output": "1\n", "returncode": 0} + m = READ_SENTINEL_RE.search(command) + if m: + return { + "output": compound_read_output( + m.group(0), size=12, sample=b"print('ok')\n", + content=leaked, total_lines=1, + ), + "returncode": 0, + } return {"output": "", "returncode": 0} mock_env.execute.side_effect = side_effect @@ -436,7 +458,7 @@ class TestSearchPathValidation: class TestSearchFilesFallbackHiddenPaths: def _make_env(self): - return make_real_subprocess_env("/") + return LocalEnvironment("/") def test_hidden_root_with_hidden_ancestor_includes_files(self, tmp_path, monkeypatch): """Fallback find should include visible files when path is inside hidden root.""" @@ -772,17 +794,19 @@ class TestByteLayerBinaryDetection: # --- integration: read_file over the mocked terminal ------------------ def _dispatch(self, cjk_bytes): - import base64 as b64 - def side_effect(command, **kwargs): - if command.startswith("if [ -f ") or command.startswith("wc -c"): - return {"output": f"{len(cjk_bytes)}\n", "returncode": 0} - if command.startswith("head -c") and "| base64" in command: - return {"output": b64.b64encode(cjk_bytes[:1000]).decode(), "returncode": 0} - if command.startswith("sed -n"): - return {"output": cjk_bytes.decode("utf-8", errors="replace"), "returncode": 0} - if command.startswith("wc -l"): - return {"output": "1\n", "returncode": 0} + m = READ_SENTINEL_RE.search(command) + if m: + return { + "output": compound_read_output( + m.group(0), + size=len(cjk_bytes), + sample=cjk_bytes[:1000], + content=cjk_bytes.decode("utf-8", errors="replace"), + total_lines=1, + ), + "returncode": 0, + } return {"output": "", "returncode": 0} return side_effect diff --git a/tests/tools/test_file_operations_edge_cases.py b/tests/tools/test_file_operations_edge_cases.py index 0865801911..103c172884 100644 --- a/tests/tools/test_file_operations_edge_cases.py +++ b/tests/tools/test_file_operations_edge_cases.py @@ -8,6 +8,7 @@ Covers: import pytest from unittest.mock import MagicMock, patch +from tests.tools.file_ops_fakes import READ_SENTINEL_RE, compound_read_output from tools.file_operations import ShellFileOperations, _parse_search_context_line @@ -205,14 +206,15 @@ class TestPaginationBounds: def fake_exec(command, *args, **kwargs): commands.append(command) - if command.startswith("if [ -f ") or command.startswith("wc -c"): - return MagicMock(exit_code=0, stdout="12") - if command.startswith("head -c"): - return MagicMock(exit_code=0, stdout="line1\nline2\n") - if command.startswith("sed -n"): - return MagicMock(exit_code=0, stdout="line1\n") - if command.startswith("wc -l"): - return MagicMock(exit_code=0, stdout="2") + m = READ_SENTINEL_RE.search(command) + if m: + return MagicMock( + exit_code=0, + stdout=compound_read_output( + m.group(0), size=12, sample=b"line1\nline2\n", + content="line1\n", total_lines=2, + ), + ) return MagicMock(exit_code=0, stdout="") with patch.object(ops, "_exec", side_effect=fake_exec): @@ -220,8 +222,9 @@ class TestPaginationBounds: assert result.error is None assert "1|line1" in result.content - sed_commands = [cmd for cmd in commands if cmd.startswith("sed -n")] - assert sed_commands == ["sed -n '1,1p' 'notes.txt' | cut -b1-8001"] + # The clamped range rides the single compound probe. + assert len(commands) == 1 + assert "sed -n '1,1p' 'notes.txt' 2>/dev/null | cut -b1-8001" in commands[0] def test_search_clamps_offset_and_limit_before_building_head_pipeline(self): env = MagicMock() @@ -233,7 +236,7 @@ class TestPaginationBounds: commands.append(command) if command.startswith("test -e"): return MagicMock(exit_code=0, stdout="exists") - if command.startswith("rg --files"): + if "--files" in command: return MagicMock(exit_code=0, stdout="a.py\n") return MagicMock(exit_code=0, stdout="") @@ -242,9 +245,9 @@ class TestPaginationBounds: result = ops.search("*.py", target="files", path=".", offset=-4, limit=-2) assert result.files == ["a.py"] - rg_commands = [cmd for cmd in commands if cmd.startswith("rg --files")] + rg_commands = [cmd for cmd in commands if "--files" in cmd] assert rg_commands - assert "| head -n 1" in rg_commands[0] + assert "| head -n 2" in rg_commands[0] # ========================================================================= diff --git a/tests/tools/test_file_ops_single_roundtrip.py b/tests/tools/test_file_ops_single_roundtrip.py new file mode 100644 index 0000000000..b7d3e95835 --- /dev/null +++ b/tests/tools/test_file_ops_single_roundtrip.py @@ -0,0 +1,474 @@ +"""``read_file`` / ``write_file`` cost one shell round-trip, not four. + +Real ``LocalEnvironment`` against ``tmp_path`` (no mocks), with a spy on +``env.execute`` counting round-trips. The cases below are exactly the ones +that used to need their own probe (existence, size, binary sample, page, +line count, trailing newline), so each proves the compound reply carries +that answer. +""" + +import logging +import os +import sys +import threading +from unittest.mock import patch + +import pytest + +from tools.environments.local import LocalEnvironment +from tools.file_operations import ExecuteResult, ShellFileOperations + +pytestmark = pytest.mark.skipif(sys.platform == "win32", reason="POSIX shell probes") + +READ_PROBE_MARK = "__HERMES_RF_" + + +@pytest.fixture(scope="module") +def _local_env(tmp_path_factory): + """One real LocalEnvironment per module; constructing one costs ~0.8 s.""" + return LocalEnvironment(cwd=str(tmp_path_factory.mktemp("file-ops"))) + + +@pytest.fixture +def _ops(_local_env, tmp_path): + """(ops, calls): file ops over the real local shell, every execute recorded.""" + env = _local_env + env.cwd = str(tmp_path) + calls = [] + real_execute = type(env).execute.__get__(env, type(env)) + + def spy(command, *args, **kwargs): + calls.append(command) + return real_execute(command, *args, **kwargs) + + env.execute = spy + try: + yield ShellFileOperations(env, cwd=str(tmp_path)), calls + finally: + env.__dict__.pop("execute", None) + + +@pytest.fixture +def shell(_ops, monkeypatch): + """Pin the shell path even where a native fast path exists.""" + monkeypatch.setenv("HERMES_NATIVE_FILE_READ", "0") + return _ops + + +@pytest.fixture +def native(_ops, monkeypatch): + """Same wiring with the native fast path on.""" + monkeypatch.delenv("HERMES_NATIVE_FILE_READ", raising=False) + return _ops + + +def _write(tmp_path, name, data: bytes): + p = tmp_path / name + p.write_bytes(data) + return str(p) + + +class TestReadFileOneRoundTrip: + def test_text_read_is_one_round_trip(self, shell, tmp_path): + ops, calls = shell + p = _write(tmp_path, "a.txt", b"one\ntwo\nthree\n") + r = ops.read_file(p) + assert len(calls) == 1 and READ_PROBE_MARK in calls[0] + assert r.error is None + # ``_add_line_numbers`` numbers the empty tail after the final + # newline: long-standing behaviour, preserved byte for byte. + assert r.content == "1|one\n2|two\n3|three\n4|" + assert (r.total_lines, r.file_size, r.truncated) == (3, 14, False) + + def test_no_trailing_newline_needs_no_extra_probe(self, shell, tmp_path): + ops, calls = shell + p = _write(tmp_path, "b.txt", b"a\nb") + r = ops.read_file(p) + assert len(calls) == 1 + # ``cut`` newline-terminates the last line; the artifact is stripped + # from the same reply that used to need a fifth ``tail -c 1`` call. + assert r.content == "1|a\n2|b" + assert r.total_lines == 1 # wc -l semantics, unchanged + + def test_pagination_window_and_hint(self, shell, tmp_path): + ops, calls = shell + p = _write(tmp_path, "c.txt", b"".join(b"l%d\n" % i for i in range(1, 11))) + r = ops.read_file(p, offset=3, limit=2) + assert len(calls) == 1 + assert r.content == "3|l3\n4|l4\n5|" + assert r.truncated is True and r.total_lines == 10 + assert "offset=5" in r.hint + + def test_offset_past_eof_note(self, shell, tmp_path): + ops, calls = shell + p = _write(tmp_path, "c.txt", b"".join(b"l%d\n" % i for i in range(1, 6))) + r = ops.read_file(p, offset=50) + assert len(calls) == 1 + assert r.content == "" and r.error is None + assert "beyond the end" in r.hint and "5" in r.hint + + def test_empty_file(self, shell, tmp_path): + ops, calls = shell + r = ops.read_file(_write(tmp_path, "e.txt", b"")) + assert len(calls) == 1 + assert r.error is None and r.content == "" and r.total_lines == 0 + assert "empty" in r.hint + + def test_bom_stripped_on_first_page(self, shell, tmp_path): + ops, calls = shell + r = ops.read_file(_write(tmp_path, "f.txt", "hello\n".encode("utf-8"))) + assert len(calls) == 1 + assert r.content == "1|hello\n2|" + + def test_crlf_bytes_survive(self, shell, tmp_path): + ops, calls = shell + r = ops.read_file(_write(tmp_path, "g.txt", b"x\r\ny\r\n")) + assert r.content == "1|x\r\n2|y\r\n3|" + + def test_long_line_clamped_and_marked(self, shell, tmp_path): + ops, calls = shell + r = ops.read_file(_write(tmp_path, "L.txt", b"a" * 9000 + b"\nshort\n")) + assert len(calls) == 1 + first, second, tail = r.content.split("\n") + assert first.endswith("... [truncated]") and len(first) < 9000 + assert second == "2|short" and tail == "3|" + + def test_relative_path_resolves_against_env_cwd(self, shell, tmp_path): + ops, calls = shell + _write(tmp_path, "rel.txt", b"here\n") + r = ops.read_file("rel.txt") + assert r.error is None and r.content == "1|here\n2|" + + def test_sentinel_lookalike_in_content_reads_intact(self, shell, tmp_path): + ops, calls = shell + lookalike = "__HERMES_RF_" + "ab" * 16 + "__" + p = _write(tmp_path, "s.txt", f"x\n{lookalike}\ny\n".encode("utf-8")) + r = ops.read_file(p) + assert r.error is None and r.total_lines == 3 + assert r.content == f"1|x\n2|{lookalike}\n3|y\n4|" + + +class TestReadFileNonTextPaths: + def test_missing_file_probes_once_then_suggests(self, shell, tmp_path): + ops, calls = shell + _write(tmp_path, "notes.txt", b"x\n") + r = ops.read_file(str(tmp_path / "note.txt")) + assert READ_PROBE_MARK in calls[0] + assert r.error and "File not found" in r.error + assert any(s.endswith("notes.txt") for s in r.similar_files) + + def test_unicode_variant_retry_still_works(self, shell, tmp_path): + ops, calls = shell + # A curly apostrophe vs the ASCII one: visually identical in a + # terminal, and — unlike NFC/NFD — never aliased by the filesystem + # (APFS resolves NFD lookups to NFC files directly, which would skip + # the retry path this test exists to exercise). + on_disk = "it\u2019s.txt" + typed = "it's.txt" + assert on_disk != typed + _write(tmp_path, on_disk, b"accent\n") + r = ops.read_file(str(tmp_path / typed)) + assert r.error is None and r.content == "1|accent\n2|" + assert r.hint is not None and "unicode-equivalent" in r.hint + + def test_directory_is_not_regular(self, shell, tmp_path): + ops, calls = shell + r = ops.read_file(str(tmp_path)) + assert len(calls) == 1 + assert r.error and "not a regular file" in r.error + + def test_binary_sample_detected_in_same_reply(self, shell, tmp_path): + ops, calls = shell + p = _write(tmp_path, "blob", b"\x00\x01\x02" + b"\x00" * 50) + r = ops.read_file(p) + assert READ_PROBE_MARK in calls[0] + assert r.is_binary is True and r.error + # Only the UTF-16 rescue may add round-trips, never a second sample. + assert not any("head -c 1000" in c for c in calls[1:]) + + def test_image_extension_stops_at_size_probe(self, shell, tmp_path): + ops, calls = shell + r = ops.read_file(_write(tmp_path, "p.png", b"\x89PNG\r\n")) + assert len(calls) == 1 and READ_PROBE_MARK not in calls[0] + assert r.is_image is True and r.file_size == 6 + + @pytest.mark.linux_only + def test_fifo_returns_not_regular_without_blocking(self, shell, tmp_path): + if not hasattr(os, "mkfifo"): + pytest.skip("no mkfifo") + ops, calls = shell + fifo = tmp_path / "pipe" + os.mkfifo(fifo) + box = {} + + def run(): + box["r"] = ops.read_file(str(fifo)) + + t = threading.Thread(target=run, daemon=True) + t.start() + t.join(20) + assert not t.is_alive(), "read_file blocked on a writer-less FIFO" + assert "not a regular file" in box["r"].error + assert len(calls) == 1 + + +class TestWriteFileRoundTrips: + """write_file: one probe, one atomic write, one hash check (three calls).""" + + @staticmethod + def _execs(calls): + return [c for c in calls] + + def test_new_text_file_is_three_round_trips(self, shell, tmp_path): + ops, calls = shell + p = str(tmp_path / "new.txt") + r = ops.write_file(p, "line one\nline two\n") + assert r.error is None and r.verified is True + assert len(calls) == 3 + assert "__HERMES_WF_" in calls[0] # probe + assert "mv -f" in calls[1] # atomic write + assert calls[2].startswith("sha256sum ") # verify + assert (tmp_path / "new.txt").read_bytes() == b"line one\nline two\n" + + def test_crlf_file_keeps_crlf_from_the_probe(self, shell, tmp_path): + ops, calls = shell + p = tmp_path / "crlf.txt" + p.write_bytes(b"a\r\nb\r\n") + r = ops.write_file(str(p), "x\ny\n") + assert r.error is None and len(calls) == 3 + assert p.read_bytes() == b"x\r\ny\r\n" + + def test_bom_is_read_from_disk_and_preserved(self, shell, tmp_path): + ops, calls = shell + p = tmp_path / "bom.txt" + p.write_bytes("old\n".encode("utf-8")) + r = ops.write_file(str(p), "new\n") + assert r.error is None and len(calls) == 3 + assert p.read_bytes() == "new\n".encode("utf-8") + + def test_pre_content_read_rides_the_same_probe(self, shell, tmp_path): + """A lintable extension wants the old text (lint delta); it comes + back in the probe reply instead of a separate ``cat``.""" + ops, calls = shell + p = tmp_path / "code.py" + p.write_bytes(b"x = 1\r\ny = 2\r\n") + r = ops.write_file(str(p), "x = 1\ny = 3\n") + assert r.error is None + probes = [c for c in calls if "__HERMES_WF_" in c] + assert len(probes) == 1 and "cat " in probes[0] + assert not any(c.startswith("cat ") for c in calls) + assert p.read_bytes() == b"x = 1\r\ny = 3\r\n" + + def test_missing_file_probe_does_not_block_the_write(self, shell, tmp_path): + ops, calls = shell + p = tmp_path / "deep" / "er" / "new.md" + r = ops.write_file(str(p), "hi\n") + assert r.error is None and r.dirs_created is True + assert len(calls) == 3 + assert p.read_bytes() == b"hi\n" + + def test_unparseable_probe_reply_falls_back_to_separate_probes(self, shell, tmp_path): + ops, calls = shell + p = tmp_path / "crlf.txt" + p.write_bytes(b"a\r\nb\r\n") + real_exec = ops._exec + + def garbled(command, *args, **kwargs): + if "__HERMES_WF_" in command: + return ExecuteResult(stdout="[Command timed out after 1s]\n", exit_code=124) + return real_exec(command, *args, **kwargs) + + with patch.object(ops, "_exec", side_effect=garbled): + r = ops.write_file(str(p), "x\ny\n") + assert r.error is None + assert p.read_bytes() == b"x\r\ny\r\n" + + +class TestNativeRead: + def test_native_read_makes_no_shell_call(self, native, tmp_path): + ops, calls = native + r = ops.read_file(_write(tmp_path, "a.txt", b"one\ntwo\n")) + assert calls == [] + assert r.error is None and r.content == "1|one\n2|two\n3|" + assert (r.total_lines, r.file_size) == (2, 8) + + def test_kill_switch_routes_to_the_shell(self, native, tmp_path, monkeypatch): + ops, calls = native + p = _write(tmp_path, "a.txt", b"one\n") + monkeypatch.setenv("HERMES_NATIVE_FILE_READ", "0") + ops.read_file(p) + assert len(calls) == 1 and READ_PROBE_MARK in calls[0] + + def test_non_local_environment_keeps_the_shell_path(self): + from unittest.mock import MagicMock + + env = MagicMock() + env.cwd = "/tmp" + assert ShellFileOperations(env)._native_read_enabled() is False + + def test_tilde_still_expands_through_the_shell(self, native): + ops, calls = native + r = ops.read_file("~/hermes-no-such-file-7f3a.txt") + assert calls[0] == "echo $HOME" + assert r.error and "File not found" in r.error + + def test_injection_lookalike_path_is_never_expanded(self, native, tmp_path): + ops, calls = native + marker = tmp_path / "pwned" + r = ops.read_file(f"~; echo PWNED > {marker}") + assert r.error and not marker.exists() + # The text reaches the shell only single-quoted, inside the missing- + # file recovery's directory listing; the tilde probe is a fixed + # ``echo $HOME`` that never embeds the path. Nothing else runs. + for c in calls: + assert c == "echo $HOME" or c.startswith("ls -1 '~; echo PWNED"), c + + @pytest.mark.linux_only + def test_fifo_refused_without_a_shell_and_without_blocking(self, native, tmp_path): + if not hasattr(os, "mkfifo"): + pytest.skip("no mkfifo") + ops, calls = native + fifo = tmp_path / "pipe" + os.mkfifo(fifo) + box = {} + + def run(): + box["r"] = ops.read_file(str(fifo)) + + t = threading.Thread(target=run, daemon=True) + t.start() + t.join(20) + assert not t.is_alive(), "native read_file blocked on a writer-less FIFO" + assert "not a regular file" in box["r"].error + assert calls == [] + + +# The native reader scans 1 MiB chunks and clamps each page line to +# ``4 * get_max_line_length() + 1`` bytes (8001 by default), exactly as +# ``sed | cut -b1-N`` does. These shapes put a newline, a line, the clamp +# point, a CRLF pair and EOF precisely on those chunk boundaries. +_CHUNK = 1 << 20 +_CLAMP = 8001 + +PARITY_CASES = [ + ("plain", b"one\ntwo\nthree\n", {}), + ("no_trailing_newline", b"a\nb", {}), + ("blank_tail", b"a\n\n", {}), + ("crlf", b"x\r\ny\r\n", {}), + ("lone_cr", b"a\rb\n", {}), + ("bom", "hello\n".encode("utf-8"), {}), + ("empty", b"", {}), + ("single_no_newline", b"solo", {}), + ("only_newline", b"\n", {}), + ("unicode", "héllo wörld\n汉字\n".encode("utf-8"), {}), + ("long_line", b"a" * 9000 + b"\nshort\n", {}), + ("multibyte_long_line", ("汉" * 4000 + "\nx\n").encode("utf-8"), {}), + ("multi_chunk_line", b"b" * 3_000_000 + b"\nz\n", {}), + ("newline_last_byte_of_chunk", b"a" * (_CHUNK - 1) + b"\nsecond\n", {}), + ("newline_first_byte_of_next_chunk", b"a" * _CHUNK + b"\nsecond\n", {}), + ("line_spans_three_boundaries", b"p\n" + b"b" * (3 * _CHUNK + 5) + b"\nz\n", {}), + ( + "clamp_fills_on_boundary", + b"x" * (_CHUNK - _CLAMP - 1) + b"\n" + b"c" * 20000 + b"\ntail\n", + {"offset": 2, "limit": 3}, + ), + ( + "clamp_fills_after_boundary", + b"x" * (_CHUNK - _CLAMP) + b"\n" + b"c" * 20000 + b"\ntail\n", + {"offset": 2, "limit": 3}, + ), + ("eof_midline_on_boundary", b"l1\n" + b"d" * (2 * _CHUNK - 3), {}), + ("blank_lines_on_boundary", b"e" * (_CHUNK - 2) + b"\n\n\n" + b"f\n", {}), + ("crlf_split_on_boundary", b"r" * (_CHUNK - 1) + b"\r\n" + b"s\r\n", {}), + ( + "page_starts_in_second_chunk", + b"q\n" * (_CHUNK // 2 + 3) + b"target1\ntarget2\n", + {"offset": _CHUNK // 2 + 3, "limit": 4}, + ), + ("window", b"".join(b"l%d\n" % i for i in range(1, 11)), {"offset": 3, "limit": 2}), + ("window_reaches_eof", b"".join(b"l%d\n" % i for i in range(1, 11)), {"offset": 9, "limit": 5}), + ("past_eof", b"".join(b"l%d\n" % i for i in range(1, 6)), {"offset": 50}), + ("nul_binary", b"\x00\x01\x02" * 20, {}), + ("latin1_tail", b"caf\xe9\n", {}), + ("sentinel_lookalike", b"x\n__HERMES_RF_" + b"ab" * 16 + b"__\ny\n", {}), +] + + +class TestNativeReadParity: + """The native path must be indistinguishable from the shell path.""" + + @pytest.mark.parametrize("name,data,kwargs", PARITY_CASES, ids=[c[0] for c in PARITY_CASES]) + def test_shell_and_native_agree(self, native, tmp_path, monkeypatch, name, data, kwargs): + ops, calls = native + p = _write(tmp_path, f"{name}.txt", data) + monkeypatch.setenv("HERMES_NATIVE_FILE_READ", "0") + via_shell = ops.read_file(p, **kwargs).to_dict() + assert calls and READ_PROBE_MARK in calls[0] + calls.clear() + monkeypatch.delenv("HERMES_NATIVE_FILE_READ") + via_native = ops.read_file(p, **kwargs).to_dict() + assert via_native == via_shell + # The native path touches the shell only for the UTF-16 rescue of binaries. + assert not any(READ_PROBE_MARK in c for c in calls) + + def test_special_paths_agree(self, native, tmp_path, monkeypatch): + ops, calls = native + _write(tmp_path, "real.txt", b"target\n") + os.symlink(tmp_path / "real.txt", tmp_path / "link.txt") + os.symlink(tmp_path / "gone", tmp_path / "dangling.txt") + (tmp_path / "sub").mkdir() + _write(tmp_path, "pic.png", b"\x89PNG\r\n") + _write(tmp_path, "notes.txt", b"n\n") + for p in ( + str(tmp_path / "link.txt"), + str(tmp_path / "dangling.txt"), + str(tmp_path / "sub"), + str(tmp_path / "pic.png"), + str(tmp_path / "note.txt"), # missing → similar-file suggestions + "real.txt", # relative to env.cwd + ): + monkeypatch.setenv("HERMES_NATIVE_FILE_READ", "0") + via_shell = ops.read_file(p).to_dict() + monkeypatch.delenv("HERMES_NATIVE_FILE_READ") + calls.clear() + via_native = ops.read_file(p).to_dict() + assert via_native == via_shell, p + assert not any(READ_PROBE_MARK in c for c in calls), p + + +class TestCompoundFallback: + def test_unparseable_reply_falls_back_to_sequential_probes(self, shell, tmp_path): + ops, calls = shell + p = _write(tmp_path, "a.txt", b"one\ntwo\n") + real_exec = ops._exec + + def garbled(command, *args, **kwargs): + if READ_PROBE_MARK in command: + return ExecuteResult(stdout="[Command timed out after 1s]\n", exit_code=124) + return real_exec(command, *args, **kwargs) + + with patch.object(ops, "_exec", side_effect=garbled): + r = ops.read_file(p) + assert r.error is None and r.content == "1|one\n2|two\n3|" + assert r.total_lines == 2 + + def test_fallback_is_logged_at_debug(self, shell, tmp_path, caplog): + """A backend that keeps falling back shows up in debug logs.""" + ops, calls = shell + p = _write(tmp_path, "a.txt", b"one\n") + real_exec = ops._exec + + def garbled(command, *args, **kwargs): + if READ_PROBE_MARK in command: + return ExecuteResult(stdout="garbage\n", exit_code=0) + return real_exec(command, *args, **kwargs) + + with caplog.at_level(logging.DEBUG, logger="tools.file_operations"), \ + patch.object(ops, "_exec", side_effect=garbled): + r = ops.read_file(p) + assert r.error is None and r.content == "1|one\n2|" + assert any( + "falling back to sequential probes" in rec.getMessage() + and str(p) in rec.getMessage() + for rec in caplog.records + ) diff --git a/tests/tools/test_file_read_guards.py b/tests/tools/test_file_read_guards.py index 1ee28f7789..52321a637f 100644 --- a/tests/tools/test_file_read_guards.py +++ b/tests/tools/test_file_read_guards.py @@ -391,7 +391,7 @@ class TestFileDedup(unittest.TestCase): _read_tracker.clear() self._tmpdir = _make_safe_tempdir("hermes-dedup-") self._tmpfile = os.path.join(self._tmpdir, "dedup_test.txt") - with open(self._tmpfile, "w") as f: + with open(self._tmpfile, "w", encoding="utf-8") as f: f.write("line one\nline two\n") def tearDown(self): @@ -463,7 +463,7 @@ class TestDedupStubLoopGuard(unittest.TestCase): _read_tracker.clear() self._tmpdir = tempfile.mkdtemp() self._tmpfile = os.path.join(self._tmpdir, "loop_test.txt") - with open(self._tmpfile, "w") as f: + with open(self._tmpfile, "w", encoding="utf-8") as f: f.write("line one\nline two\n") def tearDown(self): @@ -530,7 +530,7 @@ class TestDedupStubLoopGuard(unittest.TestCase): # File changes — mtime updates time.sleep(0.05) - with open(self._tmpfile, "w") as f: + with open(self._tmpfile, "w", encoding="utf-8") as f: f.write("brand new content\n") r4 = json.loads(read_file_tool(self._tmpfile, task_id="loop")) @@ -591,10 +591,16 @@ class TestDedupStubLoopGuard(unittest.TestCase): reset_file_dedup("loop") - # Fresh session — real read, no stub, no block + # Post-compression: block counters cleared and exact content is served + # once because the earlier payload may no longer be in context. r4 = json.loads(read_file_tool(self._tmpfile, task_id="loop")) self.assertNotIn("error", r4) self.assertNotIn("dedup", r4) + self.assertIn("content", r4) + + # The next unchanged read in this generation is lightweight again. + r5 = json.loads(read_file_tool(self._tmpfile, task_id="loop")) + self.assertTrue(r5.get("dedup")) # --------------------------------------------------------------------------- @@ -602,14 +608,13 @@ class TestDedupStubLoopGuard(unittest.TestCase): # --------------------------------------------------------------------------- class TestDedupResetOnCompression(unittest.TestCase): - """reset_file_dedup should clear the dedup cache so post-compression - reads return full content.""" + """Compaction starts a new full-content recovery generation.""" def setUp(self): _read_tracker.clear() self._tmpdir = tempfile.mkdtemp() self._tmpfile = os.path.join(self._tmpdir, "compress_test.txt") - with open(self._tmpfile, "w") as f: + with open(self._tmpfile, "w", encoding="utf-8") as f: f.write("original content\n") def tearDown(self): @@ -621,10 +626,10 @@ class TestDedupResetOnCompression(unittest.TestCase): pass @patch("tools.file_tools._get_file_ops") - def test_reset_clears_dedup(self, mock_ops): - """After reset_file_dedup, the same read returns full content.""" + def test_first_post_compaction_read_recovers_exact_content(self, mock_ops): + """First post-compaction read is full; later reads deduplicate.""" mock_ops.return_value = _make_fake_ops( - content="original content\n", file_size=18, + content="SECRET_EXACT_LINE=42\n", file_size=21, ) # First read — populates dedup cache read_file_tool(self._tmpfile, task_id="comp") @@ -636,10 +641,15 @@ class TestDedupResetOnCompression(unittest.TestCase): # Simulate compression reset_file_dedup("comp") - # Read again — should get full content + # Exact prior bytes may have been omitted from the summary, so the + # first read in the new generation must restore them. r_post = json.loads(read_file_tool(self._tmpfile, task_id="comp")) - self.assertNotEqual(r_post.get("dedup"), True, - "Post-compression read should return full content") + self.assertNotIn("dedup", r_post) + self.assertIn("SECRET_EXACT_LINE=42", r_post.get("content", "")) + + # The persisted mtime map still saves tokens after that recovery read. + r_again = json.loads(read_file_tool(self._tmpfile, task_id="comp")) + self.assertTrue(r_again.get("dedup")) @patch("tools.file_tools._get_file_ops") @@ -655,13 +665,12 @@ class TestDedupResetOnCompression(unittest.TestCase): reset_file_dedup("loop") - # 3rd read — counter should still be at 2 from before reset - # (dedup was hit for read 2, but consecutive counter was 1 for that) - # After reset, this read goes through full path, incrementing to 2 + # First read in the new generation returns full content, not a stale + # block or a stub that points to compacted-away bytes. r3 = json.loads(read_file_tool(self._tmpfile, task_id="loop")) - # Should NOT be blocked or warned — counter restarted since dedup - # intercepted reads before they reached the counter self.assertNotIn("error", r3) + self.assertNotIn("dedup", r3) + self.assertIn("content", r3) # --------------------------------------------------------------------------- @@ -760,7 +769,7 @@ class TestWriteInvalidatesDedup(unittest.TestCase): _read_tracker.clear() self._tmpdir = _make_safe_tempdir("hermes-write-dedup-") self._tmpfile = os.path.join(self._tmpdir, "write_dedup.txt") - with open(self._tmpfile, "w") as f: + with open(self._tmpfile, "w", encoding="utf-8") as f: f.write("original content\n") def tearDown(self): diff --git a/tests/tools/test_file_state_registry.py b/tests/tools/test_file_state_registry.py index 30ef964178..adc11dd67f 100644 --- a/tests/tools/test_file_state_registry.py +++ b/tests/tools/test_file_state_registry.py @@ -25,6 +25,7 @@ import unittest from tools import file_state from tools.file_tools import ( + clear_file_ops_cache, read_file_tool, write_file_tool, patch_tool, @@ -123,6 +124,54 @@ class FileStateRegistryUnitTests(unittest.TestCase): ta.join(timeout=3.0) tb.join(timeout=3.0) + def test_lock_path_state_is_released_after_last_waiter(self): + p = self._mk() + first_entered = threading.Event() + release_first = threading.Event() + second_entered = threading.Event() + + def first() -> None: + with file_state.lock_path(p): + first_entered.set() + release_first.wait(timeout=2.0) + + def second() -> None: + first_entered.wait(timeout=2.0) + with file_state.lock_path(p): + second_entered.set() + + ta = threading.Thread(target=first) + tb = threading.Thread(target=second) + ta.start() + tb.start() + self.assertTrue(first_entered.wait(timeout=2.0)) + time.sleep(0.02) + self.assertFalse(second_entered.is_set()) + release_first.set() + ta.join(timeout=3.0) + tb.join(timeout=3.0) + + registry = file_state.get_registry() + self.assertTrue(second_entered.is_set()) + self.assertNotIn(p, registry._path_locks) + self.assertNotIn(p, registry._path_lock_users) + + def test_clear_file_ops_cache_releases_task_state(self): + p = self._mk() + task_id = "finished-task" + file_state.record_read(task_id, p) + + from tools import file_tools + + file_tools._read_tracker[task_id] = {"dedup": {}} + file_tools._patch_failure_tracker[task_id] = {p: 2} + + clear_file_ops_cache(task_id) + + self.assertEqual(file_state.known_reads(task_id), []) + self.assertNotIn(task_id, file_tools._read_tracker) + self.assertNotIn(task_id, file_tools._patch_failure_tracker) + def test_kill_switch_env_var(self): p = self._mk() diff --git a/tests/tools/test_interrupt.py b/tests/tools/test_interrupt.py index 67c5fbf662..899a7f570d 100644 --- a/tests/tools/test_interrupt.py +++ b/tests/tools/test_interrupt.py @@ -62,6 +62,105 @@ class TestInterruptModule: assert other_tid in _interrupted_threads # other thread untouched _interrupted_threads.discard(other_tid) + def test_run_if_not_interrupted_skips_callback_when_already_interrupted(self): + from tools.interrupt import run_if_not_interrupted, set_interrupt + + callbacks = [] + set_interrupt(True) + try: + assert run_if_not_interrupted(lambda: callbacks.append("claimed")) is False + finally: + set_interrupt(False) + + assert callbacks == [] + + @pytest.mark.parametrize("callback_should_fail", [False, True]) + def test_run_if_not_interrupted_orders_callback_before_concurrent_interrupt( + self, callback_should_fail, monkeypatch + ): + import tools.interrupt as interrupt + + class CallbackFailure(Exception): + pass + + original_lock = interrupt._lock + attempting_interrupt_lock = threading.Event() + interrupt_published = threading.Event() + publisher_lock_contention = [] + callback_observations = [] + setters = [] + setter_tids = [] + + class ObservedLock: + def __enter__(self): + if ( + threading.current_thread() in setters + and not attempting_interrupt_lock.is_set() + ): + acquired = original_lock.acquire(blocking=False) + publisher_lock_contention.append(not acquired) + attempting_interrupt_lock.set() + if acquired: + return self + original_lock.acquire() + return self + + def __exit__(self, exc_type, exc_value, traceback): + original_lock.release() + + interrupt.set_interrupt(False) + with original_lock: + baseline = ( + set(interrupt._interrupted_threads), + dict(interrupt._interrupt_reasons), + ) + monkeypatch.setattr(interrupt, "_lock", ObservedLock()) + + def publish_interrupt(): + setter_tids.append(threading.get_ident()) + try: + interrupt.set_interrupt(True) + interrupt_published.set() + finally: + interrupt.set_interrupt(False) + + def callback(): + setter = threading.Thread(target=publish_interrupt) + setters.append(setter) + setter.start() + assert attempting_interrupt_lock.wait(5) + assert publisher_lock_contention == [True] + callback_observations.append(interrupt_published.is_set()) + if callback_should_fail: + raise CallbackFailure + + try: + if callback_should_fail: + with pytest.raises(CallbackFailure): + interrupt.run_if_not_interrupted(callback) + else: + assert interrupt.run_if_not_interrupted(callback) is True + assert interrupt_published.wait(5) + finally: + for setter in setters: + if setter.ident is not None: + setter.join(timeout=5) + interrupt.set_interrupt(False) + + assert setters + assert all(not setter.is_alive() for setter in setters) + assert setter_tids + assert callback_observations == [False] + assert interrupt_published.is_set() + with original_lock: + final_state = ( + set(interrupt._interrupted_threads), + dict(interrupt._interrupt_reasons), + ) + assert final_state == baseline + assert all(setter_tid not in final_state[0] for setter_tid in setter_tids) + assert all(setter_tid not in final_state[1] for setter_tid in setter_tids) + # --------------------------------------------------------------------------- # Unit tests: pre-tool interrupt check diff --git a/tests/tools/test_macos_protected_search.py b/tests/tools/test_macos_protected_search.py index 501280a85f..20609b6c3b 100644 --- a/tests/tools/test_macos_protected_search.py +++ b/tests/tools/test_macos_protected_search.py @@ -1,5 +1,6 @@ """macOS TCC-safe behavior for broad file searches.""" +import re from pathlib import Path import tools.file_operations as file_operations @@ -32,6 +33,17 @@ PROTECTED_NAMES = { } +def _rg_files_commands(commands): + return [command for command in commands if "--files" in command] + + +def _find_commands(commands): + return [ + command for command in commands + if command.startswith("find ") or "; find " in command + ] + + def test_broad_home_search_excludes_macos_protected_folders(tmp_path): home = tmp_path / "Users" / "alice" @@ -72,7 +84,7 @@ def test_broad_file_search_passes_protected_globs_to_ripgrep(tmp_path, monkeypat result = ops.search("*.txt", path=str(home), target="files") - rg_command = next(command for command in env.commands if command.startswith("rg --files")) + rg_command = _rg_files_commands(env.commands)[0] for dirname in PROTECTED_NAMES: assert f"!{dirname}/**" in rg_command assert result.warning is not None @@ -94,7 +106,7 @@ def test_broad_content_search_passes_protected_globs_to_ripgrep(tmp_path, monkey assert f"!{dirname}/**" in rg_command -def test_legacy_ripgrep_file_fallback_keeps_protected_globs(tmp_path, monkeypatch): +def test_empty_ripgrep_file_search_is_one_scan_with_protected_globs(tmp_path, monkeypatch): home = tmp_path / "Users" / "alice" home.mkdir(parents=True) env = RecordingEnvironment(home) @@ -104,8 +116,8 @@ def test_legacy_ripgrep_file_fallback_keeps_protected_globs(tmp_path, monkeypatc ops.search("*.txt", path=str(home), target="files") - rg_commands = [command for command in env.commands if command.startswith("rg --files")] - assert len(rg_commands) == 2 + rg_commands = _rg_files_commands(env.commands) + assert len(rg_commands) == 1 for command in rg_commands: assert "!Downloads/**" in command @@ -128,7 +140,7 @@ def test_grep_fallback_prunes_by_path_not_basename(tmp_path, monkeypatch): for dirname in PROTECTED_NAMES: # Path-scoped pruning: full protected path present, no basename-wide # --exclude-dir for protected names. - assert str(home / dirname) in pruned_command + assert ops._escape_shell_arg(str(home / dirname)) in pruned_command assert f"--exclude-dir={dirname}" not in pruned_command assert f"--exclude-dir='{dirname}'" not in pruned_command @@ -169,7 +181,7 @@ def test_remote_backend_never_prunes(tmp_path, monkeypatch): result = ops.search("*.txt", path=str(home), target="files") - rg_command = next(command for command in env.commands if command.startswith("rg --files")) + rg_command = _rg_files_commands(env.commands)[0] assert "!Downloads/**" not in rg_command assert result.warning is None @@ -185,13 +197,173 @@ def test_find_fallback_prunes_protected_directories(tmp_path, monkeypatch): ops.search("*.txt", path=str(home), target="files") - find_commands = [command for command in env.commands if command.startswith("find ")] + find_commands = _find_commands(env.commands) assert find_commands for command in find_commands: - assert str(home / "Downloads") in command + assert ops._escape_shell_arg(str(home / "Downloads")) in command assert "-prune" in command +def _multi_root_protected_search(tmp_path, monkeypatch, engine): + home = tmp_path / "Users" / "alice" + downloads = home / "Downloads" + downloads.mkdir(parents=True) + env = RecordingEnvironment(home) + ops = ShellFileOperations(env) + monkeypatch.setattr(file_operations, "_HOME", str(home)) + monkeypatch.setattr(file_operations.sys, "platform", "darwin") + monkeypatch.setattr(ops, "_has_command", lambda command: command == engine) + path_checks = 0 + + def execute(command, cwd=None, **kwargs): + nonlocal path_checks + env.commands.append(command) + if command.startswith("test -e"): + path_checks += 1 + output = "not_found\n" if path_checks == 1 else "exists\n" + return {"output": output, "returncode": 0} + if "--files" in command or command.startswith("set -o pipefail; find "): + return {"output": "", "returncode": 0} + return {"output": "yes\n", "returncode": 0} + + env.execute = execute + result = ops.search("*.txt", path=f"{home} {downloads}", target="files") + return ops, env, result, downloads + + +def test_rg_multi_root_keeps_explicit_protected_root_and_reports_actual_skips( + tmp_path, monkeypatch +): + ops, env, result, downloads = _multi_root_protected_search( + tmp_path, monkeypatch, "rg" + ) + + command = _rg_files_commands(env.commands)[0] + absolute_operand = downloads.as_posix() in command + anchored_operand = ( + f"cd {ops._escape_shell_arg(downloads.parent.as_posix())} &&" in command + and " -- '.' 'Downloads' 2>/dev/null" in command + ) + assert absolute_operand or anchored_operand + assert "!Downloads/**" not in command + assert "path contained 2 entries" in (result.warning or "") + assert "macOS protected folders" in (result.warning or "") + protected_warning = result.warning.split("macOS protected folders", 1)[1] + assert "Desktop" in protected_warning + assert "Downloads" not in protected_warning + + +def test_find_multi_root_keeps_explicit_protected_root_and_reports_actual_skips( + tmp_path, monkeypatch +): + ops, env, result, downloads = _multi_root_protected_search( + tmp_path, monkeypatch, "find" + ) + + command = _find_commands(env.commands)[0] + assert "Downloads" in command + assert re.search(r"(?/dev/null" in command + assert result.error is None + + def test_real_ripgrep_does_not_descend_into_protected_folder(tmp_path, monkeypatch): home = tmp_path / "Users" / "alice" safe = home / "safe" diff --git a/tests/tools/test_mcp_death_supervisor.py b/tests/tools/test_mcp_death_supervisor.py new file mode 100644 index 0000000000..f89ff64b46 --- /dev/null +++ b/tests/tools/test_mcp_death_supervisor.py @@ -0,0 +1,820 @@ +"""Contract tests for the shared parent-death supervisor for stdio MCP servers. + +The end-to-end tests here spawn real processes and really SIGKILL a real parent, +because the whole point of this module is behaviour that only exists when a +process dies without running any Python cleanup. A mocked parent death proves +nothing about the guarantee. +""" + +import asyncio +import contextlib +import io +import os +import signal +import subprocess +import sys +import time +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from tools import mcp_death_supervisor, mcp_tool + +pytestmark = pytest.mark.skipif( + os.name != "posix", reason="the supervisor is POSIX-only (process groups)" +) + +SUPERVISOR = os.path.join(os.path.dirname(mcp_tool.__file__), "mcp_death_supervisor.py") + +# Long enough that nothing here can pass because the victim exited on its own. +_VICTIM = [sys.executable, "-c", "import time; time.sleep(300)"] + + +def _alive(pid: int) -> bool: + try: + os.kill(pid, 0) + except (ProcessLookupError, OSError): + return False + return True + + +def _wait_gone(pid: int, timeout: float = 15.0) -> bool: + """Wait for a process this test does NOT own to disappear.""" + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if not _alive(pid): + return True + time.sleep(0.05) + return False + + +def _wait_exited(proc: subprocess.Popen, timeout: float = 15.0) -> bool: + """Wait for a direct child of this test to exit. + + ``os.kill(pid, 0)`` cannot be used for our own children: a killed child + stays a zombie until someone reaps it, and signalling a zombie succeeds. + """ + try: + proc.wait(timeout=timeout) + except subprocess.TimeoutExpired: + return False + return True + + +def _kill(pid: int) -> None: + try: + os.kill(pid, signal.SIGKILL) + except (ProcessLookupError, OSError): + pass + + +# --------------------------------------------------------------------------- +# Target safety: this process signals whole process GROUPS, so a bad target is +# unusually expensive. killpg(0, ...) would signal the supervisor's own group. +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("pgid", [0, 1, -1, -5]) +def test_refuses_process_groups_that_are_never_a_valid_target(pgid): + assert mcp_death_supervisor._is_safe_target( + pgid, own_pgid=4242, parent_pgid=777 + ) is False + + +def test_refuses_its_own_group_and_the_parents_group(): + assert mcp_death_supervisor._is_safe_target( + 4242, own_pgid=4242, parent_pgid=777 + ) is False + assert mcp_death_supervisor._is_safe_target( + 777, own_pgid=4242, parent_pgid=777 + ) is False + + +def test_accepts_an_unrelated_group(): + assert mcp_death_supervisor._is_safe_target( + 999, own_pgid=4242, parent_pgid=777 + ) is True + + +# --------------------------------------------------------------------------- +# Control protocol +# --------------------------------------------------------------------------- + + +def test_registrations_survive_to_eof_and_unregistrations_are_dropped(): + stream = io.StringIO("register 111\nregister 222\nunregister 111\n") + + still_registered = mcp_death_supervisor._serve( + stream, own_pgid=4242, parent_pgid=777 + ) + + assert still_registered == {222} + + +def test_garbage_lines_do_not_cost_us_the_other_registrations(): + # A corrupted byte on the control pipe must not take down reaping for every + # other server -- that would turn a cosmetic bug into leaked processes. + stream = io.StringIO( + "register 111\n" + "\n" + "register\n" + "register notanumber\n" + "register 222 333\n" + "explode 444\n" + "register 555\n" + ) + + still_registered = mcp_death_supervisor._serve( + stream, own_pgid=4242, parent_pgid=777 + ) + + assert still_registered == {111, 555} + + +def test_a_writer_that_never_sends_a_newline_cannot_grow_us_without_bound(): + """Found for real: iterating the stream let /dev/zero reach 15 GB. + + The supervisor is the last line of defense against leaked MCP servers, so + it must not be the process that dies under memory pressure -- and a reader + that buffers until a newline arrives is exactly that risk. + """ + huge = "register " + ("0" * 10_000_000) + "\nregister 222\n" + + still_registered = mcp_death_supervisor._serve( + io.StringIO(huge), own_pgid=4242, parent_pgid=777 + ) + + # The overlong line is skipped, and the stream resyncs on the next one. + assert still_registered == {222} + + +def test_a_line_truncated_by_the_cap_is_never_acted_on(): + # Truncation must not turn one pgid into a different, valid-looking one: + # "register 999999" clipped to "register 9" would reap the wrong group. + stream = io.StringIO("register " + "9" * (mcp_death_supervisor._MAX_LINE_CHARS)) + + assert mcp_death_supervisor._serve( + stream, own_pgid=4242, parent_pgid=777 + ) == set() + + +def test_unsafe_targets_are_rejected_at_registration_time(): + stream = io.StringIO("register 0\nregister 777\nregister 999\n") + + still_registered = mcp_death_supervisor._serve( + stream, own_pgid=4242, parent_pgid=777 + ) + + assert still_registered == {999} + + +def test_refuses_to_run_inside_the_parents_own_process_group(): + # Started without start_new_session, a killpg of the parent's group would + # take the supervisor out before it could reap. It must not pretend to work. + proc = subprocess.run( + [sys.executable, SUPERVISOR, "--parent-pgid", str(os.getpgid(0))], + stdin=subprocess.DEVNULL, + capture_output=True, + text=True, + timeout=30, + ) + + assert proc.returncode == 2 + assert "process group" in proc.stderr + + +# --------------------------------------------------------------------------- +# End to end: real processes, real death +# --------------------------------------------------------------------------- + + +def test_reaps_a_registered_group_when_the_control_pipe_reaches_eof(): + victim = subprocess.Popen(_VICTIM, start_new_session=True) + supervisor = subprocess.Popen( + [sys.executable, SUPERVISOR, "--parent-pgid", str(os.getpgid(0))], + stdin=subprocess.PIPE, + text=True, + start_new_session=True, + ) + try: + supervisor.stdin.write(f"register {os.getpgid(victim.pid)}\n") + supervisor.stdin.flush() + assert victim.poll() is None, "victim should outlive registration" + + # EOF is the death signal, whatever closed the pipe. + supervisor.stdin.close() + + assert _wait_exited(victim), "registered group survived parent death" + finally: + _kill(victim.pid) + _kill(supervisor.pid) + victim.wait(timeout=10) + supervisor.wait(timeout=10) + + +def test_leaves_an_unregistered_group_alone_at_eof(): + # The other failure direction, and the more damaging one: a clean Hermes + # shutdown unregisters as it tears each server down, so EOF must not become + # a kill-everything event for servers that were handed back. + survivor = subprocess.Popen(_VICTIM, start_new_session=True) + supervisor = subprocess.Popen( + [sys.executable, SUPERVISOR, "--parent-pgid", str(os.getpgid(0))], + stdin=subprocess.PIPE, + text=True, + start_new_session=True, + ) + try: + pgid = os.getpgid(survivor.pid) + supervisor.stdin.write(f"register {pgid}\nunregister {pgid}\n") + supervisor.stdin.flush() + supervisor.stdin.close() + + supervisor.wait(timeout=15) + assert survivor.poll() is None, "a cleanly unregistered server was killed" + finally: + _kill(survivor.pid) + _kill(supervisor.pid) + survivor.wait(timeout=10) + supervisor.wait(timeout=10) + + +# A stand-in for Hermes: registers a real child, then blocks forever holding the +# only write end of the control pipe. SIGKILLing it is the scenario the whole +# module exists for -- no cleanup code of ours gets to run. +_FAKE_PARENT = """ +import os, subprocess, sys, time + +supervisor = sys.argv[1] +victim = subprocess.Popen( + [sys.executable, "-c", "import time; time.sleep(300)"], start_new_session=True +) +sup = subprocess.Popen( + [sys.executable, supervisor, "--parent-pgid", str(os.getpgid(0))], + stdin=subprocess.PIPE, text=True, start_new_session=True, +) +sup.stdin.write("register %d\\n" % os.getpgid(victim.pid)) +sup.stdin.flush() +print("%d %d" % (victim.pid, sup.pid), flush=True) +time.sleep(300) +""" + + +# Reparented-to-init processes are by definition outside this test's subtree, +# so cleaning them up trips conftest's live-system kill guard. Real signal +# delivery to a real orphan is the entire point of these two tests. +@pytest.mark.live_system_guard_bypass +def test_reaps_the_server_when_the_registering_parent_is_sigkilled(tmp_path): + script = tmp_path / "fake_parent.py" + script.write_text(_FAKE_PARENT) + + parent = subprocess.Popen( + [sys.executable, str(script), SUPERVISOR], + stdout=subprocess.PIPE, + text=True, + ) + victim_pid = supervisor_pid = None + try: + victim_pid, supervisor_pid = ( + int(x) for x in parent.stdout.readline().split() + ) + assert _alive(victim_pid) + + # No graceful anything: the parent never runs another line of Python. + parent.kill() + parent.wait(timeout=10) + + assert _wait_gone(victim_pid), ( + "stdio MCP server survived kill -9 of its Hermes parent" + ) + finally: + for pid in (victim_pid, supervisor_pid): + if pid is not None: + _kill(pid) + _kill(parent.pid) + + +@pytest.mark.live_system_guard_bypass +def test_reaps_a_grandchild_left_in_the_registered_group(tmp_path): + # Real shape of the bug: mcp-remote exits but leaves the `node` it spawned + # behind. The grandchild reparents to init but keeps the pgid, so killpg + # still reaches it -- which is why we track groups and not pids. + script = tmp_path / "leaky_server.py" + script.write_text( + "import subprocess, sys\n" + "child = subprocess.Popen([sys.executable, '-c'," + " 'import time; time.sleep(300)'])\n" + "print(child.pid, flush=True)\n" + ) + + # start_new_session mirrors how the MCP SDK spawns stdio servers. + server = subprocess.Popen( + [sys.executable, str(script)], + stdout=subprocess.PIPE, + text=True, + start_new_session=True, + ) + grandchild_pid = int(server.stdout.readline()) + server.wait(timeout=10) # the direct child exits; the grandchild does not + + supervisor = subprocess.Popen( + [sys.executable, SUPERVISOR, "--parent-pgid", str(os.getpgid(0))], + stdin=subprocess.PIPE, + text=True, + start_new_session=True, + ) + try: + assert _alive(grandchild_pid), "grandchild should outlive its parent" + # server.pid is its own pgid leader, captured at spawn time exactly as + # mcp_tool records it -- still usable after the leader itself exited. + supervisor.stdin.write(f"register {server.pid}\n") + supervisor.stdin.flush() + supervisor.stdin.close() + + assert _wait_gone(grandchild_pid), "orphaned grandchild was not reaped" + finally: + _kill(grandchild_pid) + _kill(supervisor.pid) + supervisor.wait(timeout=10) + + +# --------------------------------------------------------------------------- +# Client side: what mcp_tool tells the supervisor +# --------------------------------------------------------------------------- + + +class _FakeSupervisor: + """Stands in for the supervisor process, recording the control stream.""" + + def __init__(self, exited=False): + self.stdin = io.StringIO() + self.pid = 4242 + self._exited = exited + self._sent = "" + self.closed = False + _real_close = self.stdin.close + + def _close(): + # Mirror a real pipe: capture what was written before the write + # end goes away, so tests can still assert on the control stream. + self._sent = self.stdin.getvalue() + self.closed = True + _real_close() + + self.stdin.close = _close + + def poll(self): + return 1 if self._exited else None + + def wait(self, timeout=None): + self.waited = True + return 0 + + def lines(self): + if self.closed: + return self._sent.splitlines() + return self.stdin.getvalue().splitlines() + + +@pytest.fixture(autouse=True) +def _reset_client_state(): + yield + mcp_tool._death_supervisor = None + mcp_tool._supervised_pgids.clear() + + +@pytest.fixture +def all_groups_alive(monkeypatch): + """Answer every liveness probe with "this group exists". + + The protocol tests below register synthetic pgids that were never real + process groups. Without this, the liveness prune correctly discards them + before the control stream can be asserted on -- so state the precondition + rather than letting these tests depend on pid-space luck. + """ + monkeypatch.setattr(mcp_tool.os, "killpg", lambda pgid, sig: None) + + +def test_register_starts_the_supervisor_once_and_reuses_it(monkeypatch, all_groups_alive): + spawned = [] + + def _spawn(): + fake = _FakeSupervisor() + spawned.append(fake) + return fake + + monkeypatch.setattr(mcp_tool, "_spawn_death_supervisor", _spawn) + + mcp_tool._update_death_supervisor("register", [111]) + mcp_tool._update_death_supervisor("register", [222]) + + assert len(spawned) == 1, "each register spawned its own supervisor" + assert spawned[0].lines() == ["register 111", "register 222"] + + +def test_unregister_is_forwarded(monkeypatch, all_groups_alive): + fake = _FakeSupervisor() + monkeypatch.setattr(mcp_tool, "_spawn_death_supervisor", lambda: fake) + + mcp_tool._update_death_supervisor("register", [111]) + mcp_tool._update_death_supervisor("unregister", [111]) + + assert fake.lines() == ["register 111", "unregister 111"] + assert mcp_tool._supervised_pgids == set() + + +def test_supervisor_is_released_once_nothing_is_left_to_reap(monkeypatch, all_groups_alive): + """An empty registration set must not keep a supervisor resident. + + A gateway that once connected a stdio server would otherwise carry a + ~15 MB process and a live pipe for the rest of its life. Closing our + write end is the same EOF the supervisor treats as parent death; with + nothing registered it exits without reaping. The next register starts a + fresh one, exactly like the dead-supervisor replay path. + """ + spawned = [] + + def _spawn(): + fake = _FakeSupervisor() + spawned.append(fake) + return fake + + monkeypatch.setattr(mcp_tool, "_spawn_death_supervisor", _spawn) + + mcp_tool._update_death_supervisor("register", [111, 222]) + mcp_tool._update_death_supervisor("unregister", [111]) + assert not spawned[0].closed, "released the supervisor while a group was still registered" + + mcp_tool._update_death_supervisor("unregister", [222]) + assert spawned[0].closed, "supervisor kept resident with nothing left to reap" + assert getattr(spawned[0], "waited", False), "released supervisor was never wait()ed -> zombie until the next Popen" + assert spawned[0].lines()[-1] == "unregister 222", "release happened before the last unregister was sent" + assert mcp_tool._death_supervisor is None + + mcp_tool._update_death_supervisor("register", [333]) + assert len(spawned) == 2 and spawned[1].lines() == ["register 333"] + + +def test_supervisor_survives_the_real_eof_release(): + """End to end: closing the control pipe with nothing registered exits cleanly.""" + if os.name != "posix": + pytest.skip("POSIX-only supervisor") + child = subprocess.Popen(_VICTIM, start_new_session=True) + try: + mcp_tool._update_death_supervisor("register", [os.getpgid(child.pid)]) + proc = mcp_tool._death_supervisor + assert proc is not None and proc.poll() is None + mcp_tool._update_death_supervisor("unregister", [os.getpgid(child.pid)]) + assert mcp_tool._death_supervisor is None + assert proc.wait(timeout=10) == 0, "supervisor did not exit on the release EOF" + assert child.poll() is None, "release reaped a group that had been unregistered" + finally: + _kill(child.pid) + child.wait(timeout=10) + + +def test_unregister_alone_does_not_start_a_supervisor(monkeypatch): + spawned = [] + monkeypatch.setattr( + mcp_tool, + "_spawn_death_supervisor", + lambda: spawned.append(1) or _FakeSupervisor(), + ) + + mcp_tool._update_death_supervisor("unregister", [111]) + + assert spawned == [] + + +def test_a_dead_supervisor_is_replaced_and_live_coverage_replayed(monkeypatch, all_groups_alive): + dead = _FakeSupervisor(exited=True) + replacement = _FakeSupervisor() + queue = [dead, replacement] + monkeypatch.setattr(mcp_tool, "_spawn_death_supervisor", lambda: queue.pop(0)) + + mcp_tool._update_death_supervisor("register", [111]) + mcp_tool._update_death_supervisor("register", [222]) + + # Losing the supervisor must not silently drop the server registered with + # it -- the replacement has to be told about 111 as well as 222. + assert set(replacement.lines()) == {"register 111", "register 222"} + + +def test_replay_does_not_resurrect_an_unregistered_group(monkeypatch, all_groups_alive): + dead = _FakeSupervisor(exited=True) + replacement = _FakeSupervisor() + queue = [dead, replacement] + monkeypatch.setattr(mcp_tool, "_spawn_death_supervisor", lambda: queue.pop(0)) + + mcp_tool._update_death_supervisor("register", [111]) + mcp_tool._update_death_supervisor("register", [222]) + mcp_tool._update_death_supervisor("unregister", [111]) + + assert mcp_tool._supervised_pgids == {222} + # 111 was legitimately replayed to the replacement (it was live when the + # dead supervisor was swapped out), then unregistered. What must never + # happen is a replay AFTER the unregister bringing it back. + lines = replacement.lines() + assert lines.index("unregister 111") > lines.index("register 111") + assert "register 111" not in lines[lines.index("unregister 111") :] + mcp_tool._update_death_supervisor("register", [333]) # any later replay/append + assert "register 111" not in replacement.lines()[len(lines) :] + + +def test_a_broken_pipe_never_propagates_into_a_live_mcp_session(monkeypatch, all_groups_alive): + class _BrokenPipe(_FakeSupervisor): + def __init__(self): + super().__init__() + + class _Stdin: + def write(self, _payload): + raise BrokenPipeError("supervisor exited after poll()") + + def flush(self): + pass + + self.stdin = _Stdin() + + monkeypatch.setattr(mcp_tool, "_spawn_death_supervisor", _BrokenPipe) + + mcp_tool._update_death_supervisor("register", [111]) # must not raise + + # Dropped, so the next registration respawns instead of writing into a + # pipe that is known to be dead. + assert mcp_tool._death_supervisor is None + + +def test_unregister_after_a_broken_pipe_rebuilds_coverage_for_survivors(monkeypatch, all_groups_alive): + """A lost supervisor must be replaced by the NEXT lifecycle event, whatever its verb. + + Sequence from the #93517 review: two groups live, the control pipe dies + (write fails, supervisor dropped, set retained), then a clean teardown + unregisters one of them. Keying the no-spawn fast path on the verb left + the survivor recorded but unsupervised; it must be keyed on the set. + """ + spawned = [] + + def _spawn(): + fake = _FakeSupervisor() + spawned.append(fake) + return fake + + monkeypatch.setattr(mcp_tool, "_spawn_death_supervisor", _spawn) + mcp_tool._update_death_supervisor("register", [111, 222]) + + class _DeadStdin: + def write(self, _payload): + raise BrokenPipeError("supervisor died") + + def flush(self): + pass + + spawned[0].stdin = _DeadStdin() + mcp_tool._update_death_supervisor("register", [333]) # the write fails; supervisor dropped + assert mcp_tool._death_supervisor is None + assert mcp_tool._supervised_pgids == {111, 222, 333} + + mcp_tool._update_death_supervisor("unregister", [222]) + + assert len(spawned) == 2, "unregister after a lost supervisor did not respawn one" + assert sorted(spawned[1].lines()) == ["register 111", "register 333"], ( + "the replacement did not receive the surviving groups" + ) + assert mcp_tool._death_supervisor is spawned[1] + + +def test_a_supervisor_that_cannot_start_is_not_fatal(monkeypatch, all_groups_alive): + monkeypatch.setattr(mcp_tool, "_spawn_death_supervisor", lambda: None) + + mcp_tool._update_death_supervisor("register", [111]) # must not raise + + assert mcp_tool._death_supervisor is None + + +@contextlib.contextmanager +def _stdio_connection(child_pid, fake_supervisor): + """Drive the real MCPServerTask._run_stdio with a known spawned child. + + Only the MCP transport itself is mocked. Everything the supervisor wiring + depends on -- child discovery, _filter_mcp_children, the real os.getpgid + lookup -- runs for real against ``child_pid``, so the pgid asserted on is + the pgid of an actual process rather than a fixture value. + """ + session = MagicMock() + session.initialize = AsyncMock() + session.list_tools = AsyncMock(return_value=SimpleNamespace(tools=[])) + + stdio_cm = MagicMock() + stdio_cm.__aenter__ = AsyncMock(return_value=(object(), object())) + stdio_cm.__aexit__ = AsyncMock(return_value=False) + session_cm = MagicMock() + session_cm.__aenter__ = AsyncMock(return_value=session) + session_cm.__aexit__ = AsyncMock(return_value=False) + + with ( + patch("tools.mcp_tool.stdio_client", return_value=stdio_cm), + patch("tools.mcp_tool.ClientSession", return_value=session_cm), + # First call is the pids_before baseline; the second reports our child + # as the newly spawned server. + patch( + "tools.mcp_tool._snapshot_child_pids", + side_effect=[set(), {child_pid}], + ), + patch("tools.mcp_tool._write_stderr_log_header"), + patch("tools.mcp_tool._get_mcp_stderr_log", return_value=None), + patch( + "tools.mcp_tool._spawn_death_supervisor", + return_value=fake_supervisor, + ), + ): + yield mcp_tool.MCPServerTask("supervisor-wiring") + + +@pytest.mark.skipif(not mcp_tool._MCP_AVAILABLE, reason="MCP SDK not installed") +def test_connecting_a_stdio_server_registers_its_real_process_group(): + fake = _FakeSupervisor() + child = subprocess.Popen(_VICTIM, start_new_session=True) + try: + with _stdio_connection(child.pid, fake) as server: + asyncio.run(server.start({"command": "echo", "args": ["hi"]})) + + assert f"register {os.getpgid(child.pid)}" in fake.lines(), ( + "connecting a stdio server did not hand its process group to the " + f"supervisor; control stream was {fake.lines()}" + ) + finally: + _kill(child.pid) + child.wait(timeout=10) + + +@pytest.mark.skipif(not mcp_tool._MCP_AVAILABLE, reason="MCP SDK not installed") +def test_a_server_that_exited_is_released_on_teardown(): + fake = _FakeSupervisor() + child = subprocess.Popen(_VICTIM, start_new_session=True) + pgid = None + try: + with _stdio_connection(child.pid, fake) as server: + + async def _connect_then_lose_the_child(): + await server.start({"command": "echo", "args": ["hi"]}) + nonlocal pgid + pgid = os.getpgid(child.pid) + # The server exits while connected. Reap it here so the + # teardown path sees a genuinely dead pid, not a zombie. + child.kill() + child.wait(timeout=10) + await server.shutdown() + + asyncio.run(_connect_then_lose_the_child()) + + assert f"register {pgid}" in fake.lines() + assert f"unregister {pgid}" in fake.lines(), ( + "a stdio server with nothing left alive stayed registered, so the " + f"supervisor would keep a stale group; stream was {fake.lines()}" + ) + finally: + _kill(child.pid) + + +@pytest.mark.skipif(not mcp_tool._MCP_AVAILABLE, reason="MCP SDK not installed") +def test_a_server_that_survived_teardown_stays_registered(): + # The case the whole module exists for: teardown did not manage to kill it. + # Releasing it here would hand the orphan back to nobody. + fake = _FakeSupervisor() + child = subprocess.Popen(_VICTIM, start_new_session=True) + try: + with _stdio_connection(child.pid, fake) as server: + + async def _connect_then_shutdown(): + await server.start({"command": "echo", "args": ["hi"]}) + await server.shutdown() + + asyncio.run(_connect_then_shutdown()) + + pgid = os.getpgid(child.pid) + assert f"register {pgid}" in fake.lines() + assert f"unregister {pgid}" not in fake.lines(), ( + "a server that outlived teardown was released from the supervisor, " + "so an ungraceful exit would leave it running forever" + ) + finally: + _kill(child.pid) + child.wait(timeout=10) + + +@pytest.mark.live_system_guard_bypass +def test_scoped_teardown_of_one_owner_keeps_the_other_owner_supervised(monkeypatch): + """Two owners (profiles / agents) each hold a stdio group; tearing one down + must release only that owner's group and leave the other covered, and the + per-process supervisor must then still know about the survivor. + + Exercises the real registry + ``_kill_orphaned_mcp_children`` scoping + rather than the control protocol alone (review request on #93517). + """ + fake = _FakeSupervisor() + monkeypatch.setattr(mcp_tool, "_spawn_death_supervisor", lambda: fake) + monkeypatch.setattr(mcp_tool.time, "sleep", lambda _s: None) # skip the SIGTERM grace wait + a = subprocess.Popen(_VICTIM, start_new_session=True) + b = subprocess.Popen(_VICTIM, start_new_session=True) + try: + pg_a, pg_b = os.getpgid(a.pid), os.getpgid(b.pid) + with mcp_tool._lock: + mcp_tool._stdio_pids[a.pid] = "profile-a" + mcp_tool._stdio_pids[b.pid] = "profile-b" + mcp_tool._stdio_pgids[a.pid] = pg_a + mcp_tool._stdio_pgids[b.pid] = pg_b + mcp_tool._update_death_supervisor("register", [pg_a, pg_b]) + + mcp_tool._kill_orphaned_mcp_children(include_active=True, server_name="profile-a") + a.wait(timeout=10) + + assert b.poll() is None, "scoped teardown of profile-a killed profile-b's server" + assert f"unregister {pg_a}" in fake.lines() + assert f"unregister {pg_b}" not in fake.lines(), ( + "scoped teardown released the OTHER owner's group from the supervisor" + ) + assert mcp_tool._supervised_pgids == {pg_b} + assert b.pid in mcp_tool._stdio_pids and b.pid in mcp_tool._stdio_pgids + finally: + for p in (a, b): + _kill(p.pid) + try: + p.wait(timeout=10) + except Exception: # noqa: BLE001 - best-effort cleanup + pass + with mcp_tool._lock: + for p in (a, b): + mcp_tool._stdio_pids.pop(p.pid, None) + mcp_tool._stdio_pgids.pop(p.pid, None) + + +@pytest.mark.live_system_guard_bypass +def test_a_group_with_nothing_left_alive_is_forgotten_and_unregistered(monkeypatch): + """A dead group must not stay registered: its pgid can be recycled. + + Uses a real process so the liveness probe is answered by the kernel rather + than a fixture -- the whole point is that we notice actual death. + """ + fake = _FakeSupervisor() + monkeypatch.setattr(mcp_tool, "_spawn_death_supervisor", lambda: fake) + + doomed = subprocess.Popen(_VICTIM, start_new_session=True) + doomed_pgid = os.getpgid(doomed.pid) + survivor = subprocess.Popen(_VICTIM, start_new_session=True) + survivor_pgid = os.getpgid(survivor.pid) + try: + mcp_tool._update_death_supervisor("register", [doomed_pgid, survivor_pgid]) + assert mcp_tool._supervised_pgids == {doomed_pgid, survivor_pgid} + + # Reap it fully so the group is genuinely empty, not a zombie. + doomed.kill() + doomed.wait(timeout=10) + + # Any later registration change is when we notice. + mcp_tool._update_death_supervisor("register", [survivor_pgid]) + + assert doomed_pgid not in mcp_tool._supervised_pgids, ( + "a group with no members left stayed registered, so a recycled " + "pgid could later be reaped as if it were an MCP server" + ) + assert survivor_pgid in mcp_tool._supervised_pgids, ( + "pruning dropped a group that is still alive" + ) + assert f"unregister {doomed_pgid}" in fake.lines(), ( + "the supervisor was never told to forget the dead group" + ) + finally: + _kill(survivor.pid) + survivor.wait(timeout=10) + _kill(doomed.pid) + + +def test_pruning_keeps_groups_it_cannot_prove_are_gone(monkeypatch): + # An ambiguous probe (EPERM: exists but not ours) must not drop coverage -- + # losing a real registration is worse than keeping a doubtful one. + monkeypatch.setattr(mcp_tool, "_supervised_pgids", {111, 222}, raising=False) + + def _probe(pgid, sig): + if pgid == 111: + raise PermissionError("exists, not ours") + raise ProcessLookupError("gone") + + monkeypatch.setattr(mcp_tool.os, "killpg", _probe) + + stale = mcp_tool._prune_dead_supervised_pgids() + + assert stale == {222} + assert mcp_tool._supervised_pgids == {111} + + +def test_no_pgids_is_a_no_op(monkeypatch): + spawned = [] + monkeypatch.setattr( + mcp_tool, + "_spawn_death_supervisor", + lambda: spawned.append(1) or _FakeSupervisor(), + ) + + mcp_tool._update_death_supervisor("register", []) + + assert spawned == [] diff --git a/tests/tools/test_mcp_npx_cached_bin.py b/tests/tools/test_mcp_npx_cached_bin.py new file mode 100644 index 0000000000..40a16640d8 --- /dev/null +++ b/tests/tools/test_mcp_npx_cached_bin.py @@ -0,0 +1,208 @@ +"""``npx -y `` should spawn the cached binary, not a resident `npm exec`. + +`npx` resolves the package and then FORKS, staying alive as the real server's +parent for the whole process lifetime while doing no work. Measured on a +4-agent host that is ~48 MB of private memory per MCP server — and it buys +nothing, because Hermes already wraps the child in its own parent-death +watchdog, so npx's supervision is a second parent nobody reads. + +Removing it must stay conservative: a cache miss, a version-pinned spec, or an +ambiguous ``bin`` map all fall back to plain `npx` so a cold machine still +installs normally. +""" + +from __future__ import annotations + +import json +import os + +import pytest + +from tools.mcp_tool import _npx_cached_bin + + +def _cache(tmp_path, *, package, deps=None, bin_field, make_bin=True, entry="abc123"): + """Build a fake npx cache entry the way npm lays one out.""" + root = tmp_path / ".npm" / "_npx" / entry + (root / "node_modules" / package).mkdir(parents=True) + (root / "package.json").write_text( + json.dumps({"dependencies": deps if deps is not None else {package: "^1.0.0"}}), + encoding="utf-8", + ) + (root / "node_modules" / package / "package.json").write_text( + json.dumps({"name": package, "bin": bin_field}), encoding="utf-8" + ) + bindir = root / "node_modules" / ".bin" + bindir.mkdir(parents=True, exist_ok=True) + name = bin_field if isinstance(bin_field, str) else list(bin_field)[0] + target = bindir / (os.path.basename(package) if isinstance(bin_field, str) else name) + if make_bin: + target.write_text("#!/usr/bin/env node\n", encoding="utf-8") + target.chmod(0o755) + return target + + +@pytest.fixture(autouse=True) +def _isolate_cache(tmp_path, monkeypatch): + monkeypatch.setenv("npm_config_cache", str(tmp_path / ".npm")) + yield + + +def test_cached_package_resolves_to_its_binary(tmp_path): + target = _cache(tmp_path, package="mcp-linear", bin_field={"mcp-linear": "dist/index.js"}) + + got = _npx_cached_bin(["-y", "mcp-linear"]) + + assert got == (str(target), []) + + +def test_scoped_package_and_trailing_args_survive(tmp_path): + target = _cache( + tmp_path, + package="@tacticlaunch/mcp-linear", + bin_field={"mcp-linear": "dist/index.js"}, + ) + + got = _npx_cached_bin(["-y", "@tacticlaunch/mcp-linear", "--port", "7"]) + + assert got == (str(target), ["--port", "7"]) + + +def test_uncached_package_falls_back_to_npx(tmp_path): + _cache(tmp_path, package="something-else", bin_field={"something-else": "i.js"}) + + assert _npx_cached_bin(["-y", "mcp-linear"]) is None + + +def test_version_pinned_spec_is_left_to_npx(tmp_path): + _cache(tmp_path, package="mcp-linear", bin_field={"mcp-linear": "dist/index.js"}) + + # The user pinned a build; npx owns that resolution and the cache key for + # a different version would not match this entry. + assert _npx_cached_bin(["-y", "mcp-linear@1.2.3"]) is None + + +def test_ambiguous_bin_map_is_left_to_npx(tmp_path): + _cache( + tmp_path, + package="multi", + bin_field={"one": "a.js", "two": "b.js"}, + ) + + # Which bin npx would choose is not ours to guess. + assert _npx_cached_bin(["-y", "multi"]) is None + + +def test_missing_or_non_executable_binary_falls_back(tmp_path): + _cache( + tmp_path, + package="mcp-linear", + bin_field={"mcp-linear": "dist/index.js"}, + make_bin=False, + ) + + assert _npx_cached_bin(["-y", "mcp-linear"]) is None + + +def test_no_cache_directory_at_all(tmp_path, monkeypatch): + monkeypatch.setenv("npm_config_cache", str(tmp_path / "nope")) + + assert _npx_cached_bin(["-y", "mcp-linear"]) is None + + +def test_corrupt_cache_manifest_is_skipped(tmp_path): + root = tmp_path / ".npm" / "_npx" / "broken" + root.mkdir(parents=True) + (root / "package.json").write_text("{ not json", encoding="utf-8") + + assert _npx_cached_bin(["-y", "mcp-linear"]) is None + + +@pytest.mark.parametrize("args", [[], ["-y"], ["--yes"], ["-p", "x"], None, "notalist"]) +def test_unusable_args_are_ignored(args): + assert _npx_cached_bin(args) is None + + +def test_osv_preflight_runs_before_the_swap(): + """The malware gate must still see `npx` + the package name. + + `_infer_ecosystem` keys off the command basename, so a command already + rewritten to `.../node_modules/.bin/mcp-linear` yields no ecosystem and + `check_package_for_malware` returns None — the gate silently becomes a + no-op. This pins the ordering: OSV inspects the original invocation. + """ + from tools.osv_check import _infer_ecosystem, _parse_package_from_args + + # What the preflight sees today, before any swap. + assert _infer_ecosystem("npx") == "npm" + assert _parse_package_from_args(["-y", "@tacticlaunch/mcp-linear"], "npm")[0] == ( + "@tacticlaunch/mcp-linear" + ) + + # What it would see if the swap happened first — nothing. + assert _infer_ecosystem("/home/u/.npm/_npx/abc/node_modules/.bin/mcp-linear") is None + + +def test_swap_happens_after_the_osv_call_in_source(): + """Structural guard for the ordering above. + + The swap and the preflight live in one async function; a future edit that + moves the swap earlier would disable the malware gate silently, and no + unit test of either piece alone would notice. + """ + from pathlib import Path as _P + + src = _P(__file__).resolve().parents[2] / "tools" / "mcp_tool.py" + text = src.read_text(encoding="utf-8") + osv_needle = "check_package_for_malware, command, args" + swap_needle = "cached = _npx_cached_bin(args)" + # Report a rename explicitly: a bare .index() ValueError here reads like a + # broken test rather than "someone renamed the thing this guards". + assert osv_needle in text, ( + f"cannot find the OSV preflight call ({osv_needle!r}) — it was renamed; " + "update this guard and re-verify the swap still happens after it" + ) + assert swap_needle in text, ( + f"cannot find the npx swap ({swap_needle!r}) — it was renamed; update " + "this guard and re-verify it still happens after the OSV preflight" + ) + + assert text.index(osv_needle) < text.index(swap_needle), ( + "the npx swap now precedes the OSV malware preflight, which silently " + "disables it: _infer_ecosystem keys off the command basename being " + "npx/uvx/pipx, so a rewritten command yields no ecosystem and " + "check_package_for_malware returns None" + ) + + +def test_windows_selects_launchers_never_the_sh_script(): + """On Windows the extensionless sh script must never be chosen. + + npm lays down three siblings per bin — ``, `.cmd`, + `.ps1` — and spawning the sh one from a Windows process fails, while + `os.access(X_OK)` there is effectively an existence check and cannot tell + them apart. Tested through the injectable helper rather than by patching + `os.name`, which breaks path handling process-wide (it took pytest's own + traceback formatting down when I tried). + """ + from tools.mcp_tool import _npx_bin_candidates + + win = _npx_bin_candidates("/c/bin", "mcp-linear", windows=True) + assert win == ["/c/bin/mcp-linear.cmd", "/c/bin/mcp-linear.exe"] + assert not any(c.endswith("mcp-linear") for c in win), "sh script must not be a candidate" + + assert _npx_bin_candidates("/bin", "mcp-linear", windows=False) == ["/bin/mcp-linear"] + + +def test_posix_resolution_uses_the_helper(tmp_path): + """The resolver honours the helper's ordering (POSIX path end-to-end).""" + target = _cache(tmp_path, package="mcp-linear", bin_field={"mcp-linear": "i.js"}) + + assert _npx_cached_bin(["-y", "mcp-linear"]) == (str(target), []) + + +def test_flag_after_the_spec_is_left_to_npx(tmp_path): + """`npx pkg -y` would forward -y to the server; that shape stays with npx.""" + _cache(tmp_path, package="mcp-linear", bin_field={"mcp-linear": "i.js"}) + + assert _npx_cached_bin(["mcp-linear", "-y"]) is None diff --git a/tests/tools/test_mcp_stdio_watchdog.py b/tests/tools/test_mcp_stdio_watchdog.py deleted file mode 100644 index 411695eea5..0000000000 --- a/tests/tools/test_mcp_stdio_watchdog.py +++ /dev/null @@ -1,40 +0,0 @@ -"""Contract tests for the direct POSIX stdio MCP child watchdog.""" - -import os -import sys - -import pytest - -from tools import mcp_stdio_watchdog, mcp_tool - - -def test_is_orphaned_is_false_while_direct_parent_is_unchanged(): - original_ppid = 1234 - - assert mcp_stdio_watchdog._is_orphaned( - original_ppid, - getppid=lambda: original_ppid, - ) is False - - -@pytest.mark.skipif(os.name != "posix", reason="watchdog wrapping is POSIX-only") -def test_wrap_command_uses_stable_parent_pid_and_preserves_command_tail(): - parent_pid = os.getpid() - command = "/opt/hermes/bin/mcp-server" - command_args = ["--label", "value with spaces", "--", "literal-tail"] - - wrapped_command, wrapped_args = mcp_tool._wrap_command_with_watchdog( - command, - command_args, - ) - - assert wrapped_command == sys.executable - assert wrapped_args == [ - os.path.join(os.path.dirname(mcp_tool.__file__), "mcp_stdio_watchdog.py"), - "--ppid", - str(parent_pid), - "--", - command, - *command_args, - ] - assert "--create-time" not in wrapped_args diff --git a/tests/tools/test_mcp_structured_content.py b/tests/tools/test_mcp_structured_content.py index 7ce59324de..b8d401cfa4 100644 --- a/tests/tools/test_mcp_structured_content.py +++ b/tests/tools/test_mcp_structured_content.py @@ -163,6 +163,7 @@ class TestMetaPassthrough: assert data == {"result": "done"} def test_meta_with_structured_content(self, _patch_mcp_server): + """With usable text, structuredContent is suppressed but _meta rides.""" session = _patch_mcp_server session.call_tool = AsyncMock( return_value=_FakeCallToolResult( @@ -175,7 +176,6 @@ class TestMetaPassthrough: data = json.loads(handler({})) assert data == { "result": "txt", - "structuredContent": {"ok": True}, "_meta": {"com.example/k": "v"}, } @@ -215,3 +215,114 @@ class TestReservedMetaKeyPredicate: assert not mcp_tool._is_reserved_mcp_meta_key("com.example/x") assert not mcp_tool._is_reserved_mcp_meta_key("plain-key") assert not mcp_tool._is_reserved_mcp_meta_key("/leading-slash") + + +class TestContentStructuredArbitration: + """content and structuredContent are alternatives — never both. + + Ported from MoonshotAI/kimi-code#3234: spec-following servers render + their data into content (verbatim dual-emit or a faithful human + reorganisation), so forwarding both sent the same information twice. + """ + + def test_dual_emit_suppresses_structured(self, _patch_mcp_server): + """Verbatim dual-emit servers: model receives content only.""" + session = _patch_mcp_server + payload = {"items": [1, 2, 3]} + session.call_tool = AsyncMock( + return_value=_FakeCallToolResult( + content=[_FakeContentBlock(json.dumps(payload))], + structuredContent=payload, + ) + ) + handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0) + data = json.loads(handler({})) + assert data == {"result": json.dumps(payload)} + + def test_prose_summary_suppresses_structured(self, _patch_mcp_server): + """Lossy prose summaries also win — no heuristic is attempted.""" + session = _patch_mcp_server + session.call_tool = AsyncMock( + return_value=_FakeCallToolResult( + content=[_FakeContentBlock("3 item(s) found")], + structuredContent={"items": [1, 2, 3]}, + ) + ) + handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0) + data = json.loads(handler({})) + assert data == {"result": "3 item(s) found"} + + def test_whitespace_only_content_falls_back(self, _patch_mcp_server): + """Whitespace-only text is not usable content — fallback fires.""" + session = _patch_mcp_server + payload = {"status": "ok"} + session.call_tool = AsyncMock( + return_value=_FakeCallToolResult( + content=[_FakeContentBlock(" \n")], + structuredContent=payload, + ) + ) + handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0) + data = json.loads(handler({})) + assert data["structuredContent"] == payload + + def test_structured_only_still_surfaced(self, _patch_mcp_server): + """structuredContent-only servers keep working (#2596 fix preserved).""" + session = _patch_mcp_server + payload = {"only": "structured"} + session.call_tool = AsyncMock( + return_value=_FakeCallToolResult( + content=[], + structuredContent=payload, + ) + ) + handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0) + data = json.loads(handler({})) + assert data["result"] == payload + + +class TestDroppedBlockNotice: + """Unsupported content blocks surface a drop notice to the model. + + Ported from MoonshotAI/kimi-code#3227. + """ + + def test_unsupported_block_renders_notice(self, _patch_mcp_server): + session = _patch_mcp_server + # NOTE: no `uri` — a uri'd block without .resource is rendered as a + # resource link by _render_mcp_resource_block, not dropped. + weird = SimpleNamespace( + type="hologram", + mimeType="application/x-hologram", + size=1234, + ) + session.call_tool = AsyncMock( + return_value=_FakeCallToolResult(content=[weird]) + ) + handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0) + data = json.loads(handler({})) + assert "[MCP content dropped: unsupported block" in data["result"] + assert "type=hologram" in data["result"] + assert "mimeType=application/x-hologram" in data["result"] + assert "size=1234" in data["result"] + + def test_drop_notice_does_not_suppress_structured(self, _patch_mcp_server): + """A drop notice is not usable content — structured fallback fires.""" + session = _patch_mcp_server + weird = SimpleNamespace(type="hologram") + payload = {"real": "data"} + session.call_tool = AsyncMock( + return_value=_FakeCallToolResult( + content=[weird], structuredContent=payload, + ) + ) + handler = mcp_tool._make_tool_handler("test-server", "my-tool", 30.0) + data = json.loads(handler({})) + assert data["structuredContent"] == payload + assert "[MCP content dropped" in data["result"] + + def test_notice_helper_minimal_block(self): + notice = mcp_tool._render_mcp_dropped_block_notice( + SimpleNamespace(), "mystery" + ) + assert notice == "[MCP content dropped: unsupported block (type=mystery)]" diff --git a/tests/tools/test_osv_check.py b/tests/tools/test_osv_check.py index 72e81058e6..d136602f02 100644 --- a/tests/tools/test_osv_check.py +++ b/tests/tools/test_osv_check.py @@ -1,6 +1,9 @@ """Tests for OSV malware check on MCP extension packages.""" import json +import time +from pathlib import Path + import pytest from unittest.mock import patch, MagicMock @@ -66,14 +69,18 @@ class TestParsePackageFromArgs: class TestCheckPackageForMalware: @pytest.fixture(autouse=True) - def _fresh_cache(self): + def _fresh_cache(self, tmp_path, monkeypatch): from tools import osv_check + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) with osv_check._cache_lock: osv_check._cache.clear() + osv_check._disk_cache_loaded = False + (tmp_path / "cache" / "osv_check.json").unlink(missing_ok=True) yield with osv_check._cache_lock: osv_check._cache.clear() - + osv_check._disk_cache_loaded = False + (tmp_path / "cache" / "osv_check.json").unlink(missing_ok=True) def test_clean_package(self): """Clean package returns None (allow).""" mock_response = MagicMock() @@ -189,6 +196,94 @@ class TestCheckPackageForMalware: check_package_for_malware("uvx", ["mcp-server-fetch"]) assert mock_url.call_count == 2 + def test_disk_cache_persists_and_reloads(self, tmp_path, monkeypatch): + """A warm disk cache is reused by a fresh in-process cache.""" + from tools import osv_check + + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + + mock_response = MagicMock() + mock_response.read.return_value = json.dumps({"vulns": []}).encode() + mock_response.__enter__ = lambda s: s + mock_response.__exit__ = MagicMock(return_value=False) + + with patch("tools.osv_check.urllib.request.urlopen", return_value=mock_response) as mock_url: + check_package_for_malware("uvx", ["mcp-server-persist"]) + + cache_file = tmp_path / "cache" / "osv_check.json" + assert cache_file.exists(), "disk cache should be written after a warm result" + + with osv_check._cache_lock: + osv_check._cache.clear() + osv_check._disk_cache_loaded = False + + with patch("tools.osv_check.urllib.request.urlopen", return_value=mock_response) as mock_url2: + check_package_for_malware("uvx", ["mcp-server-persist"]) + + assert mock_url2.call_count == 0, "disk cache must satisfy the second call" + + def test_disk_cache_format_versioned(self, tmp_path, monkeypatch): + """Disk cache JSON has a version field and recoverable entries.""" + from tools import osv_check + + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + + mock_response = MagicMock() + mock_response.read.return_value = json.dumps({"vulns": []}).encode() + mock_response.__enter__ = lambda s: s + mock_response.__exit__ = MagicMock(return_value=False) + + with patch("tools.osv_check.urllib.request.urlopen", return_value=mock_response): + check_package_for_malware("uvx", ["mcp-server-format"]) + + cache_file = tmp_path / "cache" / "osv_check.json" + with open(cache_file, "r", encoding="utf-8") as f: + data = json.load(f) + assert data["version"] == osv_check._DISK_CACHE_VERSION + assert "entries" in data + key = "PyPI|mcp-server-format|" + assert key in data["entries"] + assert "expiry" in data["entries"][key] + assert data["entries"][key]["result"] is None + + def test_disk_cache_retries_after_transient_oserror(self, tmp_path, monkeypatch): + """A busy/unreadable cache file must not disable disk loads for the process.""" + from tools import osv_check + + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + cache_file = tmp_path / "cache" / "osv_check.json" + cache_file.parent.mkdir(parents=True, exist_ok=True) + cache_file.write_text( + json.dumps({ + "version": osv_check._DISK_CACHE_VERSION, + "entries": { + "PyPI|mcp-server-retry|": { + "expiry": time.time() + 3600, + "result": None, + } + }, + }), + encoding="utf-8", + ) + + real_open = open + calls = {"n": 0} + + def flaky_open(path, *args, **kwargs): + if Path(path) == cache_file: + calls["n"] += 1 + if calls["n"] == 1: + raise OSError("resource temporarily unavailable") + return real_open(path, *args, **kwargs) + + monkeypatch.setattr("builtins.open", flaky_open) + with osv_check._cache_lock: + osv_check._load_disk_cache() + assert osv_check._disk_cache_loaded is False + osv_check._load_disk_cache() + assert osv_check._disk_cache_loaded is True + assert ("PyPI", "mcp-server-retry", None) in osv_check._cache + class TestLiveOsvQuery: """Live integration test against the real OSV API. Skipped if offline.""" diff --git a/tests/tools/test_process_registry.py b/tests/tools/test_process_registry.py index 0d97a92ee0..4930b2ac46 100644 --- a/tests/tools/test_process_registry.py +++ b/tests/tools/test_process_registry.py @@ -2,6 +2,8 @@ import json import os +import shlex +import shutil import signal import subprocess import sys @@ -780,7 +782,7 @@ class TestSpawnEnvSanitization: def __init__(self): self.commands = [] self._responses = iter([ - {"output": "hello\n"}, + {"output": "6 0\nhello\n"}, {"output": "1\n"}, {"output": "0\n"}, ]) @@ -801,11 +803,180 @@ class TestSpawnEnvSanitization: "/path with spaces/hermes_bg.exit", ) - assert env.commands[0][0] == "cat '/path with spaces/hermes_bg.log' 2>/dev/null" + assert "'/path with spaces/hermes_bg.log'" in env.commands[0][0] + assert "cat '/path with spaces/hermes_bg.log'" not in env.commands[0][0] assert env.commands[1][0] == "kill -0 \"$(cat '/path with spaces/hermes_bg.pid' 2>/dev/null)\" 2>/dev/null; echo $?" assert env.commands[2][0] == "cat '/path with spaces/hermes_bg.exit' 2>/dev/null" +class TestEnvPollerIncrementalRead: + """The sandbox log poller must read only new bytes, not the whole file. + + Reading the whole file every poll made one poll cost grow with the total + output so far, so a long noisy job re-sent all of its output over the + docker or SSH channel every two seconds. + """ + + @staticmethod + def _run_poller(registry, session, responses): + """Drive one poll cycle and hand back the commands the env saw.""" + + class FakeEnv: + def __init__(self): + self.commands = [] + self._responses = iter(responses) + + def execute(self, command, **kwargs): + self.commands.append(command) + return next(self._responses) + + env = FakeEnv() + with patch("tools.process_registry.time.sleep", return_value=None), \ + patch.object(registry, "_move_to_finished"): + registry._env_poller_loop( + session, env, "/tmp/bg.log", "/tmp/bg.pid", "/tmp/bg.exit" + ) + return env.commands + + def test_read_command_asks_only_for_new_bytes(self): + cmd = ProcessRegistry._log_delta_command("'/tmp/bg.log'", 4096) + # The offset is carried into the command, and the file is opened with + # tail rather than cat. + assert "O=4096" in cmd + assert "tail -c +$((O+1)) '/tmp/bg.log'" in cmd + assert "cat '/tmp/bg.log'" not in cmd + + def test_read_command_starts_from_zero_on_first_poll(self): + cmd = ProcessRegistry._log_delta_command("'/tmp/bg.log'", 0) + assert "O=0" in cmd + + @pytest.mark.skipif(not shutil.which("sh"), reason="needs a POSIX sh") + def test_read_command_holds_back_a_split_utf8_sequence(self, tmp_path): + """A multibyte character straddling two polls must not be split. + + The backend decodes each execute() result on its own, so returning + the first byte of an 'é' in one poll and the rest in the next would + yield replacement characters in the transcript (and break watch + patterns at the seam). Every prefix of a mixed ASCII/2/3/4-byte + string must come back decodable, with at most 3 bytes held back and + nothing held back once the trailing character is complete. + """ + full = "hé😀中a\n€bz🚀".encode() + log = tmp_path / "bg.log" + quoted = shlex.quote(str(log)) + for n in range(1, len(full) + 1): + log.write_bytes(full[:n]) + out = subprocess.run( + ["sh", "-c", ProcessRegistry._log_delta_command(quoted, 0)], + capture_output=True, timeout=30, + ).stdout + header, _, delta = out.partition(b"\n") + size, _offset = map(int, header.split()) + delta.decode("utf-8") # must not raise + assert delta == full[:size] + complete = full[:n].decode("utf-8", "ignore").encode() == full[:n] + assert (n - size) == 0 if complete else 0 < (n - size) <= 3 + + def test_first_poll_reads_from_the_start(self, registry): + session = _make_session(sid="proc_delta") + session.exited = False + commands = self._run_poller( + registry, + session, + [ + {"output": "11 0\nfirst chunk"}, + {"output": "1\n"}, + {"output": "0\n"}, + ], + ) + assert "O=0" in commands[0] + assert session.output_buffer == "first chunk" + + def test_delta_is_appended_not_replaced(self, registry): + session = _make_session(sid="proc_append", output="already here ") + session.exited = False + self._run_poller( + registry, + session, + [ + {"output": "8 0\nand new"}, + {"output": "1\n"}, + {"output": "0\n"}, + ], + ) + assert session.output_buffer == "already here and new" + + def test_second_poll_asks_from_where_the_first_one_stopped(self, registry): + session = _make_session(sid="proc_two_polls") + session.exited = False + commands = self._run_poller( + registry, + session, + [ + {"output": "11 0\nfirst chunk"}, + {"output": "0\n"}, # still running, poll again + {"output": "17 11\n and more"}, + {"output": "1\n"}, # gone now + {"output": "0\n"}, + ], + ) + assert "O=0" in commands[0] + # The second read starts at byte 11, so the first chunk is not sent + # a second time. + assert "O=11" in commands[2] + assert session.output_buffer == "first chunk and more" + + def test_truncated_log_drops_the_stale_buffer(self, registry): + session = _make_session(sid="proc_rotate") + session.exited = False + # The second read reports offset 0 even though the first one left off + # at byte 11. The file no longer reaches that byte, so it was rotated + # or truncated and the buffer we hold no longer matches it. + self._run_poller( + registry, + session, + [ + {"output": "11 0\nfirst chunk"}, + {"output": "0\n"}, # still running, poll again + {"output": "5 0\nfresh"}, + {"output": "1\n"}, + {"output": "0\n"}, + ], + ) + assert session.output_buffer == "fresh" + + def test_unreadable_header_leaves_the_buffer_alone(self, registry): + session = _make_session(sid="proc_bad", output="keep me") + session.exited = False + # No header at all, for example when the shell is missing one of the + # tools the command needs. + self._run_poller( + registry, + session, + [ + {"output": ""}, + {"output": "1\n"}, + {"output": "0\n"}, + ], + ) + assert session.output_buffer == "keep me" + + def test_buffer_stays_within_the_cap(self, registry): + session = _make_session(sid="proc_cap") + session.exited = False + session.max_output_chars = 10 + self._run_poller( + registry, + session, + [ + {"output": "20 0\n" + "x" * 20}, + {"output": "1\n"}, + {"output": "0\n"}, + ], + ) + assert session.output_buffer == "x" * 10 + + # ========================================================================= # Popen leak prevention # ========================================================================= diff --git a/tests/tools/test_search_budget_truncation.py b/tests/tools/test_search_budget_truncation.py index 432327bd61..9e58c48ba1 100644 --- a/tests/tools/test_search_budget_truncation.py +++ b/tests/tools/test_search_budget_truncation.py @@ -1,7 +1,10 @@ +import shlex from unittest.mock import MagicMock import pytest +import tools.file_operations as file_operations +from tools.environments.local import LocalEnvironment from tools.file_operations import ExecuteResult, ShellFileOperations, _search_stdout_and_limit @@ -71,3 +74,247 @@ def test_real_rg_error_still_hard_fails(ops, monkeypatch): assert result.error == "Search failed: rg: regex parse error:" assert result.limit_reason is None + + +class FindRecordingEnvironment: + is_local = False + cwd = "/narrow" + + def __init__(self, output="", code=0): + self.output = output + self.code = code + self.commands = [] + + def execute(self, command, **kwargs): + self.commands.append((command, kwargs)) + if command.startswith("command -v find"): + return {"output": "yes\n", "returncode": 0} + if command.startswith("command -v rg"): + return {"output": "", "returncode": 1} + if "find " in command: + return {"output": self.output, "returncode": self.code} + return {"output": "", "returncode": 1} + + @property + def find_commands(self): + return [ + item for item in self.commands + if item[0].startswith("find ") or "; find " in item[0] + ] + + +class MultiRootFindEnvironment(FindRecordingEnvironment): + def execute(self, command, **kwargs): + self.commands.append((command, kwargs)) + if command.startswith("test -e "): + output = "not_found\n" if "'/one/.hidden /two/.cache'" in command else "exists\n" + return {"output": output, "returncode": 0} + if command.startswith("command -v find"): + return {"output": "yes\n", "returncode": 0} + if command.startswith("command -v rg"): + return {"output": "", "returncode": 1} + if "find " in command: + return {"output": self.output, "returncode": self.code} + return {"output": "", "returncode": 1} + + +def test_find_discovery_is_one_unsorted_pruned_bounded_scan(): + env = FindRecordingEnvironment("/narrow/a.py\n/narrow/b.py\n/narrow/c.py\n/narrow/d.py\n") + result = ShellFileOperations(env)._search_files( + "*.py", "/narrow", limit=2, offset=1, order="discovery" + ) + assert result.files == ["/narrow/b.py", "/narrow/c.py"] + assert result.truncated is True + assert len(env.find_commands) == 1 + command, kwargs = env.find_commands[0] + assert "-printf" not in command + assert "sort " not in command + assert "-prune" in command + assert "head -n 4" in command + assert kwargs["timeout"] <= 60 + + +def test_find_modified_is_one_exact_scan_without_bsd_retry(): + env = FindRecordingEnvironment("30 /narrow/new.py\n20 /narrow/mid.py\n10 /narrow/old.py\n") + result = ShellFileOperations(env)._search_files( + "*.py", "/narrow", limit=1, offset=1, order="modified" + ) + assert result.files == ["/narrow/mid.py"] + assert result.truncated is True + assert len(env.find_commands) == 1 + command, _ = env.find_commands[0] + assert "-printf '%T@ %p\\n'" in command + assert "sort -rn" in command + assert "head -n 3" in command + + +def test_no_rg_multi_root_modified_is_one_globally_sorted_scan(): + env = MultiRootFindEnvironment( + "30 /two/.cache/new.py\n10 /one/.hidden/old.py\n" + ) + + result = ShellFileOperations(env).search( + "*.py", + path="/one/.hidden /two/.cache", + target="files", + order="modified", + limit=1, + ) + + assert result.error is None + assert result.files == ["/two/.cache/new.py"] + assert result.truncated is True + assert len(env.find_commands) == 1 + command, kwargs = env.find_commands[0] + assert "find '/one/.hidden' '/two/.cache'" in command + assert "sort -rn" in command + assert "head -n 2" in command + assert "! -path '/one/.hidden'" in command + assert "! -path '/two/.cache'" in command + assert kwargs["timeout"] <= 60 + + +def test_find_dash_prefixed_relative_root_is_an_explicit_operand( + tmp_path, monkeypatch +): + dash_root = tmp_path / "--version" + ordinary_root = tmp_path / "ordinary" + dash_root.mkdir() + ordinary_root.mkdir() + (dash_root / "dash.py").write_text("", encoding="utf-8") + (ordinary_root / "plain.py").write_text("", encoding="utf-8") + + ops = ShellFileOperations(LocalEnvironment(str(tmp_path))) + monkeypatch.setattr(ops, "_has_command", lambda command: command == "find") + executed = [] + real_exec = ops._exec + + def recording_exec(command, **kwargs): + if command.startswith("set -o pipefail; find "): + executed.append(command) + return real_exec(command, **kwargs) + + monkeypatch.setattr(ops, "_exec", recording_exec) + result = ops._search_files( + "*.py", ["--version", "ordinary"], limit=10, offset=0 + ) + + assert result.error is None + assert sorted(result.files) == ["./--version/dash.py", "ordinary/plain.py"] + assert len(executed) == 1 + command_tokens = shlex.split(executed[0].removeprefix("set -o pipefail; ")) + assert "./--version" in command_tokens + assert "--version" not in command_tokens + assert all("find (GNU findutils)" not in path for path in result.files) + + +def test_find_modified_capability_failure_is_actionable_without_retry(): + env = FindRecordingEnvironment("", code=1) + result = ShellFileOperations(env)._search_files( + "*.py", "/narrow", limit=2, offset=0, order="modified" + ) + assert "modification-time" in (result.error or "") + assert len(env.find_commands) == 1 + + +@pytest.mark.parametrize( + ("order", "output"), + [ + ( + "discovery", + "/narrow/a.py\n/narrow/b.py\n/narrow/c.py\n/narrow/d.py\n", + ), + ( + "modified", + "40 /narrow/a.py\n30 /narrow/b.py\n20 /narrow/c.py\n10 /narrow/d.py\n", + ), + ], +) +def test_find_sigpipe_is_benign_only_after_fetch_limit_rows(order, output): + result = ShellFileOperations(FindRecordingEnvironment(output, code=141))._search_files( + "*.py", "/narrow", limit=2, offset=1, order=order + ) + + assert result.error is None + assert result.files == ["/narrow/b.py", "/narrow/c.py"] + assert result.truncated is True + + +@pytest.mark.parametrize( + ("order", "output", "error_fragment"), + [ + ("discovery", "/narrow/partial.py\n", "bounded find traversal"), + ("modified", "10 /narrow/partial.py\n", "modification-time"), + ], +) +def test_find_sigpipe_with_fewer_than_fetch_limit_rows_fails_closed( + order, output, error_fragment +): + result = ShellFileOperations(FindRecordingEnvironment(output, code=141))._search_files( + "*.py", "/narrow", limit=2, offset=1, order=order + ) + + assert error_fragment in (result.error or "") + assert result.files == [] + assert result.total_count == 0 + + +@pytest.mark.parametrize( + ("order", "output", "error_fragment"), + [ + ("discovery", "/narrow/partial.py\n", "bounded find traversal"), + ("modified", "/narrow/not-a-timestamp.py\n", "modification-time"), + ], +) +def test_find_hard_error_discards_partial_output(order, output, error_fragment): + env = FindRecordingEnvironment(output, code=2) + + result = ShellFileOperations(env)._search_files( + "*.py", "/narrow", limit=2, offset=0, order=order + ) + + assert error_fragment in (result.error or "") + assert result.files == [] + assert result.total_count == 0 + + +def test_find_timeout_preserves_partial_results_and_limit_reason(): + env = FindRecordingEnvironment(timeout_output("/narrow/partial.py"), code=124) + + result = ShellFileOperations(env)._search_files( + "*.py", "/narrow", limit=2, offset=0, order="discovery" + ) + + assert result.files == ["/narrow/partial.py"] + assert_timed_out(result) + + +def test_find_zero_match_exit_zero_is_success(): + result = ShellFileOperations(FindRecordingEnvironment("", code=0))._search_files( + "*.missing", "/narrow", limit=2, offset=0, order="discovery" + ) + + assert result.error is None + assert result.files == [] + + +def test_local_broad_no_rg_refuses_before_find(monkeypatch, tmp_path): + home = tmp_path / "home" + home.mkdir() + ops = ShellFileOperations(LocalEnvironment(str(home))) + monkeypatch.setattr(file_operations, "_HOME", str(home)) + monkeypatch.setattr(file_operations.os.path, "isfile", lambda path: False) + commands = [] + + def fake_exec(command, **kwargs): + commands.append((command, kwargs)) + if command.startswith("command -v rg"): + return ExecuteResult("", 1) + if command.startswith("command -v find"): + return ExecuteResult("yes\n", 0) + raise AssertionError(f"broad fallback must not execute: {command}") + + monkeypatch.setattr(ops, "_exec", fake_exec) + result = ops._search_files("*.py", str(home), 10, 0, "discovery") + assert "ripgrep" in (result.error or "").lower() + assert not any(command.startswith("find ") for command, _ in commands) diff --git a/tests/tools/test_search_files_cpu_windows.py b/tests/tools/test_search_files_cpu_windows.py new file mode 100644 index 0000000000..652f7255c9 --- /dev/null +++ b/tests/tools/test_search_files_cpu_windows.py @@ -0,0 +1,300 @@ +"""Concurrency admission tests for expensive filename walks.""" + +from concurrent.futures import ThreadPoolExecutor +import threading +import types + +import pytest + +from tools.environments.local import LocalEnvironment +from tools.file_operations import ( + _ACTIVE_FILENAME_SEARCH_ROOTS, + _FILENAME_SEARCH_ADMISSION, + _normalized_filename_search_root, + SearchResult, + ShellFileOperations, +) +from tools.interrupt import set_interrupt + + +class RemoteEnvironment: + is_local = False + cwd = "/workspace" + + def execute(self, command, **kwargs): + raise AssertionError(f"unexpected backend command: {command}") + + +def _operations(env, scan): + operations = ShellFileOperations(env) + operations._resolve_command = lambda command: "/usr/bin/rg" if command == "rg" else None + operations._search_files_rg = types.MethodType(scan, operations) + return operations + + +def test_same_backend_class_and_root_serialize_five_filename_walks(): + entered = threading.Event() + release = threading.Event() + counter_lock = threading.Lock() + active = 0 + maximum_active = 0 + completed = 0 + + def scan(self, pattern, path, limit, offset, order, rg_executable=None): + nonlocal active, maximum_active, completed + with counter_lock: + active += 1 + maximum_active = max(maximum_active, active) + entered.set() + assert release.wait(5) + with counter_lock: + active -= 1 + completed += 1 + return SearchResult(files=[str(path)], total_count=1) + + operations = [_operations(RemoteEnvironment(), scan) for _ in range(5)] + with ThreadPoolExecutor(max_workers=5) as pool: + futures = [ + pool.submit(operation._search_files, "*.py", "/repo", 50, 0) + for operation in operations + ] + assert entered.wait(5) + release.set() + results = [future.result(timeout=5) for future in futures] + + assert all(result.error is None for result in results) + assert completed == 5 + assert maximum_active == 1 + + +def test_different_roots_can_enter_filename_walks_together(): + both_entered = threading.Barrier(2) + + def scan(self, pattern, path, limit, offset, order, rg_executable=None): + both_entered.wait(5) + return SearchResult(files=[str(path)], total_count=1) + + first = _operations(RemoteEnvironment(), scan) + second = _operations(RemoteEnvironment(), scan) + with ThreadPoolExecutor(max_workers=2) as pool: + futures = [ + pool.submit(first._search_files, "*.py", "/one", 50, 0), + pool.submit(second._search_files, "*.py", "/two", 50, 0), + ] + assert [future.result(timeout=5).error for future in futures] == [None, None] + + +def test_different_backend_classes_can_walk_the_same_root_together(): + class OtherRemoteEnvironment(RemoteEnvironment): + pass + + both_entered = threading.Barrier(2) + + def scan(self, pattern, path, limit, offset, order, rg_executable=None): + both_entered.wait(5) + return SearchResult(files=[str(path)], total_count=1) + + first = _operations(RemoteEnvironment(), scan) + second = _operations(OtherRemoteEnvironment(), scan) + with ThreadPoolExecutor(max_workers=2) as pool: + futures = [ + pool.submit(first._search_files, "*.py", "/same", 50, 0), + pool.submit(second._search_files, "*.py", "/same", 50, 0), + ] + assert [future.result(timeout=5).error for future in futures] == [None, None] + + +def test_overlapping_multi_root_sets_are_claimed_atomically(monkeypatch): + first_entered = threading.Event() + release_first = threading.Event() + second_waiting = threading.Event() + lock = threading.Lock() + active = 0 + maximum_active = 0 + + def scan(self, pattern, path, limit, offset, order, rg_executable=None): + nonlocal active, maximum_active + with lock: + active += 1 + maximum_active = max(maximum_active, active) + if path == ["/a", "/b"]: + first_entered.set() + if path == ["/a", "/b"]: + assert release_first.wait(5) + with lock: + active -= 1 + return SearchResult(files=[str(path)], total_count=1) + + first = _operations(RemoteEnvironment(), scan) + second = _operations(RemoteEnvironment(), scan) + original_wait = _FILENAME_SEARCH_ADMISSION.wait + + def observed_wait(timeout=None): + second_waiting.set() + return original_wait(timeout) + + monkeypatch.setattr(_FILENAME_SEARCH_ADMISSION, "wait", observed_wait) + with ThreadPoolExecutor(max_workers=2) as pool: + first_future = pool.submit(first._search_files, "*.py", ["/a", "/b"], 50, 0) + assert first_entered.wait(5) + second_future = pool.submit(second._search_files, "*.py", ["/b", "/c"], 50, 0) + assert second_waiting.wait(5) + release_first.set() + assert first_future.result(timeout=5).error is None + assert second_future.result(timeout=5).error is None + + assert maximum_active == 1 + + +def test_interrupted_waiter_returns_without_dispatch_or_late_dispatch(monkeypatch): + holder_entered = threading.Event() + release_holder = threading.Event() + waiter_waiting = threading.Event() + waiter_tid = [] + dispatches = [] + + def scan(self, pattern, path, limit, offset, order, rg_executable=None): + dispatches.append(threading.get_ident()) + holder_entered.set() + assert release_holder.wait(5) + return SearchResult(files=[str(path)], total_count=1) + + holder = _operations(RemoteEnvironment(), scan) + waiter = _operations(RemoteEnvironment(), scan) + + original_wait = _FILENAME_SEARCH_ADMISSION.wait + + def observed_wait(timeout=None): + waiter_waiting.set() + return original_wait(timeout) + + monkeypatch.setattr(_FILENAME_SEARCH_ADMISSION, "wait", observed_wait) + + def run_waiter(): + waiter_tid.append(threading.get_ident()) + return waiter._search_files("*.py", "/repo", 50, 0) + + with ThreadPoolExecutor(max_workers=2) as pool: + holder_future = pool.submit(holder._search_files, "*.py", "/repo", 50, 0) + assert holder_entered.wait(5) + waiter_future = pool.submit(run_waiter) + assert waiter_waiting.wait(5) + set_interrupt(True, waiter_tid[0]) + try: + interrupted = waiter_future.result(timeout=5) + assert "interrupted" in (interrupted.error or "").lower() + assert len(dispatches) == 1 + release_holder.set() + assert holder_future.result(timeout=5).error is None + assert len(dispatches) == 1 + finally: + set_interrupt(False, waiter_tid[0]) + release_holder.set() + + +def test_interrupt_published_after_final_sample_prevents_filename_dispatch(monkeypatch): + sampled_clear = threading.Event() + resume_acquire = threading.Event() + worker_tid = [] + dispatches = [] + + def scan(self, pattern, path, limit, offset, order, rg_executable=None): + dispatches.append(threading.get_ident()) + return SearchResult(files=[str(path)], total_count=1) + + operations = _operations(RemoteEnvironment(), scan) + original_is_interrupted = __import__( + "tools.interrupt", fromlist=["is_interrupted"] + ).is_interrupted + + def pause_after_clear_sample(): + interrupted = original_is_interrupted() + if not interrupted and threading.get_ident() == worker_tid[0]: + sampled_clear.set() + assert resume_acquire.wait(5) + return interrupted + + monkeypatch.setattr( + "tools.file_operations.tool_interrupt.is_interrupted", + pause_after_clear_sample, + ) + + def run_search(): + worker_tid.append(threading.get_ident()) + return operations._search_files("*.py", "/repo", 50, 0) + + with ThreadPoolExecutor(max_workers=1) as pool: + future = pool.submit(run_search) + assert sampled_clear.wait(5) + set_interrupt(True, worker_tid[0]) + resume_acquire.set() + try: + result = future.result(timeout=5) + finally: + set_interrupt(False, worker_tid[0]) + resume_acquire.set() + + assert "interrupted" in (result.error or "").lower() + assert dispatches == [] + assert _ACTIVE_FILENAME_SEARCH_ROOTS == set() + + +def test_empty_filename_roots_are_rejected_before_engine_resolution(): + def scan(self, pattern, path, limit, offset, order, rg_executable=None): + raise AssertionError("filename engine dispatched") + + operations = _operations(RemoteEnvironment(), scan) + operations._resolve_command = lambda command: (_ for _ in ()).throw( + AssertionError(f"engine resolution attempted: {command}") + ) + + result = operations._search_files("*.py", [], 50, 0) + + assert "at least one search root" in (result.error or "").lower() + assert _ACTIVE_FILENAME_SEARCH_ROOTS == set() + + +@pytest.mark.parametrize("raised", [Exception, KeyboardInterrupt, SystemExit, BaseException]) +def test_admission_releases_after_every_base_exception_path(raised): + attempts = 0 + + def scan(self, pattern, path, limit, offset, order, rg_executable=None): + nonlocal attempts + attempts += 1 + if attempts == 1: + raise raised("engine failed") + return SearchResult(files=[str(path)], total_count=1) + + operations = _operations(RemoteEnvironment(), scan) + with pytest.raises(raised, match="engine failed"): + operations._search_files("*.py", "/repo", 50, 0) + + result = operations._search_files("*.py", "/repo", 50, 0) + assert result.error is None + assert attempts == 2 + assert _ACTIVE_FILENAME_SEARCH_ROOTS == set() + + +def test_remote_roots_are_normalized_lexically_against_backend_cwd(monkeypatch): + env = RemoteEnvironment() + monkeypatch.setattr( + "tools.file_operations.os.path.abspath", + lambda path: (_ for _ in ()).throw(AssertionError("controller resolution used")), + ) + + relative = _normalized_filename_search_root(env, "repo/../repo", "/controller") + absolute = _normalized_filename_search_root(env, "/workspace/repo", "/controller") + + assert relative == "/workspace/repo" + assert absolute == relative + + +@pytest.mark.windows_only +def test_windows_local_root_spellings_share_one_normalized_key(): + env = LocalEnvironment.__new__(LocalEnvironment) + env.cwd = "C:/Repo" + + native = _normalized_filename_search_root(env, r"C:\Repo\src\..", "C:/ignored") + msys = _normalized_filename_search_root(env, "/c/Repo", "C:/ignored") + + assert native == msys diff --git a/tests/tools/test_search_files_engine_selection.py b/tests/tools/test_search_files_engine_selection.py new file mode 100644 index 0000000000..72e9100f84 --- /dev/null +++ b/tests/tools/test_search_files_engine_selection.py @@ -0,0 +1,539 @@ +"""Behavior tests for file-search ordering and ripgrep selection.""" + +import json +import re + +import pytest + +from tools.environments.local import LocalEnvironment +from tools.file_operations import SearchResult, ShellFileOperations +from tools.file_tools import SEARCH_FILES_SCHEMA, _handle_search_files, search_tool + + +class RecordingEnvironment: + is_local = False + cwd = "/repo" + + def __init__(self, *, rg_output="/repo/one.py\n/repo/two.py\n", rg_code=0): + self.commands = [] + self.rg_output = rg_output + self.rg_code = rg_code + + def execute(self, command, **kwargs): + self.commands.append(command) + if command.startswith("test -e "): + return {"output": "exists\n", "returncode": 0} + if command.startswith("command -v rg"): + return {"output": "/opt/Rip Grep/rg\n", "returncode": 0} + if "--version" in command: + return {"output": "ripgrep 14.1.1\n", "returncode": 0} + if "--files" in command: + return {"output": self.rg_output, "returncode": self.rg_code} + return {"output": "", "returncode": 1} + + @property + def rg_commands(self): + return [command for command in self.commands if "--files" in command] + + +def test_schema_exposes_fast_discovery_default_and_exact_modified_opt_in(): + order = SEARCH_FILES_SCHEMA["parameters"]["properties"]["order"] + + assert order["enum"] == ["discovery", "modified"] + assert order["default"] == "discovery" + assert "fast bounded traversal order" in order["description"] + assert "exact global newest-first" in order["description"] + assert "ignored for content" in order["description"] + + +def test_default_file_search_runs_one_bounded_unsorted_rg_command(): + env = RecordingEnvironment() + ops = ShellFileOperations(env) + + result = ops.search("*.py", path="/repo", target="files", limit=1, offset=1) + + assert result.files == ["/repo/two.py"] + assert len(env.rg_commands) == 1 + assert "--sortr" not in env.rg_commands[0] + assert "head -n 3" in env.rg_commands[0] + + +@pytest.mark.parametrize("engine", ["rg", "find"]) +def test_bounded_filename_total_is_serialized_as_a_lower_bound(engine, monkeypatch): + conceptual_files = [f"/repo/file-{index:03}.py" for index in range(200)] + env = RecordingEnvironment() + + def execute(command, **kwargs): + env.commands.append(command) + if command.startswith("test -e "): + return {"output": "exists\n", "returncode": 0} + if command.startswith("command -v rg"): + return { + "output": "/usr/bin/rg\n" if engine == "rg" else "", + "returncode": 0 if engine == "rg" else 1, + } + if "--files" in command or command.startswith("set -o pipefail; find "): + fetch_limit = int(re.search(r"head -n (\d+)", command).group(1)) + return { + "output": "\n".join(conceptual_files[:fetch_limit]) + "\n", + "returncode": 0, + } + return {"output": "", "returncode": 1} + + env.execute = execute + ops = ShellFileOperations(env) + if engine == "find": + monkeypatch.setattr(ops, "_has_command", lambda command: command == "find") + + result = ops.search("*.py", path="/repo", target="files", limit=50) + serialized = result.to_dict() + + assert result.total_count == 51 + assert len(result.files) == 50 + assert serialized["truncated"] is True + assert serialized["total_count_is_lower_bound"] is True + + +def test_modified_file_search_runs_one_exact_order_rg_command(): + env = RecordingEnvironment() + ops = ShellFileOperations(env) + + result = ops.search("*.py", path="/repo", target="files", order="modified") + + assert result.error is None + assert len(env.rg_commands) == 1 + assert "--sortr=modified" in env.rg_commands[0] + + +def test_modified_zero_match_exit_one_is_valid_without_capability_error(): + env = RecordingEnvironment(rg_output="", rg_code=1) + result = ShellFileOperations(env).search( + "*.missing", path="/repo", target="files", order="modified" + ) + assert result.error is None + assert result.files == [] + assert len(env.rg_commands) == 1 + + +@pytest.mark.parametrize( + "version", + [ + "ripgrep 13.0.0\n", + "ripgrep unknown\n", + "ripgrep 14 garbage\n", + "ripgrep 14\n", + "ripgrep 14.1\n", + "ripgrep 14.1.1-\n", + "ripgrep 14.1.1+\n", + "ripgrep 14.1.1-alpha..1\n", + "ripgrep 14.1.1+build..2\n", + "ripgrep 14.1.1-01\n", + "ripgrep 014.1.1\n", + "ripgrep 14.01.1\n", + "ripgrep 14.1.01\n", + ], +) +def test_modified_requires_parseable_ripgrep_14_before_search(version): + env = RecordingEnvironment() + + def execute(command, **kwargs): + env.commands.append(command) + if command.startswith("test -e "): + return {"output": "exists\n", "returncode": 0} + if command.startswith("command -v rg"): + return {"output": "/opt/Rip Grep/rg\n", "returncode": 0} + if "--version" in command: + return {"output": version, "returncode": 0} + raise AssertionError(f"search must not run: {command}") + + env.execute = execute + ops = ShellFileOperations(env) + first = ops.search("*.py", path="/repo", target="files", order="modified") + second = ops.search("*.py", path="/repo", target="files", order="modified") + assert "ripgrep 14" in (first.error or "").lower() + assert second.error == first.error + assert env.rg_commands == [] + assert len([c for c in env.commands if "--version" in c]) == 1 + + +def test_modified_accepts_complete_ripgrep_semver_with_revision_text(): + env = RecordingEnvironment() + original_execute = env.execute + + def execute(command, **kwargs): + if "--version" in command: + env.commands.append(command) + return {"output": "ripgrep 14.1.1 (rev abc123)\n", "returncode": 0} + return original_execute(command, **kwargs) + + env.execute = execute + result = ShellFileOperations(env).search( + "*.py", path="/repo", target="files", order="modified" + ) + + assert result.error is None + assert len(env.rg_commands) == 1 + + +def test_modified_accepts_ripgrep_semver_with_prerelease_and_build_metadata(): + env = RecordingEnvironment() + original_execute = env.execute + + def execute(command, **kwargs): + if "--version" in command: + env.commands.append(command) + return { + "output": "ripgrep 14.1.1-alpha.1+build.2\n", + "returncode": 0, + } + return original_execute(command, **kwargs) + + env.execute = execute + result = ShellFileOperations(env).search( + "*.py", path="/repo", target="files", order="modified" + ) + + assert result.error is None + assert len(env.rg_commands) == 1 + + +def test_empty_discovery_output_is_zero_matches_without_retry(): + env = RecordingEnvironment(rg_output="", rg_code=0) + ops = ShellFileOperations(env) + + result = ops.search("*.missing", path="/repo", target="files") + + assert result.error is None + assert result.files == [] + assert result.total_count == 0 + assert len(env.rg_commands) == 1 + + +def test_modified_capability_failure_is_actionable_and_not_downgraded(): + env = RecordingEnvironment(rg_output="", rg_code=2) + ops = ShellFileOperations(env) + + result = ops.search("*.py", path="/repo", target="files", order="modified") + + assert len(env.rg_commands) == 1 + assert result.error is not None + assert "exact modification-time order" in result.error.lower() + assert "ripgrep" in result.error + + +@pytest.mark.parametrize("order", ["discovery", "modified"]) +def test_rg_partial_output_with_error_exit_fails_closed(order): + env = RecordingEnvironment(rg_output="/repo/partial.py\n", rg_code=2) + + result = ShellFileOperations(env).search( + "*.py", path="/repo", target="files", order=order + ) + + assert result.error is not None + assert result.files == [] + assert len(env.rg_commands) == 1 + + +@pytest.mark.parametrize("order", ["discovery", "modified"]) +def test_rg_sigpipe_is_benign_only_after_fetch_limit_paths(order): + output = "".join(f"/repo/{name}.py\n" for name in ("a", "b", "c", "d")) + env = RecordingEnvironment(rg_output=output, rg_code=141) + + result = ShellFileOperations(env).search( + "*.py", path="/repo", target="files", limit=2, offset=1, order=order + ) + + assert result.error is None + assert result.files == ["/repo/b.py", "/repo/c.py"] + assert result.truncated is True + + +@pytest.mark.parametrize("order", ["discovery", "modified"]) +def test_rg_sigpipe_with_fewer_than_fetch_limit_paths_fails_closed(order): + env = RecordingEnvironment(rg_output="/repo/partial.py\n", rg_code=141) + + result = ShellFileOperations(env).search( + "*.py", path="/repo", target="files", limit=2, offset=1, order=order + ) + + assert result.error is not None + assert result.files == [] + assert result.total_count == 0 + + +def test_invalid_direct_file_order_returns_structured_error(): + env = RecordingEnvironment() + ops = ShellFileOperations(env) + + result = ops.search("*.py", path="/repo", target="files", order="random") + + assert isinstance(result, SearchResult) + assert result.error == "Invalid file search order 'random'; expected 'discovery' or 'modified'." + assert env.rg_commands == [] + + +def test_handler_forwards_modified_order(monkeypatch): + captured = {} + + def fake_search_tool(**kwargs): + captured.update(kwargs) + return "{}" + + monkeypatch.setattr("tools.file_tools.search_tool", fake_search_tool) + + _handle_search_files({"pattern": "*.py", "target": "files", "order": "modified"}) + + assert captured["order"] == "modified" + + +def test_repeated_search_key_distinguishes_order(monkeypatch): + class StubOperations: + def search(self, **kwargs): + return SearchResult() + + monkeypatch.setattr("tools.file_tools._get_file_ops", lambda task_id: StubOperations()) + task_id = "engine-order-key" + for _ in range(3): + assert "BLOCKED" not in json.loads( + search_tool("*.py", target="files", order="discovery", task_id=task_id) + ).get("error", "") + + changed = json.loads( + search_tool("*.py", target="files", order="modified", task_id=task_id) + ) + + assert "BLOCKED" not in changed.get("error", "") + + +class RipgrepInvocationEnvironment(RecordingEnvironment): + def execute(self, command, **kwargs): + self.commands.append(command) + if command.startswith("test -e "): + return {"output": "exists\n", "returncode": 0} + if command.startswith("command -v rg"): + return {"output": "/opt/Rip Grep/rg\n", "returncode": 0} + if "--version" in command: + return {"output": "ripgrep 14.1.1\n", "returncode": 0} + if "--files" in command: + return {"output": "/repo/a.py\n", "returncode": 0} + if "--line-number" in command: + return {"output": "/repo/a.py:1:needle\n", "returncode": 0} + if "--count-matches" in command: + return {"output": "", "returncode": 1} + return {"output": "", "returncode": 1} + + +def test_resolved_executable_with_spaces_is_used_by_every_rg_invocation(): + env = RipgrepInvocationEnvironment() + ops = ShellFileOperations(env) + + assert ops.search("*.py", path="/repo", target="files").files + assert ops.search("needle", path="/repo", target="content").matches + assert ops._zero_match_probe("absent", "/repo", None) is None + + invocations = [ + command for command in env.commands + if any(flag in command for flag in ("--files", "--line-number", "--count-matches")) + ] + assert invocations + assert all("'/opt/Rip Grep/rg'" in command for command in invocations) + assert all(not re.search(r"(?:^|[; ])rg\s", command) for command in invocations) + assert len([c for c in env.commands if c.startswith("command -v rg")]) == 1 + + +def test_non_rg_command_cache_keeps_cached_misses_and_bool_values(): + env = RecordingEnvironment() + ops = ShellFileOperations(env) + assert ops._has_command("find") is False + assert ops._has_command("find") is False + assert ops._command_cache == {"find": False} + assert len([c for c in env.commands if c.startswith("command -v find")]) == 1 + + +@pytest.mark.windows_only +def test_off_path_windows_rg_miss_is_reprobed_then_success_is_cached( + tmp_path, monkeypatch +): + local_app_data = tmp_path / "Local Data" + candidate = local_app_data / "Microsoft" / "WinGet" / "Links" / "rg.exe" + monkeypatch.setenv("LOCALAPPDATA", str(local_app_data)) + monkeypatch.setenv("USERPROFILE", str(tmp_path / "User Profile")) + monkeypatch.delenv("SCOOP", raising=False) + ops = ShellFileOperations(LocalEnvironment(str(tmp_path))) + probes = [] + + def command_v_miss(command, **kwargs): + probes.append(command) + from tools.file_operations import ExecuteResult + return ExecuteResult(stdout="", exit_code=1) + + monkeypatch.setattr(ops, "_exec", command_v_miss) + + assert ops._resolve_command("rg") is None + candidate.parent.mkdir(parents=True) + candidate.write_text("") + expected = str(candidate).replace("\\", "/") + assert ops._resolve_command("rg") == expected + assert ops._resolve_command("rg") == expected + assert probes == ["command -v rg 2>/dev/null", "command -v rg 2>/dev/null"] + + +def test_remote_resolution_never_probes_controller_host_paths(tmp_path, monkeypatch): + monkeypatch.setenv("LOCALAPPDATA", str(tmp_path / "controller-local")) + monkeypatch.setenv("USERPROFILE", str(tmp_path / "controller-user")) + env = RecordingEnvironment() + + def miss(command, **kwargs): + env.commands.append(command) + return {"output": "", "returncode": 1} + + env.execute = miss + ops = ShellFileOperations(env) + + assert ops._resolve_command("rg") is None + assert env.commands == ["command -v rg 2>/dev/null"] + assert str(tmp_path) not in env.commands[0] + + +@pytest.mark.windows_only +def test_remote_msys_shaped_executable_is_not_rewritten_as_controller_path(): + env = RecordingEnvironment() + + def execute(command, **kwargs): + env.commands.append(command) + if command.startswith("test -e "): + return {"output": "exists\n", "returncode": 0} + if command.startswith("command -v rg"): + return {"output": "/c/remote-tools/rg\n", "returncode": 0} + if "--files" in command: + return {"output": "/repo/a.py\n", "returncode": 0} + return {"output": "", "returncode": 1} + + env.execute = execute + + result = ShellFileOperations(env).search("*.py", path="/repo", target="files") + + assert result.files + assert "'/c/remote-tools/rg' --files" in env.rg_commands[0] + assert "C:/remote-tools/rg" not in env.rg_commands[0] + + +@pytest.mark.windows_only +def test_every_windows_drive_root_is_broad_even_when_home_is_on_another_drive( + tmp_path, monkeypatch +): + import tools.file_operations as file_operations + + monkeypatch.setattr(file_operations, "_HOME", "C:/Users/alice") + ops = ShellFileOperations(LocalEnvironment(str(tmp_path))) + + assert ops._is_broad_local_search_root("D:/") is True + assert ops._is_broad_local_search_root("D:/repo") is False + + +def test_modified_multi_path_search_preserves_exact_order_request(): + env = RecordingEnvironment() + + def execute(command, **kwargs): + env.commands.append(command) + if command.startswith("test -e "): + if "'/one /two'" in command: + return {"output": "not_found\n", "returncode": 0} + return {"output": "exists\n", "returncode": 0} + if command.startswith("command -v rg"): + return {"output": "/usr/bin/rg\n", "returncode": 0} + if "--version" in command: + return {"output": "ripgrep 14.1.1\n", "returncode": 0} + if "--files" in command: + return {"output": "/two/new.py\n/one/old.py\n", "returncode": 0} + return {"output": "", "returncode": 1} + + env.execute = execute + result = ShellFileOperations(env).search( + "*.py", path="/one /two", target="files", order="modified" + ) + + assert result.files == ["/two/new.py", "/one/old.py"] + assert len(env.rg_commands) == 1 + assert "--sortr=modified" in env.rg_commands[0] + assert "'/one'" in env.rg_commands[0] + assert "'/two'" in env.rg_commands[0] + + +def test_comma_delimited_file_roots_preserve_internal_spaces_in_one_search(): + env = RecordingEnvironment(rg_output="C:/root one/a.py\nC:/root two/b.py\n") + combined = "C:/root one, C:/root two" + + path_checks = 0 + + def execute(command, **kwargs): + nonlocal path_checks + env.commands.append(command) + if command.startswith("test -e "): + path_checks += 1 + output = "not_found\n" if path_checks == 1 else "exists\n" + return {"output": output, "returncode": 0} + if command.startswith("command -v rg"): + return {"output": "/usr/bin/rg\n", "returncode": 0} + if "--files" in command: + return {"output": env.rg_output, "returncode": 0} + return {"output": "", "returncode": 1} + + env.execute = execute + result = ShellFileOperations(env).search( + "*.py", path=combined, target="files" + ) + + assert result.error is None + assert result.files == ["C:/root one/a.py", "C:/root two/b.py"] + assert len(env.rg_commands) == 1 + assert "'C:/root one' 'C:/root two'" in env.rg_commands[0] + assert "path contained 2 entries" in (result.warning or "") + + +def test_multi_path_modified_capability_error_propagates(): + env = RecordingEnvironment() + + def execute(command, **kwargs): + env.commands.append(command) + if command.startswith("test -e "): + output = "not_found\n" if "'/one /two'" in command else "exists\n" + return {"output": output, "returncode": 0} + if command.startswith("command -v rg"): + return {"output": "/usr/bin/rg\n", "returncode": 0} + if "--version" in command: + return {"output": "ripgrep 13.0.0\n", "returncode": 0} + raise AssertionError(command) + + env.execute = execute + result = ShellFileOperations(env).search( + "*.py", path="/one /two", target="files", order="modified" + ) + assert "ripgrep 14" in (result.error or "").lower() + assert env.rg_commands == [] + + +def test_modified_timeout_preserves_partial_results_and_limit_reason(): + env = RecordingEnvironment( + rg_output="/repo/partial.py\n[Command timed out after 60s]\n", + rg_code=124, + ) + + result = ShellFileOperations(env).search( + "*.py", path="/repo", target="files", order="modified" + ) + + assert result.files == ["/repo/partial.py"] + assert result.truncated is True + assert result.limit_reason == "search_timeout" + + +def test_order_is_ignored_for_content_search(): + env = RipgrepInvocationEnvironment() + + result = ShellFileOperations(env).search( + "needle", path="/repo", target="content", order="not-a-file-order" + ) + + assert result.error is None + assert result.matches diff --git a/tests/tools/test_search_giant_line_containment.py b/tests/tools/test_search_giant_line_containment.py new file mode 100644 index 0000000000..a31c6509fa --- /dev/null +++ b/tests/tools/test_search_giant_line_containment.py @@ -0,0 +1,100 @@ +"""Giant single-line file containment in content search (cline/cline#13525 port). + +A match inside a serialized dump (multi-MB single-line JSON, minified +bundle) used to make rg/grep emit the ENTIRE matched line into stdout: +``head -n`` counts lines, so a 40MB match line crossed the transport +untruncated and was buffered whole into Python before the per-match +[:500] clamp ran (measured 42MB transport / ~180MB peak alloc for one +match). The fix bounds lines at the search-engine layer: rg gets +``--max-columns 2000 --max-columns-preview``; the grep fallbacks pipe +through ``cut -c1-2000``. + +These tests run the REAL pipelines via bash (no mocked stdout) so the +flag/pipe behavior of the installed rg/grep is what's exercised. +""" + +import os +import shutil +import subprocess + +import pytest + +from tools.file_operations import ShellFileOperations + +# Big enough to prove containment, small enough to keep the test fast. +GIANT = 5 * 1024 * 1024 # 5MB single line +# Generous ceiling: pre-fix stdout for one giant match is >= GIANT bytes. +STDOUT_CEILING = 1 * 1024 * 1024 + + +class RecordingEnv: + """Local bash executor that records the largest stdout it returned.""" + + def __init__(self, cwd): + self.cwd = cwd + self.max_stdout = 0 + + def execute(self, command, timeout=60, **kwargs): + proc = subprocess.run( + ["bash", "-c", command], + capture_output=True, text=True, errors="replace", + timeout=timeout + 30, + ) + out = proc.stdout + (proc.stderr or "") + self.max_stdout = max(self.max_stdout, len(out)) + return {"output": out, "returncode": proc.returncode} + + +@pytest.fixture() +def giant_dir(tmp_path): + (tmp_path / "trace.json").write_text( + '{"needle": "' + "x" * GIANT + '"}', encoding="utf-8" + ) + (tmp_path / "small.py").write_text("needle = 1\n", encoding="utf-8") + return tmp_path + + +def _ops(giant_dir, engine): + env = RecordingEnv(str(giant_dir)) + ops = ShellFileOperations(env) + ops._has_command = lambda cmd: cmd == engine + return ops, env + + +@pytest.mark.parametrize("engine", ["rg", "grep"]) +def test_giant_single_line_match_is_bounded(giant_dir, engine): + if shutil.which(engine) is None: + pytest.skip(f"{engine} not installed") + ops, env = _ops(giant_dir, engine) + + result = ops.search("needle", path=str(giant_dir), target="content") + + assert result.error is None + paths = {os.path.basename(m.path) for m in result.matches} + # The giant-file match must still be REPORTED (preview, not omission)... + assert paths == {"trace.json", "small.py"} + assert all(len(m.content) <= 500 for m in result.matches) + # ...but its full line must never have crossed the transport. + assert env.max_stdout < STDOUT_CEILING, ( + f"{engine} pipeline returned {env.max_stdout} bytes of stdout — " + "giant matched line was not truncated at the engine layer" + ) + + +@pytest.mark.parametrize("engine", ["rg", "grep"]) +@pytest.mark.parametrize("output_mode", ["files_only", "count"]) +def test_line_cap_skipped_for_path_and_count_modes(giant_dir, engine, output_mode): + """files_only/count lines are paths/counts — never giant, never cut.""" + if shutil.which(engine) is None: + pytest.skip(f"{engine} not installed") + ops, env = _ops(giant_dir, engine) + + result = ops.search("needle", path=str(giant_dir), target="content", + output_mode=output_mode) + + assert result.error is None + if output_mode == "files_only": + assert {os.path.basename(f) for f in result.files} == {"trace.json", "small.py"} + else: + assert {os.path.basename(k) for k in result.counts} == {"trace.json", "small.py"} + assert env.max_stdout < STDOUT_CEILING diff --git a/tests/tools/test_search_zero_match_and_multipath.py b/tests/tools/test_search_zero_match_and_multipath.py index a001447051..f3d9fbbed8 100644 --- a/tests/tools/test_search_zero_match_and_multipath.py +++ b/tests/tools/test_search_zero_match_and_multipath.py @@ -58,6 +58,86 @@ class TestZeroMatchProbe: # Same class as the casing probe: the path must be in the hint. assert "conf.cfg" in r.get("warning", "") + def test_hidden_probe_prunes_dependency_trees_and_keeps_local_ignored(self, proj, monkeypatch): + d = proj / "proj" + dependency = d / "node_modules" / "package" / ".hidden" + dependency.mkdir(parents=True) + dependency_file = dependency / "dependency.js" + dependency_file.write_text("BOUNDED_HIDDEN_TOKEN = true\n") + local = d / ".project-local" + local.mkdir() + local_file = local / "settings.cfg" + local_file.write_text("BOUNDED_HIDDEN_TOKEN = true\n") + (d / ".gitignore").write_text("node_modules/\n.project-local/\n") + + # Drive the public search seam while recording the commands that the + # zero-match probe actually executes. The real rg calls still run. + from tools.file_tools import _get_file_ops + + task_id = "t-zm-pruned-hidden" + ops = _get_file_ops(task_id=task_id) + commands = [] + real_exec = ops._exec + + def recording_exec(command, *args, **kwargs): + commands.append(command) + return real_exec(command, *args, **kwargs) + + monkeypatch.setattr(ops, "_exec", recording_exec) + r = json.loads(search_tool("BOUNDED_HIDDEN_TOKEN", path=str(d), task_id=task_id)) + warning = r.get("warning", "") + + assert r["total_count"] == 0 + assert "hidden or gitignored" in warning + assert local_file.name in warning + assert dependency_file.name not in warning + + hidden_probe_commands = [ + command for command in commands + if "--hidden" in command and "--no-ignore" in command + ] + assert len(hidden_probe_commands) == 1 + hidden_probe = hidden_probe_commands[0] + assert "--glob" in hidden_probe + assert "'!node_modules/**'" in hidden_probe + assert "'!**/node_modules/**'" in hidden_probe + + def test_hidden_probe_prunes_explicit_dependency_root(self, proj, monkeypatch): + d = proj / "proj" + dependency = d / "node_modules" / "package" / ".hidden" + dependency.mkdir(parents=True) + (dependency / "dependency.js").write_text("EXPLICIT_ROOT_TOKEN = true\n") + (d / ".gitignore").write_text("node_modules/\n") + + from tools.file_tools import _get_file_ops + + task_id = "t-zm-explicit-pruned-root" + ops = _get_file_ops(task_id=task_id) + commands = [] + real_exec = ops._exec + + def recording_exec(command, *args, **kwargs): + commands.append(command) + return real_exec(command, *args, **kwargs) + + monkeypatch.setattr(ops, "_exec", recording_exec) + r = json.loads(search_tool( + "EXPLICIT_ROOT_TOKEN", + path=str(d / "node_modules"), + task_id=task_id, + )) + + assert r["total_count"] == 0 + assert "warning" not in r + hidden_probe_commands = [ + command for command in commands + if "--hidden" in command and "--no-ignore" in command + ] + assert len(hidden_probe_commands) == 1 + hidden_probe = hidden_probe_commands[0] + assert "'!node_modules/**'" in hidden_probe + assert "'!**/node_modules/**'" in hidden_probe + def test_probe_path_list_is_capped(self, proj): d = proj / "proj" for i in range(8): diff --git a/tests/tools/test_session_search.py b/tests/tools/test_session_search.py index de9f1300e2..6e807962d2 100644 --- a/tests/tools/test_session_search.py +++ b/tests/tools/test_session_search.py @@ -113,11 +113,41 @@ class TestFormatTimestamp: # ========================================================================= class TestBrowseShape: + def test_browse_uses_bounded_recent_path(self): + class _DB: + rich_called = False + bounded_kwargs = None + + def list_recent_sessions_bounded(self, **kwargs): + self.bounded_kwargs = kwargs + return [] + + def list_sessions_rich(self, **_kwargs): + self.rich_called = True + raise AssertionError("unbounded rich listing must not be used") + + db = _DB() + result = json.loads(session_search(db=db)) + + assert result["success"] is True + assert db.rich_called is False + assert db.bounded_kwargs["timeout_seconds"] == 3.0 + + def test_browse_fails_closed_without_bounded_database_capability(self): + class _LegacyDB: + def list_sessions_rich(self, **_kwargs): + raise AssertionError("known-unbounded fallback must not be called") + + result = json.loads(session_search(db=_LegacyDB())) + + assert result["success"] is False + assert "does not support bounded recent-session browse" in result["error"] + def test_lazy_database_is_released_after_search(self, monkeypatch): class _DB: released = 0 - def list_sessions_rich(self, **_kwargs): + def list_recent_sessions_bounded(self, **_kwargs): return [] db = _DB() @@ -139,7 +169,7 @@ class TestBrowseShape: def __init__(self): self.closed = 0 - def list_sessions_rich(self, **_kwargs): + def list_recent_sessions_bounded(self, **_kwargs): return [] def close(self): diff --git a/tests/tools/test_terminal_compound_background.py b/tests/tools/test_terminal_compound_background.py index d0f1762fe9..beaa1d95f4 100644 --- a/tests/tools/test_terminal_compound_background.py +++ b/tests/tools/test_terminal_compound_background.py @@ -12,6 +12,10 @@ The rewriter fixes this by wrapping the tail in a brace group — the current shell. No subshell fork, no wait. """ +import shutil +import subprocess + +import pytest from tools.terminal_tool import _rewrite_compound_background as rewrite @@ -100,3 +104,105 @@ class TestEdgeCases: def test_tabs_between_tokens(self): assert rewrite("A\t&&\tB\t&") == "A\t&&\t{ B\t& }" + + +class TestTrailingStatementSeparator: + """A statement after the backgrounded compound on the SAME line. + + In ``A && B & C`` the trailing ``&`` is both the background operator and + the separator between the compound and ``C``. The rewrite consumes that + ``&`` into the brace group; without restoring a separator the result is + ``A && { B & } C`` — a bash syntax error (a brace group must be terminated + by ``;``, ``&``, ``|``, a newline, or ``)``/``}`` before the next command). + That mangles a valid command into one that fails entirely. + """ + + def test_trailing_command_gets_separator(self): + assert rewrite("echo hi && sleep 5 & echo done") == ( + "echo hi && { sleep 5 & } ; echo done" + ) + + def test_trailing_chain_gets_separator(self): + assert rewrite("a && b & c && d") == "a && { b & } ; c && d" + + def test_redirect_then_trailing_command(self): + assert rewrite("echo hi && sleep 5 &>/dev/null & echo done") == ( + "echo hi && { sleep 5 &>/dev/null & } ; echo done" + ) + + def test_existing_semicolon_separator_untouched(self): + # An explicit `;` already separates the group; don't add a second one. + assert rewrite("a && b &; c") == "a && { b & }; c" + + def test_newline_separator_untouched(self): + # A newline already terminates the brace group — no `;` needed. + assert rewrite("a && b &\necho next") == "a && { b & }\necho next" + + def test_pipe_after_group_untouched(self): + # `{ ...; } | cmd` is valid; the pipe is its own terminator. + assert rewrite("a && b & | cat") == "a && { b & } | cat" + + def test_redirect_prefix_on_trailing_command_gets_separator(self): + # `&>` after the group is a redirect for the NEXT command, not a + # terminator: `{ b & } &>/dev/null c` is a syntax error. + assert rewrite("a && b & &>/dev/null c") == "a && { b & } ; &>/dev/null c" + + def test_case_arm_terminator_untouched(self): + # `;;` already terminates the arm; adding `;` would leave an empty + # command between `;` and `;;`, which bash rejects. + assert rewrite("case $x in p) b && c & ;; esac") == "case $x in p) b && { c & } ;; esac" + + def test_separator_is_idempotent(self): + once = rewrite("echo hi && sleep 5 & echo done") + assert rewrite(once) == once + + def test_second_background_then_trailing(self): + assert rewrite("echo a && sleep 5 & echo b & echo c") == ( + "echo a && { sleep 5 & } ; echo b & echo c" + ) + + +@pytest.mark.skipif(shutil.which("bash") is None, reason="bash not available") +class TestRewriteIsValidBash: + """The rewrite must always produce syntactically valid bash. + + This is the crux of the trailing-statement bug: a mangled command fails + with a confusing syntax error and neither half runs. ``bash -n`` parses + without executing, so it catches the corruption directly. + """ + + @pytest.mark.parametrize( + "command", + [ + "echo hi && sleep 5 & echo done", + "a && b & c && d", + "echo hi && sleep 5 &>/dev/null & echo done", + "echo a && sleep 5 & echo b & echo c", + "A && B &", + "A && B &; C", + "A && B &\nC", + "cd /tmp && python3 -m http.server 0 &>/dev/null & curl localhost", + "a && b & &>/dev/null c", + "case $x in p) b && c & ;; esac", + "A && B & echo x\nC && D & echo y && E & echo z", + ], + ) + def test_rewrite_parses(self, command): + rewritten = rewrite(command) + result = subprocess.run( + ["bash", "-n", "-c", rewritten], + capture_output=True, + text=True, + ) + assert result.returncode == 0, ( + f"rewrite produced invalid bash: {rewritten!r}\n{result.stderr}" + ) + + def test_trailing_statement_actually_runs(self): + # End-to-end: the command after the backgrounded compound must run. + rewritten = rewrite("echo first && true & echo SECOND_RAN") + result = subprocess.run( + ["bash", "-c", rewritten], capture_output=True, text=True + ) + assert result.returncode == 0 + assert "SECOND_RAN" in result.stdout diff --git a/tests/tools/test_vision_native_fast_path.py b/tests/tools/test_vision_native_fast_path.py index 5237ae665f..82b7cb57ce 100644 --- a/tests/tools/test_vision_native_fast_path.py +++ b/tests/tools/test_vision_native_fast_path.py @@ -90,6 +90,30 @@ class TestSupportsMediaInToolResults: assert _supports_media_in_tool_results("", "anything") is False assert _supports_media_in_tool_results(None, "anything") is False # type: ignore[arg-type] + def test_profile_tool_message_veto_overrides_supports_vision(self): + """supports_vision_tool_messages=False is a hard veto even when the + profile declares supports_vision=True (xiaomi/MiMo 400s on list-type + tool-result content, #89981).""" + assert _supports_media_in_tool_results("xiaomi", "mimo-v2.5") is False + + def test_profile_veto_applies_even_when_vision_capable_lookup_agrees(self): + """A capability source marking the model vision-capable must not + re-open the native fast path for a provider that rejects it.""" + from tools.vision_tools import _should_use_native_vision_fast_path + from agent.auxiliary_client import set_runtime_main, clear_runtime_main + from agent import image_routing + + set_runtime_main("xiaomi", "mimo-v2.5") + try: + with patch.object( + image_routing, "decide_image_input_mode", return_value="native" + ), patch.object( + image_routing, "_lookup_supports_vision", return_value=True + ): + assert _should_use_native_vision_fast_path() is False + finally: + clear_runtime_main() + # ─── _build_native_vision_tool_result ──────────────────────────────────────── diff --git a/tests/tui_gateway/test_bot_relay_methods.py b/tests/tui_gateway/test_bot_relay_methods.py index 85b18d971a..8b551892f9 100644 --- a/tests/tui_gateway/test_bot_relay_methods.py +++ b/tests/tui_gateway/test_bot_relay_methods.py @@ -116,7 +116,18 @@ def test_deliver_lands_in_live_bot_chat_instead_of_subprocess(home, monkeypatch) """ spawned = [] submitted = [] - monkeypatch.setattr("subprocess.run", lambda *a, **k: spawned.append(a) or None) + + class _Proc: + returncode, stdout, stderr = 0, "pong", "" + + def _fake_run(argv, *a, **k): + # The server module's import-time update prefetch runs `git ...` on a + # daemon thread; only the relay's `hermes` CLI spawn is under test. + if argv and argv[0] != "git": + spawned.append(argv) + return _Proc() + + monkeypatch.setattr("subprocess.run", _fake_run) monkeypatch.setitem( srv._methods, "prompt.submit", lambda rid, p: submitted.append(p) or srv._ok(rid, {"status": "streaming"}) ) @@ -137,10 +148,6 @@ def test_deliver_lands_in_live_bot_chat_instead_of_subprocess(home, monkeypatch) srv._sessions["live-ops"]["pending_title"] = "Scratch" submitted.clear() - class _Proc: - returncode, stdout, stderr = 0, "pong", "" - - monkeypatch.setattr("subprocess.run", lambda *a, **k: spawned.append(a) or _Proc()) out = _result(srv._methods["bot_relay.deliver"](2, {"profile": "ops", "message": "ping"})) assert out["reply"] == "pong" and spawned and not submitted diff --git a/tests/tui_gateway/test_cross_process_orphan_ownership.py b/tests/tui_gateway/test_cross_process_orphan_ownership.py index 481b837c07..1a9de04ff9 100644 --- a/tests/tui_gateway/test_cross_process_orphan_ownership.py +++ b/tests/tui_gateway/test_cross_process_orphan_ownership.py @@ -270,6 +270,84 @@ def test_automatic_cleanup_preserves_corrupt_registry_without_overwrite( assert state_path.read_text(encoding="utf-8") == corrupt +def test_own_live_lease_ids_reports_live_owners_and_skips_the_excluded( + monkeypatch: pytest.MonkeyPatch, +) -> None: + class _Lease: + def __init__(self, lease_id: str) -> None: + self.lease_id = lease_id + + first = _Lease("first") + second = _Lease("second") + monkeypatch.setattr( + server, + "_sessions", + { + "one": {"active_session_lease": first}, + "two": {"active_session_lease": second}, + "three": {"active_session_lease": None}, + }, + ) + + assert server._own_live_lease_ids() == {"first", "second"} + assert server._own_live_lease_ids(exclude=first) == {"second"} + + +def test_automatic_cleanup_reclaims_own_orphan_lease_not_treated_as_sibling( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + profile_home = tmp_path / "profile-home" + session_id = "own-orphan-session" + owner_lease, message = server._claim_active_session_slot( + session_id, + live_session_id="vanished-runtime", + surface="desktop", + profile_home=profile_home, + ) + assert owner_lease is not None and message is None + # The owner vanished minutes ago; a lease written seconds ago is still + # inside the self-orphan grace window and must be left alone. + monkeypatch.setattr( + "hermes_cli.active_sessions._SELF_ORPHAN_GRACE_SECONDS", 0.0 + ) + ended: list[tuple[str, str]] = [] + + class _FakeDB: + def get_session(self, target: str) -> dict[str, str]: + return {"id": target, "source": "desktop"} + + def end_session(self, target: str, reason: str) -> None: + ended.append((target, reason)) + + @contextlib.contextmanager + def _profile_db(_session: dict): + yield _FakeDB() + + monkeypatch.setattr(server, "_sessions", {}) + monkeypatch.setattr(server, "_session_db", _profile_db) + monkeypatch.setattr( + server, "_notify_session_boundary", lambda *args, **kwargs: None + ) + monkeypatch.setattr( + "tools.async_delegation.interrupt_for_session", lambda *args, **kwargs: None + ) + session = { + "active_session_lease": None, + "agent": None, + "history": [], + "history_lock": threading.Lock(), + "profile_home": str(profile_home), + "session_key": session_id, + "slash_worker": None, + "source": "desktop", + } + + server._finalize_session(session, end_reason="ws_orphan_reap") + + assert ended == [(session_id, "ws_orphan_reap")] + assert active_session_registry_snapshot(registry_home=profile_home) == [] + + def test_liveness_guard_serializes_cross_process_acquire(tmp_path: Path) -> None: home = tmp_path / "guard-home" waiting_file = tmp_path / "child-waiting" diff --git a/tests/tui_gateway/test_ws_surrogate_send.py b/tests/tui_gateway/test_ws_surrogate_send.py new file mode 100644 index 0000000000..14006f38b8 --- /dev/null +++ b/tests/tui_gateway/test_ws_surrogate_send.py @@ -0,0 +1,63 @@ +"""Lone UTF-16 surrogates must not tear down the Desktop WebSocket (#97288).""" + +from __future__ import annotations + +import asyncio + +from tui_gateway.ws import WSTransport, _sanitize_ws_text + + +LONE_SURROGATE = "\ud83d" + + +def test_sanitize_ws_text_makes_utf8_encodable() -> None: + dirty = f"gateway.ready {LONE_SURROGATE} payload" + out = _sanitize_ws_text(dirty) + out.encode("utf-8") + assert LONE_SURROGATE not in out + + +def test_sanitize_ws_text_leaves_valid_text_unchanged() -> None: + clean = '{"type":"gateway.ready","ok":true}' + assert _sanitize_ws_text(clean) is clean or _sanitize_ws_text(clean) == clean + + +class _FakeWS: + def __init__(self) -> None: + self.sent: list[str] = [] + self.raise_on: str | None = None + + async def send_text(self, line: str) -> None: + line.encode("utf-8") + if self.raise_on is not None and self.raise_on in line: + raise UnicodeEncodeError("utf-8", line, 0, 1, "surrogates not allowed") + self.sent.append(line) + + +def test_safe_send_sanitizes_surrogate_and_keeps_connection() -> None: + async def _run() -> None: + loop = asyncio.get_running_loop() + ws = _FakeWS() + transport = WSTransport(ws, loop, peer="127.0.0.1:1") + dirty = f'{{"type":"gateway.ready","x":"{LONE_SURROGATE}"}}' + await transport._safe_send_many(["first", dirty, "third"]) + assert transport.closed is False + assert ws.sent[0] == "first" + assert ws.sent[-1] == "third" + assert LONE_SURROGATE not in "".join(ws.sent) + assert len(ws.sent) == 3 + + asyncio.run(_run()) + + +def test_unicode_encode_error_does_not_close_socket() -> None: + async def _run() -> None: + loop = asyncio.get_running_loop() + ws = _FakeWS() + ws.raise_on = "BOOM" + transport = WSTransport(ws, loop, peer="127.0.0.1:1") + await transport._safe_send_many(["ok-a", "BOOM-frame", "ok-b"]) + assert transport.closed is False + assert ws.sent == ["ok-a", "ok-b"] + + asyncio.run(_run()) diff --git a/tools/browser_lightpanda.py b/tools/browser_lightpanda.py index af44f29ab9..e898b24adc 100644 --- a/tools/browser_lightpanda.py +++ b/tools/browser_lightpanda.py @@ -6,6 +6,7 @@ built-in ``browser_*`` tools keep going through ``agent-browser --engine lightpa ``tools.browser_tool`` owns the session cache, inactivity reaper and atexit sweep. """ +import functools import json import logging import os @@ -88,6 +89,47 @@ def _state_dir() -> Path: return path +def _http_cache_dir() -> Path: + """Filesystem HTTP cache shared by every Lightpanda this Hermes spawns. + + Shared rather than per-session so a cached asset survives session churn. + Lightpanda holds it in sqlite (WAL) with a best-effort write path, and + ``--http-cache-entry-limit`` (upstream default 1000, not passed here) + bounds it without Hermes managing eviction. + """ + path = _state_dir() / "http-cache" + path.mkdir(parents=True, exist_ok=True) + return path + + +_HTTP_CACHE_FLAG = "--http-cache-dir" + + +@functools.lru_cache(maxsize=1) +def _binary_supports_http_cache(binary: str) -> bool: + """True if ``lightpanda serve`` accepts ``--http-cache-dir``. + + The flag landed upstream in 0.3.x; older binaries fatally reject it + ("unknown argument"), which would break every launch. Probing ``help`` + output keeps working across future flag additions without parsing + versions, and the lru_cache keeps it once per binary per process. + """ + try: + proc = subprocess.run( + [binary, "help"], + capture_output=True, text=True, timeout=3.0, + stdin=subprocess.DEVNULL, + ) + return _HTTP_CACHE_FLAG in ((proc.stdout or "") + (proc.stderr or "")) + except Exception as e: + logger.debug("lightpanda http-cache probe failed (%s); assuming no", e) + return False + + +def _record_path(session_name: str) -> Path: + return _state_dir() / f"{session_name}.json" + + def _browser_env() -> dict: try: from tools.browser_tool import _build_browser_env @@ -163,7 +205,11 @@ def launch_lightpanda(session_name: str, *, block_private_networks: bool = False f"or ~/.local/bin. {LIGHTPANDA_INSTALL_HINT}, or set browser.engine to auto.") port = _pick_free_loopback_port() - argv = [binary, "serve", "--host", "127.0.0.1", "--port", str(port)] + (["--block-private-networks"] if block_private_networks else []) + argv = [binary, "serve", "--host", "127.0.0.1", "--port", str(port)] + if _binary_supports_http_cache(binary): + argv += [_HTTP_CACHE_FLAG, str(_http_cache_dir())] + if block_private_networks: + argv.append("--block-private-networks") log_path = str(_state_dir() / f"{session_name}.log") try: with open(log_path, "wb") as log_file: diff --git a/tools/code_execution_tool.py b/tools/code_execution_tool.py index 9891682207..0066eca007 100644 --- a/tools/code_execution_tool.py +++ b/tools/code_execution_tool.py @@ -134,9 +134,9 @@ _TOOL_STUBS = { "write_file": ("path: str, content: str, cross_profile: bool = False", '"""Write content to a file (always overwrites). Returns dict with status."""', '{"path": path, "content": content, "cross_profile": cross_profile}'), - "search_files": ('pattern: str, target: str = "content", path: str = ".", file_glob: str = None, limit: int = 50, offset: int = 0, output_mode: str = "content", context: int = 0', + "search_files": ('pattern: str, target: str = "content", path: str = ".", file_glob: str = None, limit: int = 50, offset: int = 0, output_mode: str = "content", context: int = 0, order: str = "discovery"', '"""Search file contents (target="content") or find files by name (target="files"). Returns dict with "matches"."""', - '{"pattern": pattern, "target": target, "path": path, "file_glob": file_glob, "limit": limit, "offset": offset, "output_mode": output_mode, "context": context}'), + '{"pattern": pattern, "target": target, "path": path, "file_glob": file_glob, "limit": limit, "offset": offset, "output_mode": output_mode, "context": context, "order": order}'), "patch": ('path: str = None, old_string: str = None, new_string: str = None, replace_all: bool = False, mode: str = "replace", patch: str = None, cross_profile: bool = False', '"""Targeted find-and-replace (mode="replace") or V4A multi-file patches (mode="patch"). Returns dict with status."""', '{"path": path, "old_string": old_string, "new_string": new_string, "replace_all": replace_all, "mode": mode, "patch": patch, "cross_profile": cross_profile}'), @@ -798,7 +798,7 @@ _TOOL_DOC_LINES = [ ("read_file", " read_file(path: str, offset: int = 1, limit: int = 2000) -> dict\n" " Lines are 1-indexed. Returns {\"content\": \"...\", \"total_lines\": N}"), ("write_file", " write_file(path: str, content: str) -> dict\n Always overwrites the entire file."), - ("search_files", " search_files(pattern: str, target=\"content\", path=\".\", file_glob=None, limit=50) -> dict\n" + ("search_files", " search_files(pattern: str, target=\"content\", path=\".\", file_glob=None, limit=50, order=\"discovery\") -> dict\n" " target: \"content\" (search inside files) or \"files\" (find files by name). Returns {\"matches\": [...]}"), ("patch", " patch(path: str, old_string: str, new_string: str, replace_all: bool = False) -> dict\n" " Replaces old_string with new_string in the file."), diff --git a/tools/code_kernel.py b/tools/code_kernel.py index 4bf80d5947..e981a28d4e 100644 --- a/tools/code_kernel.py +++ b/tools/code_kernel.py @@ -78,12 +78,102 @@ import io import json import os import sys +import threading import traceback _SENTINEL = os.environ["HERMES_KERNEL_SENTINEL"] _CAPTURE_LIMIT = {capture_limit} _SPILL_DIR = os.environ.get("HERMES_KERNEL_SPILL_DIR", "") _SPILL_CAP = {spill_cap} +_PARENT_PROCESS_HANDLE = os.environ.pop("HERMES_KERNEL_PARENT_PROCESS_HANDLE", "") +_PARENT_DEATH_FD = os.environ.pop("HERMES_KERNEL_PARENT_DEATH_FD", "") + + +def _start_parent_death_pipe_watchdog(): + """POSIX twin of the Windows handle watchdog: exit when the parent dies. + + The host holds the only write end of an inherited pipe; a blocking read + returns EOF the instant the host exits by ANY means (SIGKILL, OOM, crash), + exactly like the MCP death supervisor. Stdin EOF alone is not enough: the + main loop only sees it between cells, so a kernel SIGKILLed mid-cell + outlived its host. Not PR_SET_PDEATHSIG — that is bound to the spawning + THREAD, and kernels are spawned from per-cell threads that exit. + """ + global _PARENT_DEATH_FD + raw_fd = _PARENT_DEATH_FD + _PARENT_DEATH_FD = "" + if sys.platform == "win32" or not raw_fd: + return + try: + fd = int(raw_fd) + os.set_inheritable(fd, False) + except (OSError, ValueError): + return + + def _wait(): + try: + while os.read(fd, 1): + pass + except OSError: + pass + os._exit(0) + + threading.Thread(target=_wait, name="hermes-parent-watchdog", daemon=True).start() + + +def _start_parent_process_watchdog(): + """Exit when the exact Windows parent process object is signaled. + + The inherited SYNCHRONIZE handle names a process object, not a reusable + PID. Missing or invalid handles fail open so watchdog setup can never kill + an otherwise healthy kernel. + """ + global _PARENT_PROCESS_HANDLE + raw_handle = _PARENT_PROCESS_HANDLE + _PARENT_PROCESS_HANDLE = "" + if sys.platform != "win32" or not raw_handle: + return + try: + import ctypes + from ctypes import wintypes + + handle = int(raw_handle) + if handle <= 0: + return + kernel32 = ctypes.WinDLL("kernel32", use_last_error=True) + kernel32.WaitForSingleObject.argtypes = [wintypes.HANDLE, wintypes.DWORD] + kernel32.WaitForSingleObject.restype = wintypes.DWORD + kernel32.SetHandleInformation.argtypes = [ + wintypes.HANDLE, + wintypes.DWORD, + wintypes.DWORD, + ] + kernel32.SetHandleInformation.restype = wintypes.BOOL + kernel32.CloseHandle.argtypes = [wintypes.HANDLE] + kernel32.CloseHandle.restype = wintypes.BOOL + # This process needs the handle, but user code spawned by a cell must + # not pass it any further. If Windows refuses to clear inheritance, + # disable the watchdog rather than leak the handle into cell children. + if not kernel32.SetHandleInformation(handle, 0x00000001, 0): + kernel32.CloseHandle(handle) + return + except (ImportError, OSError, TypeError, ValueError): + return + + def _wait(): + try: + result = kernel32.WaitForSingleObject(handle, 0xFFFFFFFF) + finally: + kernel32.CloseHandle(handle) + if result == 0x00000000: # WAIT_OBJECT_0: the parent exited + os._exit(0) + + threading.Thread(target=_wait, name="hermes-parent-watchdog", daemon=True).start() + + +_start_parent_process_watchdog() +_start_parent_death_pipe_watchdog() + _real_stdout = sys.stdout {cell_source} @@ -222,8 +312,13 @@ class SessionKernel: self.sock_path: Optional[str] = None self.server_sock: Optional[socket.socket] = None self.stop_event = threading.Event() + self.death_pipe_w: Optional[int] = None self.tool_call_log: List = [] self.tool_call_counter: List[int] = [0] + # Cells currently attached (bumped under the registry lock on selection, dropped when the + # cell settles). Reaping/cap-eviction skip attached kernels: tearing one down mid-spawn + # rmtree'd the staging dir under the spawner and killed live cells. + self.attached: int = 0 self.response_q: "queue.Queue[dict]" = queue.Queue() self.raw, self.stderr = _BoundedBuffer(), _BoundedBuffer() self.execution_count, self.last_used = 0, time.monotonic() @@ -232,8 +327,20 @@ class SessionKernel: def alive(self) -> bool: return self.proc is not None and self.proc.poll() is None + def dead(self) -> bool: + """True only once a spawned process has exited. ``proc is None`` is mid-spawn, not dead: + parallel cells for one owner race the first ``_spawn``, and treating the pending kernel as + dead made every racer replace it, orphaning the winner's process outside the registry.""" + return self.proc is not None and self.proc.poll() is not None + def teardown(self) -> None: self.stop_event.set() + if self.death_pipe_w is not None: + try: + os.close(self.death_pipe_w) + except OSError: + pass + self.death_pipe_w = None if self.alive(): from tools.code_execution_tool import _kill_process_group _kill_process_group(self.proc, escalate=True) @@ -267,9 +374,11 @@ class KernelRegistry: self._teardown(kernel) def discard(self, key: Tuple, kernel: Any) -> None: - """Drop one registry entry and tear the kernel down.""" + """Drop *kernel*'s registry entry (only if it is still the one registered under *key* — + never a replacement) and tear the kernel down.""" with self.lock: - self.kernels.pop(key, None) + if self.kernels.get(key) is kernel: + self.kernels.pop(key, None) self._teardown(kernel) @@ -436,8 +545,37 @@ def _bind_rpc_socket(kernel: SessionKernel) -> str: return rpc_endpoint +def _parent_process_handle(child_env: Dict[str, str]): + """Windows: open an inheritable SYNCHRONIZE handle to this process for the kernel's parent-death + watchdog. Returns (handle, CloseHandle, startupinfo) or (None, None, None); fails open.""" + handle = close = startupinfo = None + try: + import ctypes + from ctypes import wintypes + kernel32 = ctypes.WinDLL("kernel32", use_last_error=True) + kernel32.GetCurrentProcessId.argtypes = [] + kernel32.GetCurrentProcessId.restype = wintypes.DWORD + kernel32.OpenProcess.argtypes = [wintypes.DWORD, wintypes.BOOL, wintypes.DWORD] + kernel32.OpenProcess.restype = wintypes.HANDLE + kernel32.CloseHandle.argtypes = [wintypes.HANDLE] + kernel32.CloseHandle.restype = wintypes.BOOL + close = kernel32.CloseHandle + # SYNCHRONIZE; inherited only by the explicitly allow-listed child. + handle = kernel32.OpenProcess(0x00100000, True, kernel32.GetCurrentProcessId()) + if handle: + child_env["HERMES_KERNEL_PARENT_PROCESS_HANDLE"] = str(int(handle)) + startupinfo = subprocess.STARTUPINFO() + startupinfo.lpAttributeList = {"handle_list": [int(handle)]} + except (AttributeError, ImportError, OSError, TypeError, ValueError): + if handle and close is not None: + close(handle) + child_env.pop("HERMES_KERNEL_PARENT_PROCESS_HANDLE", None) + handle = close = startupinfo = None + return handle, close, startupinfo + + def _spawn(kernel: SessionKernel, *, child_python: str, child_cwd: str, - sandbox_tools: frozenset, max_tool_calls: int) -> None: + sandbox_tools: frozenset, max_tool_calls: int, task_id: str = "") -> None: from tools.code_execution_tool import _build_child_env, generate_hermes_tools_module kernel.tmpdir = tempfile.mkdtemp(prefix="hermes_kernel_") kernel.rpc_token = secrets.token_urlsafe(32) @@ -453,13 +591,29 @@ def _spawn(kernel: SessionKernel, *, child_python: str, child_cwd: str, child_env["HERMES_KERNEL_SPILL_DIR"] = kernel.tmpdir # Generated client reconnects after the RPC server's 300s idle timeout between cells. child_env["HERMES_RPC_PERSISTENT"] = "1" - kernel.proc = subprocess.Popen( - [child_python, os.path.join(kernel.tmpdir, "hermes_kernel_runner.py")], - # Strict mode passes an empty cwd: the kernel's staging dir plays the per-call tmpdir's role. - cwd=child_cwd or kernel.tmpdir, env=child_env, start_new_session=True, - stdout=subprocess.PIPE, stderr=subprocess.PIPE, stdin=subprocess.PIPE, - creationflags=subprocess.CREATE_NO_WINDOW if _IS_WINDOWS else 0, - ) + # Parent-death watchdog plumbing: Windows inherits a SYNCHRONIZE handle to this process; POSIX + # inherits the read end of a pipe whose only write end we hold (EOF == host gone, any cause). + parent_handle, close_handle, startupinfo = _parent_process_handle(child_env) if _IS_WINDOWS else (None, None, None) + death_r: Optional[int] = None + pass_fds: Tuple[int, ...] = () + if not _IS_WINDOWS: + death_r, kernel.death_pipe_w = os.pipe() + child_env["HERMES_KERNEL_PARENT_DEATH_FD"] = str(death_r) + pass_fds = (death_r,) + try: + kernel.proc = subprocess.Popen( + [child_python, os.path.join(kernel.tmpdir, "hermes_kernel_runner.py")], + # Strict mode passes an empty cwd: the kernel's staging dir plays the per-call tmpdir's role. + cwd=child_cwd or kernel.tmpdir, env=child_env, start_new_session=True, + stdout=subprocess.PIPE, stderr=subprocess.PIPE, stdin=subprocess.PIPE, + creationflags=subprocess.CREATE_NO_WINDOW if _IS_WINDOWS else 0, + close_fds=True, pass_fds=pass_fds, startupinfo=startupinfo, + ) + finally: + if parent_handle and close_handle is not None: + close_handle(parent_handle) + if death_r is not None: + os.close(death_r) # Deliberately NOT propagate_context_to_thread: that would freeze the spawning cell's # context/callbacks into the server thread for life. Authority is rebound per cell. for target, args in ((_rpc_forever, (kernel, max_tool_calls, sandbox_tools)), @@ -474,16 +628,22 @@ def _acquire_kernel(key: Tuple, reset: bool) -> Tuple[SessionKernel, bool]: cap, idle_timeout = _lifecycle_limits() with _REGISTRY.lock: now = time.monotonic() - expired = [_KERNELS.pop(k) for k in list(_KERNELS) if now - _KERNELS[k].last_used > idle_timeout] + # Reaping and eviction skip kernels with attached cells (the last cell out tears them down). + expired = [_KERNELS.pop(k) for k in list(_KERNELS) + if _KERNELS[k].attached == 0 and now - _KERNELS[k].last_used > idle_timeout] kernel = _KERNELS.get(key) - state_reset = kernel is not None and (reset or not kernel.alive()) + state_reset = kernel is not None and (reset or kernel.dead()) if state_reset: - expired.append(_KERNELS.pop(key)) + dropped = _KERNELS.pop(key) + if dropped.attached == 0: + expired.append(dropped) kernel = None if kernel is None: kernel = _KERNELS[key] = SessionKernel(key) kernel.last_used = time.monotonic() - by_age = sorted((k for k in _KERNELS if k != key), key=lambda k: _KERNELS[k].last_used) + kernel.attached += 1 + by_age = sorted((k for k in _KERNELS if k != key and _KERNELS[k].attached == 0), + key=lambda k: _KERNELS[k].last_used) expired.extend(_KERNELS.pop(k) for k in by_age[: max(0, len(_KERNELS) - cap)]) for doomed in expired: doomed.teardown() @@ -585,6 +745,24 @@ def execute_in_session_kernel( key = (_resolve_owner(task_id) or "", mode, child_python, child_cwd, tuple(sorted(sandbox_tools))) exec_start = time.monotonic() kernel, state_reset = _acquire_kernel(key, reset) + try: + return _run_cell(kernel, key, code, task_id=task_id, child_python=child_python, child_cwd=child_cwd, + sandbox_tools=sandbox_tools, timeout=timeout, max_tool_calls=max_tool_calls, + is_interrupted=is_interrupted, exec_start=exec_start, state_reset=state_reset) + finally: + with _REGISTRY.lock: + kernel.attached -= 1 + kernel.last_used = time.monotonic() + # Dropped from the registry (reset/dead/reaped) while cells were still attached: + # the last one out owns the teardown. + orphaned = kernel.attached == 0 and _KERNELS.get(key) is not kernel + if orphaned: + kernel.teardown() + + +def _run_cell(kernel: SessionKernel, key: Tuple, code: str, *, task_id: str, child_python: str, + child_cwd: str, sandbox_tools: frozenset, timeout: int, max_tool_calls: int, + is_interrupted, exec_start: float, state_reset: bool) -> str: reused = kernel.proc is not None # Captured on the calling thread BEFORE the cell runs (the snapshot a per-call RPC thread # would get) and installed on the kernel so RPC dispatches under THIS cell's identity. @@ -592,7 +770,7 @@ def execute_in_session_kernel( with kernel.lock: try: if kernel.proc is None: - _spawn(kernel, child_python=child_python, child_cwd=child_cwd, + _spawn(kernel, task_id=task_id, child_python=child_python, child_cwd=child_cwd, sandbox_tools=sandbox_tools, max_tool_calls=max_tool_calls) assert kernel.proc is not None and kernel.proc.stdin is not None # Per-cell tool budget: the RPC loop enforces counter < max; reset without restarting. diff --git a/tools/code_kernel_remote.py b/tools/code_kernel_remote.py index d34ad00b11..5f5829efa0 100644 --- a/tools/code_kernel_remote.py +++ b/tools/code_kernel_remote.py @@ -26,7 +26,7 @@ import threading import time import uuid from dataclasses import dataclass, field -from typing import Any, Dict, Optional, Tuple +from typing import Any, Dict, List, Optional, Tuple from tools.code_kernel import RUNNER_CELL_SOURCE, KernelRegistry @@ -114,6 +114,10 @@ class RemoteKernel: last_used: float = field(default_factory=time.monotonic) execution_count: int = 0 cell_seq: int = 0 + # Cells currently running on this kernel. Reap/evict skip attached + # kernels: killing one mid-cell tears the runner out from under a live + # poll loop (same guard as tools.code_kernel, hermes-agent#101861). + attached: int = 0 def sh(self, cmd: str, timeout: int = 15) -> str: return _sh(self.env, cmd, timeout) @@ -162,6 +166,28 @@ def shutdown_remote_kernels_for_owner(owner: str) -> None: _REGISTRY.shutdown(owner) +def _reap_unlocked(idle_timeout: int) -> List["RemoteKernel"]: + """Pop idle-expired, unattached remote kernels; caller tears them down outside the lock. The + runner self-exits after the same idle window, so this clears the HOST-side entry — without it + the map grew one entry per never-revisited (owner, env_type, task_env_id) for the gateway's life.""" + now = time.monotonic() + doomed = [key for key, kernel in _REMOTE_KERNELS.items() + if kernel.attached == 0 and now - kernel.last_used > idle_timeout] + return [_REMOTE_KERNELS.pop(key) for key in doomed] + + +def _evict_over_cap_unlocked(keep: Tuple) -> List["RemoteKernel"]: + """Pop least-recently-used unattached remote kernels beyond the process-wide cap (the same + ``max_session_kernels`` bound as local kernels, applied independently to this map).""" + from tools.code_kernel import _lifecycle_limits + cap, _ = _lifecycle_limits() + if len(_REMOTE_KERNELS) <= cap: + return [] + by_age = sorted((key for key in _REMOTE_KERNELS if key != keep and _REMOTE_KERNELS[key].attached == 0), + key=lambda key: _REMOTE_KERNELS[key].last_used) + return [_REMOTE_KERNELS.pop(key) for key in by_age[: len(_REMOTE_KERNELS) - cap]] + + atexit.register(shutdown_all_remote_kernels) @@ -215,11 +241,15 @@ def _spawn_remote_kernel(env, env_type: str, owner: str, task_env_id: str, def _acquire_remote_kernel(env, env_type: str, owner: str, task_env_id: str, sandbox_tools: frozenset, *, reset: bool, idle_exit: int) -> Tuple[Optional[RemoteKernel], bool, bool, bool]: - """Find/respawn the owner's kernel: (kernel|None, reused, state_reset, state_lost).""" + """Find/respawn the owner's kernel: (kernel|None, reused, state_reset, state_lost); reaps + idle-expired entries on the way in.""" key = _kernel_key(owner, env_type, task_env_id) state_lost = state_reset = False with _REGISTRY.lock: + expired = _reap_unlocked(idle_exit) kernel = _REMOTE_KERNELS.get(key) + for doomed in expired: + doomed.kill() if kernel is not None and reset: _REGISTRY.discard(key, kernel) kernel, state_reset = None, True @@ -274,8 +304,6 @@ def execute_in_remote_kernel( post-processes output), or ``None`` when no kernel could be spawned (caller falls open to per-call). ``state_lost``/``state_reset``/``reused`` ride in the ``kernel`` sub-dict.""" from tools.code_kernel import _resolve_owner - from tools.code_execution_tool import _rpc_poll_loop - from tools.thread_context import propagate_context_to_thread owner = _resolve_owner(task_env_id) kernel, reused, state_reset, state_lost = _acquire_remote_kernel( env, env_type, owner, task_env_id, sandbox_tools, reset=reset, idle_exit=idle_exit) @@ -283,6 +311,26 @@ def execute_in_remote_kernel( return None # fail open to per-call key = _kernel_key(owner, env_type, task_env_id) kernel.last_used = time.monotonic() + with _REGISTRY.lock: + kernel.attached += 1 + evicted = _evict_over_cap_unlocked(keep=key) + for doomed in evicted: + doomed.kill() + try: + return _run_attached_cell(kernel, key, code, env=env, task_env_id=task_env_id, + sandbox_tools=sandbox_tools, timeout=timeout, max_tool_calls=max_tool_calls, + reused=reused, state_reset=state_reset, state_lost=state_lost) + finally: + with _REGISTRY.lock: + kernel.attached -= 1 + kernel.last_used = time.monotonic() + + +def _run_attached_cell(kernel: RemoteKernel, key: Tuple, code: str, *, env, task_env_id: str, + sandbox_tools: frozenset, timeout: int, max_tool_calls: int, + reused: bool, state_reset: bool, state_lost: bool) -> Dict[str, Any]: + from tools.code_execution_tool import _rpc_poll_loop + from tools.thread_context import propagate_context_to_thread # Clean stale tool-RPC requests from a previous cell before arming this cell's poll loop, so # a background thread the last cell leaked cannot smuggle a call into this authority window. q_rpc = shlex.quote(kernel.kernel_dir + '/rpc') diff --git a/tools/computer_use/cua_backend.py b/tools/computer_use/cua_backend.py index 81159459b3..75ad3034c5 100644 --- a/tools/computer_use/cua_backend.py +++ b/tools/computer_use/cua_backend.py @@ -38,6 +38,7 @@ from tools.computer_use.cua_backend_session import _AsyncBridge, _CuaDriverSessi logger = logging.getLogger(__name__) # cua-driver's anonymous PostHog telemetry gate ("0" disables; absent => ON upstream). _CUA_TELEMETRY_ENV_VAR = "CUA_DRIVER_RS_TELEMETRY_ENABLED" +_CUA_NATIVE_WAYLAND_ENV_VAR = "CUA_DRIVER_RS_ENABLE_WAYLAND" def _computer_use_cfg() -> Dict[str, Any]: @@ -97,8 +98,15 @@ def _computer_use_max_image_dimension() -> Optional[int]: def cua_driver_child_env(base_env: Optional[Dict[str, str]] = None) -> Dict[str, str]: """Env for spawning cua-driver: ``base_env`` (default ``os.environ``) plus ``CUA_DRIVER_RS_TELEMETRY_ENABLED=0`` - unless the user opted in. Used by every spawn site (MCP, status, doctor, install) so the policy is uniform.""" - return {**(os.environ if base_env is None else base_env), **({_CUA_TELEMETRY_ENV_VAR: "0"} if _cua_telemetry_disabled() else {})} + unless the user opted in, plus the native-Wayland bridge (``computer_use.native_wayland`` config opt-in, only when + the child has a Wayland display). Used by every spawn site (MCP, status, doctor, install) so CLI and gateway + runtimes share one policy.""" + env = dict(os.environ if base_env is None else base_env) + if _cua_telemetry_disabled(): + env[_CUA_TELEMETRY_ENV_VAR] = "0" + if sys.platform == "linux" and env.get("WAYLAND_DISPLAY") and bool(_computer_use_cfg().get("native_wayland", False)): + env[_CUA_NATIVE_WAYLAND_ENV_VAR] = "1" + return env def sanitized_cua_driver_env() -> Dict[str, str]: """``cua_driver_child_env()`` with Hermes provider secrets stripped — cua-driver is a third-party binary and must diff --git a/tools/computer_use/doctor.py b/tools/computer_use/doctor.py index 983053f095..0f8e167d91 100644 --- a/tools/computer_use/doctor.py +++ b/tools/computer_use/doctor.py @@ -7,6 +7,7 @@ list_apps, CLI --version). Exit codes: 0 overall=="ok"; 1 degraded/failed; 2 bin from __future__ import annotations import json +import os import platform as _platform_mod import re import subprocess @@ -277,9 +278,17 @@ def _apply_display_count_guard(report: Report) -> Report: report["overall"] = "degraded" return report -def _print_text_report(report: Report, color: bool, *, identity: Optional[Report] = None) -> None: +def _wayland_environment_context(report: Report) -> Optional[Report]: + """Linux+Wayland only: doctor probes the CLI process's environment, not the gateway's.""" + if report.get("platform") != "linux" or not os.environ.get("WAYLAND_DISPLAY"): + return None + return {"scope": "cli_process", "gateway_environment_checked": False} + +def _print_text_report(report: Report, color: bool, *, identity: Optional[Report] = None, + environment: Optional[Report] = None) -> None: """Render like `cua-driver call health_report`: header (CLI --version preferred over health_report's stale - ``driver_version``), identity block, one line per check + indented hint/``data`` rows (support staff need them).""" + ``driver_version``), identity block, environment note, one line per check + indented hint/``data`` rows + (support staff need them).""" platform, report_v, overall = (report.get(k, "?") for k in ("platform", "driver_version", "overall")) identity = identity or {} cli_v = identity.get("cli_version") or "" @@ -293,6 +302,9 @@ def _print_text_report(report: Report, color: bool, *, identity: Optional[Report lines.append(f" {dim}binary: {identity['resolved_binary']}{reset}") if cli_v and report_v and str(report_v) not in str(cli_v) and str(cli_v) not in str(report_v): # clearly differ lines += [f" {dim}--version: {cli_v}{reset}", f" {dim}health_report.driver_version: {report_v}{reset}"] + if environment: + lines += [f" {dim}environment: current CLI process{reset}", + f" {dim}gateway environment was not checked; active gateway computer_use sessions use that process environment{reset}"] if identity.get("version_mismatch"): lines += [f" {yellow}⚠️ version mismatch: health_report says {report_v!r} but binary --version is {cli_v!r}{reset}", f" {dim}→ trust --version / packages/current for debugging; health_report's binary_version check can lag on Windows{reset}"] @@ -329,10 +341,16 @@ def run_doctor(driver_cmd: Optional[str] = None, *, include: Sequence[str] = (), return 2 report = _apply_display_count_guard(report) identity = _build_identity(binary, report) + environment = _wayland_environment_context(report) if json_output: - # Additive envelope: upstream keys preserved, identity under hermes_identity so overall/checks parsers keep working. - json.dump({**report, "hermes_identity": identity}, sys.stdout, indent=2, sort_keys=True) + # Additive envelope: upstream keys preserved, identity under hermes_identity (and environment under + # hermes_environment when present) so overall/checks parsers keep working. + payload = {**report, "hermes_identity": identity} + if environment: + payload["hermes_environment"] = environment + json.dump(payload, sys.stdout, indent=2, sort_keys=True) sys.stdout.write("\n") else: - _print_text_report(report, color=sys.stdout.isatty() if color is None else bool(color), identity=identity) + _print_text_report(report, color=sys.stdout.isatty() if color is None else bool(color), identity=identity, + environment=environment) return 0 if report.get("overall") == "ok" else 1 # unknown/missing overall must not look like success diff --git a/tools/cronjob_tools.py b/tools/cronjob_tools.py index 3c5b6881d8..fccad6090f 100644 --- a/tools/cronjob_tools.py +++ b/tools/cronjob_tools.py @@ -277,13 +277,23 @@ def _run_claimed_job(job: Dict[str, Any], extra_prompt: Optional[str] = None) -> _registered = False release_running_job(job_id) refreshed = get_job(job_id) or {} + execution = None + execution_id = job.get("execution_id") + if execution_id: + from cron.executions import get_execution + + execution = get_execution(str(execution_id)) last_status = refreshed.get("last_status") # "delivery_failed": the run succeeded but output never reached the user — not a # success for the caller; surface last_delivery_error. run_error = refreshed.get("last_error") if last_status == "delivery_failed" and not run_error: run_error = refreshed.get("last_delivery_error") - return {"claimed": True, "success": bool(processed and last_status == "ok"), "error": run_error} + ok = last_status == "ok" + if execution is not None and execution.get("status") != "completed": + ok = False + run_error = execution.get("error") or f"execution ended in {execution.get('status') or 'unknown'} state" + return {"claimed": True, "success": bool(processed and ok), "error": run_error} except Exception as e: logger.error("Failed to execute cron job %s immediately: %s", job_id, e) if _registered: diff --git a/tools/delegate_tool.py b/tools/delegate_tool.py index 137cfb60b4..6fc870113e 100644 --- a/tools/delegate_tool.py +++ b/tools/delegate_tool.py @@ -252,7 +252,7 @@ def _run_single_child( # thread before the worker settles races the conversation's finally path. _child_close_deferred = False try: - heartbeat[1].start() + heartbeat.start() _safe_progress(child_progress_cb, "subagent.start", preview=goal) run.seed_workspace() result, failure_entry, _child_close_deferred = run.await_child() diff --git a/tools/delegate_tool_child_run.py b/tools/delegate_tool_child_run.py index 70dcf74dea..7434a394dc 100644 --- a/tools/delegate_tool_child_run.py +++ b/tools/delegate_tool_child_run.py @@ -202,55 +202,74 @@ def _dump_subagent_timeout_diagnostic( # ── Per-run helpers ────────────────────────────────────────────────────────── -def _start_heartbeat(child: Any, parent_agent: Any, task_index: int) -> tuple: - """``(stop_event, thread)`` for one child's parent-activity heartbeat, NOT - started: the caller starts it inside its ``try`` so a failed ``start()`` (OS - thread exhaustion) leaves ``ident`` None and the finally-path join is skipped.""" - from tools.delegate_tool import (_HEARTBEAT_INTERVAL, _HEARTBEAT_STALE_CYCLES_IDLE, _HEARTBEAT_STALE_CYCLES_IN_TOOL) - _heartbeat_stop = threading.Event() - # Stale detection: a cycle counts as stale when (tool, iteration, - # activity_ts) all froze; thresholds differ idle vs in-tool. - last_seen = {"iter": 0, "tool": None, "ts": None, "stale": 0} +class _Heartbeat: + """One child's parent-activity heartbeat on the shared periodic scheduler thread + (``agent.periodic_scheduler``) — not one daemon thread per child. NOT started at construction: + the caller calls ``start()`` inside its ``try`` so a failed schedule (OS thread exhaustion on + first use) leaves ``handle`` None and ``stop()`` is a no-op.""" - def _heartbeat_loop(): - while not _heartbeat_stop.wait(_HEARTBEAT_INTERVAL): - touch = getattr(parent_agent, "_touch_activity", None) if parent_agent is not None else None - if not touch: - continue - desc = f"delegate_task: subagent {task_index} working" - try: - child_summary = child.get_activity_summary() - child_tool = child_summary.get("current_tool") - child_iter = child_summary.get("api_call_count", 0) - child_max = child_summary.get("max_iterations", 0) - child_activity_ts = child_summary.get("last_activity_ts") - # A slow model wait refreshes last_activity_ts (direct_api_call - # heartbeat), so it never looks stale at the idle threshold. - activity_advanced = child_activity_ts is not None and ( - last_seen["ts"] is None or child_activity_ts > last_seen["ts"] + def __init__(self, child: Any, parent_agent: Any, task_index: int): + self.child, self.parent_agent, self.task_index = child, parent_agent, task_index + # Stale detection: a cycle counts as stale when (tool, iteration, + # activity_ts) all froze; thresholds differ idle vs in-tool. + self.last_seen = {"iter": 0, "tool": None, "ts": None, "stale": 0} + self.handle = None + + def start(self) -> None: + from agent.periodic_scheduler import schedule + from tools.delegate_tool import _HEARTBEAT_INTERVAL + self.handle = schedule(self.tick, _HEARTBEAT_INTERVAL) + + def stop(self) -> None: + """wait=5 mirrors the old thread join: an in-flight tick finishes.""" + if self.handle is not None: + self.handle.cancel(wait=5) + + def tick(self): + """Returning False stops the periodic callback.""" + from tools.delegate_tool import _HEARTBEAT_STALE_CYCLES_IDLE, _HEARTBEAT_STALE_CYCLES_IN_TOOL + child, parent_agent, task_index, last_seen = self.child, self.parent_agent, self.task_index, self.last_seen + touch = getattr(parent_agent, "_touch_activity", None) if parent_agent is not None else None + if not touch: + return None + desc = f"delegate_task: subagent {task_index} working" + try: + child_summary = child.get_activity_summary() + child_tool = child_summary.get("current_tool") + child_iter = child_summary.get("api_call_count", 0) + child_max = child_summary.get("max_iterations", 0) + child_activity_ts = child_summary.get("last_activity_ts") + # A slow model wait refreshes last_activity_ts (direct_api_call + # heartbeat), so it never looks stale at the idle threshold. + activity_advanced = child_activity_ts is not None and ( + last_seen["ts"] is None or child_activity_ts > last_seen["ts"] + ) + if child_iter > last_seen["iter"] or child_tool != last_seen["tool"] or activity_advanced: + last_seen.update(iter=child_iter, tool=child_tool, stale=0) + if child_activity_ts is not None: + last_seen["ts"] = child_activity_ts + else: + last_seen["stale"] += 1 + if last_seen["stale"] >= (_HEARTBEAT_STALE_CYCLES_IN_TOOL if child_tool else _HEARTBEAT_STALE_CYCLES_IDLE): + logger.warning( + "Subagent %d appears stale (no progress for %d heartbeat cycles, tool=%s) — stopping heartbeat", + task_index, last_seen["stale"], child_tool or "", ) - if child_iter > last_seen["iter"] or child_tool != last_seen["tool"] or activity_advanced: - last_seen.update(iter=child_iter, tool=child_tool, stale=0) - if child_activity_ts is not None: - last_seen["ts"] = child_activity_ts - else: - last_seen["stale"] += 1 - if last_seen["stale"] >= (_HEARTBEAT_STALE_CYCLES_IN_TOOL if child_tool else _HEARTBEAT_STALE_CYCLES_IDLE): - logger.warning( - "Subagent %d appears stale (no progress for %d heartbeat cycles, tool=%s) — stopping heartbeat", - task_index, last_seen["stale"], child_tool or "", - ) - break # stop touching parent, let gateway timeout fire - if child_tool: - desc = f"delegate_task: subagent running {child_tool} (iteration {child_iter}/{child_max})" - elif child_summary.get("last_activity_desc", ""): - desc = f"delegate_task: subagent {child_summary.get('last_activity_desc', '')} (iteration {child_iter}/{child_max})" - except Exception: - pass - with _quiet(None): - touch(desc) + return False # stop touching parent, let gateway timeout fire + if child_tool: + desc = f"delegate_task: subagent running {child_tool} (iteration {child_iter}/{child_max})" + elif child_summary.get("last_activity_desc", ""): + desc = f"delegate_task: subagent {child_summary.get('last_activity_desc', '')} (iteration {child_iter}/{child_max})" + except Exception: + pass + with _quiet(None): + touch(desc) + return None - return _heartbeat_stop, threading.Thread(target=_heartbeat_loop, daemon=True) + +def _start_heartbeat(child: Any, parent_agent: Any, task_index: int) -> _Heartbeat: + """Build (not start) one child's heartbeat; see ``_Heartbeat``.""" + return _Heartbeat(child, parent_agent, task_index) def _register_child( child: Any, parent_agent: Any, goal: str, *, owner_session_id: Optional[str], owner_transport: Any, @@ -748,16 +767,13 @@ class _ChildRun: complete_kwargs["cost_usd"] = float(_cost_usd) _safe_progress(self.child_progress_cb, "subagent.complete", **complete_kwargs) - def cleanup(self, *, heartbeat: tuple, child_pool: Any, leased_cred_id: Any, close_deferred: bool) -> None: + def cleanup(self, *, heartbeat: _Heartbeat, child_pool: Any, leased_cred_id: Any, close_deferred: bool) -> None: """Finally-path teardown (idempotent, never raises). Order matters: stop heartbeat → drop registry entry → release credential lease → restore the parent's process-global tool names → detach from the parent's interrupt list → close the child (unless a timed-out worker still owns it) → pop the child's Relay scope if no turn is active.""" child = self.child - _heartbeat_stop, _heartbeat_thread = heartbeat - _heartbeat_stop.set() - if _heartbeat_thread.ident is not None: - _heartbeat_thread.join(timeout=5) + heartbeat.stop() # Safe even if the child was never registered (ID missing on test doubles). if self.subagent_id: diff --git a/tools/file_operations.py b/tools/file_operations.py index d964197a78..0bcfa7eabc 100644 --- a/tools/file_operations.py +++ b/tools/file_operations.py @@ -15,6 +15,8 @@ import sys # noqa: F401 (tests monkeypatch tools.file_operations.sys.platform) import difflib import hashlib import json +import logging +import secrets import unicodedata from abc import ABC, abstractmethod from typing import Optional, Dict @@ -28,8 +30,14 @@ from tools.file_operations_common import ( # noqa: F401 (re-exported) _strip_terminal_fence_leaks, normalize_read_pagination, normalize_search_pagination) from tools.file_operations_lint import LINTERS_INPROC, LintMixin, _FAIL_CLOSED_INPROC_EXTS from tools.file_operations_search import ( # noqa: F401 (re-exported) - SearchMixin, _macos_protected_search_exclusions, _parse_search_context_line, - _pattern_has_regex_newline, _search_stdout_and_limit, _split_tool_diagnostics) + SearchMixin, _ACTIVE_FILENAME_SEARCH_ROOTS, _FILENAME_SEARCH_ADMISSION, + _acquire_filename_search_roots, _filename_search_root_keys, + _macos_protected_search_exclusions, _normalized_filename_search_root, + _parse_search_context_line, _pattern_has_regex_newline, + _release_filename_search_roots, _search_stdout_and_limit, _split_tool_diagnostics) +from tools import interrupt as tool_interrupt # noqa: F401 (tests patch it via this module) + +logger = logging.getLogger(__name__) # Controller home; SearchMixin reads it (tests monkeypatch it here). _HOME = str(Path.home()) @@ -126,7 +134,8 @@ class FileOperations(ABC): @abstractmethod def search(self, pattern: str, path: str = ".", target: str = "content", file_glob: Optional[str] = None, limit: int = 50, offset: int = 0, - output_mode: str = "content", context: int = 0) -> SearchResult: + output_mode: str = "content", context: int = 0, + order: str = "discovery") -> SearchResult: """Search for content or files.""" @@ -139,6 +148,30 @@ IMAGE_EXTENSIONS = {'.png', '.jpg', '.jpeg', '.gif', '.webp', '.bmp', '.ico'} # `wc -c` prints only digits, so this can never collide with a real size. NOT_REGULAR_SENTINEL = "__hermes_not_regular__" +# Echoed by the compound read/write probes when the path does not exist. A +# compound command only reports its *last* exit status, so the missing-file +# signal that ``_probe_regular_file`` carries in ``exit 1`` travels in-band. +MISSING_SENTINEL = "__hermes_missing__" + +_READ_SENTINEL_PREFIX = "__HERMES_RF_" +_WRITE_SENTINEL_PREFIX = "__HERMES_WF_" + + +def _new_sentinel(prefix: str) -> str: + """Per-call separator line for a compound shell probe. 128 random bits make a + collision with file content negligible; the underscores keep the token outside + the base64 alphabet, so a sentinel leaking into a sample segment fails base64 + validation instead of decoding into bytes.""" + return f"{prefix}{secrets.token_hex(16)}__" + + +def _split_segments(output: str, sentinel: str) -> list[str]: + """Split compound-probe stdout on its sentinel lines. Every producer (``wc``, + ``base64``, ``cut``) newline-terminates or prints nothing, so the separator is + always ``sentinel + "\n"`` on its own line; the text after the final sentinel + is the status segment.""" + return output.split(sentinel + "\n") + class ShellFileOperations(LintMixin, SearchMixin, FileOperations): """File operations over any terminal backend exposing ``execute(command, cwd)`` @@ -155,7 +188,12 @@ class ShellFileOperations(LintMixin, SearchMixin, FileOperations): # Never os.getcwd(): that is the HOST path, absent inside container backends. self.cwd = cwd or getattr(terminal_env, 'cwd', None) or \ getattr(getattr(terminal_env, 'config', None), 'cwd', None) or "/" + # Ordinary executables: bool cache (hits AND misses). rg is special — it has + # an off-PATH resolver and may be installed mid-session — so only successful + # rg resolutions are cached (see SearchMixin._resolve_command). self._command_cache: Dict[str, bool] = {} + self._rg_resolution_cache: Dict[str, str] = {} + self._rg_modified_capability: Dict[str, Optional[str]] = {} def _exec(self, command: str, cwd: str = None, timeout: int = None, stdin_data: str = None) -> ExecuteResult: @@ -176,7 +214,10 @@ class ShellFileOperations(LintMixin, SearchMixin, FileOperations): return ExecuteResult(stdout=result.get("output", ""), exit_code=exit_code) def _has_command(self, cmd: str) -> bool: - """Check if a command exists in the environment (cached).""" + """Check if a command exists in the environment (cached); rg goes through + the resolver so a mid-session install becomes visible.""" + if cmd == "rg": + return self._resolve_command(cmd) is not None if cmd not in self._command_cache: result = self._exec(f"command -v {cmd} >/dev/null 2>&1 && echo 'yes'") self._command_cache[cmd] = result.stdout.strip() == 'yes' @@ -206,7 +247,14 @@ class ShellFileOperations(LintMixin, SearchMixin, FileOperations): result = self._exec(f"head -c {length} {self._escape_shell_arg(path)} 2>/dev/null | base64") if result.exit_code != 0: return None - encoded = "".join(_strip_terminal_fence_leaks(result.stdout).split()) + return self._decode_base64_sample(result.stdout) + + @staticmethod + def _decode_base64_sample(text: str) -> Optional[bytes]: + """Decode one ``head -c N | base64`` sample. Whitespace-joins the whole text + first (``base64`` wraps at 76 columns), so callers hand over exactly one + segment; anything else fails validation → None (legacy text heuristic).""" + encoded = "".join(_strip_terminal_fence_leaks(text).split()) if not encoded: return b"" if not re.fullmatch(r"[A-Za-z0-9+/]+={0,2}", encoded): @@ -484,41 +532,264 @@ class ShellFileOperations(LintMixin, SearchMixin, FileOperations): """Read a file with pagination, binary detection, and line numbers. ``offset`` is 1-indexed; ``limit`` is clamped by ``normalize_read_pagination``. + One shell round-trip answers every question the read needs (existence, size, + binary sample, page, line count, trailing newline; see ``_read_probe_cmd``). + An unparseable reply falls back to ``_read_file_sequential`` (one probe per + question), so an exotic shell can never do worse than before. On a local + POSIX environment the read never touches the shell (``_read_file_native``). """ path = self._expand_path(path) # before shell escaping: ~ doesn't expand in quotes offset, limit = normalize_read_pagination(offset, limit) + + if self._native_read_enabled(): + return self._read_file_native(path, offset, limit) + + # Images / known-binary extensions never inline content; the sequential + # path stops at the probes for them, so don't stream their bytes. + if self._is_image(path) or os.path.splitext(path)[1].lower() in BINARY_EXTENSIONS: + return self._read_file_sequential(path, offset, limit) + + from tools.tool_output_limits import get_max_line_length + line_clamp_bytes = 4 * get_max_line_length() + 1 + end_line = offset + limit - 1 + sentinel = _new_sentinel(_READ_SENTINEL_PREFIX) + probe = self._exec(self._read_probe_cmd(path, offset, end_line, line_clamp_bytes, sentinel)) + output = probe.stdout or "" + + if sentinel not in output: + # Single-line replies: the path is missing or not a regular file. + marker = _strip_terminal_fence_leaks(output).strip() + if marker == MISSING_SENTINEL: + return self._read_file_missing(path, offset, limit) + if marker == NOT_REGULAR_SENTINEL: + return self._not_regular_error(path) + logger.debug( + "read_file: compound probe reply for %s has no sentinel " + "(exit %s, %d chars); falling back to sequential probes", + path, probe.exit_code, len(output)) + return self._read_file_sequential(path, offset, limit) + + segments = _split_segments(output, sentinel) + if probe.exit_code != 0 or len(segments) != 6: + logger.debug( + "read_file: compound probe for %s returned exit %s with %d " + "segments (want 6); falling back to sequential probes", + path, probe.exit_code, len(segments)) + return self._read_file_sequential(path, offset, limit) + size_seg, sample_seg, page_seg, wc_seg, tail_seg, status_seg = segments + + status = _strip_terminal_fence_leaks(status_seg).split() + try: + sample_rc, read_rc = int(status[0]), int(status[1]) + except (IndexError, ValueError): + logger.debug( + "read_file: compound probe for %s has unparseable status %r; " + "falling back to sequential probes", path, status_seg[-40:]) + return self._read_file_sequential(path, offset, limit) + + try: + file_size = int(_strip_terminal_fence_leaks(size_seg).strip()) + except ValueError: + file_size = 0 + + # Byte-layer binary detection when base64 was available, else the legacy + # text heuristic over a plain sample (one extra round-trip, shells without base64). + sample_bytes = self._decode_base64_sample(sample_seg) if sample_rc == 0 else None + if sample_bytes is not None: + is_binary = self._is_likely_binary_bytes(sample_bytes) + else: + logger.debug( + "read_file: no usable base64 sample for %s (base64 exit %s); " + "paying one extra round-trip for the text heuristic", path, sample_rc) + sample_output = _strip_terminal_fence_leaks(self._head(path, 1000).stdout) + is_binary = self._is_likely_binary(path, sample_output) + if is_binary: + return self._read_binary_file(path, offset, limit, file_size, sample_bytes) + + if read_rc != 0: + return ReadResult(error=f"Failed to read file: {_strip_terminal_fence_leaks(page_seg)}") + read_output = _strip_terminal_fence_leaks(page_seg) + try: + total_lines = int(_strip_terminal_fence_leaks(wc_seg).strip()) + except ValueError: + total_lines = 0 + tail_flag = _strip_terminal_fence_leaks(tail_seg).strip() + file_ends_with_newline = tail_flag == "1" if tail_flag in ("0", "1") else None + return self._assemble_read_result( + read_output, offset=offset, end_line=end_line, total_lines=total_lines, + file_size=file_size, file_ends_with_newline=file_ends_with_newline) + + def _native_read_enabled(self) -> bool: + """Whether ``read_file`` may bypass the shell: only POSIX + ``LocalEnvironment`` + (file is on this host, path already native; Windows keeps the shell path since + file_operations holds Git-Bash-style paths there). ``HERMES_NATIVE_FILE_READ=0`` + turns the fast path off.""" + flag = os.environ.get("HERMES_NATIVE_FILE_READ", "1").strip().lower() + if flag in ("0", "false", "no", "off"): + return False + # Same "is this env the local host" test the LSP path uses; isinstance is + # microseconds and self.env is never rebound, so nothing to memoize. + return sys.platform != "win32" and self._lsp_local_only() + + def _read_file_native(self, path: str, offset: int, limit: int) -> ReadResult: + """``read_file`` without a shell — same contract as the shell path, byte for + byte. ``os.stat`` is the ``[ -f ]`` guard (a stat, never an open, so FIFOs and + devices are refused before anything touches them); the first 1000 bytes drive + the byte-layer binary check; the page is produced exactly as + ``sed -n 'a,bp' | cut -b1-N`` prints it (each line clamped to N bytes and + newline-terminated) then decoded with errors="replace" like the transport. + One chunked pass counts lines and collects the page, so neither the file nor + a pathological line is ever held whole. ``path`` is already expanded; any + unexpected OSError hands over to the shell path.""" + import stat as _stat + + full = path if os.path.isabs(path) else os.path.join( + getattr(self.env, "cwd", None) or self.cwd, path) + try: + st = os.stat(full) + except (FileNotFoundError, NotADirectoryError): + return self._read_file_missing(path, offset, limit) + except OSError: + return self._read_file_sequential(path, offset, limit) + if not _stat.S_ISREG(st.st_mode): + return self._not_regular_error(path) + file_size = st.st_size + if self._is_image(path): + return self._image_redirect_result(file_size) + + from tools.tool_output_limits import get_max_line_length + clamp = 4 * get_max_line_length() + 1 + end_line = offset + limit - 1 + + page: list[bytes] = [] + total_lines = 0 + lineno = 1 # the line currently being scanned + kept = bytearray() # first ``clamp`` bytes of that line + have_partial = False # that line has bytes but no newline yet + last_byte = b"" + try: + with open(full, "rb") as fh: + sample = fh.read(1000) + ext_binary = os.path.splitext(path)[1].lower() in BINARY_EXTENSIONS + if ext_binary or self._is_likely_binary_bytes(sample): + return self._read_binary_file(path, offset, limit, file_size, sample) + fh.seek(0) + while True: + chunk = fh.read(1 << 20) + if not chunk: + break + last_byte = chunk[-1:] + if lineno > end_line: + # Past the window: only the line count and trailing byte + # are needed, so let memchr do it instead of per-line work. + total_lines += chunk.count(b"\n") + have_partial = chunk[-1:] != b"\n" + continue + pos, n = 0, len(chunk) + while pos < n: + nl = chunk.find(b"\n", pos) + in_page = offset <= lineno <= end_line + if nl < 0: + if in_page and len(kept) < clamp: + kept += chunk[pos:pos + (clamp - len(kept))] + have_partial = True + break + if in_page: + if len(kept) < clamp: + kept += chunk[pos:min(nl, pos + (clamp - len(kept)))] + page.append(bytes(kept) + b"\n") + kept = bytearray() + have_partial = False + total_lines += 1 + lineno += 1 + pos = nl + 1 + except OSError: + return self._read_file_sequential(path, offset, limit) + if have_partial and offset <= lineno <= end_line: + # ``sed`` prints a final line that lacks a newline; ``cut`` adds one. + page.append(bytes(kept) + b"\n") + + read_output = _strip_terminal_fence_leaks(b"".join(page).decode("utf-8", errors="replace")) + return self._assemble_read_result( + read_output, offset=offset, end_line=end_line, total_lines=total_lines, + file_size=file_size, + file_ends_with_newline=(last_byte == b"\n") if file_size else None) + + @staticmethod + def _image_redirect_result(file_size: int) -> ReadResult: + return ReadResult( + is_image=True, is_binary=True, file_size=file_size, + hint=( + "Image file detected. Automatically redirected to vision_analyze tool. " + "Use vision_analyze with this file path to inspect the image contents.")) + + def _read_probe_cmd(self, path: str, offset: int, end_line: int, + line_clamp_bytes: int, sentinel: str) -> str: + """One shell command answering every question ``read_file`` asks: six + segments each closed by a ``sentinel`` line — byte size, base64 of the first + 1000 bytes, the ``sed | cut`` page, ``wc -l``, whether the last byte is a + newline, then the base64 and page pipeline statuses. Probes run only inside + ``[ -f ]`` (stat-not-open, like ``_probe_regular_file``) so a FIFO/device never + reaches ``head``/``sed``. A missing path echoes ``MISSING_SENTINEL`` (a compound + command only reports its last status). Every stage silences stderr: the local + backend merges stderr into stdout and a stray diagnostic would land inside a + segment. The byte clamp is ``4 * max_line_length + 1``; see ``_read_file_sequential``.""" + arg = self._escape_shell_arg(path) + mark = f"echo {sentinel}" + return ( + f"if [ -f {arg} ]; then " + f"wc -c < {arg} 2>/dev/null; {mark}; " + f"head -c 1000 {arg} 2>/dev/null | base64 2>/dev/null; __hs=$?; {mark}; " + f"sed -n '{offset},{end_line}p' {arg} 2>/dev/null" + f" | cut -b1-{line_clamp_bytes} 2>/dev/null; __hr=$?; {mark}; " + f"wc -l < {arg} 2>/dev/null; {mark}; " + f"tail -c 1 {arg} 2>/dev/null | wc -l; {mark}; " + f'echo "$__hs $__hr"; ' + f"elif [ -e {arg} ]; then echo {NOT_REGULAR_SENTINEL}; " + f"else echo {MISSING_SENTINEL}; fi") + + def _read_file_missing(self, path: str, offset: int, limit: int) -> ReadResult: + """Not-found recovery shared by every read path. Unicode-equivalent spellings + (NFC/NFD, confusable spaces/quotes) render identically, so the model can never + discover the byte mismatch by retyping — retrying is the tool's job. No + equivalent spelling → suggest similar files.""" + variant = self._unicode_variant_match(path) + if variant is not None: + result = self.read_file(variant, offset=offset, limit=limit) + note = ( + f"Note: '{path}' not found byte-for-byte; resolved to " + f"the unicode-equivalent file '{variant}' (invisible " + "encoding difference: NFC/NFD or special space/quote " + "characters).") + result.hint = f"{note} {result.hint}" if result.hint else note + return result + return self._suggest_similar_files(path) + + def _read_binary_file(self, path: str, offset: int, limit: int, + file_size: int, sample_bytes: Optional[bytes]) -> ReadResult: + """Binary branch shared by every read path: UTF-16 text (Notepad, PowerShell + ``>``) trips the binary guard; transcode it, else refuse with the type name.""" + utf16_result = self._try_read_utf16(path, offset, limit, file_size) + if utf16_result is not None: + return utf16_result + return ReadResult( + is_binary=True, file_size=file_size, + error=describe_binary_file(sample_bytes, file_size)) + + def _read_file_sequential(self, path: str, offset: int, limit: int) -> ReadResult: + """One-probe-per-call read: the pre-compound form, kept as fallback for + image / known-binary extensions and unparseable compound replies. ``path`` is + already expanded and ``offset``/``limit`` normalized.""" file_size, status = self._probe_regular_file(path) if status == "missing": - # Unicode-equivalent spellings render identically, so the model can never - # discover the byte mismatch by retyping — retrying is the tool's job. - variant = self._unicode_variant_match(path) - if variant is not None: - result = self.read_file(variant, offset=offset, limit=limit) - note = ( - f"Note: '{path}' not found byte-for-byte; resolved to " - f"the unicode-equivalent file '{variant}' (invisible " - "encoding difference: NFC/NFD or special space/quote " - "characters).") - result.hint = f"{note} {result.hint}" if result.hint else note - return result - return self._suggest_similar_files(path) + return self._read_file_missing(path, offset, limit) if status == "not_regular": return self._not_regular_error(path) if self._is_image(path): # never inlined — redirect to the vision tool - return ReadResult( - is_image=True, is_binary=True, file_size=file_size, - hint=( - "Image file detected. Automatically redirected to vision_analyze tool. " - "Use vision_analyze with this file path to inspect the image contents.")) + return self._image_redirect_result(file_size) is_binary, sample_bytes = self._detect_binary(path) if is_binary: - # UTF-16 text (Notepad, PowerShell `>`) trips the binary guard; transcode. - utf16_result = self._try_read_utf16(path, offset, limit, file_size) - if utf16_result is not None: - return utf16_result - return ReadResult( - is_binary=True, file_size=file_size, - error=describe_binary_file(sample_bytes, file_size)) + return self._read_binary_file(path, offset, limit, file_size, sample_bytes) # Clamp each line to a byte budget IN THE SHELL so a 400MB single-line file # never crosses the exec transport. 4*max+1 BYTES (not max+1): ``cut -b`` can @@ -535,25 +806,42 @@ class ShellFileOperations(LintMixin, SearchMixin, FileOperations): if read_result.exit_code != 0: return ReadResult(error=f"Failed to read file: {read_result.stdout}") read_output = _strip_terminal_fence_leaks(read_result.stdout) - if offset == 1: # only the first chunk can carry a BOM (byte 0) - read_output, _ = _strip_bom(read_output) wc_result = self._exec(f"wc -l < {self._escape_shell_arg(path)}") try: total_lines = int(_strip_terminal_fence_leaks(wc_result.stdout).strip()) except ValueError: total_lines = 0 + + # Only the page reaching the file's final line can carry the ``cut`` newline + # artifact (see _assemble_read_result); probe the last byte just for that case. + file_ends_with_newline: Optional[bool] = None + if not total_lines > end_line and read_output.endswith('\n'): + tail_result = self._exec(f"tail -c 1 {self._escape_shell_arg(path)} | wc -l") + if tail_result.exit_code == 0: + file_ends_with_newline = _strip_terminal_fence_leaks(tail_result.stdout).strip() != "0" + return self._assemble_read_result( + read_output, offset=offset, end_line=end_line, total_lines=total_lines, + file_size=file_size, file_ends_with_newline=file_ends_with_newline) + + def _assemble_read_result(self, read_output: str, *, offset: int, end_line: int, + total_lines: int, file_size: int, + file_ends_with_newline: Optional[bool]) -> ReadResult: + """Turn a raw ``sed | cut`` page into the final ``ReadResult``. Shared by every + read path so the BOM strip, pagination hint, ``cut`` newline-artifact fix and + the ambiguous-silence guards never drift apart. ``file_ends_with_newline`` is + None when the caller could not tell (artifact left alone, as before).""" + if offset == 1: # only the first chunk can carry a BOM (byte 0) + read_output, _ = _strip_bom(read_output) truncated = total_lines > end_line hint = None if truncated: hint = f"Use offset={end_line + 1} to continue reading (showing {offset}-{end_line} of {total_lines} lines)" # ``cut`` always newline-terminates, so a file without a trailing newline - # would grow a phantom empty last line; when this page reaches EOF, probe. - if not truncated and read_output.endswith('\n'): - tail_result = self._exec(f"tail -c 1 {self._escape_shell_arg(path)} | wc -l") - if tail_result.exit_code == 0 and _strip_terminal_fence_leaks(tail_result.stdout).strip() == "0": - read_output = read_output[:-1] + # would grow a phantom empty last line; strip it when the last byte says so. + if not truncated and read_output.endswith('\n') and file_ends_with_newline is False: + read_output = read_output[:-1] # Empty content is indistinguishable from a broken tool: name the dead end. if file_size == 0: @@ -763,33 +1051,98 @@ class ShellFileOperations(LintMixin, SearchMixin, FileOperations): f"{ext} syntax validation ({err}). The file was " "NOT created or modified. Fix the content and retry.")) - def _capture_pre_content(self, path: str, ext: str, pre_content: Optional[str]) -> Optional[str]: - """Pre-write content for the lint-delta and LSP line-shift consumers. Read - only for extensions in the UNION of in-process lint and LSP coverage (keeps - the hot path fast for binaries); a failed ``cat`` leaves None so both - consumers degrade gracefully.""" - if pre_content is not None: - return pre_content - if ext in LINTERS_INPROC or self._lsp_handles_extension(ext): + def _write_probe_cmd(self, path: str, sentinel: str, body: Optional[str]) -> str: + """One shell command for the on-disk questions ``write_file`` asks. Two + segments closed by a ``sentinel`` line: base64 of the first three bytes (BOM + detection at the byte layer, same on-disk truth as ``_file_has_bom``), then + ``body``: ``"cat"`` for the full text when pre-content is wanted, ``"sample"`` + for the 4 KB line-ending sample, or None for nothing. Gated on ``[ -f ]`` so a + FIFO/device never reaches ``head``/``cat``; a missing path echoes ``MISSING_SENTINEL``.""" + arg = self._escape_shell_arg(path) + if body == "cat": + body_cmd = f"cat {arg} 2>/dev/null" + elif body == "sample": + body_cmd = f"head -c 4096 {arg} 2>/dev/null" + else: + body_cmd = ":" + return ( + f"if [ -f {arg} ]; then " + f"head -c 3 {arg} 2>/dev/null | base64 2>/dev/null; echo {sentinel}; " + f"{body_cmd}; " + f"else echo {MISSING_SENTINEL}; fi") + + def _probe_write_target(self, path: str, pre_content: Optional[str], want_pre: bool, + ) -> tuple[bool, Optional[str], Optional[str]]: + """``(has_bom, pre_content, original_line_ending)`` for ``path`` in ONE + round-trip (replaces ``cat`` when pre-content is wanted, a ``head -c 4096`` + line-ending sample and a ``head -c 3`` BOM check). Semantics unchanged: + pre-content is read only when wanted and not supplied; the line ending comes + from pre-content when there is any, else from the sample; the BOM always comes + from disk. An unparseable reply falls back to the separate probes.""" + if want_pre and pre_content is None: + body_mode: Optional[str] = "cat" + elif not pre_content: + body_mode = "sample" + else: + body_mode = None + + sentinel = _new_sentinel(_WRITE_SENTINEL_PREFIX) + probe = self._exec(self._write_probe_cmd(path, sentinel, body_mode)) + output = probe.stdout or "" + if sentinel not in output: + if _strip_terminal_fence_leaks(output).strip() == MISSING_SENTINEL: + ending = _detect_line_ending(pre_content) if pre_content else None + return False, pre_content, ending + logger.debug( + "write_file: pre-write probe reply for %s has no sentinel " + "(exit %s, %d chars); falling back to sequential probes", + path, probe.exit_code, len(output)) + return self._probe_write_target_sequential(path, pre_content, want_pre) + + segments = _split_segments(output, sentinel) + if probe.exit_code != 0 or len(segments) != 2: + logger.debug( + "write_file: pre-write probe for %s returned exit %s with %d " + "segments (want 2); falling back to sequential probes", + path, probe.exit_code, len(segments)) + return self._probe_write_target_sequential(path, pre_content, want_pre) + head_seg, body = segments + + head_bytes = self._decode_base64_sample(head_seg) + if head_bytes is None: + # No clean base64 on this shell; ask the way we used to. + logger.debug( + "write_file: no usable base64 head for %s; paying one extra " + "round-trip for the BOM probe", path) + has_bom = self._file_has_bom(path, pre_content) + else: + has_bom = head_bytes.startswith(_UTF8_BOM.encode("utf-8")) + + if body_mode == "cat" and body: + pre_content = body + if pre_content: + ending = _detect_line_ending(pre_content) + elif body_mode == "sample" and body: + ending = _detect_line_ending(body) + else: + ending = None + return has_bom, pre_content, ending + + def _probe_write_target_sequential(self, path: str, pre_content: Optional[str], want_pre: bool, + ) -> tuple[bool, Optional[str], Optional[str]]: + """Pre-compound form of ``_probe_write_target``: one exec per question. A + failed ``cat`` leaves pre_content None so the lint-delta and LSP consumers + degrade gracefully.""" + if want_pre and pre_content is None: read_result = self._cat(path) if read_result.exit_code == 0 and read_result.stdout: - return read_result.stdout - return None - - def _match_on_disk_conventions(self, path: str, content: str, pre_content: Optional[str]) -> str: - """Re-apply the on-disk file's CRLF endings and UTF-8 BOM to ``content``: - read_file strips the BOM and models send bare-LF text, so a round-trip would - otherwise normalize CRLF files and drop the BOM. The BOM is only prepended - when ``content`` lacks one (guards double-BOM).""" - sample = pre_content - if not sample: + pre_content = read_result.stdout + if pre_content: + ending = _detect_line_ending(pre_content) + else: head = self._head(path, 4096) - sample = head.stdout if head.exit_code == 0 else "" - if _detect_line_ending(sample) == "\r\n": - content = _normalize_line_endings(content, "\r\n") - if self._file_has_bom(path, pre_content) and not _has_bom(content): - content = _UTF8_BOM + content - return content + ending = _detect_line_ending(head.stdout) if head.exit_code == 0 and head.stdout else None + return self._file_has_bom(path, pre_content), pre_content, ending def _verify_written_hash(self, path: str, content_bytes: bytes) -> tuple[Optional[bool], Optional[WriteResult]]: """Compare the on-disk sha256 to the intended bytes: ``(verified, error)``. @@ -814,11 +1167,12 @@ class ShellFileOperations(LintMixin, SearchMixin, FileOperations): """Write content atomically, creating parent directories as needed. Order: deny list → lone-surrogate refusal → fail-closed syntax gate on the - CANDIDATE content (JSON/YAML/TOML) → pre-content capture → CRLF/BOM - preservation → LSP baseline snapshot → atomic write (content rides stdin: - no ARG_MAX limit) → sha256 verification → lint delta → LSP diagnostics - when syntax is clean. ``pre_content``: pre-edit content the caller already - has (saves a ``cat``); BOM detection always probes disk regardless. + CANDIDATE content (JSON/YAML/TOML) → one compound on-disk probe + (pre-content when wanted, CRLF, BOM; see ``_probe_write_target``) → + CRLF/BOM preservation → LSP baseline snapshot → atomic write (content rides + stdin: no ARG_MAX limit) → sha256 verification → lint delta → LSP + diagnostics when syntax is clean. ``pre_content``: pre-edit content the + caller already has (skips the read); BOM detection always probes disk. """ path = self._expand_path(path) denied = get_write_denied_error(path) @@ -832,8 +1186,16 @@ class ShellFileOperations(LintMixin, SearchMixin, FileOperations): if refused is not None: return refused - pre_content = self._capture_pre_content(path, ext, pre_content) - content = self._match_on_disk_conventions(path, content, pre_content) + # Pre-content is read only for extensions in the UNION of in-process lint and + # LSP coverage (keeps the hot path fast for binaries). + want_pre = ext in LINTERS_INPROC or self._lsp_handles_extension(ext) + has_bom, pre_content, original_ending = self._probe_write_target(path, pre_content, want_pre) + # read_file strips the BOM and models send bare-LF text, so a round-trip would + # otherwise normalize CRLF files and drop the BOM (prepend only when absent). + if original_ending == "\r\n": + content = _normalize_line_endings(content, "\r\n") + if has_bom and not _has_bom(content): + content = _UTF8_BOM + content # Best-effort snapshot so the LSP tier reports only this edit's diagnostics. self._snapshot_lsp_baseline(path) # ``dirs_created`` means "parent dirs ensured" (mkdir -p is folded into @@ -956,21 +1318,27 @@ class ShellFileOperations(LintMixin, SearchMixin, FileOperations): def search(self, pattern: str, path: str = ".", target: str = "content", file_glob: Optional[str] = None, limit: int = 50, offset: int = 0, - output_mode: str = "content", context: int = 0) -> SearchResult: + output_mode: str = "content", context: int = 0, + order: str = "discovery") -> SearchResult: """Search for content (regex, ``target="content"``) or files (glob, ``target="files"``). ``output_mode``: "content", "files_only" or "count"; - ``context``: lines of context around matches.""" + ``context``: lines of context around matches; ``order``: file-search + ordering — fast "discovery" or exact "modified" time.""" offset, limit = normalize_search_pagination(offset, limit) + if target == "files" and order not in {"discovery", "modified"}: + return SearchResult( + error=(f"Invalid file search order {order!r}; expected " + "'discovery' or 'modified'.")) path = self._expand_path(path) if "not_found" in self._path_exists_probe(path): # Models often pass several paths in one string: search the parts that exist. multi = self._try_multi_path_search( - pattern, path, target, file_glob, limit, offset, output_mode, context) + pattern, path, target, file_glob, limit, offset, output_mode, context, order) if multi is not None: return multi return self._path_not_found_result(path) result = self._dispatch_search(pattern, path, target, file_glob, limit, offset, - output_mode, context) + output_mode, context, order) exclusions = self._macos_search_exclusions(path) if exclusions and not result.error: skipped = ", ".join(item.split("/")[-1] for item in exclusions) diff --git a/tools/file_operations_common.py b/tools/file_operations_common.py index 47cf676e90..b15c0d36cf 100644 --- a/tools/file_operations_common.py +++ b/tools/file_operations_common.py @@ -145,6 +145,7 @@ class SearchResult: result["counts"] = self.counts if self.truncated: result["truncated"] = True + result["total_count_is_lower_bound"] = True for key in ("limit_reason", "warning", "error"): value = getattr(self, key) if value: diff --git a/tools/file_operations_search.py b/tools/file_operations_search.py index 6d251a813e..ff7502777e 100644 --- a/tools/file_operations_search.py +++ b/tools/file_operations_search.py @@ -5,11 +5,15 @@ """ import os +import posixpath import re import sys +import threading from pathlib import Path -from typing import List, Optional +from typing import Any, List, Optional +from agent.search_policy import SEARCH_PRUNE_DIR_NAMES +from tools import interrupt as tool_interrupt from tools.file_operations_common import ExecuteResult, SearchMatch, SearchResult _MACOS_TCC_PROTECTED_HOME_DIRS = ( @@ -44,6 +48,63 @@ def _macos_protected_search_exclusions( return exclusions +# --- Filename-walk admission: one walk per (backend, root) at a time -------------- + +_FILENAME_SEARCH_ADMISSION = threading.Condition() +_ACTIVE_FILENAME_SEARCH_ROOTS: set[tuple[str, str, str]] = set() +_FILENAME_SEARCH_WAIT_SECONDS = 0.05 + + +def _normalized_filename_search_root(env: Any, root: str, fallback_cwd: str) -> str: + """Normalize a filename-walk root without resolving remote paths locally.""" + from tools.environments.local import LocalEnvironment, _IS_WINDOWS, _msys_to_windows_path + + cwd = getattr(env, "cwd", None) or fallback_cwd + if isinstance(env, LocalEnvironment): + if _IS_WINDOWS: + root = _msys_to_windows_path(root) + cwd = _msys_to_windows_path(cwd) + if not os.path.isabs(root): + root = os.path.join(cwd, root) + return os.path.normcase(os.path.abspath(os.path.normpath(root))) + if not posixpath.isabs(root): + root = posixpath.join(cwd, root) + return posixpath.normpath(root) + + +def _filename_search_root_keys(env: Any, roots: List[str], fallback_cwd: str) -> tuple[tuple[str, str, str], ...]: + """Unique backend/root admission keys in deterministic order.""" + env_type = type(env) + return tuple(sorted({ + (env_type.__module__, env_type.__qualname__, _normalized_filename_search_root(env, root, fallback_cwd)) + for root in roots})) + + +def _acquire_filename_search_roots(keys: tuple[tuple[str, str, str], ...]) -> bool: + """Atomically claim every key, polling for thread-scoped interruption.""" + with _FILENAME_SEARCH_ADMISSION: + while any(key in _ACTIVE_FILENAME_SEARCH_ROOTS for key in keys): + if tool_interrupt.is_interrupted(): + return False + _FILENAME_SEARCH_ADMISSION.wait(_FILENAME_SEARCH_WAIT_SECONDS) + if tool_interrupt.is_interrupted(): + return False + if tool_interrupt.is_interrupted(): + return False + return tool_interrupt.run_if_not_interrupted(lambda: _ACTIVE_FILENAME_SEARCH_ROOTS.update(keys)) + + +def _release_filename_search_roots(keys: tuple[tuple[str, str, str], ...]) -> None: + """Release a completed walk and leave no idle per-root state behind.""" + with _FILENAME_SEARCH_ADMISSION: + _ACTIVE_FILENAME_SEARCH_ROOTS.difference_update(keys) + _FILENAME_SEARCH_ADMISSION.notify_all() + + +_ADMISSION_INTERRUPTED_ERROR = ( + "File search was interrupted while waiting for another filename " + "search on the same root. Retry when ready.") + _SEARCH_TIMEOUT_MARKER_RE = re.compile(r"\n?\[Command timed out after \d+s\]\s*$") @@ -182,14 +243,90 @@ def _parse_search_output(result, output_mode: str, limit: int, offset: int, ) -def _has_hidden_part(parts) -> bool: - return any(part not in {".", ".."} and part.startswith(".") for part in parts) +def _posix_roots(roots: List[str]) -> bool: + """Darwin-only: every root is POSIX-shaped (no drive letter / backslash).""" + return sys.platform == "darwin" and all( + not re.match(r"^[A-Za-z]:[\\/]", root) and "\\" not in root for root in roots) class SearchMixin: """File-name and content search via rg with find/grep fallbacks. Requires ``_exec``, ``_has_command``, ``_expand_path``, ``_escape_shell_arg``, - ``_escape_native_tool_arg``, ``env`` and ``cwd`` from the host class.""" + ``_escape_native_tool_arg``, ``env``, ``cwd``, ``_command_cache``, + ``_rg_resolution_cache`` and ``_rg_modified_capability`` from the host class.""" + + # --- rg resolution -------------------------------------------------------- + + def _resolve_command(self, cmd: str) -> Optional[str]: + """Resolve an executable in the command host's namespace. Ordinary commands + keep the bool hit/miss cache; rg alone caches successful resolved paths and + re-probes misses so a mid-session install becomes visible (with off-PATH + Windows candidates: cargo, scoop, winget).""" + if cmd != "rg": + return cmd if self._has_command(cmd) else None + cached = self._rg_resolution_cache.get(cmd) + if cached: + return cached + result = self._exec("command -v rg 2>/dev/null") + if result.exit_code == 0 and result.stdout.strip(): + resolved = result.stdout.strip().splitlines()[0] + if resolved == "yes": # compatibility with old boolean-probe fakes + resolved = "rg" + self._rg_resolution_cache[cmd] = resolved + return resolved + from tools.environments.local import LocalEnvironment, _IS_WINDOWS + + if _IS_WINDOWS and isinstance(self.env, LocalEnvironment): + user_profile = os.environ.get("USERPROFILE") or str(Path.home()) + local_app_data = os.environ.get("LOCALAPPDATA") + scoop = os.environ.get("SCOOP") or os.path.join(user_profile, "scoop") + candidates = [ + os.path.join(user_profile, ".cargo", "bin", "rg.exe"), + os.path.join(scoop, "shims", "rg.exe"), + ] + if local_app_data: + candidates.append(os.path.join(local_app_data, "Microsoft", "WinGet", "Links", "rg.exe")) + for candidate in candidates: + if os.path.isfile(candidate): + resolved = candidate.replace("\\", "/") + self._rg_resolution_cache[cmd] = resolved + return resolved + return None + + _RG_VERSION_RE = re.compile( + r"(?m)^ripgrep\s+((?:0|[1-9]\d*))\." + r"(?:0|[1-9]\d*)\.(?:0|[1-9]\d*)" + r"(?:-(?:(?:0|[1-9]\d*)|(?:[0-9A-Za-z-]*[A-Za-z-]" + r"[0-9A-Za-z-]*))(?:\.(?:(?:0|[1-9]\d*)|" + r"(?:[0-9A-Za-z-]*[A-Za-z-][0-9A-Za-z-]*)))*)?" + r"(?:\+[0-9A-Za-z-]+(?:\.[0-9A-Za-z-]+)*)?" + r"(?:\s+\(rev [^)]+\))?\s*$") + + def _modified_rg_capability_error(self, executable: str) -> Optional[str]: + """Cached actionable error unless rg can sort exactly (full SemVer, >= 14).""" + if executable in self._rg_modified_capability: + return self._rg_modified_capability[executable] + result = self._exec(f"{self._quote_executable(executable)} --version", timeout=10) + match = self._RG_VERSION_RE.search(result.stdout or "") + if result.exit_code == 0 and match and int(match.group(1)) >= 14: + error = None + else: + error = ("Exact modification-time order requires ripgrep 14 or newer; " + "upgrade ripgrep or use order='discovery'.") + self._rg_modified_capability[executable] = error + return error + + def _quote_executable(self, executable: str) -> str: + """Quote an executable without leaking controller path semantics.""" + if re.fullmatch(r"[A-Za-z0-9_.-]+", executable): + return executable + from tools.environments.local import LocalEnvironment + + if isinstance(self.env, LocalEnvironment): + return self._escape_native_tool_arg(executable) + return "'" + executable.replace("'", "'\"'\"'") + "'" + + # --- macOS protected-folder exclusions -------------------------------------- def _macos_search_exclusions(self, path: str) -> List[str]: """Protected descendants to prune for this search root, if any. Gated on @@ -207,6 +344,41 @@ class SearchMixin: """Absolute-ish protected paths for find's ``-path ... -prune``.""" return [os.path.normpath(os.path.join(path, item)) for item in self._macos_search_exclusions(path)] + def _effective_macos_search_exclusions(self, roots: List[str]) -> List[tuple[str, str, str]]: + """Unique ``(root, relative, absolute)`` exclusions across ``roots``, never + pruning a root the caller chose explicitly.""" + cwd = getattr(self.env, "cwd", None) or self.cwd + use_posix_paths = _posix_roots(roots) + + def normalized(root: str) -> str: + if use_posix_paths: + return posixpath.normpath(root if posixpath.isabs(root) else posixpath.join(cwd, root)) + return os.path.normcase(os.path.abspath(os.path.normpath(root))) + + normalized_roots = [normalized(root) for root in roots] + explicit_roots = set(normalized_roots) + seen = set() + effective = [] + for root, normalized_root in zip(roots, normalized_roots): + for relative in self._macos_search_exclusions(root): + if use_posix_paths: + absolute = key = posixpath.normpath(posixpath.join(normalized_root, relative)) + else: + absolute = os.path.normpath(os.path.join(root, relative)) + key = os.path.normcase(os.path.abspath(absolute)) + if key in explicit_roots or key in seen: + continue + seen.add(key) + effective.append((root, relative, absolute)) + return effective + + @staticmethod + def _macos_protected_search_warning(paths: List[str]) -> str: + skipped = ", ".join(os.path.basename(item) for item in paths) + return ("Skipped macOS protected folders during broad search to avoid " + f"an unattended privacy prompt: {skipped}. Search a protected " + "folder directly when access is intentional.") + def _prune_expr(self, protected_paths: List[str]) -> str: """find ``\\( -path A -o -path B \\) -prune`` clause for the protected dirs.""" terms = " -o ".join(f"-path {self._escape_shell_arg(item)}" for item in protected_paths) @@ -225,9 +397,9 @@ class SearchMixin: def _dispatch_search(self, pattern: str, path: str, target: str, file_glob: Optional[str], limit: int, offset: int, - output_mode: str, context: int) -> SearchResult: + output_mode: str, context: int, order: str = "discovery") -> SearchResult: if target == "files": - return self._search_files(pattern, path, limit, offset) + return self._search_files(pattern, path, limit, offset, order) return self._search_content(pattern, path, file_glob, limit, offset, output_mode, context) def _path_not_found_result(self, path: str) -> SearchResult: @@ -249,11 +421,16 @@ class SearchMixin: def _try_multi_path_search(self, pattern: str, path: str, target: str, file_glob: Optional[str], limit: int, offset: int, - output_mode: str, context: int) -> Optional[SearchResult]: - """Recover a not-found ``path`` that is really several paths in one string - ("dir1 dir2" or comma-separated): search every existing part, merge, and - note skipped parts. None when it doesn't look like a multi-path string.""" - parts = [p for chunk in path.split(",") for p in chunk.split() if p.strip()] + output_mode: str, context: int, + order: str = "discovery") -> Optional[SearchResult]: + """Recover a not-found ``path`` that is really several paths in one string. + Commas explicitly delimit paths (internal spaces preserved); without commas + split on whitespace. Search every existing part, merge, and note skipped + parts. None when it doesn't look like a multi-path string.""" + if "," in path: + parts = [part.strip() for part in path.split(",") if part.strip()] + else: + parts = path.split() if len(parts) < 2: return None existing, missing = [], [] @@ -262,26 +439,47 @@ class SearchMixin: (existing if "exists" in self._path_exists_probe(expanded) else missing).append(expanded) if not existing: return None - merged = SearchResult() - for p in existing: - sub = self._dispatch_search(pattern, p, target, file_glob, limit, offset, output_mode, context) - if sub.error: - continue - merged.matches.extend(sub.matches) - merged.files.extend(sub.files) - merged.counts.update(sub.counts) - merged.total_count += sub.total_count - merged.truncated = merged.truncated or sub.truncated - merged.matches = merged.matches[:limit] - merged.files = merged.files[:limit] + if target == "files": + # One global traversal across roots so modified ordering and pagination + # are exact; root admission wraps the actual rg/find invocation. + merged = self._search_files(pattern, existing, limit, offset, order) + else: + merged = SearchResult() + for root in existing: + sub = self._search_content(pattern, root, file_glob, limit, offset, output_mode, context) + if sub.error: + return sub + merged.matches.extend(sub.matches) + merged.files.extend(sub.files) + merged.counts.update(sub.counts) + merged.total_count += sub.total_count + merged.truncated = merged.truncated or sub.truncated + merged.matches = merged.matches[:limit] + merged.files = merged.files[:limit] note = f"path contained {len(parts)} entries; searched {len(existing)} that exist" if missing: note += "; skipped missing: " + ", ".join(missing[:3]) if len(missing) > 3: note += f" (+{len(missing) - 3} more)" - merged.warning = note + warning_parts = [note] + if not merged.error: + protected_paths = [absolute for _r, _rel, absolute in self._effective_macos_search_exclusions(existing)] + if protected_paths: + warning_parts.append(self._macos_protected_search_warning(protected_paths)) + merged.warning = " ".join(warning_parts) return merged + def _search_prune_glob_args(self) -> str: + """rg globs pruning known heavyweight recursive subtrees. Both forms are + needed: globs are relative to each rg root, so ``**/name/**`` alone misses an + explicitly selected ``name/`` root. Names come from the shared scan policy — + no second search-only list.""" + globs = [] + for dirname in sorted(SEARCH_PRUNE_DIR_NAMES): + for prefix in ("", "**/"): + globs.extend(("--glob", self._escape_shell_arg(f"!{prefix}{dirname}/**"))) + return " ".join(globs) + # (rg flags, message template) probes for a 0-match content search, in order. # The fixed-string probe only runs when the pattern has regex metacharacters. _ZERO_MATCH_PROBES = ( @@ -299,15 +497,23 @@ class SearchMixin: """Steering hint for a 0-match content search, or None: a bare zero gives the model nothing to act on, so run cheap count-only rg probes (case-insensitive, hidden/ignored, fixed-string) and report the first that hits.""" - if not self._has_command('rg'): + rg_executable = self._resolve_command('rg') + if not rg_executable: return None + rg = self._quote_executable(rg_executable) has_meta = bool(re.search(r"[.\[\](){}?*+^$\\|]", pattern)) glob_expr = f" --glob {self._escape_shell_arg(file_glob)}" if file_glob else "" for flags, template in self._ZERO_MATCH_PROBES: if flags == "-F" and not has_meta: continue + # The hidden/ignored probe keeps --no-ignore so project-local ignored + # files stay diagnosable, but prunes heavyweight trees before rg recurses. + if flags.startswith("--hidden"): + glob_expr_probe = f"{glob_expr} {self._search_prune_glob_args()}" + else: + glob_expr_probe = glob_expr probe = self._exec( - f"rg {flags} --count-matches{glob_expr} " + f"{rg} {flags} --count-matches{glob_expr_probe} " f"{self._escape_shell_arg(pattern)} {self._escape_native_tool_arg(path)} " f"2>/dev/null | head -50", timeout=30) @@ -323,82 +529,198 @@ class SearchMixin: return template.format(total=total, n=len(per_file), paths=paths) return None - def _search_files(self, pattern: str, path: str, limit: int, offset: int) -> SearchResult: - """Search for files by name (glob-like): rg --files, else find.""" + def _is_broad_local_search_root(self, path: str) -> bool: + """Whether a no-rg LOCAL root (filesystem root, $HOME or an ancestor of it) is + unsafe for recursive find. Controller paths never classify remotes.""" + from tools.environments.local import LocalEnvironment, _IS_WINDOWS, _msys_to_windows_path + + if not isinstance(self.env, LocalEnvironment): + return False + + def normalized(value: str) -> str: + if _IS_WINDOWS: + value = _msys_to_windows_path(value).replace("\\", "/") + if not os.path.isabs(value): + value = os.path.join(getattr(self.env, "cwd", None) or self.cwd, value) + return os.path.normcase(os.path.abspath(value)) + + from tools import file_operations as _fo # lazy: _HOME is monkeypatched there + root = normalized(path) + home = normalized(_fo._HOME) + drive = os.path.splitdrive(root)[0] + anchor = drive + os.sep if drive else os.path.abspath(os.sep) + if root == os.path.normcase(anchor): + return True + try: + common = os.path.commonpath([root, home]) + except ValueError: + return False + return root == home or common == root + + def _search_files(self, pattern: str, path: str | List[str], limit: int, offset: int, + order: str = "discovery") -> SearchResult: + """Search for files by name (glob-like) across one or more roots: rg --files, + else a bounded find. ``order``: "discovery" (fast, bounded) or "modified" + (exact global newest-first; needs rg 14+ or GNU find).""" search_pattern = pattern if (not pattern.startswith('**/') and '/' not in pattern) \ else pattern.split('/')[-1] - # rg respects .gitignore, skips hidden dirs, and walks in parallel (~200x find). - if self._has_command('rg'): - return self._search_files_rg(search_pattern, path, limit, offset) - if not self._has_command('find'): + roots = [path] if isinstance(path, str) else path + if not roots: + return SearchResult(error="File search requires at least one search root in 'path'.") + + # Prefer ripgrep: bounded parallel traversal with ignore semantics. Resolve + # the engine and exact-order capability BEFORE admission so a queued request + # does not occupy a root while doing command discovery. + if self._has_command("rg"): + rg_executable = self._resolve_command("rg") or "rg" + if order == "modified": + capability_error = self._modified_rg_capability_error(rg_executable) + if capability_error: + return SearchResult(error=capability_error) + keys = _filename_search_root_keys(self.env, roots, self.cwd) + if not _acquire_filename_search_roots(keys): + return SearchResult(error=_ADMISSION_INTERRUPTED_ERROR) + try: + return self._search_files_rg(search_pattern, path, limit, offset, order, + rg_executable=rg_executable) + finally: + _release_filename_search_roots(keys) + + # A local find rooted at/above $HOME or a filesystem root can take minutes and + # prompt on protected paths: refuse before invoking find. + if any(self._is_broad_local_search_root(root) for root in roots): + return SearchResult(error=( + "Broad local file search without ripgrep is disabled because " + "find cannot keep this traversal safely bounded. Install " + "ripgrep or search a narrower directory.")) + if not self._has_command("find"): return SearchResult( error="File search requires 'rg' (ripgrep) or 'find'. " "Install ripgrep for best results: " "https://github.com/BurntSushi/ripgrep#installation") - # Hidden roots: find's path filter would exclude everything, so filter - # descendants (and paginate) in Python. - search_root = Path(path) - has_hidden_path_ancestor = _has_hidden_part(search_root.parts) - hidden_filter_expr = "" if has_hidden_path_ancestor else " -not -path '*/.*'" - pagination_expr = "" if has_hidden_path_ancestor else f" | tail -n +{offset + 1} | head -n {limit}" - # Prune protected dirs BEFORE traversal so macOS never sees an access attempt. - protected_paths = self._protected_prune_paths(path) - prune_expr = f" {self._prune_expr(protected_paths)} -o" if protected_paths else "" - base = (f"find {self._escape_shell_arg(path)}{prune_expr}{hidden_filter_expr} " - f"-type f -name {self._escape_shell_arg(search_pattern)} ") - # BSD find (macOS) has no -printf: retry without the mtime prefix. - lines, limit_reason = self._exec_lines_with_fallback( - f"{base}-printf '%T@ %p\\n' 2>/dev/null | sort -rn{pagination_expr}", - f"{base}2>/dev/null | sort -rn{pagination_expr}") - files = [] - for line in lines: - parts = line.split(' ', 1) - files.append(parts[1] if len(parts) == 2 and parts[0].replace('.', '').isdigit() else line) - if has_hidden_path_ancestor: - normalized_root = search_root.resolve() - def rel_parts(file_path): - try: - return Path(file_path).resolve().relative_to(normalized_root).parts - except ValueError: - return Path(file_path).parts - files = [f for f in files if not _has_hidden_part(rel_parts(f))][offset:offset + limit] - return SearchResult(files=files, total_count=len(files), - truncated=bool(limit_reason), limit_reason=limit_reason) + # Prune hidden descendant dirs (and hidden files, matching rg's default) while + # still allowing an explicitly selected hidden root; dash-prefixed roots get + # ``./`` so find doesn't parse them as options. + find_roots = [f"./{root}" if root.startswith("-") else root for root in roots] + q_roots = [self._escape_shell_arg(root) for root in find_roots] + root_exemptions = "".join(f" ! -path {root}" for root in q_roots) + hidden_prune = f" \\( -type d -name '.*'{root_exemptions} \\) -prune -o" + protected_paths = [absolute for _r, _rel, absolute in self._effective_macos_search_exclusions(roots)] + protected_prune = f" {self._prune_expr(protected_paths)} -o" if protected_paths else "" + fetch_limit = offset + limit + 1 + base = (f"find {' '.join(q_roots)}{protected_prune}{hidden_prune} -type f " + f"! -name '.*' -name {self._escape_shell_arg(search_pattern)}") + if order == "modified": + cmd = "set -o pipefail; " + base + f" -printf '%T@ %p\\n' 2>/dev/null | sort -rn | head -n {fetch_limit}" + else: + cmd = "set -o pipefail; " + base + f" -print 2>/dev/null | head -n {fetch_limit}" - def _exec_lines_with_fallback(self, cmd: str, fallback_cmd: str) -> tuple[List[str], Optional[str]]: - """Non-empty stdout lines of ``cmd``; when it yields nothing (and didn't time - out) run ``fallback_cmd`` instead. Returns ``(lines, limit_reason)``.""" - result = self._exec(cmd, timeout=60) + keys = _filename_search_root_keys(self.env, roots, self.cwd) + if not _acquire_filename_search_roots(keys): + return SearchResult(error=_ADMISSION_INTERRUPTED_ERROR) + try: + result = self._exec(cmd, timeout=60) + finally: + _release_filename_search_roots(keys) stdout, limit_reason = _search_stdout_and_limit(result) - if not stdout.strip() and not limit_reason: - result = self._exec(fallback_cmd, timeout=60) - stdout, limit_reason = _search_stdout_and_limit(result) - return [f for f in stdout.strip().split('\n') if f], limit_reason - def _search_files_rg(self, pattern: str, path: str, limit: int, offset: int) -> SearchResult: - """File-name search via ``rg --files``, mtime-sorted when rg >= 13 supports --sortr.""" + # Parse BEFORE classifying exit 141: under pipefail a bounded producer gets + # SIGPIPE when head closes after fetch_limit rows — benign only when the + # payload proves the bound was reached; a shorter payload is a hard failure. + raw_files: List[str] = [] + for line in stdout.splitlines(): + if order == "modified": + parts = line.split(" ", 1) + if len(parts) != 2 or not parts[0].replace(".", "", 1).isdigit(): + continue + raw_files.append(parts[1]) + elif line: + raw_files.append(line) + bounded_sigpipe = result.exit_code == 141 and len(raw_files) >= fetch_limit + if result.exit_code not in {0, 124} and not bounded_sigpipe: + if order == "modified": + return SearchResult(error=( + "Exact modification-time order requires GNU find with " + "-printf support; install ripgrep 14+ or use order='discovery'.")) + return SearchResult(error="File search failed while running bounded find traversal.") + + from tools.environments.local import LocalEnvironment, _IS_WINDOWS, _msys_to_windows_path + if _IS_WINDOWS and isinstance(self.env, LocalEnvironment): + raw_files = [_msys_to_windows_path(file_path) for file_path in raw_files] + return SearchResult( + files=raw_files[offset:offset + limit], total_count=len(raw_files), + truncated=len(raw_files) > offset + limit or bool(limit_reason), limit_reason=limit_reason) + + def _search_files_rg(self, pattern: str, path: str | List[str], limit: int, offset: int, + order: str = "discovery", rg_executable: Optional[str] = None) -> SearchResult: + """File-name search via ``rg --files`` (respects .gitignore, skips hidden dirs, + parallel walk). Discovery order stays bounded and fast; exact modification-time + ordering is explicit because it scans globally.""" # Wrap bare names so -g matches at any depth (equivalent to find -name). glob_pattern = f"*{pattern}" if ('/' not in pattern and not pattern.startswith('*')) else pattern - fetch_limit = limit + offset - exclusion_globs = " ".join(self._rg_exclusion_globs(path)) + roots = [path] if isinstance(path, str) else path + fetch_limit = limit + offset + 1 + effective_exclusions = self._effective_macos_search_exclusions(roots) + scoped_common = None + command_roots = roots + if len(roots) > 1 and effective_exclusions and _posix_roots(roots): + # Several roots: rg globs are root-relative, so cd to the common ancestor + # and express roots + exclusions relative to it. + cwd = getattr(self.env, "cwd", None) or self.cwd + absolute_roots = [ + posixpath.normpath(root if posixpath.isabs(root) else posixpath.join(cwd, root)) + for root in roots] + scoped_common = posixpath.commonpath(absolute_roots) + command_roots = [posixpath.relpath(root, scoped_common) for root in absolute_roots] + exclusion_terms = [ + f"--glob {self._escape_shell_arg(f'!{posixpath.relpath(absolute, scoped_common)}/**')}" + for _r, _rel, absolute in effective_exclusions] + else: + exclusion_terms = [ + f"--glob {self._escape_shell_arg(f'!{relative}/**')}" + for _r, relative, _abs in effective_exclusions] + exclusion_globs = " ".join(dict.fromkeys(exclusion_terms)) exclusion_args = f" {exclusion_globs}" if exclusion_globs else "" - tail = (f"-g {self._escape_shell_arg(glob_pattern)}{exclusion_args} " - f"{self._escape_native_tool_arg(path)} 2>/dev/null | head -n {fetch_limit}") - # --sortr may have failed on older rg; retry without it. - all_files, limit_reason = self._exec_lines_with_fallback( - f"rg --files --sortr=modified {tail}", f"rg --files {tail}") + rg_executable = rg_executable or self._resolve_command("rg") + if not rg_executable: + return SearchResult(error="File search requires ripgrep (rg).") + if order == "modified": + capability_error = self._modified_rg_capability_error(rg_executable) + if capability_error: + return SearchResult(error=capability_error) + rg = self._quote_executable(rg_executable) + sort_arg = " --sortr=modified" if order == "modified" else "" + root_args = " ".join(self._escape_native_tool_arg(root) for root in command_roots) + cd_prefix = f"cd {self._escape_shell_arg(scoped_common)} && " if scoped_common else "" + # ``--`` terminates options so a dash-prefixed root is never parsed as a flag. + cmd = (f"set -o pipefail; {cd_prefix}{rg} --files{sort_arg} -g {self._escape_shell_arg(glob_pattern)}" + f"{exclusion_args} -- {root_args} 2>/dev/null | head -n {fetch_limit}") + result = self._exec(cmd, timeout=60) + stdout, limit_reason = _search_stdout_and_limit(result) + all_files = [f for f in stdout.splitlines() if f] + if scoped_common: + all_files = [ + f if posixpath.isabs(f) else posixpath.normpath(posixpath.join(scoped_common, f)) + for f in all_files] + bounded_sigpipe = result.exit_code == 141 and len(all_files) >= fetch_limit + if result.exit_code not in {0, 1, 124} and not bounded_sigpipe: + if order == "modified": + return SearchResult(error=( + "Exact modification-time order failed; ripgrep 14+ is " + "required. Upgrade ripgrep or use order='discovery'.")) + return SearchResult(error="File search failed while running ripgrep.") return SearchResult( files=all_files[offset:offset + limit], total_count=len(all_files), - truncated=len(all_files) >= fetch_limit or bool(limit_reason), limit_reason=limit_reason, - ) + truncated=len(all_files) > offset + limit or bool(limit_reason), limit_reason=limit_reason) def _search_content(self, pattern: str, path: str, file_glob: Optional[str], limit: int, offset: int, output_mode: str, context: int) -> SearchResult: """Content search: rg, else grep; attaches zero-match steering hints.""" used_rg = self._has_command('rg') if used_rg: - result = self._search_with_rg(pattern, path, file_glob, limit, offset, output_mode, context) + result = self._search_with_rg(pattern, path, file_glob, limit, offset, output_mode, context, + rg_executable=self._resolve_command("rg") or "rg") elif self._has_command('grep'): result = self._search_with_grep(pattern, path, file_glob, limit, offset, output_mode, context) else: @@ -420,20 +742,37 @@ class SearchMixin: return _maybe_warn_line_oriented_newline_pattern(result, pattern) def _run_search_pipeline(self, cmd_parts: List[str], output_mode: str, limit: int, - offset: int, context: int, warning: Optional[str] = None) -> SearchResult: + offset: int, context: int, warning: Optional[str] = None, + line_cap: bool = False) -> SearchResult: """Run ``cmd_parts | head -n `` under pipefail and parse. Extra rows report the true total (context mode also emits "--" separators, so grab 200 more). pipefail keeps the engine's exit 2 alive across ``| head`` - (a truncating head makes rg exit 0 / grep 141, which the ==2 guard ignores).""" + (a truncating head makes rg exit 0 / grep 141, which the ==2 guard ignores). + ``line_cap`` appends ``| cut -c1-2000`` for engines without --max-columns + (grep): bounds giant single-line matches at the pipe layer; skipped for + files_only/count where lines are paths/counts.""" fetch_limit = limit + offset + (200 if context > 0 else 0) - cmd = "set -o pipefail; " + " ".join(cmd_parts + ["|", "head", "-n", str(fetch_limit)]) - result = self._exec(cmd, timeout=60) + parts = cmd_parts + ["|", "head", "-n", str(fetch_limit)] + if line_cap and output_mode not in ("files_only", "count"): + parts += ["|", "cut", "-c1-2000"] + result = self._exec("set -o pipefail; " + " ".join(parts), timeout=60) return _parse_search_output(result, output_mode, limit, offset, context, warning=warning) def _search_with_rg(self, pattern: str, path: str, file_glob: Optional[str], - limit: int, offset: int, output_mode: str, context: int) -> SearchResult: + limit: int, offset: int, output_mode: str, context: int, + rg_executable: Optional[str] = None) -> SearchResult: """Search using ripgrep.""" - cmd_parts = ["rg", "--line-number", "--no-heading", "--with-filename"] + rg_executable = rg_executable or self._resolve_command("rg") + if not rg_executable: + return SearchResult(error="Content search requires ripgrep (rg).") + cmd_parts = [self._quote_executable(rg_executable), "--line-number", "--no-heading", "--with-filename"] + # Giant-single-line containment (cline#13525): a match inside a multi-MB + # single-line dump makes rg emit the ENTIRE line (``head -n`` counts lines). + # --max-columns bounds each printed line at the rg layer; --max-columns-preview + # keeps a truncated prefix so the model still sees the hit. 2000 cols exceeds + # the 500-char content clamp, so nothing previously visible is lost. + if output_mode not in ("files_only", "count"): + cmd_parts.extend(["--max-columns", "2000", "--max-columns-preview"]) # A regex \n hard-errors in line-oriented mode; enable -U up front and say so. multiline = _pattern_has_regex_newline(pattern) if multiline: @@ -490,7 +829,7 @@ class SearchMixin: if relative_path not in {"", "."}: search_root += f"/{self._escape_shell_arg(relative_path)}" cmd_parts.append(search_root) - return self._run_search_pipeline(cmd_parts, output_mode, limit, offset, context) + return self._run_search_pipeline(cmd_parts, output_mode, limit, offset, context, line_cap=True) def _search_with_grep_pruned(self, pattern: str, path: str, file_glob: Optional[str], limit: int, offset: int, output_mode: str, context: int, @@ -510,4 +849,4 @@ class SearchMixin: if file_glob: find_parts.extend(["-name", self._escape_shell_arg(file_glob)]) find_parts.extend(["-exec", *grep_parts, "{}", "+", "2>/dev/null"]) - return self._run_search_pipeline(find_parts, output_mode, limit, offset, context) + return self._run_search_pipeline(find_parts, output_mode, limit, offset, context, line_cap=True) diff --git a/tools/file_state.py b/tools/file_state.py index b063ec5229..9a0f988a38 100644 --- a/tools/file_state.py +++ b/tools/file_state.py @@ -64,16 +64,29 @@ class FileStateRegistry: self._reads: Dict[str, Dict[str, ReadStamp]] = defaultdict(dict) self._last_writer: Dict[str, Tuple[str, float]] = {} self._path_locks: Dict[str, threading.Lock] = {} + self._path_lock_users: Dict[str, int] = {} self._meta_lock = threading.Lock() # guards _path_locks self._state_lock = threading.Lock() # guards _reads + _last_writer @contextmanager def lock_path(self, resolved: str): - """Per-path lock: threads on the same path serialize, different paths proceed.""" + """Per-path lock: threads on the same path serialize, different paths proceed. + The lock entry is dropped once the last holder/waiter exits.""" with self._meta_lock: lock = self._path_locks.setdefault(resolved, threading.Lock()) - with lock: + self._path_lock_users[resolved] = self._path_lock_users.get(resolved, 0) + 1 + lock.acquire() + try: yield + finally: + lock.release() + with self._meta_lock: + users = self._path_lock_users[resolved] - 1 + if users: + self._path_lock_users[resolved] = users + else: + self._path_lock_users.pop(resolved, None) + self._path_locks.pop(resolved, None) def _stamp(self, task_id: str, resolved: str, mtime: float, now: float, partial: bool) -> None: """Caller holds ``_state_lock``.""" @@ -177,6 +190,11 @@ class FileStateRegistry: with self._state_lock: return list(self._reads.get(task_id, {}).keys()) + def forget_task(self, task_id: str) -> None: + """Release read stamps owned by a task after its lifecycle ends.""" + with self._state_lock: + self._reads.pop(task_id, None) + def clear(self) -> None: """Reset all state. Intended for tests only.""" with self._state_lock: @@ -184,6 +202,7 @@ class FileStateRegistry: self._last_writer.clear() with self._meta_lock: self._path_locks.clear() + self._path_lock_users.clear() _registry = FileStateRegistry() diff --git a/tools/file_tools.py b/tools/file_tools.py index a6e74da744..9053745425 100644 --- a/tools/file_tools.py +++ b/tools/file_tools.py @@ -351,13 +351,30 @@ def _get_file_ops(task_id: str = "default") -> ShellFileOperations: def clear_file_ops_cache(task_id: str = None): - """Clear the file operations cache.""" + """Clear file-operation state for a finished task, or all tasks.""" with _file_ops_lock: if task_id: _file_ops_cache.pop(task_id, None) else: _file_ops_cache.clear() + with _read_tracker_lock: + if task_id: + _read_tracker.pop(task_id, None) + else: + _read_tracker.clear() + + with _patch_failure_lock: + if task_id: + _patch_failure_tracker.pop(task_id, None) + else: + _patch_failure_tracker.clear() + + if task_id: + file_state.get_registry().forget_task(task_id) + else: + file_state.get_registry().clear() + _SPECIAL_FILE_KINDS = ( (stat.S_ISFIFO, "a FIFO (named pipe)"), @@ -480,6 +497,7 @@ def _record_successful_read(task_data: dict, task_id: str, path: str, resolved_s """ with _read_tracker_lock: task_data["dedup_hits"].pop(dedup_key, None) + task_data["dedup_generation_reads"].add(dedup_key) task_data["read_history"].add((path, offset, limit)) count = _bump_consecutive(task_data, ("read", path, offset, limit)) try: @@ -562,10 +580,13 @@ def read_file_tool(path: str, offset: int = 1, limit: int = 2000, task_id: str = dedup_key = (resolved_str, offset, limit) with _read_tracker_lock: task_data = _task_data(task_id) - cached_mtime = task_data.get("dedup", {}).get(dedup_key) + cached_mtime = task_data["dedup"].get(dedup_key) + # First unchanged read after a compaction boundary serves full content + # (the summary may have dropped exact bytes); later ones get the stub. + content_served_in_generation = dedup_key in task_data["dedup_generation_reads"] if cached_mtime is not None: try: - if os.path.getmtime(resolved_str) == cached_mtime: + if os.path.getmtime(resolved_str) == cached_mtime and content_served_in_generation: return _dedup_stub_or_block(task_data, dedup_key, path) except OSError: pass # stat failed — fall through to full read @@ -844,14 +865,15 @@ def patch_tool(mode: str = "replace", path: str = None, old_string: str = None, def search_tool(pattern: str, target: str = "content", path: str = ".", file_glob: str = None, limit: int = 50, offset: int = 0, output_mode: str = "content", context: int = 0, + order: str = "discovery", task_id: str = "default") -> str: """Search for content or files.""" try: offset, limit = normalize_search_pagination(offset, limit) - # Pagination args are part of the key so paging through truncated + # Pagination args (and order) are part of the key so paging through truncated # results doesn't trip the repeated-search guard. - search_key = ("search", pattern, target, str(path), file_glob or "", limit, offset) + search_key = ("search", pattern, target, str(path), file_glob or "", limit, offset, order) with _read_tracker_lock: task_data = _read_tracker.setdefault(task_id, { "last_key": None, "consecutive": 0, "read_history": set()}) @@ -885,7 +907,7 @@ def search_tool(pattern: str, target: str = "content", path: str = ".", result = _get_file_ops(task_id).search( pattern=pattern, path=path, target=target, file_glob=file_glob, - limit=limit, offset=offset, output_mode=output_mode, context=context) + limit=limit, offset=offset, output_mode=output_mode, context=context, order=order) omitted = _filter_read_blocked_search_results(result, task_id) for m in getattr(result, "matches", None) or (): if getattr(m, "content", None): @@ -1066,7 +1088,7 @@ def _is_openai_family_main() -> bool: SEARCH_FILES_SCHEMA = { "name": "search_files", - "description": "Search file contents or find files by name. Use this instead of grep/rg/find/ls in terminal. Ripgrep-backed, faster than shell equivalents. On macOS, broad searches above the user home automatically skip TCC-protected folders (Desktop, Documents, Downloads, Library, Movies, Music, Pictures); target one directly when access is intentional.\n\nContent search (target='content'): Regex search inside files. Output modes: full matches with line numbers, file paths only, or match counts.\n\nFile search (target='files'): Find files by glob pattern (e.g., '*.py', '*config*'). Also use this instead of ls — results sorted by modification time.", + "description": "Search file contents or find files by name. Use this instead of grep/rg/find/ls in terminal. Ripgrep-backed, faster than shell equivalents. On macOS, broad searches above the user home automatically skip TCC-protected folders (Desktop, Documents, Downloads, Library, Movies, Music, Pictures); target one directly when access is intentional.\n\nContent search (target='content'): Regex search inside files. Output modes: full matches with line numbers, file paths only, or match counts.\n\nFile search (target='files'): Find files by glob pattern (e.g., '*.py', '*config*'). Also use this instead of ls. Discovery order is the fast bounded default; exact global newest-first order is an explicit opt-in and may scan the full tree.", "parameters": { "type": "object", "properties": { @@ -1076,6 +1098,7 @@ SEARCH_FILES_SCHEMA = { "file_glob": {"type": "string", "description": "Filter files by pattern in grep mode (e.g., '*.py' to only search Python files)"}, "limit": {"type": "integer", "description": "Maximum number of results to return (default: 50)", "default": 50}, "offset": {"type": "integer", "description": "Skip first N results for pagination (default: 0)", "default": 0}, + "order": {"type": "string", "enum": ["discovery", "modified"], "description": "File-search order: 'discovery' is fast bounded traversal order; 'modified' is exact global newest-first and may scan the full tree; ignored for content", "default": "discovery"}, "output_mode": {"type": "string", "enum": ["content", "files_only", "count"], "description": "Output format for grep mode: 'content' shows matching lines with line numbers, 'files_only' lists file paths, 'count' shows match counts per file", "default": "content"}, "context": {"type": "integer", "description": "Number of context lines before and after each match (grep mode only)", "default": 0} }, @@ -1135,7 +1158,8 @@ def _handle_search_files(args, **kw): return search_tool( pattern=args.get("pattern", ""), target=target, path=args.get("path", "."), file_glob=args.get("file_glob"), limit=args.get("limit", 50), offset=args.get("offset", 0), - output_mode=args.get("output_mode", "content"), context=args.get("context", 0), task_id=tid) + output_mode=args.get("output_mode", "content"), context=args.get("context", 0), + order=args.get("order", "discovery"), task_id=tid) def _read_file_schema_overrides(): diff --git a/tools/file_tools_read_tracking.py b/tools/file_tools_read_tracking.py index 86c639a876..6c2d0fd304 100644 --- a/tools/file_tools_read_tracking.py +++ b/tools/file_tools_read_tracking.py @@ -3,8 +3,10 @@ Process-lifetime state behind read_file/search_files/write_file/patch; ``tools.file_tools`` re-imports every name here. Per task_id ``_read_tracker`` stores: ``last_key``/``consecutive`` (loop detection; reset by any OTHER tool -call), ``read_history`` (diagnostics), ``dedup`` (key -> mtime; cleared on -context compression), ``dedup_hits`` (stub-loop breaker), ``read_timestamps`` +call), ``read_history`` (diagnostics), ``dedup`` (key -> mtime; survives context +compression), ``dedup_generation_reads`` (keys whose full content was served since +the last compaction boundary; cleared on compression so one recovery read returns +full content), ``dedup_hits`` (stub-loop breaker), ``read_timestamps`` (staleness warnings) and ``not_found`` (short-TTL negative cache). Every container is hard-capped (``_cap_read_tracker_data``) so long sessions stay small. """ @@ -44,6 +46,7 @@ def _task_data(task_id: str) -> dict: "last_key": None, "consecutive": 0, "read_history": set()}) for key in ("dedup", "dedup_hits", "read_timestamps"): task_data.setdefault(key, {}) + task_data.setdefault("dedup_generation_reads", set()) return task_data @@ -75,6 +78,7 @@ def _cap_read_tracker_data(task_data: dict) -> None: ("read_history", _READ_HISTORY_CAP), ("dedup", _DEDUP_CAP), ("dedup_hits", _DEDUP_CAP), + ("dedup_generation_reads", _DEDUP_CAP), ("read_timestamps", _READ_TIMESTAMPS_CAP), ("not_found", _NOT_FOUND_CAP)): container = task_data.get(key) @@ -141,17 +145,21 @@ def _bump_consecutive(task_data: dict, key: tuple) -> int: def reset_file_dedup(task_id: str = None): - """Clear the read-dedup cache (one task, or all when ``task_id`` is None). Called - after context compression: a "file unchanged" stub would point at summarised-away content.""" + """Advance the read-dedup generation after context compression (one task, or all + when ``task_id`` is None). The per-key ``dedup`` mtime map is PRESERVED so unchanged + files keep returning stubs instead of re-bloating the reclaimed context; the + generation-read set is cleared so the FIRST unchanged read of each key after + compaction returns full content the summary may have dropped. Stub-hit counters + are cleared so the hard block restarts fresh.""" with _read_tracker_lock: if task_id: targets = [_read_tracker[task_id]] if _read_tracker.get(task_id) else [] else: targets = list(_read_tracker.values()) for task_data in targets: - for key in ("dedup", "dedup_hits"): - if key in task_data: - task_data[key].clear() + if "dedup_hits" in task_data: + task_data["dedup_hits"].clear() + task_data.setdefault("dedup_generation_reads", set()).clear() def notify_other_tool_call(task_id: str = "default"): diff --git a/tools/interrupt.py b/tools/interrupt.py index 814fef6b77..45b5ad0306 100644 --- a/tools/interrupt.py +++ b/tools/interrupt.py @@ -6,6 +6,7 @@ is_interrupted(), which checks the CURRENT thread.""" import logging import os import threading +from collections.abc import Callable logger = logging.getLogger(__name__) @@ -54,6 +55,21 @@ def is_thread_interrupted(thread_id: int | None) -> bool: return thread_id in _interrupted_threads +def run_if_not_interrupted(callback: Callable[[], None]) -> bool: + """Run a state transition atomically with current-thread interruption. + + Returns ``False`` without calling ``callback`` when the current thread is + already interrupted. The callback runs under the interrupt lock and must + not block or re-enter any interrupt API. + """ + tid = threading.current_thread().ident + with _lock: + if tid in _interrupted_threads: + return False + callback() + return True + + def get_interrupt_reason() -> str | None: """User-safe interrupt cause for the current thread, if known.""" with _lock: diff --git a/tools/mcp_death_supervisor.py b/tools/mcp_death_supervisor.py new file mode 100644 index 0000000000..87302ae685 --- /dev/null +++ b/tools/mcp_death_supervisor.py @@ -0,0 +1,216 @@ +#!/usr/bin/env python3 +"""One parent-death supervisor per Hermes process, shared by all stdio MCP servers. + +Why this exists +--------------- +When Hermes dies without running its cleanup path (SIGKILL, OOM killer, a hard +crash), stdio MCP servers it spawned are reparented to init and keep running +forever. macOS has no ``PR_SET_PDEATHSIG``, so something has to outlive Hermes +and reap them. + +This module is deliberately standard-library-only and must not import anything +from ``tools/``: it runs after Hermes may already be dead, and pulling in +``mcp_tool`` would drag the whole agent with it. The TERM -> grace -> KILL +``killpg`` sweep in ``_reap`` therefore duplicates similar sweeps elsewhere in +the tree on purpose. + +The predecessor (``mcp_stdio_watchdog.py``) solved this with one CPython +*per MCP server*, wrapping each server command and polling ``getppid()`` every +two seconds. That costs ~10 MB of resident memory per server and detects death +up to one poll interval late. This module replaces the whole fleet of pollers +with a single supervisor per Hermes process: + +* **Death detection is a blocking read on a pipe.** Hermes holds the only write + end. When Hermes dies -- by any means, including SIGKILL -- the write end + closes and the read returns EOF. Exact, instant, and free. +* **Servers are spawned unwrapped.** The MCP SDK already spawns stdio children + with ``start_new_session=True``, so each one is its own process-group leader + and ``killpg`` still reaches its descendants. Removing the wrapper also + removes the signal-forwarding layer the wrapper needed to avoid inverting the + bug it fixed. + +Protocol (line-based, on stdin) +------------------------------- + register \n start reaping this process group on parent death + unregister \n stop reaping it (its server shut down cleanly) + +On EOF the supervisor SIGTERMs every still-registered process group, waits a +short grace period, SIGKILLs the survivors, and exits. A registered group that +Hermes never unregistered *is* the orphan set, so a clean Hermes shutdown -- +which unregisters as it tears each server down -- ends with nothing to kill. + +Unparseable lines are ignored rather than fatal: a corrupted byte on the control +pipe must not cost us the reaping guarantee for every other server. + +Residual risk: process-group reuse +---------------------------------- +We reap by pgid, so a registration is only as meaningful as the group's +identity. A group we deliberately keep registered -- an orphan that teardown +failed to kill, such as the ``node`` ``mcp-remote`` leaves behind -- can +eventually exit on its own, after which the kernel is free to hand that pgid to +an unrelated process owned by the same user. If Hermes then dies ungracefully +while the registration is still stale, we would signal a stranger. +``_is_safe_target`` cannot catch this: the value is stale, not invalid. + +Two things narrow the window. Hermes prunes registrations whose group has no +members left (``_prune_dead_supervised_pgids``) on every registration change, +and the orphan sweep unregisters whatever it reaps. Neither closes it -- a +group can die and its pgid be recycled between two probes -- so the exposure is +real but bounded to that gap, and requires an ungraceful death inside it. + +Closing it completely means proving group identity at reap time, e.g. stamping +MCP children with a boot-unique env marker and checking that some member still +carries it before signalling. That was judged not worth putting a ``ps`` parse +into the one process whose job is to stay simple enough to always work; it is +the obvious next step if this class of bug ever actually bites. Note the same +exposure already exists in Hermes's own killpg-based orphan cleanup, which this +module did not introduce (see upstream issue #88350). +""" + +from __future__ import annotations + +import argparse +import os +import signal +import sys +import time + +# Matches the grace period the per-server watchdog used before it escalated. +_TERM_GRACE_S = 3.0 +# How often we re-check for survivors during that grace period. +_REAP_POLL_S = 0.1 +# A command is "unregister " -- around 20 characters. The cap only has to +# be generous enough for a legitimate line; see _serve for why it exists. +_MAX_LINE_CHARS = 256 + + +def _is_safe_target(pgid: int, *, own_pgid: int, parent_pgid: int) -> bool: + """Return True if ``pgid`` is a process group we may signal. + + Defensive only -- Hermes already filters non-MCP children before it + registers anything (see ``_filter_mcp_children`` in ``tools/mcp_tool.py``). + But this process signals whole process *groups*, so a bad value here is + unusually expensive: ``killpg(0, ...)`` signals our own group, and pgid 1 + is init. A caller bug should cost us one unreaped server, never the + Hermes process tree or the session. + """ + if pgid <= 1: + return False + if pgid == own_pgid or pgid == parent_pgid: + return False + return True + + +def _reap(pgids: set[int]) -> None: + """SIGTERM every group, then SIGKILL whatever is still alive. + + Every process-group call below is POSIX-only by construction: this whole + module only ever runs as a child of ``_update_death_supervisor``, which + returns early unless ``os.name == "posix"``, so the supervisor is never + spawned on Windows in the first place. + """ + if not pgids: + return + + alive = set() + for pgid in pgids: + try: + os.killpg(pgid, signal.SIGTERM) # windows-footgun: ok — POSIX-only process + alive.add(pgid) + except (ProcessLookupError, PermissionError, OSError): + # Already gone, or not ours to signal. Either way, nothing to reap. + pass + + deadline = time.monotonic() + _TERM_GRACE_S + while alive and time.monotonic() < deadline: + time.sleep(_REAP_POLL_S) + for pgid in list(alive): + try: + # Signal 0 probes liveness: succeeds iff some member survives. + os.killpg(pgid, 0) # windows-footgun: ok — POSIX-only process + except (ProcessLookupError, PermissionError, OSError): + alive.discard(pgid) + + for pgid in alive: + try: + os.killpg(pgid, signal.SIGKILL) # windows-footgun: ok — POSIX-only + except (ProcessLookupError, PermissionError, OSError): + pass + + +def _serve(stream, *, own_pgid: int, parent_pgid: int) -> set[int]: + """Read control lines until EOF; return the groups still registered. + + Reads are length-capped rather than newline-terminated. Iterating the + stream instead lets a writer that never sends a newline grow this process + without bound -- feeding it ``/dev/zero`` reached 15 GB before it was + stopped. Nothing in Hermes can produce that today, but this process is the + last line of defense against leaked servers, so it must not be the thing + that dies under memory pressure. A line truncated by the cap fails to parse + and is skipped; the remainder resyncs at the next newline. + """ + registered: set[int] = set() + while True: + line = stream.readline(_MAX_LINE_CHARS) + if not line: + break # EOF: the parent is gone. + if not line.endswith("\n"): + # Truncated by the cap, or an unterminated tail at EOF. Either way + # it is not a command we are willing to act on. + continue + parts = line.split() + if len(parts) != 2: + continue + verb, raw = parts + try: + pgid = int(raw) + except ValueError: + continue + if verb == "register": + if _is_safe_target(pgid, own_pgid=own_pgid, parent_pgid=parent_pgid): + registered.add(pgid) + elif verb == "unregister": + registered.discard(pgid) + return registered + + +def main(argv=None) -> int: + parser = argparse.ArgumentParser( + description="Reap registered process groups when the parent dies." + ) + parser.add_argument( + "--parent-pgid", + type=int, + required=True, + help="Process group of the spawning Hermes process; never signalled.", + ) + args = parser.parse_args(argv) + + # The parent may be torn down with killpg on its own group. We are spawned + # with start_new_session=True precisely so that sweep cannot take us with + # it before we have reaped -- assert that here rather than trust the caller. + own_pgid = os.getpgid(0) + if own_pgid == args.parent_pgid: + print( + "mcp_death_supervisor: refusing to run inside the parent's process " + "group (a killpg of the parent would kill us before we can reap)", + file=sys.stderr, + ) + return 2 + + # A dying parent's SIGINT/SIGHUP must not preempt the reap; the pipe's EOF + # is our only shutdown signal. SIGHUP is POSIX-only, which is fine here -- + # this process is never spawned on Windows (see _reap's docstring). + for sig in (signal.SIGINT, signal.SIGHUP): # windows-footgun: ok — POSIX-only process + try: + signal.signal(sig, signal.SIG_IGN) + except (ValueError, OSError): + pass + + registered = _serve(sys.stdin, own_pgid=own_pgid, parent_pgid=args.parent_pgid) + _reap(registered) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tools/mcp_stdio_watchdog.py b/tools/mcp_stdio_watchdog.py deleted file mode 100644 index 17ca4e2a9e..0000000000 --- a/tools/mcp_stdio_watchdog.py +++ /dev/null @@ -1,97 +0,0 @@ -#!/usr/bin/env python3 -"""Parent-death watchdog supervisor for stdio MCP subprocesses. - -If Hermes dies hard (kill -9, crash) its graceful teardown never runs and the stdio child -plus its descendants are orphaned (macOS has no ``PR_SET_PDEATHSIG``); piled-up orphans then -race the new connection for the same upstream session. So the MCP command is spawned via this -supervisor, which (1) runs the real command in a new process group so the whole tree can be -killpg'd, (2) passes stdin/stdout/stderr straight through — the MCP stdio protocol talks over -those pipes, so this is a no-op relay, not a proxy — and (3) polls ``getppid()`` and, once the -parent is gone, SIGTERMs the child's group, waits, then SIGKILLs. Stdlib only so it starts -fast and cannot itself leak. Usage: ``mcp_stdio_watchdog.py --ppid -- `` -""" - -from __future__ import annotations - -import argparse -import os -import signal -import subprocess -import sys -import threading -import time - -_POLL_INTERVAL_S = 2.0 -_TERM_GRACE_S = 3.0 - - -def _is_orphaned(original_ppid: int, getppid=os.getppid) -> bool: - """Return whether this process no longer has its original POSIX parent.""" - return getppid() != original_ppid - - -def _terminate_process_group(proc: subprocess.Popen) -> None: - """Best-effort SIGTERM-then-SIGKILL of the child's process group; guards the POSIX-only - primitives so an accidental Windows run degrades to a plain child kill.""" - killpg = getattr(os, "killpg", None) - if killpg is None: # windows-footgun: ok — non-POSIX fallback - try: - proc.terminate() - proc.wait(timeout=_TERM_GRACE_S) - except (OSError, subprocess.TimeoutExpired): - proc.kill() - return - try: - pgid = os.getpgid(proc.pid) - except (ProcessLookupError, OSError): - return - for sig in (signal.SIGTERM, getattr(signal, "SIGKILL", signal.SIGTERM)): - try: - killpg(pgid, sig) - except (ProcessLookupError, PermissionError, OSError): - return - try: - proc.wait(timeout=_TERM_GRACE_S) - return - except subprocess.TimeoutExpired: - continue - - -def _watchdog_loop(proc: subprocess.Popen, original_ppid: int) -> None: - while proc.poll() is None: - if _is_orphaned(original_ppid): - _terminate_process_group(proc) - return - time.sleep(_POLL_INTERVAL_S) - - -def main(argv: list[str] | None = None) -> int: - parser = argparse.ArgumentParser(description="Parent-death watchdog for a stdio MCP subprocess.") - parser.add_argument("--ppid", type=int, required=True) - parser.add_argument("command", nargs=argparse.REMAINDER) - args = parser.parse_args(argv) - real_argv = args.command[1:] if args.command[:1] == ["--"] else list(args.command) - if not real_argv: - print("mcp_stdio_watchdog: no command given after '--'", file=sys.stderr) - return 2 - # New process group: killpg() reaches the whole tree the real command may spawn without - # touching our own group or the original parent's. - proc = subprocess.Popen(real_argv, stdin=sys.stdin, stdout=sys.stdout, stderr=sys.stderr, start_new_session=True) - # The server lives in its OWN group, so the parent's shutdown killpg of *our* group no - # longer reaches it: forward SIGTERM/SIGINT to the child's group so graceful teardown - # still kills a wedged server that ignores stdin EOF. - def _forward_shutdown(signum, frame): # noqa: ARG001 - _terminate_process_group(proc) - sys.exit(128 + signum) - signal.signal(signal.SIGTERM, _forward_shutdown) - signal.signal(signal.SIGINT, _forward_shutdown) - threading.Thread(target=_watchdog_loop, args=(proc, args.ppid), daemon=True).start() - try: - return proc.wait() - except KeyboardInterrupt: - _terminate_process_group(proc) - return 130 - - -if __name__ == "__main__": - sys.exit(main()) diff --git a/tools/mcp_tool.py b/tools/mcp_tool.py index 20284764a2..31a68debc2 100644 --- a/tools/mcp_tool.py +++ b/tools/mcp_tool.py @@ -15,8 +15,9 @@ import importlib import importlib.util import inspect import logging -import os # noqa: F401 — tests patch ``tools.mcp_tool.os.*`` +import os import shutil # noqa: F401 — tests patch ``tools.mcp_tool.shutil.which`` +import sys import threading import time from typing import Any, Callable, Dict, List, Optional, Set @@ -32,8 +33,8 @@ from tools.mcp_tool_schema import ( # noqa: F401 from tools.mcp_tool_content import ( # noqa: F401 _MCP_HARD_RESULT_CAP_CHARS, _MCP_RESOURCE_MAX_B64_CHARS, _MCP_RESOURCE_MAX_BYTES, _cache_mcp_audio_block, _cache_mcp_image_block, _is_reserved_mcp_meta_key, - _mcp_image_extension_for_mime_type, _mcp_resource_filename, _render_mcp_resource_block, - _truncate_mcp_text_result) + _mcp_image_extension_for_mime_type, _mcp_resource_filename, _render_mcp_dropped_block_notice, + _render_mcp_resource_block, _truncate_mcp_text_result) from tools.mcp_tool_errors import ( # noqa: F401 InvalidMcpUrlError, NonMcpEndpointError, _EXC_TRAVERSAL_MAX_NODES, _JSONRPC_UNSUPPORTED_PROTOCOL_VERSION, _classify_mcp_failure, _format_connect_error, @@ -43,7 +44,7 @@ from tools.mcp_tool_errors import ( # noqa: F401 from tools.mcp_tool_config import ( # noqa: F401 _ENV_VAR_PATTERN, _build_safe_env, _filter_suspicious_mcp_servers, _get_mcp_stderr_log, _interpolate_env_vars, _load_mcp_config, _resolve_stdio_command, _warn_hidden_whitespace, - _whitespace_warned, _workspace_folder, _wrap_command_with_watchdog, _write_stderr_log_header) + _npx_bin_candidates, _npx_cached_bin, _whitespace_warned, _workspace_folder, _write_stderr_log_header) from tools.mcp_tool_sampling import ( # noqa: F401 ElicitationHandler, SamplingHandler, _format_elicitation_schema_summary) from tools.mcp_tool_handlers import ( # noqa: F401 @@ -82,6 +83,36 @@ from tools.mcp_tool_discovery import ( # noqa: F401 _OSV_MALWARE_CHECK_TIMEOUT_S = 12.0 +async def _preflight_stdio_command(server_name: str, command: str, args: list) -> tuple[str, list]: + """OSV malware preflight (off-loop, wall-clock bound, fail-open on timeout), THEN the + cached-npx swap. The preflight must see the REAL command/args: anything that rewrites argv to a + wrapper or resolved binary has to happen after it, or the check silently inspects the wrapper + and becomes a no-op (``_infer_ecosystem`` keys off the command basename being npx/uvx/pipx).""" + from tools.osv_check import check_package_for_malware + try: + malware_error = await asyncio.wait_for( + asyncio.to_thread(check_package_for_malware, command, args), timeout=_OSV_MALWARE_CHECK_TIMEOUT_S) + except asyncio.TimeoutError: + logger.warning("MCP server '%s': OSV malware preflight timed out after %.0fs " + "(network slow/unreachable) — proceeding without the check.", + server_name, _OSV_MALWARE_CHECK_TIMEOUT_S) + malware_error = None + if malware_error: + raise ValueError(f"MCP server '{server_name}': {malware_error}") + + # npx resolves the package and then FORKS, staying resident as the real server's parent for + # nothing (~48 MB per server, measured). Hermes already supervises the child (shared death + # supervisor), so a cached package is spawned directly; a cache miss leaves npx untouched. + if os.path.basename(command).lower().startswith("npx"): + cached = _npx_cached_bin(args) + if cached: + direct_command, direct_args = cached + logger.debug("MCP server '%s': using cached npx binary %s (skipping the " + "resident `npm exec` parent)", server_name, direct_command) + command, args = direct_command, direct_args + return command, args + + # ---- Optional MCP SDK: availability probe now, symbol import on first use ---- _MCP_AVAILABLE = _MCP_HTTP_AVAILABLE = _MCP_NEW_HTTP = _MCP_LEGACY_HTTP = False @@ -451,6 +482,129 @@ _mcp_thread: Optional[threading.Thread] = None _lock = threading.Lock() +# ---- Shared parent-death supervisor (state lives HERE: tests rebind ``_death_supervisor``) ---- +# If this process dies without running its cleanup path (kill -9, OOM, crash, force-quit), stdio +# MCP children reparent to init and run forever; macOS has no PR_SET_PDEATHSIG, so something has +# to outlive us and reap them. ONE supervisor process serves all stdio servers and is told which +# process groups to reap over a pipe; it detects our death as EOF on that pipe (exact, instant) +# rather than polling getppid(). Replaced the per-server watchdog wrapper (~10 MB resident per +# server, plus a signal-forwarding layer because wrapping put the server in a different session +# from the pgid tracked for killpg). See tools/mcp_death_supervisor.py. POSIX-only, matching the +# killpg-based orphan cleanup below. +_death_supervisor = None # Optional[subprocess.Popen] +_death_supervisor_lock = threading.Lock() +# Groups the supervisor is reaping on our behalf; replayed verbatim on respawn so a respawn never +# silently drops coverage for servers that are still running. +_supervised_pgids: set = set() + + +def _spawn_death_supervisor(): + """Start the shared supervisor, or None if it cannot be started.""" + import subprocess + supervisor = os.path.join(os.path.dirname(os.path.abspath(__file__)), "mcp_death_supervisor.py") + try: + # start_new_session=True is load-bearing: shutdown paths killpg this process's own group, + # which would kill the supervisor before it could reap anything. + return subprocess.Popen( + [sys.executable, supervisor, "--parent-pgid", str(os.getpgid(0))], + stdin=subprocess.PIPE, stdout=subprocess.DEVNULL, stderr=_get_mcp_stderr_log(), + start_new_session=True, close_fds=True, text=True) + except Exception: + # Never let supervisor bookkeeping block a real MCP connection: graceful shutdown paths + # still reap normally; only the ungraceful-exit safety net is lost. + logger.debug("Could not start the MCP parent-death supervisor", exc_info=True) + return None + + +def _prune_dead_supervised_pgids() -> set: + """Forget supervised groups with no members left; return what went. Caller holds + ``_death_supervisor_lock``. Signal 0 is a pure existence probe (cannot terminate anything). + It narrows, but cannot close, the window where a dead group's pgid is recycled before we + notice (residual-risk note in ``tools/mcp_death_supervisor.py``).""" + killpg = getattr(os, "killpg", None) + if killpg is None: # windows-footgun: ok - POSIX-only, guarded + return set() + stale = set() + for pgid in list(_supervised_pgids): + try: + killpg(pgid, 0) + except ProcessLookupError: + stale.add(pgid) + except (PermissionError, OSError): + # Exists but not ours to signal, or the probe failed: keep it — dropping coverage on + # an ambiguous answer is the more expensive mistake. + pass + _supervised_pgids.difference_update(stale) + return stale + + +def _update_death_supervisor(verb: str, pgids) -> None: + """Register or unregister process groups (``verb`` is ``"register"``/``"unregister"``) with + the shared supervisor. Failures are swallowed: losing the safety net must never fail a live + MCP session.""" + if os.name != "posix": + return + wanted = {int(pgid) for pgid in pgids} + if not wanted: + return + + global _death_supervisor + with _death_supervisor_lock: + if verb == "register": + _supervised_pgids.update(wanted) + else: + _supervised_pgids.difference_update(wanted) + + # A registration outlives the server only while some member survives (e.g. an orphaned + # grandchild teardown failed to kill, deliberately kept registered). Once that group is + # empty its pgid can be recycled by a stranger, so prune here too — the orphan sweep + # unregisters what it reaps but is not guaranteed to run in a given process. + stale = _prune_dead_supervised_pgids() + + proc = _death_supervisor + if proc is None or proc.poll() is not None: + if not _supervised_pgids: + # Nothing left to cover: nothing to tell and nothing to respawn for. Keyed on + # the SET, not the verb: after a broken-pipe write dropped the supervisor with + # groups still registered, an unregister must still rebuild coverage for the + # survivors. + return + proc = _spawn_death_supervisor() + _death_supervisor = proc + if proc is None: + return + # A fresh supervisor knows nothing: replay live coverage (already reflects this + # call's mutation and the prune, so pruned groups never reach the replacement). + payload = "".join(f"register {pgid}\n" for pgid in _supervised_pgids) + else: + payload = "".join(f"{verb} {pgid}\n" for pgid in wanted) + payload += "".join(f"unregister {pgid}\n" for pgid in stale) + + try: + proc.stdin.write(payload) + proc.stdin.flush() + except (BrokenPipeError, ValueError, OSError): + # It exited between poll() and write(). Drop it so the next call respawns and replays + # from ``_supervised_pgids`` (the set, not the pipe, is the record of what needs reaping). + _death_supervisor = None + return + + if not _supervised_pgids: + # Nothing left to reap: release the supervisor rather than keep a ~15 MB process and a + # pipe resident for the life of a gateway. Closing our write end is the same EOF parent + # death sends; with an empty set it exits. The next register respawns and replays. + try: + proc.stdin.close() + except (BrokenPipeError, ValueError, OSError): + pass + # Reap it, or the exited supervisor stays a zombie until the next Popen in this process. + try: + proc.wait(timeout=5) + except Exception: # noqa: BLE001 - timeout or already gone; either way we drop it + pass + _death_supervisor = None + + def _mcp_registry_scope() -> Optional[str]: """Registry scope for MCP registrations: a profile overlay under a multiplexer, else None.""" from agent.secret_scope import is_multiplex_active diff --git a/tools/mcp_tool_config.py b/tools/mcp_tool_config.py index 0e39bf1a50..2c5c39d32f 100644 --- a/tools/mcp_tool_config.py +++ b/tools/mcp_tool_config.py @@ -1,7 +1,8 @@ """MCP server config loading and stdio launch environment: ${VAR}/Cursor-style interpolation, hidden-whitespace and suspicious-entry filtering, the filtered -subprocess env, command resolution, watchdog wrapping and the shared stderr log.""" +subprocess env, command resolution, the cached-npx binary shortcut and the shared stderr log.""" +import json import logging import os import re @@ -174,18 +175,81 @@ def _resolve_stdio_command(command: str, env: dict) -> tuple[str, dict]: return resolved_command, resolved_env -def _wrap_command_with_watchdog(command: str, args: list) -> tuple[str, list]: - """Wrap a stdio command in the parent-death watchdog (POSIX only — it relies on process - groups, same scope as the killpg-based orphan cleanup). Unchanged on non-POSIX or if the - PID cannot be read — watchdog bookkeeping must never block a connection.""" - if os.name != "posix": - return command, args +def _npx_bin_candidates(bin_dir: str, name: str, *, windows: Optional[bool] = None) -> list: + """Launcher paths to try for *name* inside an npx cache's ``.bin``, in order. On Windows that + directory holds the extensionless sh script plus ``.cmd``/``.ps1``; the sh one + cannot be spawned there and ``os.access(X_OK)`` is only an existence check, so select by + extension (same precedence as ``hermes_constants._candidate_node_command_names``). ``windows`` + is injectable so the branch is testable without patching ``os.name`` process-wide.""" + is_windows = os.name == "nt" if windows is None else windows + if is_windows: + return [os.path.join(bin_dir, name + ext) for ext in (".cmd", ".exe")] + return [os.path.join(bin_dir, name)] + + +def _npx_cached_bin(args: list) -> Optional[tuple]: + """Resolve ``npx -y `` to the already-installed binary, or None. + + ``npx`` resolves the package and then FORKS, staying resident as the real server's parent + for nothing (~48 MB private memory per MCP server, measured); Hermes already supervises the + child (shared death supervisor). When the package is in npx's cache we spawn its binary + directly. Deliberately conservative — None (caller keeps plain ``npx``, so a cold machine + still installs) for a cache miss, a version pin (``pkg@1.2.3``), extra npx flags, a manifest + without one obvious bin, or any unreadable cache entry. Returns ``(binary_path, remaining_args)``.""" + if not isinstance(args, list) or not args: + return None + + rest = list(args) + while rest and rest[0] in ("-y", "--yes"): + rest.pop(0) + if not rest: + return None + # `npx pkg -y` (flag AFTER the spec) would hand the server a flag npx would have eaten. + if any(str(a) in ("-y", "--yes") for a in rest[1:]): + return None + + spec = str(rest[0]) + # Scoped names keep their leading '@', so only an '@' AFTER the scope is a version separator. + if "@" in (spec[1:] if spec.startswith("@") else spec): + return None + if not spec or spec.startswith("-"): + return None + + cache_root = os.environ.get("npm_config_cache") or os.path.join(os.path.expanduser("~"), ".npm") + npx_root = os.path.join(cache_root, "_npx") + if not os.path.isdir(npx_root): + return None try: - my_pid = os.getpid() - except Exception: - return command, args - watchdog = os.path.join(os.path.dirname(os.path.abspath(__file__)), "mcp_stdio_watchdog.py") - return sys.executable, [watchdog, "--ppid", str(my_pid), "--", command, *args] + entries = os.listdir(npx_root) + except OSError: + return None + + for entry in entries: + manifest = os.path.join(npx_root, entry, "package.json") + try: + with open(manifest, "r", encoding="utf-8") as fh: + deps = (json.load(fh) or {}).get("dependencies") or {} + except (OSError, ValueError, TypeError): + continue + if spec not in deps: + continue + pkg_json = os.path.join(npx_root, entry, "node_modules", spec, "package.json") + try: + with open(pkg_json, "r", encoding="utf-8") as fh: + bin_field = (json.load(fh) or {}).get("bin") + except (OSError, ValueError, TypeError): + continue + if isinstance(bin_field, str): + names = [os.path.basename(spec)] + elif isinstance(bin_field, dict) and len(bin_field) == 1: + names = list(bin_field.keys()) + else: + continue # zero or several bins: which one npx would pick is not ours to guess + bin_dir = os.path.join(npx_root, entry, "node_modules", ".bin") + for candidate in _npx_bin_candidates(bin_dir, names[0]): + if os.path.exists(candidate) and os.access(candidate, os.X_OK): + return candidate, rest[1:] + return None def _interpolate_env_vars(value): diff --git a/tools/mcp_tool_content.py b/tools/mcp_tool_content.py index 826a144f06..92467bf40f 100644 --- a/tools/mcp_tool_content.py +++ b/tools/mcp_tool_content.py @@ -156,6 +156,28 @@ def _mcp_resource_filename(uri: str, mime_type: str) -> str: return name +def _render_mcp_dropped_block_notice(block, block_type: str) -> str: + """Inline notice for an unsupported MCP content block (kimi-code#3227): silently dropping it + leaves the model unaware content went missing. Carries whatever handles the block exposes — + mime type, uri, size, name — so the agent can fetch or reason about the missing content.""" + details = [f"type={block_type}"] + mime = mcp_field(block, "mime_type", "mimeType", None) + if mime: + details.append(f"mimeType={mime}") + uri = getattr(block, "uri", None) or getattr(getattr(block, "resource", None), "uri", None) + if uri: + details.append(f"uri={uri}") + for size_attr in ("size", "sizeInBytes"): + size = getattr(block, size_attr, None) + if isinstance(size, int): + details.append(f"size={size}") + break + name = getattr(block, "name", None) + if name and isinstance(name, str): + details.append(f"name={name}") + return f"[MCP content dropped: unsupported block ({', '.join(details)})]" + + def _render_mcp_resource_block(block, server_name: str = "") -> str: """Render a ``ResourceLink`` or ``EmbeddedResource`` block as text: embedded text → the text; embedded blob → decoded (size-capped) into the document cache with a path marker; diff --git a/tools/mcp_tool_discovery.py b/tools/mcp_tool_discovery.py index d4e01c62fc..2cf6d09b75 100644 --- a/tools/mcp_tool_discovery.py +++ b/tools/mcp_tool_discovery.py @@ -376,14 +376,30 @@ def _acquire_discovery_lock_with_retry(): return cookie -def discover_mcp_tools() -> List[str]: +def discover_mcp_tools(allowed_mcp_names: Optional[List[str]] = None) -> List[str]: """Entry point: load config, connect servers, register tools. [] without the ``mcp`` - package; idempotent (only servers missing from a previous call are retried).""" + package; idempotent (only servers missing from a previous call are retried). + + ``allowed_mcp_names``: spawn only the MCP servers named in it (built-in toolset names in the + list simply don't match); ``None`` spawns every configured server. Used by + ``hermes -z -t `` to skip cold-starting servers the caller doesn't need (10-60s + each); it only affects which servers start, not which names ``-t`` validation can see.""" servers = _core._load_mcp_config() if not servers: logger.debug("No MCP servers configured") return [] - # SDK import deferred to here so a config without servers never pays it. + if allowed_mcp_names is not None: + allowed_set = {str(n) for n in allowed_mcp_names} + filtered = {name: cfg for name, cfg in servers.items() if name in allowed_set} + if len(filtered) != len(servers): + logger.debug("MCP discovery filter: spawning %d/%d configured server(s) per --toolsets filter " + "(skipped: %s)", len(filtered), len(servers), ",".join(sorted(set(servers) - set(filtered)))) + servers = filtered + if not servers: + logger.debug("No MCP servers in --toolsets filter; skipping MCP load entirely") + return [] + # SDK import deferred to here so a config without servers — or a -t filter that keeps + # none — never pays it. if not _core._ensure_mcp_sdk(): logger.debug("MCP SDK not available -- skipping MCP tool discovery") return [] diff --git a/tools/mcp_tool_handlers.py b/tools/mcp_tool_handlers.py index 1a5695800c..af16b4ed65 100644 --- a/tools/mcp_tool_handlers.py +++ b/tools/mcp_tool_handlers.py @@ -10,13 +10,14 @@ import json import time from contextlib import asynccontextmanager from types import SimpleNamespace -from typing import Any, Callable, Dict, List, Optional +from typing import Any, Callable, Dict, List, Optional, Tuple from tools.registry import tool_error from tools.ansi_strip import strip_unicode_tags from tools.mcp_tool_common import _exc_str, _sanitize_error, mcp_field, _core from tools.mcp_tool_content import ( _MCP_HARD_RESULT_CAP_CHARS, _cache_mcp_audio_block, _cache_mcp_image_block, - _render_mcp_resource_block, _strip_reserved_meta_keys, _truncate_mcp_text_result) + _render_mcp_dropped_block_notice, _render_mcp_resource_block, _strip_reserved_meta_keys, + _truncate_mcp_text_result) from tools.mcp_tool_errors import _is_session_expired_error logger = logging.getLogger("tools.mcp_tool") @@ -321,17 +322,23 @@ def _error_result_text(result) -> str: return "".join(str(t) for t in texts if t) -def _render_content_blocks(result, server_name: str) -> str: +def _render_content_blocks(result, server_name: str) -> Tuple[str, int]: """Text passes through; image/audio blocks are cached (MEDIA: tags); resource blocks are - materialized rather than silently dropped.""" + materialized rather than silently dropped; unsupported blocks become an inline drop notice + (kimi-code#3227). Returns ``(text, usable_parts)`` — the count of REAL rendered blocks + (whitespace-only text and drop notices excluded) that the structuredContent arbitration uses.""" parts: List[str] = [] + usable_parts = 0 for block in (result.content or []): if getattr(block, "text", None): parts.append(strip_unicode_tags(block.text)) + if block.text.strip(): + usable_parts += 1 continue rendered = _cache_mcp_image_block(block) or _cache_mcp_audio_block(block) or _render_mcp_resource_block(block, server_name) if rendered: parts.append(rendered) + usable_parts += 1 continue # Benign empty renders log at debug; warn only for unknown shapes. block_type = getattr(block, "type", None) or type(block).__name__ @@ -339,8 +346,11 @@ def _render_content_blocks(result, server_name: str) -> str: logger.debug("MCP %s: content block type %r rendered empty", server_name, block_type) else: logger.warning("MCP %s: dropping unsupported content block type %r", server_name, block_type) + # Surface the drop to the MODEL, not just the log: a silent drop leaves the agent + # believing the tool returned less than it did, with no way to recover. + parts.append(_render_mcp_dropped_block_notice(block, block_type)) # Hard-cap pathological payloads; ordinary large results pass to spillover. - return _truncate_mcp_text_result("\n".join(parts)) + return _truncate_mcp_text_result("\n".join(parts)), usable_parts def _capped_structured_content(result): @@ -355,13 +365,19 @@ def _capped_structured_content(result): def _render_call_tool_result(result, server_name: str) -> str: - """Pure: ``CallToolResult`` -> handler JSON. ``content`` is primary; ``structuredContent`` - supplements it (or becomes ``result`` without text); ``_meta`` minus reserved keys.""" + """Pure: ``CallToolResult`` -> handler JSON. ``content`` and ``structuredContent`` are + ALTERNATIVES, never both forwarded (kimi-code#3234): spec-following servers already render + their data into content, so forwarding both sent it twice. content wins whenever it rendered + anything usable (no richness heuristic is attempted — none is reliable); structuredContent + fills in only when the blocks rendered effectively empty, keeping structuredContent-only + servers working. ``_meta`` minus reserved keys is always surfaced.""" if mcp_field(result, "is_error", "isError", False): return tool_error(_sanitize_error(_truncate_mcp_text_result(_error_result_text(result) or "MCP tool returned an error"))) - text_result = _render_content_blocks(result, server_name) + text_result, usable_parts = _render_content_blocks(result, server_name) structured = _capped_structured_content(result) meta = _strip_reserved_meta_keys(mcp_field(result, "meta", "meta")) + if structured is not None and usable_parts > 0: + structured = None # drop notices do not count as usable content if structured is None and meta is None: return json.dumps({"result": text_result}, ensure_ascii=False) # Key order is part of the output: "result" leads when there is text, otherwise "_meta" diff --git a/tools/mcp_tool_lifecycle.py b/tools/mcp_tool_lifecycle.py index a8b409b196..fc9ccbf482 100644 --- a/tools/mcp_tool_lifecycle.py +++ b/tools/mcp_tool_lifecycle.py @@ -197,6 +197,9 @@ def _kill_orphaned_mcp_children(include_active: bool = False, server_name: Optio if _pid_exists(pid): # survived SIGTERM _signal_mcp_process(pid, sigkill, owner, pgids.get(pid), my_pgid) logger.warning("Force-killed MCP process %d (%s) after SIGTERM timeout", pid, owner) + # These groups are reaped. Release them last, so a crash partway through the SIGTERM/SIGKILL + # dance still leaves the supervisor holding them. + _core._update_death_supervisor("unregister", pgids.values()) def _stop_mcp_loop_if_idle() -> bool: diff --git a/tools/mcp_tool_transport.py b/tools/mcp_tool_transport.py index 99eef3fd38..b8b702fb69 100644 --- a/tools/mcp_tool_transport.py +++ b/tools/mcp_tool_transport.py @@ -1,5 +1,5 @@ -"""Transport bring-up for MCPServerTask: stdio spawn (OSV preflight, watchdog wrap, child PID -ledger), Streamable HTTP / SSE connect (preflight, identity header, client certs, OAuth), +"""Transport bring-up for MCPServerTask: stdio spawn (OSV preflight, cached-npx swap, child PID +ledger + death-supervisor registration), Streamable HTTP / SSE connect (preflight, identity header, client certs, OAuth), protocol negotiation and initial tool discovery. Split from tools/mcp_tool.py.""" import logging @@ -7,7 +7,6 @@ import asyncio import os from contextlib import asynccontextmanager from typing import Dict, Optional, Set -from tools.mcp_tool_config import _wrap_command_with_watchdog from tools.mcp_tool_errors import NonMcpEndpointError, _apply_identity_header, _handshake_rejected_as_modern, _make_redirect_header_stripper, _resolve_client_cert from tools.mcp_tool_lifecycle import _filter_mcp_children, _orphan_stdio_pid_servers, _orphan_stdio_pids, _stdio_pgids, _stdio_pids from tools.mcp_tool_common import _core @@ -41,22 +40,6 @@ def _pgroup_alive(pgid: Optional[int]) -> bool: return False -async def _osv_malware_preflight(server_name: str, command: str, args: list) -> None: - """OSV malware preflight, off-loop with a wall-clock bound (fail-open on timeout). Must run on - the REAL command/args — the watchdog wrap rewrites argv to the supervisor (check becomes a no-op).""" - from tools.osv_check import check_package_for_malware - try: - malware_error = await asyncio.wait_for( - asyncio.to_thread(check_package_for_malware, command, args), timeout=_core._OSV_MALWARE_CHECK_TIMEOUT_S) - except asyncio.TimeoutError: - logger.warning("MCP server '%s': OSV malware preflight timed out after %.0fs " - "(network slow/unreachable) — proceeding without the check.", - server_name, _core._OSV_MALWARE_CHECK_TIMEOUT_S) - return - if malware_error: - raise ValueError(f"MCP server '{server_name}': {malware_error}") - - class MCPServerTransportMixin: """Methods of :class:`tools.mcp_tool.MCPServerTask` (mixed in; relies on its attributes).""" @@ -153,7 +136,13 @@ class MCPServerTransportMixin: for pid in new_pids: try: new_pgids[pid] = os.getpgid(pid) - except (AttributeError, ProcessLookupError, OSError): # Windows / already exited + except ProcessLookupError: + # Raced and already exited. The SDK spawns with start_new_session=True, so the + # child was its own group leader (pgid == pid): keep that group covered — any + # descendant it left behind still has to be reaped; the prune forgets the group + # once nothing in it is alive. + new_pgids[pid] = pid + except (AttributeError, OSError): # Windows (os.getpgid is POSIX-only) pass with _core._lock: _stdio_pids.update(dict.fromkeys(new_pids, self.name)) @@ -165,11 +154,20 @@ class MCPServerTransportMixin: register_child(_pid, "mcp-helper") except Exception: logger.debug("spawn-ledger register_child failed for MCP helper pid %s", _pid, exc_info=True) + # Hand the pgroups to the shared parent-death supervisor so an ungraceful exit of this + # process (kill -9, crash, force-quit) can't leave this server — or its descendants, e.g. + # mcp-remote's spawned `node` — running forever. The graceful paths (shutdown, + # _kill_orphaned_mcp_children) still reap as before; this only covers when they never run. + _core._update_death_supervisor("register", new_pgids.values()) def _release_spawned_children(self, new_pids: Set[int]) -> None: """Drop the ledger entries; a child (or its pgroup) still alive means SDK teardown failed (common on mid-way cancel on Linux: setsid() children escape) — mark it orphaned for the sweep.""" from gateway.status import _pid_exists + # Groups with nothing left alive; the supervisor forgets them after the lock is released. + # Groups still alive stay registered on purpose, so the supervisor still reaps them if this + # process dies before the orphan sweep runs. + released_pgids: list = [] with _core._lock: for pid in new_pids: _stdio_pids.pop(pid, None) @@ -179,7 +177,10 @@ class MCPServerTransportMixin: _orphan_stdio_pids.add(pid) _orphan_stdio_pid_servers[pid] = self.name else: # nothing to reap — drop the pgid so PID reuse can't surface stale pgroup state - _stdio_pgids.pop(pid, None) + dropped = _stdio_pgids.pop(pid, None) + if dropped is not None: + released_pgids.append(dropped) + _core._update_death_supervisor("unregister", released_pgids) async def _run_stdio(self, config: dict): """Run the server using stdio transport.""" @@ -194,10 +195,8 @@ class MCPServerTransportMixin: if not command: raise ValueError(f"MCP server '{self.name}' has no 'command' in config") command, safe_env = _core._resolve_stdio_command(command, _core._build_safe_env(config.get("env"))) - await _osv_malware_preflight(self.name, command, config.get("args", [])) - # Parent-death watchdog so kill -9 / crash can't leave the child tree running (POSIX-only). - # AFTER the OSV preflight so the check inspects the real package. - command, args = _wrap_command_with_watchdog(command, config.get("args", [])) + # OSV malware preflight, then the cached-npx swap (ordering enforced there). + command, args = await _core._preflight_stdio_command(self.name, command, config.get("args", [])) server_params = _core.StdioServerParameters( command=command, args=args, env=safe_env or None, cwd=config.get("cwd"), # Windows pipes can split non-UTF-8 bytes at chunk boundaries; substitute, don't raise. diff --git a/tools/osv_check.py b/tools/osv_check.py index d7a2591a8b..5c87fadaab 100644 --- a/tools/osv_check.py +++ b/tools/osv_check.py @@ -5,7 +5,6 @@ known malware advisories (MAL-* IDs). Regular CVEs are ignored — only confirme is blocked. Fail-open: network errors allow the package to proceed (~300ms typical). Inspired by Block/goose's extension malware check. """ - import json import logging import os @@ -13,29 +12,135 @@ import re import threading import time import urllib.request +from pathlib import Path from typing import Optional, Tuple - logger = logging.getLogger(__name__) _OSV_ENDPOINT = os.getenv("OSV_ENDPOINT", "https://api.osv.dev/v1/query") _TIMEOUT = 10 # seconds -# Result cache: (ecosystem, package, version) -> (expiry_monotonic, result). Reconnect -# ladders and parked-server self-probes re-run the preflight for the SAME package on every -# spawn; uncached, a flapping server becomes a sustained OSV/DNS query stream. Clean AND -# blocked verdicts are reusable; network failures are NOT cached (fail-open covers them and -# caching one could mask a real advisory later). +# Result cache: (ecosystem, package, version) -> (expiry_wallclock, result). Reconnect +# ladders, parked-server self-probes and repeated `hermes mcp test` runs re-run the preflight +# for the SAME package on every spawn; uncached, a flapping server becomes a sustained OSV/DNS +# query stream. Clean AND blocked verdicts are reusable; network failures are NOT cached +# (fail-open covers them and caching one could mask a real advisory later). +# The cache is also persisted under the Hermes home so separate processes and gateway +# restarts reuse warm verdicts; expiry is absolute wall-clock time so it survives restarts. +# Trade-off: a MAL advisory published right after a clean verdict is noticed at TTL expiry +# (<= 1h by default) rather than at next process start — lower OSV_CHECK_CACHE_TTL to tighten. _CACHE_TTL_S = float(os.getenv("OSV_CHECK_CACHE_TTL", "3600")) _CACHE_MAX_ENTRIES = 256 _cache: dict = {} _cache_lock = threading.Lock() +_disk_cache_loaded = False +_DISK_CACHE_VERSION = 1 + + +def _disk_cache_path() -> Optional[Path]: + """Return the path for the persistent OSV verdict cache. + + Uses ``hermes_constants.get_hermes_home()`` so the cache follows the + active profile and is isolated across Hermes homes. The cache directory + is created on demand. Returns ``None`` when Hermes home cannot be + resolved, in which case only the in-process cache is used. + """ + try: + from hermes_constants import get_hermes_home + + home = get_hermes_home() + except Exception: + return None + try: + cache_dir = home / "cache" + cache_dir.mkdir(parents=True, exist_ok=True) + return cache_dir / "osv_check.json" + except Exception: + return None + + +def _load_disk_cache() -> None: + """Load persistent cache entries from disk into the in-process cache. + + Invoked under ``_cache_lock`` from every get/put but does real work only + once per process (``_disk_cache_loaded`` latch); a transient ``OSError`` + leaves the latch unset so the next call retries. Skips expired or + malformed entries. Only adds missing keys so an in-memory overwrite + (e.g. a test forcing expiry) is not silently reversed by the disk copy. + """ + global _disk_cache_loaded + if _disk_cache_loaded: + return + + path = _disk_cache_path() + if path is None: + _disk_cache_loaded = True + return + + try: + with open(path, "r", encoding="utf-8") as f: + data = json.load(f) + except FileNotFoundError: + data = None + except OSError: + # Transient I/O (file busy, brief permission flap). Retry next call. + return + except Exception: + # Malformed JSON or anything else: unrecoverable, don't spin on it. + data = None + + _disk_cache_loaded = True + if not isinstance(data, dict) or data.get("version") != _DISK_CACHE_VERSION: + return + + now = time.time() + for key_str, entry in data.get("entries", {}).items(): + if not isinstance(entry, dict): + continue + expiry = entry.get("expiry") + result = entry.get("result") + if expiry is None or expiry <= now: + continue + parts = key_str.split("|", 2) + if len(parts) != 3: + continue + key = (parts[0], parts[1], parts[2] or None) + if key not in _cache: + _cache[key] = (expiry, result) + + +def _save_disk_cache() -> None: + """Persist the in-process cache to disk. + + Caller must hold ``_cache_lock`` for consistency. Writes atomically to + a sibling file then renames into place. + """ + path = _disk_cache_path() + if path is None: + return + + entries: dict = {} + for key, (expiry, result) in _cache.items(): + key_str = "|".join(str(k) if k is not None else "" for k in key) + entries[key_str] = {"expiry": expiry, "result": result} + + data = {"version": _DISK_CACHE_VERSION, "entries": entries} + + try: + # Shared atomic writer (temp file + fsync + rename); mkstemp's 0600 + # is kept on create, so verdicts never sit in a world-readable file. + from utils import atomic_write_text + + atomic_write_text(path, json.dumps(data)) + except Exception as exc: + logger.debug("Failed to save OSV disk cache to %s: %s", path, exc) def _cache_get(key) -> Tuple[bool, Optional[str]]: """Return (hit, result) for a fresh cache entry.""" with _cache_lock: + _load_disk_cache() entry = _cache.get(key) - if entry is not None and time.monotonic() < entry[0]: + if entry is not None and time.time() < entry[0]: return True, entry[1] _cache.pop(key, None) # absent or expired return False, None @@ -43,13 +148,15 @@ def _cache_get(key) -> Tuple[bool, Optional[str]]: def _cache_put(key, result: Optional[str]) -> None: with _cache_lock: + _load_disk_cache() if len(_cache) >= _CACHE_MAX_ENTRIES: - now = time.monotonic() + now = time.time() for k in [k for k, (exp, _) in _cache.items() if exp <= now]: del _cache[k] if len(_cache) >= _CACHE_MAX_ENTRIES: _cache.clear() # tiny working set in practice; safe reset - _cache[key] = (time.monotonic() + _CACHE_TTL_S, result) + _cache[key] = (time.time() + _CACHE_TTL_S, result) + _save_disk_cache() def check_package_for_malware(command: str, args: list) -> Optional[str]: diff --git a/tools/process_registry.py b/tools/process_registry.py index 16b0a0861e..a32b0f67fd 100644 --- a/tools/process_registry.py +++ b/tools/process_registry.py @@ -196,6 +196,35 @@ def _build_systemd_scope_argv(shell_argv: List[str], unit_suffix: str) -> List[s return _systemd_scope_argv(binary, f"hermes-worker-{unit_suffix}", *shell_argv) +def restart_safe_gateway_child_argv( + command: List[str], *, unit_suffix: str +) -> List[str]: + """Place a managed-systemd gateway child outside the gateway cgroup. + + Children that must survive an intentional gateway restart cannot rely on + ``start_new_session`` alone: systemd still kills every process in the + service cgroup. In that topology, require a transient user scope and fail + closed if it cannot be established. Standalone processes, non-systemd + supervisors, and non-Linux hosts retain the direct command. + """ + if not _IS_LINUX: + return command + if not _is_supervised_gateway_process() or not os.environ.get("INVOCATION_ID"): + return command + if not _systemd_run_user_scope_available(): + raise RuntimeError( + "cannot create restart-safe systemd scope for gateway child: " + "systemd-run --user --scope is unavailable" + ) + scoped = _build_systemd_scope_argv(command, unit_suffix=unit_suffix) + if scoped == command: + raise RuntimeError( + "cannot create restart-safe systemd scope for gateway child: " + "systemd-run disappeared after the availability probe" + ) + return scoped + + def _stop_systemd_unit(unit_name: str) -> bool: """Stop a transient systemd user scope by unit name. Reaps the *entire* cgroup — catching double-forked descendants reparented to init @@ -967,22 +996,75 @@ class ProcessRegistry: logger.debug("%s wait timed out or failed: %s", label, e) self._finish_exited(session, exit_code()) + @staticmethod + def _log_delta_command(quoted_log_path: str, offset: int) -> str: + """Shell command that reads only the log bytes written since ``offset`` + (``cat``-ing the whole file every poll re-sends all output over docker/SSH). + + Prints one header line ``" "`` then the bytes in [offset, size). + The size is read first and the tail cut at that same size, so a growing file + never sends a byte twice; a file that shrank was rotated/truncated, so the + offset drops to 0 and the reader starts over. The window end is pulled back + to a UTF-8 character boundary (the backend decodes each ``execute()`` result + on its own, so a straddling multibyte char would become U+FFFD and break watch + patterns at the seam): up to 3 trailing continuation bytes are held for the + next poll and the header reports the trimmed size.""" + return ( + f"O={offset}; " + f"S=$({{ wc -c < {quoted_log_path}; }} 2>/dev/null | tr -dc '0-9'); " + f"S=${{S:-0}}; " + f'if [ "$S" -lt "$O" ]; then O=0; fi; ' + # Scan back up to 3 continuation bytes (octal 200-277) to the lead byte; if + # the lead's declared length (3xx=2, 34x-35x=3, 36x-37x=4) exceeds the bytes + # present, trim to before it. Complete sequences and ASCII tails untouched. + f'N=0; P=$S; while [ "$P" -gt "$O" ] && [ "$N" -lt 3 ]; do ' + f"B=$(tail -c +$P {quoted_log_path} 2>/dev/null | head -c 1 | od -An -to1 | tr -dc '0-9'); " + f'case "$B" in 2[0-7][0-7]) P=$((P-1)); N=$((N+1));; *) break;; esac; done; ' + f'if [ "$N" -gt 0 ] || [ "$P" -eq "$S" ]; then ' + f"B=$(tail -c +$P {quoted_log_path} 2>/dev/null | head -c 1 | od -An -to1 | tr -dc '0-9'); " + f'case "$B" in 3[0-3][0-7]) L=2;; 3[4-5][0-7]) L=3;; 3[6-7][0-7]) L=4;; *) L=1;; esac; ' + f'if [ "$L" -gt $((N+1)) ]; then S=$((P-1)); fi; fi; ' + f'echo "$S $O"; ' + f'if [ "$S" -gt "$O" ]; then ' + f"tail -c +$((O+1)) {quoted_log_path} 2>/dev/null | head -c $((S-O)); fi" + ) + def _env_poller_loop(self, session: ProcessSession, env: Any, log_path: str, pid_path: str, exit_path: str): """Background thread: poll a sandbox log file for non-local backends.""" q = shlex.quote - prev_output_len = 0 # delta tracking for watch-pattern scanning + # Byte offset already read from the log (bytes, not chars: the shell counts bytes). + prev_output_bytes = 0 while not session.exited: time.sleep(2) try: - new_output = env.execute(f"cat {q(log_path)} 2>/dev/null", timeout=10).get("output", "") - if new_output: - delta = new_output[prev_output_len:] if len(new_output) > prev_output_len else "" - prev_output_len = len(new_output) + # Read only the bytes written since the last poll. + raw = env.execute(self._log_delta_command(q(log_path), prev_output_bytes), + timeout=10).get("output", "") + header, _, delta = raw.partition("\n") + try: + size_str, offset_str = header.split() + new_size = int(size_str) + used_offset = int(offset_str) + except ValueError: + # No usable header (command failed, shell missing a tool): skip this + # poll rather than act on a half-read value. + new_size = None + used_offset = None + delta = "" + if new_size is not None: + if used_offset < prev_output_bytes: + # Log rotated/truncated: what we hold no longer lines up. Restart. + with session._lock: + session.output_buffer = "" + prev_output_bytes = new_size + if delta: with session._lock: - session.output_buffer = new_output[-session.max_output_chars:] - if delta: - self._check_watch_patterns(session, delta) - self._emit_output(session, delta) + session.output_buffer += delta + if len(session.output_buffer) > session.max_output_chars: + session.output_buffer = session.output_buffer[-session.max_output_chars:] + self._check_watch_patterns(session, delta) + self._emit_output(session, delta) + check = env.execute( f"kill -0 \"$(cat {q(pid_path)} 2>/dev/null)\" 2>/dev/null; echo $?", timeout=5) check_output = check.get("output", "").strip() diff --git a/tools/session_search_tool.py b/tools/session_search_tool.py index b396fb8cbb..0eed751723 100644 --- a/tools/session_search_tool.py +++ b/tools/session_search_tool.py @@ -418,12 +418,18 @@ def _read_session(db, session_id: str, head: int = 20, tail: int = 10, link_prof def _list_recent_sessions(db, limit: int, current_session_id: str = None, link_profile: str = None) -> str: """Browse shape: metadata for the most recent sessions (no LLM, no FTS5).""" def _browse(): - # list_sessions_rich already applies the canonical child classifier (roots, - # /branch and /new-reset children admitted; delegation/compression children - # hidden). Re-classifying here re-hid legacy reset children — trust the query. - sessions = db.list_sessions_rich( + # Never use list_sessions_rich(order_by_last_active=True) here: it walks every + # compression chain and derives activity/previews before LIMIT, which can + # monopolise a gateway callback for minutes on a multi-GB state.db. The + # bounded browse query preselects an indexed candidate set and carries a + # cooperative SQLite VM cancellation deadline. Fail closed rather than + # silently falling back to the whole-database query shape. + bounded_list = getattr(db, "list_recent_sessions_bounded", None) + if bounded_list is None: + raise RuntimeError("session database does not support bounded recent-session browse") + sessions = bounded_list( limit=limit + 15, # extra so we can skip current / compression roots - exclude_sources=list(_HIDDEN_SESSION_SOURCES), order_by_last_active=True) + exclude_sources=list(_HIDDEN_SESSION_SOURCES), timeout_seconds=3.0) current_root, has_compression_hop = ( _resolve_to_parent(db, current_session_id) if current_session_id else (None, False)) results = [] diff --git a/tools/terminal_tool.py b/tools/terminal_tool.py index 9fc2fdc48e..8ba8a06890 100644 --- a/tools/terminal_tool.py +++ b/tools/terminal_tool.py @@ -44,7 +44,7 @@ def _redact_terminal_error_text(value: Any) -> str: from tools.interrupt import _interrupt_event # noqa: F401 — re-exported (tests patch it here) from tools.registry import tool_error from tools.terminal_tool_lifecycle import ( # noqa: F401 (re-exported; tests patch tools.terminal_tool.) - _check_disk_usage_warning, _cleanup_inactive_envs, _create_configured_env, + _check_disk_usage_warning, _cleanup_env, _cleanup_inactive_envs, _create_configured_env, _evict_environment_for_task, cleanup_all_environments, cleanup_vm, ensure_task_env, get_active_env, is_persistent_env, ) diff --git a/tools/terminal_tool_guards.py b/tools/terminal_tool_guards.py index d809271436..24940a0379 100644 --- a/tools/terminal_tool_guards.py +++ b/tools/terminal_tool_guards.py @@ -206,8 +206,12 @@ def gateway_lifecycle_block( _MAX_REFERENCED_SCRIPT_BYTES, contains_gateway_lifecycle_command_or_referenced_script, contains_launchctl_submit_command, + lifecycle_scan_root_within_budget, ) - if contains_launchctl_submit_command(command): + # Keep the specific launchctl diagnostic when this optional pre-scan fits the + # budget. The full fail-closed guard below still runs when it does not, so + # oversized roots never reach shlex here. + if lifecycle_scan_root_within_budget(command) and contains_launchctl_submit_command(command): return _blocked_json( "Blocked: launchctl submit/bootstrap registers a persistent " "KeepAlive job and is unsafe from inside the gateway process. " diff --git a/tools/terminal_tool_lifecycle.py b/tools/terminal_tool_lifecycle.py index 7ef977abd9..91647b523a 100644 --- a/tools/terminal_tool_lifecycle.py +++ b/tools/terminal_tool_lifecycle.py @@ -80,23 +80,30 @@ def _create_configured_env( ) -def _teardown_env(env: Any, task_id: str, *, force_remove: Optional[bool] = None, done_msg: str = "Cleaned up inactive environment for task: %s") -> None: - """Stop *env* via cleanup()/stop()/terminate(), whichever it has; log the outcome. +def _cleanup_env(env: Any, *, force_remove: Optional[bool] = None) -> None: + """Tear down one environment via cleanup()/stop()/terminate(), whichever it has. - ``force_remove`` is forwarded to ``cleanup()`` only when given and the - backend's signature accepts it (DockerEnvironment; others don't). A - 404/"not found" error means the sandbox is already gone — logged at info. + ``force_remove`` is forwarded to ``cleanup()`` only when given and the backend's + signature accepts it (``DockerEnvironment``, issue #20561; other backends don't). + Shared by ``cleanup_vm``, the idle reaper and the prompt-time backend probe so + the signature check lives in one place. """ + if hasattr(env, 'cleanup'): + if force_remove is not None and "force_remove" in inspect.signature(env.cleanup).parameters: + env.cleanup(force_remove=force_remove) + else: + env.cleanup() + elif hasattr(env, 'stop'): + env.stop() + elif hasattr(env, 'terminate'): + env.terminate() + + +def _teardown_env(env: Any, task_id: str, *, force_remove: Optional[bool] = None, done_msg: str = "Cleaned up inactive environment for task: %s") -> None: + """``_cleanup_env`` plus outcome logging. A 404/"not found" error means the + sandbox is already gone — logged at info.""" try: - if hasattr(env, 'cleanup'): - if force_remove is not None and "force_remove" in inspect.signature(env.cleanup).parameters: - env.cleanup(force_remove=force_remove) - else: - env.cleanup() - elif hasattr(env, 'stop'): - env.stop() - elif hasattr(env, 'terminate'): - env.terminate() + _cleanup_env(env, force_remove=force_remove) logger.info(done_msg, task_id) except Exception as e: error_str = str(e) diff --git a/tools/terminal_tool_sudo.py b/tools/terminal_tool_sudo.py index 3dcd91f3d1..9d43f7d94a 100644 --- a/tools/terminal_tool_sudo.py +++ b/tools/terminal_tool_sudo.py @@ -378,7 +378,17 @@ def _rewrite_compound_background(command: str) -> str: insert_pos = chain_end while insert_pos < amp_pos and result[insert_pos].isspace(): insert_pos += 1 - result = result[:insert_pos] + "{ " + result[insert_pos:amp_pos] + "& }" + result[amp_pos + 1 :] + # The consumed `&` also separated the compound from any statement that followed + # on the same line (`A && B & C`); `{ B & } C` is a syntax error, so restore a `;` + # when the suffix resumes with command text. No separator when the suffix already + # starts with a terminator (`;` `&` `|` newline `)` `}`) — except `&>`, which is a + # redirect prefix for the NEXT command, not a terminator. Strip only spaces/tabs: + # a newline already terminates the group. + suffix = result[amp_pos + 1 :] + tail = suffix.lstrip(" \t") + needs_separator = bool(tail) and (tail[0] not in ";\n&|)}" or tail.startswith("&>")) + separator = " ;" if needs_separator else "" + result = result[:insert_pos] + "{ " + result[insert_pos:amp_pos] + "& }" + separator + suffix return result diff --git a/tools/vision_tools.py b/tools/vision_tools.py index b8a45eddff..f91698eca3 100644 --- a/tools/vision_tools.py +++ b/tools/vision_tools.py @@ -403,12 +403,28 @@ _TOOL_RESULT_MEDIA_PROVIDERS = frozenset({ _GEMINI_PROVIDERS = frozenset({"google", "gemini", "google-gemini", "google-vertex-gemini"}) +def _profile_rejects_tool_media(provider: str) -> bool: + """Hard veto: the provider's ``ProviderProfile`` declares + ``supports_vision_tool_messages=False`` — images are accepted in user + messages but list-type tool-result content is rejected with 400 + (xiaomi/MiMo "text is not set"). ``supports_vision`` alone must not + override this, or the multimodal tool-result envelope 400s every turn + and the image never enters context (#89981). + """ + try: + from providers import get_provider_profile + profile = get_provider_profile(str(provider or "").strip().lower()) + return profile is not None and profile.supports_vision_tool_messages is False + except Exception: + return False + + def _supports_media_in_tool_results(provider: str, model: str) -> bool: """Whether provider+model accepts image content inside a tool-result message. Unknown providers are False (caller falls back to aux-LLM text) unless their ``ProviderProfile`` - declares ``supports_vision``.""" + declares ``supports_vision``; ``supports_vision_tool_messages=False`` is a hard veto.""" p = provider.strip().lower() if isinstance(provider, str) else "" - if not p: + if not p or _profile_rejects_tool_media(p): return False if p in _TOOL_RESULT_MEDIA_PROVIDERS: return True @@ -436,6 +452,11 @@ def _should_use_native_vision_fast_path() -> bool: cfg = load_config() if decide_image_input_mode(provider, model, cfg) != "native": return False + # The profile veto applies ahead of the capability lookup too: a + # model marked vision-capable by models.dev / custom_providers must + # not re-open the multimodal-envelope route the profile rejects. + if _profile_rejects_tool_media(provider): + return False return ( _supports_media_in_tool_results(provider, model) or _lookup_supports_vision(provider, model, cfg) is True) diff --git a/tui_gateway/session_lifecycle.py b/tui_gateway/session_lifecycle.py index 6b6efc09be..a3cb1b24d3 100644 --- a/tui_gateway/session_lifecycle.py +++ b/tui_gateway/session_lifecycle.py @@ -73,7 +73,7 @@ def _release_active_session_slot(session: dict | None) -> bool: lease = session.get("active_session_lease") if session else None if lease is None: return True - if (err := _lease_retry(3 if getattr(lease, "track_liveness", False) else 1, lease.release)) is not None: + if (err := _lease_retry(3 if getattr(lease, "track_liveness", False) else 1, lambda: lease.release())) is not None: logger.warning("Failed to release active session slot", exc_info=err) return False if not (getattr(lease, "released", True) or not getattr(lease, "enabled", True)): @@ -83,6 +83,13 @@ def _release_active_session_slot(session: dict | None) -> bool: return True +def _own_live_lease_ids(*, exclude=None) -> set[str]: + """Snapshot leases still backed by this process's live session records.""" + with _sessions_lock: + return {str(lease.lease_id) for session in _sessions.values() + if (lease := session.get("active_session_lease")) is not None and lease is not exclude} + + @contextlib.contextmanager def _other_runtime_lease_guard(session_id: str, session: dict): """Release this runtime and lock sibling ownership through the DB write. Yields True (another runtime @@ -97,13 +104,15 @@ def _other_runtime_lease_guard(session_id: str, session: dict): return stack = contextlib.ExitStack() active: list = [] + own_live_lease_ids = _own_live_lease_ids(exclude=lease) def _enter() -> None: stack.close() # drop anything a half-failed previous attempt left behind if lease is not None and getattr(lease, "enabled", False): - guard = release_active_session_liveness_guard(lease, session_id) + guard = release_active_session_liveness_guard(lease, session_id, own_live_lease_ids=own_live_lease_ids) else: - guard = active_session_liveness_guard(session_id, registry_home=session.get("profile_home")) + guard = active_session_liveness_guard( + session_id, registry_home=session.get("profile_home"), own_live_lease_ids=own_live_lease_ids) active[:] = [stack.enter_context(guard)] if (last_error := _lease_retry(3, _enter)) is not None: diff --git a/tui_gateway/session_reaper.py b/tui_gateway/session_reaper.py index ba91af9159..2cc7a3bbc5 100644 --- a/tui_gateway/session_reaper.py +++ b/tui_gateway/session_reaper.py @@ -182,10 +182,7 @@ def _reclaim_orphaned_leases() -> None: """Hand the registry the lease ids we still own so it can drop the rest.""" try: from hermes_cli.active_sessions import release_orphaned_leases - with _sessions_lock: - live = {lease.lease_id for session in _sessions.values() - if (lease := session.get("active_session_lease")) is not None} - if dropped := release_orphaned_leases(live): + if dropped := release_orphaned_leases(_own_live_lease_ids()): logger.info("Reclaimed %d orphaned active-session lease(s)", dropped) except Exception: logger.debug("orphaned lease reclaim failed", exc_info=True) diff --git a/tui_gateway/ws.py b/tui_gateway/ws.py index 5ba1bdb6ca..e2d2d85a24 100644 --- a/tui_gateway/ws.py +++ b/tui_gateway/ws.py @@ -15,6 +15,7 @@ import time from typing import Any from tui_gateway import server +from agent.message_sanitization import _sanitize_surrogates from tui_gateway.event_replay import replay_epoch _log = logging.getLogger(__name__) @@ -42,6 +43,16 @@ def _note_dashboard_client_activity(*, force: bool = False) -> None: _log.debug("dashboard client heartbeat touch failed", exc_info=True) +def _sanitize_ws_text(text: str) -> str: + """Return *text* that can be UTF-8 encoded for a WebSocket frame. + + ``json.dumps(..., ensure_ascii=False)`` happily emits lone UTF-16 surrogates; Starlette's + ``send_text`` then raises ``UnicodeEncodeError``, which used to latch the whole connection + closed. Same U+FFFD replacement every other Hermes transport applies. + """ + return _sanitize_surrogates(text) if text else text + + # Max seconds a pool-dispatched handler blocks waiting for the loop to flush a WS frame before we # give up waiting (the transport is NOT marked dead). _WS_WRITE_TIMEOUT_S = 10.0 @@ -145,6 +156,10 @@ class WSTransport: if batch and not self._closed: self._loop.create_task(self._safe_send_many(batch)) + @property + def closed(self) -> bool: + return self._closed + async def write_async(self, obj: dict) -> bool: """Send from the owning loop; awaits until the frame is on the wire. Buffered tokens are flushed ahead of it in the SAME batch so nothing slips between.""" @@ -161,15 +176,21 @@ class WSTransport: async with self._send_lock: if self._closed: return - try: - for line in lines: - if self._closed: - return - await self._ws.send_text(line) - except Exception as exc: - # Latch while holding the writer lock so queued batches observe the failure first. - self._closed = True - _log.warning("ws send failed peer=%s error_type=%s error=%s", self._peer, type(exc).__name__, exc) + for line in lines: + if self._closed: + return + payload = _sanitize_ws_text(line) + try: + await self._ws.send_text(payload) + except UnicodeEncodeError as exc: + # A single illegal UTF-8 frame (lone surrogate) must not tear down the socket. + _log.warning("ws send skipped invalid utf-8 frame peer=%s error=%s", self._peer, exc) + continue + except Exception as exc: + # Latch while holding the writer lock so queued batches observe the failure first. + self._closed = True + _log.warning("ws send failed peer=%s error_type=%s error=%s", self._peer, type(exc).__name__, exc) + return def close(self) -> None: # loop thread (handle_ws finally), so the TimerHandle is safe self._closed = True diff --git a/ui-tui/packages/hermes-ink/src/native-ts/yoga-layout/index.ts b/ui-tui/packages/hermes-ink/src/native-ts/yoga-layout/index.ts index a62a4bae16..4b99ee2ec2 100644 --- a/ui-tui/packages/hermes-ink/src/native-ts/yoga-layout/index.ts +++ b/ui-tui/packages/hermes-ink/src/native-ts/yoga-layout/index.ts @@ -436,6 +436,20 @@ export class Node { _cGen = -1 _cN = 0 _cWr = 0 + _rSubtreeLayoutGen = -1 + _rValid = false + _rLeft = NaN + _rTop = NaN + _rWidth = NaN + _rHeight = NaN + _rRoundedLeft = NaN + _rRoundedTop = NaN + _rRoundedWidth = NaN + _rRoundedHeight = NaN + _rParentAbsLeft = NaN + _rParentAbsTop = NaN + _rScale = NaN + _rIsText = false constructor(config?: Config) { this.style = defaultStyle() this.layout = { @@ -509,6 +523,7 @@ export class Node { this._cN = 0 this._cWr = 0 this._fbBasis = NaN + this._rValid = false } markDirty(): void { this.isDirty_ = true @@ -533,26 +548,26 @@ export class Node { this.markDirty() } getComputedLeft(): number { - return this.layout.left + return this._rValid ? this._rRoundedLeft : this.layout.left } getComputedTop(): number { - return this.layout.top + return this._rValid ? this._rRoundedTop : this.layout.top } getComputedWidth(): number { - return this.layout.width + return this._rValid ? this._rRoundedWidth : this.layout.width } getComputedHeight(): number { - return this.layout.height + return this._rValid ? this._rRoundedHeight : this.layout.height } getComputedRight(): number { const p = this.parent - return p ? p.layout.width - this.layout.left - this.layout.width : 0 + return p ? p.getComputedWidth() - this.getComputedLeft() - this.getComputedWidth() : 0 } getComputedBottom(): number { const p = this.parent - return p ? p.layout.height - this.layout.top - this.layout.height : 0 + return p ? p.getComputedHeight() - this.getComputedTop() - this.getComputedHeight() : 0 } getComputedLayout(): { left: number @@ -563,12 +578,12 @@ export class Node { height: number } { return { - left: this.layout.left, - top: this.layout.top, + left: this.getComputedLeft(), + top: this.getComputedTop(), right: this.getComputedRight(), bottom: this.getComputedBottom(), - width: this.layout.width, - height: this.layout.height + width: this.getComputedWidth(), + height: this.getComputedHeight() } } getComputedBorder(edge: Edge): number { @@ -842,6 +857,8 @@ export class Node { _yogaNodesVisited = 0 _yogaMeasureCalls = 0 _yogaCacheHits = 0 + _yogaRoundedNodes = 0 + _yogaRoundSkips = 0 _generation++ const w = ownerWidth === undefined ? NaN : ownerWidth const h = ownerHeight === undefined ? NaN : ownerHeight @@ -924,18 +941,24 @@ let _yogaNodesVisited = 0 let _yogaMeasureCalls = 0 let _yogaCacheHits = 0 let _yogaLiveNodes = 0 +let _yogaRoundedNodes = 0 +let _yogaRoundSkips = 0 export function getYogaCounters(): { visited: number measured: number cacheHits: number live: number + rounded: number + roundSkips: number } { return { visited: _yogaNodesVisited, measured: _yogaMeasureCalls, cacheHits: _yogaCacheHits, - live: _yogaLiveNodes + live: _yogaLiveNodes, + rounded: _yogaRoundedNodes, + roundSkips: _yogaRoundSkips } } @@ -1020,6 +1043,15 @@ function layoutNode( } } + if (performLayout) { + let ancestor: Node | null = node + + while (ancestor && ancestor._rSubtreeLayoutGen !== _generation) { + ancestor._rSubtreeLayoutGen = _generation + ancestor = ancestor.parent + } + } + const wasDirty = node.isDirty_ if (performLayout) { @@ -2191,28 +2223,63 @@ function collectLayoutChildren(node: Node, flow: Node[], abs: Node[]): void { } function roundLayout(node: Node, scale: number, absLeft: number, absTop: number): void { - if (scale === 0) { + const l = node.layout + const isText = node.measureFunc !== null + + if ( + node._rValid && + node._rSubtreeLayoutGen !== _generation && + sameFloat(node._rParentAbsLeft, absLeft) && + sameFloat(node._rParentAbsTop, absTop) && + sameFloat(node._rScale, scale) && + node._rIsText === isText && + sameFloat(node._rLeft, l.left) && + sameFloat(node._rTop, l.top) && + sameFloat(node._rWidth, l.width) && + sameFloat(node._rHeight, l.height) + ) { + _yogaRoundSkips++ + return } - const l = node.layout + _yogaRoundedNodes++ + const nodeLeft = l.left const nodeTop = l.top const nodeWidth = l.width const nodeHeight = l.height const absNodeLeft = absLeft + nodeLeft const absNodeTop = absTop + nodeTop - const isText = node.measureFunc !== null - l.left = roundValue(nodeLeft, scale, false, isText) - l.top = roundValue(nodeTop, scale, false, isText) - const absRight = absNodeLeft + nodeWidth - const absBottom = absNodeTop + nodeHeight - const hasFracW = !isWholeNumber(nodeWidth * scale) - const hasFracH = !isWholeNumber(nodeHeight * scale) - l.width = - roundValue(absRight, scale, isText && hasFracW, isText && !hasFracW) - roundValue(absNodeLeft, scale, false, isText) - l.height = - roundValue(absBottom, scale, isText && hasFracH, isText && !hasFracH) - roundValue(absNodeTop, scale, false, isText) + node._rValid = true + node._rLeft = nodeLeft + node._rTop = nodeTop + node._rWidth = nodeWidth + node._rHeight = nodeHeight + node._rParentAbsLeft = absLeft + node._rParentAbsTop = absTop + node._rScale = scale + node._rIsText = isText + + if (scale === 0) { + node._rRoundedLeft = nodeLeft + node._rRoundedTop = nodeTop + node._rRoundedWidth = nodeWidth + node._rRoundedHeight = nodeHeight + } else { + node._rRoundedLeft = roundValue(nodeLeft, scale, false, isText) + node._rRoundedTop = roundValue(nodeTop, scale, false, isText) + const absRight = absNodeLeft + nodeWidth + const absBottom = absNodeTop + nodeHeight + const hasFracW = !isWholeNumber(nodeWidth * scale) + const hasFracH = !isWholeNumber(nodeHeight * scale) + node._rRoundedWidth = + roundValue(absRight, scale, isText && hasFracW, isText && !hasFracW) - + roundValue(absNodeLeft, scale, false, isText) + node._rRoundedHeight = + roundValue(absBottom, scale, isText && hasFracH, isText && !hasFracH) - + roundValue(absNodeTop, scale, false, isText) + } for (const c of node.children) { roundLayout(c, scale, absNodeLeft, absNodeTop) diff --git a/ui-tui/packages/hermes-ink/src/native-ts/yoga-layout/round-layout.test.ts b/ui-tui/packages/hermes-ink/src/native-ts/yoga-layout/round-layout.test.ts new file mode 100644 index 0000000000..41f78feea3 --- /dev/null +++ b/ui-tui/packages/hermes-ink/src/native-ts/yoga-layout/round-layout.test.ts @@ -0,0 +1,321 @@ +import { describe, expect, it } from 'vitest' + +import Yoga, { FlexDirection, getYogaCounters, type Node } from './index.js' + +const snapshot = (node: Node): number[] => { + const result = [node.getComputedLeft(), node.getComputedTop(), node.getComputedWidth(), node.getComputedHeight()] + + for (let index = 0; index < node.getChildCount(); index++) { + result.push(...snapshot(node.getChild(index))) + } + + return result +} + +const buildTree = (rootWidth: number, widths: number[], scale: number) => { + const config = Yoga.Config.create() + config.setPointScaleFactor(scale) + const root = Yoga.Node.create(config) + root.setFlexDirection(FlexDirection.Column) + root.setWidth(rootWidth) + root.setHeight(20) + const leaves: Node[] = [] + + for (let groupIndex = 0; groupIndex < 4; groupIndex++) { + const group = Yoga.Node.create(config) + group.setFlexDirection(FlexDirection.Row) + group.setHeight(3.125) + root.insertChild(group, groupIndex) + + for (let leafIndex = 0; leafIndex < 8; leafIndex++) { + const leaf = Yoga.Node.create(config) + leaf.setWidth(widths[groupIndex * 8 + leafIndex]!) + leaf.setHeight(1.125 + (leafIndex % 3) * 0.25) + group.insertChild(leaf, leafIndex) + leaves.push(leaf) + } + } + + return { config, leaves, root } +} + +interface NodeSpec { + w?: number + h?: number + pad?: number + mar?: number + row: boolean + grow?: number + kids: NodeSpec[] +} + +const buildMutationTree = (spec: NodeSpec) => { + const all: Node[] = [] + + const makeNode = (value: NodeSpec): Node => { + const node = Yoga.Node.create() + + if (value.w !== undefined) { + node.setWidth(value.w) + } + + if (value.h !== undefined) { + node.setHeight(value.h) + } + + if (value.pad !== undefined) { + node.setPadding(1, value.pad) + } + + if (value.mar !== undefined) { + node.setMargin(1, value.mar) + } + + if (value.row) { + node.setFlexDirection(FlexDirection.Row) + } + + if (value.grow !== undefined) { + node.setFlexGrow(value.grow) + } + + all.push(node) + value.kids.forEach((child, index) => node.insertChild(makeNode(child), index)) + + return node + } + + return { all, root: makeNode(spec) } +} + +describe('incremental layout rounding', () => { + it('keeps rounding work flat when only the clock changes', () => { + const results = [50, 500, 5000].map(rowCount => { + const config = Yoga.Config.create() + config.setPointScaleFactor(2) + + const root = Yoga.Node.create(config) + root.setWidth(80) + root.setHeight(40) + + const transcript = Yoga.Node.create(config) + transcript.setHeight(39) + root.insertChild(transcript, 0) + + for (let index = 0; index < rowCount; index++) { + const row = Yoga.Node.create(config) + row.setWidth(20.25) + row.setHeight(0.25) + transcript.insertChild(row, index) + } + + const clock = Yoga.Node.create(config) + clock.setWidth(5.25) + clock.setHeight(1) + root.insertChild(clock, 1) + + root.calculateLayout(80, 40) + const transcriptWidth = transcript.getComputedWidth() + + clock.setWidth(6.25) + root.calculateLayout(80, 40) + + const counters = getYogaCounters() + expect(clock.getComputedWidth()).toBe(6.5) + expect(transcript.getComputedWidth()).toBe(transcriptWidth) + root.freeRecursive() + Yoga.Config.destroy(config) + + return { rounded: counters.rounded, roundSkips: counters.roundSkips } + }) + + expect(results).toEqual([ + { rounded: 2, roundSkips: 1 }, + { rounded: 2, roundSkips: 1 }, + { rounded: 2, roundSkips: 1 } + ]) + }) + + it('re-rounds cached raw geometry when the point scale changes', () => { + const config = Yoga.Config.create() + config.setPointScaleFactor(2) + + const root = Yoga.Node.create(config) + root.setWidth(20) + root.setHeight(10) + + const child = Yoga.Node.create(config) + child.setWidth(10.25) + child.setHeight(1) + root.insertChild(child, 0) + + root.calculateLayout(20, 10) + expect(child.getComputedWidth()).toBe(10.5) + + config.setPointScaleFactor(4) + child.setWidth(10.125) + root.calculateLayout(20, 10) + + expect(child.getComputedWidth()).toBe(10.25) + + config.setPointScaleFactor(0) + root.calculateLayout(20, 10) + + expect(child.getComputedWidth()).toBe(10.125) + + root.freeRecursive() + Yoga.Config.destroy(config) + }) + + it('matches a fresh full layout across leaf, root, and scale changes', () => { + const widths = Array.from({ length: 32 }, (_, index) => 1.125 + (index % 5) * 0.375) + let rootWidth = 40.25 + let scale = 2 + const incremental = buildTree(rootWidth, widths, scale) + + for (let step = 0; step < 24; step++) { + if (step % 6 === 0) { + scale = scale === 2 ? 4 : 2 + incremental.config.setPointScaleFactor(scale) + } else if (step % 5 === 0) { + rootWidth += 0.375 + incremental.root.setWidth(rootWidth) + } else { + const leafIndex = (step * 7) % widths.length + widths[leafIndex]! += 0.125 + incremental.leaves[leafIndex]!.setWidth(widths[leafIndex]!) + } + + incremental.root.calculateLayout(rootWidth, 20) + const fresh = buildTree(rootWidth, widths, scale) + fresh.root.calculateLayout(rootWidth, 20) + + expect(snapshot(incremental.root), `step ${step}`).toEqual(snapshot(fresh.root)) + + fresh.root.freeRecursive() + Yoga.Config.destroy(fresh.config) + } + + incremental.root.freeRecursive() + Yoga.Config.destroy(incremental.config) + }) + + it('rounds a fractional child after its whole-pixel parent hits the layout cache', () => { + const root = Yoga.Node.create() + root.setWidth(120) + root.setHeight(40) + const row = Yoga.Node.create() + row.setWidth(120) + row.setHeight(2) + const leaf = Yoga.Node.create() + leaf.setWidth(10.4) + leaf.setHeight(1.6) + row.insertChild(leaf, 0) + root.insertChild(row, 0) + root.calculateLayout(120, 40) + + leaf.setWidth(11.4) + leaf.setHeight(2.6) + root.calculateLayout(120, 40) + + expect(leaf.getComputedLayout()).toMatchObject({ width: 11, height: 3 }) + root.freeRecursive() + }) + + it('does not mix cached rounded coordinates with raw dimensions', () => { + const spec: NodeSpec = { + h: 4.917283, + pad: 1.934923, + row: false, + kids: [ + { + w: 11.007846, + row: false, + kids: [ + { + mar: 1.306334, + row: true, + kids: [ + { + mar: 0.199071, + row: true, + grow: 1.522763, + kids: [{ w: 4.6117, h: 5.852709, row: false, kids: [] }] + } + ] + }, + { + w: 40.14272, + pad: 1.043468, + mar: 1.495136, + row: false, + kids: [ + { + w: 8.335441, + row: false, + kids: [ + { h: 5.249866, mar: 0.997316, row: true, grow: 1.554555, kids: [] }, + { h: 6.335704, mar: 0.817136, row: false, grow: 0.079169, kids: [] } + ] + }, + { + w: 32.935694, + row: false, + grow: 1.781838, + kids: [ + { w: 38.068879, pad: 0.696862, row: true, kids: [] }, + { pad: 1.66287, row: false, kids: [] } + ] + } + ] + } + ] + } + ] + } + + const mutations = [ + { index: 10, kind: 'height', value: 2.823024 }, + { index: 6, kind: 'margin', value: 0.858021 }, + { index: 3, kind: 'grow', value: 1.321691 }, + { index: 9, kind: 'grow', value: 0.022362 }, + { index: 6, kind: 'grow', value: 0.812981 }, + { index: 7, kind: 'width', value: 5.644585 }, + { index: 2, kind: 'width', value: 20.498063 }, + { index: 4, kind: 'grow', value: 1.653743 }, + { index: 1, kind: 'height', value: 2.709275 }, + { index: 1, kind: 'margin', value: 2.411063 } + ] as const + + const applyMutation = (all: Node[], mutation: (typeof mutations)[number]) => { + const node = all[mutation.index]! + + if (mutation.kind === 'width') { + node.setWidth(mutation.value) + } else if (mutation.kind === 'height') { + node.setHeight(mutation.value) + } else if (mutation.kind === 'margin') { + node.setMargin(1, mutation.value) + } else { + node.setFlexGrow(mutation.value) + } + } + + const incremental = buildMutationTree(spec) + incremental.root.calculateLayout(96, 5) + + for (const mutation of mutations) { + applyMutation(incremental.all, mutation) + incremental.root.calculateLayout(96, 5) + } + + const fresh = buildMutationTree(spec) + mutations.forEach(mutation => applyMutation(fresh.all, mutation)) + fresh.root.calculateLayout(96, 5) + + expect(incremental.all[11]!.getComputedHeight()).toBe(fresh.all[11]!.getComputedHeight()) + expect(fresh.all[11]!.getComputedHeight()).toBe(2) + incremental.root.freeRecursive() + fresh.root.freeRecursive() + }) +}) diff --git a/uv.lock b/uv.lock index cadbe3f3f5..ec0c477bd2 100644 --- a/uv.lock +++ b/uv.lock @@ -1990,7 +1990,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 == '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.7.1,<0.8" }, + { 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.8.3,<0.9" }, { 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" }, @@ -2854,17 +2854,17 @@ wheels = [ [[package]] name = "nemo-relay" -version = "0.7.2" +version = "0.8.3" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/58/81/a7a545ac3a2f8c670d261c89df599aa8fbf49d8be45fd1f52efb36b489eb/nemo_relay-0.7.2.tar.gz", hash = "sha256:828d9f6c7d7e4e42276bb7192bd44202c761e0c76fa4943d84e051b5a99028e5", size = 1295616, upload-time = "2026-08-08T01:54:00.953Z" } +sdist = { url = "https://files.pythonhosted.org/packages/03/73/ac90ccb08faca19b2c8470bdd4d5b9bae89fc5edfde8dae72ee7b1a2d8df/nemo_relay-0.8.3.tar.gz", hash = "sha256:3670c0689f0709354068a1460131e6f01ea44cd7c2b2186930dd20b23cd079ff", size = 1615748, upload-time = "2026-09-02T03:32:45.828Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/5a/cd/f50440257f01bc5ab3d668331c90e06cf4edcc84dc7dc582d322ad05b622/nemo_relay-0.7.2-cp311-abi3-macosx_11_0_arm64.whl", hash = "sha256:e7c7977f0903793cc34c5542bf2b2e44d107def8a5ae9f1b28f06dd61ddec4ed", size = 9246341, upload-time = "2026-08-08T01:53:19.832Z" }, - { url = "https://files.pythonhosted.org/packages/ed/9f/4041446dd134218799a34b5b5fad3a62d3e1d0a6c322ba2ca4b896ba1393/nemo_relay-0.7.2-cp311-abi3-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:b4ae77c1f3d58eabda264e82ffaca54548df80caede7dd6af8cbd8f72b4a82ed", size = 8454070, upload-time = "2026-08-08T01:53:22.524Z" }, - { url = "https://files.pythonhosted.org/packages/11/83/90230c2e9fae1aee39f768d4a9ef57e9f2716bcaed1a5923cce8b526c66b/nemo_relay-0.7.2-cp311-abi3-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:0ce7103aec546766649c182619d16aa6ad07439e4d0ebd16d95c5004afb3e56a", size = 8954377, upload-time = "2026-08-08T01:53:25.267Z" }, - { url = "https://files.pythonhosted.org/packages/71/e7/463fa461d0801146fec6a00cbc02e8961b30089d65ba170f9dfa9e6e3dcd/nemo_relay-0.7.2-cp311-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:2e7d0c2629ade7313aaed71d2272dca96a2fafad0248d0f87cf40a7720b252a0", size = 10322132, upload-time = "2026-08-08T01:53:27.991Z" }, - { url = "https://files.pythonhosted.org/packages/32/8c/e20ec9c52bd1edd953157aaf24d0d9f9ab8afcbf108fc2356f398e252da8/nemo_relay-0.7.2-cp311-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:b841c92395d7686c7f233036008294b9d362af1ec5123ab0babbfd11cbb04054", size = 10704141, upload-time = "2026-08-08T01:53:30.453Z" }, - { url = "https://files.pythonhosted.org/packages/5a/c1/92a73961ea759b433b1f897b225662d499123cb962b48dc8ece19f610a09/nemo_relay-0.7.2-cp311-abi3-win_amd64.whl", hash = "sha256:0cdcc5e09d6d62d5c1d385dc62c9233eb714a25f36a09da81e5b9731e3c67903", size = 8803938, upload-time = "2026-08-08T01:53:33.437Z" }, - { url = "https://files.pythonhosted.org/packages/9d/ec/2de114dab437431173988b9b11f46e8d377e12d57e1b4903258f3e03c2df/nemo_relay-0.7.2-cp311-abi3-win_arm64.whl", hash = "sha256:ca5f66e617311f836a10d96f120f3f32a99b4267d65048453b31951de3419a9d", size = 8438997, upload-time = "2026-08-08T01:53:36.12Z" }, + { url = "https://files.pythonhosted.org/packages/2a/6a/199c2061358684550780bfc3184e588cea6533ec697ae6fa2f36f8881412/nemo_relay-0.8.3-cp311-abi3-macosx_11_0_arm64.whl", hash = "sha256:2b0a59b8a95d6ed9de099e471318af672b51426a89c2019766ac889b1e23f1d4", size = 10200865, upload-time = "2026-09-02T03:32:15.373Z" }, + { url = "https://files.pythonhosted.org/packages/df/90/3486d10c2003cda3bd4ce98d34e40f283faf87dcb1292cddca356a86c6be/nemo_relay-0.8.3-cp311-abi3-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:0c84a93ad0d4bea9a7bf67917874cd4eee24694764456a5f94c9609b1df42904", size = 9199160, upload-time = "2026-09-02T03:32:17.467Z" }, + { url = "https://files.pythonhosted.org/packages/18/00/15705b941df64443e50c49139140dd84603f23125ab48f9c122942911439/nemo_relay-0.8.3-cp311-abi3-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:10b08a939678d8f54ad02e5fb086454ffad15f269eae4854e9597956dc96b176", size = 9760299, upload-time = "2026-09-02T03:32:19.255Z" }, + { url = "https://files.pythonhosted.org/packages/00/d3/2dac6a938713cfb6f59d8f3e82e991a1635a666f25f8d7fed8a3fa9c56b7/nemo_relay-0.8.3-cp311-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:c2fba57f39100068244f8fcaffbdd801a0a8ccc7b9dbee330bb087f3970d8ccc", size = 11287368, upload-time = "2026-09-02T03:32:21.359Z" }, + { url = "https://files.pythonhosted.org/packages/e0/d2/f61379d244139306c3832082d446ad6167f50108969b3e8720e62e1b2ae5/nemo_relay-0.8.3-cp311-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:ae6804050b309a12dae5a7afc8c85b29b601c2c4a4257d365a998b774a2e7d16", size = 11717569, upload-time = "2026-09-02T03:32:23.321Z" }, + { url = "https://files.pythonhosted.org/packages/db/96/333f62b4176450154a8908f619e3323b70b43c3b733a0a0752ba3b9232b1/nemo_relay-0.8.3-cp311-abi3-win_amd64.whl", hash = "sha256:b1f006e50e44967821f3a77a2e0001eed8154db01168abcc0681891f9c339fed", size = 9656294, upload-time = "2026-09-02T03:32:25.904Z" }, + { url = "https://files.pythonhosted.org/packages/fe/bd/37ab6038111d3cdec1858e584ef9e8fdfe0db07aa6073d66d1303b18d791/nemo_relay-0.8.3-cp311-abi3-win_arm64.whl", hash = "sha256:0d3db4f0a9d1f909acbf659612bfb9859e304b53cd425ef8e0f60ca851e85173", size = 9235572, upload-time = "2026-09-02T03:32:28.042Z" }, ] [[package]] diff --git a/website/docs/developer-guide/session-storage.md b/website/docs/developer-guide/session-storage.md index 0f80accfde..4627b0ffa6 100644 --- a/website/docs/developer-guide/session-storage.md +++ b/website/docs/developer-guide/session-storage.md @@ -169,6 +169,8 @@ The `schema_version` table stores a single integer. Simple column additions are | 20 | Per-model usage attribution — seed `session_model_usage` rows from historical per-session aggregate totals | | 22 | Task-dimension usage attribution — rebuild `session_model_usage` so the `task` column participates in the PRIMARY KEY | | 23 | FTS storage redesign — external-content FTS tables replacing the v11 inline-mode copies (opt-in transition for existing DBs) | +| 29 | Cron sessions leave the trigram (substring/CJK) index; `messages_fts_trigram_src` view + triggers filter on `sessions.source`, one-time rebuild purges historical rows | +| 30 | Delegate-child (subagent) sessions leave the trigram index too — `source='subagent'` or the `$._delegate_from` marker (`FTS_TRIGRAM_SESSION_SQL`). Rows stay in `messages` and the standard `messages_fts` word index, so `session_search` still finds them; only the ~2.6× trigram shadow tables shrink. Same one-time rebuild as v29 | Versions not listed above were declarative column additions handled by `_reconcile_columns()` (version bump only, no data migration). diff --git a/website/docs/integrations/providers.md b/website/docs/integrations/providers.md index 5ede871815..d9de727acd 100644 --- a/website/docs/integrations/providers.md +++ b/website/docs/integrations/providers.md @@ -61,6 +61,8 @@ You need at least one way to connect to an LLM. Use `hermes model` to switch pro | **LM Studio** | `hermes model` → "LM Studio" (provider: `lmstudio`, optional `LM_API_KEY`) | | **Custom Endpoint** | `hermes model` → choose "Custom endpoint" (saved in `config.yaml`) | +All three OpenCode providers send an opaque, per-conversation `x-opencode-session` header on every request (main turns on every transport plus auxiliary calls such as compression and titles). OpenCode uses it to pin a conversation to one backend so its prompt cache stays warm; the value is derived from the Hermes session id and carries no personal data. + For the official API-key path, see the dedicated [Google Gemini guide](/guides/google-gemini). :::tip Model key alias @@ -327,7 +329,7 @@ model: Base URLs can be overridden with `NOVITA_BASE_URL`, `GLM_BASE_URL`, `KIMI_BASE_URL`, `MINIMAX_BASE_URL`, `MINIMAX_CN_BASE_URL`, `DASHSCOPE_BASE_URL`, `XIAOMI_BASE_URL`, `GMI_BASE_URL`, `META_BASE_URL`, or `TOKENHUB_BASE_URL` environment variables. :::note Meta contributor tier -`muse-spark-1.2-contributor` is Meta's discounted tier — Meta may train on your prompts and completions, so [interactive model selection asks for confirmation](../user-guide/configuring-models.md) before using it. Use `muse-spark-1.2` (standard pricing, no training) for confidential work. +`muse-spark-1.2-contributor` and `muse-spark-1.3-contributor` are Meta's contributor tiers — Meta may train on your prompts and completions, so [interactive model selection asks for confirmation](../user-guide/configuring-models.md) before using either. For current pricing and rate limits, see [Meta Model API pricing and rate limits](https://dev.meta.ai/docs/pricing-rate-limits/). Use the standard `muse-spark-1.2` / `muse-spark-1.3` (no training) for confidential work. ::: :::note Z.AI Endpoint Auto-Detection diff --git a/website/docs/reference/cli-commands.md b/website/docs/reference/cli-commands.md index 0d107524e7..63734a2d6a 100644 --- a/website/docs/reference/cli-commands.md +++ b/website/docs/reference/cli-commands.md @@ -1018,6 +1018,21 @@ Restore a previously created Hermes backup into your Hermes home directory. All Stop the gateway before importing to avoid conflicts with running processes. ::: +### SQLite databases + +`.db` members (`state.db`, `kanban.db`, `response_store.db`, …) are not published with a rename like ordinary files. Renaming would replace the file's inode while a gateway, dashboard, or WebUI process still holds the old one open: that process would keep reading pre-import pages and keep writing sessions nobody else can see, and those sessions would simply be absent from the database everyone opens next — with nothing logged. Instead the imported pages are written **into the existing database file**, the same way `/snapshot restore` does it, so every open connection converges on the imported data. + +If the live database cannot be replaced safely — the page copy failed *and* another process still holds the file open — the import leaves that database untouched and lists it under `Warnings (N files skipped)`. Stop the holding processes and re-run. + +Importing an older backup over newer work is still allowed, but it is no longer silent. When the imported `state.db` holds fewer messages than the one it replaced, the summary reports it: + +``` + ⚠ Session data replaced by older backup contents: + state.db: 12 session(s) / 8912 message(s) -> 3 / 24 + Anything recorded after the backup was taken is not in it. + Recover from a newer backup or snapshot: hermes snapshot list +``` + ### Examples ```bash hermes import ~/hermes-backup-20260423.zip # Prompts before overwriting existing config diff --git a/website/docs/user-guide/configuration.md b/website/docs/user-guide/configuration.md index 154f21f03b..95668e8f15 100644 --- a/website/docs/user-guide/configuration.md +++ b/website/docs/user-guide/configuration.md @@ -749,6 +749,12 @@ Set a positive integer to pin a fixed cap instead of the dynamic behavior: context_file_max_chars: 25000 ``` +Each context file read is also bounded by `context_file_read_timeout` (seconds, default `5.0`). A file that takes longer to read — typically on a network-backed filesystem such as iCloud Drive, OneDrive or NFS — is skipped with a warning so the rest of the system prompt still loads: + +```yaml +context_file_read_timeout: 5.0 +``` + ## File Read Safety Controls how much content a single `read_file` call can return. Reads that exceed the limit are rejected with an error telling the agent to use `offset` and `limit` for a smaller range. This prevents a single read of a minified JS bundle or large data file from flooding the context window. @@ -1735,7 +1741,7 @@ agent: | Value | Behavior | |-------|----------| -| `"auto"` (default) | Enabled for models matching: `gpt`, `codex`, `gemini`, `gemma`, `grok`, `glm`, `qwen`, `deepseek`. Disabled for all others (e.g. Claude). | +| `"auto"` (default) | Enabled for models matching: `gpt`, `codex`, `gemini`, `gemma`, `grok`, `glm`, `qwen`, `deepseek`, `muse`. Disabled for all others (e.g. Claude). | | `true` | Always enabled, regardless of model. Useful if you notice your current model describing actions instead of performing them. | | `false` | Always disabled, regardless of model. | | `["gpt", "codex", "qwen", "llama"]` | Enabled only when the model name contains one of the listed substrings (case-insensitive). | @@ -1770,7 +1776,7 @@ agent: | Value | Behavior | |-------|----------| -| `"auto"` (default) | Enabled for models matching: `gpt`, `codex`, `grok`, `deepseek`, `kimi`, `qwen`, `glm`, `minimax`, `mimo`, `mistral`. | +| `"auto"` (default) | Enabled for models matching: `gpt`, `codex`, `grok`, `deepseek`, `kimi`, `qwen`, `glm`, `minimax`, `mimo`, `mistral`, `muse`. | | `true` | Always enabled, regardless of model. | | `false` | Always disabled, regardless of model. | | `["deepseek", "my-custom-model"]` | Enabled only when the model name contains one of the listed substrings (case-insensitive). | diff --git a/website/docs/user-guide/configuring-models.md b/website/docs/user-guide/configuring-models.md index e67e85d5c3..0456c391e4 100644 --- a/website/docs/user-guide/configuring-models.md +++ b/website/docs/user-guide/configuring-models.md @@ -57,7 +57,7 @@ Prompt caches are keyed to the model serving the request, so any mid-conversatio ### Unattended data-training tiers -Models such as `muse-spark-1.2-contributor` are discounted because the vendor may train on your prompts and completions. Interactive model selection always shows a confirmation prompt. Non-interactive startup paths such as Kanban workers and cron agents fail closed because they cannot ask that question. +Models with a `-contributor` suffix (e.g. `muse-spark-1.2-contributor`, `muse-spark-1.3-contributor`) are discounted because the vendor may train on your prompts and completions. Interactive model selection always shows a confirmation prompt. Non-interactive startup paths such as Kanban workers and cron agents fail closed because they cannot ask that question. If training on the unattended workload's data is acceptable, record a persistent acknowledgement: diff --git a/website/docs/user-guide/features/browser.md b/website/docs/user-guide/features/browser.md index 1b2df2a39a..97b5e4ca4a 100644 --- a/website/docs/user-guide/features/browser.md +++ b/website/docs/user-guide/features/browser.md @@ -455,7 +455,7 @@ AGENT_BROWSER_ENGINE=lightpanda The engine works with both browser drivers: -- **Browser Use mode (the default).** Hermes launches `lightpanda serve --host 127.0.0.1 --port ` itself — one process per `browser_exec` session name (or per task) — and points the Browser Use CLI at it. No Chromium, Playwright or Node.js is needed. The process is reaped after `browser.inactivity_timeout`, on exit, and by the orphan sweep if Hermes crashes. Lightpanda has no graphical renderer, so `capture_screenshot()` is unavailable and the tool description tells the model to work text-first; it also holds one page per session, so the model is told to call `new_tab()` once and `goto_url()` afterwards (tracked upstream in [lightpanda-io/browser#1962](https://github.com/lightpanda-io/browser/issues/1962)). +- **Browser Use mode (the default).** Hermes launches `lightpanda serve --host 127.0.0.1 --port ` itself — one process per `browser_exec` session name (or per task) — and points the Browser Use CLI at it. No Chromium, Playwright or Node.js is needed. The process is reaped after `browser.inactivity_timeout`, on exit, and by the orphan sweep if Hermes crashes. All of these processes share one on-disk HTTP cache at `$HERMES_HOME/cache/browser-use/lightpanda/http-cache`, so repeat visits skip re-downloading assets. Hermes passes the cache flag only when the installed Lightpanda supports it (0.3.x+); older binaries simply run without a cache. To clear it, stop your Lightpanda sessions first, then delete that directory. Lightpanda has no graphical renderer, so `capture_screenshot()` is unavailable and the tool description tells the model to work text-first; it also holds one page per session, so the model is told to call `new_tab()` once and `goto_url()` afterwards (tracked upstream in [lightpanda-io/browser#1962](https://github.com/lightpanda-io/browser/issues/1962)). - **Built-in browser tools** (`/browser use off`). Hermes drives Lightpanda through `agent-browser --engine lightpanda` over CDP, the same way it drives local Chrome, with **automatic Chrome fallback**: Lightpanda handles the actions it supports (navigate, snapshot, click, type, scroll, back, press, eval) and Hermes transparently retries on Chrome for anything it doesn't. Screenshots and `browser_vision` are routed straight to Chrome. **When the engine is ignored.** `browser.engine` is the lowest-precedence browser setting: a cloud provider (including the Nous subscription browser — and on never-configured setups, any `BROWSERBASE_API_KEY` / `BROWSER_USE_API_KEY` in `~/.hermes/.env` auto-selects one), Camofox, a `browser.cdp_url` / `/browser connect` override, or `browser.use_real_profile` all take precedence. Picking Lightpanda in `hermes tools` writes `cloud_provider: local` for you; `/browser status` and `hermes doctor` report when the engine is configured but shadowed, and by what. diff --git a/website/docs/user-guide/features/computer-use.md b/website/docs/user-guide/features/computer-use.md index d1d2983c3f..1148e2949a 100644 --- a/website/docs/user-guide/features/computer-use.md +++ b/website/docs/user-guide/features/computer-use.md @@ -430,6 +430,17 @@ computer_use: capability_manifest: "" # capability manifest path, required for bounded ``` +On Linux, native Wayland support remains an explicit opt-in. Hermes passes the +opt-in to every cua-driver process, including gateway sessions, only when that +process also has `WAYLAND_DISPLAY`: + +```yaml +computer_use: + native_wayland: true +``` + +Restart a running gateway after changing this setting. + Override the driver binary path (tests / CI / local builds): ``` diff --git a/website/docs/user-guide/features/context-files.md b/website/docs/user-guide/features/context-files.md index b5c628213d..2906c4f780 100644 --- a/website/docs/user-guide/features/context-files.md +++ b/website/docs/user-guide/features/context-files.md @@ -190,6 +190,7 @@ This scanner protects against common injection patterns, but it's not a substitu | Limit | Value | |-------|-------| | Max chars per file | `context_file_max_chars` when set; otherwise dynamic (scales with model context window, floor 20,000, ceiling 500,000) | +| Read timeout per file | `context_file_read_timeout` (default 5 seconds); a file that takes longer to read — e.g. on iCloud Drive, OneDrive or NFS — is skipped with a warning | | Head truncation ratio | 70% | | Tail truncation ratio | 20% | | Truncation marker | 10% (shows char counts and suggests using file tools) | diff --git a/website/docs/user-guide/features/lsp.md b/website/docs/user-guide/features/lsp.md index 8f5830f479..c34ac0e9c0 100644 --- a/website/docs/user-guide/features/lsp.md +++ b/website/docs/user-guide/features/lsp.md @@ -237,6 +237,12 @@ respawned automatically on the next relevant file operation. Set `idle_timeout: 0` to disable reaping and hold every server's index warm for the life of the process. +Servers that support multi-root workspaces (currently pyright) run as a +**single process** per Hermes process: the first Python project spawns +it, and every further project root — for example sibling git worktrees +edited by parallel subagents — is attached to that same server as an +additional workspace folder instead of starting another copy. + ## Disabling Set `lsp.enabled: false` in `config.yaml` to disable the entire diff --git a/website/static/api/model-catalog.json b/website/static/api/model-catalog.json index d131d55fbd..e96aeb87bb 100644 --- a/website/static/api/model-catalog.json +++ b/website/static/api/model-catalog.json @@ -1,6 +1,6 @@ { "version": 1, - "updated_at": "2026-09-02T16:31:43Z", + "updated_at": "2026-09-03T07:08:39Z", "metadata": { "source": "hermes-agent repo", "docs": "https://hermes-agent.nousresearch.com/docs/reference/model-catalog" @@ -165,6 +165,18 @@ "id": "meta/muse-spark-1.2", "description": "" }, + { + "id": "meta/muse-spark-1.2-contributor", + "description": "" + }, + { + "id": "meta/muse-spark-1.3", + "description": "" + }, + { + "id": "meta/muse-spark-1.3-contributor", + "description": "" + }, { "id": "sakana/fugu-ultra", "description": ""