Skip to content

Commit 4fddab3

Browse files
committed
fix: support loading sessions over web transports
1 parent 5ae0734 commit 4fddab3

7 files changed

Lines changed: 251 additions & 16 deletions

File tree

docs/web-transport.md

Lines changed: 27 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -55,6 +55,31 @@ header, then opens the connection-scoped SSE stream. When a new `sessionId`
5555
appears it opens that session-scoped stream too. A single SSE attempt is made per
5656
stream; reconnect/retry is the caller's responsibility (v1 of the RFD).
5757

58+
### Loading an existing session
59+
60+
Both HTTP and WebSocket support `load_session()` when the agent advertises
61+
`loadSession` and implements session persistence:
62+
63+
```python
64+
init = await conn.initialize(protocol_version=1)
65+
if init.agent_capabilities.load_session:
66+
await conn.load_session(session_id="saved-session-id", cwd="/workspace", mcp_servers=[])
67+
await conn.prompt(session_id="saved-session-id", prompt=[...])
68+
```
69+
70+
For HTTP, history replay and the load response use the connection SSE stream.
71+
The client correlates the response with the session ID from the load request,
72+
then opens the session SSE stream for further prompts and agent callbacks. This
73+
also works with an empty history: a load response need not contain `sessionId`.
74+
Replay is consumed as it arrives, so histories larger than the SSE buffer do not
75+
wait for a session stream to open. WebSocket uses its existing bidirectional
76+
connection for both replay and subsequent messages.
77+
78+
A failed load returns its JSON-RPC error on the connection stream and can be
79+
retried. The server removes streams provisioned only for failed loads, while
80+
preserving established sessions and overlapping loads. It does not change the
81+
agent's load response or automatically enable the agent's `loadSession` capability.
82+
5883
## Server
5984

6085
The server uses Starlette for HTTP requests, responses, routing, streaming,
@@ -122,12 +147,13 @@ has one incoming queue and one SSE buffer per stream. The incoming queue lets
122147
POST return `202` while the agent handles the request. Output goes directly to
123148
the relevant SSE buffer; there is no intermediate transport pair or pump task.
124149

125-
HTTP output needs three routing rules:
150+
HTTP output follows these routing rules:
126151

127152
| Message | Destination | Why |
128153
| --- | --- | --- |
129154
| `initialize` response | POST body, via one Future | Establishes the connection before GET streams open |
130155
| Response containing a new `sessionId` | Connection SSE stream | The client needs the ID before it can open the session stream |
156+
| `session/load` replay and response | Connection SSE stream | Replay precedes the response; the client gets the session ID from the original request |
131157
| Other messages | Session SSE stream when known, otherwise connection stream | Responses use their request's recorded session; requests/notifications carry `sessionId` |
132158

133159
`OutboundStream` retains a bounded buffer, backpressure, and close handling.

src/acp/http/client.py

Lines changed: 26 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -28,8 +28,11 @@
2828
CONNECTION_ID_HEADER,
2929
CONTENT_TYPE_JSON,
3030
CONTENT_TYPE_SSE,
31+
LOAD_SESSION_METHOD,
3132
SESSION_ID_HEADER,
3233
is_initialize_request,
34+
is_response_message,
35+
message_id_key,
3336
method_requires_session_header,
3437
session_id_from_message,
3538
)
@@ -78,6 +81,7 @@ def __init__(
7881
self._inbox: asyncio.Queue[Any] = asyncio.Queue()
7982
self._stream_tasks: set[asyncio.Task[None]] = set()
8083
self._session_streams: set[str] = set()
84+
self._pending_loads: dict[str, str] = {}
8185

8286
# -- Transport protocol -------------------------------------------------
8387

@@ -87,7 +91,16 @@ async def send(self, message: dict[str, Any]) -> None:
8791
if is_initialize_request(message):
8892
await self._send_initialize(message)
8993
return
90-
await self._send_post(message)
94+
key = message_id_key(message.get("id"))
95+
session_id = session_id_from_message(message)
96+
if message.get("method") == LOAD_SESSION_METHOD and key is not None and session_id is not None:
97+
self._pending_loads[key] = session_id
98+
try:
99+
await self._send_post(message)
100+
except BaseException:
101+
if key is not None:
102+
self._pending_loads.pop(key, None)
103+
raise
91104

92105
async def receive(self) -> dict[str, Any] | None:
93106
item = await self._inbox.get()
@@ -99,6 +112,7 @@ async def close(self) -> None:
99112
if self._closed:
100113
return
101114
self._closed = True
115+
self._pending_loads.clear()
102116
for task in list(self._stream_tasks):
103117
task.cancel()
104118
for task in list(self._stream_tasks):
@@ -148,7 +162,7 @@ async def _send_post(self, message: dict[str, Any]) -> None:
148162
# Some servers may answer initialize-like 200 bodies; for 200 with a body enqueue it.
149163
if response.status_code == 200 and response.content:
150164
with contextlib.suppress(Exception):
151-
self._inbox.put_nowait(response.json())
165+
self._handle_incoming(response.json())
152166

153167
def _open_stream(self, *, session_id: str | None) -> None:
154168
if self._closed:
@@ -197,13 +211,20 @@ def _on_stream_closed(self, session_id: str | None) -> None:
197211
self._session_streams.discard(session_id)
198212
return
199213
if not self._closed:
214+
self._pending_loads.clear()
200215
self._inbox.put_nowait(_EOF)
201216

202217
def _handle_incoming(self, message: dict[str, Any]) -> None:
203-
# Open a session-scoped stream when any message carries a new sessionId
204-
# (e.g. a session/new or session/load result on the connection stream).
218+
# Load responses may be empty or null; the session ID is in the request.
219+
if is_response_message(message):
220+
key = message_id_key(message.get("id"))
221+
loaded = self._pending_loads.pop(key, None) if key is not None else None
222+
if loaded is not None and "result" in message:
223+
self._open_stream(session_id=loaded)
205224
session_id = session_id_from_message(message)
206-
if session_id is not None and session_id not in self._session_streams:
225+
# Replay stays on the connection stream until load succeeds. Avoid
226+
# opening a session GET that a failed load would immediately tear down.
227+
if session_id is not None and session_id not in self._pending_loads.values():
207228
self._open_stream(session_id=session_id)
208229
self._inbox.put_nowait(message)
209230

src/acp/http/protocol.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
"CONTENT_TYPE_JSON",
1818
"CONTENT_TYPE_SSE",
1919
"INITIALIZE_METHOD",
20+
"LOAD_SESSION_METHOD",
2021
"SESSION_ID_HEADER",
2122
"is_initialize_request",
2223
"is_response_message",
@@ -39,6 +40,7 @@
3940
ACP_ENDPOINT_PATH = "/acp"
4041

4142
INITIALIZE_METHOD = AGENT_METHODS["initialize"]
43+
LOAD_SESSION_METHOD = AGENT_METHODS["session_load"]
4244

4345
# Agent methods that operate on an *already-established* session and therefore
4446
# require the ``Acp-Session-Id`` header on POST + session-scoped routing of their

src/acp/http/server.py

Lines changed: 36 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -33,6 +33,7 @@
3333
from ..agent.connection import AgentSideConnection
3434
from .protocol import (
3535
CONNECTION_ID_HEADER,
36+
LOAD_SESSION_METHOD,
3637
is_initialize_request,
3738
is_response_message,
3839
message_id_key,
@@ -138,6 +139,8 @@ def __init__(self, initialize_id: Any) -> None:
138139
self.connection_stream = OutboundStream()
139140
self.session_streams: dict[str, OutboundStream] = {}
140141
self._pending_routes: dict[str, str] = {}
142+
self._pending_loads: dict[str, str] = {}
143+
self._provisional_sessions: set[str] = set()
141144

142145
async def receive(self) -> dict[str, Any] | None:
143146
return await self._incoming.get()
@@ -151,27 +154,51 @@ async def send(self, message: dict[str, Any]) -> None:
151154
if key == self._initialize_id and not self.initialize_response.done():
152155
self.initialize_response.set_result(message)
153156
return
154-
session_id = self._pending_routes.pop(key, None) if key is not None else None
155-
established = session_id_from_result(message.get("result"))
156-
if established is not None:
157-
self.session_streams.setdefault(established, OutboundStream())
158-
# The client must learn the session ID before opening its stream.
159-
session_id = None
157+
session_id = self._route_response(message, key)
158+
# Replay and load responses share the connection stream, including on
159+
# reload. This preserves replay order and never waits for a session GET.
160+
if session_id in self._pending_loads.values():
161+
session_id = None
160162
stream = (
161163
self.session_streams.get(session_id, self.connection_stream)
162164
if session_id is not None
163165
else self.connection_stream
164166
)
165167
await stream.push(message)
166168

169+
def _route_response(self, message: dict[str, Any], key: str | None) -> str | None:
170+
loaded = self._pending_loads.pop(key, None) if key is not None else None
171+
if loaded is not None:
172+
if "result" in message:
173+
self._provisional_sessions.discard(loaded)
174+
elif loaded in self._provisional_sessions and loaded not in self._pending_loads.values():
175+
self._provisional_sessions.remove(loaded)
176+
self.session_streams.pop(loaded).close()
177+
return None
178+
session_id = self._pending_routes.pop(key, None) if key is not None else None
179+
established = session_id_from_result(message.get("result"))
180+
if established is not None:
181+
self.session_streams.setdefault(established, OutboundStream())
182+
self._provisional_sessions.discard(established)
183+
# The client must learn the session ID before opening its stream.
184+
return None
185+
return session_id
186+
167187
async def deliver_to_agent(self, message: dict[str, Any]) -> None:
168188
if self._closed:
169189
raise ConnectionError("Transport closed")
170190
if "id" in message and "method" in message:
171191
session_id = session_id_from_params(message.get("params"))
172192
key = message_id_key(message["id"])
173193
if session_id is not None and key is not None:
174-
self._pending_routes[key] = session_id
194+
if message["method"] == LOAD_SESSION_METHOD:
195+
self._pending_loads[key] = session_id
196+
if session_id not in self.session_streams:
197+
# Allow clients to open a GET as soon as replay starts.
198+
self.session_streams[session_id] = OutboundStream()
199+
self._provisional_sessions.add(session_id)
200+
else:
201+
self._pending_routes[key] = session_id
175202
self._incoming.put_nowait(dict(message))
176203

177204
async def close(self) -> None:
@@ -181,6 +208,8 @@ async def close(self) -> None:
181208
self._incoming.put_nowait(None)
182209
self.initialize_response.cancel()
183210
self._pending_routes.clear()
211+
self._pending_loads.clear()
212+
self._provisional_sessions.clear()
184213
self.connection_stream.close()
185214
for stream in self.session_streams.values():
186215
stream.close()

tests/http/test_http_client.py

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -211,3 +211,35 @@ def handler(request: httpx.Request) -> httpx.Response:
211211
finally:
212212
await transport.close()
213213
await client.aclose()
214+
215+
216+
@pytest.mark.asyncio
217+
@pytest.mark.parametrize("result", [{}, None])
218+
async def test_load_opens_session_stream_from_request_id(result: Any) -> None:
219+
server = FakeServer()
220+
transport, client = _make_transport(server)
221+
try:
222+
await transport.send({"jsonrpc": "2.0", "id": 0, "method": "initialize", "params": {}})
223+
await asyncio.wait_for(transport.receive(), timeout=1)
224+
await transport.send({
225+
"jsonrpc": "2.0",
226+
"id": "load-1",
227+
"method": "session/load",
228+
"params": {"sessionId": "saved", "cwd": "/", "mcpServers": []},
229+
})
230+
replay = {"jsonrpc": "2.0", "method": "session/update", "params": {"sessionId": "saved"}}
231+
server.push_conn(replay)
232+
assert await asyncio.wait_for(transport.receive(), timeout=1) == replay
233+
await asyncio.sleep(0)
234+
assert "saved" not in server.session_streams
235+
236+
# A standard load response carries no sessionId, including null results.
237+
response = {"jsonrpc": "2.0", "id": "load-1", "result": result}
238+
server.push_conn(response)
239+
assert await asyncio.wait_for(transport.receive(), timeout=1) == response
240+
live = {"jsonrpc": "2.0", "id": 2, "result": {"stopReason": "end_turn"}}
241+
server.push_session("saved", live)
242+
assert await asyncio.wait_for(transport.receive(), timeout=1) == live
243+
finally:
244+
await transport.close()
245+
await client.aclose()

tests/http/test_http_server.py

Lines changed: 42 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@
1010

1111
from acp.exceptions import RequestError
1212
from acp.http.protocol import CONNECTION_ID_HEADER
13-
from acp.http.server import AcpServer
13+
from acp.http.server import AcpServer, _HttpTransport
1414
from acp.schema import NewSessionResponse, PromptResponse
1515
from tests.conftest import TestAgent
1616

@@ -297,3 +297,44 @@ async def prompt(self, session_id: str, prompt: Any = None, **kwargs: Any) -> Pr
297297
finally:
298298
await connection_stream.aclose()
299299
await server.close()
300+
301+
302+
@pytest.mark.asyncio
303+
@pytest.mark.parametrize("first_succeeds", [False, True])
304+
@pytest.mark.parametrize("second_succeeds", [False, True])
305+
async def test_overlapping_loads_preserve_successful_session_streams(
306+
first_succeeds: bool, second_succeeds: bool
307+
) -> None:
308+
transport = _HttpTransport(0)
309+
connection_stream = transport.connection_stream.iterate()
310+
try:
311+
for request_id in (1, 2):
312+
await transport.deliver_to_agent({
313+
"jsonrpc": "2.0",
314+
"id": request_id,
315+
"method": "session/load",
316+
"params": {"sessionId": "saved", "cwd": "/", "mcpServers": []},
317+
})
318+
# A client may attach its session GET before load completes.
319+
session_stream = transport.session_streams["saved"]
320+
for request_id, succeeds in enumerate((first_succeeds, second_succeeds), start=1):
321+
response = {"jsonrpc": "2.0", "id": request_id}
322+
response.update({"result": {}} if succeeds else {"error": {"code": -32000, "message": "load failed"}})
323+
await transport.send(response)
324+
assert await asyncio.wait_for(anext(connection_stream), timeout=1) == response
325+
if request_id == 1:
326+
assert transport.session_streams["saved"] is session_stream
327+
328+
if first_succeeds or second_succeeds:
329+
assert transport.session_streams["saved"] is session_stream
330+
live = {"jsonrpc": "2.0", "method": "session/update", "params": {"sessionId": "saved"}}
331+
await transport.send(live)
332+
messages = session_stream.iterate()
333+
assert await asyncio.wait_for(anext(messages), timeout=1) == live
334+
await messages.aclose()
335+
else:
336+
assert "saved" not in transport.session_streams
337+
assert [message async for message in session_stream.iterate()] == []
338+
finally:
339+
await connection_stream.aclose()
340+
await transport.close()

0 commit comments

Comments
 (0)