Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
78 changes: 78 additions & 0 deletions run_agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -1095,6 +1095,25 @@ def _qwen_portal_headers() -> dict:
}


def _deep_merge_extra_body(base: Dict[str, Any], overlay: Dict[str, Any]) -> Dict[str, Any]:
"""Merge ``overlay`` on top of ``base`` one level deep.

OpenRouter-style nested keys like ``extra_body["provider"]`` need a nested
merge so a fallback chain entry setting ``provider.order`` /
``allow_fallbacks`` does not clobber an existing
``provider.require_parameters`` baked into ``request_overrides`` by the
primary configuration. Non-dict values follow shallow merge (overlay wins).
"""
merged = dict(base)
for key, overlay_value in overlay.items():
existing = merged.get(key)
if isinstance(existing, dict) and isinstance(overlay_value, dict):
merged[key] = {**existing, **overlay_value}
else:
merged[key] = overlay_value
return merged


class AIAgent:
"""
AI Agent with tool calling capabilities.
Expand All @@ -1108,6 +1127,17 @@ class AIAgent:
"have been dropped to keep the conversation alive. See issue #15236.]"
)

@staticmethod
def _snapshot_request_overrides(value: Any) -> Dict[str, Any]:
"""Shallow-copy ``value`` for use as a ``request_overrides`` snapshot.

Non-dict values normalize to ``{}`` so the snapshot is always safe to
restore back into ``self.request_overrides`` without TypeError. Used
in ``__init__``, ``switch_model``, and ``_try_activate_fallback`` to
keep the three snapshot sites in lockstep.
"""
return dict(value) if isinstance(value, dict) else {}

@property
def base_url(self) -> str:
return self._base_url
Expand Down Expand Up @@ -2482,6 +2512,12 @@ def __init__(
"client_kwargs": dict(self._client_kwargs),
"use_prompt_caching": self._use_prompt_caching,
"use_native_cache_layout": self._use_native_cache_layout,
# Snapshot request_overrides so fallback-entry-specific
# extra_body forwarded by _try_activate_fallback() is cleared
# when the primary route is restored next turn (#26460).
"request_overrides": self._snapshot_request_overrides(
getattr(self, "request_overrides", None)
),
# Context engine state that _try_activate_fallback() overwrites.
# Use getattr for model/base_url/api_key/provider since plugin
# engines may not have these (they're ContextCompressor-specific).
Expand Down Expand Up @@ -2765,6 +2801,9 @@ def switch_model(self, new_model, new_provider, api_key='', base_url='', api_mod
"client_kwargs": dict(self._client_kwargs),
"use_prompt_caching": self._use_prompt_caching,
"use_native_cache_layout": self._use_native_cache_layout,
"request_overrides": self._snapshot_request_overrides(
getattr(self, "request_overrides", None)
),
"compressor_model": getattr(_cc, "model", self.model) if _cc else self.model,
"compressor_base_url": getattr(_cc, "base_url", self.base_url) if _cc else self.base_url,
"compressor_api_key": getattr(_cc, "api_key", "") if _cc else "",
Expand Down Expand Up @@ -8832,6 +8871,35 @@ def _try_activate_fallback(self, reason: "FailoverReason | None" = None) -> bool
self._transport_cache.clear()
self._fallback_activated = True

# Forward fallback-entry-specific request metadata (e.g.
# OpenRouter ``extra_body.provider`` routing) into
# request_overrides so it scopes to the active fallback route
# only. The chat_completions transport already merges
# ``request_overrides["extra_body"]`` into the outbound
# extra_body (see agent/transports/chat_completions.py); this
# makes per-fallback routing config-honored without leaking
# into unrelated requests. _restore_primary_runtime() resets
# request_overrides from the snapshot so the override clears
# when the primary route comes back. See #26460.
fb_extra_body = fb.get("extra_body")
if isinstance(fb_extra_body, dict) and fb_extra_body:
base_overrides = self._snapshot_request_overrides(
getattr(self, "request_overrides", None)
)
existing_eb = base_overrides.get("extra_body")
# Deep-merge one level so a fallback entry's nested keys
# (e.g. ``provider.order``) don't clobber unrelated nested
# keys baked into the primary override (e.g.
# ``provider.require_parameters``). Fallback wins on leaf
# collision, primary wins on absent leaves.
merged_eb = (
_deep_merge_extra_body(existing_eb, fb_extra_body)
if isinstance(existing_eb, dict)
else dict(fb_extra_body)
)
Comment on lines +8889 to +8899
base_overrides["extra_body"] = merged_eb
self.request_overrides = base_overrides

# Honor per-provider / per-model request_timeout_seconds for the
# fallback target (same knob the primary client uses). None = use
# SDK default.
Expand Down Expand Up @@ -8956,6 +9024,16 @@ def _restore_primary_runtime(self) -> bool:
self.api_key = rt["api_key"]
self._client_kwargs = dict(rt["client_kwargs"])
self._use_prompt_caching = rt["use_prompt_caching"]
# Drop any fallback-entry-specific request_overrides that
# _try_activate_fallback() merged in (e.g. OpenRouter
# extra_body.provider routing). Older sessions saved before
# #26460 won't have the snapshot key — leave the current
# request_overrides untouched in that case so any overrides
# set after snapshot capture are preserved.
if "request_overrides" in rt:
self.request_overrides = self._snapshot_request_overrides(
rt["request_overrides"]
)
# Default to native layout when the restored snapshot predates the
# native-vs-proxy split (older sessions saved before this PR).
self._use_native_cache_layout = rt.get(
Expand Down
208 changes: 208 additions & 0 deletions tests/run_agent/test_provider_fallback.py
Original file line number Diff line number Diff line change
Expand Up @@ -305,3 +305,211 @@ def test_returns_false_when_only_self_matching_entries(self):

assert ok is False
mock_resolve.assert_not_called()


# ── Fallback-entry extra_body forwarding (#26460) ─────────────────────────


class TestFallbackEntryExtraBodyForwarded:
"""When a fallback entry carries ``extra_body`` (e.g. OpenRouter
``provider.order`` for request-scoped routing), activating that entry
must merge it into ``request_overrides`` so the chat_completions
transport forwards it on the very next request — without mutating the
agent's global ``provider_routing`` knobs (``providers_allowed`` etc.)
and without persisting past the next ``_restore_primary_runtime``.
See issue #26460."""

def _activate(self, agent):
with (
patch(
"agent.auxiliary_client.resolve_provider_client",
return_value=(_mock_client(), agent._fallback_chain[0]["model"]),
),
patch(
"hermes_cli.model_normalize.normalize_model_for_provider",
side_effect=lambda m, p: m,
),
):
return agent._try_activate_fallback()

def test_fallback_extra_body_forwarded_to_request_overrides(self):
fb_extra = {
"provider": {
"order": ["baidu/fp8", "gmicloud/fp8", "deepinfra/fp4"],
"allow_fallbacks": False,
}
}
fbs = [{
"provider": "openrouter",
"model": "z-ai/glm-5.1",
"key_env": "OPENROUTER_API_KEY",
"extra_body": fb_extra,
}]
agent = _make_agent(fallback_model=fbs)
assert "extra_body" not in (agent.request_overrides or {})

assert self._activate(agent) is True
eb = agent.request_overrides.get("extra_body")
assert isinstance(eb, dict)
assert eb["provider"] == fb_extra["provider"]

def test_global_provider_routing_unchanged(self):
"""Activating a fallback entry's request-scoped extra_body must not
touch the agent-level provider_routing knobs — those still apply to
unrelated OpenRouter requests."""
fbs = [{
"provider": "openrouter",
"model": "z-ai/glm-5.1",
"extra_body": {"provider": {"order": ["baidu/fp8"], "allow_fallbacks": False}},
}]
agent = _make_agent(fallback_model=fbs)
# Snapshot the global routing knobs before activation.
before = (
list(agent.providers_allowed or []),
list(agent.providers_ignored or []),
list(agent.providers_order or []),
agent.provider_sort,
)

assert self._activate(agent) is True

after = (
list(agent.providers_allowed or []),
list(agent.providers_ignored or []),
list(agent.providers_order or []),
agent.provider_sort,
)
assert before == after

def test_fallback_without_extra_body_does_not_inject_key(self):
fbs = [{"provider": "openrouter", "model": "z-ai/glm-5.1"}]
agent = _make_agent(fallback_model=fbs)
agent.request_overrides = {}

assert self._activate(agent) is True
assert "extra_body" not in agent.request_overrides

def test_fallback_extra_body_merges_with_existing_extra_body(self):
"""If primary config already populated request_overrides['extra_body'],
the fallback's extra_body should merge on top (fallback wins on
key collision) rather than replace the whole dict."""
fbs = [{
"provider": "openrouter",
"model": "z-ai/glm-5.1",
"extra_body": {"provider": {"order": ["baidu/fp8"], "allow_fallbacks": False}},
}]
agent = _make_agent(fallback_model=fbs)
agent.request_overrides = {"extra_body": {"some_user_field": 42}}

