mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
Fix watsonx orchestrate edge cases
This commit is contained in:
parent
7b3a110f34
commit
566ea5bc8e
3 changed files with 88 additions and 5 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue