fix: forward static WXO headers

This commit is contained in:
aiedwardyi 2026-08-25 21:27:34 +09:00
parent a52bf355c6
commit eba758223d
No known key found for this signature in database
5 changed files with 92 additions and 12 deletions

View file

@ -127,6 +127,7 @@ class A2ACompletionBridgeHandler:
litellm_params: dict[str, Any],
api_base: str | None = None,
agent_extra_headers: dict[str, str] | None = None,
agent_static_headers: Mapping[str, str] | None = None,
*,
_skip_a2a_provider_routing: bool = False,
) -> dict[str, object]:
@ -140,6 +141,7 @@ class A2ACompletionBridgeHandler:
api_base: API base URL from agent_card_params
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
admin extra_headers) to forward on the upstream HTTP call.
agent_static_headers: Configured headers for provider-specific routing.
Returns:
A2A SendMessageResponse dict
@ -161,6 +163,7 @@ class A2ACompletionBridgeHandler:
"api_base": api_base,
"litellm_params": litellm_params,
"agent_extra_headers": agent_extra_headers,
"agent_static_headers": agent_static_headers,
}
if litellm_params.get("timeout") is not None:
provider_kwargs["timeout"] = litellm_params["timeout"]
@ -194,6 +197,7 @@ class A2ACompletionBridgeHandler:
litellm_params: dict[str, Any],
api_base: str | None = None,
agent_extra_headers: dict[str, str] | None = None,
agent_static_headers: Mapping[str, str] | None = None,
*,
_skip_a2a_provider_routing: bool = False,
) -> AsyncIterator[dict[str, object]]:
@ -213,6 +217,7 @@ class A2ACompletionBridgeHandler:
api_base: API base URL from agent_card_params
agent_extra_headers: Per-request headers (from x-a2a-{agent}-* rewrite and
admin extra_headers) to forward on the upstream HTTP call.
agent_static_headers: Configured headers for provider-specific routing.
Yields:
A2A streaming response events
@ -234,6 +239,7 @@ class A2ACompletionBridgeHandler:
"api_base": api_base,
"litellm_params": litellm_params,
"agent_extra_headers": agent_extra_headers,
"agent_static_headers": agent_static_headers,
}
if litellm_params.get("timeout") is not None:
provider_kwargs["timeout"] = litellm_params["timeout"]

View file

@ -32,6 +32,7 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig):
request_id=request_id,
params=params,
litellm_params=litellm_params,
static_headers=kwargs.get("agent_static_headers"),
)
async def handle_streaming(
@ -52,5 +53,6 @@ class WatsonxOrchestrateA2AConfig(BaseA2AProviderConfig):
request_id=request_id,
params=params,
litellm_params=litellm_params,
static_headers=kwargs.get("agent_static_headers"),
):
yield chunk

View file

@ -6,7 +6,7 @@ import asyncio
import hashlib
import json
import time
from collections.abc import AsyncIterator
from collections.abc import AsyncIterator, Mapping
from typing import Any, Final, NamedTuple, Protocol
import httpx
@ -26,9 +26,36 @@ _IBM_CLOUD_IAM_URL: Final = "https://iam.cloud.ibm.com/identity/token"
_POLL_INTERVAL_S: Final = 2.0
_MAX_POLL_ATTEMPTS: Final = 90
_TOKEN_CACHE_TTL_BUFFER_S: Final = 60
_WXO_RESERVED_HEADERS: Final = frozenset({"accept", "authorization", "content-type"})
_token_cache: Final[dict[str, tuple[str, float]]] = {}
def _build_wxo_headers(
token: str,
accept: str,
static_headers: Mapping[str, str] | None = None,
) -> dict[str, str]:
headers: dict[str, str] = (
{
key: value
for key, value in static_headers.items()
if isinstance(key, str)
and isinstance(value, str)
and key.lower() not in _WXO_RESERVED_HEADERS
}
if isinstance(static_headers, Mapping)
else {}
)
headers.update(
{
"Authorization": f"Bearer {token}",
"Content-Type": "application/json",
"Accept": accept,
}
)
return headers
class WXORequestParams(NamedTuple):
cp4d_host: str
instance_id: str
@ -277,6 +304,7 @@ class WatsonxOrchestrateHandler:
request_id: str,
params: dict[str, object],
litellm_params: WXOLitellmParams,
static_headers: Mapping[str, str] | None = None,
) -> dict[str, object]:
wxo: Final = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params)
@ -289,11 +317,11 @@ class WatsonxOrchestrateHandler:
client=client,
)
base_url: Final = WatsonxOrchestrateTransformation.get_api_base_url(wxo.cp4d_host, wxo.instance_id)
auth_headers: Final = {
"Authorization": f"Bearer {token}",
"Content-Type": "application/json",
"Accept": "application/json",
}
auth_headers: Final = _build_wxo_headers(
token=token,
accept="application/json",
static_headers=static_headers,
)
text: Final = WatsonxOrchestrateTransformation.extract_text_from_a2a_params(params)
body: Final = WatsonxOrchestrateTransformation.build_wxo_run_body(
@ -326,6 +354,7 @@ class WatsonxOrchestrateHandler:
litellm_params: WXOLitellmParams,
chunk_size: int = 50,
delay_ms: int = 10,
static_headers: Mapping[str, str] | None = None,
) -> AsyncIterator[dict[str, object]]:
wxo: Final = WatsonxOrchestrateHandler._extract_litellm_params(litellm_params)
@ -338,11 +367,11 @@ class WatsonxOrchestrateHandler:
client=client,
)
base_url: Final = WatsonxOrchestrateTransformation.get_api_base_url(wxo.cp4d_host, wxo.instance_id)
auth_headers: Final = {
"Authorization": f"Bearer {token}",
"Content-Type": "application/json",
"Accept": "text/event-stream, application/json",
}
auth_headers: Final = _build_wxo_headers(
token=token,
accept="text/event-stream, application/json",
static_headers=static_headers,
)
text: Final = WatsonxOrchestrateTransformation.extract_text_from_a2a_params(params)
body: Final = WatsonxOrchestrateTransformation.build_wxo_run_body(
wxo_agent_id=wxo.wxo_agent_id, text=text, thread_id=wxo.thread_id
@ -366,6 +395,7 @@ class WatsonxOrchestrateHandler:
request_id=request_id,
params=params,
litellm_params=litellm_params,
static_headers=static_headers,
)
response_text: Final = WatsonxOrchestrateTransformation.extract_text_from_a2a_message_response(result)
async for chunk in WatsonxOrchestrateTransformation.fake_streaming_from_text(

View file

@ -188,6 +188,7 @@ async def _route_registered_provider(
litellm_params=provider_params,
api_base=api_base,
agent_extra_headers=agent_extra_headers,
agent_static_headers=static_headers,
)
completion_stream: Final = A2AModelResponseIterator(
streaming_response=streaming_response,
@ -210,6 +211,7 @@ async def _route_registered_provider(
litellm_params=provider_params,
api_base=api_base,
agent_extra_headers=agent_extra_headers,
agent_static_headers=static_headers,
)
error_value: Final = response.get("error")
if isinstance(error_value, dict):

View file

@ -6,9 +6,11 @@ from pathlib import Path
import httpx
import pytest
from litellm.a2a_protocol.providers.config_manager import A2AProviderConfigManager
from litellm.a2a_protocol.providers.watsonx_orchestrate import handler as wxo_handler
from litellm.a2a_protocol.providers.watsonx_orchestrate.config import (
WatsonxOrchestrateA2AConfig,
)
from litellm.a2a_protocol.providers.watsonx_orchestrate.handler import (
WatsonxOrchestrateHandler,
)
@ -343,6 +345,44 @@ async def test_poll_run_raises_asyncio_timeout_when_never_terminal():
assert client.get_calls == 2
def test_build_wxo_headers_preserves_auth_headers():
headers = wxo_handler._build_wxo_headers(
token="token",
accept="application/json",
static_headers={
"x-tenant-id": "tenant-1",
"Authorization": "caller-token",
"content-type": "text/plain",
},
)
assert headers["x-tenant-id"] == "tenant-1"
assert headers["Authorization"] == "Bearer token"
assert headers["Content-Type"] == "application/json"
assert headers["Accept"] == "application/json"
assert "content-type" not in headers
@pytest.mark.asyncio
async def test_wxo_config_forwards_static_headers(monkeypatch):
captured = {}
async def fake_handle_non_streaming(**kwargs):
captured.update(kwargs)
return {"result": {}}
monkeypatch.setattr(WatsonxOrchestrateHandler, "handle_non_streaming", fake_handle_non_streaming)
await WatsonxOrchestrateA2AConfig().handle_non_streaming(
request_id="req-1",
params={},
litellm_params={"model": "agent"},
agent_static_headers={"x-tenant-id": "tenant-1"},
)
assert captured["static_headers"] == {"x-tenant-id": "tenant-1"}
@pytest.mark.asyncio
async def test_handle_streaming_polls_non_sse_json_until_complete(monkeypatch):
client = _JsonStreamClient({"status": "running", "run_id": "run-1"})