diff --git a/litellm/integrations/arize/arize.py b/litellm/integrations/arize/arize.py index 2e5b17185f2..4ab8d9796b2 100644 --- a/litellm/integrations/arize/arize.py +++ b/litellm/integrations/arize/arize.py @@ -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 diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 180929bcfd4..feaeaec27a3 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -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): diff --git a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py index 34d0ded618c..00ab05aba77 100644 --- a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py +++ b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py @@ -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", diff --git a/litellm/proxy/common_utils/callback_config_validation.py b/litellm/proxy/common_utils/callback_config_validation.py index 30c4ab31d6f..6a0fcbe0bb3 100644 --- a/litellm/proxy/common_utils/callback_config_validation.py +++ b/litellm/proxy/common_utils/callback_config_validation.py @@ -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. diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 346181d9d9e..bb0b3d2a687 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -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 diff --git a/tests/test_litellm/integrations/arize/test_arize.py b/tests/test_litellm/integrations/arize/test_arize.py index cdafd856b49..5fde627680d 100644 --- a/tests/test_litellm/integrations/arize/test_arize.py +++ b/tests/test_litellm/integrations/arize/test_arize.py @@ -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 diff --git a/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py b/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py index 521244c6ede..a9a6509f6dd 100644 --- a/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py +++ b/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py @@ -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" diff --git a/tests/test_litellm/proxy/common_utils/test_callback_config_validation.py b/tests/test_litellm/proxy/common_utils/test_callback_config_validation.py index a707c92dc15..33ae986a511 100644 --- a/tests/test_litellm/proxy/common_utils/test_callback_config_validation.py +++ b/tests/test_litellm/proxy/common_utils/test_callback_config_validation.py @@ -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 diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_callback_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_callback_endpoints.py index acc7c8ca21a..b6eebcb2ef3 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_callback_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_callback_endpoints.py @@ -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" diff --git a/ui/litellm-dashboard/src/components/callback_info_helpers.tsx b/ui/litellm-dashboard/src/components/callback_info_helpers.tsx index b43a85d18c5..f5138d55b5d 100644 --- a/ui/litellm-dashboard/src/components/callback_info_helpers.tsx +++ b/ui/litellm-dashboard/src/components/callback_info_helpers.tsx @@ -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", }, diff --git a/ui/litellm-dashboard/src/components/team/LoggingSettings.test.tsx b/ui/litellm-dashboard/src/components/team/LoggingSettings.test.tsx index f7ca6d516f5..5ce8ff71267 100644 --- a/ui/litellm-dashboard/src/components/team/LoggingSettings.test.tsx +++ b/ui/litellm-dashboard/src/components/team/LoggingSettings.test.tsx @@ -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(); + + 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();