Fix watsonx orchestrate edge cases

This commit is contained in:
Cursor Agent 2026-06-01 09:26:02 +00:00
parent 7b3a110f34
commit 566ea5bc8e
No known key found for this signature in database
3 changed files with 88 additions and 5 deletions

View file

@ -104,7 +104,7 @@ class WatsonxOrchestrateHandler:
else:
ttl_s = WatsonxOrchestrateHandler._cp4d_token_ttl_seconds(expiration)
expires_at = now + max(ttl_s - _TOKEN_CACHE_TTL_BUFFER_S, 60)
expires_at = now + max(ttl_s - _TOKEN_CACHE_TTL_BUFFER_S, 0)
_token_cache[cache_key] = (token, expires_at)
return token

View file

@ -42,9 +42,8 @@ class WatsonxOrchestrateTransformation:
for part in parts:
if not isinstance(part, dict):
continue
if part.get("kind") == "text" and part.get("text"):
texts.append(part["text"])
elif "text" in part and part["text"]:
kind = part.get("kind")
if kind in (None, "", "text") and part.get("text"):
texts.append(part["text"])
return " ".join(texts) or ""
@ -72,12 +71,15 @@ class WatsonxOrchestrateTransformation:
return body
@staticmethod
def extract_text_from_wxo_result(result: Dict[str, Any]) -> str:
def extract_text_from_wxo_result(result: Any) -> str:
"""
Extract response text from a WXO run result.
WXO can return text in several locations; checks in priority order per the API spec.
"""
if not isinstance(result, dict):
return ""
# Primary: last_message.content[0].text
try:
text = result["last_message"]["content"][0]["text"]

View file

@ -14,6 +14,35 @@ from litellm.a2a_protocol.providers.watsonx_orchestrate.transformation import (
)
class _JsonResponse:
def __init__(self, payload):
self.payload = payload
def raise_for_status(self):
pass
def json(self):
return self.payload
class _ShortTtlTokenClient:
def __init__(self):
self.calls = 0
async def post(self, *args, **kwargs):
self.calls += 1
return _JsonResponse({"access_token": f"token-{self.calls}", "expires_in": 30})
class _SSELines:
def __init__(self, lines):
self.lines = lines
async def aiter_lines(self):
for line in self.lines:
yield line
class TestWatsonxOrchestrateTransformation:
def test_get_api_base_url(self):
url = WatsonxOrchestrateTransformation.get_api_base_url(
@ -39,6 +68,24 @@ class TestWatsonxOrchestrateTransformation:
== "Hello world"
)
def test_extract_text_from_a2a_params_ignores_non_text_parts_with_text(self):
params = {
"message": {
"role": "user",
"parts": [
{"kind": "data", "text": "metadata label", "data": {}},
{"kind": "file", "text": "file label", "file": {}},
{"kind": "text", "text": "Hello"},
{"text": "legacy"},
{"kind": "", "text": "empty-kind"},
],
}
}
assert (
WatsonxOrchestrateTransformation.extract_text_from_a2a_params(params)
== "Hello legacy empty-kind"
)
def test_build_wxo_run_body_with_thread(self):
body = WatsonxOrchestrateTransformation.build_wxo_run_body(
wxo_agent_id="agent-uuid",
@ -115,6 +162,40 @@ def test_cp4d_token_ttl_from_absolute_expiration():
assert WatsonxOrchestrateHandler._cp4d_token_ttl_seconds(1_749_999_000, wall) == 0
@pytest.mark.asyncio
async def test_accumulate_wxo_sse_text_ignores_non_dict_json_events():
response = _SSELines(
[
"data: null",
"data: true",
'data: {"results": "streamed text"}',
]
)
assert await WatsonxOrchestrateHandler._accumulate_wxo_sse_text(response) == (
"streamed text"
)
@pytest.mark.asyncio
async def test_short_lived_tokens_are_not_served_from_cache():
client = _ShortTtlTokenClient()
token_1 = await WatsonxOrchestrateHandler._get_bearer_token(
cp4d_host="https://cpd.example.com",
auth_mode="ibm_cloud",
api_key="short-ttl-cache-key",
client=client,
)
token_2 = await WatsonxOrchestrateHandler._get_bearer_token(
cp4d_host="https://cpd.example.com",
auth_mode="ibm_cloud",
api_key="short-ttl-cache-key",
client=client,
)
assert token_1 == "token-1"
assert token_2 == "token-2"
assert client.calls == 2
def test_config_manager_returns_wxo_provider():
config = A2AProviderConfigManager.get_provider_config(
custom_llm_provider="watsonx_orchestrate"