feat(otel): excluded_services opt-out for datastore spans on tenant destinations

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-26 01:05:39 +00:00
parent e0fb89bc82
commit 661b1dbcf4
10 changed files with 1046 additions and 11 deletions

View file

@ -205,6 +205,14 @@ nothing here imports outside it:
`config.yaml` — the latter reach the config through the logger's constructor
kwargs. `baggage_team_metadata_keys` is empty by default, so none of a team's
free-form metadata is promoted until each sub-key is explicitly allowlisted.
`excluded_services` withholds datastore spans from key/team `callback_vars`
destinations while the operator's own exporters keep them: set
`LITELLM_OTEL_EXCLUDED_SERVICES` (comma-separated) or `excluded_services`
(a YAML list) under `callback_settings.otel`, naming the datastore services
to withhold (`redis`, `postgres`, `batch_write_to_db`, `redis_*`, or their
`db.system.name` spellings `redis` / `postgresql`). A span is withheld when
its `db.system.name` / `db.system` attribute is in the set, so request root,
auth, guardrail and model spans can never be excluded.
- [`baggage.py`](./model/baggage.py) — the single definition of which request-identity
values are promoted into Baggage (so child spans inherit them) and under which
attribute keys.

View file

@ -1,5 +1,6 @@
"""Typed configuration for the OpenTelemetry instrumentation."""
from collections.abc import Mapping
from enum import Enum
from functools import lru_cache
from typing import Annotated, Any, Final
@ -12,6 +13,7 @@ from litellm.integrations.otel.model.baggage import (
DEFAULT_BAGGAGE_METADATA_KEYS,
DEFAULT_BAGGAGE_TEAM_METADATA_KEYS,
)
from litellm.integrations.otel.model.spans import POSTGRESQL, db_system
from litellm.types.utils import OtelSpanScope
#: Master feature-flag env var. The logger is inert until this is truthy.
@ -173,6 +175,19 @@ class OpenTelemetryV2Config(BaseSettings):
"key/team destinations are not affected."
),
)
excluded_services: Annotated[frozenset[str], NoDecode] = Field(
default_factory=frozenset,
validation_alias=AliasChoices("excluded_services", "LITELLM_OTEL_EXCLUDED_SERVICES"),
description=(
"Datastore services whose spans are withheld from key/team ``callback_vars`` "
"OTel destinations (the operator's own exporters still receive them). Accepted "
"values are the datastore ``ServiceTypes`` names (``redis``, ``postgres``, "
"``batch_write_to_db``, ``redis_*``) or their ``db.system.name`` spellings "
"(``redis``, ``postgresql``); stored normalized to ``db.system.name`` values. "
"Configure via the ``LITELLM_OTEL_EXCLUDED_SERVICES`` env var (comma-separated) "
"or ``callback_settings.otel.excluded_services`` in config.yaml (a YAML list)."
),
)
# ----- explicit multi-destination / vocabulary configuration ------------ #
@ -267,6 +282,7 @@ class OpenTelemetryV2Config(BaseSettings):
"baggage_metadata_keys",
"baggage_team_metadata_keys",
"mapper_names",
"excluded_services",
mode="before",
)
@classmethod
@ -315,6 +331,7 @@ class OpenTelemetryV2Config(BaseSettings):
if self.legacy_compat and "legacy" not in names:
names.append("legacy")
self.mapper_names = names
self.excluded_services = _normalize_excluded_services(self.excluded_services)
return self
@property
@ -333,3 +350,32 @@ class OpenTelemetryV2Config(BaseSettings):
@classmethod
def from_env(cls) -> "OpenTelemetryV2Config":
return cls()
def _normalize_excluded_services(services: frozenset[str]) -> frozenset[str]:
"""Fold each accepted spelling to its ``db.system.name`` value.
``postgres`` and ``postgresql`` name the same system, as do every
``ServiceTypes`` member that ``db_system`` maps. Anything else means the
operator pointed the setting at a span family it cannot cover.
"""
return frozenset(_db_system_for_excluded_service(service) for service in services)
def _db_system_for_excluded_service(service: str) -> str:
resolved: Final = db_system(service) if service != POSTGRESQL else POSTGRESQL
if resolved is None:
raise ValueError(f"excluded_services: {service!r} is not a datastore service; allowed: postgres, redis")
return resolved
def validate_otel_v2_callback_settings(settings: object) -> None:
"""Parse ``callback_settings.otel`` so a malformed block fails proxy boot.
Logger construction is lazy and swallows init errors, so without this a bad
value in the shared settings only surfaces as a dropped callback at request
time.
"""
if not is_otel_v2_enabled() or not isinstance(settings, Mapping):
return
OpenTelemetryV2Config(**settings)

View file

@ -418,6 +418,13 @@ def _is_database_span(attributes: Mapping[str, AttributeValue]) -> bool:
return any(key in attributes for key in _DB_SYSTEM_KEYS)
def _is_excluded_database_span(attributes: Mapping[str, AttributeValue], excluded: frozenset[str]) -> bool:
if not excluded:
return False
system: Final = attributes.get(DB.SYSTEM_NAME) or attributes.get(DB.SYSTEM_LEGACY)
return isinstance(system, str) and system in excluded
def _is_tenant_owned_span(attributes: Mapping[str, AttributeValue]) -> bool:
return any(key in attributes for key in _TENANT_OWNED_KEYS)
@ -549,10 +556,12 @@ class TenantFanOutSpanProcessor(SpanProcessor):
processor_factory: 'Callable[["OtelDestination"], SpanProcessor | None] | None' = None,
shutdown_drain_seconds: float = _SHUTDOWN_DRAIN_SECONDS,
operator_sinks: 'Mapping[_SinkKey, "OtelSpanScope"]' = MappingProxyType({}),
excluded_db_systems: frozenset[str] = frozenset(),
pending_drains: int = _MAX_PENDING_DRAINS,
drain_pool: _DrainPool | None = None,
) -> None:
self._operator_sinks: Final = operator_sinks
self._excluded_db_systems: Final = excluded_db_systems
self._drain_seconds: Final = shutdown_drain_seconds
self._lock: Final = threading.Condition()
self._closed = False # guarded by ``_lock``: an unlocked read races the teardown it gates
@ -567,9 +576,12 @@ class TenantFanOutSpanProcessor(SpanProcessor):
def on_end(self, span: ReadableSpan) -> None:
suppressed: Final = suppressed_backends()
attributes: Final = span.attributes or _NO_ATTRIBUTES
for destination in request_destinations():
if self._operator_already_writes(span, destination, suppressed) or not _in_scope(
span, destination.span_scope
if (
self._operator_already_writes(span, destination, suppressed)
or not _in_scope(span, destination.span_scope)
or _is_excluded_database_span(attributes, self._excluded_db_systems)
):
continue
processor = self._acquire(destination)
@ -1169,7 +1181,12 @@ def attach_tenant_fan_out(provider: TracerProvider, *configs: OpenTelemetryV2Con
with _FAN_OUT_ATTACH_LOCK:
if any(isinstance(processor, TenantFanOutSpanProcessor) for processor in _attached_processors(provider)):
return
provider.add_span_processor(TenantFanOutSpanProcessor(operator_sinks=operator_sink_scopes(*configs)))
provider.add_span_processor(
TenantFanOutSpanProcessor(
operator_sinks=operator_sink_scopes(*configs),
excluded_db_systems=frozenset().union(*(config.excluded_services for config in configs)),
)
)
def deliverable_destinations(

View file

@ -174,6 +174,11 @@ def initialize_callbacks_on_proxy(
imported_list.append(code_interpreter_interception_obj)
continue
if isinstance(callback, str) and callback == "otel":
from litellm.integrations.otel.model.config import validate_otel_v2_callback_settings
validate_otel_v2_callback_settings(callback_specific_params.get("otel"))
# check if callback is a custom logger compatible callback
if isinstance(callback, str):
callback = LoggingCallbackManager._add_custom_callback_generic_api_str(callback)

View file

@ -0,0 +1,528 @@
"""OTLP/HTTP trace sink: records exported spans and exposes them over a control API.
Accepts ``application/x-protobuf`` ``ExportTraceServiceRequest`` bodies and OTLP
``http/json`` bodies on any path. Tests read spans through ``recorded_spans`` and
steer the sink through ``configure``; the process can also be frozen with
``SIGSTOP``/``SIGCONT`` after reading its pid from ``/__pid``.
"""
from __future__ import annotations
import argparse
import datetime
import json
import os
import signal
import socket
import ssl
import subprocess
import sys
import threading
import time
from collections.abc import Iterator, Mapping, Sequence
from contextlib import ExitStack, contextmanager
from dataclasses import dataclass, field
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from typing import Final
from urllib.parse import urlparse
import httpx
import psutil
from pydantic import JsonValue, TypeAdapter
from typing_extensions import ReadOnly, TypedDict
INTERNAL_MARKERS: Final = ("gen_ai.operation.name", "mcp.method.name", "litellm.guardrail_name")
class Span(TypedDict):
trace_id: ReadOnly[str]
span_id: ReadOnly[str]
parent_span_id: ReadOnly[str]
kind: ReadOnly[int]
name: ReadOnly[str]
attributes: ReadOnly[Mapping[str, JsonValue]]
resource: ReadOnly[Mapping[str, JsonValue]]
class _SpanListing(TypedDict):
next: ReadOnly[int]
spans: ReadOnly[list[Span]]
_SPAN_LISTING: Final = TypeAdapter(_SpanListing)
def _proto_spans(body: bytes) -> list[Span]:
from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest
from opentelemetry.proto.common.v1.common_pb2 import AnyValue
def scalar(value: AnyValue) -> JsonValue:
match value.WhichOneof("value"):
case "string_value":
return value.string_value
case "bool_value":
return value.bool_value
case "int_value":
return int(value.int_value)
case "double_value":
return value.double_value
case "bytes_value":
return value.bytes_value.decode("utf-8", errors="replace")
case "array_value":
return [scalar(item) for item in value.array_value.values]
case "kvlist_value":
return {pair.key: scalar(pair.value) for pair in value.kvlist_value.values}
case _:
return None
request: Final = ExportTraceServiceRequest()
request.ParseFromString(body)
return [
Span(
trace_id=span.trace_id.hex(),
span_id=span.span_id.hex(),
parent_span_id=span.parent_span_id.hex(),
kind=span.kind,
name=span.name,
attributes={attribute.key: scalar(attribute.value) for attribute in span.attributes},
resource={attribute.key: scalar(attribute.value) for attribute in resource.resource.attributes},
)
for resource in request.resource_spans
for scope in resource.scope_spans
for span in scope.spans
]
def _json_spans(body: bytes) -> list[Span]:
payload: Final = json.loads(body)
def scalar(value: object) -> JsonValue:
if not isinstance(value, dict):
return value if isinstance(value, (str, int, float, bool)) or value is None else str(value)
for key in ("stringValue", "intValue", "doubleValue", "boolValue", "bytesValue"):
if key in value:
return value[key]
if "arrayValue" in value:
return [scalar(item) for item in value["arrayValue"].get("values", [])]
if "kvlistValue" in value:
return {pair["key"]: scalar(pair["value"]) for pair in value["kvlistValue"].get("values", [])}
return None
return [
Span(
trace_id=str(span.get("traceId", "")),
span_id=str(span.get("spanId", "")),
parent_span_id=str(span.get("parentSpanId", "")),
kind=int(span.get("kind", 0)),
name=str(span.get("name", "")),
attributes={attribute["key"]: scalar(attribute.get("value")) for attribute in span.get("attributes", [])},
resource={
attribute["key"]: scalar(attribute.get("value"))
for attribute in resource.get("resource", {}).get("attributes", [])
},
)
for resource in payload.get("resourceSpans", [])
for scope in resource.get("scopeSpans", [])
for span in scope.get("spans", [])
]
def decode_spans(body: bytes, content_type: str) -> list[Span]:
if "protobuf" in content_type:
return _proto_spans(body)
return _json_spans(body)
def span_class(span: Span) -> str:
if span["kind"] == 2:
return "root"
if any(marker in span["attributes"] for marker in INTERNAL_MARKERS):
return "tenant"
return "internal"
def spans_for_trace(spans: tuple[Span, ...], trace_id: str) -> tuple[Span, ...]:
return tuple(span for span in spans if span["trace_id"] == trace_id)
@dataclass(slots=True)
class _State:
spans: list[Span] = field(default_factory=list)
requests: list[dict[str, JsonValue]] = field(default_factory=list)
status: int = 200
delay_seconds: float = 0.0
pause: threading.Event = field(default_factory=threading.Event)
def __post_init__(self) -> None:
self.pause.set()
class _Handler(BaseHTTPRequestHandler):
state: _State
protocol_version = "HTTP/1.1"
def _read_body(self) -> bytes:
return self.rfile.read(int(self.headers.get("content-length", "0")))
def _send_json(self, payload: object, status: int = 200) -> None:
body: Final = json.dumps(payload).encode()
self.send_response(status)
self.send_header("content-type", "application/json")
self.send_header("content-length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def _record(self) -> None:
body: Final = self._read_body()
self.state.pause.wait(timeout=120)
if self.state.delay_seconds > 0:
time.sleep(self.state.delay_seconds)
recorded: Final = decode_spans(body, self.headers.get("content-type", ""))
self.state.spans.extend(recorded)
self.state.requests.append(
{
"path": self.path,
"count": len(recorded),
"host": self.headers.get("host", ""),
"headers": dict(self.headers),
}
)
self._send_json({"recorded": len(recorded)}, status=self.state.status)
do_POST = _record
do_PUT = _record
def do_GET(self) -> None:
parsed: Final = urlparse(self.path)
if parsed.path == "/__spans":
since: Final = int(dict(part.split("=", 1) for part in parsed.query.split("&") if part).get("since", "0"))
self._send_json({"next": len(self.state.spans), "spans": self.state.spans[since:]})
return
if parsed.path == "/__pid":
self._send_json({"pid": os.getpid()})
return
if parsed.path == "/__requests":
self._send_json({"requests": self.state.requests})
return
self._send_json({"error": "unknown"}, status=404)
def do_DELETE(self) -> None:
if urlparse(self.path).path == "/__spans":
self.state.spans.clear()
self.state.requests.clear()
self._send_json({"cleared": True})
return
self._send_json({"error": "unknown"}, status=404)
def do_PATCH(self) -> None:
if urlparse(self.path).path != "/__control":
self._send_json({"error": "unknown"}, status=404)
return
fields: Final = json.loads(self._read_body() or b"{}")
if "status" in fields:
self.state.status = int(fields["status"])
if "delay_seconds" in fields:
self.state.delay_seconds = float(fields["delay_seconds"])
if fields.get("paused") is True:
self.state.pause.clear()
if fields.get("paused") is False:
self.state.pause.set()
self._send_json({"status": self.state.status, "delay_seconds": self.state.delay_seconds})
def log_message(self, format: str, *args: object) -> None:
pass
class _ConnectHandler(_Handler):
tunnel_context: ssl.SSLContext
def do_CONNECT(self) -> None:
self.state.requests.append({"connect": self.path})
self.connection.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n")
wrapped: Final = self.tunnel_context.wrap_socket(self.connection, server_side=True)
self.close_connection = True
type(self)(wrapped, self.client_address, self.server)
_MITM_HOSTS: Final = ("otlp.nr-data.net", "otlp.eu01.nr-data.net")
def _mitm_context(directory: Path) -> ssl.SSLContext:
from cryptography import x509
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from cryptography.x509.oid import NameOID
directory.mkdir(parents=True, exist_ok=True)
now: Final = datetime.datetime.now(datetime.timezone.utc)
window: Final = datetime.timedelta(days=2)
ca_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
ca_name: Final = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, "otlp-sink test CA")])
ca_cert: Final = (
x509.CertificateBuilder()
.subject_name(ca_name)
.issuer_name(ca_name)
.public_key(ca_key.public_key())
.serial_number(x509.random_serial_number())
.not_valid_before(now - window)
.not_valid_after(now + window)
.add_extension(x509.BasicConstraints(ca=True, path_length=None), critical=True)
.sign(ca_key, hashes.SHA256())
)
leaf_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048)
leaf_cert: Final = (
x509.CertificateBuilder()
.subject_name(x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, _MITM_HOSTS[0])]))
.issuer_name(ca_cert.subject)
.public_key(leaf_key.public_key())
.serial_number(x509.random_serial_number())
.not_valid_before(now - window)
.not_valid_after(now + window)
.add_extension(x509.SubjectAlternativeName([x509.DNSName(host) for host in _MITM_HOSTS]), critical=False)
.sign(ca_key, hashes.SHA256())
)
ca_pem: Final = directory / "ca.pem"
ca_pem.write_bytes(ca_cert.public_bytes(serialization.Encoding.PEM))
leaf_pem: Final = directory / "leaf.pem"
leaf_pem.write_bytes(leaf_cert.public_bytes(serialization.Encoding.PEM))
leaf_key_pem: Final = directory / "leaf-key.pem"
leaf_key_pem.write_bytes(
leaf_key.private_bytes(
serialization.Encoding.PEM,
serialization.PrivateFormat.TraditionalOpenSSL,
serialization.NoEncryption(),
)
)
context: Final = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
context.load_cert_chain(str(leaf_pem), str(leaf_key_pem))
return context
def _grpc_trace_server(state: _State, port: int) -> object:
from concurrent import futures
import grpc
from opentelemetry.proto.collector.trace.v1 import trace_service_pb2, trace_service_pb2_grpc
class _TraceService(trace_service_pb2_grpc.TraceServiceServicer):
def Export(self, request: object, context: grpc.ServicerContext) -> object:
state.pause.wait(timeout=120)
if state.delay_seconds > 0:
time.sleep(state.delay_seconds)
recorded: Final = _proto_spans(request.SerializeToString())
state.spans.extend(recorded)
state.requests.append(
{
"grpc": "Export",
"metadata": {key: value for key, value in context.invocation_metadata()},
"count": len(recorded),
}
)
return trace_service_pb2.ExportTraceServiceResponse()
server: Final = grpc.server(futures.ThreadPoolExecutor(max_workers=4))
trace_service_pb2_grpc.add_TraceServiceServicer_to_server(_TraceService(), server)
server.add_insecure_port(f"127.0.0.1:{port}")
server.start()
return server
def recorded_spans(url: str, since: int = 0) -> tuple[int, tuple[Span, ...]]:
response: Final = httpx.get(f"{url}/__spans", params={"since": since}, trust_env=False, timeout=15)
response.raise_for_status()
listing: Final = _SPAN_LISTING.validate_python(response.json())
return listing["next"], tuple(listing["spans"])
def configure_sink(url: str, **fields: JsonValue) -> None:
httpx.request("PATCH", f"{url}/__control", json=dict(fields), trust_env=False, timeout=15).raise_for_status()
def reset_sink(url: str) -> None:
httpx.delete(f"{url}/__spans", trust_env=False, timeout=15).raise_for_status()
def sink_pid(url: str) -> int:
return int(httpx.get(f"{url}/__pid", trust_env=False, timeout=15).json()["pid"])
_REQUEST_LISTING: Final = TypeAdapter(list[dict[str, JsonValue]])
def recorded_requests(url: str) -> tuple[Mapping[str, JsonValue], ...]:
response: Final = httpx.get(f"{url}/__requests", trust_env=False, timeout=15)
response.raise_for_status()
return tuple(_REQUEST_LISTING.validate_python(response.json()["requests"]))
@dataclass(frozen=True, slots=True)
class SpanSinks:
operator: str
tenant: str
arize: str
@dataclass(frozen=True, slots=True)
class GrpcSink:
url: str
control_url: str
@dataclass(frozen=True, slots=True)
class ConnectSink:
proxy_url: str
control_url: str
ca_pem: str
def _free_port() -> int:
with socket.socket() as reserve:
reserve.bind(("127.0.0.1", 0))
return int(reserve.getsockname()[1])
def _pid_reachable(url: str) -> bool:
try:
return httpx.get(f"{url}/__pid", trust_env=False, timeout=2).status_code == 200
except httpx.TransportError:
return False
@contextmanager
def owned_sinks(directory: Path) -> Iterator[SpanSinks]:
from integration._support.process import group_members, signal_group, stop_root_process
directory.mkdir(parents=True, exist_ok=True)
ports: Final = tuple(_free_port() for _ in range(3))
root: Final = Path(__file__).resolve().parents[3]
with ExitStack() as stack:
processes: Final = tuple(
subprocess.Popen(
[sys.executable, "-P", "-m", "integration._support.otlp_sink", "--port", str(port)],
cwd=root,
stdout=stack.enter_context((directory / f"otlp-sink-{port}.log").open("w")),
stderr=subprocess.STDOUT,
start_new_session=True,
)
for port in ports
)
try:
urls: Final = tuple(f"http://127.0.0.1:{port}" for port in ports)
deadline: Final = time.monotonic() + 30
while True:
alive: Final = all(process.poll() is None for process in processes)
assert alive, "OTLP sink exited before readiness"
if all(_pid_reachable(url) for url in urls):
break
assert time.monotonic() < deadline, "OTLP sink readiness deadline exceeded"
time.sleep(0.05)
yield SpanSinks(operator=urls[0], tenant=urls[1], arize=urls[2])
finally:
for process in processes:
stopped: Final = stop_root_process(process)
residual: Final = group_members(process.pid)
if residual:
signal_group(process.pid, signal.SIGKILL)
psutil.wait_procs(residual, timeout=5)
survivors: Final = group_members(process.pid)
assert not survivors and stopped, "OTLP sink required forced cleanup"
@contextmanager
def _spawn_sink(directory: Path, log_name: str, argv: Sequence[str]) -> Iterator[None]:
from integration._support.process import group_members, signal_group, stop_root_process
directory.mkdir(parents=True, exist_ok=True)
root: Final = Path(__file__).resolve().parents[3]
with (directory / log_name).open("w") as log:
process: Final = subprocess.Popen(
[sys.executable, "-P", "-m", "integration._support.otlp_sink", *argv],
cwd=root,
stdout=log,
stderr=subprocess.STDOUT,
start_new_session=True,
)
try:
yield
finally:
stopped: Final = stop_root_process(process)
residual: Final = group_members(process.pid)
if residual:
signal_group(process.pid, signal.SIGKILL)
psutil.wait_procs(residual, timeout=5)
survivors: Final = group_members(process.pid)
assert not survivors and stopped, "OTLP sink required forced cleanup"
def _await_sink(url: str) -> None:
deadline: Final = time.monotonic() + 30
while not _pid_reachable(url):
assert time.monotonic() < deadline, "OTLP sink readiness deadline exceeded"
time.sleep(0.05)
@contextmanager
def owned_grpc_sink(directory: Path) -> Iterator[GrpcSink]:
http_port: Final = _free_port()
grpc_port: Final = _free_port()
with _spawn_sink(
directory, "otlp-grpc-sink.log", ["--port", str(http_port), "--grpc-port", str(grpc_port)]
):
control_url: Final = f"http://127.0.0.1:{http_port}"
_await_sink(control_url)
yield GrpcSink(url=f"http://127.0.0.1:{grpc_port}", control_url=control_url)
@contextmanager
def owned_connect_sink(directory: Path) -> Iterator[ConnectSink]:
http_port: Final = _free_port()
tunnel_port: Final = _free_port()
ca_dir: Final = directory / "mitm"
with _spawn_sink(
directory,
"otlp-connect-sink.log",
["--port", str(http_port), "--connect-port", str(tunnel_port), "--ca-dir", str(ca_dir)],
):
control_url: Final = f"http://127.0.0.1:{http_port}"
_await_sink(control_url)
yield ConnectSink(
proxy_url=f"http://127.0.0.1:{tunnel_port}",
control_url=control_url,
ca_pem=str(ca_dir / "ca.pem"),
)
def main() -> None:
parser: Final = argparse.ArgumentParser()
parser.add_argument("--port", type=int, required=True)
parser.add_argument("--grpc-port", type=int, default=0)
parser.add_argument("--connect-port", type=int, default=0)
parser.add_argument("--ca-dir", type=Path, default=None)
arguments: Final = parser.parse_args()
bound_state: Final = _State()
class BoundHandler(_Handler):
state = bound_state
if arguments.grpc_port:
grpc_server: Final = _grpc_trace_server(bound_state, arguments.grpc_port)
assert grpc_server is not None
if arguments.connect_port:
assert arguments.ca_dir is not None, "--connect-port needs --ca-dir"
bound_context: Final = _mitm_context(arguments.ca_dir)
class BoundConnectHandler(_ConnectHandler):
state = bound_state
tunnel_context = bound_context # pyright: ignore[reportIncompatibleVariableOverride] # bound context, not a new field
tunnel: Final = ThreadingHTTPServer(("127.0.0.1", arguments.connect_port), BoundConnectHandler)
tunnel.daemon_threads = True
threading.Thread(target=tunnel.serve_forever, daemon=True).start()
server: Final = ThreadingHTTPServer(("127.0.0.1", arguments.port), BoundHandler)
server.daemon_threads = True
server.serve_forever()
if __name__ == "__main__":
main()

View file

@ -193,6 +193,46 @@ class Provider:
self.scripts[name] = deque(int(str(value)) for value in statuses)
return JSONResponse({"configured": len(statuses)})
async def responses(self, request: Request) -> Response:
body: Final = JSON_OBJECT.validate_json(await request.body())
self.observations.put(Observation(request.url.path, request.headers.get("authorization", ""), body))
response_id: Final = f"resp-{uuid.uuid4().hex[:24]}"
model: Final = body.get("model") if isinstance(body.get("model"), str) else "audit-chat"
def envelope(with_usage: bool) -> dict[str, JsonValue]:
return {
"id": response_id,
"object": "response",
"created_at": 1,
"status": "completed",
"model": model,
"output": [
{
"type": "message",
"id": f"msg-{response_id}",
"status": "completed",
"role": "assistant",
"content": [
{"type": "output_text", "text": "integration upstream response", "annotations": []}
],
}
],
"usage": (
{"input_tokens": 20, "output_tokens": 20, "total_tokens": 40} if with_usage else None
),
}
if body.get("stream") is True:
events: Final = (
{"type": "response.created", "response": envelope(False)},
{"type": "response.output_text.delta", "delta": "integration upstream "},
{"type": "response.output_text.delta", "delta": "response"},
{"type": "response.completed", "response": envelope(True)},
)
stream_body: Final = "".join(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n" for event in events)
return Response(content=stream_body.encode(), media_type="text/event-stream")
return JSONResponse(envelope(True))
async def observed(self, _request: Request) -> Response:
values: Final = tuple(self.observations.get() for _ in range(self.observations.qsize()))
return JSONResponse(
@ -239,14 +279,6 @@ class Provider:
response: Final = self.scenario_store.get(scenario_id)
if response is None:
return JSONResponse({"error": "Unknown scenario"}, status_code=404)
if request.method == "POST" and "json" in request.headers.get("content-type", ""):
raw_body: Final = await request.body()
if raw_body:
body: Final = JSON_OBJECT.validate_json(raw_body)
if isinstance(body, dict):
self.observations.put(
Observation(request.url.path, request.headers.get("authorization", ""), body)
)
if isinstance(response, RoutedResponse):
route_key: Final = f"{request.method} /{'/'.join(segments[1:])}"
route: Final = next(
@ -371,6 +403,7 @@ class Provider:
Route("/v1/completions", completions, methods=["POST"]),
Route("/v1/embeddings", embeddings, methods=["POST"]),
Route("/v1/moderations", moderations, methods=["POST"]),
Route("/v1/responses", self.responses, methods=["POST"]),
Route("/vector_stores/{vector_store_id}/search", self.vector_store_search, methods=["POST"]),
Route("/{path:path}", self.scripted, methods=["POST"]),
Route("/{path:path}", self.scripted, methods=["GET"]),

View file

@ -0,0 +1,53 @@
from __future__ import annotations
import uuid
from collections.abc import Callable, Iterator, Mapping
from pathlib import Path
from typing import Final
from urllib.parse import urlparse
import pytest
import yaml
from integration._support.otlp_sink import SpanSinks, owned_sinks
from pydantic import JsonValue
AuditConfigWriter = Callable[[Path, Mapping[str, JsonValue]], Path]
@pytest.fixture(scope="module")
def audit_sinks(tmp_path_factory: pytest.TempPathFactory) -> Iterator[SpanSinks]:
directory: Final = tmp_path_factory.mktemp("otel-audit-sinks")
with owned_sinks(directory) as sinks:
yield sinks
@pytest.fixture(scope="module")
def otel_audit_config(audit_sinks: SpanSinks) -> AuditConfigWriter:
tenant_host: Final = urlparse(audit_sinks.tenant).netloc
def write(directory: Path, litellm_settings: Mapping[str, JsonValue] = {}) -> Path:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["litellm_settings"] = {
**config.get("litellm_settings", {}),
"callbacks": ["otel"],
"provider_url_destination_allowed_hosts": [tenant_host],
**dict(litellm_settings),
}
config["callback_settings"] = {
"otel": {"exporter": "http/json", "endpoint": audit_sinks.operator, "use_simple_processor": True}
}
config["general_settings"] = {**config.get("general_settings", {}), "user_api_key_cache_ttl": 2}
path: Final = directory / f"otel-audit-{uuid.uuid4().hex}.yaml"
path.write_text(yaml.safe_dump(config))
return path
return write
@pytest.fixture(scope="module")
def langfuse_vars(audit_sinks: SpanSinks) -> dict[str, JsonValue]:
return {
"langfuse_public_key": "pk-lf-audit",
"langfuse_secret_key": "sk-lf-audit",
"langfuse_host": audit_sinks.tenant,
}

View file

@ -0,0 +1,268 @@
from __future__ import annotations
import time
import uuid
from collections.abc import Callable, Iterator, Mapping
from pathlib import Path
from typing import Final
import httpx
import pytest
import yaml
from integration._support.client import (
Gateway,
Scenario,
eventually,
gateway_from_environment,
)
from integration._support.otlp_sink import (
Span,
SpanSinks,
recorded_spans,
spans_for_trace,
)
from integration._support.process import owned_proxy, owned_proxy_process
from pydantic import JsonValue
AuditConfigWriter = Callable[[Path, Mapping[str, JsonValue]], Path]
DB_SYSTEM_KEYS: Final = frozenset({"db.system.name", "db.system"})
@pytest.fixture(scope="module")
def gateway(audit_sinks: SpanSinks) -> Iterator[Gateway]:
with gateway_from_environment() as base:
yield base
def _config_with(
directory: Path,
otel_audit_config: AuditConfigWriter,
*,
otel: Mapping[str, JsonValue] = {},
extra: Callable[[dict], None] | None = None,
) -> Path:
config: Final = yaml.safe_load(otel_audit_config(directory, {}).read_text())
config["callback_settings"]["otel"].update(dict(otel))
if extra is not None:
extra(config)
path: Final = directory / f"otel-excl-{uuid.uuid4().hex}.yaml"
path.write_text(yaml.safe_dump(config))
return path
def _add_callback(gateway: Gateway, team_id: str, callback_vars: Mapping[str, JsonValue]) -> httpx.Response:
return gateway.request(
"POST",
f"/team/{team_id}/callback",
{"callback_name": "langfuse_otel", "callback_vars": dict(callback_vars)},
)
def _drive(candidate: Gateway, langfuse_vars: Mapping[str, JsonValue]) -> httpx.Response:
with candidate.scenario() as scenario:
model: Final = scenario.model(model="openai/audit-chat", api_base=f"{candidate.upstream_url}/v1")
team_id: Final = scenario.team()
callback: Final = _add_callback(candidate, team_id, langfuse_vars)
assert callback.status_code == 200, callback.text
key: Final = scenario.key(team_id=team_id)
response: Final = candidate.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": f"otel-excl-{uuid.uuid4().hex}"}]},
key=key,
)
assert response.status_code == 200, response.text
return response
def _trace_id(sink_url: str, response: httpx.Response, seconds: float = 40) -> str:
call_id: Final = response.headers.get("x-litellm-call-id")
response_id: Final = response.json().get("id")
def look() -> str | None:
_, spans = recorded_spans(sink_url)
return next(
(
str(span["trace_id"])
for span in spans
if (call_id is not None and span["attributes"].get("litellm.call_id") == call_id)
or (response_id is not None and span["attributes"].get("gen_ai.response.id") == response_id)
),
None,
)
found: Final = eventually(look, lambda value: value is not None, seconds=seconds)
assert found is not None
return found
def _trace_spans(sink_url: str, trace_id: str, seconds: float = 30) -> tuple[Span, ...]:
"""The trace's spans once the post-call tail has landed.
The spend-writer and other post-response spans flush after the request
answers, so absence assertions poll for the whole window instead of
settling at the first glimpse of the root span.
"""
deadline: Final = time.monotonic() + seconds
group: tuple[Span, ...] = () # rebind-ok: drains samples until the post-call tail lands
while time.monotonic() < deadline:
_, spans = recorded_spans(sink_url)
group = spans_for_trace(spans, trace_id)
time.sleep(0.5)
assert group, f"trace {trace_id} never reached {sink_url}"
return group
def _await_db_span(sink_url: str, trace_id: str | None, needle: str, seconds: float = 40, since: int = 0) -> None:
def seen() -> bool:
_, spans = recorded_spans(sink_url, since)
group: Final = spans if trace_id is None else spans_for_trace(spans, trace_id)
return any(
needle in str(span["name"]) or needle in {str(span["attributes"].get(k)) for k in DB_SYSTEM_KEYS}
for span in group
)
landed: Final = eventually(seen, bool, seconds=seconds)
assert landed, f"{needle} span never landed at {sink_url}"
def _db_systems(spans: tuple[Span, ...]) -> set[str]:
return {
str(span["attributes"][key])
for span in spans
for key in DB_SYSTEM_KEYS
if key in span["attributes"]
}
def _assert_core_spans_present(spans: tuple[Span, ...]) -> None:
attributes_by_span: Final = tuple(span["attributes"] for span in spans)
assert any(span["kind"] == 2 for span in spans), "request root span missing"
assert any("gen_ai.operation.name" in attrs for attrs in attributes_by_span), "model span missing"
assert any("litellm.guardrail.name" in attrs for attrs in attributes_by_span), "guardrail span missing"
names: Final = sorted(str(span["name"]) for span in spans)
assert any(name.startswith("auth") for name in names), f"auth span missing in {names}"
def _guardrail_block(config: dict) -> None:
config["guardrails"] = [
{
"guardrail_name": f"excl-filter-{uuid.uuid4().hex[:8]}",
"litellm_params": {
"guardrail": "litellm_content_filter",
"mode": "pre_call",
"default_on": True,
"patterns": [
{
"pattern_type": "regex",
"pattern_name": "excl_secret",
"pattern": "TOPSECRET\\d{9}",
"action": "BLOCK",
}
],
},
}
]
def test_excluded_services_drops_db_spans_at_tenant_only(
gateway: Gateway,
audit_sinks: SpanSinks,
otel_audit_config: AuditConfigWriter,
langfuse_vars: dict[str, JsonValue],
tmp_path: Path,
) -> None:
config: Final = _config_with(
tmp_path, otel_audit_config, otel={"excluded_services": ["redis", "postgres"]}, extra=_guardrail_block
)
with owned_proxy(gateway, tmp_path, {"LITELLM_OTEL_V2": "1"}, config=config, workers=2) as candidate:
traffic: Final = _drive(candidate, langfuse_vars)
tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic)
tenant_spans: Final = _trace_spans(audit_sinks.tenant, tenant_trace)
_assert_core_spans_present(tenant_spans)
assert _db_systems(tenant_spans) == set(), (
f"db spans reached tenant: {sorted(str(s['name']) for s in tenant_spans)}"
)
operator_trace: Final = _trace_id(audit_sinks.operator, traffic)
operator_systems: Final = _db_systems(_trace_spans(audit_sinks.operator, operator_trace))
assert operator_trace == tenant_trace
assert {"redis", "postgresql"} <= operator_systems, f"operator lost db spans: {operator_systems}"
def test_env_excluded_services_drops_only_redis(
gateway: Gateway,
audit_sinks: SpanSinks,
otel_audit_config: AuditConfigWriter,
langfuse_vars: dict[str, JsonValue],
tmp_path: Path,
) -> None:
config: Final = _config_with(tmp_path, otel_audit_config)
with owned_proxy(
gateway, tmp_path, {"LITELLM_OTEL_V2": "1", "LITELLM_OTEL_EXCLUDED_SERVICES": "redis"}, config=config, workers=2
) as candidate:
start, _ = recorded_spans(audit_sinks.tenant)
traffic: Final = _drive(candidate, langfuse_vars)
_await_db_span(audit_sinks.tenant, None, "postgresql", seconds=60, since=start)
_, tenant_spans = recorded_spans(audit_sinks.tenant, start)
systems: Final = _db_systems(tenant_spans)
assert "postgresql" in systems, f"postgresql spans missing at tenant: {systems}"
assert "redis" not in systems, f"redis spans reached tenant: {sorted(str(s['name']) for s in tenant_spans)}"
def test_config_excluded_services_wins_over_env(
gateway: Gateway,
audit_sinks: SpanSinks,
otel_audit_config: AuditConfigWriter,
langfuse_vars: dict[str, JsonValue],
tmp_path: Path,
) -> None:
config: Final = _config_with(tmp_path, otel_audit_config, otel={"excluded_services": ["postgres"]})
with owned_proxy(
gateway, tmp_path, {"LITELLM_OTEL_V2": "1", "LITELLM_OTEL_EXCLUDED_SERVICES": "redis"}, config=config, workers=2
) as candidate:
start, _ = recorded_spans(audit_sinks.tenant)
traffic: Final = _drive(candidate, langfuse_vars)
tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic)
_await_db_span(audit_sinks.tenant, tenant_trace, "redis", seconds=60)
tenant_spans: Final = _trace_spans(audit_sinks.tenant, tenant_trace, seconds=15)
_, all_tenant = recorded_spans(audit_sinks.tenant, start)
systems: Final = _db_systems(tenant_spans)
assert "redis" in systems, f"redis spans missing at tenant: {systems}"
assert "postgresql" not in _db_systems(all_tenant), f"postgresql spans reached tenant: {_db_systems(all_tenant)}"
def test_bogus_excluded_service_fails_proxy_start(
gateway: Gateway, otel_audit_config: AuditConfigWriter, tmp_path: Path
) -> None:
config: Final = _config_with(tmp_path, otel_audit_config, otel={"excluded_services": ["auth"]})
with pytest.raises(AssertionError, match="readiness"):
with owned_proxy_process(gateway, tmp_path, {"LITELLM_OTEL_V2": "1"}, config=config, workers=2):
pass
logs: Final = [path.read_text() for path in tmp_path.glob("owned-proxy-*.log")]
assert logs, "no owned proxy log written"
text: Final = "\n".join(logs)
assert "'auth' is not a datastore service" in text, text[-3000:]
assert "postgres, redis" in text, text[-3000:]
def test_postgres_exclusion_covers_batch_write_to_db(
gateway: Gateway,
audit_sinks: SpanSinks,
otel_audit_config: AuditConfigWriter,
langfuse_vars: dict[str, JsonValue],
tmp_path: Path,
) -> None:
config: Final = _config_with(tmp_path, otel_audit_config, otel={"excluded_services": ["postgres"]})
with owned_proxy(gateway, tmp_path, {"LITELLM_OTEL_V2": "1"}, config=config, workers=2) as candidate:
op_start, _ = recorded_spans(audit_sinks.operator)
ten_start, _ = recorded_spans(audit_sinks.tenant)
traffic: Final = _drive(candidate, langfuse_vars)
_await_db_span(audit_sinks.operator, None, "batch_write_to_db", seconds=60, since=op_start)
tenant_trace: Final = _trace_id(audit_sinks.tenant, traffic)
_await_db_span(audit_sinks.tenant, tenant_trace, "redis", seconds=60)
tenant_spans: Final = _trace_spans(audit_sinks.tenant, tenant_trace, seconds=15)
_, all_tenant = recorded_spans(audit_sinks.tenant, ten_start)
names: Final = sorted(str(span["name"]) for span in all_tenant)
assert "redis" in _db_systems(tenant_spans), f"redis spans missing at tenant: {names}"
assert not any("batch_write_to_db" in name for name in names), f"spend writer reached tenant: {names}"

