From 653f82d60b315fa104895a49df1785f87c7eca6e Mon Sep 17 00:00:00 2001 From: Frost Ming Date: Mon, 21 Sep 2026 10:46:49 +0800 Subject: [PATCH 1/2] feat: complete missing v1 protocol methods --- docs/quickstart.md | 41 ++++ scripts/gen_signature.py | 19 +- src/acp/__init__.py | 58 ++++++ src/acp/agent/connection.py | 57 ++++- src/acp/agent/router.py | 67 ++++++ src/acp/client/connection.py | 276 ++++++++++++++++++++++++- src/acp/client/router.py | 18 ++ src/acp/interfaces.py | 153 +++++++++++++- tests/conftest.py | 18 ++ tests/test_gen_all.py | 38 ++++ tests/test_protocol_method_coverage.py | 210 +++++++++++++++++++ 11 files changed, 947 insertions(+), 8 deletions(-) create mode 100644 tests/test_protocol_method_coverage.py diff --git a/docs/quickstart.md b/docs/quickstart.md index 319de39..1724c11 100644 --- a/docs/quickstart.md +++ b/docs/quickstart.md @@ -187,6 +187,47 @@ are available as `TerminalAuthMethod`, replacing the incorrect `EnvVarAuthMethod name. Accepted elicitation content validates scalar values and string lists; nested objects are not valid form values. +## Additional v1 protocol methods + +The SDK exposes all agent and client methods in the bundled `schema-v1.23.0` +method tables. `session/delete` and `logout` are stable and do not require +`use_unstable_protocol=True`. The other methods in the table below are unstable; +enable that flag on the receiving connection (or `run_agent`) to route them. +Check the peer's advertised capabilities before calling these methods. + +| Wire method | Python method | Receiver | +| --- | --- | --- | +| `session/delete` | `delete_session(session_id=...)` | Agent | +| `providers/list`, `providers/set`, `providers/disable` | `list_providers`, `set_provider`, `disable_provider` | Agent | +| `logout` | `logout()` | Agent | +| `mcp/connect`, `mcp/disconnect` | `connect_mcp`, `disconnect_mcp` | Client | +| `mcp/message` (request) | `mcp_message` | Either peer | +| `mcp/message` (notification) | `notify_mcp` | Either peer | +| `nes/start`, `nes/suggest`, `nes/close` | `start_nes`, `suggest_nes`, `close_nes` | Agent | +| `nes/accept`, `nes/reject` | `accept_nes`, `reject_nes` | Agent | +| `document/didOpen`, `document/didChange`, `document/didClose`, `document/didSave`, `document/didFocus` | `did_open`, `did_change`, `did_close`, `did_save`, `did_focus` | Agent | + +For example, an agent that advertises session deletion implements: + +```python +from acp import DeleteSessionResponse + +async def delete_session(self, session_id: str, **kwargs) -> DeleteSessionResponse: + await self.session_store.delete(session_id) + return DeleteSessionResponse() +``` + +A client then calls `await connection.delete_session(session_id=session_id)`. +Deletion removes stored session data; `close_session` only closes the active +session. The SDK dispatches these calls to your implementation; it does not +provide session storage or provider management itself. Missing request handlers +return the JSON-RPC method-not-found error. + +MCP requests return the inner JSON result unchanged, including `null`. Use +`notify_mcp` for one-way MCP messages. Both APIs accept `connection_id`, `method`, +and optional `params`. These methods share the same connections and routers +across stdio, HTTP, and WebSocket transports. + ## Optional — Talk to the Gemini CLI _Have the Gemini CLI installed? Run the bridge to exercise permission flows._ diff --git a/scripts/gen_signature.py b/scripts/gen_signature.py index d516abf..8bbb4ac 100644 --- a/scripts/gen_signature.py +++ b/scripts/gen_signature.py @@ -134,10 +134,8 @@ def _format_annotation(self, annotation: t.Any) -> ast.expr: for argument in rest: formatted = ast.BinOp(left=formatted, op=ast.BitOr(), right=self._format_annotation(argument)) return formatted - if origin is t.Literal and annotation in self._literals.values(): - name = next(name for name, value in self._literals.items() if value is annotation) - self._add_schema_import(name) - return ast.Name(id=name) + if origin is t.Literal: + return self._format_literal(annotation) elif ( inspect.isclass(annotation) and issubclass(annotation, BaseModel) @@ -166,6 +164,19 @@ def _format_annotation(self, annotation: t.Any) -> ast.expr: self._add_typing_import("Any") return ast.Name(id="Any") + def _format_literal(self, annotation: t.Any) -> ast.expr: + if annotation in self._literals.values(): + name = next(name for name, value in self._literals.items() if value is annotation) + self._add_schema_import(name) + return ast.Name(id=name) + self._add_typing_import("Literal") + values = [ast.Constant(value=value) for value in t.get_args(annotation)] + return ast.Subscript( + value=ast.Name(id="Literal"), + slice=values[0] if len(values) == 1 else ast.Tuple(elts=values, ctx=ast.Load()), + ctx=ast.Load(), + ) + def gen_signature(source_dir: Path) -> None: global schema diff --git a/src/acp/__init__.py b/src/acp/__init__.py index ded77bb..02928a1 100644 --- a/src/acp/__init__.py +++ b/src/acp/__init__.py @@ -14,11 +14,16 @@ ) from .schema import ( AcceptElicitationResponse, + AcceptNesNotification, AuthenticateRequest, AuthenticateResponse, CancelElicitationResponse, CancelNotification, + CloseNesRequest, + CloseNesResponse, CompleteElicitationNotification, + ConnectMcpRequest, + ConnectMcpResponse, CreateElicitationRequest, CreateElicitationResponse, CreateFormElicitationRequest, @@ -31,6 +36,17 @@ CreateUrlRequestElicitationRequest, CreateUrlSessionElicitationRequest, DeclineElicitationResponse, + DeleteSessionRequest, + DeleteSessionResponse, + DidChangeDocumentNotification, + DidCloseDocumentNotification, + DidFocusDocumentNotification, + DidOpenDocumentNotification, + DidSaveDocumentNotification, + DisableProviderRequest, + DisableProviderResponse, + DisconnectMcpRequest, + DisconnectMcpResponse, ElicitationBooleanPropertySchema, ElicitationCapabilities, ElicitationFormCapabilities, @@ -50,8 +66,14 @@ InitializeResponse, KillTerminalRequest, KillTerminalResponse, + ListProvidersRequest, + ListProvidersResponse, LoadSessionRequest, LoadSessionResponse, + LogoutRequest, + LogoutResponse, + MessageMcpNotification, + MessageMcpRequest, NewSessionRequest, NewSessionResponse, OtherElicitationResponse, @@ -59,15 +81,22 @@ PromptResponse, ReadTextFileRequest, ReadTextFileResponse, + RejectNesNotification, ReleaseTerminalRequest, ReleaseTerminalResponse, RequestPermissionRequest, RequestPermissionResponse, SessionNotification, + SetProviderRequest, + SetProviderResponse, SetSessionConfigOptionResponse, SetSessionConfigOptionSelectRequest, SetSessionModeRequest, SetSessionModeResponse, + StartNesRequest, + StartNesResponse, + SuggestNesRequest, + SuggestNesResponse, TerminalOutputRequest, TerminalOutputResponse, WaitForTerminalExitRequest, @@ -97,6 +126,35 @@ "AGENT_METHODS", "CLIENT_METHODS", # types + "AcceptNesNotification", + "CloseNesRequest", + "CloseNesResponse", + "ConnectMcpRequest", + "ConnectMcpResponse", + "DeleteSessionRequest", + "DeleteSessionResponse", + "DidChangeDocumentNotification", + "DidCloseDocumentNotification", + "DidFocusDocumentNotification", + "DidOpenDocumentNotification", + "DidSaveDocumentNotification", + "DisableProviderRequest", + "DisableProviderResponse", + "DisconnectMcpRequest", + "DisconnectMcpResponse", + "ListProvidersRequest", + "ListProvidersResponse", + "LogoutRequest", + "LogoutResponse", + "MessageMcpNotification", + "MessageMcpRequest", + "RejectNesNotification", + "SetProviderRequest", + "SetProviderResponse", + "StartNesRequest", + "StartNesResponse", + "SuggestNesRequest", + "SuggestNesResponse", "InitializeRequest", "InitializeResponse", "NewSessionRequest", diff --git a/src/acp/agent/connection.py b/src/acp/agent/connection.py index f521236..37e7f46 100644 --- a/src/acp/agent/connection.py +++ b/src/acp/agent/connection.py @@ -21,6 +21,8 @@ CancelElicitationResponse, CompleteElicitationNotification, ConfigOptionUpdate, + ConnectMcpRequest, + ConnectMcpResponse, CreateElicitationResponse, CreateFormElicitationRequest, CreateFormRequestElicitationRequest, @@ -32,6 +34,8 @@ CreateUrlSessionElicitationRequest, CurrentModeUpdate, DeclineElicitationResponse, + DisconnectMcpRequest, + DisconnectMcpResponse, ElicitationFormRequestMode, ElicitationFormSessionMode, ElicitationMode, @@ -40,6 +44,8 @@ EnvVariable, KillTerminalRequest, KillTerminalResponse, + MessageMcpNotification, + MessageMcpRequest, PermissionOption, ReadTextFileRequest, ReadTextFileResponse, @@ -64,7 +70,15 @@ WriteTextFileRequest, WriteTextFileResponse, ) -from ..utils import compatible_class, notify_model, param_model, request_model, request_optional_model, serialize_params +from ..utils import ( + compatible_class, + notify_model, + param_model, + request_model, + request_model_from_dict, + request_optional_model, + serialize_params, +) from .router import build_agent_router __all__ = ["AgentSideConnection"] @@ -278,6 +292,47 @@ async def complete_elicitation(self, elicitation_id: str, **kwargs: Any) -> None CompleteElicitationNotification(elicitation_id=elicitation_id, field_meta=kwargs or None), ) + @param_model(ConnectMcpRequest) + async def connect_mcp(self, server_id: str, **kwargs: Any) -> ConnectMcpResponse: + return await request_model( + self._conn, + CLIENT_METHODS["mcp_connect"], + ConnectMcpRequest(server_id=server_id, field_meta=kwargs or None), + ConnectMcpResponse, + ) + + @param_model(DisconnectMcpRequest) + async def disconnect_mcp(self, connection_id: str, **kwargs: Any) -> DisconnectMcpResponse: + return await request_model_from_dict( + self._conn, + CLIENT_METHODS["mcp_disconnect"], + DisconnectMcpRequest(connection_id=connection_id, field_meta=kwargs or None), + DisconnectMcpResponse, + ) + + @param_model(MessageMcpRequest) + async def mcp_message( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> Any: + return await self._conn.send_request( + CLIENT_METHODS["mcp_message"], + serialize_params( + MessageMcpRequest(connection_id=connection_id, method=method, params=params, field_meta=kwargs or None) + ), + ) + + @param_model(MessageMcpNotification) + async def notify_mcp( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> None: + await notify_model( + self._conn, + CLIENT_METHODS["mcp_message"], + MessageMcpNotification( + connection_id=connection_id, method=method, params=params, field_meta=kwargs or None + ), + ) + async def ext_method(self, method: str, params: dict[str, Any]) -> dict[str, Any]: return await self._conn.send_request(f"_{method}", params) diff --git a/src/acp/agent/router.py b/src/acp/agent/router.py index 7dd58a5..01b15a7 100644 --- a/src/acp/agent/router.py +++ b/src/acp/agent/router.py @@ -9,19 +9,36 @@ from ..meta import AGENT_METHODS from ..router import MessageRouter, Route, _resolve_handler, _warn_legacy_handler from ..schema import ( + AcceptNesNotification, AuthenticateRequest, CancelNotification, + CloseNesRequest, CloseSessionRequest, + DeleteSessionRequest, + DidChangeDocumentNotification, + DidCloseDocumentNotification, + DidFocusDocumentNotification, + DidOpenDocumentNotification, + DidSaveDocumentNotification, + DisableProviderRequest, ForkSessionRequest, InitializeRequest, + ListProvidersRequest, ListSessionsRequest, LoadSessionRequest, + LogoutRequest, + MessageMcpNotification, + MessageMcpRequest, NewSessionRequest, PromptRequest, + RejectNesNotification, ResumeSessionRequest, + SetProviderRequest, SetSessionConfigOptionBooleanRequest, SetSessionConfigOptionSelectRequest, SetSessionModeRequest, + StartNesRequest, + SuggestNesRequest, ) from ..utils import model_to_kwargs, normalize_result @@ -102,6 +119,56 @@ def build_agent_router(agent: Agent, use_unstable_protocol: bool = False) -> Mes router.route_notification(AGENT_METHODS["session_cancel"], CancelNotification, agent, "cancel") + router.route_request( + AGENT_METHODS["session_delete"], + DeleteSessionRequest, + agent, + "delete_session", + adapt_result=normalize_result, + ) + router.route_request(AGENT_METHODS["providers_list"], ListProvidersRequest, agent, "list_providers", unstable=True) + router.route_request( + AGENT_METHODS["providers_set"], + SetProviderRequest, + agent, + "set_provider", + unstable=True, + adapt_result=normalize_result, + ) + router.route_request( + AGENT_METHODS["providers_disable"], + DisableProviderRequest, + agent, + "disable_provider", + unstable=True, + adapt_result=normalize_result, + ) + router.route_request(AGENT_METHODS["logout"], LogoutRequest, agent, "logout", adapt_result=normalize_result) + router.route_request(AGENT_METHODS["mcp_message"], MessageMcpRequest, agent, "mcp_message", unstable=True) + router.route_notification(AGENT_METHODS["mcp_message"], MessageMcpNotification, agent, "notify_mcp", unstable=True) + router.route_request(AGENT_METHODS["nes_start"], StartNesRequest, agent, "start_nes", unstable=True) + router.route_request(AGENT_METHODS["nes_suggest"], SuggestNesRequest, agent, "suggest_nes", unstable=True) + router.route_request( + AGENT_METHODS["nes_close"], CloseNesRequest, agent, "close_nes", unstable=True, adapt_result=normalize_result + ) + router.route_notification(AGENT_METHODS["nes_accept"], AcceptNesNotification, agent, "accept_nes", unstable=True) + router.route_notification(AGENT_METHODS["nes_reject"], RejectNesNotification, agent, "reject_nes", unstable=True) + router.route_notification( + AGENT_METHODS["document_did_open"], DidOpenDocumentNotification, agent, "did_open", unstable=True + ) + router.route_notification( + AGENT_METHODS["document_did_change"], DidChangeDocumentNotification, agent, "did_change", unstable=True + ) + router.route_notification( + AGENT_METHODS["document_did_close"], DidCloseDocumentNotification, agent, "did_close", unstable=True + ) + router.route_notification( + AGENT_METHODS["document_did_save"], DidSaveDocumentNotification, agent, "did_save", unstable=True + ) + router.route_notification( + AGENT_METHODS["document_did_focus"], DidFocusDocumentNotification, agent, "did_focus", unstable=True + ) + @router.handle_extension_request async def _handle_extension_request(name: str, payload: dict[str, Any]) -> Any: ext = getattr(agent, "ext_method", None) diff --git a/src/acp/client/connection.py b/src/acp/client/connection.py index 000a5da..4db0c85 100644 --- a/src/acp/client/connection.py +++ b/src/acp/client/connection.py @@ -3,7 +3,7 @@ import asyncio from collections.abc import Callable from contextvars import ContextVar -from typing import Any, cast, final +from typing import Any, Literal, cast, final from .._transport import Transport from ..connection import Connection @@ -12,14 +12,26 @@ from ..meta import AGENT_METHODS, CLIENT_METHODS from ..router import _resolve_handler, _warn_legacy_handler from ..schema import ( + AcceptNesNotification, AcpMcpServer, AudioContentBlock, AuthenticateRequest, AuthenticateResponse, CancelNotification, ClientCapabilities, + CloseNesRequest, + CloseNesResponse, CloseSessionRequest, CloseSessionResponse, + DeleteSessionRequest, + DeleteSessionResponse, + DidChangeDocumentNotification, + DidCloseDocumentNotification, + DidFocusDocumentNotification, + DidOpenDocumentNotification, + DidSaveDocumentNotification, + DisableProviderRequest, + DisableProviderResponse, EmbeddedResourceContentBlock, ForkSessionRequest, ForkSessionResponse, @@ -28,28 +40,55 @@ Implementation, InitializeRequest, InitializeResponse, + ListProvidersRequest, + ListProvidersResponse, ListSessionsRequest, ListSessionsResponse, LoadSessionRequest, LoadSessionResponse, + LogoutRequest, + LogoutResponse, McpServerStdio, + MessageMcpNotification, + MessageMcpRequest, + NesRepository, + NesSuggestContext, NewSessionRequest, NewSessionResponse, + Position, PromptRequest, PromptResponse, + Range, + RejectNesNotification, ResourceContentBlock, ResumeSessionRequest, ResumeSessionResponse, SessionNotification, + SetProviderRequest, + SetProviderResponse, SetSessionConfigOptionBooleanRequest, SetSessionConfigOptionResponse, SetSessionConfigOptionSelectRequest, SetSessionModeRequest, SetSessionModeResponse, SseMcpServer, + StartNesRequest, + StartNesResponse, + SuggestNesRequest, + SuggestNesResponse, TextContentBlock, + TextDocumentContentChangeEvent, + WorkspaceFolder, +) +from ..utils import ( + compatible_class, + notify_model, + param_model, + param_models, + request_model, + request_model_from_dict, + serialize_params, ) -from ..utils import compatible_class, notify_model, param_model, param_models, request_model, request_model_from_dict from .router import build_client_router __all__ = ["ClientSideConnection"] @@ -336,6 +375,239 @@ async def cancel(self, session_id: str, **kwargs: Any) -> None: CancelNotification(session_id=session_id, field_meta=kwargs or None), ) + @param_model(DeleteSessionRequest) + async def delete_session(self, session_id: str, **kwargs: Any) -> DeleteSessionResponse: + return await request_model_from_dict( + self._conn, + AGENT_METHODS["session_delete"], + DeleteSessionRequest(session_id=session_id, field_meta=kwargs or None), + DeleteSessionResponse, + ) + + @param_model(ListProvidersRequest) + async def list_providers(self, **kwargs: Any) -> ListProvidersResponse: + return await request_model( + self._conn, + AGENT_METHODS["providers_list"], + ListProvidersRequest(field_meta=kwargs or None), + ListProvidersResponse, + ) + + @param_model(SetProviderRequest) + async def set_provider( + self, + provider_id: str, + api_type: Literal["anthropic"] + | Literal["openai"] + | Literal["azure"] + | Literal["vertex"] + | Literal["bedrock"] + | str, + base_url: str, + headers: dict[str, str] | None = None, + **kwargs: Any, + ) -> SetProviderResponse: + return await request_model_from_dict( + self._conn, + AGENT_METHODS["providers_set"], + SetProviderRequest( + provider_id=provider_id, + api_type=api_type, + base_url=base_url, + headers=headers, + field_meta=kwargs or None, + ), + SetProviderResponse, + ) + + @param_model(DisableProviderRequest) + async def disable_provider(self, provider_id: str, **kwargs: Any) -> DisableProviderResponse: + return await request_model_from_dict( + self._conn, + AGENT_METHODS["providers_disable"], + DisableProviderRequest(provider_id=provider_id, field_meta=kwargs or None), + DisableProviderResponse, + ) + + @param_model(LogoutRequest) + async def logout(self, **kwargs: Any) -> LogoutResponse: + return await request_model_from_dict( + self._conn, AGENT_METHODS["logout"], LogoutRequest(field_meta=kwargs or None), LogoutResponse + ) + + @param_model(MessageMcpRequest) + async def mcp_message( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> Any: + return await self._conn.send_request( + AGENT_METHODS["mcp_message"], + serialize_params( + MessageMcpRequest(connection_id=connection_id, method=method, params=params, field_meta=kwargs or None) + ), + ) + + @param_model(MessageMcpNotification) + async def notify_mcp( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> None: + await notify_model( + self._conn, + AGENT_METHODS["mcp_message"], + MessageMcpNotification( + connection_id=connection_id, method=method, params=params, field_meta=kwargs or None + ), + ) + + @param_model(StartNesRequest) + async def start_nes( + self, + workspace_uri: str | None = None, + workspace_folders: list[WorkspaceFolder] | None = None, + repository: NesRepository | None = None, + **kwargs: Any, + ) -> StartNesResponse: + return await request_model( + self._conn, + AGENT_METHODS["nes_start"], + StartNesRequest( + workspace_uri=workspace_uri, + workspace_folders=workspace_folders, + repository=repository, + field_meta=kwargs or None, + ), + StartNesResponse, + ) + + @param_model(SuggestNesRequest) + async def suggest_nes( + self, + session_id: str, + uri: str, + version: int, + position: Position, + trigger_kind: Literal["automatic", "diagnostic", "manual"], + selection: Range | None = None, + context: NesSuggestContext | None = None, + **kwargs: Any, + ) -> SuggestNesResponse: + return await request_model( + self._conn, + AGENT_METHODS["nes_suggest"], + SuggestNesRequest( + session_id=session_id, + uri=uri, + version=version, + position=position, + selection=selection, + trigger_kind=trigger_kind, + context=context, + field_meta=kwargs or None, + ), + SuggestNesResponse, + ) + + @param_model(CloseNesRequest) + async def close_nes(self, session_id: str, **kwargs: Any) -> CloseNesResponse: + return await request_model_from_dict( + self._conn, + AGENT_METHODS["nes_close"], + CloseNesRequest(session_id=session_id, field_meta=kwargs or None), + CloseNesResponse, + ) + + @param_model(AcceptNesNotification) + async def accept_nes(self, session_id: str, id: str, **kwargs: Any) -> None: # noqa: A002 + await notify_model( + self._conn, + AGENT_METHODS["nes_accept"], + AcceptNesNotification(session_id=session_id, id=id, field_meta=kwargs or None), + ) + + @param_model(RejectNesNotification) + async def reject_nes( + self, + session_id: str, + id: str, # noqa: A002 + reason: Literal["rejected", "ignored", "replaced", "cancelled"] | None = None, + **kwargs: Any, + ) -> None: + await notify_model( + self._conn, + AGENT_METHODS["nes_reject"], + RejectNesNotification(session_id=session_id, id=id, reason=reason, field_meta=kwargs or None), + ) + + @param_model(DidOpenDocumentNotification) + async def did_open( + self, session_id: str, uri: str, language_id: str, version: int, text: str, **kwargs: Any + ) -> None: + await notify_model( + self._conn, + AGENT_METHODS["document_did_open"], + DidOpenDocumentNotification( + session_id=session_id, + uri=uri, + language_id=language_id, + version=version, + text=text, + field_meta=kwargs or None, + ), + ) + + @param_model(DidChangeDocumentNotification) + async def did_change( + self, + session_id: str, + uri: str, + version: int, + content_changes: list[TextDocumentContentChangeEvent], + **kwargs: Any, + ) -> None: + await notify_model( + self._conn, + AGENT_METHODS["document_did_change"], + DidChangeDocumentNotification( + session_id=session_id, + uri=uri, + version=version, + content_changes=content_changes, + field_meta=kwargs or None, + ), + ) + + @param_model(DidCloseDocumentNotification) + async def did_close(self, session_id: str, uri: str, **kwargs: Any) -> None: + await notify_model( + self._conn, + AGENT_METHODS["document_did_close"], + DidCloseDocumentNotification(session_id=session_id, uri=uri, field_meta=kwargs or None), + ) + + @param_model(DidSaveDocumentNotification) + async def did_save(self, session_id: str, uri: str, **kwargs: Any) -> None: + await notify_model( + self._conn, + AGENT_METHODS["document_did_save"], + DidSaveDocumentNotification(session_id=session_id, uri=uri, field_meta=kwargs or None), + ) + + @param_model(DidFocusDocumentNotification) + async def did_focus( + self, session_id: str, uri: str, version: int, position: Position, visible_range: Range, **kwargs: Any + ) -> None: + await notify_model( + self._conn, + AGENT_METHODS["document_did_focus"], + DidFocusDocumentNotification( + session_id=session_id, + uri=uri, + version=version, + position=position, + visible_range=visible_range, + field_meta=kwargs or None, + ), + ) + async def ext_method(self, method: str, params: dict[str, Any]) -> dict[str, Any]: return await self._conn.send_request(f"_{method}", params) diff --git a/src/acp/client/router.py b/src/acp/client/router.py index 204a67b..b7bc825 100644 --- a/src/acp/client/router.py +++ b/src/acp/client/router.py @@ -10,17 +10,21 @@ from ..router import MessageRouter, Route, _resolve_handler, _warn_legacy_handler from ..schema import ( CompleteElicitationNotification, + ConnectMcpRequest, CreateElicitationRequest, CreateFormRequestElicitationRequest, CreateFormSessionElicitationRequest, CreateTerminalRequest, CreateUrlRequestElicitationRequest, CreateUrlSessionElicitationRequest, + DisconnectMcpRequest, ElicitationFormRequestMode, ElicitationFormSessionMode, ElicitationUrlRequestMode, ElicitationUrlSessionMode, KillTerminalRequest, + MessageMcpNotification, + MessageMcpRequest, ReadTextFileRequest, ReleaseTerminalRequest, RequestPermissionRequest, @@ -162,6 +166,20 @@ def build_client_router(client: Client, use_unstable_protocol: bool = False) -> router.route_notification(CLIENT_METHODS["session_update"], SessionNotification, client, "session_update") + router.route_request(CLIENT_METHODS["mcp_connect"], ConnectMcpRequest, client, "connect_mcp", unstable=True) + router.route_request( + CLIENT_METHODS["mcp_disconnect"], + DisconnectMcpRequest, + client, + "disconnect_mcp", + unstable=True, + adapt_result=normalize_result, + ) + router.route_request(CLIENT_METHODS["mcp_message"], MessageMcpRequest, client, "mcp_message", unstable=True) + router.route_notification( + CLIENT_METHODS["mcp_message"], MessageMcpNotification, client, "notify_mcp", unstable=True + ) + @router.handle_extension_request async def _handle_extension_request(name: str, payload: dict[str, Any]) -> Any: ext = getattr(client, "ext_method", None) diff --git a/src/acp/interfaces.py b/src/acp/interfaces.py index 00ff919..5dba20a 100644 --- a/src/acp/interfaces.py +++ b/src/acp/interfaces.py @@ -1,8 +1,9 @@ from __future__ import annotations -from typing import Any, Protocol +from typing import Any, Literal, Protocol from .schema import ( + AcceptNesNotification, AcpMcpServer, AgentMessageChunk, AgentPlanContentUpdate, @@ -15,14 +16,29 @@ AvailableCommandsUpdate, CancelNotification, ClientCapabilities, + CloseNesRequest, + CloseNesResponse, CloseSessionRequest, CloseSessionResponse, CompleteElicitationNotification, ConfigOptionUpdate, + ConnectMcpRequest, + ConnectMcpResponse, CreateElicitationResponse, CreateTerminalRequest, CreateTerminalResponse, CurrentModeUpdate, + DeleteSessionRequest, + DeleteSessionResponse, + DidChangeDocumentNotification, + DidCloseDocumentNotification, + DidFocusDocumentNotification, + DidOpenDocumentNotification, + DidSaveDocumentNotification, + DisableProviderRequest, + DisableProviderResponse, + DisconnectMcpRequest, + DisconnectMcpResponse, ElicitationMode, EmbeddedResourceContentBlock, EnvVariable, @@ -35,18 +51,29 @@ InitializeResponse, KillTerminalRequest, KillTerminalResponse, + ListProvidersRequest, + ListProvidersResponse, ListSessionsRequest, ListSessionsResponse, LoadSessionRequest, LoadSessionResponse, + LogoutRequest, + LogoutResponse, McpServerStdio, + MessageMcpNotification, + MessageMcpRequest, + NesRepository, + NesSuggestContext, NewSessionRequest, NewSessionResponse, PermissionOption, + Position, PromptRequest, PromptResponse, + Range, ReadTextFileRequest, ReadTextFileResponse, + RejectNesNotification, ReleaseTerminalRequest, ReleaseTerminalResponse, RequestPermissionRequest, @@ -59,15 +86,22 @@ SessionUpdateCompactionSummaryChunk, SessionUpdateCompactionUpdate, SessionUpdateNotice, + SetProviderRequest, + SetProviderResponse, SetSessionConfigOptionBooleanRequest, SetSessionConfigOptionResponse, SetSessionConfigOptionSelectRequest, SetSessionModeRequest, SetSessionModeResponse, SseMcpServer, + StartNesRequest, + StartNesResponse, + SuggestNesRequest, + SuggestNesResponse, TerminalOutputRequest, TerminalOutputResponse, TextContentBlock, + TextDocumentContentChangeEvent, ToolCallProgress, ToolCallStart, ToolCallUpdate, @@ -75,6 +109,7 @@ UserMessageChunk, WaitForTerminalExitRequest, WaitForTerminalExitResponse, + WorkspaceFolder, WriteTextFileRequest, WriteTextFileResponse, ) @@ -84,6 +119,22 @@ class Client(Protocol): + @param_model(ConnectMcpRequest) + async def connect_mcp(self, server_id: str, **kwargs: Any) -> ConnectMcpResponse: ... + + @param_model(DisconnectMcpRequest) + async def disconnect_mcp(self, connection_id: str, **kwargs: Any) -> DisconnectMcpResponse: ... + + @param_model(MessageMcpRequest) + async def mcp_message( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> Any: ... + + @param_model(MessageMcpNotification) + async def notify_mcp( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> None: ... + @param_model(RequestPermissionRequest) async def request_permission( self, session_id: str, tool_call: ToolCallUpdate, options: list[PermissionOption], **kwargs: Any @@ -165,6 +216,106 @@ def on_connect(self, conn: Agent) -> None: ... class Agent(Protocol): + @param_model(DeleteSessionRequest) + async def delete_session(self, session_id: str, **kwargs: Any) -> DeleteSessionResponse: ... + + @param_model(ListProvidersRequest) + async def list_providers(self, **kwargs: Any) -> ListProvidersResponse: ... + + @param_model(SetProviderRequest) + async def set_provider( + self, + provider_id: str, + api_type: Literal["anthropic"] + | Literal["openai"] + | Literal["azure"] + | Literal["vertex"] + | Literal["bedrock"] + | str, + base_url: str, + headers: dict[str, str] | None = None, + **kwargs: Any, + ) -> SetProviderResponse: ... + + @param_model(DisableProviderRequest) + async def disable_provider(self, provider_id: str, **kwargs: Any) -> DisableProviderResponse: ... + + @param_model(LogoutRequest) + async def logout(self, **kwargs: Any) -> LogoutResponse: ... + + @param_model(MessageMcpRequest) + async def mcp_message( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> Any: ... + + @param_model(MessageMcpNotification) + async def notify_mcp( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> None: ... + + @param_model(StartNesRequest) + async def start_nes( + self, + workspace_uri: str | None = None, + workspace_folders: list[WorkspaceFolder] | None = None, + repository: NesRepository | None = None, + **kwargs: Any, + ) -> StartNesResponse: ... + + @param_model(SuggestNesRequest) + async def suggest_nes( + self, + session_id: str, + uri: str, + version: int, + position: Position, + trigger_kind: Literal["automatic", "diagnostic", "manual"], + selection: Range | None = None, + context: NesSuggestContext | None = None, + **kwargs: Any, + ) -> SuggestNesResponse: ... + + @param_model(CloseNesRequest) + async def close_nes(self, session_id: str, **kwargs: Any) -> CloseNesResponse: ... + + @param_model(AcceptNesNotification) + async def accept_nes(self, session_id: str, id: str, **kwargs: Any) -> None: ... # noqa: A002 + + @param_model(RejectNesNotification) + async def reject_nes( + self, + session_id: str, + id: str, # noqa: A002 + reason: Literal["rejected", "ignored", "replaced", "cancelled"] | None = None, + **kwargs: Any, + ) -> None: ... + + @param_model(DidOpenDocumentNotification) + async def did_open( + self, session_id: str, uri: str, language_id: str, version: int, text: str, **kwargs: Any + ) -> None: ... + + @param_model(DidChangeDocumentNotification) + async def did_change( + self, + session_id: str, + uri: str, + version: int, + content_changes: list[TextDocumentContentChangeEvent], + **kwargs: Any, + ) -> None: ... + + @param_model(DidCloseDocumentNotification) + async def did_close(self, session_id: str, uri: str, **kwargs: Any) -> None: ... + + @param_model(DidSaveDocumentNotification) + async def did_save(self, session_id: str, uri: str, **kwargs: Any) -> None: ... + + @param_model(DidFocusDocumentNotification) + async def did_focus( + self, session_id: str, uri: str, version: int, position: Position, visible_range: Range, **kwargs: Any + ) -> None: ... + @param_model(InitializeRequest) async def initialize( self, diff --git a/tests/conftest.py b/tests/conftest.py index 2dc476b..38c501f 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -10,8 +10,10 @@ AcceptElicitationResponse, AuthenticateResponse, CompleteElicitationNotification, + ConnectMcpResponse, CreateElicitationResponse, CreateTerminalResponse, + DisconnectMcpResponse, ElicitationMode, InitializeResponse, KillTerminalResponse, @@ -235,6 +237,22 @@ async def complete_elicitation(self, elicitation_id: str, **kwargs: Any) -> None CompleteElicitationNotification(elicitation_id=elicitation_id, field_meta=kwargs or None) ) + async def connect_mcp(self, server_id: str, **kwargs: Any) -> ConnectMcpResponse: + raise RequestError.method_not_found("mcp/connect") + + async def disconnect_mcp(self, connection_id: str, **kwargs: Any) -> DisconnectMcpResponse: + raise RequestError.method_not_found("mcp/disconnect") + + async def mcp_message( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> Any: + raise RequestError.method_not_found("mcp/message") + + async def notify_mcp( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> None: + raise RequestError.method_not_found("mcp/message") + async def ext_method(self, method: str, params: dict) -> dict: self.ext_calls.append((method, params)) if method == "example.com/ping": diff --git a/tests/test_gen_all.py b/tests/test_gen_all.py index dec76f3..6091f37 100644 --- a/tests/test_gen_all.py +++ b/tests/test_gen_all.py @@ -70,3 +70,41 @@ def test_codegen_check_is_clean_and_read_only() -> None: assert generate_schema(check=True, protocol_version=2) assert generate_meta(check=True, protocol_version=2) assert {output: output.read_bytes() for output in outputs} == before + + +def test_signature_generation_preserves_inline_literal_values() -> None: + import ast + from typing import Literal, get_type_hints + + from scripts.gen_signature import NodeTransformer + + tree = ast.parse( + "from typing import Any\n" + "from schema import SetProviderRequest, SuggestNesRequest\n" + "class Methods:\n" + " @param_model(SetProviderRequest)\n" + " async def set_provider(self, **kwargs: Any): ...\n" + " @param_model(SuggestNesRequest)\n" + " async def suggest_nes(self, **kwargs: Any): ...\n" + ) + NodeTransformer().visit(tree) + ast.fix_missing_locations(tree) + # Evaluate the annotation itself to catch unquoted literal values and ensure + # generation adds the required typing import. + typing_import = tree.body[0] + assert isinstance(typing_import, ast.ImportFrom) + assert "Literal" in {alias.name for alias in typing_import.names} + methods = tree.body[-1] + assert isinstance(methods, ast.ClassDef) + suggest = methods.body[-1] + assert isinstance(suggest, ast.AsyncFunctionDef) + trigger = next(arg for arg in suggest.args.args if arg.arg == "trigger_kind") + annotation = ast.unparse(trigger.annotation) + + class Annotated: + __annotations__ = {"trigger": annotation} + + assert ( + get_type_hints(Annotated, globalns={"Literal": Literal})["trigger"] + == Literal["automatic", "diagnostic", "manual"] + ) diff --git a/tests/test_protocol_method_coverage.py b/tests/test_protocol_method_coverage.py new file mode 100644 index 0000000..2b15f15 --- /dev/null +++ b/tests/test_protocol_method_coverage.py @@ -0,0 +1,210 @@ +"""Exercise v1 methods added after the original session/terminal API.""" + +import asyncio +import inspect +from typing import Any, cast +from unittest.mock import AsyncMock + +import pytest +from pydantic import ValidationError + +from acp import schema +from acp.agent.connection import AgentSideConnection +from acp.agent.router import build_agent_router +from acp.client.connection import ClientSideConnection +from acp.client.router import build_client_router +from acp.exceptions import RequestError +from acp.interfaces import Agent, Client +from acp.meta import AGENT_METHODS, CLIENT_METHODS + +REQUESTS = [ + ("delete_session", "session/delete", {"session_id": "s"}, schema.DeleteSessionResponse()), + ("list_providers", "providers/list", {}, schema.ListProvidersResponse(providers=[])), + ( + "set_provider", + "providers/set", + {"provider_id": "p", "api_type": "custom", "base_url": "https://example.com", "headers": {"X-Test": "test"}}, + schema.SetProviderResponse(), + ), + ("disable_provider", "providers/disable", {"provider_id": "p"}, schema.DisableProviderResponse()), + ("logout", "logout", {}, schema.LogoutResponse()), + ("start_nes", "nes/start", {"workspace_uri": "file:///workspace"}, schema.StartNesResponse(session_id="n")), + ( + "suggest_nes", + "nes/suggest", + { + "session_id": "n", + "uri": "file:///a", + "version": 1, + "position": schema.Position(line=0, character=0), + "trigger_kind": "manual", + }, + schema.SuggestNesResponse(suggestions=[]), + ), + ("close_nes", "nes/close", {"session_id": "n"}, schema.CloseNesResponse()), +] +POSITION = schema.Position(line=0, character=0) +NOTIFICATIONS = [ + ("accept_nes", "nes/accept", {"session_id": "n", "id": "edit"}), + ("reject_nes", "nes/reject", {"session_id": "n", "id": "edit", "reason": "rejected"}), + ( + "did_open", + "document/didOpen", + {"session_id": "n", "uri": "file:///a", "language_id": "python", "version": 1, "text": "hello"}, + ), + ( + "did_change", + "document/didChange", + { + "session_id": "n", + "uri": "file:///a", + "version": 2, + "content_changes": [schema.TextDocumentContentChangeEvent(text="world")], + }, + ), + ("did_close", "document/didClose", {"session_id": "n", "uri": "file:///a"}), + ("did_save", "document/didSave", {"session_id": "n", "uri": "file:///a", "text": "world"}), + ( + "did_focus", + "document/didFocus", + { + "session_id": "n", + "uri": "file:///a", + "version": 2, + "position": POSITION, + "visible_range": schema.Range(start=POSITION, end=POSITION), + }, + ), +] + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("name", "wire_method", "params", "response"), REQUESTS) +async def test_agent_request_roundtrip(connect, agent, name, wire_method, params, response): + handler = AsyncMock(return_value=response) + setattr(agent, name, handler) + _, connection = connect(use_unstable_protocol=True) + result = await getattr(connection, name)(**params, trace="test") + assert result == response + assert handler.await_args is not None + assert handler.await_args.kwargs.items() >= (params | {"trace": "test"}).items() + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("name", "wire_method", "params"), NOTIFICATIONS) +async def test_agent_notification_roundtrip(connect, agent, name, wire_method, params): + received = asyncio.Event() + calls = [] + + async def handler(**kwargs): + calls.append(kwargs) + received.set() + + setattr(agent, name, handler) + _, connection = connect(use_unstable_protocol=True) + await getattr(connection, name)(**params, trace="test") + await asyncio.wait_for(received.wait(), 2) + assert calls[0].items() >= (params | {"trace": "test"}).items() + + +@pytest.mark.asyncio +async def test_mcp_connection_lifecycle(connect, client): + client.connect_mcp = AsyncMock(return_value=schema.ConnectMcpResponse(connection_id="c")) + client.disconnect_mcp = AsyncMock(return_value=None) + connection, _ = connect(use_unstable_protocol=True) + assert (await connection.connect_mcp(server_id="m", trace="test")).connection_id == "c" + client.connect_mcp.assert_awaited_once_with(server_id="m", trace="test") + assert await connection.disconnect_mcp(connection_id="c") == schema.DisconnectMcpResponse() + client.disconnect_mcp.assert_awaited_once_with(connection_id="c") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("response", [None, {}, {"tools": []}, [1, "x"], False, 42]) +@pytest.mark.parametrize("direction", ["agent", "client"]) +async def test_mcp_requests_and_notifications_share_method(connect, agent, client, direction, response): + receiver = agent if direction == "agent" else client + receiver.mcp_message = AsyncMock(return_value=response) + received = asyncio.Event() + + async def notify_mcp(**kwargs): + assert kwargs == {"connection_id": "c", "method": "notifications/initialized", "params": None, "trace": "test"} + received.set() + + receiver.notify_mcp = notify_mcp + agent_side, client_side = connect(use_unstable_protocol=True) + sender = client_side if direction == "agent" else agent_side + assert ( + await sender.mcp_message(connection_id="c", method="tools/list", params={"cursor": "next"}, trace="test") + == response + ) + receiver.mcp_message.assert_awaited_once_with( + connection_id="c", method="tools/list", params={"cursor": "next"}, trace="test" + ) + await sender.notify_mcp(connection_id="c", method="notifications/initialized", trace="test") + await asyncio.wait_for(received.wait(), 2) + + +@pytest.mark.asyncio +async def test_delete_session_empty_and_legacy_response(connect, agent): + agent.delete_session = AsyncMock(return_value=None) + _, connection = connect(use_unstable_protocol=True) + with pytest.warns(DeprecationWarning): + result = await connection.deleteSession(schema.DeleteSessionRequest(session_id="s")) + assert result == schema.DeleteSessionResponse() + agent.delete_session.assert_awaited_once_with(session_id="s") + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("name", "wire_method", "params", "response"), REQUESTS) +async def test_unimplemented_requests_report_method_not_found(name, wire_method, params, response): + router = build_agent_router(cast(Agent, object()), use_unstable_protocol=True) + with pytest.raises(RequestError) as exc: + await router(wire_method, {}, False) + assert isinstance(exc.value, RequestError) + assert exc.value.code == -32601 + + +@pytest.mark.asyncio +async def test_delete_session_validation(): + class Handler: + async def delete_session(self, session_id: str, **kwargs: Any): + return None + + router = build_agent_router(cast(Agent, Handler())) + assert await router("session/delete", {"sessionId": "s"}, False) == {} + with pytest.raises(ValidationError): + await router("session/delete", {}, False) + + +@pytest.mark.parametrize( + ("builder", "methods", "connection", "interface"), + [ + (build_agent_router, AGENT_METHODS, ClientSideConnection, Agent), + (build_client_router, CLIENT_METHODS, AgentSideConnection, Client), + ], +) +def test_all_schema_methods_have_routes_and_senders(builder, methods, connection, interface): + router = builder(object(), use_unstable_protocol=True) + assert set(router._requests) | set(router._notifications) == set(methods.values()) + source = inspect.getsource(connection) + for key in methods: + assert f'["{key}"]' in source + for name, method in inspect.getmembers(connection, inspect.isfunction): + if hasattr(method, "__param_model__"): + assert hasattr(interface, name) or any(char.isupper() for char in name) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("name", "params", "response"), + [ + ("delete_session", {"session_id": "s"}, schema.DeleteSessionResponse()), + ("logout", {}, schema.LogoutResponse()), + ], +) +async def test_stable_methods_work_without_unstable_flag(connect, agent, name, params, response): + handler = AsyncMock(return_value=response) + setattr(agent, name, handler) + _, connection = connect() + assert await getattr(connection, name)(**params) == response + handler.assert_awaited_once_with(**params) From ccb7cfbe39b0cb8043de85cb4936ebf374931f21 Mon Sep 17 00:00:00 2001 From: Frost Ming Date: Mon, 21 Sep 2026 10:56:02 +0800 Subject: [PATCH 2/2] test: fix protocol handler mocks on Python 3.10 --- tests/test_protocol_method_coverage.py | 14 ++++++++++++-- 1 file changed, 12 insertions(+), 2 deletions(-) diff --git a/tests/test_protocol_method_coverage.py b/tests/test_protocol_method_coverage.py index 2b15f15..3c5b821 100644 --- a/tests/test_protocol_method_coverage.py +++ b/tests/test_protocol_method_coverage.py @@ -82,7 +82,12 @@ @pytest.mark.parametrize(("name", "wire_method", "params", "response"), REQUESTS) async def test_agent_request_roundtrip(connect, agent, name, wire_method, params, response): handler = AsyncMock(return_value=response) - setattr(agent, name, handler) + + # Python 3.10 cannot inspect the signature of a bare AsyncMock. + async def handle(**kwargs): + return await handler(**kwargs) + + setattr(agent, name, handle) _, connection = connect(use_unstable_protocol=True) result = await getattr(connection, name)(**params, trace="test") assert result == response @@ -204,7 +209,12 @@ def test_all_schema_methods_have_routes_and_senders(builder, methods, connection ) async def test_stable_methods_work_without_unstable_flag(connect, agent, name, params, response): handler = AsyncMock(return_value=response) - setattr(agent, name, handler) + + # Python 3.10 cannot inspect the signature of a bare AsyncMock. + async def handle(**kwargs): + return await handler(**kwargs) + + setattr(agent, name, handle) _, connection = connect() assert await getattr(connection, name)(**params) == response handler.assert_awaited_once_with(**params)