mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
fix(a2a/wxo): scope streaming transport fallback to initial POST only
Narrow the httpx.TransportError fallback in handle_streaming so it only covers the initial POST /runs/stream. Errors during polling or SSE consumption now propagate instead of triggering handle_non_streaming, which would have submitted a duplicate WXO run for the same request.
This commit is contained in:
parent
a931b54e59
commit
ff2c00597f
2 changed files with 132 additions and 33 deletions
|
|
@ -325,41 +325,10 @@ class WatsonxOrchestrateHandler:
|
|||
stream=True,
|
||||
)
|
||||
response.raise_for_status()
|
||||
content_type = response.headers.get("content-type", "").lower()
|
||||
|
||||
if "text/event-stream" not in content_type:
|
||||
response_body = await response.aread()
|
||||
result = json.loads(response_body)
|
||||
result = await WatsonxOrchestrateHandler._get_successful_run_data(
|
||||
run_data=result,
|
||||
base_url=base_url,
|
||||
auth_headers=auth_headers,
|
||||
client=client,
|
||||
)
|
||||
accumulated_text = (
|
||||
WatsonxOrchestrateTransformation.extract_text_from_wxo_result(
|
||||
result
|
||||
)
|
||||
)
|
||||
else:
|
||||
accumulated_text = (
|
||||
await WatsonxOrchestrateHandler._accumulate_wxo_sse_text(response)
|
||||
)
|
||||
|
||||
async for (
|
||||
chunk
|
||||
) in WatsonxOrchestrateTransformation.fake_streaming_from_text(
|
||||
text=accumulated_text,
|
||||
request_id=request_id,
|
||||
chunk_size=chunk_size,
|
||||
delay_ms=delay_ms,
|
||||
):
|
||||
yield chunk
|
||||
|
||||
except httpx.TransportError as exc:
|
||||
verbose_logger.warning(
|
||||
f"WXO: Streaming transport failed ({exc!r}), "
|
||||
"falling back to non-streaming + fake streaming",
|
||||
f"WXO: Streaming request failed before a run was submitted "
|
||||
f"({exc!r}), falling back to non-streaming + fake streaming",
|
||||
exc_info=True,
|
||||
)
|
||||
result = await WatsonxOrchestrateHandler.handle_non_streaming(
|
||||
|
|
@ -381,3 +350,30 @@ class WatsonxOrchestrateHandler:
|
|||
delay_ms=delay_ms,
|
||||
):
|
||||
yield chunk
|
||||
return
|
||||
|
||||
content_type = response.headers.get("content-type", "").lower()
|
||||
if "text/event-stream" not in content_type:
|
||||
response_body = await response.aread()
|
||||
result = json.loads(response_body)
|
||||
result = await WatsonxOrchestrateHandler._get_successful_run_data(
|
||||
run_data=result,
|
||||
base_url=base_url,
|
||||
auth_headers=auth_headers,
|
||||
client=client,
|
||||
)
|
||||
accumulated_text = (
|
||||
WatsonxOrchestrateTransformation.extract_text_from_wxo_result(result)
|
||||
)
|
||||
else:
|
||||
accumulated_text = await WatsonxOrchestrateHandler._accumulate_wxo_sse_text(
|
||||
response
|
||||
)
|
||||
|
||||
async for chunk in WatsonxOrchestrateTransformation.fake_streaming_from_text(
|
||||
text=accumulated_text,
|
||||
request_id=request_id,
|
||||
chunk_size=chunk_size,
|
||||
delay_ms=delay_ms,
|
||||
):
|
||||
yield chunk
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ import os
|
|||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../../.."))
|
||||
|
|
@ -370,6 +371,108 @@ async def test_handle_streaming_does_not_fallback_on_invalid_json(monkeypatch):
|
|||
assert not any(url.endswith("/runs") for url in client.post_urls)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_streaming_does_not_resubmit_run_on_poll_transport_error(
|
||||
monkeypatch,
|
||||
):
|
||||
class _RunSubmissionClient:
|
||||
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 _JsonStreamResponse({"status": "running", "run_id": "run-1"})
|
||||
if url.endswith("/runs"):
|
||||
return _JsonResponse({"status": "completed", "results": "duplicate"})
|
||||
raise AssertionError(url)
|
||||
|
||||
client = _RunSubmissionClient()
|
||||
|
||||
async def poll_run(base_url, run_id, auth_headers, client, **kwargs):
|
||||
raise httpx.ConnectError("connection reset during poll")
|
||||
|
||||
monkeypatch.setattr(
|
||||
WatsonxOrchestrateHandler,
|
||||
"_http_client",
|
||||
lambda timeout=90.0: client,
|
||||
)
|
||||
monkeypatch.setattr(WatsonxOrchestrateHandler, "_poll_run", poll_run)
|
||||
|
||||
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": "poll-transport-error-cache-key",
|
||||
"auth_mode": "ibm_cloud",
|
||||
}
|
||||
|
||||
with pytest.raises(httpx.TransportError):
|
||||
async for _ in WatsonxOrchestrateHandler.handle_streaming(
|
||||
request_id="req-1",
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
delay_ms=0,
|
||||
):
|
||||
pass
|
||||
|
||||
assert not any(url.endswith("/runs") for url in client.post_urls)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_streaming_falls_back_when_initial_post_fails(monkeypatch):
|
||||
class _StreamPostFailsClient:
|
||||
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"):
|
||||
raise httpx.ConnectError("cannot reach stream endpoint")
|
||||
if url.endswith("/runs"):
|
||||
return _JsonResponse({"status": "completed", "results": "fallback"})
|
||||
raise AssertionError(url)
|
||||
|
||||
client = _StreamPostFailsClient()
|
||||
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": "stream-post-fails-cache-key",
|
||||
"auth_mode": "ibm_cloud",
|
||||
}
|
||||
|
||||
events = [
|
||||
event
|
||||
async for event in WatsonxOrchestrateHandler.handle_streaming(
|
||||
request_id="req-1",
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
delay_ms=0,
|
||||
)
|
||||
]
|
||||
artifact_text = "".join(
|
||||
event["result"]["artifact"]["parts"][0]["text"]
|
||||
for event in events
|
||||
if event["result"].get("kind") == "artifact-update"
|
||||
)
|
||||
|
||||
assert artifact_text == "fallback"
|
||||
assert sum(url.endswith("/runs") for url in client.post_urls) == 1
|
||||
|
||||
|
||||
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