View file

@ -90,6 +90,36 @@ def test_baggage_processor_allowlist_uses_config_keys():
assert LiteLLM.TEAM_ALIAS not in span.attributes # not in this allowlist
@pytest.mark.parametrize(
"given,expected",
[
(["redis"], frozenset({"redis"})),
(["postgres"], frozenset({"postgresql"})),
(["postgresql"], frozenset({"postgresql"})),
(["batch_write_to_db"], frozenset({"postgresql"})),
(["redis_spend_update_queue"], frozenset({"redis"})),
(["redis", "postgres"], frozenset({"redis", "postgresql"})),
],
)
def test_excluded_services_normalize_to_db_system_names(given, expected):
assert OpenTelemetryV2Config(excluded_services=given).excluded_services == expected
def test_excluded_services_from_env_csv(monkeypatch):
monkeypatch.setenv("LITELLM_OTEL_EXCLUDED_SERVICES", "redis, postgres")
assert OpenTelemetryV2Config().excluded_services == frozenset({"redis", "postgresql"})
def test_excluded_services_config_wins_over_env(monkeypatch):
monkeypatch.setenv("LITELLM_OTEL_EXCLUDED_SERVICES", "redis")
assert OpenTelemetryV2Config(excluded_services=["postgres"]).excluded_services == frozenset({"postgresql"})
def test_excluded_services_rejects_a_non_datastore_service():
with pytest.raises(Exception, match="'auth' is not a datastore service; allowed: postgres, redis"):
OpenTelemetryV2Config(excluded_services=["auth"])
# --------------------------------------------------------------------------- #
# Area 2 — pass-through LLM span parents to the ambient server span
# --------------------------------------------------------------------------- #

View file

@ -514,6 +514,53 @@ class TestFanOut:
for child in ("auth /v1/chat/completions", "chat gpt-4"):
assert by_name[child].parent.span_id == root.context.span_id
def test_excluded_services_drop_only_the_datastore_spans_at_the_tenant(self):
"""The exclusion is per ``db.system.*`` value: a span naming an excluded
datastore never reaches the tenant, while every span of the request's
own work (root, auth, guardrail, model) still does, and the operator's
own exporter keeps the full tree."""
dest_exporter, operator_exporter = InMemorySpanExporter(), InMemorySpanExporter()
provider = TracerProvider()
provider.add_span_processor(SimpleSpanProcessor(operator_exporter))
provider.add_span_processor(
TenantFanOutSpanProcessor(
processor_factory=lambda _d: SimpleSpanProcessor(dest_exporter),
excluded_db_systems=frozenset({"redis", "postgresql"}),
)
)
tracer = get_tracer(provider, "litellm")
def run():
set_request_destinations((LANGFUSE_DEST,))
with tracer.start_as_current_span("POST /v1/chat/completions"):
with tracer.start_as_current_span("auth /v1/chat/completions"):
pass
with tracer.start_as_current_span("execute_guardrail pii"):
pass
with tracer.start_as_current_span("redis async_get_cache") as redis_span:
redis_span.set_attribute("db.system.name", "redis")
with tracer.start_as_current_span("batch_write_to_db _PROXY_track_cost_callback") as spend_span:
spend_span.set_attribute("db.system", "postgresql")
with tracer.start_as_current_span("chat gpt-4"):
pass
in_fresh_context(run)
assert {s.name for s in dest_exporter.get_finished_spans()} == {
"POST /v1/chat/completions",
"auth /v1/chat/completions",
"execute_guardrail pii",
"chat gpt-4",
}
assert {s.name for s in operator_exporter.get_finished_spans()} == {
"POST /v1/chat/completions",
"auth /v1/chat/completions",
"execute_guardrail pii",
"redis async_get_cache",
"batch_write_to_db _PROXY_track_cost_callback",
"chat gpt-4",
}
def test_a_team_naming_two_backends_gets_the_trace_at_both(self):
"""The fan-out rides one provider, so it cannot skip a destination on the
grounds that some other backend owns it: nothing else would deliver it."""