feat(a2a): add watsonx Orchestrate agent provider

Bridge A2A message/send to WXO runs API (CP4D and IBM Cloud IAM auth),
with dashboard agent type metadata and unit tests for transformations.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Sameer Kankute 2026-06-01 12:36:26 +05:30
parent 28c0d8579b
commit 3b681b1b14
No known key found for this signature in database
7 changed files with 786 additions and 0 deletions

View file

@ -48,4 +48,11 @@ class A2AProviderConfigManager:
return BedrockAgentCoreA2AConfig()
if custom_llm_provider == "watsonx_orchestrate":
from litellm.a2a_protocol.providers.watsonx_orchestrate.config import (
WatsonxOrchestrateA2AConfig,
)
return WatsonxOrchestrateA2AConfig()
return None

View file

@ -0,0 +1,3 @@
"""
IBM watsonx Orchestrate (WXO) A2A provider.
"""

View file

@ -0,0 +1,55 @@
"""
A2A provider configuration for IBM watsonx Orchestrate (WXO).
"""
from typing import Any, AsyncIterator, Dict, Optional
from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig
from litellm.a2a_protocol.providers.watsonx_orchestrate.handler import (
WatsonxOrchestrateHandler,
)
class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig):
"""A2A bridge for IBM watsonx Orchestrate (REST runs API + poll/SSE)."""
async def handle_non_streaming(
self,
request_id: str,
params: Dict[str, Any],
api_base: Optional[str] = None,
**kwargs: Any,
) -> Dict[str, Any]:
"""Handle a non-streaming A2A request via WXO runs API."""
litellm_params = kwargs.get("litellm_params")
if not litellm_params:
raise ValueError(
"litellm_params is required for WatsonxOrchestrateA2AConfig "
"(must contain cp4d_host, instance_id, wxo_agent_id, api_key)"
)
return await WatsonxOrchestrateHandler.handle_non_streaming(
request_id=request_id,
params=params,
litellm_params=litellm_params,
)
async def handle_streaming(
self,
request_id: str,
params: Dict[str, Any],
api_base: Optional[str] = None,
**kwargs: Any,
) -> AsyncIterator[Dict[str, Any]]:
"""Handle a streaming A2A request via WXO streaming runs API."""
litellm_params = kwargs.get("litellm_params")
if not litellm_params:
raise ValueError(
"litellm_params is required for WatsonxOrchestrateA2AConfig "
"(must contain cp4d_host, instance_id, wxo_agent_id, api_key)"
)
async for chunk in WatsonxOrchestrateHandler.handle_streaming(
request_id=request_id,
params=params,
litellm_params=litellm_params,
):
yield chunk

View file

@ -0,0 +1,352 @@
"""
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 json
from typing import Any, AsyncIterator, Dict, Optional, 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.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
class WatsonxOrchestrateHandler:
"""
Handler for IBM watsonx Orchestrate agent requests.
"""
@staticmethod
async def _get_bearer_token(
cp4d_host: str,
auth_mode: str,
api_key: str,
username: Optional[str] = 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},
)
if auth_mode == "ibm_cloud":
verbose_logger.debug("WXO: Authenticating via IBM Cloud IAM")
response = await client.post(
_IBM_CLOUD_IAM_URL,
data={
"grant_type": "urn:ibm:params:oauth:grant-type:apikey",
"apikey": api_key,
},
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'"
)
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"])
@staticmethod
async def _poll_run(
base_url: str,
run_id: str,
auth_headers: Dict[str, str],
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):
await asyncio.sleep(interval_s)
response = await client.get(url, headers=auth_headers)
response.raise_for_status()
result: Dict[str, Any] = response.json()
status = result.get("status", "")
verbose_logger.debug(
f"WXO: Poll {attempt + 1}/{max_attempts} run='{run_id}' status='{status}'"
)
if status in WatsonxOrchestrateTransformation.TERMINAL_STATES:
return result
raise TimeoutError(
f"WXO run '{run_id}' did not reach a terminal state after "
f"{max_attempts * interval_s:.0f}s"
)
@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 ""
api_key = litellm_params.get("api_key") or ""
username = litellm_params.get("username") or None
auth_mode = litellm_params.get("auth_mode") or "cp4d"
thread_id = litellm_params.get("thread_id") or None
if not cp4d_host:
raise ValueError("'cp4d_host' is required in litellm_params for WXO agents")
if not instance_id:
raise ValueError(
"'instance_id' is required in litellm_params for WXO agents"
)
if not wxo_agent_id:
raise ValueError(
"'wxo_agent_id' is required in litellm_params for WXO agents"
)
if not api_key:
raise ValueError("'api_key' is required in litellm_params for WXO agents")
return (
cp4d_host,
instance_id,
wxo_agent_id,
api_key,
username,
auth_mode,
thread_id,
)
@staticmethod
async def handle_non_streaming(
request_id: str,
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,
wxo_agent_id,
api_key,
username,
auth_mode,
thread_id,
) = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params)
token = await WatsonxOrchestrateHandler._get_bearer_token(
cp4d_host=cp4d_host,
auth_mode=auth_mode,
api_key=api_key,
username=username,
)
base_url = WatsonxOrchestrateTransformation.get_api_base_url(
cp4d_host, instance_id
)
auth_headers = {
"Authorization": f"Bearer {token}",
"Content-Type": "application/json",
"Accept": "application/json",
}
text = WatsonxOrchestrateTransformation.extract_text_from_a2a_params(params)
body = WatsonxOrchestrateTransformation.build_wxo_run_body(
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,
headers=auth_headers,
)
run_response.raise_for_status()
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,
)
status = run_data.get("status", "")
if status not in WatsonxOrchestrateTransformation.SUCCESS_STATES:
raise RuntimeError(
f"WXO run ended with non-success status '{status}': {run_data}"
)
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
)
@staticmethod
async def handle_streaming(
request_id: str,
params: Dict[str, Any],
litellm_params: Dict[str, Any],
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,
wxo_agent_id,
api_key,
username,
auth_mode,
thread_id,
) = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params)
token = await WatsonxOrchestrateHandler._get_bearer_token(
cp4d_host=cp4d_host,
auth_mode=auth_mode,
api_key=api_key,
username=username,
)
base_url = WatsonxOrchestrateTransformation.get_api_base_url(
cp4d_host, instance_id
)
auth_headers = {
"Authorization": f"Bearer {token}",
"Content-Type": "application/json",
}
text = WatsonxOrchestrateTransformation.extract_text_from_a2a_params(params)
body = WatsonxOrchestrateTransformation.build_wxo_run_body(
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
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
async for (
chunk
) in WatsonxOrchestrateTransformation.fake_streaming_from_text(
text=accumulated_text,
request_id=request_id,
chunk_size=chunk_size,
delay_ms=delay_ms,
):
yield chunk
except Exception as exc:
verbose_logger.warning(
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
async for (
chunk
) in WatsonxOrchestrateTransformation.fake_streaming_from_text(
text=response_text,
request_id=request_id,
chunk_size=chunk_size,
delay_ms=delay_ms,
):
yield chunk

View file

@ -0,0 +1,202 @@
"""
Transformation layer for IBM watsonx Orchestrate (WXO) agent provider.
WXO uses a REST API (not A2A/JSON-RPC) with an async-poll execution model:
POST /v1/orchestrate/runs → submit run, get run_id
GET /v1/orchestrate/runs/{id} → poll until terminal state
POST /v1/orchestrate/runs/stream → native SSE streaming
"""
import asyncio
from typing import Any, AsyncIterator, Dict, Optional
from uuid import uuid4
from litellm._logging import verbose_logger
class WatsonxOrchestrateTransformation:
"""
Handles request/response transformation between A2A and the WXO REST API.
"""
TERMINAL_STATES = frozenset(
{"completed", "succeeded", "failed", "error", "cancelled"}
)
SUCCESS_STATES = frozenset({"completed", "succeeded"})
@staticmethod
def get_api_base_url(cp4d_host: str, instance_id: str) -> str:
"""Build the WXO API base URL from host and instance ID."""
return f"{cp4d_host.rstrip('/')}/orchestrate/cpd/instances/{instance_id}"
@staticmethod
def extract_text_from_a2a_params(params: Dict[str, Any]) -> str:
"""
Extract user message text from A2A MessageSendParams.
A2A format: params.message.parts[*] where part.kind == "text"
"""
message = params.get("message", {})
parts = message.get("parts", [])
texts = []
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"]:
texts.append(part["text"])
return " ".join(texts) or ""
@staticmethod
def build_wxo_run_body(
wxo_agent_id: str,
text: str,
thread_id: Optional[str] = None,
) -> Dict[str, Any]:
"""Build the WXO POST /v1/orchestrate/runs request body."""
body: Dict[str, Any] = {
"agent_id": wxo_agent_id,
"message": {
"role": "user",
"content": [
{
"response_type": "conversational_search",
"text": text,
}
],
},
}
if thread_id:
body["thread_id"] = thread_id
return body
@staticmethod
def extract_text_from_wxo_result(result: Dict[str, 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.
"""
# Primary: last_message.content[0].text
try:
text = result["last_message"]["content"][0]["text"]
if text:
return str(text)
except (KeyError, IndexError, TypeError):
pass
# Secondary: result.data.message.content[0].text
try:
text = result["result"]["data"]["message"]["content"][0]["text"]
if text:
return str(text)
except (KeyError, IndexError, TypeError):
pass
# Tertiary: results as a raw string
results = result.get("results")
if results and isinstance(results, str):
return results
return ""
@staticmethod
def build_a2a_message_response(request_id: str, text: str) -> Dict[str, Any]:
"""
Build a standard A2A non-streaming SendMessageResponse (kind=message).
"""
return {
"jsonrpc": "2.0",
"id": request_id,
"result": {
"kind": "message",
"role": "agent",
"parts": [{"kind": "text", "text": text}],
"messageId": str(uuid4()),
},
}
@staticmethod
async def fake_streaming_from_text(
text: str,
request_id: str,
chunk_size: int = 50,
delay_ms: int = 10,
) -> AsyncIterator[Dict[str, Any]]:
"""
Emit standard A2A streaming events from a completed text response.
Event sequence:
1. task (kind="task", state="submitted")
2. status-update (kind="status-update", state="working")
3. artifact-update chunks
4. status-update (kind="status-update", state="completed", final=True)
"""
task_id = str(uuid4())
context_id = str(uuid4())
artifact_id = str(uuid4())
# 1. Task submitted
yield {
"jsonrpc": "2.0",
"id": request_id,
"result": {
"contextId": context_id,
"id": task_id,
"kind": "task",
"status": {"state": "submitted"},
},
}
# 2. Working
yield {
"jsonrpc": "2.0",
"id": request_id,
"result": {
"contextId": context_id,
"final": False,
"kind": "status-update",
"status": {"state": "working"},
"taskId": task_id,
},
}
await asyncio.sleep(delay_ms / 1000.0)
# 3. Artifact chunks (always emit at least one chunk, even for empty text)
text_to_chunk = text or ""
for i in range(0, max(len(text_to_chunk), 1), chunk_size):
chunk_text = text_to_chunk[i : i + chunk_size]
is_last = (i + chunk_size) >= max(len(text_to_chunk), 1)
yield {
"jsonrpc": "2.0",
"id": request_id,
"result": {
"contextId": context_id,
"kind": "artifact-update",
"taskId": task_id,
"artifact": {
"artifactId": artifact_id,
"parts": [{"kind": "text", "text": chunk_text}],
},
},
}
if not is_last:
await asyncio.sleep(delay_ms / 1000.0)
# 4. Completed
yield {
"jsonrpc": "2.0",
"id": request_id,
"result": {
"contextId": context_id,
"final": True,
"kind": "status-update",
"status": {"state": "completed"},
"taskId": task_id,
},
}
verbose_logger.debug(
f"WXO: Fake streaming completed for request_id={request_id}"
)

View file

@ -189,6 +189,78 @@
"litellm_params_template": {
"custom_llm_provider": "vertex_ai"
}
},
{
"agent_type": "watsonx_orchestrate",
"agent_type_display_name": "watsonx Orchestrate",
"description": "Connect to IBM watsonx Orchestrate agents via CP4D or IBM Cloud IAM",
"logo_url": "/ui/assets/logos/watsonx.svg",
"credential_fields": [
{
"key": "cp4d_host",
"label": "CP4D Host URL",
"placeholder": "https://cpd-cpd.apps.example.com",
"tooltip": "Your CP4D cluster base URL (e.g. https://cpd-cpd.apps.example.com). For IBM Cloud WXO, use the service endpoint.",
"required": true,
"field_type": "text",
"default_value": null,
"include_in_litellm_params": true
},
{
"key": "instance_id",
"label": "WXO Instance ID",
"placeholder": "1769134113217795",
"tooltip": "The numeric watsonx Orchestrate instance ID. Find it in the WXO service URL: /orchestrate/cpd/instances/<INSTANCE_ID>",
"required": true,
"field_type": "text",
"default_value": null,
"include_in_litellm_params": true
},
{
"key": "wxo_agent_id",
"label": "WXO Agent ID",
"placeholder": "588c8cdf-60f4-454b-8468-8702b19dca46",
"tooltip": "UUID of the agent in watsonx Orchestrate. Find it via the WXO console or GET /v1/orchestrate/agents.",
"required": true,
"field_type": "text",
"default_value": null,
"include_in_litellm_params": true
},
{
"key": "auth_mode",
"label": "Authentication Mode",
"placeholder": null,
"tooltip": "cp4d: on-prem / CloudPak for Data (requires username). ibm_cloud: IBM Cloud IAM (api_key only).",
"required": false,
"field_type": "select",
"options": ["cp4d", "ibm_cloud"],
"default_value": "cp4d",
"include_in_litellm_params": true
},
{
"key": "username",
"label": "Username (CP4D only)",
"placeholder": "admin",
"tooltip": "Your CP4D username. Required when auth_mode is 'cp4d'.",
"required": false,
"field_type": "text",
"default_value": null,
"include_in_litellm_params": true
},
{
"key": "api_key",
"label": "API Key",
"placeholder": null,
"tooltip": "CP4D API key (auth_mode=cp4d) or IBM Cloud API key (auth_mode=ibm_cloud).",
"required": true,
"field_type": "password",
"default_value": null,
"include_in_litellm_params": true
}
],
"litellm_params_template": {
"custom_llm_provider": "watsonx_orchestrate"
}
}
]

