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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 35 additions & 0 deletions docs/quickstart.md
Original file line number Diff line number Diff line change
Expand Up @@ -228,6 +228,41 @@ 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`.

`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. 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

_Have the Gemini CLI installed? Run the bridge to exercise permission flows._
Expand Down
63 changes: 60 additions & 3 deletions scripts/gen_signature.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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)

Expand All @@ -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"
Expand Down
72 changes: 72 additions & 0 deletions src/acp/_protocol_adapters.py
Original file line number Diff line number Diff line change
@@ -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)
180 changes: 2 additions & 178 deletions src/acp/agent/router.py
Original file line number Diff line number Diff line change
@@ -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)
Loading
Loading