From 661b1dbcf45ea1c3f0d2e035937ec09006ef709c Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 26 Sep 2026 01:05:39 +0000 Subject: [PATCH] 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> --- litellm/integrations/otel/README.md | 8 + litellm/integrations/otel/model/config.py | 46 ++ .../integrations/otel/plumbing/providers.py | 23 +- litellm/proxy/common_utils/callback_utils.py | 5 + tests/integration/_support/otlp_sink.py | 528 ++++++++++++++++++ tests/integration/_support/upstream.py | 49 +- tests/integration/observability/conftest.py | 53 ++ .../test_otel_excluded_services.py | 268 +++++++++ ..._v2_config_baggage_parenting_guardrails.py | 30 + .../otel/test_otel_v2_destinations.py | 47 ++ 10 files changed, 1046 insertions(+), 11 deletions(-) create mode 100644 tests/integration/_support/otlp_sink.py create mode 100644 tests/integration/observability/conftest.py create mode 100644 tests/integration/observability/test_otel_excluded_services.py diff --git a/litellm/integrations/otel/README.md b/litellm/integrations/otel/README.md index d8dfabe23d6..33fa38f8848 100644 --- a/litellm/integrations/otel/README.md +++ b/litellm/integrations/otel/README.md @@ -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. diff --git a/litellm/integrations/otel/model/config.py b/litellm/integrations/otel/model/config.py index 5447a8ee80a..ee169f968b5 100644 --- a/litellm/integrations/otel/model/config.py +++ b/litellm/integrations/otel/model/config.py @@ -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) diff --git a/litellm/integrations/otel/plumbing/providers.py b/litellm/integrations/otel/plumbing/providers.py index 8bac36aad76..050ebd51f37 100644 --- a/litellm/integrations/otel/plumbing/providers.py +++ b/litellm/integrations/otel/plumbing/providers.py @@ -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( diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index cb8b51d092e..cbe5dca2edc 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -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) diff --git a/tests/integration/_support/otlp_sink.py b/tests/integration/_support/otlp_sink.py new file mode 100644 index 00000000000..eeabe9d886f --- /dev/null +++ b/tests/integration/_support/otlp_sink.py @@ -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() diff --git a/tests/integration/_support/upstream.py b/tests/integration/_support/upstream.py index e645126032b..27a18915ff3 100644 --- a/tests/integration/_support/upstream.py +++ b/tests/integration/_support/upstream.py @@ -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"]), diff --git a/tests/integration/observability/conftest.py b/tests/integration/observability/conftest.py new file mode 100644 index 00000000000..f09a703f047 --- /dev/null +++ b/tests/integration/observability/conftest.py @@ -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, + } diff --git a/tests/integration/observability/test_otel_excluded_services.py b/tests/integration/observability/test_otel_excluded_services.py new file mode 100644 index 00000000000..ee9455fcac8 --- /dev/null +++ b/tests/integration/observability/test_otel_excluded_services.py @@ -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}" diff --git a/tests/unit/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py b/tests/unit/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py index dcaff3c911a..2184db533a5 100644 --- a/tests/unit/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py +++ b/tests/unit/integrations/otel/test_otel_v2_config_baggage_parenting_guardrails.py @@ -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 # --------------------------------------------------------------------------- # diff --git a/tests/unit/integrations/otel/test_otel_v2_destinations.py b/tests/unit/integrations/otel/test_otel_v2_destinations.py index 9cb3dbb9deb..ebc4747a502 100644 --- a/tests/unit/integrations/otel/test_otel_v2_destinations.py +++ b/tests/unit/integrations/otel/test_otel_v2_destinations.py @@ -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."""