fix: close registered A2A review gaps

This commit is contained in:
aiedwardyi 2026-08-24 18:17:58 +09:00
parent 07cea813c6
commit 32cb4c076b
No known key found for this signature in database
8 changed files with 200 additions and 33 deletions

View file

@ -59,7 +59,12 @@ class A2ACompletionBridgeHandler:
message: Final = params.get("message", {})
# Transform A2A message to OpenAI format
openai_messages: Final = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
supplied_messages: Final = params.get("messages")
openai_messages: Final = (
supplied_messages
if isinstance(supplied_messages, list)
else A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(message)
)
# Get completion params
custom_llm_provider: Final = litellm_params.get("custom_llm_provider")
@ -149,10 +154,12 @@ class A2ACompletionBridgeHandler:
if a2a_provider_config is not None:
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=params,
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,
)
@ -218,10 +225,12 @@ class A2ACompletionBridgeHandler:
if a2a_provider_config is not None:
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=params,
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,
):
@ -261,6 +270,7 @@ class A2ACompletionBridgeHandler:
# 3. Accumulate content and emit artifact update
accumulated_text = ""
accumulated_tool_calls: Final[list[object]] = [] # mutable-ok: collect streaming tool-call deltas
chunk_count = 0
async for chunk in response:
chunk_count += 1
@ -271,6 +281,9 @@ class A2ACompletionBridgeHandler:
choice = chunk.choices[0]
if hasattr(choice, "delta") and choice.delta:
content = choice.delta.content or ""
tool_calls = getattr(choice.delta, "tool_calls", None)
if isinstance(tool_calls, (list, tuple)):
accumulated_tool_calls.extend(tool_calls)
if content:
accumulated_text += content
@ -289,6 +302,8 @@ class A2ACompletionBridgeHandler:
state="completed",
final=True,
)
if accumulated_tool_calls:
completed_event["result"]["tool_calls"] = accumulated_tool_calls
yield completed_event
verbose_logger.info(

View file

@ -202,12 +202,16 @@ class A2ACompletionBridgeTransformation:
if finish_reason:
a2a_message["finish_reason"] = finish_reason
usage: Final = getattr(response, "usage", None)
# Build A2A response
a2a_response: Final = {
"jsonrpc": "2.0",
"id": request_id,
"result": a2a_message,
}
if usage is not None:
a2a_response["usage"] = usage.model_dump(exclude_none=True) if hasattr(usage, "model_dump") else usage
verbose_logger.debug("OpenAI -> A2A transform: content_length=%s", len(content))

View file

@ -71,15 +71,16 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
# Determine finish reason
finish_reason: Final = self._get_finish_reason(chunk)
tool_calls: Final = self._get_tool_calls(chunk)
# Return generic streaming chunk
return GenericStreamingChunk(
text=text,
is_finished=bool(finish_reason),
finish_reason=finish_reason or "",
is_finished=bool(finish_reason or tool_calls),
finish_reason=finish_reason or ("tool_calls" if tool_calls else ""),
usage=None,
index=0,
tool_use=None,
tool_use=tool_calls,
)
except Exception:
# Return empty chunk on parse error
@ -92,9 +93,7 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
tool_use=None,
)
def _handle_string_chunk(
self, str_line: str | dict
) -> GenericStreamingChunk | ModelResponseStream:
def _handle_string_chunk(self, str_line: str | dict) -> GenericStreamingChunk | ModelResponseStream:
if isinstance(str_line, dict):
return self.chunk_parser(chunk=str_line)
return super()._handle_string_chunk(str_line=str_line)
@ -118,3 +117,15 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
return "stop"
return None
def _get_tool_calls(self, chunk: dict) -> list[dict] | 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
message = result.get("message")
if isinstance(message, dict) and isinstance(message.get("tool_calls"), list):
return message["tool_calls"]
return None

View file

@ -80,9 +80,7 @@ def _get_agent_request_headers(data: Mapping[str, object]) -> dict[str, str]:
metadata = data.get("litellm_metadata")
raw_headers = metadata.get("headers") if isinstance(metadata, Mapping) else None
return (
{str(key).lower(): str(value) for key, value in raw_headers.items()}
if isinstance(raw_headers, Mapping)
else {}
{str(key).lower(): str(value) for key, value in raw_headers.items()} if isinstance(raw_headers, Mapping) else {}
)
@ -95,10 +93,12 @@ class _A2AMessage(TypedDict):
role: ReadOnly[str]
parts: ReadOnly[tuple[_A2ATextPart, ...]]
messageId: ReadOnly[str]
contextId: ReadOnly[str | None]
class _A2AParams(TypedDict):
message: ReadOnly[_A2AMessage]
messages: ReadOnly[list[AllMessageValues]]
async def _route_registered_provider(
@ -120,20 +120,27 @@ async def _route_registered_provider(
messages: Final = _MESSAGES_ADAPTER.validate_python(raw_messages)
stream: Final = data.get("stream") is True
request_id: Final = str(uuid4())
raw_session_id: Final = data.get("litellm_session_id")
metadata: Final = data.get("metadata")
session_id: Final = (
raw_session_id
if isinstance(raw_session_id, str)
else metadata.get("session_id")
if isinstance(metadata, Mapping) and isinstance(metadata.get("session_id"), str)
else None
)
params: Final[_A2AParams] = {
"message": {
"role": "user",
"parts": ({"kind": "text", "text": convert_messages_to_prompt(messages)},),
"messageId": str(uuid4()),
}
"contextId": session_id,
},
"messages": messages,
}
provider_params: Final = {
**_OBJECT_DICT_ADAPTER.validate_python(litellm_params),
**{
key: data[key]
for key in _FORWARDED_REQUEST_PARAMS
if key in data and data[key] is not None
},
**{key: data[key] for key in _FORWARDED_REQUEST_PARAMS if key in data and data[key] is not None},
}
bridge_params: Final = _OBJECT_DICT_ADAPTER.validate_python(params)
configured_headers: Final = litellm_params.get("extra_headers") or litellm_params.get("headers")
@ -161,6 +168,7 @@ async def _route_registered_provider(
logging_obj.litellm_params.update(pricing_params)
logging_obj.model_call_details["litellm_params"].update(pricing_params)
logging_obj.custom_pricing = True
provider_params["no-log"] = True
if stream:
streaming_response: Final = A2ACompletionBridgeHandler.handle_streaming(
@ -226,9 +234,12 @@ async def _route_registered_provider(
)
],
)
usage: Final = response.get("usage")
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)
if isinstance(logging_obj, Logging):
logging_obj.model_call_details["usage"] = usage
if isinstance(logging_obj, Logging):
@ -253,9 +264,7 @@ def _merge_agent_guardrails(
if not agent_guardrails:
return data
configured_guardrails: list[object] = (
agent_guardrails if isinstance(agent_guardrails, list) else [agent_guardrails]
)
configured_guardrails: list[object] = agent_guardrails if isinstance(agent_guardrails, list) else [agent_guardrails]
metadata_key: Final = "litellm_metadata" if "litellm_metadata" in data else "metadata"
metadata = data.get(metadata_key)
metadata_guardrails = metadata.get("guardrails") if isinstance(metadata, dict) else None
@ -296,6 +305,49 @@ async def merge_a2a_agent_guardrails_before_hooks(data: Mapping[str, object]) ->
return _merge_agent_guardrails(data, agent.litellm_params.get("guardrails"))
async def authorize_a2a_agent_before_hooks(
data: Mapping[str, object],
user_api_key_dict: UserAPIKeyAuth | None,
) -> Mapping[str, object]:
model_name: Final = data.get("model")
if not isinstance(model_name, str) or not model_name.startswith("a2a/"):
return data
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import AgentRequestHandler
from litellm.proxy.common_utils.registry_read_through import get_agent_with_read_through
agent = await get_agent_with_read_through(model_name[4:])
if agent is None:
return data
is_admin: Final = user_api_key_dict is not None and (
user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
)
if not is_admin:
is_allowed: Final = await AgentRequestHandler.is_agent_allowed(
agent_id=agent.agent_id,
user_api_key_auth=user_api_key_dict,
)
if not is_allowed:
raise HTTPException(
status_code=403,
detail=f"Agent '{agent.agent_name}' is not allowed for your key/team. Contact proxy admin for access.",
)
if (agent.litellm_params or {}).get("require_trace_id_on_calls_to_agent"):
_enforce_inbound_trace_id(data, agent.agent_id)
if isinstance(data, dict):
data["agent_id"] = agent.agent_id
metadata = data.get("metadata")
if not isinstance(metadata, dict):
metadata = {}
data["metadata"] = metadata
metadata["agent_id"] = agent.agent_id
return data
def _get_agent_dynamic_headers(
data: Mapping[str, object],
agent_id: str,
@ -409,6 +461,16 @@ async def route_a2a_agent_request(
registered_params_value.get("custom_llm_provider") if registered_params_value else None
)
registered_provider: Final = registered_provider_value if isinstance(registered_provider_value, str) else None
from litellm.a2a_protocol.litellm_completion_bridge.handler import A2A_USER_API_KEY_HASH_PARAM
registered_params_for_route: Final[Mapping[str, object]] = (
{
**registered_params_value,
A2A_USER_API_KEY_HASH_PARAM: user_api_key_dict.api_key,
}
if registered_params_value and user_api_key_dict is not None and user_api_key_dict.api_key
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
registered_model: Final = registered_params_value.get("model") if registered_params_value else None
@ -454,7 +516,7 @@ async def route_a2a_agent_request(
data=routed_data,
model_name=model_name,
api_base=api_base,
litellm_params=registered_params_value,
litellm_params=registered_params_for_route,
static_headers=registered_static_headers,
dynamic_headers=registered_dynamic_headers,
)

View file

@ -37,9 +37,7 @@ async def append_agents_to_model_group(
if agent is not None:
agent_params: Final = agent.litellm_params
provider_value: Final = agent_params.get("custom_llm_provider") if agent_params else None
custom_llm_provider: Final = (
provider_value if isinstance(provider_value, str) else "a2a"
)
custom_llm_provider: Final = provider_value if isinstance(provider_value, str) else "a2a"
model_groups.append(
ModelGroupInfoProxy(
model_group=f"a2a/{agent.agent_name}",
@ -79,9 +77,7 @@ async def append_agents_to_model_info(
if agent is not None:
agent_params: Final = agent.litellm_params
provider_value: Final = agent_params.get("custom_llm_provider") if agent_params else None
custom_llm_provider: Final = (
provider_value if isinstance(provider_value, str) else "a2a"
)
custom_llm_provider: Final = provider_value if isinstance(provider_value, str) else "a2a"
models.append(
{
"model_name": f"a2a/{agent.agent_name}",

View file

@ -1836,6 +1836,16 @@ class ProxyBaseLLMRequestProcessing:
## LOGGING OBJECT ## - initialize logging object for logging success/failure events for call
## IMPORTANT Note: - initialize this before running pre-call checks. Ensures we log rejected requests to langfuse.
from litellm.proxy.agent_endpoints.a2a_routing import (
authorize_a2a_agent_before_hooks,
merge_a2a_agent_guardrails_before_hooks,
)
self.data = await authorize_a2a_agent_before_hooks(
data=self.data,
user_api_key_dict=user_api_key_dict,
)
logging_obj, self.data = litellm.utils.function_setup(
original_function=route_type,
rules_obj=litellm.utils.Rules(),
@ -1845,10 +1855,6 @@ class ProxyBaseLLMRequestProcessing:
self.data["litellm_logging_obj"] = logging_obj
from litellm.proxy.agent_endpoints.a2a_routing import (
merge_a2a_agent_guardrails_before_hooks,
)
self.data = await merge_a2a_agent_guardrails_before_hooks(self.data)
# Merge model-level guardrails before pre_call_hook so DB/UI-configured

View file

@ -27,6 +27,27 @@ async def test_async_iterator_accepts_decoded_a2a_events():
assert chunk["text"] == "Hello"
@pytest.mark.asyncio
async def test_async_iterator_preserves_tool_calls():
tool_calls = [
{
"id": "call-1",
"type": "function",
"function": {"name": "lookup", "arguments": "{}"},
}
]
async def _events():
yield {"jsonrpc": "2.0", "result": {"tool_calls": tool_calls}}
iterator = A2AModelResponseIterator(streaming_response=_events(), sync_stream=False)
chunk = await iterator.__aiter__().__anext__()
assert chunk["tool_use"] == tool_calls
assert chunk["finish_reason"] == "tool_calls"
@pytest.mark.asyncio
async def test_async_iterator_propagates_jsonrpc_errors():
async def _events():

View file

@ -290,6 +290,58 @@ async def test_route_a2a_registered_provider_preserves_identity_headers():
assert "x-litellm-user-id" not in {key.lower() for key in headers if key != "X-LiteLLM-User-Id"}
@pytest.mark.asyncio
async def test_route_a2a_registered_provider_preserves_messages_and_session():
from litellm.a2a_protocol.litellm_completion_bridge.handler import A2A_USER_API_KEY_HASH_PARAM
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": "langflow", "model": "flow"},
)
data = {
"model": "a2a/test-agent",
"messages": [
{"role": "system", "content": "Be concise"},
{"role": "user", "content": "Hello"},
],
"litellm_session_id": "session-1",
}
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(
data,
"acompletion",
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
)
await call
bridge_kwargs = bridge.await_args.kwargs
assert bridge_kwargs["params"]["messages"] == data["messages"]
assert bridge_kwargs["params"]["message"]["contextId"] == "session-1"
assert bridge_kwargs["litellm_params"][A2A_USER_API_KEY_HASH_PARAM] == "hashed-key"
@pytest.mark.asyncio
async def test_route_a2a_requires_inbound_trace_id():
from litellm.types.agents import AgentResponse