Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/a2a/version_convert.py: 13%
204 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# a2a-sdk's compat conversions (pb2_v10, ParseDict, MessageToDict, to_compat_*)
3# are protobuf-generated/untyped, so every conversion call here takes Unknown-typed
4# arguments. This module is the A2A 0.3<->1.0 boundary; the rule is off file-wide
5# rather than scattering per-line ignores across every SDK call.
6"""
7Normalize A2A JSON-RPC payloads to the protocol version LiteLLM serves for an agent.
9LiteLLM fronts upstream agents and lets an admin pin the protocol version it speaks
10to clients (``0.3`` or ``1.0``) per agent. Upstream responses may arrive in either
11wire shape, so every response, stream event, forwarded request and extended card is
12converted to the served version here. Conversion is shape-detecting (we infer the
13payload's current version rather than trusting a stored one) and best-effort: any
14failure falls back to returning the input unchanged so a conversion bug can never
15break an otherwise-valid response.
17The two wire shapes:
19- ``0.3``: JSON dump of the compat pydantic types, discriminated by a ``kind`` field
20 (``message`` / ``task`` / ``status-update`` / ``artifact-update``). A send result is
21 the bare object.
22- ``1.0``: protobuf JSON (``MessageToDict``), a oneof envelope keyed by
23 ``message`` / ``task`` / ``statusUpdate`` / ``artifactUpdate`` with no ``kind``. A
24 ``Task`` result is a bare object without ``kind``.
25"""
27from collections.abc import Callable
28from typing import Final, Literal
30from pydantic import BaseModel
32from litellm._logging import verbose_proxy_logger
33from litellm.proxy.a2a.agent_card import normalize_protocol_version
35A2AVersion = Literal["0.3", "1.0"]
36RequestId = str | int | None
37JsonDict = dict[str, object]
39_V1_SEND_ENVELOPE_KEYS: Final = frozenset({"message", "task"})
40_V1_STREAM_ENVELOPE_KEYS: Final = frozenset({"message", "task", "statusUpdate", "artifactUpdate"})
43def _dump_03(model: BaseModel) -> JsonDict:
44 """Dump a compat (0.3) pydantic model to its camelCase wire dict."""
45 return model.model_dump(by_alias=True, exclude_none=True, mode="json")
48def _best_effort(convert: Callable[[], JsonDict], fallback: JsonDict, *, label: str) -> JsonDict:
49 """Run a conversion, returning ``fallback`` unchanged if it raises."""
50 try:
51 return convert()
52 except Exception as e: # noqa: BLE001 - best-effort passthrough
53 verbose_proxy_logger.debug("A2A %s conversion failed: %s", label, e)
54 return fallback
57def normalize_jsonrpc_response(content: JsonDict, target: A2AVersion, *, method: str) -> JsonDict:
58 """Convert a JSON-RPC response's ``result`` to ``target``.
60 Errors and non-dict results pass through untouched.
61 """
62 if content.get("error") is not None:
63 return content
64 result: Final = content.get("result")
65 if not isinstance(result, dict):
66 return content
68 converted: Final = _convert_result(result, target, method=method, request_id=_as_request_id(content.get("id")))
69 if converted is result:
70 return content
71 return {**content, "result": converted}
74def normalize_stream_event(event: JsonDict, target: A2AVersion, *, request_id: RequestId) -> JsonDict:
75 """Convert a single streamed JSON-RPC event's ``result`` to ``target``."""
76 if event.get("error") is not None:
77 return event
78 result: Final = event.get("result")
79 if not isinstance(result, dict):
80 return event
82 converted: Final = _convert_stream_result(result, target, request_id=request_id)
83 if converted is result:
84 return event
85 return {**event, "result": converted}
88def normalize_request_params(params: JsonDict, served: A2AVersion, *, method: str) -> JsonDict:
89 """Down-convert forwarded request ``params`` from the served version to 0.3.
91 Upstream agents in this proxy pivot on 0.3 wire format, so when LiteLLM serves
92 1.0 the inbound params must be lowered before forwarding. A no-op when the served
93 version is already 0.3.
94 """
95 if served == "0.3":
96 return params
97 return _best_effort(
98 lambda: _lower_request_params(params, method=method),
99 params,
100 label=f"request params ({method})",
101 )
104def _detect_card_version(card: JsonDict) -> A2AVersion:
105 """Infer the wire version of an agent card dict.
107 ``protocolVersion`` is the authoritative indicator; semver values normalize to
108 their major.minor (``"0.3.0"`` -> ``"0.3"``). Fall back to presence of
109 ``supportedInterfaces`` (a 1.0-only field) only when the explicit field is
110 absent or unrecognized; cards carrying neither signal are treated as 0.3.
111 """
112 normalized: Final = normalize_protocol_version(card.get("protocolVersion"))
113 if normalized is not None:
114 return normalized
115 return "1.0" if "supportedInterfaces" in card else "0.3"
118def normalize_agent_card(card: JsonDict, target: A2AVersion) -> JsonDict:
119 """Convert an extended agent card to ``target``.
121 When lowering to 0.3, ``additionalInterfaces`` is stripped so the conversion never
122 re-exposes upstream backend URLs that the LiteLLM-fronting merge deliberately drops.
123 """
124 if not isinstance(card, dict):
125 return card
127 current: Final = _detect_card_version(card)
128 if current == target and not (target == "0.3" and "supportedInterfaces" in card):
129 return card
130 return _best_effort(lambda: _convert_agent_card(card, target), card, label="agent card")
133def _as_request_id(value: object) -> RequestId:
134 return value if isinstance(value, (str, int)) else None
137def _convert_result(
138 result: JsonDict,
139 target: A2AVersion,
140 *,
141 method: str,
142 request_id: RequestId,
143) -> JsonDict:
144 if method == "message/send":
145 return _convert_send_result(result, target, request_id=request_id)
146 if method in ("tasks/get", "tasks/cancel"):
147 return _convert_task(result, target)
148 if method == "tasks/list":
149 return _convert_list_tasks_result(result, target)
150 return result
153def _detect_send_version(result: JsonDict) -> A2AVersion | None:
154 if "kind" in result:
155 return "0.3"
156 if result.keys() & _V1_SEND_ENVELOPE_KEYS:
157 return "1.0"
158 return None
161def _convert_send_result(result: JsonDict, target: A2AVersion, *, request_id: RequestId) -> JsonDict:
162 current: Final = _detect_send_version(result)
163 if current is None or current == target:
164 return result
165 return _best_effort(
166 lambda: _send_result_to(result, target, request_id),
167 result,
168 label="send result",
169 )
172def _send_result_to(result: JsonDict, target: A2AVersion, request_id: RequestId) -> JsonDict:
173 from a2a.compat.v0_3.conversions import (
174 MessageToDict,
175 ParseDict,
176 pb2_v10,
177 to_compat_send_message_response,
178 to_core_send_message_response,
179 types_v03,
180 )
182 if target == "1.0":
183 compat_result: Final = _validate_message_or_task(result)
184 response: Final = types_v03.SendMessageResponse(
185 root=types_v03.SendMessageSuccessResponse(
186 id=str(request_id) if request_id is not None else "",
187 result=compat_result, # pyright: ignore[reportArgumentType]
188 )
189 )
190 return MessageToDict(
191 to_core_send_message_response(response),
192 preserving_proto_field_name=False,
193 )
195 pb: Final = pb2_v10.SendMessageResponse()
196 ParseDict(result, pb, ignore_unknown_fields=True)
197 return _dump_03(to_compat_send_message_response(pb, request_id).root.result)
200def _convert_task(result: JsonDict, target: A2AVersion) -> JsonDict:
201 current: Final[A2AVersion] = "0.3" if "kind" in result else "1.0"
202 if current == target:
203 return result
204 return _best_effort(lambda: _task_to(result, target), result, label="task")
207def _detect_list_tasks_version(result: JsonDict) -> A2AVersion | None:
208 tasks: Final = result.get("tasks")
209 if not isinstance(tasks, list) or not tasks:
210 return None
211 first: Final = tasks[0]
212 if not isinstance(first, dict):
213 return None
214 return "0.3" if "kind" in first else "1.0"
217def _convert_list_tasks_result(result: JsonDict, target: A2AVersion) -> JsonDict:
218 current: Final = _detect_list_tasks_version(result)
219 if current is None or current == target:
220 return result
221 return _best_effort(
222 lambda: _list_tasks_result_to(result, target),
223 result,
224 label="list tasks result",
225 )
228def _list_tasks_result_to(result: JsonDict, target: A2AVersion) -> JsonDict:
229 tasks: Final = result.get("tasks")
230 if not isinstance(tasks, list):
231 return result
232 return {
233 **result,
234 "tasks": [_task_to(item, target) if isinstance(item, dict) else item for item in tasks],
235 }
238def _task_to(result: JsonDict, target: A2AVersion) -> JsonDict:
239 from a2a.compat.v0_3.conversions import (
240 MessageToDict,
241 ParseDict,
242 pb2_v10,
243 to_compat_task,
244 to_core_task,
245 types_v03,
246 )
248 if target == "1.0":
249 core: Final = to_core_task(types_v03.Task.model_validate(result))
250 return MessageToDict(core, preserving_proto_field_name=False)
252 pb: Final = pb2_v10.Task()
253 ParseDict(result, pb, ignore_unknown_fields=True)
254 return _dump_03(to_compat_task(pb))
257def _detect_stream_version(result: JsonDict) -> A2AVersion | None:
258 if "kind" in result:
259 return "0.3"
260 if result.keys() & _V1_STREAM_ENVELOPE_KEYS:
261 return "1.0"
262 return None
265def _convert_stream_result(result: JsonDict, target: A2AVersion, *, request_id: RequestId) -> JsonDict:
266 current: Final = _detect_stream_version(result)
267 if current is None or current == target:
268 return result
269 return _best_effort(
270 lambda: _stream_result_to(result, target, request_id),
271 result,
272 label="stream event",
273 )
276def _stream_result_to(result: JsonDict, target: A2AVersion, request_id: RequestId) -> JsonDict:
277 from a2a.compat.v0_3.conversions import (
278 MessageToDict,
279 ParseDict,
280 pb2_v10,
281 to_compat_stream_response,
282 to_core_stream_response,
283 types_v03,
284 )
286 if target == "1.0":
287 event: Final = _validate_stream_event(result)
288 wrapper: Final = types_v03.SendStreamingMessageSuccessResponse(
289 id=str(request_id) if request_id is not None else "",
290 result=event, # pyright: ignore[reportArgumentType]
291 )
292 return MessageToDict(to_core_stream_response(wrapper), preserving_proto_field_name=False)
294 pb: Final = pb2_v10.StreamResponse()
295 ParseDict(result, pb, ignore_unknown_fields=True)
296 return _dump_03(to_compat_stream_response(pb, request_id).result)
299def _convert_agent_card(card: JsonDict, target: A2AVersion) -> JsonDict:
300 from a2a.compat.v0_3.conversions import (
301 MessageToDict,
302 ParseDict,
303 pb2_v10,
304 to_compat_agent_card,
305 to_core_agent_card,
306 types_v03,
307 )
309 if target == "0.3":
310 pb: Final = pb2_v10.AgentCard()
311 ParseDict(card, pb, ignore_unknown_fields=True)
312 lowered: Final = _dump_03(to_compat_agent_card(pb))
313 lowered.pop("additionalInterfaces", None)
314 return lowered
316 core: Final = to_core_agent_card(types_v03.AgentCard.model_validate(card))
317 return MessageToDict(core, preserving_proto_field_name=False)
320def _validate_message_or_task(result: JsonDict) -> BaseModel:
321 from a2a.compat.v0_3.conversions import types_v03
323 if result.get("kind") == "task":
324 return types_v03.Task.model_validate(result)
325 return types_v03.Message.model_validate(result)
328def _validate_stream_event(result: JsonDict) -> BaseModel:
329 from a2a.compat.v0_3.conversions import types_v03
331 kind: Final = result.get("kind")
332 if kind == "task":
333 return types_v03.Task.model_validate(result)
334 if kind == "status-update":
335 return types_v03.TaskStatusUpdateEvent.model_validate(result)
336 if kind == "artifact-update":
337 return types_v03.TaskArtifactUpdateEvent.model_validate(result)
338 return types_v03.Message.model_validate(result)
341def _lower_request_params(params: JsonDict, *, method: str) -> JsonDict:
342 if method == "tasks/list":
343 return _lower_list_tasks_params(params)
345 from a2a.compat.v0_3.conversions import (
346 ParseDict,
347 pb2_v10,
348 to_compat_cancel_task_request,
349 to_compat_create_task_push_notification_config_request,
350 to_compat_delete_task_push_notification_config_request,
351 to_compat_get_task_push_notification_config_request,
352 to_compat_get_task_request,
353 to_compat_list_task_push_notification_config_request,
354 to_compat_subscribe_to_task_request,
355 )
357 lowerings: Final[dict[str, Callable[[JsonDict], BaseModel]]] = {
358 "tasks/get": lambda p: to_compat_get_task_request(_parse(ParseDict, p, pb2_v10.GetTaskRequest()), "").params,
359 "tasks/cancel": lambda p: (
360 to_compat_cancel_task_request(_parse(ParseDict, p, pb2_v10.CancelTaskRequest()), "").params
361 ),
362 "tasks/resubscribe": lambda p: (
363 to_compat_subscribe_to_task_request(_parse(ParseDict, p, pb2_v10.SubscribeToTaskRequest()), "").params
364 ),
365 "tasks/pushNotificationConfig/set": lambda p: (
366 to_compat_create_task_push_notification_config_request(
367 _parse(
368 ParseDict,
369 _flatten_create_push_notification_params(p),
370 pb2_v10.TaskPushNotificationConfig(),
371 ),
372 "",
373 ).params
374 ),
375 "tasks/pushNotificationConfig/get": lambda p: (
376 to_compat_get_task_push_notification_config_request(
377 _parse(ParseDict, p, pb2_v10.GetTaskPushNotificationConfigRequest()), ""
378 ).params
379 ),
380 "tasks/pushNotificationConfig/list": lambda p: (
381 to_compat_list_task_push_notification_config_request(
382 _parse(ParseDict, p, pb2_v10.ListTaskPushNotificationConfigsRequest()), ""
383 ).params
384 ),
385 "tasks/pushNotificationConfig/delete": lambda p: (
386 to_compat_delete_task_push_notification_config_request(
387 _parse(ParseDict, p, pb2_v10.DeleteTaskPushNotificationConfigRequest()), ""
388 ).params
389 ),
390 }
391 lower: Final = lowerings.get(method)
392 if lower is None:
393 return params
394 return _dump_03(lower(params))
397def _lower_list_tasks_params(params: JsonDict) -> JsonDict:
398 from a2a.compat.v0_3.conversions import (
399 MessageToDict,
400 ParseDict,
401 pb2_v10,
402 types_v03,
403 )
405 proto: Final = pb2_v10.ListTasksRequest()
406 _parse(ParseDict, params, proto)
407 lowered: Final = MessageToDict(proto, preserving_proto_field_name=False)
408 status_name: Final = str(pb2_v10.TaskState.Name(proto.status))
409 valid_0_3_values: Final = frozenset(str(member.value) for member in types_v03.TaskState)
410 compat_status: Final = _proto_task_state_name_to_0_3(status_name, valid_0_3_values)
411 if compat_status is None:
412 lowered.pop("status", None)
413 else:
414 lowered["status"] = compat_status
415 return lowered
418def _proto_task_state_name_to_0_3(name: str, valid_0_3_values: frozenset[str]) -> str | None:
419 """Map a 1.0 protobuf ``TaskState`` enum name to its 0.3 wire string.
421 The ``TASK_STATE_<NAME>`` enum names line up with the 0.3 wire values once the
422 prefix is dropped and underscores become dashes, so no private SDK mapping is
423 needed. The result is validated against the 0.3 enum's own values; an unspecified
424 or unrecognized state yields ``None`` so the status filter is dropped.
425 """
426 base: Final = name.removeprefix("TASK_STATE_")
427 if base == "UNSPECIFIED":
428 return None
429 candidate: Final = base.lower().replace("_", "-")
430 return candidate if candidate in valid_0_3_values else None
433def _flatten_create_push_notification_params(params: JsonDict) -> JsonDict:
434 """Merge 1.x create envelope fields (parent/configId/config) into flat pb fields."""
435 flat: Final = dict(params)
436 config: Final = flat.pop("config", None)
437 push_config: Final = flat.pop("pushNotificationConfig", None)
438 nested: Final = config if config is not None else push_config
439 if not isinstance(nested, dict):
440 return params
441 parent: Final = flat.pop("parent", None)
442 if isinstance(parent, str) and parent.startswith("tasks/") and "taskId" not in flat:
443 flat["taskId"] = parent.removeprefix("tasks/").split("/")[0]
444 if (config_id := flat.pop("configId", None)) and "id" not in nested:
445 nested["id"] = config_id
446 flat.update(nested)
447 return flat
450def _parse(parse_dict: Callable[..., object], data: JsonDict, message: object) -> object:
451 parse_dict(data, message, ignore_unknown_fields=True)
452 return message