assert self._activate(agent) is True
eb = agent.request_overrides["extra_body"]
# Pre-existing user extra_body field preserved.
assert eb["some_user_field"] == 42
# Fallback-entry extra_body merged in.
assert eb["provider"]["order"] == ["baidu/fp8"]

def test_invalid_extra_body_type_is_ignored(self):
"""Defensive: a non-dict extra_body in the chain entry must not
crash activation or mutate request_overrides."""
fbs = [{
"provider": "openrouter",
"model": "z-ai/glm-5.1",
"extra_body": "not-a-dict", # invalid shape
}]
Comment on lines +411 to +418
agent = _make_agent(fallback_model=fbs)
agent.request_overrides = {}

assert self._activate(agent) is True
assert "extra_body" not in agent.request_overrides

def test_empty_extra_body_dict_does_not_inject_key(self):
"""An explicit ``"extra_body": {}`` on a chain entry must behave the
same as an absent key — no override injected. Guards against a
regression if the ``and fb_extra_body`` truthiness check is later
removed."""
fbs = [{
"provider": "openrouter",
"model": "z-ai/glm-5.1",
"extra_body": {}, # explicit empty
}]
agent = _make_agent(fallback_model=fbs)
agent.request_overrides = {}

assert self._activate(agent) is True
assert "extra_body" not in agent.request_overrides

def test_fallback_extra_body_deep_merges_nested_provider_dict(self):
"""When primary request_overrides already carry ``extra_body.provider``
(e.g. OpenRouter ``require_parameters: true`` baked in by config),
the fallback's ``provider.order`` / ``allow_fallbacks`` must merge
into that nested dict instead of replacing it wholesale."""
fbs = [{
"provider": "openrouter",
"model": "z-ai/glm-5.1",
"extra_body": {"provider": {"order": ["baidu/fp8"], "allow_fallbacks": False}},
}]
agent = _make_agent(fallback_model=fbs)
agent.request_overrides = {
"extra_body": {"provider": {"require_parameters": True, "data_collection": "deny"}}
}

assert self._activate(agent) is True
eb = agent.request_overrides["extra_body"]
# Pre-existing nested keys preserved.
assert eb["provider"]["require_parameters"] is True
assert eb["provider"]["data_collection"] == "deny"
# Fallback-entry nested keys merged in (fallback wins on collisions).
assert eb["provider"]["order"] == ["baidu/fp8"]
assert eb["provider"]["allow_fallbacks"] is False

def test_restore_primary_runtime_clears_fallback_extra_body(self):
fbs = [{
"provider": "openrouter",
"model": "z-ai/glm-5.1",
"extra_body": {"provider": {"order": ["baidu/fp8"], "allow_fallbacks": False}},
}]
agent = _make_agent(fallback_model=fbs)
# Snapshot baseline overrides at init (typically empty).
baseline_overrides = dict(agent.request_overrides or {})

assert self._activate(agent) is True
assert "extra_body" in agent.request_overrides

# Simulate the start of the next turn — restore from snapshot.
with patch.object(
agent, "_create_openai_client", return_value=MagicMock(),
):
assert agent._restore_primary_runtime() is True

assert agent.request_overrides == baseline_overrides

def test_primary_runtime_snapshot_includes_request_overrides(self):
"""Fix #26460 requires the snapshot to capture request_overrides at
init so restoration after fallback activation doesn't leak the
fallback-entry override into the primary route."""
agent = _make_agent(fallback_model=None)
assert "request_overrides" in agent._primary_runtime
assert isinstance(agent._primary_runtime["request_overrides"], dict)

def test_restore_from_older_snapshot_preserves_current_overrides(self):
"""Older sessions persisted ``_primary_runtime`` without the
``request_overrides`` field. On restore, code paths that need to
operate on those sessions must NOT overwrite the current
``self.request_overrides`` with ``{}`` — that would silently drop
user-set overrides applied after the snapshot was taken."""
fbs = [{"provider": "openrouter", "model": "z-ai/glm-5.1"}]
agent = _make_agent(fallback_model=fbs)
assert self._activate(agent) is True
# Simulate an older session whose snapshot predates #26460 by
# dropping the request_overrides key after activation.
agent._primary_runtime.pop("request_overrides", None)
# User set an override AFTER the (older) snapshot was taken.
agent.request_overrides = {"extra_body": {"user_field": "preserved"}}

with patch.object(
agent, "_create_openai_client", return_value=MagicMock(),
):
assert agent._restore_primary_runtime() is True

# Override survives the restore — older snapshot didn't carry the key.
assert agent.request_overrides == {"extra_body": {"user_field": "preserved"}}
Loading