|
2 | 2 |
|
3 | 3 | import asyncio |
4 | 4 | from collections.abc import Callable |
5 | | -from dataclasses import dataclass |
6 | 5 | from typing import Any, cast |
7 | 6 |
|
8 | 7 | from pydantic import BaseModel |
9 | 8 |
|
10 | 9 | from acp import meta as v1_meta |
11 | 10 | from acp import schema as v1_schema |
12 | 11 | from acp.agent.connection import AgentSideConnection as V1AgentSideConnection |
13 | | -from acp.client.connection import ClientSideConnection as V1ClientSideConnection |
14 | 12 | from acp.connection import Connection, MethodHandler |
15 | 13 | from acp.exceptions import RequestError |
16 | 14 | from acp.interfaces import Agent as V1Agent |
|
20 | 18 | from .v2._connection import open_connection |
21 | 19 | from .v2.agent import AgentFactory as V2AgentFactory |
22 | 20 | from .v2.agent import AgentSideConnection as V2AgentSideConnection |
23 | | -from .v2.client import ClientFactory as V2ClientFactory |
24 | | -from .v2.client import ClientSideConnection as V2ClientSideConnection |
25 | 21 |
|
26 | 22 | __all__ = [ |
27 | 23 | "AgentProtocolConnection", |
28 | 24 | "AgentProtocolRouter", |
29 | | - "ClientNegotiator", |
30 | | - "NegotiatedClient", |
31 | | - "NegotiatedV1", |
32 | | - "NegotiatedV2", |
33 | | - "UnsupportedProtocolVersionError", |
34 | | - "V1ClientConfig", |
35 | | - "V2ClientConfig", |
36 | 25 | ] |
37 | 26 |
|
38 | 27 | V1AgentFactory = Callable[[V1Client], V1Agent] |
39 | | -V1ClientFactory = Callable[[V1Agent], V1Client] |
40 | 28 |
|
41 | 29 |
|
42 | 30 | def _dump(model: BaseModel) -> dict[str, Any]: |
@@ -80,33 +68,6 @@ def _normalize_initialize(params: Any, selected_version: int) -> dict[str, Any]: |
80 | 68 | return _dump(request) |
81 | 69 |
|
82 | 70 |
|
83 | | -class _SwitchingHandler: |
84 | | - def __init__(self) -> None: |
85 | | - self._handler: MethodHandler | None = None |
86 | | - self._failure: BaseException | None = None |
87 | | - self._ready = asyncio.Event() |
88 | | - |
89 | | - def bind(self, handler: MethodHandler) -> None: |
90 | | - if self._handler is not None or self._failure is not None: |
91 | | - raise RuntimeError("Protocol handler has already been resolved") |
92 | | - self._handler = handler |
93 | | - self._ready.set() |
94 | | - |
95 | | - def fail(self, error: BaseException) -> None: |
96 | | - if self._handler is not None: |
97 | | - return |
98 | | - self._failure = error |
99 | | - self._ready.set() |
100 | | - |
101 | | - async def __call__(self, method: str, params: Any | None, is_notification: bool) -> Any: |
102 | | - await self._ready.wait() |
103 | | - if self._failure is not None: |
104 | | - raise self._failure |
105 | | - if self._handler is None: |
106 | | - raise RuntimeError("Protocol handler was not resolved") |
107 | | - return await self._handler(method, params, is_notification) |
108 | | - |
109 | | - |
110 | 71 | class _AgentNegotiationHandler: |
111 | 72 | def __init__( |
112 | 73 | self, |
@@ -246,130 +207,3 @@ async def run( |
246 | 207 | await connection.listen() |
247 | 208 | finally: |
248 | 209 | await asyncio.shield(connection.close()) |
249 | | - |
250 | | - |
251 | | -@dataclass(frozen=True, slots=True) |
252 | | -class V1ClientConfig: |
253 | | - client: V1ClientFactory | V1Client |
254 | | - initialize: v1_schema.InitializeRequest |
255 | | - |
256 | | - def __post_init__(self) -> None: |
257 | | - if self.initialize.protocol_version != v1_meta.PROTOCOL_VERSION: |
258 | | - raise ValueError(f"V1ClientConfig requires protocol version {v1_meta.PROTOCOL_VERSION}") |
259 | | - |
260 | | - |
261 | | -@dataclass(frozen=True, slots=True) |
262 | | -class V2ClientConfig: |
263 | | - client: V2ClientFactory | v2.Client |
264 | | - initialize: v2.schema.InitializeRequest |
265 | | - |
266 | | - def __post_init__(self) -> None: |
267 | | - if self.initialize.protocol_version != v2.PROTOCOL_VERSION: |
268 | | - raise ValueError(f"V2ClientConfig requires protocol version {v2.PROTOCOL_VERSION}") |
269 | | - |
270 | | - |
271 | | -@dataclass(frozen=True, slots=True) |
272 | | -class NegotiatedV1: |
273 | | - connection: V1ClientSideConnection |
274 | | - initialize: v1_schema.InitializeResponse |
275 | | - protocol_version: int = v1_meta.PROTOCOL_VERSION |
276 | | - |
277 | | - |
278 | | -@dataclass(frozen=True, slots=True) |
279 | | -class NegotiatedV2: |
280 | | - connection: V2ClientSideConnection |
281 | | - initialize: v2.schema.InitializeResponse |
282 | | - protocol_version: int = v2.PROTOCOL_VERSION |
283 | | - |
284 | | - |
285 | | -NegotiatedClient = NegotiatedV1 | NegotiatedV2 |
286 | | - |
287 | | - |
288 | | -class UnsupportedProtocolVersionError(ValueError): |
289 | | - def __init__(self, requested: int, offered: int, supported: frozenset[int]) -> None: |
290 | | - self.requested = requested |
291 | | - self.offered = offered |
292 | | - self.supported = supported |
293 | | - super().__init__(f"Agent selected ACP protocol {offered}; requested {requested}, supported {sorted(supported)}") |
294 | | - |
295 | | - |
296 | | -class ClientNegotiator: |
297 | | - """Send one initialize request and return the selected typed client.""" |
298 | | - |
299 | | - def __init__( |
300 | | - self, |
301 | | - input_stream: Any, |
302 | | - output_stream: Any = None, |
303 | | - *, |
304 | | - v1: V1ClientConfig | None = None, |
305 | | - v2: V2ClientConfig | None = None, |
306 | | - **connection_kwargs: Any, |
307 | | - ) -> None: |
308 | | - if v1 is None and v2 is None: |
309 | | - raise ValueError("Configure at least one ACP client version") |
310 | | - self._v1 = v1 |
311 | | - self._v2 = v2 |
312 | | - self._handler = _SwitchingHandler() |
313 | | - self._connection = open_connection( |
314 | | - self._handler, |
315 | | - input_stream, |
316 | | - output_stream, |
317 | | - **connection_kwargs, |
318 | | - ) |
319 | | - self._lock = asyncio.Lock() |
320 | | - self._resolved: NegotiatedClient | None = None |
321 | | - self._failure: BaseException | None = None |
322 | | - |
323 | | - async def negotiate(self) -> NegotiatedClient: |
324 | | - async with self._lock: |
325 | | - if self._resolved is not None: |
326 | | - return self._resolved |
327 | | - if self._failure is not None: |
328 | | - raise self._failure |
329 | | - try: |
330 | | - self._resolved = await self._negotiate_once() |
331 | | - except BaseException as error: |
332 | | - self._failure = error |
333 | | - self._handler.fail(error) |
334 | | - await self._connection.close() |
335 | | - raise |
336 | | - return self._resolved |
337 | | - |
338 | | - async def _negotiate_once(self) -> NegotiatedClient: |
339 | | - offered_request: BaseModel = ( |
340 | | - self._v2.initialize if self._v2 is not None else cast(V1ClientConfig, self._v1).initialize |
341 | | - ) |
342 | | - |
343 | | - response = await self._connection.send_request(v2.AGENT_METHODS["initialize"], _dump(offered_request)) |
344 | | - offered = _read_protocol_version(response) |
345 | | - requested = offered_request.protocol_version |
346 | | - |
347 | | - if offered == v2.PROTOCOL_VERSION and self._v2 is not None: |
348 | | - initialize = v2.schema.InitializeResponse.model_validate(response) |
349 | | - connection, handler = V2ClientSideConnection._attach(self._v2.client, self._connection) |
350 | | - connection._complete_initialization(self._v2.initialize, initialize) |
351 | | - self._handler.bind(handler) |
352 | | - return NegotiatedV2(connection, initialize) |
353 | | - if offered == v1_meta.PROTOCOL_VERSION and self._v1 is not None: |
354 | | - initialize = v1_schema.InitializeResponse.model_validate(response) |
355 | | - connection, handler = V1ClientSideConnection._attach(self._v1.client, self._connection) |
356 | | - self._handler.bind(handler) |
357 | | - return NegotiatedV1(connection, initialize) |
358 | | - supported = frozenset( |
359 | | - version |
360 | | - for version, config in ( |
361 | | - (v1_meta.PROTOCOL_VERSION, self._v1), |
362 | | - (v2.PROTOCOL_VERSION, self._v2), |
363 | | - ) |
364 | | - if config is not None |
365 | | - ) |
366 | | - raise UnsupportedProtocolVersionError(requested, offered, supported) |
367 | | - |
368 | | - async def close(self) -> None: |
369 | | - await self._connection.close() |
370 | | - |
371 | | - async def __aenter__(self) -> ClientNegotiator: |
372 | | - return self |
373 | | - |
374 | | - async def __aexit__(self, exc_type: Any, exc: Any, tb: Any) -> None: |
375 | | - await self.close() |
0 commit comments