fix: preserve A2A pricing and stream data

This commit is contained in:
aiedwardyi 2026-08-25 10:57:17 +09:00
parent 33a8945ca5
commit c1add95b62
No known key found for this signature in database
9 changed files with 185 additions and 28 deletions

View file

@ -280,6 +280,9 @@ class A2ACompletionBridgeHandler:
# 3. Forward content as artifact updates
accumulated_tool_calls: Final[list[object]] = [] # mutable-ok: collect streaming tool-call deltas
choice_texts: dict[int, str] = {}
choice_tool_calls: dict[int, list[object]] = {}
choice_finish_reasons: dict[int, str] = {}
stream_usage: object | None = None
stream_finish_reason: str | None = None
chunk_count = 0
@ -304,24 +307,34 @@ class A2ACompletionBridgeHandler:
stream_usage = dumped_usage
# 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)
choices = getattr(chunk, "choices", None) if chunk is not None else None
if isinstance(choices, (list, tuple)):
for choice_position, choice in enumerate(choices):
raw_index = getattr(choice, "index", choice_position)
choice_index = raw_index if isinstance(raw_index, int) else choice_position
choice_texts.setdefault(choice_index, "")
raw_finish_reason = getattr(choice, "finish_reason", None)
if isinstance(raw_finish_reason, str) and raw_finish_reason:
choice_finish_reasons[choice_index] = raw_finish_reason
if choice_index == 0 or stream_finish_reason is None:
stream_finish_reason = raw_finish_reason
content = ""
delta = getattr(choice, "delta", None)
if delta:
raw_content = getattr(delta, "content", None)
content = raw_content if isinstance(raw_content, str) else ""
choice_texts[choice_index] += content
tool_calls = getattr(delta, "tool_calls", None)
if isinstance(tool_calls, (list, tuple)):
accumulated_tool_calls.extend(tool_calls)
choice_tool_calls.setdefault(choice_index, []).extend(tool_calls)
if content:
artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event(
ctx=ctx,
text=content,
)
yield artifact_event
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:
@ -339,6 +352,28 @@ class A2ACompletionBridgeHandler:
completed_event["result"]["finish_reason"] = stream_finish_reason
if stream_usage is not None:
completed_event["usage"] = stream_usage
if len(choice_texts) > 1:
completed_event["result"]["choices"] = [
{
"index": choice_index,
"message": {
"kind": "message",
"role": "agent",
"parts": [{"kind": "text", "text": choice_texts[choice_index]}],
**(
{"tool_calls": choice_tool_calls[choice_index]}
if choice_tool_calls.get(choice_index)
else {}
),
},
**(
{"finish_reason": choice_finish_reasons[choice_index]}
if choice_index in choice_finish_reasons
else {}
),
}
for choice_index in sorted(choice_texts)
]
yield completed_event
verbose_logger.info(

View file

@ -1233,7 +1233,8 @@ class CustomStreamWrapper:
)
if "tool_use" in anthropic_response_obj and anthropic_response_obj["tool_use"] is not None:
completion_obj["tool_calls"] = [anthropic_response_obj["tool_use"]]
tool_use = anthropic_response_obj["tool_use"]
completion_obj["tool_calls"] = tool_use if isinstance(tool_use, list) else [tool_use]
if (
"provider_specific_fields" in anthropic_response_obj
@ -2559,6 +2560,8 @@ def convert_generic_chunk_to_model_response_stream(
) -> ModelResponseStream:
from litellm.types.utils import Delta
tool_use = chunk.get("tool_use", None)
tool_calls = tool_use if isinstance(tool_use, list) else [tool_use] if tool_use is not None else None
model_response_stream: Final = ModelResponseStream(
id=str(uuid.uuid4()),
model="",
@ -2567,7 +2570,7 @@ def convert_generic_chunk_to_model_response_stream(
index=chunk.get("index", 0),
delta=Delta(
content=chunk["text"],
tool_calls=chunk.get("tool_use", None),
tool_calls=tool_calls,
),
)
],

View file

@ -149,16 +149,18 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
return raw_usage
return raw_usage
def _get_tool_calls(self, chunk: dict) -> ChatCompletionToolCallChunk | None:
def _get_tool_calls(
self, chunk: dict
) -> ChatCompletionToolCallChunk | list[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:
return self._serialize_tool_call(tool_calls[0])
return self._serialize_tool_calls(tool_calls)
message = result.get("message")
if isinstance(message, dict) and isinstance(message.get("tool_calls"), list) and message["tool_calls"]:
return self._serialize_tool_call(message["tool_calls"][0])
return self._serialize_tool_calls(message["tool_calls"])
return None
@staticmethod
@ -171,6 +173,19 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
return tool_call.dict(exclude_none=True)
return None
@classmethod
def _serialize_tool_calls(
cls, tool_calls: list[object]
) -> ChatCompletionToolCallChunk | list[ChatCompletionToolCallChunk] | None:
serialized: Final = [
tool_call_value
for tool_call in tool_calls
if (tool_call_value := cls._serialize_tool_call(tool_call)) is not None
]
if len(serialized) == 1:
return serialized[0]
return serialized or None
async def aclose(self) -> None:
streaming_response = self.streaming_response
self.streaming_response = None

View file

@ -163,6 +163,13 @@ async def _route_registered_provider(
logging_obj: Final = data.get("litellm_logging_obj")
if isinstance(logging_obj, Logging):
provider_params["no-log"] = True
provider_model: Final = litellm_params.get("model")
if isinstance(provider_model, str):
logging_obj.model_call_details["model"] = provider_model
logging_obj.model_call_details.setdefault("litellm_params", {})["model"] = provider_model
provider_name: Final = litellm_params.get("custom_llm_provider")
if isinstance(provider_name, str):
logging_obj.model_call_details["custom_llm_provider"] = provider_name
pricing_params = {
key: litellm_params[key]
for key in _A2A_PRICING_PARAMS

View file

@ -1841,11 +1841,6 @@ class ProxyBaseLLMRequestProcessing:
merge_a2a_agent_guardrails_before_hooks,
)
self.data = await authorize_a2a_agent_before_hooks(
data=self.data,
user_api_key_dict=user_api_key_dict,
)
logging_obj, self.data = litellm.utils.function_setup(
original_function=route_type,
rules_obj=litellm.utils.Rules(),
@ -1855,6 +1850,11 @@ class ProxyBaseLLMRequestProcessing:
self.data["litellm_logging_obj"] = logging_obj
self.data = await authorize_a2a_agent_before_hooks(
data=self.data,
user_api_key_dict=user_api_key_dict,
)
self.data = await merge_a2a_agent_guardrails_before_hooks(self.data)
# Merge model-level guardrails before pre_call_hook so DB/UI-configured

View file

@ -317,7 +317,7 @@ class ModelInfo(ModelInfoBase, total=False):
class GenericStreamingChunk(TypedDict, total=False):
text: Required[str]
tool_use: ChatCompletionToolCallChunk | None
tool_use: ChatCompletionToolCallChunk | list[ChatCompletionToolCallChunk] | None
is_finished: Required[bool]
finish_reason: Required[str]
usage: Required[ChatCompletionUsageBlock | None]

View file

@ -249,6 +249,44 @@ async def test_handle_streaming_emits_proper_events():
assert events[4]["usage"]["total_tokens"] == 5
@pytest.mark.asyncio
async def test_handle_streaming_preserves_multiple_choices():
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2ACompletionBridgeHandler,
)
mock_chunk = MagicMock()
first_choice = MagicMock()
first_choice.index = 0
first_choice.finish_reason = None
first_choice.delta.content = "first"
second_choice = MagicMock()
second_choice.index = 1
second_choice.finish_reason = "length"
second_choice.delta.content = "second"
mock_chunk.choices = [first_choice, second_choice]
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()
events = [
event
async for event in A2ACompletionBridgeHandler.handle_streaming(
request_id="req-choices",
params={"message": {"role": "user", "parts": []}},
litellm_params={"custom_llm_provider": "langgraph", "model": "agent", "n": 2},
)
]
choices = events[-1]["result"]["choices"]
assert [choice["index"] for choice in choices] == [0, 1]
assert choices[0]["message"]["parts"][0]["text"] == "first"
assert choices[1]["message"]["parts"][0]["text"] == "second"
assert choices[1]["finish_reason"] == "length"
@pytest.mark.asyncio
async def test_provider_config_receives_full_message_history():
from litellm.a2a_protocol.litellm_completion_bridge.handler import (

View file

@ -49,6 +49,31 @@ async def test_async_iterator_preserves_tool_calls():
assert chunk["finish_reason"] == "tool_calls"
@pytest.mark.asyncio
async def test_async_iterator_preserves_parallel_tool_calls():
tool_calls = [
{
"id": "call-1",
"type": "function",
"function": {"name": "lookup", "arguments": "{}"},
},
{
"id": "call-2",
"type": "function",
"function": {"name": "write", "arguments": "{}"},
},
]
async def _events():
yield {"jsonrpc": "2.0", "result": {"tool_calls": tool_calls}}
iterator = A2AModelResponseIterator(streaming_response=_events(), sync_stream=False)
chunk = await iterator.__aiter__().__anext__()
assert chunk["tool_use"] == tool_calls
@pytest.mark.asyncio
async def test_async_iterator_serializes_delta_tool_calls_and_usage():
delta = Delta(

View file

@ -11,6 +11,7 @@ from fastapi import HTTPException
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints.a2a_routing import (
_route_registered_provider,
merge_a2a_agent_guardrails_before_hooks,
route_a2a_agent_request,
)
@ -642,6 +643,39 @@ async def test_route_a2a_stream_uses_registered_provider():
assert response is wrapper
@pytest.mark.asyncio
async def test_registered_provider_logging_uses_provider_model_for_builtin_pricing():
class FakeLogging:
def __init__(self) -> None:
self.model_call_details = {"litellm_params": {}}
self.litellm_params = self.model_call_details["litellm_params"]
self.custom_pricing = False
logging_obj = FakeLogging()
response = {"result": {"message": {"parts": [{"kind": "text", "text": "hello"}]}}}
with (
patch("litellm.litellm_core_utils.litellm_logging.Logging", FakeLogging),
patch(
"litellm.a2a_protocol.litellm_completion_bridge.handler.A2ACompletionBridgeHandler.handle_non_streaming",
AsyncMock(return_value=response),
),
):
await _route_registered_provider(
data={
"messages": [{"role": "user", "content": "hello"}],
"litellm_logging_obj": logging_obj,
},
model_name="a2a/agent",
api_base="https://provider.example",
litellm_params={"model": "gpt-4o", "custom_llm_provider": "openai"},
static_headers=None,
)
assert logging_obj.model_call_details["model"] == "gpt-4o"
assert logging_obj.model_call_details["custom_llm_provider"] == "openai"
assert logging_obj.model_call_details["litellm_params"]["model"] == "gpt-4o"
@pytest.mark.asyncio
async def test_route_non_a2a_model_raises_error_if_not_in_router():
"""Test that non-a2a models that aren't in router raise an error"""