diff --git a/gateway/run_agent_cache.py b/gateway/run_agent_cache.py index d8be1190dc..0255c00a33 100644 --- a/gateway/run_agent_cache.py +++ b/gateway/run_agent_cache.py @@ -1,4 +1,4 @@ -"""Agent cache, session model overrides, turn leases, run generations and conversation-scope reset methods for GatewayRunner. +"""Agent cache, session model overrides, turn leases, run generations and conversation-scope reset for GatewayRunner. Split out of ``gateway/run.py``; bound onto ``GatewayRunner`` via the MRO. ``gateway.run`` internals are imported lazily inside method bodies (import cycle), @@ -7,16 +7,16 @@ so ``patch("gateway.run.X")`` keeps intercepting them at call time. from __future__ import annotations -import logging -from typing import TYPE_CHECKING import inspect +import logging import threading import time -from contextlib import suppress +from contextlib import nullcontext, suppress +from typing import TYPE_CHECKING, Any, Dict, List, Optional + from gateway.config import Platform from gateway.session import SessionSource, build_session_context_prompt from hermes_cli.config import cfg_get -from typing import Any, Dict, List, Optional if TYPE_CHECKING: # string annotations only; never imported at runtime (cycle) from gateway.run import GatewayRunner, TurnRunner # noqa: F401 @@ -24,17 +24,23 @@ if TYPE_CHECKING: # string annotations only; never imported at runtime (cycle) # Log-record parity with the origin module. logger = logging.getLogger("gateway.run") +# Override fields layered onto runtime kwargs when non-None (partial overrides don't clobber defaults). +_OVERRIDE_APPLY_KEYS = ( + "provider", "requested_provider", "api_key", "base_url", "api_mode", "credential_pool", "capabilities", "max_tokens", +) + + +def _first_agent(entry: Any) -> Any: + """Unwrap a cache entry (``(agent, sig, ...)`` tuple or bare agent) to its agent.""" + return entry[0] if isinstance(entry, tuple) and entry else entry + class GatewayAgentCacheMixin: - """Agent cache, session model overrides, turn leases, run generations and conversation-scope reset methods for GatewayRunner.""" - - @classmethod - def _empty_honcho_cache_busting_config(cls) -> dict[str, Any]: - return {key: None for key in cls._HONCHO_CACHE_BUSTING_KEYS} + """Agent cache, session model overrides, turn leases, run generations and conversation-scope reset for GatewayRunner.""" @classmethod def _extract_honcho_cache_busting_config(cls) -> dict[str, Any]: - """Extract Honcho identity keys, memoized by honcho.json mtime.""" + """Extract Honcho identity keys, memoized by honcho.json mtime; all-None when unavailable.""" try: from plugins.memory.honcho.client import HonchoClientConfig, resolve_config_path @@ -60,7 +66,7 @@ class GatewayAgentCacheMixin: cls._HONCHO_CACHE_BUSTING_MEMO = {memo_key: values} return dict(values) except Exception: - return cls._empty_honcho_cache_busting_config() + return dict.fromkeys(cls._HONCHO_CACHE_BUSTING_KEYS) @classmethod def _extract_cache_busting_config(cls, user_config: dict | None) -> dict: @@ -75,8 +81,7 @@ class GatewayAgentCacheMixin: for section, key in cls._CACHE_BUSTING_CONFIG_KEYS: section_val = cfg.get(section) if section == "checkpoints" and isinstance(section_val, bool): - # Preserve legacy ``checkpoints: true`` behavior. A live - # toggle must still rebuild the cached agent. + # Legacy ``checkpoints: true``: a live toggle must still rebuild the cached agent. out[f"{section}.{key}"] = section_val if key == "enabled" else None elif isinstance(section_val, dict): out[f"{section}.{key}"] = section_val.get(key) @@ -89,25 +94,20 @@ class GatewayAgentCacheMixin: except Exception: out["tools.registry_generation"] = None - # Honcho identity-mapping keys live in honcho.json, not user_config. - # Only read that file when Honcho is the active memory provider. + # Honcho identity-mapping keys live in honcho.json, not user_config; only read that file + # when Honcho is the active memory provider. provider = cfg_get(cfg, "memory", "provider") if isinstance(provider, str) and provider.lower() == "honcho": out.update(cls._extract_honcho_cache_busting_config()) else: - out.update(cls._empty_honcho_cache_busting_config()) + out.update(dict.fromkeys(cls._HONCHO_CACHE_BUSTING_KEYS)) return out @staticmethod def _agent_config_signature( - model: str, - runtime: dict, - enabled_toolsets: list, - ephemeral_prompt: str, - cache_keys: dict | None = None, - user_id: str | None = None, - user_id_alt: str | None = None, + model: str, runtime: dict, enabled_toolsets: list, ephemeral_prompt: str, + cache_keys: dict | None = None, user_id: str | None = None, user_id_alt: str | None = None, skip_context_files: bool = False, ) -> str: """Compute a stable string key from agent config values. @@ -125,8 +125,6 @@ class GatewayAgentCacheMixin: _api_key = str(runtime.get("api_key", "") or "") _api_key_fingerprint = hashlib.sha256(_api_key.encode()).hexdigest() if _api_key else "" - _cache_keys_sorted = sorted((cache_keys or {}).items()) - blob = _j.dumps( [ model, @@ -137,10 +135,9 @@ class GatewayAgentCacheMixin: runtime.get("api_mode", ""), sorted((runtime.get("capabilities") or {}).items()), sorted(enabled_toolsets) if enabled_toolsets else [], - # reasoning_config excluded — it's set per-message on the - # cached agent and doesn't affect system prompt or tools. + # reasoning_config excluded — set per-message on the cached agent; no prompt/tool effect. ephemeral_prompt or "", - _cache_keys_sorted, + sorted((cache_keys or {}).items()), str(user_id or ""), str(user_id_alt or ""), # skip_context_files changes the agent's frozen system prompt (context files in vs out): @@ -152,6 +149,11 @@ class GatewayAgentCacheMixin: ) return hashlib.sha256(blob.encode()).hexdigest()[:16] + def _session_model_override(self, session_key: str) -> Optional[dict]: + """Current in-memory /model override for ``session_key`` (None when absent).""" + state = self._peek_session_state(session_key) + return state.conversation.model_override if state else None + def _rehydrate_session_model_override(self, session_key: str) -> None: """Lazily restore a persisted /model override after a gateway restart. @@ -160,11 +162,7 @@ class GatewayAgentCacheMixin: is never persisted and is re-resolved. No-op when an in-memory override or nothing exists. """ from gateway.run import _resolve_runtime_agent_kwargs_for_provider - _rehydrate_state = self._peek_session_state(session_key) - if ( - _rehydrate_state is not None - and _rehydrate_state.conversation.model_override is not None - ): + if self._session_model_override(session_key) is not None: return store = getattr(self, "session_store", None) if store is None: @@ -172,30 +170,21 @@ class GatewayAgentCacheMixin: try: persisted = store.get_model_override(session_key) except Exception: - logger.debug( - "Failed to read persisted session model override", exc_info=True - ) + logger.debug("Failed to read persisted session model override", exc_info=True) return if not persisted: return - override: Dict[str, Any] = { - "model": persisted.get("model"), - "provider": persisted.get("provider"), - "base_url": persisted.get("base_url"), - } + override: Dict[str, Any] = {k: persisted.get(k) for k in ("model", "provider", "base_url")} provider = persisted.get("provider") if provider: - # Re-resolve credentials for the persisted provider. On failure (e.g. credentials - # removed since the switch) keep the credential-less override — - # _resolve_session_agent_runtime falls back to env resolution and layers model/provider. + # Re-resolve credentials for the persisted provider. On failure (e.g. credentials removed + # since the switch) keep the credential-less override — _resolve_session_agent_runtime + # falls back to env resolution and layers model/provider. try: runtime = _resolve_runtime_agent_kwargs_for_provider(provider) - override["api_key"] = runtime.get("api_key") - override["api_mode"] = runtime.get("api_mode") - override["credential_pool"] = runtime.get("credential_pool") - override["request_overrides"] = dict( - runtime.get("request_overrides") or {} - ) + for k in ("api_key", "api_mode", "credential_pool"): + override[k] = runtime.get(k) + override["request_overrides"] = dict(runtime.get("request_overrides") or {}) override["requested_provider"] = runtime.get("requested_provider") override["capabilities"] = dict(runtime.get("capabilities") or {}) override["max_tokens"] = runtime.get("max_tokens") @@ -213,30 +202,18 @@ class GatewayAgentCacheMixin: session_key, override.get("model"), provider or "", ) - def _apply_session_model_override( - self, session_key: str, model: str, runtime_kwargs: dict - ) -> tuple: + def _apply_session_model_override(self, session_key: str, model: str, runtime_kwargs: dict) -> tuple: """Apply /model session overrides if present, returning (model, runtime_kwargs). Overrides take precedence over config.yaml defaults so the switched model is actually used; ``None`` fields are skipped so partial overrides don't clobber valid defaults. """ from gateway.run import _credential_pool_for_provider - _apply_state = self._peek_session_state(session_key) - override = _apply_state.conversation.model_override if _apply_state else None + override = self._session_model_override(session_key) if not override: return model, runtime_kwargs model = override.get("model", model) - for key in ( - "provider", - "requested_provider", - "api_key", - "base_url", - "api_mode", - "credential_pool", - "capabilities", - "max_tokens", - ): + for key in _OVERRIDE_APPLY_KEYS: val = override.get(key) if val is not None: runtime_kwargs[key] = val @@ -244,25 +221,19 @@ class GatewayAgentCacheMixin: # it (even as None) so switching to a provider without configured overrides clears a stale # value left by the default provider's runtime resolution. if "request_overrides" in override: - override_request_overrides = override.get("request_overrides") - if isinstance(override_request_overrides, dict) and override_request_overrides: - runtime_kwargs["request_overrides"] = dict(override_request_overrides) - else: - runtime_kwargs["request_overrides"] = override_request_overrides + ro = override.get("request_overrides") + runtime_kwargs["request_overrides"] = dict(ro) if isinstance(ro, dict) and ro else ro if ( runtime_kwargs.get("api_key") and runtime_kwargs.get("credential_pool") is None and override.get("provider") ): - runtime_kwargs["credential_pool"] = _credential_pool_for_provider( - override.get("provider") - ) + runtime_kwargs["credential_pool"] = _credential_pool_for_provider(override.get("provider")) return model, runtime_kwargs def _snapshot_session_model_override(self, session_key: str) -> dict: """Capture a gateway session override before a one-turn switch.""" - _snap_state = self._peek_session_state(session_key) - override = _snap_state.conversation.model_override if _snap_state else None + override = self._session_model_override(session_key) return { "had_override": override is not None, "override": dict(override) if override is not None else None, @@ -273,9 +244,7 @@ class GatewayAgentCacheMixin: if not session_key: return if snapshot.get("had_override"): - self._session_state(session_key).conversation.model_override = dict( - snapshot.get("override") or {} - ) + self._session_state(session_key).conversation.model_override = dict(snapshot.get("override") or {}) else: _rst_state = self._peek_session_state(session_key) if _rst_state is not None: @@ -284,15 +253,11 @@ class GatewayAgentCacheMixin: def _is_intentional_model_switch(self, session_key: str, agent_model: str) -> bool: """Return True if *agent_model* matches an active /model session override.""" - _ims_state = self._peek_session_state(session_key) - override = _ims_state.conversation.model_override if _ims_state else None + override = self._session_model_override(session_key) return override is not None and override.get("model") == agent_model def _release_running_agent_state( - self, - session_key: str, - *, - run_generation: Optional[int] = None, + self, session_key: str, *, run_generation: Optional[int] = None ) -> bool: """Pop ALL per-running-agent state entries for ``session_key``; True when cleared. @@ -303,9 +268,7 @@ class GatewayAgentCacheMixin: """ if not session_key: return False - if run_generation is not None and not self._is_session_run_current( - session_key, run_generation - ): + if run_generation is not None and not self._is_session_run_current(session_key, run_generation): return False state = self._peek_session_state(session_key) if state is not None: @@ -314,9 +277,7 @@ class GatewayAgentCacheMixin: try: lease.release() except Exception: - logger.debug( - "Failed to release active session slot", exc_info=True - ) + logger.debug("Failed to release active session slot", exc_info=True) # One structured reset instead of a drifting pop-list. Turn-lease tokens are deliberately NOT # cleared here — _release_turn_lease owns them. state.turn.clear() @@ -325,21 +286,29 @@ class GatewayAgentCacheMixin: self._persist_active_agents() return True + def _held_turn_lease(self, session_key: str, run_generation: int): + """Return ``(registry, turn)`` when ``session_key`` holds a lease token for ``run_generation``, else None.""" + if not session_key: + return None + registry = getattr(self, "_turn_leases", None) + state = self._peek_session_state(session_key) + if state is None or registry is None: + return None + turn = state.turn + if turn.lease_token is None or turn.lease_generation != run_generation: + return None + return registry, turn + def _release_turn_lease(self, session_key: str, run_generation: int) -> bool: """Release the turn lease acquired by (``session_key``, ``run_generation``). Token map is keyed by (routing key, run generation), so a stale unwind pops only ITS token and the registry's identity check refuses it if a newer turn holds the lease. Idempotent. """ - if not session_key: - return False - registry = getattr(self, "_turn_leases", None) - state = self._peek_session_state(session_key) - if state is None or registry is None: - return False - turn = state.turn - if turn.lease_token is None or turn.lease_generation != run_generation: + held = self._held_turn_lease(session_key, run_generation) + if held is None: return False + registry, turn = held token = turn.lease_token turn.lease_token = None turn.lease_generation = None @@ -349,24 +318,17 @@ class GatewayAgentCacheMixin: logger.debug("Failed to release turn lease", exc_info=True) return False - def _rebind_turn_lease( - self, session_key: str, run_generation: int, new_session_id: str - ) -> bool: + def _rebind_turn_lease(self, session_key: str, run_generation: int, new_session_id: str) -> bool: """Follow a mid-turn session_id rotation with the held turn lease. Compression can rotate ``session_entry.session_id`` mid-turn; the flush targets the NEW id, so the serialization boundary must follow or an alias key resolving the new id could start a concurrent turn the lease never sees. Call at every mid-turn reassignment; no-op if no token. """ - if not session_key or not new_session_id: - return False - registry = getattr(self, "_turn_leases", None) - state = self._peek_session_state(session_key) - if state is None or registry is None: - return False - turn = state.turn - if turn.lease_token is None or turn.lease_generation != run_generation: + held = self._held_turn_lease(session_key, run_generation) if new_session_id else None + if held is None: return False + registry, turn = held try: return registry.rebind(turn.lease_token, new_session_id) except Exception: @@ -386,8 +348,6 @@ class GatewayAgentCacheMixin: from gateway.run import _CONVERSATION_SCOPED_STATE if not session_key: return - # Structural clear: every conversation-scoped field resets in one - # call — no per-attribute pop-list to drift. state = self._peek_session_state(session_key) if state is not None: state.conversation.clear() @@ -399,18 +359,14 @@ class GatewayAgentCacheMixin: if isinstance(store, dict): store.pop(session_key, None) self._clear_session_boundary_security_state(session_key) - logger.debug( - "Cleared conversation scope for %s (%s)", session_key, reason - ) + logger.debug("Cleared conversation scope for %s (%s)", session_key, reason) def _clear_session_boundary_security_state(self, session_key: str) -> None: """Clear per-session control state that must not survive a boundary switch.""" if not session_key: return - pending_skills_reload_notes = getattr( - self, "_pending_skills_reload_notes", None - ) + pending_skills_reload_notes = getattr(self, "_pending_skills_reload_notes", None) if isinstance(pending_skills_reload_notes, dict): pending_skills_reload_notes.pop(session_key, None) @@ -419,19 +375,12 @@ class GatewayAgentCacheMixin: _sec_state.persistent.approvals = None _sec_state.persistent.update_prompt_pending = False - try: + with suppress(Exception): from tools import slash_confirm as _slash_confirm_mod - except Exception: - _slash_confirm_mod = None - if _slash_confirm_mod is not None: try: _slash_confirm_mod.clear(session_key) except Exception as e: - logger.debug( - "Failed to clear slash-confirm state for session boundary %s: %s", - session_key, - e, - ) + logger.debug("Failed to clear slash-confirm state for session boundary %s: %s", session_key, e) try: from tools.approval import clear_session as _clear_approval_session @@ -441,11 +390,7 @@ class GatewayAgentCacheMixin: try: _clear_approval_session(session_key) except Exception as e: - logger.debug( - "Failed to clear approval state for session boundary %s: %s", - session_key, - e, - ) + logger.debug("Failed to clear approval state for session boundary %s: %s", session_key, e) def _begin_session_run_generation(self, session_key: str) -> int: """Claim a fresh, monotonically increasing run generation token for ``session_key``. @@ -464,12 +409,7 @@ class GatewayAgentCacheMixin: """Invalidate any in-flight run token for ``session_key``.""" generation = self._begin_session_run_generation(session_key) if reason: - logger.info( - "Invalidated run generation for %s → %d (%s)", - session_key, - generation, - reason, - ) + logger.info("Invalidated run generation for %s → %d (%s)", session_key, generation, reason) return generation def _is_session_run_current(self, session_key: str, generation: int) -> bool: @@ -480,12 +420,7 @@ class GatewayAgentCacheMixin: current = state.persistent.run_generation if state is not None else 0 return int(current) == int(generation) - def _bind_adapter_run_generation( - self, - adapter: Any, - session_key: str, - generation: int | None, - ) -> None: + def _bind_adapter_run_generation(self, adapter: Any, session_key: str, generation: int | None) -> None: """Bind a gateway run generation to the adapter's active-session event.""" if not adapter or not session_key or generation is None: return @@ -497,13 +432,8 @@ class GatewayAgentCacheMixin: pass async def _interrupt_and_clear_session( - self, - session_key: str, - source: SessionSource, - *, - interrupt_reason: str, - invalidation_reason: str, - release_running_state: bool = True, + self, session_key: str, source: SessionSource, *, interrupt_reason: str, + invalidation_reason: str, release_running_state: bool = True, ) -> None: """Interrupt the current run and clear queued session state consistently.""" from gateway.run import _AGENT_PENDING_SENTINEL, _reap_gateway_turn_processes, request_hard_interrupt @@ -515,50 +445,37 @@ class GatewayAgentCacheMixin: _process_baseline = None if running_agent and running_agent is not _AGENT_PENDING_SENTINEL: request_hard_interrupt(running_agent, interrupt_reason) - _process_task_id = getattr( - running_agent, "_gateway_turn_process_task_id", "" - ) - _process_baseline = getattr( - running_agent, "_gateway_turn_process_baseline", None - ) + _process_task_id = getattr(running_agent, "_gateway_turn_process_task_id", "") + _process_baseline = getattr(running_agent, "_gateway_turn_process_baseline", None) # Bump the generation BEFORE scheduling the reap thread and capture the post-bump value: # task_id is session-scoped, so a replacement turn spawning before the reap runs bumps it # again and the closure sees a stale generation and skips — the replacement's own baseline # covers its cleanup, so nothing stays unreaped. - _generation_at_interrupt = self._invalidate_session_run_generation( - session_key, reason=invalidation_reason - ) + _generation_at_interrupt = self._invalidate_session_run_generation(session_key, reason=invalidation_reason) if _process_task_id and _process_baseline is not None: threading.Thread( target=_reap_gateway_turn_processes, args=(_process_task_id, _process_baseline), kwargs={ "source": "gateway_turn_interrupt", - "is_still_current": lambda: self._is_session_run_current( - session_key, _generation_at_interrupt - ), + "is_still_current": lambda: self._is_session_run_current(session_key, _generation_at_interrupt), }, name=f"gateway-turn-reaper-{_process_task_id[:12]}", daemon=True, ).start() adapter = self._adapter_for_source(source) - interrupt_session_activity = getattr( - type(adapter), "interrupt_session_activity", None - ) + interrupt_session_activity = getattr(type(adapter), "interrupt_session_activity", None) if adapter and callable(interrupt_session_activity): metadata = self._thread_metadata_for_source(source) try: params = inspect.signature(interrupt_session_activity).parameters accepts_metadata = "metadata" in params or any( - param.kind is inspect.Parameter.VAR_KEYWORD - for param in params.values() + param.kind is inspect.Parameter.VAR_KEYWORD for param in params.values() ) except (TypeError, ValueError): accepts_metadata = False if accepts_metadata: - await adapter.interrupt_session_activity( - session_key, source.chat_id, metadata=metadata - ) + await adapter.interrupt_session_activity(session_key, source.chat_id, metadata=metadata) else: await adapter.interrupt_session_activity(session_key, source.chat_id) if adapter and hasattr(adapter, "get_pending_message"): @@ -573,9 +490,7 @@ class GatewayAgentCacheMixin: # message rebuilds from history; the old agent keeps its flag so a hung drain still dies. self._evict_cached_agent(session_key) - async def _refresh_agent_cache_message_count( - self, session_key: str, session_id: Optional[str] - ) -> None: + async def _refresh_agent_cache_message_count(self, session_key: str, session_id: Optional[str]) -> None: """Re-baseline a cached agent's stored message_count after THIS turn. The coherence guard compares on-disk ``message_count`` against the BUILD-time snapshot and @@ -602,24 +517,19 @@ class GatewayAgentCacheMixin: cached = _cache.get(session_key) # Only re-baseline a live 3-tuple entry; skip pending sentinels, legacy 2-tuples (they opt # out of the guard), and entries evicted/rebuilt mid-turn. - if ( - isinstance(cached, tuple) - and len(cached) > 2 - and cached[0] is not _AGENT_PENDING_SENTINEL - ): - # A snapshot taken for a different session_id (same session_key, different conversation) - # belongs to a different DB row — leave it alone. - _snapshot_sid = cached[3] if len(cached) > 3 else None - if _snapshot_sid is not None and _snapshot_sid != session_id: - return - if cached[2] != _live: - if _snapshot_sid is None: - # Legacy 3-tuple: preserve the 3-element shape for callers indexing ``cached[2]``. - _cache[session_key] = (cached[0], cached[1], _live) - else: - _cache[session_key] = ( - cached[0], cached[1], _live, _snapshot_sid, - ) + if not (isinstance(cached, tuple) and len(cached) > 2 and cached[0] is not _AGENT_PENDING_SENTINEL): + return + # A snapshot taken for a different session_id (same session_key, different conversation) + # belongs to a different DB row — leave it alone. + _snapshot_sid = cached[3] if len(cached) > 3 else None + if _snapshot_sid is not None and _snapshot_sid != session_id: + return + if cached[2] != _live: + # Legacy 3-tuple keeps its 3-element shape for callers indexing ``cached[2]``. + _cache[session_key] = ( + (cached[0], cached[1], _live) if _snapshot_sid is None + else (cached[0], cached[1], _live, _snapshot_sid) + ) def _set_pending_turn_sidecar_notes(self, session_key: str, notes: List[str]) -> None: """Stage per-turn must-deliver notes for the next agent run (one-shot).""" @@ -664,9 +574,7 @@ class GatewayAgentCacheMixin: return "[Voice channel now: not connected to a voice channel]" return f"[Voice channel now: {vc_now}]" - def _pinned_session_context_prompt( - self, context, redact_pii: bool, session_key: Optional[str] - ) -> str: + def _pinned_session_context_prompt(self, context, redact_pii: bool, session_key: Optional[str]) -> str: """Return the session-context prompt, pinned per session. Key hit → pinned bytes reused VERBATIM (immune to renderer nondeterminism); key miss → @@ -681,10 +589,7 @@ class GatewayAgentCacheMixin: return _eph_pin[1] text = build_session_context_prompt(context, redact_pii=redact_pii) if session_key: - self._session_state(session_key).conversation.ephemeral_pin = ( - _eph_key, - text, - ) + self._session_state(session_key).conversation.ephemeral_pin = (_eph_key, text) return text @staticmethod @@ -747,11 +652,7 @@ class GatewayAgentCacheMixin: slack_tools, tuple(p.value for p in context.connected_platforms), tuple( - ( - p.value, - str(getattr(hc, "name", "") or ""), - str(getattr(hc, "chat_id", "") or ""), - ) + (p.value, str(getattr(hc, "name", "") or ""), str(getattr(hc, "chat_id", "") or "")) for p, hc in context.home_channels.items() ), bool(redact_pii), @@ -777,39 +678,51 @@ class GatewayAgentCacheMixin: _evict_state.conversation.ephemeral_pin = None _evict_state.conversation.vc_last = None + # Tests build runners with ``_agent_cache_lock = None``; evict lock-free then. _lock = getattr(self, "_agent_cache_lock", None) + _cache = getattr(self, "_agent_cache", None) evicted = None - if _lock: - with _lock: - evicted = self._agent_cache.pop(session_key, None) - else: - _cache = getattr(self, "_agent_cache", None) - if _cache is not None: + if _cache is not None: + with _lock or nullcontext(): evicted = _cache.pop(session_key, None) - agent = evicted[0] if isinstance(evicted, tuple) and evicted else evicted + agent = _first_agent(evicted) if agent is None or agent is _AGENT_PENDING_SENTINEL: return # Don't tear down an agent that's actively mid-turn — its client, # sandbox and child subagents are in use by the running request. - running_ids = self._running_agent_ids() - if id(agent) in running_ids: + if id(agent) in self._running_agent_ids(): return try: threading.Thread( - target=self._release_evicted_agent_soft, - args=(agent,), - daemon=True, + target=self._release_evicted_agent_soft, args=(agent,), daemon=True, name=f"agent-evict-{str(session_key)[:24]}", ).start() except Exception: - # If we can't spawn a thread (interpreter shutdown), release - # inline as a best-effort fallback. + # Can't spawn a thread (interpreter shutdown) — release inline as a best-effort fallback. with suppress(Exception): self._release_evicted_agent_soft(agent) + def _finalizable_unexpired_session_entry(self, key: str): + """Return the session-store entry for ``key`` when the expiry watcher will still finalize it. + + None when the store/entry is missing, the session is not finalizable (``mode == "none"`` + never finalizes) or it has already expired (the watcher tears those down itself). + """ + _store = getattr(self, "session_store", None) + if _store is None: + return None + try: + _store._ensure_loaded() + entry = _store._entries.get(key) + except Exception: + return None + if entry is None or not _store.is_session_finalizable(entry) or _store._is_session_expired(entry): + return None + return entry + def _commit_memory_before_soft_evict(self, agent: Any, key: str) -> None: """Fire on_session_end extraction before soft-evicting a live agent. @@ -817,31 +730,21 @@ class GatewayAgentCacheMixin: ``_session_expiry_watcher``'s job at true expiry. But the watcher tears down whatever it finds in ``_agent_cache``; if the LRU cap soft-evicts first, memory providers never see the transcript. So commit extraction here via ``commit_memory_session`` (no teardown). Only for - finalizable sessions — ``mode == "none"`` never finalizes. Best-effort: failures swallowed. + finalizable, not-yet-expired sessions. Best-effort: failures swallowed. """ if agent is None or not hasattr(agent, "commit_memory_session"): return if getattr(agent, "_memory_manager", None) is None: return # no external memory provider — nothing to commit try: - _store = getattr(self, "session_store", None) - if _store is None: - return - _store._ensure_loaded() - entry = _store._entries.get(key) - if entry is None: - return - # Compensate only when the watcher would expect this agent at expiry (finite policy, not yet - # expired). Expired sessions are torn down by the watcher; mode="none" is never finalized. - if not _store.is_session_finalizable(entry): - return - if _store._is_session_expired(entry): + if self._finalizable_unexpired_session_entry(key) is None: return messages = getattr(agent, "_session_messages", None) agent.commit_memory_session(messages if isinstance(messages, list) else None) logger.debug( "Committed on_session_end extraction before soft-evicting " - "finalizable session=%s (cache pressure, pre-expiry)", key, + "finalizable session=%s (cache pressure, pre-expiry)", + key, ) except Exception as _e: logger.debug("Pre-evict memory commit failed for %s: %s", key, _e) @@ -867,8 +770,7 @@ class GatewayAgentCacheMixin: if hasattr(agent, "release_clients"): agent.release_clients() else: - # Older agent instance (shouldn't happen in practice) — - # fall back to the legacy full-close path. + # Older agent instance (shouldn't happen in practice) — legacy full-close path. self._cleanup_agent_resources(agent) except Exception: pass @@ -908,14 +810,12 @@ class GatewayAgentCacheMixin: def _agent_cache_cap(self) -> int: """Effective LRU cap — the configured override, else the default.""" from gateway.run import _AGENT_CACHE_MAX_SIZE - configured = self._agent_cache_bounds().max_size - return configured if configured else _AGENT_CACHE_MAX_SIZE + return self._agent_cache_bounds().max_size or _AGENT_CACHE_MAX_SIZE def _agent_cache_idle_ttl(self) -> float: """Effective idle TTL in seconds — configured override, else default.""" from gateway.run import _AGENT_CACHE_IDLE_TTL_SECS - configured = self._agent_cache_bounds().idle_ttl_secs - return configured if configured else _AGENT_CACHE_IDLE_TTL_SECS + return self._agent_cache_bounds().idle_ttl_secs or _AGENT_CACHE_IDLE_TTL_SECS def _sweep_agent_cache_under_pressure(self) -> int: """Shed cached transcripts once the gateway heap nears its budget; returns count evicted. @@ -928,9 +828,7 @@ class GatewayAgentCacheMixin: """ from gateway.run import _AGENT_PENDING_SENTINEL from gateway.agent_cache_pressure import ( - plan_pressure_evictions, - read_anon_rss_mb, - transcript_persistence_caught_up, + plan_pressure_evictions, read_anon_rss_mb, transcript_persistence_caught_up ) bounds = self._agent_cache_bounds() @@ -939,8 +837,8 @@ class GatewayAgentCacheMixin: _cache = getattr(self, "_agent_cache", None) _lock = getattr(self, "_agent_cache_lock", None) if not _cache or _lock is None: - # Nothing cached — whatever is using the heap, it isn't us, and - # warning about it every tick would point at the wrong subsystem. + # Nothing cached — whatever is using the heap, it isn't us, and warning about it every + # tick would point at the wrong subsystem. return 0 rss_mb = read_anon_rss_mb() @@ -949,22 +847,16 @@ class GatewayAgentCacheMixin: running_ids = self._running_agent_ids() + def _is_live(agent: Any) -> bool: + return agent is not None and agent is not _AGENT_PENDING_SENTINEL and id(agent) not in running_ids + def _is_evictable(key: str, agent: Any) -> bool: - if agent is None or agent is _AGENT_PENDING_SENTINEL: - return False - if id(agent) in running_ids: - return False - return transcript_persistence_caught_up(agent) + return _is_live(agent) and transcript_persistence_caught_up(agent) with _lock: - ordered = [ - (key, entry[0] if isinstance(entry, tuple) and entry else entry) - for key, entry in _cache.items() - ] + ordered = [(key, _first_agent(entry)) for key, entry in _cache.items()] plan = plan_pressure_evictions( - ordered, - is_evictable=_is_evictable, - max_evictions=bounds.max_evictions_per_pass, + ordered, is_evictable=_is_evictable, max_evictions=bounds.max_evictions_per_pass, protect_recent=bounds.protect_recent, ) for key, _ in plan: @@ -972,14 +864,7 @@ class GatewayAgentCacheMixin: if not plan: _mid_turn = sum(1 for _, a in ordered if a is not None and id(a) in running_ids) - _unflushed = sum( - 1 - for _, a in ordered - if a is not None - and a is not _AGENT_PENDING_SENTINEL - and id(a) not in running_ids - and not transcript_persistence_caught_up(a) - ) + _unflushed = sum(1 for _, a in ordered if _is_live(a) and not transcript_persistence_caught_up(a)) logger.warning( "Agent cache pressure: anon RSS %dMB over budget %dMB but no " "evictable session (%d cached, %d mid-turn, %d blocked on " @@ -999,20 +884,16 @@ class GatewayAgentCacheMixin: logger.warning( "Agent cache pressure: anon RSS %dMB over budget %dMB — evicting " "%d LRU session(s): %s", - rss_mb, bounds.memory_high_mb, evicted_count, - ", ".join(key for key, _ in plan), + rss_mb, bounds.memory_high_mb, evicted_count, ", ".join(key for key, _ in plan), ) try: threading.Thread( - target=self._release_pressure_batch, - args=(plan,), - daemon=True, - name="agent-cache-pressure", + target=self._release_pressure_batch, args=(plan,), daemon=True, name="agent-cache-pressure", ).start() except Exception: self._release_pressure_batch(plan) - # NOTE: _release_pressure_batch drains `plan` in place (so the trim runs with no lingering - # agent refs) — len(plan) is 0 once the daemon thread finishes, hence the pre-captured count. + # _release_pressure_batch drains `plan` in place (so the trim runs with no lingering agent + # refs) — len(plan) is 0 once the daemon thread finishes, hence the pre-captured count. return evicted_count def _release_pressure_batch(self, plan: List[tuple]) -> None: @@ -1047,8 +928,8 @@ class GatewayAgentCacheMixin: _cache = getattr(self, "_agent_cache", None) if _cache is None: return - # OrderedDict.popitem(last=False) pops oldest; plain dict lacks the - # arg so skip enforcement if a test fixture swapped the cache type. + # OrderedDict.popitem(last=False) pops oldest; plain dict lacks the arg so skip enforcement + # if a test fixture swapped the cache type. if not hasattr(_cache, "move_to_end"): return @@ -1062,14 +943,12 @@ class GatewayAgentCacheMixin: cap = self._agent_cache_cap() excess = max(0, len(_cache) - cap) evict_plan: List[tuple] = [] # [(key, agent), ...] - if excess > 0: - ordered_keys = list(_cache.keys()) - for key in ordered_keys[:excess]: - entry = _cache.get(key) - agent = entry[0] if isinstance(entry, tuple) and entry else None - if agent is not None and id(agent) in running_ids: - continue # active mid-turn; don't evict, don't substitute - evict_plan.append((key, agent)) + for key in list(_cache.keys())[:excess]: + entry = _cache.get(key) + agent = entry[0] if isinstance(entry, tuple) and entry else None + if agent is not None and id(agent) in running_ids: + continue # active mid-turn; don't evict, don't substitute + evict_plan.append((key, agent)) for key, _ in evict_plan: _cache.pop(key, None) @@ -1083,17 +962,12 @@ class GatewayAgentCacheMixin: ) for key, agent in evict_plan: - logger.info( - "Agent cache at cap; evicting LRU session=%s (cache_size=%d)", - key, len(_cache), - ) + logger.info("Agent cache at cap; evicting LRU session=%s (cache_size=%d)", key, len(_cache)) if agent is not None: # Commit end-of-session memory, then soft-release, both on the daemon thread so the # (possibly network-bound) provider call never blocks the held cache lock. threading.Thread( - target=self._commit_then_release_soft, - args=(agent, key), - daemon=True, + target=self._commit_then_release_soft, args=(agent, key), daemon=True, name=f"agent-cache-evict-{key[:24]}", ).start() @@ -1114,38 +988,21 @@ class GatewayAgentCacheMixin: with _lock: for key, entry in list(_cache.items()): agent = entry[0] if isinstance(entry, tuple) and entry else None - if agent is None: - continue - if id(agent) in running_ids: + if agent is None or id(agent) in running_ids: continue # mid-turn — don't tear it down last_activity = getattr(agent, "_last_activity_ts", None) - if last_activity is None: + if last_activity is None or (now - last_activity) <= idle_ttl: continue - if (now - last_activity) > idle_ttl: - # If the session hasn't actually expired in the store (e.g. daily-reset fires hours - # after the last message), keep the agent cached so the expiry watcher can still find - # it and call on_session_end() with the live transcript. BUT only defer when the - # watcher will EVER finalize it: for mode == "none" (is_session_finalizable() False) - # deferring pins the agent for the gateway's lifetime — the leak this sweep relieves. - # Those fall through to soft eviction WITHOUT on_session_end, correctly (never a - # session-end boundary). Finite sessions evicted under LRU-cap pressure are covered - # by _commit_memory_before_soft_evict on the cap path. - session_entry = None - _store = getattr(self, "session_store", None) - try: - if _store is not None: - _store._ensure_loaded() - session_entry = _store._entries.get(key) - except Exception: - session_entry = None - if ( - session_entry is not None - and _store is not None - and _store.is_session_finalizable(session_entry) - and not _store._is_session_expired(session_entry) - ): - continue # keep agent — finite session hasn't expired - to_evict.append((key, agent)) + # If the session hasn't actually expired in the store (e.g. daily-reset fires hours + # after the last message), keep the agent cached so the expiry watcher can still find + # it and call on_session_end() with the live transcript. BUT only defer when the + # watcher will EVER finalize it: for mode == "none" deferring pins the agent for the + # gateway's lifetime — the leak this sweep relieves. Those fall through to soft + # eviction WITHOUT on_session_end, correctly (never a session-end boundary). Finite + # sessions evicted under LRU-cap pressure are covered by _commit_memory_before_soft_evict. + if self._finalizable_unexpired_session_entry(key) is not None: + continue # keep agent — finite session hasn't expired + to_evict.append((key, agent)) for key, _ in to_evict: _cache.pop(key, None) for key, agent in to_evict: @@ -1154,9 +1011,7 @@ class GatewayAgentCacheMixin: key, now - getattr(agent, "_last_activity_ts", now), ) threading.Thread( - target=self._release_evicted_agent_soft, - args=(agent,), - daemon=True, + target=self._release_evicted_agent_soft, args=(agent,), daemon=True, name=f"agent-cache-idle-{key[:24]}", ).start() return len(to_evict) diff --git a/gateway/run_notifications.py b/gateway/run_notifications.py index c527c53204..4e68811cac 100644 --- a/gateway/run_notifications.py +++ b/gateway/run_notifications.py @@ -259,10 +259,7 @@ class GatewayNotificationsMixin: return switched async def _deliver_media_from_response( - self, - response: str, - event: MessageEvent, - adapter, + self, response: str, event: MessageEvent, adapter, thread_metadata: Optional[Dict[str, Any]] = None, ) -> None: """Extract explicit MEDIA: tags from an already-streamed response and deliver them. @@ -335,15 +332,9 @@ class GatewayNotificationsMixin: logger.warning("Post-stream media extraction failed: %s", e) async def _deliver_queued_first_response( - self, - response: str, - source: SessionSource, - adapter, - metadata: Optional[Dict[str, Any]] = None, - event_message_id: Optional[str] = None, - text_already_delivered: bool = False, - deliver_media: bool = True, - stream_consumer=None, + self, response: str, source: SessionSource, adapter, + metadata: Optional[Dict[str, Any]] = None, event_message_id: Optional[str] = None, + text_already_delivered: bool = False, deliver_media: bool = True, stream_consumer=None, ) -> None: """Deliver a queued response using the normal text+attachment split.""" from gateway.run import _strip_response_attachments_for_direct_send @@ -403,8 +394,7 @@ class GatewayNotificationsMixin: claimed=_hermes_home / ".update_pending.claimed.json", output=_hermes_home / ".update_output.txt", exit_code=_hermes_home / ".update_exit_code", - prompt=_hermes_home / ".update_prompt.json", - response=_hermes_home / ".update_response", + prompt=_hermes_home / ".update_prompt.json", response=_hermes_home / ".update_response", ) def _resolve_update_target(self, paths: "_UpdatePaths") -> Optional["_UpdateTarget"]: @@ -422,10 +412,8 @@ class GatewayNotificationsMixin: platform = Platform(platform_str) adapter = self.adapters.get(platform) metadata = self._thread_metadata_for_target( - platform, chat_id, pending.get("thread_id"), - chat_type=pending.get("chat_type"), - reply_to_message_id=pending.get("message_id"), - adapter=adapter, + platform, chat_id, pending.get("thread_id"), chat_type=pending.get("chat_type"), + reply_to_message_id=pending.get("message_id"), adapter=adapter, ) if not adapter: return None @@ -513,10 +501,7 @@ class GatewayNotificationsMixin: state.persistent.update_prompt_pending = False async def _watch_update_progress( - self, - poll_interval: float = 2.0, - stream_interval: float = 4.0, - timeout: float = 1800.0, + self, poll_interval: float = 2.0, stream_interval: float = 4.0, timeout: float = 1800.0 ) -> None: """Watch ``hermes update --gateway``, streaming output + forwarding prompts. @@ -655,10 +640,8 @@ class GatewayNotificationsMixin: if adapter and chat_id: metadata = self._thread_metadata_for_target( - platform, chat_id, pending.get("thread_id"), - chat_type=pending.get("chat_type"), - reply_to_message_id=pending.get("message_id"), - adapter=adapter, + platform, chat_id, pending.get("thread_id"), chat_type=pending.get("chat_type"), + reply_to_message_id=pending.get("message_id"), adapter=adapter, ) from tools.ansi_strip import strip_ansi output = strip_ansi(output).strip() @@ -707,15 +690,14 @@ class GatewayNotificationsMixin: platform_cfg = self.config.platforms.get(platform) if platform_cfg is not None and not platform_cfg.gateway_restart_notification: logger.info( - "Restart notification suppressed: %s has gateway_restart_notification=false", platform_str, + "Restart notification suppressed: %s has gateway_restart_notification=false", + platform_str, ) return None metadata = self._thread_metadata_for_target( - platform, chat_id, thread_id, - chat_type=data.get("chat_type"), - reply_to_message_id=data.get("message_id"), - adapter=transport.adapter, + platform, chat_id, thread_id, chat_type=data.get("chat_type"), + reply_to_message_id=data.get("message_id"), adapter=transport.adapter, ) if data.get("delivered_via_upstream_relay") is True: metadata = dict(metadata or {}) @@ -723,8 +705,7 @@ class GatewayNotificationsMixin: if data.get(field): metadata[field] = str(data[field]) result = await transport.send( - platform, str(chat_id), - "♻ Gateway restarted successfully. Your session continues.", + platform, str(chat_id), "♻ Gateway restarted successfully. Your session continues.", metadata=_non_conversational_metadata(metadata, platform=platform), ) # adapter.send() catches provider errors (e.g. "Chat not found") and returns @@ -783,9 +764,7 @@ class GatewayNotificationsMixin: return False async def _send_home_channel_startup_notifications( - self, - *, - skip_targets: Optional[set[tuple[str, str, Optional[str]]]] = None, + self, *, skip_targets: Optional[set[tuple[str, str, Optional[str]]]] = None ) -> set[tuple[str, str, Optional[str]]]: """Notify configured home channels that the gateway is back online. @@ -921,13 +900,8 @@ class GatewayNotificationsMixin: platform_name, chat_id, chat_type, ) return SessionSource( - platform=platform, - chat_id=chat_id, - chat_type=chat_type, - thread_id=_opt("thread_id"), - user_id=_opt("user_id"), - user_name=_opt("user_name"), - scope_id=scope_id, + platform=platform, chat_id=chat_id, chat_type=chat_type, thread_id=_opt("thread_id"), + user_id=_opt("user_id"), user_name=_opt("user_name"), scope_id=scope_id, ) async def _drain_watch_notifications(self, completion_queue) -> None: @@ -1051,12 +1025,8 @@ class GatewayNotificationsMixin: if parent_session_id: metadata["gateway_session_id"] = parent_session_id synth_event = MessageEvent( - text=synth_text, - message_type=MessageType.TEXT, - source=source, - internal=True, - message_id=str(evt.get("message_id") or "").strip() or None, - metadata=metadata, + text=synth_text, message_type=MessageType.TEXT, source=source, internal=True, + message_id=str(evt.get("message_id") or "").strip() or None, metadata=metadata, ) logger.info( "Watch pattern notification — injecting for %s chat=%s thread=%s", @@ -1634,11 +1604,8 @@ class GatewayNotificationsMixin: _started = getattr(session, "started_at", None) _dur = max(0.0, time.time() - _started) if isinstance(_started, (int, float)) else None return _format_concise_process_notification( - session_id, - _redact_gateway_user_facing_secrets(getattr(session, "command", "") or ""), - session.exit_code, - new_output, - duration_seconds=_dur, + session_id, _redact_gateway_user_facing_secrets(getattr(session, "command", "") or ""), + session.exit_code, new_output, duration_seconds=_dur, ) async def _run_process_watcher(self, watcher: dict) -> None: