diff --git a/litellm/a2a_protocol/litellm_completion_bridge/handler.py b/litellm/a2a_protocol/litellm_completion_bridge/handler.py index d4708e02786..04f013e86c6 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/handler.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/handler.py @@ -89,11 +89,12 @@ class A2ACompletionBridgeHandler: "api_base": api_base, "stream": stream, } + configured_headers: Final[object] = litellm_params.get("extra_headers") or litellm_params.get("headers") # Add litellm_params (contains api_key, client_id, client_secret, tenant_id, etc.) litellm_params_to_add: Final = { k: v for k, v in litellm_params.items() - if k not in ("model", "custom_llm_provider") and k not in _AGENT_ONLY_PARAMS + if k not in ("model", "custom_llm_provider", "extra_headers", "headers") and k not in _AGENT_ONLY_PARAMS } completion_params.update(litellm_params_to_add) # Apply forward metadata AFTER the litellm_params merge so the helper @@ -105,10 +106,10 @@ class A2ACompletionBridgeHandler: params=params, ) - if agent_extra_headers: + if agent_extra_headers or configured_headers: completion_params["extra_headers"] = merge_agent_headers( dynamic_headers=agent_extra_headers, - static_headers=completion_params.get("extra_headers"), + static_headers=configured_headers if isinstance(configured_headers, Mapping) else None, ) return completion_params @@ -361,6 +362,7 @@ class A2ACompletionBridgeHandler: artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event( ctx=ctx, text=content, + index=choice_index, ) yield artifact_event finally: @@ -380,13 +382,10 @@ class A2ACompletionBridgeHandler: completed_event["result"]["finish_reason"] = stream_finish_reason if stream_usage is not None: completed_event["usage"] = stream_usage - if choice_delta_fields.get(0): - completed_event["result"].update(choice_delta_fields[0]) - if 0 in choice_logprobs: - completed_event["result"]["logprobs"] = choice_logprobs[0] if len(choice_texts) > 1: - completed_event["result"]["choices"] = [ - { + choice_payloads: list[dict[str, object]] = [] + for choice_index in sorted(choice_texts): + choice_payload: dict[str, object] = { "index": choice_index, "message": { "kind": "message", @@ -397,7 +396,6 @@ class A2ACompletionBridgeHandler: if choice_tool_calls.get(choice_index) else {} ), - **choice_delta_fields.get(choice_index, {}), }, **( {"finish_reason": choice_finish_reasons[choice_index]} @@ -406,8 +404,29 @@ class A2ACompletionBridgeHandler: ), **({"logprobs": choice_logprobs[choice_index]} if choice_index in choice_logprobs else {}), } - for choice_index in sorted(choice_texts) - ] + if choice_delta_fields.get(choice_index): + choice_payload["delta"] = choice_delta_fields[choice_index] + choice_payloads.append(choice_payload) + completed_event["result"]["choices"] = choice_payloads + else: + metadata_indices = sorted(set(choice_delta_fields) | set(choice_logprobs)) + if metadata_indices: + completed_event["result"]["choices"] = [ + { + "index": choice_index, + **( + {"delta": choice_delta_fields[choice_index]} + if choice_delta_fields.get(choice_index) + else {} + ), + **( + {"logprobs": choice_logprobs[choice_index]} + if choice_index in choice_logprobs + else {} + ), + } + for choice_index in metadata_indices + ] 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 02006e98fa0..6e11acb71a6 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/transformation.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/transformation.py @@ -262,6 +262,10 @@ class A2ACompletionBridgeTransformation: } if usage is not None: a2a_response["usage"] = usage.model_dump(exclude_none=True) if hasattr(usage, "model_dump") else usage + for field in ("system_fingerprint", "service_tier"): + value = getattr(response, field, None) + if value is not None: + a2a_response[field] = value if len(serialized_choices) > 1: a2a_response["choices"] = serialized_choices @@ -354,6 +358,7 @@ class A2ACompletionBridgeTransformation: def create_artifact_update_event( ctx: A2AStreamingContext, text: str, + index: int | None = None, ) -> dict[str, Any]: """ Create an artifact update event with content. @@ -362,15 +367,18 @@ class A2ACompletionBridgeTransformation: ctx: Streaming context text: The text content for the artifact """ + artifact: Final[dict[str, Any]] = { + "artifactId": str(uuid4()), + "name": "response", + "parts": [{"kind": "text", "text": text}], + } + if index is not None: + artifact["index"] = index return { "id": ctx.request_id, "jsonrpc": "2.0", "result": { - "artifact": { - "artifactId": str(uuid4()), - "name": "response", - "parts": [{"kind": "text", "text": text}], - }, + "artifact": artifact, "contextId": ctx.context_id, "kind": "artifact-update", "taskId": ctx.task_id, diff --git a/litellm/llms/a2a/chat/streaming_iterator.py b/litellm/llms/a2a/chat/streaming_iterator.py index 1a598bb7a50..605c8cbd375 100644 --- a/litellm/llms/a2a/chat/streaming_iterator.py +++ b/litellm/llms/a2a/chat/streaming_iterator.py @@ -71,6 +71,16 @@ class A2AModelResponseIterator(BaseModelResponseIterator): try: # Extract text from A2A response result: Final = chunk.get("result", {}) + chunk_index = 0 + if isinstance(result, Mapping): + artifact = result.get("artifact") + if isinstance(artifact, Mapping) and isinstance(artifact.get("index"), int): + chunk_index = artifact["index"] + choices = result.get("choices") + if isinstance(choices, list) and choices and isinstance(choices[0], Mapping): + raw_index = choices[0].get("index") + if isinstance(raw_index, int): + chunk_index = raw_index status: Final = result.get("status", {}) if isinstance(result, Mapping) else {} is_working_status: Final = ( isinstance(result, Mapping) @@ -132,7 +142,7 @@ class A2AModelResponseIterator(BaseModelResponseIterator): is_finished=bool(finish_reason or tool_calls), finish_reason=finish_reason or ("tool_calls" if tool_calls else ""), usage=usage, - index=0, + index=chunk_index, tool_use=tool_calls, provider_specific_fields=provider_fields or None, ) diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index bd02cfdf907..ef36296fff5 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -860,7 +860,7 @@ async def invoke_agent_a2a( _enqueue_fn: Final = getattr(logging_obj, "_enqueue_deferred_logging", None) if _enqueue_fn is not None: logging_obj._enqueue_deferred_logging = None - _enqueue_fn() + _enqueue_fn(response) response_dict: Final[dict[str, Any]] = ( response.model_dump(mode="json", exclude_none=True) diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index ff8381c6531..b1af6121839 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -37,6 +37,7 @@ _FORWARDED_REQUEST_PARAMS: Final = frozenset( "frequency_penalty", "functions", "function_call", + "guided_json", "include_server_side_tool_invocations", "logit_bias", "logprobs", @@ -305,6 +306,10 @@ async def _route_registered_provider( id=str(response.get("id") or request_id), model=model_name, choices=model_choices, + system_fingerprint=response.get("system_fingerprint") + if isinstance(response.get("system_fingerprint"), str) + else None, + service_tier=response.get("service_tier") if isinstance(response.get("service_tier"), str) else None, ) raw_usage: Final = response.get("usage") usage: Final = litellm.Usage(**raw_usage) if isinstance(raw_usage, Mapping) else raw_usage @@ -315,10 +320,10 @@ async def _route_registered_provider( if isinstance(logging_obj, Logging): - def _enqueue_logging() -> None: + def _enqueue_logging(final_response: ModelResponse | None = None) -> None: asyncio.create_task( logging_obj.dispatch_success_handlers( - model_response, + final_response if final_response is not None else model_response, cache_hit=False, prefer_async_handlers=True, ) @@ -544,7 +549,7 @@ async def route_a2a_agent_request( 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 + api_base: Final = configured_api_base if isinstance(configured_api_base, str) and configured_api_base else agent_url registered_model: Final = registered_params_value.get("model") if registered_params_value else None cardless_provider: Final = registered_provider == "watsonx_orchestrate" or ( registered_provider == "bedrock" and isinstance(registered_model, str) and "agentcore" in registered_model diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 5927214a202..04432b2ca21 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -2544,6 +2544,7 @@ class ProxyBaseLLMRequestProcessing: ProxyBaseLLMRequestProcessing._flush_deferred_async_logging( logging_obj=logging_obj, exception_raised=_exception_raised, + response=response, ) # Streaming cleanup: if an exception occurred AND the deferred @@ -3027,6 +3028,7 @@ class ProxyBaseLLMRequestProcessing: def _flush_deferred_async_logging( logging_obj: Any, exception_raised: bool, + response: Any | None = None, ) -> None: """ Fire the deferred async-success closure stored by wrapper_async, then @@ -3057,7 +3059,7 @@ class ProxyBaseLLMRequestProcessing: if exception_raised: return try: - _enqueue_fn() + _enqueue_fn(response) if response is not None else _enqueue_fn() except Exception as e: verbose_proxy_logger.exception("Error firing deferred logging: %s", e) diff --git a/litellm/utils.py b/litellm/utils.py index e5ce7157e77..fb636b094f1 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1818,11 +1818,11 @@ def client(original_function): if not _is_litellm_internal_call: if getattr(logging_obj, "_defer_async_logging", False): - def _enqueue_deferred_logging() -> None: + def _enqueue_deferred_logging(final_response=None) -> None: asyncio.create_task( _client_async_logging_helper( logging_obj=logging_obj, - result=result, + result=final_response if final_response is not None else result, start_time=start_time, end_time=end_time, is_completion_with_fallbacks=is_completion_with_fallbacks, diff --git a/tests/test_litellm/a2a_protocol/providers/bedrock_agentcore/test_bedrock_agentcore_a2a.py b/tests/test_litellm/a2a_protocol/providers/bedrock_agentcore/test_bedrock_agentcore_a2a.py index 5503a5668bf..a959386aae9 100644 --- a/tests/test_litellm/a2a_protocol/providers/bedrock_agentcore/test_bedrock_agentcore_a2a.py +++ b/tests/test_litellm/a2a_protocol/providers/bedrock_agentcore/test_bedrock_agentcore_a2a.py @@ -10,10 +10,9 @@ Verifies that: """ import json - -import pytest from unittest.mock import AsyncMock, MagicMock, patch +import pytest SAMPLE_ARN = "arn:aws:bedrock-agentcore:us-west-2:123456789:runtime/my_agent" SAMPLE_MODEL = f"bedrock/agentcore/{SAMPLE_ARN}" @@ -482,6 +481,7 @@ class TestHandlerIntegration: api_base=None, litellm_params=SAMPLE_LITELLM_PARAMS, agent_extra_headers=None, + agent_static_headers=None, ) @pytest.mark.asyncio diff --git a/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py b/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py index 11f2e75107d..29d4e6ffe36 100644 --- a/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py +++ b/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py @@ -324,10 +324,13 @@ async def test_handle_streaming_preserves_non_text_delta_fields(): ] result = events[-1]["result"] - assert result["audio"] == {"data": "abc"} - assert result["reasoning_content"] == "thinking" - assert result["provider_specific_fields"] == {"trace_id": "trace-1"} - assert result["logprobs"] == {"content": []} + choice_result = result["choices"][0] + assert choice_result["delta"] == { + "audio": {"data": "abc"}, + "reasoning_content": "thinking", + "provider_specific_fields": {"trace_id": "trace-1"}, + } + assert choice_result["logprobs"] == {"content": []} @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index 2a34f4404bc..05bc118907c 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -62,7 +62,7 @@ async def test_route_a2a_model_bypasses_router(): "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", mock_registry, ): - result = await route_request( + await route_request( data=data, llm_router=mock_router, user_model=None, @@ -593,6 +593,7 @@ async def test_route_a2a_stream_uses_registered_provider(): litellm_params={"custom_llm_provider": "pydantic_ai_agents"}, ) logging_obj = Mock(spec=Logging) + logging_obj.model_call_details = {} data = { "model": "a2a/test-agent", "messages": [{"role": "user", "content": "Hello"}], @@ -745,7 +746,7 @@ def _router_without_models(): @pytest.mark.asyncio async def test_route_a2a_model_read_through_recovers_agent_created_on_sibling_replica(monkeypatch): - import litellm.proxy.proxy_server as proxy_server + from litellm.proxy import proxy_server from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry agent_name = "a2a-sibling-replica-agent"