mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
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:
parent
ff2c00597f
commit
65a06a57e2
3 changed files with 36 additions and 45 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -60,7 +60,7 @@ class WatsonxOrchestrateTransformation:
|
|||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"response_type": "conversational_search",
|
||||
"response_type": "text",
|
||||
"text": text,
|
||||
}
|
||||
],
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue