From 9a4ab99fb48411f7343c1150248beb5f26bbfd22 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 1 Jun 2026 12:47:43 +0530 Subject: [PATCH] fix(a2a): use shared httpx client and cache WXO auth tokens Route WXO streaming through get_async_httpx_client (TLS verification enabled). Cache bearer tokens with TTL buffer. Extract A2A reply text via a dedicated helper instead of hard-coded JSON paths. Co-authored-by: Cursor --- .../providers/watsonx_orchestrate/handler.py | 224 +++++++++--------- .../watsonx_orchestrate/transformation.py | 20 ++ ...test_watsonx_orchestrate_transformation.py | 17 ++ 3 files changed, 145 insertions(+), 116 deletions(-) diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py index f4bf7f39e57..90aaf8300ef 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py @@ -1,35 +1,47 @@ """ Handler for IBM watsonx Orchestrate (WXO) agent provider. - -Authentication: - - CP4D: POST /icp4d-api/v1/authorize → {"token": "..."} - - IBM Cloud: POST https://iam.cloud.ibm.com/identity/token → {"access_token": "..."} - -Execution: - - Non-streaming: POST /v1/orchestrate/runs, then poll GET /v1/orchestrate/runs/{run_id} - - Streaming: POST /v1/orchestrate/runs/stream (SSE), falls back to poll + fake streaming """ import asyncio +import hashlib import json -from typing import Any, AsyncIterator, Dict, Optional, cast +import time +from typing import Any, AsyncIterator, Dict, Optional, Tuple, cast from litellm._logging import verbose_logger from litellm.a2a_protocol.providers.watsonx_orchestrate.transformation import ( WatsonxOrchestrateTransformation, ) -from litellm.llms.custom_httpx.http_handler import get_async_httpx_client +from litellm.llms.custom_httpx.http_handler import ( + AsyncHTTPHandler, + get_async_httpx_client, +) from litellm.types.llms.custom_http import httpxSpecialProvider _IBM_CLOUD_IAM_URL = "https://iam.cloud.ibm.com/identity/token" _POLL_INTERVAL_S = 2.0 -_MAX_POLL_ATTEMPTS = 90 # ~3 minutes at 2s per poll +_MAX_POLL_ATTEMPTS = 90 +_TOKEN_CACHE_TTL_BUFFER_S = 60 +_token_cache: Dict[str, Tuple[str, float]] = {} class WatsonxOrchestrateHandler: - """ - Handler for IBM watsonx Orchestrate agent requests. - """ + @staticmethod + def _http_client(timeout: float = 90.0) -> AsyncHTTPHandler: + return get_async_httpx_client( + llm_provider=cast(Any, httpxSpecialProvider.A2AProvider), + params={"timeout": timeout}, + ) + + @staticmethod + def _token_cache_key( + auth_mode: str, + cp4d_host: str, + api_key: str, + username: Optional[str], + ) -> str: + material = f"{auth_mode}:{cp4d_host}:{username or ''}:{api_key}" + return hashlib.sha256(material.encode()).hexdigest() @staticmethod async def _get_bearer_token( @@ -37,20 +49,20 @@ class WatsonxOrchestrateHandler: auth_mode: str, api_key: str, username: Optional[str] = None, + client: Optional[AsyncHTTPHandler] = None, ) -> str: - """ - Obtain a WXO bearer token. - - auth_mode="cp4d" → POST /icp4d-api/v1/authorize - auth_mode="ibm_cloud" → POST https://iam.cloud.ibm.com/identity/token - """ - client = get_async_httpx_client( - llm_provider=cast(Any, httpxSpecialProvider.A2AProvider), - params={"timeout": 30.0}, + cache_key = WatsonxOrchestrateHandler._token_cache_key( + auth_mode, cp4d_host, api_key, username ) + now = time.monotonic() + cached = _token_cache.get(cache_key) + if cached and cached[1] > now: + return cached[0] + + if client is None: + client = WatsonxOrchestrateHandler._http_client(timeout=30.0) if auth_mode == "ibm_cloud": - verbose_logger.debug("WXO: Authenticating via IBM Cloud IAM") response = await client.post( _IBM_CLOUD_IAM_URL, data={ @@ -60,38 +72,38 @@ class WatsonxOrchestrateHandler: headers={"Content-Type": "application/x-www-form-urlencoded"}, ) response.raise_for_status() - return str(response.json()["access_token"]) - - # Default: CP4D - if not username: - raise ValueError( - "'username' is required in litellm_params when auth_mode='cp4d'" + payload = response.json() + token = str(payload["access_token"]) + ttl_s = int(payload.get("expires_in", 3600)) + else: + if not username: + raise ValueError( + "'username' is required in litellm_params when auth_mode='cp4d'" + ) + token_url = f"{cp4d_host.rstrip('/')}/icp4d-api/v1/authorize" + response = await client.post( + token_url, + json={"username": username, "api_key": api_key}, + headers={"Content-Type": "application/json"}, ) - token_url = f"{cp4d_host.rstrip('/')}/icp4d-api/v1/authorize" - verbose_logger.debug(f"WXO: Authenticating via CP4D at {token_url}") - response = await client.post( - token_url, - json={"username": username, "api_key": api_key}, - headers={"Content-Type": "application/json"}, - ) - response.raise_for_status() - return str(response.json()["token"]) + response.raise_for_status() + payload = response.json() + token = str(payload["token"]) + ttl_s = int(payload.get("expiration", 3600)) + + expires_at = now + max(ttl_s - _TOKEN_CACHE_TTL_BUFFER_S, 60) + _token_cache[cache_key] = (token, expires_at) + return token @staticmethod async def _poll_run( base_url: str, run_id: str, auth_headers: Dict[str, str], + client: AsyncHTTPHandler, max_attempts: int = _MAX_POLL_ATTEMPTS, interval_s: float = _POLL_INTERVAL_S, ) -> Dict[str, Any]: - """ - Poll GET /v1/orchestrate/runs/{run_id} until a terminal state is reached. - """ - client = get_async_httpx_client( - llm_provider=cast(Any, httpxSpecialProvider.A2AProvider), - params={"timeout": 30.0}, - ) url = f"{base_url}/v1/orchestrate/runs/{run_id}" for attempt in range(max_attempts): @@ -111,9 +123,28 @@ class WatsonxOrchestrateHandler: f"{max_attempts * interval_s:.0f}s" ) + @staticmethod + async def _accumulate_wxo_sse_text(response: Any) -> str: + accumulated_text = "" + async for line in response.aiter_lines(): + if not line.startswith("data:"): + continue + data_str = line[5:].strip() + if not data_str or data_str == "[DONE]": + continue + try: + event = json.loads(data_str) + except json.JSONDecodeError: + continue + chunk_text = WatsonxOrchestrateTransformation.extract_text_from_wxo_result( + event + ) + if chunk_text: + accumulated_text += chunk_text + return accumulated_text + @staticmethod def _extract_litellm_params(litellm_params: Dict[str, Any]) -> tuple: - """Validate and extract required WXO params from litellm_params.""" cp4d_host = litellm_params.get("cp4d_host") or "" instance_id = litellm_params.get("instance_id") or "" wxo_agent_id = litellm_params.get("wxo_agent_id") or "" @@ -151,9 +182,6 @@ class WatsonxOrchestrateHandler: params: Dict[str, Any], litellm_params: Dict[str, Any], ) -> Dict[str, Any]: - """ - Submit a WXO run and poll for completion, then return a standard A2A message response. - """ ( cp4d_host, instance_id, @@ -164,11 +192,13 @@ class WatsonxOrchestrateHandler: thread_id, ) = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params) + client = WatsonxOrchestrateHandler._http_client(timeout=90.0) token = await WatsonxOrchestrateHandler._get_bearer_token( cp4d_host=cp4d_host, auth_mode=auth_mode, api_key=api_key, username=username, + client=client, ) base_url = WatsonxOrchestrateTransformation.get_api_base_url( cp4d_host, instance_id @@ -184,13 +214,6 @@ class WatsonxOrchestrateHandler: wxo_agent_id=wxo_agent_id, text=text, thread_id=thread_id ) - verbose_logger.info( - f"WXO: Submitting run for agent='{wxo_agent_id}' at {base_url}" - ) - client = get_async_httpx_client( - llm_provider=cast(Any, httpxSpecialProvider.A2AProvider), - params={"timeout": 90.0}, - ) run_response = await client.post( f"{base_url}/v1/orchestrate/runs", json=body, @@ -200,17 +223,15 @@ class WatsonxOrchestrateHandler: run_data: Dict[str, Any] = run_response.json() status = run_data.get("status", "") - verbose_logger.debug(f"WXO: Run submitted, initial status='{status}'") - if status not in WatsonxOrchestrateTransformation.TERMINAL_STATES: run_id = run_data.get("run_id") or run_data.get("id") or "" if not run_id: raise ValueError(f"WXO: No run_id in response: {run_data}") - verbose_logger.info(f"WXO: Polling run '{run_id}' for completion...") run_data = await WatsonxOrchestrateHandler._poll_run( base_url=base_url, run_id=run_id, auth_headers=auth_headers, + client=client, ) status = run_data.get("status", "") @@ -222,10 +243,6 @@ class WatsonxOrchestrateHandler: response_text = WatsonxOrchestrateTransformation.extract_text_from_wxo_result( run_data ) - verbose_logger.info( - f"WXO: Run completed successfully for request_id={request_id}" - ) - return WatsonxOrchestrateTransformation.build_a2a_message_response( request_id=request_id, text=response_text ) @@ -238,13 +255,6 @@ class WatsonxOrchestrateHandler: chunk_size: int = 50, delay_ms: int = 10, ) -> AsyncIterator[Dict[str, Any]]: - """ - Stream a WXO run. - - Tries native SSE via POST /v1/orchestrate/runs/stream first. - If that fails or returns non-SSE content, falls back to non-streaming - poll + fake A2A streaming events. - """ ( cp4d_host, instance_id, @@ -255,11 +265,13 @@ class WatsonxOrchestrateHandler: thread_id, ) = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params) + client = WatsonxOrchestrateHandler._http_client(timeout=120.0) token = await WatsonxOrchestrateHandler._get_bearer_token( cp4d_host=cp4d_host, auth_mode=auth_mode, api_key=api_key, username=username, + client=client, ) base_url = WatsonxOrchestrateTransformation.get_api_base_url( cp4d_host, instance_id @@ -273,46 +285,28 @@ class WatsonxOrchestrateHandler: wxo_agent_id=wxo_agent_id, text=text, thread_id=thread_id ) - verbose_logger.info(f"WXO: Submitting streaming run for agent='{wxo_agent_id}'") - try: - import httpx as _httpx + response = await client.post( + f"{base_url}/v1/orchestrate/runs/stream", + json=body, + headers=auth_headers, + stream=True, + ) + response.raise_for_status() + content_type = response.headers.get("content-type", "").lower() - accumulated_text = "" - async with _httpx.AsyncClient(verify=False, timeout=120.0) as http_client: - async with http_client.stream( - "POST", - f"{base_url}/v1/orchestrate/runs/stream", - json=body, - headers=auth_headers, - ) as stream_resp: - stream_resp.raise_for_status() - content_type = stream_resp.headers.get("content-type", "") - - if "text/event-stream" not in content_type: - # Provider returned a plain JSON response; fake-stream it - response_body = await stream_resp.aread() - result = json.loads(response_body) - accumulated_text = WatsonxOrchestrateTransformation.extract_text_from_wxo_result( - result - ) - else: - # Parse SSE lines and accumulate text from WXO events - async for line in stream_resp.aiter_lines(): - if not line.startswith("data:"): - continue - data_str = line[5:].strip() - if not data_str or data_str == "[DONE]": - continue - try: - event = json.loads(data_str) - chunk_text = WatsonxOrchestrateTransformation.extract_text_from_wxo_result( - event - ) - if chunk_text: - accumulated_text += chunk_text - except json.JSONDecodeError: - pass + if "text/event-stream" not in content_type: + response_body = await response.aread() + result = json.loads(response_body) + accumulated_text = ( + WatsonxOrchestrateTransformation.extract_text_from_wxo_result( + result + ) + ) + else: + accumulated_text = ( + await WatsonxOrchestrateHandler._accumulate_wxo_sse_text(response) + ) async for ( chunk @@ -329,18 +323,16 @@ class WatsonxOrchestrateHandler: f"WXO: Streaming request failed ({exc!r}), " "falling back to non-streaming + fake streaming" ) - # Fallback: poll then fake-stream result = await WatsonxOrchestrateHandler.handle_non_streaming( request_id=request_id, params=params, litellm_params=litellm_params, ) - # Extract the text from the A2A message response we built - response_text = "" - try: - response_text = result["result"]["parts"][0]["text"] - except (KeyError, IndexError, TypeError): - pass + response_text = ( + WatsonxOrchestrateTransformation.extract_text_from_a2a_message_response( + result + ) + ) async for ( chunk ) in WatsonxOrchestrateTransformation.fake_streaming_from_text( diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py index 01428749270..908d51bf287 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py @@ -101,6 +101,26 @@ class WatsonxOrchestrateTransformation: return "" + @staticmethod + def extract_text_from_a2a_message_response(a2a_response: Dict[str, Any]) -> str: + result = a2a_response.get("result") + if not isinstance(result, dict): + verbose_logger.warning("WXO: A2A response missing result object") + return "" + parts = result.get("parts") + if not isinstance(parts, list): + verbose_logger.warning("WXO: A2A result has no parts list") + return "" + for part in parts: + if ( + isinstance(part, dict) + and part.get("kind") == "text" + and part.get("text") + ): + return str(part["text"]) + verbose_logger.warning("WXO: A2A result parts contained no text") + return "" + @staticmethod def build_a2a_message_response(request_id: str, text: str) -> Dict[str, Any]: """ 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 a45e3952659..a4aa9972043 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 @@ -86,6 +86,23 @@ class TestWatsonxOrchestrateTransformation: assert out["result"]["kind"] == "message" assert out["result"]["parts"][0]["text"] == "answer" + def test_extract_text_from_a2a_message_response(self): + envelope = WatsonxOrchestrateTransformation.build_a2a_message_response( + "req-1", "answer" + ) + assert ( + WatsonxOrchestrateTransformation.extract_text_from_a2a_message_response( + envelope + ) + == "answer" + ) + assert ( + WatsonxOrchestrateTransformation.extract_text_from_a2a_message_response( + {"result": {}} + ) + == "" + ) + def test_config_manager_returns_wxo_provider(): config = A2AProviderConfigManager.get_provider_config(