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:
mateo-berri 2026-06-02 05:54:26 +00:00
parent a931b54e59
commit ff2c00597f
No known key found for this signature in database
2 changed files with 132 additions and 33 deletions

View file

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

View file

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