Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions src/acp/experimental/negotiation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand Down
2 changes: 2 additions & 0 deletions src/acp/experimental/v2/agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
6 changes: 5 additions & 1 deletion src/acp/http/asgi.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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).
Expand Down
4 changes: 2 additions & 2 deletions src/acp/http/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down
12 changes: 10 additions & 2 deletions src/acp/http/protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand All @@ -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
Expand Down
19 changes: 14 additions & 5 deletions src/acp/http/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -45,6 +45,7 @@
if TYPE_CHECKING:
from collections.abc import AsyncGenerator

from ..experimental.negotiation import AgentProtocolConnection, AgentProtocolRouter
from ..interfaces import Agent

__all__ = [
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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 ---------------------------------------------------------------

Expand Down Expand Up @@ -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
Expand Down
10 changes: 8 additions & 2 deletions src/acp/ws/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"]
Expand Down Expand Up @@ -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:
Expand Down
226 changes: 226 additions & 0 deletions tests/http/test_v2_loopback.py
Original file line number Diff line number Diff line change
@@ -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()