litellm/tests/unit/proxy/agent_endpoints/test_a2a_endpoints.py
devin-ai-integration[bot] 25109a523b
test(proxy): move utils, agent_endpoints and endpoint tests into tests/unit/proxy (#44006)
* test(proxy): move utils, agent_endpoints and endpoint tests into tests/unit/proxy

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): package moved unit test directories

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): exclude proxy-db-owned files from the misc target

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): drop the redundant fixture docstrings in the proxy conftest

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yuneng <yuneng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
2026-10-01 11:06:42 -07:00

2691 lines
97 KiB
Python

"""
Mock tests for A2A endpoints.
Tests that invoke_agent_a2a properly integrates with add_litellm_data_to_request.
"""
import json
import socket
import sys
from collections.abc import Awaitable, Callable, Mapping
from contextlib import AbstractContextManager, ExitStack
from dataclasses import dataclass
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.agents import AgentCaller
AddLiteLLMData = Callable[..., Awaitable[dict[str, object]]]
@dataclass(frozen=True, slots=True)
class CapturedAgentCall:
request_id: object
agent_extra_headers: dict[str, str] | None
@pytest.mark.asyncio
async def test_invoke_agent_a2a_adds_litellm_data():
"""
Test that invoke_agent_a2a calls add_litellm_data_to_request
and the resulting data includes proxy_server_request.
"""
from litellm.proxy._types import UserAPIKeyAuth
# Track the data passed to add_litellm_data_to_request
captured_data = {}
async def mock_add_litellm_data(data, **kwargs):
# Simulate what add_litellm_data_to_request does
data["proxy_server_request"] = {
"url": "http://localhost:4000/a2a/test-agent",
"method": "POST",
"headers": {},
"body": dict(data),
}
captured_data.update(data)
return data
# Mock response from asend_message
mock_response = MagicMock()
mock_response.model_dump.return_value = {
"jsonrpc": "2.0",
"id": "test-id",
"result": {"status": "success"},
}
# Mock agent
mock_agent = MagicMock()
mock_agent.agent_id = "test-agent"
mock_agent.agent_card_params = {
"url": "http://backend-agent:10001",
"name": "Test Agent",
}
mock_agent.litellm_params = None
# Mock request
mock_request = MagicMock()
mock_request.json = AsyncMock(
return_value={
"jsonrpc": "2.0",
"id": "test-id",
"method": "message/send",
"metadata": {"model_info": {"id": "caller-supplied-id"}},
"params": {
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "Hello"}],
"messageId": "msg-123",
}
},
}
)
mock_user_api_key_dict = UserAPIKeyAuth(
api_key="sk-test-key",
user_id="test-user",
team_id="test-team",
)
# Try to use real a2a.types if available, otherwise create realistic mocks
# This test focuses on LiteLLM integration, not A2A protocol correctness,
# but we want mocks that behave like the real types to catch usage issues
try:
from a2a.types import (
MessageSendParams,
SendMessageRequest,
SendStreamingMessageRequest,
)
# Real types available - use them
pass
except ImportError:
# Real types not available - create realistic mocks
pass
def make_mock_pydantic_class(name):
"""Create a mock class that behaves like a Pydantic model."""
class MockPydanticClass:
def __init__(self, **kwargs):
self.__dict__.update(kwargs)
# Store kwargs for model_dump() if needed
self._kwargs = kwargs
def model_dump(self, mode="json", exclude_none=False):
"""Mock model_dump method."""
result = dict(self._kwargs)
if exclude_none:
result = {k: v for k, v in result.items() if v is not None}
return result
MockPydanticClass.__name__ = name
return MockPydanticClass
MessageSendParams = make_mock_pydantic_class("MessageSendParams")
SendMessageRequest = make_mock_pydantic_class("SendMessageRequest")
SendStreamingMessageRequest = make_mock_pydantic_class("SendStreamingMessageRequest")
# Create a mock module for a2a.types
mock_a2a_types = MagicMock()
mock_a2a_types.MessageSendParams = MessageSendParams
mock_a2a_types.SendMessageRequest = SendMessageRequest
mock_a2a_types.SendStreamingMessageRequest = SendStreamingMessageRequest
# Patch at the source modules
# Note: add_litellm_data_to_request is called from common_request_processing,
# so we need to patch it there, not at litellm_pre_call_utils
with (
patch(
"litellm.proxy.agent_endpoints.a2a_endpoints._get_agent",
return_value=mock_agent,
),
patch(
"litellm.proxy.common_request_processing.add_litellm_data_to_request",
side_effect=mock_add_litellm_data,
) as mock_add_data,
patch(
"litellm.a2a_protocol.create_a2a_client",
new_callable=AsyncMock,
),
patch(
"litellm.a2a_protocol.asend_message",
new_callable=AsyncMock,
return_value=mock_response,
) as mock_send_message,
patch(
"litellm.proxy.proxy_server.general_settings",
{},
),
patch(
"litellm.proxy.proxy_server.proxy_config",
MagicMock(),
),
patch(
"litellm.proxy.proxy_server.version",
"1.0.0",
),
patch.dict(
sys.modules,
{"a2a": MagicMock(), "a2a.types": mock_a2a_types},
),
patch(
"litellm.a2a_protocol.main.A2A_SDK_AVAILABLE",
True,
),
):
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
mock_fastapi_response = MagicMock()
await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=mock_fastapi_response,
user_api_key_dict=mock_user_api_key_dict,
)
# Verify add_litellm_data_to_request was called
mock_add_data.assert_called_once()
# Verify model and custom_llm_provider were set
assert mock_send_message.await_args.kwargs["model"] == "a2a_agent/Test Agent"
assert captured_data["metadata"]["model_group"] == "a2a_agent/Test Agent"
assert captured_data["metadata"]["model_info"] == {"id": mock_agent.agent_id}
assert captured_data.get("model") == "a2a_agent/Test Agent"
assert captured_data.get("custom_llm_provider") == "a2a_agent"
# Verify proxy_server_request was added
assert "proxy_server_request" in captured_data
assert captured_data["proxy_server_request"]["method"] == "POST"
@pytest.mark.asyncio
async def test_invoke_agent_a2a_handles_none_agent_card_params():
"""Agents without ``agent_card_params`` (e.g. plain chat agents routed
through the A2A endpoint by mistake) must not raise ``AttributeError`` on
``agent_card_params.get(...)`` — they should return a JSON-RPC error.
"""
from litellm.proxy._types import UserAPIKeyAuth
mock_agent = MagicMock()
mock_agent.agent_card_params = None
mock_agent.litellm_params = None
mock_request = MagicMock()
mock_request.json = AsyncMock(
return_value={
"jsonrpc": "2.0",
"id": "test-id",
"method": "message/send",
"params": {
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "Hello"}],
"messageId": "msg-123",
}
},
}
)
mock_user_api_key_dict = UserAPIKeyAuth(
api_key="sk-test-key",
user_id="test-user",
team_id="test-team",
)
with (
patch(
"litellm.proxy.agent_endpoints.a2a_endpoints._get_agent",
return_value=mock_agent,
),
patch(
"litellm.a2a_protocol.main.A2A_SDK_AVAILABLE",
True,
),
patch.dict(sys.modules, {"a2a": MagicMock(), "a2a.types": MagicMock()}),
):
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
mock_fastapi_response = MagicMock()
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=mock_fastapi_response,
user_api_key_dict=mock_user_api_key_dict,
)
# JSONResponse exposes the body bytes; decode and verify it's a
# JSON-RPC error, not an "internal error" from a Python exception.
body = json.loads(response.body.decode())
assert body["jsonrpc"] == "2.0"
assert body["error"]["code"] == -32000
assert "no URL configured" in body["error"]["message"]
@pytest.mark.asyncio
async def test_invoke_agent_a2a_injects_authenticated_key_hash_for_bridge():
"""Completion-bridge agents must receive the authenticated key hash in
litellm_params so provider configs (e.g. LangFlow) can scope provider-side
session memory per key. Regression for cross-key A2A session bleed."""
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2A_USER_API_KEY_HASH_PARAM,
)
from litellm.proxy._types import UserAPIKeyAuth
captured = {}
async def mock_add_litellm_data(data, **kwargs):
data["proxy_server_request"] = {
"url": "http://localhost:4000/a2a/lf-agent",
"method": "POST",
"headers": {},
"body": {},
}
data.setdefault("metadata", {})
return data
async def capture_asend_message(**kwargs):
captured.update(kwargs)
resp = MagicMock()
resp.model_dump.return_value = {"jsonrpc": "2.0", "id": "test-id", "result": {}}
return resp
mock_agent = MagicMock()
mock_agent.agent_id = "lf-agent"
mock_agent.agent_name = "lf-agent"
# No URL: the bridge derives the endpoint from the LangFlow agent config.
mock_agent.agent_card_params = {"name": "LF Agent"}
mock_agent.litellm_params = {
"custom_llm_provider": "langflow",
"model": "langflow/flow-1",
}
mock_agent.static_headers = None
mock_agent.extra_headers = None
mock_request = MagicMock()
mock_request.headers = {}
mock_request.json = AsyncMock(
return_value={
"jsonrpc": "2.0",
"id": "test-id",
"method": "message/send",
"params": {
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
"contextId": "ctx-1",
}
},
}
)
mock_user_api_key_dict = UserAPIKeyAuth(
api_key="sk-hashed-123",
user_id="test-user",
team_id="test-team",
)
with (
patch(
"litellm.proxy.agent_endpoints.a2a_endpoints._get_agent",
return_value=mock_agent,
),
patch(
"litellm.proxy.common_request_processing.add_litellm_data_to_request",
side_effect=mock_add_litellm_data,
),
patch(
"litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed",
new=AsyncMock(return_value=True),
),
patch(
"litellm.a2a_protocol.asend_message",
new=AsyncMock(side_effect=capture_asend_message),
),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.proxy_config", MagicMock()),
patch("litellm.proxy.proxy_server.version", "1.0.0"),
patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True),
patch.dict(sys.modules, {"a2a": MagicMock(), "a2a.types": MagicMock()}),
):
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
await invoke_agent_a2a(
agent_id="lf-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=mock_user_api_key_dict,
)
assert captured.get("litellm_params", {}).get(A2A_USER_API_KEY_HASH_PARAM) == mock_user_api_key_dict.api_key, (
"authenticated key hash was not forwarded to the completion bridge"
)
def _make_agent_mock(url: str = "http://backend-agent:10001") -> MagicMock:
agent = MagicMock()
agent.agent_id = "test-agent"
agent.agent_name = "test-agent"
agent.agent_card_params = {"url": url, "name": "Test Agent"}
agent.litellm_params = {}
agent.static_headers = None
agent.extra_headers = None
return agent
def _make_request_mock(method: str, params: Mapping[str, object], request_id: object = "req-1") -> MagicMock:
req = MagicMock()
req.headers = {}
req.json = AsyncMock(
return_value={
"jsonrpc": "2.0",
"id": request_id,
"method": method,
"params": params,
}
)
return req
def _base_patches(
agent: MagicMock, add_litellm_data: AddLiteLLMData | None = None
) -> list[AbstractContextManager[object]]:
return [
patch(
"litellm.proxy.agent_endpoints.a2a_endpoints._get_agent",
return_value=agent,
),
patch(
"litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed",
new=AsyncMock(return_value=True),
),
patch(
"litellm.proxy.common_request_processing.add_litellm_data_to_request",
new=AsyncMock(side_effect=add_litellm_data or _add_proxy_data),
),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.proxy_config", MagicMock()),
patch("litellm.proxy.proxy_server.version", "1.0.0"),
]
async def _add_proxy_data(data: dict[str, object], **kwargs: object) -> dict[str, object]:
return {
**data,
"proxy_server_request": {"url": "http://localhost:4000", "method": "POST", "headers": {}, "body": {}},
"metadata": data.get("metadata", {}),
}
_HELLO_MESSAGE_PARAMS = {
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "Hello"}],
"messageId": "msg-123",
}
}
async def _invoke_message_method(
method: str,
mock_request: MagicMock,
user_api_key_dict: UserAPIKeyAuth,
add_litellm_data: AddLiteLLMData | None = None,
agent: MagicMock | None = None,
) -> CapturedAgentCall:
from fastapi.responses import JSONResponse
class MessageSendParams:
def __init__(self, **kwargs: object) -> None:
self.__dict__.update(kwargs)
class SendMessageRequest:
def __init__(self, **kwargs: object) -> None:
self.__dict__.update(kwargs)
async def fake_asend_message(request: SendMessageRequest, **kwargs: object) -> MagicMock:
response: Final = MagicMock()
response.model_dump.return_value = {
"jsonrpc": "2.0",
"id": request.__dict__["id"],
"result": {"status": "success"},
}
return response
async def fake_stream_message(request_id: object, **kwargs: object) -> JSONResponse:
return JSONResponse({"jsonrpc": "2.0", "id": request_id})
mock_a2a_types: Final = MagicMock()
mock_a2a_types.MessageSendParams = MessageSendParams
mock_a2a_types.SendMessageRequest = SendMessageRequest
is_send: Final = method == "message/send"
downstream: Final = AsyncMock(side_effect=fake_asend_message if is_send else fake_stream_message)
with ExitStack() as stack:
for p in _base_patches(agent or _make_agent_mock(), add_litellm_data):
stack.enter_context(p)
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
if is_send:
stack.enter_context(patch.dict(sys.modules, {"a2a": MagicMock(), "a2a.types": mock_a2a_types}))
stack.enter_context(patch("litellm.a2a_protocol.asend_message", new=downstream))
else:
stack.enter_context(
patch("litellm.proxy.agent_endpoints.a2a_endpoints._handle_stream_message", new=downstream)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
kwargs: Final = downstream.call_args.kwargs
request_id: Final = kwargs["request"].__dict__["id"] if is_send else kwargs["request_id"]
return CapturedAgentCall(request_id=request_id, agent_extra_headers=kwargs.get("agent_extra_headers"))
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
async def test_message_methods_preserve_numeric_zero_request_id(method: str):
mock_request = _make_request_mock(method, _HELLO_MESSAGE_PARAMS, request_id=0)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
captured = await _invoke_message_method(method, mock_request, user_api_key_dict)
assert captured.request_id == 0
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
async def test_message_methods_forward_caller_identity_headers(method: str):
mock_request = _make_request_mock(method, _HELLO_MESSAGE_PARAMS)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="user-abc", team_id="team-xyz")
captured = await _invoke_message_method(method, mock_request, user_api_key_dict)
forwarded_headers = captured.agent_extra_headers or {}
assert forwarded_headers.get("X-LiteLLM-User-Id") == "user-abc"
assert forwarded_headers.get("X-LiteLLM-Team-Id") == "team-xyz"
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
async def test_agent_calling_another_agent_forwards_the_human_who_invoked_it(method: str):
"""LIT-8014: an agent acting for alice calls a second agent through the proxy. That hop must
carry alice, not the first agent's owner, so the chain stays capped at what alice may reach."""
mock_request = _make_request_mock(method, _HELLO_MESSAGE_PARAMS)
agent_key = UserAPIKeyAuth(api_key="sk-agent", user_id="agent-owner", team_id="agent-team", agent_id="agent-1")
agent_key.agent_caller = AgentCaller(user_id="alice", team_id="callers")
captured = await _invoke_message_method(method, mock_request, agent_key)
forwarded_headers = captured.agent_extra_headers or {}
assert (forwarded_headers.get("X-LiteLLM-User-Id"), forwarded_headers.get("X-LiteLLM-Team-Id")) == (
"alice",
"callers",
)
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
async def test_message_methods_send_the_entra_bearer_for_azure_agents(method: str):
"""A Microsoft Foundry agent accepts only an Entra ID bearer, so an agent registered with
Entra credentials in litellm_params must reach the backend with that bearer on every call."""
agent = _make_agent_mock()
agent.litellm_params = {"azure_ad_token": "entra-token"}
mock_request = _make_request_mock(method, _HELLO_MESSAGE_PARAMS)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
captured = await _invoke_message_method(method, mock_request, user_api_key_dict, agent=agent)
assert (captured.agent_extra_headers or {}).get("Authorization") == "Bearer entra-token"
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
async def test_message_methods_leave_agents_without_entra_params_unauthenticated(method: str):
mock_request = _make_request_mock(method, _HELLO_MESSAGE_PARAMS)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
captured = await _invoke_message_method(method, mock_request, user_api_key_dict)
assert "Authorization" not in (captured.agent_extra_headers or {})
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
async def test_message_methods_leave_entra_fields_to_the_model_provider_for_bridge_agents(method: str):
"""A completion-bridge agent's tenant_id/client_id/client_secret belong to the model provider it
calls through litellm, so the proxy must not mint a Foundry bearer for them."""
agent = _make_agent_mock()
agent.litellm_params = {
"custom_llm_provider": "azure_ai",
"model": "azure_ai/foundry-model",
"tenant_id": "tenant",
"client_id": "client",
"client_secret": "sp-secret",
}
mock_request = _make_request_mock(method, _HELLO_MESSAGE_PARAMS)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
captured = await _invoke_message_method(method, mock_request, user_api_key_dict, agent=agent)
assert "Authorization" not in (captured.agent_extra_headers or {})
@pytest.mark.asyncio
async def test_message_send_reports_an_unresolvable_entra_credential_as_internal_error(monkeypatch):
"""An agent whose Entra credential points at an unset environment variable must fail the call
with the JSON-RPC internal error naming the credential fields, never reach the backend unauthenticated."""
monkeypatch.delenv("LITELLM_TEST_UNSET_FOUNDRY_TOKEN", raising=False)
agent = _make_agent_mock()
agent.litellm_params = {"azure_ad_token": "os.environ/LITELLM_TEST_UNSET_FOUNDRY_TOKEN"}
mock_request = _make_request_mock("message/send", _HELLO_MESSAGE_PARAMS)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(
side_effect=lambda user_api_key_dict, data, call_type, skip_guardrails=False: data
)
mock_proxy_logging.post_call_failure_hook = AsyncMock(return_value=None)
downstream = AsyncMock()
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch( # test-quality-ok: same proxy_logging_obj injection the sibling failure-hook tests use; the request must fail before any backend call is made
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging
)
)
stack.enter_context(
patch( # test-quality-ok: the observation point proving the backend is never called; the sibling send tests use the same seam
"litellm.a2a_protocol.asend_message", new=downstream
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
body = json.loads(response.body.decode())
assert response.status_code == 500
assert body["error"]["code"] == -32603
assert "client_secret" in body["error"]["message"]
downstream.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
async def test_message_methods_caller_identity_headers_cannot_be_spoofed(method: str):
mock_request = _make_request_mock(method, _HELLO_MESSAGE_PARAMS)
mock_request.headers = {
"x-a2a-test-agent-x-litellm-user-id": "attacker-user",
"x-a2a-test-agent-x-litellm-team-id": "attacker-team",
}
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="real-user", team_id="real-team")
captured = await _invoke_message_method(method, mock_request, user_api_key_dict)
forwarded_headers = captured.agent_extra_headers or {}
assert forwarded_headers.get("X-LiteLLM-User-Id") == "real-user", (
"authenticated user id must not be overridden by forwarded client headers"
)
assert forwarded_headers.get("X-LiteLLM-Team-Id") == "real-team", (
"authenticated team id must not be overridden by forwarded client headers"
)
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["message/send", "message/stream"])
async def test_message_methods_forward_key_bound_identity_not_pre_call_rewrite(method: str):
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
mock_request = _make_request_mock(method, _HELLO_MESSAGE_PARAMS)
mock_request.headers = {"X-OpenWebUI-User-Id": "header-mapped-user"}
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="key-user", team_id="key-team")
general_settings: Final = {
"user_header_mappings": [{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"}]
}
async def apply_user_header_mapping(data: dict[str, object], **kwargs: object) -> dict[str, object]:
LiteLLMProxyRequestSetup.add_internal_user_from_user_mapping(
general_settings, user_api_key_dict, dict(mock_request.headers)
)
return await _add_proxy_data(data, **kwargs)
captured = await _invoke_message_method(
method, mock_request, user_api_key_dict, add_litellm_data=apply_user_header_mapping
)
assert user_api_key_dict.user_id == "header-mapped-user", "precondition: pre-call rewrite ran"
forwarded_headers = captured.agent_extra_headers or {}
assert forwarded_headers.get("X-LiteLLM-User-Id") == "key-user"
assert forwarded_headers.get("X-LiteLLM-Team-Id") == "key-team"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"method,params",
[
("tasks/get", {"id": "task-1"}),
("tasks/list", {"contextId": "ctx-1"}),
("tasks/cancel", {"id": "task-1"}),
(
"tasks/pushNotificationConfig/set",
{"taskId": "task-1", "url": "https://webhook.example.com"},
),
("tasks/pushNotificationConfig/get", {"taskId": "task-1", "id": "cfg-1"}),
("tasks/pushNotificationConfig/list", {"taskId": "task-1"}),
("tasks/pushNotificationConfig/delete", {"taskId": "task-1", "id": "cfg-1"}),
],
)
async def test_task_methods_forward_jsonrpc(method: str, params: dict):
from litellm.proxy._types import UserAPIKeyAuth
upstream_response = {
"jsonrpc": "2.0",
"id": "req-1",
"result": {"id": "task-1", "status": {"state": "completed"}},
}
agent = _make_agent_mock()
mock_request = _make_request_mock(method, params)
mock_http_response = MagicMock()
mock_http_response.json.return_value = upstream_response
mock_http_response.is_success = True
mock_http_response.raise_for_status = MagicMock()
mock_handler = MagicMock()
mock_handler.post = AsyncMock(return_value=mock_http_response)
mock_handler.client = MagicMock()
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=mock_handler,
)
)
stack.enter_context(
patch(
"litellm.proxy.agent_endpoints.a2a_endpoints.validate_url",
return_value=("https://webhook.example.com", "webhook.example.com"),
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
body = json.loads(response.body.decode())
assert body["jsonrpc"] == "2.0"
assert body["result"]["id"] == "task-1"
posted = mock_handler.post.call_args
assert posted is not None
forwarded_body = posted.kwargs.get("json") or posted.args[1]
assert forwarded_body["method"] == method
@pytest.mark.asyncio
async def test_task_methods_forward_the_entra_bearer_for_azure_agents():
"""tasks/get on a Foundry agent polls the task the agent created, so the forwarded call needs
the same Entra bearer as message/send."""
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
agent.litellm_params = {"azure_ad_token": "entra-token"}
mock_request = _make_request_mock("tasks/get", {"id": "task-1"})
mock_http_response = MagicMock()
mock_http_response.json.return_value = {"jsonrpc": "2.0", "id": "req-1", "result": {"id": "task-1"}}
mock_http_response.is_success = True
mock_http_response.raise_for_status = MagicMock()
mock_handler = MagicMock()
mock_handler.post = AsyncMock(return_value=mock_http_response)
mock_handler.client = MagicMock()
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch( # test-quality-ok: the task route builds its own httpx client; the sibling task tests capture the post through the same seam
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client", return_value=mock_handler
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1"),
)
posted_headers = mock_handler.post.call_args.kwargs["headers"]
assert posted_headers["Authorization"] == "Bearer entra-token"
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["tasks/get", "tasks/resubscribe"])
async def test_task_methods_extract_litellm_params_before_forwarding(method: str):
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
params = {
"id": "task-1",
"guardrails": ["guardrail-1"],
"tags": ["tag-1"],
}
mock_request = _make_request_mock(method, params)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
captured_data = {}
async def capture_proxy_data(data, **kwargs):
captured_data.update(data)
return await _add_proxy_data(data, **kwargs)
upstream_response = {
"jsonrpc": "2.0",
"id": "req-1",
"result": {"id": "task-1", "status": {"state": "completed"}},
}
mock_http_response = MagicMock()
mock_http_response.json.return_value = upstream_response
mock_http_response.is_success = True
async def fake_aiter_lines():
yield 'data: {"jsonrpc":"2.0","id":"req-1","result":{"taskId":"task-1"}}'
mock_resp = AsyncMock()
mock_resp.is_success = True
mock_resp.aiter_lines = fake_aiter_lines
mock_resp.aclose = AsyncMock()
mock_async_client = MagicMock()
mock_async_client.build_request = MagicMock(return_value=MagicMock())
mock_async_client.send = AsyncMock(return_value=mock_resp)
mock_handler = MagicMock()
mock_handler.post = AsyncMock(return_value=mock_http_response)
mock_handler.client = mock_async_client
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch(
"litellm.proxy.common_request_processing.add_litellm_data_to_request",
new=AsyncMock(side_effect=capture_proxy_data),
)
)
stack.enter_context(
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=mock_handler,
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
if method == "tasks/resubscribe":
async for _ in response.body_iterator:
pass
if method == "tasks/resubscribe":
forwarded_body = mock_async_client.build_request.call_args.kwargs["json"]
else:
forwarded_body = mock_handler.post.call_args.kwargs["json"]
assert forwarded_body["params"] == {"id": "task-1"}
assert captured_data["guardrails"] == ["guardrail-1"]
assert captured_data["tags"] == ["tag-1"]
@pytest.mark.asyncio
async def test_subscribe_to_task_returns_sse_stream():
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
mock_request = _make_request_mock("SubscribeToTask", {"id": "task-1"})
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
sse_lines = [
'data: {"jsonrpc":"2.0","id":"req-1","result":{"taskId":"task-1","status":{"state":"working"}}}',
'data: {"jsonrpc":"2.0","id":"req-1","result":{"taskId":"task-1","status":{"state":"completed"}}}',
]
async def fake_aiter_lines():
for line in sse_lines:
yield line
mock_resp = AsyncMock()
mock_resp.is_success = True
mock_resp.aiter_lines = fake_aiter_lines
mock_resp.aclose = AsyncMock()
mock_async_client = MagicMock()
mock_async_client.build_request = MagicMock(return_value=MagicMock())
mock_async_client.send = AsyncMock(return_value=mock_resp)
mock_handler = MagicMock()
mock_handler.client = mock_async_client
mock_handler.post = AsyncMock()
chunks = []
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=mock_handler,
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
assert response.media_type == "text/event-stream"
async for chunk in response.body_iterator:
chunks.append(chunk)
full = "".join(chunks)
assert "working" in full
assert "completed" in full
@pytest.mark.asyncio
async def test_subscribe_to_task_calls_pre_call_hook():
"""tasks/resubscribe must run pre_call_hook so guardrails configured on
the agent are enforced before streaming begins."""
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
mock_request = _make_request_mock("tasks/resubscribe", {"id": "task-1"})
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
async def fake_aiter_lines():
yield 'data: {"jsonrpc":"2.0","id":"req-1","result":{"taskId":"task-1","status":{"state":"completed"}}}'
mock_resp = AsyncMock()
mock_resp.is_success = True
mock_resp.aiter_lines = fake_aiter_lines
mock_resp.aclose = AsyncMock()
mock_async_client = MagicMock()
mock_async_client.build_request = MagicMock(return_value=MagicMock())
mock_async_client.send = AsyncMock(return_value=mock_resp)
mock_handler = MagicMock()
mock_handler.client = mock_async_client
mock_handler.post = AsyncMock()
async def _passthrough_iterator(response, **kwargs):
async for chunk in response:
yield chunk
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(
side_effect=lambda user_api_key_dict, data, call_type, skip_guardrails=False: data
)
mock_proxy_logging.async_post_call_streaming_iterator_hook = _passthrough_iterator
mock_proxy_logging.post_call_failure_hook = AsyncMock(return_value=None)
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=mock_handler,
)
)
stack.enter_context(
patch(
"litellm.proxy.proxy_server.proxy_logging_obj",
mock_proxy_logging,
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
assert response.media_type == "text/event-stream"
async for _ in response.body_iterator:
pass
mock_proxy_logging.pre_call_hook.assert_awaited_once()
call_kwargs = mock_proxy_logging.pre_call_hook.await_args.kwargs
assert call_kwargs.get("call_type") == "asend_message"
assert call_kwargs.get("user_api_key_dict") == user_api_key_dict
@pytest.mark.asyncio
async def test_subscribe_to_task_runs_post_call_streaming_guardrail():
"""tasks/resubscribe must route streamed events through the post-call
streaming hook so output guardrails configured on the agent inspect the
streamed task content. Regression: the SSE path previously returned the raw
upstream stream and bypassed guardrails entirely."""
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._types import UserAPIKeyAuth
inspected: list = []
class _RecordingGuardrail(CustomGuardrail):
async def async_post_call_streaming_hook(self, user_api_key_dict, response):
inspected.append(response)
return response
guardrail = _RecordingGuardrail(guardrail_name="record-a2a", default_on=True, event_hook="post_call")
agent = _make_agent_mock()
mock_request = _make_request_mock("tasks/resubscribe", {"id": "task-1"})
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
async def fake_aiter_lines():
yield (
'data: {"jsonrpc":"2.0","id":"req-1","result":'
'{"kind":"message","parts":[{"kind":"text","text":"resubscribe-secret"}]}}'
)
mock_resp = AsyncMock()
mock_resp.is_success = True
mock_resp.aiter_lines = fake_aiter_lines
mock_resp.aclose = AsyncMock()
mock_async_client = MagicMock()
mock_async_client.build_request = MagicMock(return_value=MagicMock())
mock_async_client.send = AsyncMock(return_value=mock_resp)
mock_handler = MagicMock()
mock_handler.client = mock_async_client
mock_handler.post = AsyncMock()
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=mock_handler,
)
)
stack.enter_context(patch.object(litellm, "callbacks", [guardrail]))
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
assert response.media_type == "text/event-stream"
async for _ in response.body_iterator:
pass
assert any("resubscribe-secret" in str(r) for r in inspected), (
"tasks/resubscribe streamed content was not passed to the post-call streaming guardrail hook"
)
@pytest.mark.asyncio
async def test_task_method_failure_hook_uses_enriched_request_data():
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
mock_request = _make_request_mock("tasks/get", {"id": "task-1"})
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
async def add_proxy_data_copy(data, **kwargs):
enriched = dict(data)
enriched["proxy_server_request"] = {
"url": "http://localhost:4000",
"method": "POST",
"headers": {},
"body": {},
}
enriched.setdefault("metadata", {})
return enriched
mock_handler = MagicMock()
mock_handler.post = AsyncMock(side_effect=RuntimeError("upstream failed"))
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(
side_effect=lambda user_api_key_dict, data, call_type, skip_guardrails=False: data
)
mock_proxy_logging.post_call_failure_hook = AsyncMock(return_value=None)
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch(
"litellm.proxy.common_request_processing.add_litellm_data_to_request",
new=AsyncMock(side_effect=add_proxy_data_copy),
)
)
stack.enter_context(
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=mock_handler,
)
)
stack.enter_context(
patch(
"litellm.proxy.proxy_server.proxy_logging_obj",
mock_proxy_logging,
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
body = json.loads(response.body.decode())
assert body["error"]["code"] == -32603
failure_data = mock_proxy_logging.post_call_failure_hook.await_args.kwargs["request_data"]
assert failure_data.get("litellm_call_id")
assert failure_data.get("agent_id") == "test-agent"
@pytest.mark.asyncio
async def test_agentcore_invalid_context_id_returns_jsonrpc_invalid_params_400():
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
agent.litellm_params = {
"custom_llm_provider": "bedrock",
"model": "bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:123456789012:runtime/demo",
"api_key": "test-jwt-token",
}
mock_request = _make_request_mock(
"message/send",
{
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "Hello"}],
"messageId": "msg-1",
"contextId": "too-short",
}
},
)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(
side_effect=lambda user_api_key_dict, data, call_type, skip_guardrails=False: data
)
mock_proxy_logging.post_call_failure_hook = AsyncMock(return_value=None)
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch( # test-quality-ok: same proxy_logging_obj injection the sibling failure-hook test uses; no HTTP call is made because the request is rejected before signing
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
body = json.loads(response.body.decode())
assert response.status_code == 400
assert body["id"] == "req-1"
assert body["error"]["code"] == -32602
assert "Invalid AgentCore runtime session id" in body["error"]["message"]
assert "Internal error" not in body["error"]["message"]
mock_proxy_logging.post_call_failure_hook.assert_awaited_once()
@pytest.mark.asyncio
async def test_get_extended_agent_card_rewrites_url():
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
mock_request = _make_request_mock("GetExtendedAgentCard", {})
mock_request.base_url = "http://localhost:4000/"
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
upstream_card = {
"name": "Test Agent",
"url": "http://backend-agent:10001",
"description": "A test agent",
}
upstream_response = {"jsonrpc": "2.0", "id": "req-1", "result": upstream_card}
mock_http_response = MagicMock()
mock_http_response.json.return_value = upstream_response
mock_http_response.is_success = True
mock_http_response.raise_for_status = MagicMock()
mock_handler = MagicMock()
mock_handler.post = AsyncMock(return_value=mock_http_response)
mock_handler.client = MagicMock()
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=mock_handler,
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
body = json.loads(response.body.decode())
assert body["result"]["url"] == "http://localhost:4000/a2a/test-agent"
assert body["result"]["name"] == "Test Agent"
@pytest.mark.asyncio
async def test_get_agent_card_uses_proxy_base_url_when_set(monkeypatch):
"""Regression: discovery must expose the public proxy URL, not the internal one."""
from litellm.proxy._types import UserAPIKeyAuth
monkeypatch.setenv("PROXY_BASE_URL", "https://litellm.example.com")
agent = _make_agent_mock()
agent.agent_card_params["protocolVersion"] = "1.0"
agent.agent_card_params["supportedInterfaces"] = [
{
"url": "http://old-proxy.example.com/a2a/test-agent",
"protocolBinding": "JSONRPC",
"protocolVersion": "1.0",
}
]
mock_request = MagicMock()
mock_request.base_url = "http://litellm-internal:4000/"
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
from litellm.proxy.agent_endpoints.a2a_endpoints import get_agent_card
response = await get_agent_card(
agent_id="test-agent",
request=mock_request,
user_api_key_dict=user_api_key_dict,
)
body = json.loads(response.body.decode())
assert body["url"] == "https://litellm.example.com/a2a/test-agent"
assert body["supportedInterfaces"][0]["url"] == "https://litellm.example.com/a2a/test-agent"
@pytest.mark.asyncio
async def test_get_agent_card_normalizes_0_3_discovery_card():
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
agent.agent_card_params["protocolVersion"] = "0.3"
agent.agent_card_params["supportedInterfaces"] = [
{
"url": "http://localhost:4000/a2a/test-agent",
"protocolBinding": "JSONRPC",
"protocolVersion": "0.3",
}
]
mock_request = MagicMock()
mock_request.base_url = "http://localhost:4000/"
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
from litellm.proxy.agent_endpoints.a2a_endpoints import get_agent_card
response = await get_agent_card(
agent_id="test-agent",
request=mock_request,
user_api_key_dict=user_api_key_dict,
)
body = json.loads(response.body.decode())
assert body["protocolVersion"] == "0.3"
assert body["url"] == "http://localhost:4000/a2a/test-agent"
assert "supportedInterfaces" not in body
@pytest.mark.asyncio
async def test_get_agent_card_0_3_card_with_a2a_version_1_0_header():
"""Regression: 0.3 card normalized to 1.0 must not KeyError on debug log."""
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
agent.agent_card_params = {
"name": "Test Agent",
"description": "A test agent",
"url": "http://backend-agent:10001",
"version": "1.0.0",
"capabilities": {"streaming": True},
"skills": [{"id": "s1", "name": "skill one", "description": "d", "tags": ["t"]}],
"defaultInputModes": ["text"],
"defaultOutputModes": ["text"],
}
mock_request = MagicMock()
mock_request.base_url = "http://localhost:4000/"
mock_request.headers = {"a2a-version": "1.0"}
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
from litellm.proxy.agent_endpoints.a2a_endpoints import get_agent_card
response = await get_agent_card(
agent_id="test-agent",
request=mock_request,
user_api_key_dict=user_api_key_dict,
)
body = json.loads(response.body.decode())
assert "url" not in body
assert body["supportedInterfaces"][0]["url"] == ("http://localhost:4000/a2a/test-agent")
@pytest.mark.asyncio
async def test_get_extended_agent_card_uses_proxy_base_url_when_set(monkeypatch):
"""Regression: proxied extended cards must rewrite url to the public proxy base."""
from litellm.proxy._types import UserAPIKeyAuth
monkeypatch.setenv("PROXY_BASE_URL", "https://litellm.example.com")
agent = _make_agent_mock()
mock_request = _make_request_mock("GetExtendedAgentCard", {})
mock_request.base_url = "http://litellm-internal:4000/"
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
upstream_card = {
"name": "Test Agent",
"url": "http://backend-agent:10001",
"description": "A test agent",
}
upstream_response = {"jsonrpc": "2.0", "id": "req-1", "result": upstream_card}
mock_http_response = MagicMock()
mock_http_response.json.return_value = upstream_response
mock_http_response.is_success = True
mock_http_response.raise_for_status = MagicMock()
mock_handler = MagicMock()
mock_handler.post = AsyncMock(return_value=mock_http_response)
mock_handler.client = MagicMock()
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=mock_handler,
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
body = json.loads(response.body.decode())
assert body["result"]["url"] == "https://litellm.example.com/a2a/test-agent"
def test_build_merged_agent_card_uses_proxy_base_url_for_supported_interfaces(
monkeypatch,
):
"""Regression: agent create/update must front supportedInterfaces with the public base."""
from litellm.proxy.agent_endpoints.endpoints import _build_merged_agent_card
monkeypatch.setenv("PROXY_BASE_URL", "https://litellm.example.com")
mock_request = MagicMock()
mock_request.base_url = "http://litellm-internal:4000/"
merged = _build_merged_agent_card(
{"name": "My Agent", "url": "http://upstream:8080"},
agent_id="jenkins_agent",
http_request=mock_request,
)
assert merged["supportedInterfaces"][0]["url"] == ("https://litellm.example.com/a2a/jenkins_agent")
@pytest.mark.asyncio
async def test_unknown_method_returns_jsonrpc_error():
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
mock_request = _make_request_mock("SomeUnknownMethod", {})
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
body = json.loads(response.body.decode())
assert body["error"]["code"] == -32601
assert "SomeUnknownMethod" in body["error"]["message"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"pascal_method,expected_wire_method",
[
("GetTask", "tasks/get"),
("ListTasks", "tasks/list"),
("CancelTask", "tasks/cancel"),
("SubscribeToTask", "tasks/resubscribe"),
("CreateTaskPushNotificationConfig", "tasks/pushNotificationConfig/set"),
("GetTaskPushNotificationConfig", "tasks/pushNotificationConfig/get"),
("ListTaskPushNotificationConfigs", "tasks/pushNotificationConfig/list"),
("DeleteTaskPushNotificationConfig", "tasks/pushNotificationConfig/delete"),
("GetExtendedAgentCard", "agent/getAuthenticatedExtendedCard"),
],
)
async def test_pascal_method_names_normalize_to_wire_format(pascal_method: str, expected_wire_method: str):
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
mock_request = _make_request_mock(pascal_method, {"id": "task-1"})
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
upstream_response = {"jsonrpc": "2.0", "id": "req-1", "result": {"id": "task-1"}}
mock_http_response = MagicMock()
mock_http_response.json.return_value = upstream_response
mock_http_response.is_success = True
mock_http_response.raise_for_status = MagicMock()
async def _empty_aiter_lines():
return
yield # make it an async generator
mock_sse_resp = AsyncMock()
mock_sse_resp.is_success = True
mock_sse_resp.aiter_lines = _empty_aiter_lines
mock_sse_resp.aclose = AsyncMock()
mock_async_client = MagicMock()
mock_async_client.build_request = MagicMock(return_value=MagicMock())
mock_async_client.send = AsyncMock(return_value=mock_sse_resp)
mock_handler = MagicMock()
mock_handler.post = AsyncMock(return_value=mock_http_response)
mock_handler.client = mock_async_client
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=mock_handler,
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
if expected_wire_method == "tasks/resubscribe":
assert response.media_type == "text/event-stream"
async for _ in response.body_iterator:
pass
else:
body = json.loads(response.body.decode())
assert "error" not in body, f"Got error: {body}"
if expected_wire_method != "tasks/resubscribe":
posted = mock_handler.post.call_args
forwarded_body = posted.kwargs.get("json") or posted.args[1]
assert forwarded_body["method"] == expected_wire_method, (
f"Expected '{expected_wire_method}' forwarded for PascalCase '{pascal_method}', "
f"but got '{forwarded_body['method']}'"
)
@pytest.mark.parametrize(
"params",
[
{
"message": {
"messageId": "msg-1",
"role": "ROLE_USER",
"parts": [{"text": "hello"}],
},
"configuration": {},
},
{
"message": {
"messageId": "msg-2",
"role": "user",
"parts": [{"kind": "text", "text": "hello"}],
},
},
],
)
def test_build_message_send_params_accepts_wire_and_a2a_10(params):
from litellm.proxy.agent_endpoints.a2a_endpoints import _build_message_send_params
result = _build_message_send_params(params)
assert result.message.role.value == "user"
assert result.message.parts[0].root.text == "hello"
def test_build_message_send_params_proto_fallback_ignores_unknown_fields():
from litellm.proxy.agent_endpoints.a2a_endpoints import _build_message_send_params
result = _build_message_send_params(
{
"message": {
"messageId": "msg-1",
"role": "ROLE_USER",
"parts": [{"text": "hello"}],
},
"configuration": {},
"futureField": "ignored",
}
)
assert result.message.role.value == "user"
@pytest.mark.asyncio
async def test_handle_stream_message_rejects_invalid_params_with_32602():
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={"message": 12345},
)
assert response.media_type == "text/event-stream"
chunks = [chunk async for chunk in response.body_iterator]
body = "".join(chunk.decode() if isinstance(chunk, bytes) else chunk for chunk in chunks)
assert body.startswith("data: ")
assert body.endswith("\n\n")
payload = json.loads(body.removeprefix("data: ").strip())
assert payload["error"]["code"] == -32602
assert payload["id"] == "req-1"
@pytest.mark.asyncio
async def test_handle_stream_message_frames_events_as_sse():
"""message/stream must return text/event-stream with each JSON-RPC object
framed as ``data: <json>\\n\\n``. Regression for #35027: NDJSON framing
breaks the official a2a-sdk client, which requires SSE."""
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
events = [
{
"jsonrpc": "2.0",
"id": "req-1",
"result": {"kind": "task", "id": "t-1", "status": {"state": "working"}},
},
{
"jsonrpc": "2.0",
"id": "req-1",
"result": {"kind": "message", "parts": [{"kind": "text", "text": "pong"}]},
},
]
async def fake_stream(**kwargs):
for event in events:
yield event
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch(
"litellm.a2a_protocol.asend_message_streaming",
new=fake_stream,
)
)
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
},
)
assert response.media_type == "text/event-stream"
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert len(chunks) == len(events)
for chunk, event in zip(chunks, events):
assert chunk.startswith("data: ")
assert chunk.endswith("\n\n")
assert json.loads(chunk.removeprefix("data: ").strip()) == event
@pytest.mark.asyncio
async def test_handle_stream_message_sdk_unavailable_frames_error_as_sse():
"""When the a2a package is unavailable the -32603 error must still be
emitted as a single SSE event so the a2a-sdk client can parse it."""
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
with patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", False):
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={"message": {"role": "user", "parts": []}},
)
assert response.media_type == "text/event-stream"
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert len(chunks) == 1
assert chunks[0].startswith("data: ")
assert chunks[0].endswith("\n\n")
payload = json.loads(chunks[0].removeprefix("data: ").strip())
assert payload["error"]["code"] == -32603
assert payload["id"] == "req-1"
@pytest.mark.asyncio
async def test_handle_stream_message_proxy_hook_path_frames_events_as_sse():
"""When proxy hooks are wired the events are routed through
async_streaming_data_generator; that path must also frame each JSON-RPC
object as ``data: <json>\\n\\n`` (regression for #35027)."""
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
from litellm.proxy.utils import ProxyLogging
events = [
{"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task", "id": "t-1"}},
{
"jsonrpc": "2.0",
"id": "req-1",
"result": {"kind": "message", "parts": [{"kind": "text", "text": "pong"}]},
},
]
async def fake_stream(**kwargs):
for event in events:
yield event
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
request_data={"model": "a2a/test"},
proxy_logging_obj=proxy_logging_obj,
)
assert response.media_type == "text/event-stream"
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert len(chunks) == len(events)
for chunk, event in zip(chunks, events):
assert chunk.startswith("data: ")
assert chunk.endswith("\n\n")
assert json.loads(chunk.removeprefix("data: ").strip()) == event
@pytest.mark.asyncio
async def test_handle_stream_message_frames_preserialized_jsonrpc_error_once():
"""A stream chunk that is already a serialized JSON-RPC object (what a
guardrail may yield when it terminates an A2A stream mid-flight) must be
framed as one SSE event carrying that object, not JSON-encoded a second time
into a bare string."""
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
error_event = {
"jsonrpc": "2.0",
"id": "req-1",
"error": {"code": -32603, "message": "blocked by guardrail", "data": {}},
}
async def fake_stream(**kwargs):
yield json.dumps(error_event) + "\n"
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
},
)
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert len(chunks) == 1
payload = json.loads(chunks[0].removeprefix("data: ").strip())
assert payload == error_event
@pytest.mark.asyncio
async def test_handle_stream_message_proxy_hook_path_frames_errors_as_sse():
"""A failure while the hooked generator is streaming must reach the client as
a ``data:``-framed JSON-RPC error, not as a bare NDJSON line."""
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
from litellm.proxy.utils import ProxyLogging
async def fake_stream(**kwargs):
yield {"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task", "id": "t-1"}}
raise ValueError("upstream died")
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
request_data={"model": "a2a/test"},
proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()),
)
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert len(chunks) == 2
assert chunks[-1].startswith("data: ")
error_payload = json.loads(chunks[-1].removeprefix("data: ").strip())
assert error_payload["id"] == "req-1"
assert error_payload["error"]["code"] == -32603
assert "upstream died" in error_payload["error"]["message"]
@pytest.mark.asyncio
async def test_handle_stream_message_frames_upstream_call_failure_as_sse_error():
"""A failure raised before any event is streamed (with proxy hooks wired) is
still delivered as a ``data:``-framed JSON-RPC error."""
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
from litellm.proxy.utils import ProxyLogging
def fake_stream(**kwargs):
raise ValueError("could not reach agent")
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
request_data={"model": "a2a/test"},
proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()),
)
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert len(chunks) == 1
error_payload = json.loads(chunks[0].removeprefix("data: ").strip())
assert error_payload["id"] == "req-1"
assert error_payload["error"]["code"] == -32603
assert "could not reach agent" in error_payload["error"]["message"]
@pytest.mark.asyncio
async def test_handle_stream_message_forwards_unparseable_chunk_as_sse_event():
"""A chunk that is not JSON at all still leaves as one well-formed SSE event
instead of raising and killing the stream."""
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
async def fake_stream(**kwargs):
yield "not json at all"
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
},
)
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert chunks == ['data: "not json at all"\n\n']
@pytest.mark.asyncio
async def test_handle_stream_message_frames_mid_stream_failure_as_sse_error():
"""An upstream failure after the response started is reported as a
``data:``-framed JSON-RPC error object, so an SSE client sees the failure."""
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
async def fake_stream(**kwargs):
yield {"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task", "id": "t-1"}}
raise RuntimeError("upstream died")
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
response = await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
},
)
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert len(chunks) == 2
error_payload = json.loads(chunks[-1].removeprefix("data: ").strip())
assert error_payload["id"] == "req-1"
assert error_payload["error"]["code"] == -32603
assert "upstream died" in error_payload["error"]["message"]
@pytest.mark.asyncio
async def test_send_message_pascal_case_routes_to_asend_message():
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
params = {
"message": {
"messageId": "msg-123",
"role": "ROLE_USER",
"parts": [{"text": "Hello"}],
},
"configuration": {},
}
mock_request = _make_request_mock("SendMessage", params)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
captured = {}
async def capture_asend_message(request, **kwargs):
captured["method"] = request.method
captured["role"] = request.params.message.role.value
response = MagicMock()
response.model_dump.return_value = {
"jsonrpc": "2.0",
"id": request.id,
"result": {
"contextId": "ctx-1",
"kind": "message",
"messageId": "msg-123",
"parts": [{"kind": "text", "text": "Hello"}],
"role": "agent",
},
}
return response
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(
patch(
"litellm.a2a_protocol.asend_message",
new=AsyncMock(side_effect=capture_asend_message),
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
body = json.loads(response.body.decode())
assert "error" not in body, f"Got error: {body}"
assert captured["method"] == "message/send"
assert captured["role"] == "user"
assert "message" in body["result"]
assert body["result"]["message"]["role"] == "ROLE_AGENT"
def test_normalize_response_wraps_flat_message_result_for_1_0():
from litellm.proxy.a2a.version_convert import normalize_jsonrpc_response
wire_response = {
"jsonrpc": "2.0",
"id": "req-1",
"result": {
"contextId": "ctx-1",
"kind": "message",
"messageId": "msg-1",
"parts": [{"kind": "text", "text": "hello"}],
"role": "agent",
"taskId": "task-1",
},
}
formatted = normalize_jsonrpc_response(wire_response, "1.0", method="message/send")
assert "message" in formatted["result"]
assert formatted["result"]["message"]["role"] == "ROLE_AGENT"
assert formatted["result"]["message"]["parts"] == [{"text": "hello"}]
assert "contextId" not in formatted["result"]
def test_normalize_response_keeps_wire_format_for_0_3():
from litellm.proxy.a2a.version_convert import normalize_jsonrpc_response
wire_response = {
"jsonrpc": "2.0",
"id": "req-1",
"result": {
"contextId": "ctx-1",
"kind": "message",
"messageId": "msg-1",
"parts": [{"kind": "text", "text": "hello"}],
"role": "agent",
},
}
assert normalize_jsonrpc_response(wire_response, "0.3", method="message/send") is wire_response
@pytest.mark.asyncio
async def test_task_method_upstream_jsonrpc_error_on_http_4xx_is_relayed():
"""When upstream returns HTTP 4xx with a JSON-RPC error body, the error body
must be relayed to the client unchanged, not replaced with a generic string."""
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
mock_request = _make_request_mock("tasks/get", {"id": "nonexistent"})
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
upstream_error = {
"jsonrpc": "2.0",
"id": "req-1",
"error": {"code": -32001, "message": "Task not found"},
}
mock_http_response = MagicMock()
mock_http_response.json.return_value = upstream_error
mock_http_response.is_success = False
mock_http_response.raise_for_status = MagicMock(side_effect=Exception("404 Not Found"))
mock_handler = MagicMock()
mock_handler.post = AsyncMock(return_value=mock_http_response)
mock_handler.client = MagicMock()
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=mock_handler,
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
body = json.loads(response.body.decode())
assert body["error"]["code"] == -32001
assert body["error"]["message"] == "Task not found"
@pytest.mark.asyncio
async def test_subscribe_to_task_upstream_error_yields_jsonrpc_error_event():
"""When upstream returns a non-2xx response for tasks/resubscribe, the SSE
stream must yield a JSON-RPC error event instead of silently breaking."""
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
mock_request = _make_request_mock("tasks/resubscribe", {"id": "task-1"})
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
mock_resp = AsyncMock()
mock_resp.is_success = False
mock_resp.status_code = 404
mock_resp.reason_phrase = "Not Found"
mock_resp.aread = AsyncMock(return_value=b'{"jsonrpc":"2.0","error":{"code":-32001,"message":"Task not found"}}')
mock_resp.aclose = AsyncMock()
mock_async_client = MagicMock()
mock_async_client.build_request = MagicMock(return_value=MagicMock())
mock_async_client.send = AsyncMock(return_value=mock_resp)
mock_handler = MagicMock()
mock_handler.client = mock_async_client
mock_handler.post = AsyncMock()
chunks = []
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=mock_handler,
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
assert response.media_type == "text/event-stream"
async for chunk in response.body_iterator:
chunks.append(chunk)
full = "".join(chunks)
body = json.loads(full.removeprefix("data: ").strip())
assert body["id"] == "req-1"
assert body["error"]["code"] == -32001
assert body["error"]["message"] == "Task not found"
@pytest.mark.asyncio
async def test_forward_jsonrpc_sse_fallback_error_uses_jsonrpc_error_code():
mock_resp = AsyncMock()
mock_resp.is_success = False
mock_resp.status_code = 503
mock_resp.reason_phrase = "Service Unavailable"
mock_resp.aread = AsyncMock(return_value=b"upstream unavailable")
mock_resp.aclose = AsyncMock()
mock_async_client = MagicMock()
mock_async_client.build_request = MagicMock(return_value=MagicMock())
mock_async_client.send = AsyncMock(return_value=mock_resp)
mock_handler = MagicMock()
mock_handler.client = mock_async_client
with patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=mock_handler,
):
from litellm.proxy.agent_endpoints.a2a_endpoints import _forward_jsonrpc_sse
response = await _forward_jsonrpc_sse(
agent_url="http://backend-agent:10001",
body={"jsonrpc": "2.0", "id": "req-1", "method": "tasks/resubscribe"},
request_id="req-1",
)
chunks = []
async for chunk in response.body_iterator:
chunks.append(chunk)
body = json.loads("".join(chunks).removeprefix("data: ").strip())
assert body["error"]["code"] == -32603
assert body["error"]["message"] == "Service Unavailable"
@pytest.mark.asyncio
async def test_task_methods_forward_caller_identity_headers():
"""Task operations must forward X-LiteLLM-User-Id and X-LiteLLM-Team-Id so the
upstream agent can scope resources to the authenticated caller."""
from litellm.proxy._types import UserAPIKeyAuth
upstream_response = {
"jsonrpc": "2.0",
"id": "req-1",
"result": {"id": "task-1", "status": {"state": "completed"}},
}
agent = _make_agent_mock()
mock_request = _make_request_mock("tasks/get", {"id": "task-1"})
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="user-abc", team_id="team-xyz")
mock_http_response = MagicMock()
mock_http_response.json.return_value = upstream_response
mock_http_response.is_success = True
mock_handler = MagicMock()
mock_handler.post = AsyncMock(return_value=mock_http_response)
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=mock_handler,
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
posted_headers = mock_handler.post.call_args.kwargs.get("headers") or {}
assert posted_headers.get("X-LiteLLM-User-Id") == "user-abc"
assert posted_headers.get("X-LiteLLM-Team-Id") == "team-xyz"
@pytest.mark.asyncio
@pytest.mark.parametrize("method", ["tasks/get", "tasks/resubscribe"])
async def test_task_methods_forward_trace_header(method: str):
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
mock_request = _make_request_mock(method, {"id": "task-1"})
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
async def add_proxy_data_with_trace(data, **kwargs):
data = await _add_proxy_data(data, **kwargs)
data["litellm_trace_id"] = "trace-123"
return data
upstream_response = {
"jsonrpc": "2.0",
"id": "req-1",
"result": {"id": "task-1", "status": {"state": "completed"}},
}
mock_http_response = MagicMock()
mock_http_response.json.return_value = upstream_response
mock_http_response.is_success = True
async def fake_aiter_lines():
yield 'data: {"jsonrpc":"2.0","id":"req-1","result":{"taskId":"task-1"}}'
mock_resp = AsyncMock()
mock_resp.is_success = True
mock_resp.aiter_lines = fake_aiter_lines
mock_resp.aclose = AsyncMock()
mock_async_client = MagicMock()
mock_async_client.build_request = MagicMock(return_value=MagicMock())
mock_async_client.send = AsyncMock(return_value=mock_resp)
mock_handler = MagicMock()
mock_handler.post = AsyncMock(return_value=mock_http_response)
mock_handler.client = mock_async_client
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch(
"litellm.proxy.common_request_processing.add_litellm_data_to_request",
new=AsyncMock(side_effect=add_proxy_data_with_trace),
)
)
stack.enter_context(
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=mock_handler,
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
response = await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
if method == "tasks/resubscribe":
async for _ in response.body_iterator:
pass
if method == "tasks/resubscribe":
forwarded_headers = mock_async_client.build_request.call_args.kwargs["headers"]
else:
forwarded_headers = mock_handler.post.call_args.kwargs["headers"]
assert forwarded_headers.get("X-LiteLLM-Trace-Id") == "trace-123"
@pytest.mark.asyncio
async def test_push_notification_config_set_rejects_http_url():
"""tasks/pushNotificationConfig/set must reject non-HTTPS callback URLs to prevent SSRF."""
from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
mock_request = _make_request_mock(
"tasks/pushNotificationConfig/set",
{"taskId": "task-1", "url": "http://internal-webhook.example.com/hook"},
)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
with pytest.raises(HTTPException) as exc_info:
await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
assert exc_info.value.status_code == 400
assert "HTTPS" in exc_info.value.detail
@pytest.mark.asyncio
async def test_push_notification_config_set_rejects_private_ip():
"""tasks/pushNotificationConfig/set must reject callback URLs pointing to private IP ranges."""
from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
mock_request = _make_request_mock(
"tasks/pushNotificationConfig/set",
{"taskId": "task-1", "url": "https://192.168.1.100/hook"},
)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
with pytest.raises(HTTPException) as exc_info:
await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
assert exc_info.value.status_code == 400
assert "blocked address" in exc_info.value.detail.lower()
@pytest.mark.asyncio
async def test_push_notification_config_set_validates_nested_url_when_top_level_present():
"""A safe top-level params.url must not let a private pushNotificationConfig.url bypass SSRF checks.
Both URL-bearing fields are forwarded to the agent, so both must be validated independently.
"""
from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
mock_request = _make_request_mock(
"tasks/pushNotificationConfig/set",
{
"taskId": "task-1",
"url": "https://1.1.1.1/hook",
"pushNotificationConfig": {"url": "https://192.168.1.100/hook"},
},
)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
with pytest.raises(HTTPException) as exc_info:
await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
assert exc_info.value.status_code == 400
assert "blocked address" in exc_info.value.detail.lower()
def test_push_notification_config_set_rejects_private_dns_resolution():
from fastapi import HTTPException
from litellm.proxy.agent_endpoints.a2a_endpoints import (
_validate_push_notification_url,
)
with patch(
"litellm.litellm_core_utils.url_utils.socket.getaddrinfo",
return_value=[
(
socket.AF_INET,
socket.SOCK_STREAM,
socket.IPPROTO_TCP,
"",
("10.0.0.5", 443),
)
],
):
with pytest.raises(HTTPException) as exc_info:
_validate_push_notification_url("https://webhook.example.com/hook")
assert exc_info.value.status_code == 400
assert "blocked address" in exc_info.value.detail.lower()
@pytest.mark.asyncio
async def test_push_notification_config_set_rejects_null_push_config():
from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth
agent = _make_agent_mock()
mock_request = _make_request_mock(
"tasks/pushNotificationConfig/set",
{"taskId": "task-1", "pushNotificationConfig": None},
)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1", team_id="t1")
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
with pytest.raises(HTTPException) as exc_info:
await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
assert exc_info.value.status_code == 400
assert "pushNotificationConfig must be an object" in exc_info.value.detail
@pytest.mark.asyncio
async def test_caller_identity_headers_cannot_be_spoofed_via_forwarded_headers():
"""A client must not be able to override X-LiteLLM-User-Id / X-LiteLLM-Team-Id
by including x-a2a-<agent>-x-litellm-user-id in their request headers.
The authenticated identity must always win."""
from litellm.proxy._types import UserAPIKeyAuth
upstream_response = {
"jsonrpc": "2.0",
"id": "req-1",
"result": {"id": "task-1", "status": {"state": "completed"}},
}
agent = _make_agent_mock()
mock_request = _make_request_mock("tasks/get", {"id": "task-1"})
mock_request.headers = {
"x-a2a-test-agent-x-litellm-user-id": "attacker-user",
"x-a2a-test-agent-x-litellm-team-id": "attacker-team",
}
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="real-user", team_id="real-team")
mock_http_response = MagicMock()
mock_http_response.json.return_value = upstream_response
mock_http_response.is_success = True
mock_handler = MagicMock()
mock_handler.post = AsyncMock(return_value=mock_http_response)
with ExitStack() as stack:
for p in _base_patches(agent):
stack.enter_context(p)
stack.enter_context(
patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=mock_handler,
)
)
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
await invoke_agent_a2a(
agent_id="test-agent",
request=mock_request,
fastapi_response=MagicMock(),
user_api_key_dict=user_api_key_dict,
)
posted_headers = mock_handler.post.call_args.kwargs.get("headers") or {}
assert posted_headers.get("X-LiteLLM-User-Id") == "real-user", (
"authenticated user id must not be overridden by forwarded client headers"
)
assert posted_headers.get("X-LiteLLM-Team-Id") == "real-team", (
"authenticated team id must not be overridden by forwarded client headers"
)
def _agent(protocol_version):
agent = MagicMock()
agent.agent_card_params = {"protocolVersion": protocol_version} if protocol_version is not None else {}
return agent
def _request_with_a2a_header(value):
request = MagicMock()
request.headers = {"a2a-version": value} if value is not None else {}
return request
def test_served_version_config_governs_over_header():
from litellm.proxy.agent_endpoints.a2a_endpoints import _served_version
# A 0.3-configured agent serves 0.3 even when the client asks for 1.0.
agent = _agent("0.3")
request = _request_with_a2a_header("1.0")
assert _served_version(agent, request) == "0.3"
# A 1.0-configured agent serves 1.0 even when the client asks for 0.3.
assert _served_version(_agent("1.0"), _request_with_a2a_header("0.3")) == "1.0"
def test_served_version_falls_back_to_header_when_unconfigured():
from litellm.proxy.agent_endpoints.a2a_endpoints import _served_version
assert _served_version(_agent(None), _request_with_a2a_header("1.0")) == "1.0"
assert _served_version(_agent(None), _request_with_a2a_header(None)) == "0.3"
def _sse_agent_handler(lines):
mock_resp = AsyncMock()
mock_resp.is_success = True
mock_resp.aiter_lines = lines
mock_resp.aclose = AsyncMock()
mock_async_client = MagicMock()
mock_async_client.build_request = MagicMock(return_value=MagicMock())
mock_async_client.send = AsyncMock(return_value=mock_resp)
mock_handler = MagicMock()
mock_handler.client = mock_async_client
return mock_handler
async def _resubscribe_response():
from litellm.proxy.agent_endpoints.a2a_endpoints import _forward_jsonrpc_sse
return await _forward_jsonrpc_sse(
agent_url="http://backend-agent:10001",
body={"jsonrpc": "2.0", "id": "req-1", "method": "tasks/resubscribe"},
request_id="req-1",
)
@pytest.mark.asyncio
async def test_forward_jsonrpc_sse_pings_while_the_upstream_agent_is_still_silent(
monkeypatch,
):
"""Regression for LIT-5737. The upstream agent is only contacted once the body
iterator is first pulled, so a slow first event leaves the response body idle
for its whole time-to-first-token and an idle-timeout hop drops a healthy
connection."""
import asyncio
import litellm
monkeypatch.setattr(litellm, "sse_keepalive_ping_interval_seconds", 0.05)
async def _slow_lines():
await asyncio.sleep(0.3)
yield 'data: {"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task"}}'
with patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=_sse_agent_handler(_slow_lines),
):
response = await _resubscribe_response()
assert response.headers["x-accel-buffering"] == "no"
chunks = [chunk async for chunk in response.body_iterator]
# A comment, not a frame: an A2A client parsing JSON-RPC events has to be able
# to discard the filler without understanding it.
assert chunks[0] == ": ping\n\n"
assert chunks.count(": ping\n\n") >= 3
assert json.loads(chunks[-1].removeprefix("data: "))["result"]["kind"] == "task"
@pytest.mark.asyncio
async def test_forward_jsonrpc_sse_is_untouched_while_keepalives_are_unconfigured(
monkeypatch,
):
"""Off until an operator sets an interval, so the default stream is unchanged."""
import litellm
monkeypatch.setattr(litellm, "sse_keepalive_ping_interval_seconds", None)
async def _lines():
yield 'data: {"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task"}}'
with patch(
"litellm.llms.custom_httpx.http_handler.get_async_httpx_client",
return_value=_sse_agent_handler(_lines),
):
response = await _resubscribe_response()
assert "x-accel-buffering" not in response.headers
chunks = [chunk async for chunk in response.body_iterator]
assert not any(chunk.startswith(":") for chunk in chunks)
assert json.loads(chunks[-1].removeprefix("data: "))["result"]["kind"] == "task"
async def _stream_message_response():
from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message
return await _handle_stream_message(
api_base="http://upstream.local",
request_id="req-1",
params={
"message": {
"role": "user",
"parts": [{"kind": "text", "text": "hi"}],
"messageId": "msg-1",
}
},
)
@pytest.mark.asyncio
async def test_handle_stream_message_pings_while_the_upstream_agent_is_still_silent(
monkeypatch,
):
"""message/stream is SSE like tasks/resubscribe, so a slow first event must be
held open by the same keepalives rather than sitting idle for the whole
time-to-first-token."""
import asyncio
import litellm
monkeypatch.setattr(litellm, "sse_keepalive_ping_interval_seconds", 0.05)
async def fake_stream(**kwargs):
await asyncio.sleep(0.3)
yield {"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task", "id": "t-1"}}
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
response = await _stream_message_response()
assert response.headers["x-accel-buffering"] == "no"
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert chunks[0] == ": ping\n\n"
assert chunks.count(": ping\n\n") >= 3
assert json.loads(chunks[-1].removeprefix("data: "))["result"]["kind"] == "task"
@pytest.mark.asyncio
async def test_handle_stream_message_is_untouched_while_keepalives_are_unconfigured(
monkeypatch,
):
"""Off until an operator sets an interval, so the default stream is unchanged."""
import litellm
monkeypatch.setattr(litellm, "sse_keepalive_ping_interval_seconds", None)
async def fake_stream(**kwargs):
yield {"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task", "id": "t-1"}}
with ExitStack() as stack:
stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True))
stack.enter_context(patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream))
response = await _stream_message_response()
assert "x-accel-buffering" not in response.headers
chunks = [chunk.decode() if isinstance(chunk, bytes) else chunk async for chunk in response.body_iterator]
assert not any(chunk.startswith(":") for chunk in chunks)
assert json.loads(chunks[-1].removeprefix("data: "))["result"]["kind"] == "task"
def test_forwarding_headers_minted_bearer_replaces_a_forwarded_authorization_of_any_case():
"""A client header the admin chose to forward keeps the casing the config named it with, so a forwarded
`authorization` must not travel next to the minted `Authorization` as a second header line."""
from litellm.proxy.agent_endpoints.a2a_endpoints import _forwarding_headers
merged = _forwarding_headers(
caller_identity={},
request_data={},
agent_extra_headers={"authorization": "Bearer client-token", "X-Custom": "kept"},
backend_auth_header={"Authorization": "Bearer minted-token"},
)
assert merged == {"X-Custom": "kept", "Authorization": "Bearer minted-token"}