fix(auxiliary): scope runtime state to each turn

This commit is contained in:
Teknium
2026-07-17 08:37:59 -07:00
parent 89130bf1f7
commit 73057ed161
6 changed files with 264 additions and 37 deletions
+67 -11
View File
@@ -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()
+3 -1
View File
@@ -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
+24 -4
View File
@@ -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(
+22 -15
View File
@@ -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:
"""
+143 -6
View File
@@ -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()
@@ -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)