Merge pull request #40669 from BerriAI/litellm_otel_passthrough_trace_propagation

fix(otel): propagate W3C trace context on HTTP and WebSocket passthrough
This commit is contained in:
yucheng-berri 2026-09-16 16:02:46 -07:00 • committed by GitHub
commit 659c4ce5a1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 359 additions and 31 deletions

View file

@ -2965,16 +2965,13 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
)
propagator: Final = TraceContextTextMapPropagator()
carrier: Final = {"traceparent": _traceparent}
carrier: Final = {key: headers[key] for key in ("traceparent", "tracestate") if headers.get(key) is not None}
_parent_context: Final = propagator.extract(carrier=carrier)
return _parent_context
def _get_span_context(self, kwargs, default_span: Span | None = None):
from opentelemetry import context, trace
from opentelemetry.trace.propagation.tracecontext import (
TraceContextTextMapPropagator,
)
litellm_params: Final = kwargs.get("litellm_params", {}) or {}
proxy_server_request: Final = litellm_params.get("proxy_server_request", {}) or {}
@ -2998,11 +2995,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
# Priority 2: HTTP traceparent header
if traceparent is not None:
verbose_logger.debug("OpenTelemetry: Using traceparent header for context propagation")
carrier: Final = {"traceparent": traceparent}
return (
TraceContextTextMapPropagator().extract(carrier=carrier),
None,
)
return self.get_traceparent_from_header(headers=headers), None
# Priority 3: Active span from global context (auto-detection)
try:

View file

@ -26,6 +26,7 @@ if TYPE_CHECKING:
from litellm.integrations.otel.model.destination import OtelDestination
_PROPAGATOR: Final = TraceContextTextMapPropagator()
_W3C_TRACE_HEADERS: Final = frozenset(("traceparent", "tracestate"))
# The request's root span — the FastAPI-owned SERVER span — captured ONCE when the
# proxy first resolves it, so request-level spans (the LLM call, guardrails) can
@ -310,6 +311,37 @@ def extract_traceparent(headers: Mapping[str, str]) -> Context | None:
return _PROPAGATOR.extract(carrier)
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)
current: Final = get_current()
if is_recordable_span(get_current_span(current)):
return current
return None
def inject_trace_context(headers: Mapping[str, str], parent_span: object = None) -> dict[str, str]:
"""``headers`` plus W3C ``traceparent``/``tracestate`` for this request's span.
Parent preference: ``parent_span`` (the request span auth stashed on the key), then
the anchored request root span, then the ambient active span. Only trace context is
injected, never Baggage. Unchanged when no valid span exists anywhere.
"""
context: Final = _outgoing_trace_context(parent_span)
if context is None:
return dict(headers) # mutable-ok: OpenTelemetry propagator requires a mutable carrier
carrier: Final = { # mutable-ok: OpenTelemetry propagator requires a mutable carrier
key: value for key, value in headers.items() if key.lower() not in _W3C_TRACE_HEADERS
}
_PROPAGATOR.inject(carrier, context=context)
return carrier
# The OTLP destinations this request's key or team pointed its traces at, resolved
# once during auth. A ``ContextVar`` for the same reason the root span above is one:
# it rides the request task's context into the ``asyncio.create_task`` children that

View file

@ -986,6 +986,7 @@ async def pass_through_request(
headers=headers,
forward_headers=forward_headers,
)
upstream_headers: Final = _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)
@ -1019,7 +1020,7 @@ async def pass_through_request(
verbose_proxy_logger.debug(
"Pass through endpoint sending request to \nURL %s\nheaders: %s\nbody: %s\n",
url,
headers,
upstream_headers,
_parsed_body,
)
@ -1257,7 +1258,7 @@ async def pass_through_request(
additional_args={
"complete_input_dict": _parsed_body,
"api_base": str(logging_url),
"headers": headers,
"headers": upstream_headers,
},
)
stream = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body(
@ -1274,7 +1275,7 @@ async def pass_through_request(
request=request,
async_client=async_client,
url=url,
headers=headers,
headers=upstream_headers,
requested_query_params=requested_query_params,
stream=True,
)
@ -1286,7 +1287,7 @@ async def pass_through_request(
request.method,
url,
params=requested_query_params,
headers=headers,
headers=upstream_headers,
content=state_raw_body,
)
if state_raw_body is not None
@ -1294,7 +1295,7 @@ async def pass_through_request(
request.method,
url,
params=requested_query_params,
headers=headers,
headers=upstream_headers,
json=_parsed_body,
)
)
@ -1371,7 +1372,7 @@ async def pass_through_request(
raw_body_request: Final = async_client.build_request(
request.method,
url,
headers=headers,
headers=upstream_headers,
params=requested_query_params,
content=state_raw_body,
)
@ -1381,7 +1382,7 @@ async def pass_through_request(
request=request,
async_client=async_client,
url=url,
headers=headers,
headers=upstream_headers,
requested_query_params=requested_query_params,
_parsed_body=_parsed_body,
forward_multipart=is_multipart,
@ -2158,6 +2159,17 @@ def _upstream_close_to_relay(task_results: Iterable[object]) -> Close | None:
return upstream_close
_WEBSOCKET_FORWARDED_HEADERS: Final = frozenset(("authorization", "x-api-key", "x-goog-user-project"))
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, parent_span=parent_span)
async def websocket_passthrough_request(
websocket: WebSocket,
target: str,
@ -2200,20 +2212,15 @@ async def websocket_passthrough_request(
await websocket.accept()
verbose_proxy_logger.debug("WebSocket passthrough (%s): WebSocket connection accepted", endpoint)
# Prepare headers for the upstream connection
upstream_headers: Final = custom_headers.copy()
if forward_headers:
# Forward relevant headers from the incoming request
incoming_headers: Final = dict(websocket.headers)
for header_name, header_value in incoming_headers.items():
# Only forward certain headers to avoid conflicts
if header_name.lower() in [
"authorization",
"x-api-key",
"x-goog-user-project",
]:
upstream_headers[header_name] = header_value
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 websocket.headers.items()
if forward_headers and header_name.lower() in _WEBSOCKET_FORWARDED_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

@ -5,6 +5,7 @@ builders, and the registry validator's failure paths. Needs the OTel SDK."""
import json
import threading
from collections.abc import Iterator
from contextvars import Context as ContextVarContext
from dataclasses import replace
from http.server import BaseHTTPRequestHandler, HTTPServer, ThreadingHTTPServer
@ -15,6 +16,8 @@ pytest.importorskip("opentelemetry")
from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ( # noqa: E402
ExportTraceServiceRequest,
)
from opentelemetry import baggage # noqa: E402
from opentelemetry.context import attach, detach # noqa: E402
from opentelemetry.sdk.metrics import MeterProvider # noqa: E402
from opentelemetry.sdk.metrics.export import InMemoryMetricReader # noqa: E402
from opentelemetry.sdk.trace import TracerProvider # noqa: E402
@ -26,7 +29,10 @@ from opentelemetry.sdk.trace.export import ( # noqa: E402
from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( # noqa: E402
InMemorySpanExporter,
)
from opentelemetry.trace import SpanKind # noqa: E402
from opentelemetry.trace import SpanKind, get_current_span # noqa: E402
from opentelemetry.trace.propagation.tracecontext import ( # noqa: E402
TraceContextTextMapPropagator,
)
from litellm.integrations.otel.plumbing import context as ctx_mod # noqa: E402
from litellm.integrations.otel.plumbing import providers # noqa: E402
@ -464,6 +470,115 @@ def test_extract_traceparent():
assert ctx_mod.extract_traceparent({"x": "y"}) is None
def _test_tracer():
exporter = InMemorySpanExporter()
provider = TracerProvider()
provider.add_span_processor(SimpleSpanProcessor(exporter))
return provider.get_tracer("test")
def test_inject_trace_context_prefers_request_root_span():
def run():
tracer = _test_tracer()
with tracer.start_as_current_span("root") as root:
ctx_mod.set_request_root_span(root)
result = ctx_mod.inject_trace_context(
{"traceparent": "00-11111111111111111111111111111111-2222222222222222-01"}
)
propagated = get_current_span(TraceContextTextMapPropagator().extract(result))
return result, root, propagated
result, root, propagated = ContextVarContext().run(run)
assert result["traceparent"] != "00-11111111111111111111111111111111-2222222222222222-01"
assert propagated.get_span_context().trace_id == root.get_span_context().trace_id
assert propagated.get_span_context().span_id == root.get_span_context().span_id
def test_inject_trace_context_uses_ambient_span_without_request_root():
def run():
tracer = _test_tracer()
with tracer.start_as_current_span("ambient") as ambient:
result = ctx_mod.inject_trace_context({})
propagated = get_current_span(TraceContextTextMapPropagator().extract(result))
return ambient, propagated
ambient, propagated = ContextVarContext().run(run)
assert propagated.get_span_context().trace_id == ambient.get_span_context().trace_id
assert propagated.get_span_context().span_id == ambient.get_span_context().span_id
def test_inject_trace_context_replaces_stale_trace_headers():
def run():
tracer = _test_tracer()
with tracer.start_as_current_span("ambient") as ambient:
headers = {
"Traceparent": "00-" + "a" * 32 + "-" + "b" * 16 + "-01",
"Tracestate": "vendor=old",
"x-keep": "1",
}
result = ctx_mod.inject_trace_context(headers)
propagated = get_current_span(TraceContextTextMapPropagator().extract(result))
return result, ambient, propagated
result, ambient, propagated = ContextVarContext().run(run)
assert sum(key.lower() == "traceparent" for key in result) == 1
assert not any(key.lower() == "tracestate" for key in result)
assert result["x-keep"] == "1"
assert propagated.get_span_context().trace_id == ambient.get_span_context().trace_id
def test_inject_trace_context_prefers_explicit_parent_span_over_root_and_ambient():
def run():
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
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():
headers = {"x-custom": "value"}
result = ContextVarContext().run(lambda: ctx_mod.inject_trace_context(headers))
assert result == headers
assert "traceparent" not in result
assert result is not headers
def test_inject_trace_context_does_not_forward_baggage():
def run():
tracer = _test_tracer()
with tracer.start_as_current_span("ambient"):
token = attach(baggage.set_baggage("litellm.team.id", "team"))
try:
return ctx_mod.inject_trace_context({})
finally:
detach(token)
result = ContextVarContext().run(run)
assert "baggage" not in result
def test_set_request_baggage_empty_returns_context():
assert ctx_mod.set_request_baggage({}) is not None

View file

@ -5424,6 +5424,65 @@ class TestGetSpanContextLitellmMetadataFallback(unittest.TestCase):
self.assertIsNone(detected_span)
class TestInboundTraceContextKeepsCallerTracestate(unittest.TestCase):
"""The request span built from inbound W3C headers must carry the caller's
tracestate so outbound propagation (passthrough) re-emits it instead of
dropping it alongside the stripped stale header."""
CALLER_TRACEPARENT = "00-" + "a" * 32 + "-" + "b" * 16 + "-01"
CALLER_TRACESTATE = "vendor=abc,other=xyz"
def _otel(self):
provider = TracerProvider()
provider.add_span_processor(SimpleSpanProcessor(InMemorySpanExporter()))
otel = OpenTelemetry()
otel.tracer = provider.get_tracer(__name__)
return otel
def test_request_span_propagates_caller_tracestate_downstream(self):
from opentelemetry.trace.propagation.tracecontext import TraceContextTextMapPropagator
from litellm.integrations.otel.plumbing.context import inject_trace_context
inbound = {"traceparent": self.CALLER_TRACEPARENT, "tracestate": self.CALLER_TRACESTATE}
span = self._otel().create_litellm_proxy_request_started_span(
start_time=datetime.now(timezone.utc), headers=inbound
)
outbound = inject_trace_context(inbound, parent_span=span)
span.end()
propagated = trace.get_current_span(TraceContextTextMapPropagator().extract(outbound)).get_span_context()
self.assertEqual(outbound["tracestate"], self.CALLER_TRACESTATE)
self.assertEqual(propagated.trace_id, span.get_span_context().trace_id)
self.assertEqual(propagated.span_id, span.get_span_context().span_id)
self.assertNotEqual(outbound["traceparent"], self.CALLER_TRACEPARENT)
def test_request_span_without_caller_tracestate_emits_none(self):
from litellm.integrations.otel.plumbing.context import inject_trace_context
inbound = {"traceparent": self.CALLER_TRACEPARENT}
span = self._otel().create_litellm_proxy_request_started_span(
start_time=datetime.now(timezone.utc), headers=inbound
)
outbound = inject_trace_context(inbound, parent_span=span)
span.end()
self.assertNotIn("tracestate", outbound)
self.assertNotEqual(outbound["traceparent"], self.CALLER_TRACEPARENT)
def test_span_context_from_header_keeps_caller_tracestate(self):
kwargs = {
"litellm_params": {
"proxy_server_request": {
"headers": {"traceparent": self.CALLER_TRACEPARENT, "tracestate": self.CALLER_TRACESTATE}
}
}
}
ctx, detected_span = self._otel()._get_span_context(kwargs)
self.assertIsNone(detected_span)
self.assertEqual(trace.get_current_span(ctx).get_span_context().trace_state.to_header(), self.CALLER_TRACESTATE)
class TestEndProxySpanLitellmMetadataFallback(unittest.TestCase):
"""
Tests for _end_proxy_span_from_kwargs() falling back to litellm_metadata.

View file

@ -2,6 +2,7 @@ import asyncio
import json
import logging
import os
import sys
from collections.abc import Callable
from contextlib import ExitStack, contextmanager
from io import BytesIO
@ -29,6 +30,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
resolve_pass_through_request_timeout,
resolve_llm_passthrough_timeout,
websocket_passthrough_request,
_with_trace_context,
)
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -46,6 +48,15 @@ import litellm
MESSAGE_START_SSE_FRAME = b'event: message_start\ndata: {"type": "message_start"}\n\n'
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"}, parent_span=None)
assert headers == {"authorization": "x"}
assert "traceparent" not in headers
# Test is_multipart
def test_is_multipart():
# Test with multipart content type
@ -4270,6 +4281,47 @@ def _relay_client_request(method="GET"):
return mock_request
@pytest.mark.asyncio
@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
captured: dict[str, httpx.Headers] = {}
def transport_handler(upstream_request: httpx.Request) -> httpx.Response:
captured["headers"] = upstream_request.headers
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, {})
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()
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
async def test_pass_through_request_relays_non_json_body_without_buffering():
"""
@ -4866,6 +4918,76 @@ async def test_websocket_passthrough_forwards_non_ascii_first_frame():
assert all(call.kwargs.get("code") != 1011 for call in websocket.close.await_args_list)
@pytest.mark.asyncio
@pytest.mark.parametrize("forward_headers", [True, False])
@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
from starlette.websockets import WebSocketState
captured: dict[str, dict[str, str]] = {}
upstream_ws = FakeUpstreamWebSocket(b"{}")
def fake_connect(target, additional_headers):
captured["headers"] = additional_headers
return FakeUpstreamConnect(upstream_ws)
websocket = MagicMock()
websocket.accept = AsyncMock()
websocket.send_text = AsyncMock()
websocket.send_bytes = AsyncMock()
websocket.receive = AsyncMock(return_value={"type": "websocket.disconnect"})
websocket.close = AsyncMock()
websocket.headers = {"authorization": "Bearer client"}
websocket.client_state = WebSocketState.CONNECTED
websocket.application_state = WebSocketState.CONNECTED
tracer = TracerProvider().get_tracer("test")
mock_proxy_logging = MagicMock()
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_success_hook = AsyncMock()
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_worker = MagicMock()
mock_worker.ensure_initialized_and_enqueue = MagicMock(
side_effect=lambda async_coroutine: async_coroutine.close()
)
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging)
monkeypatch.setattr(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.connect",
fake_connect,
)
monkeypatch.setattr(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.GLOBAL_LOGGING_WORKER",
mock_worker,
)
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=user_api_key_dict,
forward_headers=forward_headers,
endpoint="/realtime",
accept_websocket=True,
)
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)
class ClosingUpstreamWebSocket:
def __init__(self, close_exc: Exception):
self._close_exc = close_exc