mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
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:
parent
7ba073aa26
commit
05ededf8a0
4 changed files with 74 additions and 46 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue