diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index 62ca6b0254e..22e2e34c586 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Any, ClassVar, Final, Optional 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 @@ -913,9 +913,19 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac if turn_off_message_logging is False and not excluded_fields: return model_call_details + redacted_model_call_details: Final = ( + { + **model_call_details, + "messages": [{"role": "user", "content": REDACTED_BY_LITELLM}], + "prompt": "", + "input": "", + } + if turn_off_message_logging and not self.redacts_messages_itself() + else model_call_details.copy() + ) standard_logging_object: Final = model_call_details.get("standard_logging_object") if standard_logging_object is None: - return model_call_details.copy() + return redacted_model_call_details # Make a copy of just the standard_logging_object to avoid modifying the original standard_logging_object_copy: Final = { @@ -963,7 +973,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac else EMPTY_MAPPING ) return { - **model_call_details, + **redacted_model_call_details, **redacted_params, "standard_logging_object": standard_logging_object_copy, } diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index ae608244b38..8e01c6793f8 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -93,6 +93,8 @@ from litellm.litellm_core_utils.model_param_helper import ModelParamHelper from litellm.litellm_core_utils.redact_messages import ( redact_message_input_output_from_custom_logger, redact_message_input_output_from_logging, + redact_model_call_details_for_custom_logger, + redact_response_for_custom_logger, redact_streaming_responses_for_custom_logger, should_redact_message_logging, ) @@ -2856,13 +2858,16 @@ class Logging(LiteLLMLoggingBaseClass): ): # custom logger class if self.stream and complete_streaming_response is None: callback.log_stream_event( - kwargs=redact_streaming_responses_for_custom_logger( - model_call_details=callback.redact_standard_logging_payload_from_model_call_details( - model_call_details=self.model_call_details + kwargs=redact_model_call_details_for_custom_logger( + model_call_details=redact_streaming_responses_for_custom_logger( + model_call_details=callback.redact_standard_logging_payload_from_model_call_details( + model_call_details=self.model_call_details + ), + custom_logger=callback, ), custom_logger=callback, ), - response_obj=result, + response_obj=redact_response_for_custom_logger(result=result, custom_logger=callback), start_time=start_time, end_time=end_time, ) @@ -2874,13 +2879,16 @@ class Logging(LiteLLMLoggingBaseClass): result = self.model_call_details["complete_response"] callback.log_success_event( - kwargs=redact_streaming_responses_for_custom_logger( - model_call_details=callback.redact_standard_logging_payload_from_model_call_details( - model_call_details=self.model_call_details + kwargs=redact_model_call_details_for_custom_logger( + model_call_details=redact_streaming_responses_for_custom_logger( + model_call_details=callback.redact_standard_logging_payload_from_model_call_details( + model_call_details=self.model_call_details + ), + custom_logger=callback, ), custom_logger=callback, ), - response_obj=result, + response_obj=redact_response_for_custom_logger(result=result, custom_logger=callback), start_time=start_time, end_time=end_time, ) @@ -3195,26 +3203,32 @@ class Logging(LiteLLMLoggingBaseClass): model_call_details = redact_streaming_responses_for_custom_logger( model_call_details=model_call_details, custom_logger=callback ) + model_call_details = redact_model_call_details_for_custom_logger( + model_call_details=model_call_details, custom_logger=callback + ) ################################## if self.stream is True: if "async_complete_streaming_response" in model_call_details: await callback.async_log_success_event( kwargs=model_call_details, - response_obj=model_call_details["async_complete_streaming_response"], + response_obj=redact_response_for_custom_logger( + result=model_call_details["async_complete_streaming_response"], + custom_logger=callback, + ), start_time=start_time, end_time=end_time, ) else: await callback.async_log_stream_event( # [TODO]: move this to being an async log stream event function kwargs=model_call_details, - response_obj=result, + response_obj=redact_response_for_custom_logger(result=result, custom_logger=callback), start_time=start_time, end_time=end_time, ) else: await callback.async_log_success_event( kwargs=model_call_details, - response_obj=result, + response_obj=redact_response_for_custom_logger(result=result, custom_logger=callback), start_time=start_time, end_time=end_time, ) @@ -3502,9 +3516,12 @@ class Logging(LiteLLMLoggingBaseClass): callback.log_failure_event( start_time=start_time, end_time=end_time, - response_obj=result, - kwargs=callback.redact_standard_logging_payload_from_model_call_details( - model_call_details=self.model_call_details + response_obj=redact_response_for_custom_logger(result=result, custom_logger=callback), + kwargs=redact_model_call_details_for_custom_logger( + model_call_details=callback.redact_standard_logging_payload_from_model_call_details( + model_call_details=self.model_call_details + ), + custom_logger=callback, ), ) if callback == "langfuse": @@ -3636,10 +3653,13 @@ class Logging(LiteLLMLoggingBaseClass): continue if isinstance(callback, CustomLogger): # custom logger class await callback.async_log_failure_event( - kwargs=callback.redact_standard_logging_payload_from_model_call_details( - model_call_details=self.model_call_details + kwargs=redact_model_call_details_for_custom_logger( + model_call_details=callback.redact_standard_logging_payload_from_model_call_details( + model_call_details=self.model_call_details + ), + custom_logger=callback, ), - response_obj=result, + response_obj=redact_response_for_custom_logger(result=result, custom_logger=callback), 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 95777ed7f87..89eb6fcc38c 100644 --- a/litellm/litellm_core_utils/redact_messages.py +++ b/litellm/litellm_core_utils/redact_messages.py @@ -68,6 +68,36 @@ def redact_streaming_responses_for_custom_logger(model_call_details: dict, custo return {**model_call_details, **redacted_entries} +def redact_response_for_custom_logger(result: object, custom_logger: CustomLogger) -> object: + opted_out: Final = ( + getattr(custom_logger, "message_logging", True) is not True + or getattr(custom_logger, "turn_off_message_logging", False) is True + ) + if not opted_out or custom_logger.redacts_messages_itself(): + return result + return perform_redaction(model_call_details={}, result=result, redact_streaming_responses=False) + + +def redact_model_call_details_for_custom_logger( + model_call_details: dict[str, object], custom_logger: CustomLogger +) -> dict[str, object]: + opted_out: Final = ( + getattr(custom_logger, "message_logging", True) is not True + or getattr(custom_logger, "turn_off_message_logging", False) is True + ) + if not opted_out or custom_logger.redacts_messages_itself(): + return model_call_details + response_keys: Final = ("response", "original_response", "complete_response") + return { + **model_call_details, + **{ + key: redact_response_for_custom_logger(result=model_call_details[key], custom_logger=custom_logger) + for key in response_keys + if model_call_details.get(key) is not None + }, + } + + def _redacted_streaming_response_copy(streaming_response): redacted_response: Final = copy.deepcopy(streaming_response) _redact_streaming_response(redacted_response) diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 9604f72daa0..c5d7ba9f08a 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -48,19 +48,32 @@ def logging_obj(): class _RecordingCustomLogger(CustomLogger): def __init__(self, turn_off_message_logging: bool): super().__init__(turn_off_message_logging=turn_off_message_logging) - self.received_kwargs: list[dict] = [] + self.received_kwargs: tuple[dict[str, object], ...] = () + self.received_response_objs: tuple[object, ...] = () - def log_stream_event(self, kwargs: dict, response_obj: object, start_time, end_time): - self.received_kwargs.append(kwargs) + def log_stream_event( + self, kwargs: dict[str, object], response_obj: object, start_time: datetime.datetime, end_time: datetime.datetime + ): + self.received_kwargs += (kwargs,) + self.received_response_objs += (response_obj,) - def log_success_event(self, kwargs: dict, response_obj: object, start_time, end_time): - self.received_kwargs.append(kwargs) + def log_success_event( + self, kwargs: dict[str, object], response_obj: object, start_time: datetime.datetime, end_time: datetime.datetime + ): + self.received_kwargs += (kwargs,) + self.received_response_objs += (response_obj,) - def log_failure_event(self, kwargs: dict, response_obj: object, start_time, end_time): - self.received_kwargs.append(kwargs) + def log_failure_event( + self, kwargs: dict[str, object], response_obj: object, start_time: datetime.datetime, end_time: datetime.datetime + ): + self.received_kwargs += (kwargs,) + self.received_response_objs += (response_obj,) - async def async_log_failure_event(self, kwargs: dict, response_obj: object, start_time, end_time): - self.received_kwargs.append(kwargs) + async def async_log_failure_event( + self, kwargs: dict[str, object], response_obj: object, start_time: datetime.datetime, end_time: datetime.datetime + ): + self.received_kwargs += (kwargs,) + self.received_response_objs += (response_obj,) def test_get_combined_callback_list_preserves_insertion_order(logging_obj): @@ -1038,6 +1051,9 @@ def test_success_handler_redacts_custom_logger_payload_per_callback(logging_obj) plain_logger = _RecordingCustomLogger(turn_off_message_logging=False) logging_obj.stream = False logging_obj.model_call_details["litellm_params"] = {} + logging_obj.model_call_details["original_response"] = { + "choices": [{"message": {"content": "original response"}}] + } standard_logging_object = { "messages": [{"role": "user", "content": "original message"}], "response": {"choices": []}, @@ -1060,6 +1076,34 @@ def test_success_handler_redacts_custom_logger_payload_per_callback(logging_obj) assert ( logging_obj.model_call_details["standard_logging_object"]["messages"][0]["content"] == "original message" ) + assert redacting_logger.received_kwargs[0]["messages"] == [{"role": "user", "content": "redacted-by-litellm"}] + assert ( + redacting_logger.received_kwargs[0]["original_response"]["choices"][0]["message"]["content"] + == "redacted-by-litellm" + ) + assert ( + plain_logger.received_kwargs[0]["original_response"]["choices"][0]["message"]["content"] == "original response" + ) + assert redacting_logger.received_response_objs[0] == {"text": "redacted-by-litellm"} + assert plain_logger.received_response_objs[0] == {"id": "response"} + + +def test_redact_standard_logging_payload_redacts_top_level_fields_without_standard_payload(): + logger = CustomLogger(turn_off_message_logging=True) + model_call_details = { + "messages": [{"role": "user", "content": "original message"}], + "input": "original input", + "prompt": "original prompt", + } + + redacted_details = logger.redact_standard_logging_payload_from_model_call_details(model_call_details) + + assert redacted_details["messages"] == [{"role": "user", "content": "redacted-by-litellm"}] + assert redacted_details["input"] == "" + assert redacted_details["prompt"] == "" + assert model_call_details["messages"][0]["content"] == "original message" + assert model_call_details["input"] == "original input" + assert model_call_details["prompt"] == "original prompt" def test_failure_handler_redacts_custom_logger_payload_per_callback(logging_obj):