From 55f92cac6e0211f3532e6ab4210e8ddb874891b6 Mon Sep 17 00:00:00 2001 From: Frost Ming Date: Mon, 21 Sep 2026 15:47:29 +0800 Subject: [PATCH 1/2] refactor: derive routers from protocol metadata --- docs/quickstart.md | 28 +++++ src/acp/_protocol_adapters.py | 72 +++++++++++ src/acp/agent/router.py | 180 +-------------------------- src/acp/client/router.py | 191 +---------------------------- src/acp/interfaces.py | 152 ++++++++++++++++------- src/acp/router.py | 85 +++++++++++-- src/acp/utils.py | 85 ++++++++++++- tests/test_elicitation_catchall.py | 2 +- tests/test_gen_all.py | 10 +- tests/test_protocol_router.py | 179 +++++++++++++++++++++++++++ 10 files changed, 555 insertions(+), 429 deletions(-) create mode 100644 src/acp/_protocol_adapters.py create mode 100644 tests/test_protocol_router.py diff --git a/docs/quickstart.md b/docs/quickstart.md index 1724c11..9556ee5 100644 --- a/docs/quickstart.md +++ b/docs/quickstart.md @@ -228,6 +228,34 @@ MCP requests return the inner JSON result unchanged, including `null`. Use and optional `params`. These methods share the same connections and routers across stdio, HTTP, and WebSocket transports. +## Maintaining protocol routes + +The `Agent` and `Client` protocols in `src/acp/interfaces.py` are the source of +truth for incoming routes. Declare routing metadata alongside each method's +parameter model: + +```python +@param_model( + DeleteSessionRequest, + method=AGENT_METHODS["session_delete"], + adapt_result=normalize_result, +) +async def delete_session(self, session_id: str, **kwargs: Any) -> DeleteSessionResponse: ... +``` + +`build_agent_router` and `build_client_router` read these declarations through +`MessageRouter.from_protocol`. New protocol declarations automatically become +routes; implementation-only methods are not exposed. Use `kind="notification"` +for notifications and `unstable=True` for methods requiring explicit opt-in. +Optional handlers can declare `optional=True` and `default_result`. + +For multiple wire models, `param_models` accepts the same routing options. +`validate_params` and `adapt_params` provide custom validation and conversion +when the wire representation differs from the Python signature, as with config +options and elicitation. Legacy handlers still receive the validated request +model. Connection methods retain model-only decorators for signature generation +and legacy calls; they do not repeat the routing options. + ## Optional — Talk to the Gemini CLI _Have the Gemini CLI installed? Run the bridge to exercise permission flows._ diff --git a/src/acp/_protocol_adapters.py b/src/acp/_protocol_adapters.py new file mode 100644 index 0000000..318735c --- /dev/null +++ b/src/acp/_protocol_adapters.py @@ -0,0 +1,72 @@ +from __future__ import annotations + +from typing import Any, cast + +from pydantic import BaseModel, TypeAdapter + +from .exceptions import RequestError +from .schema import ( + CreateElicitationRequest, + CreateFormRequestElicitationRequest, + CreateFormSessionElicitationRequest, + CreateUrlRequestElicitationRequest, + CreateUrlSessionElicitationRequest, + ElicitationFormRequestMode, + ElicitationFormSessionMode, + ElicitationUrlRequestMode, + ElicitationUrlSessionMode, + SetSessionConfigOptionBooleanRequest, + SetSessionConfigOptionSelectRequest, +) + +_CREATE_ELICITATION_REQUEST_ADAPTER = TypeAdapter(CreateElicitationRequest) + + +def validate_create_elicitation_request(params: Any) -> CreateElicitationRequest: + return _CREATE_ELICITATION_REQUEST_ADAPTER.validate_python(params) + + +def _mode_from_create_elicitation_request( + request: CreateElicitationRequest, +) -> ElicitationFormSessionMode | ElicitationFormRequestMode | ElicitationUrlSessionMode | ElicitationUrlRequestMode: + if isinstance(request, CreateFormSessionElicitationRequest): + return ElicitationFormSessionMode( + session_id=request.session_id, + tool_call_id=request.tool_call_id, + requested_schema=request.requested_schema, + ) + if isinstance(request, CreateFormRequestElicitationRequest): + return ElicitationFormRequestMode( + request_id=request.request_id, + requested_schema=request.requested_schema, + ) + + if isinstance(request, CreateUrlSessionElicitationRequest): + return ElicitationUrlSessionMode( + session_id=request.session_id, + tool_call_id=request.tool_call_id, + elicitation_id=request.elicitation_id, + url=request.url, + ) + if isinstance(request, CreateUrlRequestElicitationRequest): + return ElicitationUrlRequestMode( + request_id=request.request_id, + elicitation_id=request.elicitation_id, + url=request.url, + ) + raise RequestError.invalid_params({"details": f"Unsupported elicitation mode: {request.mode!r}"}) + + +def elicitation_to_kwargs(request: BaseModel) -> dict[str, Any]: + # The validator has already resolved the wire union, including custom modes. + request = cast(CreateElicitationRequest, request) + kwargs = {"message": request.message, "mode": _mode_from_create_elicitation_request(request)} + if request.field_meta: + kwargs.update(request.field_meta) + return kwargs + + +def validate_set_config_option_request(params: Any) -> BaseModel: + if isinstance(params, dict) and params.get("type") == "boolean": + return SetSessionConfigOptionBooleanRequest.model_validate(params) + return SetSessionConfigOptionSelectRequest.model_validate(params) diff --git a/src/acp/agent/router.py b/src/acp/agent/router.py index 01b15a7..5d58a58 100644 --- a/src/acp/agent/router.py +++ b/src/acp/agent/router.py @@ -1,186 +1,10 @@ from __future__ import annotations -from typing import Any - -from pydantic import BaseModel - -from ..exceptions import RequestError from ..interfaces import Agent -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 +from ..router import MessageRouter __all__ = ["build_agent_router"] -_SET_CONFIG_OPTION_MODELS = (SetSessionConfigOptionBooleanRequest, SetSessionConfigOptionSelectRequest) - - -def _validate_set_config_option_request(params: Any) -> BaseModel: - if isinstance(params, dict) and params.get("type") == "boolean": - return SetSessionConfigOptionBooleanRequest.model_validate(params) - return SetSessionConfigOptionSelectRequest.model_validate(params) - - -def _make_set_config_option_handler(agent: Agent) -> Any: - func, attr, legacy_api = _resolve_handler(agent, "set_config_option") - if func is None: - return None - - async def wrapper(params: Any) -> Any: - if legacy_api: - _warn_legacy_handler(agent, attr) - request = _validate_set_config_option_request(params) - if legacy_api: - return await func(request) - return await func(**model_to_kwargs(request, _SET_CONFIG_OPTION_MODELS)) - - return wrapper - - def build_agent_router(agent: Agent, use_unstable_protocol: bool = False) -> MessageRouter: - router = MessageRouter(use_unstable_protocol=use_unstable_protocol) - - router.route_request(AGENT_METHODS["initialize"], InitializeRequest, agent, "initialize") - router.route_request(AGENT_METHODS["session_new"], NewSessionRequest, agent, "new_session") - router.route_request( - AGENT_METHODS["session_load"], - LoadSessionRequest, - agent, - "load_session", - adapt_result=normalize_result, - ) - router.route_request(AGENT_METHODS["session_list"], ListSessionsRequest, agent, "list_sessions") - router.route_request( - AGENT_METHODS["session_close"], - CloseSessionRequest, - agent, - "close_session", - adapt_result=normalize_result, - unstable=True, - ) - router.route_request( - AGENT_METHODS["session_set_mode"], - SetSessionModeRequest, - agent, - "set_session_mode", - adapt_result=normalize_result, - ) - router.route_request(AGENT_METHODS["session_prompt"], PromptRequest, agent, "prompt") - router.add_route( - Route( - method=AGENT_METHODS["session_set_config_option"], - func=_make_set_config_option_handler(agent), - kind="request", - adapt_result=normalize_result, - ) - ) - router.route_request( - AGENT_METHODS["authenticate"], - AuthenticateRequest, - agent, - "authenticate", - adapt_result=normalize_result, - ) - router.route_request(AGENT_METHODS["session_fork"], ForkSessionRequest, agent, "fork_session", unstable=True) - router.route_request(AGENT_METHODS["session_resume"], ResumeSessionRequest, agent, "resume_session", unstable=True) - - 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) - if ext is None: - raise RequestError.method_not_found(f"_{name}") - return await ext(name, payload) - - @router.handle_extension_notification - async def _handle_extension_notification(name: str, payload: dict[str, Any]) -> None: - ext = getattr(agent, "ext_notification", None) - if ext is None: - return - await ext(name, payload) - - return router + return MessageRouter.from_protocol(Agent, agent, use_unstable_protocol=use_unstable_protocol) diff --git a/src/acp/client/router.py b/src/acp/client/router.py index b7bc825..ecff0bc 100644 --- a/src/acp/client/router.py +++ b/src/acp/client/router.py @@ -1,197 +1,10 @@ from __future__ import annotations -from typing import Any - -from pydantic import TypeAdapter - -from ..exceptions import RequestError from ..interfaces import Client -from ..meta import CLIENT_METHODS -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, - SessionNotification, - TerminalOutputRequest, - WaitForTerminalExitRequest, - WriteTextFileRequest, -) -from ..utils import normalize_result +from ..router import MessageRouter __all__ = ["build_client_router"] -_CREATE_ELICITATION_REQUEST_ADAPTER = TypeAdapter(CreateElicitationRequest) - - -def _validate_create_elicitation_request(params: Any) -> CreateElicitationRequest: - return _CREATE_ELICITATION_REQUEST_ADAPTER.validate_python(params) - - -def _mode_from_create_elicitation_request( - request: CreateElicitationRequest, -) -> ElicitationFormSessionMode | ElicitationFormRequestMode | ElicitationUrlSessionMode | ElicitationUrlRequestMode: - if isinstance(request, CreateFormSessionElicitationRequest): - return ElicitationFormSessionMode( - session_id=request.session_id, - tool_call_id=request.tool_call_id, - requested_schema=request.requested_schema, - ) - if isinstance(request, CreateFormRequestElicitationRequest): - return ElicitationFormRequestMode( - request_id=request.request_id, - requested_schema=request.requested_schema, - ) - - if isinstance(request, CreateUrlSessionElicitationRequest): - return ElicitationUrlSessionMode( - session_id=request.session_id, - tool_call_id=request.tool_call_id, - elicitation_id=request.elicitation_id, - url=request.url, - ) - if isinstance(request, CreateUrlRequestElicitationRequest): - return ElicitationUrlRequestMode( - request_id=request.request_id, - elicitation_id=request.elicitation_id, - url=request.url, - ) - raise RequestError.invalid_params({"details": f"Unsupported elicitation mode: {request.mode!r}"}) - - -def _make_create_elicitation_handler(client: Client) -> Any: - func, attr, legacy_api = _resolve_handler(client, "create_elicitation") - if func is None: - return None - - async def wrapper(params: Any) -> Any: - if legacy_api: - _warn_legacy_handler(client, attr) - request = _validate_create_elicitation_request(params) - if legacy_api: - return await func(request) - kwargs = {"message": request.message, "mode": _mode_from_create_elicitation_request(request)} - if request.field_meta: - kwargs.update(request.field_meta) - return await func(**kwargs) - - return wrapper def build_client_router(client: Client, use_unstable_protocol: bool = False) -> MessageRouter: - router = MessageRouter(use_unstable_protocol=use_unstable_protocol) - - router.route_request(CLIENT_METHODS["fs_write_text_file"], WriteTextFileRequest, client, "write_text_file") - router.route_request(CLIENT_METHODS["fs_read_text_file"], ReadTextFileRequest, client, "read_text_file") - router.route_request( - CLIENT_METHODS["session_request_permission"], - RequestPermissionRequest, - client, - "request_permission", - ) - router.route_request( - CLIENT_METHODS["terminal_create"], - CreateTerminalRequest, - client, - "create_terminal", - optional=True, - default_result=None, - ) - router.route_request( - CLIENT_METHODS["terminal_output"], - TerminalOutputRequest, - client, - "terminal_output", - optional=True, - default_result=None, - ) - router.route_request( - CLIENT_METHODS["terminal_release"], - ReleaseTerminalRequest, - client, - "release_terminal", - optional=True, - default_result={}, - adapt_result=normalize_result, - ) - router.route_request( - CLIENT_METHODS["terminal_wait_for_exit"], - WaitForTerminalExitRequest, - client, - "wait_for_terminal_exit", - optional=True, - default_result=None, - ) - router.route_request( - CLIENT_METHODS["terminal_kill"], - KillTerminalRequest, - client, - "kill_terminal", - optional=True, - default_result={}, - adapt_result=normalize_result, - ) - - router.add_route( - Route( - method=CLIENT_METHODS["elicitation_create"], - func=_make_create_elicitation_handler(client), - kind="request", - adapt_result=normalize_result, - warn_unstable=not use_unstable_protocol, - ) - ) - router.route_notification( - CLIENT_METHODS["elicitation_complete"], - CompleteElicitationNotification, - client, - "complete_elicitation", - unstable=True, - ) - - 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) - if ext is None: - raise RequestError.method_not_found(f"_{name}") - return await ext(name, payload) - - @router.handle_extension_notification - async def _handle_extension_notification(name: str, payload: dict[str, Any]) -> None: - ext = getattr(client, "ext_notification", None) - if ext is None: - return - await ext(name, payload) - - return router + return MessageRouter.from_protocol(Client, client, use_unstable_protocol=use_unstable_protocol) diff --git a/src/acp/interfaces.py b/src/acp/interfaces.py index 5dba20a..13b2f27 100644 --- a/src/acp/interfaces.py +++ b/src/acp/interfaces.py @@ -2,6 +2,12 @@ from typing import Any, Literal, Protocol +from ._protocol_adapters import ( + elicitation_to_kwargs, + validate_create_elicitation_request, + validate_set_config_option_request, +) +from .meta import AGENT_METHODS, CLIENT_METHODS from .schema import ( AcceptNesNotification, AcpMcpServer, @@ -25,8 +31,12 @@ ConnectMcpRequest, ConnectMcpResponse, CreateElicitationResponse, + CreateFormRequestElicitationRequest, + CreateFormSessionElicitationRequest, CreateTerminalRequest, CreateTerminalResponse, + CreateUrlRequestElicitationRequest, + CreateUrlSessionElicitationRequest, CurrentModeUpdate, DeleteSessionRequest, DeleteSessionResponse, @@ -113,34 +123,36 @@ WriteTextFileRequest, WriteTextFileResponse, ) -from .utils import param_model, param_models +from .utils import normalize_result, param_model, param_models __all__ = ["Agent", "Client"] class Client(Protocol): - @param_model(ConnectMcpRequest) + @param_model(ConnectMcpRequest, method=CLIENT_METHODS["mcp_connect"], unstable=True) async def connect_mcp(self, server_id: str, **kwargs: Any) -> ConnectMcpResponse: ... - @param_model(DisconnectMcpRequest) + @param_model( + DisconnectMcpRequest, method=CLIENT_METHODS["mcp_disconnect"], unstable=True, adapt_result=normalize_result + ) async def disconnect_mcp(self, connection_id: str, **kwargs: Any) -> DisconnectMcpResponse: ... - @param_model(MessageMcpRequest) + @param_model(MessageMcpRequest, method=CLIENT_METHODS["mcp_message"], unstable=True) async def mcp_message( self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any ) -> Any: ... - @param_model(MessageMcpNotification) + @param_model(MessageMcpNotification, method=CLIENT_METHODS["mcp_message"], kind="notification", unstable=True) async def notify_mcp( self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any ) -> None: ... - @param_model(RequestPermissionRequest) + @param_model(RequestPermissionRequest, method=CLIENT_METHODS["session_request_permission"]) async def request_permission( self, session_id: str, tool_call: ToolCallUpdate, options: list[PermissionOption], **kwargs: Any ) -> RequestPermissionResponse: ... - @param_model(SessionNotification) + @param_model(SessionNotification, method=CLIENT_METHODS["session_update"], kind="notification") async def session_update( self, session_id: str, @@ -163,17 +175,17 @@ async def session_update( **kwargs: Any, ) -> None: ... - @param_model(WriteTextFileRequest) + @param_model(WriteTextFileRequest, method=CLIENT_METHODS["fs_write_text_file"]) async def write_text_file( self, session_id: str, path: str, content: str, **kwargs: Any ) -> WriteTextFileResponse | None: ... - @param_model(ReadTextFileRequest) + @param_model(ReadTextFileRequest, method=CLIENT_METHODS["fs_read_text_file"]) async def read_text_file( self, session_id: str, path: str, line: int | None = None, limit: int | None = None, **kwargs: Any ) -> ReadTextFileResponse: ... - @param_model(CreateTerminalRequest) + @param_model(CreateTerminalRequest, method=CLIENT_METHODS["terminal_create"], optional=True, default_result=None) async def create_terminal( self, session_id: str, @@ -185,27 +197,57 @@ async def create_terminal( **kwargs: Any, ) -> CreateTerminalResponse: ... - @param_model(TerminalOutputRequest) + @param_model(TerminalOutputRequest, method=CLIENT_METHODS["terminal_output"], optional=True, default_result=None) async def terminal_output(self, session_id: str, terminal_id: str, **kwargs: Any) -> TerminalOutputResponse: ... - @param_model(ReleaseTerminalRequest) + @param_model( + ReleaseTerminalRequest, + method=CLIENT_METHODS["terminal_release"], + optional=True, + default_result={}, + adapt_result=normalize_result, + ) async def release_terminal( self, session_id: str, terminal_id: str, **kwargs: Any ) -> ReleaseTerminalResponse | None: ... - @param_model(WaitForTerminalExitRequest) + @param_model( + WaitForTerminalExitRequest, method=CLIENT_METHODS["terminal_wait_for_exit"], optional=True, default_result=None + ) async def wait_for_terminal_exit( self, session_id: str, terminal_id: str, **kwargs: Any ) -> WaitForTerminalExitResponse: ... - @param_model(KillTerminalRequest) + @param_model( + KillTerminalRequest, + method=CLIENT_METHODS["terminal_kill"], + optional=True, + default_result={}, + adapt_result=normalize_result, + ) async def kill_terminal(self, session_id: str, terminal_id: str, **kwargs: Any) -> KillTerminalResponse | None: ... + @param_models( + CreateFormSessionElicitationRequest, + CreateFormRequestElicitationRequest, + CreateUrlSessionElicitationRequest, + CreateUrlRequestElicitationRequest, + method=CLIENT_METHODS["elicitation_create"], + unstable=True, + validate_params=validate_create_elicitation_request, + adapt_params=elicitation_to_kwargs, + adapt_result=normalize_result, + ) async def create_elicitation( self, message: str, mode: ElicitationMode, **kwargs: Any ) -> CreateElicitationResponse: ... - @param_model(CompleteElicitationNotification) + @param_model( + CompleteElicitationNotification, + method=CLIENT_METHODS["elicitation_complete"], + kind="notification", + unstable=True, + ) async def complete_elicitation(self, elicitation_id: str, **kwargs: Any) -> None: ... async def ext_method(self, method: str, params: dict[str, Any]) -> dict[str, Any]: ... @@ -216,13 +258,15 @@ def on_connect(self, conn: Agent) -> None: ... class Agent(Protocol): - @param_model(DeleteSessionRequest) + @param_model(DeleteSessionRequest, method=AGENT_METHODS["session_delete"], adapt_result=normalize_result) async def delete_session(self, session_id: str, **kwargs: Any) -> DeleteSessionResponse: ... - @param_model(ListProvidersRequest) + @param_model(ListProvidersRequest, method=AGENT_METHODS["providers_list"], unstable=True) async def list_providers(self, **kwargs: Any) -> ListProvidersResponse: ... - @param_model(SetProviderRequest) + @param_model( + SetProviderRequest, method=AGENT_METHODS["providers_set"], unstable=True, adapt_result=normalize_result + ) async def set_provider( self, provider_id: str, @@ -237,23 +281,25 @@ async def set_provider( **kwargs: Any, ) -> SetProviderResponse: ... - @param_model(DisableProviderRequest) + @param_model( + DisableProviderRequest, method=AGENT_METHODS["providers_disable"], unstable=True, adapt_result=normalize_result + ) async def disable_provider(self, provider_id: str, **kwargs: Any) -> DisableProviderResponse: ... - @param_model(LogoutRequest) + @param_model(LogoutRequest, method=AGENT_METHODS["logout"], adapt_result=normalize_result) async def logout(self, **kwargs: Any) -> LogoutResponse: ... - @param_model(MessageMcpRequest) + @param_model(MessageMcpRequest, method=AGENT_METHODS["mcp_message"], unstable=True) async def mcp_message( self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any ) -> Any: ... - @param_model(MessageMcpNotification) + @param_model(MessageMcpNotification, method=AGENT_METHODS["mcp_message"], kind="notification", unstable=True) async def notify_mcp( self, connection_id: str, method: str, params: dict[str, Any] | None = None, **kwargs: Any ) -> None: ... - @param_model(StartNesRequest) + @param_model(StartNesRequest, method=AGENT_METHODS["nes_start"], unstable=True) async def start_nes( self, workspace_uri: str | None = None, @@ -262,7 +308,7 @@ async def start_nes( **kwargs: Any, ) -> StartNesResponse: ... - @param_model(SuggestNesRequest) + @param_model(SuggestNesRequest, method=AGENT_METHODS["nes_suggest"], unstable=True) async def suggest_nes( self, session_id: str, @@ -275,13 +321,13 @@ async def suggest_nes( **kwargs: Any, ) -> SuggestNesResponse: ... - @param_model(CloseNesRequest) + @param_model(CloseNesRequest, method=AGENT_METHODS["nes_close"], unstable=True, adapt_result=normalize_result) async def close_nes(self, session_id: str, **kwargs: Any) -> CloseNesResponse: ... - @param_model(AcceptNesNotification) + @param_model(AcceptNesNotification, method=AGENT_METHODS["nes_accept"], kind="notification", unstable=True) async def accept_nes(self, session_id: str, id: str, **kwargs: Any) -> None: ... # noqa: A002 - @param_model(RejectNesNotification) + @param_model(RejectNesNotification, method=AGENT_METHODS["nes_reject"], kind="notification", unstable=True) async def reject_nes( self, session_id: str, @@ -290,12 +336,16 @@ async def reject_nes( **kwargs: Any, ) -> None: ... - @param_model(DidOpenDocumentNotification) + @param_model( + DidOpenDocumentNotification, method=AGENT_METHODS["document_did_open"], kind="notification", unstable=True + ) async def did_open( self, session_id: str, uri: str, language_id: str, version: int, text: str, **kwargs: Any ) -> None: ... - @param_model(DidChangeDocumentNotification) + @param_model( + DidChangeDocumentNotification, method=AGENT_METHODS["document_did_change"], kind="notification", unstable=True + ) async def did_change( self, session_id: str, @@ -305,18 +355,24 @@ async def did_change( **kwargs: Any, ) -> None: ... - @param_model(DidCloseDocumentNotification) + @param_model( + DidCloseDocumentNotification, method=AGENT_METHODS["document_did_close"], kind="notification", unstable=True + ) async def did_close(self, session_id: str, uri: str, **kwargs: Any) -> None: ... - @param_model(DidSaveDocumentNotification) + @param_model( + DidSaveDocumentNotification, method=AGENT_METHODS["document_did_save"], kind="notification", unstable=True + ) async def did_save(self, session_id: str, uri: str, **kwargs: Any) -> None: ... - @param_model(DidFocusDocumentNotification) + @param_model( + DidFocusDocumentNotification, method=AGENT_METHODS["document_did_focus"], kind="notification", unstable=True + ) async def did_focus( self, session_id: str, uri: str, version: int, position: Position, visible_range: Range, **kwargs: Any ) -> None: ... - @param_model(InitializeRequest) + @param_model(InitializeRequest, method=AGENT_METHODS["initialize"]) async def initialize( self, protocol_version: int, @@ -325,7 +381,7 @@ async def initialize( **kwargs: Any, ) -> InitializeResponse: ... - @param_model(NewSessionRequest) + @param_model(NewSessionRequest, method=AGENT_METHODS["session_new"]) async def new_session( self, cwd: str, @@ -334,7 +390,7 @@ async def new_session( **kwargs: Any, ) -> NewSessionResponse: ... - @param_model(LoadSessionRequest) + @param_model(LoadSessionRequest, method=AGENT_METHODS["session_load"], adapt_result=normalize_result) async def load_session( self, cwd: str, @@ -344,23 +400,29 @@ async def load_session( **kwargs: Any, ) -> LoadSessionResponse | None: ... - @param_model(ListSessionsRequest) + @param_model(ListSessionsRequest, method=AGENT_METHODS["session_list"]) async def list_sessions( self, cwd: str | None = None, cursor: str | None = None, **kwargs: Any ) -> ListSessionsResponse: ... - @param_model(SetSessionModeRequest) + @param_model(SetSessionModeRequest, method=AGENT_METHODS["session_set_mode"], adapt_result=normalize_result) async def set_session_mode(self, session_id: str, mode_id: str, **kwargs: Any) -> SetSessionModeResponse | None: ... - @param_models(SetSessionConfigOptionBooleanRequest, SetSessionConfigOptionSelectRequest) + @param_models( + SetSessionConfigOptionBooleanRequest, + SetSessionConfigOptionSelectRequest, + method=AGENT_METHODS["session_set_config_option"], + validate_params=validate_set_config_option_request, + adapt_result=normalize_result, + ) async def set_config_option( self, config_id: str, session_id: str, value: str | bool, **kwargs: Any ) -> SetSessionConfigOptionResponse | None: ... - @param_model(AuthenticateRequest) + @param_model(AuthenticateRequest, method=AGENT_METHODS["authenticate"], adapt_result=normalize_result) async def authenticate(self, method_id: str, **kwargs: Any) -> AuthenticateResponse | None: ... - @param_model(PromptRequest) + @param_model(PromptRequest, method=AGENT_METHODS["session_prompt"]) async def prompt( self, session_id: str, @@ -374,7 +436,7 @@ async def prompt( **kwargs: Any, ) -> PromptResponse: ... - @param_model(ForkSessionRequest) + @param_model(ForkSessionRequest, method=AGENT_METHODS["session_fork"], unstable=True) async def fork_session( self, session_id: str, @@ -384,7 +446,7 @@ async def fork_session( **kwargs: Any, ) -> ForkSessionResponse: ... - @param_model(ResumeSessionRequest) + @param_model(ResumeSessionRequest, method=AGENT_METHODS["session_resume"], unstable=True) async def resume_session( self, session_id: str, @@ -394,10 +456,12 @@ async def resume_session( **kwargs: Any, ) -> ResumeSessionResponse: ... - @param_model(CloseSessionRequest) + @param_model( + CloseSessionRequest, method=AGENT_METHODS["session_close"], unstable=True, adapt_result=normalize_result + ) async def close_session(self, session_id: str, **kwargs: Any) -> CloseSessionResponse | None: ... - @param_model(CancelNotification) + @param_model(CancelNotification, method=AGENT_METHODS["session_cancel"], kind="notification") async def cancel(self, session_id: str, **kwargs: Any) -> None: ... async def ext_method(self, method: str, params: dict[str, Any]) -> dict[str, Any]: ... diff --git a/src/acp/router.py b/src/acp/router.py index 3069deb..8c2ce10 100644 --- a/src/acp/router.py +++ b/src/acp/router.py @@ -4,11 +4,11 @@ import warnings from collections.abc import Awaitable, Callable from dataclasses import dataclass -from typing import Any, Literal, TypeVar +from typing import Any, Literal, TypeVar, Union -from pydantic import BaseModel +from pydantic import BaseModel, TypeAdapter -from acp.utils import to_camel_case +from acp.utils import _RouteMetadata, model_to_kwargs, to_camel_case from .exceptions import RequestError @@ -90,21 +90,86 @@ def add_route(self, route: Route) -> None: else: self._notifications[route.method] = route - def _make_func(self, model: type[BaseModel], obj: Any, attr: str) -> AsyncHandler | None: + @classmethod + def from_protocol(cls, protocol: type, obj: Any, *, use_unstable_protocol: bool = False) -> MessageRouter: + """Build routes from decorated protocol members, including inherited ones. + + The protocol is the allowlist: implementation-only methods are never + exposed as RPC routes. Request and notification routes remain distinct + even when they use the same wire method (such as ``mcp/message``). + """ + router = cls(use_unstable_protocol=use_unstable_protocol) + for attr, member in inspect.getmembers(protocol): + metadata = getattr(member, "__route__", None) + if not isinstance(metadata, _RouteMetadata): + continue + validate = metadata.validate_params + if validate is None: + validate = ( + metadata.models[0].model_validate + if len(metadata.models) == 1 + else TypeAdapter(Union[metadata.models]).validate_python # noqa: UP007 - runtime tuple of models + ) + route = Route( + method=metadata.method, + func=router._make_func( + metadata.models[0], + obj, + attr, + validate_params=validate, + adapt_params=metadata.adapt_params, + models=metadata.models, + ), + kind=metadata.kind, + optional=metadata.optional, + default_result=metadata.default_result, + adapt_result=metadata.adapt_result, + warn_unstable=metadata.unstable and not use_unstable_protocol, + ) + routes = router._requests if route.kind == "request" else router._notifications + if route.method in routes: + raise ValueError(f"Duplicate {route.kind} route: {route.method}") + router.add_route(route) + + @router.handle_extension_request + async def _handle_extension_request(name: str, payload: dict[str, Any]) -> Any: + ext = getattr(obj, "ext_method", None) + if ext is None: + raise RequestError.method_not_found(f"_{name}") + return await ext(name, payload) + + @router.handle_extension_notification + async def _handle_extension_notification(name: str, payload: dict[str, Any]) -> None: + ext = getattr(obj, "ext_notification", None) + if ext is not None: + await ext(name, payload) + + return router + + def _make_func( + self, + model: type[BaseModel], + obj: Any, + attr: str, + *, + validate_params: Callable[[Any], BaseModel] | None = None, + adapt_params: Callable[[BaseModel], dict[str, Any]] | None = None, + models: tuple[type[BaseModel], ...] | None = None, + ) -> AsyncHandler | None: func, attr, legacy_api = _resolve_handler(obj, attr) if func is None: return None + validate = validate_params or model.model_validate + param_models = models or (model,) async def wrapper(params: Any) -> Any: if legacy_api: _warn_legacy_handler(obj, attr) - model_obj = model.model_validate(params) + model_obj = validate(params) if legacy_api: - return await func(model_obj) # type: ignore[arg-type] - params = {k: getattr(model_obj, k) for k in model.model_fields if k != "field_meta"} - if meta := getattr(model_obj, "field_meta", None): - params.update(meta) - return await func(**params) # type: ignore[arg-type] + return await func(model_obj) + kwargs = adapt_params(model_obj) if adapt_params else model_to_kwargs(model_obj, param_models) + return await func(**kwargs) return wrapper diff --git a/src/acp/utils.py b/src/acp/utils.py index fc78af8..0df283b 100644 --- a/src/acp/utils.py +++ b/src/acp/utils.py @@ -3,7 +3,8 @@ import functools import warnings from collections.abc import Callable -from typing import Any, TypeVar +from dataclasses import dataclass +from typing import Any, Literal, TypeVar from pydantic import BaseModel @@ -125,25 +126,97 @@ async def notify_model(conn: Connection, method: str, params: BaseModel) -> None await conn.send_notification(method, serialize_params(params)) -def param_model(param_cls: type[BaseModel]) -> Callable[[MethodT], MethodT]: - """Decorator to map the method parameters to a Pydantic model. - It is just a marker and does nothing at runtime. +@dataclass(frozen=True) +class _RouteMetadata: + method: str + models: MultiParamModelSpec + kind: Literal["request", "notification"] = "request" + unstable: bool = False + optional: bool = False + default_result: Any = None + adapt_result: Callable[[Any], Any] | None = None + validate_params: Callable[[Any], BaseModel] | None = None + adapt_params: Callable[[BaseModel], dict[str, Any]] | None = None + + +def param_model( + param_cls: type[BaseModel], + *, + method: str | None = None, + kind: Literal["request", "notification"] = "request", + unstable: bool = False, + optional: bool = False, + default_result: Any = None, + adapt_result: Callable[[Any], Any] | None = None, + validate_params: Callable[[Any], BaseModel] | None = None, + adapt_params: Callable[[BaseModel], dict[str, Any]] | None = None, +) -> Callable[[MethodT], MethodT]: + """Mark a parameter model and optionally declare a protocol route. + + Routing metadata belongs on the Agent/Client protocol declaration. Connection + methods can keep using the model-only form for legacy-call compatibility and + signature generation. The decorated function itself is never wrapped. """ + decorate = param_models( + param_cls, + method=method, + kind=kind, + unstable=unstable, + optional=optional, + default_result=default_result, + adapt_result=adapt_result, + validate_params=validate_params, + adapt_params=adapt_params, + ) def decorator(func: MethodT) -> MethodT: + decorate(func) + delattr(func, "__param_models__") func.__param_model__ = param_cls # type: ignore[attr-defined] return func return decorator -def param_models(*param_cls: type[BaseModel]) -> Callable[[MethodT], MethodT]: - """Decorator to mark a method as accepting multiple legacy parameter models.""" +def param_models( + *param_cls: type[BaseModel], + method: str | None = None, + kind: Literal["request", "notification"] = "request", + unstable: bool = False, + optional: bool = False, + default_result: Any = None, + adapt_result: Callable[[Any], Any] | None = None, + validate_params: Callable[[Any], BaseModel] | None = None, + adapt_params: Callable[[BaseModel], dict[str, Any]] | None = None, +) -> Callable[[MethodT], MethodT]: + """Mark multiple parameter models, with optional protocol routing metadata. + + A route validates the model union and passes fields shared by its models to + the handler. ``validate_params`` and ``adapt_params`` override these steps + for wire formats that differ from the public Python API. + """ if not param_cls: raise ValueError("param_models() requires at least one model class") + metadata = ( + None + if method is None + else _RouteMetadata( + method=method, + models=param_cls, + kind=kind, + unstable=unstable, + optional=optional, + default_result=default_result, + adapt_result=adapt_result, + validate_params=validate_params, + adapt_params=adapt_params, + ) + ) def decorator(func: MethodT) -> MethodT: func.__param_models__ = param_cls # type: ignore[attr-defined] + if metadata is not None: + func.__route__ = metadata # type: ignore[attr-defined] return func return decorator diff --git a/tests/test_elicitation_catchall.py b/tests/test_elicitation_catchall.py index aa94849..7e77e61 100644 --- a/tests/test_elicitation_catchall.py +++ b/tests/test_elicitation_catchall.py @@ -8,7 +8,7 @@ import pytest from pydantic import TypeAdapter, ValidationError -from acp.client.router import _mode_from_create_elicitation_request +from acp._protocol_adapters import _mode_from_create_elicitation_request from acp.exceptions import RequestError from acp.schema import ( AcceptElicitationResponse, diff --git a/tests/test_gen_all.py b/tests/test_gen_all.py index 6091f37..98c6e09 100644 --- a/tests/test_gen_all.py +++ b/tests/test_gen_all.py @@ -82,7 +82,7 @@ def test_signature_generation_preserves_inline_literal_values() -> None: "from typing import Any\n" "from schema import SetProviderRequest, SuggestNesRequest\n" "class Methods:\n" - " @param_model(SetProviderRequest)\n" + ' @param_model(SetProviderRequest, method="providers/set", unstable=True)\n' " async def set_provider(self, **kwargs: Any): ...\n" " @param_model(SuggestNesRequest)\n" " async def suggest_nes(self, **kwargs: Any): ...\n" @@ -96,6 +96,14 @@ def test_signature_generation_preserves_inline_literal_values() -> None: assert "Literal" in {alias.name for alias in typing_import.names} methods = tree.body[-1] assert isinstance(methods, ast.ClassDef) + provider = methods.body[0] + assert isinstance(provider, ast.AsyncFunctionDef) + decorator = provider.decorator_list[0] + assert isinstance(decorator, ast.Call) + assert {keyword.arg: ast.literal_eval(keyword.value) for keyword in decorator.keywords} == { + "method": "providers/set", + "unstable": True, + } suggest = methods.body[-1] assert isinstance(suggest, ast.AsyncFunctionDef) trigger = next(arg for arg in suggest.args.args if arg.arg == "trigger_kind") diff --git a/tests/test_protocol_router.py b/tests/test_protocol_router.py new file mode 100644 index 0000000..4789c85 --- /dev/null +++ b/tests/test_protocol_router.py @@ -0,0 +1,179 @@ +"""Protocols are the source of truth for both signatures and wire routing.""" + +from typing import Any, Literal, Protocol, cast + +import pytest +from pydantic import BaseModel, Field, ValidationError + +from acp.agent.router import build_agent_router +from acp.client.router import build_client_router +from acp.exceptions import RequestError +from acp.interfaces import Agent, Client +from acp.router import MessageRouter +from acp.schema import SetSessionConfigOptionBooleanRequest +from acp.utils import normalize_result, param_model, param_models + + +class Params(BaseModel): + value: str + field_meta: dict[str, Any] | None = Field(None, alias="_meta") + + +class Parent(Protocol): + @param_model(Params, method="example/message", adapt_result=normalize_result) + async def message(self, value: str, **kwargs: Any) -> Any: ... + + +class Example(Parent, Protocol): + @param_model(Params, method="example/message", kind="notification") + async def notify(self, value: str, **kwargs: Any) -> None: ... + + @param_model(Params, method="example/preview", unstable=True) + async def preview(self, value: str, **kwargs: Any) -> Any: ... + + @param_model(Params, method="example/optional", optional=True, default_result={}) + async def optional(self, value: str, **kwargs: Any) -> Any: ... + + @param_model(Params) + async def local(self, value: str, **kwargs: Any) -> Any: ... + + +@pytest.mark.asyncio +async def test_routes_come_from_protocol_including_inherited_members(): + calls = [] + + class Implementation: + async def message(self, value: str, **kwargs: Any): + calls.append(("request", value, kwargs)) + + async def notify(self, value: str, **kwargs: Any): + calls.append(("notification", value, kwargs)) + + async def preview(self, value: str, **kwargs: Any): + return value + + @param_model(Params, method="example/private") + async def private(self, value: str, **kwargs: Any): + raise AssertionError("Implementation-only routes must not be exposed") + + router = MessageRouter.from_protocol(Example, Implementation()) + assert await router("example/message", {"value": "hello", "_meta": {"trace": 1}}, False) == {} + assert await router("example/message", {"value": "bye"}, True) is None + assert calls == [("request", "hello", {"trace": 1}), ("notification", "bye", {})] + assert await router("example/optional", {}, False) == {} + with pytest.raises(ValidationError): + await router("example/message", {}, False) + for method in ["example/private", "local"]: + with pytest.raises(RequestError): + await router(method, {"value": "hidden"}, False) + with pytest.warns(UserWarning, match="unstable"), pytest.raises(RequestError): + await router("example/preview", {"value": "preview"}, False) + enabled = MessageRouter.from_protocol(Example, Implementation(), use_unstable_protocol=True) + assert await enabled("example/preview", {"value": "preview"}, False) == "preview" + + +@pytest.mark.asyncio +async def test_union_route_uses_shared_fields_and_preserves_legacy_model(): + class Text(BaseModel): + value: str + kind: Literal["text"] = "text" + + class Number(BaseModel): + value: int + kind: Literal["number"] = "number" + precision: int = 0 + + class UnionProtocol(Protocol): + @param_models(Text, Number, method="example/union") + async def set_value(self, value: str | int, kind: str) -> Any: ... + + class Modern: + async def set_value(self, value: str | int, kind: str): + return value, kind + + class Legacy: + async def setValue(self, params): + return params + + payload = {"kind": "number", "value": 3, "precision": 2} + router = MessageRouter.from_protocol(UnionProtocol, Modern()) + assert await router("example/union", payload, False) == (3, "number") + legacy = MessageRouter.from_protocol(UnionProtocol, Legacy()) + with pytest.warns(DeprecationWarning): + result = await legacy("example/union", payload, False) + assert isinstance(result, Number) + assert result.precision == 2 + + +def test_duplicate_routes_are_rejected(): + class Duplicate(Parent, Protocol): + @param_model(Params, method="example/message") + async def other(self, value: str, **kwargs: Any) -> Any: ... + + with pytest.raises(ValueError, match="Duplicate request route"): + MessageRouter.from_protocol(Duplicate, object()) + + +@pytest.mark.asyncio +async def test_legacy_config_adapter_receives_boolean_request(): + class Legacy: + async def setConfigOption(self, params): + assert isinstance(params, SetSessionConfigOptionBooleanRequest) + assert params.value is False + return None + + router = build_agent_router(cast(Agent, Legacy())) + with pytest.warns(DeprecationWarning): + assert ( + await router( + "session/set_config_option", + {"sessionId": "s", "configId": "flag", "type": "boolean", "value": False}, + False, + ) + == {} + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("legacy", [False, True]) +async def test_elicitation_custom_mode_keeps_existing_dispatch_behavior(legacy): + class Handler: + async def createElicitation(self, params): + return {"mode": params.mode} + + class Modern: + async def create_elicitation(self, **kwargs): + raise AssertionError("Unsupported modes must fail before invoking the handler") + + router = build_client_router(cast(Client, Handler() if legacy else Modern()), use_unstable_protocol=True) + payload = {"mode": "x-voice", "message": "hi", "sessionId": "s"} + if legacy: + with pytest.warns(DeprecationWarning): + assert await router("elicitation/create", payload, False) == {"mode": "x-voice"} + else: + with pytest.raises(RequestError) as error: + await router("elicitation/create", payload, False) + assert isinstance(error.value, RequestError) + assert error.value.code == -32602 + + +@pytest.mark.asyncio +async def test_missing_handlers_and_extensions(): + router = MessageRouter.from_protocol(Example, object()) + with pytest.raises(RequestError): + await router("example/message", {"value": "hi"}, False) + with pytest.raises(RequestError): + await router("_example/extension", {}, False) + assert await router("_example/extension", {}, True) is None + + class Extensions: + async def ext_method(self, name, payload): + return [name, payload] + + async def ext_notification(self, name, payload): + assert name == "example/extension" + assert payload == {"value": 1} + + router = MessageRouter.from_protocol(Example, Extensions()) + assert await router("_example/extension", {"value": 1}, False) == ["example/extension", {"value": 1}] + assert await router("_example/extension", {"value": 1}, True) is None From c3ffad8d75aa981e7ac8dc15ae521acfb3907061 Mon Sep 17 00:00:00 2001 From: Frost Ming Date: Mon, 21 Sep 2026 16:24:46 +0800 Subject: [PATCH 2/2] refactor: replace param_models with param_model for single model handling Signed-off-by: Frost Ming --- docs/quickstart.md | 11 ++- scripts/gen_signature.py | 63 ++++++++++++- src/acp/client/connection.py | 3 +- src/acp/interfaces.py | 17 ++-- src/acp/router.py | 22 ++--- src/acp/utils.py | 168 ++++++++++------------------------ tests/test_gen_all.py | 61 ++++++++++++ tests/test_param_model.py | 126 +++++++++++++++++++++++++ tests/test_protocol_router.py | 4 +- 9 files changed, 319 insertions(+), 156 deletions(-) create mode 100644 tests/test_param_model.py diff --git a/docs/quickstart.md b/docs/quickstart.md index 9556ee5..15cfb05 100644 --- a/docs/quickstart.md +++ b/docs/quickstart.md @@ -249,11 +249,18 @@ routes; implementation-only methods are not exposed. Use `kind="notification"` for notifications and `unstable=True` for methods requiring explicit opt-in. Optional handlers can declare `optional=True` and `default_result`. -For multiple wire models, `param_models` accepts the same routing options. +`param_model` takes exactly one type expression: a model, a union such as +`ModelA | ModelB`, or `Annotated[ModelA | ModelB, Field(discriminator="type")]`. +The router validates the original type with Pydantic's `TypeAdapter`, preserving +`Annotated` validation metadata. Union handlers receive the fields common to all +branches by default. The former `param_models(A, B, ...)` form is replaced by +`param_model(A | B, ...)`. + `validate_params` and `adapt_params` provide custom validation and conversion when the wire representation differs from the Python signature, as with config options and elicitation. Legacy handlers still receive the validated request -model. Connection methods retain model-only decorators for signature generation +model. Signature generation expands single-model fields and preserves handwritten +union signatures. Connection methods retain model-only decorators for signature generation and legacy calls; they do not repeat the routing options. ## Optional — Talk to the Gemini CLI diff --git a/scripts/gen_signature.py b/scripts/gen_signature.py index 8bbb4ac..cec2e43 100644 --- a/scripts/gen_signature.py +++ b/scripts/gen_signature.py @@ -38,6 +38,24 @@ def __init__(self) -> 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._current_model_name: str | None = None + self._type_aliases: dict[str, ast.expr] = {} + self._schema_names: dict[str, str] = {} + self._annotated_names = {"Annotated"} + self._schema_modules = {"schema"} + + def visit_Module(self, node: ast.Module) -> ast.AST: + for statement in node.body: + if isinstance(statement, ast.Assign): + for target in statement.targets: + if isinstance(target, ast.Name): + self._type_aliases[target.id] = statement.value + elif ( + isinstance(statement, ast.AnnAssign) + and isinstance(statement.target, ast.Name) + and statement.value is not None + ): + self._type_aliases[statement.target.id] = statement.value + return self.generic_visit(node) def _add_typing_import(self, name: str) -> None: if not self._type_import_node: @@ -66,10 +84,46 @@ def transform(self, source_file: Path) -> None: def visit_ImportFrom(self, node: ast.ImportFrom) -> ast.AST: if node.module == "schema": self._schema_import_node = node + self._schema_names.update({alias.asname or alias.name: alias.name for alias in node.names}) elif node.module == "typing": self._type_import_node = node + self._annotated_names.update( + alias.asname or alias.name for alias in node.names if alias.name == "Annotated" + ) + elif node.module is None: + self._schema_modules.update(alias.asname or alias.name for alias in node.names if alias.name == "schema") return node + def _single_param_model(self, expression: ast.expr, seen: frozenset[str] = frozenset()) -> t.Any: + """Resolve single models without evaluating source code or union metadata. + + Union signatures keep their handwritten parameters. + Annotated single models can still expand their underlying model fields. + """ + if isinstance(expression, ast.Name): + name = expression.id + if name in seen: + 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) + 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) + elif isinstance(expression, ast.Subscript): + name = ast.unparse(expression.value) + if name not in self._annotated_names and name != "typing.Annotated": + return None + if not isinstance(expression.slice, ast.Tuple): + return None + return self._single_param_model(expression.slice.elts[0], seen) + else: + return None + while t.get_origin(model) is t.Annotated: + model = t.get_args(model)[0] + return model if inspect.isclass(model) and issubclass(model, BaseModel) else None + def visit_FunctionDef(self, node: ast.FunctionDef) -> ast.AST: return self.visit_func(node) @@ -89,9 +143,12 @@ def visit_func(self, node: ast.FunctionDef | ast.AsyncFunctionDef) -> ast.AST: ) if not decorator: return self.generic_visit(node) - model_name = t.cast(ast.Name, decorator.args[0]).id - model = t.cast(type[schema.BaseModel], getattr(schema, model_name)) - self._current_model_name = model_name + if not decorator.args: + return self.generic_visit(node) + model = self._single_param_model(decorator.args[0]) + if model is None: + return self.generic_visit(node) + self._current_model_name = model.__name__ try: param_defaults = [ self._to_param_def(name, field) for name, field in model.model_fields.items() if name != "field_meta" diff --git a/src/acp/client/connection.py b/src/acp/client/connection.py index 4db0c85..3cd73d2 100644 --- a/src/acp/client/connection.py +++ b/src/acp/client/connection.py @@ -84,7 +84,6 @@ compatible_class, notify_model, param_model, - param_models, request_model, request_model_from_dict, serialize_params, @@ -262,7 +261,7 @@ async def set_session_mode(self, session_id: str, mode_id: str, **kwargs: Any) - SetSessionModeResponse, ) - @param_models(SetSessionConfigOptionBooleanRequest, SetSessionConfigOptionSelectRequest) + @param_model(SetSessionConfigOptionBooleanRequest | SetSessionConfigOptionSelectRequest) async def set_config_option( self, config_id: str, session_id: str, value: str | bool, **kwargs: Any ) -> SetSessionConfigOptionResponse: diff --git a/src/acp/interfaces.py b/src/acp/interfaces.py index 13b2f27..63ddf88 100644 --- a/src/acp/interfaces.py +++ b/src/acp/interfaces.py @@ -123,7 +123,7 @@ WriteTextFileRequest, WriteTextFileResponse, ) -from .utils import normalize_result, param_model, param_models +from .utils import normalize_result, param_model __all__ = ["Agent", "Client"] @@ -227,11 +227,11 @@ async def wait_for_terminal_exit( ) async def kill_terminal(self, session_id: str, terminal_id: str, **kwargs: Any) -> KillTerminalResponse | None: ... - @param_models( - CreateFormSessionElicitationRequest, - CreateFormRequestElicitationRequest, - CreateUrlSessionElicitationRequest, - CreateUrlRequestElicitationRequest, + @param_model( + CreateFormSessionElicitationRequest + | CreateFormRequestElicitationRequest + | CreateUrlSessionElicitationRequest + | CreateUrlRequestElicitationRequest, method=CLIENT_METHODS["elicitation_create"], unstable=True, validate_params=validate_create_elicitation_request, @@ -408,9 +408,8 @@ async def list_sessions( @param_model(SetSessionModeRequest, method=AGENT_METHODS["session_set_mode"], adapt_result=normalize_result) async def set_session_mode(self, session_id: str, mode_id: str, **kwargs: Any) -> SetSessionModeResponse | None: ... - @param_models( - SetSessionConfigOptionBooleanRequest, - SetSessionConfigOptionSelectRequest, + @param_model( + SetSessionConfigOptionBooleanRequest | SetSessionConfigOptionSelectRequest, method=AGENT_METHODS["session_set_config_option"], validate_params=validate_set_config_option_request, adapt_result=normalize_result, diff --git a/src/acp/router.py b/src/acp/router.py index 8c2ce10..f83794f 100644 --- a/src/acp/router.py +++ b/src/acp/router.py @@ -4,7 +4,7 @@ import warnings from collections.abc import Awaitable, Callable from dataclasses import dataclass -from typing import Any, Literal, TypeVar, Union +from typing import Any, Literal, TypeVar from pydantic import BaseModel, TypeAdapter @@ -103,22 +103,14 @@ def from_protocol(cls, protocol: type, obj: Any, *, use_unstable_protocol: bool metadata = getattr(member, "__route__", None) if not isinstance(metadata, _RouteMetadata): continue - validate = metadata.validate_params - if validate is None: - validate = ( - metadata.models[0].model_validate - if len(metadata.models) == 1 - else TypeAdapter(Union[metadata.models]).validate_python # noqa: UP007 - runtime tuple of models - ) route = Route( method=metadata.method, func=router._make_func( - metadata.models[0], + metadata.request_type, obj, attr, - validate_params=validate, + validate_params=metadata.validate_params, adapt_params=metadata.adapt_params, - models=metadata.models, ), kind=metadata.kind, optional=metadata.optional, @@ -148,19 +140,17 @@ async def _handle_extension_notification(name: str, payload: dict[str, Any]) -> def _make_func( self, - model: type[BaseModel], + request_type: Any, obj: Any, attr: str, *, validate_params: Callable[[Any], BaseModel] | None = None, adapt_params: Callable[[BaseModel], dict[str, Any]] | None = None, - models: tuple[type[BaseModel], ...] | None = None, ) -> AsyncHandler | None: func, attr, legacy_api = _resolve_handler(obj, attr) if func is None: return None - validate = validate_params or model.model_validate - param_models = models or (model,) + validate = validate_params or TypeAdapter(request_type).validate_python async def wrapper(params: Any) -> Any: if legacy_api: @@ -168,7 +158,7 @@ async def wrapper(params: Any) -> Any: model_obj = validate(params) if legacy_api: return await func(model_obj) - kwargs = adapt_params(model_obj) if adapt_params else model_to_kwargs(model_obj, param_models) + kwargs = adapt_params(model_obj) if adapt_params else model_to_kwargs(model_obj, request_type) return await func(**kwargs) return wrapper diff --git a/src/acp/utils.py b/src/acp/utils.py index 0df283b..ab628a8 100644 --- a/src/acp/utils.py +++ b/src/acp/utils.py @@ -1,10 +1,11 @@ from __future__ import annotations import functools +import types import warnings from collections.abc import Callable from dataclasses import dataclass -from typing import Any, Literal, TypeVar +from typing import Annotated, Any, Literal, TypeVar, Union, get_args, get_origin from pydantic import BaseModel @@ -27,24 +28,37 @@ MethodT = TypeVar("MethodT", bound=Callable) ClassT = TypeVar("ClassT", bound=type) T = TypeVar("T") -MultiParamModelSpec = tuple[type[BaseModel], ...] -def _param_models_name(models: MultiParamModelSpec) -> str: - return " | ".join(model_type.__name__ for model_type in models) +def _param_model_types(request_type: Any) -> tuple[type[BaseModel], ...]: + """Extract model classes for field expansion, retaining the original type elsewhere. - -def _param_models_field_names(models: MultiParamModelSpec) -> tuple[str, ...]: + Annotated metadata belongs to validation, so only unwrap it for introspection. + All union branches must be models; scalar and container request types are not + supported by the field-expanded Python API. + """ + origin = get_origin(request_type) + if origin is Annotated: + return _param_model_types(get_args(request_type)[0]) + if origin in (Union, types.UnionType): + return tuple(dict.fromkeys(model for branch in get_args(request_type) for model in _param_model_types(branch))) + if isinstance(request_type, type) and issubclass(request_type, BaseModel): + return (request_type,) + raise TypeError("param_model() expects a BaseModel type, a union of model types, or Annotated models") + + +def _param_model_field_names(request_type: Any) -> tuple[str, ...]: + models = _param_model_types(request_type) shared_fields = set(models[0].model_fields) for model_type in models[1:]: shared_fields &= set(model_type.model_fields) return tuple(field_name for field_name in models[0].model_fields if field_name in shared_fields) -def model_to_kwargs(model_obj: BaseModel, models: MultiParamModelSpec) -> dict[str, Any]: +def model_to_kwargs(model_obj: BaseModel, request_type: Any) -> dict[str, Any]: kwargs = { field_name: getattr(model_obj, field_name) - for field_name in _param_models_field_names(models) + for field_name in _param_model_field_names(request_type) if field_name != "field_meta" } if meta := getattr(model_obj, "field_meta", None): @@ -129,7 +143,7 @@ async def notify_model(conn: Connection, method: str, params: BaseModel) -> None @dataclass(frozen=True) class _RouteMetadata: method: str - models: MultiParamModelSpec + request_type: Any kind: Literal["request", "notification"] = "request" unstable: bool = False optional: bool = False @@ -140,7 +154,7 @@ class _RouteMetadata: def param_model( - param_cls: type[BaseModel], + request_type: Any, *, method: str | None = None, kind: Literal["request", "notification"] = "request", @@ -151,58 +165,24 @@ def param_model( validate_params: Callable[[Any], BaseModel] | None = None, adapt_params: Callable[[BaseModel], dict[str, Any]] | None = None, ) -> Callable[[MethodT], MethodT]: - """Mark a parameter model and optionally declare a protocol route. + """Declare one request type: a model, a model union, or Annotated models. Routing metadata belongs on the Agent/Client protocol declaration. Connection - methods can keep using the model-only form for legacy-call compatibility and - signature generation. The decorated function itself is never wrapped. - """ - decorate = param_models( - param_cls, - method=method, - kind=kind, - unstable=unstable, - optional=optional, - default_result=default_result, - adapt_result=adapt_result, - validate_params=validate_params, - adapt_params=adapt_params, - ) - - def decorator(func: MethodT) -> MethodT: - decorate(func) - delattr(func, "__param_models__") - func.__param_model__ = param_cls # type: ignore[attr-defined] - return func - - return decorator - - -def param_models( - *param_cls: type[BaseModel], - method: str | None = None, - kind: Literal["request", "notification"] = "request", - unstable: bool = False, - optional: bool = False, - default_result: Any = None, - adapt_result: Callable[[Any], Any] | None = None, - validate_params: Callable[[Any], BaseModel] | None = None, - adapt_params: Callable[[BaseModel], dict[str, Any]] | None = None, -) -> Callable[[MethodT], MethodT]: - """Mark multiple parameter models, with optional protocol routing metadata. + methods can use the type-only form for signature generation and legacy calls. + Unions pass their common fields to keyword-based handlers; ``adapt_params`` + overrides that conversion. ``validate_params`` overrides TypeAdapter-based + validation when the wire format requires custom branch selection. - A route validates the model union and passes fields shared by its models to - the handler. ``validate_params`` and ``adapt_params`` override these steps - for wire formats that differ from the public Python API. + The original type expression, including Annotated metadata, is preserved. + This decorator never wraps the function. """ - if not param_cls: - raise ValueError("param_models() requires at least one model class") + _param_model_types(request_type) metadata = ( None if method is None else _RouteMetadata( method=method, - models=param_cls, + request_type=request_type, kind=kind, unstable=unstable, optional=optional, @@ -214,7 +194,7 @@ def param_models( ) def decorator(func: MethodT) -> MethodT: - func.__param_models__ = param_cls # type: ignore[attr-defined] + func.__param_model__ = request_type # type: ignore[attr-defined] if metadata is not None: func.__route__ = metadata # type: ignore[attr-defined] return func @@ -228,55 +208,9 @@ def to_camel_case(snake_str: str) -> str: return components[0] + "".join(x.title() for x in components[1:]) -def _make_legacy_func(func: Callable[..., T], model: type[BaseModel]) -> Callable[[Any, BaseModel], T]: - @functools.wraps(func) - def wrapped(self, params: BaseModel) -> T: - warnings.warn( - f"Calling {func.__name__} with {model.__name__} parameter is " # type: ignore[attr-defined] - "deprecated, please update to the new API style.", - DeprecationWarning, - stacklevel=3, - ) - kwargs = { - field_name: getattr(params, field_name) for field_name in model.model_fields if field_name != "field_meta" - } - if meta := getattr(params, "field_meta", None): - kwargs.update(meta) - return func(self, **kwargs) # type: ignore[arg-type] - - return wrapped - - -def _make_compatible_func(func: Callable[..., T], model: type[BaseModel]) -> Callable[..., T]: - @functools.wraps(func) - def wrapped(self, *args: Any, **kwargs: Any) -> T: - param = None - if not kwargs and len(args) == 1: - param = args[0] - elif not args and len(kwargs) == 1: - param = kwargs.get("params") - if isinstance(param, model): - warnings.warn( - f"Calling {func.__name__} with {model.__name__} parameter " # type: ignore[attr-defined] - "is deprecated, please update to the new API style.", - DeprecationWarning, - stacklevel=3, - ) - kwargs = { - field_name: getattr(param, field_name) - for field_name in model.model_fields - if field_name != "field_meta" - } - if meta := getattr(param, "field_meta", None): - kwargs.update(meta) - return func(self, **kwargs) # type: ignore[arg-type] - return func(self, *args, **kwargs) - - return wrapped - - -def _make_multi_legacy_func(func: Callable[..., T], models: MultiParamModelSpec) -> Callable[[Any, BaseModel], T]: - model_name = _param_models_name(models) +def _make_legacy_func(func: Callable[..., T], request_type: Any) -> Callable[[Any, BaseModel], T]: + models = _param_model_types(request_type) + model_name = " | ".join(model.__name__ for model in models) @functools.wraps(func) def wrapped(self, params: BaseModel) -> T: @@ -286,13 +220,14 @@ def wrapped(self, params: BaseModel) -> T: DeprecationWarning, stacklevel=3, ) - return func(self, **model_to_kwargs(params, models)) # type: ignore[arg-type] + return func(self, **model_to_kwargs(params, request_type)) # type: ignore[arg-type] return wrapped -def _make_multi_compatible_func(func: Callable[..., T], models: MultiParamModelSpec) -> Callable[..., T]: - model_name = _param_models_name(models) +def _make_compatible_func(func: Callable[..., T], request_type: Any) -> Callable[..., T]: + models = _param_model_types(request_type) + model_name = " | ".join(model.__name__ for model in models) @functools.wraps(func) def wrapped(self, *args: Any, **kwargs: Any) -> T: @@ -308,7 +243,7 @@ def wrapped(self, *args: Any, **kwargs: Any) -> T: DeprecationWarning, stacklevel=3, ) - return func(self, **model_to_kwargs(param, models)) # type: ignore[arg-type] + return func(self, **model_to_kwargs(param, request_type)) # type: ignore[arg-type] return func(self, *args, **kwargs) return wrapped @@ -320,22 +255,11 @@ def compatible_class(cls: ClassT) -> ClassT: func = getattr(cls, attr) if not callable(func): continue - model = getattr(func, "__param_model__", None) - models = getattr(func, "__param_models__", None) - if model is None and models is None: + request_type = getattr(func, "__param_model__", None) + if request_type is None: continue if "_" in attr: - if models is not None: - setattr(cls, to_camel_case(attr), _make_multi_legacy_func(func, models)) - else: - if model is None: - continue - setattr(cls, to_camel_case(attr), _make_legacy_func(func, model)) + setattr(cls, to_camel_case(attr), _make_legacy_func(func, request_type)) else: - if models is not None: - setattr(cls, attr, _make_multi_compatible_func(func, models)) - else: - if model is None: - continue - setattr(cls, attr, _make_compatible_func(func, model)) + setattr(cls, attr, _make_compatible_func(func, request_type)) return cls diff --git a/tests/test_gen_all.py b/tests/test_gen_all.py index 98c6e09..5f3d324 100644 --- a/tests/test_gen_all.py +++ b/tests/test_gen_all.py @@ -1,5 +1,7 @@ from pathlib import Path +import pytest + from acp.schema import ReadTextFileRequest from scripts.gen_all import resolve_ref, schema_source_paths from scripts.gen_meta import generate_meta @@ -116,3 +118,62 @@ class Annotated: get_type_hints(Annotated, globalns={"Literal": Literal})["trigger"] == Literal["automatic", "diagnostic", "manual"] ) + + +@pytest.mark.parametrize( + "expression", + [ + "SetSessionConfigOptionBooleanRequest | SetSessionConfigOptionSelectRequest", + "Union[SetSessionConfigOptionBooleanRequest, SetSessionConfigOptionSelectRequest]", + "Annotated[SetSessionConfigOptionBooleanRequest | SetSessionConfigOptionSelectRequest, Field(discriminator='type')]", + "RequestAlias", + "CreateElicitationRequest", + ], +) +def test_signature_generation_preserves_union_signatures(expression) -> None: + import ast + + from scripts.gen_signature import NodeTransformer + + tree = ast.parse( + "from typing import Annotated, Any, Union\n" + "from schema import SetSessionConfigOptionBooleanRequest, SetSessionConfigOptionSelectRequest, CreateElicitationRequest\n" + "RequestAlias = SetSessionConfigOptionBooleanRequest | SetSessionConfigOptionSelectRequest\n" + "class Methods:\n" + f" @param_model({expression}, method='example/request')\n" + " async def request(self, config_id: str, session_id: str, value: str | bool, **kwargs: Any): ...\n" + ) + before = ast.dump(tree) + NodeTransformer().visit(tree) + assert ast.dump(tree) == before + + +@pytest.mark.parametrize( + "expression", + [ + "DeleteSessionRequest", + "Annotated[DeleteSessionRequest, Field(title='Delete')]", + "RequestAlias", + "schema.DeleteSessionRequest", + ], +) +def test_signature_generation_expands_single_model_types(expression) -> None: + import ast + + from scripts.gen_signature import NodeTransformer + + tree = ast.parse( + "from typing import Annotated, Any, TypeAlias\n" + "from schema import DeleteSessionRequest\n" + "RequestAlias: TypeAlias = Annotated[DeleteSessionRequest, Field(title='Delete')]\n" + "class Methods:\n" + f" @param_model({expression}, method='session/delete')\n" + " async def delete_session(self, **kwargs: Any): ...\n" + ) + NodeTransformer().visit(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", "session_id"] + assert ast.unparse(method.decorator_list[0]) == f"param_model({expression}, method='session/delete')" diff --git a/tests/test_param_model.py b/tests/test_param_model.py new file mode 100644 index 0000000..0d82ed3 --- /dev/null +++ b/tests/test_param_model.py @@ -0,0 +1,126 @@ +from typing import Annotated, Any, Literal, Protocol, Union + +import pytest +from pydantic import BaseModel, BeforeValidator, Field, ValidationError + +from acp.router import MessageRouter +from acp.utils import compatible_class, param_model + + +class TextRequest(BaseModel): + kind: Literal["text"] = "text" + value: str + field_meta: dict[str, Any] | None = Field(None, alias="_meta") + + +class NumberRequest(BaseModel): + kind: Literal["number"] = "number" + value: int + precision: int = 0 + field_meta: dict[str, Any] | None = Field(None, alias="_meta") + + +TaggedRequest = Annotated[TextRequest | NumberRequest, Field(discriminator="kind")] + + +@pytest.mark.parametrize( + "request_type", + [ + TextRequest, + TextRequest | NumberRequest, + Union[TextRequest, NumberRequest], # noqa: UP007 - cover the typing.Union spelling + TaggedRequest, + Annotated[TextRequest, Field(title="Text")], + Annotated[TextRequest, Field(title="Text")] | NumberRequest, + ], +) +def test_decorator_preserves_original_type_and_function(request_type): + async def handler(**kwargs): + return kwargs + + decorated = param_model(request_type, method="example/request")(handler) + assert decorated is handler + assert decorated.__param_model__ is request_type + assert decorated.__route__.request_type is request_type + assert not hasattr(decorated, "__param_models__") + + +@pytest.mark.parametrize("request_type", [str, Any, list[TextRequest], TextRequest | str, TextRequest | None]) +def test_decorator_rejects_non_model_branches(request_type): + with pytest.raises(TypeError, match="expects a BaseModel"): + param_model(request_type) + + +def test_decorator_accepts_only_one_type_expression(): + with pytest.raises(TypeError): + param_model(TextRequest, NumberRequest) # type: ignore[call-arg] + + +@pytest.mark.asyncio +async def test_discriminated_union_keeps_validation_metadata_and_shared_fields(): + class Interface(Protocol): + @param_model(TaggedRequest, method="example/request") + async def request(self, value: str | int, kind: str, **kwargs: Any) -> Any: ... + + class Handler: + async def request(self, value: str | int, kind: str, **kwargs: Any): + return value, kind, kwargs + + router = MessageRouter.from_protocol(Interface, Handler()) + assert await router( + "example/request", + { + "kind": "number", + "value": 3, + "precision": 2, + "_meta": {"trace": "test"}, + }, + False, + ) == (3, "number", {"trace": "test"}) + # An ordinary union would infer a branch from defaults. A tagged union must + # require its discriminator, proving the Annotated metadata was retained. + with pytest.raises(ValidationError) as error: + await router("example/request", {"value": "hello"}, False) + assert isinstance(error.value, ValidationError) + assert error.value.errors()[0]["type"] == "union_tag_not_found" + + +@pytest.mark.asyncio +async def test_annotated_single_model_keeps_custom_validator(): + def uppercase(payload): + return {**payload, "value": payload["value"].upper()} + + class Interface(Protocol): + @param_model(Annotated[TextRequest, BeforeValidator(uppercase)], method="example/request") + async def request(self, value: str, kind: str, **kwargs: Any) -> str: ... + + class Handler: + async def request(self, value: str, kind: str, **kwargs: Any): + return value + + router = MessageRouter.from_protocol(Interface, Handler()) + assert await router("example/request", {"value": "hello"}, False) == "HELLO" + + +@pytest.mark.parametrize("request_type", [TextRequest | NumberRequest, TaggedRequest]) +def test_union_legacy_calls_use_shared_fields_and_meta(request_type): + @compatible_class + class Calls: + @param_model(request_type) + def set_value(self, value, kind, **kwargs): + return value, kind, kwargs + + @param_model(request_type) + def configure(self, value, kind, **kwargs): + return value, kind, kwargs + + calls = Calls() + request = NumberRequest(value=4, precision=2, _meta={"trace": "test"}) + expected = (4, "number", {"trace": "test"}) + with pytest.warns(DeprecationWarning): + assert getattr(calls, "setValue")(request) == expected # noqa: B009 - dynamically added alias + with pytest.warns(DeprecationWarning): + assert calls.configure(request) == expected + with pytest.warns(DeprecationWarning): + assert calls.configure(params=request) == expected + assert calls.configure(value=4, kind="number", trace="test") == expected diff --git a/tests/test_protocol_router.py b/tests/test_protocol_router.py index 4789c85..6d564d0 100644 --- a/tests/test_protocol_router.py +++ b/tests/test_protocol_router.py @@ -11,7 +11,7 @@ from acp.interfaces import Agent, Client from acp.router import MessageRouter from acp.schema import SetSessionConfigOptionBooleanRequest -from acp.utils import normalize_result, param_model, param_models +from acp.utils import normalize_result, param_model class Params(BaseModel): @@ -84,7 +84,7 @@ class Number(BaseModel): precision: int = 0 class UnionProtocol(Protocol): - @param_models(Text, Number, method="example/union") + @param_model(Text | Number, method="example/union") async def set_value(self, value: str | int, kind: str) -> Any: ... class Modern: