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:
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user