mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
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:
parent
28c0d8579b
commit
3b681b1b14
7 changed files with 786 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -0,0 +1,3 @@
|
|||
"""
|
||||
IBM watsonx Orchestrate (WXO) A2A provider.
|
||||
"""
|
||||
55
litellm/a2a_protocol/providers/watsonx_orchestrate/config.py
Normal file
55
litellm/a2a_protocol/providers/watsonx_orchestrate/config.py
Normal 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
|
||||
352
litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py
Normal file
352
litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py
Normal 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
|
||||
|
|
@ -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}"
|
||||
)
|
||||
|
|
@ -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"
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
Loading…
Add table
Reference in a new issue