mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
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 <cursoragent@cursor.com>
This commit is contained in:
parent
3b681b1b14
commit
9a4ab99fb4
3 changed files with 145 additions and 116 deletions
|
|
@ -1,35 +1,47 @@
|
|||
"""
|
||||
Handler for IBM watsonx Orchestrate (WXO) agent provider.
|
||||
|
||||
Authentication:
|
||||
- CP4D: POST <cp4d_host>/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 <cp4d_host>/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(
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue