mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
test(a2a): cover streamed 0.3 replies through the real client
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
633813e6ee
commit
ba709d3de5
1 changed files with 125 additions and 3 deletions
|
|
@ -222,6 +222,77 @@ _LANGGRAPH_TASK_REPLY = {
|
|||
}
|
||||
|
||||
|
||||
_LANGGRAPH_STREAM_EVENTS = [
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": "reply",
|
||||
"result": {
|
||||
"kind": "task",
|
||||
"id": "run-1:task-1",
|
||||
"contextId": "thread-1",
|
||||
"history": [
|
||||
{
|
||||
"kind": "message",
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "hi"}],
|
||||
"messageId": "m-user",
|
||||
"taskId": "run-1:task-1",
|
||||
"contextId": "thread-1",
|
||||
}
|
||||
],
|
||||
"status": {"state": "submitted"},
|
||||
},
|
||||
},
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": "reply",
|
||||
"result": {
|
||||
"kind": "status-update",
|
||||
"taskId": "run-1:task-1",
|
||||
"contextId": "thread-1",
|
||||
"status": {
|
||||
"state": "working",
|
||||
"message": {
|
||||
"kind": "message",
|
||||
"role": "agent",
|
||||
"parts": [{"kind": "text", "text": "langgraph echo: hi"}],
|
||||
"messageId": "m-agent",
|
||||
"taskId": "run-1:task-1",
|
||||
"contextId": "thread-1",
|
||||
},
|
||||
},
|
||||
"final": False,
|
||||
},
|
||||
},
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": "reply",
|
||||
"result": {
|
||||
"kind": "artifact-update",
|
||||
"taskId": "run-1:task-1",
|
||||
"contextId": "thread-1",
|
||||
"artifact": {
|
||||
"artifactId": "art-1",
|
||||
"name": "Assistant Response",
|
||||
"parts": [{"kind": "text", "text": "langgraph echo: hi"}],
|
||||
},
|
||||
"lastChunk": True,
|
||||
},
|
||||
},
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": "reply",
|
||||
"result": {
|
||||
"kind": "status-update",
|
||||
"taskId": "run-1:task-1",
|
||||
"contextId": "thread-1",
|
||||
"status": {"state": "completed"},
|
||||
"final": True,
|
||||
},
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
_LOWERCASE_BINDING_CARD = {
|
||||
"name": "langgraph-agent",
|
||||
"version": "1.0.0",
|
||||
|
|
@ -246,9 +317,10 @@ _UPPERCASE_BINDING_CARD = {
|
|||
class _RequestRecorder:
|
||||
"""Records the headers httpx put on the wire, per outbound request."""
|
||||
|
||||
def __init__(self, card=_AGENT_CARD, rpc_reply=_RPC_REPLY):
|
||||
def __init__(self, card=_AGENT_CARD, rpc_reply=_RPC_REPLY, rpc_stream_events=None):
|
||||
self.card = card
|
||||
self.rpc_reply = rpc_reply
|
||||
self.rpc_stream_events = rpc_stream_events
|
||||
self.card_requests = []
|
||||
self.card_urls = []
|
||||
self.rpc_requests = []
|
||||
|
|
@ -263,6 +335,9 @@ class _RequestRecorder:
|
|||
return httpx.Response(200, json=self.card)
|
||||
self.rpc_requests.append(headers)
|
||||
self.rpc_bodies.append(json.loads(request.content))
|
||||
if self.rpc_stream_events is not None:
|
||||
body = "".join(f"data: {json.dumps(event)}\n\n" for event in self.rpc_stream_events)
|
||||
return httpx.Response(200, headers={"content-type": "text/event-stream"}, content=body)
|
||||
return httpx.Response(200, json=self.rpc_reply)
|
||||
|
||||
|
||||
|
|
@ -271,7 +346,10 @@ def _a2a_client_cache_key(timeout: float, provider: str = httpxSpecialProvider.A
|
|||
|
||||
|
||||
async def _seed_shared_a2a_client(
|
||||
card=_AGENT_CARD, rpc_reply=_RPC_REPLY, provider: str = httpxSpecialProvider.A2AProvider
|
||||
card=_AGENT_CARD,
|
||||
rpc_reply=_RPC_REPLY,
|
||||
rpc_stream_events=None,
|
||||
provider: str = httpxSpecialProvider.A2AProvider,
|
||||
) -> _RequestRecorder:
|
||||
"""Put the one A2A client the cache will hand out behind a mock transport.
|
||||
|
||||
|
|
@ -279,7 +357,7 @@ async def _seed_shared_a2a_client(
|
|||
it. The injected client is a real httpx.AsyncClient, so the merge of per-request
|
||||
headers over client defaults, which is what these tests are about, stays real.
|
||||
"""
|
||||
recorder = _RequestRecorder(card=card, rpc_reply=rpc_reply)
|
||||
recorder = _RequestRecorder(card=card, rpc_reply=rpc_reply, rpc_stream_events=rpc_stream_events)
|
||||
handler = AsyncHTTPHandler(timeout=DEFAULT_A2A_AGENT_TIMEOUT)
|
||||
owned_client = handler.client
|
||||
handler.client = httpx.AsyncClient(transport=httpx.MockTransport(recorder))
|
||||
|
|
@ -431,6 +509,50 @@ async def test_protocol_version_override_round_trips_the_langgraph_dialect(card,
|
|||
assert recorder.rpc_bodies[-1]["method"] == "message/send"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("card", [_LOWERCASE_BINDING_CARD, _UPPERCASE_BINDING_CARD])
|
||||
@pytest.mark.asyncio
|
||||
async def test_protocol_version_override_streams_the_langgraph_dialect(card, isolated_client_cache):
|
||||
recorder = await _seed_shared_a2a_client(card=card, rpc_stream_events=_LANGGRAPH_STREAM_EVENTS)
|
||||
|
||||
a2a_client = await create_a2a_client(
|
||||
base_url="http://127.0.0.1:9",
|
||||
streaming=True,
|
||||
protocol_version="0.3",
|
||||
)
|
||||
request = SendStreamingMessageRequest(
|
||||
id="reply",
|
||||
params=MessageSendParams(
|
||||
message={"messageId": "m-user", "role": "user", "parts": [{"kind": "text", "text": "hi"}]}
|
||||
),
|
||||
)
|
||||
streamed = [chunk async for chunk in _stream_messages(a2a_client, request)]
|
||||
|
||||
artifact_event = next(chunk.root.result for chunk in streamed if chunk.root.result.kind == "artifact-update")
|
||||
completed_status = next(
|
||||
chunk.root.result for chunk in streamed if chunk.root.result.kind == "status-update" and chunk.root.result.final
|
||||
)
|
||||
assert artifact_event.artifact.parts[0].root.text == "langgraph echo: hi"
|
||||
assert completed_status.status.state.value == "completed"
|
||||
assert recorder.rpc_bodies[-1]["method"] == "message/stream"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_uppercase_binding_card_stream_without_override_fails(isolated_client_cache):
|
||||
await _seed_shared_a2a_client(card=_UPPERCASE_BINDING_CARD, rpc_stream_events=_LANGGRAPH_STREAM_EVENTS)
|
||||
|
||||
a2a_client = await create_a2a_client(base_url="http://127.0.0.1:9", streaming=True)
|
||||
request = SendStreamingMessageRequest(
|
||||
id="reply",
|
||||
params=MessageSendParams(
|
||||
message={"messageId": "m-user", "role": "user", "parts": [{"kind": "text", "text": "hi"}]}
|
||||
),
|
||||
)
|
||||
|
||||
with pytest.raises(Exception, match='has no field named "kind"'):
|
||||
async for _ in _stream_messages(a2a_client, request):
|
||||
pass
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_asend_message_resolves_float_protocol_version_from_agent_params(isolated_client_cache):
|
||||
recorder = await _seed_shared_a2a_client(card=_UPPERCASE_BINDING_CARD, rpc_reply=_LANGGRAPH_TASK_REPLY)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue