From 7dd00bb47d8064f8a9400a35c7e04c3ceece1e57 Mon Sep 17 00:00:00 2001 From: Frederick Date: Fri, 24 Jul 2026 11:37:18 -0700 Subject: [PATCH] fix(api_server): close divergence gaps from gateway/run.py MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Three parity fixes between the API server and the native gateway's agent-runtime resolution, integrated with the provider-aware request routing that landed in #70853: - Session-persisted model is honored: POST /api/sessions {"model": ...} stores a model that the chat handlers previously fetched and threw away. A stored value that matches a model_routes alias goes through the route path (route provider/credentials apply); a raw model string threads through as session_model, pinning the session's turns ahead of per-request body values but below an explicit session /model override. - Empty-model recovery: provider-catalog default when config has no model.default but a provider resolved, plus last-known-good model recovery (#35314) keyed on gateway_session_key only (never ephemeral session_id — no unbounded growth from one-off requests). - Provider auth failures surface as controlled responses: RuntimeError from _resolve_runtime_agent_kwargs() is re-raised as a dedicated _ProviderAuthResolutionError at the call site, caught narrowly in _run_agent() and the /v1/runs executor to return run.py's response shape instead of an undifferentiated 500 (session-chat endpoints previously returned a raw aiohttp 500 with no JSON body). Salvaged from PR #57947 by @FvanW; session-model route-alias resolution from PR #59941 by @kaishi00. Co-authored-by: kaishi00 --- contributors/emails/fred.vanwagenen@gmail.com | 1 + gateway/platforms/api_server.py | 210 +++++++++++++++++- tests/gateway/test_api_server.py | 138 ++++++++++++ tests/gateway/test_api_server_runs.py | 35 +++ tests/gateway/test_session_api.py | 179 +++++++++++++++ 5 files changed, 552 insertions(+), 11 deletions(-) create mode 100644 contributors/emails/fred.vanwagenen@gmail.com diff --git a/contributors/emails/fred.vanwagenen@gmail.com b/contributors/emails/fred.vanwagenen@gmail.com new file mode 100644 index 000000000000..e6175ab6667b --- /dev/null +++ b/contributors/emails/fred.vanwagenen@gmail.com @@ -0,0 +1 @@ +FvanW diff --git a/gateway/platforms/api_server.py b/gateway/platforms/api_server.py index ff2a1bcd3862..15760f25f868 100644 --- a/gateway/platforms/api_server.py +++ b/gateway/platforms/api_server.py @@ -1149,6 +1149,22 @@ except Exception: # pragma: no cover - scanner is optional hardening _scan_cron_prompt = None +class _ProviderAuthResolutionError(RuntimeError): + """Raised only when gateway.run._resolve_runtime_agent_kwargs() fails + to resolve provider credentials. + + That function is the sole raiser of RuntimeError(format_runtime_ + provider_error(...)) anywhere in _create_agent()'s call graph. + Re-raising it as this dedicated subclass -- instead of catching bare + RuntimeError around the much wider _create_agent()+run_conversation() + span -- lets callers distinguish "provider auth/credential failure" + from any other RuntimeError a provider adapter or run_conversation() + might legitimately raise (e.g. run_agent.py's "Failed to recreate + closed OpenAI client"), which a bare `except RuntimeError` there would + otherwise mislabel as an auth failure. + """ + + class APIServerAdapter(BasePlatformAdapter): """ OpenAI-compatible HTTP API server adapter. @@ -1237,6 +1253,13 @@ class APIServerAdapter(BasePlatformAdapter): # in-flight run by run_id. self._run_approval_sessions: Dict[str, str] = {} self._session_db: Optional[Any] = None # Lazy-init SessionDB for session continuity + # Last-known-good resolved model per session (keyed by gateway_session_key + # ONLY — never session_id, which rotates/is ephemeral for one-off API + # server requests; "*" is the process-wide fallback), mirroring + # GatewayRunner._last_resolved_model in run.py — recovers from a + # transient empty model resolution (#35314) instead of building an + # agent with model="" that 400s every call until manual retry. + self._last_resolved_model: Dict[str, str] = {} self._session_db_lock: Optional[asyncio.Lock] = None # Single-flight for lazy init # Concurrency cap shared across all agent-serving endpoints # (/v1/chat/completions, /v1/responses, /v1/runs). Read from @@ -2056,6 +2079,7 @@ class APIServerAdapter(BasePlatformAdapter): requested_provider: Optional[str] = None, model_options: Optional[Dict[str, Any]] = None, route: Optional[Dict[str, Any]] = None, + session_model: Optional[str] = None, ) -> Any: """ Create an AIAgent instance using the gateway's runtime config. @@ -2076,6 +2100,13 @@ class APIServerAdapter(BasePlatformAdapter): 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 at creation time (``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. Precedence: session ``/model`` override → + ``session_model`` → route alias / per-request selection → global. """ from run_agent import AIAgent from gateway.run import ( @@ -2088,7 +2119,18 @@ class APIServerAdapter(BasePlatformAdapter): ) from hermes_cli.tools_config import _get_platform_tools - runtime_kwargs = _resolve_runtime_agent_kwargs() + # Catch RuntimeError ONLY around this call, not the wider + # _create_agent()+run_conversation() span -- + # _resolve_runtime_agent_kwargs() is the sole raiser of + # RuntimeError(format_runtime_provider_error(...)) for provider + # auth/credential failure. Re-raising as + # _ProviderAuthResolutionError lets _run_agent() (and + # _handle_runs()) distinguish this from an unrelated RuntimeError + # elsewhere in the call graph. + try: + runtime_kwargs = _resolve_runtime_agent_kwargs() + except RuntimeError as exc: + raise _ProviderAuthResolutionError(str(exc)) from exc reasoning_config = GatewayRunner._load_reasoning_config() model = _resolve_gateway_model() @@ -2148,29 +2190,51 @@ class APIServerAdapter(BasePlatformAdapter): return None # Final precedence mirrors the gateway contract: - # session /model override → model_routes mapping selected by the request - # model alias → direct per-request provider/model → global defaults. - # model_options stay request-scoped regardless of which selection wins. + # session /model override → session-persisted model (POST + # /api/sessions {"model": ...}) → model_routes mapping selected by + # the request model alias → direct per-request provider/model → + # global defaults. model_options stay request-scoped regardless + # of which selection wins. session_key = gateway_session_key or session_id + session_row_model = _clean_request_string(session_model) session_override = self._session_model_override_for(session_key) if session_override: - session_model = _clean_request_string(session_override.get("model")) or model + override_model = _clean_request_string(session_override.get("model")) or model session_provider = _clean_request_string(session_override.get("provider")) current_provider = _clean_request_string(runtime_kwargs.get("provider")) provider_runtime = _resolve_provider_runtime( session_provider or current_provider, - target_model=session_model, + target_model=override_model, required=False, ) if provider_runtime: _apply_runtime_agent_overrides(runtime_kwargs, provider_runtime) _apply_runtime_agent_overrides(runtime_kwargs, session_override) - model = session_model + model = override_model if route or request_model or request_provider: logger.debug( "api_server request selection skipped: session /model override wins for %s", session_key or "", ) + elif session_row_model: + # Session-persisted model (raw string that resolved to no route + # alias). Pins this session's turns ahead of per-request body + # values — a session's chosen model is a standing selection, + # matching the native gateway's session-model semantics. + current_provider = _clean_request_string(runtime_kwargs.get("provider")) + provider_runtime = _resolve_provider_runtime( + current_provider, + target_model=session_row_model, + required=False, + ) + if provider_runtime: + _apply_runtime_agent_overrides(runtime_kwargs, provider_runtime) + model = session_row_model + if request_model or request_provider: + logger.debug( + "api_server request selection skipped: session-persisted model wins for %s", + session_key or "", + ) else: if route is not None: # The request's ``model`` field selected this route, so its @@ -2210,6 +2274,56 @@ class APIServerAdapter(BasePlatformAdapter): route_provider or "", request_provider or "", ) + + # When the config has no model.default but a provider was resolved + # (e.g. user ran `hermes auth add openai-codex` without `hermes model`), + # fall back to the provider's first catalog model so the API call + # doesn't fail with "model must be a non-empty string". Mirrors + # run.py::_resolve_session_agent_runtime. Runs after the selection + # block above so a route/session/request override that already + # resolved a model is never treated as "empty" here. + if not model and runtime_kwargs.get("provider"): + try: + from hermes_cli.models import get_default_model_for_provider + model = get_default_model_for_provider(runtime_kwargs["provider"]) + if model: + logger.info( + "No model configured — defaulting to %s for provider %s", + model, runtime_kwargs["provider"], + ) + except Exception: + pass + + # Final safety net (#35314): if resolution still produced an empty + # model — e.g. a transient config-cache miss — reuse the last model + # successfully resolved for this session (or, failing that, the most + # recent one resolved process-wide). Building an agent with model="" + # makes every API call fail HTTP 400 until a manual retry. Mirrors + # run.py::_resolve_session_agent_runtime. + # + # Cache key is gateway_session_key ONLY, never session_id — unlike + # run.py's native gateway (stable, long-lived chat scopes), the API + # server hands out a fresh UUID session_id per one-off request + # (/v1/responses, /v1/runs when no explicit session is supplied). + # Keying on session_id would leave one permanent dict entry per + # stateless request, growing unbounded for the life of the process. + _resolved_key = gateway_session_key or "" + if not model: + _recovered = (self._last_resolved_model.get(_resolved_key) + or self._last_resolved_model.get("*")) + if _recovered: + logger.warning( + "Empty model resolved for session=%s — recovering " + "last-known-good model %s (config read likely returned " + "empty; see #35314)", + _resolved_key, _recovered, + ) + model = _recovered + elif model: + if _resolved_key: + self._last_resolved_model[_resolved_key] = model + self._last_resolved_model["*"] = model + user_config = _load_gateway_config() enabled_toolsets = sorted(_get_platform_tools(user_config, "api_server")) @@ -2884,7 +2998,7 @@ class APIServerAdapter(BasePlatformAdapter): if key_err is not None: return key_err session_id = request.match_info["session_id"] - _, err = await self._get_existing_session_or_404(session_id) + session, err = await self._get_existing_session_or_404(session_id) if err: return err body, err = await self._read_json_body(request) @@ -2896,7 +3010,16 @@ class APIServerAdapter(BasePlatformAdapter): 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) - route = self._resolve_route(body.get("model")) + # Session-persisted model (POST /api/sessions {"model": ...}) — + # previously fetched and discarded here, so a session's chosen model + # silently had no effect on any chat turn. If the stored value is a + # model_routes alias, apply it through the route path so route + # provider/credentials come along; a raw model string threads + # through as session_model. + stored_model = session.get("model") if isinstance(session, dict) else None + stored_route = self._resolve_route(stored_model) + route = stored_route or self._resolve_route(body.get("model")) + session_model = stored_model if (stored_model and stored_route is None) else None agent_overrides = _request_agent_overrides(body, virtual_model=self._model_name) selection_error = self._request_route_conflict_error( session_id=session_id, @@ -2915,6 +3038,7 @@ class APIServerAdapter(BasePlatformAdapter): session_id=session_id, gateway_session_key=gateway_session_key, route=route, + session_model=session_model, **agent_overrides, ) effective_session_id = result.get("session_id") if isinstance(result, dict) else session_id @@ -2939,7 +3063,7 @@ class APIServerAdapter(BasePlatformAdapter): if key_err is not None: return key_err session_id = request.match_info["session_id"] - _, err = await self._get_existing_session_or_404(session_id) + session, err = await self._get_existing_session_or_404(session_id) if err: return err body, err = await self._read_json_body(request) @@ -2951,7 +3075,12 @@ class APIServerAdapter(BasePlatformAdapter): 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) - route = self._resolve_route(body.get("model")) + # Session-persisted model — see _handle_session_chat for the + # non-streaming twin of this resolution. + stored_model = session.get("model") if isinstance(session, dict) else None + stored_route = self._resolve_route(stored_model) + route = stored_route or self._resolve_route(body.get("model")) + session_model = stored_model if (stored_model and stored_route is None) else None agent_overrides = _request_agent_overrides(body, virtual_model=self._model_name) selection_error = self._request_route_conflict_error( session_id=session_id, @@ -3017,6 +3146,7 @@ class APIServerAdapter(BasePlatformAdapter): tool_progress_callback=_tool_progress, gateway_session_key=gateway_session_key, route=route, + session_model=session_model, **agent_overrides, ) final_response = _resolve_media_to_data_urls(result.get("final_response", "") if isinstance(result, dict) else "") @@ -5150,6 +5280,7 @@ class APIServerAdapter(BasePlatformAdapter): requested_provider: Optional[str] = None, model_options: Optional[Dict[str, Any]] = None, route: Optional[Dict[str, Any]] = None, + session_model: Optional[str] = None, ) -> tuple: """ Create an agent and run a conversation in a thread executor. @@ -5161,6 +5292,10 @@ class APIServerAdapter(BasePlatformAdapter): 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 + row. It is used only when the persisted value did not resolve to a + ``model_routes`` alias — see ``_create_agent`` for precedence. + 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 @@ -5194,6 +5329,7 @@ class APIServerAdapter(BasePlatformAdapter): requested_provider=requested_provider, model_options=model_options, route=route, + session_model=session_model, ) if agent_ref is not None: agent_ref[0] = agent @@ -5227,6 +5363,33 @@ class APIServerAdapter(BasePlatformAdapter): if _compacted_in_place or _session_rotated: result["_compressed"] = True return result, usage + except _ProviderAuthResolutionError as exc: + # Only _ProviderAuthResolutionError — raised exclusively + # where _resolve_runtime_agent_kwargs() is called inside + # _create_agent() — means a provider auth/credential + # failure. Catching bare RuntimeError here would + # mislabel unrelated RuntimeErrors from + # run_conversation() (e.g. "Failed to recreate closed + # OpenAI client") as auth failures. Matches run.py's + # response shape (final_response text, no HTTP error). + # Previously this propagated unhandled: + # /v1/chat/completions caught it as an undifferentiated + # "Internal server error" 500, and + # /api/sessions/{id}/chat[/stream] didn't catch it at + # all (raw aiohttp 500, no JSON body). Handling it + # here, once, covers every _run_agent() caller; + # /v1/runs has its own branch in its executor. + logger.warning("Provider authentication failed for session=%s: %s", + session_id or "", exc) + return ( + { + "final_response": f"⚠️ Provider authentication failed: {exc}", + "messages": [], + "api_calls": 0, + "tools": [], + }, + {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}, + ) finally: clear_session_vars(tokens) @@ -5619,6 +5782,31 @@ class APIServerAdapter(BasePlatformAdapter): except Exception: pass raise + except _ProviderAuthResolutionError as exc: + # /v1/runs builds its own agent via _create_agent() and does + # not route through _run_agent() (see that method's own + # _ProviderAuthResolutionError branch), so it needs its own + # handling to surface the same distinguished, controlled + # message the other endpoints give a provider auth/credential + # failure, instead of falling through to the generic + # except-Exception branch below. + logger.warning("Provider authentication failed for run=%s: %s", run_id, exc) + error_msg = f"⚠️ Provider authentication failed: {exc}" + self._set_run_status( + run_id, + "failed", + error=error_msg, + last_event="run.failed", + ) + try: + _put_event_if_active({ + "event": "run.failed", + "run_id": run_id, + "timestamp": time.time(), + "error": error_msg, + }) + except Exception: + pass except Exception as exc: logger.exception("[api_server] run %s failed", run_id) self._set_run_status( diff --git a/tests/gateway/test_api_server.py b/tests/gateway/test_api_server.py index f14320a96102..347dc423a86d 100644 --- a/tests/gateway/test_api_server.py +++ b/tests/gateway/test_api_server.py @@ -5524,3 +5524,141 @@ class TestRouteWithoutModelKeepsDefault: assert captured["model"] == "global/model" assert captured["api_key"] == "sk-route" + + +# --------------------------------------------------------------------------- +# Empty-model recovery + provider-auth error typing in _create_agent +# (salvaged from PR #57947 by @FvanW) +# --------------------------------------------------------------------------- + + +class TestCreateAgentModelRecovery: + def test_create_agent_defaults_to_provider_catalog_model_when_empty(self, monkeypatch): + """api_server.py had no equivalent of run.py's provider-catalog + default when model resolves empty but a provider did resolve (e.g. + `hermes auth add openai-codex` without `hermes model`) — + AIAgent(model="") 400s every call.""" + captured = {} + + class FakeAgent: + def __init__(self, **kwargs): + captured.update(kwargs) + + _patch_create_agent_runtime(monkeypatch, captured, FakeAgent) + monkeypatch.setattr( + "gateway.run._resolve_runtime_agent_kwargs", + lambda: {"provider": "openai-codex", "base_url": "https://example.test/v1", + "api_mode": "codex_responses"}, + ) + monkeypatch.setattr("gateway.run._resolve_gateway_model", lambda: "") + monkeypatch.setattr( + "hermes_cli.models.get_default_model_for_provider", + lambda provider: "gpt-5.5-codex" if provider == "openai-codex" else None, + ) + + adapter = APIServerAdapter(PlatformConfig(enabled=True)) + monkeypatch.setattr(adapter, "_ensure_session_db", lambda: None) + + agent = adapter._create_agent(session_id="api-session") + + assert isinstance(agent, FakeAgent) + assert captured["model"] == "gpt-5.5-codex" + + def test_create_agent_recovers_last_known_good_model_when_empty(self, monkeypatch): + """Last-known-good recovery (#35314): a transient config-cache miss + producing an empty model would build AIAgent(model="") and fail every + call until manual retry, instead of reusing the model that just + worked.""" + captured = [] + + class FakeAgent: + def __init__(self, **kwargs): + captured.append(dict(kwargs)) + + _patch_create_agent_runtime(monkeypatch, {}, FakeAgent) + monkeypatch.setattr("run_agent.AIAgent", FakeAgent) + + adapter = APIServerAdapter(PlatformConfig(enabled=True)) + monkeypatch.setattr(adapter, "_ensure_session_db", lambda: None) + + # Turn 1: model resolves fine — populates the last-known-good cache + # (keyed on gateway_session_key). + monkeypatch.setattr("gateway.run._resolve_gateway_model", lambda: "minimax/minimax-m3") + adapter._create_agent(session_id="api-session", gateway_session_key="stable-chan-1") + assert captured[0]["model"] == "minimax/minimax-m3" + assert adapter._last_resolved_model["stable-chan-1"] == "minimax/minimax-m3" + + # Turn 2: transient empty resolution, no provider catalog default — + # must recover the model from turn 1, not build model="". + monkeypatch.setattr("gateway.run._resolve_gateway_model", lambda: "") + monkeypatch.setattr( + "gateway.run._resolve_runtime_agent_kwargs", + lambda: {"provider": None, "base_url": None, "api_mode": None}, + ) + adapter._create_agent(session_id="another-session", gateway_session_key="stable-chan-1") + assert captured[1]["model"] == "minimax/minimax-m3" + + def test_last_resolved_model_cache_does_not_grow_per_ephemeral_session_id(self, monkeypatch): + """/v1/responses and /v1/runs hand out a fresh UUID session_id per + one-off request (no gateway_session_key). Keying the last-known-good + cache on session_id would leave one permanent dict entry per + stateless request, growing unbounded for the life of the process. + Only gateway_session_key may create an entry.""" + class FakeAgent: + def __init__(self, **kwargs): + pass + + _patch_create_agent_runtime(monkeypatch, {}, FakeAgent) + + adapter = APIServerAdapter(PlatformConfig(enabled=True)) + monkeypatch.setattr(adapter, "_ensure_session_db", lambda: None) + + import uuid as _uuid + for _ in range(50): + adapter._create_agent(session_id=str(_uuid.uuid4())) + + # Only the "*" process-wide fallback may exist — never one entry per + # ephemeral session_id. + assert list(adapter._last_resolved_model.keys()) == ["*"] + + def test_create_agent_wraps_runtime_credential_failure_as_provider_auth_error(self, monkeypatch): + """_create_agent() must convert a RuntimeError from + gateway.run._resolve_runtime_agent_kwargs() — the sole raiser of + RuntimeError(format_runtime_provider_error(...)) in this call graph — + into _ProviderAuthResolutionError right at the call site, independent + of any caller's exception handling.""" + from gateway.platforms.api_server import _ProviderAuthResolutionError + + monkeypatch.setattr( + "gateway.run._resolve_runtime_agent_kwargs", + lambda: (_ for _ in ()).throw(RuntimeError("No credentials found for provider 'nous'")), + ) + + adapter = APIServerAdapter(PlatformConfig(enabled=True)) + monkeypatch.setattr(adapter, "_ensure_session_db", lambda: None) + + with pytest.raises(_ProviderAuthResolutionError, match="No credentials found for provider 'nous'"): + adapter._create_agent(session_id="api-session") + + def test_create_agent_session_model_pins_ahead_of_request(self, monkeypatch): + """Session-persisted model beats per-request body values but yields + to an explicit session /model override.""" + captured = {} + + class FakeAgent: + def __init__(self, **kwargs): + captured.update(kwargs) + + _patch_create_agent_runtime(monkeypatch, captured, FakeAgent) + + adapter = APIServerAdapter(PlatformConfig(enabled=True)) + 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="session-row/model", + requested_model="request/model", + ) + + assert captured["model"] == "session-row/model" diff --git a/tests/gateway/test_api_server_runs.py b/tests/gateway/test_api_server_runs.py index f98d38e94c23..5ca7f221fa37 100644 --- a/tests/gateway/test_api_server_runs.py +++ b/tests/gateway/test_api_server_runs.py @@ -948,3 +948,38 @@ class TestStopRun: body = await events_resp.text() # Stream should have received run.failed and closed assert "run.failed" in body or "stream closed" in body + + +class TestRunsProviderAuthFailure: + @pytest.mark.asyncio + async def test_status_reports_provider_auth_failure_distinctly(self, adapter): + """/v1/runs builds its own agent via _create_agent() and does not + route through _run_agent(), so the controlled "Provider + authentication failed" message added there does not cover this + endpoint. _handle_runs()'s own _ProviderAuthResolutionError branch + must give the same distinguished message instead of the generic + except-Exception "run failed" text.""" + from gateway.platforms.api_server import _ProviderAuthResolutionError + + app = _create_runs_app(adapter) + async with TestClient(TestServer(app)) as cli: + with patch.object(adapter, "_create_agent") as mock_create: + mock_create.side_effect = _ProviderAuthResolutionError( + "No credentials found for provider 'nous'" + ) + + resp = await cli.post("/v1/runs", json={"input": "hello"}) + assert resp.status == 202 + data = await resp.json() + run_id = data["run_id"] + + for _ in range(40): + status_resp = await cli.get(f"/v1/runs/{run_id}") + status = await status_resp.json() + if status["status"] == "failed": + break + await asyncio.sleep(0.05) + + assert status["status"] == "failed" + assert status["error"] == "⚠️ Provider authentication failed: No credentials found for provider 'nous'" + assert status["last_event"] == "run.failed" diff --git a/tests/gateway/test_session_api.py b/tests/gateway/test_session_api.py index ac6314f2d28d..c2a98c6012fa 100644 --- a/tests/gateway/test_session_api.py +++ b/tests/gateway/test_session_api.py @@ -438,3 +438,182 @@ async def test_session_header_rejected_without_api_key(adapter, session_db): assert resp.status == 403 data = await resp.json() assert "X-Hermes-Session-Key requires API key" in data["error"]["message"] + + +# --------------------------------------------------------------------------- +# Session-persisted model threading + provider-auth failure surfacing +# (salvaged from PR #57947 by @FvanW and PR #59941 by @kaishi00) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_session_chat_threads_session_model_to_run_agent(auth_adapter, session_db): + """POST /api/sessions persists a per-session model, but the chat handler + previously fetched the session record and threw it away — the session's + chosen model silently had no effect on any chat turn.""" + session_id = session_db.create_session("model-pinned-session", "api_server", model="claude-sonnet-4-6") + + mock_run = AsyncMock(return_value=({"final_response": "ok", "session_id": session_id}, {"total_tokens": 1})) + 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": "hi"}, + headers={"Authorization": "Bearer sk-test"}, + ) + assert resp.status == 200 + + mock_run.assert_awaited_once() + _, kwargs = mock_run.call_args + assert kwargs["session_model"] == "claude-sonnet-4-6" + + +@pytest.mark.asyncio +async def test_session_chat_stream_threads_session_model_to_run_agent(adapter, session_db): + """Streaming twin of the session-model threading test above.""" + session_id = session_db.create_session("model-pinned-stream-session", "api_server", model="gpt-5.5") + + mock_run = AsyncMock(return_value=({"final_response": "ok", "session_id": session_id}, {"total_tokens": 1})) + 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/stream", + json={"message": "hi"}, + ) + assert resp.status == 200 + await resp.read() + + mock_run.assert_awaited_once() + _, kwargs = mock_run.call_args + assert kwargs["session_model"] == "gpt-5.5" + + +@pytest.mark.asyncio +async def test_session_chat_resolves_stored_model_route_alias(session_db, monkeypatch): + """A session-persisted model that matches a model_routes alias must go + through the route path (so route provider/credentials apply) and NOT be + passed as a raw session_model (idea from PR #59941 by @kaishi00).""" + adapter = APIServerAdapter( + PlatformConfig( + enabled=True, + extra={"model_routes": {"alias": {"model": "route/model", "provider": "openrouter"}}}, + ) + ) + adapter._session_db = session_db + session_id = session_db.create_session("route-pinned-session", "api_server", model="alias") + + mock_run = AsyncMock(return_value=({"final_response": "ok", "session_id": session_id}, {"total_tokens": 1})) + 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": "hi"}, + ) + assert resp.status == 200 + + _, kwargs = mock_run.call_args + assert kwargs["route"] == {"model": "route/model", "provider": "openrouter"} + assert kwargs["session_model"] is None + + +@pytest.mark.asyncio +async def test_run_agent_returns_controlled_response_on_provider_auth_failure(adapter, monkeypatch): + """_resolve_runtime_agent_kwargs() (inside _create_agent()) raises + RuntimeError on provider auth/credential failure. Previously this + propagated unhandled out of _run_agent(): /v1/chat/completions caught it + as a generic 500, and /api/sessions/{id}/chat didn't catch it at all + (raw aiohttp 500, no JSON body). Must now return run.py's controlled + response shape instead of raising. Exercises the REAL boundary + (gateway.run._resolve_runtime_agent_kwargs, the sole raiser).""" + monkeypatch.setattr( + "gateway.run._resolve_runtime_agent_kwargs", + lambda: (_ for _ in ()).throw( + RuntimeError("No credentials found for provider 'nous' — run `hermes auth add nous`") + ), + ) + + result, usage = await adapter._run_agent( + user_message="hello", + conversation_history=[], + session_id="request-session", + ) + + assert result == { + "final_response": "⚠️ Provider authentication failed: No credentials found for provider 'nous' — run `hermes auth add nous`", + "messages": [], + "api_calls": 0, + "tools": [], + } + assert usage == {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0} + + +@pytest.mark.asyncio +async def test_run_agent_does_not_swallow_unrelated_exceptions(adapter, monkeypatch): + """The _ProviderAuthResolutionError catch must stay narrow — a TypeError + elsewhere in _create_agent()/run_conversation() must still propagate.""" + def fake_create_agent(**kwargs): + raise TypeError("unrelated bug: unexpected keyword argument") + + monkeypatch.setattr(adapter, "_create_agent", fake_create_agent) + + with pytest.raises(TypeError, match="unrelated bug"): + await adapter._run_agent( + user_message="hello", + conversation_history=[], + session_id="request-session", + ) + + +@pytest.mark.asyncio +async def test_run_agent_does_not_swallow_unrelated_runtime_error_from_run_conversation(adapter, monkeypatch): + """agent.run_conversation() can legitimately raise a RuntimeError + unrelated to provider auth (e.g. run_agent.py's "Failed to recreate + closed OpenAI client"). A bare `except RuntimeError` around the whole + _create_agent()+run_conversation() span would mislabel it as + "Provider authentication failed". Only _ProviderAuthResolutionError — + raised exclusively inside _create_agent() at the + _resolve_runtime_agent_kwargs() call site — may trigger the controlled + response; this unrelated RuntimeError must propagate unhandled.""" + class _FakeAgent: + def run_conversation(self, **kwargs): + raise RuntimeError("Failed to recreate closed OpenAI client") + + monkeypatch.setattr(adapter, "_create_agent", lambda **kwargs: _FakeAgent()) + + with pytest.raises(RuntimeError, match="Failed to recreate closed OpenAI client"): + await adapter._run_agent( + user_message="hello", + conversation_history=[], + session_id="request-session", + ) + + +@pytest.mark.asyncio +async def test_session_chat_surfaces_controlled_response_on_provider_auth_failure(auth_adapter, session_db, monkeypatch): + """End-to-end: POST /api/sessions/{id}/chat previously had zero wrapping + around _run_agent() — an unhandled RuntimeError produced a raw aiohttp + 500 with no JSON body. Must now return 200 with the controlled error + message as the assistant content. Exercises the real + gateway.run._resolve_runtime_agent_kwargs() boundary, not a mocked + _create_agent().""" + session_id = session_db.create_session("auth-fail-session", "api_server") + + monkeypatch.setattr( + "gateway.run._resolve_runtime_agent_kwargs", + lambda: (_ for _ in ()).throw(RuntimeError("Auth failed: token expired")), + ) + + app = _create_session_app(auth_adapter) + async with TestClient(TestServer(app)) as cli: + resp = await cli.post( + f"/api/sessions/{session_id}/chat", + json={"message": "hi"}, + headers={"Authorization": "Bearer sk-test"}, + ) + assert resp.status == 200 + payload = await resp.json() + + assert payload["message"]["content"] == "⚠️ Provider authentication failed: Auth failed: token expired"