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

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. 

8 

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. 

16 

17The two wire shapes: 

18 

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""" 

26 

27from collections.abc import Callable 

28from typing import Final, Literal 

29 

30from pydantic import BaseModel 

31 

32from litellm._logging import verbose_proxy_logger 

33from litellm.proxy.a2a.agent_card import normalize_protocol_version 

34 

35A2AVersion = Literal["0.3", "1.0"] 

36RequestId = str | int | None 

37JsonDict = dict[str, object] 

38 

39_V1_SEND_ENVELOPE_KEYS: Final = frozenset({"message", "task"}) 

40_V1_STREAM_ENVELOPE_KEYS: Final = frozenset({"message", "task", "statusUpdate", "artifactUpdate"}) 

41 

42 

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") 

46 

47 

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 

55 

56 

57def normalize_jsonrpc_response(content: JsonDict, target: A2AVersion, *, method: str) -> JsonDict: 

58 """Convert a JSON-RPC response's ``result`` to ``target``. 

59 

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 

67 

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} 

72 

73 

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 

81 

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} 

86 

87 

88def normalize_request_params(params: JsonDict, served: A2AVersion, *, method: str) -> JsonDict: 

89 """Down-convert forwarded request ``params`` from the served version to 0.3. 

90 

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 ) 

102 

103 

104def _detect_card_version(card: JsonDict) -> A2AVersion: 

105 """Infer the wire version of an agent card dict. 

106 

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" 

116 

117 

118def normalize_agent_card(card: JsonDict, target: A2AVersion) -> JsonDict: 

119 """Convert an extended agent card to ``target``. 

120 

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 

126 

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") 

131 

132 

133def _as_request_id(value: object) -> RequestId: 

134 return value if isinstance(value, (str, int)) else None 

135 

136 

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 

151 

152 

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 

159 

160 

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 ) 

170 

171 

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 ) 

181 

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 ) 

194 

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) 

198 

199 

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") 

205 

206 

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" 

215 

216 

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 ) 

226 

227 

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 } 

236 

237 

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 ) 

247 

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) 

251 

252 pb: Final = pb2_v10.Task() 

253 ParseDict(result, pb, ignore_unknown_fields=True) 

254 return _dump_03(to_compat_task(pb)) 

255 

256 

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 

263 

264 

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 ) 

274 

275 

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 ) 

285 

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) 

293 

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) 

297 

298 

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 ) 

308 

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 

315 

316 core: Final = to_core_agent_card(types_v03.AgentCard.model_validate(card)) 

317 return MessageToDict(core, preserving_proto_field_name=False) 

318 

319 

320def _validate_message_or_task(result: JsonDict) -> BaseModel: 

321 from a2a.compat.v0_3.conversions import types_v03 

322 

323 if result.get("kind") == "task": 

324 return types_v03.Task.model_validate(result) 

325 return types_v03.Message.model_validate(result) 

326 

327 

328def _validate_stream_event(result: JsonDict) -> BaseModel: 

329 from a2a.compat.v0_3.conversions import types_v03 

330 

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) 

339 

340 

341def _lower_request_params(params: JsonDict, *, method: str) -> JsonDict: 

342 if method == "tasks/list": 

343 return _lower_list_tasks_params(params) 

344 

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 ) 

356 

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)) 

395 

396 

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 ) 

404 

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 

416 

417 

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. 

420 

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 

431 

432 

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 

448 

449 

450def _parse(parse_dict: Callable[..., object], data: JsonDict, message: object) -> object: 

451 parse_dict(data, message, ignore_unknown_fields=True) 

452 return message