diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 4ec31925ea2..c4f635b69b5 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -1,5 +1,6 @@ import asyncio import copy +import re import time from collections import OrderedDict from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union @@ -28,6 +29,14 @@ _SPECIAL_HEADERS_CACHE = frozenset( v.value.lower() for v in SpecialHeaders._member_map_.values() ) +# Matches any header of the form x--session-id (case-insensitive). +# Excludes the two explicit litellm headers which are handled with higher priority. +_GENERIC_SESSION_ID_HEADER_RE = re.compile(r"^x-.+-session-id$", re.IGNORECASE) +_EXPLICIT_SESSION_HEADERS = frozenset({"x-litellm-trace-id", "x-litellm-session-id"}) +# Session-id values must be non-empty strings of alphanumerics, hyphens, or underscores +# (covers UUIDs and most common session-id formats). +_SESSION_ID_VALUE_RE = re.compile(r"^[a-zA-Z0-9_\-]{8,}$") + def _sanitize_for_log(value: Any) -> str: """ @@ -43,6 +52,8 @@ def _sanitize_for_log(value: Any) -> str: text = repr(value) # Strip CR/LF characters commonly used for log injection return text.replace("\r", "").replace("\n", "") + + from litellm.router import Router from litellm.secret_managers.main import get_secret_bool from litellm.types.llms.anthropic import ANTHROPIC_API_HEADERS @@ -113,13 +124,43 @@ def _get_metadata_variable_name(request: Request) -> str: return "metadata" +def _extract_generic_session_id_from_headers( + normalized: Dict[str, str], +) -> Optional[str]: + """ + Scan a normalised (lower-cased keys) header dict for any header that looks + like ``x--session-id`` and whose value is a plausible session/trace + identifier (alphanumeric + hyphens/underscores, at least 8 chars). + + The two explicit LiteLLM headers (``x-litellm-trace-id`` / + ``x-litellm-session-id``) are excluded here because they are handled with + higher priority by the caller. + + Example: ``x-claude-code-session-id: e96634a3-fa28-4083-b354-55542e2dca01`` + """ + for key, value in normalized.items(): + if ( + key not in _EXPLICIT_SESSION_HEADERS + and _GENERIC_SESSION_ID_HEADER_RE.match(key) + and isinstance(value, str) + and _SESSION_ID_VALUE_RE.match(value) + ): + return value + return None + + def get_chain_id_from_headers(headers: Optional[Dict[str, str]]) -> Optional[str]: """ Extract chain id for call chaining from request headers. - x-litellm-trace-id and x-litellm-session-id are interchangeable; when both - are present, x-litellm-trace-id takes precedence. Header keys are matched - case-insensitively so this works with raw header dicts from any transport. + Priority order: + 1. ``x-litellm-trace-id`` (explicit, highest priority) + 2. ``x-litellm-session-id`` (explicit) + 3. Any ``x--session-id`` header whose value looks like a session id + (alphanumeric / UUID, at least 8 chars). E.g. ``x-claude-code-session-id``. + + Header keys are matched case-insensitively so this works with raw header + dicts from any transport. Used by MCP (and other paths that have raw_headers but no Request) to set litellm_trace_id/litellm_session_id for spend logs and logging consistency. @@ -127,8 +168,10 @@ def get_chain_id_from_headers(headers: Optional[Dict[str, str]]) -> Optional[str if not headers: return None normalized = {k.lower(): v for k, v in headers.items() if isinstance(k, str)} - return normalized.get("x-litellm-trace-id") or normalized.get( - "x-litellm-session-id" + return ( + normalized.get("x-litellm-trace-id") + or normalized.get("x-litellm-session-id") + or _extract_generic_session_id_from_headers(normalized) ) @@ -220,12 +263,12 @@ def _get_dynamic_logging_metadata( user_api_key_dict: UserAPIKeyAuth, proxy_config: ProxyConfig ) -> Optional[TeamCallbackMetadata]: callback_settings_obj: Optional[TeamCallbackMetadata] = None - key_dynamic_logging_settings: Optional[ - dict - ] = KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(user_api_key_dict) - team_dynamic_logging_settings: Optional[ - dict - ] = KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(user_api_key_dict) + key_dynamic_logging_settings: Optional[dict] = ( + KeyAndTeamLoggingSettings.get_key_dynamic_logging_settings(user_api_key_dict) + ) + team_dynamic_logging_settings: Optional[dict] = ( + KeyAndTeamLoggingSettings.get_team_dynamic_logging_settings(user_api_key_dict) + ) ######################################################################################### # Key-based callbacks ######################################################################################### @@ -647,10 +690,8 @@ class LiteLLMProxyRequestSetup: ######################################################################################### agent_id_from_header = headers.get("x-litellm-agent-id") - # x-litellm-trace-id and x-litellm-session-id are interchangeable for call chaining - chain_id = headers.get("x-litellm-trace-id") or headers.get( - "x-litellm-session-id" - ) + # Explicit litellm headers take precedence; fall back to any x-*-session-id header. + chain_id = get_chain_id_from_headers(dict(headers)) if agent_id_from_header: metadata_from_headers["agent_id"] = agent_id_from_header @@ -779,11 +820,11 @@ class LiteLLMProxyRequestSetup: ## KEY-LEVEL SPEND LOGS / TAGS if "tags" in key_metadata and key_metadata["tags"] is not None: - data[_metadata_variable_name][ - "tags" - ] = LiteLLMProxyRequestSetup._merge_tags( - request_tags=data[_metadata_variable_name].get("tags"), - tags_to_add=key_metadata["tags"], + data[_metadata_variable_name]["tags"] = ( + LiteLLMProxyRequestSetup._merge_tags( + request_tags=data[_metadata_variable_name].get("tags"), + tags_to_add=key_metadata["tags"], + ) ) if "disable_global_guardrails" in key_metadata and isinstance( key_metadata["disable_global_guardrails"], bool @@ -1079,9 +1120,9 @@ async def add_litellm_data_to_request( # noqa: PLR0915 data[_metadata_variable_name]["litellm_api_version"] = version if general_settings is not None: - data[_metadata_variable_name][ - "global_max_parallel_requests" - ] = general_settings.get("global_max_parallel_requests", None) + data[_metadata_variable_name]["global_max_parallel_requests"] = ( + general_settings.get("global_max_parallel_requests", None) + ) ### KEY-LEVEL Controls key_metadata = user_api_key_dict.metadata @@ -1881,7 +1922,9 @@ async def move_guardrails_to_metadata( ) # Only check policy engine if no local config (avoid import + registry lookup) - if not (has_key_config or has_team_config or has_project_config or has_request_config): + if not ( + has_key_config or has_team_config or has_project_config or has_request_config + ): from litellm.proxy.policy_engine.policy_registry import get_policy_registry if not get_policy_registry().is_initialized(): diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index cf7e71b14d4..8763b12f37b 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -1191,6 +1191,57 @@ def test_add_litellm_metadata_from_request_headers_both_headers_trace_id_precede assert data["litellm_trace_id"] == "trace-value" +def test_add_litellm_metadata_from_request_headers_generic_session_id_header(): + """A generic x--session-id header is used when no explicit litellm header is set.""" + headers = {"x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01"} + data = {"metadata": {}} + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert data["metadata"]["session_id"] == "e96634a3-fa28-4083-b354-55542e2dca01" + assert data["litellm_session_id"] == "e96634a3-fa28-4083-b354-55542e2dca01" + assert data["litellm_trace_id"] == "e96634a3-fa28-4083-b354-55542e2dca01" + + +def test_add_litellm_metadata_from_request_headers_explicit_header_beats_generic(): + """Explicit x-litellm-trace-id wins over a generic x-*-session-id header.""" + headers = { + "x-litellm-trace-id": "explicit-trace-id-value", + "x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01", + } + data = {"metadata": {}} + LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers( + headers=headers, data=data, _metadata_variable_name="metadata" + ) + assert data["litellm_session_id"] == "explicit-trace-id-value" + assert data["litellm_trace_id"] == "explicit-trace-id-value" + + +def test_get_chain_id_from_headers_generic_vendor_session_id(): + """get_chain_id_from_headers picks up any x--session-id with a valid value.""" + from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers + + assert ( + get_chain_id_from_headers( + {"x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01"} + ) + == "e96634a3-fa28-4083-b354-55542e2dca01" + ) + # Short / non-alphanumeric values should be ignored + assert get_chain_id_from_headers({"x-foo-session-id": "short"}) is None + assert get_chain_id_from_headers({"x-foo-session-id": "has spaces!!"}) is None + # Explicit headers still take precedence + assert ( + get_chain_id_from_headers( + { + "x-litellm-trace-id": "explicit-id-value", + "x-claude-code-session-id": "e96634a3-fa28-4083-b354-55542e2dca01", + } + ) + == "explicit-id-value" + ) + + def test_get_internal_user_header_from_mapping_returns_expected_header(): mappings = [ {"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"},