fix(auxiliary): replan cache sections on the async fallback path too

_call_fallback_candidate_sync replans messages/tools for each resolved
destination, but its async mirror still shipped the caller's decorated
sections verbatim — the primary destination's markers (including a
direct-native tool marker) leaked to fallback candidates with different
cache contracts, and the relay saw the display label instead of the
resolved provider/api_mode. Mirror the sync path: resolve the
destination, replan both sections, thread provider/api_mode through
_relay_async_completion, and replan again for the auth-refresh retry
client. Mutation-checked: the new parity test fails on the verbatim
pass-through shape.

Follow-up to #76032 (#20880).
This commit is contained in:
kshitij
2026-08-01 15:20:30 +05:30
parent e7340ea281
commit 536754919d
2 changed files with 102 additions and 11 deletions
+39 -11
View File
@@ -4468,43 +4468,71 @@ async def _call_fallback_candidate_async(
task or "call", fb_label, fb_timeout, effective_timeout,
)
effective_timeout = fb_timeout
fb_base = str(getattr(fb_client, "base_url", "") or "")
destination = _fallback_destination(task, fb_client, fb_model, fb_label)
fallback_messages, fallback_tools = _replan_synchronous_cache_sections(
messages,
tools,
destination=destination,
)
fb_kwargs = _build_call_kwargs(
fb_label, fb_model, messages,
destination.provider, destination.model, fallback_messages,
temperature=temperature, max_tokens=max_tokens,
tools=tools, timeout=effective_timeout,
tools=fallback_tools, timeout=effective_timeout,
extra_body=effective_extra_body, reasoning_config=reasoning_config,
base_url=fb_base, task=task)
base_url=destination.base_url, task=task)
try:
return _validate_llm_response(
await _relay_async_completion(
fb_client,
fb_kwargs,
provider=fb_label,
provider=destination.provider,
api_mode=destination.api_mode,
),
task,
)
except Exception as fb_err:
if not _is_auth_error(fb_err):
raise
fb_provider = _auth_refresh_provider_for_route(fb_label, fb_base)
fb_provider = _auth_refresh_provider_for_route(
destination.provider, destination.base_url
)
if fb_provider not in {"auto", "", None} and _refresh_provider_credentials(fb_provider):
retry_client, retry_model = _get_cached_client(
fb_provider, fb_model, async_mode=True)
fb_provider,
destination.model,
async_mode=True,
base_url=destination.base_url or None,
api_mode=destination.api_mode,
)
if retry_client is not None:
retry_destination = _FallbackDestination(
fb_provider,
destination.base_url
or str(getattr(retry_client, "base_url", "") or ""),
destination.api_mode,
retry_model or destination.model,
)
retry_messages, retry_tools = _replan_synchronous_cache_sections(
messages,
tools,
destination=retry_destination,
)
retry_kwargs = _build_call_kwargs(
fb_provider, retry_model or fb_model, messages,
retry_destination.provider,
retry_destination.model,
retry_messages,
temperature=temperature, max_tokens=max_tokens,
tools=tools, timeout=effective_timeout,
tools=retry_tools, timeout=effective_timeout,
extra_body=effective_extra_body,
reasoning_config=reasoning_config,
base_url=str(getattr(retry_client, "base_url", "") or fb_base), task=task)
base_url=retry_destination.base_url, task=task)
try:
return _validate_llm_response(
await _relay_async_completion(
retry_client,
retry_kwargs,
provider=fb_provider,
provider=retry_destination.provider,
api_mode=retry_destination.api_mode,
),
task,
)
+63
View File
@@ -4213,3 +4213,66 @@ class TestSynchronousFallbackCachePlans:
for part in (message.get("content") if isinstance(message.get("content"), list) else [])
)
class TestAsynchronousFallbackCachePlans:
@pytest.mark.asyncio
async def test_async_fallback_replans_cache_sections_like_sync(self, monkeypatch):
"""Async mirror parity: per-destination cache replan, not verbatim pass-through."""
from agent.auxiliary_client import (
_call_fallback_candidate_async,
_try_configured_fallback_chain,
)
entry = {
"provider": "anthropic",
"model": "claude-sonnet-4-6",
"base_url": "https://api.anthropic.com",
"api_mode": "anthropic_messages",
}
client = MagicMock()
client.base_url = entry["base_url"]
async def _create(**kwargs):
return _DummyResponse()
client.chat.completions.create = MagicMock(side_effect=_create)
monkeypatch.setattr(
"agent.auxiliary_client.resolve_provider_client",
lambda provider, model=None, **kwargs: (client, model),
)
monkeypatch.setattr(
"agent.auxiliary_client._get_auxiliary_task_config",
lambda task: {"fallback_chain": [entry]},
)
fallback_client, fallback_model, label = _try_configured_fallback_chain(
task="moa_aggregator",
failed_provider="primary",
)
tools = [{
"type": "function",
"function": {
"name": "lookup",
"parameters": {"type": "object", "properties": {}},
},
}]
await _call_fallback_candidate_async(
fallback_client,
fallback_model,
label,
task="moa_aggregator",
messages=[
{"role": "system", "content": "stable prefix"},
{"role": "user", "content": "lookup"},
],
temperature=None,
max_tokens=None,
tools=tools,
effective_timeout=30.0,
effective_extra_body={},
reasoning_config=None,
)
wire_tools = client.chat.completions.create.call_args.kwargs["tools"]
assert "cache_control" in wire_tools[-1]
assert "cache_control" not in tools[-1]