fix(proxy): keep root trace/session fields on equal ids and usable caller values

The _caller_trace_field gating added for caller-metadata precedence was
presence-aware, which regressed three pre-call behaviors versus main:

- Equal ids in W3C headers and body metadata suppressed the
  traceparent/baggage branch entirely, leaving the root
  litellm_trace_id/litellm_session_id unset so router fallbacks minted a
  fresh uuid4 per request. The header branch now also fires when the
  caller value equals the header-derived id, so equal ids stamp the root
  fields exactly like main.
- A truthy but non-string metadata value (e.g. session_id 4815162342)
  counted as caller-supplied and suppressed the baggage fallback, so
  downstream str-only consumers (code interpreter sandbox reuse) got a
  new sandbox per turn. _caller_trace_field now counts only non-empty
  string values, and the generate policy falls through to generation
  when the caller value is not usable.
- On litellm_metadata routes a usable caller session id suppressed the
  missing_session_id policies while never landing on the root field, so
  the root session stayed unset. The policy now promotes the caller's
  usable session id to litellm_session_id instead of leaving it stranded.
This commit is contained in:
Yucheng He 2026-10-04 01:28:32 -07:00
parent 4d3940cdb6
commit 99cf509980
2 changed files with 129 additions and 28 deletions

View file

@ -150,17 +150,25 @@ def _session_id_from_baggage(baggage: str) -> str | None:
return None
def _caller_trace_field(data: Mapping[str, object], metadata_variable_name: str, field: str) -> object | None:
def _caller_trace_field(data: Mapping[str, object], metadata_variable_name: str, field: str) -> str | None:
"""The caller's value for a trace-control field, counted only when it is a
usable id: a non-empty string. An explicitly empty/unusable value on the
active metadata container still shadows the promoted requester value, but
neither ever counts as "the caller supplied this field" on its own, so a
numeric session id or an empty string cannot suppress the W3C header
fallback or satisfy a missing-session-id policy."""
active: Final = data.get(metadata_variable_name)
if isinstance(active, Mapping) and field in active:
active_map: Final = cast(Mapping[str, object], active) # cast-ok: isinstance above, free-form JSON values
return active_map[field] or None
active_value: Final = active_map[field]
return active_value if isinstance(active_value, str) and active_value else None
promoted: Final = metadata_variable_name == "litellm_metadata" and field in LITELLM_TRACE_CONTROL_METADATA_FIELDS
requester: Final = data.get("metadata")
if not promoted or not isinstance(requester, Mapping):
return None
requester_map: Final = cast(Mapping[str, object], requester) # cast-ok: isinstance above, free-form JSON values
return requester_map.get(field) or None
requester_value: Final = requester_map.get(field)
return requester_value if isinstance(requester_value, str) and requester_value else None
def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None:
@ -841,7 +849,16 @@ def apply_missing_session_id_policy(
):
metadata["session_id"] = body_session_id
return
if data.get("litellm_session_id") or _caller_trace_field(data, _metadata_variable_name, "session_id") is not None:
caller_session_id: Final = _caller_trace_field(data, _metadata_variable_name, "session_id")
if caller_session_id is not None:
# The caller supplied a usable session id, so generate/reject must not
# fire. Surface it on the root field as well: consumers that read
# ``litellm_session_id`` (router fallbacks, spend logs, sandbox reuse)
# otherwise see no session at all and mint a fresh uuid4 per request.
if not data.get("litellm_session_id"):
data["litellm_session_id"] = caller_session_id # rebind-ok: data is an out-param
return
if data.get("litellm_session_id"):
return
match policy:
case "generate":
@ -1597,42 +1614,44 @@ class LiteLLMProxyRequestSetup:
# (https://www.w3.org/TR/trace-context/, https://www.w3.org/TR/baggage/).
# Lower priority than everything above - only fires when neither the
# explicit litellm headers, the Anthropic-metadata path, nor the
# caller's own request metadata set the field - but lets a caller's
# existing traceparent/baggage headers (from real OTel instrumentation)
# correlate with litellm's own logs instead of generating an unrelated
# trace_id.
# caller's own request metadata set the field to a DIFFERENT usable id
# - but lets a caller's existing traceparent/baggage headers (from
# real OTel instrumentation) correlate with litellm's own logs instead
# of generating an unrelated trace_id.
normalized_headers: Final = MappingProxyType({k.lower(): v for k, v in headers.items() if isinstance(k, str)})
if (
"litellm_trace_id" not in data
and _caller_trace_field(
cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object
_metadata_variable_name,
"trace_id",
)
is None
):
if "litellm_trace_id" not in data:
traceparent: Final = normalized_headers.get("traceparent")
if isinstance(traceparent, str):
trace_id_from_traceparent: Final = _trace_id_from_traceparent(traceparent)
if trace_id_from_traceparent:
# The caller's metadata wins over the header fallback unless
# both carry the same id: stamping the root field then claims
# nothing the caller did not already ask for, and keeps the
# W3C-correlated root trace id instead of a generated uuid4.
caller_trace_id: Final = _caller_trace_field(
cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object
_metadata_variable_name,
"trace_id",
)
if trace_id_from_traceparent and (
caller_trace_id is None or caller_trace_id == trace_id_from_traceparent
):
metadata_from_headers["trace_id"] = trace_id_from_traceparent
data["litellm_trace_id"] = trace_id_from_traceparent # rebind-ok: data is an out-param
verbose_proxy_logger.debug(
"Extracted trace_id from W3C traceparent header: %s", trace_id_from_traceparent
)
if (
"litellm_session_id" not in data
and _caller_trace_field(
cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object
_metadata_variable_name,
"session_id",
)
is None
):
if "litellm_session_id" not in data:
baggage: Final = normalized_headers.get("baggage")
if isinstance(baggage, str):
session_id_from_baggage: Final = _session_id_from_baggage(baggage)
if session_id_from_baggage:
caller_session_id: Final = _caller_trace_field(
cast(Mapping[str, object], data), # cast-ok: request body is a str-keyed JSON object
_metadata_variable_name,
"session_id",
)
if session_id_from_baggage and (
caller_session_id is None or caller_session_id == session_id_from_baggage
):
metadata_from_headers["session_id"] = session_id_from_baggage
data["litellm_session_id"] = session_id_from_baggage # rebind-ok: data is an out-param
verbose_proxy_logger.debug("Extracted session_id from W3C baggage header")

View file

@ -3743,6 +3743,36 @@ def test_add_litellm_metadata_from_request_headers_promoted_metadata_beats_heade
assert f"litellm_{field}" not in data
def test_add_litellm_metadata_from_request_headers_equal_ids_still_stamp_root_fields():
headers = {
"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01",
"baggage": "session.id=matching-session-42",
}
data = {
"metadata": {"trace_id": "4bf92f3577b34da6a3ce929d0e0e4736", "session_id": "matching-session-42"},
}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=data, _metadata_variable_name="metadata"
)
assert data["litellm_trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736"
assert data["litellm_session_id"] == "matching-session-42"
assert data["metadata"]["trace_id"] == "4bf92f3577b34da6a3ce929d0e0e4736"
assert data["metadata"]["session_id"] == "matching-session-42"
@pytest.mark.parametrize("non_string_session_id", [4815162342, True, {"session": "nested"}])
def test_add_litellm_metadata_from_request_headers_non_string_body_session_id_falls_back_to_baggage(
non_string_session_id: object,
):
headers = {"baggage": "session.id=header-session-42"}
data = {"metadata": {"session_id": non_string_session_id}}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=data, _metadata_variable_name="metadata"
)
assert data["litellm_session_id"] == "header-session-42"
assert data["metadata"]["session_id"] == "header-session-42"
def _otel_span_with_trace_id(trace_id: int) -> NonRecordingSpan:
return NonRecordingSpan(SpanContext(trace_id=trace_id, span_id=0x00F067AA0BA902B7, is_remote=False))
@ -8456,6 +8486,58 @@ async def test_missing_session_id_generate_reuses_promoted_caller_trace_id(path:
)
@pytest.mark.asyncio
@pytest.mark.parametrize("policy", ["generate", "reject"])
async def test_missing_session_id_policy_promotes_caller_session_to_root_field(policy: str):
"""A caller-supplied usable session id satisfies the missing_session_id policies on
litellm_metadata routes (where it is not yet the managed metadata field) and must also
land on the root ``litellm_session_id`` field: consumers that read the root field
(router fallbacks, spend logs, sandbox reuse) otherwise mint a fresh uuid4 per request."""
request = _request_for("/v1/responses")
updated = await add_litellm_data_to_request(
data={
"model": "gpt-4o",
"input": "hi",
"litellm_trace_id": "root-trace-42",
"metadata": {"session_id": "caller-session-42"},
},
request=request,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": policy},
)
assert updated["litellm_trace_id"] == "root-trace-42"
assert updated["litellm_session_id"] == "caller-session-42"
assert updated["litellm_metadata"]["session_id"] == "caller-session-42"
assert SESSION_ID_GENERATED_METADATA_KEY not in updated["litellm_metadata"]
@pytest.mark.asyncio
async def test_missing_session_id_generate_ignores_non_string_caller_session_id():
"""A non-string session id is not a usable session: the generate policy must fall through
to generation instead of letting an unusable value strand the root session field."""
request = _request_for("/v1/responses")
updated = await add_litellm_data_to_request(
data={
"model": "gpt-4o",
"input": "hi",
"litellm_trace_id": "root-trace-42",
"metadata": {"session_id": 4815162342},
},
request=request,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": "generate"},
)
assert updated["litellm_session_id"] == "root-trace-42"
assert updated["litellm_metadata"]["session_id"] == "root-trace-42"
assert updated["litellm_metadata"][SESSION_ID_GENERATED_METADATA_KEY] is True
@pytest.mark.asyncio
@pytest.mark.parametrize("policy", ["generate", "reject"])
async def test_missing_session_id_policy_keeps_client_supplied_session_id(policy: str):