diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py index 02413465796..a60f6420cca 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py @@ -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 diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py index 908d51bf287..b634098f463 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py @@ -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"] diff --git a/tests/test_litellm/a2a_protocol/providers/watsonx_orchestrate/test_watsonx_orchestrate_transformation.py b/tests/test_litellm/a2a_protocol/providers/watsonx_orchestrate/test_watsonx_orchestrate_transformation.py index 3e0375f7b90..8480c81d21f 100644 --- a/tests/test_litellm/a2a_protocol/providers/watsonx_orchestrate/test_watsonx_orchestrate_transformation.py +++ b/tests/test_litellm/a2a_protocol/providers/watsonx_orchestrate/test_watsonx_orchestrate_transformation.py @@ -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"