From 5d238be2ca4ec5abdd8785902f7d6c7e11bfde9f Mon Sep 17 00:00:00 2001 From: Jack <4762467+apollo-orbit-dev@users.noreply.github.com> Date: Sat, 27 Jun 2026 12:07:50 -0500 Subject: [PATCH] fix(gateway): carry request_overrides through /model session overrides Follow-up to the previous commit (which fixed the default/fallback provider path). A mid-session `/model` switch stores a per-session override bundle in `_session_model_overrides` that omitted `request_overrides`, and the two consumers (`_resolve_session_agent_runtime` fast path and `_apply_session_model_override`) only copied provider/api_key/base_url/api_mode. So switching *to* a custom provider via `/model` did not apply its `extra_body`. - `ModelSwitchResult` gains a `request_overrides` field, derived for the switched provider via `_get_named_custom_provider` / `_custom_provider_request_overrides` (the same overrides `resolve_runtime_provider` surfaces for the default path). - Both `/model` override-storage sites in slash_commands.py persist it. - Both consumers apply it; `_apply_session_model_override` also clears a stale value when switching to a provider that has none. Extends tests/gateway/test_turn_request_overrides.py (3 new cases). --- gateway/run.py | 17 ++++--- hermes_cli/model_switch.py | 16 +++++++ tests/gateway/test_turn_request_overrides.py | 49 ++++++++++++++++++++ 3 files changed, 76 insertions(+), 6 deletions(-) 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