From 4d8a3af88e2bdca39d1d03d6a196b12f266370f6 Mon Sep 17 00:00:00 2001 From: Frost Ming Date: Mon, 21 Sep 2026 16:46:45 +0800 Subject: [PATCH] refactor(v2)!: expand parameters and derive routes from protocols --- docs/experimental-v2.md | 99 ++++-- scripts/gen_all.py | 5 + scripts/gen_signature.py | 45 ++- src/acp/experimental/v2/__init__.py | 3 + src/acp/experimental/v2/_methods.py | 205 ++---------- src/acp/experimental/v2/_params.py | 65 ++++ src/acp/experimental/v2/_router.py | 23 +- src/acp/experimental/v2/agent.py | 148 +++++++-- src/acp/experimental/v2/client.py | 446 +++++++++++++++++++++----- src/acp/experimental/v2/interfaces.py | 314 ++++++++++++++++++ tests/test_gen_all.py | 38 +++ tests/test_protocol_negotiation.py | 74 ++++- tests/test_v2_routing.py | 141 ++++++++ tests/test_v2_runtime.py | 224 +++++++++---- 14 files changed, 1441 insertions(+), 389 deletions(-) create mode 100644 src/acp/experimental/v2/_params.py create mode 100644 src/acp/experimental/v2/interfaces.py create mode 100644 tests/test_v2_routing.py diff --git a/docs/experimental-v2.md b/docs/experimental-v2.md index ab005a8..7c1149a 100644 --- a/docs/experimental-v2.md +++ b/docs/experimental-v2.md @@ -5,40 +5,97 @@ The bindings use `schema-v2.0.0-alpha.5`. -The v2 runtime is separate from the stable v1 API. Its methods accept and return -generated request and response models directly. Install update handlers on the -client before opening a session because updates are independent connection -traffic: +The v2 runtime is separate from the stable v1 API. Like v1, connection methods +and agent/client handlers accept expanded, snake-case parameters. Responses and +nested values (content blocks, updates, capabilities) use `v2.schema` models. +Extra keyword arguments carry request `_meta`. + +Install update handlers before opening a session because updates are independent +connection traffic: ```python +from typing import Any from acp.experimental import v2 -class MyClient: +class MyClient(v2.Client): async def session_update( - self, - notification: v2.schema.UpdateSessionNotification, + self, session_id: str, update: Any, **kwargs: Any, ) -> None: - handle_update(notification) + handle_update(session_id, update) connection = v2.connect_to_agent(MyClient(), transport) initialized = await connection.initialize( - v2.schema.InitializeRequest( - protocol_version=v2.PROTOCOL_VERSION, - info=v2.schema.Implementation(name="my-client", version="1.0.0"), - ) -) -session = await connection.new_session( - v2.schema.NewSessionRequest(cwd="/workspace") + protocol_version=v2.PROTOCOL_VERSION, + info=v2.schema.Implementation(name="my-client", version="1.0.0"), ) +session = await connection.new_session(cwd="/workspace") accepted = await connection.prompt( - v2.schema.PromptRequest( - session_id=session.session_id, - prompt=[v2.schema.TextContentBlock(text="Hello")], - ) + session_id=session.session_id, + prompt=[v2.schema.TextContentBlock(text="Hello")], +) +``` + +Implement an agent with the same expanded handler style: + +```python +class MyAgent(v2.Agent): + 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="my-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="session-1") + + +await v2.run_agent(MyAgent()) +``` + +`v2.Agent` and `v2.Client` describe the v2 handler signatures; subclassing is +optional. Implement only the methods you support. Unimplemented requests return +method-not-found, and unimplemented notifications are ignored. The v2 protocols +are separate from v1 because initialization, prompt responses, permissions, and +session updates have different contracts. Both versions use `param_model` +metadata to derive their routes. V2 retains strict request/response validation +and requires successful initialization before other traffic. + +Previously, v2 methods accepted a whole request model. Replace +`connection.new_session(v2.schema.NewSessionRequest(cwd="/workspace"))` with +`connection.new_session(cwd="/workspace")`, and expand handler parameters likewise. + +Union requests also use expanded parameters: + +```python +await connection.set_config_option(config_id="thinking", session_id=session.session_id, value=True) +# type defaults to "boolean" for bool values and "id" otherwise. +await connection.set_config_option( + config_id="vendor/limit", session_id=session.session_id, value=10, type="vendor/number", +) +await agent_connection.create_elicitation( + message="Sign in", mode="url", session_id=session.session_id, + elicitation_id="sign-in-1", url="https://example.com/login", ) ``` +For elicitation, `session_id` selects session scope; otherwise `request_id` +selects request scope (including `None`). Pass `requested_schema` for form mode, +or `elicitation_id` and `url` for URL mode. Handlers receive the validated +branch's fields, including `type` for config options and `mode` for elicitation. + `session/prompt` returns after the agent inserts the user message into the ACP conversation, without waiting for processing to finish. The response requires a non-null `message_id`. Agents return `v2.schema.PromptResponse(message_id=...)` @@ -48,7 +105,7 @@ the same ID. That update may arrive before or after the response; use and do not carry a prompt identifier. Agents can send `v2.schema.SessionNotice(severity="warning", title="Context is nearly full")` -in an `UpdateSessionNotification`. V2 notices require no client capability and +with `await agent_connection.session_update(session_id=session_id, update=notice)`. V2 notices require no client capability and are live advisory events, outside retained session history. Clients may ignore them. Titles must be non-empty, and severity also accepts custom or future strings. @@ -59,7 +116,7 @@ tool name, while omitting `name` leaves it unchanged. This also applies to terminal updates and patch metadata. When applying received patches, use `update.model_dump(by_alias=True, exclude_unset=True)` to retain that distinction. -Setting `replay_from=v2.schema.ReplayFromStartVariant()` on a `ResumeSessionRequest` +Setting `replay_from=v2.schema.ReplayFromStartVariant()` on `connection.resume_session(...)` requests all retained conversation history; agents need not retain every message. Accepted elicitation content validates scalar values and string lists; nested objects are not valid form values. diff --git a/scripts/gen_all.py b/scripts/gen_all.py index 4e0375c..5e6bb76 100644 --- a/scripts/gen_all.py +++ b/scripts/gen_all.py @@ -94,6 +94,8 @@ def main() -> None: gen_meta.generate_meta(protocol_version=protocol_version) if protocol_version == 1: gen_signature.gen_signature(ROOT / "src" / "acp") + else: + gen_signature.gen_signature(ROOT / "src" / "acp" / "experimental" / "v2", protocol_version=2) if args.format_output: format_generated_files(protocol_version) @@ -116,6 +118,9 @@ def format_generated_files(protocol_version: int) -> None: files = [ ROOT / "src" / "acp" / "experimental" / "v2" / "schema.py", ROOT / "src" / "acp" / "experimental" / "v2" / "meta.py", + ROOT / "src" / "acp" / "experimental" / "v2" / "interfaces.py", + ROOT / "src" / "acp" / "experimental" / "v2" / "agent.py", + ROOT / "src" / "acp" / "experimental" / "v2" / "client.py", ] subprocess.check_call([sys.executable, "-m", "ruff", "check", "--fix", *(str(path) for path in files)]) # noqa: S603 subprocess.check_call([sys.executable, "-m", "ruff", "format", *(str(path) for path in files)]) # noqa: S603 diff --git a/scripts/gen_signature.py b/scripts/gen_signature.py index cec2e43..ef6c171 100644 --- a/scripts/gen_signature.py +++ b/scripts/gen_signature.py @@ -6,7 +6,7 @@ import typing as t from pathlib import Path -from pydantic import BaseModel +from pydantic import AnyUrl, BaseModel from pydantic.fields import FieldInfo from pydantic_core import PydanticUndefined @@ -33,10 +33,14 @@ def _load_schema_module() -> t.Any: class NodeTransformer(ast.NodeTransformer): - def __init__(self) -> None: + def __init__(self, schema_module: t.Any = None) -> None: + self._schema = schema_module if schema_module is not None else schema + self._qualified_schema: str | None = None self._type_import_node: ast.ImportFrom | None = None self._schema_import_node: ast.ImportFrom | None = None - self._literals = {name: value for name, value in schema.__dict__.items() if t.get_origin(value) is t.Literal} + self._literals = { + name: value for name, value in self._schema.__dict__.items() if t.get_origin(value) is t.Literal + } self._current_model_name: str | None = None self._type_aliases: dict[str, ast.expr] = {} self._schema_names: dict[str, str] = {} @@ -92,6 +96,9 @@ def visit_ImportFrom(self, node: ast.ImportFrom) -> ast.AST: ) elif node.module is None: self._schema_modules.update(alias.asname or alias.name for alias in node.names if alias.name == "schema") + self._qualified_schema = next( + (alias.asname or alias.name for alias in node.names if alias.name == "schema"), None + ) return node def _single_param_model(self, expression: ast.expr, seen: frozenset[str] = frozenset()) -> t.Any: @@ -106,11 +113,11 @@ def _single_param_model(self, expression: ast.expr, seen: frozenset[str] = froze return None if name in self._type_aliases: return self._single_param_model(self._type_aliases[name], seen | {name}) - model = getattr(schema, self._schema_names.get(name, name), None) + model = getattr(self._schema, self._schema_names.get(name, name), None) elif isinstance(expression, ast.Attribute) and isinstance(expression.value, ast.Name): if expression.value.id not in self._schema_modules: return None - model = getattr(schema, expression.attr, None) + model = getattr(self._schema, expression.attr, None) elif isinstance(expression, ast.Subscript): name = ast.unparse(expression.value) if name not in self._annotated_names and name != "typing.Annotated": @@ -165,7 +172,7 @@ def visit_func(self, node: ast.FunctionDef | ast.AsyncFunctionDef) -> ast.AST: def _to_param_def(self, name: str, field: FieldInfo) -> tuple[ast.arg, ast.expr | None]: arg = ast.arg(arg=name) ann = field.annotation - override_optional = (self._current_model_name, name) in SIGNATURE_OPTIONAL_FIELDS + override_optional = self._schema is schema and (self._current_model_name, name) in SIGNATURE_OPTIONAL_FIELDS if override_optional: if ann is not None: ann = ann | None @@ -196,10 +203,10 @@ def _format_annotation(self, annotation: t.Any) -> ast.expr: elif ( inspect.isclass(annotation) and issubclass(annotation, BaseModel) - and annotation.__module__ == schema.__name__ + and annotation.__module__ == self._schema.__name__ ): self._add_schema_import(annotation.__name__) - return ast.Name(id=annotation.__name__) + return self._schema_reference(annotation.__name__) elif args := t.get_args(annotation): return ast.Subscript( value=self._format_annotation(origin), @@ -208,6 +215,11 @@ def _format_annotation(self, annotation: t.Any) -> ast.expr: else self._format_annotation(args[0]), ctx=ast.Load(), ) + return self._format_scalar_annotation(annotation) + + def _format_scalar_annotation(self, annotation: t.Any) -> ast.expr: + if annotation is AnyUrl: + return ast.parse("str | AnyUrl", mode="eval").body elif annotation.__module__ == "typing": name = annotation.__name__ self._add_typing_import(name) @@ -221,11 +233,16 @@ def _format_annotation(self, annotation: t.Any) -> ast.expr: self._add_typing_import("Any") return ast.Name(id="Any") + def _schema_reference(self, name: str) -> ast.expr: + if self._qualified_schema: + return ast.Attribute(value=ast.Name(id=self._qualified_schema), attr=name) + return ast.Name(id=name) + 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) + return self._schema_reference(name) self._add_typing_import("Literal") values = [ast.Constant(value=value) for value in t.get_args(annotation)] return ast.Subscript( @@ -235,9 +252,15 @@ def _format_literal(self, annotation: t.Any) -> ast.expr: ) -def gen_signature(source_dir: Path) -> None: +def gen_signature(source_dir: Path, *, protocol_version: int = 1) -> None: global schema schema = _load_schema_module() + if protocol_version == 2: + from acp.experimental.v2 import schema as version_schema + else: + version_schema = schema for source_file in source_dir.rglob("*.py"): - transformer = NodeTransformer() + if protocol_version == 1 and "experimental" in source_file.relative_to(source_dir).parts: + continue + transformer = NodeTransformer(version_schema) transformer.transform(source_file) diff --git a/src/acp/experimental/v2/__init__.py b/src/acp/experimental/v2/__init__.py index d0af723..724b1d3 100644 --- a/src/acp/experimental/v2/__init__.py +++ b/src/acp/experimental/v2/__init__.py @@ -3,11 +3,14 @@ from . import schema from .agent import AgentSideConnection, run_agent from .client import ClientSideConnection, connect_to_agent +from .interfaces import Agent, Client from .meta import PROTOCOL_VERSION __all__ = [ "PROTOCOL_VERSION", + "Agent", "AgentSideConnection", + "Client", "ClientSideConnection", "connect_to_agent", "run_agent", diff --git a/src/acp/experimental/v2/_methods.py b/src/acp/experimental/v2/_methods.py index ba72864..3ecd787 100644 --- a/src/acp/experimental/v2/_methods.py +++ b/src/acp/experimental/v2/_methods.py @@ -1,12 +1,12 @@ from __future__ import annotations from dataclasses import dataclass -from typing import Any +from functools import cache +from typing import Any, get_type_hints from pydantic import TypeAdapter -from . import schema -from .meta import AGENT_METHODS, CLIENT_METHODS +from .interfaces import Agent, Client @dataclass(frozen=True, slots=True) @@ -25,175 +25,34 @@ class NotificationSpec: params: TypeAdapter[Any] -def request( - method: str, - handler: str, - request_type: Any, - response_type: Any, - *, - empty_response: bool = False, -) -> RequestSpec: - return RequestSpec( - method=method, - handler=handler, - request=TypeAdapter(request_type), - response=TypeAdapter(response_type), - empty_response=empty_response, - ) - - -def notification(method: str, handler: str, params_type: Any) -> NotificationSpec: - return NotificationSpec(method=method, handler=handler, params=TypeAdapter(params_type)) - - -SetConfigOptionRequest = ( - schema.SetSessionConfigOptionIdRequest - | schema.SetSessionConfigOptionBooleanRequest - | schema.SetSessionConfigOptionOtherRequest -) - -CreateElicitationRequest = ( - schema.CreateOtherSessionElicitationRequest - | schema.CreateOtherRequestElicitationRequest - | schema.CreateFormSessionElicitationRequest - | schema.CreateFormRequestElicitationRequest - | schema.CreateUrlSessionElicitationRequest - | schema.CreateUrlRequestElicitationRequest -) - -CreateElicitationResponse = ( - schema.AcceptElicitationResponse - | schema.DeclineElicitationResponse - | schema.CancelElicitationResponse - | schema.OtherElicitationResponse -) - - -AGENT_REQUESTS = ( - request(AGENT_METHODS["initialize"], "initialize", schema.InitializeRequest, schema.InitializeResponse), - request( - AGENT_METHODS["auth_login"], - "login", - schema.LoginAuthRequest, - schema.LoginAuthResponse, - empty_response=True, - ), - request( - AGENT_METHODS["providers_list"], "list_providers", schema.ListProvidersRequest, schema.ListProvidersResponse - ), - request( - AGENT_METHODS["providers_set"], - "set_provider", - schema.SetProviderRequest, - schema.SetProviderResponse, - empty_response=True, - ), - request( - AGENT_METHODS["providers_disable"], - "disable_provider", - schema.DisableProviderRequest, - schema.DisableProviderResponse, - empty_response=True, - ), - request(AGENT_METHODS["session_new"], "new_session", schema.NewSessionRequest, schema.NewSessionResponse), - request( - AGENT_METHODS["session_set_config_option"], - "set_config_option", - SetConfigOptionRequest, - schema.SetSessionConfigOptionResponse, - ), - request( - AGENT_METHODS["session_prompt"], - "prompt", - schema.PromptRequest, - schema.PromptResponse, - ), - request(AGENT_METHODS["mcp_message"], "mcp_message", schema.MessageMcpRequest, Any), - request(AGENT_METHODS["session_list"], "list_sessions", schema.ListSessionsRequest, schema.ListSessionsResponse), - request( - AGENT_METHODS["session_delete"], - "delete_session", - schema.DeleteSessionRequest, - schema.DeleteSessionResponse, - empty_response=True, - ), - request(AGENT_METHODS["session_fork"], "fork_session", schema.ForkSessionRequest, schema.ForkSessionResponse), - request( - AGENT_METHODS["session_resume"], "resume_session", schema.ResumeSessionRequest, schema.ResumeSessionResponse - ), - request( - AGENT_METHODS["session_close"], - "close_session", - schema.CloseSessionRequest, - schema.CloseSessionResponse, - empty_response=True, - ), - request( - AGENT_METHODS["auth_logout"], - "logout", - schema.LogoutAuthRequest, - schema.LogoutAuthResponse, - empty_response=True, - ), - request(AGENT_METHODS["nes_start"], "start_nes", schema.StartNesRequest, schema.StartNesResponse), - request(AGENT_METHODS["nes_suggest"], "suggest_nes", schema.SuggestNesRequest, schema.SuggestNesResponse), - request( - AGENT_METHODS["nes_close"], - "close_nes", - schema.CloseNesRequest, - schema.CloseNesResponse, - empty_response=True, - ), -) - -AGENT_NOTIFICATIONS = ( - notification(AGENT_METHODS["session_cancel"], "cancel_session", schema.CancelSessionNotification), - notification(AGENT_METHODS["mcp_message"], "notify_mcp", schema.MessageMcpNotification), - notification(AGENT_METHODS["document_did_open"], "did_open", schema.DidOpenDocumentNotification), - notification(AGENT_METHODS["document_did_change"], "did_change", schema.DidChangeDocumentNotification), - notification(AGENT_METHODS["document_did_close"], "did_close", schema.DidCloseDocumentNotification), - notification(AGENT_METHODS["document_did_save"], "did_save", schema.DidSaveDocumentNotification), - notification(AGENT_METHODS["document_did_focus"], "did_focus", schema.DidFocusDocumentNotification), - notification(AGENT_METHODS["nes_accept"], "accept_nes", schema.AcceptNesNotification), - notification(AGENT_METHODS["nes_reject"], "reject_nes", schema.RejectNesNotification), -) - -CLIENT_REQUESTS = ( - request( - CLIENT_METHODS["session_request_permission"], - "request_permission", - schema.RequestPermissionRequest, - schema.RequestPermissionResponse, - ), - request(CLIENT_METHODS["mcp_connect"], "connect_mcp", schema.ConnectMcpRequest, schema.ConnectMcpResponse), - request(CLIENT_METHODS["mcp_message"], "mcp_message", schema.MessageMcpRequest, Any), - request( - CLIENT_METHODS["mcp_disconnect"], - "disconnect_mcp", - schema.DisconnectMcpRequest, - schema.DisconnectMcpResponse, - empty_response=True, - ), - request( - CLIENT_METHODS["elicitation_create"], - "create_elicitation", - CreateElicitationRequest, - CreateElicitationResponse, - ), -) - -CLIENT_NOTIFICATIONS = ( - notification(CLIENT_METHODS["session_update"], "session_update", schema.UpdateSessionNotification), - notification(CLIENT_METHODS["mcp_message"], "notify_mcp", schema.MessageMcpNotification), - notification( - CLIENT_METHODS["elicitation_complete"], - "complete_elicitation", - schema.CompleteElicitationNotification, - ), -) - - +@cache +def protocol_specs(protocol: type) -> tuple[tuple[RequestSpec, ...], tuple[NotificationSpec, ...]]: + """Compile v2 wire validation from the public protocol declarations.""" + requests: dict[str, RequestSpec] = {} + notifications: dict[str, NotificationSpec] = {} + for name in dir(protocol): + handler = getattr(protocol, name) + metadata = getattr(handler, "__route__", None) + if metadata is None: + continue + if metadata.kind == "notification": + if metadata.method in notifications: + raise ValueError(f"Duplicate notification: {metadata.method}") + notifications[metadata.method] = NotificationSpec(metadata.method, name, TypeAdapter(metadata.request_type)) + else: + if metadata.method in requests: + raise ValueError(f"Duplicate request: {metadata.method}") + requests[metadata.method] = RequestSpec( + metadata.method, + name, + TypeAdapter(metadata.request_type), + TypeAdapter(get_type_hints(handler, include_extras=True)["return"]), + empty_response=metadata.default_result == {}, + ) + return tuple(requests.values()), tuple(notifications.values()) + + +AGENT_REQUESTS, AGENT_NOTIFICATIONS = protocol_specs(Agent) +CLIENT_REQUESTS, CLIENT_NOTIFICATIONS = protocol_specs(Client) AGENT_REQUESTS_BY_METHOD = {spec.method: spec for spec in AGENT_REQUESTS} -AGENT_NOTIFICATIONS_BY_METHOD = {spec.method: spec for spec in AGENT_NOTIFICATIONS} CLIENT_REQUESTS_BY_METHOD = {spec.method: spec for spec in CLIENT_REQUESTS} -CLIENT_NOTIFICATIONS_BY_METHOD = {spec.method: spec for spec in CLIENT_NOTIFICATIONS} diff --git a/src/acp/experimental/v2/_params.py b/src/acp/experimental/v2/_params.py new file mode 100644 index 0000000..6b21a3b --- /dev/null +++ b/src/acp/experimental/v2/_params.py @@ -0,0 +1,65 @@ +from __future__ import annotations + +from typing import Any, TypeVar + +from pydantic import AnyUrl, BaseModel, TypeAdapter + +from . import schema +from .interfaces import CreateElicitationRequest, SetConfigOptionRequest + +ModelT = TypeVar("ModelT", bound=BaseModel) +_CONFIG = TypeAdapter(SetConfigOptionRequest) +_ELICITATION = TypeAdapter(CreateElicitationRequest) + + +def build_request(model: type[ModelT], fields: dict[str, Any], meta: dict[str, Any]) -> ModelT: + """Build a request without changing nested models' explicitly set fields.""" + params = { + name: value for name, value in fields.items() if value is not None or model.model_fields[name].is_required() + } + if meta: + params["field_meta"] = meta + return model.model_validate(params) + + +def build_config_request( + config_id: str, session_id: str, value: Any, config_type: str | None, meta: dict[str, Any] +) -> SetConfigOptionRequest: + return _CONFIG.validate_python({ + "configId": config_id, + "sessionId": session_id, + "value": value, + "type": config_type if config_type is not None else ("boolean" if isinstance(value, bool) else "id"), + "_meta": meta or None, + }) + + +def build_elicitation_request( + *, + message: str, + mode: str, + session_id: str | None, + request_id: int | str | None, + tool_call_id: str | None, + requested_schema: schema.ElicitationSchema | None, + elicitation_id: str | None, + url: str | AnyUrl | None, + meta: dict[str, Any], +) -> CreateElicitationRequest: + if session_id is not None and request_id is not None: + raise ValueError("Specify either session_id or request_id") + if session_id is None and tool_call_id is not None: + raise ValueError("tool_call_id requires session_id") + params: dict[str, Any] = {"message": message, "mode": mode} + if session_id is not None: + params["sessionId"] = session_id + if tool_call_id is not None: + params["toolCallId"] = tool_call_id + else: + params["requestId"] = request_id + for name, value in (("requestedSchema", requested_schema), ("elicitationId", elicitation_id), ("url", url)): + if value is not None: + params[name] = value + if meta: + params["_meta"] = meta + return _ELICITATION.validate_python(params) diff --git a/src/acp/experimental/v2/_router.py b/src/acp/experimental/v2/_router.py index 57c0860..4afb0e7 100644 --- a/src/acp/experimental/v2/_router.py +++ b/src/acp/experimental/v2/_router.py @@ -4,8 +4,9 @@ from typing import Any from acp.exceptions import RequestError +from acp.utils import model_to_kwargs -from ._methods import NotificationSpec, RequestSpec +from ._methods import NotificationSpec, RequestSpec, protocol_specs ExtensionRequest = Callable[[str, Any], Awaitable[Any]] ExtensionNotification = Callable[[str, Any], Awaitable[None]] @@ -15,31 +16,39 @@ class MethodRouter: def __init__( self, target: Any, - requests: tuple[RequestSpec, ...], - notifications: tuple[NotificationSpec, ...], + protocol: type, ) -> None: self._target = target + self._protocol = protocol + requests, notifications = protocol_specs(protocol) self._requests = {spec.method: spec for spec in requests} self._notifications = {spec.method: spec for spec in notifications} + def _handler(self, name: str) -> Any: + handler = getattr(self._target, name, None) + if getattr(handler, "__func__", handler) is getattr(self._protocol, name, None): + return None + return handler + def request_spec(self, method: str) -> RequestSpec | None: return self._requests.get(method) async def handle_request(self, spec: RequestSpec, params: Any) -> Any: - handler = getattr(self._target, spec.handler, None) + handler = self._handler(spec.handler) if handler is None: raise RequestError.method_not_found(spec.method) request = spec.request.validate_python(params) - response = await handler(request) + response = await handler(**model_to_kwargs(request, type(request))) if response is None and spec.empty_response: response = {} return spec.response.validate_python(response) async def handle_notification(self, spec: NotificationSpec, params: Any) -> None: - handler = getattr(self._target, spec.handler, None) + handler = self._handler(spec.handler) if handler is None: return - await handler(spec.params.validate_python(params)) + request = spec.params.validate_python(params) + await handler(**model_to_kwargs(request, type(request))) async def __call__(self, method: str, params: Any | None, is_notification: bool) -> Any: if method.startswith("_"): diff --git a/src/acp/experimental/v2/agent.py b/src/acp/experimental/v2/agent.py index ddc68c0..da96072 100644 --- a/src/acp/experimental/v2/agent.py +++ b/src/acp/experimental/v2/agent.py @@ -4,21 +4,20 @@ from collections.abc import Callable from typing import Any -from pydantic import BaseModel +from pydantic import AnyUrl, BaseModel from acp.connection import Connection, MethodHandler +from acp.utils import param_model from . import schema from ._connection import open_connection from ._initialization import InitializationState from ._methods import ( - AGENT_NOTIFICATIONS, - AGENT_REQUESTS, CLIENT_REQUESTS_BY_METHOD, - CreateElicitationRequest, - CreateElicitationResponse, ) +from ._params import build_elicitation_request, build_request from ._router import MethodRouter +from .interfaces import Agent, CreateElicitationRequest, CreateElicitationResponse from .meta import CLIENT_METHODS __all__ = ["AgentSideConnection", "run_agent"] @@ -30,7 +29,7 @@ def _dump(model: BaseModel) -> dict[str, Any]: class _AgentRouter: def __init__(self, agent: object, state: InitializationState) -> None: - self._router = MethodRouter(agent, AGENT_REQUESTS, AGENT_NOTIFICATIONS) + self._router = MethodRouter(agent, Agent) self._state = state async def __call__(self, method: str, params: Any | None, is_notification: bool) -> Any: @@ -91,38 +90,139 @@ def attach( async def _listen(self) -> None: await self._conn.main_loop() + @param_model(schema.RequestPermissionRequest) async def request_permission( self, - request: schema.RequestPermissionRequest, + session_id: str, + title: str, + options: list[schema.PermissionOption], + description: str | None = None, + subject: schema.ToolCallPermissionSubjectVariant + | schema.CommandPermissionSubjectVariant + | schema.OtherPermissionSubject + | None = None, + **kwargs: Any, ) -> schema.RequestPermissionResponse: return await self._request( CLIENT_METHODS["session_request_permission"], - request, + build_request( + schema.RequestPermissionRequest, + { + "session_id": session_id, + "title": title, + "options": options, + "description": description, + "subject": subject, + }, + kwargs, + ), ) - async def session_update(self, notification: schema.UpdateSessionNotification) -> None: - await self._notify(CLIENT_METHODS["session_update"], notification) + @param_model(schema.UpdateSessionNotification) + async def session_update( + self, + session_id: str, + update: schema.UserMessageChunk + | schema.UserMessageUpdate + | schema.AgentMessageChunk + | schema.AgentMessageUpdate + | schema.AgentThoughtChunk + | schema.AgentThoughtUpdate + | schema.ToolCallContentChunkUpdate + | schema.SessionToolCallUpdate + | schema.SessionTerminalUpdate + | schema.SessionTerminalOutputChunk + | schema.SessionPlanUpdate + | schema.SessionPlanRemovedUpdate + | schema.AvailableCommandsUpdate + | schema.ConfigOptionUpdate + | schema.SessionInfoUpdate + | schema.UsageUpdate + | schema.SessionNotice + | schema.SessionCompactionUpdate + | schema.SessionCompactionSummaryChunk + | schema.OtherSessionUpdate + | schema.RunningSessionStateUpdate + | schema.IdleSessionStateUpdate + | schema.RequiresActionSessionStateUpdate + | schema.OtherSessionStateUpdate, + **kwargs: Any, + ) -> None: + await self._notify( + CLIENT_METHODS["session_update"], + build_request(schema.UpdateSessionNotification, {"session_id": session_id, "update": update}, kwargs), + ) + + @param_model(schema.ConnectMcpRequest) + async def connect_mcp(self, server_id: str, **kwargs: Any) -> schema.ConnectMcpResponse: + return await self._request( + CLIENT_METHODS["mcp_connect"], build_request(schema.ConnectMcpRequest, {"server_id": server_id}, kwargs) + ) - async def connect_mcp(self, request: schema.ConnectMcpRequest) -> schema.ConnectMcpResponse: - return await self._request(CLIENT_METHODS["mcp_connect"], request) + @param_model(schema.MessageMcpRequest) + async def mcp_message( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> Any: + return await self._request( + CLIENT_METHODS["mcp_message"], + build_request( + schema.MessageMcpRequest, {"connection_id": connection_id, "method": method, "params": params}, kwargs + ), + ) - async def mcp_message(self, message: schema.MessageMcpRequest) -> Any: - return await self._request(CLIENT_METHODS["mcp_message"], message) + @param_model(schema.MessageMcpNotification) + async def notify_mcp( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> None: + await self._notify( + CLIENT_METHODS["mcp_message"], + build_request( + schema.MessageMcpNotification, + {"connection_id": connection_id, "method": method, "params": params}, + kwargs, + ), + ) - async def notify_mcp(self, notification: schema.MessageMcpNotification) -> None: - await self._notify(CLIENT_METHODS["mcp_message"], notification) + @param_model(schema.DisconnectMcpRequest) + async def disconnect_mcp(self, connection_id: str, **kwargs: Any) -> schema.DisconnectMcpResponse: + return await self._request( + CLIENT_METHODS["mcp_disconnect"], + build_request(schema.DisconnectMcpRequest, {"connection_id": connection_id}, kwargs), + ) - async def disconnect_mcp( + @param_model(CreateElicitationRequest) + async def create_elicitation( self, - request: schema.DisconnectMcpRequest, - ) -> schema.DisconnectMcpResponse: - return await self._request(CLIENT_METHODS["mcp_disconnect"], request) - - async def create_elicitation(self, request: CreateElicitationRequest) -> CreateElicitationResponse: + message: str, + mode: str, + *, + session_id: str | None = None, + request_id: int | str | None = None, + tool_call_id: str | None = None, + requested_schema: schema.ElicitationSchema | None = None, + elicitation_id: str | None = None, + url: str | AnyUrl | None = None, + **kwargs: Any, + ) -> CreateElicitationResponse: + request = build_elicitation_request( + message=message, + mode=mode, + session_id=session_id, + request_id=request_id, + tool_call_id=tool_call_id, + requested_schema=requested_schema, + elicitation_id=elicitation_id, + url=url, + meta=kwargs, + ) return await self._request(CLIENT_METHODS["elicitation_create"], request) - async def complete_elicitation(self, notification: schema.CompleteElicitationNotification) -> None: - await self._notify(CLIENT_METHODS["elicitation_complete"], notification) + @param_model(schema.CompleteElicitationNotification) + async def complete_elicitation(self, elicitation_id: str, **kwargs: Any) -> None: + await self._notify( + CLIENT_METHODS["elicitation_complete"], + build_request(schema.CompleteElicitationNotification, {"elicitation_id": elicitation_id}, kwargs), + ) async def send_extension_request(self, method: str, params: Any = None) -> Any: await self._state.require(method) diff --git a/src/acp/experimental/v2/client.py b/src/acp/experimental/v2/client.py index c744b93..ce82acd 100644 --- a/src/acp/experimental/v2/client.py +++ b/src/acp/experimental/v2/client.py @@ -1,20 +1,21 @@ from __future__ import annotations -from typing import Any +from typing import Any, Literal -from pydantic import BaseModel +from pydantic import AnyUrl, BaseModel + +from acp.utils import param_model from . import schema from ._connection import open_connection from ._initialization import InitializationState from ._methods import ( AGENT_REQUESTS_BY_METHOD, - CLIENT_NOTIFICATIONS, - CLIENT_REQUESTS, - SetConfigOptionRequest, ) +from ._params import build_config_request, build_request from ._router import MethodRouter from .agent import _dump, _extension_method +from .interfaces import Client, SetConfigOptionRequest from .meta import AGENT_METHODS __all__ = ["ClientSideConnection", "connect_to_agent"] @@ -22,7 +23,7 @@ class _ClientRouter: def __init__(self, client: object, state: InitializationState) -> None: - self._router = MethodRouter(client, CLIENT_REQUESTS, CLIENT_NOTIFICATIONS) + self._router = MethodRouter(client, Client) self._state = state async def __call__(self, method: str, params: Any | None, is_notification: bool) -> Any: @@ -46,7 +47,19 @@ def __init__( if on_connect := getattr(client, "on_connect", None): on_connect(self) - async def initialize(self, request: schema.InitializeRequest) -> schema.InitializeResponse: + @param_model(schema.InitializeRequest) + async def initialize( + self, + protocol_version: int, + info: schema.Implementation, + capabilities: schema.ClientCapabilities | None = None, + **kwargs: Any, + ) -> schema.InitializeResponse: + request = build_request( + schema.InitializeRequest, + {"protocol_version": protocol_version, "info": info, "capabilities": capabilities}, + kwargs, + ) self._state.begin(request) try: response = await self._conn.send_request(AGENT_METHODS["initialize"], _dump(request)) @@ -58,83 +71,358 @@ async def initialize(self, request: schema.InitializeRequest) -> schema.Initiali raise return parsed - async def login(self, request: schema.LoginAuthRequest) -> schema.LoginAuthResponse: - return await self._request(AGENT_METHODS["auth_login"], request) - - async def logout(self, request: schema.LogoutAuthRequest) -> schema.LogoutAuthResponse: - return await self._request(AGENT_METHODS["auth_logout"], request) - - async def list_providers(self, request: schema.ListProvidersRequest) -> schema.ListProvidersResponse: - return await self._request(AGENT_METHODS["providers_list"], request) - - async def set_provider(self, request: schema.SetProviderRequest) -> schema.SetProviderResponse: - return await self._request(AGENT_METHODS["providers_set"], request) - - async def disable_provider(self, request: schema.DisableProviderRequest) -> schema.DisableProviderResponse: - return await self._request(AGENT_METHODS["providers_disable"], request) - - async def new_session(self, request: schema.NewSessionRequest) -> schema.NewSessionResponse: - return await self._request(AGENT_METHODS["session_new"], request) - - async def list_sessions(self, request: schema.ListSessionsRequest) -> schema.ListSessionsResponse: - return await self._request(AGENT_METHODS["session_list"], request) - - async def delete_session(self, request: schema.DeleteSessionRequest) -> schema.DeleteSessionResponse: - return await self._request(AGENT_METHODS["session_delete"], request) - - async def fork_session(self, request: schema.ForkSessionRequest) -> schema.ForkSessionResponse: - return await self._request(AGENT_METHODS["session_fork"], request) + @param_model(schema.LoginAuthRequest) + async def login(self, method_id: str, **kwargs: Any) -> schema.LoginAuthResponse: + return await self._request( + AGENT_METHODS["auth_login"], build_request(schema.LoginAuthRequest, {"method_id": method_id}, kwargs) + ) - async def resume_session(self, request: schema.ResumeSessionRequest) -> schema.ResumeSessionResponse: - return await self._request(AGENT_METHODS["session_resume"], request) + @param_model(schema.LogoutAuthRequest) + async def logout(self, **kwargs: Any) -> schema.LogoutAuthResponse: + return await self._request(AGENT_METHODS["auth_logout"], build_request(schema.LogoutAuthRequest, {}, kwargs)) - async def close_session(self, request: schema.CloseSessionRequest) -> schema.CloseSessionResponse: - return await self._request(AGENT_METHODS["session_close"], request) + @param_model(schema.ListProvidersRequest) + async def list_providers(self, **kwargs: Any) -> schema.ListProvidersResponse: + return await self._request( + AGENT_METHODS["providers_list"], build_request(schema.ListProvidersRequest, {}, kwargs) + ) - async def set_config_option(self, request: SetConfigOptionRequest) -> schema.SetSessionConfigOptionResponse: + @param_model(schema.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 | AnyUrl, + headers: dict[str, str] | None = None, + **kwargs: Any, + ) -> schema.SetProviderResponse: + return await self._request( + AGENT_METHODS["providers_set"], + build_request( + schema.SetProviderRequest, + {"provider_id": provider_id, "api_type": api_type, "base_url": base_url, "headers": headers}, + kwargs, + ), + ) + + @param_model(schema.DisableProviderRequest) + async def disable_provider(self, provider_id: str, **kwargs: Any) -> schema.DisableProviderResponse: + return await self._request( + AGENT_METHODS["providers_disable"], + build_request(schema.DisableProviderRequest, {"provider_id": provider_id}, kwargs), + ) + + @param_model(schema.NewSessionRequest) + async def new_session( + self, + cwd: str, + additional_directories: list[str] | None = None, + mcp_servers: list[schema.HttpMcpServer | schema.AcpMcpServer | schema.StdioMcpServer | schema.OtherMcpServer] + | None = None, + **kwargs: Any, + ) -> schema.NewSessionResponse: + return await self._request( + AGENT_METHODS["session_new"], + build_request( + schema.NewSessionRequest, + {"cwd": cwd, "additional_directories": additional_directories, "mcp_servers": mcp_servers}, + kwargs, + ), + ) + + @param_model(schema.ListSessionsRequest) + async def list_sessions( + self, cwd: str | None = None, cursor: str | None = None, **kwargs: Any + ) -> schema.ListSessionsResponse: + return await self._request( + AGENT_METHODS["session_list"], + build_request(schema.ListSessionsRequest, {"cwd": cwd, "cursor": cursor}, kwargs), + ) + + @param_model(schema.DeleteSessionRequest) + async def delete_session(self, session_id: str, **kwargs: Any) -> schema.DeleteSessionResponse: + return await self._request( + AGENT_METHODS["session_delete"], + build_request(schema.DeleteSessionRequest, {"session_id": session_id}, kwargs), + ) + + @param_model(schema.ForkSessionRequest) + async def fork_session( + self, + session_id: str, + cwd: str, + additional_directories: list[str] | None = None, + mcp_servers: list[schema.HttpMcpServer | schema.AcpMcpServer | schema.StdioMcpServer | schema.OtherMcpServer] + | None = None, + **kwargs: Any, + ) -> schema.ForkSessionResponse: + return await self._request( + AGENT_METHODS["session_fork"], + build_request( + schema.ForkSessionRequest, + { + "session_id": session_id, + "cwd": cwd, + "additional_directories": additional_directories, + "mcp_servers": mcp_servers, + }, + kwargs, + ), + ) + + @param_model(schema.ResumeSessionRequest) + async def resume_session( + self, + session_id: str, + cwd: str, + additional_directories: list[str] | None = None, + mcp_servers: list[schema.HttpMcpServer | schema.AcpMcpServer | schema.StdioMcpServer | schema.OtherMcpServer] + | None = None, + replay_from: schema.ReplayFromStartVariant | schema.OtherReplayFrom | None = None, + **kwargs: Any, + ) -> schema.ResumeSessionResponse: + return await self._request( + AGENT_METHODS["session_resume"], + build_request( + schema.ResumeSessionRequest, + { + "session_id": session_id, + "cwd": cwd, + "additional_directories": additional_directories, + "mcp_servers": mcp_servers, + "replay_from": replay_from, + }, + kwargs, + ), + ) + + @param_model(schema.CloseSessionRequest) + async def close_session(self, session_id: str, **kwargs: Any) -> schema.CloseSessionResponse: + return await self._request( + AGENT_METHODS["session_close"], + build_request(schema.CloseSessionRequest, {"session_id": session_id}, kwargs), + ) + + @param_model(SetConfigOptionRequest) + async def set_config_option( + self, + config_id: str, + session_id: str, + value: Any, + *, + type: str | None = None, # noqa: A002 + **kwargs: Any, + ) -> schema.SetSessionConfigOptionResponse: + request = build_config_request(config_id, session_id, value, type, kwargs) return await self._request(AGENT_METHODS["session_set_config_option"], request) - async def prompt(self, request: schema.PromptRequest) -> schema.PromptResponse: - return await self._request(AGENT_METHODS["session_prompt"], request) - - async def cancel_session(self, notification: schema.CancelSessionNotification) -> None: - await self._notify(AGENT_METHODS["session_cancel"], notification) - - async def mcp_message(self, message: schema.MessageMcpRequest) -> Any: - return await self._request(AGENT_METHODS["mcp_message"], message) - - async def notify_mcp(self, notification: schema.MessageMcpNotification) -> None: - await self._notify(AGENT_METHODS["mcp_message"], notification) - - async def start_nes(self, request: schema.StartNesRequest) -> schema.StartNesResponse: - return await self._request(AGENT_METHODS["nes_start"], request) - - async def suggest_nes(self, request: schema.SuggestNesRequest) -> schema.SuggestNesResponse: - return await self._request(AGENT_METHODS["nes_suggest"], request) - - async def accept_nes(self, notification: schema.AcceptNesNotification) -> None: - await self._notify(AGENT_METHODS["nes_accept"], notification) - - async def reject_nes(self, notification: schema.RejectNesNotification) -> None: - await self._notify(AGENT_METHODS["nes_reject"], notification) - - async def close_nes(self, request: schema.CloseNesRequest) -> schema.CloseNesResponse: - return await self._request(AGENT_METHODS["nes_close"], request) - - async def did_open(self, notification: schema.DidOpenDocumentNotification) -> None: - await self._notify(AGENT_METHODS["document_did_open"], notification) - - async def did_change(self, notification: schema.DidChangeDocumentNotification) -> None: - await self._notify(AGENT_METHODS["document_did_change"], notification) - - async def did_close(self, notification: schema.DidCloseDocumentNotification) -> None: - await self._notify(AGENT_METHODS["document_did_close"], notification) - - async def did_save(self, notification: schema.DidSaveDocumentNotification) -> None: - await self._notify(AGENT_METHODS["document_did_save"], notification) - - async def did_focus(self, notification: schema.DidFocusDocumentNotification) -> None: - await self._notify(AGENT_METHODS["document_did_focus"], notification) + @param_model(schema.PromptRequest) + async def prompt( + self, + session_id: str, + prompt: list[ + schema.TextContentBlock + | schema.ImageContentBlock + | schema.AudioContentBlock + | schema.ResourceContentBlock + | schema.EmbeddedResourceContentBlock + | schema.OtherContentBlock + ], + **kwargs: Any, + ) -> schema.PromptResponse: + return await self._request( + AGENT_METHODS["session_prompt"], + build_request(schema.PromptRequest, {"session_id": session_id, "prompt": prompt}, kwargs), + ) + + @param_model(schema.CancelSessionNotification) + async def cancel_session(self, session_id: str, **kwargs: Any) -> None: + await self._notify( + AGENT_METHODS["session_cancel"], + build_request(schema.CancelSessionNotification, {"session_id": session_id}, kwargs), + ) + + @param_model(schema.MessageMcpRequest) + async def mcp_message( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> Any: + return await self._request( + AGENT_METHODS["mcp_message"], + build_request( + schema.MessageMcpRequest, {"connection_id": connection_id, "method": method, "params": params}, kwargs + ), + ) + + @param_model(schema.MessageMcpNotification) + async def notify_mcp( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> None: + await self._notify( + AGENT_METHODS["mcp_message"], + build_request( + schema.MessageMcpNotification, + {"connection_id": connection_id, "method": method, "params": params}, + kwargs, + ), + ) + + @param_model(schema.StartNesRequest) + async def start_nes( + self, + workspace_uri: str | AnyUrl | None = None, + workspace_folders: list[schema.WorkspaceFolder] | None = None, + repository: schema.NesRepository | None = None, + **kwargs: Any, + ) -> schema.StartNesResponse: + return await self._request( + AGENT_METHODS["nes_start"], + build_request( + schema.StartNesRequest, + {"workspace_uri": workspace_uri, "workspace_folders": workspace_folders, "repository": repository}, + kwargs, + ), + ) + + @param_model(schema.SuggestNesRequest) + async def suggest_nes( + self, + session_id: str, + uri: str | AnyUrl, + version: int, + position: schema.Position, + trigger_kind: Literal["automatic"] | Literal["diagnostic"] | Literal["manual"] | str, + selection: schema.Range | None = None, + context: schema.NesSuggestContext | None = None, + **kwargs: Any, + ) -> schema.SuggestNesResponse: + return await self._request( + AGENT_METHODS["nes_suggest"], + build_request( + schema.SuggestNesRequest, + { + "session_id": session_id, + "uri": uri, + "version": version, + "position": position, + "trigger_kind": trigger_kind, + "selection": selection, + "context": context, + }, + kwargs, + ), + ) + + @param_model(schema.AcceptNesNotification) + async def accept_nes(self, session_id: str, suggestion_id: str, **kwargs: Any) -> None: + await self._notify( + AGENT_METHODS["nes_accept"], + build_request( + schema.AcceptNesNotification, {"session_id": session_id, "suggestion_id": suggestion_id}, kwargs + ), + ) + + @param_model(schema.RejectNesNotification) + async def reject_nes( + self, + session_id: str, + suggestion_id: str, + reason: Literal["rejected"] + | Literal["ignored"] + | Literal["replaced"] + | Literal["cancelled"] + | str + | None = None, + **kwargs: Any, + ) -> None: + await self._notify( + AGENT_METHODS["nes_reject"], + build_request( + schema.RejectNesNotification, + {"session_id": session_id, "suggestion_id": suggestion_id, "reason": reason}, + kwargs, + ), + ) + + @param_model(schema.CloseNesRequest) + async def close_nes(self, session_id: str, **kwargs: Any) -> schema.CloseNesResponse: + return await self._request( + AGENT_METHODS["nes_close"], build_request(schema.CloseNesRequest, {"session_id": session_id}, kwargs) + ) + + @param_model(schema.DidOpenDocumentNotification) + async def did_open( + self, session_id: str, uri: str | AnyUrl, language_id: str, version: int, text: str, **kwargs: Any + ) -> None: + await self._notify( + AGENT_METHODS["document_did_open"], + build_request( + schema.DidOpenDocumentNotification, + {"session_id": session_id, "uri": uri, "language_id": language_id, "version": version, "text": text}, + kwargs, + ), + ) + + @param_model(schema.DidChangeDocumentNotification) + async def did_change( + self, + session_id: str, + uri: str | AnyUrl, + version: int, + content_changes: list[schema.TextDocumentContentChangeEvent], + **kwargs: Any, + ) -> None: + await self._notify( + AGENT_METHODS["document_did_change"], + build_request( + schema.DidChangeDocumentNotification, + {"session_id": session_id, "uri": uri, "version": version, "content_changes": content_changes}, + kwargs, + ), + ) + + @param_model(schema.DidCloseDocumentNotification) + async def did_close(self, session_id: str, uri: str | AnyUrl, **kwargs: Any) -> None: + await self._notify( + AGENT_METHODS["document_did_close"], + build_request(schema.DidCloseDocumentNotification, {"session_id": session_id, "uri": uri}, kwargs), + ) + + @param_model(schema.DidSaveDocumentNotification) + async def did_save(self, session_id: str, uri: str | AnyUrl, **kwargs: Any) -> None: + await self._notify( + AGENT_METHODS["document_did_save"], + build_request(schema.DidSaveDocumentNotification, {"session_id": session_id, "uri": uri}, kwargs), + ) + + @param_model(schema.DidFocusDocumentNotification) + async def did_focus( + self, + session_id: str, + uri: str | AnyUrl, + version: int, + position: schema.Position, + visible_range: schema.Range, + **kwargs: Any, + ) -> None: + await self._notify( + AGENT_METHODS["document_did_focus"], + build_request( + schema.DidFocusDocumentNotification, + { + "session_id": session_id, + "uri": uri, + "version": version, + "position": position, + "visible_range": visible_range, + }, + kwargs, + ), + ) async def send_extension_request(self, method: str, params: Any = None) -> Any: await self._state.require(method) diff --git a/src/acp/experimental/v2/interfaces.py b/src/acp/experimental/v2/interfaces.py new file mode 100644 index 0000000..36692e1 --- /dev/null +++ b/src/acp/experimental/v2/interfaces.py @@ -0,0 +1,314 @@ +from __future__ import annotations + +from typing import Any, Literal, Protocol + +from pydantic import AnyUrl + +from acp.utils import param_model + +from . import schema +from .meta import AGENT_METHODS, CLIENT_METHODS + +SetConfigOptionRequest = ( + schema.SetSessionConfigOptionIdRequest + | schema.SetSessionConfigOptionBooleanRequest + | schema.SetSessionConfigOptionOtherRequest +) + +CreateElicitationRequest = ( + schema.CreateOtherSessionElicitationRequest + | schema.CreateOtherRequestElicitationRequest + | schema.CreateFormSessionElicitationRequest + | schema.CreateFormRequestElicitationRequest + | schema.CreateUrlSessionElicitationRequest + | schema.CreateUrlRequestElicitationRequest +) + +CreateElicitationResponse = ( + schema.AcceptElicitationResponse + | schema.DeclineElicitationResponse + | schema.CancelElicitationResponse + | schema.OtherElicitationResponse +) + + +class Agent(Protocol): + """Agent handlers for ACP v2; only override the methods you support.""" + + @param_model(schema.InitializeRequest, method=AGENT_METHODS["initialize"]) + async def initialize( + self, + protocol_version: int, + info: schema.Implementation, + capabilities: schema.ClientCapabilities | None = None, + **kwargs: Any, + ) -> schema.InitializeResponse: ... + + @param_model(schema.LoginAuthRequest, method=AGENT_METHODS["auth_login"], default_result={}) + async def login(self, method_id: str, **kwargs: Any) -> schema.LoginAuthResponse: ... + + @param_model(schema.LogoutAuthRequest, method=AGENT_METHODS["auth_logout"], default_result={}) + async def logout(self, **kwargs: Any) -> schema.LogoutAuthResponse: ... + + @param_model(schema.ListProvidersRequest, method=AGENT_METHODS["providers_list"]) + async def list_providers(self, **kwargs: Any) -> schema.ListProvidersResponse: ... + + @param_model(schema.SetProviderRequest, method=AGENT_METHODS["providers_set"], default_result={}) + async def set_provider( + self, + provider_id: str, + api_type: Literal["anthropic"] + | Literal["openai"] + | Literal["azure"] + | Literal["vertex"] + | Literal["bedrock"] + | str, + base_url: str | AnyUrl, + headers: dict[str, str] | None = None, + **kwargs: Any, + ) -> schema.SetProviderResponse: ... + + @param_model(schema.DisableProviderRequest, method=AGENT_METHODS["providers_disable"], default_result={}) + async def disable_provider(self, provider_id: str, **kwargs: Any) -> schema.DisableProviderResponse: ... + + @param_model(schema.NewSessionRequest, method=AGENT_METHODS["session_new"]) + async def new_session( + self, + cwd: str, + additional_directories: list[str] | None = None, + mcp_servers: list[schema.HttpMcpServer | schema.AcpMcpServer | schema.StdioMcpServer | schema.OtherMcpServer] + | None = None, + **kwargs: Any, + ) -> schema.NewSessionResponse: ... + + @param_model(schema.ListSessionsRequest, method=AGENT_METHODS["session_list"]) + async def list_sessions( + self, cwd: str | None = None, cursor: str | None = None, **kwargs: Any + ) -> schema.ListSessionsResponse: ... + + @param_model(schema.DeleteSessionRequest, method=AGENT_METHODS["session_delete"], default_result={}) + async def delete_session(self, session_id: str, **kwargs: Any) -> schema.DeleteSessionResponse: ... + + @param_model(schema.ForkSessionRequest, method=AGENT_METHODS["session_fork"]) + async def fork_session( + self, + session_id: str, + cwd: str, + additional_directories: list[str] | None = None, + mcp_servers: list[schema.HttpMcpServer | schema.AcpMcpServer | schema.StdioMcpServer | schema.OtherMcpServer] + | None = None, + **kwargs: Any, + ) -> schema.ForkSessionResponse: ... + + @param_model(schema.ResumeSessionRequest, method=AGENT_METHODS["session_resume"]) + async def resume_session( + self, + session_id: str, + cwd: str, + additional_directories: list[str] | None = None, + mcp_servers: list[schema.HttpMcpServer | schema.AcpMcpServer | schema.StdioMcpServer | schema.OtherMcpServer] + | None = None, + replay_from: schema.ReplayFromStartVariant | schema.OtherReplayFrom | None = None, + **kwargs: Any, + ) -> schema.ResumeSessionResponse: ... + + @param_model(schema.CloseSessionRequest, method=AGENT_METHODS["session_close"], default_result={}) + async def close_session(self, session_id: str, **kwargs: Any) -> schema.CloseSessionResponse: ... + + @param_model(SetConfigOptionRequest, method=AGENT_METHODS["session_set_config_option"]) + async def set_config_option( + self, + config_id: str, + session_id: str, + value: Any, + *, + type: str | None = None, # noqa: A002 + **kwargs: Any, + ) -> schema.SetSessionConfigOptionResponse: ... + + @param_model(schema.PromptRequest, method=AGENT_METHODS["session_prompt"]) + async def prompt( + self, + session_id: str, + prompt: list[ + schema.TextContentBlock + | schema.ImageContentBlock + | schema.AudioContentBlock + | schema.ResourceContentBlock + | schema.EmbeddedResourceContentBlock + | schema.OtherContentBlock + ], + **kwargs: Any, + ) -> schema.PromptResponse: ... + + @param_model(schema.CancelSessionNotification, method=AGENT_METHODS["session_cancel"], kind="notification") + async def cancel_session(self, session_id: str, **kwargs: Any) -> None: ... + + @param_model(schema.MessageMcpRequest, method=AGENT_METHODS["mcp_message"]) + async def mcp_message( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> Any: ... + + @param_model(schema.MessageMcpNotification, method=AGENT_METHODS["mcp_message"], kind="notification") + async def notify_mcp( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> None: ... + + @param_model(schema.StartNesRequest, method=AGENT_METHODS["nes_start"]) + async def start_nes( + self, + workspace_uri: str | AnyUrl | None = None, + workspace_folders: list[schema.WorkspaceFolder] | None = None, + repository: schema.NesRepository | None = None, + **kwargs: Any, + ) -> schema.StartNesResponse: ... + + @param_model(schema.SuggestNesRequest, method=AGENT_METHODS["nes_suggest"]) + async def suggest_nes( + self, + session_id: str, + uri: str | AnyUrl, + version: int, + position: schema.Position, + trigger_kind: Literal["automatic"] | Literal["diagnostic"] | Literal["manual"] | str, + selection: schema.Range | None = None, + context: schema.NesSuggestContext | None = None, + **kwargs: Any, + ) -> schema.SuggestNesResponse: ... + + @param_model(schema.AcceptNesNotification, method=AGENT_METHODS["nes_accept"], kind="notification") + async def accept_nes(self, session_id: str, suggestion_id: str, **kwargs: Any) -> None: ... + + @param_model(schema.RejectNesNotification, method=AGENT_METHODS["nes_reject"], kind="notification") + async def reject_nes( + self, + session_id: str, + suggestion_id: str, + reason: Literal["rejected"] + | Literal["ignored"] + | Literal["replaced"] + | Literal["cancelled"] + | str + | None = None, + **kwargs: Any, + ) -> None: ... + + @param_model(schema.CloseNesRequest, method=AGENT_METHODS["nes_close"], default_result={}) + async def close_nes(self, session_id: str, **kwargs: Any) -> schema.CloseNesResponse: ... + + @param_model(schema.DidOpenDocumentNotification, method=AGENT_METHODS["document_did_open"], kind="notification") + async def did_open( + self, session_id: str, uri: str | AnyUrl, language_id: str, version: int, text: str, **kwargs: Any + ) -> None: ... + + @param_model(schema.DidChangeDocumentNotification, method=AGENT_METHODS["document_did_change"], kind="notification") + async def did_change( + self, + session_id: str, + uri: str | AnyUrl, + version: int, + content_changes: list[schema.TextDocumentContentChangeEvent], + **kwargs: Any, + ) -> None: ... + + @param_model(schema.DidCloseDocumentNotification, method=AGENT_METHODS["document_did_close"], kind="notification") + async def did_close(self, session_id: str, uri: str | AnyUrl, **kwargs: Any) -> None: ... + + @param_model(schema.DidSaveDocumentNotification, method=AGENT_METHODS["document_did_save"], kind="notification") + async def did_save(self, session_id: str, uri: str | AnyUrl, **kwargs: Any) -> None: ... + + @param_model(schema.DidFocusDocumentNotification, method=AGENT_METHODS["document_did_focus"], kind="notification") + async def did_focus( + self, + session_id: str, + uri: str | AnyUrl, + version: int, + position: schema.Position, + visible_range: schema.Range, + **kwargs: Any, + ) -> None: ... + + +class Client(Protocol): + """Client handlers for ACP v2; only override the methods you support.""" + + @param_model(schema.RequestPermissionRequest, method=CLIENT_METHODS["session_request_permission"]) + async def request_permission( + self, + session_id: str, + title: str, + options: list[schema.PermissionOption], + description: str | None = None, + subject: schema.ToolCallPermissionSubjectVariant + | schema.CommandPermissionSubjectVariant + | schema.OtherPermissionSubject + | None = None, + **kwargs: Any, + ) -> schema.RequestPermissionResponse: ... + + @param_model(schema.UpdateSessionNotification, method=CLIENT_METHODS["session_update"], kind="notification") + async def session_update( + self, + session_id: str, + update: schema.UserMessageChunk + | schema.UserMessageUpdate + | schema.AgentMessageChunk + | schema.AgentMessageUpdate + | schema.AgentThoughtChunk + | schema.AgentThoughtUpdate + | schema.ToolCallContentChunkUpdate + | schema.SessionToolCallUpdate + | schema.SessionTerminalUpdate + | schema.SessionTerminalOutputChunk + | schema.SessionPlanUpdate + | schema.SessionPlanRemovedUpdate + | schema.AvailableCommandsUpdate + | schema.ConfigOptionUpdate + | schema.SessionInfoUpdate + | schema.UsageUpdate + | schema.SessionNotice + | schema.SessionCompactionUpdate + | schema.SessionCompactionSummaryChunk + | schema.OtherSessionUpdate + | schema.RunningSessionStateUpdate + | schema.IdleSessionStateUpdate + | schema.RequiresActionSessionStateUpdate + | schema.OtherSessionStateUpdate, + **kwargs: Any, + ) -> None: ... + + @param_model(schema.ConnectMcpRequest, method=CLIENT_METHODS["mcp_connect"]) + async def connect_mcp(self, server_id: str, **kwargs: Any) -> schema.ConnectMcpResponse: ... + + @param_model(schema.MessageMcpRequest, method=CLIENT_METHODS["mcp_message"]) + async def mcp_message( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> Any: ... + + @param_model(schema.MessageMcpNotification, method=CLIENT_METHODS["mcp_message"], kind="notification") + async def notify_mcp( + self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any + ) -> None: ... + + @param_model(schema.DisconnectMcpRequest, method=CLIENT_METHODS["mcp_disconnect"], default_result={}) + async def disconnect_mcp(self, connection_id: str, **kwargs: Any) -> schema.DisconnectMcpResponse: ... + + @param_model(CreateElicitationRequest, method=CLIENT_METHODS["elicitation_create"]) + async def create_elicitation( + self, + message: str, + mode: str, + *, + session_id: str | None = None, + request_id: int | str | None = None, + tool_call_id: str | None = None, + requested_schema: schema.ElicitationSchema | None = None, + elicitation_id: str | None = None, + url: str | AnyUrl | None = None, + **kwargs: Any, + ) -> CreateElicitationResponse: ... + + @param_model( + schema.CompleteElicitationNotification, method=CLIENT_METHODS["elicitation_complete"], kind="notification" + ) + async def complete_elicitation(self, elicitation_id: str, **kwargs: Any) -> None: ... diff --git a/tests/test_gen_all.py b/tests/test_gen_all.py index 5f3d324..5a9240d 100644 --- a/tests/test_gen_all.py +++ b/tests/test_gen_all.py @@ -177,3 +177,41 @@ def test_signature_generation_expands_single_model_types(expression) -> None: assert isinstance(method, ast.AsyncFunctionDef) assert [arg.arg for arg in method.args.args] == ["self", "session_id"] assert ast.unparse(method.decorator_list[0]) == f"param_model({expression}, method='session/delete')" + + +def test_v2_signature_generation_uses_v2_schema_and_qualified_types() -> None: + import ast + + from acp.experimental import v2 + from scripts.gen_signature import NodeTransformer + + tree = ast.parse( + "from typing import Any\n" + "from . import schema\n" + "class Methods:\n" + " @param_model(schema.InitializeRequest)\n" + " async def initialize(self, **kwargs: Any): ...\n" + ) + NodeTransformer(v2.schema).visit(tree) + ast.fix_missing_locations(tree) + methods = tree.body[-1] + assert isinstance(methods, ast.ClassDef) + method = methods.body[0] + assert isinstance(method, ast.AsyncFunctionDef) + assert [arg.arg for arg in method.args.args] == ["self", "protocol_version", "info", "capabilities"] + assert method.args.args[2].annotation is not None + assert method.args.args[3].annotation is not None + assert ast.unparse(method.args.args[2].annotation) == "schema.Implementation" + assert ast.unparse(method.args.args[3].annotation) == "schema.ClientCapabilities | None" + + +def test_v1_signature_generation_does_not_rewrite_experimental_files(tmp_path) -> None: + from scripts.gen_signature import gen_signature + + experimental = tmp_path / "experimental" / "v2" + experimental.mkdir(parents=True) + target = experimental / "interfaces.py" + source = "from . import schema\n@param_model(schema.InitializeRequest)\nasync def initialize(self): ...\n" + target.write_text(source) + gen_signature(tmp_path) + assert target.read_text() == source diff --git a/tests/test_protocol_negotiation.py b/tests/test_protocol_negotiation.py index 93a043c..06e763e 100644 --- a/tests/test_protocol_negotiation.py +++ b/tests/test_protocol_negotiation.py @@ -41,7 +41,9 @@ class V2Agent: def __init__(self) -> None: self.initialize_calls = 0 - async def initialize(self, request: v2.schema.InitializeRequest) -> v2.schema.InitializeResponse: + async def initialize( + self, protocol_version: int, info: v2.schema.Implementation, **kwargs: Any + ) -> v2.schema.InitializeResponse: self.initialize_calls += 1 return v2.schema.InitializeResponse( protocol_version=v2.PROTOCOL_VERSION, @@ -53,7 +55,36 @@ class UpdateClient: def __init__(self) -> None: self.updates: asyncio.Queue[v2.schema.UpdateSessionNotification] = asyncio.Queue() - async def session_update(self, notification: v2.schema.UpdateSessionNotification) -> None: + async def session_update( + self, + session_id: str, + update: v2.schema.UserMessageChunk + | v2.schema.UserMessageUpdate + | v2.schema.AgentMessageChunk + | v2.schema.AgentMessageUpdate + | v2.schema.AgentThoughtChunk + | v2.schema.AgentThoughtUpdate + | v2.schema.ToolCallContentChunkUpdate + | v2.schema.SessionToolCallUpdate + | v2.schema.SessionTerminalUpdate + | v2.schema.SessionTerminalOutputChunk + | v2.schema.SessionPlanUpdate + | v2.schema.SessionPlanRemovedUpdate + | v2.schema.AvailableCommandsUpdate + | v2.schema.ConfigOptionUpdate + | v2.schema.SessionInfoUpdate + | v2.schema.UsageUpdate + | v2.schema.SessionNotice + | v2.schema.SessionCompactionUpdate + | v2.schema.SessionCompactionSummaryChunk + | v2.schema.OtherSessionUpdate + | v2.schema.RunningSessionStateUpdate + | v2.schema.IdleSessionStateUpdate + | v2.schema.RequiresActionSessionStateUpdate + | v2.schema.OtherSessionStateUpdate, + **kwargs: Any, + ) -> None: + notification = v2.schema.UpdateSessionNotification(session_id=session_id, update=update, **kwargs) await self.updates.put(notification) @@ -62,13 +93,21 @@ def __init__(self, connection: v2.AgentSideConnection) -> None: super().__init__() self.connection = connection - async def prompt(self, request: v2.schema.PromptRequest) -> v2.schema.PromptResponse: - await self.connection.session_update( - v2.schema.UpdateSessionNotification( - session_id=request.session_id, - update=v2.schema.IdleSessionStateUpdate(), - ) - ) + async def prompt( + self, + session_id: str, + prompt: list[ + v2.schema.TextContentBlock + | v2.schema.ImageContentBlock + | v2.schema.AudioContentBlock + | v2.schema.ResourceContentBlock + | v2.schema.EmbeddedResourceContentBlock + | v2.schema.OtherContentBlock + ], + **kwargs: Any, + ) -> v2.schema.PromptResponse: + request = v2.schema.PromptRequest(session_id=session_id, prompt=prompt, **kwargs) + await self.connection.session_update(session_id=request.session_id, update=v2.schema.IdleSessionStateUpdate()) return v2.schema.PromptResponse(message_id="user-message-1") @@ -89,7 +128,9 @@ async def test_agent_protocol_router_selects_v2() -> None: client_connection = v2.ClientSideConnection(Client(), client_transport, observers=[wire.append]) try: - initialized = await client_connection.initialize(v2_initialize()) + initialized = await client_connection.initialize( + protocol_version=v2.PROTOCOL_VERSION, info=v2.schema.Implementation(name="v2-client", version="1.0.0") + ) assert initialized.info.name == "v2-agent" assert v1_agent.initialize_calls == 0 @@ -172,14 +213,13 @@ def create_agent(connection: v2.AgentSideConnection) -> RoutedV2Agent: client_connection_2 = v2.ClientSideConnection(client_2, client_2_transport) try: - await client_connection_1.initialize(v2_initialize()) - await client_connection_2.initialize(v2_initialize()) - await client_connection_1.prompt( - v2.schema.PromptRequest( - session_id="session-1", - prompt=[v2.schema.TextContentBlock(text="hello")], - ) + await client_connection_1.initialize( + protocol_version=v2.PROTOCOL_VERSION, info=v2.schema.Implementation(name="v2-client", version="1.0.0") + ) + await client_connection_2.initialize( + protocol_version=v2.PROTOCOL_VERSION, info=v2.schema.Implementation(name="v2-client", version="1.0.0") ) + await client_connection_1.prompt(session_id="session-1", prompt=[v2.schema.TextContentBlock(text="hello")]) update = await asyncio.wait_for(client_1.updates.get(), timeout=1) assert update.session_id == "session-1" diff --git a/tests/test_v2_routing.py b/tests/test_v2_routing.py new file mode 100644 index 0000000..976d073 --- /dev/null +++ b/tests/test_v2_routing.py @@ -0,0 +1,141 @@ +from __future__ import annotations + +import inspect +from typing import Any + +import pytest +from pydantic import ValidationError + +from acp.exceptions import RequestError +from acp.experimental import v2 +from acp.experimental.v2._methods import protocol_specs +from acp.experimental.v2._router import MethodRouter +from acp.experimental.v2.meta import AGENT_METHODS, CLIENT_METHODS + + +@pytest.mark.parametrize( + ("protocol", "connection", "methods"), + [(v2.Agent, v2.ClientSideConnection, AGENT_METHODS), (v2.Client, v2.AgentSideConnection, CLIENT_METHODS)], +) +def test_protocols_cover_wire_methods_and_match_connection_signatures(protocol, connection, methods) -> None: + requests, notifications = protocol_specs(protocol) + assert {spec.method for spec in (*requests, *notifications)} == set(methods.values()) + for spec in (*requests, *notifications): + assert inspect.signature(getattr(protocol, spec.handler)) == inspect.signature( + getattr(connection, spec.handler) + ) + assert "mcp/message" in {spec.method for spec in requests} + assert "mcp/message" in {spec.method for spec in notifications} + + +@pytest.mark.asyncio +async def test_protocol_stubs_are_not_treated_as_implemented_handlers() -> None: + class Agent(v2.Agent): + pass + + router = MethodRouter(Agent(), v2.Agent) + with pytest.raises(RequestError) as error: + await router("session/new", {"cwd": "/workspace"}, False) + assert isinstance(error.value, RequestError) + assert error.value.code == -32601 + assert await router("session/cancel", {"sessionId": "s"}, True) is None + + +@pytest.mark.asyncio +async def test_v2_route_validation_and_empty_responses_remain_strict() -> None: + class Agent: + def __init__(self) -> None: + self.calls = 0 + + async def prompt(self, **kwargs: Any) -> Any: + self.calls += 1 + return {"stopReason": "end_turn"} # v1 response must not pass v2 validation. + + async def logout(self, **kwargs: Any) -> None: + pass + + async def mcp_message(self, **kwargs: Any) -> Any: + return None + + agent = Agent() + router = MethodRouter(agent, v2.Agent) + with pytest.raises(ValidationError): + await router("session/prompt", {"sessionId": "s", "prompt": [{"type": "text"}]}, False) + assert agent.calls == 0 + with pytest.raises(ValidationError): + await router("session/prompt", {"sessionId": "s", "prompt": []}, False) + assert agent.calls == 1 + assert isinstance(await router("auth/logout", {}, False), v2.schema.LogoutAuthResponse) + assert await router("mcp/message", {"connectionId": "m", "method": "ping"}, False) is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("tag", "value"), + [("id", "fast"), ("boolean", True), ("vendor/number", 10)], +) +async def test_config_union_flattens_selected_branch_and_metadata(tag, value) -> None: + class Agent: + async def set_config_option(self, config_id, session_id, value, *, type, **kwargs): # noqa: A002 + assert (config_id, session_id, type) == ("option", "s", tag) + assert kwargs == {"trace": "t"} + return {"configOptions": []} + + router = MethodRouter(Agent(), v2.Agent) + result = await router( + "session/set_config_option", + { + "configId": "option", + "sessionId": "s", + "type": tag, + "value": value, + "_meta": {"trace": "t"}, + }, + False, + ) + assert isinstance(result, v2.schema.SetSessionConfigOptionResponse) + with pytest.raises(ValidationError): + await router( + "session/set_config_option", + { + "configId": "option", + "sessionId": "s", + "type": "boolean", + "value": {"bad": True}, + }, + False, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("scope", [{"sessionId": "s", "toolCallId": "tool"}, {"requestId": 7}, {"requestId": None}]) +@pytest.mark.parametrize("mode", ["form", "url", "vendor/custom"]) +async def test_elicitation_retains_branch_specific_fields(scope, mode) -> None: + captured = {} + + class Client: + async def create_elicitation(self, **kwargs): + captured.update(kwargs) + return {"action": "decline"} + + params = {"message": "Input", "mode": mode, **scope, "_meta": {"trace": "t"}} + if mode == "form": + params["requestedSchema"] = {"type": "object", "properties": {}} + if mode == "url": + params.update(elicitationId="e", url="https://example.com/") + result = await MethodRouter(Client(), v2.Client)("elicitation/create", params, False) + assert isinstance(result, v2.schema.DeclineElicitationResponse) + assert captured["mode"] == mode + assert captured["trace"] == "t" + if "sessionId" in scope: + assert captured["session_id"] == "s" + assert captured["tool_call_id"] == "tool" + assert "request_id" not in captured + else: + assert captured["request_id"] == scope["requestId"] + assert "session_id" not in captured + if mode == "form": + assert isinstance(captured["requested_schema"], v2.schema.ElicitationSchema) + if mode == "url": + assert captured["elicitation_id"] == "e" + assert str(captured["url"]) == "https://example.com/" diff --git a/tests/test_v2_runtime.py b/tests/test_v2_runtime.py index 78b3377..6300f5d 100644 --- a/tests/test_v2_runtime.py +++ b/tests/test_v2_runtime.py @@ -18,7 +18,36 @@ class SessionClient: def __init__(self) -> None: self.updates: asyncio.Queue[v2.schema.UpdateSessionNotification] = asyncio.Queue() - async def session_update(self, notification: v2.schema.UpdateSessionNotification) -> None: + async def session_update( + self, + session_id: str, + update: v2.schema.UserMessageChunk + | v2.schema.UserMessageUpdate + | v2.schema.AgentMessageChunk + | v2.schema.AgentMessageUpdate + | v2.schema.AgentThoughtChunk + | v2.schema.AgentThoughtUpdate + | v2.schema.ToolCallContentChunkUpdate + | v2.schema.SessionToolCallUpdate + | v2.schema.SessionTerminalUpdate + | v2.schema.SessionTerminalOutputChunk + | v2.schema.SessionPlanUpdate + | v2.schema.SessionPlanRemovedUpdate + | v2.schema.AvailableCommandsUpdate + | v2.schema.ConfigOptionUpdate + | v2.schema.SessionInfoUpdate + | v2.schema.UsageUpdate + | v2.schema.SessionNotice + | v2.schema.SessionCompactionUpdate + | v2.schema.SessionCompactionSummaryChunk + | v2.schema.OtherSessionUpdate + | v2.schema.RunningSessionStateUpdate + | v2.schema.IdleSessionStateUpdate + | v2.schema.RequiresActionSessionStateUpdate + | v2.schema.OtherSessionStateUpdate, + **kwargs: Any, + ) -> None: + notification = v2.schema.UpdateSessionNotification(session_id=session_id, update=update, **kwargs) await self.updates.put(notification) @@ -28,16 +57,19 @@ def __init__(self, *, response_version: int = v2.PROTOCOL_VERSION) -> None: self.initialize_calls = 0 self.client_name: str | None = None - async def initialize(self, request: v2.schema.InitializeRequest) -> v2.schema.InitializeResponse: + async def initialize( + self, protocol_version: int, info: v2.schema.Implementation, **kwargs: Any + ) -> v2.schema.InitializeResponse: self.initialize_calls += 1 - self.client_name = request.info.name + self.client_name = info.name return v2.schema.InitializeResponse( protocol_version=self.response_version, info=v2.schema.Implementation(name="test-agent", version="1.0.0"), capabilities=v2.schema.AgentCapabilities(session=v2.schema.SessionCapabilities()), ) - async def new_session(self, request: v2.schema.NewSessionRequest) -> v2.schema.NewSessionResponse: + async def new_session(self, cwd: str, **kwargs: Any) -> v2.schema.NewSessionResponse: + request = v2.schema.NewSessionRequest(cwd=cwd, **kwargs) return v2.schema.NewSessionResponse(session_id=f"session:{request.cwd}") @@ -54,22 +86,28 @@ def __init__(self, *, echo_before_response: bool) -> None: def on_connect(self, connection: v2.AgentSideConnection) -> None: self.connection = connection - async def new_session(self, request: v2.schema.NewSessionRequest) -> v2.schema.NewSessionResponse: - response = await super().new_session(request) - await self.connection.session_update( - v2.schema.UpdateSessionNotification( - session_id=response.session_id, - update=v2.schema.IdleSessionStateUpdate(), - ) - ) + async def new_session(self, cwd: str, **kwargs: Any) -> v2.schema.NewSessionResponse: + request = v2.schema.NewSessionRequest(cwd=cwd, **kwargs) + response = await super().new_session(cwd=request.cwd) + await self.connection.session_update(session_id=response.session_id, update=v2.schema.IdleSessionStateUpdate()) return response - async def prompt(self, request: v2.schema.PromptRequest) -> v2.schema.PromptResponse: + async def prompt( + self, + session_id: str, + prompt: list[ + v2.schema.TextContentBlock + | v2.schema.ImageContentBlock + | v2.schema.AudioContentBlock + | v2.schema.ResourceContentBlock + | v2.schema.EmbeddedResourceContentBlock + | v2.schema.OtherContentBlock + ], + **kwargs: Any, + ) -> v2.schema.PromptResponse: + request = v2.schema.PromptRequest(session_id=session_id, prompt=prompt, **kwargs) await self.connection.session_update( - v2.schema.UpdateSessionNotification( - session_id=request.session_id, - update=v2.schema.RunningSessionStateUpdate(), - ) + session_id=request.session_id, update=v2.schema.RunningSessionStateUpdate() ) if self.echo_before_response: await self.echo_prompt(request) @@ -77,10 +115,8 @@ async def prompt(self, request: v2.schema.PromptRequest) -> v2.schema.PromptResp async def echo_prompt(self, request: v2.schema.PromptRequest) -> None: await self.connection.session_update( - v2.schema.UpdateSessionNotification( - session_id=request.session_id, - update=v2.schema.UserMessageUpdate(message_id="user-message-1", content=request.prompt), - ) + session_id=request.session_id, + update=v2.schema.UserMessageUpdate(message_id="user-message-1", content=request.prompt), ) @@ -93,7 +129,9 @@ class ExtensionAgent: def __init__(self) -> None: self.notifications: asyncio.Queue[tuple[str, Any]] = asyncio.Queue() - async def initialize(self, request: v2.schema.InitializeRequest) -> v2.schema.InitializeResponse: + async def initialize( + self, protocol_version: int, info: v2.schema.Implementation, **kwargs: Any + ) -> v2.schema.InitializeResponse: return v2.schema.InitializeResponse( protocol_version=v2.PROTOCOL_VERSION, info=v2.schema.Implementation(name="extension-agent", version="1.0.0"), @@ -102,10 +140,12 @@ async def initialize(self, request: v2.schema.InitializeRequest) -> v2.schema.In async def handle_extension_request(self, method: str, params: Any) -> Any: return {"method": method, "params": params} - async def cancel_session(self, notification: v2.schema.CancelSessionNotification) -> None: + async def cancel_session(self, session_id: str, **kwargs: Any) -> None: + notification = v2.schema.CancelSessionNotification(session_id=session_id, **kwargs) await self.notifications.put(("cancel", notification)) - async def notify_mcp(self, notification: v2.schema.MessageMcpNotification) -> None: + async def notify_mcp(self, connection_id: str, method: str, **kwargs: Any) -> None: + notification = v2.schema.MessageMcpNotification(connection_id=connection_id, method=method, **kwargs) await self.notifications.put(("mcp", notification)) @@ -124,8 +164,10 @@ async def test_v2_runtime_initializes_and_routes_generated_models() -> None: client_connection = v2.ClientSideConnection(Client(), client_transport) try: - initialized = await client_connection.initialize(initialize_request()) - session = await client_connection.new_session(v2.schema.NewSessionRequest(cwd="/workspace")) + initialized = await client_connection.initialize( + protocol_version=v2.PROTOCOL_VERSION, info=v2.schema.Implementation(name="test-client", version="1.0.0") + ) + session = await client_connection.new_session(cwd="/workspace") assert initialized.protocol_version == v2.PROTOCOL_VERSION assert session.session_id == "session:/workspace" @@ -144,7 +186,7 @@ async def test_v2_runtime_rejects_calls_before_initialize() -> None: try: with pytest.raises(RequestError, match="Invalid request"): - await client_connection.new_session(v2.schema.NewSessionRequest(cwd="/workspace")) + await client_connection.new_session(cwd="/workspace") finally: await client_connection.close() await agent_connection.close() @@ -158,7 +200,9 @@ async def test_callable_agent_is_not_treated_as_a_factory() -> None: client_connection = v2.ClientSideConnection(Client(), client_transport) try: - initialized = await client_connection.initialize(initialize_request()) + initialized = await client_connection.initialize( + protocol_version=v2.PROTOCOL_VERSION, info=v2.schema.Implementation(name="test-client", version="1.0.0") + ) assert initialized.protocol_version == v2.PROTOCOL_VERSION assert agent.initialize_calls == 1 @@ -176,7 +220,9 @@ async def test_v2_runtime_rejects_a_different_protocol_version() -> None: try: with pytest.raises(RequestError) as error: - await client_connection.initialize(initialize_request(protocol_version=1)) + await client_connection.initialize( + protocol_version=1, info=v2.schema.Implementation(name="test-client", version="1.0.0") + ) assert isinstance(error.value, RequestError) assert error.value.code == -32602 @@ -194,7 +240,9 @@ async def test_v2_runtime_rejects_a_mismatched_initialize_response() -> None: try: with pytest.raises(RequestError) as error: - await client_connection.initialize(initialize_request()) + await client_connection.initialize( + protocol_version=v2.PROTOCOL_VERSION, info=v2.schema.Implementation(name="test-client", version="1.0.0") + ) assert isinstance(error.value, RequestError) assert error.value.code == -32600 @@ -213,21 +261,20 @@ async def test_session_updates_are_delivered_independently_from_prompt(echo_befo client_connection = v2.ClientSideConnection(client, client_transport) try: - await client_connection.initialize(initialize_request()) - session = await client_connection.new_session(v2.schema.NewSessionRequest(cwd="/workspace")) + await client_connection.initialize( + protocol_version=v2.PROTOCOL_VERSION, info=v2.schema.Implementation(name="test-client", version="1.0.0") + ) + session = await client_connection.new_session(cwd="/workspace") request = v2.schema.PromptRequest( session_id=session.session_id, prompt=[v2.schema.TextContentBlock(text="hello")], ) - response = await client_connection.prompt(request) + response = await client_connection.prompt(session_id=request.session_id, prompt=request.prompt) if not echo_before_response: await agent.echo_prompt(request) # Completion is independent traffic, sent after prompt acceptance. await agent_connection.session_update( - v2.schema.UpdateSessionNotification( - session_id=session.session_id, - update=v2.schema.IdleSessionStateUpdate(stop_reason="end_turn"), - ) + session_id=session.session_id, update=v2.schema.IdleSessionStateUpdate(stop_reason="end_turn") ) ready = await asyncio.wait_for(client.updates.get(), timeout=1) @@ -271,10 +318,10 @@ async def test_notice_and_compaction_updates_reach_client(update) -> None: v2.AgentSideConnection(Agent(), agent_transport) as agent_connection, v2.ClientSideConnection(client, client_transport) as client_connection, ): - await client_connection.initialize(initialize_request()) - await agent_connection.session_update( - v2.schema.UpdateSessionNotification(session_id="session-1", update=update) + await client_connection.initialize( + protocol_version=v2.PROTOCOL_VERSION, info=v2.schema.Implementation(name="test-client", version="1.0.0") ) + await agent_connection.session_update(session_id="session-1", update=update) received = await asyncio.wait_for(client.updates.get(), timeout=1) assert received.session_id == "session-1" assert received.update == update @@ -303,11 +350,11 @@ async def test_session_patches_preserve_omitted_and_cleared_fields(updates) -> N v2.AgentSideConnection(Agent(), agent_transport) as agent_connection, v2.ClientSideConnection(client, client_transport) as client_connection, ): - await client_connection.initialize(initialize_request()) + await client_connection.initialize( + protocol_version=v2.PROTOCOL_VERSION, info=v2.schema.Implementation(name="test-client", version="1.0.0") + ) for update in updates: - await agent_connection.session_update( - v2.schema.UpdateSessionNotification(session_id="session-1", update=update) - ) + await agent_connection.session_update(session_id="session-1", update=update) received = await asyncio.wait_for(client.updates.get(), timeout=1) assert received.update.model_dump(by_alias=True, exclude_unset=True) == update.model_dump( by_alias=True, exclude_unset=True @@ -320,6 +367,8 @@ def test_v2_public_entry_point_is_explicit() -> None: assert exported["PROTOCOL_VERSION"] == 2 assert exported["schema"] is v2.schema assert set(exported) == { + "Agent", + "Client", "AgentSideConnection", "ClientSideConnection", "PROTOCOL_VERSION", @@ -338,7 +387,9 @@ async def test_extension_and_notification_names_are_explicit() -> None: client_connection = v2.ClientSideConnection(ExtensionClient(), client_transport) try: - await client_connection.initialize(initialize_request()) + await client_connection.initialize( + protocol_version=v2.PROTOCOL_VERSION, info=v2.schema.Implementation(name="test-client", version="1.0.0") + ) assert await client_connection.send_extension_request("_vendor/do", {"value": 1}) == { "method": "_vendor/do", @@ -351,10 +402,8 @@ async def test_extension_and_notification_names_are_explicit() -> None: with pytest.raises(ValueError, match="must start with '_'"): await client_connection.send_extension_request("vendor/do") - await client_connection.cancel_session(v2.schema.CancelSessionNotification(session_id="session-1")) - await client_connection.notify_mcp( - v2.schema.MessageMcpNotification(connection_id="mcp-1", method="notifications/progress") - ) + await client_connection.cancel_session(session_id="session-1") + await client_connection.notify_mcp(connection_id="mcp-1", method="notifications/progress") cancel_kind, cancel = await asyncio.wait_for(agent.notifications.get(), timeout=1) mcp_kind, mcp = await asyncio.wait_for(agent.notifications.get(), timeout=1) @@ -372,13 +421,10 @@ async def test_unhandled_notifications_are_ignored(caplog: pytest.LogCaptureFixt client_connection = v2.ClientSideConnection(object(), client_transport) try: - await client_connection.initialize(initialize_request()) - await agent_connection.session_update( - v2.schema.UpdateSessionNotification( - session_id="session-1", - update=v2.schema.IdleSessionStateUpdate(), - ) + await client_connection.initialize( + protocol_version=v2.PROTOCOL_VERSION, info=v2.schema.Implementation(name="test-client", version="1.0.0") ) + await agent_connection.session_update(session_id="session-1", update=v2.schema.IdleSessionStateUpdate()) await client_connection.send_extension_notification("_vendor/event") await asyncio.sleep(0) await asyncio.sleep(0) @@ -396,12 +442,76 @@ async def test_missing_request_handler_returns_method_not_found() -> None: client_connection = v2.ClientSideConnection(object(), client_transport) try: - await client_connection.initialize(initialize_request()) + await client_connection.initialize( + protocol_version=v2.PROTOCOL_VERSION, info=v2.schema.Implementation(name="test-client", version="1.0.0") + ) with pytest.raises(RequestError) as error: - await client_connection.new_session(v2.schema.NewSessionRequest(cwd="/workspace")) + await client_connection.new_session(cwd="/workspace") assert isinstance(error.value, RequestError) assert error.value.code == -32601 finally: await client_connection.close() await agent_connection.close() + + +@pytest.mark.asyncio +async def test_expanded_union_calls_round_trip_with_metadata() -> None: + configurations: list[tuple[Any, Any, Any]] = [] + elicitations: list[dict[str, Any]] = [] + + class ConfigAgent(Agent): + async def set_config_option(self, config_id, session_id, value, *, type, **kwargs): # noqa: A002 + assert (config_id, session_id, kwargs) == ("option", "s", {"trace": "config"}) + configurations.append((type, value, kwargs)) + return v2.schema.SetSessionConfigOptionResponse(config_options=[]) + + class ElicitationClient: + async def create_elicitation(self, message, mode, **kwargs): + assert message == "Input" + elicitations.append({"mode": mode, **kwargs}) + return v2.schema.DeclineElicitationResponse() + + client_transport, agent_transport = memory_transport_pair() + async with ( + v2.AgentSideConnection(ConfigAgent(), agent_transport) as agent_connection, + v2.ClientSideConnection(ElicitationClient(), client_transport) as client_connection, + ): + await client_connection.initialize( + protocol_version=v2.PROTOCOL_VERSION, + info=v2.schema.Implementation(name="test", version="1"), + ) + for value, tag in [(True, None), ("fast", None), (10, "vendor/number")]: + await client_connection.set_config_option("option", "s", value, type=tag, trace="config") + assert [(tag, value) for tag, value, _ in configurations] == [ + ("boolean", True), + ("id", "fast"), + ("vendor/number", 10), + ] + await agent_connection.create_elicitation( + "Input", + "form", + session_id="s", + tool_call_id="tool", + requested_schema=v2.schema.ElicitationSchema(properties={}), + trace="form", + ) + await agent_connection.create_elicitation( + "Input", + "url", + request_id=7, + elicitation_id="e", + url="https://example.com/", + trace="url", + ) + await agent_connection.create_elicitation("Input", "vendor/custom", request_id=None) + assert elicitations[0]["session_id"] == "s" + assert elicitations[0]["tool_call_id"] == "tool" + assert elicitations[0]["trace"] == "form" + assert isinstance(elicitations[0]["requested_schema"], v2.schema.ElicitationSchema) + assert elicitations[1]["request_id"] == 7 + assert elicitations[1]["elicitation_id"] == "e" + assert elicitations[1]["trace"] == "url" + assert elicitations[2]["request_id"] is None + with pytest.raises(ValueError, match="either session_id or request_id"): + await agent_connection.create_elicitation("Input", "vendor/custom", session_id="s", request_id=7)