Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 31 additions & 2 deletions gateway/platforms/api_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -1241,6 +1241,7 @@ def _create_agent(
tool_complete_callback=None,
gateway_session_key: Optional[str] = None,
route: Optional[Dict[str, Any]] = None,
session_model: Optional[str] = None,
) -> Any:
"""
Create an AIAgent instance using the gateway's runtime config.
Expand All @@ -1261,6 +1262,11 @@ def _create_agent(
routing). When set — and no session ``/model`` override exists for
this session — its model/provider/api_key/base_url override the
global defaults for this agent instance only.

``session_model`` is the raw model persisted on a native API session
row (``POST /api/sessions {"model": ...}``) when that value does not
resolve to a ``model_routes`` alias. Session-chat handlers pass either
``route`` (alias hit) or ``session_model`` (raw model), never both.
"""
from run_agent import AIAgent
from gateway.run import (
Expand Down Expand Up @@ -1333,6 +1339,16 @@ def _create_agent(
gateway_session_key or session_id,
)

# Native session API model pinning. POST /api/sessions already
# persists a ``model`` field, but the chat handlers previously threw
# the session row away before constructing the agent. If that stored
# model is a configured model_routes alias, the handler passes the
# resolved ``route`` above so provider credentials/base_url are applied.
# If it is a raw model string, apply it here on top of the default
# runtime, mirroring the intended "session chooses the model" contract.
if session_model and not route and not session_override:
model = session_model

user_config = _load_gateway_config()
enabled_toolsets = sorted(_get_platform_tools(user_config, "api_server"))

Expand Down Expand Up @@ -1882,7 +1898,7 @@ async def _handle_session_chat(self, request: "web.Request") -> "web.Response":
if key_err is not None:
return key_err
session_id = request.match_info["session_id"]
_, err = self._get_existing_session_or_404(session_id)
session, err = self._get_existing_session_or_404(session_id)
if err:
return err
body, err = await self._read_json_body(request)
Expand All @@ -1895,12 +1911,16 @@ async def _handle_session_chat(self, request: "web.Request") -> "web.Response":
if system_prompt is not None and not isinstance(system_prompt, str):
return web.json_response(_openai_error("system_message must be a string", code="invalid_system_message"), status=400)
history = self._conversation_history_for_session(session_id)
stored_model = session.get("model") if isinstance(session, dict) else None
stored_route = self._resolve_route(stored_model)
result, usage = await self._run_agent(
user_message=user_message,
conversation_history=history,
ephemeral_system_prompt=system_prompt,
session_id=session_id,
gateway_session_key=gateway_session_key,
route=stored_route,
session_model=stored_model if stored_model and stored_route is None else None,
)
effective_session_id = result.get("session_id") if isinstance(result, dict) else session_id
final_response = _resolve_media_to_data_urls(result.get("final_response", "") if isinstance(result, dict) else "")
Expand All @@ -1926,7 +1946,7 @@ async def _handle_session_chat_stream(self, request: "web.Request") -> "web.Stre
if key_err is not None:
return key_err
session_id = request.match_info["session_id"]
_, err = self._get_existing_session_or_404(session_id)
session, err = self._get_existing_session_or_404(session_id)
if err:
return err
body, err = await self._read_json_body(request)
Expand All @@ -1938,6 +1958,8 @@ async def _handle_session_chat_stream(self, request: "web.Request") -> "web.Stre
system_prompt = body.get("system_message") or body.get("instructions")
if system_prompt is not None and not isinstance(system_prompt, str):
return web.json_response(_openai_error("system_message must be a string", code="invalid_system_message"), status=400)
stored_model = session.get("model") if isinstance(session, dict) else None
stored_route = self._resolve_route(stored_model)

loop = asyncio.get_running_loop()
queue: "asyncio.Queue[Optional[tuple[str, Dict[str, Any]]]]" = asyncio.Queue()
Expand Down Expand Up @@ -1992,6 +2014,8 @@ async def _run_and_signal() -> None:
stream_delta_callback=_delta,
tool_progress_callback=_tool_progress,
gateway_session_key=gateway_session_key,
route=stored_route,
session_model=stored_model if stored_model and stored_route is None else None,
)
final_response = _resolve_media_to_data_urls(result.get("final_response", "") if isinstance(result, dict) else "")
effective_session_id = result.get("session_id", session_id) if isinstance(result, dict) else session_id
Expand Down Expand Up @@ -4031,6 +4055,7 @@ async def _run_agent(
agent_ref: Optional[list] = None,
gateway_session_key: Optional[str] = None,
route: Optional[Dict[str, Any]] = None,
session_model: Optional[str] = None,
) -> tuple:
"""
Create an agent and run a conversation in a thread executor.
Expand All @@ -4042,6 +4067,9 @@ async def _run_agent(
request's ``model`` field) that overrides the global model/provider
for this specific request.

*session_model* is a raw model persisted on a native API session. It
is used only when the persisted value did not resolve to a route.

If *agent_ref* is a one-element list, the AIAgent instance is stored
at ``agent_ref[0]`` before ``run_conversation`` begins. This allows
callers (e.g. the SSE writer) to call ``agent.interrupt()`` from
Expand All @@ -4067,6 +4095,7 @@ def _run():
tool_complete_callback=tool_complete_callback,
gateway_session_key=gateway_session_key,
route=route,
session_model=session_model,
)
if agent_ref is not None:
agent_ref[0] = agent
Expand Down
17 changes: 17 additions & 0 deletions tests/gateway/test_api_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -4038,6 +4038,23 @@ def __init__(self, **kwargs):
assert captured["model"] == "global/model"
assert captured["api_key"] == "sk-global"

def test_raw_session_model_overrides_global_model(self, monkeypatch):
captured = {}

class FakeAgent:
def __init__(self, **kwargs):
captured.update(kwargs)

_patch_create_agent_runtime(monkeypatch, captured, FakeAgent)
adapter = _make_routing_adapter({})
monkeypatch.setattr(adapter, "_ensure_session_db", lambda: None)
monkeypatch.setattr(adapter, "_session_model_override_for", lambda *_: None)

adapter._create_agent(session_id="s1", session_model="claude-sonnet-4-6")

assert captured["model"] == "claude-sonnet-4-6"
assert captured["api_key"] == "sk-global"

def test_session_model_override_beats_route(self, monkeypatch):
"""A user-issued /model on the session must win over static route config."""
captured = {}
Expand Down
79 changes: 79 additions & 0 deletions tests/gateway/test_session_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -251,6 +251,85 @@ async def test_session_chat_loads_history_and_preserves_session_headers(auth_ada
]


@pytest.mark.asyncio
async def test_session_chat_resolves_persisted_model_route(session_db):
"""A model_routes alias stored on the session must route like /v1 requests.

