diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index d4162369a35..5f0289f833c 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -4,12 +4,16 @@ import re import traceback from collections.abc import AsyncGenerator, Mapping, Sequence from types import MappingProxyType -from typing import TYPE_CHECKING, Any, ClassVar, Final, Optional +from typing import TYPE_CHECKING, Any, ClassVar, Final, Optional, cast from pydantic import BaseModel from litellm._logging import verbose_logger -from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER, EMPTY_MAPPING +from litellm.constants import ( + DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER, + EMPTY_MAPPING, + REDACTED_BY_LITELLM, +) from litellm.types.integrations.argilla import ArgillaItem from litellm.types.integrations.custom_logger import AgenticLoopPlan from litellm.types.llms.openai import AllMessageValues, ChatCompletionRequest @@ -24,6 +28,7 @@ from litellm.types.utils import ( StandardAuditLogPayload, StandardCallbackDynamicParams, StandardLoggingPayload, + StandardLoggingPayloadErrorInformation, ) if TYPE_CHECKING: @@ -62,6 +67,22 @@ _BASE64_INLINE_PATTERN: Final = re.compile( ) +def _redacted_failure_error_fields(standard_logging_object: Mapping[str, object]) -> dict[str, object]: + from litellm.litellm_core_utils.redact_messages import redact_error_information + + fields: Final[dict[str, object]] = {} # mutable-ok: merged into the standard_logging_object copy below + if standard_logging_object.get("error_str"): + fields["error_str"] = REDACTED_BY_LITELLM + error_information: Final = standard_logging_object.get("error_information") + if isinstance(error_information, Mapping): + fields["error_information"] = redact_error_information( + cast( # cast-ok: same TypedDict shape as the input mapping + StandardLoggingPayloadErrorInformation, error_information + ) + ) + return fields + + class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callback#callback-class # Class variables or attributes server_fulfilled_tool_names: ClassVar[frozenset[str]] = frozenset() @@ -900,7 +921,9 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac This method handles two features: 1. turn_off_message_logging: When True, redacts messages and responses (unless the callback - redacts them itself, see `redacts_messages_itself`) + redacts them itself, see `redacts_messages_itself`), and redacts `error_str`, + `error_information`'s message/traceback and `traceback_exception` independent of + `redacts_messages_itself` 2. standard_logging_payload_excluded_fields: Removes specified fields entirely Return a modified copy of the provided logging payload. @@ -960,6 +983,9 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac model_response_dict: Final = model_response.model_dump() standard_logging_object_copy["response"] = model_response_dict + if turn_off_message_logging: + standard_logging_object_copy.update(_redacted_failure_error_fields(standard_logging_object_copy)) + params: Final = model_call_details.get("litellm_params") request: Final = params.get("proxy_server_request") if isinstance(params, dict) else None redacted_params: Final = ( @@ -967,9 +993,15 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac if turn_off_message_logging and isinstance(params, dict) and isinstance(request, dict) else EMPTY_MAPPING ) + redacted_failure_fields: Final = ( + MappingProxyType({"traceback_exception": REDACTED_BY_LITELLM}) + if turn_off_message_logging and model_call_details.get("traceback_exception") + else EMPTY_MAPPING + ) return { **model_call_details, **redacted_params, + **redacted_failure_fields, "standard_logging_object": standard_logging_object_copy, } diff --git a/litellm/integrations/datadog/datadog.py b/litellm/integrations/datadog/datadog.py index 2ca8b0ed236..49958eb63ea 100644 --- a/litellm/integrations/datadog/datadog.py +++ b/litellm/integrations/datadog/datadog.py @@ -303,12 +303,21 @@ class DataDogLogger( from litellm.litellm_core_utils.litellm_logging import ( StandardLoggingPayloadSetup, ) + from litellm.litellm_core_utils.redact_messages import ( + redact_error_information, + should_redact_failed_request, + ) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps - error_information: Final = StandardLoggingPayloadSetup.get_error_information( + _error_information_raw: Final = StandardLoggingPayloadSetup.get_error_information( original_exception=original_exception, traceback_str=traceback_str, ) + error_information: Final = ( + redact_error_information(_error_information_raw) + if should_redact_failed_request(request_data) + else _error_information_raw + ) _code: Final = error_information.get("error_code") or "" status_code: int | None = None if _code and str(_code).strip().isdigit(): diff --git a/litellm/integrations/mlflow.py b/litellm/integrations/mlflow.py index 8731e96440f..c3da1aecc8f 100644 --- a/litellm/integrations/mlflow.py +++ b/litellm/integrations/mlflow.py @@ -3,7 +3,9 @@ import threading from typing import TYPE_CHECKING, Any, Final from litellm._logging import verbose_logger +from litellm.constants import REDACTED_BY_LITELLM from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.redact_messages import should_redact_message_logging if TYPE_CHECKING: from litellm.types.utils import StandardLoggingPayload @@ -88,7 +90,19 @@ class MlflowLogger(CustomLogger): # Record exception info as event if exception := kwargs.get("exception"): - span.add_event(SpanEvent.from_exception(exception)) + if should_redact_message_logging(kwargs): + span.add_event( + SpanEvent( + name="exception", + attributes={ # mutable-ok: mlflow SpanEvent expects a plain dict of attributes + "exception.type": type(exception).__name__, + "exception.message": REDACTED_BY_LITELLM, + "exception.stacktrace": REDACTED_BY_LITELLM, + }, + ) + ) + else: + span.add_event(SpanEvent.from_exception(exception)) self._extract_and_set_chat_attributes(span, kwargs, response_obj) self._end_span_or_trace( diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 8d588896b2f..2b3c9e398b3 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Any, Final, TypedDict, cast import litellm from litellm._logging import verbose_logger +from litellm.constants import REDACTED_BY_LITELLM from litellm.integrations._types.open_inference import ( OpenInferenceSpanKindValues, SpanAttributes, @@ -378,6 +379,23 @@ class OpenTelemetryConfig: ) +def _server_span_failure_redact(span: "Span") -> bool: + """Redaction decision for re-stamping a SERVER span the failure hook may + have already marked: an existing error message/stack-trace attribute keeps + its request-aware value, so a restamp with the global-only probe can't leak + an opt-in's raw text or clobber a valid opt-out.""" + from litellm.integrations._types.open_inference import ErrorAttributes + from litellm.litellm_core_utils.redact_messages import should_redact_message_logging + + attributes: Final = getattr(span, "attributes", None) or {} + if ErrorAttributes.ERROR_MESSAGE in attributes or ErrorAttributes.ERROR_STACK_TRACE in attributes: + return ( + attributes.get(ErrorAttributes.ERROR_MESSAGE) == REDACTED_BY_LITELLM + or attributes.get(ErrorAttributes.ERROR_STACK_TRACE) == REDACTED_BY_LITELLM + ) + return should_redact_message_logging({}) # mutable-ok: read-only probe for the global flag + + class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): def __init__( self, @@ -930,17 +948,26 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): from litellm.litellm_core_utils.litellm_logging import ( StandardLoggingPayloadSetup, ) + from litellm.litellm_core_utils.redact_messages import ( + redact_error_information, + should_redact_failed_request, + ) - error_information: Final = StandardLoggingPayloadSetup.get_error_information( + redact: Final = should_redact_failed_request(request_data) + _error_information_raw: Final = StandardLoggingPayloadSetup.get_error_information( original_exception=original_exception, traceback_str=traceback_str, ) + error_information: Final = ( + redact_error_information(_error_information_raw) if redact else _error_information_raw + ) self._record_exception_on_span( span=parent_otel_span, kwargs={ "exception": original_exception, "standard_logging_object": {"error_information": error_information}, }, + redact_content=redact, ) # _record_exception_on_span only stamps when error_code is set; @@ -963,7 +990,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): self.safe_set_attribute( span=exception_logging_span, key="exception", - value=str(original_exception), + value=REDACTED_BY_LITELLM if redact else str(original_exception), ) self._set_team_attributes_on_span( span=exception_logging_span, @@ -2207,25 +2234,40 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): ): parent_otel_span.end(end_time=self._to_ns(end_time)) - def _record_exception_on_span(self, span: Span, kwargs: dict): + def _record_exception_on_span(self, span: Span, kwargs: dict, redact_content: bool | None = None): """ Record exception information on the span using OTEL standard methods. This extracts error information from StandardLoggingPayload and: 1. Uses span.record_exception() for the actual exception object (OTEL standard) 2. Sets structured error attributes from StandardLoggingPayloadErrorInformation + + ``redact_content`` overrides the message-redaction decision; when unset it is + resolved from kwargs via ``should_redact_message_logging``. """ try: from litellm.integrations._types.open_inference import ( ErrorAttributes, ) + from litellm.litellm_core_utils.redact_messages import should_redact_message_logging + + redact: Final = should_redact_message_logging(kwargs) if redact_content is None else redact_content # Get the exception object if available exception: Final = kwargs.get("exception") # Record the exception using OTEL's standard method if exception is not None: - span.record_exception(exception) + if redact: + span.record_exception( + exception, + attributes={ # mutable-ok: record_exception accepts a dict of event attributes + "exception.message": REDACTED_BY_LITELLM, + "exception.stacktrace": REDACTED_BY_LITELLM, + }, + ) + else: + span.record_exception(exception) # Get StandardLoggingPayload for structured error information standard_logging_payload: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object") @@ -3571,12 +3613,19 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): from litellm.litellm_core_utils.litellm_logging import ( StandardLoggingPayloadSetup, ) + from litellm.litellm_core_utils.redact_messages import redact_error_information + redact: Final = _server_span_failure_redact(span) error_information: Final = StandardLoggingPayloadSetup.get_error_information(original_exception=exception) error_information["error_code"] = str(status_code) self._record_exception_on_span( span=span, - kwargs={"standard_logging_object": {"error_information": error_information}}, + kwargs={ # mutable-ok: _record_exception_on_span reads this kwargs dict + "standard_logging_object": { # mutable-ok: standard_logging_object shape the recorder expects + "error_information": redact_error_information(error_information) if redact else error_information + } + }, + redact_content=redact, ) def set_preprocessing_duration_attribute(self, span: Span | None, container: object) -> None: diff --git a/litellm/integrations/otel/logger.py b/litellm/integrations/otel/logger.py index 55eb8e8fb71..e0a9f4df3e6 100644 --- a/litellm/integrations/otel/logger.py +++ b/litellm/integrations/otel/logger.py @@ -26,6 +26,7 @@ from typing_extensions import TypedDict, Unpack import litellm from litellm._logging import verbose_logger +from litellm.constants import REDACTED_BY_LITELLM from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.otel.emitter import SpanEmitter, stamp_error from litellm.integrations.otel.mappers import resolve_mappers @@ -48,7 +49,7 @@ from litellm.integrations.otel.model.payloads import ( is_mcp_list_tools, is_mcp_tool_call, ) -from litellm.integrations.otel.model.semconv import Error +from litellm.integrations.otel.model.semconv import Error, LiteLLMError from litellm.integrations.otel.model.spans import SpanRole, span_role_for_service from litellm.integrations.otel.model.utils import to_ns from litellm.integrations.otel.plumbing.context import ( @@ -77,6 +78,10 @@ from litellm.integrations.otel.plumbing.providers import ( resolve_meter_provider, ) from litellm.integrations.otel.plumbing.routing import TenantTracerCache +from litellm.litellm_core_utils.redact_messages import ( + should_redact_failed_request, + should_redact_message_logging, +) if TYPE_CHECKING: from opentelemetry.metrics import MeterProvider @@ -92,6 +97,21 @@ if TYPE_CHECKING: LITELLM_TRACER_NAME: Final = "litellm" _published_v2_provider: ApiTracerProvider | None = None +_GLOBAL_REDACTION_PROBE: Final[dict[str, object]] = {} # mutable-ok: read-only global-redaction probe + + +def _span_failure_redact(span: "Span") -> bool: + """Redaction decision for re-stamping a span the failure hook may have + already marked: an existing ``error.message``/stack-trace attribute keeps + its request-aware value, so a later restamp with the global probe can't + leak an opt-in's raw text or clobber a valid opt-out.""" + attributes: Final = getattr(span, "attributes", None) or {} + if Error.MESSAGE in attributes or LiteLLMError.STACK_TRACE in attributes: + return ( + attributes.get(Error.MESSAGE) == REDACTED_BY_LITELLM + or attributes.get(LiteLLMError.STACK_TRACE) == REDACTED_BY_LITELLM + ) + return should_redact_message_logging(_GLOBAL_REDACTION_PROBE) def _span_error_from_exception( @@ -99,6 +119,7 @@ def _span_error_from_exception( *, status_code: int | None = None, traceback_str: str | None = None, + redact_content: bool = False, ) -> SpanError: """A ``SpanError`` for a proxy-level failure that never produced a ``StandardLoggingPayload`` (auth / validation / malformed-body rejections), @@ -113,9 +134,13 @@ def _span_error_from_exception( ) return SpanError( error_type=info.get("error_class") or info.get("error_code") or None, - message=info.get("error_message") or None, + message=REDACTED_BY_LITELLM + if redact_content and info.get("error_message") + else (info.get("error_message") or None), code=str(status_code) if status_code is not None else (info.get("error_code") or None), - stack_trace=info.get("traceback") or None, + stack_trace=REDACTED_BY_LITELLM + if redact_content and info.get("traceback") + else (info.get("traceback") or None), llm_provider=info.get("llm_provider") or None, ) @@ -734,14 +759,25 @@ class OpenTelemetryV2(CustomLogger): pass @contextmanager - def start_phase_span(self, name: str) -> "Iterator[Span]": + def start_phase_span(self, name: str, *, redact_content: bool = False) -> "Iterator[Span]": span: Final = self._emitter.start_span(SpanRole.SERVICE, name) - with use_span(span, end_on_exit=True): + redact: Final = redact_content or should_redact_message_logging(_GLOBAL_REDACTION_PROBE) + with use_span( + span, + end_on_exit=True, + record_exception=not redact, + set_status_on_exception=not redact, + ): try: yield span except Exception as exc: if is_recordable_span(span): - stamp_error(span, _span_error_from_exception(exc), record_event=False, set_status=False) + stamp_error( + span, + _span_error_from_exception(exc, redact_content=redact), + record_event=redact, + set_status=redact, + ) raise async def async_pre_call_hook( @@ -782,7 +818,11 @@ class OpenTelemetryV2(CustomLogger): already_stamped: Final = Error.TYPE in (getattr(span, "attributes", None) or ()) stamp_error( span, - _span_error_from_exception(exception, status_code=status_code), + _span_error_from_exception( + exception, + status_code=status_code, + redact_content=_span_failure_redact(span), + ), record_event=not already_stamped, ) @@ -808,7 +848,14 @@ class OpenTelemetryV2(CustomLogger): span: Final = mcp_message_transport_span() or request_root_span() or user_api_key_dict.parent_otel_span if span is None or not is_recordable_span(span): return - stamp_error(span, _span_error_from_exception(original_exception, traceback_str=traceback_str)) + stamp_error( + span, + _span_error_from_exception( + original_exception, + traceback_str=traceback_str, + redact_content=should_redact_failed_request(request_data), + ), + ) return def emit_guardrail_span(self, entry: "StandardLoggingGuardrailInformation") -> None: @@ -990,12 +1037,12 @@ def fan_out_provider() -> ApiTracerProvider: @contextmanager -def phase_span(name: str) -> "Iterator[Span | None]": +def phase_span(name: str, *, redact_content: bool = False) -> "Iterator[Span | None]": logger: Final = _registered_v2_logger() if logger is None: yield None return - with logger.start_phase_span(name) as span: + with logger.start_phase_span(name, redact_content=redact_content) as span: yield span diff --git a/litellm/integrations/otel/runtime.py b/litellm/integrations/otel/runtime.py index 13903597e1a..1ca60c72639 100644 --- a/litellm/integrations/otel/runtime.py +++ b/litellm/integrations/otel/runtime.py @@ -17,7 +17,7 @@ if TYPE_CHECKING: @cache -def _otel_runtime() -> "tuple[Callable[[str], AbstractContextManager[Span | None]], Callable[..., None]] | None": +def _otel_runtime() -> "tuple[Callable[..., AbstractContextManager[Span | None]], Callable[..., None]] | None": """Resolve the SDK-backed hooks once and cache the outcome, absence included. CPython never caches a failed import, so without this memoization every call @@ -32,7 +32,7 @@ def _otel_runtime() -> "tuple[Callable[[str], AbstractContextManager[Span | None @contextmanager -def phase_span(name: str) -> "Iterator[Span | None]": +def phase_span(name: str, *, redact_content: bool = False) -> "Iterator[Span | None]": """Run a request phase inside a live active span so its DB/service calls nest. Yields ``None`` (a plain no-op) when the OTel SDK is unavailable or V2 is not @@ -42,7 +42,7 @@ def phase_span(name: str) -> "Iterator[Span | None]": if runtime is None: yield None return - with runtime[0](name) as span: + with runtime[0](name, redact_content=redact_content) as span: yield span diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 2162a200565..27d27bcee49 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -3782,7 +3782,9 @@ class Logging(LiteLLMLoggingBaseClass): start_time=start_time, end_time=end_time, response_obj=result, - kwargs=self.model_call_details, + kwargs=callback.redact_standard_logging_payload_from_model_call_details( + model_call_details=self.model_call_details + ), ) if callback == "langfuse": global langFuseLogger @@ -3900,8 +3902,10 @@ class Logging(LiteLLMLoggingBaseClass): global_callbacks=litellm._async_failure_callback, ) - result: Final = None # result sent to all loggers, init this to None incase it's not created - + result: Final = redact_message_input_output_from_logging( + model_call_details=self.model_call_details, + result=None, + ) self.has_run_logging(event_type="async_failure") for callback in callbacks: try: @@ -3915,7 +3919,9 @@ class Logging(LiteLLMLoggingBaseClass): continue if isinstance(callback, CustomLogger): # custom logger class await callback.async_log_failure_event( - kwargs=self.model_call_details, + kwargs=callback.redact_standard_logging_payload_from_model_call_details( + model_call_details=self.model_call_details + ), response_obj=result, start_time=start_time, end_time=end_time, diff --git a/litellm/litellm_core_utils/redact_messages.py b/litellm/litellm_core_utils/redact_messages.py index 8c77ef32cfb..80318a07775 100644 --- a/litellm/litellm_core_utils/redact_messages.py +++ b/litellm/litellm_core_utils/redact_messages.py @@ -11,7 +11,8 @@ import asyncio import copy import inspect from collections.abc import Mapping -from typing import TYPE_CHECKING, Any, Final +from types import MappingProxyType +from typing import TYPE_CHECKING, Any, Final, cast import litellm from litellm.constants import REDACTED_BY_LITELLM @@ -20,13 +21,19 @@ from litellm.litellm_core_utils.classifier_logging import without_classifier_aud from litellm.litellm_core_utils.core_helpers import ( get_metadata_variable_name_from_kwargs, ) +from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( + iter_client_callback_metadata_dicts, +) from litellm.litellm_core_utils.served_output_texts import SERVED_OUTPUT_TEXTS_KEY from litellm.llms.vertex_ai.common_utils import ( redact_vertex_ai_metadata_from_litellm_params, redact_vertex_ai_metadata_from_logged_object, ) from litellm.secret_managers.main import str_to_bool -from litellm.types.utils import StandardCallbackDynamicParams +from litellm.types.utils import ( + StandardCallbackDynamicParams, + StandardLoggingPayloadErrorInformation, +) if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import ( @@ -183,6 +190,100 @@ def redacted_standard_logging_payload(payload: Mapping[str, object]) -> Mapping[ return _redact_standard_logging_object(payload) +_REDACTED_ERROR_FIELDS: Final = ("error_message", "traceback") + +_MESSAGE_REDACTION_ENABLE_HEADERS: Final = ( + "litellm-enable-message-redaction", # old header. maintain backwards compatibility + "x-litellm-enable-message-redaction", # new header +) + + +def _is_non_empty_str(value: object) -> bool: + return isinstance(value, str) and bool(value) + + +def redact_error_information( + error_information: StandardLoggingPayloadErrorInformation, +) -> StandardLoggingPayloadErrorInformation: + """ + Return a copy of ``StandardLoggingPayloadErrorInformation`` with the fields that can + quote the prompt (``error_message``, ``traceback``) replaced by ``REDACTED_BY_LITELLM`` + when they are non-empty strings. Every other field is carried over unchanged. + """ + redacted_fields: Final = MappingProxyType( + { + field: REDACTED_BY_LITELLM + for field in _REDACTED_ERROR_FIELDS + if _is_non_empty_str(error_information.get(field)) + } + ) + return cast( # cast-ok: same TypedDict shape as the input, two fields narrowed to the sentinel + StandardLoggingPayloadErrorInformation, + {**error_information, **redacted_fields}, # mutable-ok: callers pop/update keys on the fresh payload dict + ) + + +def _request_turn_off_message_logging(request_data: Mapping[str, object]) -> object: + """``turn_off_message_logging`` resolves like ``initialize_standard_callback_dynamic_params``: + the top-level value when present, else the first client-metadata slot carrying it.""" + return ( + request_data["turn_off_message_logging"] + if "turn_off_message_logging" in request_data + else next( + ( + slot["turn_off_message_logging"] + for _, slot in iter_client_callback_metadata_dicts(dict(request_data)) + if "turn_off_message_logging" in slot + ), + None, + ) + ) + + +def request_opts_into_message_redaction(headers: Mapping[str, str], request_data: Mapping[str, object]) -> bool: + """Opt-in signals only: usable at auth time before the key's + ``allow_client_message_redaction_opt_out`` permission is known, so the disable + header and the global flag are deliberately not consulted.""" + return ( + any(bool(headers.get(header)) for header in _MESSAGE_REDACTION_ENABLE_HEADERS) + or _request_turn_off_message_logging(request_data) is True + ) + + +def should_redact_failed_request(request_data: Mapping[str, object]) -> bool: + """ + Message-redaction decision for proxy ``async_post_call_failure_hook`` callbacks, which + only receive ``request_data`` (the Logging object is popped before hooks run). + ``litellm_metadata`` is included only when present so + ``get_metadata_variable_name_from_kwargs`` resolves ``metadata`` for chat routes. + """ + dynamic_param: Final = _request_turn_off_message_logging(request_data) + litellm_params: Final = MappingProxyType( + { + key: request_data.get(key) + for key in ("metadata", "litellm_metadata") + if key == "metadata" or key in request_data + } + ) + return should_redact_message_logging( + { # mutable-ok: the model_call_details shape the decision helper reads + "litellm_params": litellm_params, + "standard_callback_dynamic_params": { # mutable-ok: dynamic-params slot the helper reads + "turn_off_message_logging": dynamic_param + }, + } + ) + + +def maybe_redact_error_information( + error_information: StandardLoggingPayloadErrorInformation, + request_data: Mapping[str, object], +) -> StandardLoggingPayloadErrorInformation: + if should_redact_failed_request(request_data): + return redact_error_information(error_information) + return error_information + + def _redact_standard_logging_object(payload: Mapping[str, object]) -> dict[str, object]: standard_logging_object: Final = copy.deepcopy(without_classifier_audit(payload)) redacted_str: Final = REDACTED_BY_LITELLM @@ -207,6 +308,16 @@ def _redact_standard_logging_object(payload: Mapping[str, object]) -> dict[str, else: # For other formats (empty dict, None, etc.), use simple text format standard_logging_object["response"] = {"text": redacted_str} + + if standard_logging_object.get("error_str"): + standard_logging_object["error_str"] = redacted_str + error_information: Final = standard_logging_object.get("error_information") + if isinstance(error_information, Mapping): + standard_logging_object["error_information"] = redact_error_information( + cast( # cast-ok: same TypedDict shape as the input mapping + StandardLoggingPayloadErrorInformation, error_information + ) + ) return standard_logging_object @@ -272,6 +383,8 @@ def perform_redaction(model_call_details: dict, result, redact_streaming_respons standard_logging_object: Final = model_call_details.get("standard_logging_object") if isinstance(standard_logging_object, Mapping): model_call_details["standard_logging_object"] = _redact_standard_logging_object(standard_logging_object) + if isinstance(model_call_details.get("traceback_exception"), str) and model_call_details["traceback_exception"]: + model_call_details["traceback_exception"] = REDACTED_BY_LITELLM redact_vertex_ai_metadata_from_litellm_params(model_call_details) # Redact streaming response @@ -352,13 +465,8 @@ def should_redact_message_logging(model_call_details: dict) -> bool: # User explicitly disabled redaction via header return False - possible_enable_headers: Final = [ - "litellm-enable-message-redaction", # old header. maintain backwards compatibility - "x-litellm-enable-message-redaction", # new header - ] - is_redaction_enabled_via_header = False - for header in possible_enable_headers: + for header in _MESSAGE_REDACTION_ENABLE_HEADERS: if bool(request_headers.get(header, False)): is_redaction_enabled_via_header = True break diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 5194f62cf78..c634858f8de 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -38,6 +38,7 @@ from litellm.integrations.otel.model.config import is_otel_v2_enabled from litellm.integrations.otel.runtime import phase_span, seed_request_identity from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value +from litellm.litellm_core_utils.redact_messages import request_opts_into_message_redaction from litellm.proxy._types import * from litellm.proxy.agent_endpoints.auth.agent_caller import agent_caller_from_headers from litellm.proxy.auth.auth_checks import ( @@ -3406,7 +3407,13 @@ async def user_api_key_auth( # Run the whole auth phase inside a live ``auth`` span so the DB lookups it # triggers (key/user/team object reads) nest under it instead of flattening # onto the server span. No-op when OTel V2 isn't active. - with phase_span(f"auth {route}"), spend_counter_batch_scope(_spend_counter_redis_cache()): + with ( + phase_span( + f"auth {route}", + redact_content=request_opts_into_message_redaction(_safe_get_request_headers(request), request_data), + ), + spend_counter_batch_scope(_spend_counter_redis_cache()), + ): try: user_api_key_auth_obj: Final = await _user_api_key_auth_builder( request=request, diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 6a2ec120060..7ba95fcc157 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -16,6 +16,7 @@ from litellm.litellm_core_utils.core_helpers import ( ) from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import guardrail_information_cost +from litellm.litellm_core_utils.redact_messages import maybe_redact_error_information from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_checks import ( get_key_object, @@ -172,6 +173,9 @@ class _ProxyDBLogger(CustomLogger): original_exception=original_exception, traceback_str=traceback_str, ) + _error_information = maybe_redact_error_information( + error_information=_error_information, request_data=request_data + ) if should_suppress_spend_log_tracebacks(): # Drop the traceback key entirely so the per-row Metadata pane in # the UI (which renders the JSON blob verbatim) doesn't show a diff --git a/tests/integration/observability/test_failure_redaction.py b/tests/integration/observability/test_failure_redaction.py new file mode 100644 index 00000000000..87660dfbd2a --- /dev/null +++ b/tests/integration/observability/test_failure_redaction.py @@ -0,0 +1,1110 @@ +import asyncio +import json +import threading +import time +import uuid +from collections.abc import Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import anthropic +import httpx +import openai +import pytest +import yaml +from integration._support.client import Gateway, JsonValue, eventually, gateway_from_environment, object_value +from integration._support.database import read_rows +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, Wire, wire_server + + +@dataclass(frozen=True, slots=True) +class Rig: + proxy: Gateway + provider: Wire + sink: Wire + batches: list[Request] + + def failure_events(self, model: str) -> tuple[dict[str, JsonValue], ...]: + self.batches.extend(self.sink.drain()) + return tuple( + object_value(event) + for batch in self.batches + for event in json.loads(batch.body) + if model in json.dumps(event) + ) + + +def _prompt_text(body: dict[str, JsonValue]) -> str: + messages: Final = body.get("messages") + if isinstance(messages, list) and messages: + last: Final = messages[-1] + if isinstance(last, dict): + content: Final = last.get("content") + if isinstance(content, str): + return content + if isinstance(content, list): + return "".join( + str(part["text"]) + for part in content + if isinstance(part, dict) and isinstance(part.get("text"), str) + ) + input_value: Final = body.get("input") + if isinstance(input_value, str): + return input_value + if isinstance(input_value, list): + return json.dumps(input_value) + return json.dumps(body)[:200] + + +def _success_reply(target: str, text: str) -> Reply: + if target.endswith("/messages"): + return Reply( + body=json.dumps( + { + "id": "msg_" + uuid.uuid4().hex, + "type": "message", + "role": "assistant", + "model": "claude-sonnet-5", + "content": [{"type": "text", "text": text}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 3, "output_tokens": 2}, + } + ).encode() + ) + if target.endswith("/responses"): + return Reply( + body=json.dumps( + { + "id": "resp_" + uuid.uuid4().hex, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "id": "msg_" + uuid.uuid4().hex, + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": text, "annotations": []}], + } + ], + "usage": {"input_tokens": 3, "output_tokens": 2, "total_tokens": 5}, + } + ).encode() + ) + return Reply( + body=json.dumps( + { + "id": "chatcmpl-" + uuid.uuid4().hex, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": text}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5}, + } + ).encode() + ) + + +def _provider(request: Request) -> Reply: + try: + body: Final = json.loads(request.body) + except json.JSONDecodeError: + return Reply(status=404, body=b"{}") + text: Final = _prompt_text(body) + if text.startswith("ok "): + return _success_reply(request.target, text) + status: Final = 401 if text.startswith("auth401 ") else 400 + error: Final = {"type": "invalid_request_error", "message": f"Unsupported content: {text}"} + if request.target.endswith("/messages"): + return Reply(status=status, body=json.dumps({"type": "error", "error": error}).encode()) + return Reply(status=status, body=json.dumps({"error": error}).encode()) + + +_SINK_OUTAGE: Final = threading.Event() +_SINK_SLOW: Final = threading.Event() + + +def _sink(request: Request) -> Reply: + if _SINK_OUTAGE.is_set(): + return Reply(status=503, body=b'{"error":"sink down"}') + if _SINK_SLOW.is_set(): + time.sleep(2) + return Reply() + + +@contextmanager +def _booted_rig( + root: Path, + provider: Wire, + sink: Wire, + settings: Mapping[str, JsonValue], +) -> Iterator[Rig]: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["general_settings"]["allow_client_side_credentials"] = True + config["litellm_settings"].update({"DEFAULT_FLUSH_INTERVAL_SECONDS": 1, **settings}) + path: Final = root / "failure_redaction.yaml" + path.write_text(yaml.safe_dump(config)) + with ( + gateway_from_environment() as gateway, + owned_proxy(gateway, root, {"GENERIC_LOGGER_ENDPOINT": sink.url}, config=path, workers=2) as proxy, + ): + yield Rig(proxy, provider, sink, []) # mutable-ok: sink drain consumes batches, later polls keep earlier ones + + +@pytest.fixture(scope="module") +def provider() -> Iterator[Wire]: + with wire_server(_provider) as wire: + yield wire + + +@pytest.fixture(scope="module") +def sink() -> Iterator[Wire]: + with wire_server(_sink) as wire: + yield wire + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory, provider: Wire, sink: Wire) -> Iterator[Rig]: + with _booted_rig( + tmp_path_factory.mktemp("failure_redaction"), + provider, + sink, + { + "callbacks": ["generic_api"], + "turn_off_message_logging": True, + "standard_logging_payload_excluded_fields": ["hidden_params"], + }, + ) as booted: + yield booted + + +@pytest.fixture(scope="module") +def rig_off(tmp_path_factory: pytest.TempPathFactory, provider: Wire, sink: Wire) -> Iterator[Rig]: + with _booted_rig( + tmp_path_factory.mktemp("failure_redaction_off"), provider, sink, {"callbacks": ["generic_api"]} + ) as booted: + yield booted + + +@pytest.fixture(scope="module") +def rig_failure_callback(tmp_path_factory: pytest.TempPathFactory, provider: Wire, sink: Wire) -> Iterator[Rig]: + with _booted_rig( + tmp_path_factory.mktemp("failure_redaction_failure_cb"), + provider, + sink, + {"failure_callback": ["generic_api"], "turn_off_message_logging": True}, + ) as booted: + yield booted + + +@pytest.fixture(scope="module") +def rig_success_callback(tmp_path_factory: pytest.TempPathFactory, provider: Wire, sink: Wire) -> Iterator[Rig]: + with _booted_rig( + tmp_path_factory.mktemp("failure_redaction_success_cb"), + provider, + sink, + {"success_callback": ["generic_api"], "turn_off_message_logging": True}, + ) as booted: + yield booted + + +def _secret_prompt() -> str: + return "confidential-prompt-" + uuid.uuid4().hex + + +def _single_failure_event(rig: Rig, model: str) -> dict[str, JsonValue]: + events: Final = eventually( + lambda: tuple(event for event in rig.failure_events(model) if event.get("status") == "failure"), + lambda values: len(values) >= 1, + seconds=30, + ) + assert len(events) == 1, events + return events[0] + + +def _spend_row(call_id: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT status, messages, response, proxy_server_request, metadata FROM "LiteLLM_SpendLogs" ' + "WHERE request_id=%s", + (call_id,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + return rows[0] + + +def _chat(rig: Rig, model: str, messages: list[JsonValue], key: str | None = None) -> httpx.Response: + return rig.proxy.request("POST", "/v1/chat/completions", {"model": model, "messages": messages}, key=key) + + +def _error_information(event: dict[str, JsonValue]) -> dict[str, JsonValue]: + return object_value(event["error_information"]) + + +def test_provider_error_echoing_the_prompt_is_redacted_in_callbacks_and_spend_logs(rig: Rig) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model( + api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key", num_retries=0 + ) + response: Final = _chat(rig, model, [{"role": "user", "content": secret}]) + assert response.status_code == 400, response.text + assert any(secret.encode() in request.body for request in rig.provider.drain()) + event: Final = _single_failure_event(rig, model) + assert secret not in json.dumps(event), json.dumps(event) + error_information: Final = _error_information(event) + assert error_information["error_class"] == "BadRequestError", error_information + assert error_information["error_code"] == "400", error_information + assert error_information["llm_provider"] == "openai", error_information + assert "hidden_params" not in event, sorted(event) + row: Final = _spend_row(response.headers["x-litellm-call-id"]) + assert row["status"] == "failure", row + assert secret not in json.dumps(row, default=str), row + persisted: Final = object_value( + json.loads(row["metadata"]) if isinstance(row["metadata"], str) else row["metadata"] + ) + assert object_value(persisted["error_information"])["error_class"] == "BadRequestError", persisted + + +def test_transformation_error_quoting_the_prompt_is_redacted_in_callbacks(rig: Rig) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model( + model="anthropic/claude-sonnet-5", api_base=rig.provider.url, api_key="synthetic-anthropic-key" + ) + response: Final = _chat( + rig, + model, + [ + {"role": "user", "content": secret}, + { + "role": "assistant", + "content": None, + "tool_calls": [{"id": 12345, "type": "function", "function": {"name": None, "arguments": "{}"}}], + }, + {"role": "tool", "tool_call_id": 12345, "content": secret}, + ], + ) + assert response.status_code == 400, response.text + assert rig.provider.drain() == () + event: Final = _single_failure_event(rig, model) + assert secret not in json.dumps(event), json.dumps(event) + error_information: Final = _error_information(event) + assert error_information["error_code"] == "400", error_information + assert error_information["error_class"], error_information + + +def test_proxy_only_rejection_does_not_leak_the_prompt_to_callbacks(rig: Rig) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + allowed: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + denied: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[allowed]) + rig.provider.drain() + response: Final = _chat(rig, denied, [{"role": "user", "content": secret}], key=key) + assert response.status_code in (401, 403), response.text + assert rig.provider.drain() == () + event: Final = _single_failure_event(rig, denied) + assert secret not in json.dumps(event), json.dumps(event) + assert "hidden_params" not in event, sorted(event) + assert _error_information(event)["error_code"] == str(response.status_code), event + + +def _openai(rig: Rig, key: str | None = None) -> openai.OpenAI: + return openai.OpenAI( + base_url=str(rig.proxy.client.base_url) + "/v1", + api_key=key or rig.proxy.key, + max_retries=0, + ) + + +def _async_openai(rig: Rig, key: str | None = None) -> openai.AsyncOpenAI: + return openai.AsyncOpenAI( + base_url=str(rig.proxy.client.base_url) + "/v1", + api_key=key or rig.proxy.key, + max_retries=0, + ) + + +def _anthropic(rig: Rig, key: str | None = None) -> anthropic.Anthropic: + return anthropic.Anthropic(base_url=str(rig.proxy.client.base_url), api_key=key or rig.proxy.key, max_retries=0) + + +def _async_anthropic(rig: Rig, key: str | None = None) -> anthropic.AsyncAnthropic: + return anthropic.AsyncAnthropic( + base_url=str(rig.proxy.client.base_url), api_key=key or rig.proxy.key, max_retries=0 + ) + + +def _call_id(response: httpx.Response) -> str: + call_id: Final = response.headers.get("x-litellm-call-id") + assert call_id, dict(response.headers) + return call_id + + +def _error_call_id(error: openai.APIStatusError | anthropic.APIStatusError) -> str: + call_id: Final = error.response.headers.get("x-litellm-call-id") + assert call_id, dict(error.response.headers) + return call_id + + +def _assert_failure_redacted(rig: Rig, model: str, secret: str, call_id: str) -> dict[str, JsonValue]: + event: Final = _single_failure_event(rig, model) + assert secret not in json.dumps(event), json.dumps(event) + error_information: Final = _error_information(event) + assert error_information["error_class"], error_information + assert error_information["llm_provider"], error_information + row: Final = _spend_row(call_id) + assert row["status"] == "failure", row + assert secret not in json.dumps(row, default=str), row + return event + + +def _assert_failure_raw(rig: Rig, model: str, secret: str, call_id: str | None = None) -> dict[str, JsonValue]: + event: Final = _single_failure_event(rig, model) + assert secret in json.dumps(event), json.dumps(event) + if call_id is not None: + row: Final = _spend_row(call_id) + assert secret in json.dumps(row, default=str), row + return event + + +def _denied_body(path: str, model: str, secret: str) -> dict[str, JsonValue]: + if path == "/v1/messages": + return {"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": secret}]} + if path == "/v1/responses": + return {"model": model, "input": secret} + return {"model": model, "messages": [{"role": "user", "content": secret}]} + + +# --- A. Endpoint x stream x client, provider 400 echoing the prompt ------------------- + +_SDK_CLIENTS: Final = ("openai_sync", "openai_async", "httpx") + + +def _chat_call(rig: Rig, client: str, model: str, secret: str, stream: bool) -> str: + messages: Final = [{"role": "user", "content": secret}] + if client == "httpx": + response: Final = rig.proxy.request( + "POST", "/v1/chat/completions", {"model": model, "messages": messages, "stream": stream} + ) + response.read() + assert response.status_code == 400, response.text + return _call_id(response) + if client == "openai_sync": + with pytest.raises(openai.BadRequestError) as caught: + _openai(rig).chat.completions.create(model=model, messages=messages, stream=stream) + assert caught.value.status_code == 400, caught.value + return _error_call_id(caught.value) + + async def fire() -> str: + with pytest.raises(openai.BadRequestError) as caught: + await _async_openai(rig).chat.completions.create(model=model, messages=messages, stream=stream) + return _error_call_id(caught.value) + + return asyncio.run(fire()) + + +@pytest.mark.parametrize("client", _SDK_CLIENTS, ids=lambda value: value) +def test_a1_chat_provider_error_redacted(rig: Rig, client: str) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + call_id: Final = _chat_call(rig, client, model, secret, stream=False) + assert any(secret.encode() in request.body for request in rig.provider.drain()) + event: Final = _assert_failure_redacted(rig, model, secret, call_id) + assert object_value(event["error_information"])["error_class"] == "BadRequestError", event + + +@pytest.mark.parametrize("client", _SDK_CLIENTS, ids=lambda value: value) +def test_a2_chat_stream_provider_error_redacted(rig: Rig, client: str) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + call_id: Final = _chat_call(rig, client, model, secret, stream=True) + assert any(secret.encode() in request.body for request in rig.provider.drain()) + _assert_failure_redacted(rig, model, secret, call_id) + + +def _consume_messages_stream(client: anthropic.Anthropic, model: str, messages: list[dict[str, str]]) -> None: + with client.messages.stream(model=model, max_tokens=16, messages=messages) as events: + for _ in events: + pass + + +async def _consume_messages_stream_async( + client: anthropic.AsyncAnthropic, model: str, messages: list[dict[str, str]] +) -> None: + async with client.messages.stream(model=model, max_tokens=16, messages=messages) as events: + async for _ in events: + pass + + +def _messages_call(rig: Rig, client: str, model: str, secret: str, stream: bool) -> str: + messages: Final = [{"role": "user", "content": secret}] + if client == "httpx": + response: Final = rig.proxy.request( + "POST", + "/v1/messages", + {"model": model, "max_tokens": 16, "messages": messages, "stream": stream}, + ) + response.read() + assert response.status_code == 400, response.text + return _call_id(response) + if client == "openai_sync": + anthropic_client: Final = _anthropic(rig) + if stream: + with pytest.raises(anthropic.BadRequestError) as caught: + _consume_messages_stream(anthropic_client, model, messages) + return _error_call_id(caught.value) + with pytest.raises(anthropic.BadRequestError) as caught: + anthropic_client.messages.create(model=model, max_tokens=16, messages=messages) + return _error_call_id(caught.value) + + async def fire() -> str: + client_async: Final = _async_anthropic(rig) + if stream: + with pytest.raises(anthropic.BadRequestError) as caught: + await _consume_messages_stream_async(client_async, model, messages) + return _error_call_id(caught.value) + with pytest.raises(anthropic.BadRequestError) as caught: + await client_async.messages.create(model=model, max_tokens=16, messages=messages) + return _error_call_id(caught.value) + + return asyncio.run(fire()) + + +@pytest.mark.parametrize("client", _SDK_CLIENTS, ids=lambda value: value) +def test_a3_messages_provider_error_redacted(rig: Rig, client: str) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model( + model="anthropic/claude-sonnet-5", api_base=rig.provider.url, api_key="synthetic-anthropic-key" + ) + call_id: Final = _messages_call(rig, client, model, secret, stream=False) + assert any(secret.encode() in request.body for request in rig.provider.drain()) + _assert_failure_redacted(rig, model, secret, call_id) + + +@pytest.mark.parametrize("client", _SDK_CLIENTS, ids=lambda value: value) +def test_a4_messages_stream_provider_error_redacted(rig: Rig, client: str) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model( + model="anthropic/claude-sonnet-5", api_base=rig.provider.url, api_key="synthetic-anthropic-key" + ) + call_id: Final = _messages_call(rig, client, model, secret, stream=True) + assert any(secret.encode() in request.body for request in rig.provider.drain()) + _assert_failure_redacted(rig, model, secret, call_id) + + +def _responses_call(rig: Rig, client: str, model: str, secret: str, stream: bool) -> str: + if client == "httpx": + response: Final = rig.proxy.request( + "POST", "/v1/responses", {"model": model, "input": secret, "stream": stream} + ) + assert response.status_code == 400, response.text + return _call_id(response) + if client == "openai_sync": + with pytest.raises(openai.BadRequestError) as caught: + _openai(rig).responses.create(model=model, input=secret, stream=stream) + return _error_call_id(caught.value) + + async def fire() -> str: + with pytest.raises(openai.BadRequestError) as caught: + await _async_openai(rig).responses.create(model=model, input=secret, stream=stream) + return _error_call_id(caught.value) + + return asyncio.run(fire()) + + +@pytest.mark.parametrize("client", _SDK_CLIENTS, ids=lambda value: value) +def test_a5_responses_provider_error_redacted(rig: Rig, client: str) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + call_id: Final = _responses_call(rig, client, model, secret, stream=False) + assert any(secret.encode() in request.body for request in rig.provider.drain()) + _assert_failure_redacted(rig, model, secret, call_id) + + +@pytest.mark.parametrize("client", _SDK_CLIENTS, ids=lambda value: value) +def test_a6_responses_stream_provider_error_redacted(rig: Rig, client: str) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + call_id: Final = _responses_call(rig, client, model, secret, stream=True) + assert any(secret.encode() in request.body for request in rig.provider.drain()) + _assert_failure_redacted(rig, model, secret, call_id) + + +def test_a7_messages_transformation_error_redacted(rig: Rig) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model( + model="anthropic/claude-sonnet-5", api_base=rig.provider.url, api_key="synthetic-anthropic-key" + ) + response: Final = rig.proxy.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 16, + "messages": [ + {"role": "user", "content": secret}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": 12345, "type": "function", "function": {"name": None, "arguments": "{}"}} + ], + }, + {"role": "tool", "tool_call_id": 12345, "content": secret}, + ], + }, + ) + assert response.status_code == 400, response.text + _assert_failure_redacted(rig, model, secret, _call_id(response)) + + +@pytest.mark.parametrize( + "path", ("/v1/chat/completions", "/v1/messages", "/v1/responses"), ids=lambda v: v.split("/")[-1] +) +def test_a8_proxy_only_rejection_redacted(rig: Rig, path: str) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + allowed: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + denied: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[allowed]) + rig.provider.drain() + response: Final = rig.proxy.request("POST", path, _denied_body(path, denied, secret), key=key) + assert response.status_code in (401, 403), response.text + assert rig.provider.drain() == () + event: Final = _single_failure_event(rig, denied) + assert secret not in json.dumps(event), json.dumps(event) + assert _error_information(event)["error_code"] == str(response.status_code), event + + +def test_a9_unknown_model_failure_redacted(rig: Rig) -> None: + secret: Final = _secret_prompt() + model: Final = "unknown-" + uuid.uuid4().hex + rig.provider.drain() + response: Final = _chat(rig, model, [{"role": "user", "content": secret}]) + assert response.status_code == 400, response.text + assert rig.provider.drain() == () + event: Final = eventually(lambda: rig.failure_events(model), lambda values: len(values) >= 1, seconds=20)[0] + assert secret not in json.dumps(event), json.dumps(event) + row: Final = _spend_row(_call_id(response)) + assert secret not in json.dumps(row, default=str), row + + +# --- B. Redaction source modes (YAML global off unless noted) ------------------------- + + +@pytest.mark.parametrize( + ("headers", "body"), + ( + pytest.param(None, {}, id="b1_no_signal"), + pytest.param({"x-litellm-enable-message-redaction": "true"}, {}, id="b2_enable_header"), + pytest.param({"litellm-enable-message-redaction": "true"}, {}, id="b3_legacy_enable_header"), + pytest.param(None, {"turn_off_message_logging": True}, id="b4_body_param"), + pytest.param(None, {"metadata": {"turn_off_message_logging": True}}, id="b5_metadata_param"), + pytest.param(None, {"litellm_metadata": {"turn_off_message_logging": True}}, id="b5b_litellm_metadata"), + ), +) +def test_b_request_level_opt_in_redacts( + rig_off: Rig, headers: Mapping[str, str] | None, body: Mapping[str, JsonValue] +) -> None: + secret: Final = _secret_prompt() + with rig_off.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig_off.provider.url + "/v1", api_key="synthetic-provider-key") + response: Final = rig_off.proxy.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": secret}], **body}, + headers=headers, + ) + assert response.status_code == 400, response.text + assert any(secret.encode() in request.body for request in rig_off.provider.drain()) + if headers is None and not body: + _assert_failure_raw(rig_off, model, secret, _call_id(response)) + else: + _assert_failure_redacted(rig_off, model, secret, _call_id(response)) + + +def test_b8_global_on_permitted_key_opt_out_keeps_raw(rig: Rig) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model], metadata={"allow_client_message_redaction_opt_out": True}) + response: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": secret}]}, + key=key, + headers={"litellm-disable-message-redaction": "true"}, + ) + assert response.status_code == 400, response.text + _assert_failure_raw(rig, model, secret, _call_id(response)) + + +def test_b9_global_on_disable_header_without_permission_stays_redacted(rig: Rig) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + response: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": secret}]}, + key=key, + headers={"litellm-disable-message-redaction": "true"}, + ) + assert response.status_code == 400, response.text + _assert_failure_redacted(rig, model, secret, _call_id(response)) + + +def test_b10_global_on_team_permission_opt_out_keeps_raw(rig: Rig) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + team: Final = scenario.team(metadata={"allow_client_message_redaction_opt_out": True}) + key: Final = scenario.key(models=[model], team_id=team) + response: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": secret}]}, + key=key, + headers={"litellm-disable-message-redaction": "true"}, + ) + assert response.status_code == 400, response.text + _assert_failure_raw(rig, model, secret, _call_id(response)) + + +def _logging_callback_vars(flag: bool) -> dict[str, JsonValue]: + return {"logging": [{"callback_name": "generic_api", "callback_vars": {"turn_off_message_logging": flag}}]} + + +def _event_for_call(rig: Rig, model: str, call_id: str) -> dict[str, JsonValue]: + return eventually( + lambda: tuple(event for event in rig.failure_events(model) if event.get("litellm_call_id") == call_id), + lambda values: len(values) == 1, + seconds=90, + )[0] + + +def _assert_callback_vars_decision(rig: Rig, model: str, flag: bool, key: str) -> None: + ok_secret: Final = f"ok {_secret_prompt()}" + succeeded: Final = rig.proxy.request( + "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": ok_secret}]}, key=key + ) + assert succeeded.status_code == 200, succeeded.text + success_event: Final = _event_for_call(rig, model, _call_id(succeeded)) + assert (ok_secret not in json.dumps(success_event)) == flag, json.dumps(success_event) + fail_secret: Final = _secret_prompt() + failed: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": fail_secret}]}, + key=key, + ) + assert failed.status_code == 400, failed.text + if flag: + _assert_failure_redacted(rig, model, fail_secret, _call_id(failed)) + else: + _assert_failure_raw(rig, model, fail_secret, _call_id(failed)) + + +@pytest.mark.timeout(280) +@pytest.mark.parametrize("global_flag", ["on", "off"]) +@pytest.mark.parametrize("flag", [True, False], ids=["vars_true", "vars_false"]) +def test_b6_key_logging_callback_vars_drive_the_failure_decision( + request: pytest.FixtureRequest, global_flag: str, flag: bool +) -> None: + rig: Final = request.getfixturevalue("rig" if global_flag == "on" else "rig_off") + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model], metadata=_logging_callback_vars(flag)) + _assert_callback_vars_decision(rig, model, flag, key) + + +@pytest.mark.timeout(280) +@pytest.mark.parametrize("global_flag", ["on", "off"]) +@pytest.mark.parametrize("flag", [True, False], ids=["vars_true", "vars_false"]) +def test_b7_team_logging_callback_vars_drive_the_failure_decision( + request: pytest.FixtureRequest, global_flag: str, flag: bool +) -> None: + rig: Final = request.getfixturevalue("rig" if global_flag == "on" else "rig_off") + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + team: Final = scenario.team(metadata=_logging_callback_vars(flag)) + key: Final = scenario.key(models=[model], team_id=team) + _assert_callback_vars_decision(rig, model, flag, key) + + +# --- C. Callback registration modes (YAML global on) ---------------------------------- + + +def test_c2_failure_callback_registration_redacts(rig_failure_callback: Rig) -> None: + rig: Final = rig_failure_callback + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + response: Final = _chat(rig, model, [{"role": "user", "content": secret}]) + assert response.status_code == 400, response.text + _assert_failure_redacted(rig, model, secret, _call_id(response)) + + +def test_c3_success_callback_only_emits_no_failure_event(rig_success_callback: Rig) -> None: + rig: Final = rig_success_callback + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + response: Final = _chat(rig, model, [{"role": "user", "content": secret}]) + assert response.status_code == 400, response.text + events: Final = eventually( + lambda: rig.failure_events(model), + lambda values: len(values) >= 1, + seconds=8, + return_last_on_timeout=True, + ) + assert events == (), events + + +def test_c4_excluded_fields_stripped_on_failure(rig: Rig) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + response: Final = _chat(rig, model, [{"role": "user", "content": secret}]) + assert response.status_code == 400, response.text + event: Final = _assert_failure_redacted(rig, model, secret, _call_id(response)) + assert "hidden_params" not in event, sorted(event) + + +# --- D. Cache-hit twin ----------------------------------------------------------------- + + +@pytest.mark.timeout(280) +def test_d1_cache_hit_success_then_failure_redacted(rig: Rig) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + rig.provider.drain() + ok_text: Final = "ok " + uuid.uuid4().hex + first: Final = _chat(rig, model, [{"role": "user", "content": ok_text}]) + assert first.status_code == 200, first.text + second: Final = _chat(rig, model, [{"role": "user", "content": ok_text}]) + assert second.status_code == 200, second.text + assert second.json()["id"] == first.json()["id"], "second call did not hit the response cache" + response: Final = _chat(rig, model, [{"role": "user", "content": secret}]) + assert response.status_code == 400, response.text + _assert_failure_redacted(rig, model, secret, _call_id(response)) + + +# --- E. Sad paths ----------------------------------------------------------------------- + +_HOSTILE_HEADERS: Final = ("", "0", "false", "1", "a,b", "h" * 5000) + + +@pytest.mark.parametrize("value", _HOSTILE_HEADERS, ids=lambda v: v[:8] or "empty") +@pytest.mark.timeout(280) +def test_e1_hostile_enable_header_unauthenticated_and_authenticated(rig_off: Rig, value: str) -> None: + secret: Final = _secret_prompt() + with rig_off.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig_off.provider.url + "/v1", api_key="synthetic-provider-key") + rejected: Final = rig_off.proxy.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": secret}]}, + key="sk-bogus", + headers={"x-litellm-enable-message-redaction": value}, + ) + assert rejected.status_code == 401, rejected.text + call_ids: list[str] = [] # mutable-ok: collect per-request call ids across the loop + for _ in range(2): + response: Final = rig_off.proxy.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": secret}], "cache": {"no-cache": True}}, + headers={"x-litellm-enable-message-redaction": value}, + ) + assert response.status_code == 400, response.text + call_ids.append(_call_id(response)) + rig_off.provider.drain() + for call_id in call_ids: + events: Final = eventually( + lambda cid=call_id: tuple( + event + for event in rig_off.failure_events(model) + if event.get("litellm_call_id") == cid + and object_value(event.get("error_information")).get("error_class") != "KeyNotFoundError" + ), + lambda values: len(values) == 1, + seconds=90, + ) + if value: + assert secret not in json.dumps(events[0]), (value, json.dumps(events[0])[:400]) + else: + assert secret in json.dumps(events[0]), (value, json.dumps(events[0])[:400]) + + +@pytest.mark.parametrize("value", (1, [True], "", "v" * 5000, None), ids=lambda v: type(v).__name__) +def test_e2_odd_turn_off_message_logging_values_never_500(rig_off: Rig, value: JsonValue) -> None: + secret: Final = _secret_prompt() + with rig_off.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig_off.provider.url + "/v1", api_key="synthetic-provider-key") + response: Final = rig_off.proxy.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": secret}], + "turn_off_message_logging": value, + }, + ) + assert response.status_code == 400, response.text + event: Final = _single_failure_event(rig_off, model) + assert event["status"] == "failure", event + + +def test_e3_sink_rejections_keep_proxy_serving_and_later_events_still_redacted(rig: Rig) -> None: + secret: Final = _secret_prompt() + _SINK_OUTAGE.set() + try: + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + during: Final = _chat(rig, model, [{"role": "user", "content": "down " + secret}]) + assert during.status_code == 400, during.text + ok: Final = _chat(rig, model, [{"role": "user", "content": "ok alive " + uuid.uuid4().hex}]) + assert ok.status_code == 200, ok.text + rig.sink.drain() + finally: + _SINK_OUTAGE.clear() + with rig.proxy.scenario() as scenario: + model_two: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + after: Final = _chat(rig, model_two, [{"role": "user", "content": "after " + secret}]) + assert after.status_code == 400, after.text + event: Final = _single_failure_event(rig, model_two) + assert secret not in json.dumps(event), json.dumps(event) + + +def test_e4_header_only_unknown_model_spend_row_never_carries_marker(rig_off: Rig) -> None: + secret: Final = _secret_prompt() + response: Final = rig_off.proxy.request( + "POST", + "/v1/chat/completions", + {"model": "unknown-" + uuid.uuid4().hex, "messages": [{"role": "user", "content": secret}]}, + headers={"x-litellm-enable-message-redaction": "true"}, + ) + assert response.status_code == 400, response.text + row: Final = _spend_row(_call_id(response)) + assert secret not in json.dumps(row, default=str), row + + +@pytest.mark.parametrize("callback_fixture", ("rig", "rig_failure_callback"), ids=("callbacks", "failure_callback")) +@pytest.mark.parametrize("text", ("auth401 ", ""), ids=("provider_401", "provider_400")) +def test_e5_provider_401_and_400_echo_redacted( + request: pytest.FixtureRequest, callback_fixture: str, text: str +) -> None: + rig: Final = request.getfixturevalue(callback_fixture) + secret: Final = text + _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + response: Final = _chat(rig, model, [{"role": "user", "content": secret}]) + assert response.status_code in (400, 401), response.text + _assert_failure_redacted(rig, model, secret, _call_id(response)) + + +@pytest.mark.timeout(280) +def test_e6_malformed_callback_config_fails_boot_identically( + tmp_path_factory: pytest.TempPathFactory, provider: Wire, sink: Wire +) -> None: + with pytest.raises(AssertionError, match="exited before readiness"): + with _booted_rig( + tmp_path_factory.mktemp("failure_redaction_bad_cb"), + provider, + sink, + { + "callbacks": ["generic_api", "not_a_callback"], + "failure_callback": None, + "turn_off_message_logging": True, + }, + ): + raise AssertionError("malformed callback config must not boot") + + +def test_e7_spend_logs_detail_endpoint_redacted(rig: Rig) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + response: Final = _chat(rig, model, [{"role": "user", "content": secret}]) + assert response.status_code == 400, response.text + call_id: Final = _call_id(response) + _assert_failure_redacted(rig, model, secret, call_id) + detail: Final = eventually( + lambda: rig.proxy.get(f"/spend/logs/ui/{call_id}"), + lambda value: bool(value), + seconds=30, + ) + assert secret not in json.dumps(detail), json.dumps(detail)[:2000] + + +# --- F. Edge ---------------------------------------------------------------------------- + +_F1_SLOTS: Final = ("top", "metadata", "litellm_metadata") + + +def _redaction_body(model: str, secret: str, slot: str, value: JsonValue) -> dict[str, JsonValue]: + body: Final[dict[str, JsonValue]] = { + "model": model, + "messages": [{"role": "user", "content": secret}], + } + if slot == "top": + if value != "MISSING": + body["turn_off_message_logging"] = value + else: + body[slot] = {} if value == "MISSING" else {"turn_off_message_logging": value} + return body + + +@pytest.mark.parametrize("slot", _F1_SLOTS) +@pytest.mark.parametrize("value", ("MISSING", None, ""), ids=("missing", "null", "empty")) +def test_f1_turn_off_edge_values_global_on(rig: Rig, slot: str, value: JsonValue) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + response: Final = rig.proxy.request("POST", "/v1/chat/completions", _redaction_body(model, secret, slot, value)) + assert response.status_code == 400, response.text + _assert_failure_redacted(rig, model, secret, _call_id(response)) + + +@pytest.mark.parametrize("slot", _F1_SLOTS) +@pytest.mark.parametrize("value", ("MISSING", None, ""), ids=("missing", "null", "empty")) +def test_f1_turn_off_edge_values_global_off(rig_off: Rig, slot: str, value: JsonValue) -> None: + secret: Final = _secret_prompt() + with rig_off.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig_off.provider.url + "/v1", api_key="synthetic-provider-key") + response: Final = rig_off.proxy.request( + "POST", "/v1/chat/completions", _redaction_body(model, secret, slot, value) + ) + assert response.status_code == 400, response.text + _assert_failure_raw(rig_off, model, secret, _call_id(response)) + + +def test_f2_permitted_opt_out_top_level_false_beats_metadata_true(rig: Rig) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model], metadata={"allow_client_message_redaction_opt_out": True}) + response: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": secret}], + "turn_off_message_logging": False, + "metadata": {"turn_off_message_logging": True}, + }, + key=key, + ) + assert response.status_code == 400, response.text + _assert_failure_raw(rig, model, secret, _call_id(response)) + + +def test_f3_team_opt_out_permission_flip_takes_effect(rig: Rig) -> None: + first_secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + team: Final = scenario.team(metadata={"allow_client_message_redaction_opt_out": True}) + key: Final = scenario.key(models=[model], team_id=team) + allowed: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": first_secret}]}, + key=key, + headers={"litellm-disable-message-redaction": "true"}, + ) + assert allowed.status_code == 400, allowed.text + _assert_failure_raw(rig, model, first_secret, _call_id(allowed)) + rig.proxy.post("/team/update", {"team_id": team, "metadata": {}}) + + def now_redacted() -> tuple[bool, ...]: + secret: Final = _secret_prompt() + fired: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": secret}]}, + key=key, + headers={"litellm-disable-message-redaction": "true"}, + ) + assert fired.status_code == 400, fired.text + events: Final = eventually( + lambda: tuple( + event for event in rig.failure_events(model) if event.get("litellm_call_id") == _call_id(fired) + ), + lambda values: len(values) == 1, + seconds=20, + ) + return (secret not in json.dumps(events[0]),) + + converged: Final = eventually(now_redacted, lambda values: values[0], seconds=120) + assert converged[0] + + +@pytest.mark.timeout(280) +def test_f3_key_logging_callback_vars_flip_takes_effect(rig: Rig) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model], metadata=_logging_callback_vars(True)) + denied: Final = rig.proxy.request( + "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": secret}]}, key=key + ) + assert denied.status_code == 400, denied.text + _assert_failure_redacted(rig, model, secret, _call_id(denied)) + rig.proxy.post("/key/update", {"key": key, "metadata": _logging_callback_vars(False)}) + + def now_raw() -> tuple[bool, ...]: + probe: Final = _secret_prompt() + fired: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": probe}]}, + key=key, + ) + assert fired.status_code == 400, fired.text + events: Final = eventually( + lambda: tuple( + event for event in rig.failure_events(model) if event.get("litellm_call_id") == _call_id(fired) + ), + lambda values: len(values) == 1, + seconds=20, + ) + return (probe in json.dumps(events[0]),) + + converged: Final = eventually(now_raw, lambda values: values[0], seconds=120) + assert converged[0] + + +def test_f4_five_identical_failures_log_once_each(rig: Rig) -> None: + secret: Final = _secret_prompt() + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + call_ids: Final = tuple(_call_id(_chat(rig, model, [{"role": "user", "content": secret}])) for _ in range(5)) + assert len(set(call_ids)) == 5 + for call_id in call_ids: + eventually( + lambda cid=call_id: tuple( + event for event in rig.failure_events(model) if event.get("litellm_call_id") == cid + ), + lambda values: len(values) == 1, + seconds=30, + ) + assert _spend_row(call_id)["status"] == "failure", call_id diff --git a/tests/integration/observability/test_failure_redaction_chaos.py b/tests/integration/observability/test_failure_redaction_chaos.py new file mode 100644 index 00000000000..d49f23ac7c6 --- /dev/null +++ b/tests/integration/observability/test_failure_redaction_chaos.py @@ -0,0 +1,322 @@ +import json +import signal +import threading +import time +import uuid +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import httpx +import psutil +import pytest +import yaml +from integration._support.client import Gateway, JsonValue, eventually, gateway_from_environment, object_value +from integration._support.process import OwnedProxy, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server + +BURST: Final = 30 + + +def _provider(request: Request) -> Reply: + try: + body: Final = json.loads(request.body) + except json.JSONDecodeError: + return Reply(status=404, body=b"{}") + text: Final = ( + body["messages"][-1]["content"] if "messages" in body else str(body.get("input", json.dumps(body)[:200])) + ) + return Reply( + status=400, + body=json.dumps( + {"error": {"type": "invalid_request_error", "message": f"Unsupported content: {text}"}} + ).encode(), + ) + + +@dataclass(frozen=True, slots=True) +class Rig: + proxy: Gateway + process: OwnedProxy + provider: Wire + sink: Wire + model: str + outage: threading.Event + slow: threading.Event + batches: list[Request] + + def failure_events(self) -> tuple[dict[str, JsonValue], ...]: + self.batches.extend(self.sink.drain()) # mutable-ok: drain consumes, polls keep earlier batches + return tuple( + object_value(event) + for batch in self.batches + for event in json.loads(batch.body) + if self.model in json.dumps(event) + ) + + +def _bodies(model: str, marker: str) -> tuple[tuple[str, dict[str, JsonValue]], ...]: + chat: Final = tuple( + ( + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": f"burst {marker} {index}"}], + **({"stream": True} if index % 2 else {}), + }, + ) + for index in range(BURST // 3 * 2) + ) + messages: Final = tuple( + ( + "/v1/messages", + { + "model": model, + "max_tokens": 16, + "messages": [{"role": "user", "content": f"burst {marker} m{index}"}], + }, + ) + for index in range(BURST // 6) + ) + responses: Final = tuple( + ("/v1/responses", {"model": model, "input": f"burst {marker} r{index}"}) + for index in range(BURST - len(chat) - len(messages)) + ) + return chat + messages + responses + + +def _json_body_ok(body: bytes) -> bool: + try: + json.loads(body) + except json.JSONDecodeError: + return False + return True + + +def _fire(rig: Rig, bodies: tuple[tuple[str, dict[str, JsonValue]], ...]) -> tuple[tuple[int, str | None], ...]: + def call(item: tuple[str, dict[str, JsonValue]]) -> tuple[int, str | None]: + try: + response: Final = rig.proxy.request("POST", item[0], item[1]) + response.read() + return response.status_code, response.headers.get("x-litellm-call-id") + except httpx.HTTPError: + return -1, None + + with ThreadPoolExecutor(max_workers=8) as pool: + return tuple(pool.map(call, bodies)) + + +@pytest.mark.timeout(280) +def test_g1_sink_outage_mid_burst_lands_each_call_id_once_redacted( + tmp_path: Path, +) -> None: + marker: Final = uuid.uuid4().hex + outage: Final = threading.Event() + + def sink(request: Request) -> Reply: + if outage.is_set(): + return Reply(status=503, body=b'{"error":"sink down"}') + return Reply() + + with wire_server(_provider) as provider, wire_server(sink) as endpoint: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"].update( + {"callbacks": ["generic_api"], "DEFAULT_FLUSH_INTERVAL_SECONDS": 1, "turn_off_message_logging": True} + ) + path: Final = tmp_path / "chaos.yaml" + path.write_text(yaml.safe_dump(config)) + with ( + gateway_from_environment() as gateway, + owned_proxy_process( + gateway, tmp_path, {"GENERIC_LOGGER_ENDPOINT": endpoint.url}, config=path, workers=2 + ) as owned, + owned.gateway.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + rig: Final = Rig(owned.gateway, owned, provider, endpoint, model, outage, threading.Event(), []) + bodies: Final = _bodies(model, marker) + outage.set() + first_half: Final = _fire(rig, bodies[: BURST // 2]) + outage.clear() + second_half: Final = _fire(rig, bodies[BURST // 2 :]) + outcomes: Final = first_half + second_half + answered: Final = tuple(outcome for outcome in outcomes if outcome[0] >= 0) + assert all(status == 400 for status, _ in answered), outcomes + events: Final = eventually(rig.failure_events, lambda values: len(values) >= len(second_half), seconds=70) + seen: Final = tuple(str(event.get("litellm_call_id")) for event in events) + assert len(seen) == len(set(seen)), ("duplicate failure events", seen) + for event in events: + assert marker not in json.dumps(event), json.dumps(event)[:400] + + +@pytest.mark.timeout(280) +def test_g2_slow_sink_no_deadlock_no_duplicates(tmp_path: Path) -> None: + marker: Final = uuid.uuid4().hex + slow: Final = threading.Event() + + def sink(request: Request) -> Reply: + if slow.is_set(): + time.sleep(1) + return Reply() + + with wire_server(_provider) as provider, wire_server(sink) as endpoint: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"].update( + {"callbacks": ["generic_api"], "DEFAULT_FLUSH_INTERVAL_SECONDS": 1, "turn_off_message_logging": True} + ) + path: Final = tmp_path / "chaos_slow.yaml" + path.write_text(yaml.safe_dump(config)) + with ( + gateway_from_environment() as gateway, + owned_proxy_process( + gateway, tmp_path, {"GENERIC_LOGGER_ENDPOINT": endpoint.url}, config=path, workers=2 + ) as owned, + owned.gateway.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + rig: Final = Rig(owned.gateway, owned, provider, endpoint, model, threading.Event(), slow, []) + slow.set() + bodies: Final = _bodies(model, marker)[:6] + outcomes: Final = _fire(rig, bodies) + assert all(status == 400 for status, _ in outcomes), outcomes + call_ids: Final = tuple(cid for _, cid in outcomes if cid) + events: Final = eventually(rig.failure_events, lambda values: len(values) >= len(bodies), seconds=120) + landed: Final = tuple(str(event.get("litellm_call_id")) for event in events) + assert len(landed) == len(set(landed)), ("duplicate failure events", landed) + for event in events: + assert marker not in json.dumps(event), json.dumps(event)[:400] + assert set(call_ids) <= set(landed), (call_ids, landed) + + +@pytest.mark.timeout(280) +def test_g3_proxy_restart_mid_burst(tmp_path: Path) -> None: + marker: Final = uuid.uuid4().hex + with wire_server(_provider) as provider, wire_server(lambda _: Reply()) as endpoint: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"].update( + {"callbacks": ["generic_api"], "DEFAULT_FLUSH_INTERVAL_SECONDS": 1, "turn_off_message_logging": True} + ) + path: Final = tmp_path / "chaos_restart.yaml" + path.write_text(yaml.safe_dump(config)) + overrides: Final = {"GENERIC_LOGGER_ENDPOINT": endpoint.url} + with gateway_from_environment() as gateway: + bodies: Final = _bodies("restart-model", marker) + with owned_proxy_process(gateway, tmp_path, overrides, config=path, workers=2) as owned_one: + owned_one.gateway.post( + "/model/new", + { + "model_name": "restart-model", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": provider.url + "/v1", + "api_key": "synthetic-provider-key", + }, + }, + ) + first: Final = _fire( + Rig( + owned_one.gateway, + owned_one, + provider, + endpoint, + "restart-model", + threading.Event(), + threading.Event(), + [], + ), + bodies[: BURST // 2], + ) + with owned_proxy_process(gateway, tmp_path, overrides, config=path, workers=2) as owned_two: + rig_two: Final = Rig( + owned_two.gateway, + owned_two, + provider, + endpoint, + "restart-model", + threading.Event(), + threading.Event(), + [], + ) + + probes: list[tuple[int, str]] = [] # mutable-ok: readiness polls are real served requests + + def served() -> tuple[int, ...]: + try: + probe: Final = rig_two.proxy.request("POST", *bodies[BURST // 2]) + except httpx.HTTPError: + return (-1,) + probes.append((probe.status_code, probe.headers.get("x-litellm-call-id") or "")) + return (probe.status_code,) + + eventually(served, lambda statuses: statuses[0] == 400, seconds=60) + second: Final = _fire(rig_two, bodies[BURST // 2 :]) + answered: Final = first + second + tuple(probes) + assert all(status in (400, 500) for status, _ in answered), answered + post_restart_ids: Final = frozenset(cid for status, cid in second + tuple(probes) if status == 400 and cid) + assert post_restart_ids, answered + assert any(marker.encode() in request.body for request in provider.drain()) + collected: list[Request] = [] # mutable-ok: drain consumes batches, later polls keep earlier ones + + def landed() -> tuple[str, ...]: + collected.extend(endpoint.drain()) + return tuple( + str(event.get("litellm_call_id")) + for batch in collected + if _json_body_ok(batch.body) + for event in json.loads(batch.body) + ) + + landed_ids: Final = eventually(landed, lambda ids: post_restart_ids <= set(ids), seconds=60) + assert all(not batch.body.strip() for batch in collected if not _json_body_ok(batch.body)), collected + assert len(landed_ids) == len(set(landed_ids)), ("duplicate events after restart", landed_ids) + events: Final = tuple( + object_value(event) + for batch in collected + if _json_body_ok(batch.body) + for event in json.loads(batch.body) + ) + for event in events: + assert marker not in json.dumps(event), json.dumps(event)[:400] + + +@pytest.mark.timeout(280) +def test_g4_worker_kill_keeps_serving_redacted(tmp_path: Path) -> None: + marker: Final = uuid.uuid4().hex + with wire_server(_provider) as provider, wire_server(lambda _: Reply()) as endpoint: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"].update( + {"callbacks": ["generic_api"], "DEFAULT_FLUSH_INTERVAL_SECONDS": 1, "turn_off_message_logging": True} + ) + path: Final = tmp_path / "chaos_worker.yaml" + path.write_text(yaml.safe_dump(config)) + with ( + gateway_from_environment() as gateway, + owned_proxy_process( + gateway, tmp_path, {"GENERIC_LOGGER_ENDPOINT": endpoint.url}, config=path, workers=2 + ) as owned, + owned.gateway.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + rig: Final = Rig(owned.gateway, owned, provider, endpoint, model, threading.Event(), threading.Event(), []) + bodies: Final = _bodies(model, marker)[:12] + children: Final = psutil.Process(owned.process.pid).children(recursive=True) + assert children, "no uvicorn worker children found" + with ThreadPoolExecutor(max_workers=8) as pool: + futures: Final = tuple( + pool.submit(lambda b: rig.proxy.request("POST", b[0], b[1]), body) for body in bodies + ) + eventually(lambda: provider.received.qsize() >= 3, bool, seconds=30) + children[0].send_signal(signal.SIGKILL) + statuses: list[int] = [] # mutable-ok: collect per-request outcomes from concurrent futures + for future in futures: + try: + statuses.append(future.result().status_code) + except httpx.HTTPError: + statuses.append(-1) + assert all(status == 400 for status in statuses if status >= 0), statuses + events: Final = eventually(rig.failure_events, lambda values: len(values) >= 1, seconds=70) + landed: Final = tuple(str(event.get("litellm_call_id")) for event in events) + assert len(landed) == len(set(landed)), ("duplicate failure events", landed) + for event in events: + assert marker not in json.dumps(event), json.dumps(event)[:400] diff --git a/tests/integration/observability/test_failure_redaction_datadog.py b/tests/integration/observability/test_failure_redaction_datadog.py new file mode 100644 index 00000000000..7a6d1309178 --- /dev/null +++ b/tests/integration/observability/test_failure_redaction_datadog.py @@ -0,0 +1,90 @@ +import gzip +import json +import uuid +from collections.abc import Iterator +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import pytest +import yaml +from integration._support.client import Gateway, JsonValue, eventually, gateway_from_environment, object_value +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, Wire, wire_server + + +def _body(batch: Request) -> bytes: + return gzip.decompress(batch.body) if batch.body[:2] == b"\x1f\x8b" else batch.body + + +def _provider(request: Request) -> Reply: + try: + body: Final = json.loads(request.body) + except json.JSONDecodeError: + return Reply(status=404, body=b"{}") + text: Final = body["messages"][-1]["content"] + return Reply( + status=400, + body=json.dumps( + {"error": {"type": "invalid_request_error", "message": f"Unsupported content: {text}"}} + ).encode(), + ) + + +@dataclass(frozen=True, slots=True) +class Rig: + proxy: Gateway + provider: Wire + sink: Wire + + def log_entries(self, model: str) -> tuple[dict[str, JsonValue], ...]: + return tuple( + object_value(entry) + for batch in self.sink.drain() + for entry in json.loads(_body(batch)) + if model in json.dumps(entry) + ) + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]: + root: Final = tmp_path_factory.mktemp("failure_redaction_datadog") + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"].update( + {"callbacks": ["datadog"], "DEFAULT_FLUSH_INTERVAL_SECONDS": 1, "turn_off_message_logging": True} + ) + path: Final = root / "datadog_failure.yaml" + path.write_text(yaml.safe_dump(config)) + with ( + wire_server(_provider) as provider, + wire_server(lambda _: Reply()) as sink, + gateway_from_environment() as gateway, + owned_proxy( + gateway, + root, + {"DD_API_KEY": "synthetic-dd-key", "DD_BASE_URL": sink.url, "DD_SITE": "localhost"}, + config=path, + workers=2, + ) as proxy, + ): + yield Rig(proxy, provider, sink) + + +def test_c5_datadog_failure_log_redacted_keeps_status(rig: Rig) -> None: + secret: Final = "dd-secret-" + uuid.uuid4().hex + with rig.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig.provider.url + "/v1", api_key="synthetic-provider-key") + response: Final = rig.proxy.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": secret}]}, + ) + assert response.status_code == 400, response.text + assert any(secret.encode() in request.body for request in rig.provider.drain()) + entries: Final = eventually(lambda: rig.log_entries(model), lambda values: len(values) >= 1, seconds=30) + entry: Final = entries[0] + assert secret not in json.dumps(entry), json.dumps(entry)[:2000] + message: Final = object_value(json.loads(str(entry.get("message", "{}")))) + assert secret not in json.dumps(message), json.dumps(message)[:2000] + error_information: Final = object_value(message.get("error_information")) + assert error_information.get("error_class"), error_information diff --git a/tests/integration/observability/test_failure_redaction_otel.py b/tests/integration/observability/test_failure_redaction_otel.py new file mode 100644 index 00000000000..840650a4ede --- /dev/null +++ b/tests/integration/observability/test_failure_redaction_otel.py @@ -0,0 +1,327 @@ +import json +import uuid +from collections.abc import Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import httpx +import pytest +import yaml +from google.protobuf.json_format import MessageToDict +from integration._support.client import Gateway, JsonValue, eventually, gateway_from_environment, object_value +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, Wire, wire_server +from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest + + +def _provider(request: Request) -> Reply: + try: + body: Final = json.loads(request.body) + except json.JSONDecodeError: + return Reply(status=404, body=b"{}") + text: Final = body["messages"][-1]["content"] + return Reply( + status=400, + body=json.dumps( + {"error": {"type": "invalid_request_error", "message": f"Unsupported content: {text}"}} + ).encode(), + ) + + +def _decode(body: bytes) -> dict[str, JsonValue]: + if body[:1] == b"{": + return json.loads(body) + request: Final = ExportTraceServiceRequest() + request.ParseFromString(body) + return object_value(MessageToDict(request)) + + +@dataclass(frozen=True, slots=True) +class Spans: + wire: Wire + batches: list[Request] + + def all(self) -> tuple[dict[str, JsonValue], ...]: + self.batches.extend(self.wire.drain()) + return tuple( + span + for batch in self.batches + for resource in _decode(batch.body).get("resourceSpans", ()) + for scope in resource.get("scopeSpans", ()) + for span in scope.get("spans", ()) + ) + + def named(self, name: str, model: str) -> tuple[dict[str, JsonValue], ...]: + return tuple(span for span in self.all() if span.get("name") == name and model in json.dumps(span)) + + def in_trace(self, name: str, trace_id: str) -> tuple[dict[str, JsonValue], ...]: + return tuple(span for span in self.all() if span.get("name") == name and span.get("traceId") == trace_id) + + +def _span_attributes(span: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + return { + str(attribute["key"]): object_value(attribute["value"]).get("stringValue") + or object_value(attribute["value"]).get("intValue") + for attribute in span.get("attributes", ()) + if isinstance(attribute, dict) + } + + +def _exception_events(span: Mapping[str, JsonValue]) -> tuple[dict[str, JsonValue], ...]: + return tuple( + object_value(event) + for event in span.get("events", ()) + if isinstance(event, dict) and event.get("name") == "exception" + ) + + +def _exception_texts(span: Mapping[str, JsonValue]) -> str: + return json.dumps(_exception_events(span)) + + +def _error_attribute(span: Mapping[str, JsonValue]) -> str: + attributes: Final = _span_attributes(span) + return str(attributes.get("error.message", "")) + + +@dataclass(frozen=True, slots=True) +class Rig: + proxy: Gateway + provider: Wire + sink: Spans + + +@contextmanager +def _otel_rig(root: Path, provider: Wire, sink: Wire, v2: bool, global_on: bool) -> Iterator[Rig]: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + settings: Final[dict[str, JsonValue]] = {"callbacks": ["otel"]} + if global_on: + settings["turn_off_message_logging"] = True + config["litellm_settings"].update(settings) + if v2: + config["callback_settings"] = { + "otel": {"exporter": "http/json", "endpoint": sink.url, "mapper_names": ["genai"]} + } + path: Final = root / "otel_failure.yaml" + path.write_text(yaml.safe_dump(config)) + env: Final = ( + {"LITELLM_OTEL_V2": "1", "OTEL_BSP_SCHEDULE_DELAY": "300"} + if v2 + else { + "OTEL_EXPORTER": "http/json", + "OTEL_EXPORTER_OTLP_ENDPOINT": sink.url, + "OTEL_BSP_SCHEDULE_DELAY": "300", + } + ) + with ( + gateway_from_environment() as gateway, + owned_proxy(gateway, root, env, config=path, workers=2) as proxy, + ): + yield Rig(proxy, provider, Spans(sink, [])) # mutable-ok: drain consumes, polls keep earlier batches + + +@pytest.fixture(scope="module") +def provider() -> Iterator[Wire]: + with wire_server(_provider) as wire: + yield wire + + +@pytest.fixture(scope="module") +def sink() -> Iterator[Wire]: + with wire_server(lambda _: Reply()) as wire: + yield wire + + +@pytest.fixture(scope="module") +def rig_v1_on(tmp_path_factory: pytest.TempPathFactory, provider: Wire, sink: Wire) -> Iterator[Rig]: + with _otel_rig(tmp_path_factory.mktemp("otel_v1_on"), provider, sink, v2=False, global_on=True) as booted: + yield booted + + +@pytest.fixture(scope="module") +def rig_v2_on(tmp_path_factory: pytest.TempPathFactory, provider: Wire, sink: Wire) -> Iterator[Rig]: + with _otel_rig(tmp_path_factory.mktemp("otel_v2_on"), provider, sink, v2=True, global_on=True) as booted: + yield booted + + +@pytest.fixture(scope="module") +def rig_v2_off(tmp_path_factory: pytest.TempPathFactory, provider: Wire, sink: Wire) -> Iterator[Rig]: + with _otel_rig(tmp_path_factory.mktemp("otel_v2_off"), provider, sink, v2=True, global_on=False) as booted: + yield booted + + +def _secret() -> str: + return "otel-secret-" + uuid.uuid4().hex + + +def _fail(rig: Rig, model: str, secret: str, **kwargs: JsonValue) -> httpx.Response: + headers: Final = kwargs.pop("headers", None) + key: Final = kwargs.pop("key", None) + return rig.proxy.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": secret}], **kwargs}, + headers=headers if isinstance(headers, dict) else None, + key=key if isinstance(key, str) else None, + ) + + +def _llm_span(rig: Rig, call_id: str) -> dict[str, JsonValue]: + def found() -> tuple[dict[str, JsonValue], ...]: + return tuple( + span + for span in rig.sink.all() + if call_id in json.dumps(span) and str(span.get("name", "")).startswith(("chat ", "litellm_request")) + ) + + return eventually(found, lambda values: len(values) >= 1, seconds=260)[0] + + +_SERVER_SPAN_NAMES: Final = ("Received Proxy Server Request", "POST /v1/chat/completions", "POST /v1/messages") + + +def _server_span(rig: Rig, call_id: str) -> dict[str, JsonValue]: + trace_id: Final = str(_llm_span(rig, call_id)["traceId"]) + + def found() -> tuple[dict[str, JsonValue], ...]: + return tuple( + span + for span in rig.sink.all() + if span.get("traceId") == trace_id and str(span.get("name", "")) in _SERVER_SPAN_NAMES + ) + + spans: Final = eventually(found, lambda values: len(values) >= 1, seconds=60, return_last_on_timeout=True) + assert spans, [(span.get("name"), span.get("traceId")) for span in rig.sink.all()] + return spans[0] + + +def _span_for(rig: Rig, name: str, model: str) -> dict[str, JsonValue]: + spans: Final = eventually(lambda: rig.sink.named(name, model), lambda values: len(values) >= 1, seconds=260) + return spans[0] + + +def _auth_exception_span_ids(rig: Rig) -> frozenset[str]: + return frozenset( + str(span.get("spanId")) + for span in rig.sink.all() + if str(span.get("name", "")).startswith("auth") and _exception_events(span) + ) + + +def _auth_exception_span(rig: Rig, exclude: frozenset[str]) -> dict[str, JsonValue]: + def found() -> tuple[dict[str, JsonValue], ...]: + return tuple( + span + for span in rig.sink.all() + if str(span.get("name", "")).startswith("auth") + and _exception_events(span) + and str(span.get("spanId")) not in exclude + ) + + return eventually(found, lambda values: len(values) >= 1, seconds=260)[0] + + +# C6: OTEL v1 failure span redaction under global on +@pytest.mark.timeout(320) +def test_c6_v1_provider_error_spans_redacted(rig_v1_on: Rig) -> None: + secret: Final = _secret() + with rig_v1_on.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig_v1_on.provider.url + "/v1", api_key="synthetic-provider-key") + response: Final = _fail(rig_v1_on, model, secret) + assert response.status_code == 400, response.text + call_id: Final = response.headers["x-litellm-call-id"] + assert any(secret.encode() in request.body for request in rig_v1_on.provider.drain()) + server: Final = _server_span(rig_v1_on, call_id) + assert _error_attribute(server) == "redacted-by-litellm", _span_attributes(server) + assert secret not in _exception_texts(server), _exception_texts(server)[:600] + request_span: Final = _llm_span(rig_v1_on, call_id) + assert _error_attribute(request_span) == "redacted-by-litellm", _span_attributes(request_span) + assert secret not in _exception_texts(request_span), _exception_texts(request_span)[:600] + + +# C7: OTEL v2 request + server spans redacted under global on +@pytest.mark.timeout(320) +def test_c7_v2_provider_error_spans_redacted(rig_v2_on: Rig) -> None: + secret: Final = _secret() + with rig_v2_on.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig_v2_on.provider.url + "/v1", api_key="synthetic-provider-key") + response: Final = _fail(rig_v2_on, model, secret) + assert response.status_code == 400, response.text + call_id: Final = response.headers["x-litellm-call-id"] + server: Final = _server_span(rig_v2_on, call_id) + assert _error_attribute(server) == "redacted-by-litellm", _span_attributes(server) + assert secret not in _exception_texts(server), _exception_texts(server)[:600] + + +# C8: v2 server span restamp keeps request opt-in under global off +@pytest.mark.timeout(320) +def test_c8_v2_header_opt_in_restamped_server_span_redacted(rig_v2_off: Rig) -> None: + secret: Final = _secret() + with rig_v2_off.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig_v2_off.provider.url + "/v1", api_key="synthetic-provider-key") + response: Final = _fail(rig_v2_off, model, secret, headers={"x-litellm-enable-message-redaction": "true"}) + assert response.status_code == 400, response.text + call_id: Final = response.headers["x-litellm-call-id"] + server: Final = _server_span(rig_v2_off, call_id) + assert _error_attribute(server) == "redacted-by-litellm", _span_attributes(server) + assert secret not in _exception_texts(server), _exception_texts(server)[:600] + + +# C9: v2 server span keeps permitted opt-out raw under global on +@pytest.mark.timeout(320) +def test_c9_v2_permitted_opt_out_keeps_server_span_raw(rig_v2_on: Rig) -> None: + secret: Final = _secret() + with rig_v2_on.proxy.scenario() as scenario: + model: Final = scenario.model(api_base=rig_v2_on.provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model], metadata={"allow_client_message_redaction_opt_out": True}) + response: Final = _fail( + rig_v2_on, model, secret, key=key, headers={"litellm-disable-message-redaction": "true"} + ) + assert response.status_code == 400, response.text + call_id: Final = response.headers["x-litellm-call-id"] + server: Final = _server_span(rig_v2_on, call_id) + assert secret in _error_attribute(server), _span_attributes(server) + + +# C10: v2 auth phase span honors request opt-in under global off +@pytest.mark.timeout(320) +def test_c10_v2_auth_span_honors_header_opt_in(rig_v2_off: Rig) -> None: + with rig_v2_off.proxy.scenario() as scenario: + allowed: Final = scenario.model(api_base=rig_v2_off.provider.url + "/v1", api_key="synthetic-provider-key") + denied: Final = scenario.model(api_base=rig_v2_off.provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[allowed]) + before: Final = _auth_exception_span_ids(rig_v2_off) + response: Final = rig_v2_off.proxy.request( + "POST", + "/v1/chat/completions", + {"model": denied, "messages": [{"role": "user", "content": "hi"}]}, + key=key, + headers={"x-litellm-enable-message-redaction": "true"}, + ) + assert response.status_code in (400, 401, 403, 404), response.text + auth: Final = _auth_exception_span(rig_v2_off, before) + assert allowed not in _exception_texts(auth), _exception_texts(auth)[:600] + assert "redacted-by-litellm" in _exception_texts(auth), _exception_texts(auth)[:600] + + +# C11: v2 auth phase span redacts under global on even with unpermitted disable header +@pytest.mark.timeout(320) +def test_c11_v2_auth_span_redacts_under_global_on(rig_v2_on: Rig) -> None: + with rig_v2_on.proxy.scenario() as scenario: + allowed: Final = scenario.model(api_base=rig_v2_on.provider.url + "/v1", api_key="synthetic-provider-key") + denied: Final = scenario.model(api_base=rig_v2_on.provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[allowed]) + before: Final = _auth_exception_span_ids(rig_v2_on) + response: Final = rig_v2_on.proxy.request( + "POST", + "/v1/chat/completions", + {"model": denied, "messages": [{"role": "user", "content": "hi"}]}, + key=key, + headers={"litellm-disable-message-redaction": "true"}, + ) + assert response.status_code in (400, 401, 403, 404), response.text + auth: Final = _auth_exception_span(rig_v2_on, before) + assert allowed not in _exception_texts(auth), _exception_texts(auth)[:600] + assert "redacted-by-litellm" in _exception_texts(auth), _exception_texts(auth)[:600] diff --git a/tests/unit/integrations/datadog/test_datadog_logger_batching.py b/tests/unit/integrations/datadog/test_datadog_logger_batching.py index e2707d321bf..7089cb8ba9f 100644 --- a/tests/unit/integrations/datadog/test_datadog_logger_batching.py +++ b/tests/unit/integrations/datadog/test_datadog_logger_batching.py @@ -1,4 +1,5 @@ import asyncio +import json from unittest.mock import AsyncMock, Mock, patch import httpx @@ -6,6 +7,7 @@ import pytest from httpx import Request, Response from pydantic import BaseModel, computed_field +import litellm from litellm.integrations.datadog.datadog import DataDogLogger from litellm.llms.custom_httpx.http_handler import MaskedHTTPStatusError from litellm.types.integrations.datadog import ( @@ -558,3 +560,44 @@ async def test_raised_intake_error_preserves_datadog_requeue_behavior(datadog_en await logger.async_send_batch() assert [event["message"] for event in logger.log_queue] == ['{"event": 0}', '{"event": 1}'] + + +@pytest.mark.asyncio +async def test_failure_hook_redacts_exception_payload_when_redaction_on(datadog_env, monkeypatch): + monkeypatch.setattr(litellm, "turn_off_message_logging", True) + with patch("asyncio.create_task"): + logger = DataDogLogger() + + secret = "secret-prompt-marker" + await logger.async_post_call_failure_hook( + request_data={"metadata": {}}, + original_exception=litellm.BadRequestError( + message=f"Unsupported content: {secret}", model="gpt-4o", llm_provider="openai" + ), + user_api_key_dict=type("UserKey", (), {})(), + traceback_str=f"Traceback ... {secret} ...", + ) + + message = json.loads(logger.log_queue[0]["message"]) + assert secret not in json.dumps(message) + assert message["exception"] == "redacted-by-litellm" + assert message["traceback"] == "redacted-by-litellm" + assert message["error_class"] == "BadRequestError" + + +@pytest.mark.asyncio +async def test_failure_hook_keeps_exception_payload_when_redaction_off(datadog_env, monkeypatch): + monkeypatch.setattr(litellm, "turn_off_message_logging", False) + with patch("asyncio.create_task"): + logger = DataDogLogger() + + secret = "secret-prompt-marker" + await logger.async_post_call_failure_hook( + request_data={"metadata": {}}, + original_exception=Exception(f"boom {secret}"), + user_api_key_dict=type("UserKey", (), {})(), + traceback_str="trace", + ) + + message = json.loads(logger.log_queue[0]["message"]) + assert message["exception"] == f"boom {secret}" diff --git a/tests/unit/integrations/otel/test_otel_v2_logger.py b/tests/unit/integrations/otel/test_otel_v2_logger.py index 62bf75bd083..7b24b4429f6 100644 --- a/tests/unit/integrations/otel/test_otel_v2_logger.py +++ b/tests/unit/integrations/otel/test_otel_v2_logger.py @@ -3238,3 +3238,190 @@ def test_provisional_close_then_payload_close_does_not_duplicate(): server.end() llm_spans = [s for s in exporter.get_finished_spans() if s.name.startswith("chat")] assert len(llm_spans) == 1 + + +def test_async_post_call_failure_hook_redacts_error_text_when_gated(): + """With message redaction on, the proxy-level failure span must not carry the + prompt through error.message / the exception event, while error.type and the + provider error code stay intact.""" + import litellm + from litellm.proxy._types import UserAPIKeyAuth + + litellm.turn_off_message_logging = True + try: + logger, exporter = _logger() + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) + set_request_root_span(server) + secret = "secret-prompt-marker" + result = asyncio.run( + logger.async_post_call_failure_hook( + request_data={"metadata": {}}, + original_exception=_proxy_exc(f"Unsupported content: {secret}", 400), + user_api_key_dict=UserAPIKeyAuth(), + traceback_str=f"Traceback ... {secret} ...", + ) + ) + server.end() + assert result is None + (span,) = exporter.get_finished_spans() + assert secret not in str(dict(span.attributes or {})) + assert secret not in str([dict(e.attributes or {}) for e in span.events]) + assert span.attributes["error.message"] == "redacted-by-litellm" + assert span.attributes["error.type"] == "ProxyException" + assert span.attributes["litellm.provider.error.code"] == "400" + assert span.attributes["litellm.provider.error.stack_trace"] == "redacted-by-litellm" + finally: + litellm.turn_off_message_logging = False + + +def test_async_post_call_failure_hook_keeps_error_text_when_not_gated(): + from litellm.proxy._types import UserAPIKeyAuth + + logger, exporter = _logger() + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) + set_request_root_span(server) + secret = "secret-prompt-marker" + asyncio.run( + logger.async_post_call_failure_hook( + request_data={"metadata": {}}, + original_exception=_proxy_exc(f"Unsupported content: {secret}", 400), + user_api_key_dict=UserAPIKeyAuth(), + traceback_str="trace", + ) + ) + server.end() + (span,) = exporter.get_finished_spans() + assert secret in span.attributes["error.message"] + + +def test_record_error_attributes_on_span_preserves_request_opt_in_redaction(): + """Global flag off but the request opts in via header: the failure hook + stamps redacted values on the SERVER span; the exception-handler restamp + must keep them instead of re-writing raw error text from the global probe.""" + import litellm + from litellm.proxy._types import UserAPIKeyAuth + + litellm.turn_off_message_logging = False + try: + logger, exporter = _logger() + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) + set_request_root_span(server) + secret = "secret-prompt-marker" + exc = _proxy_exc(f"Unsupported content: {secret}", 400) + asyncio.run( + logger.async_post_call_failure_hook( + request_data={"metadata": {"headers": {"x-litellm-enable-message-redaction": "true"}}}, + original_exception=exc, + user_api_key_dict=UserAPIKeyAuth(), + traceback_str=f"Traceback ... {secret} ...", + ) + ) + logger.record_error_attributes_on_span(server, exc, 400) + server.end() + (span,) = exporter.get_finished_spans() + assert secret not in str(dict(span.attributes or {})) + assert secret not in str([dict(e.attributes or {}) for e in span.events]) + assert secret not in str(span.status.description or "") + assert span.attributes["error.message"] == "redacted-by-litellm" + assert span.attributes["litellm.provider.error.code"] == "400" + finally: + litellm.turn_off_message_logging = False + + +def test_record_error_attributes_on_span_preserves_opt_out_raw_message(): + """Global flag on but the request opts out via header: the failure hook + stamps the raw error text; the restamp must not overwrite it with the + redaction marker.""" + import litellm + from litellm.proxy._types import UserAPIKeyAuth + + litellm.turn_off_message_logging = True + try: + logger, exporter = _logger() + server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME) + set_request_root_span(server) + secret = "secret-prompt-marker" + exc = _proxy_exc(f"Unsupported content: {secret}", 400) + asyncio.run( + logger.async_post_call_failure_hook( + request_data={"metadata": {"headers": {"litellm-disable-message-redaction": "true"}}}, + original_exception=exc, + user_api_key_dict=UserAPIKeyAuth(), + traceback_str=f"Traceback ... {secret} ...", + ) + ) + logger.record_error_attributes_on_span(server, exc, 400) + server.end() + (span,) = exporter.get_finished_spans() + assert secret in span.attributes["error.message"] + assert span.attributes["litellm.provider.error.code"] == "400" + finally: + litellm.turn_off_message_logging = False + + +def test_start_phase_span_does_not_record_raw_exception_when_gated(): + """With message redaction on, an exception raised inside a phase span must + not leak through use_span's automatic exception recording: exactly one + exception event lands, carrying only the redaction marker.""" + import litellm + + litellm.turn_off_message_logging = True + try: + logger, exporter = _logger() + secret = "secret-prompt-marker" + with pytest.raises(RuntimeError): + with logger.start_phase_span("auth"): + raise RuntimeError(f"auth exploded: {secret}") + (span,) = exporter.get_finished_spans() + assert secret not in str(dict(span.attributes or {})) + assert secret not in str([dict(e.attributes or {}) for e in span.events]) + assert secret not in str(span.status.description or "") + exception_events = [e for e in span.events if e.name == "exception"] + assert len(exception_events) == 1 + assert (exception_events[0].attributes or {}).get("exception.message") == "redacted-by-litellm" + finally: + litellm.turn_off_message_logging = False + + +def test_start_phase_span_redact_content_kwarg_redacts_exception(): + """A caller-provided opt-in (e.g. an enable header seen at auth time, where + the key's opt-out permission is not known yet) must redact the phase-span + exception even when the global flag is off.""" + import litellm + + litellm.turn_off_message_logging = False + try: + logger, exporter = _logger() + secret = "secret-prompt-marker" + with pytest.raises(RuntimeError): + with logger.start_phase_span("auth", redact_content=True): + raise RuntimeError(f"auth exploded: {secret}") + (span,) = exporter.get_finished_spans() + assert secret not in str(dict(span.attributes or {})) + assert secret not in str([dict(e.attributes or {}) for e in span.events]) + assert secret not in str(span.status.description or "") + exception_events = [e for e in span.events if e.name == "exception"] + assert len(exception_events) == 1 + assert (exception_events[0].attributes or {}).get("exception.message") == "redacted-by-litellm" + finally: + litellm.turn_off_message_logging = False + + +def test_start_phase_span_without_opt_in_keeps_raw_exception(): + """The default ``redact_content=False`` must leave use_span's own raw + exception recording untouched when the global flag is off.""" + import litellm + + litellm.turn_off_message_logging = False + try: + logger, exporter = _logger() + secret = "secret-prompt-marker" + with pytest.raises(RuntimeError): + with logger.start_phase_span("auth"): + raise RuntimeError(f"auth exploded: {secret}") + (span,) = exporter.get_finished_spans() + exception_events = [e for e in span.events if e.name == "exception"] + assert len(exception_events) == 1 + assert secret in str(dict(exception_events[0].attributes or {})) + finally: + litellm.turn_off_message_logging = False diff --git a/tests/unit/integrations/otel/test_runtime.py b/tests/unit/integrations/otel/test_runtime.py index d11f31b2523..239331b737b 100644 --- a/tests/unit/integrations/otel/test_runtime.py +++ b/tests/unit/integrations/otel/test_runtime.py @@ -9,6 +9,8 @@ import lock. These tests pin the import to a single resolution. import builtins +import pytest + import litellm.integrations.otel.runtime as runtime @@ -62,3 +64,36 @@ def test_wrappers_no_op_when_runtime_absent(monkeypatch): assert span is None assert runtime.seed_request_identity({"token": "sk-x"}, model="gpt-4o") is None + + +def test_phase_span_forwards_redact_content_to_registered_logger(monkeypatch): + """``redact_content`` must reach the v2 logger through the runtime shim: an + exception raised inside the forwarded span carries only the redaction marker, + even with the global flag off.""" + pytest.importorskip("opentelemetry") + + from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + + import litellm + import litellm.integrations.otel.logger as otel_logger + from litellm.integrations.otel import OpenTelemetryV2Config + from litellm.integrations.otel.logger import OpenTelemetryV2 + from litellm.integrations.otel.plumbing import providers + + cfg = OpenTelemetryV2Config(exporter="in_memory", legacy_compat=True, baggage_team_metadata_keys=[]) + exporter = InMemorySpanExporter() + logger = OpenTelemetryV2(config=cfg, tracer_provider=providers.build_tracer_provider(cfg, exporter=exporter)) + monkeypatch.setattr(otel_logger, "_registered_v2_logger", lambda: logger) + + litellm.turn_off_message_logging = False + try: + secret = "secret-prompt-marker" + with pytest.raises(RuntimeError): + with runtime.phase_span("auth", redact_content=True): + raise RuntimeError(f"auth exploded: {secret}") + (span,) = exporter.get_finished_spans() + assert secret not in str(dict(span.attributes or {})) + assert secret not in str([dict(e.attributes or {}) for e in span.events]) + assert secret not in str(span.status.description or "") + finally: + litellm.turn_off_message_logging = False diff --git a/tests/unit/integrations/test_mlflow.py b/tests/unit/integrations/test_mlflow.py index f828c34a9ff..bbaea11868e 100644 --- a/tests/unit/integrations/test_mlflow.py +++ b/tests/unit/integrations/test_mlflow.py @@ -262,3 +262,79 @@ def test_mlflow_end_span_or_trace_works_with_mlflow_2x_client(): span=child_span, outputs="out", end_time_ns=1, status="OK" ) assert client.ended_spans == [("req-2", "span-2")] + +def _failure_modules(): + modules = _mock_mlflow_modules() + + class RecordingSpanEvent: + calls = [] + + def __init__(self, name, attributes): + self.name = name + self.attributes = attributes + + @classmethod + def from_exception(cls, exception): + cls.calls.append(exception) + return cls("exception-from-exception", {"exception.message": str(exception)}) + + RecordingSpanEvent.calls = [] + modules["mlflow.entities"].SpanEvent = RecordingSpanEvent + modules["_span_event_cls"] = RecordingSpanEvent + return modules + + +def test_mlflow_failure_event_redacts_exception_when_gated(monkeypatch): + import litellm + + monkeypatch.setattr(litellm, "turn_off_message_logging", True) + modules = _failure_modules() + with patch.dict("sys.modules", modules): + from litellm.integrations.mlflow import MlflowLogger + + mlflow_logger = MlflowLogger() + span = MagicMock() + mlflow_logger._start_span_or_trace = MagicMock(return_value=span) + mlflow_logger._end_span_or_trace = MagicMock() + mlflow_logger._extract_and_set_chat_attributes = MagicMock() + + secret = "secret-prompt-marker" + mlflow_logger._handle_failure( + kwargs={"litellm_call_id": "x", "exception": Exception(f"boom {secret}")}, + response_obj=None, + start_time=datetime.utcnow(), + end_time=datetime.utcnow(), + ) + + event = span.add_event.call_args.args[0] + assert event.attributes == { + "exception.type": "Exception", + "exception.message": "redacted-by-litellm", + "exception.stacktrace": "redacted-by-litellm", + } + assert modules["_span_event_cls"].calls == [] + + +def test_mlflow_failure_event_uses_from_exception_when_not_gated(monkeypatch): + import litellm + + monkeypatch.setattr(litellm, "turn_off_message_logging", False) + modules = _failure_modules() + with patch.dict("sys.modules", modules): + from litellm.integrations.mlflow import MlflowLogger + + mlflow_logger = MlflowLogger() + span = MagicMock() + mlflow_logger._start_span_or_trace = MagicMock(return_value=span) + mlflow_logger._end_span_or_trace = MagicMock() + mlflow_logger._extract_and_set_chat_attributes = MagicMock() + + exc = Exception("boom secret-prompt-marker") + mlflow_logger._handle_failure( + kwargs={"litellm_call_id": "x", "exception": exc}, + response_obj=None, + start_time=datetime.utcnow(), + end_time=datetime.utcnow(), + ) + + assert modules["_span_event_cls"].calls == [exc] diff --git a/tests/unit/integrations/test_opentelemetry.py b/tests/unit/integrations/test_opentelemetry.py index 175bd95c263..dcb553ca18e 100644 --- a/tests/unit/integrations/test_opentelemetry.py +++ b/tests/unit/integrations/test_opentelemetry.py @@ -6762,3 +6762,130 @@ class TestOpenTelemetryNonInferenceUsage(unittest.TestCase): self.assertEqual( self._time_per_output_token_calls("aget_responses", response_obj=self.BACKGROUND_RESPONSE_OBJ), 1 ) + + +SECRET_PROMPT = "secret-prompt-marker" + + +class TestOpenTelemetryFailureHookRedaction(unittest.TestCase): + def _run_hook(self, request_data, exception, redact): + original = litellm.turn_off_message_logging + litellm.turn_off_message_logging = redact + try: + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + + otel = OpenTelemetry() + otel.tracer = tracer + server_span = tracer.start_span("Received Proxy Server Request") + + user_api_key_dict = MagicMock() + user_api_key_dict.parent_otel_span = server_span + + asyncio.run( + otel.async_post_call_failure_hook( + request_data=request_data, + original_exception=exception, + user_api_key_dict=user_api_key_dict, + traceback_str=f"Traceback ... {SECRET_PROMPT} ...", + ) + ) + finally: + litellm.turn_off_message_logging = original + + finished = {s.name: s for s in exporter.get_finished_spans()} + return finished + + def test_failure_hook_redacts_error_text_when_gated(self): + secret = SECRET_PROMPT + finished = self._run_hook( + {"metadata": {}}, + litellm.BadRequestError(message=f"Unsupported content: {secret}", model="gpt-4o", llm_provider="openai"), + redact=True, + ) + server = finished["Received Proxy Server Request"] + child = finished["Failed Proxy Server Request"] + for span in (server, child): + assert secret not in str(dict(span.attributes or {})) + for event in span.events: + assert secret not in str(dict(event.attributes or {})) + assert server.attributes["error.message"] == "redacted-by-litellm" + assert server.attributes["error.type"] == "BadRequestError" + assert child.attributes["exception"] == "redacted-by-litellm" + + def test_failure_hook_keeps_error_text_when_not_gated(self): + secret = SECRET_PROMPT + finished = self._run_hook( + {"metadata": {}}, + litellm.BadRequestError(message=f"Unsupported content: {secret}", model="gpt-4o", llm_provider="openai"), + redact=False, + ) + child = finished["Failed Proxy Server Request"] + assert secret in child.attributes["exception"] + + def test_record_error_attributes_on_span_redacts_via_global_gate(self): + original = litellm.turn_off_message_logging + litellm.turn_off_message_logging = True + try: + otel = OpenTelemetry() + span = MagicMock() + otel.record_error_attributes_on_span( + span=span, + exception=litellm.BadRequestError( + message=f"Unsupported content: {SECRET_PROMPT}", model="gpt-4o", llm_provider="openai" + ), + status_code=400, + ) + finally: + litellm.turn_off_message_logging = original + stamped = {call.args[0]: call.args[1] for call in span.set_attribute.call_args_list} + assert stamped.get("error.message") == "redacted-by-litellm" + assert stamped.get("error.type") == "BadRequestError" + + def _restamped_span(self, stamped_message, redact, exception): + original = litellm.turn_off_message_logging + litellm.turn_off_message_logging = redact + try: + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + otel = OpenTelemetry() + otel.tracer = provider.get_tracer(__name__) + span = otel.tracer.start_span("Received Proxy Server Request") + otel.safe_set_attribute(span=span, key="error.message", value=stamped_message) + otel.record_error_attributes_on_span(span=span, exception=exception, status_code=400) + span.end() + return exporter.get_finished_spans()[0] + finally: + litellm.turn_off_message_logging = original + + def test_record_error_attributes_on_span_preserves_request_opt_in_redaction(self): + """Global flag off but the request opted in: the failure hook already + stamped the redaction marker on the SERVER span; the exception-handler + restamp must keep it instead of writing raw error text.""" + span = self._restamped_span( + "redacted-by-litellm", + redact=False, + exception=litellm.BadRequestError( + message=f"Unsupported content: {SECRET_PROMPT}", model="gpt-4o", llm_provider="openai" + ), + ) + assert span.attributes["error.message"] == "redacted-by-litellm" + assert span.attributes["error.code"] == "400" + assert SECRET_PROMPT not in str(dict(span.attributes or {})) + + def test_record_error_attributes_on_span_preserves_opt_out_raw_message(self): + """Global flag on but the request opted out: the failure hook stamped + the raw error text; the restamp must not overwrite it with the + redaction marker.""" + span = self._restamped_span( + f"Unsupported content: {SECRET_PROMPT}", + redact=True, + exception=litellm.BadRequestError( + message=f"Unsupported content: {SECRET_PROMPT}", model="gpt-4o", llm_provider="openai" + ), + ) + assert SECRET_PROMPT in span.attributes["error.message"] + assert span.attributes["error.code"] == "400" diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index 60b7ed32399..05bfe3484cc 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -9081,3 +9081,141 @@ def test_signoz_dispatch_requires_an_endpoint(monkeypatch): logging_module._in_memory_loggers.clear() monkeypatch.delenv("LITELLM_OTEL_V2", raising=False) is_otel_v2_enabled.cache_clear() + + +class _CapturingFailureLogger(CustomLogger): + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.sync_kwargs = None + self.async_kwargs = None + + def log_failure_event(self, kwargs, response_obj, start_time, end_time): + self.sync_kwargs = kwargs + + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): + self.async_kwargs = kwargs + + +def _failure_logging_obj(secret): + return LitellmLogging( + model="gpt-4o", + messages=[{"role": "user", "content": secret}], + stream=False, + call_type="completion", + start_time=time.time(), + litellm_call_id="failure-redaction-test", + function_id="failure-redaction-test", + ) + + +def _assert_no_secret_leaks(kwargs, secret): + payload = kwargs.get("standard_logging_object") or {} + assert secret not in json.dumps(payload, default=str) + assert secret not in str(kwargs.get("traceback_exception")) + error_information = payload.get("error_information") or {} + assert error_information.get("error_class") == "BadRequestError" + assert error_information.get("error_code") == "400" + assert error_information.get("llm_provider") == "openai" + + +def test_sync_failure_handler_redacts_error_text_for_custom_logger(monkeypatch): + secret = "secret-prompt-marker" + capture = _CapturingFailureLogger() + monkeypatch.setattr(litellm, "turn_off_message_logging", True) + monkeypatch.setattr(litellm, "failure_callback", [capture]) + monkeypatch.setattr(litellm, "standard_logging_payload_excluded_fields", None) + logging_obj = _failure_logging_obj(secret) + + logging_obj.failure_handler( + litellm.BadRequestError(message=f"Unsupported content: {secret}", model="gpt-4o", llm_provider="openai"), + f"Traceback ... {secret} ...", + ) + + assert capture.sync_kwargs is not None + _assert_no_secret_leaks(capture.sync_kwargs, secret) + + +def test_sync_failure_handler_honours_excluded_fields(monkeypatch): + capture = _CapturingFailureLogger() + monkeypatch.setattr(litellm, "failure_callback", [capture]) + monkeypatch.setattr(litellm, "standard_logging_payload_excluded_fields", ["hidden_params"]) + logging_obj = _failure_logging_obj("anything") + + logging_obj.failure_handler(litellm.BadRequestError(message="denied", model="gpt-4o", llm_provider="openai"), "tb") + + payload = capture.sync_kwargs.get("standard_logging_object") or {} + assert "hidden_params" not in payload + + +def test_sync_failure_handler_leaves_error_text_alone_when_redaction_off(monkeypatch): + secret = "secret-prompt-marker" + capture = _CapturingFailureLogger() + monkeypatch.setattr(litellm, "turn_off_message_logging", False) + monkeypatch.setattr(litellm, "failure_callback", [capture]) + monkeypatch.setattr(litellm, "standard_logging_payload_excluded_fields", None) + logging_obj = _failure_logging_obj(secret) + + logging_obj.failure_handler( + litellm.BadRequestError(message=f"Unsupported content: {secret}", model="gpt-4o", llm_provider="openai"), + f"Traceback ... {secret} ...", + ) + + payload = capture.sync_kwargs.get("standard_logging_object") or {} + assert secret in payload.get("error_information", {}).get("error_message", "") + assert secret in str(payload.get("error_str")) + + +@pytest.mark.asyncio +async def test_async_failure_handler_redacts_error_text_for_custom_logger(monkeypatch): + secret = "secret-prompt-marker" + capture = _CapturingFailureLogger() + monkeypatch.setattr(litellm, "turn_off_message_logging", True) + monkeypatch.setattr(litellm, "_async_failure_callback", [capture]) + monkeypatch.setattr(litellm, "standard_logging_payload_excluded_fields", None) + logging_obj = _failure_logging_obj(secret) + + await logging_obj.async_failure_handler( + litellm.BadRequestError(message=f"Unsupported content: {secret}", model="gpt-4o", llm_provider="openai"), + f"Traceback ... {secret} ...", + ) + + assert capture.async_kwargs is not None + _assert_no_secret_leaks(capture.async_kwargs, secret) + + +@pytest.mark.asyncio +async def test_async_failure_handler_honours_excluded_fields(monkeypatch): + capture = _CapturingFailureLogger() + monkeypatch.setattr(litellm, "_async_failure_callback", [capture]) + monkeypatch.setattr(litellm, "standard_logging_payload_excluded_fields", ["hidden_params"]) + logging_obj = _failure_logging_obj("anything") + + await logging_obj.async_failure_handler( + litellm.BadRequestError(message="denied", model="gpt-4o", llm_provider="openai"), "tb" + ) + + payload = capture.async_kwargs.get("standard_logging_object") or {} + assert "hidden_params" not in payload + + +@pytest.mark.asyncio +async def test_per_callback_turn_off_redacts_error_fields_even_for_self_redacting_logger(monkeypatch): + secret = "secret-prompt-marker" + + class SelfRedactingLogger(_CapturingFailureLogger): + def redacts_messages_itself(self): + return True + + capture = SelfRedactingLogger(turn_off_message_logging=True) + monkeypatch.setattr(litellm, "turn_off_message_logging", False) + monkeypatch.setattr(litellm, "_async_failure_callback", [capture]) + monkeypatch.setattr(litellm, "standard_logging_payload_excluded_fields", None) + logging_obj = _failure_logging_obj(secret) + + await logging_obj.async_failure_handler( + litellm.BadRequestError(message=f"Unsupported content: {secret}", model="gpt-4o", llm_provider="openai"), + f"Traceback ... {secret} ...", + ) + + assert capture.async_kwargs is not None + _assert_no_secret_leaks(capture.async_kwargs, secret) diff --git a/tests/unit/litellm_core_utils/test_redact_messages.py b/tests/unit/litellm_core_utils/test_redact_messages.py index 76d037ce760..fa09d34df4e 100644 --- a/tests/unit/litellm_core_utils/test_redact_messages.py +++ b/tests/unit/litellm_core_utils/test_redact_messages.py @@ -1040,3 +1040,193 @@ def test_perform_redaction_drops_the_served_output_texts_from_the_callback_kwarg details: Final = {"litellm_params": {}, SERVED_OUTPUT_TEXTS_KEY: ("Card: ",)} perform_redaction(details, None) assert SERVED_OUTPUT_TEXTS_KEY not in details + + +def _error_information(**overrides): + info = { + "error_code": "400", + "error_class": "BadRequestError", + "llm_provider": "openai", + "traceback": "Traceback ... secret-prompt-marker ...", + "error_message": "Unsupported content: secret-prompt-marker", + "error_rate_limit_category": None, + "normalized_error": None, + } + info.update(overrides) + return info + + +class TestRedactErrorInformation: + def test_replaces_message_and_traceback_keeps_everything_else(self): + from litellm.litellm_core_utils.redact_messages import redact_error_information + + info = _error_information() + redacted = redact_error_information(info) + assert redacted == { + "error_code": "400", + "error_class": "BadRequestError", + "llm_provider": "openai", + "traceback": "redacted-by-litellm", + "error_message": "redacted-by-litellm", + "error_rate_limit_category": None, + "normalized_error": None, + } + + def test_empty_fields_stay_empty(self): + from litellm.litellm_core_utils.redact_messages import redact_error_information + + redacted = redact_error_information(_error_information(traceback="", error_message=None)) + assert redacted["traceback"] == "" + assert redacted["error_message"] is None + assert "secret-prompt-marker" not in str(redacted) + + def test_input_not_mutated(self): + from litellm.litellm_core_utils.redact_messages import redact_error_information + + info = _error_information() + redact_error_information(info) + assert info["error_message"] == "Unsupported content: secret-prompt-marker" + assert info["traceback"] == "Traceback ... secret-prompt-marker ..." + + +class TestFailureRedactionOnStandardLoggingObject: + def test_redacted_standard_logging_payload_covers_error_fields(self): + payload = { + "messages": [{"role": "user", "content": "secret-prompt-marker"}], + "error_str": "Error: secret-prompt-marker", + "error_information": _error_information(), + "status": "failure", + } + redacted = redacted_standard_logging_payload(payload) + assert "secret-prompt-marker" not in str(redacted) + assert redacted["error_str"] == "redacted-by-litellm" + assert redacted["error_information"] == { + "error_code": "400", + "error_class": "BadRequestError", + "llm_provider": "openai", + "traceback": "redacted-by-litellm", + "error_message": "redacted-by-litellm", + "error_rate_limit_category": None, + "normalized_error": None, + } + + def test_perform_redaction_covers_traceback_exception(self): + details = { + "litellm_params": {}, + "traceback_exception": "Traceback ... secret-prompt-marker", + "standard_logging_object": {"error_str": "secret-prompt-marker"}, + } + perform_redaction(details, None) + assert "secret-prompt-marker" not in str(details) + assert details["traceback_exception"] == "redacted-by-litellm" + assert details["standard_logging_object"]["error_str"] == "redacted-by-litellm" + + def test_perform_redaction_leaves_empty_traceback_exception_alone(self): + details = {"litellm_params": {}, "traceback_exception": ""} + perform_redaction(details, None) + assert details["traceback_exception"] == "" + + +def _request_data(metadata=None, litellm_metadata=None, turn_off_message_logging=None): + data = {"metadata": metadata if metadata is not None else {}} + if litellm_metadata is not None: + data["litellm_metadata"] = litellm_metadata + if turn_off_message_logging is not None: + data["turn_off_message_logging"] = turn_off_message_logging + return data + + +class TestShouldRedactFailedRequest: + def test_global_on(self): + from litellm.litellm_core_utils.redact_messages import should_redact_failed_request + + litellm.turn_off_message_logging = True + assert should_redact_failed_request(_request_data()) is True + + def test_global_off(self): + from litellm.litellm_core_utils.redact_messages import should_redact_failed_request + + assert should_redact_failed_request(_request_data()) is False + + def test_enable_header_in_metadata(self): + from litellm.litellm_core_utils.redact_messages import should_redact_failed_request + + request_data = _request_data(metadata={"headers": {"x-litellm-enable-message-redaction": "true"}}) + assert should_redact_failed_request(request_data) is True + + def test_disable_header_overrides_global_on(self): + from litellm.litellm_core_utils.redact_messages import should_redact_failed_request + + litellm.turn_off_message_logging = True + request_data = _request_data(metadata={"headers": {"litellm-disable-message-redaction": "true"}}) + assert should_redact_failed_request(request_data) is False + + def test_enable_header_in_litellm_metadata(self): + from litellm.litellm_core_utils.redact_messages import should_redact_failed_request + + request_data = _request_data(litellm_metadata={"headers": {"x-litellm-enable-message-redaction": "true"}}) + assert should_redact_failed_request(request_data) is True + + def test_dynamic_param_true(self): + from litellm.litellm_core_utils.redact_messages import should_redact_failed_request + + assert should_redact_failed_request(_request_data(turn_off_message_logging=True)) is True + + def test_dynamic_param_false_overrides_global_on(self): + from litellm.litellm_core_utils.redact_messages import should_redact_failed_request + + litellm.turn_off_message_logging = True + assert should_redact_failed_request(_request_data(turn_off_message_logging=False)) is False + + def test_dynamic_param_in_metadata_slot(self): + from litellm.litellm_core_utils.redact_messages import should_redact_failed_request + + assert should_redact_failed_request(_request_data(metadata={"turn_off_message_logging": True})) is True + + def test_dynamic_param_in_litellm_metadata_slot(self): + from litellm.litellm_core_utils.redact_messages import should_redact_failed_request + + assert should_redact_failed_request(_request_data(litellm_metadata={"turn_off_message_logging": True})) is True + + def test_top_level_dynamic_param_beats_metadata_slot(self): + from litellm.litellm_core_utils.redact_messages import should_redact_failed_request + + litellm.turn_off_message_logging = True + request_data = _request_data( + metadata={"turn_off_message_logging": True}, + turn_off_message_logging=False, + ) + assert should_redact_failed_request(request_data) is False + + +class TestRequestOptsIntoMessageRedaction: + @pytest.mark.parametrize("header", ["litellm-enable-message-redaction", "x-litellm-enable-message-redaction"]) + def test_enable_header(self, header: str) -> None: + from litellm.litellm_core_utils.redact_messages import request_opts_into_message_redaction + + assert request_opts_into_message_redaction({header: "true"}, {}) is True + + def test_top_level_dynamic_param(self) -> None: + from litellm.litellm_core_utils.redact_messages import request_opts_into_message_redaction + + assert request_opts_into_message_redaction({}, {"turn_off_message_logging": True}) is True + + def test_metadata_slot_dynamic_param(self) -> None: + from litellm.litellm_core_utils.redact_messages import request_opts_into_message_redaction + + assert request_opts_into_message_redaction({}, {"metadata": {"turn_off_message_logging": True}}) is True + + def test_empty_inputs(self) -> None: + from litellm.litellm_core_utils.redact_messages import request_opts_into_message_redaction + + assert request_opts_into_message_redaction({}, {}) is False + + def test_disable_header_alone_is_not_an_opt_in(self) -> None: + from litellm.litellm_core_utils.redact_messages import request_opts_into_message_redaction + + assert request_opts_into_message_redaction({"litellm-disable-message-redaction": "true"}, {}) is False + + def test_dynamic_param_false(self) -> None: + from litellm.litellm_core_utils.redact_messages import request_opts_into_message_redaction + + assert request_opts_into_message_redaction({}, {"turn_off_message_logging": False}) is False diff --git a/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py b/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py new file mode 100644 index 00000000000..3feecccaaa1 --- /dev/null +++ b/tests/unit/proxy/hooks/test_proxy_track_cost_callback.py @@ -0,0 +1,61 @@ +import json +from unittest.mock import AsyncMock, MagicMock + +import pytest + +import litellm +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger + +SECRET = "secret-prompt-marker" + + +def _logger_and_writer(): + writer = MagicMock() + writer.update_database = AsyncMock() + logger = _ProxyDBLogger(spend_writer=lambda: writer) + return logger, writer + + +@pytest.mark.asyncio +async def test_failure_hook_redacts_persisted_error_information(monkeypatch): + monkeypatch.setattr(litellm, "turn_off_message_logging", True) + logger, writer = _logger_and_writer() + + request_data = {"metadata": {}} + await logger.async_post_call_failure_hook( + request_data=request_data, + original_exception=litellm.BadRequestError( + message=f"Unsupported content: {SECRET}", model="gpt-4o", llm_provider="openai" + ), + user_api_key_dict=UserAPIKeyAuth(), + traceback_str=f"Traceback ... {SECRET} ...", + ) + + persisted = request_data["litellm_params"]["metadata"]["error_information"] + assert SECRET not in json.dumps(persisted) + assert persisted["error_message"] == "redacted-by-litellm" + assert persisted["traceback"] == "redacted-by-litellm" + assert persisted["error_class"] == "BadRequestError" + assert persisted["error_code"] == "400" + writer.update_database.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_failure_hook_leaves_error_information_alone_when_redaction_off(monkeypatch): + monkeypatch.setattr(litellm, "turn_off_message_logging", False) + logger, writer = _logger_and_writer() + + request_data = {"metadata": {}} + await logger.async_post_call_failure_hook( + request_data=request_data, + original_exception=litellm.BadRequestError( + message=f"Unsupported content: {SECRET}", model="gpt-4o", llm_provider="openai" + ), + user_api_key_dict=UserAPIKeyAuth(), + traceback_str="trace", + ) + + persisted = request_data["litellm_params"]["metadata"]["error_information"] + assert SECRET in persisted["error_message"] + writer.update_database.assert_awaited_once()