From 07cea813c6a59bfe04f533e0c38679ae5b3a4490 Mon Sep 17 00:00:00 2001 From: aiedwardyi <41576951+aiedwardyi@users.noreply.github.com> Date: Mon, 24 Aug 2026 17:26:08 +0900 Subject: [PATCH] fix: harden registered A2A routing --- .../transformation.py | 21 ++ litellm/proxy/agent_endpoints/a2a_routing.py | 129 +++++++-- litellm/proxy/common_request_processing.py | 6 + .../proxy/test_route_a2a_models.py | 249 +++++++++++++++++- 4 files changed, 384 insertions(+), 21 deletions(-) diff --git a/litellm/a2a_protocol/litellm_completion_bridge/transformation.py b/litellm/a2a_protocol/litellm_completion_bridge/transformation.py index 15cf77708f9..1ac90b3d294 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/transformation.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/transformation.py @@ -173,6 +173,23 @@ class A2ACompletionBridgeTransformation: if hasattr(choice, "message") and choice.message: content = choice.message.content or "" + 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", @@ -180,6 +197,10 @@ class A2ACompletionBridgeTransformation: "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 # Build A2A response a2a_response: Final = { diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index bd8b222ce4f..6256d9acccc 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -23,7 +23,7 @@ from litellm.interactions.agents.utils import merge_agent_headers from litellm.llms.a2a.common_utils import A2AError, convert_messages_to_prompt, extract_text_from_a2a_response from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.types.llms.openai import AllMessageValues -from litellm.types.utils import Choices, Message, ModelResponse +from litellm.types.utils import Choices, CustomPricingLiteLLMParams, Message, ModelResponse if TYPE_CHECKING: from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper @@ -66,6 +66,24 @@ _FORWARDED_REQUEST_PARAMS: Final = frozenset( "web_search_options", } ) +_A2A_PRICING_PARAMS: Final = frozenset({"cost_per_query", "response_cost"}) | frozenset( + CustomPricingLiteLLMParams.model_fields +) + + +def _get_agent_request_headers(data: Mapping[str, object]) -> dict[str, str]: + proxy_request: Final = data.get("proxy_server_request") + raw_headers: object = proxy_request.get("headers") if isinstance(proxy_request, Mapping) else None + if not isinstance(raw_headers, Mapping): + metadata: Final = data.get("metadata") + if not isinstance(metadata, Mapping): + 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 {} + ) class _A2ATextPart(TypedDict): @@ -132,6 +150,18 @@ async def _route_registered_provider( if agent_extra_headers: provider_params["extra_headers"] = agent_extra_headers + logging_obj: Final = data.get("litellm_logging_obj") + if isinstance(logging_obj, Logging): + pricing_params = { + key: litellm_params[key] + for key in _A2A_PRICING_PARAMS + if key in litellm_params and litellm_params[key] is not None + } + if pricing_params: + logging_obj.litellm_params.update(pricing_params) + logging_obj.model_call_details["litellm_params"].update(pricing_params) + logging_obj.custom_pricing = True + if stream: streaming_response: Final = A2ACompletionBridgeHandler.handle_streaming( request_id=request_id, @@ -145,7 +175,6 @@ async def _route_registered_provider( sync_stream=False, model=model_name, ) - logging_obj: Final = data.get("litellm_logging_obj") if not isinstance(logging_obj, Logging): raise TypeError("litellm_logging_obj is required for streaming A2A requests") return CustomStreamWrapper( @@ -172,19 +201,35 @@ async def _route_registered_provider( message=f"A2A error: {error_message if isinstance(error_message, str) else 'Unknown error'}", ) + result: Final = response.get("result") + 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 - Choices(finish_reason="stop", index=0, message=Message(content=text, role="assistant")) + Choices( + finish_reason=( + finish_reason + if isinstance(finish_reason, str) + else "tool_calls" + if normalized_tool_calls + else "stop" + ), + index=0, + message=Message(content=text, role="assistant", tool_calls=normalized_tool_calls), + ) ], ) usage: Final = response.get("usage") if usage is not None: setattr(model_response, "usage", usage) - logging_obj: Final = data.get("litellm_logging_obj") if isinstance(logging_obj, Logging): def _enqueue_logging() -> None: @@ -211,7 +256,8 @@ def _merge_agent_guardrails( configured_guardrails: list[object] = ( agent_guardrails if isinstance(agent_guardrails, list) else [agent_guardrails] ) - metadata = data.get("metadata") + 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 root_guardrails = data.get("guardrails") existing_guardrails: list[object] = [] @@ -227,36 +273,42 @@ def _merge_agent_guardrails( if isinstance(data, dict): data["guardrails"] = merged_guardrails if isinstance(metadata, dict): - metadata["guardrails"] = merged_guardrails + data[metadata_key] = {**metadata, "guardrails": merged_guardrails} return data merged_data = dict(data) merged_data["guardrails"] = merged_guardrails if isinstance(metadata, dict): - merged_data["metadata"] = {**metadata, "guardrails": merged_guardrails} + merged_data[metadata_key] = {**metadata, "guardrails": merged_guardrails} return merged_data +async def merge_a2a_agent_guardrails_before_hooks(data: Mapping[str, object]) -> 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.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 or not agent.litellm_params: + return data + return _merge_agent_guardrails(data, agent.litellm_params.get("guardrails")) + + def _get_agent_dynamic_headers( data: Mapping[str, object], agent_id: str, agent_name: str, extra_headers: list[str] | None, ) -> dict[str, str]: - proxy_request: Final = data.get("proxy_server_request") - raw_headers: object = proxy_request.get("headers") if isinstance(proxy_request, Mapping) else None - if not isinstance(raw_headers, Mapping): - metadata: Final = data.get("metadata") - raw_headers = metadata.get("headers") if isinstance(metadata, Mapping) else None - normalized_headers: Final = ( - {str(key).lower(): str(value) for key, value in raw_headers.items()} - if isinstance(raw_headers, Mapping) - else {} - ) + normalized_headers: Final = _get_agent_request_headers(data) dynamic_headers: dict[str, str] = {} for header_name in extra_headers or []: header_name_str: Final = str(header_name) + if header_name_str.lower().startswith("x-litellm-"): + continue value: Final = normalized_headers.get(header_name_str.lower()) if value is not None: dynamic_headers[header_name_str] = value @@ -266,11 +318,32 @@ def _get_agent_dynamic_headers( for key, value in normalized_headers.items(): if key.startswith(prefix): header_name: Final = key[len(prefix) :] - if header_name: + if header_name and not header_name.lower().startswith("x-litellm-"): dynamic_headers[header_name] = value return dynamic_headers +def _get_agent_identity_headers(user_api_key_dict: UserAPIKeyAuth | None) -> dict[str, str]: + if user_api_key_dict is None: + return {} + headers: dict[str, str] = {} + if user_api_key_dict.user_id: + headers["X-LiteLLM-User-Id"] = user_api_key_dict.user_id + if user_api_key_dict.team_id: + headers["X-LiteLLM-Team-Id"] = user_api_key_dict.team_id + return headers + + +def _enforce_inbound_trace_id(data: Mapping[str, object], agent_id: str) -> None: + from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers + + if not get_chain_id_from_headers(_get_agent_request_headers(data)): + raise HTTPException( + status_code=400, + detail=f"Agent '{agent_id}' requires x-litellm-trace-id header on all inbound requests.", + ) + + async def route_a2a_agent_request( data: Mapping[str, object], route_type: str, @@ -325,6 +398,9 @@ async def route_a2a_agent_request( detail=f"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) + # Get API base URL from agent config agent_card_params: Final = agent.agent_card_params agent_url: Final = agent_card_params.get("url") if agent_card_params else None @@ -336,7 +412,7 @@ async def route_a2a_agent_request( 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 - cardless_provider: Final = ( + 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: @@ -354,6 +430,19 @@ async def route_a2a_agent_request( agent_name=agent.agent_name, extra_headers=agent.extra_headers, ) + registered_static_headers: Mapping[str, str] | None = agent.static_headers + if registered_params_value and registered_params_value.get("databricks_oauth"): + from litellm.proxy.agent_endpoints.databricks_oauth import resolve_databricks_app_auth_header + + databricks_headers = await resolve_databricks_app_auth_header(dict(registered_params_value)) + registered_static_headers = merge_agent_headers( + dynamic_headers=registered_static_headers, + static_headers=databricks_headers, + ) + registered_static_headers = merge_agent_headers( + dynamic_headers=registered_static_headers, + static_headers=_get_agent_identity_headers(user_api_key_dict), + ) if ( registered_provider and registered_provider != "a2a" @@ -366,7 +455,7 @@ async def route_a2a_agent_request( model_name=model_name, api_base=api_base, litellm_params=registered_params_value, - static_headers=agent.static_headers, + static_headers=registered_static_headers, dynamic_headers=registered_dynamic_headers, ) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index dbbf9cb673e..f7b12357af4 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1845,6 +1845,12 @@ 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 # guardrails actually execute on pre_call. Without this, guardrails set # via litellm_params.guardrails are only honored on post_call paths diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index d96ee37dbdf..a30b7d422cf 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -7,8 +7,13 @@ Maps to: litellm/proxy/agent_endpoints/a2a_routing.py from unittest.mock import AsyncMock, Mock, patch import pytest +from fastapi import HTTPException -from litellm.proxy.agent_endpoints.a2a_routing import route_a2a_agent_request +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.agent_endpoints.a2a_routing import ( + merge_a2a_agent_guardrails_before_hooks, + route_a2a_agent_request, +) from litellm.proxy.route_llm_request import route_request @@ -189,6 +194,248 @@ 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_cardless_watsonx_orchestrate_uses_registered_model(): + 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": "watsonx_orchestrate", + "model": "agent", + "cp4d_host": "https://wxo.example.com", + "instance_id": "instance", + }, + ) + 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"] is None + + +@pytest.mark.asyncio +async def test_route_a2a_registered_provider_preserves_identity_headers(): + 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", + "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"}], + "proxy_server_request": { + "headers": { + "x-a2a-test-agent-x-litellm-user-id": "attacker", + "x-a2a-test-agent-x-litellm-team-id": "attacker-team", + } + }, + }, + "acompletion", + user_api_key_dict=UserAPIKeyAuth(user_id="trusted-user", team_id="trusted-team"), + ) + await call + + headers = bridge.await_args.kwargs["agent_extra_headers"] + assert headers["X-LiteLLM-User-Id"] == "trusted-user" + assert headers["X-LiteLLM-Team-Id"] == "trusted-team" + 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_requires_inbound_trace_id(): + 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", + "require_trace_id_on_calls_to_agent": True, + }, + ) + + 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), + ), + ): + with pytest.raises(HTTPException, match="requires x-litellm-trace-id"): + await route_a2a_agent_request( + {"model": "a2a/test-agent", "messages": [{"role": "user", "content": "Hello"}]}, + "acompletion", + ) + + +@pytest.mark.asyncio +async def test_route_a2a_resolves_databricks_oauth_headers(): + 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": "databricks", "databricks_oauth": {"client_id": "id"}}, + ) + 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.proxy.agent_endpoints.databricks_oauth.resolve_databricks_app_auth_header", + AsyncMock(return_value={"Authorization": "Bearer minted"}), + ), + 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["agent_extra_headers"]["Authorization"] == "Bearer minted" + + +@pytest.mark.asyncio +async def test_registered_provider_response_preserves_tool_calls(): + 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", + "result": { + "kind": "message", + "parts": [], + "tool_calls": [ + { + "id": "call-1", + "type": "function", + "function": {"name": "lookup", "arguments": "{}"}, + } + ], + "finish_reason": "tool_calls", + }, + } + + 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 response.choices[0].finish_reason == "tool_calls" + assert response.choices[0].message.tool_calls[0].id == "call-1" + + +@pytest.mark.asyncio +async def test_a2a_agent_guardrails_merge_before_hooks(): + 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={"guardrails": ["agent-guardrail"]}, + ) + with patch( + "litellm.proxy.common_utils.registry_read_through.get_agent_with_read_through", + AsyncMock(return_value=agent), + ): + merged = await merge_a2a_agent_guardrails_before_hooks( + {"model": "a2a/test-agent", "guardrails": ["request-guardrail"]} + ) + + assert merged["guardrails"] == ["request-guardrail", "agent-guardrail"] + + @pytest.mark.asyncio async def test_route_a2a_stream_uses_registered_provider(): from litellm.litellm_core_utils.litellm_logging import Logging