diff --git a/gateway/run.py b/gateway/run.py index ffd8264c12..df62736448 100644 --- a/gateway/run.py +++ b/gateway/run.py @@ -8348,6 +8348,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew "api_mode": override.get("api_mode"), "max_tokens": override.get("max_tokens"), "credential_pool": override.get("credential_pool"), + "request_overrides": override.get("request_overrides"), } if override_runtime.get("api_key"): if override_runtime.get("credential_pool") is None: @@ -27501,12 +27502,16 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew val = override.get(key) if val is not None: runtime_kwargs[key] = val - override_request_overrides = override.get("request_overrides") - if isinstance(override_request_overrides, dict): - runtime_kwargs["request_overrides"] = _deep_merge_request_overrides( - runtime_kwargs.get("request_overrides"), - override_request_overrides, - ) + # request_overrides reflects the switched-to provider; apply whenever + # the override recorded it (even as None) so switching to a provider + # without configured overrides clears a stale value left by the + # default provider's runtime resolution. + if "request_overrides" in override: + override_request_overrides = override.get("request_overrides") + if isinstance(override_request_overrides, dict) and override_request_overrides: + runtime_kwargs["request_overrides"] = dict(override_request_overrides) + else: + runtime_kwargs["request_overrides"] = override_request_overrides if ( runtime_kwargs.get("api_key") and runtime_kwargs.get("credential_pool") is None diff --git a/hermes_cli/model_switch.py b/hermes_cli/model_switch.py index f26e3f19ec..c50df9e104 100644 --- a/hermes_cli/model_switch.py +++ b/hermes_cli/model_switch.py @@ -2185,6 +2185,22 @@ def switch_model( if hermes_warn: warnings.append(hermes_warn) + # Carry the switched provider's request_overrides (e.g. a custom_providers + # ``extra_body`` such as chat_template_kwargs) so a ``/model`` switch to a + # custom provider applies it on the gateway, matching the default-provider + # path. resolve_runtime_provider surfaces these for named custom providers. + request_overrides = None + try: + from hermes_cli.runtime_provider import ( + _get_named_custom_provider, + _custom_provider_request_overrides, + ) + _cp_for_ro = _get_named_custom_provider(target_provider) + if _cp_for_ro: + request_overrides = _custom_provider_request_overrides(_cp_for_ro) or None + except Exception: + request_overrides = None + # --- Build result --- return ModelSwitchResult( success=True, diff --git a/tests/gateway/test_turn_request_overrides.py b/tests/gateway/test_turn_request_overrides.py index e65568d644..c985125176 100644 --- a/tests/gateway/test_turn_request_overrides.py +++ b/tests/gateway/test_turn_request_overrides.py @@ -89,3 +89,52 @@ def test_resolve_runtime_agent_kwargs_carries_request_overrides(monkeypatch): ) rk = gateway_run._resolve_runtime_agent_kwargs() assert rk["request_overrides"] == PROVIDER_OVERRIDES + + +# --- /model session-override follow-up: request_overrides must survive a switch --- + +def test_session_override_applies_request_overrides(): + """A /model switch to a custom provider carries its extra_body into runtime.""" + runner = object.__new__(GatewayRunner) + runner._session_model_overrides = { + "sess1": { + "model": "thinkmodel", + "provider": "custom", + "api_key": "k", + "base_url": "http://10.0.0.1:8000/v1", + "api_mode": "chat_completions", + "request_overrides": PROVIDER_OVERRIDES, + } + } + rk = _runtime_kwargs() # default resolution carried no overrides + model, out = runner._apply_session_model_override("sess1", "oldmodel", rk) + assert model == "thinkmodel" + assert out["request_overrides"] == PROVIDER_OVERRIDES + + +def test_session_override_clears_stale_request_overrides(): + """Switching to a provider with no overrides clears a stale value.""" + runner = object.__new__(GatewayRunner) + runner._session_model_overrides = { + "sess1": { + "model": "plain", + "provider": "openrouter", + "api_key": "k", + "base_url": "https://openrouter.ai/api/v1", + "api_mode": "chat_completions", + "request_overrides": None, + } + } + rk = _runtime_kwargs(request_overrides=PROVIDER_OVERRIDES) # stale, from default + _, out = runner._apply_session_model_override("sess1", "old", rk) + assert out.get("request_overrides") is None + + +def test_session_override_absent_is_noop(): + """No override for the session leaves runtime_kwargs untouched.""" + runner = object.__new__(GatewayRunner) + runner._session_model_overrides = {} + rk = _runtime_kwargs(request_overrides=PROVIDER_OVERRIDES) + model, out = runner._apply_session_model_override("nope", "keepme", rk) + assert model == "keepme" + assert out["request_overrides"] == PROVIDER_OVERRIDES