simplify: use full traceparent value as session_id instead of regex extraction

Drop the _TRACEPARENT_RE regex and _extract_trace_id_from_traceparent
helper. Just check for the 'traceparent' header directly and use the
full value as the session_id.

Co-authored-by: Krrish Dholakia <krrish-berri-2@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-05-05 00:32:57 +00:00
parent c2aa2f2255
commit 5e52132381
No known key found for this signature in database
2 changed files with 14 additions and 72 deletions

View file

@ -39,13 +39,6 @@ _EXPLICIT_SESSION_HEADERS = frozenset({"x-litellm-trace-id", "x-litellm-session-
# (covers UUIDs and most common session-id formats).
_SESSION_ID_VALUE_RE = re.compile(r"^[a-zA-Z0-9_\-]{8,}$")
# W3C Trace Context traceparent format:
# {version}-{trace-id}-{parent-id}-{trace-flags}
# e.g. 00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01
_TRACEPARENT_RE = re.compile(
r"^[0-9a-f]{2}-([0-9a-f]{32})-[0-9a-f]{16}-[0-9a-f]{2}$", re.IGNORECASE
)
def _sanitize_for_log(value: Any) -> str:
"""
@ -316,27 +309,6 @@ def _extract_generic_session_id_from_headers(
return None
def _extract_trace_id_from_traceparent(
normalized: Dict[str, str],
) -> Optional[str]:
"""
Extract the ``trace-id`` component from a W3C ``traceparent`` header.
The traceparent format is ``{version}-{trace-id}-{parent-id}-{trace-flags}``,
e.g. ``00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01``.
Returns the 32-hex-char trace-id if the header is present and well-formed,
otherwise ``None``.
"""
traceparent = normalized.get("traceparent")
if not traceparent or not isinstance(traceparent, str):
return None
match = _TRACEPARENT_RE.match(traceparent.strip())
if match:
return match.group(1)
return None
def get_chain_id_from_headers(headers: Optional[Dict[str, str]]) -> Optional[str]:
"""
Extract chain id for call chaining from request headers.
@ -346,9 +318,9 @@ def get_chain_id_from_headers(headers: Optional[Dict[str, str]]) -> Optional[str
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``.
4. W3C ``traceparent`` header — the 32-hex-char trace-id is extracted and
used as the session id. This allows chaining LLM calls that are part
of one agent / distributed-trace interaction.
4. W3C ``traceparent`` header — the full value is used as the session id.
This allows chaining LLM calls that are part of one agent /
distributed-trace interaction.
Header keys are matched case-insensitively so this works with raw header
dicts from any transport.
@ -363,7 +335,8 @@ def get_chain_id_from_headers(headers: Optional[Dict[str, str]]) -> Optional[str
normalized.get("x-litellm-trace-id")
or normalized.get("x-litellm-session-id")
or _extract_generic_session_id_from_headers(normalized)
or _extract_trace_id_from_traceparent(normalized)
or normalized.get("traceparent")
or None
)

View file

@ -2277,14 +2277,11 @@ def test_get_chain_id_from_headers_generic_vendor_session_id():
def test_get_chain_id_from_headers_traceparent():
"""get_chain_id_from_headers extracts trace-id from a W3C traceparent header."""
"""get_chain_id_from_headers uses the full traceparent value as session id."""
from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers
traceparent = "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01"
assert (
get_chain_id_from_headers({"traceparent": traceparent})
== "0af7651916cd43dd8448eb211c80319c"
)
assert get_chain_id_from_headers({"traceparent": traceparent}) == traceparent
def test_get_chain_id_from_headers_traceparent_case_insensitive():
@ -2292,19 +2289,14 @@ def test_get_chain_id_from_headers_traceparent_case_insensitive():
from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers
traceparent = "00-abcdef1234567890abcdef1234567890-1234567890abcdef-00"
assert (
get_chain_id_from_headers({"Traceparent": traceparent})
== "abcdef1234567890abcdef1234567890"
)
assert get_chain_id_from_headers({"Traceparent": traceparent}) == traceparent
def test_get_chain_id_from_headers_traceparent_malformed_ignored():
"""Malformed traceparent values are ignored."""
def test_get_chain_id_from_headers_traceparent_empty_ignored():
"""Empty traceparent value is ignored."""
from litellm.proxy.litellm_pre_call_utils import get_chain_id_from_headers
assert get_chain_id_from_headers({"traceparent": "not-a-valid-traceparent"}) is None
assert get_chain_id_from_headers({"traceparent": ""}) is None
assert get_chain_id_from_headers({"traceparent": "00-short-abc-01"}) is None
def test_get_chain_id_from_headers_traceparent_lower_priority_than_explicit():
@ -2345,33 +2337,10 @@ def test_add_litellm_metadata_from_request_headers_traceparent_sets_session_id()
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=data, _metadata_variable_name="metadata"
)
expected_trace_id = "0af7651916cd43dd8448eb211c80319c"
assert data["metadata"]["session_id"] == expected_trace_id
assert data["metadata"]["trace_id"] == expected_trace_id
assert data["litellm_session_id"] == expected_trace_id
assert data["litellm_trace_id"] == expected_trace_id
def test_extract_trace_id_from_traceparent_directly():
"""Unit test for _extract_trace_id_from_traceparent helper."""
from litellm.proxy.litellm_pre_call_utils import _extract_trace_id_from_traceparent
assert (
_extract_trace_id_from_traceparent(
{"traceparent": "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01"}
)
== "0af7651916cd43dd8448eb211c80319c"
)
assert _extract_trace_id_from_traceparent({}) is None
assert _extract_trace_id_from_traceparent({"traceparent": ""}) is None
assert _extract_trace_id_from_traceparent({"traceparent": "invalid-format"}) is None
# Uppercase hex should also work
assert (
_extract_trace_id_from_traceparent(
{"traceparent": "00-0AF7651916CD43DD8448EB211C80319C-B7AD6B7169203331-01"}
)
== "0AF7651916CD43DD8448EB211C80319C"
)
assert data["metadata"]["session_id"] == traceparent
assert data["metadata"]["trace_id"] == traceparent
assert data["litellm_session_id"] == traceparent
assert data["litellm_trace_id"] == traceparent
def test_get_internal_user_header_from_mapping_returns_expected_header():