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:
devin-ai-integration[bot] 2026-09-22 11:21:40 -05:00 • committed by GitHub
parent 226fa7845a
commit b673ee61e7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 487 additions and 78 deletions

View file

@ -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

View file

@ -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):

View file

@ -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",

View file

@ -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.

View file

@ -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

View file

@ -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

View file

@ -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"

View file

@ -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

View file

@ -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"

View file

@ -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",
},

View file

@ -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();