litellm/tests/test_litellm/proxy/a2a/test_version_convert.py
Sameer Kankute 8e30cfbeb1
feat(a2a): support a2a-sdk 1.x proxy routing for 0.3 and 1.0 agents (#30950)
* feat(a2a): support a2a-sdk 1.x proxy routing for 0.3 and 1.0 agents

Bump a2a-sdk to 1.x and wire send/stream through compat conversions so the proxy accepts A2A 1.0 JSON-RPC while preserving 0.3 wire clients.

Co-authored-by: Cursor <cursoragent@cursor.com>

* Add user controlled protocol version in agents

* Fix exeception mapping

* Fix a2a base url

* Add e2e test for a2a

* Fix lint

* Fix lint

* fix(a2a): harden card version detection and header isolation coverage

Use protocolVersion when inferring agent card wire format, assert distinct httpx cache keys in the header-isolation test, and suppress targeted basedpyright errors for optional SDK imports.

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(a2a): suppress reportArgumentType for SDK compat types and fix streaming trace ID

- Add pyright: ignore[reportArgumentType] to SendMessageSuccessResponse id= and
  result= args in _send_message, and SendStreamingMessageResponse root= in
  _stream_messages, where a2a-sdk compat types diverge from basedpyright's
  inferred signature, reducing the reportArgumentType count back within budget.
- Fix streaming trace ID in astream_a2a_message to use str(request.id) when
  available instead of always generating a new uuid4(), restoring JSON-RPC
  request-ID correlation for observability.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* style(a2a): expand SendStreamingMessageResponse for black formatting

Move pyright: ignore comment to the root= argument line so Black
accepts the expanded multi-line form.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* fix(a2a): fix 2 reportArgumentType errors without suppression

- main.py: narrow logging_obj from object|None to Optional[Logging] via
  isinstance check before A2AStreamingIterator call, fixing the
  "Logging | object" argument type mismatch at line 699.
- a2a_endpoints.py: extract response_dict with explicit isinstance(dict)
  guard before passing to normalize_jsonrpc_response, fixing the
  "LLMResponseTypes | dict[str, Any]" type mismatch at line 835.
- Remove spurious pyright: ignore comments added in previous commits that
  were not suppressing the actual errors.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* fix(a2a): rewrite upstream URL for 1.0 agent cards in getAuthenticatedExtendedCard

1.0 upstream agent cards store the endpoint URL in supportedInterfaces[0].url
rather than a top-level url field. The previous guard only rewrote url when
it existed at the top level, so after normalize_agent_card lowered a 1.0 card
to 0.3 the upstream internal address leaked into the url field of the 0.3
response.

Fix: rewrite both url and supportedInterfaces[0].url to the proxy address
before calling normalize_agent_card, ensuring the upstream address is never
visible to downstream clients regardless of the upstream card's wire format.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* fix: extend _served_version to all PascalCase methods; add direct httpx-client isolation proof

- _served_version now checks `_PASCAL_TO_WIRE` membership instead of two
  hardcoded names, so GetTask/CancelTask/etc. are promoted to 1.0 wire format
  alongside SendMessage — prevents mixed wire formats mid-session
- test_create_a2a_client_uses_fresh_httpx_client now asserts
  a2a_client_a._litellm_httpx_client is not a2a_client_b._litellm_httpx_client
  (direct proof that header bleed cannot occur), in addition to the cache-key
  inequality check

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* fix: id:0 silently dropped in version_convert; explicit continue in stream retry

- version_convert.py: replace `request_id or ""` with
  `str(request_id) if request_id is not None else ""` in both
  _send_result_to and _stream_result_to; id=0 is valid JSON-RPC and
  must not be coerced to "" which breaks response correlation
- main.py: add explicit `continue` after the A2ALocalhostURLError retry
  in _execute_a2a_stream_with_retry so the control flow (retry → next
  iteration → stream_succeeded guard) is unambiguous

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>

* fix: preserve a2a retry and discovery card urls

* Fix black

* Fix test

* fix(a2a): avoid KeyError in discovery log after 0.3→1.0 card normalization

When a 0.3-style agent card is normalized to 1.0, the top-level url key is
replaced by supportedInterfaces; log the already-computed proxy_url instead.

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(a2a): preserve taskId when lowering push notification config set params

Flatten 1.x create envelope fields before parsing into TaskPushNotificationConfig so 1.0 clients forwarding to 0.3 upstream keep taskId and config.

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(a2a): ignore unknown fields in message/send proto fallback

ParseDict in _build_message_send_params now matches other inbound paths so 1.0 clients with extra proto fields are not rejected with -32602.

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(a2a): normalize tasks/list params and response across protocol versions

Convert list task entries on the response path and lower ListTasksRequest params including status filters when forwarding 1.0 clients to 0.3 upstream.

Co-authored-by: Cursor <cursoragent@cursor.com>

* fix(a2a): avoid reportArgumentType in _lower_list_tasks_params; use local var instead of _parse return

Co-authored-by: Cursor <cursoragent@cursor.com>

* refactor(a2a): drop private SDK symbol in tasks/list status lowering

_lower_list_tasks_params imported _CORE_TO_COMPAT_TASK_STATE, a private
a2a-sdk symbol that could disappear on a patch release and silently break
status-filter lowering. Derive the 0.3 wire string from the public
protobuf enum name instead (TASK_STATE_<NAME> maps to the 0.3 value once
the prefix is dropped and underscores become dashes) and validate the
result against the 0.3 TaskState enum's own values via a fully-typed pure
helper. Behavior is unchanged for every state; unspecified or unrecognized
states still drop the filter. Adds parametrized regression tests covering
dashed wire values (input-required, auth-required) and the unspecified drop.

* fix(a2a): drop redundant push-notification envelope key; unify MessageToDict import

_flatten_create_push_notification_params used `config or pushNotificationConfig`,
which short-circuits so a co-present pushNotificationConfig key was never popped and
leaked into the flattened params. Pop both keys unconditionally and prefer config
when present. Adds a regression test on the helper that fails on the old leak.

Also import MessageToDict from a2a.compat.v0_3.conversions in _lower_list_tasks_params
to match every other conversion helper in the module instead of pulling it straight
from google.protobuf.json_format.

* fix(a2a): reject invalid message/stream params early with -32602

_handle_stream_message built MessageSendParams lazily inside the
stream_response() generator, so malformed 1.0 params surfaced as a generic
-32603 after the 200 status line was already committed. The non-streaming
path validates up front and returns -32602 (Invalid params). Validate
eagerly before returning the StreamingResponse and emit -32602 on failure
so both paths reject malformed params identically. Adds a regression test
asserting the streamed error code is -32602.

* fix(a2a): raise clear error when non-streaming send ends on an update event

_send_message fed the SDK iterator's last event straight into
SendMessageSuccessResponse, whose result only accepts Message or Task. A
non-standard upstream whose final event is a TaskStatusUpdateEvent or
TaskArtifactUpdateEvent made the response construction raise an opaque
pydantic ValidationError. Guard the converted result and raise a clear
RuntimeError instead, consistent with the no-response guard above it.
Adds regression tests for the Message happy path and the update-event
rejection via an injected fake client.

* test(a2a): lock in clean merged agent-card URL without PROXY_BASE_URL

Regression coverage proving _build_merged_agent_card produces no double
slash in supportedInterfaces[0].url when PROXY_BASE_URL is unset and
request.base_url carries a trailing slash. get_custom_url routes through
join_paths, which rstrips the base, so the f-string join stays clean.

* style(a2a): modernize type annotations to satisfy strict ruff budget

After merging the black->ruff-format migration from base, the A2A files
owned by this PR still used Optional[X]/quoted annotations that pushed
UP037/UP045 over their lowered ceilings. Convert to X | None, drop the
now-unnecessary quoted local annotation in _send_message, and remove the
imports left unused by the rewrite. Type semantics are unchanged.

* style(a2a): type a2a_endpoints dict params as dict[str, Any]

The merge with the formatter-migration baseline tightened the
reportUnknownArgumentType ceiling; bare dict annotations made every value
Unknown and pushed the codebase total over cap. Annotate the JSON-RPC
params, body, metadata, and litellm_params dicts as dict[str, Any] so
their values are typed, dropping the unknown-argument count back under the
ceiling. No behavior change.

* fix(a2a): guard localhost retry against a missing agent card

handle_a2a_localhost_retry rewrote the card URL and called create_client
with whatever agent_card it received. The caller resolves the card from
the SDK client (Optional), so a None card reached set_agent_card_url and
create_client, surfacing an opaque SDK error instead of a clear one. Add
an early RuntimeError guard mirroring the httpx-client check, drop the now
always-true card None-check on the stash line, and cover it with a
regression test.

* style(a2a): disable reportUnknownArgumentType in a2a-sdk boundary modules

The lint env type-checks without the optional a2a-sdk/protobuf installed, so
every call into the protobuf-generated compat conversions counts as an
Unknown-typed argument and the new A2A code pushed the codebase
reportUnknownArgumentType total over its ceiling. These three modules are
the A2A SDK boundary; turn the rule off file-wide with a documented reason
instead of scattering dozens of per-line ignores across every SDK call.

* fix(a2a): tolerate unknown fields when lowering 1.0->0.3; align streaming trace id

Two issues greptile flagged:

version_convert: the 1.0->0.3 lowering paths (_send_result_to, _task_to,
_stream_result_to) called ParseDict without ignore_unknown_fields=True, so a
1.0 upstream response carrying vendor extensions raised and best-effort fell
back to passing the un-lowered 1.0 shape to a 0.3 client. Set the flag to match
the agent-card path and every inbound path; unknown fields are now dropped and
the result is correctly lowered.

main.py: asend_message_streaming derived X-LiteLLM-Trace-Id from the JSON-RPC
request id, unlike asend_message which uses the logging object's
litellm_trace_id. Prefer the logging trace id (then request id, then a uuid) so
streamed and non-streamed calls correlate under the same trace.

Adds regression tests for both, including the stream-event lowering path.

* style(a2a): apply ruff format to a2a protocol and proxy modules

Co-authored-by: Cursor <cursoragent@cursor.com>

---------

Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
2026-06-29 09:32:39 +05:30

315 lines
9.6 KiB
Python

"""Unit tests for A2A protocol version normalization in
litellm/proxy/a2a/version_convert.py.
These assert the conversion actually changes wire shape in the right direction and
preserves core fields on a round trip, so a mutation that no-ops or flips the direction
fails the suite.
"""
import pytest
from litellm.proxy.a2a.version_convert import (
normalize_agent_card,
normalize_jsonrpc_response,
normalize_request_params,
normalize_stream_event,
)
a2a = pytest.importorskip("a2a.compat.v0_3.conversions")
def _rpc(result: dict, request_id: str = "1") -> dict:
return {"jsonrpc": "2.0", "id": request_id, "result": result}
V03_MESSAGE = {
"kind": "message",
"messageId": "m1",
"role": "agent",
"parts": [{"kind": "text", "text": "hi"}],
}
V03_TASK = {
"kind": "task",
"id": "t1",
"contextId": "c1",
"status": {"state": "completed"},
}
V03_STATUS_UPDATE = {
"kind": "status-update",
"taskId": "t1",
"contextId": "c1",
"status": {"state": "working"},
"final": False,
}
V03_ARTIFACT_UPDATE = {
"kind": "artifact-update",
"taskId": "t1",
"contextId": "c1",
"artifact": {"artifactId": "a1", "parts": [{"kind": "text", "text": "out"}]},
}
def test_send_result_0_3_to_1_0_wraps_in_envelope():
out = normalize_jsonrpc_response(_rpc(V03_MESSAGE), "1.0", method="message/send")
assert "message" in out["result"]
assert "kind" not in out["result"]
assert out["result"]["message"]["messageId"] == "m1"
def test_send_result_1_0_to_0_3_unwraps_to_bare_kind():
v1 = normalize_jsonrpc_response(_rpc(V03_MESSAGE), "1.0", method="message/send")
out = normalize_jsonrpc_response(v1, "0.3", method="message/send")
assert out["result"]["kind"] == "message"
assert out["result"]["messageId"] == "m1"
assert out["result"]["parts"][0]["text"] == "hi"
def test_send_result_same_version_is_identity_passthrough():
rpc = _rpc(V03_MESSAGE)
out = normalize_jsonrpc_response(rpc, "0.3", method="message/send")
assert out is rpc
def test_task_result_round_trip_preserves_ids():
v1 = normalize_jsonrpc_response(_rpc(V03_TASK), "1.0", method="tasks/get")
assert "kind" not in v1["result"]
assert v1["result"]["id"] == "t1"
back = normalize_jsonrpc_response(v1, "0.3", method="tasks/get")
assert back["result"]["kind"] == "task"
assert back["result"]["id"] == "t1"
assert back["result"]["contextId"] == "c1"
def test_error_response_passes_through_untouched():
err = {"jsonrpc": "2.0", "id": "1", "error": {"code": -32600, "message": "bad"}}
assert normalize_jsonrpc_response(err, "1.0", method="message/send") is err
def test_malformed_result_falls_back_to_passthrough():
# A 0.3 message missing required fields can't validate; conversion must not raise.
rpc = _rpc({"kind": "message"})
out = normalize_jsonrpc_response(rpc, "1.0", method="message/send")
assert out["result"] == {"kind": "message"}
def test_unknown_shape_passes_through():
rpc = _rpc({"unexpected": "shape"})
out = normalize_jsonrpc_response(rpc, "1.0", method="message/send")
assert out is rpc
@pytest.mark.parametrize("event", [V03_STATUS_UPDATE, V03_ARTIFACT_UPDATE])
def test_stream_event_round_trip_preserves_kind(event):
v1 = normalize_stream_event(_rpc(event), "1.0", request_id="1")
assert "kind" not in v1["result"]
back = normalize_stream_event(v1, "0.3", request_id="1")
assert back["result"]["kind"] == event["kind"]
assert back["result"]["taskId"] == "t1"
def test_stream_event_envelope_key_for_status_update():
v1 = normalize_stream_event(_rpc(V03_STATUS_UPDATE), "1.0", request_id="1")
assert "statusUpdate" in v1["result"]
def test_request_params_lowering_is_noop_for_0_3():
params = {"id": "t1", "historyLength": 5}
assert normalize_request_params(params, "0.3", method="tasks/get") is params
def test_request_params_lowering_get_task_to_0_3():
out = normalize_request_params(
{"id": "t1", "historyLength": 5}, "1.0", method="tasks/get"
)
assert out["id"] == "t1"
assert out["historyLength"] == 5
def test_request_params_lowering_create_push_notification_config_preserves_task_id():
out = normalize_request_params(
{
"parent": "tasks/task-1",
"configId": "cfg-1",
"config": {"url": "https://webhook.example.com"},
},
"1.0",
method="tasks/pushNotificationConfig/set",
)
assert out["taskId"] == "task-1"
assert out["pushNotificationConfig"]["url"] == "https://webhook.example.com"
assert out["pushNotificationConfig"]["id"] == "cfg-1"
def test_flatten_create_push_notification_drops_redundant_envelope_key():
from litellm.proxy.a2a.version_convert import (
_flatten_create_push_notification_params,
)
flat = _flatten_create_push_notification_params(
{
"parent": "tasks/task-1",
"config": {"url": "https://chosen.example.com"},
"pushNotificationConfig": {"url": "https://ignored.example.com"},
}
)
assert flat["url"] == "https://chosen.example.com"
assert "pushNotificationConfig" not in flat
assert "config" not in flat
def test_request_params_lowering_list_tasks_to_0_3():
out = normalize_request_params(
{
"contextId": "ctx-1",
"pageSize": 10,
"status": "TASK_STATE_COMPLETED",
},
"1.0",
method="tasks/list",
)
assert out["contextId"] == "ctx-1"
assert out["pageSize"] == 10
assert out["status"] == "completed"
@pytest.mark.parametrize(
"proto_status, expected",
[
("TASK_STATE_COMPLETED", "completed"),
("TASK_STATE_INPUT_REQUIRED", "input-required"),
("TASK_STATE_AUTH_REQUIRED", "auth-required"),
("TASK_STATE_CANCELED", "canceled"),
],
)
def test_list_tasks_status_filter_lowers_to_0_3_wire_value(proto_status, expected):
out = normalize_request_params(
{"status": proto_status},
"1.0",
method="tasks/list",
)
assert out["status"] == expected
def test_list_tasks_unspecified_status_is_dropped():
out = normalize_request_params(
{"contextId": "ctx-1", "status": "TASK_STATE_UNSPECIFIED"},
"1.0",
method="tasks/list",
)
assert "status" not in out
assert out["contextId"] == "ctx-1"
@pytest.mark.parametrize(
"method, result",
[
(
"message/send",
{
"task": {
"id": "t1",
"contextId": "c1",
"status": {"state": "completed"},
},
"vendorExtraField": "x",
},
),
(
"tasks/get",
{
"id": "t1",
"contextId": "c1",
"status": {"state": "completed"},
"vendorExtraField": "x",
},
),
],
)
def test_lowering_1_0_to_0_3_tolerates_unknown_upstream_fields(method, result):
out = normalize_jsonrpc_response(_rpc(result), "0.3", method=method)
lowered = out["result"]
assert lowered["kind"] == "task"
assert lowered["id"] == "t1"
assert "vendorExtraField" not in lowered
def test_stream_event_lowering_1_0_to_0_3_tolerates_unknown_fields():
event = {
"task": {"id": "t1", "contextId": "c1", "status": {"state": "completed"}},
"vendorExtraField": "x",
}
out = normalize_stream_event(_rpc(event), "0.3", request_id="1")
lowered = out["result"]
assert lowered["kind"] == "task"
assert lowered["id"] == "t1"
def test_list_tasks_result_round_trip_preserves_task_ids():
rpc = _rpc(
{
"tasks": [
{
"kind": "task",
"id": "t1",
"contextId": "c1",
"status": {"state": "completed"},
}
],
"nextPageToken": "tok",
}
)
v1 = normalize_jsonrpc_response(rpc, "1.0", method="tasks/list")
assert "kind" not in v1["result"]["tasks"][0]
assert v1["result"]["tasks"][0]["id"] == "t1"
back = normalize_jsonrpc_response(v1, "0.3", method="tasks/list")
assert back["result"]["tasks"][0]["kind"] == "task"
assert back["result"]["tasks"][0]["id"] == "t1"
def _extended_card_1_0() -> dict:
return {
"name": "Card",
"description": "d",
"version": "1.0.0",
"supportedInterfaces": [
{
"url": "https://upstream.example",
"protocolBinding": "JSONRPC",
"protocolVersion": "0.3",
},
{
"url": "http://internal:9999",
"protocolBinding": "JSONRPC",
"protocolVersion": "0.3",
},
],
}
def test_agent_card_lowered_to_0_3_drops_additional_interfaces():
# A 1.0 card with multiple interfaces would lower into a 0.3 card carrying the
# secondary backend URLs in ``additionalInterfaces``; those must be stripped so
# the conversion never re-exposes an upstream backend to A2A clients.
out = normalize_agent_card(_extended_card_1_0(), "0.3")
assert out["url"] == "https://upstream.example"
assert "additionalInterfaces" not in out
assert "supportedInterfaces" not in out
assert "http://internal:9999" not in str(out)
def test_agent_card_with_0_3_pin_and_supported_interfaces_is_lowered():
card = _extended_card_1_0()
card["protocolVersion"] = "0.3"
out = normalize_agent_card(card, "0.3")
assert out["protocolVersion"] == "0.3"
assert "supportedInterfaces" not in out
def test_agent_card_same_version_passthrough():
card = _extended_card_1_0()
assert normalize_agent_card(card, "1.0") is card