diff --git a/litellm/a2a_protocol/providers/config_manager.py b/litellm/a2a_protocol/providers/config_manager.py index d684efd4756..ecb8f66bdeb 100644 --- a/litellm/a2a_protocol/providers/config_manager.py +++ b/litellm/a2a_protocol/providers/config_manager.py @@ -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 diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/__init__.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/__init__.py new file mode 100644 index 00000000000..096bcc01214 --- /dev/null +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/__init__.py @@ -0,0 +1,3 @@ +""" +IBM watsonx Orchestrate (WXO) A2A provider. +""" diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py new file mode 100644 index 00000000000..dbd4a0558f7 --- /dev/null +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py @@ -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 diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py new file mode 100644 index 00000000000..f4bf7f39e57 --- /dev/null +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py @@ -0,0 +1,352 @@ +""" +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 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 /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 diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py new file mode 100644 index 00000000000..01428749270 --- /dev/null +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py @@ -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}" + ) diff --git a/litellm/proxy/public_endpoints/agent_create_fields.json b/litellm/proxy/public_endpoints/agent_create_fields.json index 931c9a43498..e58bd97cce7 100644 --- a/litellm/proxy/public_endpoints/agent_create_fields.json +++ b/litellm/proxy/public_endpoints/agent_create_fields.json @@ -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/", + "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" + } } ] 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 new file mode 100644 index 00000000000..a45e3952659 --- /dev/null +++ b/tests/test_litellm/a2a_protocol/providers/watsonx_orchestrate/test_watsonx_orchestrate_transformation.py @@ -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"