mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
365 lines
14 KiB
Python
365 lines
14 KiB
Python
"""
|
|
Test A2A completion bridge streaming transformation to proper A2A format.
|
|
|
|
Tests that the completion bridge emits proper A2A streaming events:
|
|
1. Task event (kind: "task") - Initial task with status "submitted"
|
|
2. Status update (kind: "status-update") - Status "working"
|
|
3. Artifact update (kind: "artifact-update") - Content delivery
|
|
4. Status update (kind: "status-update") - Final "completed" status
|
|
"""
|
|
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
|
|
|
import pytest
|
|
|
|
|
|
class TestA2AStreamingTransformation:
|
|
"""Test the A2A streaming transformation creates proper events."""
|
|
|
|
def test_a2a_metadata_forwarded_to_completion_params(self):
|
|
from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
|
|
A2ACompletionBridgeTransformation,
|
|
)
|
|
|
|
message = {
|
|
"role": "user",
|
|
"parts": [{"text": "Reply to ticket #4823"}],
|
|
"metadata": {"skillId": "draft_reply"},
|
|
}
|
|
openai_messages = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
|
|
# Metadata is forwarded on the run payload only, not duplicated on messages.
|
|
assert "metadata" not in openai_messages[0]
|
|
|
|
completion_params: dict = {
|
|
"model": "langgraph/agent",
|
|
"messages": openai_messages,
|
|
}
|
|
A2ACompletionBridgeTransformation.apply_forward_metadata_to_completion_params(
|
|
completion_params=completion_params,
|
|
a2a_message=message,
|
|
params={"metadata": {"trace": "abc"}},
|
|
)
|
|
assert completion_params["extra_body"]["metadata"] == {
|
|
"trace": "abc",
|
|
"skillId": "draft_reply",
|
|
}
|
|
|
|
def test_configured_metadata_wins_over_forwarded_a2a_metadata(self):
|
|
from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
|
|
A2ACompletionBridgeTransformation,
|
|
)
|
|
|
|
# Agent-owner-configured run metadata in ``extra_body``.
|
|
completion_params: dict = {
|
|
"model": "langgraph/agent",
|
|
"messages": [],
|
|
"extra_body": {
|
|
"metadata": {"owner_tag": "prod", "trace": "server-set"},
|
|
"other": "keep",
|
|
},
|
|
}
|
|
# Client tries to overwrite ``trace`` and inject a new key.
|
|
message = {
|
|
"role": "user",
|
|
"parts": [{"text": "hi"}],
|
|
"metadata": {"trace": "client-spoof", "skillId": "draft_reply"},
|
|
}
|
|
A2ACompletionBridgeTransformation.apply_forward_metadata_to_completion_params(
|
|
completion_params=completion_params,
|
|
a2a_message=message,
|
|
params={"metadata": {"trace": "client-spoof-2"}},
|
|
)
|
|
assert completion_params["extra_body"]["other"] == "keep"
|
|
assert completion_params["extra_body"]["metadata"] == {
|
|
"owner_tag": "prod",
|
|
"trace": "server-set",
|
|
"skillId": "draft_reply",
|
|
}
|
|
|
|
def test_langgraph_transform_preserves_message_metadata(self):
|
|
from litellm.llms.langgraph.chat.transformation import LangGraphConfig
|
|
|
|
config = LangGraphConfig()
|
|
request = config.transform_request(
|
|
model="langgraph/agent",
|
|
messages=[
|
|
{
|
|
"role": "user",
|
|
"content": "Reply to ticket #4823",
|
|
"metadata": {"skillId": "draft_reply"},
|
|
}
|
|
],
|
|
optional_params={},
|
|
litellm_params={"stream": False},
|
|
headers={},
|
|
)
|
|
assert request["input"]["messages"][-1]["metadata"] == {
|
|
"skillId": "draft_reply",
|
|
}
|
|
|
|
def test_create_task_event(self):
|
|
"""Test that create_task_event produces proper A2A task event structure."""
|
|
from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
|
|
A2ACompletionBridgeTransformation,
|
|
A2AStreamingContext,
|
|
)
|
|
|
|
input_message = {
|
|
"role": "user",
|
|
"parts": [{"kind": "text", "text": "Hello"}],
|
|
"messageId": "msg-123",
|
|
}
|
|
ctx = A2AStreamingContext(request_id="req-456", input_message=input_message)
|
|
|
|
event = A2ACompletionBridgeTransformation.create_task_event(ctx)
|
|
|
|
# Validate structure
|
|
assert event["jsonrpc"] == "2.0"
|
|
assert event["id"] == "req-456"
|
|
assert event["result"]["kind"] == "task"
|
|
assert event["result"]["status"]["state"] == "submitted"
|
|
assert "contextId" in event["result"]
|
|
assert "id" in event["result"] # task id
|
|
assert "history" in event["result"]
|
|
assert len(event["result"]["history"]) == 1
|
|
assert event["result"]["history"][0]["role"] == "user"
|
|
|
|
def test_create_status_update_working(self):
|
|
"""Test that create_status_update_event produces proper working status."""
|
|
from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
|
|
A2ACompletionBridgeTransformation,
|
|
A2AStreamingContext,
|
|
)
|
|
|
|
ctx = A2AStreamingContext(
|
|
request_id="req-456",
|
|
input_message={"role": "user", "parts": []},
|
|
)
|
|
|
|
event = A2ACompletionBridgeTransformation.create_status_update_event(
|
|
ctx=ctx,
|
|
state="working",
|
|
final=False,
|
|
message_text="Processing...",
|
|
)
|
|
|
|
assert event["result"]["kind"] == "status-update"
|
|
assert event["result"]["status"]["state"] == "working"
|
|
assert event["result"]["final"] is False
|
|
assert "taskId" in event["result"]
|
|
assert "contextId" in event["result"]
|
|
assert "timestamp" in event["result"]["status"]
|
|
|
|
def test_create_artifact_update(self):
|
|
"""Test that create_artifact_update_event produces proper artifact event."""
|
|
from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
|
|
A2ACompletionBridgeTransformation,
|
|
A2AStreamingContext,
|
|
)
|
|
|
|
ctx = A2AStreamingContext(
|
|
request_id="req-456",
|
|
input_message={"role": "user", "parts": []},
|
|
)
|
|
|
|
event = A2ACompletionBridgeTransformation.create_artifact_update_event(
|
|
ctx=ctx,
|
|
text="Hello, I am an AI assistant.",
|
|
)
|
|
|
|
assert event["result"]["kind"] == "artifact-update"
|
|
assert "artifact" in event["result"]
|
|
assert "artifactId" in event["result"]["artifact"]
|
|
assert event["result"]["artifact"]["name"] == "response"
|
|
assert event["result"]["artifact"]["parts"][0]["kind"] == "text"
|
|
assert event["result"]["artifact"]["parts"][0]["text"] == "Hello, I am an AI assistant."
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_handle_streaming_emits_proper_events():
|
|
"""Test that handle_streaming emits events in correct order with proper structure."""
|
|
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
|
|
A2ACompletionBridgeHandler,
|
|
)
|
|
|
|
# Mock litellm.acompletion to return a streaming response
|
|
mock_chunk1 = MagicMock()
|
|
mock_chunk1.choices = [MagicMock()]
|
|
mock_chunk1.choices[0].delta = MagicMock()
|
|
mock_chunk1.choices[0].delta.content = "Hello"
|
|
|
|
mock_chunk2 = MagicMock()
|
|
mock_chunk2.choices = [MagicMock()]
|
|
mock_chunk2.choices[0].delta = MagicMock()
|
|
mock_chunk2.choices[0].delta.content = " world"
|
|
|
|
async def mock_streaming_response():
|
|
yield mock_chunk1
|
|
yield mock_chunk2
|
|
|
|
with patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion:
|
|
mock_acompletion.return_value = mock_streaming_response()
|
|
|
|
params = {
|
|
"message": {
|
|
"role": "user",
|
|
"parts": [{"kind": "text", "text": "Hi"}],
|
|
"messageId": "msg-123",
|
|
}
|
|
}
|
|
|
|
events = []
|
|
async for event in A2ACompletionBridgeHandler.handle_streaming(
|
|
request_id="req-456",
|
|
params=params,
|
|
litellm_params={"custom_llm_provider": "langgraph", "model": "agent"},
|
|
api_base="http://localhost:2024",
|
|
):
|
|
events.append(event)
|
|
|
|
# Should have 4 events: task, working, artifact, completed
|
|
assert len(events) == 4
|
|
|
|
# Event 1: task submitted
|
|
assert events[0]["result"]["kind"] == "task"
|
|
assert events[0]["result"]["status"]["state"] == "submitted"
|
|
|
|
# Event 2: status working
|
|
assert events[1]["result"]["kind"] == "status-update"
|
|
assert events[1]["result"]["status"]["state"] == "working"
|
|
assert events[1]["result"]["final"] is False
|
|
|
|
# Event 3: artifact update with accumulated content
|
|
assert events[2]["result"]["kind"] == "artifact-update"
|
|
assert events[2]["result"]["artifact"]["parts"][0]["text"] == "Hello world"
|
|
|
|
# Event 4: status completed
|
|
assert events[3]["result"]["kind"] == "status-update"
|
|
assert events[3]["result"]["status"]["state"] == "completed"
|
|
assert events[3]["result"]["final"] is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_handle_streaming_forwards_api_key():
|
|
"""Test that handle_streaming forwards api_key from litellm_params to acompletion."""
|
|
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
|
|
A2ACompletionBridgeHandler,
|
|
)
|
|
|
|
mock_chunk = MagicMock()
|
|
mock_chunk.choices = [MagicMock()]
|
|
mock_chunk.choices[0].delta = MagicMock()
|
|
mock_chunk.choices[0].delta.content = "Response"
|
|
|
|
async def mock_streaming_response():
|
|
yield mock_chunk
|
|
|
|
with patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion:
|
|
mock_acompletion.return_value = mock_streaming_response()
|
|
|
|
params = {
|
|
"message": {
|
|
"role": "user",
|
|
"parts": [{"kind": "text", "text": "Hi"}],
|
|
"messageId": "msg-123",
|
|
}
|
|
}
|
|
|
|
events = []
|
|
async for event in A2ACompletionBridgeHandler.handle_streaming(
|
|
request_id="req-456",
|
|
params=params,
|
|
litellm_params={
|
|
"custom_llm_provider": "azure_ai",
|
|
"model": "agents/asst_123",
|
|
"api_key": "test-api-key-12345",
|
|
},
|
|
api_base="https://example.azure.com/",
|
|
):
|
|
events.append(event)
|
|
|
|
# Verify acompletion was called with api_key
|
|
mock_acompletion.assert_called_once()
|
|
call_kwargs = mock_acompletion.call_args.kwargs
|
|
assert call_kwargs["api_key"] == "test-api-key-12345"
|
|
assert call_kwargs["api_base"] == "https://example.azure.com/"
|
|
assert call_kwargs["model"] == "azure_ai/agents/asst_123"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_handle_non_streaming_forwards_api_key():
|
|
"""Test that handle_non_streaming forwards api_key from litellm_params to acompletion."""
|
|
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
|
|
A2ACompletionBridgeHandler,
|
|
)
|
|
|
|
mock_response = MagicMock()
|
|
mock_response.choices = [MagicMock()]
|
|
mock_response.choices[0].message = MagicMock()
|
|
mock_response.choices[0].message.content = "Hello!"
|
|
mock_response.id = "resp-123"
|
|
|
|
with patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion:
|
|
mock_acompletion.return_value = mock_response
|
|
|
|
params = {
|
|
"message": {
|
|
"role": "user",
|
|
"parts": [{"kind": "text", "text": "Hi"}],
|
|
"messageId": "msg-123",
|
|
}
|
|
}
|
|
|
|
await A2ACompletionBridgeHandler.handle_non_streaming(
|
|
request_id="req-456",
|
|
params=params,
|
|
litellm_params={
|
|
"custom_llm_provider": "azure_ai",
|
|
"model": "agents/asst_456",
|
|
"api_key": "my-secret-api-key",
|
|
},
|
|
api_base="https://my-azure.com/",
|
|
)
|
|
|
|
# Verify acompletion was called with api_key
|
|
mock_acompletion.assert_called_once()
|
|
call_kwargs = mock_acompletion.call_args.kwargs
|
|
assert call_kwargs["api_key"] == "my-secret-api-key"
|
|
assert call_kwargs["api_base"] == "https://my-azure.com/"
|
|
assert call_kwargs["model"] == "azure_ai/agents/asst_456"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_handle_streaming_keeps_agent_card_path_out_of_the_completion_call():
|
|
"""agent_card_path describes where an A2A agent serves its card; a completion-bridge agent carrying
|
|
it must not pass it to litellm.acompletion, where an unknown kwarg breaks the provider call."""
|
|
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
|
|
A2ACompletionBridgeHandler,
|
|
)
|
|
|
|
async def mock_streaming_response():
|
|
chunk = MagicMock()
|
|
chunk.choices = [MagicMock()]
|
|
chunk.choices[0].delta = MagicMock()
|
|
chunk.choices[0].delta.content = "Hello"
|
|
yield chunk
|
|
|
|
with patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion:
|
|
mock_acompletion.return_value = mock_streaming_response()
|
|
|
|
events = [
|
|
event
|
|
async for event in A2ACompletionBridgeHandler.handle_streaming(
|
|
request_id="req-card-path",
|
|
params={"message": {"role": "user", "parts": [{"kind": "text", "text": "Hi"}], "messageId": "m1"}},
|
|
litellm_params={
|
|
"custom_llm_provider": "langgraph",
|
|
"model": "agent",
|
|
"agent_card_path": "agentCard/v1.0",
|
|
},
|
|
api_base="http://localhost:2024",
|
|
)
|
|
]
|
|
|
|
assert len(events) == 4
|
|
assert "agent_card_path" not in mock_acompletion.call_args.kwargs
|