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:
@@ -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
|
||||
|
||||
@@ -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"):
|
||||
|
||||
@@ -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,
|
||||
|
||||
220
tests/gateway/test_custom_provider_request_overrides.py
Normal file
220
tests/gateway/test_custom_provider_request_overrides.py
Normal 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"}},
|
||||
}
|
||||
94
tests/gateway/test_model_command_request_overrides.py
Normal file
94
tests/gateway/test_model_command_request_overrides.py
Normal 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"}},
|
||||
}
|
||||
Reference in New Issue
Block a user