mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
Fix WXO streaming fallback error handling
This commit is contained in:
parent
566ea5bc8e
commit
3fcbc4ceac
2 changed files with 67 additions and 3 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue