fix(otel): parent passthrough trace propagation on the legacy request span

Pass user_api_key_dict.parent_otel_span into the outgoing W3C injection so the
legacy otel callback propagates its litellm_request span, falling back to the
otel_v2 request root span and then the ambient span. Extend the mapped unit
tests to assert the propagated trace and span ids over real captured headers
for HTTP and WebSocket passthrough with forwarding on and off.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-16 07:24:52 +00:00
parent 7ba073aa26
commit 05ededf8a0
4 changed files with 74 additions and 46 deletions

View file

@ -310,7 +310,10 @@ def extract_traceparent(headers: Mapping[str, str]) -> Context | None:
return _PROPAGATOR.extract(carrier)
def _outgoing_trace_context(inbound_headers: Mapping[str, str] | None = None) -> Context | None:
def _outgoing_trace_context(parent_span: object) -> Context | None:
if isinstance(parent_span, Span) and is_recordable_span(parent_span):
return context_from_span(parent_span)
root: Final = request_root_span()
if root is not None:
return context_from_span(root)
@ -318,27 +321,19 @@ def _outgoing_trace_context(inbound_headers: Mapping[str, str] | None = None) ->
current: Final = get_current()
if is_recordable_span(get_current_span(current)):
return current
if inbound_headers is None:
return None
inbound_context: Final = extract_traceparent(inbound_headers)
if inbound_context is None or not is_recordable_span(get_current_span(inbound_context)):
return None
return inbound_context
return None
def inject_trace_context(
headers: Mapping[str, str],
inbound_headers: Mapping[str, str] | None = None,
) -> dict[str, str]:
def inject_trace_context(headers: Mapping[str, str], parent_span: object = None) -> dict[str, str]:
"""``headers`` plus W3C ``traceparent``/``tracestate`` for the current request's span.
Parent preference: the anchored request root span, then the ambient active span,
then the trace context the caller sent inbound. Only trace context is injected,
never Baggage, so per-request identity baggage cannot leak upstream. Unchanged
when no valid span context exists anywhere.
Parent preference: the request span auth stashed on the key (the legacy
``litellm_request`` SERVER span, or the FastAPI server span under otel_v2), then
the anchored request root span, then the ambient active span. Only trace context
is injected, never Baggage, so per-request identity baggage cannot leak upstream.
Unchanged when no valid span context exists anywhere.
"""
context: Final = _outgoing_trace_context(inbound_headers)
context: Final = _outgoing_trace_context(parent_span)
if context is None:
return dict(headers) # mutable-ok: OpenTelemetry propagator requires a mutable carrier
carrier: Final = dict(headers) # mutable-ok: OpenTelemetry propagator requires a mutable carrier

View file

@ -985,7 +985,7 @@ async def pass_through_request(
headers=headers,
forward_headers=forward_headers,
)
headers = _with_trace_context(headers, inbound_headers=_safe_get_request_headers(request))
headers = _with_trace_context(headers, parent_span=user_api_key_dict.parent_otel_span)
requested_query_params: dict | None = query_params or dict(request.query_params)
@ -2161,12 +2161,12 @@ def _upstream_close_to_relay(task_results: Iterable[object]) -> Close | None:
_WEBSOCKET_FORWARDED_HEADERS: Final = frozenset(("authorization", "x-api-key", "x-goog-user-project"))
def _with_trace_context(headers: Mapping[str, str], inbound_headers: Mapping[str, str]) -> dict[str, str]:
def _with_trace_context(headers: Mapping[str, str], parent_span: object) -> dict[str, str]:
try:
from litellm.integrations.otel.plumbing.context import inject_trace_context
except ImportError:
return dict(headers) # mutable-ok: matches inject_trace_context's carrier return type
return inject_trace_context(headers, inbound_headers=inbound_headers)
return inject_trace_context(headers, parent_span=parent_span)
async def websocket_passthrough_request(
@ -2211,16 +2211,15 @@ async def websocket_passthrough_request(
await websocket.accept()
verbose_proxy_logger.debug("WebSocket passthrough (%s): WebSocket connection accepted", endpoint)
incoming_headers: Final = dict(websocket.headers) # mutable-ok: propagator carrier
forwarded_headers: Final = { # mutable-ok: propagator carrier
forwarded_headers: Final = { # mutable-ok: one-shot upstream header dict, read as a Mapping
**custom_headers,
**{
header_name: header_value
for header_name, header_value in incoming_headers.items()
for header_name, header_value in websocket.headers.items()
if forward_headers and header_name.lower() in _WEBSOCKET_FORWARDED_HEADERS
},
}
upstream_headers: Final = _with_trace_context(forwarded_headers, inbound_headers=incoming_headers)
upstream_headers: Final = _with_trace_context(forwarded_headers, parent_span=user_api_key_dict.parent_otel_span)
# Initialize logging object similar to HTTP passthrough
team_callbacks: Final = _resolve_team_callback_wiring(

View file

@ -507,16 +507,32 @@ def test_inject_trace_context_uses_ambient_span_without_request_root():
assert propagated.get_span_context().span_id == ambient.get_span_context().span_id
def test_inject_trace_context_forwards_valid_inbound_context_without_span():
inbound = {"traceparent": "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01"}
def test_inject_trace_context_prefers_explicit_parent_span_over_root_and_ambient():
def run():
result = ctx_mod.inject_trace_context({}, inbound_headers=inbound)
return get_current_span(TraceContextTextMapPropagator().extract(result))
tracer = _test_tracer()
parent = tracer.start_span("litellm_request")
with tracer.start_as_current_span("ambient") as ambient:
ctx_mod.set_request_root_span(ambient)
result = ctx_mod.inject_trace_context({}, parent_span=parent)
propagated = get_current_span(TraceContextTextMapPropagator().extract(result))
return parent, ambient, propagated
propagated = ContextVarContext().run(run)
assert propagated.get_span_context().trace_id == int("0af7651916cd43dd8448eb211c80319c", 16)
assert propagated.get_span_context().span_id == int("b7ad6b7169203331", 16)
parent, ambient, propagated = ContextVarContext().run(run)
assert propagated.get_span_context().trace_id == parent.get_span_context().trace_id
assert propagated.get_span_context().span_id == parent.get_span_context().span_id
assert propagated.get_span_context().span_id != ambient.get_span_context().span_id
def test_inject_trace_context_skips_unusable_parent_span():
def run():
tracer = _test_tracer()
with tracer.start_as_current_span("ambient") as ambient:
result = ctx_mod.inject_trace_context({}, parent_span=object())
propagated = get_current_span(TraceContextTextMapPropagator().extract(result))
return ambient, propagated
ambient, propagated = ContextVarContext().run(run)
assert propagated.get_span_context().span_id == ambient.get_span_context().span_id
def test_inject_trace_context_returns_headers_unchanged_without_context():

View file

@ -51,7 +51,7 @@ MESSAGE_START_SSE_FRAME = b'event: message_start\ndata: {"type": "message_start"
def test_with_trace_context_without_opentelemetry(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setitem(sys.modules, "litellm.integrations.otel.plumbing.context", None)
headers = _with_trace_context({"authorization": "x"}, {})
headers = _with_trace_context({"authorization": "x"}, parent_span=None)
assert headers == {"authorization": "x"}
assert "traceparent" not in headers
@ -4282,11 +4282,11 @@ def _relay_client_request(method="GET"):
@pytest.mark.asyncio
async def test_pass_through_request_propagates_active_trace_context():
@pytest.mark.parametrize("span_source", ["auth_parent_span", "ambient_span"])
async def test_pass_through_request_propagates_active_trace_context(span_source: str):
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.trace import get_current_span
from opentelemetry.trace.propagation.tracecontext import TraceContextTextMapPropagator
from litellm.proxy._types import UserAPIKeyAuth
captured: dict[str, httpx.Headers] = {}
@ -4295,17 +4295,23 @@ async def test_pass_through_request_propagates_active_trace_context():
return httpx.Response(200, json={"ok": True}, request=upstream_request)
fake_client, cleanup = _inject_fake_passthrough_client(httpx.MockTransport(transport_handler), timeout=None)
tracer = TracerProvider().get_tracer("test")
try:
with ExitStack() as stack:
_enter_relay_logging_mocks(stack, {})
tracer = TracerProvider().get_tracer("test")
with tracer.start_as_current_span("passthrough") as span:
response = await pass_through_request(
request=_relay_client_request(method="POST"),
target="http://internal-api.test/v1/generate",
custom_headers={},
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
)
if span_source == "auth_parent_span":
span = tracer.start_span("litellm_request")
stack.callback(span.end)
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", parent_otel_span=span)
else:
span = stack.enter_context(tracer.start_as_current_span("passthrough"))
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test")
response = await pass_through_request(
request=_relay_client_request(method="POST"),
target="http://internal-api.test/v1/generate",
custom_headers={},
user_api_key_dict=user_api_key_dict,
)
finally:
cleanup()
await fake_client.aclose()
@ -4313,6 +4319,7 @@ async def test_pass_through_request_propagates_active_trace_context():
assert response.status_code == 200
propagated = get_current_span(TraceContextTextMapPropagator().extract(captured["headers"]))
assert propagated.get_span_context().trace_id == span.get_span_context().trace_id
assert propagated.get_span_context().span_id == span.get_span_context().span_id
@pytest.mark.asyncio
@ -4913,7 +4920,10 @@ async def test_websocket_passthrough_forwards_non_ascii_first_frame():
@pytest.mark.asyncio
@pytest.mark.parametrize("forward_headers", [True, False])
async def test_websocket_passthrough_propagates_active_trace_context(monkeypatch, forward_headers: bool):
@pytest.mark.parametrize("span_source", ["auth_parent_span", "ambient_span"])
async def test_websocket_passthrough_propagates_active_trace_context(
monkeypatch, forward_headers: bool, span_source: str
):
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.trace import get_current_span
from opentelemetry.trace.propagation.tracecontext import TraceContextTextMapPropagator
@ -4954,12 +4964,19 @@ async def test_websocket_passthrough_propagates_active_trace_context(monkeypatch
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.GLOBAL_LOGGING_WORKER",
mock_worker,
)
with tracer.start_as_current_span("websocket_passthrough") as span:
with ExitStack() as stack:
if span_source == "auth_parent_span":
span = tracer.start_span("litellm_request")
stack.callback(span.end)
user_api_key_dict = UserAPIKeyAuth(parent_otel_span=span)
else:
span = stack.enter_context(tracer.start_as_current_span("websocket_passthrough"))
user_api_key_dict = UserAPIKeyAuth()
await websocket_passthrough_request(
websocket=websocket,
target="wss://upstream.example.test/v1/realtime",
custom_headers={},
user_api_key_dict=UserAPIKeyAuth(),
user_api_key_dict=user_api_key_dict,
forward_headers=forward_headers,
endpoint="/realtime",
accept_websocket=True,
@ -4967,6 +4984,7 @@ async def test_websocket_passthrough_propagates_active_trace_context(monkeypatch
propagated = get_current_span(TraceContextTextMapPropagator().extract(captured["headers"]))
assert propagated.get_span_context().trace_id == span.get_span_context().trace_id
assert propagated.get_span_context().span_id == span.get_span_context().span_id
assert captured["headers"].get("authorization") == ("Bearer client" if forward_headers else None)