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).
This commit is contained in:
Jack
2026-06-27 12:07:50 -05:00
committed by Teknium
parent fc00e36c6b
commit 5d238be2ca
3 changed files with 76 additions and 6 deletions

View File

@@ -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

View File

@@ -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,

View File

@@ -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