mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-25 01:02:15 +00:00
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:
parent
666f6b01b4
commit
fc0055497c
13 changed files with 734 additions and 14 deletions
|
|
@ -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
|
||||
|
|
|
|||
98
litellm/litellm_core_utils/provider_affinity.py
Normal file
98
litellm/litellm_core_utils/provider_affinity.py
Normal 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
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
107
tests/test_litellm/litellm_core_utils/test_provider_affinity.py
Normal file
107
tests/test_litellm/litellm_core_utils/test_provider_affinity.py
Normal 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"},
|
||||
)
|
||||
== {}
|
||||
)
|
||||
|
||||
|
|
@ -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"
|
||||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
4
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
4
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue