mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix: close second A2A review round
This commit is contained in:
parent
d2f05d1cf1
commit
33a8945ca5
7 changed files with 353 additions and 88 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue