141 lines
5.2 KiB
Python
141 lines
5.2 KiB
Python
"""Regression tests: the gateway must preserve a custom provider's
|
|
``request_overrides`` on per-turn agent config.
|
|
|
|
A ``custom_providers`` entry can carry an ``extra_body`` (e.g.
|
|
``chat_template_kwargs`` to toggle a local model's thinking).
|
|
``resolve_runtime_provider`` surfaces it as ``request_overrides`` on the
|
|
resolved runtime dict, but the gateway used to rebuild the runtime from a
|
|
fixed key whitelist that omitted it -- so the provider's configured
|
|
``extra_body`` never reached the model on the gateway path, and only
|
|
``/fast`` service-tier overrides survived.
|
|
"""
|
|
|
|
import pytest
|
|
|
|
from gateway.run import GatewayRunner
|
|
|
|
|
|
PROVIDER_OVERRIDES = {"extra_body": {"chat_template_kwargs": {"enable_thinking": True}}}
|
|
|
|
|
|
def _runtime_kwargs(**extra):
|
|
base = {
|
|
"api_key": "no-key-required",
|
|
"base_url": "http://10.0.0.1:8000/v1",
|
|
"provider": "custom",
|
|
"api_mode": "chat_completions",
|
|
"command": None,
|
|
"args": [],
|
|
"credential_pool": None,
|
|
"max_tokens": None,
|
|
}
|
|
base.update(extra)
|
|
return base
|
|
|
|
|
|
def _runner(service_tier=None):
|
|
runner = object.__new__(GatewayRunner)
|
|
runner._service_tier = service_tier
|
|
return runner
|
|
|
|
|
|
def test_provider_request_overrides_preserved_without_service_tier():
|
|
"""No /fast: the provider's extra_body must pass straight through."""
|
|
runner = _runner(service_tier=None)
|
|
rk = _runtime_kwargs(request_overrides=PROVIDER_OVERRIDES)
|
|
route = runner._resolve_turn_agent_config("hi", "main", rk)
|
|
assert route["request_overrides"] == PROVIDER_OVERRIDES
|
|
# A copy, not an alias into runtime_kwargs.
|
|
assert route["request_overrides"] is not rk["request_overrides"]
|
|
|
|
|
|
def test_provider_request_overrides_merged_under_fast_mode(monkeypatch):
|
|
"""/fast active: provider extra_body AND the service-tier marker both survive."""
|
|
monkeypatch.setattr(
|
|
"hermes_cli.models.resolve_fast_mode_overrides",
|
|
lambda model_id, **_route: {"service_tier": "priority"},
|
|
)
|
|
runner = _runner(service_tier="priority")
|
|
rk = _runtime_kwargs(request_overrides=PROVIDER_OVERRIDES)
|
|
route = runner._resolve_turn_agent_config("hi", "main", rk)
|
|
assert route["request_overrides"]["extra_body"] == PROVIDER_OVERRIDES["extra_body"]
|
|
assert route["request_overrides"]["service_tier"] == "priority"
|
|
|
|
|
|
def test_no_provider_overrides_yields_empty():
|
|
"""Regression: absent provider overrides, behaviour is unchanged ({})."""
|
|
runner = _runner(service_tier=None)
|
|
route = runner._resolve_turn_agent_config("hi", "main", _runtime_kwargs())
|
|
assert route["request_overrides"] == {}
|
|
|
|
|
|
def test_resolve_runtime_agent_kwargs_carries_request_overrides(monkeypatch):
|
|
"""The module-level runtime resolver must not drop request_overrides."""
|
|
import gateway.run as gateway_run
|
|
|
|
fake_runtime = {
|
|
"api_key": "k",
|
|
"base_url": "http://10.0.0.1:8000/v1",
|
|
"provider": "custom",
|
|
"api_mode": "chat_completions",
|
|
"request_overrides": PROVIDER_OVERRIDES,
|
|
}
|
|
monkeypatch.setattr(
|
|
"hermes_cli.runtime_provider.resolve_runtime_provider",
|
|
lambda *a, **k: dict(fake_runtime),
|
|
)
|
|
monkeypatch.setattr(
|
|
"hermes_cli.runtime_provider._get_model_config", lambda: {}
|
|
)
|
|
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
|