From eba758223d33fc777b9a5a206738135f255b4b78 Mon Sep 17 00:00:00 2001 From: aiedwardyi <41576951+aiedwardyi@users.noreply.github.com> Date: Tue, 25 Aug 2026 21:27:34 +0900 Subject: [PATCH] fix: forward static WXO headers --- .../litellm_completion_bridge/handler.py | 6 +++ .../providers/watsonx_orchestrate/config.py | 2 + .../providers/watsonx_orchestrate/handler.py | 52 +++++++++++++++---- litellm/proxy/agent_endpoints/a2a_routing.py | 2 + ...test_watsonx_orchestrate_transformation.py | 42 ++++++++++++++- 5 files changed, 92 insertions(+), 12 deletions(-) diff --git a/litellm/a2a_protocol/litellm_completion_bridge/handler.py b/litellm/a2a_protocol/litellm_completion_bridge/handler.py index f99a91729d4..d4708e02786 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/handler.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/handler.py @@ -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"] diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py index ca84d3e07b4..1b8024ae075 100644 --- a/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py +++ b/litellm/a2a_protocol/providers/watsonx_orchestrate/config.py @@ -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 diff --git a/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py b/litellm/a2a_protocol/providers/watsonx_orchestrate/handler.py index c66b07c321c..f648cca40cd 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 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( diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 38cd5e9f4f2..ff8381c6531 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -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): 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 43dfdaba02d..869769fe630 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 @@ -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"})