From 32cb4c076b740d97c89e944e1b4b652573b1fdc4 Mon Sep 17 00:00:00 2001 From: aiedwardyi <41576951+aiedwardyi@users.noreply.github.com> Date: Mon, 24 Aug 2026 18:17:58 +0900 Subject: [PATCH] fix: close registered A2A review gaps --- .../litellm_completion_bridge/handler.py | 21 ++++- .../transformation.py | 4 + litellm/llms/a2a/chat/streaming_iterator.py | 23 +++-- litellm/proxy/agent_endpoints/a2a_routing.py | 90 ++++++++++++++++--- .../agent_endpoints/model_list_helpers.py | 8 +- litellm/proxy/common_request_processing.py | 14 ++- .../a2a/chat/test_a2a_streaming_iterator.py | 21 +++++ .../proxy/test_route_a2a_models.py | 52 +++++++++++ 8 files changed, 200 insertions(+), 33 deletions(-) diff --git a/litellm/a2a_protocol/litellm_completion_bridge/handler.py b/litellm/a2a_protocol/litellm_completion_bridge/handler.py index a62a2b0c724..4a3f7d608ce 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/handler.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/handler.py @@ -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( diff --git a/litellm/a2a_protocol/litellm_completion_bridge/transformation.py b/litellm/a2a_protocol/litellm_completion_bridge/transformation.py index 1ac90b3d294..c87b1367377 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/transformation.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/transformation.py @@ -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)) diff --git a/litellm/llms/a2a/chat/streaming_iterator.py b/litellm/llms/a2a/chat/streaming_iterator.py index 1c3fa0c6c92..9da95139c4e 100644 --- a/litellm/llms/a2a/chat/streaming_iterator.py +++ b/litellm/llms/a2a/chat/streaming_iterator.py @@ -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 diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 6256d9acccc..69c3aebf4f6 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -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, ) diff --git a/litellm/proxy/agent_endpoints/model_list_helpers.py b/litellm/proxy/agent_endpoints/model_list_helpers.py index 67d721f9faa..3fa62d0275c 100644 --- a/litellm/proxy/agent_endpoints/model_list_helpers.py +++ b/litellm/proxy/agent_endpoints/model_list_helpers.py @@ -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}", diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index f7b12357af4..16a10a16e70 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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 diff --git a/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py b/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py index b301b2f3c6e..f436fe27a57 100644 --- a/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py +++ b/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py @@ -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(): diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index a30b7d422cf..8cbeaf6f6d1 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -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