From 3d6df166f4ea34244ce8dbfeae37c144937ce7a6 Mon Sep 17 00:00:00 2001 From: aiedwardyi <41576951+aiedwardyi@users.noreply.github.com> Date: Wed, 26 Aug 2026 17:18:50 +0900 Subject: [PATCH] fix: close final A2A review gaps --- .../providers/watsonx_orchestrate/handler.py | 41 +++++++++++-------- .../litellm_core_utils/streaming_handler.py | 3 ++ litellm/llms/a2a/chat/streaming_iterator.py | 17 ++++++++ litellm/proxy/agent_endpoints/a2a_routing.py | 20 ++++++++- litellm/proxy/common_request_processing.py | 7 ++++ litellm/types/utils.py | 1 + ...test_watsonx_orchestrate_transformation.py | 12 ++++++ .../a2a/chat/test_a2a_streaming_iterator.py | 12 ++++++ .../proxy/test_common_request_processing.py | 4 +- .../proxy/test_route_a2a_models.py | 25 +++++++++++ 10 files changed, 121 insertions(+), 21 deletions(-) diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py index 0af22578ccc..f4eb65715fa 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py @@ -268,23 +268,28 @@ class WatsonxOrchestrateHandler: return run_data @staticmethod - async def _accumulate_wxo_sse_text(response: Any) -> str: - source: Final[_WXOView] = {"sse_source": response} - accumulated_text = "" - async for line in source["sse_source"].aiter_lines(): - if not line.startswith("data:"): - continue - data_str = line[5:].strip() - if not data_str or data_str == "[DONE]": - continue - try: - event = WatsonxOrchestrateHandler._decode_run_event(data_str) - except json.JSONDecodeError: - continue - chunk_text = WatsonxOrchestrateTransformation.extract_text_from_wxo_result(event) - if chunk_text: - accumulated_text += chunk_text - return accumulated_text + async def _accumulate_wxo_sse_text(response: Any, timeout: float | None = None) -> str: + async def _collect() -> str: + source: Final[_WXOView] = {"sse_source": response} + accumulated_text = "" + async for line in source["sse_source"].aiter_lines(): + if not line.startswith("data:"): + continue + data_str = line[5:].strip() + if not data_str or data_str == "[DONE]": + continue + try: + event = WatsonxOrchestrateHandler._decode_run_event(data_str) + except json.JSONDecodeError: + continue + chunk_text = WatsonxOrchestrateTransformation.extract_text_from_wxo_result(event) + if chunk_text: + accumulated_text += chunk_text + return accumulated_text + + if timeout is None: + return await _collect() + return await asyncio.wait_for(_collect(), timeout=max(timeout, 0)) @staticmethod def _extract_litellm_params(litellm_params: WXOLitellmParams) -> WXORequestParams: @@ -438,7 +443,7 @@ class WatsonxOrchestrateHandler: ) accumulated_text = WatsonxOrchestrateTransformation.extract_text_from_wxo_result(result) else: - accumulated_text = await WatsonxOrchestrateHandler._accumulate_wxo_sse_text(response) + accumulated_text = await WatsonxOrchestrateHandler._accumulate_wxo_sse_text(response, timeout=timeout) async for chunk in WatsonxOrchestrateTransformation.fake_streaming_from_text( text=accumulated_text, diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 7cfaddc0b0e..bbb3856208c 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -1227,6 +1227,9 @@ class CustomStreamWrapper: if not _chunk_has_content and (not isinstance(chunk, dict) or "provider_specific_fields" not in chunk): raise StopIteration anthropic_response_obj: Final[GChunk] = cast(GChunk, chunk) + chunk_id = anthropic_response_obj.get("id") + if isinstance(chunk_id, str) and chunk_id.strip(): + model_response = self.set_model_id(chunk_id, model_response) completion_obj["content"] = anthropic_response_obj["text"] chunk_index = anthropic_response_obj.get("index") if isinstance(chunk_index, int): diff --git a/litellm/llms/a2a/chat/streaming_iterator.py b/litellm/llms/a2a/chat/streaming_iterator.py index 099a9f67047..c4f8d1b2edd 100644 --- a/litellm/llms/a2a/chat/streaming_iterator.py +++ b/litellm/llms/a2a/chat/streaming_iterator.py @@ -4,6 +4,7 @@ A2A Streaming Response Iterator from collections.abc import Mapping from typing import Any, Final +from uuid import uuid4 import litellm from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator @@ -33,6 +34,7 @@ class A2AModelResponseIterator(BaseModelResponseIterator): json_mode=json_mode, ) self.model = model + self.response_id: str | None = None def chunk_parser(self, chunk: dict) -> GenericStreamingChunk | ModelResponseStream: """ @@ -69,6 +71,16 @@ class A2AModelResponseIterator(BaseModelResponseIterator): raise A2AError(status_code=500, message=f"A2A error: {error_message}") try: + if self.response_id is None: + raw_response_id = chunk.get("id") + raw_result = chunk.get("result") + if not isinstance(raw_response_id, str) and isinstance(raw_result, Mapping): + raw_response_id = raw_result.get("id") + self.response_id = ( + raw_response_id + if isinstance(raw_response_id, str) and raw_response_id.strip() + else f"chatcmpl-{uuid4().hex}" + ) # Extract text from A2A response result: Final = chunk.get("result", {}) chunk_index = 0 @@ -181,6 +193,7 @@ class A2AModelResponseIterator(BaseModelResponseIterator): if streaming_choices: return ModelResponseStream( choices=streaming_choices, + id=self.response_id, usage=usage, provider_specific_fields=provider_fields or None, ) @@ -188,6 +201,7 @@ class A2AModelResponseIterator(BaseModelResponseIterator): # Return generic streaming chunk return GenericStreamingChunk( text=text, + id=self.response_id, is_finished=bool(finish_reason or tool_calls), finish_reason=finish_reason or ("tool_calls" if tool_calls else ""), usage=usage, @@ -197,8 +211,11 @@ class A2AModelResponseIterator(BaseModelResponseIterator): ) except Exception: # Return empty chunk on parse error + if self.response_id is None: + self.response_id = f"chatcmpl-{uuid4().hex}" return GenericStreamingChunk( text="", + id=self.response_id, is_finished=False, finish_reason="", usage=None, diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 3e14d31f008..7a5e06330da 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -260,6 +260,7 @@ async def _route_registered_provider( "provider_specific_fields", "reasoning_content", "reasoning_items", + "refusal", "thinking_blocks", ): value = message_payload.get(field) @@ -319,7 +320,24 @@ async def _route_registered_provider( 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 + usage = litellm.Usage(**raw_usage) if isinstance(raw_usage, Mapping) else raw_usage + if usage is None and native_provider: + try: + from litellm.utils import token_counter + + prompt_tokens: Final = token_counter(model="gpt-3.5-turbo", messages=messages) + completion_tokens: Final = token_counter( + model="gpt-3.5-turbo", + text=extract_text_from_a2a_response(response), + count_response_tokens=True, + ) + usage = litellm.Usage( + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + total_tokens=prompt_tokens + completion_tokens, + ) + except Exception: # noqa: BLE001 - token estimation must not fail the response + pass if usage is not None: model_response.usage = usage if isinstance(logging_obj, Logging): diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 035bc770aed..c6f5cfbe5db 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1890,6 +1890,13 @@ class ProxyBaseLLMRequestProcessing: llm_router=llm_router, trust_client_model_info=False, ) + if isinstance(self.data.get("model"), str) and self.data["model"].startswith("a2a/"): + self.data = await proxy_logging_obj.pre_call_hook( + user_api_key_dict=user_api_key_dict, + data=self.data, + call_type=route_type, + guardrails_only=True, + ) # 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 diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 61d7aca0430..e3914511574 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -317,6 +317,7 @@ class ModelInfo(ModelInfoBase, total=False): class GenericStreamingChunk(TypedDict, total=False): text: Required[str] + id: str tool_use: ChatCompletionToolCallChunk | list[ChatCompletionToolCallChunk] | None is_finished: Required[bool] finish_reason: Required[str] 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 fa8b11ea182..7dadfb9c063 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 @@ -48,6 +48,12 @@ class _SSELines: yield line +class _HangingSSELines: + async def aiter_lines(self): + await asyncio.Event().wait() + yield "" + + class _InvalidJsonStreamResponse: headers = {"content-type": "application/json"} @@ -233,6 +239,12 @@ async def test_accumulate_wxo_sse_text_ignores_non_dict_json_events(): ) +@pytest.mark.asyncio +async def test_accumulate_wxo_sse_text_respects_timeout(): + with pytest.raises(asyncio.TimeoutError): + await WatsonxOrchestrateHandler._accumulate_wxo_sse_text(_HangingSSELines(), timeout=0.001) + + @pytest.mark.asyncio async def test_short_lived_tokens_are_not_served_from_cache(): client = _ShortTtlTokenClient() 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 8518cd3d2f7..1643604ad84 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 @@ -28,6 +28,18 @@ async def test_async_iterator_accepts_decoded_a2a_events(): assert chunk["text"] == "Hello" +def test_chunk_parser_reuses_response_id_for_idless_artifacts(): + iterator = A2AModelResponseIterator(streaming_response=[], sync_stream=False) + first = iterator.chunk_parser( + {"result": {"kind": "artifact-update", "artifact": {"parts": [{"kind": "text", "text": "one"}]}}} + ) + second = iterator.chunk_parser( + {"result": {"kind": "artifact-update", "artifact": {"parts": [{"kind": "text", "text": "two"}]}}} + ) + + assert first["id"] == second["id"] + + @pytest.mark.asyncio async def test_async_iterator_ignores_status_message_text(): async def _events(): diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 5d552d0a673..b755255fb48 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -6252,7 +6252,7 @@ class TestPerRequestModelGroupAlias: @pytest.mark.asyncio - async def test_a2a_reroute_does_not_repeat_pre_call_hook(self, monkeypatch): + async def test_a2a_reroute_runs_target_guardrails(self, monkeypatch): processing_obj = ProxyBaseLLMRequestProcessing( data={"model": "source", "messages": [{"role": "user", "content": "hello"}]} ) @@ -6310,7 +6310,7 @@ class TestPerRequestModelGroupAlias: ) assert returned_data["model"] == "a2a/agent" - assert hook_modes == [False] + assert hook_modes == [False, True] class TestInjectCostIntoUsageDict: diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index 05bc118907c..f0d0ff084d6 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -116,6 +116,7 @@ async def test_route_a2a_model_uses_registered_provider(): "kind": "message", "role": "agent", "parts": [{"kind": "text", "text": "Hello back"}], + "refusal": "I cannot complete that request.", "messageId": "message-id", }, } @@ -143,6 +144,7 @@ async def test_route_a2a_model_uses_registered_provider(): bridge.assert_awaited_once() generic_completion.assert_not_called() assert response.choices[0].message.content == "Hello back" + assert response.choices[0].message.refusal == "I cannot complete that request." bridge_kwargs = bridge.await_args.kwargs assert bridge_kwargs["litellm_params"]["max_tokens"] == 32 assert bridge_kwargs["litellm_params"]["temperature"] == 0.2 @@ -677,6 +679,29 @@ async def test_registered_provider_logging_uses_provider_model_for_builtin_prici assert logging_obj.model_call_details["litellm_params"]["model"] == "gpt-4o" +@pytest.mark.asyncio +async def test_native_registered_provider_estimates_usage_when_missing(monkeypatch): + response = {"result": {"message": {"parts": [{"kind": "text", "text": "hello"}]}}} + counter = Mock(side_effect=[3, 2]) + monkeypatch.setattr("litellm.utils.token_counter", counter) + + with patch( + "litellm.a2a_protocol.litellm_completion_bridge.handler.A2ACompletionBridgeHandler.handle_non_streaming", + AsyncMock(return_value=response), + ): + result = await _route_registered_provider( + data={"messages": [{"role": "user", "content": "hello"}]}, + model_name="a2a/agent", + api_base="https://provider.example", + litellm_params={"model": "agent", "custom_llm_provider": "pydantic_ai_agents"}, + static_headers=None, + ) + + assert result.usage.prompt_tokens == 3 + assert result.usage.completion_tokens == 2 + assert result.usage.total_tokens == 5 + + @pytest.mark.asyncio async def test_route_non_a2a_model_raises_error_if_not_in_router(): """Test that non-a2a models that aren't in router raise an error"""