Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/agent_endpoints/a2a_endpoints.py: 14%

400 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1# pyright: reportUnknownArgumentType=false 

2# This module forwards JSON-RPC payloads through the untyped a2a-sdk compat 

3# conversions (pb2_v10/ParseDict/MessageToDict/to_compat_*), so SDK and decoded-JSON 

4# values flow in as Unknown. The rule is off file-wide rather than scattering per-line 

5# ignores across every SDK and JSON-RPC call. 

6""" 

7A2A Protocol endpoints for LiteLLM Proxy. 

8 

9Allows clients to invoke agents through LiteLLM using the A2A protocol. 

10The A2A SDK can point to LiteLLM's URL and invoke agents registered with LiteLLM. 

11""" 

12 

13import json 

14from collections.abc import AsyncGenerator, Mapping 

15from copy import deepcopy 

16from types import MappingProxyType 

17from typing import TYPE_CHECKING, Any, Final, Protocol 

18from urllib.parse import urlparse 

19 

20from fastapi import APIRouter, Depends, HTTPException, Request, Response 

21from fastapi.responses import JSONResponse, StreamingResponse 

22from pydantic import ValidationError 

23 

24import litellm 

25from litellm._logging import verbose_proxy_logger 

26from litellm.litellm_core_utils.url_utils import SSRFError, validate_url 

27from litellm.llms.a2a.common_utils import resolve_a2a_hop_auth_header 

28from litellm.proxy._types import UserAPIKeyAuth 

29from litellm.proxy.a2a.version_convert import ( 

30 A2AVersion, 

31 normalize_agent_card, 

32 normalize_jsonrpc_response, 

33 normalize_request_params, 

34 normalize_stream_event, 

35) 

36from litellm.proxy.agent_endpoints.databricks_oauth import ( 

37 DATABRICKS_OAUTH_PARAM, 

38 resolve_databricks_app_auth_header, 

39) 

40from litellm.proxy.agent_endpoints.utils import merge_agent_headers 

41from litellm.proxy.auth.user_api_key_auth import user_api_key_auth 

42from litellm.proxy.common_utils.sse_keepalive import ( 

43 SSE_COMMENT_PING, 

44 coerce_keepalive_interval, 

45 wrap_sse_stream_with_keepalive_pings, 

46) 

47from litellm.proxy.utils import ProxyLogging, get_custom_url 

48from litellm.types.utils import all_litellm_params 

49 

50if TYPE_CHECKING: 50 ↛ 51line 50 didn't jump to line 51 because the condition on line 50 was never true

51 from a2a.compat.v0_3.types import MessageSendParams 

52 

53 from litellm.types.agents import AgentResponse 

54 

55router: Final = APIRouter() 

56 

57# Mirrors the native seam's own headers: a reverse proxy that batches the whole 

58# stream would swallow the keepalives this route sends to defeat idle timeouts. 

59_SSE_KEEPALIVE_HEADERS: Final[Mapping[str, str]] = MappingProxyType( 

60 { 

61 "Cache-Control": "no-cache", 

62 "X-Accel-Buffering": "no", 

63 } 

64) 

65 

66_PASCAL_TO_WIRE: Final[Mapping[str, str]] = { 

67 "SendMessage": "message/send", 

68 "SendStreamingMessage": "message/stream", 

69 "GetTask": "tasks/get", 

70 "ListTasks": "tasks/list", 

71 "CancelTask": "tasks/cancel", 

72 "SubscribeToTask": "tasks/resubscribe", 

73 "CreateTaskPushNotificationConfig": "tasks/pushNotificationConfig/set", 

74 "GetTaskPushNotificationConfig": "tasks/pushNotificationConfig/get", 

75 "ListTaskPushNotificationConfigs": "tasks/pushNotificationConfig/list", 

76 "DeleteTaskPushNotificationConfig": "tasks/pushNotificationConfig/delete", 

77 "GetExtendedAgentCard": "agent/getAuthenticatedExtendedCard", 

78} 

79 

80 

81def _sse_event(payload: object) -> str: 

82 """Frame a JSON-RPC object as a single A2A SSE event (``data: <json>\\n\\n``).""" 

83 return f"data: {json.dumps(payload)}\n\n" 

84 

85 

86def _to_jsonrpc_object(chunk: object) -> object: 

87 """Coerce a streamed chunk to the JSON-RPC object it carries. 

88 

89 Chunks arrive as SDK models, plain dicts, or, when a guardrail terminates a 

90 stream, as an already serialized JSON-RPC object. 

91 """ 

92 if isinstance(chunk, (str, bytes, bytearray)): 

93 try: 

94 return json.loads(chunk) 

95 except (json.JSONDecodeError, UnicodeDecodeError): 

96 return chunk 

97 if hasattr(chunk, "model_dump"): 

98 return chunk.model_dump(mode="json", exclude_none=True) 

99 return chunk 

100 

101 

102def _build_message_send_params(params: dict[str, Any]) -> "MessageSendParams": 

103 """Build MessageSendParams from wire (0.3) or A2A 1.0 JSON-RPC params.""" 

104 from a2a.compat.v0_3.types import MessageSendParams 

105 

106 try: 

107 return MessageSendParams(**params) 

108 except ValidationError: 

109 from a2a.compat.v0_3.conversions import pb2_v10, to_compat_send_message_request 

110 from google.protobuf.json_format import ParseDict, ParseError 

111 

112 pb: Final = pb2_v10.SendMessageRequest() 

113 try: 

114 ParseDict(params, pb, ignore_unknown_fields=True) 

115 except ParseError as e: 

116 raise ValueError(f"Invalid message/send params: {e}") from e 

117 return to_compat_send_message_request(pb, "").params 

118 

119 

120def _served_version(agent: "AgentResponse", request: Request, original_method: str | None = None) -> A2AVersion: 

121 """Protocol version LiteLLM serves for this agent. 

122 

123 The agent's configured version governs. For agents that pin no version, fall back 

124 to the client's signal: PascalCase JSON-RPC methods and an ``a2a-version: 1.x`` 

125 header both mark a 1.0 caller; otherwise default to 0.3. 

126 """ 

127 configured: Final = (agent.agent_card_params or {}).get("protocolVersion") 

128 if configured in ("0.3", "1.0"): 

129 return configured 

130 if original_method in _PASCAL_TO_WIRE: 

131 return "1.0" 

132 return "1.0" if request.headers.get("a2a-version", "").startswith("1.") else "0.3" 

133 

134 

135def _validate_push_notification_url(url: str) -> None: 

136 parsed: Final = urlparse(url) 

137 if parsed.scheme != "https": 

138 raise HTTPException( 

139 status_code=400, 

140 detail="Push notification URL must use HTTPS", 

141 ) 

142 try: 

143 validate_url(url) 

144 except (SSRFError, ValueError) as e: 

145 raise HTTPException(status_code=400, detail=str(e)) from e 

146 

147 

148def _caller_identity_headers(user_api_key_dict: UserAPIKeyAuth) -> Mapping[str, str]: 

149 """The human behind this call. An agent key acting for an invoking user forwards that user, not 

150 itself, so a chain of agents stays capped at what the original caller may reach.""" 

151 caller: Final = user_api_key_dict.agent_caller 

152 user_id: Final = caller.user_id if caller is not None else user_api_key_dict.user_id 

153 team_id: Final = caller.team_id if caller is not None else user_api_key_dict.team_id 

154 return MappingProxyType( 

155 { 

156 name: value 

157 for name, value in ( 

158 ("X-LiteLLM-User-Id", user_id), 

159 ("X-LiteLLM-Team-Id", team_id), 

160 ) 

161 if value 

162 } 

163 ) 

164 

165 

166async def _resolve_backend_auth_header( 

167 litellm_params: dict[str, object], 

168 custom_llm_provider: object, 

169) -> Mapping[str, str] | None: 

170 if litellm_params.get(DATABRICKS_OAUTH_PARAM): 

171 return await resolve_databricks_app_auth_header(litellm_params) 

172 return await resolve_a2a_hop_auth_header(litellm_params, custom_llm_provider) 

173 

174 

175def _forwarding_headers( 

176 caller_identity: Mapping[str, str], 

177 request_data: Mapping[str, object], 

178 agent_extra_headers: Mapping[str, str] | None, 

179 backend_auth_header: Mapping[str, str] | None, 

180) -> dict[str, str] | None: 

181 backend_auth: Final = tuple(backend_auth_header.items()) if backend_auth_header else () 

182 minted_names: Final = frozenset(name.lower() for name, _ in backend_auth) 

183 passthrough: Final = tuple( 

184 (name, value) 

185 for name, value in (agent_extra_headers.items() if agent_extra_headers else ()) 

186 if not name.lower().startswith("x-litellm-") and name.lower() not in minted_names 

187 ) 

188 trace_id: Final = request_data.get("litellm_trace_id") 

189 trace: Final = (("X-LiteLLM-Trace-Id", str(trace_id)),) if trace_id else () 

190 merged: Final = dict((*passthrough, *caller_identity.items(), *trace, *backend_auth)) 

191 return merged or None 

192 

193 

194def _jsonrpc_error( 

195 request_id: object, 

196 code: int, 

197 message: str, 

198 status_code: int = 400, 

199) -> JSONResponse: 

200 """Create a JSON-RPC 2.0 error response.""" 

201 return JSONResponse( 

202 content={ 

203 "jsonrpc": "2.0", 

204 "id": request_id, 

205 "error": {"code": code, "message": message}, 

206 }, 

207 status_code=status_code, 

208 ) 

209 

210 

211async def _get_agent(agent_id: str) -> "AgentResponse | None": 

212 """Look up an agent by ID or name. Returns None if not found.""" 

213 from litellm.proxy.common_utils.registry_read_through import ( 

214 get_agent_with_read_through, 

215 ) 

216 

217 return await get_agent_with_read_through(agent_id) 

218 

219 

220def _enforce_inbound_trace_id(agent: "AgentResponse", request: Request) -> None: 

221 """Raise 400 if agent requires x-litellm-trace-id on inbound calls and it is missing.""" 

222 agent_litellm_params: Final = agent.litellm_params or {} 

223 if not agent_litellm_params.get("require_trace_id_on_calls_to_agent"): 

224 return 

225 

226 from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers 

227 

228 headers_dict: Final = dict(request.headers) 

229 trace_id: Final = get_chain_id_from_headers(headers_dict) 

230 if not trace_id: 

231 raise HTTPException( 

232 status_code=400, 

233 detail=(f"Agent '{agent.agent_id}' requires x-litellm-trace-id header on all inbound requests."), 

234 ) 

235 

236 

237class _JsonRpcResponse(Protocol): 

238 def json(self) -> dict[str, object]: ... 238 ↛ exitline 238 didn't return from function 'json' because

239 

240 

241def _jsonrpc_body(response: _JsonRpcResponse) -> dict[str, object]: 

242 """The decoded JSON-RPC body of ``response``.""" 

243 return response.json() 

244 

245 

246async def _forward_jsonrpc( 

247 agent_url: str, 

248 body: dict[str, object], 

249 extra_headers: Mapping[str, str] | None = None, 

250) -> dict[str, object]: 

251 from litellm.llms.custom_httpx.http_handler import get_async_httpx_client 

252 from litellm.types.llms.custom_http import httpxSpecialProvider 

253 

254 headers: Final = {"Content-Type": "application/json", **(extra_headers or {})} 

255 handler: Final = get_async_httpx_client( 

256 llm_provider=httpxSpecialProvider.A2A, 

257 params={"timeout": 60.0}, 

258 ) 

259 resp: Final = await handler.post(agent_url, json=body, headers=headers) 

260 try: 

261 result: Final = _jsonrpc_body(resp) 

262 except Exception: 

263 resp.raise_for_status() 

264 raise 

265 if not resp.is_success and "error" not in result: 

266 resp.raise_for_status() 

267 return result 

268 

269 

270async def _a2a_sse_event_source( 

271 agent_url: str, 

272 body: Mapping[str, object], 

273 request_id: str | int | None = None, 

274 extra_headers: Mapping[str, str] | None = None, 

275 served_version: A2AVersion = "0.3", 

276) -> AsyncGenerator[Mapping[str, object], None]: 

277 """Stream an upstream A2A SSE response as parsed JSON-RPC event dicts. 

278 

279 Upstream HTTP/JSON-RPC errors are surfaced as a single JSON-RPC error event 

280 so the caller can relay them instead of breaking the stream. 

281 """ 

282 from litellm.llms.custom_httpx.http_handler import get_async_httpx_client 

283 from litellm.types.agents import _normalize_a2a_jsonrpc_response 

284 from litellm.types.llms.custom_http import httpxSpecialProvider 

285 

286 headers: Final = { 

287 "Content-Type": "application/json", 

288 "Accept": "text/event-stream", 

289 **(extra_headers or {}), 

290 } 

291 handler: Final = get_async_httpx_client( 

292 llm_provider=httpxSpecialProvider.A2A, 

293 params={"timeout": None}, 

294 ) 

295 async_client: Final = handler.client 

296 req: Final = async_client.build_request("POST", agent_url, json=body, headers=headers) 

297 resp: Final = await async_client.send(req, stream=True) 

298 try: 

299 if not resp.is_success: 

300 error_body: Final = await resp.aread() 

301 error_event: Mapping[str, object] | None = None 

302 try: 

303 parsed: Final = json.loads(error_body) 

304 if isinstance(parsed, dict) and "error" in parsed: 

305 error_event = _normalize_a2a_jsonrpc_response(parsed, request_id=request_id) 

306 except Exception: 

307 error_event = None 

308 yield error_event or { 

309 "jsonrpc": "2.0", 

310 "id": request_id, 

311 "error": {"code": -32603, "message": resp.reason_phrase}, 

312 } 

313 return 

314 async for line in resp.aiter_lines(): 

315 stripped = line.strip() 

316 if not stripped.startswith("data:"): 

317 continue 

318 payload = stripped[len("data:") :].strip() 

319 if not payload: 

320 continue 

321 try: 

322 event = json.loads(payload) 

323 except Exception: 

324 continue 

325 if isinstance(event, dict): 

326 event = normalize_stream_event(event, served_version, request_id=request_id) 

327 yield event 

328 finally: 

329 await resp.aclose() 

330 

331 

332def _sse_streaming_response(generator: AsyncGenerator[str, None]) -> StreamingResponse: 

333 # The upstream agent is only contacted once this generator is first pulled, so 

334 # a slow first event leaves the response body idle for its whole 

335 # time-to-first-token and an intermediary with an idle read timeout drops a 

336 # healthy connection. Off until an operator sets an interval, and the 

337 # buffering hint only goes out when there are keepalives to protect. 

338 keepalive_interval: Final = coerce_keepalive_interval(litellm.sse_keepalive_ping_interval_seconds) 

339 if keepalive_interval is None: 

340 return StreamingResponse(generator, media_type="text/event-stream") 

341 return StreamingResponse( 

342 wrap_sse_stream_with_keepalive_pings(generator, keepalive_interval, ping_chunk=SSE_COMMENT_PING), 

343 media_type="text/event-stream", 

344 headers=_SSE_KEEPALIVE_HEADERS, 

345 ) 

346 

347 

348async def _forward_jsonrpc_sse( 

349 agent_url: str, 

350 body: Mapping[str, object], 

351 request_id: str | int | None = None, 

352 extra_headers: Mapping[str, str] | None = None, 

353 proxy_logging_obj: ProxyLogging | None = None, 

354 user_api_key_dict: UserAPIKeyAuth | None = None, 

355 request_data: dict[str, object] | None = None, 

356 served_version: A2AVersion = "0.3", 

357) -> StreamingResponse: 

358 event_source: Final = _a2a_sse_event_source( 

359 agent_url, 

360 body, 

361 request_id=request_id, 

362 extra_headers=extra_headers, 

363 served_version=served_version, 

364 ) 

365 

366 def _serialize_chunk(chunk: object) -> str: 

367 return f"data: {json.dumps(chunk)}\n\n" 

368 

369 def _serialize_error(proxy_exc: object) -> str: 

370 return ( 

371 "data: " 

372 + json.dumps( 

373 { 

374 "jsonrpc": "2.0", 

375 "id": request_id, 

376 "error": { 

377 "code": -32603, 

378 "message": getattr(proxy_exc, "message", str(proxy_exc)), 

379 }, 

380 } 

381 ) 

382 + "\n\n" 

383 ) 

384 

385 if proxy_logging_obj is not None and user_api_key_dict is not None and request_data is not None: 

386 # Route streamed events through the shared streaming generator so the 

387 # post-call streaming hook (and therefore agent guardrails) inspects 

388 # tasks/resubscribe output the same way message/stream does. 

389 from litellm.proxy.common_request_processing import ( 

390 ProxyBaseLLMRequestProcessing, 

391 ) 

392 

393 generator: AsyncGenerator[str, None] = ProxyBaseLLMRequestProcessing.async_streaming_data_generator( 

394 response=event_source, 

395 user_api_key_dict=user_api_key_dict, 

396 request_data=request_data, 

397 proxy_logging_obj=proxy_logging_obj, 

398 serialize_chunk=_serialize_chunk, 

399 serialize_error=_serialize_error, 

400 ) 

401 else: 

402 

403 async def _passthrough() -> AsyncGenerator[str, None]: 

404 async for chunk in event_source: 

405 yield _serialize_chunk(chunk) 

406 

407 generator = _passthrough() 

408 

409 return _sse_streaming_response(generator) 

410 

411 

412async def _handle_stream_message( 

413 api_base: str | None, 

414 request_id: str | int, 

415 params: dict[str, object], 

416 litellm_params: dict[str, object] | None = None, 

417 agent_id: str | None = None, 

418 metadata: dict[str, object] | None = None, 

419 proxy_server_request: dict[str, object] | None = None, 

420 *, 

421 agent_extra_headers: dict[str, str] | None = None, 

422 user_api_key_dict: UserAPIKeyAuth | None = None, 

423 request_data: dict[str, object] | None = None, 

424 proxy_logging_obj: ProxyLogging | None = None, 

425 served_version: A2AVersion = "0.3", 

426) -> StreamingResponse: 

427 """Handle message/stream method via SDK functions. 

428 

429 The A2A JSON-RPC binding streams responses as SSE (text/event-stream) with 

430 each JSON-RPC object framed as ``data: <json>\n\n``, matching the official 

431 a2a-sdk client which rejects any other Content-Type. When user_api_key_dict, 

432 request_data, and proxy_logging_obj are provided, events are routed through 

433 common_request_processing.async_streaming_data_generator so proxy hooks and 

434 cost injection apply. 

435 """ 

436 from litellm.a2a_protocol import asend_message_streaming 

437 from litellm.a2a_protocol.main import A2A_SDK_AVAILABLE 

438 

439 if not A2A_SDK_AVAILABLE: 

440 

441 async def _error_stream(): 

442 yield _sse_event( 

443 { 

444 "jsonrpc": "2.0", 

445 "id": request_id, 

446 "error": { 

447 "code": -32603, 

448 "message": "Server error: 'a2a' package not installed", 

449 }, 

450 } 

451 ) 

452 

453 return StreamingResponse(_error_stream(), media_type="text/event-stream") 

454 

455 from a2a.compat.v0_3.types import SendStreamingMessageRequest 

456 

457 use_proxy_hooks = user_api_key_dict is not None and request_data is not None and proxy_logging_obj is not None 

458 

459 try: 

460 message_send_params: Final = _build_message_send_params(params) 

461 except (ValidationError, ValueError) as e: 

462 invalid_params_message: Final = f"Invalid params: {e}" 

463 

464 async def _invalid_params_stream(): 

465 yield _sse_event( 

466 { 

467 "jsonrpc": "2.0", 

468 "id": request_id, 

469 "error": {"code": -32602, "message": invalid_params_message}, 

470 } 

471 ) 

472 

473 return StreamingResponse(_invalid_params_stream(), media_type="text/event-stream") 

474 

475 def _sse_chunk(chunk: object) -> str: 

476 obj = _to_jsonrpc_object(chunk) 

477 if isinstance(obj, dict): 

478 obj = normalize_stream_event(obj, served_version, request_id=request_id) 

479 return _sse_event(obj) 

480 

481 async def stream_response(): 

482 try: 

483 a2a_request: Final = SendStreamingMessageRequest( 

484 id=request_id, 

485 params=message_send_params, 

486 ) 

487 a2a_stream: Final = asend_message_streaming( 

488 request=a2a_request, 

489 api_base=api_base, 

490 litellm_params=litellm_params, 

491 agent_id=agent_id, 

492 metadata=metadata, 

493 proxy_server_request=proxy_server_request, 

494 agent_extra_headers=agent_extra_headers, 

495 ) 

496 

497 if ( 

498 use_proxy_hooks 

499 and user_api_key_dict is not None 

500 and request_data is not None 

501 and proxy_logging_obj is not None 

502 ): 

503 from litellm.proxy.common_request_processing import ( 

504 ProxyBaseLLMRequestProcessing, 

505 ) 

506 

507 def _sse_error(proxy_exc: object) -> str: 

508 return _sse_event( 

509 { 

510 "jsonrpc": "2.0", 

511 "id": request_id, 

512 "error": { 

513 "code": -32603, 

514 "message": getattr( 

515 proxy_exc, 

516 "message", 

517 f"Streaming error: {proxy_exc}", 

518 ), 

519 }, 

520 } 

521 ) 

522 

523 async for line in ProxyBaseLLMRequestProcessing.async_streaming_data_generator( 

524 response=a2a_stream, 

525 user_api_key_dict=user_api_key_dict, 

526 request_data=request_data, 

527 proxy_logging_obj=proxy_logging_obj, 

528 serialize_chunk=_sse_chunk, 

529 serialize_error=_sse_error, 

530 ): 

531 yield line 

532 else: 

533 async for chunk in a2a_stream: 

534 yield _sse_chunk(chunk) 

535 except Exception as e: 

536 verbose_proxy_logger.exception("Error streaming A2A response: %s", e) 

537 if ( 

538 use_proxy_hooks 

539 and proxy_logging_obj is not None 

540 and user_api_key_dict is not None 

541 and request_data is not None 

542 ): 

543 transformed_exception: Final = await proxy_logging_obj.post_call_failure_hook( 

544 user_api_key_dict=user_api_key_dict, 

545 original_exception=e, 

546 request_data=request_data, 

547 ) 

548 if transformed_exception is not None: 

549 e = transformed_exception 

550 if isinstance(e, HTTPException): 

551 raise 

552 yield _sse_event( 

553 { 

554 "jsonrpc": "2.0", 

555 "id": request_id, 

556 "error": { 

557 "code": -32603, 

558 "message": f"Streaming error: {e}", 

559 }, 

560 } 

561 ) 

562 

563 return _sse_streaming_response(stream_response()) 

564 

565 

566@router.get( 

567 "/a2a/{agent_id}/.well-known/agent-card.json", 

568 tags=["[beta] A2A Agents"], 

569 dependencies=[Depends(user_api_key_auth)], 

570) 

571@router.get( 

572 "/a2a/{agent_id}/.well-known/agent.json", 

573 tags=["[beta] A2A Agents"], 

574 dependencies=[Depends(user_api_key_auth)], 

575) 

576async def get_agent_card( 

577 agent_id: str, 

578 request: Request, 

579 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

580): 

581 """ 

582 Get the agent card for an agent (A2A discovery endpoint). 

583 

584 Supports both standard paths: 

585 - /.well-known/agent-card.json 

586 - /.well-known/agent.json 

587 

588 The URL in the agent card is rewritten to point to the LiteLLM proxy, 

589 so all subsequent A2A calls go through LiteLLM for logging and cost tracking. 

590 """ 

591 from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( 

592 AgentRequestHandler, 

593 ) 

594 

595 try: 

596 agent: Final = await _get_agent(agent_id) 

597 if agent is None: 597 ↛ 601line 597 didn't jump to line 601 because the condition on line 597 was always true

598 raise HTTPException(status_code=404, detail=f"Agent '{agent_id}' not found") 

599 

600 # Check agent permission (skip for admin users) 

601 is_allowed: Final = await AgentRequestHandler.is_agent_allowed( 

602 agent_id=agent.agent_id, 

603 user_api_key_auth=user_api_key_dict, 

604 ) 

605 if not is_allowed: 

606 raise HTTPException( 

607 status_code=403, 

608 detail=f"Agent '{agent_id}' is not allowed for your key/team. Contact proxy admin for access.", 

609 ) 

610 

611 if not agent.agent_card_params: 

612 raise HTTPException( 

613 status_code=404, 

614 detail=f"Agent '{agent_id}' has no agent card configured", 

615 ) 

616 

617 proxy_url: Final = get_custom_url(str(request.base_url), route=f"a2a/{agent_id}") 

618 agent_card = deepcopy(agent.agent_card_params) 

619 agent_card["url"] = proxy_url 

620 interfaces: Final = agent_card.get("supportedInterfaces") 

621 if isinstance(interfaces, list) and interfaces: 

622 interfaces[0]["url"] = proxy_url 

623 served_version: Final = _served_version(agent, request) 

624 agent_card = normalize_agent_card(agent_card, served_version) 

625 

626 verbose_proxy_logger.debug("Returning agent card for '%s' with proxy URL: %s", agent_id, proxy_url) 

627 return JSONResponse(content=agent_card) 

628 

629 except HTTPException: 

630 raise 

631 except Exception as e: 

632 verbose_proxy_logger.exception("Error getting agent card: %s", e) 

633 raise HTTPException(status_code=500, detail=str(e)) 

634 

635 

636@router.post( 

637 "/a2a/{agent_id}", 

638 tags=["[beta] A2A Agents"], 

639 dependencies=[Depends(user_api_key_auth)], 

640) 

641@router.post( 

642 "/a2a/{agent_id}/message/send", 

643 tags=["[beta] A2A Agents"], 

644 dependencies=[Depends(user_api_key_auth)], 

645) 

646@router.post( 

647 "/v1/a2a/{agent_id}/message/send", 

648 tags=["[beta] A2A Agents"], 

649 dependencies=[Depends(user_api_key_auth)], 

650) 

651async def invoke_agent_a2a( 

652 agent_id: str, 

653 request: Request, 

654 fastapi_response: Response, 

655 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

656): 

657 """ 

658 Invoke an agent using the A2A protocol (JSON-RPC 2.0). 

659 

660 Supported methods: 

661 - message/send: Send a message and get a response 

662 - message/stream: Send a message and stream the response 

663 """ 

664 from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( 

665 AgentRequestHandler, 

666 ) 

667 from litellm.proxy.proxy_server import ( 

668 general_settings, 

669 proxy_config, 

670 proxy_logging_obj, 

671 version, 

672 ) 

673 

674 body: dict[str, Any] = {} 

675 request_data: dict[str, Any] = body 

676 try: 

677 body = await request.json() 

678 request_data = body 

679 

680 verbose_proxy_logger.debug("A2A request for agent '%s': %s", agent_id, body) 

681 

682 # Validate JSON-RPC format 

683 if body.get("jsonrpc") != "2.0": 

684 return _jsonrpc_error(body.get("id"), -32600, "Invalid Request: jsonrpc must be '2.0'") 

685 

686 request_id: Final[Any | None] = body.get("id") 

687 original_method: Final[str | None] = body.get("method") 

688 method: str | None = original_method 

689 params = body.get("params", {}) 

690 

691 if method: 

692 method = _PASCAL_TO_WIRE.get(method, method) 

693 

694 if isinstance(params, dict): 

695 # extract any litellm params from the params - eg. 'guardrails' 

696 # ``metadata`` is intentionally excluded: it's a first-class A2A 

697 # ``MessageSendParams`` field that the completion bridge forwards 

698 # downstream via ``get_forward_metadata``. Stripping it here would 

699 # collide with litellm's spend-tracking ``metadata`` kwarg and 

700 # silently drop the caller's A2A request-level metadata. 

701 params_to_remove: Final = [] 

702 for key, value in params.items(): 

703 if key in all_litellm_params and key not in {"id", "metadata"}: 

704 params_to_remove.append(key) 

705 body[key] = value 

706 for key in params_to_remove: 

707 params.pop(key) 

708 

709 # Find the agent 

710 agent: Final = await _get_agent(agent_id) 

711 if agent is None: 

712 return _jsonrpc_error(request_id, -32000, f"Agent '{agent_id}' not found", 404) 

713 

714 served_version: Final = _served_version(agent, request, original_method) 

715 

716 is_allowed: Final = await AgentRequestHandler.is_agent_allowed( 

717 agent_id=agent.agent_id, 

718 user_api_key_auth=user_api_key_dict, 

719 ) 

720 if not is_allowed: 

721 raise HTTPException( 

722 status_code=403, 

723 detail=f"Agent '{agent_id}' is not allowed for your key/team. Contact proxy admin for access.", 

724 ) 

725 

726 _enforce_inbound_trace_id(agent, request) 

727 

728 # Get backend URL and agent name 

729 agent_card_params: Final = agent.agent_card_params or {} 

730 agent_url: Final = agent_card_params.get("url") 

731 agent_name: Final = agent_card_params.get("name", agent_id) 

732 

733 # Get litellm_params (may include custom_llm_provider for completion bridge) 

734 litellm_params: dict[str, object] = agent.litellm_params or {} 

735 custom_llm_provider: Final = litellm_params.get("custom_llm_provider") 

736 

737 # Hand the authenticated key hash to the completion bridge so provider 

738 # configs can scope provider-side session state per key (e.g. LangFlow 

739 # session memory) instead of trusting the client-supplied A2A contextId. 

740 if custom_llm_provider and user_api_key_dict.api_key: 

741 from litellm.a2a_protocol.litellm_completion_bridge.handler import ( 

742 A2A_USER_API_KEY_HASH_PARAM, 

743 ) 

744 

745 litellm_params = { 

746 **litellm_params, 

747 A2A_USER_API_KEY_HASH_PARAM: user_api_key_dict.api_key, 

748 } 

749 

750 # URL is required unless using completion bridge with a provider that derives endpoint from model 

751 # (e.g., bedrock/agentcore derives endpoint from ARN in model string) 

752 if not agent_url and not custom_llm_provider: 

753 return _jsonrpc_error(request_id, -32000, f"Agent '{agent_id}' has no URL configured", 500) 

754 

755 verbose_proxy_logger.info( 

756 "Proxying A2A request to agent '%s' at %s", agent_id, agent_url or "completion-bridge" 

757 ) 

758 

759 # Set up data dict for litellm processing 

760 if "metadata" not in body: 

761 body["metadata"] = {} 

762 body["metadata"]["agent_id"] = agent.agent_id 

763 body["agent_id"] = agent.agent_id 

764 

765 body.update( 

766 { 

767 "model": f"a2a_agent/{agent_name}", 

768 "custom_llm_provider": "a2a_agent", 

769 } 

770 ) 

771 

772 # Add litellm data (user_api_key, user_id, team_id, etc.) 

773 from litellm.proxy.common_request_processing import ( 

774 ProxyBaseLLMRequestProcessing, 

775 ) 

776 

777 caller_identity: Final = _caller_identity_headers(user_api_key_dict) 

778 processor: Final = ProxyBaseLLMRequestProcessing(data=body) 

779 data, logging_obj = await processor.common_processing_pre_call_logic( 

780 request=request, 

781 general_settings=general_settings, 

782 user_api_key_dict=user_api_key_dict, 

783 proxy_logging_obj=proxy_logging_obj, 

784 proxy_config=proxy_config, 

785 route_type="asend_message", 

786 version=version, 

787 ) 

788 request_data = data 

789 

790 # Build merged headers for the backend agent 

791 static_headers: Final[Mapping[str, str]] = dict(agent.static_headers or {}) 

792 

793 raw_headers: Final = dict(request.headers) 

794 normalized: Final = {k.lower(): v for k, v in raw_headers.items()} 

795 

796 dynamic_headers: Final[dict[str, str]] = {} 

797 

798 # 1. Admin-configured extra_headers: forward named headers from client request 

799 if agent.extra_headers: 

800 for header_name in agent.extra_headers: 

801 header_name_str = str(header_name) 

802 val = normalized.get(header_name_str.lower()) 

803 if val is not None: 

804 dynamic_headers[header_name_str] = val 

805 

806 # 2. Convention-based forwarding: x-a2a-{agent_id_or_name}-{header_name} 

807 # Matches both agent_id (UUID) and agent_name (alias), case-insensitive. 

808 for alias in (agent.agent_id.lower(), agent.agent_name.lower()): 

809 prefix = f"x-a2a-{alias}-" 

810 for key, val in normalized.items(): 

811 if key.startswith(prefix): 

812 header_name = key[len(prefix) :] 

813 if header_name: 

814 dynamic_headers[header_name] = val 

815 

816 agent_extra_headers: Final = _forwarding_headers( 

817 caller_identity=caller_identity, 

818 request_data=data, 

819 agent_extra_headers=merge_agent_headers( 

820 dynamic_headers=dynamic_headers or None, 

821 static_headers=static_headers or None, 

822 ), 

823 backend_auth_header=await _resolve_backend_auth_header(litellm_params, custom_llm_provider), 

824 ) 

825 

826 # Merge agent-level guardrails into data so post_call_success_hook and 

827 # _handle_stream_message both pick them up. A2A agents use model 

828 # a2a_agent/*, which is not an llm_router deployment, so 

829 # _check_and_merge_model_level_guardrails() skips them. 

830 _agent_guardrails = litellm_params.get("guardrails") 

831 if _agent_guardrails: 

832 if not isinstance(_agent_guardrails, list): 

833 _agent_guardrails = [_agent_guardrails] 

834 _existing_guardrails: list = data.get("guardrails") or [] 

835 if not isinstance(_existing_guardrails, list): 

836 _existing_guardrails = [_existing_guardrails] 

837 data["guardrails"] = _existing_guardrails + [g for g in _agent_guardrails if g not in _existing_guardrails] 

838 

839 # Route through SDK functions 

840 if method == "message/send": 

841 from litellm.a2a_protocol import asend_message 

842 from litellm.a2a_protocol.main import A2A_SDK_AVAILABLE 

843 

844 if not A2A_SDK_AVAILABLE: 

845 return _jsonrpc_error( 

846 request_id, 

847 -32603, 

848 "Server error: 'a2a' package not installed. Please install 'a2a-sdk'.", 

849 500, 

850 ) 

851 from a2a.compat.v0_3.types import SendMessageRequest 

852 

853 try: 

854 message_send_params: Final = _build_message_send_params(params) 

855 except (ValidationError, ValueError) as e: 

856 return _jsonrpc_error(request_id, -32602, f"Invalid params: {e}") 

857 

858 a2a_request: Final = SendMessageRequest( 

859 id=request_id if request_id is not None else "", 

860 params=message_send_params, 

861 ) 

862 # Defer spend-log until after post_call_success_hook so guardrail 

863 # results written by the unified_guardrail hook are captured. 

864 logging_obj._defer_async_logging = True 

865 response = await asend_message( 

866 request=a2a_request, 

867 api_base=agent_url, 

868 litellm_params=litellm_params, 

869 agent_id=agent.agent_id, 

870 metadata=data.get("metadata", {}), 

871 proxy_server_request=data.get("proxy_server_request"), 

872 litellm_logging_obj=logging_obj, 

873 agent_extra_headers=agent_extra_headers, 

874 ) 

875 

876 try: 

877 response = await proxy_logging_obj.post_call_success_hook( 

878 user_api_key_dict=user_api_key_dict, 

879 data=data, 

880 response=response, 

881 ) 

882 finally: 

883 _enqueue_fn: Final = getattr(logging_obj, "_enqueue_deferred_logging", None) 

884 if _enqueue_fn is not None: 

885 logging_obj._enqueue_deferred_logging = None 

886 _enqueue_fn() 

887 

888 response_dict: Final[dict[str, Any]] = ( 

889 response.model_dump(mode="json", exclude_none=True) 

890 if hasattr(response, "model_dump") 

891 else response 

892 if isinstance(response, dict) 

893 else {} 

894 ) 

895 return JSONResponse( 

896 content=normalize_jsonrpc_response( 

897 response_dict, 

898 served_version, 

899 method="message/send", 

900 ) 

901 ) 

902 

903 elif method == "message/stream": 

904 return await _handle_stream_message( 

905 api_base=agent_url, 

906 request_id=request_id if request_id is not None else "", 

907 params=params, 

908 litellm_params=litellm_params, 

909 agent_id=agent.agent_id, 

910 metadata=data.get("metadata", {}), 

911 proxy_server_request=data.get("proxy_server_request"), 

912 agent_extra_headers=agent_extra_headers, 

913 user_api_key_dict=user_api_key_dict, 

914 request_data=data, 

915 proxy_logging_obj=proxy_logging_obj, 

916 served_version=served_version, 

917 ) 

918 elif method in { 

919 "tasks/get", 

920 "tasks/list", 

921 "tasks/cancel", 

922 "tasks/pushNotificationConfig/set", 

923 "tasks/pushNotificationConfig/get", 

924 "tasks/pushNotificationConfig/list", 

925 "tasks/pushNotificationConfig/delete", 

926 "agent/getAuthenticatedExtendedCard", 

927 }: 

928 if not agent_url: 

929 return _jsonrpc_error(request_id, -32000, f"Agent '{agent_id}' has no URL configured", 500) 

930 if isinstance(params, dict): 

931 params = normalize_request_params(params, served_version, method=method) 

932 if method == "tasks/pushNotificationConfig/set": 

933 if not isinstance(params, dict): 

934 raise HTTPException( 

935 status_code=400, 

936 detail="params must be an object", 

937 ) 

938 push_config: Final = params.get("pushNotificationConfig", {}) 

939 if "pushNotificationConfig" in params and not isinstance(push_config, dict): 

940 raise HTTPException( 

941 status_code=400, 

942 detail="pushNotificationConfig must be an object", 

943 ) 

944 for callback_url in (params.get("url"), push_config.get("url")): 

945 if not callback_url: 

946 continue 

947 if not isinstance(callback_url, str): 

948 raise HTTPException( 

949 status_code=400, 

950 detail="Push notification URL must be a string", 

951 ) 

952 _validate_push_notification_url(callback_url) 

953 forward_body: dict[str, object] = { 

954 "jsonrpc": "2.0", 

955 "id": request_id, 

956 "method": method, 

957 "params": params, 

958 } 

959 result = await _forward_jsonrpc(agent_url, forward_body, extra_headers=agent_extra_headers) 

960 if method == "agent/getAuthenticatedExtendedCard": 

961 card: Final = result.get("result") 

962 if isinstance(card, dict): 

963 proxy_url: Final = get_custom_url(str(request.base_url), route=f"a2a/{agent_id}") 

964 # Rewrite the upstream agent URL in both 0.3 (top-level `url`) 

965 # and 1.0 (`supportedInterfaces[0].url`) wire formats so that 

966 # downstream clients never see the upstream internal address. 

967 if "url" in card: 

968 card["url"] = proxy_url 

969 interfaces: Final = card.get("supportedInterfaces") 

970 if isinstance(interfaces, list) and interfaces: 

971 interfaces[0]["url"] = proxy_url 

972 result["result"] = normalize_agent_card(card, served_version) 

973 else: 

974 result = normalize_jsonrpc_response(result, served_version, method=method) 

975 from litellm.types.agents import LiteLLMSendMessageResponse 

976 

977 response = LiteLLMSendMessageResponse.from_dict(result, request_id=request_id) 

978 response = await proxy_logging_obj.post_call_success_hook( 

979 user_api_key_dict=user_api_key_dict, 

980 data=data, 

981 response=response, 

982 ) 

983 return JSONResponse( 

984 content=( 

985 response.model_dump(mode="json", exclude_none=True) if hasattr(response, "model_dump") else response 

986 ) 

987 ) 

988 

989 elif method == "tasks/resubscribe": 

990 if not agent_url: 

991 return _jsonrpc_error(request_id, -32000, f"Agent '{agent_id}' has no URL configured", 500) 

992 if isinstance(params, dict): 

993 params = normalize_request_params(params, served_version, method=method) 

994 forward_body = { 

995 "jsonrpc": "2.0", 

996 "id": request_id, 

997 "method": method, 

998 "params": params, 

999 } 

1000 return await _forward_jsonrpc_sse( 

1001 agent_url, 

1002 forward_body, 

1003 request_id=request_id, 

1004 extra_headers=agent_extra_headers, 

1005 proxy_logging_obj=proxy_logging_obj, 

1006 user_api_key_dict=user_api_key_dict, 

1007 request_data=data, 

1008 served_version=served_version, 

1009 ) 

1010 

1011 else: 

1012 return _jsonrpc_error(request_id, -32601, f"Method '{method}' not found") 

1013 

1014 except HTTPException: 

1015 raise 

1016 except Exception as e: 

1017 verbose_proxy_logger.exception("Error invoking agent: %s", e) 

1018 try: 

1019 await proxy_logging_obj.post_call_failure_hook( 

1020 user_api_key_dict=user_api_key_dict, 

1021 original_exception=e, 

1022 request_data=request_data, 

1023 ) 

1024 except Exception: 

1025 pass 

1026 if isinstance(e, litellm.BadRequestError): 1026 ↛ 1027line 1026 didn't jump to line 1027 because the condition on line 1026 was never true

1027 return _jsonrpc_error(body.get("id"), -32602, e.message, 400) 

1028 return _jsonrpc_error(body.get("id"), -32603, f"Internal error: {e}", 500)