Skip to content

Commit a11bb6e

Browse files
authored
refactor: derive routers from protocol metadata (#149)
* refactor: derive routers from protocol metadata * refactor: replace param_models with param_model for single model handling Signed-off-by: Frost Ming <me@frostming.com>
1 parent 51573d7 commit a11bb6e

13 files changed

Lines changed: 808 additions & 519 deletions

‎docs/quickstart.md‎

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -228,6 +228,41 @@ MCP requests return the inner JSON result unchanged, including `null`. Use
228228
and optional `params`. These methods share the same connections and routers
229229
across stdio, HTTP, and WebSocket transports.
230230

231+
## Maintaining protocol routes
232+
233+
The `Agent` and `Client` protocols in `src/acp/interfaces.py` are the source of
234+
truth for incoming routes. Declare routing metadata alongside each method's
235+
parameter model:
236+
237+
```python
238+
@param_model(
239+
DeleteSessionRequest,
240+
method=AGENT_METHODS["session_delete"],
241+
adapt_result=normalize_result,
242+
)
243+
async def delete_session(self, session_id: str, **kwargs: Any) -> DeleteSessionResponse: ...
244+
```
245+
246+
`build_agent_router` and `build_client_router` read these declarations through
247+
`MessageRouter.from_protocol`. New protocol declarations automatically become
248+
routes; implementation-only methods are not exposed. Use `kind="notification"`
249+
for notifications and `unstable=True` for methods requiring explicit opt-in.
250+
Optional handlers can declare `optional=True` and `default_result`.
251+
252+
`param_model` takes exactly one type expression: a model, a union such as
253+
`ModelA | ModelB`, or `Annotated[ModelA | ModelB, Field(discriminator="type")]`.
254+
The router validates the original type with Pydantic's `TypeAdapter`, preserving
255+
`Annotated` validation metadata. Union handlers receive the fields common to all
256+
branches by default. The former `param_models(A, B, ...)` form is replaced by
257+
`param_model(A | B, ...)`.
258+
259+
`validate_params` and `adapt_params` provide custom validation and conversion
260+
when the wire representation differs from the Python signature, as with config
261+
options and elicitation. Legacy handlers still receive the validated request
262+
model. Signature generation expands single-model fields and preserves handwritten
263+
union signatures. Connection methods retain model-only decorators for signature generation
264+
and legacy calls; they do not repeat the routing options.
265+
231266
## Optional — Talk to the Gemini CLI
232267

233268
_Have the Gemini CLI installed? Run the bridge to exercise permission flows._

‎scripts/gen_signature.py‎

Lines changed: 60 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,24 @@ def __init__(self) -> None:
3838
self._schema_import_node: ast.ImportFrom | None = None
3939
self._literals = {name: value for name, value in schema.__dict__.items() if t.get_origin(value) is t.Literal}
4040
self._current_model_name: str | None = None
41+
self._type_aliases: dict[str, ast.expr] = {}
42+
self._schema_names: dict[str, str] = {}
43+
self._annotated_names = {"Annotated"}
44+
self._schema_modules = {"schema"}
45+
46+
def visit_Module(self, node: ast.Module) -> ast.AST:
47+
for statement in node.body:
48+
if isinstance(statement, ast.Assign):
49+
for target in statement.targets:
50+
if isinstance(target, ast.Name):
51+
self._type_aliases[target.id] = statement.value
52+
elif (
53+
isinstance(statement, ast.AnnAssign)
54+
and isinstance(statement.target, ast.Name)
55+
and statement.value is not None
56+
):
57+
self._type_aliases[statement.target.id] = statement.value
58+
return self.generic_visit(node)
4159

4260
def _add_typing_import(self, name: str) -> None:
4361
if not self._type_import_node:
@@ -66,10 +84,46 @@ def transform(self, source_file: Path) -> None:
6684
def visit_ImportFrom(self, node: ast.ImportFrom) -> ast.AST:
6785
if node.module == "schema":
6886
self._schema_import_node = node
87+
self._schema_names.update({alias.asname or alias.name: alias.name for alias in node.names})
6988
elif node.module == "typing":
7089
self._type_import_node = node
90+
self._annotated_names.update(
91+
alias.asname or alias.name for alias in node.names if alias.name == "Annotated"
92+
)
93+
elif node.module is None:
94+
self._schema_modules.update(alias.asname or alias.name for alias in node.names if alias.name == "schema")
7195
return node
7296

97+
def _single_param_model(self, expression: ast.expr, seen: frozenset[str] = frozenset()) -> t.Any:
98+
"""Resolve single models without evaluating source code or union metadata.
99+
100+
Union signatures keep their handwritten parameters.
101+
Annotated single models can still expand their underlying model fields.
102+
"""
103+
if isinstance(expression, ast.Name):
104+
name = expression.id
105+
if name in seen:
106+
return None
107+
if name in self._type_aliases:
108+
return self._single_param_model(self._type_aliases[name], seen | {name})
109+
model = getattr(schema, self._schema_names.get(name, name), None)
110+
elif isinstance(expression, ast.Attribute) and isinstance(expression.value, ast.Name):
111+
if expression.value.id not in self._schema_modules:
112+
return None
113+
model = getattr(schema, expression.attr, None)
114+
elif isinstance(expression, ast.Subscript):
115+
name = ast.unparse(expression.value)
116+
if name not in self._annotated_names and name != "typing.Annotated":
117+
return None
118+
if not isinstance(expression.slice, ast.Tuple):
119+
return None
120+
return self._single_param_model(expression.slice.elts[0], seen)
121+
else:
122+
return None
123+
while t.get_origin(model) is t.Annotated:
124+
model = t.get_args(model)[0]
125+
return model if inspect.isclass(model) and issubclass(model, BaseModel) else None
126+
73127
def visit_FunctionDef(self, node: ast.FunctionDef) -> ast.AST:
74128
return self.visit_func(node)
75129

@@ -89,9 +143,12 @@ def visit_func(self, node: ast.FunctionDef | ast.AsyncFunctionDef) -> ast.AST:
89143
)
90144
if not decorator:
91145
return self.generic_visit(node)
92-
model_name = t.cast(ast.Name, decorator.args[0]).id
93-
model = t.cast(type[schema.BaseModel], getattr(schema, model_name))
94-
self._current_model_name = model_name
146+
if not decorator.args:
147+
return self.generic_visit(node)
148+
model = self._single_param_model(decorator.args[0])
149+
if model is None:
150+
return self.generic_visit(node)
151+
self._current_model_name = model.__name__
95152
try:
96153
param_defaults = [
97154
self._to_param_def(name, field) for name, field in model.model_fields.items() if name != "field_meta"

‎src/acp/_protocol_adapters.py‎

Lines changed: 72 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,72 @@
1+
from __future__ import annotations
2+
3+
from typing import Any, cast
4+
5+
from pydantic import BaseModel, TypeAdapter
6+
7+
from .exceptions import RequestError
8+
from .schema import (
9+
CreateElicitationRequest,
10+
CreateFormRequestElicitationRequest,
11+
CreateFormSessionElicitationRequest,
12+
CreateUrlRequestElicitationRequest,
13+
CreateUrlSessionElicitationRequest,
14+
ElicitationFormRequestMode,
15+
ElicitationFormSessionMode,
16+
ElicitationUrlRequestMode,
17+
ElicitationUrlSessionMode,
18+
SetSessionConfigOptionBooleanRequest,
19+
SetSessionConfigOptionSelectRequest,
20+
)
21+
22+
_CREATE_ELICITATION_REQUEST_ADAPTER = TypeAdapter(CreateElicitationRequest)
23+
24+
25+
def validate_create_elicitation_request(params: Any) -> CreateElicitationRequest:
26+
return _CREATE_ELICITATION_REQUEST_ADAPTER.validate_python(params)
27+
28+
29+
def _mode_from_create_elicitation_request(
30+
request: CreateElicitationRequest,
31+
) -> ElicitationFormSessionMode | ElicitationFormRequestMode | ElicitationUrlSessionMode | ElicitationUrlRequestMode:
32+
if isinstance(request, CreateFormSessionElicitationRequest):
33+
return ElicitationFormSessionMode(
34+
session_id=request.session_id,
35+
tool_call_id=request.tool_call_id,
36+
requested_schema=request.requested_schema,
37+
)
38+
if isinstance(request, CreateFormRequestElicitationRequest):
39+
return ElicitationFormRequestMode(
40+
request_id=request.request_id,
41+
requested_schema=request.requested_schema,
42+
)
43+
44+
if isinstance(request, CreateUrlSessionElicitationRequest):
45+
return ElicitationUrlSessionMode(
46+
session_id=request.session_id,
47+
tool_call_id=request.tool_call_id,
48+
elicitation_id=request.elicitation_id,
49+
url=request.url,
50+
)
51+
if isinstance(request, CreateUrlRequestElicitationRequest):
52+
return ElicitationUrlRequestMode(
53+
request_id=request.request_id,
54+
elicitation_id=request.elicitation_id,
55+
url=request.url,
56+
)
57+
raise RequestError.invalid_params({"details": f"Unsupported elicitation mode: {request.mode!r}"})
58+
59+
60+
def elicitation_to_kwargs(request: BaseModel) -> dict[str, Any]:
61+
# The validator has already resolved the wire union, including custom modes.
62+
request = cast(CreateElicitationRequest, request)
63+
kwargs = {"message": request.message, "mode": _mode_from_create_elicitation_request(request)}
64+
if request.field_meta:
65+
kwargs.update(request.field_meta)
66+
return kwargs
67+
68+
69+
def validate_set_config_option_request(params: Any) -> BaseModel:
70+
if isinstance(params, dict) and params.get("type") == "boolean":
71+
return SetSessionConfigOptionBooleanRequest.model_validate(params)
72+
return SetSessionConfigOptionSelectRequest.model_validate(params)

‎src/acp/agent/router.py‎

Lines changed: 2 additions & 178 deletions
Original file line numberDiff line numberDiff line change
@@ -1,186 +1,10 @@
11
from __future__ import annotations
22

3-
from typing import Any
4-
5-
from pydantic import BaseModel
6-
7-
from ..exceptions import RequestError
83
from ..interfaces import Agent
9-
from ..meta import AGENT_METHODS
10-
from ..router import MessageRouter, Route, _resolve_handler, _warn_legacy_handler
11-
from ..schema import (
12-
AcceptNesNotification,
13-
AuthenticateRequest,
14-
CancelNotification,
15-
CloseNesRequest,
16-
CloseSessionRequest,
17-
DeleteSessionRequest,
18-
DidChangeDocumentNotification,
19-
DidCloseDocumentNotification,
20-
DidFocusDocumentNotification,
21-
DidOpenDocumentNotification,
22-
DidSaveDocumentNotification,
23-
DisableProviderRequest,
24-
ForkSessionRequest,
25-
InitializeRequest,
26-
ListProvidersRequest,
27-
ListSessionsRequest,
28-
LoadSessionRequest,
29-
LogoutRequest,
30-
MessageMcpNotification,
31-
MessageMcpRequest,
32-
NewSessionRequest,
33-
PromptRequest,
34-
RejectNesNotification,
35-
ResumeSessionRequest,
36-
SetProviderRequest,
37-
SetSessionConfigOptionBooleanRequest,
38-
SetSessionConfigOptionSelectRequest,
39-
SetSessionModeRequest,
40-
StartNesRequest,
41-
SuggestNesRequest,
42-
)
43-
from ..utils import model_to_kwargs, normalize_result
4+
from ..router import MessageRouter
445

456
__all__ = ["build_agent_router"]
467

478

48-
_SET_CONFIG_OPTION_MODELS = (SetSessionConfigOptionBooleanRequest, SetSessionConfigOptionSelectRequest)
49-
50-
51-
def _validate_set_config_option_request(params: Any) -> BaseModel:
52-
if isinstance(params, dict) and params.get("type") == "boolean":
53-
return SetSessionConfigOptionBooleanRequest.model_validate(params)
54-
return SetSessionConfigOptionSelectRequest.model_validate(params)
55-
56-
57-
def _make_set_config_option_handler(agent: Agent) -> Any:
58-
func, attr, legacy_api = _resolve_handler(agent, "set_config_option")
59-
if func is None:
60-
return None
61-
62-
async def wrapper(params: Any) -> Any:
63-
if legacy_api:
64-
_warn_legacy_handler(agent, attr)
65-
request = _validate_set_config_option_request(params)
66-
if legacy_api:
67-
return await func(request)
68-
return await func(**model_to_kwargs(request, _SET_CONFIG_OPTION_MODELS))
69-
70-
return wrapper
71-
72-
739
def build_agent_router(agent: Agent, use_unstable_protocol: bool = False) -> MessageRouter:
74-
router = MessageRouter(use_unstable_protocol=use_unstable_protocol)
75-
76-
router.route_request(AGENT_METHODS["initialize"], InitializeRequest, agent, "initialize")
77-
router.route_request(AGENT_METHODS["session_new"], NewSessionRequest, agent, "new_session")
78-
router.route_request(
79-
AGENT_METHODS["session_load"],
80-
LoadSessionRequest,
81-
agent,
82-
"load_session",
83-
adapt_result=normalize_result,
84-
)
85-
router.route_request(AGENT_METHODS["session_list"], ListSessionsRequest, agent, "list_sessions")
86-
router.route_request(
87-
AGENT_METHODS["session_close"],
88-
CloseSessionRequest,
89-
agent,
90-
"close_session",
91-
adapt_result=normalize_result,
92-
unstable=True,
93-
)
94-
router.route_request(
95-
AGENT_METHODS["session_set_mode"],
96-
SetSessionModeRequest,
97-
agent,
98-
"set_session_mode",
99-
adapt_result=normalize_result,
100-
)
101-
router.route_request(AGENT_METHODS["session_prompt"], PromptRequest, agent, "prompt")
102-
router.add_route(
103-
Route(
104-
method=AGENT_METHODS["session_set_config_option"],
105-
func=_make_set_config_option_handler(agent),
106-
kind="request",
107-
adapt_result=normalize_result,
108-
)
109-
)
110-
router.route_request(
111-
AGENT_METHODS["authenticate"],
112-
AuthenticateRequest,
113-
agent,
114-
"authenticate",
115-
adapt_result=normalize_result,
116-
)
117-
router.route_request(AGENT_METHODS["session_fork"], ForkSessionRequest, agent, "fork_session", unstable=True)
118-
router.route_request(AGENT_METHODS["session_resume"], ResumeSessionRequest, agent, "resume_session", unstable=True)
119-
120-
router.route_notification(AGENT_METHODS["session_cancel"], CancelNotification, agent, "cancel")
121-
122-
router.route_request(
123-
AGENT_METHODS["session_delete"],
124-
DeleteSessionRequest,
125-
agent,
126-
"delete_session",
127-
adapt_result=normalize_result,
128-
)
129-
router.route_request(AGENT_METHODS["providers_list"], ListProvidersRequest, agent, "list_providers", unstable=True)
130-
router.route_request(
131-
AGENT_METHODS["providers_set"],
132-
SetProviderRequest,
133-
agent,
134-
"set_provider",
135-
unstable=True,
136-
adapt_result=normalize_result,
137-
)
138-
router.route_request(
139-
AGENT_METHODS["providers_disable"],
140-
DisableProviderRequest,
141-
agent,
142-
"disable_provider",
143-
unstable=True,
144-
adapt_result=normalize_result,
145-
)
146-
router.route_request(AGENT_METHODS["logout"], LogoutRequest, agent, "logout", adapt_result=normalize_result)
147-
router.route_request(AGENT_METHODS["mcp_message"], MessageMcpRequest, agent, "mcp_message", unstable=True)
148-
router.route_notification(AGENT_METHODS["mcp_message"], MessageMcpNotification, agent, "notify_mcp", unstable=True)
149-
router.route_request(AGENT_METHODS["nes_start"], StartNesRequest, agent, "start_nes", unstable=True)
150-
router.route_request(AGENT_METHODS["nes_suggest"], SuggestNesRequest, agent, "suggest_nes", unstable=True)
151-
router.route_request(
152-
AGENT_METHODS["nes_close"], CloseNesRequest, agent, "close_nes", unstable=True, adapt_result=normalize_result
153-
)
154-
router.route_notification(AGENT_METHODS["nes_accept"], AcceptNesNotification, agent, "accept_nes", unstable=True)
155-
router.route_notification(AGENT_METHODS["nes_reject"], RejectNesNotification, agent, "reject_nes", unstable=True)
156-
router.route_notification(
157-
AGENT_METHODS["document_did_open"], DidOpenDocumentNotification, agent, "did_open", unstable=True
158-
)
159-
router.route_notification(
160-
AGENT_METHODS["document_did_change"], DidChangeDocumentNotification, agent, "did_change", unstable=True
161-
)
162-
router.route_notification(
163-
AGENT_METHODS["document_did_close"], DidCloseDocumentNotification, agent, "did_close", unstable=True
164-
)
165-
router.route_notification(
166-
AGENT_METHODS["document_did_save"], DidSaveDocumentNotification, agent, "did_save", unstable=True
167-
)
168-
router.route_notification(
169-
AGENT_METHODS["document_did_focus"], DidFocusDocumentNotification, agent, "did_focus", unstable=True
170-
)
171-
172-
@router.handle_extension_request
173-
async def _handle_extension_request(name: str, payload: dict[str, Any]) -> Any:
174-
ext = getattr(agent, "ext_method", None)
175-
if ext is None:
176-
raise RequestError.method_not_found(f"_{name}")
177-
return await ext(name, payload)
178-
179-
@router.handle_extension_notification
180-
async def _handle_extension_notification(name: str, payload: dict[str, Any]) -> None:
181-
ext = getattr(agent, "ext_notification", None)
182-
if ext is None:
183-
return
184-
await ext(name, payload)
185-
186-
return router
10+
return MessageRouter.from_protocol(Agent, agent, use_unstable_protocol=use_unstable_protocol)

0 commit comments

Comments
 (0)