POST /api/sessions accepts and persists a ``model`` field. If that value is
a configured API-server alias, session chat must pass the resolved route to
_run_agent so provider credentials/base_url are switched too. Passing the
alias as a raw AIAgent model would run it on the wrong provider.
"""
route = {"model": "step-3.7-flash", "provider": "custom:stepfun"}
adapter = APIServerAdapter(
PlatformConfig(enabled=True, extra={"key": "sk-test", "model_routes": {"aria-voice": route}})
)
adapter._session_db = session_db
session_id = session_db.create_session("voice-session", "api_server", model="aria-voice")

mock_run = AsyncMock(return_value=({"final_response": "routed", "session_id": session_id}, {"total_tokens": 3}))
app = _create_session_app(adapter)
with patch.object(adapter, "_run_agent", mock_run):
async with TestClient(TestServer(app)) as cli:
resp = await cli.post(
f"/api/sessions/{session_id}/chat",
json={"message": "hello"},
headers={"Authorization": "Bearer sk-test"},
)
assert resp.status == 200, await resp.text()

kwargs = mock_run.call_args.kwargs
assert kwargs["route"] == route
assert kwargs["session_model"] is None


@pytest.mark.asyncio
async def test_session_chat_threads_persisted_raw_model(auth_adapter, session_db):
"""A non-alias model persisted on the session still pins the chat turn."""
session_id = session_db.create_session("raw-model-session", "api_server", model="claude-sonnet-4-6")

mock_run = AsyncMock(return_value=({"final_response": "raw", "session_id": session_id}, {"total_tokens": 3}))
app = _create_session_app(auth_adapter)
with patch.object(auth_adapter, "_run_agent", mock_run):
async with TestClient(TestServer(app)) as cli:
resp = await cli.post(
f"/api/sessions/{session_id}/chat",
json={"message": "hello"},
headers={"Authorization": "Bearer sk-test"},
)
assert resp.status == 200, await resp.text()

kwargs = mock_run.call_args.kwargs
assert kwargs["route"] is None
assert kwargs["session_model"] == "claude-sonnet-4-6"


@pytest.mark.asyncio
async def test_session_chat_stream_resolves_persisted_model_route(session_db):
"""Streaming native session chat uses the same persisted route as sync chat."""
route = {"model": "step-3.7-flash", "provider": "custom:stepfun"}
adapter = APIServerAdapter(PlatformConfig(enabled=True, extra={"model_routes": {"aria-voice": route}}))
adapter._session_db = session_db
session_id = session_db.create_session("voice-stream-session", "api_server", model="aria-voice")
captured_kwargs = {}

async def fake_run(**kwargs):
captured_kwargs.update(kwargs)
kwargs["stream_delta_callback"]("routed")
return {"final_response": "routed", "session_id": session_id}, {"total_tokens": 4}

app = _create_session_app(adapter)
with patch.object(adapter, "_run_agent", side_effect=fake_run):
async with TestClient(TestServer(app)) as cli:
resp = await cli.post(f"/api/sessions/{session_id}/chat/stream", json={"message": "hello"})
assert resp.status == 200, await resp.text()
body = await resp.text()

assert "event: assistant.completed" in body
assert captured_kwargs["route"] == route
assert captured_kwargs["session_model"] is None


@pytest.mark.asyncio
async def test_session_chat_accepts_multimodal_message(auth_adapter, session_db):
session_id = session_db.create_session("image-session", "api_server")
Expand Down
Loading