fix(proxy): count promoted caller trace ids on litellm_metadata routes

Addresses review. On LITELLM_METADATA_ROUTES the proxy promotes the caller's
trace control fields (trace_id, session_id, ...) from `metadata` into
`litellm_metadata`, but only after the header fallback runs. Checking only
`litellm_metadata` let traceparent and baggage fill those fields first, and the
promotion then skipped them because they were already set, so the trace was
recorded under the header ids

The check now also reads `metadata` for fields in
LITELLM_TRACE_CONTROL_METADATA_FIELDS on those routes, and
apply_missing_session_id_policy uses the same check, so `reject` no longer
refuses a request whose session id is about to be promoted

The earlier test asserting that a session_id in `metadata` must not block
baggage on /v1/responses had the premise wrong, since that value is promoted.
It is replaced by tests that assert the promoted caller ids win on
/v1/responses and /v1/messages, with no policy, reject, and generate
This commit is contained in:
Filipe Andujar 2026-09-24 16:29:48 -03:00
parent bae13227fa
commit c2eb718ddd
2 changed files with 44 additions and 12 deletions

View file

@ -150,8 +150,13 @@ def _session_id_from_baggage(baggage: str) -> str | None:
return None
def _client_steered_trace_field(metadata: object, field: str) -> bool:
return isinstance(metadata, Mapping) and bool(cast(Mapping[str, object], metadata).get(field))
def _caller_set_trace_field(data: Mapping[str, object], metadata_variable_name: str, field: str) -> bool:
promoted: Final = metadata_variable_name == "litellm_metadata" and field in LITELLM_TRACE_CONTROL_METADATA_FIELDS
sources: Final = (metadata_variable_name, "metadata") if promoted else (metadata_variable_name,)
return any(
isinstance(container := data.get(name), Mapping) and bool(cast(Mapping[str, object], container).get(field))
for name in sources
)
def _stampable_key_hash(user_api_key_dict: UserAPIKeyAuth) -> str | None:
@ -829,7 +834,7 @@ def apply_missing_session_id_policy(
):
metadata["session_id"] = body_session_id
return
if data.get("litellm_session_id") or metadata.get("session_id"):
if data.get("litellm_session_id") or _caller_set_trace_field(data, _metadata_variable_name, "session_id"):
return
match policy:
case "generate":
@ -1584,8 +1589,8 @@ class LiteLLMProxyRequestSetup:
# 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 not _client_steered_trace_field(
cast(object, data.get(_metadata_variable_name)), "trace_id"
if "litellm_trace_id" not in data and not _caller_set_trace_field(
cast(Mapping[str, object], data), _metadata_variable_name, "trace_id"
):
traceparent: Final = normalized_headers.get("traceparent")
if isinstance(traceparent, str):
@ -1596,8 +1601,8 @@ class LiteLLMProxyRequestSetup:
verbose_proxy_logger.debug(
"Extracted trace_id from W3C traceparent header: %s", trace_id_from_traceparent
)
if "litellm_session_id" not in data and not _client_steered_trace_field(
cast(object, data.get(_metadata_variable_name)), "session_id"
if "litellm_session_id" not in data and not _caller_set_trace_field(
cast(Mapping[str, object], data), _metadata_variable_name, "session_id"
):
baggage: Final = normalized_headers.get("baggage")
if isinstance(baggage, str):

View file

@ -3704,13 +3704,18 @@ def test_add_litellm_metadata_from_request_headers_empty_body_session_id_falls_b
assert data["metadata"]["session_id"] == "baggage-session-42"
def test_add_litellm_metadata_from_request_headers_provider_metadata_does_not_block_baggage():
headers = {"baggage": "session.id=baggage-session-42"}
data = {"metadata": {"session_id": "provider-facing-value"}, "litellm_metadata": {}}
@pytest.mark.parametrize("field", ["trace_id", "session_id"])
def test_add_litellm_metadata_from_request_headers_promoted_metadata_beats_headers(field: str):
headers = {
"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01",
"baggage": "session.id=baggage-session-42",
}
data = {"metadata": {field: "caller-chosen"}, "litellm_metadata": {}}
LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers(
headers=headers, data=data, _metadata_variable_name="litellm_metadata"
)
assert data["litellm_session_id"] == "baggage-session-42"
assert field not in data["litellm_metadata"]
assert f"litellm_{field}" not in data
def _otel_span_with_trace_id(trace_id: int) -> NonRecordingSpan:
@ -8200,7 +8205,7 @@ async def test_missing_session_id_omit_keeps_client_supplied_session_id():
("path", "client_body"),
[
("/v1/chat/completions", {"model": "gpt-4o", "messages": [], "metadata": {"session_id": ""}}),
("/v1/responses", {"model": "gpt-4o", "input": "hi", "metadata": {"session_id": "provider-facing"}}),
("/v1/responses", {"model": "gpt-4o", "input": "hi"}),
],
)
async def test_missing_session_id_reject_accepts_baggage_session_id(path: str, client_body: dict[str, object]):
@ -8218,6 +8223,28 @@ async def test_missing_session_id_reject_accepts_baggage_session_id(path: str, c
assert updated["litellm_session_id"] == "baggage-session-42"
@pytest.mark.asyncio
@pytest.mark.parametrize("path", ["/v1/responses", "/v1/messages"])
@pytest.mark.parametrize("policy", [None, "reject", "generate"])
async def test_promoted_caller_trace_ids_beat_traceparent_and_baggage(path: str, policy: str | None):
request = _request_for(path)
request.headers = {
"traceparent": "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01",
"baggage": "session.id=baggage-session-42",
}
updated = await add_litellm_data_to_request(
data={"model": "gpt-4o", "metadata": {"trace_id": "caller-trace", "session_id": "caller-session"}},
request=request,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
proxy_config=MagicMock(),
general_settings={"missing_session_id": policy} if policy else {},
)
assert updated["litellm_metadata"]["trace_id"] == "caller-trace"
assert updated["litellm_metadata"]["session_id"] == "caller-session"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"client_body",