mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix: address A2A streaming review feedback
This commit is contained in:
parent
a5997bb908
commit
d2bfb5ed1b
9 changed files with 241 additions and 33 deletions
|
|
@ -46,6 +46,21 @@ class A2ACompletionBridgeHandler:
|
|||
Static methods for handling A2A requests via LiteLLM completion.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _merge_stream_values(previous: object, current: object) -> object:
|
||||
if isinstance(previous, Mapping) and isinstance(current, Mapping):
|
||||
merged = dict(previous)
|
||||
for key, value in current.items():
|
||||
merged[key] = (
|
||||
A2ACompletionBridgeHandler._merge_stream_values(merged[key], value)
|
||||
if key in merged
|
||||
else value
|
||||
)
|
||||
return merged
|
||||
if isinstance(previous, list) and isinstance(current, list):
|
||||
return [*previous, *current]
|
||||
return current
|
||||
|
||||
@staticmethod
|
||||
def _build_completion_params(
|
||||
params: dict[str, Any],
|
||||
|
|
@ -94,7 +109,8 @@ class A2ACompletionBridgeHandler:
|
|||
litellm_params_to_add: Final = {
|
||||
k: v
|
||||
for k, v in litellm_params.items()
|
||||
if k not in ("model", "custom_llm_provider", "extra_headers", "headers") and k not in _AGENT_ONLY_PARAMS
|
||||
if k not in ("model", "custom_llm_provider", "extra_headers", "headers", "api_base", "stream")
|
||||
and k not in _AGENT_ONLY_PARAMS
|
||||
}
|
||||
completion_params.update(litellm_params_to_add)
|
||||
# Apply forward metadata AFTER the litellm_params merge so the helper
|
||||
|
|
@ -290,8 +306,9 @@ class A2ACompletionBridgeHandler:
|
|||
choice_texts: dict[int, str] = {}
|
||||
choice_tool_calls: dict[int, list[object]] = {}
|
||||
choice_delta_fields: dict[int, dict[str, object]] = {}
|
||||
choice_logprobs: dict[int, object] = {}
|
||||
choice_logprobs: dict[int, dict[str, object]] = {}
|
||||
choice_finish_reasons: dict[int, str] = {}
|
||||
stream_metadata: dict[str, str] = {}
|
||||
stream_usage: object | None = None
|
||||
stream_finish_reason: str | None = None
|
||||
chunk_count = 0
|
||||
|
|
@ -315,6 +332,14 @@ class A2ACompletionBridgeHandler:
|
|||
if isinstance(dumped_usage, Mapping):
|
||||
stream_usage = dumped_usage
|
||||
|
||||
for metadata_name in ("system_fingerprint", "service_tier"):
|
||||
metadata_value = getattr(chunk, metadata_name, None)
|
||||
if not isinstance(metadata_value, str):
|
||||
chunk_fields = A2ACompletionBridgeTransformation._model_dump(chunk)
|
||||
metadata_value = chunk_fields.get(metadata_name)
|
||||
if isinstance(metadata_value, str) and metadata_value:
|
||||
stream_metadata[metadata_name] = metadata_value
|
||||
|
||||
# Extract delta content
|
||||
choices = getattr(chunk, "choices", None) if chunk is not None else None
|
||||
if isinstance(choices, (list, tuple)):
|
||||
|
|
@ -356,7 +381,12 @@ class A2ACompletionBridgeHandler:
|
|||
raw_logprobs = getattr(choice, "logprobs", None)
|
||||
serialized_logprobs = A2ACompletionBridgeTransformation._model_dump(raw_logprobs)
|
||||
if serialized_logprobs:
|
||||
choice_logprobs[choice_index] = serialized_logprobs
|
||||
previous_logprobs = choice_logprobs.get(choice_index, {})
|
||||
merged_logprobs = A2ACompletionBridgeHandler._merge_stream_values(
|
||||
previous_logprobs, serialized_logprobs
|
||||
)
|
||||
if isinstance(merged_logprobs, dict):
|
||||
choice_logprobs[choice_index] = merged_logprobs
|
||||
|
||||
if content:
|
||||
artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event(
|
||||
|
|
@ -382,9 +412,18 @@ 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:
|
||||
for metadata_name, metadata_value in stream_metadata.items():
|
||||
completed_event[metadata_name] = metadata_value
|
||||
choice_indices = sorted(
|
||||
set(choice_texts)
|
||||
| set(choice_tool_calls)
|
||||
| set(choice_delta_fields)
|
||||
| set(choice_logprobs)
|
||||
| set(choice_finish_reasons)
|
||||
)
|
||||
if choice_indices:
|
||||
choice_payloads: list[dict[str, object]] = []
|
||||
for choice_index in sorted(choice_texts):
|
||||
for choice_index in choice_indices:
|
||||
choice_payload: dict[str, object] = {
|
||||
"index": choice_index,
|
||||
"message": {
|
||||
|
|
@ -408,25 +447,6 @@ class A2ACompletionBridgeHandler:
|
|||
choice_payload["delta"] = choice_delta_fields[choice_index]
|
||||
choice_payloads.append(choice_payload)
|
||||
completed_event["result"]["choices"] = choice_payloads
|
||||
else:
|
||||
metadata_indices = sorted(set(choice_delta_fields) | set(choice_logprobs))
|
||||
if metadata_indices:
|
||||
completed_event["result"]["choices"] = [
|
||||
{
|
||||
"index": choice_index,
|
||||
**(
|
||||
{"delta": choice_delta_fields[choice_index]}
|
||||
if choice_delta_fields.get(choice_index)
|
||||
else {}
|
||||
),
|
||||
**(
|
||||
{"logprobs": choice_logprobs[choice_index]}
|
||||
if choice_index in choice_logprobs
|
||||
else {}
|
||||
),
|
||||
}
|
||||
for choice_index in metadata_indices
|
||||
]
|
||||
yield completed_event
|
||||
|
||||
verbose_logger.info(
|
||||
|
|
|
|||
|
|
@ -34,6 +34,7 @@ class A2AStreamingContext:
|
|||
self.request_id = request_id
|
||||
self.task_id = str(uuid4())
|
||||
self.context_id = str(uuid4())
|
||||
self.artifact_id = str(uuid4())
|
||||
self.input_message = input_message
|
||||
self.accumulated_text = ""
|
||||
self.has_emitted_task = False
|
||||
|
|
@ -368,7 +369,7 @@ class A2ACompletionBridgeTransformation:
|
|||
text: The text content for the artifact
|
||||
"""
|
||||
artifact: Final[dict[str, Any]] = {
|
||||
"artifactId": str(uuid4()),
|
||||
"artifactId": ctx.artifact_id,
|
||||
"name": "response",
|
||||
"parts": [{"kind": "text", "text": text}],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig
|
|||
from litellm.a2a_protocol.providers.watsonx_orchestrate.handler import (
|
||||
WatsonxOrchestrateHandler,
|
||||
)
|
||||
from litellm.interactions.agents.utils import merge_agent_headers
|
||||
|
||||
|
||||
class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig):
|
||||
|
|
@ -28,11 +29,16 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig):
|
|||
"litellm_params is required for WatsonxOrchestrateA2AConfig "
|
||||
"(must contain cp4d_host, instance_id, wxo_agent_id, api_key)"
|
||||
)
|
||||
forwarded_headers: Final = merge_agent_headers(
|
||||
dynamic_headers=kwargs.get("agent_extra_headers"),
|
||||
static_headers=kwargs.get("agent_static_headers"),
|
||||
)
|
||||
return await WatsonxOrchestrateHandler.handle_non_streaming(
|
||||
request_id=request_id,
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
static_headers=kwargs.get("agent_static_headers"),
|
||||
static_headers=forwarded_headers,
|
||||
timeout=kwargs.get("timeout"),
|
||||
)
|
||||
|
||||
async def handle_streaming(
|
||||
|
|
@ -49,10 +55,15 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig):
|
|||
"litellm_params is required for WatsonxOrchestrateA2AConfig "
|
||||
"(must contain cp4d_host, instance_id, wxo_agent_id, api_key)"
|
||||
)
|
||||
forwarded_headers: Final = merge_agent_headers(
|
||||
dynamic_headers=kwargs.get("agent_extra_headers"),
|
||||
static_headers=kwargs.get("agent_static_headers"),
|
||||
)
|
||||
async for chunk in WatsonxOrchestrateHandler.handle_streaming(
|
||||
request_id=request_id,
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
static_headers=kwargs.get("agent_static_headers"),
|
||||
static_headers=forwarded_headers,
|
||||
timeout=kwargs.get("timeout"),
|
||||
):
|
||||
yield chunk
|
||||
|
|
|
|||
|
|
@ -305,10 +305,13 @@ class WatsonxOrchestrateHandler:
|
|||
params: dict[str, object],
|
||||
litellm_params: WXOLitellmParams,
|
||||
static_headers: Mapping[str, str] | None = None,
|
||||
timeout: float | None = None,
|
||||
) -> dict[str, object]:
|
||||
wxo: Final = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params)
|
||||
|
||||
client: Final = WatsonxOrchestrateHandler._http_client(timeout=90.0)
|
||||
client: Final = WatsonxOrchestrateHandler._http_client(
|
||||
timeout=timeout if timeout is not None else 90.0
|
||||
)
|
||||
token: Final = await WatsonxOrchestrateHandler._get_bearer_token(
|
||||
cp4d_host=wxo.cp4d_host,
|
||||
auth_mode=wxo.auth_mode,
|
||||
|
|
@ -355,10 +358,13 @@ class WatsonxOrchestrateHandler:
|
|||
chunk_size: int = 50,
|
||||
delay_ms: int = 10,
|
||||
static_headers: Mapping[str, str] | None = None,
|
||||
timeout: float | None = None,
|
||||
) -> AsyncIterator[dict[str, object]]:
|
||||
wxo: Final = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params)
|
||||
|
||||
client: Final = WatsonxOrchestrateHandler._http_client(timeout=120.0)
|
||||
client: Final = WatsonxOrchestrateHandler._http_client(
|
||||
timeout=timeout if timeout is not None else 120.0
|
||||
)
|
||||
token: Final = await WatsonxOrchestrateHandler._get_bearer_token(
|
||||
cp4d_host=wxo.cp4d_host,
|
||||
auth_mode=wxo.auth_mode,
|
||||
|
|
@ -396,6 +402,7 @@ class WatsonxOrchestrateHandler:
|
|||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
static_headers=static_headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
response_text: Final = WatsonxOrchestrateTransformation.extract_text_from_a2a_message_response(result)
|
||||
async for chunk in WatsonxOrchestrateTransformation.fake_streaming_from_text(
|
||||
|
|
|
|||
|
|
@ -3,12 +3,12 @@ A2A Streaming Response Iterator
|
|||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
from typing import Any, 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
|
||||
from litellm.types.utils import Delta, GenericStreamingChunk, ModelResponseStream, StreamingChoices
|
||||
|
||||
from ..common_utils import A2AError, extract_text_from_a2a_response
|
||||
|
||||
|
|
@ -90,6 +90,13 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
|
|||
)
|
||||
text: Final = "" if is_working_status else extract_text_from_a2a_response(chunk)
|
||||
provider_fields: dict[str, object] = {}
|
||||
provider_fields.update(
|
||||
{
|
||||
key: value
|
||||
for key, value in chunk.items()
|
||||
if key in {"system_fingerprint", "service_tier"} and value is not None
|
||||
}
|
||||
)
|
||||
if isinstance(result, Mapping) and not is_working_status:
|
||||
control_fields = {
|
||||
"artifacts",
|
||||
|
|
@ -136,6 +143,48 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
|
|||
tool_calls: Final = self._get_tool_calls(chunk)
|
||||
usage: Final = self._get_usage(chunk)
|
||||
|
||||
if isinstance(result, Mapping):
|
||||
choices = result.get("choices")
|
||||
if isinstance(choices, list) and choices:
|
||||
streaming_choices: list[StreamingChoices] = []
|
||||
for choice_position, raw_choice in enumerate(choices):
|
||||
if not isinstance(raw_choice, Mapping):
|
||||
continue
|
||||
raw_index = raw_choice.get("index", choice_position)
|
||||
choice_index = raw_index if isinstance(raw_index, int) else choice_position
|
||||
delta_fields: dict[str, Any] = {}
|
||||
raw_delta = raw_choice.get("delta")
|
||||
if isinstance(raw_delta, Mapping):
|
||||
delta_fields.update(raw_delta)
|
||||
raw_message = raw_choice.get("message")
|
||||
if isinstance(raw_message, Mapping):
|
||||
message_text = extract_text_from_a2a_response({"result": {"message": raw_message}})
|
||||
if message_text and "content" not in delta_fields:
|
||||
delta_fields["content"] = message_text
|
||||
message_tool_calls = raw_message.get("tool_calls")
|
||||
if message_tool_calls and "tool_calls" not in delta_fields:
|
||||
delta_fields["tool_calls"] = message_tool_calls
|
||||
raw_finish_reason = raw_choice.get("finish_reason")
|
||||
choice_finish_reason = (
|
||||
raw_finish_reason
|
||||
if isinstance(raw_finish_reason, str) and raw_finish_reason
|
||||
else finish_reason
|
||||
)
|
||||
streaming_choices.append(
|
||||
StreamingChoices(
|
||||
index=choice_index,
|
||||
delta=Delta(**delta_fields),
|
||||
finish_reason=choice_finish_reason,
|
||||
logprobs=raw_choice.get("logprobs"),
|
||||
)
|
||||
)
|
||||
if streaming_choices:
|
||||
return ModelResponseStream(
|
||||
choices=streaming_choices,
|
||||
usage=usage,
|
||||
provider_specific_fields=provider_fields or None,
|
||||
)
|
||||
|
||||
# Return generic streaming chunk
|
||||
return GenericStreamingChunk(
|
||||
text=text,
|
||||
|
|
|
|||
|
|
@ -1872,12 +1872,25 @@ class ProxyBaseLLMRequestProcessing:
|
|||
trust_client_model_info=False,
|
||||
)
|
||||
|
||||
authorized_model = self.data.get("model")
|
||||
self.data = await proxy_logging_obj.pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
data=self.data,
|
||||
call_type=route_type,
|
||||
)
|
||||
|
||||
if self.data.get("model") != authorized_model:
|
||||
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)
|
||||
self.data = _check_and_merge_model_level_guardrails(
|
||||
data=self.data,
|
||||
llm_router=llm_router,
|
||||
trust_client_model_info=False,
|
||||
)
|
||||
|
||||
# Refresh AFTER pre_call_hook: guardrails (e.g. Presidio PII masking) may
|
||||
# have mutated `self.data` in place, and the audit-trail snapshot taken in
|
||||
# add_litellm_data_to_request predates that mutation.
|
||||
|
|
|
|||
|
|
@ -364,7 +364,7 @@ def test_build_wxo_headers_preserves_auth_headers():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_wxo_config_forwards_static_headers(monkeypatch):
|
||||
async def test_wxo_config_forwards_headers_and_timeout(monkeypatch):
|
||||
captured = {}
|
||||
|
||||
async def fake_handle_non_streaming(**kwargs):
|
||||
|
|
@ -377,10 +377,16 @@ async def test_wxo_config_forwards_static_headers(monkeypatch):
|
|||
request_id="req-1",
|
||||
params={},
|
||||
litellm_params={"model": "agent"},
|
||||
agent_extra_headers={"x-request-id": "request-1"},
|
||||
agent_static_headers={"x-tenant-id": "tenant-1"},
|
||||
timeout=12,
|
||||
)
|
||||
|
||||
assert captured["static_headers"] == {"x-tenant-id": "tenant-1"}
|
||||
assert captured["static_headers"] == {
|
||||
"x-request-id": "request-1",
|
||||
"x-tenant-id": "tenant-1",
|
||||
}
|
||||
assert captured["timeout"] == 12
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -240,6 +240,10 @@ async def test_handle_streaming_emits_proper_events():
|
|||
# Event 4: second artifact update
|
||||
assert events[3]["result"]["kind"] == "artifact-update"
|
||||
assert events[3]["result"]["artifact"]["parts"][0]["text"] == " world"
|
||||
assert (
|
||||
events[2]["result"]["artifact"]["artifactId"]
|
||||
== events[3]["result"]["artifact"]["artifactId"]
|
||||
)
|
||||
|
||||
# Event 5: status completed
|
||||
assert events[4]["result"]["kind"] == "status-update"
|
||||
|
|
@ -249,6 +253,72 @@ async def test_handle_streaming_emits_proper_events():
|
|||
assert events[4]["usage"]["total_tokens"] == 5
|
||||
|
||||
|
||||
def test_build_completion_params_keeps_bridge_routing_fields():
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
|
||||
A2ACompletionBridgeHandler,
|
||||
)
|
||||
|
||||
params = A2ACompletionBridgeHandler._build_completion_params(
|
||||
params={"message": {"role": "user", "parts": []}},
|
||||
litellm_params={
|
||||
"custom_llm_provider": "openai",
|
||||
"model": "agent",
|
||||
"api_base": "https://untrusted.example",
|
||||
"stream": False,
|
||||
},
|
||||
api_base="https://configured.example",
|
||||
agent_extra_headers=None,
|
||||
stream=True,
|
||||
)
|
||||
|
||||
assert params["api_base"] == "https://configured.example"
|
||||
assert params["stream"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_streaming_accumulates_logprobs_and_provider_metadata():
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
|
||||
A2ACompletionBridgeHandler,
|
||||
)
|
||||
|
||||
chunks = []
|
||||
for token in ("a", "b"):
|
||||
choice = MagicMock()
|
||||
choice.index = 0
|
||||
choice.finish_reason = None
|
||||
choice.delta.content = token
|
||||
choice.logprobs = {"content": [{"token": token}]}
|
||||
chunk = MagicMock()
|
||||
chunk.choices = [choice]
|
||||
chunk.system_fingerprint = "fp-1"
|
||||
chunk.service_tier = "scale"
|
||||
chunks.append(chunk)
|
||||
chunks[-1].choices[0].finish_reason = "stop"
|
||||
|
||||
async def mock_streaming_response():
|
||||
for chunk in chunks:
|
||||
yield 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-metadata",
|
||||
params={"message": {"role": "user", "parts": []}},
|
||||
litellm_params={"custom_llm_provider": "openai", "model": "agent"},
|
||||
)
|
||||
]
|
||||
|
||||
result = events[-1]
|
||||
assert result["system_fingerprint"] == "fp-1"
|
||||
assert result["service_tier"] == "scale"
|
||||
assert result["result"]["choices"][0]["logprobs"]["content"] == [
|
||||
{"token": "a"},
|
||||
{"token": "b"},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_streaming_preserves_multiple_choices():
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
|
||||
|
|
|
|||
|
|
@ -120,6 +120,37 @@ async def test_async_iterator_preserves_parallel_tool_calls():
|
|||
assert chunk["tool_use"] == tool_calls
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_iterator_preserves_every_terminal_choice():
|
||||
async def _events():
|
||||
yield {
|
||||
"jsonrpc": "2.0",
|
||||
"result": {
|
||||
"kind": "status-update",
|
||||
"status": {"state": "completed"},
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"parts": [{"kind": "text", "text": "first"}]},
|
||||
"finish_reason": "stop",
|
||||
},
|
||||
{
|
||||
"index": 1,
|
||||
"message": {"parts": [{"kind": "text", "text": "second"}]},
|
||||
"finish_reason": "length",
|
||||
},
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
iterator = A2AModelResponseIterator(streaming_response=_events(), sync_stream=False)
|
||||
chunk = await iterator.__aiter__().__anext__()
|
||||
|
||||
assert [choice.index for choice in chunk.choices] == [0, 1]
|
||||
assert [choice.delta.content for choice in chunk.choices] == ["first", "second"]
|
||||
assert [choice.finish_reason for choice in chunk.choices] == ["stop", "length"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_iterator_serializes_delta_tool_calls_and_usage():
|
||||
delta = Delta(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue