From 65a06a57e25469f3f0fe703429b535c93e097843 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 2 Jun 2026 06:09:40 +0000 Subject: [PATCH] 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. --- .../providers/watsonx_orchestrate/handler.py | 77 ++++++++----------- .../watsonx_orchestrate/transformation.py | 2 +- ...test_watsonx_orchestrate_transformation.py | 2 +- 3 files changed, 36 insertions(+), 45 deletions(-) diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py index 74b027d77ff..e96f6bec24d 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py @@ -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: diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py index b634098f463..824e9dbcdd2 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/transformation.py @@ -60,7 +60,7 @@ class WatsonxOrchestrateTransformation: "role": "user", "content": [ { - "response_type": "conversational_search", + "response_type": "text", "text": text, } ], 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 index df2876900b6..5e1520c6867 100644 --- 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 @@ -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(