refactor(a2a/wxo): type run-param extraction and use text response_type

Return a typed WXORequestParams NamedTuple from _extract_litellm_params
instead of a positional tuple so call sites read params by name, and send
the user message with response_type 'text' so the run body is valid across
all WXO agent configurations rather than the search-specific type.
This commit is contained in:
mateo-berri 2026-06-02 06:09:40 +00:00
parent ff2c00597f
commit 65a06a57e2
No known key found for this signature in database
3 changed files with 36 additions and 45 deletions

View file

@ -6,7 +6,7 @@ import asyncio
import hashlib
import json
import time
from typing import Any, AsyncIterator, Dict, Optional, Tuple, cast
from typing import Any, AsyncIterator, Dict, NamedTuple, Optional, Tuple, cast
import httpx
@ -27,6 +27,16 @@ _TOKEN_CACHE_TTL_BUFFER_S = 60
_token_cache: Dict[str, Tuple[str, float]] = {}
class WXORequestParams(NamedTuple):
cp4d_host: str
instance_id: str
wxo_agent_id: str
api_key: str
username: Optional[str]
auth_mode: str
thread_id: Optional[str]
class WatsonxOrchestrateHandler:
@staticmethod
def _http_client(timeout: float = 90.0) -> AsyncHTTPHandler:
@ -186,14 +196,11 @@ class WatsonxOrchestrateHandler:
return accumulated_text
@staticmethod
def _extract_litellm_params(litellm_params: Dict[str, Any]) -> tuple:
def _extract_litellm_params(litellm_params: Dict[str, Any]) -> WXORequestParams:
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")
@ -208,14 +215,14 @@ class WatsonxOrchestrateHandler:
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,
return WXORequestParams(
cp4d_host=cp4d_host,
instance_id=instance_id,
wxo_agent_id=wxo_agent_id,
api_key=api_key,
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,
)
@staticmethod
@ -224,26 +231,18 @@ class WatsonxOrchestrateHandler:
params: Dict[str, Any],
litellm_params: Dict[str, Any],
) -> Dict[str, Any]:
(
cp4d_host,
instance_id,
wxo_agent_id,
api_key,
username,
auth_mode,
thread_id,
) = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params)
wxo = 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,
cp4d_host=wxo.cp4d_host,
auth_mode=wxo.auth_mode,
api_key=wxo.api_key,
username=wxo.username,
client=client,
)
base_url = WatsonxOrchestrateTransformation.get_api_base_url(
cp4d_host, instance_id
wxo.cp4d_host, wxo.instance_id
)
auth_headers = {
"Authorization": f"Bearer {token}",
@ -253,7 +252,7 @@ class WatsonxOrchestrateHandler:
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
wxo_agent_id=wxo.wxo_agent_id, text=text, thread_id=wxo.thread_id
)
run_response = await client.post(
@ -286,26 +285,18 @@ class WatsonxOrchestrateHandler:
chunk_size: int = 50,
delay_ms: int = 10,
) -> AsyncIterator[Dict[str, Any]]:
(
cp4d_host,
instance_id,
wxo_agent_id,
api_key,
username,
auth_mode,
thread_id,
) = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params)
wxo = 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,
cp4d_host=wxo.cp4d_host,
auth_mode=wxo.auth_mode,
api_key=wxo.api_key,
username=wxo.username,
client=client,
)
base_url = WatsonxOrchestrateTransformation.get_api_base_url(
cp4d_host, instance_id
wxo.cp4d_host, wxo.instance_id
)
auth_headers = {
"Authorization": f"Bearer {token}",
@ -314,7 +305,7 @@ class WatsonxOrchestrateHandler:
}
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
wxo_agent_id=wxo.wxo_agent_id, text=text, thread_id=wxo.thread_id
)
try:

View file

@ -60,7 +60,7 @@ class WatsonxOrchestrateTransformation:
"role": "user",
"content": [
{
"response_type": "conversational_search",
"response_type": "text",
"text": text,
}
],

View file

@ -149,7 +149,7 @@ class TestWatsonxOrchestrateTransformation:
)
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]["response_type"] == "text"
assert body["message"]["content"][0]["text"] == "Hi"
@pytest.mark.parametrize(