feat: add configurable provider affinity header mapping (#41033)

* feat: add configurable provider affinity header mapping

* fix: sync provider affinity API types

* fix: harden provider affinity header mapping

* fix: avoid provider affinity import cycle

* fix: preserve input callback header mutations

* fix: address provider affinity code scanning findings

* fix: satisfy provider affinity type discipline gate

* test: cover omitted pre-call argument isolation

* fix: resolve remaining provider affinity codeql alerts

* fix: redact provider affinity headers after calls

* refactor: drop provider affinity header log redaction

* fix: reject control characters in affinity session ids as a bad request

* chore: regenerate the openapi snapshot on python 3.12 and reuse the session marker constant

* fix(responses): read the affinity session from the named metadata argument

---------

Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
togear 2026-09-23 04:31:33 +08:00 • committed by GitHub
parent 666f6b01b4
commit fc0055497c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 734 additions and 14 deletions

View file

@ -24,10 +24,7 @@ AWS_CREDENTIAL_KWARGS_KEYS: Final = frozenset(
}
)
# Keys `completion()` forwards from its own kwargs into `get_litellm_params`,
# which are otherwise invisible to it because that call site passes explicit
# named arguments rather than `**kwargs`.
FORWARDED_KWARGS_KEYS: Final = AWS_CREDENTIAL_KWARGS_KEYS
PROVIDER_AFFINITY_HEADER_KWARG_KEY: Final = "provider_affinity_header"
# Pre-define optional kwargs keys as frozenset for O(1) lookups
# These are extracted from kwargs only if present, avoiding unnecessary .get() calls
@ -63,6 +60,7 @@ OPTIONAL_KWARGS_KEYS: Final = (
"itpm",
"otpm",
"use_xai_oauth",
PROVIDER_AFFINITY_HEADER_KWARG_KEY,
}
)
| AWS_CREDENTIAL_KWARGS_KEYS

View file

@ -0,0 +1,98 @@
import re
from collections.abc import Mapping
from typing import Final
from litellm.constants import SESSION_ID_GENERATED_METADATA_KEY
_HTTP_HEADER_NAME_PATTERN: Final = re.compile(r"^[!#$%&'*+\-.^_`|~0-9A-Za-z]+$")
_FORBIDDEN_AFFINITY_HEADERS: Final = frozenset(
{
"api-key",
"authorization",
"connection",
"content-length",
"content-type",
"cookie",
"host",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"proxy-connection",
"set-cookie",
"te",
"trailer",
"transfer-encoding",
"upgrade",
"www-authenticate",
"x-api-key",
"x-goog-api-key",
}
)
def validate_provider_affinity_header_name(header: str) -> str:
if not _HTTP_HEADER_NAME_PATTERN.fullmatch(header):
raise ValueError("provider_affinity_header must be a valid HTTP header name")
if header.lower() in _FORBIDDEN_AFFINITY_HEADERS:
raise ValueError("provider_affinity_header cannot be an authentication, cookie, or transport header")
return header
def _get_value(value: object, key: str) -> object | None:
if isinstance(value, Mapping):
return value.get(key)
return getattr(value, key, None)
def _get_provider_affinity_header_name(litellm_params: object | None) -> str | None:
header: Final = _get_value(litellm_params, "provider_affinity_header") if litellm_params is not None else None
if header is None:
return None
if not isinstance(header, str):
raise TypeError("provider_affinity_header must be a string")
return validate_provider_affinity_header_name(header)
def get_stable_session_id(litellm_params: object | None) -> str | None:
if litellm_params is None:
return None
direct_session_id: Final = _get_value(litellm_params, "session_id")
if direct_session_id:
return str(direct_session_id)
metadata_values: Final[tuple[object, ...]] = tuple(
value for key in ("metadata", "litellm_metadata") if (value := _get_value(litellm_params, key)) is not None
)
has_generated_session_id: Final = any(
isinstance(metadata, Mapping) and metadata.get(SESSION_ID_GENERATED_METADATA_KEY)
for metadata in metadata_values
)
litellm_session_id: Final = _get_value(litellm_params, "litellm_session_id")
if litellm_session_id and not has_generated_session_id:
return str(litellm_session_id)
for metadata in metadata_values:
if (
isinstance(metadata, Mapping)
and not metadata.get(SESSION_ID_GENERATED_METADATA_KEY)
and (value := metadata.get("session_id"))
):
return str(value)
return None
def add_provider_affinity_header( # mutable-ok: downstream handlers add auth and signing headers
headers: Mapping[str, object], litellm_params: object | None
) -> dict[str, object]: # mutable-ok: downstream handlers add auth and signing headers
header_name: Final = _get_provider_affinity_header_name(litellm_params)
if header_name is None or any(key.lower() == header_name.lower() for key in headers):
return dict(headers) # mutable-ok: downstream handlers add auth and signing headers
session_id: Final = get_stable_session_id(litellm_params)
if session_id is None:
return dict(headers) # mutable-ok: downstream handlers add auth and signing headers
if any(character in session_id for character in ("\r", "\n", "\0")):
raise ValueError("session_id cannot contain HTTP header control characters")
return {**headers, header_name: session_id} # mutable-ok: downstream handlers add auth and signing headers

