mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix: address A2A review feedback
This commit is contained in:
parent
32cb4c076b
commit
d2f05d1cf1
7 changed files with 228 additions and 87 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue