feat(litellm_pre_call_utils.py): use claude code session for litellm session id

allows claude code logs to be stitched together, making it easy to know they were all part of the same conversation
This commit is contained in:
Krrish Dholakia 2026-04-15 15:57:46 -07:00
parent 11e22bdd78
commit 37d6a18ce5
2 changed files with 118 additions and 24 deletions

View file

@ -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-<something>-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-<vendor>-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-<vendor>-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():

View file

@ -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-<vendor>-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-<vendor>-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"},