fix: preserve named custom provider request_overrides in gateway and /model switches

Carry provider-derived request_overrides through runtime resolution,
fallback projection, session /model state, restart rehydration, and
turn-route merge so named custom providers keep extra_body and related
overrides.
This commit is contained in:
CharZhou
2026-07-20 08:56:38 +08:00
committed by Teknium
parent 556777ddb1
commit 863aac9012
5 changed files with 349 additions and 2 deletions

View File

@@ -2969,6 +2969,7 @@ def _resolve_runtime_agent_kwargs() -> dict:
"command": runtime.get("command"),
"args": list(runtime.get("args") or []),
"credential_pool": runtime.get("credential_pool"),
"request_overrides": dict(runtime.get("request_overrides") or {}),
"max_tokens": max_tokens,
}
@@ -3109,9 +3110,23 @@ def _resolve_runtime_agent_kwargs_for_provider(provider: str) -> dict:
"command": runtime.get("command"),
"args": list(runtime.get("args") or []),
"credential_pool": runtime.get("credential_pool"),
"request_overrides": dict(runtime.get("request_overrides") or {}),
}
def _deep_merge_request_overrides(base: Optional[dict], override: Optional[dict]) -> dict:
"""Merge request_overrides dicts, deep-merging nested dictionaries."""
from hermes_cli.config import _deep_merge
base_dict = dict(base or {})
override_dict = dict(override or {})
if not base_dict:
return override_dict
if not override_dict:
return base_dict
return _deep_merge(base_dict, override_dict)
def _credential_pool_for_provider(provider: Optional[str]):
"""Return the live credential pool for a provider id (e.g. ``custom:hyper``)."""
if not provider or not str(provider).strip():
@@ -3167,6 +3182,7 @@ def _try_resolve_fallback_provider() -> dict | None:
"command": runtime.get("command"),
"args": list(runtime.get("args") or []),
"credential_pool": runtime.get("credential_pool"),
"request_overrides": dict(runtime.get("request_overrides") or {}),
"model": entry.get("model"),
}
except Exception as fb_exc:
@@ -8459,6 +8475,7 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
"credential_pool": runtime_kwargs.get("credential_pool"),
"max_tokens": runtime_kwargs.get("max_tokens"),
}
base_request_overrides = dict(runtime_kwargs.get("request_overrides") or {})
route = {
"model": model,
"runtime": runtime,
@@ -8475,14 +8492,17 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
service_tier = getattr(self, "_service_tier", None)
if not service_tier:
route["request_overrides"] = {}
route["request_overrides"] = base_request_overrides
return route
try:
overrides = resolve_fast_mode_overrides(route["model"])
except Exception:
overrides = None
route["request_overrides"] = overrides or {}
route["request_overrides"] = _deep_merge_request_overrides(
base_request_overrides,
overrides or {},
)
return route
def _sync_session_model_from_agent(self, session_id: str, agent: Any) -> None:
@@ -27409,6 +27429,9 @@ class GatewayRunner(GatewayAuthorizationMixin, GatewayKanbanWatchersMixin, Gatew
override["api_key"] = runtime.get("api_key")
override["api_mode"] = runtime.get("api_mode")
override["credential_pool"] = runtime.get("credential_pool")
override["request_overrides"] = dict(
runtime.get("request_overrides") or {}
)
if not override.get("base_url"):
override["base_url"] = runtime.get("base_url")
except Exception:
@@ -27443,6 +27466,12 @@ 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,
)
if (
runtime_kwargs.get("api_key")
and runtime_kwargs.get("credential_pool") is None

View File

@@ -2009,6 +2009,7 @@ class GatewaySlashCommandsMixin:
"api_key": result.api_key,
"base_url": result.base_url,
"api_mode": result.api_mode,
"request_overrides": dict(result.request_overrides or {}),
}
# Write-through the non-secret parts to the session
@@ -2321,6 +2322,7 @@ class GatewaySlashCommandsMixin:
"api_key": result.api_key,
"base_url": result.base_url,
"api_mode": result.api_mode,
"request_overrides": dict(result.request_overrides or {}),
}
if one_turn:
if not hasattr(self, "_pending_one_turn_model_restores"):

View File

