diff --git a/agent/agent_runtime_helpers.py b/agent/agent_runtime_helpers.py index 20699283e5..7b5b9b081a 100644 --- a/agent/agent_runtime_helpers.py +++ b/agent/agent_runtime_helpers.py @@ -2854,20 +2854,42 @@ def _apply_switched_provider_request_overrides(agent, new_provider): A ``custom_providers`` entry can carry an ``extra_body`` (e.g. ``chat_template_kwargs`` to toggle a local model's thinking). The gateway rebuild path carries this via ``request_overrides``; an *in-place* swap - (CLI / TUI ``/model``) must re-derive it for the new provider, otherwise the - previous provider's ``extra_body`` lingers. Non-provider overrides - (``service_tier`` / ``speed`` from ``/fast``) are preserved. + (CLI / TUI ``/model``) must re-derive it for the switched-to provider, + otherwise the previous provider's ``extra_body`` lingers. + + The switched-to entry is matched by **provider key, base_url, and model** — + the same condition ``agent_init._merge_custom_provider_extra_body`` applies + at build time — via the shared ``_custom_provider_extra_body_for_agent`` + matcher. Matching by name alone would let a *different* model selected at the + same named endpoint inherit an ``extra_body`` configured for another model. + A stale ``extra_body`` is always cleared when the switched-to provider/model + resolves none; non-provider overrides (``service_tier`` / ``speed`` from + ``/fast``) are preserved. """ - from hermes_cli.runtime_provider import ( - _get_named_custom_provider, - _custom_provider_request_overrides, + from agent.agent_init import _custom_provider_extra_body_for_agent + + # Prefer the init-time cache (agent_init stores ``agent._custom_providers`` + # right where it runs its own _merge_custom_provider_extra_body); fall back + # to a fresh load only if a caller built the agent without it. + custom_providers = getattr(agent, "_custom_providers", None) + if custom_providers is None: + try: + from hermes_cli.config import load_config, get_compatible_custom_providers + custom_providers = get_compatible_custom_providers(load_config()) + except Exception: + custom_providers = [] + + new_extra_body = _custom_provider_extra_body_for_agent( + provider=new_provider, + model=getattr(agent, "model", "") or "", + base_url=getattr(agent, "base_url", "") or "", + custom_providers=custom_providers or [], ) - cp = _get_named_custom_provider(new_provider) - new_ro = _custom_provider_request_overrides(cp) if cp else None + overrides = dict(getattr(agent, "request_overrides", {}) or {}) - overrides.pop("extra_body", None) - if new_ro and new_ro.get("extra_body"): - overrides["extra_body"] = new_ro["extra_body"] + overrides.pop("extra_body", None) # always drop the previous provider's extra_body + if new_extra_body: + overrides["extra_body"] = dict(new_extra_body) agent.request_overrides = overrides diff --git a/tests/agent/test_switch_model_request_overrides.py b/tests/agent/test_switch_model_request_overrides.py index 30543e896d..512b71f9ed 100644 --- a/tests/agent/test_switch_model_request_overrides.py +++ b/tests/agent/test_switch_model_request_overrides.py @@ -5,6 +5,11 @@ Before the fix, agent_runtime_helpers.switch_model() swapped model/provider/ base_url/api_key in place but never touched request_overrides, so a /model switch to a thinking-enabled custom provider in the TUI/CLI kept the old provider's extra_body. + +The switched-to entry is matched by provider key + base_url + model (the same +condition agent_init._merge_custom_provider_extra_body uses at build time), so a +*different* model selected at the same named endpoint does not inherit an +extra_body configured for another model. """ import agent.agent_runtime_helpers as arh @@ -14,39 +19,99 @@ class _Agent: pass -def test_switch_applies_new_provider_extra_body(monkeypatch): +# Two entries share the same named endpoint / base_url but pin different models — +# the exact case a name-only match got wrong. +CUSTOM_PROVIDERS = [ + { + "name": "main-think", + "base_url": "http://10.0.0.1:8000/v1", + "model": "think-model", + "extra_body": {"chat_template_kwargs": {"enable_thinking": True}}, + }, + { + "name": "main-plain", + "base_url": "http://10.0.0.1:8000/v1", + "model": "plain-model", + "extra_body": {"chat_template_kwargs": {"enable_thinking": False}}, + }, +] + + +def _agent(*, model, base_url, request_overrides, custom_providers=CUSTOM_PROVIDERS): a = _Agent() - a.request_overrides = {"service_tier": "priority"} # pre-existing /fast override - monkeypatch.setattr( - "hermes_cli.runtime_provider._get_named_custom_provider", - lambda name: {"name": "main-think", - "extra_body": {"chat_template_kwargs": {"enable_thinking": True}}}, + # switch_model() sets these on the live agent before calling the helper. + a.model = model + a.base_url = base_url + a.provider = "custom" + a.request_overrides = request_overrides + a._custom_providers = custom_providers # init-time cache the helper reads + return a + + +def test_switch_applies_matched_provider_extra_body(): + """Switching to the matching provider+model applies its extra_body and + preserves non-provider overrides (service_tier/speed from /fast).""" + a = _agent( + model="think-model", + base_url="http://10.0.0.1:8000/v1", + request_overrides={"service_tier": "priority"}, ) arh._apply_switched_provider_request_overrides(a, "custom:main-think") assert a.request_overrides["extra_body"] == {"chat_template_kwargs": {"enable_thinking": True}} assert a.request_overrides["service_tier"] == "priority" # preserved -def test_switch_to_noncustom_clears_stale_extra_body(monkeypatch): - a = _Agent() - a.request_overrides = { - "extra_body": {"chat_template_kwargs": {"enable_thinking": True}}, - "service_tier": "priority", - } - monkeypatch.setattr( - "hermes_cli.runtime_provider._get_named_custom_provider", lambda name: None +def test_switch_to_noncustom_clears_stale_extra_body(): + """Switching to a built-in provider clears the previous provider's extra_body.""" + a = _agent( + model="claude-x", + base_url="https://api.anthropic.com", + request_overrides={ + "extra_body": {"chat_template_kwargs": {"enable_thinking": True}}, + "service_tier": "priority", + }, ) arh._apply_switched_provider_request_overrides(a, "anthropic") assert "extra_body" not in a.request_overrides # stale extra_body cleared assert a.request_overrides["service_tier"] == "priority" # preserved -def test_switch_from_none_overrides(monkeypatch): - a = _Agent() - a.request_overrides = None - monkeypatch.setattr( - "hermes_cli.runtime_provider._get_named_custom_provider", - lambda name: {"name": "main", "extra_body": {"chat_template_kwargs": {"enable_thinking": False}}}, +def test_switch_from_none_overrides(): + """A None request_overrides is handled and gets the matched extra_body.""" + a = _agent( + model="plain-model", + base_url="http://10.0.0.1:8000/v1", + request_overrides=None, ) - arh._apply_switched_provider_request_overrides(a, "custom:main") + arh._apply_switched_provider_request_overrides(a, "custom:main-plain") assert a.request_overrides == {"extra_body": {"chat_template_kwargs": {"enable_thinking": False}}} + + +def test_switch_to_different_model_same_endpoint_does_not_inherit(): + """Review regression: selecting a *different* model while naming a custom + provider must NOT inherit that provider's extra_body when the models differ. + + 'main-think' pins 'think-model'. Selecting 'plain-model' under + custom:main-think must not carry enable_thinking=True — the model-aware + matcher rejects the mismatch and the stale extra_body is cleared. (A + name-only match would have wrongly carried it over.) + """ + a = _agent( + model="plain-model", # differs from main-think's pinned 'think-model' + base_url="http://10.0.0.1:8000/v1", + request_overrides={"extra_body": {"chat_template_kwargs": {"enable_thinking": True}}}, + ) + arh._apply_switched_provider_request_overrides(a, "custom:main-think") + assert "extra_body" not in a.request_overrides # not inherited; stale cleared + + +def test_switch_endpoint_mismatch_does_not_inherit(): + """A matching provider *name* but a different base_url must not match either + (endpoint identity is part of the condition).""" + a = _agent( + model="think-model", + base_url="http://10.9.9.9:8000/v1", # different endpoint than the entry + request_overrides={"extra_body": {"chat_template_kwargs": {"enable_thinking": True}}}, + ) + arh._apply_switched_provider_request_overrides(a, "custom:main-think") + assert "extra_body" not in a.request_overrides # base_url mismatch -> cleared