|
1 | 1 | from __future__ import annotations |
2 | 2 |
|
3 | | -from typing import Any |
4 | | - |
5 | | -from pydantic import BaseModel |
6 | | - |
7 | | -from ..exceptions import RequestError |
8 | 3 | 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 |
44 | 5 |
|
45 | 6 | __all__ = ["build_agent_router"] |
46 | 7 |
|
47 | 8 |
|
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 | | - |
73 | 9 | 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