|
2 | 2 | # SPDX-License-Identifier: Apache-2.0 |
3 | 3 |
|
4 | 4 | # ruff: noqa: N802 |
| 5 | +import asyncio |
5 | 6 | from abc import ABC, abstractmethod |
6 | 7 | from collections.abc import AsyncIterable, Callable |
7 | 8 |
|
@@ -320,3 +321,137 @@ async def GetExtendedAgentCard( |
320 | 321 | if self.card_modifier: |
321 | 322 | card_to_serve = self.card_modifier(card_to_serve) |
322 | 323 | return card_to_serve |
| 324 | + |
| 325 | + |
| 326 | +class SRPCSharedHandler(a2a_pb2_slimrpc.A2AServiceSharedServicer): |
| 327 | + """Maps incoming broadcast-live SlimRPC SendLiveMessage calls to the request handler. |
| 328 | +
|
| 329 | + Peer agent StreamResponse events are translated to StreamRequest items and |
| 330 | + merged into the unified inbound stream before on_live_message_send is called, |
| 331 | + so application code sees a single mixed stream of client and peer messages. |
| 332 | + """ |
| 333 | + |
| 334 | + def __init__( |
| 335 | + self, |
| 336 | + agent_card: AgentCard, |
| 337 | + request_handler: RequestHandler, |
| 338 | + context_builder: CallContextBuilder | None = None, |
| 339 | + card_modifier: Callable[[AgentCard], AgentCard] | None = None, |
| 340 | + ) -> None: |
| 341 | + self.agent_card = agent_card |
| 342 | + self.request_handler = request_handler |
| 343 | + self.context_builder = context_builder or DefaultCallContextBuilder() |
| 344 | + self.card_modifier = card_modifier |
| 345 | + |
| 346 | + async def raise_error_response(self, error: A2AError) -> None: |
| 347 | + code = _SLIM_ERROR_CODE_MAP.get(type(error), code_pb2.UNKNOWN) |
| 348 | + raise slim_bindings.RpcError.Rpc( |
| 349 | + code=code, |
| 350 | + message=f"{type(error).__name__}: {error.message}", |
| 351 | + details=None, |
| 352 | + ) |
| 353 | + |
| 354 | + async def SendLiveMessage( |
| 355 | + self, |
| 356 | + request_stream: slim_bindings.RequestStream, |
| 357 | + context: slim_bindings.Context, |
| 358 | + sink: slim_bindings.ResponseSink, |
| 359 | + peer_responses: AsyncIterable, |
| 360 | + ) -> AsyncIterable[a2a_pb2.StreamResponse]: |
| 361 | + """Handles broadcast SendLiveMessage: merges client stream with translated peer events.""" |
| 362 | + server_context = self.context_builder.build(context) |
| 363 | + |
| 364 | + queue: asyncio.Queue[a2a_pb2.StreamRequest | None] = asyncio.Queue() |
| 365 | + |
| 366 | + async def _feed_client() -> None: |
| 367 | + try: |
| 368 | + while True: |
| 369 | + msg = await request_stream.next_async() |
| 370 | + if msg.is_end(): |
| 371 | + break |
| 372 | + if msg.is_error(): |
| 373 | + await queue.put(None) |
| 374 | + return |
| 375 | + if msg.is_data(): |
| 376 | + req = a2a_pb2.StreamRequest.FromString(msg[0]) |
| 377 | + if not server_context.tenant: |
| 378 | + server_context.tenant = req.tenant |
| 379 | + await queue.put(req) |
| 380 | + finally: |
| 381 | + await queue.put(None) |
| 382 | + |
| 383 | + async def _feed_peers() -> None: |
| 384 | + try: |
| 385 | + async for source, stream_response in peer_responses: |
| 386 | + src_str = str(source) |
| 387 | + translated = _translate_peer_response(src_str, stream_response) |
| 388 | + if translated is not None: |
| 389 | + await queue.put(translated) |
| 390 | + finally: |
| 391 | + await queue.put(None) |
| 392 | + |
| 393 | + async def _merged_stream() -> AsyncIterable[a2a_pb2.StreamRequest]: |
| 394 | + client_task = asyncio.ensure_future(_feed_client()) |
| 395 | + peer_task = asyncio.ensure_future(_feed_peers()) |
| 396 | + pending = 2 |
| 397 | + try: |
| 398 | + while pending > 0: |
| 399 | + item = await queue.get() |
| 400 | + if item is None: |
| 401 | + pending -= 1 |
| 402 | + else: |
| 403 | + yield item |
| 404 | + finally: |
| 405 | + client_task.cancel() |
| 406 | + peer_task.cancel() |
| 407 | + |
| 408 | + try: |
| 409 | + async for event in self.request_handler.on_live_message_send( |
| 410 | + _merged_stream(), server_context |
| 411 | + ): |
| 412 | + yield event |
| 413 | + except A2AError as e: |
| 414 | + await self.raise_error_response(e) |
| 415 | + |
| 416 | + |
| 417 | +def _translate_peer_response( |
| 418 | + source: str, |
| 419 | + response: a2a_pb2.StreamResponse, |
| 420 | +) -> a2a_pb2.StreamRequest | None: |
| 421 | + """Translate a peer StreamResponse to a StreamRequest per spec Section 5.""" |
| 422 | + which = response.WhichOneof("payload") |
| 423 | + meta = {"slim-src": source} |
| 424 | + |
| 425 | + if which == "task": |
| 426 | + task = response.task |
| 427 | + meta["slim-peer-task-id"] = task.id |
| 428 | + msg = a2a_pb2.Message( |
| 429 | + role=a2a_pb2.Role.ROLE_USER, |
| 430 | + parts=[a2a_pb2.Part(text=f"peer task started: {task.id}")], |
| 431 | + ) |
| 432 | + return a2a_pb2.StreamRequest(message=msg, metadata=meta) |
| 433 | + |
| 434 | + if which == "status_update": |
| 435 | + update = response.status_update |
| 436 | + meta["slim-peer-task-id"] = update.task_id |
| 437 | + if update.status.HasField("message"): |
| 438 | + return a2a_pb2.StreamRequest(message=update.status.message, metadata=meta) |
| 439 | + state_name = a2a_pb2.TaskState.Name(update.status.state) |
| 440 | + meta["slim-peer-state"] = state_name |
| 441 | + msg = a2a_pb2.Message( |
| 442 | + role=a2a_pb2.Role.ROLE_USER, |
| 443 | + parts=[a2a_pb2.Part(text=f"peer task state: {state_name}")], |
| 444 | + ) |
| 445 | + return a2a_pb2.StreamRequest(message=msg, metadata=meta) |
| 446 | + |
| 447 | + if which == "artifact_update": |
| 448 | + update = response.artifact_update |
| 449 | + meta["slim-peer-task-id"] = update.task_id |
| 450 | + return a2a_pb2.StreamRequest(artifact_update=update, metadata=meta) |
| 451 | + |
| 452 | + if which == "message_update": |
| 453 | + update = response.message_update |
| 454 | + meta["slim-peer-task-id"] = update.task_id |
| 455 | + return a2a_pb2.StreamRequest(message=update.message, metadata=meta) |
| 456 | + |
| 457 | + return None |
0 commit comments