From 73057ed1616c7c9973a0bc9ed4d913b93c969c4f Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Fri, 17 Jul 2026 08:37:59 -0700 Subject: [PATCH] fix(auxiliary): scope runtime state to each turn --- agent/auxiliary_client.py | 78 +++++++-- agent/image_routing.py | 4 +- gateway/run.py | 28 +++- run_agent.py | 37 +++-- .../agent/test_auxiliary_runtime_cache_key.py | 149 +++++++++++++++++- .../test_image_input_routing_runtime.py | 5 + 6 files changed, 264 insertions(+), 37 deletions(-) diff --git a/agent/auxiliary_client.py b/agent/auxiliary_client.py index fd86ee2d77..91523d76f8 100644 --- a/agent/auxiliary_client.py +++ b/agent/auxiliary_client.py @@ -42,6 +42,7 @@ Payment / credit exhaustion fallback: import contextlib import contextvars +import hashlib import inspect import json import logging @@ -2381,7 +2382,7 @@ def set_runtime_main( api_key: Any = "", api_mode: str = "", auth_mode: str = "", -) -> None: +) -> contextvars.Token: """Record the current context's live main runtime for auxiliary routing. Context-local state prevents concurrent gateway sessions from overwriting @@ -2404,7 +2405,7 @@ def set_runtime_main( } # Publish authoritative context before updating locked compatibility # mirrors; concurrent sessions never read those mirrors at runtime. - _RUNTIME_MAIN_CONTEXT.set(runtime) + token = _RUNTIME_MAIN_CONTEXT.set(runtime) with _RUNTIME_MAIN_COMPAT_LOCK: ( _RUNTIME_MAIN_PROVIDER, @@ -2417,6 +2418,30 @@ def set_runtime_main( _RUNTIME_MAIN_COMPAT_SNAPSHOT = tuple( runtime[field] for field in _MAIN_RUNTIME_FIELDS ) + return token + + +def reset_runtime_main(token: contextvars.Token) -> None: + """Restore the runtime binding that preceded one scoped turn.""" + if token is None: + return + try: + _RUNTIME_MAIN_CONTEXT.reset(token) + except (RuntimeError, ValueError): + # A token cannot be reset from another copied Context. Background + # workers inherit values, not ownership of the parent's token. + pass + + +@contextlib.contextmanager +def scoped_runtime_main(main_runtime: Optional[Dict[str, Any]]): + """Temporarily bind an explicit runtime without touching legacy mirrors.""" + runtime = _normalize_main_runtime(main_runtime) + token = _RUNTIME_MAIN_CONTEXT.set(runtime or None) + try: + yield runtime + finally: + _RUNTIME_MAIN_CONTEXT.reset(token) def clear_runtime_main() -> None: @@ -5542,6 +5567,7 @@ def resolve_vision_provider_client( base_url: Optional[str] = None, api_key: Optional[str] = None, async_mode: bool = False, + main_runtime: Optional[Dict[str, Any]] = None, ) -> Tuple[Optional[str], Optional[Any], Optional[str]]: """Resolve the client actually used for vision tasks. @@ -5550,6 +5576,7 @@ def resolve_vision_provider_client( backends, so users can intentionally force experimental providers. Auto mode stays conservative and only tries vision backends known to work today. """ + runtime = _normalize_main_runtime(main_runtime) requested, resolved_model, resolved_base_url, resolved_api_key, resolved_api_mode = _resolve_task_provider_model( "vision", provider, model, base_url, api_key ) @@ -5575,6 +5602,7 @@ def resolve_vision_provider_client( explicit_base_url=resolved_base_url, explicit_api_key=resolved_api_key, api_mode=resolved_api_mode, + main_runtime=runtime, ) if client is None: return provider_for_base_override, None, None @@ -5598,8 +5626,8 @@ def resolve_vision_provider_client( # live from the catalog — tried when # DEEPINFRA_API_KEY is set) # 5. Stop - main_provider = _read_main_provider() - main_model = _read_main_model() + main_provider = str(runtime.get("provider") or _read_main_provider()) + main_model = str(runtime.get("model") or _read_main_model()) if main_provider and main_provider not in {"auto", ""}: # A provider-specific vision default wins over the user's chat model: # static overrides (xiaomi/zai) and catalog-backed discovery (the @@ -5662,13 +5690,13 @@ def resolve_vision_provider_client( rpc_api_key = None rpc_api_mode = resolved_api_mode if main_provider == "custom" or main_provider.startswith("custom:"): - runtime_base_url = _runtime_main_value("base_url") + runtime_base_url = runtime.get("base_url") if runtime_base_url: rpc_base_url = runtime_base_url - rpc_api_key = _runtime_main_value("api_key") or None + rpc_api_key = runtime.get("api_key") or None rpc_api_mode = ( resolved_api_mode - or _runtime_main_value("api_mode") + or runtime.get("api_mode") or None ) else: @@ -5684,6 +5712,7 @@ def resolve_vision_provider_client( api_mode=rpc_api_mode, explicit_base_url=rpc_base_url, explicit_api_key=rpc_api_key, + main_runtime=runtime, is_vision=True) if rpc_client is not None: logger.info( @@ -5726,6 +5755,7 @@ def resolve_vision_provider_client( base_url=_zai_url, api_key=resolved_api_key or None, api_mode="chat_completions", + main_runtime=runtime, is_vision=True, ) if client is not None: @@ -5733,6 +5763,7 @@ def resolve_vision_provider_client( # Fallback: try without explicit base_url (old behavior) client, final_model = _get_cached_client(requested, resolved_model, async_mode, api_mode=resolved_api_mode, + main_runtime=runtime, is_vision=True) if client is None: return requested, None, None @@ -5740,6 +5771,7 @@ def resolve_vision_provider_client( client, final_model = _get_cached_client(requested, resolved_model, async_mode, api_mode=resolved_api_mode, + main_runtime=runtime, is_vision=True) if client is None: return requested, None, None @@ -5833,6 +5865,9 @@ def _runtime_cache_discriminator(field: str, value: Any) -> Any: """Return a hashable, secret-safe runtime cache-key component.""" if field == "api_key" and callable(value): return _CallableCacheDiscriminator(value) + if field == "api_key" and isinstance(value, str) and value: + digest = hashlib.blake2b(value.encode("utf-8"), digest_size=16).digest() + return ("api-key-digest", digest) return value @@ -5867,8 +5902,9 @@ def _client_cache_key( # APIConnectionError that fails the sibling advisor (root cause of the run2 # double-advisor "Connection error" collapse). Keying on model gives each # model its own client, so concurrent fan-out calls never cross-close. - model_key = model or "" - return (provider, async_mode, base_url or "", api_key or "", api_mode or "", runtime_key, is_vision, task_key, pool_hint, model_key) + model_key = model or runtime.get("model", "") + api_key_key = _runtime_cache_discriminator("api_key", api_key or "") + return (provider, async_mode, base_url or "", api_key_key, api_mode or "", runtime_key, is_vision, task_key, pool_hint, model_key) def _store_cached_client(cache_key: tuple, client: Any, default_model: Optional[str], *, bound_loop: Any = None) -> None: @@ -6151,13 +6187,20 @@ def _get_cached_client( if cache_key not in _client_cache: # Safety belt: if the cache has grown beyond the max, evict # the oldest entries (FIFO — dict preserves insertion order). + # Do not close an evicted client here: another caller may be + # mid-request with the object it obtained from this cache. + # Dropping the cache reference lets normal refcount/GC cleanup + # happen after in-flight users release it. while len(_client_cache) >= _CLIENT_CACHE_MAX_SIZE: - evict_key, evict_entry = next(iter(_client_cache.items())) - _close_cached_client(evict_entry[0]) + evict_key = next(iter(_client_cache)) del _client_cache[evict_key] _client_cache[cache_key] = (client, default_model, bound_loop) else: + built_client = client client, default_model, _ = _client_cache[cache_key] + # This concurrently built loser was never exposed to a caller, + # so it is safe to close immediately. + _close_cached_client(built_client) return client, model or default_model @@ -6910,6 +6953,11 @@ def call_llm( Raises: RuntimeError: If no provider is configured. """ + # Capture one immutable runtime snapshot for keying, resolution, retries, + # and fallbacks. Reading ambient state independently in each phase lets a + # concurrent /model switch produce a key for one runtime and a client for + # another. + main_runtime = _normalize_main_runtime(main_runtime) resolved_provider, resolved_model, resolved_base_url, resolved_api_key, resolved_api_mode = _resolve_task_provider_model( task, provider, model, base_url, api_key) if api_mode: @@ -6924,6 +6972,7 @@ def call_llm( base_url=resolved_base_url or base_url, api_key=resolved_api_key or api_key, async_mode=False, + main_runtime=main_runtime, ) if client is None and resolved_provider != "auto" and not resolved_base_url: logger.warning( @@ -6934,6 +6983,7 @@ def call_llm( provider="auto", model=resolved_model, async_mode=False, + main_runtime=main_runtime, ) if client is None: raise RuntimeError( @@ -7536,6 +7586,9 @@ async def async_call_llm( Same as call_llm() but async. See call_llm() for full documentation. """ + # Keep every async phase on the same runtime identity, even if another + # session switches models while this task is awaiting network I/O. + main_runtime = _normalize_main_runtime(main_runtime) resolved_provider, resolved_model, resolved_base_url, resolved_api_key, resolved_api_mode = _resolve_task_provider_model( task, provider, model, base_url, api_key) effective_extra_body = _get_task_extra_body(task) @@ -7548,6 +7601,7 @@ async def async_call_llm( base_url=resolved_base_url or base_url, api_key=resolved_api_key or api_key, async_mode=True, + main_runtime=main_runtime, ) if client is None and resolved_provider != "auto" and not resolved_base_url: logger.warning( @@ -7558,6 +7612,7 @@ async def async_call_llm( provider="auto", model=resolved_model, async_mode=True, + main_runtime=main_runtime, ) if client is None: raise RuntimeError( @@ -7573,6 +7628,7 @@ async def async_call_llm( base_url=resolved_base_url, api_key=resolved_api_key, api_mode=resolved_api_mode, + main_runtime=main_runtime, ) if client is None: _explicit = (resolved_provider or "").strip().lower() diff --git a/agent/image_routing.py b/agent/image_routing.py index 9299710c2d..8cfa085527 100644 --- a/agent/image_routing.py +++ b/agent/image_routing.py @@ -270,7 +270,9 @@ def _resolve_inference_base_url( from agent.auxiliary_client import _runtime_main_value runtime = str(_runtime_main_value("base_url") or "").strip() - if runtime: + runtime_provider = str(_runtime_main_value("provider") or "").strip().lower() + requested_provider = str(provider or "").strip().lower() + if runtime and (not requested_provider or requested_provider == runtime_provider): return runtime except Exception: pass diff --git a/gateway/run.py b/gateway/run.py index 275a8ed6de..61fcd96eb7 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -10860,10 +10860,30 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew "Image routing: text (mode=%s). Pre-analyzing %d image(s) via vision_analyze.", _img_mode, len(image_paths), ) - message_text = await self._enrich_message_with_vision( - message_text, - image_paths, - ) + # Vision enrichment runs before AIAgent.run_conversation(), + # so bind this session's resolved runtime explicitly rather + # than consulting process-global compatibility mirrors. + vision_runtime = None + try: + turn_model, runtime_kwargs = self._resolve_session_agent_runtime( + source=source, + session_key=session_key, + ) + vision_runtime = dict(runtime_kwargs or {}) + vision_runtime["model"] = turn_model + except Exception: + logger.debug( + "vision enrichment: session runtime resolution failed", + exc_info=True, + ) + + from agent.auxiliary_client import scoped_runtime_main + + with scoped_runtime_main(vision_runtime): + message_text = await self._enrich_message_with_vision( + message_text, + image_paths, + ) if audio_paths: message_text, _successful_transcripts = await self._enrich_message_with_transcription( diff --git a/run_agent.py b/run_agent.py index de4a61c947..e48e83f8f7 100644 --- a/run_agent.py +++ b/run_agent.py @@ -6202,21 +6202,28 @@ class AIAgent: acct_token = set_accounting_context( getattr(self, "_session_db", None), getattr(self, "session_id", None) ) - try: - return run_conversation( - self, - user_message, - system_message, - conversation_history, - task_id, - stream_callback, - persist_user_message, - persist_user_timestamp=persist_user_timestamp, - moa_config=moa_config, - ) - finally: - reset_accounting_context(acct_token) - reset_conversation_context(token) + from agent.auxiliary_client import scoped_runtime_main + + # The outer token restores the caller's Context even though turn setup + # replaces the value with the live runtime after fallback restoration. + # Keep the scope local instead of storing ContextVar tokens on the agent, + # which may be observed from another thread. + with scoped_runtime_main({}): + try: + return run_conversation( + self, + user_message, + system_message, + conversation_history, + task_id, + stream_callback, + persist_user_message, + persist_user_timestamp=persist_user_timestamp, + moa_config=moa_config, + ) + finally: + reset_accounting_context(acct_token) + reset_conversation_context(token) def chat(self, message: str, stream_callback: Optional[callable] = None) -> str: """ diff --git a/tests/agent/test_auxiliary_runtime_cache_key.py b/tests/agent/test_auxiliary_runtime_cache_key.py index 18d8ab28d5..b7957f1f82 100644 --- a/tests/agent/test_auxiliary_runtime_cache_key.py +++ b/tests/agent/test_auxiliary_runtime_cache_key.py @@ -5,9 +5,11 @@ #56889, which isolates callers that pass different explicit ``model=`` values. """ +import asyncio from concurrent.futures import ThreadPoolExecutor from threading import Barrier -from unittest.mock import MagicMock, patch +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -114,6 +116,37 @@ def test_context_without_runtime_does_not_fall_back_to_other_session_globals(): assert contextvars.Context().run(fresh_context) == {} +def test_runtime_context_token_restores_previous_value_after_turn(): + """Turn-scoped runtime binding must not leak into later work in the same context.""" + token = aux.set_runtime_main(**_runtime("turn-model")) + assert aux._normalize_main_runtime(None)["model"] == "turn-model" + + aux.reset_runtime_main(token) + + assert aux._normalize_main_runtime(None) == {} + + +def test_aiagent_wrapper_resets_runtime_context_after_turn(): + """Every production run_conversation exit restores the caller's Context.""" + from run_agent import AIAgent + + agent = SimpleNamespace( + _conversation_root_id=lambda: "root-session", + _session_db=None, + session_id="session-id", + ) + + def fake_turn(*_args, **_kwargs): + aux.set_runtime_main(**_runtime("wrapped-turn")) + return {"final_response": "ok"} + + with patch("agent.conversation_loop.run_conversation", side_effect=fake_turn): + result = AIAgent.run_conversation(agent, "hello") + + assert result["final_response"] == "ok" + assert aux._normalize_main_runtime(None) == {} + + def test_legacy_patched_globals_are_visible_only_without_an_active_runtime(): """Direct legacy patches work, but never override context-local session state.""" with patch.object(aux, "_RUNTIME_MAIN_PROVIDER", "custom:legacy"), patch.object( @@ -175,6 +208,88 @@ def test_explicit_model_cache_isolation_remains_independent_of_runtime_key(): assert first != second +def test_pinned_provider_without_model_inherits_live_runtime_model_in_cache_key(): + """A pinned provider with model=auto must follow the switched main model.""" + first = aux._client_cache_key( + "openrouter", + async_mode=False, + main_runtime=_runtime("old-model", provider="openrouter"), + ) + second = aux._client_cache_key( + "openrouter", + async_mode=False, + main_runtime=_runtime("new-model", provider="openrouter"), + ) + + assert first != second + + +def test_explicit_vision_runtime_wins_over_stale_ambient_runtime(): + """Vision resolution must use the immutable runtime supplied by its caller.""" + aux.set_runtime_main(**_runtime("ambient-old")) + explicit = _runtime("explicit-new") + captured = {} + + def fake_resolve(provider, model, **kwargs): + captured.update(provider=provider, model=model, **kwargs) + return MagicMock(), model + + with patch.object( + aux, "_resolve_task_provider_model", return_value=("auto", None, None, None, None) + ), patch.object(aux, "_main_model_supports_vision", return_value=True), patch.object( + aux, "resolve_provider_client", side_effect=fake_resolve + ): + provider, _client, model = aux.resolve_vision_provider_client( + main_runtime=explicit + ) + + assert provider == "custom:llama-swap" + assert model == "explicit-new" + assert captured["explicit_base_url"] == "http://llama-swap.test/v1" + + +def test_image_routing_does_not_borrow_base_url_from_different_provider(): + """An explicit provider must not inherit another runtime's custom endpoint.""" + from agent.image_routing import _resolve_inference_base_url + + aux.set_runtime_main(**_runtime("custom-model")) + cfg = { + "model": { + "provider": "openrouter", + "base_url": "https://openrouter.ai/api/v1", + } + } + + assert ( + _resolve_inference_base_url(cfg, "openrouter") + == "https://openrouter.ai/api/v1" + ) + + +def test_async_initial_cache_lookup_receives_explicit_runtime_snapshot(): + """The first async lookup must not drop main_runtime and only pass it on fallback.""" + runtime = _runtime("async-new") + response = MagicMock() + response.choices = [MagicMock(message=MagicMock(content="ok"))] + client = MagicMock() + client.chat.completions.create = AsyncMock(return_value=response) + + with patch.object( + aux, + "_resolve_task_provider_model", + return_value=("openrouter", None, None, None, None), + ), patch.object(aux, "_get_cached_client", return_value=(client, "async-new")) as get_client: + asyncio.run( + aux.async_call_llm( + task="approval", + main_runtime=runtime, + messages=[{"role": "user", "content": "approve?"}], + ) + ) + + assert get_client.call_args.kwargs["main_runtime"] == aux._normalize_main_runtime(runtime) + + def test_unhashable_callable_runtime_api_keys_are_safe_secret_free_discriminators(): """Callable token providers remain cacheable without leaking returned tokens.""" @@ -204,8 +319,31 @@ def test_unhashable_callable_runtime_api_keys_are_safe_secret_free_discriminator assert "second-super-secret-token" not in rendered -def test_fifo_eviction_closes_oldest_sync_client_once_after_65_entries(): - """The bounded cache closes, rather than merely drops, its FIFO sync victim.""" +def test_string_api_keys_are_not_retained_in_cache_key_repr(): + """String credentials discriminate clients without living in cache-key memory.""" + first_secret = "first-literal-super-secret" + second_secret = "second-literal-super-secret" + first = aux._client_cache_key( + "auto", + async_mode=False, + api_key=first_secret, + main_runtime={**_runtime("same"), "api_key": first_secret}, + ) + second = aux._client_cache_key( + "auto", + async_mode=False, + api_key=second_secret, + main_runtime={**_runtime("same"), "api_key": second_secret}, + ) + + assert first != second + rendered = repr((first, second)) + assert first_secret not in rendered + assert second_secret not in rendered + + +def test_fifo_eviction_does_not_close_client_that_may_have_an_inflight_call(): + """A bounded-cache eviction must not invalidate another caller's client.""" clients = [] def fake_resolve(_provider, model, _async_mode, **_kwargs): @@ -218,11 +356,10 @@ def test_fifo_eviction_closes_oldest_sync_client_once_after_65_entries(): aux._get_cached_client("custom", model=f"model-{index}") assert len(aux._client_cache) == 64 - clients[0].close.assert_called_once_with() - for client in clients[1:]: + for client in clients: client.close.assert_not_called() aux.shutdown_cached_clients() - clients[0].close.assert_called_once_with() + clients[0].close.assert_not_called() for client in clients[1:]: client.close.assert_called_once_with() diff --git a/tests/gateway/test_image_input_routing_runtime.py b/tests/gateway/test_image_input_routing_runtime.py index 5bf34d3901..bc50935807 100644 --- a/tests/gateway/test_image_input_routing_runtime.py +++ b/tests/gateway/test_image_input_routing_runtime.py @@ -123,8 +123,13 @@ async def test_prepare_image_routing_falls_back_to_text_for_text_only_session_ov monkeypatch.setattr("agent.image_routing._lookup_supports_vision", fake_supports) async def fake_enrich(user_text, image_paths): + from agent import auxiliary_client as aux + assert user_text == "look" assert image_paths == ["/tmp/cashback.png"] + runtime = aux._normalize_main_runtime(None) + assert runtime["provider"] == "xiaomi" + assert runtime["model"] == "mimo-v2.5-pro" return "[vision summary]\n\nlook" monkeypatch.setattr(runner, "_enrich_message_with_vision", fake_enrich)