fix: close second A2A review round

This commit is contained in:
aiedwardyi 2026-08-24 23:32:53 +09:00
parent d2f05d1cf1
commit 33a8945ca5
No known key found for this signature in database
7 changed files with 353 additions and 88 deletions

View file

@ -154,7 +154,7 @@ class A2ACompletionBridgeHandler:
if a2a_provider_config is not None:
verbose_logger.info("A2A: Using provider config for %s", custom_llm_provider)
provider_params: Final = {key: value for key, value in params.items() if key != "messages"}
provider_params: Final = dict(params)
provider_kwargs: Final[dict[str, Any]] = {
"request_id": request_id,
"params": provider_params,
@ -227,7 +227,7 @@ class A2ACompletionBridgeHandler:
if a2a_provider_config is not None:
verbose_logger.info("A2A: Using provider config for %s (streaming)", custom_llm_provider)
provider_params: Final = {key: value for key, value in params.items() if key != "messages"}
provider_params: Final = dict(params)
provider_kwargs: Final[dict[str, Any]] = {
"request_id": request_id,
"params": provider_params,
@ -237,8 +237,14 @@ class A2ACompletionBridgeHandler:
}
if litellm_params.get("timeout") is not None:
provider_kwargs["timeout"] = litellm_params["timeout"]
async for chunk in a2a_provider_config.handle_streaming(**provider_kwargs):
yield chunk
provider_stream: Final = a2a_provider_config.handle_streaming(**provider_kwargs)
try:
async for chunk in provider_stream:
yield chunk
finally:
close_provider_stream = getattr(provider_stream, "aclose", None)
if close_provider_stream is not None:
await close_provider_stream()
return
@ -274,26 +280,52 @@ class A2ACompletionBridgeHandler:
# 3. Forward content as artifact updates
accumulated_tool_calls: Final[list[object]] = [] # mutable-ok: collect streaming tool-call deltas
stream_usage: object | None = None
stream_finish_reason: str | None = None
chunk_count = 0
async for chunk in response:
chunk_count += 1
try:
async for chunk in response:
chunk_count += 1
# Extract delta content
content = ""
if chunk is not None and hasattr(chunk, "choices") and chunk.choices:
choice = chunk.choices[0]
if hasattr(choice, "delta") and choice.delta:
content = choice.delta.content or ""
tool_calls = getattr(choice.delta, "tool_calls", None)
if isinstance(tool_calls, (list, tuple)):
accumulated_tool_calls.extend(tool_calls)
raw_usage = getattr(chunk, "usage", None)
if isinstance(raw_usage, Mapping):
stream_usage = raw_usage
else:
dump_usage = getattr(raw_usage, "model_dump", None)
if callable(dump_usage):
dumped_usage = dump_usage(exclude_none=True)
if isinstance(dumped_usage, Mapping):
stream_usage = dumped_usage
else:
dict_usage = getattr(raw_usage, "dict", None)
if callable(dict_usage):
dumped_usage = dict_usage(exclude_none=True)
if isinstance(dumped_usage, Mapping):
stream_usage = dumped_usage
if content:
artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event(
ctx=ctx,
text=content,
)
yield artifact_event
# Extract delta content
content = ""
if chunk is not None and hasattr(chunk, "choices") and chunk.choices:
choice = chunk.choices[0]
raw_finish_reason = getattr(choice, "finish_reason", None)
if isinstance(raw_finish_reason, str) and raw_finish_reason:
stream_finish_reason = raw_finish_reason
if hasattr(choice, "delta") and choice.delta:
content = choice.delta.content or ""
tool_calls = getattr(choice.delta, "tool_calls", None)
if isinstance(tool_calls, (list, tuple)):
accumulated_tool_calls.extend(tool_calls)
if content:
artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event(
ctx=ctx,
text=content,
)
yield artifact_event
finally:
close_response = getattr(response, "aclose", None)
if close_response is not None:
await close_response()
# 4. Emit final status update (kind: "status-update", status: "completed", final: true)
completed_event: Final = A2ACompletionBridgeTransformation.create_status_update_event(
@ -303,6 +335,10 @@ class A2ACompletionBridgeHandler:
)
if accumulated_tool_calls:
completed_event["result"]["tool_calls"] = accumulated_tool_calls
if stream_finish_reason:
completed_event["result"]["finish_reason"] = stream_finish_reason
if stream_usage is not None:
completed_event["usage"] = stream_usage
yield completed_event
verbose_logger.info(

View file

@ -151,6 +151,22 @@ class A2ACompletionBridgeTransformation:
return [openai_message]
@staticmethod
def _model_dump(value: Any) -> dict[str, Any]:
if isinstance(value, dict):
return value
dump = getattr(value, "model_dump", None)
if callable(dump):
dumped = dump(exclude_none=True)
if isinstance(dumped, dict):
return dumped
dump = getattr(value, "dict", None)
if callable(dump):
dumped = dump(exclude_none=True)
if isinstance(dumped, dict):
return dumped
return {}
@staticmethod
def openai_response_to_a2a_response(
response: Any,
@ -170,16 +186,19 @@ class A2ACompletionBridgeTransformation:
raw_choices: Final = getattr(response, "choices", None)
if raw_choices:
for choice in raw_choices:
content: Final = (
getattr(getattr(choice, "message", None), "content", None) or ""
)
raw_message = getattr(choice, "message", None)
message_fields: Final = A2ACompletionBridgeTransformation._model_dump(raw_message)
raw_content = message_fields.get("content")
if raw_content is None:
raw_content = getattr(raw_message, "content", None)
content: Final = raw_content if isinstance(raw_content, str) else ""
message: Final = {
"kind": "message",
"role": "agent",
"parts": [{"kind": "text", "text": content}],
"messageId": uuid4().hex,
}
raw_tool_calls = getattr(getattr(choice, "message", None), "tool_calls", None)
raw_tool_calls = message_fields.get("tool_calls")
if raw_tool_calls:
message["tool_calls"] = [
call.model_dump(exclude_none=True)
@ -189,10 +208,38 @@ class A2ACompletionBridgeTransformation:
else call
for call in raw_tool_calls
]
finish_reason: Final = getattr(choice, "finish_reason", None)
for field in (
"annotations",
"audio",
"function_call",
"images",
"provider_specific_fields",
"reasoning_content",
"reasoning_items",
"thinking_blocks",
):
value = message_fields.get(field)
if value is not None:
message[field] = value
choice_fields: Final = A2ACompletionBridgeTransformation._model_dump(choice)
finish_reason: Final = choice_fields.get("finish_reason")
if finish_reason is None:
raw_finish_reason = getattr(choice, "finish_reason", None)
finish_reason = raw_finish_reason if isinstance(raw_finish_reason, str) else None
if finish_reason:
message["finish_reason"] = finish_reason
serialized_choices.append({"index": len(serialized_choices), "message": message})
choice_payload: Final[dict[str, Any]] = {
"index": len(serialized_choices),
"message": message,
}
logprobs = choice_fields.get("logprobs")
if logprobs is None:
raw_logprobs = getattr(choice, "logprobs", None)
logprobs = raw_logprobs if isinstance(raw_logprobs, dict) else None
if logprobs is not None:
choice_payload["logprobs"] = logprobs
message["logprobs"] = logprobs
serialized_choices.append(choice_payload)
a2a_message: Final = (
serialized_choices[0]["message"]

View file

@ -2,8 +2,10 @@
A2A Streaming Response Iterator
"""
from collections.abc import Mapping
from typing import Final
import litellm
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.types.llms.openai import ChatCompletionToolCallChunk
from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
@ -73,13 +75,14 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
# Determine finish reason
finish_reason: Final = self._get_finish_reason(chunk)
tool_calls: Final = self._get_tool_calls(chunk)
usage: Final = self._get_usage(chunk)
# Return generic streaming chunk
return GenericStreamingChunk(
text=text,
is_finished=bool(finish_reason or tool_calls),
finish_reason=finish_reason or ("tool_calls" if tool_calls else ""),
usage=None,
usage=usage,
index=0,
tool_use=tool_calls,
)
@ -105,6 +108,14 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
# Check for task completion
if isinstance(result, dict):
explicit_finish_reason: Final = result.get("finish_reason")
if isinstance(explicit_finish_reason, str) and explicit_finish_reason:
return explicit_finish_reason
message: Final = result.get("message")
if isinstance(message, dict):
message_finish_reason: Final = message.get("finish_reason")
if isinstance(message_finish_reason, str) and message_finish_reason:
return message_finish_reason
status: Final = result.get("status", {})
if isinstance(status, dict):
state: Final = status.get("state")
@ -119,16 +130,53 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
return None
def _get_usage(self, chunk: dict) -> object | None:
raw_usage: object | None = chunk.get("usage")
result: Final = chunk.get("result", {})
if raw_usage is None and isinstance(result, dict):
raw_usage = result.get("usage")
if raw_usage is None:
return None
if isinstance(raw_usage, Mapping):
try:
return litellm.Usage(**raw_usage)
except Exception:
return raw_usage
if hasattr(raw_usage, "model_dump"):
try:
return litellm.Usage(**raw_usage.model_dump(exclude_none=True))
except Exception:
return raw_usage
return raw_usage
def _get_tool_calls(self, chunk: dict) -> ChatCompletionToolCallChunk | None:
result: Final = chunk.get("result", {})
if not isinstance(result, dict):
return None
tool_calls = result.get("tool_calls")
if isinstance(tool_calls, list) and tool_calls:
first_tool_call: Final = tool_calls[0]
return first_tool_call if isinstance(first_tool_call, dict) else None
return self._serialize_tool_call(tool_calls[0])
message = result.get("message")
if isinstance(message, dict) and isinstance(message.get("tool_calls"), list) and message["tool_calls"]:
first_tool_call = message["tool_calls"][0]
return first_tool_call if isinstance(first_tool_call, dict) else None
return self._serialize_tool_call(message["tool_calls"][0])
return None
@staticmethod
def _serialize_tool_call(tool_call: object) -> ChatCompletionToolCallChunk | None:
if isinstance(tool_call, dict):
return tool_call
if hasattr(tool_call, "model_dump"):
return tool_call.model_dump(exclude_none=True)
if hasattr(tool_call, "dict"):
return tool_call.dict(exclude_none=True)
return None
async def aclose(self) -> None:
streaming_response = self.streaming_response
self.streaming_response = None
try:
await super().aclose()
finally:
close_stream = getattr(streaming_response, "aclose", None)
if close_stream is not None:
await close_stream()

View file

@ -52,6 +52,7 @@ _FORWARDED_REQUEST_PARAMS: Final = frozenset(
"response_format",
"seed",
"service_tier",
"safety_identifier",
"stop",
"store",
"temperature",
@ -64,6 +65,8 @@ _FORWARDED_REQUEST_PARAMS: Final = frozenset(
"user",
"verbosity",
"web_search_options",
"output_config",
"prompt_cache_key",
}
)
_A2A_PRICING_PARAMS: Final = frozenset({"cost_per_query", "response_cost"}) | frozenset(
@ -159,6 +162,7 @@ async def _route_registered_provider(
logging_obj: Final = data.get("litellm_logging_obj")
if isinstance(logging_obj, Logging):
provider_params["no-log"] = True
pricing_params = {
key: litellm_params[key]
for key in _A2A_PRICING_PARAMS
@ -168,7 +172,6 @@ async def _route_registered_provider(
logging_obj.litellm_params.update(pricing_params)
logging_obj.model_call_details["litellm_params"].update(pricing_params)
logging_obj.custom_pricing = True
provider_params["no-log"] = True
if stream:
streaming_response: Final = A2ACompletionBridgeHandler.handle_streaming(
@ -214,64 +217,80 @@ async def _route_registered_provider(
nested_message: Final = result_dict.get("message")
response_message: Final = nested_message if isinstance(nested_message, Mapping) else result_dict
response_choices: Final = response.get("choices")
choice_payloads: Final = (
response_choices
if isinstance(response_choices, list)
else result_dict.get("choices")
)
choice_payloads: Final = response_choices if isinstance(response_choices, list) else result_dict.get("choices")
def _serialize_value(value: object) -> object:
if hasattr(value, "model_dump"):
return value.model_dump(exclude_none=True)
if hasattr(value, "dict"):
return value.dict(exclude_none=True)
return value
def _build_message(message_payload: Mapping[str, object], content: str) -> Message:
message_kwargs: dict[str, object] = {
"content": content,
"role": "assistant",
}
raw_tool_calls = message_payload.get("tool_calls")
if isinstance(raw_tool_calls, list):
message_kwargs["tool_calls"] = raw_tool_calls
for field in (
"audio",
"annotations",
"function_call",
"images",
"provider_specific_fields",
"reasoning_content",
"reasoning_items",
"thinking_blocks",
):
value = message_payload.get(field)
if value is not None:
message_kwargs[field] = _serialize_value(value)
return Message(**message_kwargs)
if isinstance(choice_payloads, list) and choice_payloads:
model_choices = [
Choices(
finish_reason=(
choice.get("finish_reason")
if isinstance(choice, Mapping) and isinstance(choice.get("finish_reason"), str)
else choice.get("message", {}).get("finish_reason")
if isinstance(choice, Mapping)
and isinstance(choice.get("message"), Mapping)
and isinstance(choice.get("message", {}).get("finish_reason"), str)
model_choices = []
for choice_index, choice in enumerate(choice_payloads):
choice_mapping: Mapping[str, object] = choice if isinstance(choice, Mapping) else {}
raw_message = choice_mapping.get("message")
message_payload: Mapping[str, object] = raw_message if isinstance(raw_message, Mapping) else choice_mapping
choice_kwargs: dict[str, object] = {
"finish_reason": (
choice_mapping.get("finish_reason")
if isinstance(choice_mapping.get("finish_reason"), str)
else message_payload.get("finish_reason")
if isinstance(message_payload.get("finish_reason"), str)
else "stop"
),
index=choice.get("index", choice_index)
if isinstance(choice, Mapping) and isinstance(choice.get("index", choice_index), int)
"index": choice_mapping.get("index", choice_index)
if isinstance(choice_mapping.get("index", choice_index), int)
else choice_index,
message=Message(
content=extract_text_from_a2a_response(
{"result": choice.get("message", choice)}
if isinstance(choice, Mapping)
else {"result": {}}
),
role="assistant",
tool_calls=(
choice.get("message", {}).get("tool_calls")
if isinstance(choice, Mapping)
and isinstance(choice.get("message"), Mapping)
and isinstance(choice.get("message", {}).get("tool_calls"), list)
else choice.get("tool_calls")
if isinstance(choice, Mapping) and isinstance(choice.get("tool_calls"), list)
else None
),
"message": _build_message(
message_payload,
extract_text_from_a2a_response({"result": message_payload}),
),
)
for choice_index, choice in enumerate(choice_payloads)
]
}
raw_logprobs = choice_mapping.get("logprobs", message_payload.get("logprobs"))
if raw_logprobs is not None:
choice_kwargs["logprobs"] = _serialize_value(raw_logprobs)
model_choices.append(Choices(**choice_kwargs))
else:
tool_calls: Final = response_message.get("tool_calls")
normalized_tool_calls: Final = tool_calls if isinstance(tool_calls, list) else None
finish_reason: Final = response_message.get("finish_reason")
text: Final = extract_text_from_a2a_response(response)
model_choices = [
Choices(
finish_reason=(
finish_reason
if isinstance(finish_reason, str)
else "tool_calls"
if normalized_tool_calls
else "stop"
),
index=0,
message=Message(content=text, role="assistant", tool_calls=normalized_tool_calls),
)
]
choice_kwargs = {
"finish_reason": (
finish_reason if isinstance(finish_reason, str) else "tool_calls" if normalized_tool_calls else "stop"
),
"index": 0,
"message": _build_message(response_message, text),
}
raw_logprobs = response_message.get("logprobs")
if raw_logprobs is not None:
choice_kwargs["logprobs"] = _serialize_value(raw_logprobs)
model_choices = [Choices(**choice_kwargs)]
model_response: Final = ModelResponse(
id=str(response.get("id") or request_id),
model=model_name,

View file

@ -12,6 +12,8 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.types.utils import Choices, Message, ModelResponse
class TestA2AStreamingTransformation:
"""Test the A2A streaming transformation creates proper events."""
@ -26,9 +28,7 @@ class TestA2AStreamingTransformation:
"parts": [{"text": "Reply to ticket #4823"}],
"metadata": {"skillId": "draft_reply"},
}
openai_messages = (
A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
)
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]
@ -174,10 +174,7 @@ class TestA2AStreamingTransformation:
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."
)
assert event["result"]["artifact"]["parts"][0]["text"] == "Hello, I am an AI assistant."
@pytest.mark.asyncio
@ -197,6 +194,8 @@ async def test_handle_streaming_emits_proper_events():
mock_chunk2.choices = [MagicMock()]
mock_chunk2.choices[0].delta = MagicMock()
mock_chunk2.choices[0].delta.content = " world"
mock_chunk2.choices[0].finish_reason = "length"
mock_chunk2.usage = {"prompt_tokens": 2, "completion_tokens": 3, "total_tokens": 5}
async def mock_streaming_response():
yield mock_chunk1
@ -246,6 +245,66 @@ async def test_handle_streaming_emits_proper_events():
assert events[4]["result"]["kind"] == "status-update"
assert events[4]["result"]["status"]["state"] == "completed"
assert events[4]["result"]["final"] is True
assert events[4]["result"]["finish_reason"] == "length"
assert events[4]["usage"]["total_tokens"] == 5
@pytest.mark.asyncio
async def test_provider_config_receives_full_message_history():
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2ACompletionBridgeHandler,
)
provider_config = MagicMock()
provider_config.handle_non_streaming = AsyncMock(return_value={"result": {}})
messages = [
{"role": "system", "content": "Be concise"},
{"role": "user", "content": "Hello"},
]
params = {
"message": {"role": "user", "parts": []},
"messages": messages,
}
with patch(
"litellm.a2a_protocol.litellm_completion_bridge.handler.A2AProviderConfigManager.get_provider_config",
return_value=provider_config,
):
await A2ACompletionBridgeHandler.handle_non_streaming(
request_id="req-1",
params=params,
litellm_params={"custom_llm_provider": "langflow", "model": "flow"},
)
assert provider_config.handle_non_streaming.await_args.kwargs["params"]["messages"] == messages
def test_response_transform_preserves_audio_and_logprobs():
from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
A2ACompletionBridgeTransformation,
)
response = ModelResponse(
id="resp-1",
model="test-model",
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(
content="hello",
role="assistant",
audio={"data": "abc", "expires_at": 1, "transcript": "hello"},
),
logprobs={"content": []},
)
],
)
transformed = A2ACompletionBridgeTransformation.openai_response_to_a2a_response(response)
assert transformed["result"]["audio"]["data"] == "abc"
assert transformed["result"]["logprobs"] == {"content": []}
@pytest.mark.asyncio

View file

@ -4,6 +4,7 @@ import pytest
from litellm.llms.a2a.chat.streaming_iterator import A2AModelResponseIterator
from litellm.llms.a2a.common_utils import A2AError
from litellm.types.utils import Delta
@pytest.mark.asyncio
@ -48,6 +49,55 @@ async def test_async_iterator_preserves_tool_calls():
assert chunk["finish_reason"] == "tool_calls"
@pytest.mark.asyncio
async def test_async_iterator_serializes_delta_tool_calls_and_usage():
delta = Delta(
tool_calls=[
{
"id": "call-1",
"type": "function",
"function": {"name": "lookup", "arguments": "{}"},
}
]
)
async def _events():
yield {
"jsonrpc": "2.0",
"usage": {"prompt_tokens": 2, "completion_tokens": 3, "total_tokens": 5},
"result": {
"tool_calls": [delta.tool_calls[0]],
"finish_reason": "length",
},
}
iterator = A2AModelResponseIterator(streaming_response=_events(), sync_stream=False)
chunk = await iterator.__aiter__().__anext__()
assert chunk["tool_use"]["id"] == "call-1"
assert chunk["finish_reason"] == "length"
assert chunk["usage"].total_tokens == 5
@pytest.mark.asyncio
async def test_async_iterator_closes_nested_stream():
closed = False
async def _events():
nonlocal closed
try:
yield {"jsonrpc": "2.0", "result": {"kind": "artifact-update"}}
raise AssertionError("stream should be closed before a second event")
finally:
closed = True
iterator = A2AModelResponseIterator(streaming_response=_events(), sync_stream=False)
await iterator.__aiter__().__anext__()
await iterator.aclose()
assert closed is True
@pytest.mark.asyncio
async def test_async_iterator_propagates_jsonrpc_errors():
async def _events():

View file

@ -98,6 +98,9 @@ async def test_route_a2a_model_uses_registered_provider():
"temperature": 0.2,
"timeout": 12.0,
"tools": [{"type": "function", "function": {"name": "lookup"}}],
"output_config": {"format": "json"},
"prompt_cache_key": "cache-key",
"safety_identifier": "safety-id",
"proxy_server_request": {
"headers": {
"x-tenant": "tenant-1",
@ -144,6 +147,9 @@ async def test_route_a2a_model_uses_registered_provider():
assert bridge_kwargs["litellm_params"]["temperature"] == 0.2
assert bridge_kwargs["litellm_params"]["timeout"] == 12.0
assert bridge_kwargs["litellm_params"]["tools"] == data["tools"]
assert bridge_kwargs["litellm_params"]["output_config"] == data["output_config"]
assert bridge_kwargs["litellm_params"]["prompt_cache_key"] == data["prompt_cache_key"]
assert bridge_kwargs["litellm_params"]["safety_identifier"] == data["safety_identifier"]
assert bridge_kwargs["litellm_params"]["guardrails"] == ["request-guardrail", "agent-guardrail"]
assert bridge_kwargs["litellm_params"]["extra_headers"] == {
"X-Tenant": "tenant-1",