mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix: forward static WXO headers
This commit is contained in:
parent
a52bf355c6
commit
eba758223d
5 changed files with 92 additions and 12 deletions
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue