mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
e0fb89bc82
commit
661b1dbcf4
10 changed files with 1046 additions and 11 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
528
tests/integration/_support/otlp_sink.py
Normal file
528
tests/integration/_support/otlp_sink.py
Normal 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()
|
||||
|
|
@ -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"]),
|
||||
|
|
|
|||
53
tests/integration/observability/conftest.py
Normal file
53
tests/integration/observability/conftest.py
Normal 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,
|
||||
}
|
||||
268
tests/integration/observability/test_otel_excluded_services.py
Normal file
268
tests/integration/observability/test_otel_excluded_services.py
Normal 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}"
|
||||
|
|
@ -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
|
||||
# --------------------------------------------------------------------------- #
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue