fix: address A2A streaming review feedback

This commit is contained in:
aiedwardyi 2026-08-26 12:05:32 +09:00
parent a5997bb908
commit d2bfb5ed1b
No known key found for this signature in database
9 changed files with 241 additions and 33 deletions

View file

@ -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(

View file

@ -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}],
}

View file

@ -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

View file

@ -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(

View file

@ -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,

View file

@ -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.

View file

@ -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

View file

@ -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 (

View file

@ -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(