View file

@ -1711,7 +1711,7 @@ class HTTPHandler:
def get_async_httpx_client(
llm_provider: LlmProviders | httpxSpecialProvider,
llm_provider: LlmProviders | httpxSpecialProvider | str,
params: dict | None = None,
shared_session: Optional["ClientSession"] = None,
) -> AsyncHTTPHandler:

View file

@ -2857,7 +2857,7 @@ class BaseLLMHTTPHandler:
id(shared_session) if shared_session else None,
)
async_httpx_client = get_async_httpx_client(
llm_provider=litellm.LlmProviders(custom_llm_provider),
llm_provider=custom_llm_provider,
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
shared_session=shared_session,
)

View file

@ -78,8 +78,9 @@ from litellm.litellm_core_utils.chat_completion_agentic_loop import (
from litellm.litellm_core_utils.completion_timeout import CompletionTimeout
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.get_litellm_params import (
FORWARDED_KWARGS_KEYS,
AWS_CREDENTIAL_KWARGS_KEYS,
OPTIONAL_KWARGS_KEYS,
PROVIDER_AFFINITY_HEADER_KWARG_KEY,
)
from litellm.litellm_core_utils.get_provider_specific_headers import (
ProviderSpecificHeaderUtils,
@ -96,6 +97,7 @@ from litellm.litellm_core_utils.mock_functions import (
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_content_from_model_response,
)
from litellm.litellm_core_utils.provider_affinity import add_provider_affinity_header
from litellm.litellm_core_utils.request_timeout_resolver import (
get_configured_request_timeout,
)
@ -5644,8 +5646,30 @@ def completion(
gigachat_scope=kwargs.get("gigachat_scope"),
gigachat_auth_url=kwargs.get("gigachat_auth_url"),
gigachat_access_token=kwargs.get("gigachat_access_token"),
**{key: kwargs[key] for key in FORWARDED_KWARGS_KEYS if key in kwargs},
**{
key: kwargs[key]
for key in (*AWS_CREDENTIAL_KWARGS_KEYS, PROVIDER_AFFINITY_HEADER_KWARG_KEY)
if key in kwargs
},
)
if litellm_params.get("provider_affinity_header") is not None:
try:
headers = add_provider_affinity_header(
headers=headers or litellm.headers or MappingProxyType({}),
litellm_params=MappingProxyType(
{
"provider_affinity_header": litellm_params["provider_affinity_header"],
"litellm_session_id": kwargs.get("litellm_session_id"),
"session_id": kwargs.get("session_id"),
"metadata": metadata,
"litellm_metadata": kwargs.get("litellm_metadata"),
}
),
)
except ValueError as affinity_error:
raise litellm.BadRequestError(
message=str(affinity_error), model=model, llm_provider=custom_llm_provider
) from affinity_error
cast(LiteLLMLoggingObj, logging).update_environment_variables(
model=model,
user=user,

View file

@ -26,6 +26,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
update_responses_input_with_model_file_ids,
update_responses_tools_with_model_file_ids,
)
from litellm.litellm_core_utils.provider_affinity import add_provider_affinity_header
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.llms.openai_like.responses.transformation import OpenAILikeResponsesConfig
@ -1218,6 +1219,27 @@ def responses(
# get llm provider logic
litellm_params: Final = GenericLiteLLMParams(**kwargs)
try:
effective_extra_headers: Final = (
add_provider_affinity_header(
headers=extra_headers or MappingProxyType({}),
litellm_params=MappingProxyType(
{
"provider_affinity_header": litellm_params.provider_affinity_header,
"litellm_session_id": kwargs.get("litellm_session_id"),
"session_id": kwargs.get("session_id"),
"metadata": metadata,
"litellm_metadata": kwargs.get("litellm_metadata"),
}
),
)
if litellm_params.provider_affinity_header is not None
else extra_headers
)
except ValueError as affinity_error:
raise litellm.BadRequestError(
message=str(affinity_error), model=model, llm_provider=custom_llm_provider
) from affinity_error
#########################################################
# MOCK RESPONSE LOGIC
@ -1261,7 +1283,7 @@ def responses(
top_p=top_p,
truncation=truncation,
user=user,
extra_headers=extra_headers,
extra_headers=effective_extra_headers,
extra_query=extra_query,
extra_body=extra_body,
timeout=timeout,
@ -1332,7 +1354,7 @@ def responses(
safety_identifier=safety_identifier,
text_format=text_format,
allowed_openai_params=allowed_openai_params,
extra_headers=extra_headers,
extra_headers=effective_extra_headers,
extra_query=extra_query,
extra_body=extra_body,
timeout=timeout,
@ -1352,7 +1374,7 @@ def responses(
custom_llm_provider=custom_llm_provider,
_is_async=_is_async,
stream=stream,
extra_headers=extra_headers,
extra_headers=effective_extra_headers,
extra_body=extra_body,
timeout=timeout if timeout is not None else request_timeout,
allowed_openai_params=allowed_openai_params,
@ -1381,6 +1403,7 @@ def responses(
"model_info": kwargs.get("model_info"),
"data_residency": infer_openai_data_residency(custom_llm_provider, litellm_params.api_base),
"metadata": (kwargs["litellm_metadata"] if "litellm_metadata" in kwargs else kwargs.get("metadata")),
"provider_affinity_header": litellm_params.provider_affinity_header,
},
custom_llm_provider=custom_llm_provider,
)
@ -1400,7 +1423,7 @@ def responses(
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params,
logging_obj=litellm_logging_obj,
extra_headers=extra_headers,
extra_headers=effective_extra_headers,
extra_body=extra_body,
timeout=timeout or request_timeout,
_is_async=_is_async,

View file

@ -16,6 +16,7 @@ from typing_extensions import Protocol, ReadOnly, Required, TypedDict, runtime_c
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
from litellm.litellm_core_utils.provider_affinity import validate_provider_affinity_header_name
from litellm.types.router_weights import RouterWeights
if TYPE_CHECKING:
@ -396,6 +397,7 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
organization: str | None = None # for openai orgs
configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS = None
litellm_credential_name: str | None = None
provider_affinity_header: str | None = None
## LOGGING PARAMS ##
litellm_trace_id: str | None = None
@ -467,6 +469,13 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
valkey_text_field: str | None = None
valkey_embedding_field: str | None = None
@field_validator("provider_affinity_header")
@classmethod
def validate_provider_affinity_header(cls, value: str | None) -> str | None:
if value is None:
return None
return validate_provider_affinity_header_name(value)
@model_validator(mode="before")
@classmethod
def preprocess_input_data(cls, data: object) -> object:
@ -569,6 +578,7 @@ class LiteLLMParamsTypedDict(TypedDict, total=False):
stream_timeout: float | str | None
max_retries: int | None
organization: list | str | None # for openai orgs
provider_affinity_header: ReadOnly[str | None]
configurable_clientside_auth_params: (
CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS # for allowing api base switching on finetuned models
)

View file

@ -4073,6 +4073,7 @@ all_litellm_params = (
"litellm_credential_name",
"allowed_openai_params",
"litellm_session_id",
"provider_affinity_header",
"use_litellm_proxy",
"use_chat_completions_api",
"rust",

View file

@ -0,0 +1,107 @@
import pytest
from litellm.constants import SESSION_ID_GENERATED_METADATA_KEY
from litellm.litellm_core_utils.provider_affinity import (
add_provider_affinity_header,
get_stable_session_id,
)
@pytest.mark.parametrize(
("litellm_params", "expected"),
[
({"litellm_session_id": "litellm-session"}, "litellm-session"),
({"session_id": "direct-session"}, "direct-session"),
({"metadata": {"session_id": "metadata-session"}}, "metadata-session"),
({"litellm_metadata": {"session_id": "litellm-metadata-session"}}, "litellm-metadata-session"),
],
)
def test_get_stable_session_id_uses_explicit_session_sources(litellm_params: dict, expected: str):
assert get_stable_session_id(litellm_params) == expected
def test_get_stable_session_id_does_not_use_trace_id():
assert get_stable_session_id({"litellm_trace_id": "per-request-trace"}) is None
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
def test_get_stable_session_id_ignores_proxy_generated_session(metadata_key: str):
assert (
get_stable_session_id(
{
"litellm_session_id": "generated-session",
metadata_key: {
"session_id": "generated-session",
SESSION_ID_GENERATED_METADATA_KEY: True,
},
}
)
is None
)
def test_get_stable_session_id_prefers_explicit_session_over_proxy_generated_session():
assert (
get_stable_session_id(
{
"session_id": "explicit-session",
"litellm_session_id": "generated-session",
"metadata": {
"session_id": "generated-session",
SESSION_ID_GENERATED_METADATA_KEY: True,
},
}
)
== "explicit-session"
)
def test_add_provider_affinity_header_maps_session_id():
headers = add_provider_affinity_header(
headers={"Content-Type": "application/json"},
litellm_params={
"litellm_session_id": "session-123",
"provider_affinity_header": "X-Conversation-Id",
},
)
assert headers == {
"Content-Type": "application/json",
"X-Conversation-Id": "session-123",
}
def test_add_provider_affinity_header_preserves_explicit_header_case_insensitively():
headers = add_provider_affinity_header(
headers={"x-conversation-id": "explicit-session"},
litellm_params={
"litellm_session_id": "session-123",
"provider_affinity_header": "X-Conversation-Id",
},
)
assert headers == {"x-conversation-id": "explicit-session"}
@pytest.mark.parametrize("session_id", ["session\r", "session\n", "session\0"])
def test_add_provider_affinity_header_rejects_control_characters(session_id: str):
with pytest.raises(ValueError, match="session_id cannot contain HTTP header control characters"):
add_provider_affinity_header(
headers={},
litellm_params={
"litellm_session_id": session_id,
"provider_affinity_header": "X-Conversation-Id",
},
)
def test_add_provider_affinity_header_does_nothing_without_config_or_session():
assert add_provider_affinity_header({}, {"litellm_session_id": "session-123"}) == {}
assert (
add_provider_affinity_header(
{},
{"provider_affinity_header": "X-Conversation-Id"},
)
== {}
)

View file

@ -0,0 +1,389 @@
import json
from unittest.mock import MagicMock
import httpx
import pytest
import litellm
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.openai_like import dynamic_config
from litellm.llms.openai_like.json_loader import JSONProviderRegistry, SimpleProviderConfig
def _provider(*, responses: bool = False) -> SimpleProviderConfig:
endpoints = ["/v1/chat/completions"]
if responses:
endpoints.append("/v1/responses")
return SimpleProviderConfig(
"db_only_provider",
{
"base_url": "https://db-only.example/v1",
"api_key_env": "DYNAMIC_PROVIDER_API_KEY",
"supported_endpoints": endpoints,
},
)
def _chat_response_payload(content: str = "dynamic response") -> dict[str, object]:
return {
"id": "chatcmpl-dynamic-provider",
"object": "chat.completion",
"created": 1234567890,
"model": "test-model",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": content},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 2, "completion_tokens": 2, "total_tokens": 4},
}
def _responses_payload() -> dict[str, object]:
return {
"id": "resp_dynamic_provider",
"object": "response",
"created_at": 1234567890,
"status": "completed",
"model": "test-model",
"output": [],
"parallel_tool_calls": True,
"usage": {"input_tokens": 2, "output_tokens": 2, "total_tokens": 4},
"error": None,
}
@pytest.fixture(autouse=True)
def _isolate_registry_state():
original_providers = dict(JSONProviderRegistry._providers)
dynamic_config._responses_config_cache.clear()
yield
JSONProviderRegistry._providers = original_providers
dynamic_config._responses_config_cache.clear()
def test_dynamic_provider_receives_affinity_header_for_chat():
from openai import OpenAI
JSONProviderRegistry._providers = {"db_only_provider": _provider()}
requests: list[httpx.Request] = []
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(200, json=_chat_response_payload())
client = OpenAI(
api_key="test-key",
base_url="https://db-only.example/v1",
http_client=httpx.Client(transport=httpx.MockTransport(respond)),
)
try:
litellm.completion(
model="db_only_provider/test-model",
messages=[{"role": "user", "content": "hello"}],
api_key="test-key",
client=client,
extra_headers={"X-Customer-Header": "customer-value"},
litellm_session_id="session-sync",
provider_affinity_header="X-Conversation-Id",
)
finally:
client.close()
assert requests[0].headers["x-conversation-id"] == "session-sync"
assert requests[0].headers["x-customer-header"] == "customer-value"
def test_dynamic_provider_does_not_use_trace_id_for_chat_affinity():
from openai import OpenAI
JSONProviderRegistry._providers = {"db_only_provider": _provider()}
requests: list[httpx.Request] = []
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(200, json=_chat_response_payload())
client = OpenAI(
api_key="test-key",
base_url="https://db-only.example/v1",
http_client=httpx.Client(transport=httpx.MockTransport(respond)),
)
try:
litellm.completion(
model="db_only_provider/test-model",
messages=[{"role": "user", "content": "hello"}],
api_key="test-key",
client=client,
metadata={"trace_id": "per-request-trace"},
provider_affinity_header="X-Conversation-Id",
)
finally:
client.close()
assert "x-conversation-id" not in requests[0].headers
@pytest.mark.asyncio
async def test_dynamic_provider_receives_affinity_header_for_async_chat():
from openai import AsyncOpenAI
JSONProviderRegistry._providers = {"db_only_provider": _provider()}
requests: list[httpx.Request] = []
async def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(200, json=_chat_response_payload("async response"))
client = AsyncOpenAI(
api_key="test-key",
base_url="https://db-only.example/v1",
http_client=httpx.AsyncClient(transport=httpx.MockTransport(respond)),
)
try:
response = await litellm.acompletion(
model="db_only_provider/test-model",
messages=[{"role": "user", "content": "hello"}],
api_key="test-key",
client=client,
litellm_session_id="session-async",
provider_affinity_header="X-Conversation-Id",
)
finally:
await client.close()
assert response.choices[0].message.content == "async response"
assert requests[0].headers["x-conversation-id"] == "session-async"
def test_dynamic_provider_receives_affinity_header_for_streaming_chat():
from openai import OpenAI
JSONProviderRegistry._providers = {"db_only_provider": _provider()}
chunks = [
{
"id": "chatcmpl-dynamic-stream",
"object": "chat.completion.chunk",
"created": 1234567890,
"model": "test-model",
"choices": [
{
"index": 0,
"delta": {"role": "assistant", "content": "streamed"},
"finish_reason": None,
}
],
},
{
"id": "chatcmpl-dynamic-stream",
"object": "chat.completion.chunk",
"created": 1234567890,
"model": "test-model",
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
},
]
stream_body = "".join(f"data: {json.dumps(chunk)}\n\n" for chunk in chunks) + "data: [DONE]\n\n"
requests: list[httpx.Request] = []
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(200, content=stream_body, headers={"content-type": "text/event-stream"})
client = OpenAI(
api_key="test-key",
base_url="https://db-only.example/v1",
http_client=httpx.Client(transport=httpx.MockTransport(respond)),
)
try:
response_chunks = list(
litellm.completion(
model="db_only_provider/test-model",
messages=[{"role": "user", "content": "hello"}],
api_key="test-key",
client=client,
stream=True,
litellm_session_id="session-stream",
provider_affinity_header="X-Conversation-Id",
)
)
finally:
client.close()
assert any(chunk.choices[0].delta.content == "streamed" for chunk in response_chunks)
assert requests[0].headers["x-conversation-id"] == "session-stream"
def test_dynamic_provider_receives_affinity_header_for_responses():
JSONProviderRegistry._providers = {"db_only_provider": _provider(responses=True)}
logging_obj = MagicMock()
requests: list[httpx.Request] = []
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(200, json=_responses_payload())
http_client = httpx.Client(transport=httpx.MockTransport(respond))
try:
litellm.responses(
model="db_only_provider/test-model",
input="hello",
api_key="test-key",
litellm_session_id="session-responses",
provider_affinity_header="X-Conversation-Id",
litellm_logging_obj=logging_obj,
client=HTTPHandler(client=http_client),
)
finally:
http_client.close()
assert requests[0].headers["X-Conversation-Id"] == "session-responses"
assert (
logging_obj.update_from_kwargs.call_args.kwargs["litellm_params"]["provider_affinity_header"]
== "X-Conversation-Id"
)
def test_dynamic_provider_uses_metadata_session_id_for_responses():
JSONProviderRegistry._providers = {"db_only_provider": _provider(responses=True)}
requests: list[httpx.Request] = []
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(200, json=_responses_payload())
http_client = httpx.Client(transport=httpx.MockTransport(respond))
try:
litellm.responses(
model="db_only_provider/test-model",
input="hello",
api_key="test-key",
metadata={"session_id": "session-from-metadata"},
provider_affinity_header="X-Conversation-Id",
litellm_logging_obj=MagicMock(),
client=HTTPHandler(client=http_client),
)
finally:
http_client.close()
assert requests[0].headers["X-Conversation-Id"] == "session-from-metadata"
@pytest.mark.asyncio
@pytest.mark.parametrize("stream", [False, True], ids=["non-streaming", "streaming"])
async def test_dynamic_provider_receives_affinity_header_for_async_responses(stream: bool):
JSONProviderRegistry._providers = {"db_only_provider": _provider(responses=True)}
requests: list[httpx.Request] = []
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(200, json=_responses_payload())
client = AsyncHTTPHandler()
await client.close()
client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond))
try:
response = await litellm.aresponses(
model="db_only_provider/test-model",
input="hello",
api_key="test-key",
stream=stream,
litellm_session_id="session-async-responses",
provider_affinity_header="X-Conversation-Id",
client=client,
)
finally:
await client.client.aclose()
assert requests[0].headers["X-Conversation-Id"] == "session-async-responses"
if stream:
assert hasattr(response, "__aiter__")
else:
assert getattr(response, "model", None) == "test-model"
def test_control_characters_in_session_id_are_a_bad_request_for_chat():
from openai import OpenAI
JSONProviderRegistry._providers = {"db_only_provider": _provider()}
requests: list[httpx.Request] = []
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(200, json=_chat_response_payload())
client = OpenAI(
api_key="test-key",
base_url="https://db-only.example/v1",
http_client=httpx.Client(transport=httpx.MockTransport(respond)),
)
try:
with pytest.raises(litellm.BadRequestError, match="HTTP header control characters"):
litellm.completion(
model="db_only_provider/test-model",
messages=[{"role": "user", "content": "hello"}],
api_key="test-key",
client=client,
litellm_session_id="session\nsplit",
provider_affinity_header="X-Conversation-Id",
)
finally:
client.close()
assert requests == []
def test_control_characters_in_session_id_are_a_bad_request_for_responses():
JSONProviderRegistry._providers = {"db_only_provider": _provider(responses=True)}
requests: list[httpx.Request] = []
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(200, json=_responses_payload())
http_client = httpx.Client(transport=httpx.MockTransport(respond))
try:
with pytest.raises(litellm.BadRequestError, match="HTTP header control characters"):
litellm.responses(
model="db_only_provider/test-model",
input="hello",
api_key="test-key",
litellm_session_id="session\nsplit",
provider_affinity_header="X-Conversation-Id",
litellm_logging_obj=MagicMock(),
client=HTTPHandler(client=http_client),
)
finally:
http_client.close()
assert requests == []
def test_builtin_provider_receives_affinity_header():
from openai import OpenAI
requests: list[httpx.Request] = []
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(200, json=_chat_response_payload())
client = OpenAI(
api_key="test-key",
base_url="https://api.openai.com/v1",
http_client=httpx.Client(transport=httpx.MockTransport(respond)),
)
try:
litellm.completion(
model="openai/test-model",
messages=[{"role": "user", "content": "hello"}],
api_key="test-key",
client=client,
litellm_session_id="session-builtin",
provider_affinity_header="X-Conversation-Id",
)
finally:
client.close()
assert requests[0].headers["X-Conversation-Id"] == "session-builtin"

View file

@ -66,6 +66,30 @@ class TestUseResponsesApiBridgeFlag:
mock_bridge_handler.assert_called_once()
@patch.object(
import_module("litellm.responses.main").litellm_completion_transformation_handler, "response_api_handler"
)
@patch.object(
import_module("litellm.responses.main").ProviderConfigManager, "get_provider_responses_api_config"
)
def test_provider_affinity_header_is_forwarded_through_bridge(self, mock_get_config, mock_bridge_handler):
mock_get_config.return_value = litellm.OpenAIResponsesAPIConfig()
mock_bridge_handler.return_value = MagicMock()
litellm.responses(
model="openai/my-custom-model",
input="Hello",
use_chat_completions_api=True,
litellm_session_id="session-bridge",
provider_affinity_header="X-Conversation-Id",
extra_headers={"X-Customer-Header": "customer-value"},
litellm_logging_obj=MagicMock(),
)
forwarded_headers = mock_bridge_handler.call_args.kwargs["extra_headers"]
assert forwarded_headers["X-Conversation-Id"] == "session-bridge"
assert forwarded_headers["X-Customer-Header"] == "customer-value"
@patch.object(
import_module("litellm.responses.main").litellm_completion_transformation_handler, "response_api_handler"
)

View file

@ -91,7 +91,7 @@ def test_pricing_strings_are_coerced_to_float():
def test_invalid_pricing_is_rejected():
with pytest.raises(ValueError, match='validation error for ModelInfo'):
with pytest.raises(ValueError, match="validation error for ModelInfo"):
ModelInfo(id="x", input_cost_per_token="free")
@ -118,7 +118,9 @@ def test_drop_params_ignores_non_flag_non_string_values_with_a_warning(value, ca
assert f"drop_params={value!r} is not a flag value" in caplog.text
@pytest.mark.parametrize("value", [True, "true", None, "os.environ/DROP_PARAMS", "v2:gcm:ciphertext-from-a-pre-fix-row"])
@pytest.mark.parametrize(
"value", [True, "true", None, "os.environ/DROP_PARAMS", "v2:gcm:ciphertext-from-a-pre-fix-row"]
)
def test_drop_params_flags_and_strings_log_nothing(value, caplog):
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
GenericLiteLLMParams(drop_params=value)
@ -148,6 +150,46 @@ def test_aws_session_tags_reject_shapes_sts_would_refuse(aws_session_tags):
LiteLLM_Params(model="bedrock/anthropic.claude-opus-5", aws_session_tags=aws_session_tags)
def test_provider_affinity_header_is_normalized():
params = LiteLLM_Params(
model="openai/gpt-4o-mini",
provider_affinity_header="X-Conversation-Id",
)
assert params.provider_affinity_header == "X-Conversation-Id"
assert params.model_dump(exclude_none=True)["provider_affinity_header"] == "X-Conversation-Id"
@pytest.mark.parametrize(
"header",
[
"Authorization",
"Proxy-Authorization",
"Cookie",
"Set-Cookie",
"Host",
"Content-Length",
"Content-Type",
"X-API-Key",
],
)
def test_provider_affinity_header_rejects_sensitive_or_transport_headers(header: str):
with pytest.raises(ValueError, match="provider_affinity_header"):
LiteLLM_Params(
model="openai/gpt-4o-mini",
provider_affinity_header=header,
)
@pytest.mark.parametrize("header", ["", "X Conversation Id", "X-Conversation-Id\r\nInjected: true"])
def test_provider_affinity_header_rejects_invalid_header_names(header: str):
with pytest.raises(ValueError, match="provider_affinity_header"):
LiteLLM_Params(
model="openai/gpt-4o-mini",
provider_affinity_header=header,
)
def test_model_info_parses_access_windows_time_strings():
import datetime

View file

@ -31476,6 +31476,8 @@ export interface components {
output_cost_per_video_token?: number | null;
/** Output Vector Size */
output_vector_size?: number | null;
/** Provider Affinity Header */
provider_affinity_header?: string | null;
/** Quality Router Config */
quality_router_config?: {
[key: string]: unknown;
@ -42364,6 +42366,8 @@ export interface components {
output_cost_per_video_token?: number | null;
/** Output Vector Size */
output_vector_size?: number | null;
/** Provider Affinity Header */
provider_affinity_header?: string | null;
/** Quality Router Config */
quality_router_config?: {
[key: string]: unknown;