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
« 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.
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"""
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
20from fastapi import APIRouter, Depends, HTTPException, Request, Response
21from fastapi.responses import JSONResponse, StreamingResponse
22from pydantic import ValidationError
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
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
53 from litellm.types.agents import AgentResponse
55router: Final = APIRouter()
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)
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}
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"
86def _to_jsonrpc_object(chunk: object) -> object:
87 """Coerce a streamed chunk to the JSON-RPC object it carries.
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
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
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
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
120def _served_version(agent: "AgentResponse", request: Request, original_method: str | None = None) -> A2AVersion:
121 """Protocol version LiteLLM serves for this agent.
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"
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
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 )
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)
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
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 )
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 )
217 return await get_agent_with_read_through(agent_id)
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
226 from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers
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 )
237class _JsonRpcResponse(Protocol):
238 def json(self) -> dict[str, object]: ... 238 ↛ exitline 238 didn't return from function 'json' because
241def _jsonrpc_body(response: _JsonRpcResponse) -> dict[str, object]:
242 """The decoded JSON-RPC body of ``response``."""
243 return response.json()
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
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
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.
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
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()
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 )
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 )
366 def _serialize_chunk(chunk: object) -> str:
367 return f"data: {json.dumps(chunk)}\n\n"
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 )
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 )
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:
403 async def _passthrough() -> AsyncGenerator[str, None]:
404 async for chunk in event_source:
405 yield _serialize_chunk(chunk)
407 generator = _passthrough()
409 return _sse_streaming_response(generator)
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.
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
439 if not A2A_SDK_AVAILABLE:
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 )
453 return StreamingResponse(_error_stream(), media_type="text/event-stream")
455 from a2a.compat.v0_3.types import SendStreamingMessageRequest
457 use_proxy_hooks = user_api_key_dict is not None and request_data is not None and proxy_logging_obj is not None
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}"
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 )
473 return StreamingResponse(_invalid_params_stream(), media_type="text/event-stream")
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)
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 )
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 )
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 )
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 )
563 return _sse_streaming_response(stream_response())
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).
584 Supports both standard paths:
585 - /.well-known/agent-card.json
586 - /.well-known/agent.json
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 )
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")
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 )
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 )
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)
626 verbose_proxy_logger.debug("Returning agent card for '%s' with proxy URL: %s", agent_id, proxy_url)
627 return JSONResponse(content=agent_card)
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))
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).
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 )
674 body: dict[str, Any] = {}
675 request_data: dict[str, Any] = body
676 try:
677 body = await request.json()
678 request_data = body
680 verbose_proxy_logger.debug("A2A request for agent '%s': %s", agent_id, body)
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'")
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", {})
691 if method:
692 method = _PASCAL_TO_WIRE.get(method, method)
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)
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)
714 served_version: Final = _served_version(agent, request, original_method)
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 )
726 _enforce_inbound_trace_id(agent, request)
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)
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")
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 )
745 litellm_params = {
746 **litellm_params,
747 A2A_USER_API_KEY_HASH_PARAM: user_api_key_dict.api_key,
748 }
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)
755 verbose_proxy_logger.info(
756 "Proxying A2A request to agent '%s' at %s", agent_id, agent_url or "completion-bridge"
757 )
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
765 body.update(
766 {
767 "model": f"a2a_agent/{agent_name}",
768 "custom_llm_provider": "a2a_agent",
769 }
770 )
772 # Add litellm data (user_api_key, user_id, team_id, etc.)
773 from litellm.proxy.common_request_processing import (
774 ProxyBaseLLMRequestProcessing,
775 )
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
790 # Build merged headers for the backend agent
791 static_headers: Final[Mapping[str, str]] = dict(agent.static_headers or {})
793 raw_headers: Final = dict(request.headers)
794 normalized: Final = {k.lower(): v for k, v in raw_headers.items()}
796 dynamic_headers: Final[dict[str, str]] = {}
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
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
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 )
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]
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
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
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}")
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 )
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()
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 )
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
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 )
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 )
1011 else:
1012 return _jsonrpc_error(request_id, -32601, f"Method '{method}' not found")
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)