@@ -620,6 +620,7 @@ class ModelSwitchResult:
api_key: str = ""
base_url: str = ""
api_mode: str = ""
request_overrides: Optional[dict] = None
error_message: str = ""
warning_message: str = ""
provider_label: str = ""
@@ -2192,6 +2193,7 @@ def switch_model(
api_key=api_key,
base_url=base_url,
api_mode=api_mode,
request_overrides=dict(request_overrides or {}),
warning_message=" | ".join(warnings) if warnings else "",
provider_label=provider_label,
resolved_via_alias=resolved_alias,

View File

@@ -0,0 +1,220 @@
"""Regression tests for gateway preservation of provider-derived request_overrides.
Named custom providers can return request_overrides (for example
``extra_body.text.verbosity`` for OpenAI Responses). The gateway must preserve
those overrides on the runtime path and merge fast-mode overrides on top rather
than replacing them with an empty dict.
"""
from __future__ import annotations
import asyncio
import sys
import threading
import types
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import gateway.run as gateway_run
from gateway.config import Platform
from gateway.session import SessionSource
class _CapturingAgent:
last_init = None
def __init__(self, *args, **kwargs):
type(self).last_init = dict(kwargs)
self.tools = []
self.request_overrides = dict(kwargs.get("request_overrides") or {})
def run_conversation(self, user_message: str, conversation_history=None, task_id=None):
return {
"final_response": "ok",
"messages": [],
"api_calls": 1,
}
def _install_fake_agent(monkeypatch):
fake_run_agent = types.ModuleType("run_agent")
fake_run_agent.AIAgent = _CapturingAgent
monkeypatch.setitem(sys.modules, "run_agent", fake_run_agent)
def _make_runner():
runner = object.__new__(gateway_run.GatewayRunner)
runner.adapters = {}
runner.session_store = None
runner.config = None
runner._voice_mode = {}
runner._ephemeral_system_prompt = ""
runner._prefill_messages = []
runner._reasoning_config = None
runner._show_reasoning = False
runner._provider_routing = {}
runner._fallback_model = None
runner._service_tier = None
runner._running_agents = {}
runner._running_agents_ts = {}
runner._background_tasks = set()
runner._session_db = None
runner._session_model_overrides = {}
runner._session_reasoning_overrides = {}
runner._pending_model_notes = {}
runner._pending_approvals = {}
runner._agent_cache = {}
runner._agent_cache_lock = threading.Lock()
runner._get_or_create_gateway_honcho = lambda session_key: (None, None)
runner.hooks = MagicMock()
runner.hooks.emit = AsyncMock()
runner.hooks.loaded_hooks = []
return runner
def _make_source() -> SessionSource:
return SessionSource(
platform=Platform.FEISHU,
chat_id="ou_test",
chat_type="dm",
user_id="user-1",
user_name="tester",
)
def test_resolve_runtime_agent_kwargs_preserves_request_overrides(monkeypatch):
monkeypatch.setattr(
"hermes_cli.runtime_provider.resolve_runtime_provider",
lambda: {
"api_key": "***",
"base_url": "https://example.test/v1",
"provider": "custom",
"api_mode": "codex_responses",
"command": None,
"args": [],
"credential_pool": None,
"request_overrides": {
"extra_body": {"text": {"verbosity": "low"}},
},
},
)
result = gateway_run._resolve_runtime_agent_kwargs()
assert result["request_overrides"] == {
"extra_body": {"text": {"verbosity": "low"}},
}
def test_turn_route_preserves_provider_request_overrides_without_fast_mode():
runner = _make_runner()
runner._service_tier = None
runtime_kwargs = {
"api_key": "***",
"base_url": "https://example.test/v1",
"provider": "custom",
"api_mode": "codex_responses",
"command": None,
"args": [],
"credential_pool": None,
"request_overrides": {
"extra_body": {"text": {"verbosity": "low"}},
},
}
route = gateway_run.GatewayRunner._resolve_turn_agent_config(
runner,
"hi",
"gpt-5.4",
runtime_kwargs,
)
assert route["request_overrides"] == {
"extra_body": {"text": {"verbosity": "low"}},
}
def test_turn_route_merges_fast_mode_with_provider_request_overrides():
runner = _make_runner()
runner._service_tier = "priority"
runtime_kwargs = {
"api_key": "***",
"base_url": "https://example.test/v1",
"provider": "custom",
"api_mode": "codex_responses",
"command": None,
"args": [],
"credential_pool": None,
"request_overrides": {
"extra_body": {"text": {"verbosity": "low"}},
},
}
with patch(
"hermes_cli.models.resolve_fast_mode_overrides",
return_value={"service_tier": "priority"},
):
route = gateway_run.GatewayRunner._resolve_turn_agent_config(
runner,
"hi",
"gpt-5.4",
runtime_kwargs,
)
assert route["request_overrides"] == {
"extra_body": {"text": {"verbosity": "low"}},
"service_tier": "priority",
}
@pytest.mark.asyncio
async def test_run_agent_preserves_provider_request_overrides_on_gateway_path(monkeypatch):
monkeypatch.setattr(gateway_run, "_load_gateway_config", lambda: {})
monkeypatch.setattr(gateway_run, "load_dotenv", lambda *args, **kwargs: None)
monkeypatch.setattr(gateway_run, "_load_gateway_runtime_config", lambda: {})
monkeypatch.setattr(gateway_run, "_resolve_gateway_model", lambda config=None: "gpt-5.4")
monkeypatch.setattr(
gateway_run,
"_resolve_runtime_agent_kwargs",
lambda: {
"provider": "custom",
"api_mode": "codex_responses",
"base_url": "https://example.test/v1",
"api_key": "***",
"request_overrides": {
"extra_body": {"text": {"verbosity": "low"}},
},
},
)
_install_fake_agent(monkeypatch)
import hermes_cli.tools_config as tools_config
monkeypatch.setattr(tools_config, "_get_platform_tools", lambda user_config, platform_key: {"core"})
runner = _make_runner()
source = _make_source()
session_key = "agent:main:feishu:dm:ou_test"
runner.session_store = SimpleNamespace(
get_or_create_session=lambda _source: SimpleNamespace(session_id="session-1"),
load_transcript=lambda _session_id: [],
)
_CapturingAgent.last_init = None
result = await runner._run_agent(
message="hi",
context_prompt="",
history=[],
source=source,
session_id="session-1",
session_key=session_key,
)
assert result["final_response"] == "ok"
assert _CapturingAgent.last_init is not None
assert _CapturingAgent.last_init["request_overrides"] == {
"extra_body": {"text": {"verbosity": "low"}},
}

View File

@@ -0,0 +1,94 @@
"""Regression tests for gateway /model preserving named-custom request_overrides."""
import pytest
from gateway.config import Platform
from gateway.platforms.base import MessageEvent, MessageType
from gateway.run import GatewayRunner
from gateway.session import SessionSource
def _make_runner():
runner = object.__new__(GatewayRunner)
runner.adapters = {}
runner._voice_mode = {}
runner._session_model_overrides = {}
runner._pending_model_notes = {}
runner._agent_cache = {}
runner._agent_cache_lock = None
runner._session_db = None
runner._evict_cached_agent = lambda _session_key: None
runner.session_store = None
return runner
def _make_event(text="/model"):
return MessageEvent(
text=text,
message_type=MessageType.TEXT,
source=SessionSource(
platform=Platform.FEISHU,
chat_id="ou_test",
chat_type="dm",
user_id="user-1",
),
)
@pytest.mark.asyncio
async def test_handle_model_command_stores_request_overrides_for_named_custom_provider(
tmp_path,
monkeypatch,
):
import gateway.run as gateway_run
from hermes_cli.model_switch import ModelSwitchResult
hermes_home = tmp_path / ".hermes"
hermes_home.mkdir()
(hermes_home / "config.yaml").write_text(
"""
model:
default: gpt-5.4
provider: openai-codex
providers: {}
custom_providers:
- name: Local (127.0.0.1:4141)
base_url: http://127.0.0.1:4141/v1
model: rotator-openrouter-coding
extra_body:
text:
verbosity: low
""".lstrip(),
encoding="utf-8",
)
monkeypatch.setattr(gateway_run, "_hermes_home", hermes_home)
monkeypatch.setattr("agent.models_dev.fetch_models_dev", lambda: {})
monkeypatch.setattr(
"hermes_cli.model_switch.switch_model",
lambda **kw: ModelSwitchResult(
success=True,
new_model="rotator-openrouter-coding",
target_provider="custom:local-(127.0.0.1:4141)",
provider_changed=True,
api_key="no-key-required",
base_url="http://127.0.0.1:4141/v1",
api_mode="codex_responses",
request_overrides={
"extra_body": {"text": {"verbosity": "low"}},
},
provider_label="Local (127.0.0.1:4141)",
is_global=False,
),
)
runner = _make_runner()
event = _make_event("/model rotator-openrouter-coding --provider custom:local-(127.0.0.1:4141)")
result = await runner._handle_model_command(event)
assert result is not None
session_key = runner._session_key_for_source(event.source)
assert runner._session_model_overrides[session_key]["request_overrides"] == {
"extra_body": {"text": {"verbosity": "low"}},
}