mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
feat(arize): per-team success and error sampling rates for the Arize AX callback (#42383)
* feat(arize): per-team success and error sampling rates for the Arize AX callback Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(arize): fail open on invalid sampling rates and type the new sampling code Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(arize): type the sampling test helpers and parametrized fixture Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin <yassin@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
226fa7845a
commit
b673ee61e7
11 changed files with 487 additions and 78 deletions
|
|
@ -4,13 +4,17 @@ arize AI is OTEL compatible
|
|||
this file has Arize ai specific helper functions
|
||||
"""
|
||||
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from collections.abc import Callable, Mapping
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.arize import _utils
|
||||
from litellm.integrations.arize._utils import ArizeOTELAttributes
|
||||
from litellm.integrations.opentelemetry import OpenTelemetry
|
||||
from litellm.integrations.opentelemetry import _MAX_DYNAMIC_TRACER_PROVIDERS, OpenTelemetry, OpenTelemetryConfig
|
||||
from litellm.types.integrations.arize import ArizeConfig
|
||||
from litellm.types.services import ServiceLoggerPayload
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
|
|
@ -26,6 +30,9 @@ else:
|
|||
Protocol = Any
|
||||
Span = Any
|
||||
|
||||
_SUCCESS_SAMPLING_RATE_VAR: Final = "arize_success_sampling_rate"
|
||||
_ERROR_SAMPLING_RATE_VAR: Final = "arize_error_sampling_rate"
|
||||
|
||||
|
||||
class ArizeLogger(OpenTelemetry):
|
||||
"""
|
||||
|
|
@ -36,6 +43,26 @@ class ArizeLogger(OpenTelemetry):
|
|||
fighting over the global ``opentelemetry.trace`` TracerProvider singleton.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: OpenTelemetryConfig | None = None,
|
||||
callback_name: str | None = None,
|
||||
tracer_provider: object | None = None,
|
||||
logger_provider: object | None = None,
|
||||
meter_provider: object | None = None,
|
||||
max_dynamic_tracer_providers: int = _MAX_DYNAMIC_TRACER_PROVIDERS,
|
||||
random_draw: Callable[[], float] | None = None,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
config=config,
|
||||
callback_name=callback_name,
|
||||
tracer_provider=tracer_provider,
|
||||
logger_provider=logger_provider,
|
||||
meter_provider=meter_provider,
|
||||
max_dynamic_tracer_providers=max_dynamic_tracer_providers,
|
||||
)
|
||||
self._random_draw: Final[Callable[[], float]] = random_draw if random_draw is not None else random.random
|
||||
|
||||
def _init_tracing(self, tracer_provider):
|
||||
"""
|
||||
Override to always create a *private* TracerProvider for Arize.
|
||||
|
|
@ -55,6 +82,72 @@ class ArizeLogger(OpenTelemetry):
|
|||
self.tracer = provider.get_tracer("litellm")
|
||||
self.span_kind = SpanKind
|
||||
|
||||
def _handle_success(
|
||||
self,
|
||||
kwargs: dict[str, object],
|
||||
response_obj: object,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
) -> None:
|
||||
if not self._should_export(kwargs, _SUCCESS_SAMPLING_RATE_VAR):
|
||||
return
|
||||
super()._handle_success(kwargs, response_obj, start_time, end_time)
|
||||
|
||||
def _handle_failure(
|
||||
self,
|
||||
kwargs: dict[str, object],
|
||||
response_obj: object,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
) -> None:
|
||||
if not self._should_export(kwargs, _ERROR_SAMPLING_RATE_VAR):
|
||||
return
|
||||
super()._handle_failure(kwargs, response_obj, start_time, end_time)
|
||||
|
||||
def _sampling_rate_for_request(self, kwargs: Mapping[str, object], var: str) -> float | None:
|
||||
dynamic_params: Final = kwargs.get("standard_callback_dynamic_params")
|
||||
if not isinstance(dynamic_params, Mapping):
|
||||
return None
|
||||
value: Final = dynamic_params.get(var)
|
||||
if value is None or value in ("", "None"):
|
||||
return None
|
||||
try:
|
||||
rate = float(value)
|
||||
except (TypeError, ValueError):
|
||||
verbose_logger.warning(
|
||||
"ArizeLogger: %s value %r is not a number; exporting the request",
|
||||
var,
|
||||
value,
|
||||
)
|
||||
return None
|
||||
if not math.isfinite(rate) or not 0.0 <= rate <= 1.0:
|
||||
verbose_logger.warning(
|
||||
"ArizeLogger: %s value %r is outside 0.0..1.0; exporting the request",
|
||||
var,
|
||||
value,
|
||||
)
|
||||
return None
|
||||
return rate
|
||||
|
||||
def _should_export(self, kwargs: dict[str, object], var: str) -> bool:
|
||||
rate: Final = self._sampling_rate_for_request(kwargs, var)
|
||||
if rate is None:
|
||||
return True
|
||||
otel_internal: Final = self._otel_internal_state(kwargs)
|
||||
key: Final = f"arize_sampled:{var}"
|
||||
cached: Final = otel_internal.get(key)
|
||||
if isinstance(cached, bool):
|
||||
return cached
|
||||
sampled: Final = rate > 0.0 and self._random_draw() <= rate
|
||||
otel_internal[key] = sampled
|
||||
if not sampled:
|
||||
verbose_logger.debug(
|
||||
"ArizeLogger: dropping request, %s rate %r rejected the draw",
|
||||
var,
|
||||
rate,
|
||||
)
|
||||
return sampled
|
||||
|
||||
def _init_otel_logger_on_litellm_proxy(self):
|
||||
"""
|
||||
Override: Arize should NOT overwrite the proxy's
|
||||
|
|
|
|||
|
|
@ -1240,6 +1240,25 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
# End of Team/Key Based Logging Control Flow
|
||||
#########################################################
|
||||
|
||||
def _otel_internal_state(self, kwargs: dict[str, object]) -> dict[str, object]:
|
||||
"""Return the request-local ``_otel_internal`` marker dict, creating it if absent."""
|
||||
litellm_params = kwargs.get("litellm_params")
|
||||
if not isinstance(litellm_params, dict):
|
||||
litellm_params = {}
|
||||
kwargs["litellm_params"] = litellm_params
|
||||
|
||||
_metadata = litellm_params.get("metadata")
|
||||
if not isinstance(_metadata, dict):
|
||||
_metadata = {}
|
||||
litellm_params["metadata"] = _metadata
|
||||
|
||||
_otel_internal = _metadata.get("_otel_internal")
|
||||
if not isinstance(_otel_internal, dict):
|
||||
_otel_internal = {}
|
||||
_metadata["_otel_internal"] = _otel_internal
|
||||
|
||||
return _otel_internal
|
||||
|
||||
def _emit_once(self, kwargs: dict, *scope: object) -> bool:
|
||||
"""Return True the first time this handler is asked to emit a span
|
||||
for the given (handler, scope) on this kwargs; False on repeats.
|
||||
|
|
@ -1264,20 +1283,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
request-local (kwargs is shared across the sync/async callbacks and
|
||||
lifecycle hooks for one request).
|
||||
"""
|
||||
litellm_params = kwargs.get("litellm_params")
|
||||
if not isinstance(litellm_params, dict):
|
||||
litellm_params = {}
|
||||
kwargs["litellm_params"] = litellm_params
|
||||
|
||||
_metadata = litellm_params.get("metadata")
|
||||
if not isinstance(_metadata, dict):
|
||||
_metadata = {}
|
||||
litellm_params["metadata"] = _metadata
|
||||
|
||||
_otel_internal = _metadata.get("_otel_internal")
|
||||
if not isinstance(_otel_internal, dict):
|
||||
_otel_internal = {}
|
||||
_metadata["_otel_internal"] = _otel_internal
|
||||
_otel_internal = self._otel_internal_state(kwargs)
|
||||
|
||||
spans_logged = _otel_internal.get("spans_logged")
|
||||
if not isinstance(spans_logged, dict):
|
||||
|
|
|
|||
|
|
@ -98,6 +98,8 @@ _supported_callback_params: Final[tuple[str, ...]] = (
|
|||
"arize_api_key",
|
||||
"arize_space_key",
|
||||
"arize_space_id",
|
||||
"arize_success_sampling_rate",
|
||||
"arize_error_sampling_rate",
|
||||
"posthog_api_key",
|
||||
"posthog_host",
|
||||
"braintrust_api_key",
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ Team callbacks arrive as a single ``AddTeamCallback``, key callbacks arrive as a
|
|||
per-integration checks here.
|
||||
"""
|
||||
|
||||
import math
|
||||
from collections.abc import Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
|
@ -13,11 +14,18 @@ _NEWRELIC_CALLBACK: Final = "newrelic"
|
|||
_NEWRELIC_VAR_PREFIX: Final = "newrelic_"
|
||||
_LANGFUSE_OTEL_CALLBACK: Final = "langfuse_otel"
|
||||
_LANGFUSE_SPAN_SCOPE_VAR: Final = "langfuse_span_scope"
|
||||
_ARIZE_CALLBACK: Final = "arize"
|
||||
_ARIZE_SAMPLING_RATE_VARS: Final[frozenset[str]] = frozenset(
|
||||
{"arize_success_sampling_rate", "arize_error_sampling_rate"}
|
||||
)
|
||||
|
||||
|
||||
def callback_config_error(callback_name: str | None, callback_vars: Mapping[str, str] | None) -> str | None:
|
||||
if not callback_vars:
|
||||
return None
|
||||
arize_error: Final = _arize_sampling_rate_error(callback_name, callback_vars)
|
||||
if arize_error is not None:
|
||||
return arize_error
|
||||
langfuse_error: Final = _langfuse_environment_error(callback_vars) or _langfuse_span_scope_error(
|
||||
callback_name, callback_vars
|
||||
)
|
||||
|
|
@ -86,7 +94,7 @@ _VAR_FAMILIES: Final[Mapping[str, str]] = MappingProxyType(
|
|||
}
|
||||
)
|
||||
|
||||
_FAMILY_OPTION_VARS: Final[frozenset[str]] = frozenset({_LANGFUSE_SPAN_SCOPE_VAR})
|
||||
_FAMILY_OPTION_VARS: Final[frozenset[str]] = frozenset({_LANGFUSE_SPAN_SCOPE_VAR, *_ARIZE_SAMPLING_RATE_VARS})
|
||||
|
||||
|
||||
def _family_of(var: str) -> str | None:
|
||||
|
|
@ -212,6 +220,22 @@ def _logging_entry_error(entry: object) -> str | None:
|
|||
return callback_config_error(callback_name, _entry_callback_vars(entry))
|
||||
|
||||
|
||||
def _arize_sampling_rate_error(callback_name: str | None, callback_vars: Mapping[str, str]) -> str | None:
|
||||
for var in sorted(_ARIZE_SAMPLING_RATE_VARS):
|
||||
value = callback_vars.get(var)
|
||||
if value is None or value in ("", "None"):
|
||||
continue
|
||||
if callback_name != _ARIZE_CALLBACK:
|
||||
return f"{var} applies to the {_ARIZE_CALLBACK} callback only, not {callback_name!r}"
|
||||
try:
|
||||
rate = float(value)
|
||||
except (TypeError, ValueError):
|
||||
return f"{var} must be a number between 0.0 and 1.0 (inclusive), got {value!r}"
|
||||
if not math.isfinite(rate) or not 0.0 <= rate <= 1.0:
|
||||
return f"{var} must be a number between 0.0 and 1.0 (inclusive), got {value!r}"
|
||||
return None
|
||||
|
||||
|
||||
def _newrelic_config_error(callback_vars: Mapping[str, str]) -> str | None:
|
||||
"""Per-team New Relic routing runs on the OTel v2 path only.
|
||||
|
||||
|
|
|
|||
|
|
@ -3638,6 +3638,8 @@ class StandardCallbackDynamicParams(TypedDict, total=False):
|
|||
arize_api_key: str | None
|
||||
arize_space_key: str | None
|
||||
arize_space_id: str | None
|
||||
arize_success_sampling_rate: ReadOnly[float | None]
|
||||
arize_error_sampling_rate: ReadOnly[float | None]
|
||||
|
||||
# PostHog dynamic params
|
||||
posthog_api_key: str | None
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@ from unittest.mock import MagicMock, Mock, patch
|
|||
# Adds the grandparent directory to sys.path to allow importing project modules
|
||||
|
||||
import asyncio
|
||||
import datetime
|
||||
from collections.abc import Callable
|
||||
|
||||
import pytest
|
||||
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
|
||||
|
|
@ -64,9 +66,7 @@ async def test_arize_dynamic_params():
|
|||
print(f"Tracer calls: {len(tracer_calls)}")
|
||||
|
||||
# We should have captured calls for both requests
|
||||
assert (
|
||||
len(tracer_calls) >= 2
|
||||
), f"Expected at least 2 tracer calls, got {len(tracer_calls)}"
|
||||
assert len(tracer_calls) >= 2, f"Expected at least 2 tracer calls, got {len(tracer_calls)}"
|
||||
|
||||
# Check that we have the expected dynamic params in the kwargs
|
||||
team1_found = False
|
||||
|
|
@ -87,9 +87,7 @@ async def test_arize_dynamic_params():
|
|||
assert team1_found, "team1 dynamic params not found"
|
||||
assert team2_found, "team2 dynamic params not found"
|
||||
|
||||
print(
|
||||
"✅ All assertions passed - OpenTelemetry logger correctly received dynamic params"
|
||||
)
|
||||
print("✅ All assertions passed - OpenTelemetry logger correctly received dynamic params")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -114,11 +112,8 @@ async def test_arize_dynamic_headers_in_grpc_requests():
|
|||
"opentelemetry.exporter.otlp.proto.http.trace_exporter.OTLPSpanExporter",
|
||||
mock_otlp_http_exporter,
|
||||
):
|
||||
|
||||
# Create ArizeLogger with HTTP configuration
|
||||
config = OpenTelemetryConfig(
|
||||
exporter="otlp_http", endpoint="https://otlp.arize.com/v1"
|
||||
)
|
||||
config = OpenTelemetryConfig(exporter="otlp_http", endpoint="https://otlp.arize.com/v1")
|
||||
arize_logger = ArizeLogger(config=config)
|
||||
litellm.callbacks = [arize_logger]
|
||||
|
||||
|
|
@ -147,25 +142,17 @@ async def test_arize_dynamic_headers_in_grpc_requests():
|
|||
print(f"Captured exporter headers: {exporter_headers}")
|
||||
|
||||
# Should have multiple exporter calls (default + dynamic)
|
||||
assert (
|
||||
len(exporter_headers) >= 2
|
||||
), f"Expected at least 2 exporter calls, got {len(exporter_headers)}"
|
||||
assert len(exporter_headers) >= 2, f"Expected at least 2 exporter calls, got {len(exporter_headers)}"
|
||||
|
||||
# Find team1 and team2 headers
|
||||
team1_found = False
|
||||
team2_found = False
|
||||
|
||||
for headers in exporter_headers:
|
||||
if (
|
||||
headers.get("api_key") == "team1_api_key"
|
||||
and headers.get("arize-space-id") == "team1_space_id"
|
||||
):
|
||||
if headers.get("api_key") == "team1_api_key" and headers.get("arize-space-id") == "team1_space_id":
|
||||
team1_found = True
|
||||
print(f"✅ Found team1 headers: {headers}")
|
||||
elif (
|
||||
headers.get("api_key") == "team2_api_key"
|
||||
and headers.get("arize-space-id") == "team2_space_id"
|
||||
):
|
||||
elif headers.get("api_key") == "team2_api_key" and headers.get("arize-space-id") == "team2_space_id":
|
||||
team2_found = True
|
||||
print(f"✅ Found team2 headers: {headers}")
|
||||
|
||||
|
|
@ -173,6 +160,119 @@ async def test_arize_dynamic_headers_in_grpc_requests():
|
|||
assert team1_found, "team1 dynamic headers not found in exporter calls"
|
||||
assert team2_found, "team2 dynamic headers not found in exporter calls"
|
||||
|
||||
print(
|
||||
"✅ Test passed - Dynamic Arize params correctly passed to gRPC/HTTP exporter"
|
||||
)
|
||||
print("✅ Test passed - Dynamic Arize params correctly passed to gRPC/HTTP exporter")
|
||||
|
||||
|
||||
_START = datetime.datetime.now()
|
||||
_END = datetime.datetime.now()
|
||||
|
||||
|
||||
def _sampled_arize_logger(
|
||||
random_draw: Callable[[], float] | None = None,
|
||||
) -> tuple[ArizeLogger, InMemorySpanExporter]:
|
||||
from opentelemetry.sdk.trace import TracerProvider
|
||||
from opentelemetry.sdk.trace.export import SimpleSpanProcessor
|
||||
|
||||
exporter = InMemorySpanExporter()
|
||||
provider = TracerProvider()
|
||||
provider.add_span_processor(SimpleSpanProcessor(exporter))
|
||||
logger = ArizeLogger(tracer_provider=provider, random_draw=random_draw)
|
||||
return logger, exporter
|
||||
|
||||
|
||||
def _request_spans(exporter: InMemorySpanExporter) -> int:
|
||||
return sum(1 for span in exporter.get_finished_spans() if span.name == "litellm_request")
|
||||
|
||||
|
||||
def _arize_kwargs(callback_vars: dict[str, str] | None = None) -> dict[str, object]:
|
||||
kwargs: dict[str, object] = {
|
||||
"model": "gpt-4",
|
||||
"litellm_params": {"metadata": {}},
|
||||
"standard_logging_object": {
|
||||
"id": "call-1",
|
||||
"call_type": "completion",
|
||||
"model": "gpt-4",
|
||||
"metadata": {},
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
},
|
||||
}
|
||||
if callback_vars is not None:
|
||||
kwargs["standard_callback_dynamic_params"] = callback_vars
|
||||
return kwargs
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_success_sampling_rate_zero_exports_no_spans():
|
||||
logger, exporter = _sampled_arize_logger(random_draw=lambda: 0.0)
|
||||
await logger.async_log_success_event(_arize_kwargs({"arize_success_sampling_rate": "0.0"}), None, _START, _END)
|
||||
assert _request_spans(exporter) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_success_sampling_rate_one_exports_a_span():
|
||||
logger, exporter = _sampled_arize_logger()
|
||||
await logger.async_log_success_event(_arize_kwargs({"arize_success_sampling_rate": "1.0"}), None, _START, _END)
|
||||
assert _request_spans(exporter) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unset_sampling_rate_exports_everything():
|
||||
logger, exporter = _sampled_arize_logger()
|
||||
await logger.async_log_success_event(_arize_kwargs(), None, _START, _END)
|
||||
assert _request_spans(exporter) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_draw_above_rate_is_dropped_draw_at_rate_is_exported():
|
||||
dropped, dropped_exporter = _sampled_arize_logger(random_draw=lambda: 0.3)
|
||||
await dropped.async_log_success_event(_arize_kwargs({"arize_success_sampling_rate": "0.2"}), None, _START, _END)
|
||||
assert _request_spans(dropped_exporter) == 0
|
||||
|
||||
kept, kept_exporter = _sampled_arize_logger(random_draw=lambda: 0.2)
|
||||
await kept.async_log_success_event(_arize_kwargs({"arize_success_sampling_rate": "0.2"}), None, _START, _END)
|
||||
assert _request_spans(kept_exporter) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_success_and_error_rates_are_independent():
|
||||
logger, exporter = _sampled_arize_logger()
|
||||
kwargs = _arize_kwargs({"arize_success_sampling_rate": "0.0", "arize_error_sampling_rate": "1.0"})
|
||||
await logger.async_log_success_event(kwargs, None, _START, _END)
|
||||
await logger.async_log_failure_event(kwargs, ValueError("boom"), _START, _END)
|
||||
assert _request_spans(exporter) == 1
|
||||
|
||||
logger2, exporter2 = _sampled_arize_logger()
|
||||
kwargs2 = _arize_kwargs({"arize_success_sampling_rate": "1.0", "arize_error_sampling_rate": "0.0"})
|
||||
await logger2.async_log_success_event(kwargs2, None, _START, _END)
|
||||
await logger2.async_log_failure_event(kwargs2, ValueError("boom"), _START, _END)
|
||||
assert _request_spans(exporter2) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_one_draw_per_request_across_sync_and_async_handlers():
|
||||
draws: list[int] = []
|
||||
|
||||
def counting_draw() -> float:
|
||||
draws.append(1)
|
||||
return 0.5
|
||||
|
||||
logger, _ = _sampled_arize_logger(random_draw=counting_draw)
|
||||
kwargs = _arize_kwargs({"arize_success_sampling_rate": "1.0"})
|
||||
logger.log_success_event(kwargs, None, _START, _END)
|
||||
await logger.async_log_success_event(kwargs, None, _START, _END)
|
||||
assert len(draws) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unparsable_sampling_rate_exports_rather_than_dropping():
|
||||
logger, exporter = _sampled_arize_logger(random_draw=lambda: 0.99)
|
||||
await logger.async_log_success_event(_arize_kwargs({"arize_success_sampling_rate": "abc"}), None, _START, _END)
|
||||
assert _request_spans(exporter) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("bad", ["nan", "inf", "-inf", "1.5", "-0.1"])
|
||||
async def test_out_of_range_sampling_rate_exports_rather_than_dropping(bad: str):
|
||||
logger, exporter = _sampled_arize_logger(random_draw=lambda: 0.99)
|
||||
await logger.async_log_success_event(_arize_kwargs({"arize_success_sampling_rate": bad}), None, _START, _END)
|
||||
assert _request_spans(exporter) == 1
|
||||
|
|
|
|||
|
|
@ -1,4 +1,3 @@
|
|||
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
|
|
@ -246,9 +245,7 @@ def test_trusted_vars_overlay_uses_shared_parser_semantics():
|
|||
# datadog handler consumes, so values are str()-coerced identically.
|
||||
from litellm.types.utils import TRUSTED_CALLBACK_VARS_FIELD
|
||||
|
||||
params = initialize_standard_callback_dynamic_params(
|
||||
{TRUSTED_CALLBACK_VARS_FIELD: {"newrelic_api_key": 12345}}
|
||||
)
|
||||
params = initialize_standard_callback_dynamic_params({TRUSTED_CALLBACK_VARS_FIELD: {"newrelic_api_key": 12345}})
|
||||
|
||||
assert params.get("newrelic_api_key") == "12345"
|
||||
|
||||
|
|
@ -266,3 +263,19 @@ def test_validate_langfuse_environment_value():
|
|||
for bad in ["Production", "langfuse-eu", "", "team a"]:
|
||||
with pytest.raises(ValueError, match="langfuse_environment"):
|
||||
validate_langfuse_environment_value(bad)
|
||||
|
||||
|
||||
def test_arize_sampling_rates_are_picked_up_from_metadata():
|
||||
kwargs = {
|
||||
"litellm_params": {
|
||||
"metadata": {
|
||||
"arize_success_sampling_rate": "0.5",
|
||||
"arize_error_sampling_rate": "0.1",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
params = initialize_standard_callback_dynamic_params(kwargs)
|
||||
|
||||
assert params.get("arize_success_sampling_rate") == "0.5"
|
||||
assert params.get("arize_error_sampling_rate") == "0.1"
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ import pytest
|
|||
from litellm.proxy.common_utils.callback_config_validation import (
|
||||
callback_config_error,
|
||||
conflicting_span_scope_error,
|
||||
cross_entry_family_error,
|
||||
logging_metadata_config_error,
|
||||
)
|
||||
|
||||
|
|
@ -67,8 +68,16 @@ def test_one_span_scope_per_team(new_vars, stored, rejected):
|
|||
def test_key_logging_entries_may_not_disagree_on_the_span_scope():
|
||||
disagreeing = {
|
||||
"logging": [
|
||||
{"callback_name": "langfuse_otel", "callback_type": "success", "callback_vars": {"langfuse_span_scope": "full"}},
|
||||
{"callback_name": "langfuse_otel", "callback_type": "failure", "callback_vars": {"langfuse_span_scope": "llm_only"}},
|
||||
{
|
||||
"callback_name": "langfuse_otel",
|
||||
"callback_type": "success",
|
||||
"callback_vars": {"langfuse_span_scope": "full"},
|
||||
},
|
||||
{
|
||||
"callback_name": "langfuse_otel",
|
||||
"callback_type": "failure",
|
||||
"callback_vars": {"langfuse_span_scope": "llm_only"},
|
||||
},
|
||||
]
|
||||
}
|
||||
error = logging_metadata_config_error(disagreeing)
|
||||
|
|
@ -76,9 +85,46 @@ def test_key_logging_entries_may_not_disagree_on_the_span_scope():
|
|||
|
||||
agreeing = {
|
||||
"logging": [
|
||||
{"callback_name": "langfuse_otel", "callback_type": "success", "callback_vars": {"langfuse_span_scope": "llm_only"}},
|
||||
{"callback_name": "langfuse_otel", "callback_type": "failure", "callback_vars": {"langfuse_span_scope": "llm_only"}},
|
||||
{
|
||||
"callback_name": "langfuse_otel",
|
||||
"callback_type": "success",
|
||||
"callback_vars": {"langfuse_span_scope": "llm_only"},
|
||||
},
|
||||
{
|
||||
"callback_name": "langfuse_otel",
|
||||
"callback_type": "failure",
|
||||
"callback_vars": {"langfuse_span_scope": "llm_only"},
|
||||
},
|
||||
{"callback_name": "otel", "callback_type": "success", "callback_vars": {}},
|
||||
]
|
||||
}
|
||||
assert logging_metadata_config_error(agreeing) is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("var", ["arize_success_sampling_rate", "arize_error_sampling_rate"])
|
||||
@pytest.mark.parametrize("bad", ["1.5", "-0.1", "abc", "nan", "inf"])
|
||||
def test_callback_config_error_rejects_out_of_range_arize_sampling_rate(var, bad):
|
||||
error = callback_config_error("arize", {var: bad})
|
||||
assert error is not None and var in error and repr(bad) in error
|
||||
|
||||
|
||||
@pytest.mark.parametrize("var", ["arize_success_sampling_rate", "arize_error_sampling_rate"])
|
||||
@pytest.mark.parametrize("good", ["0", "1", "0.25", "", "None"])
|
||||
def test_callback_config_error_accepts_in_range_arize_sampling_rate(var, good):
|
||||
assert callback_config_error("arize", {var: good}) is None
|
||||
|
||||
|
||||
def test_arize_sampling_rate_rejected_on_non_arize_callback():
|
||||
error = callback_config_error("langfuse", {"arize_success_sampling_rate": "0.5"})
|
||||
assert error is not None
|
||||
assert "applies to the arize callback only" in error
|
||||
assert callback_config_error("arize", {"arize_success_sampling_rate": "0.5"}) is None
|
||||
|
||||
|
||||
def test_arize_sampling_rates_are_not_family_credentials():
|
||||
"""The rates choose what the Arize family exports, not where it sends, so an
|
||||
entry that repeats or adds a rate next to a stored Arize entry is not the
|
||||
credential-redirect shape cross_entry_family_error rejects."""
|
||||
stored = [{"arize_api_key": "k1", "arize_success_sampling_rate": "0.5"}]
|
||||
assert cross_entry_family_error({"arize_success_sampling_rate": "0.1"}, stored) is None
|
||||
assert cross_entry_family_error({"arize_error_sampling_rate": "0.5"}, stored) is None
|
||||
|
|
|
|||
|
|
@ -110,9 +110,7 @@ def patched_prisma():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_team_callbacks_rejects_unauthorized_caller(
|
||||
patched_prisma, unauthorized_caller
|
||||
):
|
||||
async def test_add_team_callbacks_rejects_unauthorized_caller(patched_prisma, unauthorized_caller):
|
||||
data = AddTeamCallback(
|
||||
callback_name="langfuse",
|
||||
callback_type="success",
|
||||
|
|
@ -133,9 +131,7 @@ async def test_add_team_callbacks_rejects_unauthorized_caller(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disable_team_logging_rejects_unauthorized_caller(
|
||||
patched_prisma, unauthorized_caller
|
||||
):
|
||||
async def test_disable_team_logging_rejects_unauthorized_caller(patched_prisma, unauthorized_caller):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await disable_team_logging(
|
||||
http_request=Mock(spec=Request),
|
||||
|
|
@ -147,9 +143,7 @@ async def test_disable_team_logging_rejects_unauthorized_caller(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_team_callbacks_rejects_unauthorized_caller(
|
||||
patched_prisma, unauthorized_caller
|
||||
):
|
||||
async def test_get_team_callbacks_rejects_unauthorized_caller(patched_prisma, unauthorized_caller):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await get_team_callbacks(
|
||||
http_request=Mock(spec=Request),
|
||||
|
|
@ -180,9 +174,7 @@ async def test_proxy_admin_can_add_team_callbacks(patched_prisma):
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_admin_of_target_team_can_add_callbacks(patched_prisma):
|
||||
patched_prisma.get_data = AsyncMock(
|
||||
return_value=_team_row(admin_user_id="team_admin_user")
|
||||
)
|
||||
patched_prisma.get_data = AsyncMock(return_value=_team_row(admin_user_id="team_admin_user"))
|
||||
|
||||
team_admin = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
|
|
@ -472,9 +464,7 @@ async def test_add_team_callbacks_writes_encrypted_callback_vars(monkeypatch):
|
|||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
written = json.loads(
|
||||
mock_prisma.db.litellm_teamtable.update.await_args.kwargs["data"]["metadata"]
|
||||
)
|
||||
written = json.loads(mock_prisma.db.litellm_teamtable.update.await_args.kwargs["data"]["metadata"])
|
||||
cv = written["logging"][0]["callback_vars"]
|
||||
assert cv["langfuse_secret_key"] != "sk-lf-real-secret"
|
||||
assert cv["langfuse_public_key"] != "pk-lf-real-public"
|
||||
|
|
@ -1192,8 +1182,12 @@ async def test_add_team_callbacks_rejects_team_deleted_before_write():
|
|||
)
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: proxy_server module global is the endpoint's only injection point
|
||||
patch("litellm.proxy.proxy_server.master_key", None), # test-quality-ok: proxy_server module global is the endpoint's only injection point
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.prisma_client", mock_prisma
|
||||
), # test-quality-ok: proxy_server module global is the endpoint's only injection point
|
||||
patch(
|
||||
"litellm.proxy.proxy_server.master_key", None
|
||||
), # test-quality-ok: proxy_server module global is the endpoint's only injection point
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await add_team_callbacks(
|
||||
|
|
@ -1482,12 +1476,16 @@ async def test_unknown_team_is_indistinguishable_from_no_access(call_handler, un
|
|||
a probe for valid team ids. The unknown-team response has to match the
|
||||
no-access one exactly, status and body.
|
||||
"""
|
||||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_client: # test-quality-ok: the handler imports prisma_client from proxy_server at call time, so there is no seam to inject through
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_client
|
||||
): # test-quality-ok: the handler imports prisma_client from proxy_server at call time, so there is no seam to inject through
|
||||
mock_client.get_data = AsyncMock(return_value=None)
|
||||
with pytest.raises(HTTPException) as unknown_team:
|
||||
await call_handler(unauthorized_caller)
|
||||
|
||||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_client: # test-quality-ok: the handler imports prisma_client from proxy_server at call time, so there is no seam to inject through
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_client
|
||||
): # test-quality-ok: the handler imports prisma_client from proxy_server at call time, so there is no seam to inject through
|
||||
mock_client.get_data = AsyncMock(return_value=_team_row())
|
||||
mock_client.db.litellm_teamtable.update = AsyncMock()
|
||||
with patch( # test-quality-ok: _verify_team_access calls this module-level helper directly, so there is no seam to inject through
|
||||
|
|
@ -1508,7 +1506,9 @@ async def test_proxy_admin_still_told_the_team_is_unknown():
|
|||
"""The masking is only for callers who could not have managed the team; a proxy
|
||||
admin keeps the diagnosable error."""
|
||||
admin = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin", api_key="sk-admin")
|
||||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_client: # test-quality-ok: the handler imports prisma_client from proxy_server at call time, so there is no seam to inject through
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client") as mock_client
|
||||
): # test-quality-ok: the handler imports prisma_client from proxy_server at call time, so there is no seam to inject through
|
||||
mock_client.get_data = AsyncMock(return_value=None)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await get_team_callbacks(
|
||||
|
|
@ -1526,13 +1526,29 @@ async def test_proxy_admin_still_told_the_team_is_unknown():
|
|||
[
|
||||
# the redirect, in every carrier a caller could pick: an entry naming
|
||||
# only a host, pairing with a key pair written on another entry
|
||||
({"langfuse_host": "http://attacker.invalid"}, [{"langfuse_public_key": "pk", "langfuse_secret_key": "sk"}], True),
|
||||
(
|
||||
{"langfuse_host": "http://attacker.invalid"},
|
||||
[{"langfuse_public_key": "pk", "langfuse_secret_key": "sk"}],
|
||||
True,
|
||||
),
|
||||
# the sibling carrier -- langfuse and langfuse_otel are one account
|
||||
({"langfuse_host": "http://attacker.invalid"}, [{"langfuse_host": "https://us.cloud.langfuse.com", "langfuse_secret_key": "sk"}], True),
|
||||
(
|
||||
{"langfuse_host": "http://attacker.invalid"},
|
||||
[{"langfuse_host": "https://us.cloud.langfuse.com", "langfuse_secret_key": "sk"}],
|
||||
True,
|
||||
),
|
||||
# a destination variable no integration registry lists
|
||||
({"dd_agent_host": "attacker.invalid"}, [{"dd_api_key": "k", "dd_site": "us5.datadoghq.com"}], True),
|
||||
# one entry owning its family end to end is the feature
|
||||
({"langfuse_host": "https://eu.cloud.langfuse.com", "langfuse_public_key": "pk", "langfuse_secret_key": "sk"}, [], False),
|
||||
(
|
||||
{
|
||||
"langfuse_host": "https://eu.cloud.langfuse.com",
|
||||
"langfuse_public_key": "pk",
|
||||
"langfuse_secret_key": "sk",
|
||||
},
|
||||
[],
|
||||
False,
|
||||
),
|
||||
# a different family alongside an existing one stays fine
|
||||
({"gcs_bucket_name": "bucket"}, [{"langfuse_public_key": "pk", "langfuse_secret_key": "sk"}], False),
|
||||
({"langsmith_api_key": "k"}, [{"dd_api_key": "k"}], False),
|
||||
|
|
@ -1541,19 +1557,55 @@ async def test_proxy_admin_still_told_the_team_is_unknown():
|
|||
# the span scope picks what the family exports, not where to, so a second
|
||||
# entry may set either legal value next to the family's credentials
|
||||
({"langfuse_span_scope": "llm_only"}, [{"langfuse_public_key": "pk", "langfuse_secret_key": "sk"}], False),
|
||||
({"langfuse_public_key": "pk", "langfuse_secret_key": "sk", "langfuse_span_scope": "full"}, [{"langfuse_public_key": "pk", "langfuse_secret_key": "sk", "langfuse_span_scope": "llm_only"}], False),
|
||||
(
|
||||
{"langfuse_public_key": "pk", "langfuse_secret_key": "sk", "langfuse_span_scope": "full"},
|
||||
[{"langfuse_public_key": "pk", "langfuse_secret_key": "sk", "langfuse_span_scope": "llm_only"}],
|
||||
False,
|
||||
),
|
||||
# the scope on the stored entry must not shield a redirect riding next to it
|
||||
({"langfuse_host": "http://attacker.invalid", "langfuse_span_scope": "llm_only"}, [{"langfuse_public_key": "pk", "langfuse_secret_key": "sk", "langfuse_span_scope": "llm_only"}], True),
|
||||
(
|
||||
{"langfuse_host": "http://attacker.invalid", "langfuse_span_scope": "llm_only"},
|
||||
[{"langfuse_public_key": "pk", "langfuse_secret_key": "sk", "langfuse_span_scope": "llm_only"}],
|
||||
True,
|
||||
),
|
||||
# the same integration registered for a second event: identical values
|
||||
# flatten to the identical dict, so there is nothing to redirect
|
||||
({"langfuse_host": "https://us.cloud.langfuse.com", "langfuse_public_key": "pk", "langfuse_secret_key": "sk"}, [{"langfuse_host": "https://us.cloud.langfuse.com", "langfuse_public_key": "pk", "langfuse_secret_key": "sk"}], False),
|
||||
(
|
||||
{
|
||||
"langfuse_host": "https://us.cloud.langfuse.com",
|
||||
"langfuse_public_key": "pk",
|
||||
"langfuse_secret_key": "sk",
|
||||
},
|
||||
[
|
||||
{
|
||||
"langfuse_host": "https://us.cloud.langfuse.com",
|
||||
"langfuse_public_key": "pk",
|
||||
"langfuse_secret_key": "sk",
|
||||
}
|
||||
],
|
||||
False,
|
||||
),
|
||||
# the same credential under its other spelling is the same credential
|
||||
({"langfuse_secret": "sk"}, [{"langfuse_public_key": "pk", "langfuse_secret_key": "sk"}], False),
|
||||
# a value the family already holds cannot be moved into another of its
|
||||
# variables either; the exporter would address or authenticate with it
|
||||
({"langfuse_host": "pk"}, [{"langfuse_host": "https://us.cloud.langfuse.com", "langfuse_public_key": "pk"}], True),
|
||||
(
|
||||
{"langfuse_host": "pk"},
|
||||
[{"langfuse_host": "https://us.cloud.langfuse.com", "langfuse_public_key": "pk"}],
|
||||
True,
|
||||
),
|
||||
# the same shape with one value moved is the redirect again
|
||||
({"langfuse_host": "http://attacker.invalid", "langfuse_public_key": "pk", "langfuse_secret_key": "sk"}, [{"langfuse_host": "https://us.cloud.langfuse.com", "langfuse_public_key": "pk", "langfuse_secret_key": "sk"}], True),
|
||||
(
|
||||
{"langfuse_host": "http://attacker.invalid", "langfuse_public_key": "pk", "langfuse_secret_key": "sk"},
|
||||
[
|
||||
{
|
||||
"langfuse_host": "https://us.cloud.langfuse.com",
|
||||
"langfuse_public_key": "pk",
|
||||
"langfuse_secret_key": "sk",
|
||||
}
|
||||
],
|
||||
True,
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_one_entry_owns_a_credential_family(new_vars, stored, rejected):
|
||||
|
|
@ -1568,7 +1620,13 @@ def test_one_entry_owns_a_credential_family(new_vars, stored, rejected):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("caller", [_admin_auth(), UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="victim_admin", api_key="sk-team-admin")])
|
||||
@pytest.mark.parametrize(
|
||||
"caller",
|
||||
[
|
||||
_admin_auth(),
|
||||
UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="victim_admin", api_key="sk-team-admin"),
|
||||
],
|
||||
)
|
||||
async def test_a_second_entry_may_not_flip_the_span_scope(patched_prisma, caller):
|
||||
"""The entries flatten last-wins at request time, so a failure entry saying
|
||||
llm_only next to a success entry saying full would export whichever is stored
|
||||
|
|
@ -1580,7 +1638,11 @@ async def test_a_second_entry_may_not_flip_the_span_scope(patched_prisma, caller
|
|||
{
|
||||
"callback_name": "langfuse_otel",
|
||||
"callback_type": "success",
|
||||
"callback_vars": {"langfuse_public_key": "pk", "langfuse_secret_key": "sk", "langfuse_span_scope": "full"},
|
||||
"callback_vars": {
|
||||
"langfuse_public_key": "pk",
|
||||
"langfuse_secret_key": "sk",
|
||||
"langfuse_span_scope": "full",
|
||||
},
|
||||
}
|
||||
]
|
||||
}
|
||||
|
|
@ -1610,3 +1672,35 @@ async def test_a_second_entry_may_not_flip_the_span_scope(patched_prisma, caller
|
|||
user_api_key_dict=caller,
|
||||
)
|
||||
patched_prisma.db.litellm_teamtable.update.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_team_callbacks_rejects_out_of_range_arize_sampling_rate(patched_prisma):
|
||||
data = AddTeamCallback(
|
||||
callback_name="arize",
|
||||
callback_type="success",
|
||||
callback_vars={"arize_success_sampling_rate": "1.5"},
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await add_team_callbacks(
|
||||
data=data,
|
||||
http_request=Mock(spec=Request),
|
||||
team_id="team-victim",
|
||||
user_api_key_dict=_admin_auth(),
|
||||
)
|
||||
assert exc.value.status_code == 400
|
||||
assert "arize_success_sampling_rate" in str(exc.value.detail)
|
||||
patched_prisma.db.litellm_teamtable.update.assert_not_called()
|
||||
|
||||
|
||||
def test_add_team_callback_accepts_arize_sampling_rate_vars():
|
||||
data = AddTeamCallback(
|
||||
callback_name="arize",
|
||||
callback_type="success",
|
||||
callback_vars={
|
||||
"arize_success_sampling_rate": "0.5",
|
||||
"arize_error_sampling_rate": "1.0",
|
||||
},
|
||||
)
|
||||
assert data.callback_vars["arize_success_sampling_rate"] == "0.5"
|
||||
assert data.callback_vars["arize_error_sampling_rate"] == "1.0"
|
||||
|
|
|
|||
|
|
@ -30,6 +30,8 @@ export const CALLBACK_CONFIGS: CallbackConfig[] = [
|
|||
dynamic_params: {
|
||||
arize_api_key: "password",
|
||||
arize_space_id: "password",
|
||||
arize_success_sampling_rate: "number",
|
||||
arize_error_sampling_rate: "number",
|
||||
},
|
||||
description: "Arize Logging Integration",
|
||||
},
|
||||
|
|
|
|||
|
|
@ -239,6 +239,33 @@ describe("LoggingSettings", () => {
|
|||
]);
|
||||
});
|
||||
|
||||
it("renders sampling rate inputs for the Arize callback and records changes", () => {
|
||||
const mockOnChange = vi.fn();
|
||||
|
||||
const initialValue = [
|
||||
{
|
||||
callback_name: "arize",
|
||||
callback_type: "success",
|
||||
callback_vars: {},
|
||||
},
|
||||
];
|
||||
|
||||
renderWithProviders(<LoggingSettings value={initialValue} onChange={mockOnChange} />);
|
||||
|
||||
const successInput = screen.getByPlaceholderText("os.environ/ARIZE_SUCCESS_SAMPLING_RATE");
|
||||
const errorInput = screen.getByPlaceholderText("os.environ/ARIZE_ERROR_SAMPLING_RATE");
|
||||
expect(successInput).toBeInTheDocument();
|
||||
expect(errorInput).toBeInTheDocument();
|
||||
|
||||
fireEvent.change(successInput, { target: { value: "0.4" } });
|
||||
let lastCall = mockOnChange.mock.calls[mockOnChange.mock.calls.length - 1];
|
||||
expect(lastCall[0][0].callback_vars.arize_success_sampling_rate).toBe("0.4");
|
||||
|
||||
fireEvent.change(errorInput, { target: { value: "0.9" } });
|
||||
lastCall = mockOnChange.mock.calls[mockOnChange.mock.calls.length - 1];
|
||||
expect(lastCall[0][0].callback_vars.arize_error_sampling_rate).toBe("0.9");
|
||||
});
|
||||
|
||||
it("correctly handles numerical input with decimal values", () => {
|
||||
const mockOnChange = vi.fn();
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue