diff --git a/docs/web-transport.md b/docs/web-transport.md index ee87930..922ac8f 100644 --- a/docs/web-transport.md +++ b/docs/web-transport.md @@ -75,6 +75,10 @@ Replay is consumed as it arrives, so histories larger than the SSE buffer do not wait for a session stream to open. WebSocket uses its existing bidirectional connection for both replay and subsequent messages. +`resume_session()` (when the agent advertises `sessionCapabilities.resume`) +follows the same HTTP routing without replaying history: the response uses the +connection SSE stream, and the client then opens the session SSE stream. + A failed load returns its JSON-RPC error on the connection stream and can be retried. The server removes streams provisioned only for failed loads, while preserving established sessions and overlapping loads. It does not change the @@ -153,7 +157,7 @@ HTTP output follows these routing rules: | --- | --- | --- | | `initialize` response | POST body, via one Future | Establishes the connection before GET streams open | | Response containing a new `sessionId` | Connection SSE stream | The client needs the ID before it can open the session stream | -| `session/load` replay and response | Connection SSE stream | Replay precedes the response; the client gets the session ID from the original request | +| `session/load` replay and response, `session/resume` response | Connection SSE stream | Replay precedes the response; the client gets the session ID from the original request | | Other messages | Session SSE stream when known, otherwise connection stream | Responses use their request's recorded session; requests/notifications carry `sessionId` | `OutboundStream` retains a bounded buffer, backpressure, and close handling. diff --git a/src/acp/http/client.py b/src/acp/http/client.py index cee5949..8baf6bd 100644 --- a/src/acp/http/client.py +++ b/src/acp/http/client.py @@ -28,7 +28,7 @@ CONNECTION_ID_HEADER, CONTENT_TYPE_JSON, CONTENT_TYPE_SSE, - LOAD_SESSION_METHOD, + SESSION_ATTACH_METHODS, SESSION_ID_HEADER, is_initialize_request, is_response_message, @@ -81,6 +81,7 @@ def __init__( self._inbox: asyncio.Queue[Any] = asyncio.Queue() self._stream_tasks: set[asyncio.Task[None]] = set() self._session_streams: set[str] = set() + # Pending session/load and session/resume requests, keyed by request id. self._pending_loads: dict[str, str] = {} # -- Transport protocol ------------------------------------------------- @@ -93,7 +94,7 @@ async def send(self, message: dict[str, Any]) -> None: return key = message_id_key(message.get("id")) session_id = session_id_from_message(message) - if message.get("method") == LOAD_SESSION_METHOD and key is not None and session_id is not None: + if message.get("method") in SESSION_ATTACH_METHODS and key is not None and session_id is not None: self._pending_loads[key] = session_id try: await self._send_post(message) @@ -215,7 +216,7 @@ def _on_stream_closed(self, session_id: str | None) -> None: self._inbox.put_nowait(_EOF) def _handle_incoming(self, message: dict[str, Any]) -> None: - # Load responses may be empty or null; the session ID is in the request. + # Load and resume responses may omit sessionId; the ID is in the request. if is_response_message(message): key = message_id_key(message.get("id")) loaded = self._pending_loads.pop(key, None) if key is not None else None diff --git a/src/acp/http/protocol.py b/src/acp/http/protocol.py index e61476a..fe361ff 100644 --- a/src/acp/http/protocol.py +++ b/src/acp/http/protocol.py @@ -18,6 +18,8 @@ "CONTENT_TYPE_SSE", "INITIALIZE_METHOD", "LOAD_SESSION_METHOD", + "RESUME_SESSION_METHOD", + "SESSION_ATTACH_METHODS", "SESSION_ID_HEADER", "is_initialize_request", "is_response_message", @@ -41,6 +43,11 @@ INITIALIZE_METHOD = AGENT_METHODS["initialize"] LOAD_SESSION_METHOD = AGENT_METHODS["session_load"] +RESUME_SESSION_METHOD = AGENT_METHODS["session_resume"] + +# Methods that attach the ``sessionId`` given in their request. Their responses +# need not echo the ID, so both peers take it from the request. +SESSION_ATTACH_METHODS = frozenset({LOAD_SESSION_METHOD, RESUME_SESSION_METHOD}) # Agent methods that operate on an *already-established* session and therefore # require the ``Acp-Session-Id`` header on POST + session-scoped routing of their diff --git a/src/acp/http/server.py b/src/acp/http/server.py index 1138ea4..0cd2cf2 100644 --- a/src/acp/http/server.py +++ b/src/acp/http/server.py @@ -33,7 +33,7 @@ from ..agent.connection import AgentSideConnection from .protocol import ( CONNECTION_ID_HEADER, - LOAD_SESSION_METHOD, + SESSION_ATTACH_METHODS, is_initialize_request, is_response_message, message_id_key, @@ -139,6 +139,7 @@ def __init__(self, initialize_id: Any) -> None: self.connection_stream = OutboundStream() self.session_streams: dict[str, OutboundStream] = {} self._pending_routes: dict[str, str] = {} + # Pending session/load and session/resume requests, keyed by request id. self._pending_loads: dict[str, str] = {} self._provisional_sessions: set[str] = set() @@ -155,7 +156,7 @@ async def send(self, message: dict[str, Any]) -> None: self.initialize_response.set_result(message) return session_id = self._route_response(message, key) - # Replay and load responses share the connection stream, including on + # Replay and load/resume responses share the connection stream, including on # reload. This preserves replay order and never waits for a session GET. if session_id in self._pending_loads.values(): session_id = None @@ -191,7 +192,7 @@ async def deliver_to_agent(self, message: dict[str, Any]) -> None: session_id = session_id_from_params(message.get("params")) key = message_id_key(message["id"]) if session_id is not None and key is not None: - if message["method"] == LOAD_SESSION_METHOD: + if message["method"] in SESSION_ATTACH_METHODS: self._pending_loads[key] = session_id if session_id not in self.session_streams: # Allow clients to open a GET as soon as replay starts. diff --git a/src/acp/interfaces.py b/src/acp/interfaces.py index 63ddf88..073f516 100644 --- a/src/acp/interfaces.py +++ b/src/acp/interfaces.py @@ -445,7 +445,7 @@ async def fork_session( **kwargs: Any, ) -> ForkSessionResponse: ... - @param_model(ResumeSessionRequest, method=AGENT_METHODS["session_resume"], unstable=True) + @param_model(ResumeSessionRequest, method=AGENT_METHODS["session_resume"]) async def resume_session( self, session_id: str, @@ -455,9 +455,7 @@ async def resume_session( **kwargs: Any, ) -> ResumeSessionResponse: ... - @param_model( - CloseSessionRequest, method=AGENT_METHODS["session_close"], unstable=True, adapt_result=normalize_result - ) + @param_model(CloseSessionRequest, method=AGENT_METHODS["session_close"], adapt_result=normalize_result) async def close_session(self, session_id: str, **kwargs: Any) -> CloseSessionResponse | None: ... @param_model(CancelNotification, method=AGENT_METHODS["session_cancel"], kind="notification") diff --git a/tests/http/test_loopback.py b/tests/http/test_loopback.py index 4b25def..acbdbb9 100644 --- a/tests/http/test_loopback.py +++ b/tests/http/test_loopback.py @@ -23,6 +23,7 @@ NewSessionResponse, PromptResponse, RequestPermissionResponse, + ResumeSessionResponse, TextContentBlock, ) from acp.ws.client import create_websocket_stream @@ -213,3 +214,42 @@ async def test_failed_load_can_retry_and_preserves_existing_session(protocol: st assert result.stop_reason == "end_turn" finally: await conn.close() + + +class _ResumingAgent(_LoopbackAgent): + def __init__(self) -> None: + super().__init__() + self.ask_permission = True + self.fail_resume = False + + async def resume_session(self, cwd: str, session_id: str, **kwargs: Any) -> ResumeSessionResponse: + if self.fail_resume: + raise RequestError(-32000, "resume failed") + return ResumeSessionResponse() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("protocol", ["http", "ws"]) +async def test_resume_session_supports_prompt(protocol: str, serve_asgi) -> None: + agent = _ResumingAgent() + server = await serve_asgi(_make_app(agent)) + transport = ( + create_http_stream(server.http_url) if protocol == "http" else await create_websocket_stream(server.ws_url) + ) + client = _CapturingClient() + conn = connect_to_agent(client, transport) + try: + await conn.initialize(protocol_version=1) + # A failed resume can be retried; the successful one attaches the session. + agent.fail_resume = True + with pytest.raises(RequestError, match="resume failed"): + await asyncio.wait_for(conn.resume_session(cwd="/", session_id="saved-session"), timeout=5) + agent.fail_resume = False + resumed = await asyncio.wait_for(conn.resume_session(cwd="/", session_id="saved-session"), timeout=5) + assert resumed == ResumeSessionResponse() + result = await asyncio.wait_for(conn.prompt(session_id="saved-session", prompt=[]), timeout=5) + assert result.stop_reason == "end_turn" + assert client.permission_requested + assert client.updates[-1].content.text == "hello" + finally: + await conn.close() diff --git a/tests/test_unstable.py b/tests/test_unstable.py index 7f519fc..08f03c6 100644 --- a/tests/test_unstable.py +++ b/tests/test_unstable.py @@ -66,5 +66,17 @@ async def test_call_unstable_protocol_warning(connect): with pytest.warns(UserWarning) as record: with pytest.raises(RequestError): - await agent_conn.close_session(session_id="sess") + await agent_conn.fork_session(cwd="/workspace", session_id="sess") assert len(record) == 1 + + +@pytest.mark.parametrize("agent", [UnstableAgent()]) +@pytest.mark.asyncio +async def test_stable_session_lifecycle_does_not_require_unstable_protocol(connect): + _, agent_conn = connect(use_unstable_protocol=False) + + resp = await agent_conn.resume_session(cwd="/workspace", session_id="sess") + assert isinstance(resp, ResumeSessionResponse) + + resp = await agent_conn.close_session(session_id="sess") + assert isinstance(resp, CloseSessionResponse)