From 3fcbc4ceacbc5ecdd1b204c89be691902044d0b1 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Mon, 1 Jun 2026 10:08:04 +0000 Subject: [PATCH] Fix WXO streaming fallback error handling --- .../providers/watsonx_orchestrate/handler.py | 9 ++- ...test_watsonx_orchestrate_transformation.py | 61 +++++++++++++++++++ 2 files changed, 67 insertions(+), 3 deletions(-) diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py index a60f6420cca..922052cac17 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py @@ -8,6 +8,8 @@ import json import time from typing import Any, AsyncIterator, Dict, Optional, Tuple, cast +import httpx + from litellm._logging import verbose_logger from litellm.a2a_protocol.providers.watsonx_orchestrate.transformation import ( WatsonxOrchestrateTransformation, @@ -331,10 +333,11 @@ class WatsonxOrchestrateHandler: ): yield chunk - except Exception as exc: + except httpx.TransportError as exc: verbose_logger.warning( - f"WXO: Streaming request failed ({exc!r}), " - "falling back to non-streaming + fake streaming" + f"WXO: Streaming transport failed ({exc!r}), " + "falling back to non-streaming + fake streaming", + exc_info=True, ) result = await WatsonxOrchestrateHandler.handle_non_streaming( request_id=request_id, 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 8480c81d21f..4f5f063b37c 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 @@ -1,3 +1,4 @@ +import json import os import sys @@ -43,6 +44,31 @@ class _SSELines: yield line +class _InvalidJsonStreamResponse: + headers = {"content-type": "application/json"} + + def raise_for_status(self): + pass + + async def aread(self): + return b"not-json" + + +class _InvalidJsonStreamClient: + def __init__(self): + self.post_urls = [] + + async def post(self, url, **kwargs): + self.post_urls.append(url) + if "identity/token" in url: + return _JsonResponse({"access_token": "token", "expires_in": 3600}) + if url.endswith("/runs/stream"): + return _InvalidJsonStreamResponse() + if url.endswith("/runs"): + return _JsonResponse({"status": "completed", "results": "fallback text"}) + raise AssertionError(url) + + class TestWatsonxOrchestrateTransformation: def test_get_api_base_url(self): url = WatsonxOrchestrateTransformation.get_api_base_url( @@ -196,6 +222,41 @@ async def test_short_lived_tokens_are_not_served_from_cache(): assert client.calls == 2 +@pytest.mark.asyncio +async def test_handle_streaming_does_not_fallback_on_invalid_json(monkeypatch): + client = _InvalidJsonStreamClient() + monkeypatch.setattr( + WatsonxOrchestrateHandler, + "_http_client", + lambda timeout=90.0: client, + ) + + params = { + "message": { + "parts": [ + {"kind": "text", "text": "Hello"}, + ], + } + } + litellm_params = { + "cp4d_host": "https://cpd.example.com", + "instance_id": "instance-id", + "wxo_agent_id": "agent-id", + "api_key": "invalid-json-stream-cache-key", + "auth_mode": "ibm_cloud", + } + + with pytest.raises(json.JSONDecodeError): + async for _ in WatsonxOrchestrateHandler.handle_streaming( + request_id="req-1", + params=params, + litellm_params=litellm_params, + ): + pass + + assert not any(url.endswith("/runs") for url in client.post_urls) + + def test_config_manager_returns_wxo_provider(): config = A2AProviderConfigManager.get_provider_config( custom_llm_provider="watsonx_orchestrate"