Fix WXO streaming fallback error handling

This commit is contained in:
Cursor Agent 2026-06-01 10:08:04 +00:00
parent 566ea5bc8e
commit 3fcbc4ceac
No known key found for this signature in database
2 changed files with 67 additions and 3 deletions

View file

@ -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,

View file

@ -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"