View file

@ -0,0 +1,95 @@
import os
import sys
import pytest
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.a2a_protocol.providers.config_manager import A2AProviderConfigManager
from litellm.a2a_protocol.providers.watsonx_orchestrate.transformation import (
WatsonxOrchestrateTransformation,
)
class TestWatsonxOrchestrateTransformation:
def test_get_api_base_url(self):
url = WatsonxOrchestrateTransformation.get_api_base_url(
"https://cpd.example.com/",
"1769134113217795",
)
assert (
url == "https://cpd.example.com/orchestrate/cpd/instances/1769134113217795"
)
def test_extract_text_from_a2a_params(self):
params = {
"message": {
"role": "user",
"parts": [
{"kind": "text", "text": "Hello"},
{"kind": "text", "text": "world"},
],
}
}
assert (
WatsonxOrchestrateTransformation.extract_text_from_a2a_params(params)
== "Hello world"
)
def test_build_wxo_run_body_with_thread(self):
body = WatsonxOrchestrateTransformation.build_wxo_run_body(
wxo_agent_id="agent-uuid",
text="Hi",
thread_id="thread-1",
)
assert body["agent_id"] == "agent-uuid"
assert body["thread_id"] == "thread-1"
assert body["message"]["content"][0]["response_type"] == "conversational_search"
assert body["message"]["content"][0]["text"] == "Hi"
@pytest.mark.parametrize(
"result,expected",
[
(
{
"last_message": {
"content": [{"type": "text", "text": "from last_message"}]
}
},
"from last_message",
),
(
{
"result": {
"data": {
"message": {"content": [{"text": "from nested result"}]}
}
}
},
"from nested result",
),
({"results": "raw string"}, "raw string"),
],
)
def test_extract_text_from_wxo_result(self, result, expected):
assert (
WatsonxOrchestrateTransformation.extract_text_from_wxo_result(result)
== expected
)
def test_build_a2a_message_response(self):
out = WatsonxOrchestrateTransformation.build_a2a_message_response(
"req-1", "answer"
)
assert out["jsonrpc"] == "2.0"
assert out["id"] == "req-1"
assert out["result"]["kind"] == "message"
assert out["result"]["parts"][0]["text"] == "answer"
def test_config_manager_returns_wxo_provider():
config = A2AProviderConfigManager.get_provider_config(
custom_llm_provider="watsonx_orchestrate"
)
assert config is not None
assert config.__class__.__name__ == "WatsonxOrchestrateA2AConfig"