mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
* 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>
2691 lines
97 KiB
Python
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"}
|