mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix: preserve A2A pricing and stream data
This commit is contained in:
parent
33a8945ca5
commit
c1add95b62
9 changed files with 185 additions and 28 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
),
|
||||
)
|
||||
],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue