diff --git a/litellm/a2a_protocol/litellm_completion_bridge/handler.py b/litellm/a2a_protocol/litellm_completion_bridge/handler.py index 04f013e86c6..ffc94a15890 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/handler.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/handler.py @@ -46,6 +46,21 @@ class A2ACompletionBridgeHandler: Static methods for handling A2A requests via LiteLLM completion. """ + @staticmethod + def _merge_stream_values(previous: object, current: object) -> object: + if isinstance(previous, Mapping) and isinstance(current, Mapping): + merged = dict(previous) + for key, value in current.items(): + merged[key] = ( + A2ACompletionBridgeHandler._merge_stream_values(merged[key], value) + if key in merged + else value + ) + return merged + if isinstance(previous, list) and isinstance(current, list): + return [*previous, *current] + return current + @staticmethod def _build_completion_params( params: dict[str, Any], @@ -94,7 +109,8 @@ class A2ACompletionBridgeHandler: litellm_params_to_add: Final = { k: v for k, v in litellm_params.items() - if k not in ("model", "custom_llm_provider", "extra_headers", "headers") and k not in _AGENT_ONLY_PARAMS + if k not in ("model", "custom_llm_provider", "extra_headers", "headers", "api_base", "stream") + 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 @@ -290,8 +306,9 @@ class A2ACompletionBridgeHandler: choice_texts: dict[int, str] = {} choice_tool_calls: dict[int, list[object]] = {} choice_delta_fields: dict[int, dict[str, object]] = {} - choice_logprobs: dict[int, object] = {} + choice_logprobs: dict[int, dict[str, object]] = {} choice_finish_reasons: dict[int, str] = {} + stream_metadata: dict[str, str] = {} stream_usage: object | None = None stream_finish_reason: str | None = None chunk_count = 0 @@ -315,6 +332,14 @@ class A2ACompletionBridgeHandler: if isinstance(dumped_usage, Mapping): stream_usage = dumped_usage + for metadata_name in ("system_fingerprint", "service_tier"): + metadata_value = getattr(chunk, metadata_name, None) + if not isinstance(metadata_value, str): + chunk_fields = A2ACompletionBridgeTransformation._model_dump(chunk) + metadata_value = chunk_fields.get(metadata_name) + if isinstance(metadata_value, str) and metadata_value: + stream_metadata[metadata_name] = metadata_value + # Extract delta content choices = getattr(chunk, "choices", None) if chunk is not None else None if isinstance(choices, (list, tuple)): @@ -356,7 +381,12 @@ class A2ACompletionBridgeHandler: raw_logprobs = getattr(choice, "logprobs", None) serialized_logprobs = A2ACompletionBridgeTransformation._model_dump(raw_logprobs) if serialized_logprobs: - choice_logprobs[choice_index] = serialized_logprobs + previous_logprobs = choice_logprobs.get(choice_index, {}) + merged_logprobs = A2ACompletionBridgeHandler._merge_stream_values( + previous_logprobs, serialized_logprobs + ) + if isinstance(merged_logprobs, dict): + choice_logprobs[choice_index] = merged_logprobs if content: artifact_event: Final = A2ACompletionBridgeTransformation.create_artifact_update_event( @@ -382,9 +412,18 @@ class A2ACompletionBridgeHandler: completed_event["result"]["finish_reason"] = stream_finish_reason if stream_usage is not None: completed_event["usage"] = stream_usage - if len(choice_texts) > 1: + for metadata_name, metadata_value in stream_metadata.items(): + completed_event[metadata_name] = metadata_value + choice_indices = sorted( + set(choice_texts) + | set(choice_tool_calls) + | set(choice_delta_fields) + | set(choice_logprobs) + | set(choice_finish_reasons) + ) + if choice_indices: choice_payloads: list[dict[str, object]] = [] - for choice_index in sorted(choice_texts): + for choice_index in choice_indices: choice_payload: dict[str, object] = { "index": choice_index, "message": { @@ -408,25 +447,6 @@ class A2ACompletionBridgeHandler: 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 6e11acb71a6..644ea9ab75f 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/transformation.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/transformation.py @@ -34,6 +34,7 @@ class A2AStreamingContext: self.request_id = request_id self.task_id = str(uuid4()) self.context_id = str(uuid4()) + self.artifact_id = str(uuid4()) self.input_message = input_message self.accumulated_text = "" self.has_emitted_task = False @@ -368,7 +369,7 @@ class A2ACompletionBridgeTransformation: text: The text content for the artifact """ artifact: Final[dict[str, Any]] = { - "artifactId": str(uuid4()), + "artifactId": ctx.artifact_id, "name": "response", "parts": [{"kind": "text", "text": text}], } diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py index 1b8024ae075..59d6f6003b0 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py @@ -9,6 +9,7 @@ from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig from litellm.a2a_protocol.providers.watsonx_orchestrate.handler import ( WatsonxOrchestrateHandler, ) +from litellm.interactions.agents.utils import merge_agent_headers class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig): @@ -28,11 +29,16 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig): "litellm_params is required for WatsonxOrchestrateA2AConfig " "(must contain cp4d_host, instance_id, wxo_agent_id, api_key)" ) + forwarded_headers: Final = merge_agent_headers( + dynamic_headers=kwargs.get("agent_extra_headers"), + static_headers=kwargs.get("agent_static_headers"), + ) return await WatsonxOrchestrateHandler.handle_non_streaming( request_id=request_id, params=params, litellm_params=litellm_params, - static_headers=kwargs.get("agent_static_headers"), + static_headers=forwarded_headers, + timeout=kwargs.get("timeout"), ) async def handle_streaming( @@ -49,10 +55,15 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig): "litellm_params is required for WatsonxOrchestrateA2AConfig " "(must contain cp4d_host, instance_id, wxo_agent_id, api_key)" ) + forwarded_headers: Final = merge_agent_headers( + dynamic_headers=kwargs.get("agent_extra_headers"), + static_headers=kwargs.get("agent_static_headers"), + ) async for chunk in WatsonxOrchestrateHandler.handle_streaming( request_id=request_id, params=params, litellm_params=litellm_params, - static_headers=kwargs.get("agent_static_headers"), + static_headers=forwarded_headers, + timeout=kwargs.get("timeout"), ): yield chunk diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py index f648cca40cd..2d3e5cb1569 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py @@ -305,10 +305,13 @@ class WatsonxOrchestrateHandler: params: dict[str, object], litellm_params: WXOLitellmParams, static_headers: Mapping[str, str] | None = None, + timeout: float | None = None, ) -> dict[str, object]: wxo: Final = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params) - client: Final = WatsonxOrchestrateHandler._http_client(timeout=90.0) + client: Final = WatsonxOrchestrateHandler._http_client( + timeout=timeout if timeout is not None else 90.0 + ) token: Final = await WatsonxOrchestrateHandler._get_bearer_token( cp4d_host=wxo.cp4d_host, auth_mode=wxo.auth_mode, @@ -355,10 +358,13 @@ class WatsonxOrchestrateHandler: chunk_size: int = 50, delay_ms: int = 10, static_headers: Mapping[str, str] | None = None, + timeout: float | None = None, ) -> AsyncIterator[dict[str, object]]: wxo: Final = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params) - client: Final = WatsonxOrchestrateHandler._http_client(timeout=120.0) + client: Final = WatsonxOrchestrateHandler._http_client( + timeout=timeout if timeout is not None else 120.0 + ) token: Final = await WatsonxOrchestrateHandler._get_bearer_token( cp4d_host=wxo.cp4d_host, auth_mode=wxo.auth_mode, @@ -396,6 +402,7 @@ class WatsonxOrchestrateHandler: params=params, litellm_params=litellm_params, static_headers=static_headers, + timeout=timeout, ) response_text: Final = WatsonxOrchestrateTransformation.extract_text_from_a2a_message_response(result) async for chunk in WatsonxOrchestrateTransformation.fake_streaming_from_text( diff --git a/litellm/llms/a2a/chat/streaming_iterator.py b/litellm/llms/a2a/chat/streaming_iterator.py index 605c8cbd375..099a9f67047 100644 --- a/litellm/llms/a2a/chat/streaming_iterator.py +++ b/litellm/llms/a2a/chat/streaming_iterator.py @@ -3,12 +3,12 @@ A2A Streaming Response Iterator """ from collections.abc import Mapping -from typing import Final +from typing import Any, Final import litellm 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 litellm.types.utils import Delta, GenericStreamingChunk, ModelResponseStream, StreamingChoices from ..common_utils import A2AError, extract_text_from_a2a_response @@ -90,6 +90,13 @@ class A2AModelResponseIterator(BaseModelResponseIterator): ) text: Final = "" if is_working_status else extract_text_from_a2a_response(chunk) provider_fields: dict[str, object] = {} + provider_fields.update( + { + key: value + for key, value in chunk.items() + if key in {"system_fingerprint", "service_tier"} and value is not None + } + ) if isinstance(result, Mapping) and not is_working_status: control_fields = { "artifacts", @@ -136,6 +143,48 @@ class A2AModelResponseIterator(BaseModelResponseIterator): tool_calls: Final = self._get_tool_calls(chunk) usage: Final = self._get_usage(chunk) + if isinstance(result, Mapping): + choices = result.get("choices") + if isinstance(choices, list) and choices: + streaming_choices: list[StreamingChoices] = [] + for choice_position, raw_choice in enumerate(choices): + if not isinstance(raw_choice, Mapping): + continue + raw_index = raw_choice.get("index", choice_position) + choice_index = raw_index if isinstance(raw_index, int) else choice_position + delta_fields: dict[str, Any] = {} + raw_delta = raw_choice.get("delta") + if isinstance(raw_delta, Mapping): + delta_fields.update(raw_delta) + raw_message = raw_choice.get("message") + if isinstance(raw_message, Mapping): + message_text = extract_text_from_a2a_response({"result": {"message": raw_message}}) + if message_text and "content" not in delta_fields: + delta_fields["content"] = message_text + message_tool_calls = raw_message.get("tool_calls") + if message_tool_calls and "tool_calls" not in delta_fields: + delta_fields["tool_calls"] = message_tool_calls + raw_finish_reason = raw_choice.get("finish_reason") + choice_finish_reason = ( + raw_finish_reason + if isinstance(raw_finish_reason, str) and raw_finish_reason + else finish_reason + ) + streaming_choices.append( + StreamingChoices( + index=choice_index, + delta=Delta(**delta_fields), + finish_reason=choice_finish_reason, + logprobs=raw_choice.get("logprobs"), + ) + ) + if streaming_choices: + return ModelResponseStream( + choices=streaming_choices, + usage=usage, + provider_specific_fields=provider_fields or None, + ) + # Return generic streaming chunk return GenericStreamingChunk( text=text, diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 04432b2ca21..035bc770aed 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1872,12 +1872,25 @@ class ProxyBaseLLMRequestProcessing: trust_client_model_info=False, ) + authorized_model = self.data.get("model") self.data = await proxy_logging_obj.pre_call_hook( user_api_key_dict=user_api_key_dict, data=self.data, call_type=route_type, ) + if self.data.get("model") != authorized_model: + self.data = await authorize_a2a_agent_before_hooks( + data=self.data, + user_api_key_dict=user_api_key_dict, + ) + self.data = await merge_a2a_agent_guardrails_before_hooks(self.data) + self.data = _check_and_merge_model_level_guardrails( + data=self.data, + llm_router=llm_router, + trust_client_model_info=False, + ) + # Refresh AFTER pre_call_hook: guardrails (e.g. Presidio PII masking) may # have mutated `self.data` in place, and the audit-trail snapshot taken in # add_litellm_data_to_request predates that mutation. diff --git a/tests/test_litellm/a2a_protocol/providers/watsonx_orchestrate/test_watsonx_orchestrate_transformation.py b/tests/test_litellm/a2a_protocol/providers/watsonx_orchestrate/test_watsonx_orchestrate_transformation.py index 869769fe630..fa8b11ea182 100644 --- a/tests/test_litellm/a2a_protocol/providers/watsonx_orchestrate/test_watsonx_orchestrate_transformation.py +++ b/tests/test_litellm/a2a_protocol/providers/watsonx_orchestrate/test_watsonx_orchestrate_transformation.py @@ -364,7 +364,7 @@ def test_build_wxo_headers_preserves_auth_headers(): @pytest.mark.asyncio -async def test_wxo_config_forwards_static_headers(monkeypatch): +async def test_wxo_config_forwards_headers_and_timeout(monkeypatch): captured = {} async def fake_handle_non_streaming(**kwargs): @@ -377,10 +377,16 @@ async def test_wxo_config_forwards_static_headers(monkeypatch): request_id="req-1", params={}, litellm_params={"model": "agent"}, + agent_extra_headers={"x-request-id": "request-1"}, agent_static_headers={"x-tenant-id": "tenant-1"}, + timeout=12, ) - assert captured["static_headers"] == {"x-tenant-id": "tenant-1"} + assert captured["static_headers"] == { + "x-request-id": "request-1", + "x-tenant-id": "tenant-1", + } + assert captured["timeout"] == 12 @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 29d4e6ffe36..acd584b36db 100644 --- a/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py +++ b/tests/test_litellm/a2a_protocol/test_completion_bridge_streaming.py @@ -240,6 +240,10 @@ async def test_handle_streaming_emits_proper_events(): # Event 4: second artifact update assert events[3]["result"]["kind"] == "artifact-update" assert events[3]["result"]["artifact"]["parts"][0]["text"] == " world" + assert ( + events[2]["result"]["artifact"]["artifactId"] + == events[3]["result"]["artifact"]["artifactId"] + ) # Event 5: status completed assert events[4]["result"]["kind"] == "status-update" @@ -249,6 +253,72 @@ async def test_handle_streaming_emits_proper_events(): assert events[4]["usage"]["total_tokens"] == 5 +def test_build_completion_params_keeps_bridge_routing_fields(): + from litellm.a2a_protocol.litellm_completion_bridge.handler import ( + A2ACompletionBridgeHandler, + ) + + params = A2ACompletionBridgeHandler._build_completion_params( + params={"message": {"role": "user", "parts": []}}, + litellm_params={ + "custom_llm_provider": "openai", + "model": "agent", + "api_base": "https://untrusted.example", + "stream": False, + }, + api_base="https://configured.example", + agent_extra_headers=None, + stream=True, + ) + + assert params["api_base"] == "https://configured.example" + assert params["stream"] is True + + +@pytest.mark.asyncio +async def test_handle_streaming_accumulates_logprobs_and_provider_metadata(): + from litellm.a2a_protocol.litellm_completion_bridge.handler import ( + A2ACompletionBridgeHandler, + ) + + chunks = [] + for token in ("a", "b"): + choice = MagicMock() + choice.index = 0 + choice.finish_reason = None + choice.delta.content = token + choice.logprobs = {"content": [{"token": token}]} + chunk = MagicMock() + chunk.choices = [choice] + chunk.system_fingerprint = "fp-1" + chunk.service_tier = "scale" + chunks.append(chunk) + chunks[-1].choices[0].finish_reason = "stop" + + async def mock_streaming_response(): + for chunk in chunks: + yield chunk + + with patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion: + mock_acompletion.return_value = mock_streaming_response() + events = [ + event + async for event in A2ACompletionBridgeHandler.handle_streaming( + request_id="req-metadata", + params={"message": {"role": "user", "parts": []}}, + litellm_params={"custom_llm_provider": "openai", "model": "agent"}, + ) + ] + + result = events[-1] + assert result["system_fingerprint"] == "fp-1" + assert result["service_tier"] == "scale" + assert result["result"]["choices"][0]["logprobs"]["content"] == [ + {"token": "a"}, + {"token": "b"}, + ] + + @pytest.mark.asyncio async def test_handle_streaming_preserves_multiple_choices(): from litellm.a2a_protocol.litellm_completion_bridge.handler import ( 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 9c4212ee135..8518cd3d2f7 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 @@ -120,6 +120,37 @@ async def test_async_iterator_preserves_parallel_tool_calls(): assert chunk["tool_use"] == tool_calls +@pytest.mark.asyncio +async def test_async_iterator_preserves_every_terminal_choice(): + async def _events(): + yield { + "jsonrpc": "2.0", + "result": { + "kind": "status-update", + "status": {"state": "completed"}, + "choices": [ + { + "index": 0, + "message": {"parts": [{"kind": "text", "text": "first"}]}, + "finish_reason": "stop", + }, + { + "index": 1, + "message": {"parts": [{"kind": "text", "text": "second"}]}, + "finish_reason": "length", + }, + ], + }, + } + + iterator = A2AModelResponseIterator(streaming_response=_events(), sync_stream=False) + chunk = await iterator.__aiter__().__anext__() + + assert [choice.index for choice in chunk.choices] == [0, 1] + assert [choice.delta.content for choice in chunk.choices] == ["first", "second"] + assert [choice.finish_reason for choice in chunk.choices] == ["stop", "length"] + + @pytest.mark.asyncio async def test_async_iterator_serializes_delta_tool_calls_and_usage(): delta = Delta(