diff --git a/gateway/platforms/api_server.py b/gateway/platforms/api_server.py index fee6782a27..c00d9c5c9e 100644 --- a/gateway/platforms/api_server.py +++ b/gateway/platforms/api_server.py @@ -1193,7 +1193,11 @@ class APIServerAdapter(OpenAICompatRoutesMixin, BasePlatformAdapter): self._app: Optional["web.Application"] = None self._runner: Optional["web.AppRunner"] = None self._site: Optional["web.TCPSite"] = None - self._response_store = ResponseStore() + from hermes_constants import get_hermes_home + self._response_store = ResponseStore() # this home's; a /p// route gets its own + self._response_store_home = str(get_hermes_home()) + self._response_stores: Dict[str, ResponseStore] = {} + self._response_store_lock = threading.Lock() _api_runs._initialize_run_state(self, store_factory=RunIdempotencyStore) self._session_db: Optional[Any] = None # explicit override (tests/manual wiring) self._session_dbs: Dict[str, Any] = {} # per-profile-home SessionDB cache @@ -1705,6 +1709,21 @@ class APIServerAdapter(OpenAICompatRoutesMixin, BasePlatformAdapter): return None, _invalid_request("Session key too long") return raw, None + # -- Responses state ---------------------------------------------------------------- + + def _current_response_store(self) -> "ResponseStore": + """Responses state of the routed profile's home. Conversation names are client-chosen, so one + shared store let any profile's key read, chain onto and overwrite another's (#84253).""" + from hermes_constants import get_hermes_home + home = get_hermes_home() + if str(home) == self._response_store_home: + return self._response_store + with self._response_store_lock: + store = self._response_stores.get(str(home)) + if store is None: + store = self._response_stores[str(home)] = ResponseStore(db_path=str(home / "response_store.db")) + return store + # -- Session DB ------------------------------------------------------------------- def _open_and_cache_session_db(self, home) -> Optional[Any]: @@ -4289,9 +4308,13 @@ class APIServerAdapter(OpenAICompatRoutesMixin, BasePlatformAdapter): files, #37011). """ self._mark_disconnected() - if self._response_store is not None: + # getattr: disconnect() tolerates bare __new__ fixtures (pinned in test_api_server_run_idempotency). + routed = getattr(self, "_response_stores", {}) + stores = [s for s in (getattr(self, "_response_store", None), *list(routed.values())) if s is not None] + routed.clear() + for store in stores: try: - self._response_store.close() + store.close() except Exception: logger.debug("Failed to close response store for %s", self.name, exc_info=True) _api_runs._close_run_state(self) diff --git a/gateway/platforms/api_server_openai_routes.py b/gateway/platforms/api_server_openai_routes.py index 0a261a2fd4..c64d1b7134 100644 --- a/gateway/platforms/api_server_openai_routes.py +++ b/gateway/platforms/api_server_openai_routes.py @@ -206,6 +206,8 @@ class _ResponsesStream: self.model, self.created_at, self.conversation_history = model, created_at, conversation_history self.user_message, self.instructions = user_message, instructions self.conversation, self.store, self.session_id = conversation, store, session_id + # Resolved in the request's profile scope: a snapshot written after it (disconnect) must not follow another. + self.response_store = adapter._current_response_store() self.final_text_parts: List[str] = [] self.pending_tool_calls: List[Dict[str, Any]] = [] # open function_call items, in order self.emitted_items: List[Dict[str, Any]] = [] # output items so far (terminal payload) @@ -250,13 +252,13 @@ class _ResponsesStream: def persist_snapshot(self, response_env: Dict[str, Any], *, history=None, session_id=None): if not self.store: return - self.adapter._response_store.put(self.response_id, { + self.response_store.put(self.response_id, { "response": response_env, "conversation_history": self._history_with_user() if history is None else history, "instructions": self.instructions, "session_id": session_id or self.session_id}) if self.conversation: - self.adapter._response_store.set_conversation(self.conversation, self.response_id) + self.response_store.set_conversation(self.conversation, self.response_id) def persist_incomplete_if_needed(self) -> None: """Persist an ``incomplete`` snapshot when no terminal one was written (disconnect / @@ -973,7 +975,7 @@ class OpenAICompatRoutesMixin: return _error_response("Cannot use both 'conversation' and 'previous_response_id'", 400) if conversation: # A conversation name resolves to its latest response_id (unknown = new conversation). - previous_response_id = self._response_store.get_conversation(conversation) + previous_response_id = self._current_response_store().get_conversation(conversation) input_messages: List[Dict[str, Any]] = [] if isinstance(raw_input, str): @@ -1013,7 +1015,7 @@ class OpenAICompatRoutesMixin: logger.debug("Both conversation_history and previous_response_id provided; using conversation_history") stored_session_id = None if not conversation_history and previous_response_id: - stored = self._response_store.get(previous_response_id) + stored = self._current_response_store().get(previous_response_id) if stored is None: return _error_response(f"Previous response not found: {previous_response_id}", 404) conversation_history = list(stored.get("conversation_history", [])) @@ -1113,11 +1115,12 @@ class OpenAICompatRoutesMixin: "output": self._extract_output_items(result, start_index=output_start_index), "usage": _responses_usage_payload(usage)} if store: - self._response_store.put(response_id, { + response_store = self._current_response_store() + response_store.put(response_id, { "response": response_data, "conversation_history": full_history, "instructions": instructions, "session_id": _effective_session_id}) if conversation: - self._response_store.set_conversation(conversation, response_id) + response_store.set_conversation(conversation, response_id) response_headers = {"X-Hermes-Session-Id": _effective_session_id} if gateway_session_key: response_headers["X-Hermes-Session-Key"] = gateway_session_key @@ -1130,7 +1133,7 @@ class OpenAICompatRoutesMixin: if auth_err: return auth_err response_id = request.match_info["response_id"] - stored = self._response_store.get(response_id) + stored = self._current_response_store().get(response_id) if stored is None: return _error_response(f"Response not found: {response_id}", 404) return web.json_response(stored["response"]) @@ -1142,7 +1145,7 @@ class OpenAICompatRoutesMixin: if auth_err: return auth_err response_id = request.match_info["response_id"] - if not self._response_store.delete(response_id): + if not self._current_response_store().delete(response_id): return _error_response(f"Response not found: {response_id}", 404) return web.json_response({"id": response_id, "object": "response", "deleted": True}) diff --git a/gateway/platforms/api_server_runs.py b/gateway/platforms/api_server_runs.py index 70a9479731..fc389b3316 100644 --- a/gateway/platforms/api_server_runs.py +++ b/gateway/platforms/api_server_runs.py @@ -393,7 +393,7 @@ def _resolve_conversation_history( logger.debug("Both conversation_history and previous_response_id provided; using conversation_history") stored_session_id = None if not conversation_history and previous_response_id: - stored = self._response_store.get(previous_response_id) + stored = self._current_response_store().get(previous_response_id) if stored: conversation_history = list(stored.get("conversation_history", [])) stored_session_id = stored.get("session_id") diff --git a/tests/gateway/test_api_server_response_store_profile_scope.py b/tests/gateway/test_api_server_response_store_profile_scope.py new file mode 100644 index 0000000000..0a0298d6dd --- /dev/null +++ b/tests/gateway/test_api_server_response_store_profile_scope.py @@ -0,0 +1,96 @@ +"""/v1/responses state belongs to the routed profile (#84253). + +A multiplexed API server mirrors every route under /p// with that profile's key, but kept ONE +ResponseStore at the launch home. Conversation names are client-chosen ("main", "my-project"), so another +profile's key could chain onto a profile's conversation (reading its transcript, instructions and session), +become its tip, and GET or DELETE its responses. Real profile homes, prefix middleware and route table; only +the agent run is faked. +""" +from __future__ import annotations + +import pytest + +from agent import secret_scope as ss + +ALICE, BOB, DEFAULT = "a" * 40, "b" * 40, "d" * 40 + + +@pytest.fixture(autouse=True) +def _reset_multiplex(): + ss.set_multiplex_active(False) + yield + ss.set_multiplex_active(False) + + +async def _response_id(resp) -> str: + """The id from a JSON body, or from the SSE stream's response.created event.""" + import json + if resp.content_type == "application/json": + return (await resp.json())["id"] + for line in (await resp.text()).splitlines(): + if line.startswith("data:") and '"response"' in line: + return json.loads(line[5:])["response"]["id"] + raise AssertionError("no response id in the stream") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stream", [False, True], ids=["json", "sse"]) +async def test_a_profile_key_never_reaches_another_profiles_responses(tmp_path, monkeypatch, stream): + from aiohttp import web + from aiohttp.test_utils import TestClient, TestServer + from gateway.config import GatewayConfig, PlatformConfig + from gateway.platforms import api_server as api + from gateway.platforms.api_server import APIServerAdapter + + home = tmp_path / ".hermes" + for name, key in (("alice", ALICE), ("bob", BOB)): + (home / "profiles" / name).mkdir(parents=True) + (home / "profiles" / name / ".env").write_text(f"API_SERVER_KEY={key}\n", encoding="utf-8") + monkeypatch.setenv("HERMES_HOME", str(home)) + + adapter = APIServerAdapter(PlatformConfig(enabled=True, extra={"key": DEFAULT})) + adapter.gateway_runner = type("_Runner", (), {"config": GatewayConfig(multiplex_profiles=True)})() + ss.set_multiplex_active(True) + runs = [] + + async def fake_run_agent(**kw): + who = api._api_request_profile.get() + runs.append({"profile": who, "history": kw.get("conversation_history"), + "instructions": kw.get("ephemeral_system_prompt"), "session_id": kw.get("session_id")}) + return ({"final_response": f"ok from {who}", "messages": [], "api_calls": 1, + "session_id": kw.get("session_id")}, {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}) + + adapter._run_agent = fake_run_agent + app = web.Application(middlewares=[adapter._make_profile_prefix_middleware()]) + for method, path, handler in adapter._http_route_table(): + app.router.add_route(method, path, handler) + app.router.add_route(method, f"/p/{{profile}}{path}", handler) + app["api_server_adapter"] = adapter + + def auth(key): + return {"Authorization": f"Bearer {key}"} + + try: + async with TestClient(TestServer(app)) as client: + first = await client.post("/p/alice/v1/responses", headers=auth(ALICE), json={ + "input": "SECRET-ALICE", "conversation": "main", "instructions": "ALICE PRIVATE PROMPT", + "stream": stream}) + alice_id = await _response_id(first) + alice_session = runs[-1]["session_id"] + + bob = await client.post("/p/bob/v1/responses", headers=auth(BOB), json={ + "input": "what did I say?", "conversation": "main", "stream": stream}) + assert bob.status == 200 + await bob.read() + assert runs[-1]["history"] == [] and runs[-1]["instructions"] is None + assert runs[-1]["session_id"] != alice_session + assert (await client.get(f"/p/bob/v1/responses/{alice_id}", headers=auth(BOB))).status == 404 + assert (await client.delete(f"/p/bob/v1/responses/{alice_id}", headers=auth(BOB))).status == 404 + + await (await client.post("/p/alice/v1/responses", headers=auth(ALICE), json={ + "input": "next", "conversation": "main", "stream": stream})).read() + assert "what did I say?" not in str(runs[-1]["history"]) + assert "SECRET-ALICE" in str(runs[-1]["history"]) + assert (await client.get(f"/p/alice/v1/responses/{alice_id}", headers=auth(ALICE))).status == 200 + finally: + await adapter.disconnect()