fix: address A2A review feedback

This commit is contained in:
aiedwardyi 2026-08-24 22:47:40 +09:00
parent 32cb4c076b
commit d2f05d1cf1
No known key found for this signature in database
7 changed files with 228 additions and 87 deletions

View file

@ -155,14 +155,16 @@ class A2ACompletionBridgeHandler:
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"}
return await a2a_provider_config.handle_non_streaming(
request_id=request_id,
params=provider_params,
api_base=api_base,
timeout=litellm_params.get("timeout") or 60.0,
litellm_params=litellm_params,
agent_extra_headers=agent_extra_headers,
)
provider_kwargs: Final[dict[str, Any]] = {
"request_id": request_id,
"params": provider_params,
"api_base": api_base,
"litellm_params": litellm_params,
"agent_extra_headers": agent_extra_headers,
}
if litellm_params.get("timeout") is not None:
provider_kwargs["timeout"] = litellm_params["timeout"]
return await a2a_provider_config.handle_non_streaming(**provider_kwargs)
completion_params: Final = A2ACompletionBridgeHandler._build_completion_params(
params=params,
@ -226,14 +228,16 @@ class A2ACompletionBridgeHandler:
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"}
async for chunk in a2a_provider_config.handle_streaming(
request_id=request_id,
params=provider_params,
api_base=api_base,
timeout=litellm_params.get("timeout") or 60.0,
litellm_params=litellm_params,
agent_extra_headers=agent_extra_headers,
):
provider_kwargs: Final[dict[str, Any]] = {
"request_id": request_id,
"params": provider_params,
"api_base": api_base,
"litellm_params": litellm_params,
"agent_extra_headers": agent_extra_headers,
}
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
return
@ -268,8 +272,7 @@ class A2ACompletionBridgeHandler:
# Call litellm.acompletion with streaming
response: Final = await A2ACompletionBridgeHandler._acompletion(completion_params)
# 3. Accumulate content and emit artifact update
accumulated_text = ""
# 3. Forward content as artifact updates
accumulated_tool_calls: Final[list[object]] = [] # mutable-ok: collect streaming tool-call deltas
chunk_count = 0
async for chunk in response:
@ -286,15 +289,11 @@ class A2ACompletionBridgeHandler:
accumulated_tool_calls.extend(tool_calls)
if content:
accumulated_text += content
# Emit artifact update with accumulated content
if accumulated_text:
artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event(
ctx=ctx,
text=accumulated_text,
)
yield artifact_event
artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event(
ctx=ctx,
text=content,
)
yield artifact_event
# 4. Emit final status update (kind: "status-update", status: "completed", final: true)
completed_event: Final = A2ACompletionBridgeTransformation.create_status_update_event(

View file

@ -166,41 +166,44 @@ class A2ACompletionBridgeTransformation:
Returns:
A2A SendMessageResponse dict
"""
# Extract content from response
content = ""
if hasattr(response, "choices") and response.choices:
choice: Final = response.choices[0]
if hasattr(choice, "message") and choice.message:
content = choice.message.content or ""
serialized_choices: list[dict[str, Any]] = []
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 ""
)
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)
if raw_tool_calls:
message["tool_calls"] = [
call.model_dump(exclude_none=True)
if hasattr(call, "model_dump")
else call.dict(exclude_none=True)
if hasattr(call, "dict")
else call
for call in raw_tool_calls
]
finish_reason: Final = getattr(choice, "finish_reason", None)
if finish_reason:
message["finish_reason"] = finish_reason
serialized_choices.append({"index": len(serialized_choices), "message": message})
tool_calls: list[Any] | None = None
finish_reason: str | None = None
if hasattr(response, "choices") and response.choices:
choice = response.choices[0]
finish_reason = getattr(choice, "finish_reason", None)
message = getattr(choice, "message", None)
raw_tool_calls = getattr(message, "tool_calls", None)
if raw_tool_calls:
tool_calls = [
call.model_dump(exclude_none=True)
if hasattr(call, "model_dump")
else call.dict(exclude_none=True)
if hasattr(call, "dict")
else call
for call in raw_tool_calls
]
# Build A2A message
a2a_message: Final = {
"kind": "message",
"role": "agent",
"parts": [{"kind": "text", "text": content}],
"messageId": uuid4().hex,
}
if tool_calls:
a2a_message["tool_calls"] = tool_calls
if finish_reason:
a2a_message["finish_reason"] = finish_reason
a2a_message: Final = (
serialized_choices[0]["message"]
if serialized_choices
else {
"kind": "message",
"role": "agent",
"parts": [{"kind": "text", "text": ""}],
"messageId": uuid4().hex,
}
)
usage: Final = getattr(response, "usage", None)
@ -212,8 +215,10 @@ class A2ACompletionBridgeTransformation:
}
if usage is not None:
a2a_response["usage"] = usage.model_dump(exclude_none=True) if hasattr(usage, "model_dump") else usage
if len(serialized_choices) > 1:
a2a_response["choices"] = serialized_choices
verbose_logger.debug("OpenAI -> A2A transform: content_length=%s", len(content))
verbose_logger.debug("OpenAI -> A2A transform: content_length=%s", len(a2a_message["parts"][0]["text"]))
return a2a_response

View file

@ -5,6 +5,7 @@ A2A Streaming Response Iterator
from typing import Final
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 ..common_utils import A2AError, extract_text_from_a2a_response
@ -118,14 +119,16 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
return None
def _get_tool_calls(self, chunk: dict) -> list[dict] | None:
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):
return 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
message = result.get("message")
if isinstance(message, dict) and isinstance(message.get("tool_calls"), list):
return message["tool_calls"]
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 None

View file

@ -213,14 +213,53 @@ async def _route_registered_provider(
result_dict: Final = result if isinstance(result, Mapping) else {}
nested_message: Final = result_dict.get("message")
response_message: Final = nested_message if isinstance(nested_message, Mapping) else result_dict
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_response: Final = ModelResponse(
id=str(response.get("id") or request_id),
model=model_name,
choices=[ # mutable-ok: ModelResponse requires a choices list
response_choices: Final = response.get("choices")
choice_payloads: Final = (
response_choices
if isinstance(response_choices, list)
else result_dict.get("choices")
)
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)
else "stop"
),
index=choice.get("index", choice_index)
if isinstance(choice, Mapping) and isinstance(choice.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
),
),
)
for choice_index, choice in enumerate(choice_payloads)
]
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
@ -232,12 +271,16 @@ async def _route_registered_provider(
index=0,
message=Message(content=text, role="assistant", tool_calls=normalized_tool_calls),
)
],
]
model_response: Final = ModelResponse(
id=str(response.get("id") or request_id),
model=model_name,
choices=model_choices,
)
raw_usage: Final = response.get("usage")
usage: Final = litellm.Usage(**raw_usage) if isinstance(raw_usage, Mapping) else raw_usage
if usage is not None:
setattr(model_response, "usage", usage)
model_response.usage = usage
if isinstance(logging_obj, Logging):
logging_obj.model_call_details["usage"] = usage
@ -477,7 +520,8 @@ async def route_a2a_agent_request(
cardless_provider: Final = registered_provider == "watsonx_orchestrate" or (
registered_provider == "bedrock" and isinstance(registered_model, str) and "agentcore" in registered_model
)
if (not isinstance(agent_url, str) or not agent_url) and not cardless_provider:
has_configured_api_base: Final = isinstance(configured_api_base, str) and bool(configured_api_base)
if (not isinstance(agent_url, str) or not agent_url) and not has_configured_api_base and not cardless_provider:
verbose_proxy_logger.error("[A2A] Agent '%s' has no URL configured", agent_name)
route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type)
raise ProxyModelNotFoundError(route=route_name, model_name=model_name, retryable_with_model_read_through=False)

View file

@ -222,8 +222,8 @@ async def test_handle_streaming_emits_proper_events():
):
events.append(event)
# Should have 4 events: task, working, artifact, completed
assert len(events) == 4
# Should have 5 events: task, working, two artifacts, completed
assert len(events) == 5
# Event 1: task submitted
assert events[0]["result"]["kind"] == "task"
@ -234,14 +234,18 @@ async def test_handle_streaming_emits_proper_events():
assert events[1]["result"]["status"]["state"] == "working"
assert events[1]["result"]["final"] is False
# Event 3: artifact update with accumulated content
# Event 3: first artifact update
assert events[2]["result"]["kind"] == "artifact-update"
assert events[2]["result"]["artifact"]["parts"][0]["text"] == "Hello world"
assert events[2]["result"]["artifact"]["parts"][0]["text"] == "Hello"
# Event 4: status completed
assert events[3]["result"]["kind"] == "status-update"
assert events[3]["result"]["status"]["state"] == "completed"
assert events[3]["result"]["final"] is True
# Event 4: second artifact update
assert events[3]["result"]["kind"] == "artifact-update"
assert events[3]["result"]["artifact"]["parts"][0]["text"] == " world"
# Event 5: status completed
assert events[4]["result"]["kind"] == "status-update"
assert events[4]["result"]["status"]["state"] == "completed"
assert events[4]["result"]["final"] is True
@pytest.mark.asyncio

View file

@ -44,7 +44,7 @@ async def test_async_iterator_preserves_tool_calls():
chunk = await iterator.__aiter__().__anext__()
assert chunk["tool_use"] == tool_calls
assert chunk["tool_use"] == tool_calls[0]
assert chunk["finish_reason"] == "tool_calls"

View file

@ -194,6 +194,92 @@ async def test_route_a2a_cardless_bedrock_agentcore_uses_registered_model():
assert bridge.await_args.kwargs["api_base"] is None
@pytest.mark.asyncio
async def test_route_a2a_registered_provider_uses_configured_api_base_without_card_url():
from litellm.types.agents import AgentResponse
agent = AgentResponse(
agent_id="test-agent-id",
agent_name="test-agent",
agent_card_params={},
litellm_params={
"custom_llm_provider": "langflow",
"model": "flow",
"api_base": "https://flow.example.com",
},
)
bridge_response = {
"jsonrpc": "2.0",
"id": "request-id",
"result": {"kind": "message", "parts": [{"kind": "text", "text": "Hello back"}]},
}
with (
patch(
"litellm.proxy.common_utils.registry_read_through.get_agent_with_read_through",
AsyncMock(return_value=agent),
),
patch(
"litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed",
AsyncMock(return_value=True),
),
patch(
"litellm.a2a_protocol.litellm_completion_bridge.handler.A2ACompletionBridgeHandler.handle_non_streaming",
AsyncMock(return_value=bridge_response),
) as bridge,
):
call = await route_a2a_agent_request(
{"model": "a2a/test-agent", "messages": [{"role": "user", "content": "Hello"}]},
"acompletion",
)
await call
assert bridge.await_args.kwargs["api_base"] == "https://flow.example.com"
@pytest.mark.asyncio
async def test_registered_provider_response_preserves_multiple_choices():
from litellm.types.agents import AgentResponse
agent = AgentResponse(
agent_id="test-agent-id",
agent_name="test-agent",
agent_card_params={"url": "http://agent.example.com"},
litellm_params={"custom_llm_provider": "pydantic_ai_agents"},
)
bridge_response = {
"jsonrpc": "2.0",
"id": "request-id",
"choices": [
{"index": 0, "message": {"parts": [{"kind": "text", "text": "first"}]}},
{"index": 1, "message": {"parts": [{"kind": "text", "text": "second"}]}},
],
"result": {},
}
with (
patch(
"litellm.proxy.common_utils.registry_read_through.get_agent_with_read_through",
AsyncMock(return_value=agent),
),
patch(
"litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed",
AsyncMock(return_value=True),
),
patch(
"litellm.a2a_protocol.litellm_completion_bridge.handler.A2ACompletionBridgeHandler.handle_non_streaming",
AsyncMock(return_value=bridge_response),
),
):
call = await route_a2a_agent_request(
{"model": "a2a/test-agent", "messages": [{"role": "user", "content": "Hello"}]},
"acompletion",
)
response = await call
assert [choice.message.content for choice in response.choices] == ["first", "second"]
@pytest.mark.asyncio
async def test_route_a2a_cardless_watsonx_orchestrate_uses_registered_model():
from litellm.types.agents import AgentResponse