fix(auxiliary): scope runtime state to each turn
This commit is contained in:
+67
-11
@@ -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()
|
||||
|
||||
@@ -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
@@ -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
@@ -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:
|
||||
"""
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user