From 9ac2dd561edc125598aca3777ab1d09df5a3b554 Mon Sep 17 00:00:00 2001 From: Josh Stevenson Date: Sat, 3 Oct 2026 23:48:11 -0700 Subject: [PATCH 1/2] fix(context): preserve ownership across startup, routes and resume --- agent/agent_init.py | 31 +- agent/agent_runtime_helpers.py | 301 ++++++++++-------- agent/auxiliary_client.py | 22 +- agent/chat_completion_helpers.py | 29 +- .../_context_governor/__init__.py | 44 ++- .../test_context_governor_route_budget.py | 87 +++++ .../test_ares_context_handoff_refusal.py | 244 ++++++++++++++ .../test_ares_context_initialization.py | 229 +++++++++++++ 8 files changed, 816 insertions(+), 171 deletions(-) create mode 100644 tests/agent/test_context_governor_route_budget.py create mode 100644 tests/run_agent/test_ares_context_handoff_refusal.py create mode 100644 tests/run_agent/test_ares_context_initialization.py diff --git a/agent/agent_init.py b/agent/agent_init.py index 06b16b733eb76..0b928ae4f2a57 100644 --- a/agent/agent_init.py +++ b/agent/agent_init.py @@ -2821,7 +2821,16 @@ def _parse_prune_int(raw, default): except Exception: pass - if _engine_name != "compressor": + if _engine_name == "ri-context-governor": + # Ares explicitly selects a certified, receipt-preserving engine. + # Readiness failure must not select Hermes' LLM compressor instead. + from plugins.context_engine import ( + ContextEngineActivationError, + load_context_engine_strict, + ) + + _selected_engine = load_context_engine_strict(_engine_name) + elif _engine_name != "compressor": # Try loading from plugins/context_engine// try: from plugins.context_engine import load_context_engine @@ -2905,6 +2914,7 @@ def _parse_prune_int(raw, default): api_key=getattr(agent, "api_key", ""), provider=agent.provider, api_mode=agent.api_mode, + max_tokens=agent.max_tokens, threshold_percent=compression_threshold, ) if not agent.quiet_mode: @@ -2937,8 +2947,12 @@ def _parse_prune_int(raw, default): if callable(_bind_session_state): try: _bind_session_state(session_db=session_db, session_id=agent.session_id) - except Exception: - pass + except Exception as _ce_bind_err: + if _engine_name == "ri-context-governor": + raise ContextEngineActivationError( + f"configured context engine '{_engine_name}' failed its " + f"session binding: {_ce_bind_err}" + ) from _ce_bind_err agent.compression_enabled = compression_enabled agent.compression_in_place = compression_in_place agent.context_rebase_enabled = compression_context_rebase @@ -3075,6 +3089,11 @@ def _parse_prune_int(raw, default): # Notify context engine of session start if hasattr(agent, "context_compressor") and agent.context_compressor: + _session_start_context = ( + {"session_db": session_db} + if _engine_name == "ri-context-governor" + else {} + ) try: agent.context_compressor.on_session_start( agent.session_id, @@ -3083,8 +3102,14 @@ def _parse_prune_int(raw, default): model=agent.model, context_length=getattr(agent.context_compressor, "context_length", 0), conversation_id=getattr(agent, "_gateway_session_key", None), + **_session_start_context, ) except Exception as _ce_err: + if _engine_name == "ri-context-governor": + raise ContextEngineActivationError( + f"configured context engine '{_engine_name}' failed its " + f"session start: {_ce_err}" + ) from _ce_err _ra().logger.debug("Context engine on_session_start: %s", _ce_err) agent._subdirectory_hints = SubdirectoryHintTracker( diff --git a/agent/agent_runtime_helpers.py b/agent/agent_runtime_helpers.py index be22ac87665c3..5b86b86f0c523 100644 --- a/agent/agent_runtime_helpers.py +++ b/agent/agent_runtime_helpers.py @@ -1508,6 +1508,64 @@ def drop_thinking_only_and_merge_users( +def _snapshot_context_handoff(agent): + """Return a rollback for host state until the context engine accepts a route. + + Keep engine admission and durable state with their existing owners. This + captures only reversible host fields; fallback index/cooldown progress must + survive refusal so rejected candidates are not retried indefinitely. + """ + missing = object() + snapshot = { + name: getattr(agent, name, missing) + for name in ( + "model", "provider", "requested_provider", "base_url", "api_mode", + "api_key", "client", "_anthropic_client", "_anthropic_api_key", + "_anthropic_base_url", "_is_anthropic_oauth", "_client_kwargs", + "_credential_pool", "_credential_pool_entry_id", + "_config_context_length", "_custom_providers", "_reasoning_echo_flag", + "_use_prompt_caching", "_use_native_cache_layout", "_transport_cache", + "_fallback_activated", + ) + } + # Preserve cache/kwargs identity as well as contents: clearing a live dict + # must not poison the rollback target or a reference held by its caller. + dict_contents = { + name: dict(value) for name, value in snapshot.items() + if name in {"_client_kwargs", "_transport_cache"} and isinstance(value, dict) + } + + def restore(): + rejected_clients = ( + getattr(agent, "client", None), + getattr(agent, "_anthropic_client", None), + ) + for name, value in snapshot.items(): + if value is missing: + if hasattr(agent, name): + delattr(agent, name) + else: + if name in dict_contents: + value.clear() + value.update(dict_contents[name]) + setattr(agent, name, value) + # Retire only newly attached clients, never the restored shared ones. + # Socket shutdown defers FD release to GC instead of hard-closing a + # shared client from a thread whose earlier request may still unwind. + preserved = (snapshot["client"], snapshot["_anthropic_client"]) + retired = set() + for client in rejected_clients: + if client is None or any(client is old for old in preserved) or id(client) in retired: + continue + retired.add(id(client)) + try: + agent._retire_shared_openai_client(client, reason="context_handoff_refused") + except Exception: + logger.debug("Rejected context handoff client retirement failed", exc_info=True) + + return restore + + def restore_primary_runtime(agent) -> bool: """Restore the primary runtime at the start of a new turn. @@ -1622,6 +1680,8 @@ def restore_primary_runtime(agent) -> bool: provider_fallback_active = bool( getattr(agent, "_provider_fallback_active", False) ) + restore_handoff = _snapshot_context_handoff(agent) + handoff_accepted = False try: # ── Core runtime state ── agent.model = rt["model"] @@ -1679,14 +1739,19 @@ def restore_primary_runtime(agent) -> bool: # ── Restore context engine state ── cc = agent.context_compressor - cc.update_model( + from agent.auxiliary_client import _update_compressor_model + + _update_compressor_model( + cc, model=rt["compressor_model"], context_length=rt["compressor_context_length"], base_url=rt["compressor_base_url"], api_key=rt["compressor_api_key"], provider=rt["compressor_provider"], api_mode=rt.get("compressor_api_mode", ""), + max_tokens=getattr(agent, "max_tokens", None), ) + handoff_accepted = True # ── Rebind and re-select the primary credential pool ── # A cross-provider fallback attaches the fallback provider's pool. The @@ -1824,6 +1889,8 @@ def restore_primary_runtime(agent) -> bool: pass return True except Exception as e: + if not handoff_accepted: + restore_handoff() logger.warning("Failed to restore primary runtime: %s", e) return False @@ -2768,56 +2835,10 @@ def switch_model(agent, new_model, new_provider, api_key='', base_url='', api_mo old_model = agent.model old_provider = agent.provider - # ── Snapshot all fields the swap+rebuild can mutate ── - # If the rebuild raises (bad API key, network error, build_anthropic_client - # failure, etc.) we restore these atomically so the agent isn't left with a - # new model/provider name paired with the OLD client — that mismatch causes - # HTTP 400s like "claude-sonnet-4-6 is not supported on openai-codex" on the - # next turn. Callers in cli.py / gateway/run.py / tui_gateway/server.py - # catch the re-raised exception and show the user a warning; without this - # rollback the warning is misleading because the swap partially succeeded. - # Use a sentinel so we can distinguish "attribute was unset" from - # "attribute was None" and skip the restore for genuinely-missing - # attributes (tests construct bare agents via __new__ without all fields). - _MISSING = object() - _snapshot = { - name: getattr(agent, name, _MISSING) - for name in ( - "model", - "provider", - "requested_provider", - "base_url", - "api_mode", - "api_key", - "client", - "_anthropic_client", - "_anthropic_api_key", - "_anthropic_base_url", - "_is_anthropic_oauth", - "_config_context_length", - "_reasoning_echo_flag", - ) - } - # _client_kwargs is a dict — snapshot a shallow copy so mutating the - # live dict doesn't poison the rollback target. - _snapshot["_client_kwargs"] = dict(getattr(agent, "_client_kwargs", {}) or {}) - # Snapshot the credential pool reference so a failed client rebuild can - # restore the original pool (issue #52727: pool reload is part of this - # switch and must be reversible on rollback). - _snapshot["_credential_pool"] = getattr(agent, "_credential_pool", _MISSING) - _snapshot["_credential_pool_entry_id"] = getattr( - agent, "_credential_pool_entry_id", _MISSING - ) - - def _restore_snapshot() -> None: - for _name, _value in _snapshot.items(): - if _value is _MISSING: - # Attribute did not exist before the swap — don't fabricate it. - continue - try: - setattr(agent, _name, _value) - except Exception: # noqa: BLE001 - pass + # Callers retain this agent after a failed /model swap and rely on its + # previous route surviving. Include cache/pool state and preserve that + # contract until the context engine has accepted the destination budget. + restore_handoff = _snapshot_context_handoff(agent) try: # Clear the per-config context_length override so the new model's @@ -2990,102 +3011,108 @@ def _restore_snapshot() -> None: # caller's exception handler can surface a meaningful warning. The # exception is re-raised; cli.py / gateway/run.py / tui_gateway catch # it and print "Agent swap failed; change applied to next session". - _restore_snapshot() + restore_handoff() raise - # ── LM Studio: preload before probing context length ── - _sm_custom_providers = None try: - from hermes_cli.config import ( - get_compatible_custom_providers, - get_custom_provider_context_length, - load_config, - ) + # ── LM Studio: preload before probing context length ── + _sm_custom_providers = None + try: + from hermes_cli.config import ( + get_compatible_custom_providers, + get_custom_provider_context_length, + load_config, + ) - _sm_cfg = load_config() - _sm_custom_providers = get_compatible_custom_providers(_sm_cfg) - _destination_context_intent = get_custom_provider_context_length( - model=agent.model, - base_url=agent.base_url, - custom_providers=_sm_custom_providers, + _sm_cfg = load_config() + _sm_custom_providers = get_compatible_custom_providers(_sm_cfg) + _destination_context_intent = get_custom_provider_context_length( + model=agent.model, + base_url=agent.base_url, + custom_providers=_sm_custom_providers, + ) + except Exception: + _destination_context_intent = None + agent._config_context_length = _destination_context_intent + _runtime_context_length = agent._ensure_lmstudio_runtime_loaded( + _destination_context_intent ) - except Exception: - _destination_context_intent = None - agent._config_context_length = _destination_context_intent - _runtime_context_length = agent._ensure_lmstudio_runtime_loaded( - _destination_context_intent - ) - if agent._lmstudio_load_was_unverified(_runtime_context_length): - logger.warning( - "LM Studio model activation was rejected or completed without a " - "verifiable active context length during model switch; continuing " - "with configured context" + if agent._lmstudio_load_was_unverified(_runtime_context_length): + logger.warning( + "LM Studio model activation was rejected or completed without a " + "verifiable active context length during model switch; continuing " + "with configured context" + ) + _effective_context_length = agent._effective_lmstudio_context_length( + _destination_context_intent, + _runtime_context_length, ) - _effective_context_length = agent._effective_lmstudio_context_length( - _destination_context_intent, - _runtime_context_length, - ) - # ── Re-evaluate prompt caching ── - # Refresh the custom-provider snapshot from the config just loaded above - # so the per-model ``prompt_caching`` capability lookup sees the same - # live list the context-length resolution used — without this, a flag - # added to config.yaml after session start is invisible to a /model - # switch (the policy would read the stale init-time snapshot). - if _sm_custom_providers is not None: - agent._custom_providers = _sm_custom_providers - agent._use_prompt_caching, agent._use_native_cache_layout = ( - agent._anthropic_prompt_cache_policy( - provider=new_provider, - base_url=agent.base_url, - api_mode=api_mode, - model=new_model, + # ── Re-evaluate prompt caching ── + # Refresh the custom-provider snapshot from the config just loaded above + # so the per-model ``prompt_caching`` capability lookup sees the same + # live list the context-length resolution used — without this, a flag + # added to config.yaml after session start is invisible to a /model + # switch (the policy would read the stale init-time snapshot). + if _sm_custom_providers is not None: + agent._custom_providers = _sm_custom_providers + agent._use_prompt_caching, agent._use_native_cache_layout = ( + agent._anthropic_prompt_cache_policy( + provider=new_provider, + base_url=agent.base_url, + api_mode=api_mode, + model=new_model, + ) ) - ) - # ── Update context compressor ── - if hasattr(agent, "context_compressor") and agent.context_compressor: - from agent.model_metadata import get_model_context_length - if _sm_custom_providers is None: - try: - from hermes_cli.config import get_compatible_custom_providers, load_config - _sm_custom_providers = get_compatible_custom_providers(load_config()) - except Exception: - _sm_custom_providers = None - # ``agent.api_key`` may be a callable (Azure Foundry Entra ID - # token provider). ``get_model_context_length`` expects a - # string for its live-probe paths; for Foundry the context - # length normally resolves via config or static catalogs and - # never hits a probe, but coerce to empty string defensively. - _ctx_api_key = agent.api_key if isinstance(agent.api_key, str) else "" - new_context_length = get_model_context_length( - agent.model, - base_url=agent.base_url, - api_key=_ctx_api_key, - provider=agent.provider, - config_context_length=_effective_context_length, - custom_providers=_sm_custom_providers, - ) - # Forward the per-model resolved threshold (Codex gpt-5.x autoraise - # included) to engines that accept threshold_percent; the built-in - # compressor re-resolves internally and is skipped by the guard. - from agent.auxiliary_client import ( - _effective_compression_threshold_percent, - _update_compressor_model, - ) + # ── Update context compressor ── + if hasattr(agent, "context_compressor") and agent.context_compressor: + from agent.model_metadata import get_model_context_length + if _sm_custom_providers is None: + try: + from hermes_cli.config import get_compatible_custom_providers, load_config + _sm_custom_providers = get_compatible_custom_providers(load_config()) + except Exception: + _sm_custom_providers = None + # ``agent.api_key`` may be a callable (Azure Foundry Entra ID + # token provider). ``get_model_context_length`` expects a + # string for its live-probe paths; for Foundry the context + # length normally resolves via config or static catalogs and + # never hits a probe, but coerce to empty string defensively. + _ctx_api_key = agent.api_key if isinstance(agent.api_key, str) else "" + new_context_length = get_model_context_length( + agent.model, + base_url=agent.base_url, + api_key=_ctx_api_key, + provider=agent.provider, + config_context_length=_effective_context_length, + custom_providers=_sm_custom_providers, + ) + # Forward the per-model resolved threshold (Codex gpt-5.x autoraise + # included) to engines that accept threshold_percent; the built-in + # compressor re-resolves internally and is skipped by the guard. + from agent.auxiliary_client import ( + _effective_compression_threshold_percent, + _update_compressor_model, + ) - _update_compressor_model( - agent.context_compressor, - model=agent.model, - context_length=new_context_length, - base_url=agent.base_url, - api_key=agent.api_key, # context_compressor forwards to call_llm; callable preserved - provider=agent.provider, - api_mode=agent.api_mode, - threshold_percent=_effective_compression_threshold_percent( - agent.model, agent.provider - ), - ) + _update_compressor_model( + agent.context_compressor, + model=agent.model, + context_length=new_context_length, + base_url=agent.base_url, + api_key=agent.api_key, # context_compressor forwards to call_llm; callable preserved + provider=agent.provider, + api_mode=agent.api_mode, + max_tokens=getattr(agent, "max_tokens", None), + threshold_percent=_effective_compression_threshold_percent( + agent.model, agent.provider + ), + ) + + except Exception: + restore_handoff() + raise # ── Re-resolve reasoning_config from per-model override ── # The new model may have a different reasoning_effort override. Re-read diff --git a/agent/auxiliary_client.py b/agent/auxiliary_client.py index a7b26a08a703c..bc1b0167a601c 100644 --- a/agent/auxiliary_client.py +++ b/agent/auxiliary_client.py @@ -836,15 +836,17 @@ def _update_compressor_model( api_key: Any = "", provider: str = "", api_mode: str = "", + max_tokens: Optional[int] = None, threshold_percent: Optional[float] = None, ) -> None: - """Call ``update_model``, forwarding ``threshold_percent`` when supported. + """Forward route budget fields accepted by the engine's ``update_model``. The built-in ContextCompressor re-resolves its threshold internally and does not accept the kwarg; external engines (e.g. ri-context-governor) accept it so the resolved host threshold (including the Codex gpt-5.x - autoraise) reaches their trigger. A signature guard keeps every engine - on its own contract. + autoraise) reaches their trigger. The response reserve is also forwarded + when explicitly supported, including None to clear an old reservation. + A signature guard keeps every engine on its own contract. """ _kwargs = { "model": model, @@ -854,13 +856,15 @@ def _update_compressor_model( "provider": provider, "api_mode": api_mode, } - if threshold_percent is not None: - try: - import inspect + try: + import inspect - _params = inspect.signature(compressor.update_model).parameters - except (TypeError, ValueError): - _params = {} + _params = inspect.signature(compressor.update_model).parameters + except (TypeError, ValueError): + _params = {} + if "max_tokens" in _params: + _kwargs["max_tokens"] = max_tokens + if threshold_percent is not None: if "threshold_percent" in _params: _kwargs["threshold_percent"] = threshold_percent compressor.update_model(**_kwargs) diff --git a/agent/chat_completion_helpers.py b/agent/chat_completion_helpers.py index 39a72462d3306..d1376ab63f1f2 100644 --- a/agent/chat_completion_helpers.py +++ b/agent/chat_completion_helpers.py @@ -2472,7 +2472,9 @@ def _fallback_reason_text(reason: "FailoverReason | None") -> str: return str(value or reason or "provider failure").replace("_", " ") -def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool: +def try_activate_fallback( + agent, reason: "FailoverReason | None" = None, *, _continuing_chain: bool = False +) -> bool: """Switch to the next fallback model/provider in the chain. Called when the current model is failing after retries. Swaps the @@ -2484,7 +2486,10 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool auth resolution and client construction — no duplicated provider→key mappings. """ - if reason in {FailoverReason.rate_limit, FailoverReason.billing, FailoverReason.upstream_rate_limit}: + # Restoring the primary after an inadmissible candidate must not turn the + # internal chain walk into another originating provider failure. Separate + # public calls still arm backoff even when the chain is already exhausted. + if not _continuing_chain and reason in {FailoverReason.rate_limit, FailoverReason.billing, FailoverReason.upstream_rate_limit}: # Only start cooldown when leaving the primary provider. If we're # already on a fallback and chain-switching, the primary wasn't the # source of the 429 so the cooldown should not be reset/extended. @@ -2532,11 +2537,11 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool agent._unavailable_fallback_keys = unavailable if fb_key in unavailable: logger.debug("Fallback skip: %s previously marked unavailable", fb_key) - return agent._try_activate_fallback(reason) + return try_activate_fallback(agent, reason, _continuing_chain=True) fb_provider = (fb.get("provider") or "").strip().lower() fb_model = (fb.get("model") or "").strip() if not fb_provider or not fb_model: - return agent._try_activate_fallback(reason) # skip invalid, try next + return try_activate_fallback(agent, reason, _continuing_chain=True) # skip invalid, try next local_skip_reason = _fallback_entry_unavailable_without_network(agent, fb) if local_skip_reason: @@ -2547,7 +2552,7 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool fb_model, local_skip_reason, ) - return agent._try_activate_fallback(reason) + return try_activate_fallback(agent, reason, _continuing_chain=True) # Skip entries that resolve to the same backend that just failed — # falling back to it loops the failure. Identity semantics (which axes @@ -2572,11 +2577,15 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool "as the current one (%s)", fb_provider, fb_model, current_ident.base_url or current_ident.provider, ) - return agent._try_activate_fallback(reason) + return try_activate_fallback(agent, reason, _continuing_chain=True) # Use centralized router for client construction. # raw_codex=True because the main agent needs direct responses.stream() # access for Codex providers. + from agent.agent_runtime_helpers import _snapshot_context_handoff + + restore_handoff = _snapshot_context_handoff(agent) + handoff_accepted = False try: from agent.auxiliary_client import resolve_provider_client # Pass base_url and api_key from fallback config so custom @@ -2629,7 +2638,7 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool "Fallback to %s failed: provider not configured", fb_provider) unavailable.add(fb_key) - return agent._try_activate_fallback(reason) # try next in chain + return try_activate_fallback(agent, reason, _continuing_chain=True) # try next in chain try: from hermes_cli.model_normalize import normalize_model_for_provider @@ -2844,10 +2853,12 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool api_key=getattr(agent, "api_key", ""), # callable preserved → call_llm provider=agent.provider, api_mode=agent.api_mode, + max_tokens=getattr(agent, "max_tokens", None), threshold_percent=_effective_compression_threshold_percent( agent.model, agent.provider ), ) + handoff_accepted = True # Re-resolve reasoning_config for the new fallback model (Closes #21256). # Shared chokepoint: per-model override > global reasoning_effort @@ -2906,10 +2917,12 @@ def try_activate_fallback(agent, reason: "FailoverReason | None" = None) -> bool _reset_stale_streak(agent) return True except Exception as e: + if not handoff_accepted: + restore_handoff() if fb_provider == "nous": unavailable.add(fb_key) logger.error("Failed to activate fallback %s: %s", fb_model, e) - return agent._try_activate_fallback(reason) # try next in chain + return try_activate_fallback(agent, reason, _continuing_chain=True) # try next in chain diff --git a/plugins/context_engine/_context_governor/__init__.py b/plugins/context_engine/_context_governor/__init__.py index 9c155ca3f1aae..e83137806295d 100644 --- a/plugins/context_engine/_context_governor/__init__.py +++ b/plugins/context_engine/_context_governor/__init__.py @@ -622,6 +622,14 @@ def update_model( protect_first_n: int | None = None, protect_last_n: int | None = None, ) -> None: + resolved_context_length = int(context_length or 0) + output_reserve = int(max_tokens) if max_tokens and int(max_tokens) > 0 else None + effective_window = resolved_context_length - (output_reserve or 0) + if resolved_context_length > 0 and effective_window <= 0: + raise ValueError( + "context-governor response reservation leaves no input budget " + f"(context_length={resolved_context_length}, max_tokens={output_reserve})" + ) # Persist the active agent route. The optional summary-specific fields # override these only when explicitly configured. self.model = str(model or "") @@ -635,16 +643,11 @@ def update_model( self.protect_first_n = int(protect_first_n) if protect_last_n is not None: self.protect_last_n = int(protect_last_n) - self.context_length = int(context_length or 0) - self.max_tokens = ( - int(max_tokens) if max_tokens and int(max_tokens) > 0 else None - ) + self.context_length = resolved_context_length + self.max_tokens = output_reserve # Account for output reservation in effective input budget - effective_window = self.context_length - (self.max_tokens or 0) - if effective_window <= 0: - effective_window = self.context_length self.threshold_tokens = ( - int(effective_window * self.threshold_percent) if effective_window else 0 + int(effective_window * self.threshold_percent) if effective_window > 0 else 0 ) def update_from_response(self, usage: Dict[str, Any]) -> None: @@ -3823,17 +3826,30 @@ def _request_min_net_savings_tokens(self) -> int | None: return None def _target_tokens(self, current_tokens: int | None) -> int: + target = None explicit = self._policy.get("token_budget") try: if explicit is not None and int(explicit) > 0: - return max(512, int(explicit)) + target = max(512, int(explicit)) except (TypeError, ValueError): pass - if self.context_length: - return max(512, int(self.context_length * 0.20)) - if current_tokens: - return max(512, int(current_tokens * 0.20)) - return 8000 + if target is None: + if self.context_length: + target = max(512, int(self.context_length * 0.20)) + elif current_tokens: + target = max(512, int(current_tokens * 0.20)) + else: + target = 8000 + if self.context_length > 0: + input_window = self.context_length - (self.max_tokens or 0) + if input_window <= 0: + raise ValueError("context-governor response reservation leaves no input budget") + # Configured policy remains the requested target. The active route + # is a hard ceiling, including when its input window is below the + # adapter's usual 512-token target floor. Rust still owns admission + # and may refuse a target that cannot preserve protected content. + target = min(target, input_window) + return target def _run_json( self, diff --git a/tests/agent/test_context_governor_route_budget.py b/tests/agent/test_context_governor_route_budget.py new file mode 100644 index 0000000000000..4c98ee20fc619 --- /dev/null +++ b/tests/agent/test_context_governor_route_budget.py @@ -0,0 +1,87 @@ +"""Route admission limits and signature-compatible model budget handoffs.""" + +from unittest.mock import patch + +import pytest + +from agent.auxiliary_client import _update_compressor_model +from plugins.context_engine._context_governor import ContextGovernorEngine + + +@pytest.fixture +def governor(tmp_path): + with patch("hermes_cli.config.load_config", return_value={}): + return ContextGovernorEngine( + binary=str(tmp_path / "fake-governor"), + store_dir=str(tmp_path / "governor-store"), + ) + + +@pytest.mark.parametrize( + ("policy_budget", "window", "reserve", "expected"), + [ + (128_000, 64_000, 4096, 59_904), + (128_000, 64_000, None, 64_000), + (8000, 64_000, 4096, 8000), + (128_000, 256, 128, 128), + (None, 64_000, 4096, 12_800), + (None, 256, 128, 128), + (128_000, 0, 4096, 128_000), + (128_000, 1_050_000, 4096, 128_000), + (1_050_000, 1_050_000, None, 1_050_000), + (1_050_000, 1_048_576, 4096, 1_044_480), + ], +) +def test_policy_target_cannot_exceed_known_input_window( + governor, policy_budget, window, reserve, expected +): + governor._policy["token_budget"] = policy_budget + governor.update_model("test-model", window, max_tokens=reserve) + assert governor._target_tokens(100_000) == expected + if window: + assert expected <= window - (reserve or 0) + else: + assert governor.threshold_tokens == 0 + + +@pytest.mark.parametrize("reserve", [64_000, 64_001]) +def test_exhausted_input_budget_refuses_update_without_rebinding(governor, reserve): + governor.update_model("valid-model", 64_000, max_tokens=4096) + with pytest.raises(ValueError, match="no input budget"): + governor.update_model("invalid-model", 64_000, max_tokens=reserve) + assert governor.model == "valid-model" + assert governor.max_tokens == 4096 + assert governor.context_length == 64_000 + + +def test_handoff_forwards_and_clears_the_engine_response_reserve(governor): + governor._policy["token_budget"] = 128_000 + for reserve, target in [(4096, 59_904), (None, 64_000)]: + _update_compressor_model( + governor, + model="test-model", + context_length=64_000, + max_tokens=reserve, + threshold_percent=0.75, + ) + assert governor.max_tokens == reserve + assert governor._target_tokens(100_000) == target + + +def test_legacy_engine_receives_only_its_declared_update_api(): + calls = [] + + class LegacyEngine: + def update_model( + self, model, context_length, base_url="", api_key="", provider="", api_mode="" + ): + calls.append((model, context_length)) + + _update_compressor_model( + LegacyEngine(), + model="legacy-model", + context_length=64_000, + max_tokens=4096, + threshold_percent=0.75, + ) + assert calls == [("legacy-model", 64_000)] diff --git a/tests/run_agent/test_ares_context_handoff_refusal.py b/tests/run_agent/test_ares_context_handoff_refusal.py new file mode 100644 index 0000000000000..2920d8cf66748 --- /dev/null +++ b/tests/run_agent/test_ares_context_handoff_refusal.py @@ -0,0 +1,244 @@ +"""An Ares budget refusal must leave the owning host on its previous route. + +These follow-up counterexamples use real AIAgent handoffs and the real governor +budget validator. Provider clients, metadata, activation and native commands are +fakes; no conversation or provider request runs. +""" + +from unittest.mock import MagicMock, patch + +import pytest + +from agent.error_classifier import FailoverReason +from tests.run_agent.test_ares_context_initialization import _host_init, governor + + +_HOST_FIELDS = ( + "model", "provider", "requested_provider", "base_url", "_base_url_lower", + "api_mode", "api_key", "client", "_anthropic_client", "_anthropic_api_key", + "_anthropic_base_url", "_is_anthropic_oauth", "_client_kwargs", + "_credential_pool", "_credential_pool_entry_id", "_config_context_length", + "_custom_providers", "_reasoning_echo_flag", "_use_prompt_caching", + "_use_native_cache_layout", "_transport_cache", "_fallback_activated", + "_primary_runtime", "_cached_system_prompt", "reasoning_config", + "_provider_fallback_active", "_provider_fallback_route", + "_pending_fallback_notice", +) +_ENGINE_FIELDS = ( + "model", "provider", "base_url", "api_key", "api_mode", "context_length", + "max_tokens", "threshold_tokens", "threshold_percent", "protect_first_n", + "protect_last_n", "session_id", "_lineage_session_id", "_session_db", + "_pending_admission", "last_receipt_id", "compression_count", +) + + +@pytest.fixture(autouse=True) +def _fake_catalog_prewarm(): + # Initializer prewarming uses an imported alias in its background thread. + # Both metadata entry points stay fake, in addition to the worker guard. + with ( + patch("agent.agent_init.fetch_model_metadata", return_value={}), + patch("agent.model_metadata.fetch_model_metadata", return_value={}), + ): + yield + + +def _state(obj, fields): + values = {} + for name in fields: + if not hasattr(obj, name): + continue + value = getattr(obj, name) + if isinstance(value, dict): + value = dict(value) + elif isinstance(value, list): + value = list(value) + values[name] = value + return values + + +def _sentinel_runtime(agent): + # Make silent collateral changes observable, including in-place cache + # clearing and re-reading a custom-provider snapshot during a switch. + agent._transport_cache = {"previous-route": object()} + agent._cached_system_prompt = "stable previous prompt" + agent._credential_pool = MagicMock(provider=agent.provider) + agent._credential_pool_entry_id = "previous-credential" + agent._custom_providers = [{"name": "previous-route"}] + agent._use_prompt_caching = True + agent._use_native_cache_layout = True + agent._config_context_length = 64_000 + + +def _fallback_client(): + client = MagicMock(name="FakeFallbackClient") + client.base_url = "https://api.openai.com/v1" + client.api_key = "fake-fallback-key" + return client + + +@pytest.mark.parametrize("api_mode", ["chat_completions", "anthropic_messages"]) +def test_refused_manual_switch_preserves_host_and_engine(governor, api_mode): + with _host_init(governor) as create: + agent = create(max_tokens=4096) + _sentinel_runtime(agent) + before_host = _state(agent, _HOST_FIELDS) + before_engine = _state(governor, _ENGINE_FIELDS) + previous_client = agent.client + previous_pool = agent._credential_pool + previous_cache = agent._transport_cache + new_client = MagicMock(name="RefusedClient") + retire = MagicMock() + with ( + patch("agent.model_metadata.get_model_context_length", return_value=4096), + patch("agent.credential_pool.load_pool", return_value=None), + patch("agent.anthropic_adapter.build_anthropic_client", return_value=new_client), + patch.object(agent, "_create_openai_client", return_value=new_client), + patch.object(agent, "_retire_shared_openai_client", retire), + ): + with pytest.raises(ValueError, match="no input budget"): + agent.switch_model( + "refused-model", "openai", api_key="fake-new-key", + base_url="https://api.openai.com/v1", api_mode=api_mode, + ) + assert _state(agent, _HOST_FIELDS) == before_host + assert _state(governor, _ENGINE_FIELDS) == before_engine + assert agent.client is previous_client + assert agent._credential_pool is previous_pool + assert agent._transport_cache is previous_cache + retire.assert_called_once_with(new_client, reason="context_handoff_refused") + previous_client.close.assert_not_called() + new_client.close.assert_not_called() + new_client.chat.completions.create.assert_not_called() + new_client.messages.create.assert_not_called() + + +@pytest.mark.parametrize("already_on_fallback", [False, True]) +def test_refused_exhausted_fallback_preserves_active_runtime(governor, already_on_fallback): + with _host_init(governor) as create: + agent = create(max_tokens=4096) + if already_on_fallback: + agent._fallback_chain = [{"provider": "openai", "model": "previous-fallback"}] + with ( + patch("agent.auxiliary_client.resolve_provider_client", return_value=(_fallback_client(), None)), + patch("agent.model_metadata.get_model_context_length", return_value=32_000), + patch("agent.credential_pool.load_pool", return_value=None), + ): + assert agent._try_activate_fallback() is True + agent._fallback_index = 0 + _sentinel_runtime(agent) + agent._fallback_chain = [{"provider": "openai", "model": "refused-fallback"}] + agent._fallback_model = agent._fallback_chain[0] + before_host = _state(agent, _HOST_FIELDS) + before_engine = _state(governor, _ENGINE_FIELDS) + fallback_client = _fallback_client() + with ( + patch("agent.auxiliary_client.resolve_provider_client", return_value=(fallback_client, None)), + patch("agent.model_metadata.get_model_context_length", return_value=4096), + patch("agent.credential_pool.load_pool", return_value=None), + ): + assert agent._try_activate_fallback() is False + assert _state(agent, _HOST_FIELDS) == before_host + assert _state(governor, _ENGINE_FIELDS) == before_engine + # Preserve intentional chain progress and exhaustion cooldown. Rolling + # those back would retry the same rejected route indefinitely. + assert agent._fallback_index == 1 + assert agent._rate_limited_until > 0 + fallback_client.responses.create.assert_not_called() + + +def test_refused_fallback_advances_to_a_valid_candidate(governor): + with _host_init(governor) as create: + agent = create(max_tokens=4096) + agent._fallback_chain = [ + {"provider": "openai", "model": "refused-fallback"}, + {"provider": "openai", "model": "valid-fallback"}, + ] + clients = [_fallback_client(), _fallback_client()] + with ( + patch("agent.auxiliary_client.resolve_provider_client", side_effect=[(c, None) for c in clients]), + patch("agent.model_metadata.get_model_context_length", side_effect=[4096, 32_000]), + patch("agent.credential_pool.load_pool", return_value=None), + ): + assert agent._try_activate_fallback() is True + assert agent.model == governor.model == "valid-fallback" + assert agent.client is clients[1] + assert governor.context_length == 32_000 + assert governor.max_tokens == 4096 + assert agent._fallback_index == 2 + assert agent._provider_fallback_route == ("valid-fallback", "openai") + assert len(agent._pending_fallback_notice) == 1 + assert "test-model via openrouter" in agent._pending_fallback_notice[0] + assert "using valid-fallback" in agent._pending_fallback_notice[0] + + +@pytest.mark.parametrize( + "reason", [FailoverReason.rate_limit, FailoverReason.billing, FailoverReason.upstream_rate_limit] +) +@pytest.mark.parametrize("has_valid_tail", [False, True]) +def test_one_primary_failure_arms_backoff_once_through_refused_routes( + governor, reason, has_valid_tail +): + with _host_init(governor) as create: + agent = create(max_tokens=4096) + agent._fallback_chain = [{"provider": "openai", "model": "refused-fallback"}] + if has_valid_tail: + agent._fallback_chain.append({"provider": "openai", "model": "valid-fallback"}) + clients = [_fallback_client() for _ in agent._fallback_chain] + windows = [4096, 32_000] if has_valid_tail else [4096] + before_host = _state(agent, _HOST_FIELDS) + before_engine = _state(governor, _ENGINE_FIELDS) + with ( + patch("agent.chat_completion_helpers.time.monotonic", return_value=1000.0), + patch("agent.auxiliary_client.resolve_provider_client", side_effect=[(c, None) for c in clients]), + patch("agent.model_metadata.get_model_context_length", side_effect=windows), + patch("agent.credential_pool.load_pool", return_value=None), + ): + assert agent._try_activate_fallback(reason) is has_valid_tail + assert agent._rate_limit_backoff_count == 1 + assert agent._rate_limited_until == 1060.0 + if has_valid_tail: + assert agent.model == governor.model == "valid-fallback" + assert agent._fallback_index == 2 + else: + assert _state(agent, _HOST_FIELDS) == before_host + assert _state(governor, _ENGINE_FIELDS) == before_engine + assert agent._fallback_index == 1 + # A distinct originating primary failure still advances the + # backoff even with an already exhausted fallback chain. + assert agent._try_activate_fallback(reason) is False + assert agent._rate_limit_backoff_count == 2 + assert agent._rate_limited_until == 1120.0 + + +def test_refused_primary_restore_preserves_fallback_then_allows_retry(governor): + with _host_init(governor) as create: + agent = create(max_tokens=4096) + agent._fallback_chain = [{"provider": "openai", "model": "valid-fallback"}] + agent._fallback_model = agent._fallback_chain[0] + fallback_client = _fallback_client() + with ( + patch("agent.auxiliary_client.resolve_provider_client", return_value=(fallback_client, None)), + patch("agent.model_metadata.get_model_context_length", return_value=32_000), + patch("agent.credential_pool.load_pool", return_value=None), + ): + assert agent._try_activate_fallback() is True + _sentinel_runtime(agent) + before_host = _state(agent, _HOST_FIELDS) + before_engine = _state(governor, _ENGINE_FIELDS) + original_window = agent._primary_runtime["compressor_context_length"] + agent._primary_runtime["compressor_context_length"] = 4096 + # The saved intended destination may be invalid after reserve changes; + # the active fallback must remain internally coherent on refusal. + before_host["_primary_runtime"]["compressor_context_length"] = 4096 + assert agent._restore_primary_runtime() is False + assert _state(agent, _HOST_FIELDS) == before_host + assert _state(governor, _ENGINE_FIELDS) == before_engine + assert agent.client is fallback_client + agent._primary_runtime["compressor_context_length"] = original_window + assert agent._restore_primary_runtime() is True + assert agent.model == governor.model == "test-model" + assert agent._fallback_activated is False + assert agent._provider_fallback_active is False + assert governor.context_length == original_window + assert governor.max_tokens == 4096 diff --git a/tests/run_agent/test_ares_context_initialization.py b/tests/run_agent/test_ares_context_initialization.py new file mode 100644 index 0000000000000..71ef990d02861 --- /dev/null +++ b/tests/run_agent/test_ares_context_initialization.py @@ -0,0 +1,229 @@ +"""Ares engine selection and resumed lineage through the real host initializer. + +All provider clients, activation probes and governor commands are fakes. The +session ancestry is read from an isolated real SessionDB, not an adapter stub. +""" + +import importlib +from contextlib import ExitStack, contextmanager +from unittest.mock import MagicMock, patch + +import pytest + +from hermes_state import SessionDB +from plugins.context_engine import ContextEngineActivationError +from agent.context_compressor import ContextCompressor + + +@pytest.fixture +def governor(tmp_path): + engine_type = importlib.import_module( + "plugins.context_engine.ri-context-governor" + ).RiContextGovernorEngine + with patch("hermes_cli.config.load_config", return_value={}): + engine = engine_type( + binary=str(tmp_path / "fake-governor"), + store_dir=str(tmp_path / "governor-store"), + ) + engine.probe_activation = MagicMock(return_value={"verified": True}) + engine._run_certified_json = MagicMock(return_value=[]) + engine._load_prior_session_context = MagicMock() + return engine + + +@contextmanager +def _host_init(engine, *, engine_name="ri-context-governor"): + cfg = { + "context": {"engine": engine_name}, + "compression": {"threshold": 0.75}, + "agent": {"environment_probe": False}, + } + with ExitStack() as stack: + for target, value in ( + ("hermes_cli.config.load_config", cfg), + ("hermes_cli.config.load_config_readonly", cfg), + ("plugins.context_engine._load_engine_from_dir", engine), + ("agent.model_metadata.get_model_context_length", 64_000), + ("agent.context_compressor.get_model_context_length", 64_000), + ("run_agent.get_tool_definitions", []), + ("run_agent.check_toolset_requirements", {}), + ): + stack.enter_context(patch(target, return_value=value)) + stack.enter_context(patch("run_agent.OpenAI")) + if engine_name != "ri-context-governor": + stack.enter_context( + patch("plugins.context_engine.load_context_engine", return_value=engine) + ) + from run_agent import AIAgent + + def create(**kwargs): + options = dict( + model="test-model", + provider="openrouter", + api_key="fake-provider-key", + base_url="https://provider.invalid/v1", + quiet_mode=True, + skip_context_files=True, + skip_memory=True, + skip_background_review=True, + ) + options.update(kwargs) + return AIAgent(**options) + + yield create + + +def test_explicit_ares_selection_refuses_missing_engine(): + with ( + _host_init(None) as create, + patch("agent.agent_init.ContextCompressor", wraps=ContextCompressor) as stock, + patch("hermes_cli.plugins.get_plugin_context_engine") as optional, + ): + with pytest.raises(ContextEngineActivationError, match="instantiated"): + create() + stock.assert_not_called() + optional.assert_not_called() + + +@pytest.mark.parametrize( + ("stage", "method", "message"), + [ + ("probe", "probe_activation", "activation probe"), + ("binding", "bind_session_state", "session binding"), + ("startup", "on_session_start", "session start"), + ], +) +def test_explicit_ares_selection_surfaces_startup_failure( + governor, stage, method, message +): + setattr(governor, method, MagicMock(side_effect=RuntimeError(f"{stage} failed"))) + with ( + _host_init(governor) as create, + patch("agent.agent_init.ContextCompressor") as stock, + ): + with pytest.raises(ContextEngineActivationError, match=message) as failure: + create() + assert isinstance(failure.value.__cause__, RuntimeError) + stock.assert_not_called() + + +def test_explicit_ares_selection_probes_before_session_binding(governor): + events = [] + governor.probe_activation.side_effect = lambda: events.append("probe") + original_bind = governor.bind_session_state + + def bind(**kwargs): + events.append("bind") + return original_bind(**kwargs) + + governor.bind_session_state = bind + with _host_init(governor) as create: + agent = create(session_id="new-session", max_tokens=4096) + assert agent.context_compressor is governor + assert events[0:2] == ["probe", "bind"] + governor.probe_activation.assert_called_once() + assert governor.max_tokens == 4096 + assert governor.threshold_tokens == int((64_000 - 4096) * 0.75) + + +@pytest.mark.parametrize( + ("end_reason", "model_config", "expected_root"), + [ + ("compression", {}, "root"), + ("compression", {"_branched_from": "root"}, "tip"), + ("compression", {"_delegate_from": "root"}, "tip"), + ( + "context_rebase", + { + "_context_rebase_from": "root", + "_context_rebase_transition": "transition-fixture", + "_context_epoch": 1, + }, + "root", + ), + ("context_rebase", {}, "tip"), + ], +) +def test_session_start_keeps_the_sessiondb_lineage( + governor, tmp_path, end_reason, model_config, expected_root +): + db = SessionDB(tmp_path / "session-state.db") + try: + db.create_session("root", "cli") + db.end_session("root", end_reason) + db.create_session( + "tip", "cli", parent_session_id="root", model_config=model_config + ) + with _host_init(governor) as create: + create(session_id="tip", session_db=db) + assert governor.session_id == "tip" + assert governor._session_db is db + assert governor._governor_session_id() == expected_root + finally: + db.close() + + +def test_optional_engine_missing_still_uses_stock_compressor(): + with ( + _host_init(None, engine_name="optional-engine") as create, + patch("hermes_cli.plugins.get_plugin_context_engine", return_value=None), + ): + agent = create() + from agent.context_compressor import ContextCompressor + + assert isinstance(agent.context_compressor, ContextCompressor) + + +def test_optional_engine_startup_remains_best_effort(governor): + governor.probe_activation.side_effect = AssertionError("must not strict-probe") + governor.bind_session_state = MagicMock(side_effect=RuntimeError("optional bind")) + governor.on_session_start = MagicMock(side_effect=RuntimeError("optional start")) + with _host_init(governor, engine_name="optional-engine") as create: + agent = create() + assert agent.context_compressor is governor + governor.probe_activation.assert_not_called() + assert "session_db" not in governor.on_session_start.call_args.kwargs + + +def test_manual_model_switch_keeps_the_output_reserve(governor): + governor._policy["token_budget"] = 128_000 + with _host_init(governor) as create: + agent = create(max_tokens=4096) + with patch("agent.model_metadata.get_model_context_length", return_value=32_000): + agent.switch_model( + "switched-model", + "openrouter", + api_key="fake-provider-key", + base_url="https://provider.invalid/v1", + ) + assert governor.model == "switched-model" + assert governor.max_tokens == 4096 + assert governor.context_length == 32_000 + assert governor._target_tokens(100_000) == 27_904 + + +def test_fallback_and_primary_restore_keep_the_output_reserve(governor): + governor._policy["token_budget"] = 128_000 + fallback_client = MagicMock() + fallback_client.base_url = "https://api.openai.com/v1" + fallback_client.api_key = "fake-fallback-key" + with _host_init(governor) as create: + agent = create(max_tokens=4096) + agent._fallback_chain = [{"provider": "openai", "model": "fallback-model"}] + agent._fallback_model = agent._fallback_chain[0] + with ( + patch( + "agent.auxiliary_client.resolve_provider_client", + return_value=(fallback_client, None), + ), + patch("agent.model_metadata.get_model_context_length", return_value=32_000), + ): + assert agent._try_activate_fallback() is True + assert governor.model == "fallback-model" + assert governor.max_tokens == 4096 + assert governor._target_tokens(100_000) == 27_904 + assert agent._restore_primary_runtime() is True + assert governor.model == "test-model" + assert governor.max_tokens == 4096 + assert governor.context_length == 64_000 + assert governor._target_tokens(100_000) == 59_904 From 6f9d47054f6fc218bfe6955a63803a7445e7c996 Mon Sep 17 00:00:00 2001 From: Josh Stevenson Date: Sun, 4 Oct 2026 02:38:38 -0700 Subject: [PATCH 2/2] fix(context): preserve effective MoA mode and absent-compressor restore --- agent/agent_runtime_helpers.py | 26 ++++---- .../agent/test_message_sanitization_policy.py | 2 +- tests/agent/test_moa_switch_api_mode.py | 61 +++++++++++-------- .../test_ares_context_handoff_refusal.py | 30 +++++++++ 4 files changed, 81 insertions(+), 38 deletions(-) diff --git a/agent/agent_runtime_helpers.py b/agent/agent_runtime_helpers.py index 5b86b86f0c523..16e82fb718afb 100644 --- a/agent/agent_runtime_helpers.py +++ b/agent/agent_runtime_helpers.py @@ -1741,16 +1741,20 @@ def restore_primary_runtime(agent) -> bool: cc = agent.context_compressor from agent.auxiliary_client import _update_compressor_model - _update_compressor_model( - cc, - model=rt["compressor_model"], - context_length=rt["compressor_context_length"], - base_url=rt["compressor_base_url"], - api_key=rt["compressor_api_key"], - provider=rt["compressor_provider"], - api_mode=rt.get("compressor_api_mode", ""), - max_tokens=getattr(agent, "max_tokens", None), - ) + if cc is not None: + # A host with no engine has no engine state to restore. Every + # actual engine still requires its saved destination and must + # accept it before the host handoff can commit. + _update_compressor_model( + cc, + model=rt["compressor_model"], + context_length=rt["compressor_context_length"], + base_url=rt["compressor_base_url"], + api_key=rt["compressor_api_key"], + provider=rt["compressor_provider"], + api_mode=rt.get("compressor_api_mode", ""), + max_tokens=getattr(agent, "max_tokens", None), + ) handoff_accepted = True # ── Rebind and re-select the primary credential pool ── @@ -2923,7 +2927,7 @@ def switch_model(agent, new_model, new_provider, api_key='', base_url='', api_mo # moa://local placeholder → HTTP 404 → fallback to a reference # model. Pin chat_completions here so the primary call always goes # through MoAClient.chat.completions, matching agent_init.py. - agent.api_mode = "chat_completions" + api_mode = agent.api_mode = "chat_completions" agent.api_key = api_key or "moa-virtual-provider" agent.base_url = "moa://local" agent._client_kwargs = {} diff --git a/tests/agent/test_message_sanitization_policy.py b/tests/agent/test_message_sanitization_policy.py index 683d2d8df0f6e..8bae2f5b56f5c 100644 --- a/tests/agent/test_message_sanitization_policy.py +++ b/tests/agent/test_message_sanitization_policy.py @@ -455,7 +455,7 @@ def test_restore_primary_reverts_flag(self): agent._reasoning_echo_flag = False from agent.agent_runtime_helpers import restore_primary_runtime - restore_primary_runtime(agent) + assert restore_primary_runtime(agent) is True # Flag should be restored from snapshot assert agent._reasoning_echo_flag is True diff --git a/tests/agent/test_moa_switch_api_mode.py b/tests/agent/test_moa_switch_api_mode.py index 4d1b048bdc018..743461c13714f 100644 --- a/tests/agent/test_moa_switch_api_mode.py +++ b/tests/agent/test_moa_switch_api_mode.py @@ -16,23 +16,26 @@ from __future__ import annotations -import types +from unittest.mock import Mock import pytest -def _make_fake_agent(): - """A minimal stand-in carrying only the attributes switch_model touches.""" - agent = types.SimpleNamespace() +def _make_fake_agent(primary_api_mode): + """Use the real host methods with inert client and no context engine.""" + from run_agent import AIAgent + + agent = object.__new__(AIAgent) agent.model = "minimax-m3" agent.provider = "opencode-go" - agent.api_mode = "anthropic_messages" + agent.api_mode = primary_api_mode agent.api_key = "old-key" agent.base_url = "https://old.example/v1" agent.client = object() agent._client_kwargs = {"base_url": "https://old.example/v1"} agent._config_context_length = 123456 agent._transport_cache = {} + agent.context_compressor = None agent.quiet_mode = True # switch_model re-reads reasoning_echo for the incoming model as part of the # core field swap, before the moa branch runs. On a real AIAgent this is a @@ -40,6 +43,10 @@ def _make_fake_agent(): # this test asserts on. agent._reasoning_echo_flag = False agent._read_reasoning_echo_from_config = lambda: False + if primary_api_mode == "anthropic_messages": + agent._anthropic_api_key = "old-anthropic-key" + agent._anthropic_base_url = "https://old.example" + agent._is_anthropic_oauth = False return agent @@ -47,7 +54,8 @@ def _make_fake_agent(): "incoming_api_mode", ["codex_responses", "anthropic_messages", "chat_completions", ""], ) -def test_switch_to_moa_pins_chat_completions(monkeypatch, incoming_api_mode): +@pytest.mark.parametrize("primary_api_mode", ["chat_completions", "anthropic_messages"]) +def test_switch_to_moa_pins_chat_completions(monkeypatch, incoming_api_mode, primary_api_mode): """Switching to provider=moa must force api_mode=chat_completions. No matter what transport the resolver/aggregator implies for the preset, @@ -57,27 +65,23 @@ def test_switch_to_moa_pins_chat_completions(monkeypatch, incoming_api_mode): """ from agent import agent_runtime_helpers as arh - # Neutralize the post-swap machinery that needs a real AIAgent (credential - # pool reload, context-compressor refresh, primary-runtime bookkeeping). - # We only assert the api_mode invariant set in the moa client-build branch. - monkeypatch.setattr(arh, "load_pool", lambda *a, **k: None, raising=False) + # Keep profile/config and credential reads offline while exercising the + # real host helpers and lazy MoA facade through a complete successful swap. + monkeypatch.setattr("agent.credential_pool.load_pool", lambda *a, **k: None) + monkeypatch.setattr("hermes_cli.config.load_config", lambda: {}) + monkeypatch.setattr("hermes_cli.config.load_config_readonly", lambda: {}) - agent = _make_fake_agent() - try: - arh.switch_model( - agent, - new_model="frontier", - new_provider="moa", - api_key="moa-virtual-provider", - base_url="moa://local", - api_mode=incoming_api_mode, - ) - except Exception: - # switch_model does post-swap work (compressor, pool, runtime) that may - # raise against a fake agent. The runtime-field swap — including the - # api_mode pin in the moa branch — happens before any of that, so the - # invariant we care about is already set even if a later step blew up. - pass + agent = _make_fake_agent(primary_api_mode) + cache_policy = Mock(wraps=agent._anthropic_prompt_cache_policy) + monkeypatch.setattr(agent, "_anthropic_prompt_cache_policy", cache_policy) + arh.switch_model( + agent, + new_model="frontier", + new_provider="moa", + api_key="moa-virtual-provider", + base_url="moa://local", + api_mode=incoming_api_mode, + ) assert agent.provider == "moa" assert agent.base_url == "moa://local" @@ -88,3 +92,8 @@ def test_switch_to_moa_pins_chat_completions(monkeypatch, incoming_api_mode): ) # The MoAClient facade should be installed as the client. assert type(agent.client).__name__ == "MoAClient" + assert cache_policy.call_args.kwargs["api_mode"] == agent.api_mode + assert agent._primary_runtime["provider"] == agent.provider + assert agent._primary_runtime["api_mode"] == agent.api_mode + assert "anthropic_api_key" not in agent._primary_runtime + assert "anthropic_base_url" not in agent._primary_runtime diff --git a/tests/run_agent/test_ares_context_handoff_refusal.py b/tests/run_agent/test_ares_context_handoff_refusal.py index 2920d8cf66748..0dfccf810822e 100644 --- a/tests/run_agent/test_ares_context_handoff_refusal.py +++ b/tests/run_agent/test_ares_context_handoff_refusal.py @@ -242,3 +242,33 @@ def test_refused_primary_restore_preserves_fallback_then_allows_retry(governor): assert agent._provider_fallback_active is False assert governor.context_length == original_window assert governor.max_tokens == 4096 + + +@pytest.mark.parametrize("missing_field", ["compressor_model", "compressor_context_length"]) +@pytest.mark.parametrize("falsey_engine", [False, True]) +def test_active_engine_requires_restore_snapshot_fields( + governor, monkeypatch, missing_field, falsey_engine +): + with _host_init(governor) as create: + agent = create(max_tokens=4096) + agent._fallback_chain = [{"provider": "openai", "model": "valid-fallback"}] + with ( + patch("agent.auxiliary_client.resolve_provider_client", return_value=(_fallback_client(), None)), + patch("agent.model_metadata.get_model_context_length", return_value=32_000), + patch("agent.credential_pool.load_pool", return_value=None), + ): + assert agent._try_activate_fallback() is True + if falsey_engine: + monkeypatch.setattr(type(governor), "__bool__", lambda self: False, raising=False) + saved_value = agent._primary_runtime.pop(missing_field) + before_host = _state(agent, _HOST_FIELDS) + before_engine = _state(governor, _ENGINE_FIELDS) + previous_client = agent.client + assert agent._restore_primary_runtime() is False + assert _state(agent, _HOST_FIELDS) == before_host + assert _state(governor, _ENGINE_FIELDS) == before_engine + assert agent.client is previous_client + agent._primary_runtime[missing_field] = saved_value + assert agent._restore_primary_runtime() is True + assert agent.model == governor.model == "test-model" + assert agent._fallback_activated is False