fix: address A2A review feedback

This commit is contained in:
aiedwardyi 2026-08-26 11:29:13 +09:00
parent eba758223d
commit a5997bb908
No known key found for this signature in database
10 changed files with 81 additions and 33 deletions

View file

@ -89,11 +89,12 @@ class A2ACompletionBridgeHandler:
"api_base": api_base,
"stream": stream,
}
configured_headers: Final[object] = litellm_params.get("extra_headers") or litellm_params.get("headers")
# Add litellm_params (contains api_key, client_id, client_secret, tenant_id, etc.)
litellm_params_to_add: Final = {
k: v
for k, v in litellm_params.items()
if k not in ("model", "custom_llm_provider") and k not in _AGENT_ONLY_PARAMS
if k not in ("model", "custom_llm_provider", "extra_headers", "headers") 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
@ -105,10 +106,10 @@ class A2ACompletionBridgeHandler:
params=params,
)
if agent_extra_headers:
if agent_extra_headers or configured_headers:
completion_params["extra_headers"] = merge_agent_headers(
dynamic_headers=agent_extra_headers,
static_headers=completion_params.get("extra_headers"),
static_headers=configured_headers if isinstance(configured_headers, Mapping) else None,
)
return completion_params
@ -361,6 +362,7 @@ class A2ACompletionBridgeHandler:
artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event(
ctx=ctx,
text=content,
index=choice_index,
)
yield artifact_event
finally:
@ -380,13 +382,10 @@ class A2ACompletionBridgeHandler:
completed_event["result"]["finish_reason"] = stream_finish_reason
if stream_usage is not None:
completed_event["usage"] = stream_usage
if choice_delta_fields.get(0):
completed_event["result"].update(choice_delta_fields[0])
if 0 in choice_logprobs:
completed_event["result"]["logprobs"] = choice_logprobs[0]
if len(choice_texts) > 1:
completed_event["result"]["choices"] = [
{
choice_payloads: list[dict[str, object]] = []
for choice_index in sorted(choice_texts):
choice_payload: dict[str, object] = {
"index": choice_index,
"message": {
"kind": "message",
@ -397,7 +396,6 @@ class A2ACompletionBridgeHandler:
if choice_tool_calls.get(choice_index)
else {}
),
**choice_delta_fields.get(choice_index, {}),
},
**(
{"finish_reason": choice_finish_reasons[choice_index]}
@ -406,8 +404,29 @@ class A2ACompletionBridgeHandler:
),
**({"logprobs": choice_logprobs[choice_index]} if choice_index in choice_logprobs else {}),
}
for choice_index in sorted(choice_texts)
]
if choice_delta_fields.get(choice_index):
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

@ -262,6 +262,10 @@ class A2ACompletionBridgeTransformation:
}
if usage is not None:
a2a_response["usage"] = usage.model_dump(exclude_none=True) if hasattr(usage, "model_dump") else usage
for field in ("system_fingerprint", "service_tier"):
value = getattr(response, field, None)
if value is not None:
a2a_response[field] = value
if len(serialized_choices) > 1:
a2a_response["choices"] = serialized_choices
@ -354,6 +358,7 @@ class A2ACompletionBridgeTransformation:
def create_artifact_update_event(
ctx: A2AStreamingContext,
text: str,
index: int | None = None,
) -> dict[str, Any]:
"""
Create an artifact update event with content.
@ -362,15 +367,18 @@ class A2ACompletionBridgeTransformation:
ctx: Streaming context
text: The text content for the artifact
"""
artifact: Final[dict[str, Any]] = {
"artifactId": str(uuid4()),
"name": "response",
"parts": [{"kind": "text", "text": text}],
}
if index is not None:
artifact["index"] = index
return {
"id": ctx.request_id,
"jsonrpc": "2.0",
"result": {
"artifact": {
"artifactId": str(uuid4()),
"name": "response",
"parts": [{"kind": "text", "text": text}],
},
"artifact": artifact,
"contextId": ctx.context_id,
"kind": "artifact-update",
"taskId": ctx.task_id,

View file

@ -71,6 +71,16 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
try:
# Extract text from A2A response
result: Final = chunk.get("result", {})
chunk_index = 0
if isinstance(result, Mapping):
artifact = result.get("artifact")
if isinstance(artifact, Mapping) and isinstance(artifact.get("index"), int):
chunk_index = artifact["index"]
choices = result.get("choices")
if isinstance(choices, list) and choices and isinstance(choices[0], Mapping):
raw_index = choices[0].get("index")
if isinstance(raw_index, int):
chunk_index = raw_index
status: Final = result.get("status", {}) if isinstance(result, Mapping) else {}
is_working_status: Final = (
isinstance(result, Mapping)
@ -132,7 +142,7 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
is_finished=bool(finish_reason or tool_calls),
finish_reason=finish_reason or ("tool_calls" if tool_calls else ""),
usage=usage,
index=0,
index=chunk_index,
tool_use=tool_calls,
provider_specific_fields=provider_fields or None,
)

View file

@ -860,7 +860,7 @@ async def invoke_agent_a2a(
_enqueue_fn: Final = getattr(logging_obj, "_enqueue_deferred_logging", None)
if _enqueue_fn is not None:
logging_obj._enqueue_deferred_logging = None
_enqueue_fn()
_enqueue_fn(response)
response_dict: Final[dict[str, Any]] = (
response.model_dump(mode="json", exclude_none=True)

View file

@ -37,6 +37,7 @@ _FORWARDED_REQUEST_PARAMS: Final = frozenset(
"frequency_penalty",
"functions",
"function_call",
"guided_json",
"include_server_side_tool_invocations",
"logit_bias",
"logprobs",
@ -305,6 +306,10 @@ async def _route_registered_provider(
id=str(response.get("id") or request_id),
model=model_name,
choices=model_choices,
system_fingerprint=response.get("system_fingerprint")
if isinstance(response.get("system_fingerprint"), str)
else None,
service_tier=response.get("service_tier") if isinstance(response.get("service_tier"), str) else None,
)
raw_usage: Final = response.get("usage")
usage: Final = litellm.Usage(**raw_usage) if isinstance(raw_usage, Mapping) else raw_usage
@ -315,10 +320,10 @@ async def _route_registered_provider(
if isinstance(logging_obj, Logging):
def _enqueue_logging() -> None:
def _enqueue_logging(final_response: ModelResponse | None = None) -> None:
asyncio.create_task(
logging_obj.dispatch_success_handlers(
model_response,
final_response if final_response is not None else model_response,
cache_hit=False,
prefer_async_handlers=True,
)
@ -544,7 +549,7 @@ async def route_a2a_agent_request(
else registered_params_value or {}
)
configured_api_base: Final = registered_params_value.get("api_base") if registered_params_value else None
api_base: Final = configured_api_base if isinstance(configured_api_base, str) else agent_url
api_base: Final = configured_api_base if isinstance(configured_api_base, str) and configured_api_base else agent_url
registered_model: Final = registered_params_value.get("model") if registered_params_value else None
cardless_provider: Final = registered_provider == "watsonx_orchestrate" or (
registered_provider == "bedrock" and isinstance(registered_model, str) and "agentcore" in registered_model

View file

@ -2544,6 +2544,7 @@ class ProxyBaseLLMRequestProcessing:
ProxyBaseLLMRequestProcessing._flush_deferred_async_logging(
logging_obj=logging_obj,
exception_raised=_exception_raised,
response=response,
)
# Streaming cleanup: if an exception occurred AND the deferred
@ -3027,6 +3028,7 @@ class ProxyBaseLLMRequestProcessing:
def _flush_deferred_async_logging(
logging_obj: Any,
exception_raised: bool,
response: Any | None = None,
) -> None:
"""
Fire the deferred async-success closure stored by wrapper_async, then
@ -3057,7 +3059,7 @@ class ProxyBaseLLMRequestProcessing:
if exception_raised:
return
try:
_enqueue_fn()
_enqueue_fn(response) if response is not None else _enqueue_fn()
except Exception as e:
verbose_proxy_logger.exception("Error firing deferred logging: %s", e)

View file

@ -1818,11 +1818,11 @@ def client(original_function):
if not _is_litellm_internal_call:
if getattr(logging_obj, "_defer_async_logging", False):
def _enqueue_deferred_logging() -> None:
def _enqueue_deferred_logging(final_response=None) -> None:
asyncio.create_task(
_client_async_logging_helper(
logging_obj=logging_obj,
result=result,
result=final_response if final_response is not None else result,
start_time=start_time,
end_time=end_time,
is_completion_with_fallbacks=is_completion_with_fallbacks,

View file

@ -10,10 +10,9 @@ Verifies that:
"""
import json
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
SAMPLE_ARN = "arn:aws:bedrock-agentcore:us-west-2:123456789:runtime/my_agent"
SAMPLE_MODEL = f"bedrock/agentcore/{SAMPLE_ARN}"
@ -482,6 +481,7 @@ class TestHandlerIntegration:
api_base=None,
litellm_params=SAMPLE_LITELLM_PARAMS,
agent_extra_headers=None,
agent_static_headers=None,
)
@pytest.mark.asyncio

View file

@ -324,10 +324,13 @@ async def test_handle_streaming_preserves_non_text_delta_fields():
]
result = events[-1]["result"]
assert result["audio"] == {"data": "abc"}
assert result["reasoning_content"] == "thinking"
assert result["provider_specific_fields"] == {"trace_id": "trace-1"}
assert result["logprobs"] == {"content": []}
choice_result = result["choices"][0]
assert choice_result["delta"] == {
"audio": {"data": "abc"},
"reasoning_content": "thinking",
"provider_specific_fields": {"trace_id": "trace-1"},
}
assert choice_result["logprobs"] == {"content": []}
@pytest.mark.asyncio

View file

@ -62,7 +62,7 @@ async def test_route_a2a_model_bypasses_router():
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry",
mock_registry,
):
result = await route_request(
await route_request(
data=data,
llm_router=mock_router,
user_model=None,
@ -593,6 +593,7 @@ async def test_route_a2a_stream_uses_registered_provider():
litellm_params={"custom_llm_provider": "pydantic_ai_agents"},
)
logging_obj = Mock(spec=Logging)
logging_obj.model_call_details = {}
data = {
"model": "a2a/test-agent",
"messages": [{"role": "user", "content": "Hello"}],
@ -745,7 +746,7 @@ def _router_without_models():
@pytest.mark.asyncio
async def test_route_a2a_model_read_through_recovers_agent_created_on_sibling_replica(monkeypatch):
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy import proxy_server
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
agent_name = "a2a-sibling-replica-agent"