mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
commit
659c4ce5a1
6 changed files with 359 additions and 31 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue