From 7abd3364e36a1ced1c7d5ec97178fde13a094b85 Mon Sep 17 00:00:00 2001 From: Thomas Wood Date: Sat, 3 Oct 2026 11:59:46 -0400 Subject: [PATCH] feat(http): serve v2 over Streamable HTTP and WebSocket The Streamable HTTP + WebSocket transports predate the experimental v2 runtime and were hard-wired to the v1 AgentSideConnection and its method table. This completes the port so the transports are protocol-agnostic, matching the TypeScript reference (AcpServer takes an AgentConnector): - AcpServer / handle_websocket accept an AgentProtocolRouter, which negotiates v1 or v2 from the first initialize and attaches the right runtime. - AgentProtocolConnection gains a public listen() so stdio, HTTP and WebSocket drive it uniformly. - session/resume joins session/load as a replay method, so v2 resume routes replay over the connection stream exactly like v1 load. - v2 AgentSideConnection.attach now calls on_connect, matching __init__ and the v1 attach (agents like crow-cli v2 that learn their connection from on_connect would otherwise never receive it). Tests: v2 over HTTP, v2 over WebSocket, v2 resume replay, and a dual-stack test serving v1 and v2 clients on one port. Session-Id: enchanted-spectral-skunk-of-ampleness --- src/acp/experimental/negotiation.py | 3 + src/acp/experimental/v2/agent.py | 2 + src/acp/http/asgi.py | 6 +- src/acp/http/client.py | 4 +- src/acp/http/protocol.py | 12 +- src/acp/http/server.py | 19 ++- src/acp/ws/server.py | 10 +- tests/http/test_v2_loopback.py | 226 ++++++++++++++++++++++++++++ 8 files changed, 270 insertions(+), 12 deletions(-) create mode 100644 tests/http/test_v2_loopback.py diff --git a/src/acp/experimental/negotiation.py b/src/acp/experimental/negotiation.py index 8f84da3..48851f0 100644 --- a/src/acp/experimental/negotiation.py +++ b/src/acp/experimental/negotiation.py @@ -136,6 +136,9 @@ def __init__(self, connection: Connection) -> None: async def _listen(self) -> None: await self._connection.main_loop() + async def listen(self) -> None: + await self._listen() + async def close(self) -> None: await self._connection.close() diff --git a/src/acp/experimental/v2/agent.py b/src/acp/experimental/v2/agent.py index da96072..4eaf2f0 100644 --- a/src/acp/experimental/v2/agent.py +++ b/src/acp/experimental/v2/agent.py @@ -84,6 +84,8 @@ def attach( self._state = InitializationState() self._conn = connection agent = agent_factory(self) + if on_connect := getattr(agent, "on_connect", None): + on_connect(self) router = _AgentRouter(agent, self._state) return self, router diff --git a/src/acp/http/asgi.py b/src/acp/http/asgi.py index 84e59a3..5e8216c 100644 --- a/src/acp/http/asgi.py +++ b/src/acp/http/asgi.py @@ -10,6 +10,7 @@ from collections.abc import AsyncIterator from contextlib import asynccontextmanager from functools import partial +from typing import TYPE_CHECKING from starlette.applications import Starlette from starlette.requests import Request @@ -21,10 +22,13 @@ from .protocol import ACP_ENDPOINT_PATH, CONNECTION_ID_HEADER, CONTENT_TYPE_SSE, SESSION_ID_HEADER from .server import AcpServer, AgentFactory +if TYPE_CHECKING: + from ..experimental.negotiation import AgentProtocolRouter + __all__ = ["create_asgi_app"] -def create_asgi_app(agent_factory: AgentFactory, *, path: str = ACP_ENDPOINT_PATH) -> Starlette: +def create_asgi_app(agent_factory: AgentFactory | AgentProtocolRouter, *, path: str = ACP_ENDPOINT_PATH) -> Starlette: """Create a Starlette app with one agent instance per connection. The app handles POST/GET/DELETE and WebSocket at ``path`` (default: /acp). diff --git a/src/acp/http/client.py b/src/acp/http/client.py index cee5949..14ce434 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, + LOAD_SESSION_METHODS, SESSION_ID_HEADER, is_initialize_request, is_response_message, @@ -93,7 +93,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 LOAD_SESSION_METHODS and key is not None and session_id is not None: self._pending_loads[key] = session_id try: await self._send_post(message) diff --git a/src/acp/http/protocol.py b/src/acp/http/protocol.py index e61476a..24483a1 100644 --- a/src/acp/http/protocol.py +++ b/src/acp/http/protocol.py @@ -17,7 +17,7 @@ "CONTENT_TYPE_JSON", "CONTENT_TYPE_SSE", "INITIALIZE_METHOD", - "LOAD_SESSION_METHOD", + "LOAD_SESSION_METHODS", "SESSION_ID_HEADER", "is_initialize_request", "is_response_message", @@ -40,7 +40,15 @@ ACP_ENDPOINT_PATH = "/acp" INITIALIZE_METHOD = AGENT_METHODS["initialize"] -LOAD_SESSION_METHOD = AGENT_METHODS["session_load"] +# Replay methods whose responses and replayed history stay on the +# connection-scoped stream until the load completes: v1's ``session/load`` and +# v2's ``session/resume``. The session id is in the request, so these do NOT +# require the ``Acp-Session-Id`` header - a freshly attached client may not yet +# have a session-scoped stream open. +LOAD_SESSION_METHODS = frozenset({ + AGENT_METHODS["session_load"], + AGENT_METHODS["session_resume"], +}) # 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..b441d2b 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, + LOAD_SESSION_METHODS, is_initialize_request, is_response_message, message_id_key, @@ -45,6 +45,7 @@ if TYPE_CHECKING: from collections.abc import AsyncGenerator + from ..experimental.negotiation import AgentProtocolConnection, AgentProtocolRouter from ..interfaces import Agent __all__ = [ @@ -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 LOAD_SESSION_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. @@ -223,9 +224,17 @@ class AcpServer: ``AgentSideConnection`` to produce a per-connection ``Agent``. """ - def __init__(self, agent_factory: AgentFactory) -> None: + def __init__(self, agent_factory: AgentFactory | AgentProtocolRouter) -> None: self.agent_factory = agent_factory - self._connections: dict[str, tuple[AgentSideConnection, _HttpTransport]] = {} + self._connections: dict[str, tuple[AgentSideConnection | AgentProtocolConnection, _HttpTransport]] = {} + + def _connect(self, transport: _HttpTransport) -> AgentSideConnection | AgentProtocolConnection: + """Bind one agent (v1) or protocol router (v1/v2) to a transport.""" + from ..experimental.negotiation import AgentProtocolRouter + + if isinstance(self.agent_factory, AgentProtocolRouter): + return self.agent_factory.connect(transport) + return AgentSideConnection(self.agent_factory, transport) # -- POST --------------------------------------------------------------- @@ -267,7 +276,7 @@ async def handle_post( async def _handle_initialize(self, message: dict[str, Any]) -> Response: connection_id = uuid.uuid4().hex transport = _HttpTransport(message.get("id")) - conn = AgentSideConnection(self.agent_factory, transport) + conn = self._connect(transport) self._connections[connection_id] = (conn, transport) # Deliver initialize to the agent and await its response so we can return # the 200 body synchronously (initialize is the one blocking POST). If the diff --git a/src/acp/ws/server.py b/src/acp/ws/server.py index f6179ea..e5b2a51 100644 --- a/src/acp/ws/server.py +++ b/src/acp/ws/server.py @@ -12,6 +12,7 @@ from ..http.protocol import CONNECTION_ID_HEADER if TYPE_CHECKING: + from ..experimental.negotiation import AgentProtocolRouter from ..http.server import AgentFactory __all__ = ["handle_websocket"] @@ -48,11 +49,16 @@ async def close(self) -> None: await self._ws.close() -async def handle_websocket(agent_factory: AgentFactory, websocket: WebSocket) -> None: +async def handle_websocket(agent_factory: AgentFactory | AgentProtocolRouter, websocket: WebSocket) -> None: """Run one agent for the lifetime of the socket; disconnect cancels its work.""" await websocket.accept(headers=[(CONNECTION_ID_HEADER.lower().encode(), uuid.uuid4().hex.encode())]) transport = _WebSocketTransport(websocket) - conn = AgentSideConnection(agent_factory, transport, listening=False) + from ..experimental.negotiation import AgentProtocolRouter + + if isinstance(agent_factory, AgentProtocolRouter): + conn = agent_factory.connect(transport, listening=False) + else: + conn = AgentSideConnection(agent_factory, transport, listening=False) try: await conn.listen() finally: diff --git a/tests/http/test_v2_loopback.py b/tests/http/test_v2_loopback.py new file mode 100644 index 0000000..9f591aa --- /dev/null +++ b/tests/http/test_v2_loopback.py @@ -0,0 +1,226 @@ +"""v2 end-to-end loopback over the version-agnostic HTTP/WS transports. + +The v2 runtime is served through an :class:`AgentProtocolRouter` instead of a +v1 ``AgentFactory``; the ASGI server must negotiate v2 from ``initialize`` and +route ``session/update`` traffic (and ``session/resume`` replay) the same way it +already routes v1's ``session/load``. +""" + +from __future__ import annotations + +import asyncio +from typing import Any + +import pytest + +import acp +from acp.experimental import AgentProtocolRouter, v2 +from acp.http.asgi import create_asgi_app +from acp.http.client import create_http_stream +from acp.ws.client import create_websocket_stream +from tests.conftest import TestAgent, TestClient + + +class RoutedV2Agent: + def __init__(self) -> None: + self.connection: v2.AgentSideConnection | None = None + + def on_connect(self, connection: v2.AgentSideConnection) -> None: + self.connection = connection + + async def initialize( + self, + protocol_version: int, + info: v2.schema.Implementation, + capabilities: v2.schema.ClientCapabilities | None = None, + **kwargs: Any, + ) -> v2.schema.InitializeResponse: + return v2.schema.InitializeResponse( + protocol_version=v2.PROTOCOL_VERSION, + info=v2.schema.Implementation(name="v2-agent", version="1.0.0"), + ) + + async def new_session( + self, + cwd: str, + additional_directories: list[str] | None = None, + mcp_servers: list[Any] | None = None, + **kwargs: Any, + ) -> v2.schema.NewSessionResponse: + return v2.schema.NewSessionResponse(session_id="sess-v2") + + async def prompt( + self, + session_id: str, + prompt: list[Any], + **kwargs: Any, + ) -> v2.schema.PromptResponse: + assert self.connection is not None + await self.connection.session_update( + session_id=session_id, + update=v2.schema.AgentMessageChunk( + message_id="msg-1", + content=v2.schema.TextContentBlock(text="hello-v2"), + ), + ) + return v2.schema.PromptResponse(message_id="user-msg-1") + + async def resume_session( + self, + session_id: str, + cwd: str, + additional_directories: list[str] | None = None, + mcp_servers: list[Any] | None = None, + replay_from: Any = None, + **kwargs: Any, + ) -> v2.schema.ResumeSessionResponse: + assert self.connection is not None + for i in range(3): + await self.connection.session_update( + session_id=session_id, + update=v2.schema.AgentMessageChunk( + message_id=f"history-{i}", + content=v2.schema.TextContentBlock(text=f"history-{i}"), + ), + ) + return v2.schema.ResumeSessionResponse() + + +class CapturingClient: + def __init__(self) -> None: + self.updates: list[tuple[str, Any]] = [] + + async def session_update(self, session_id: str, update: Any, **kwargs: Any) -> None: + self.updates.append((session_id, update)) + + +def _make_app(router: AgentProtocolRouter) -> Any: + return create_asgi_app(router) + + +def _v2_agent_factory(connection: v2.AgentSideConnection) -> RoutedV2Agent: + return RoutedV2Agent() + + +def _router() -> AgentProtocolRouter: + return AgentProtocolRouter(v2=_v2_agent_factory) + + +async def _v2_client(transport: Any) -> tuple[v2.ClientSideConnection, CapturingClient]: + client = CapturingClient() + conn = v2.connect_to_agent(client, transport) + return conn, client + + +@pytest.mark.asyncio +async def test_v2_http_prompt_streams_update(serve_asgi) -> None: + server = await serve_asgi(_make_app(_router())) + transport = create_http_stream(server.http_url) + conn, client = await _v2_client(transport) + try: + init = await asyncio.wait_for( + conn.initialize( + protocol_version=v2.PROTOCOL_VERSION, + info=v2.schema.Implementation(name="v2-client", version="1.0.0"), + ), + timeout=10, + ) + assert init.protocol_version == v2.PROTOCOL_VERSION + new = await asyncio.wait_for(conn.new_session(cwd="."), timeout=10) + assert new.session_id == "sess-v2" + result = await asyncio.wait_for( + conn.prompt(session_id=new.session_id, prompt=[v2.schema.TextContentBlock(text="hi")]), + timeout=10, + ) + assert result.message_id == "user-msg-1" + await asyncio.sleep(0.2) + assert client.updates, "expected a session/update over SSE" + assert client.updates[0][0] == "sess-v2" + assert client.updates[0][1].content.text == "hello-v2" + finally: + await conn.close() + await transport.close() + + +@pytest.mark.asyncio +async def test_v2_websocket_prompt_streams_update(serve_asgi) -> None: + server = await serve_asgi(_make_app(_router())) + transport = await create_websocket_stream(server.ws_url) + conn, client = await _v2_client(transport) + try: + await asyncio.wait_for( + conn.initialize( + protocol_version=v2.PROTOCOL_VERSION, + info=v2.schema.Implementation(name="v2-client", version="1.0.0"), + ), + timeout=10, + ) + new = await asyncio.wait_for(conn.new_session(cwd="."), timeout=10) + await asyncio.wait_for( + conn.prompt(session_id=new.session_id, prompt=[v2.schema.TextContentBlock(text="hi")]), + timeout=10, + ) + await asyncio.sleep(0.2) + assert client.updates, "expected a session/update over WS" + assert client.updates[0][1].content.text == "hello-v2" + finally: + await conn.close() + await transport.close() + + +@pytest.mark.asyncio +async def test_v2_http_resume_replays_history(serve_asgi) -> None: + server = await serve_asgi(_make_app(_router())) + transport = create_http_stream(server.http_url) + conn, client = await _v2_client(transport) + try: + await asyncio.wait_for( + conn.initialize( + protocol_version=v2.PROTOCOL_VERSION, + info=v2.schema.Implementation(name="v2-client", version="1.0.0"), + ), + timeout=10, + ) + loaded = await asyncio.wait_for( + conn.resume_session(session_id="saved-session", cwd="/"), + timeout=10, + ) + assert loaded == v2.schema.ResumeSessionResponse() + assert [update.content.text for _, update in client.updates] == [f"history-{i}" for i in range(3)] + finally: + await conn.close() + await transport.close() + + +@pytest.mark.asyncio +async def test_dual_stack_http_serves_v1_and_v2(serve_asgi) -> None: + router = AgentProtocolRouter(v1=lambda conn: TestAgent(), v2=_v2_agent_factory) + server = await serve_asgi(create_asgi_app(router)) + + v1_transport = create_http_stream(server.http_url) + v1_conn = acp.connect_to_agent(TestClient(), v1_transport) + try: + init = await asyncio.wait_for(v1_conn.initialize(protocol_version=acp.PROTOCOL_VERSION), timeout=10) + assert init.protocol_version == acp.PROTOCOL_VERSION + new = await asyncio.wait_for(v1_conn.new_session(cwd="."), timeout=10) + assert new.session_id == "test-session-123" + finally: + await v1_conn.close() + await v1_transport.close() + + v2_transport = create_http_stream(server.http_url) + v2_conn, _ = await _v2_client(v2_transport) + try: + init = await asyncio.wait_for( + v2_conn.initialize( + protocol_version=v2.PROTOCOL_VERSION, + info=v2.schema.Implementation(name="v2-client", version="1.0.0"), + ), + timeout=10, + ) + assert init.protocol_version == v2.PROTOCOL_VERSION + new = await asyncio.wait_for(v2_conn.new_session(cwd="."), timeout=10) + assert new.session_id == "sess-v2" + finally: + await v2_conn.close() + await v2_transport.close()