diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index 7f373569b21..112d038039e 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -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 diff --git a/litellm/litellm_core_utils/provider_affinity.py b/litellm/litellm_core_utils/provider_affinity.py new file mode 100644 index 00000000000..31cd9a7ff69 --- /dev/null +++ b/litellm/litellm_core_utils/provider_affinity.py @@ -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 diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 8ea46a6b261..312f46bb021 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -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: diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index cbef44ce488..8cbc28362a8 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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, ) diff --git a/litellm/main.py b/litellm/main.py index 7d231a1bf7a..7f4b34d28a0 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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, diff --git a/litellm/responses/main.py b/litellm/responses/main.py index c5032536df4..5c4c9ea3987 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -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, diff --git a/litellm/types/router.py b/litellm/types/router.py index 2353093d435..57bd4263894 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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 ) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index f10f102d8ec..f2c8f0e7044 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -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", diff --git a/tests/test_litellm/litellm_core_utils/test_provider_affinity.py b/tests/test_litellm/litellm_core_utils/test_provider_affinity.py new file mode 100644 index 00000000000..edf4f5169b5 --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/test_provider_affinity.py @@ -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"}, + ) + == {} + ) + diff --git a/tests/test_litellm/llms/openai_like/test_provider_affinity_forwarding.py b/tests/test_litellm/llms/openai_like/test_provider_affinity_forwarding.py new file mode 100644 index 00000000000..43101234e63 --- /dev/null +++ b/tests/test_litellm/llms/openai_like/test_provider_affinity_forwarding.py @@ -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" diff --git a/tests/test_litellm/responses/test_responses_api_bridge_flag.py b/tests/test_litellm/responses/test_responses_api_bridge_flag.py index 16135106b41..642495fab86 100644 --- a/tests/test_litellm/responses/test_responses_api_bridge_flag.py +++ b/tests/test_litellm/responses/test_responses_api_bridge_flag.py @@ -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" ) diff --git a/tests/test_litellm/types/test_router.py b/tests/test_litellm/types/test_router.py index df744a77fe4..4d4c326d1ca 100644 --- a/tests/test_litellm/types/test_router.py +++ b/tests/test_litellm/types/test_router.py @@ -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 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index a53115b7af2..be1e000bd66 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -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;