diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 89702e0345f..354cde57396 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -134,8 +134,10 @@ from .initialize_dynamic_callback_params import ( from .specialty_caches.dynamic_logging_cache import DynamicLoggingCache try: - from enterprise.enterprise_callbacks.generic_api_callback import GenericAPILogger - from enterprise.enterprise_callbacks.send_emails.resend_email import ( + from litellm_enterprise.enterprise_callbacks.generic_api_callback import ( + GenericAPILogger, + ) + from litellm_enterprise.enterprise_callbacks.send_emails.resend_email import ( ResendEmailLogger, ) except Exception as e: @@ -257,9 +259,9 @@ class Logging(LiteLLMLoggingBaseClass): self.litellm_trace_id: str = litellm_trace_id or str(uuid.uuid4()) self.function_id = function_id self.streaming_chunks: List[Any] = [] # for generating complete stream response - self.sync_streaming_chunks: List[Any] = ( - [] - ) # for generating complete stream response + self.sync_streaming_chunks: List[ + Any + ] = [] # for generating complete stream response self.log_raw_request_response = log_raw_request_response # Initialize dynamic callbacks @@ -610,9 +612,9 @@ class Logging(LiteLLMLoggingBaseClass): if anthropic_cache_control_logger := AnthropicCacheControlHook.get_custom_logger_for_anthropic_cache_control_hook( non_default_params ): - self.model_call_details["prompt_integration"] = ( - anthropic_cache_control_logger.__class__.__name__ - ) + self.model_call_details[ + "prompt_integration" + ] = anthropic_cache_control_logger.__class__.__name__ return anthropic_cache_control_logger ######################################################### @@ -630,9 +632,9 @@ class Logging(LiteLLMLoggingBaseClass): ), ) ) - self.model_call_details["prompt_integration"] = ( - vector_store_custom_logger.__class__.__name__ - ) + self.model_call_details[ + "prompt_integration" + ] = vector_store_custom_logger.__class__.__name__ return vector_store_custom_logger return None @@ -684,9 +686,9 @@ class Logging(LiteLLMLoggingBaseClass): model ): # if model name was changes pre-call, overwrite the initial model call name with the new one self.model_call_details["model"] = model - self.model_call_details["litellm_params"]["api_base"] = ( - self._get_masked_api_base(additional_args.get("api_base", "")) - ) + self.model_call_details["litellm_params"][ + "api_base" + ] = self._get_masked_api_base(additional_args.get("api_base", "")) def pre_call(self, input, api_key, model=None, additional_args={}): # noqa: PLR0915 # Log the exact input to the LLM API @@ -715,10 +717,10 @@ class Logging(LiteLLMLoggingBaseClass): try: # [Non-blocking Extra Debug Information in metadata] if turn_off_message_logging is True: - _metadata["raw_request"] = ( - "redacted by litellm. \ + _metadata[ + "raw_request" + ] = "redacted by litellm. \ 'litellm.turn_off_message_logging=True'" - ) else: curl_command = self._get_request_curl_command( api_base=additional_args.get("api_base", ""), @@ -729,32 +731,32 @@ class Logging(LiteLLMLoggingBaseClass): _metadata["raw_request"] = str(curl_command) # split up, so it's easier to parse in the UI - self.model_call_details["raw_request_typed_dict"] = ( - RawRequestTypedDict( - raw_request_api_base=str( - additional_args.get("api_base") or "" - ), - raw_request_body=self._get_raw_request_body( - additional_args.get("complete_input_dict", {}) - ), - raw_request_headers=self._get_masked_headers( - additional_args.get("headers", {}) or {}, - ignore_sensitive_headers=True, - ), - error=None, - ) + self.model_call_details[ + "raw_request_typed_dict" + ] = RawRequestTypedDict( + raw_request_api_base=str( + additional_args.get("api_base") or "" + ), + raw_request_body=self._get_raw_request_body( + additional_args.get("complete_input_dict", {}) + ), + raw_request_headers=self._get_masked_headers( + additional_args.get("headers", {}) or {}, + ignore_sensitive_headers=True, + ), + error=None, ) except Exception as e: - self.model_call_details["raw_request_typed_dict"] = ( - RawRequestTypedDict( - error=str(e), - ) + self.model_call_details[ + "raw_request_typed_dict" + ] = RawRequestTypedDict( + error=str(e), ) - _metadata["raw_request"] = ( - "Unable to Log \ + _metadata[ + "raw_request" + ] = "Unable to Log \ raw request: {}".format( - str(e) - ) + str(e) ) if self.logger_fn and callable(self.logger_fn): try: @@ -1085,9 +1087,9 @@ class Logging(LiteLLMLoggingBaseClass): verbose_logger.debug( f"response_cost_failure_debug_information: {debug_info}" ) - self.model_call_details["response_cost_failure_debug_information"] = ( - debug_info - ) + self.model_call_details[ + "response_cost_failure_debug_information" + ] = debug_info return None try: @@ -1112,9 +1114,9 @@ class Logging(LiteLLMLoggingBaseClass): verbose_logger.debug( f"response_cost_failure_debug_information: {debug_info}" ) - self.model_call_details["response_cost_failure_debug_information"] = ( - debug_info - ) + self.model_call_details[ + "response_cost_failure_debug_information" + ] = debug_info return None @@ -1174,9 +1176,9 @@ class Logging(LiteLLMLoggingBaseClass): end_time = datetime.datetime.now() if self.completion_start_time is None: self.completion_start_time = end_time - self.model_call_details["completion_start_time"] = ( - self.completion_start_time - ) + self.model_call_details[ + "completion_start_time" + ] = self.completion_start_time self.model_call_details["log_event_type"] = "successful_api_call" self.model_call_details["end_time"] = end_time self.model_call_details["cache_hit"] = cache_hit @@ -1256,39 +1258,39 @@ class Logging(LiteLLMLoggingBaseClass): "response_cost" ] else: - self.model_call_details["response_cost"] = ( - self._response_cost_calculator(result=logging_result) - ) + self.model_call_details[ + "response_cost" + ] = self._response_cost_calculator(result=logging_result) ## STANDARDIZED LOGGING PAYLOAD - self.model_call_details["standard_logging_object"] = ( - get_standard_logging_object_payload( - kwargs=self.model_call_details, - init_response_obj=logging_result, - start_time=start_time, - end_time=end_time, - logging_obj=self, - status="success", - standard_built_in_tools_params=self.standard_built_in_tools_params, - ) + self.model_call_details[ + "standard_logging_object" + ] = get_standard_logging_object_payload( + kwargs=self.model_call_details, + init_response_obj=logging_result, + start_time=start_time, + end_time=end_time, + logging_obj=self, + status="success", + standard_built_in_tools_params=self.standard_built_in_tools_params, ) elif isinstance(result, dict) or isinstance(result, list): ## STANDARDIZED LOGGING PAYLOAD - self.model_call_details["standard_logging_object"] = ( - get_standard_logging_object_payload( - kwargs=self.model_call_details, - init_response_obj=result, - start_time=start_time, - end_time=end_time, - logging_obj=self, - status="success", - standard_built_in_tools_params=self.standard_built_in_tools_params, - ) + self.model_call_details[ + "standard_logging_object" + ] = get_standard_logging_object_payload( + kwargs=self.model_call_details, + init_response_obj=result, + start_time=start_time, + end_time=end_time, + logging_obj=self, + status="success", + standard_built_in_tools_params=self.standard_built_in_tools_params, ) elif standard_logging_object is not None: - self.model_call_details["standard_logging_object"] = ( - standard_logging_object - ) + self.model_call_details[ + "standard_logging_object" + ] = standard_logging_object else: # streaming chunks + image gen. self.model_call_details["response_cost"] = None @@ -1344,23 +1346,23 @@ class Logging(LiteLLMLoggingBaseClass): verbose_logger.debug( "Logging Details LiteLLM-Success Call streaming complete" ) - self.model_call_details["complete_streaming_response"] = ( - complete_streaming_response - ) - self.model_call_details["response_cost"] = ( - self._response_cost_calculator(result=complete_streaming_response) - ) + self.model_call_details[ + "complete_streaming_response" + ] = complete_streaming_response + self.model_call_details[ + "response_cost" + ] = self._response_cost_calculator(result=complete_streaming_response) ## STANDARDIZED LOGGING PAYLOAD - self.model_call_details["standard_logging_object"] = ( - get_standard_logging_object_payload( - kwargs=self.model_call_details, - init_response_obj=complete_streaming_response, - start_time=start_time, - end_time=end_time, - logging_obj=self, - status="success", - standard_built_in_tools_params=self.standard_built_in_tools_params, - ) + self.model_call_details[ + "standard_logging_object" + ] = get_standard_logging_object_payload( + kwargs=self.model_call_details, + init_response_obj=complete_streaming_response, + start_time=start_time, + end_time=end_time, + logging_obj=self, + status="success", + standard_built_in_tools_params=self.standard_built_in_tools_params, ) callbacks = self.get_combined_callback_list( dynamic_success_callbacks=self.dynamic_success_callbacks, @@ -1680,10 +1682,10 @@ class Logging(LiteLLMLoggingBaseClass): ) else: if self.stream and complete_streaming_response: - self.model_call_details["complete_response"] = ( - self.model_call_details.get( - "complete_streaming_response", {} - ) + self.model_call_details[ + "complete_response" + ] = self.model_call_details.get( + "complete_streaming_response", {} ) result = self.model_call_details["complete_response"] openMeterLogger.log_success_event( @@ -1723,10 +1725,10 @@ class Logging(LiteLLMLoggingBaseClass): ) else: if self.stream and complete_streaming_response: - self.model_call_details["complete_response"] = ( - self.model_call_details.get( - "complete_streaming_response", {} - ) + self.model_call_details[ + "complete_response" + ] = self.model_call_details.get( + "complete_streaming_response", {} ) result = self.model_call_details["complete_response"] @@ -1833,9 +1835,9 @@ class Logging(LiteLLMLoggingBaseClass): if complete_streaming_response is not None: print_verbose("Async success callbacks: Got a complete streaming response") - self.model_call_details["async_complete_streaming_response"] = ( - complete_streaming_response - ) + self.model_call_details[ + "async_complete_streaming_response" + ] = complete_streaming_response try: if self.model_call_details.get("cache_hit", False) is True: self.model_call_details["response_cost"] = 0.0 @@ -1845,10 +1847,10 @@ class Logging(LiteLLMLoggingBaseClass): model_call_details=self.model_call_details ) # base_model defaults to None if not set on model_info - self.model_call_details["response_cost"] = ( - self._response_cost_calculator( - result=complete_streaming_response - ) + self.model_call_details[ + "response_cost" + ] = self._response_cost_calculator( + result=complete_streaming_response ) verbose_logger.debug( @@ -1861,16 +1863,16 @@ class Logging(LiteLLMLoggingBaseClass): self.model_call_details["response_cost"] = None ## STANDARDIZED LOGGING PAYLOAD - self.model_call_details["standard_logging_object"] = ( - get_standard_logging_object_payload( - kwargs=self.model_call_details, - init_response_obj=complete_streaming_response, - start_time=start_time, - end_time=end_time, - logging_obj=self, - status="success", - standard_built_in_tools_params=self.standard_built_in_tools_params, - ) + self.model_call_details[ + "standard_logging_object" + ] = get_standard_logging_object_payload( + kwargs=self.model_call_details, + init_response_obj=complete_streaming_response, + start_time=start_time, + end_time=end_time, + logging_obj=self, + status="success", + standard_built_in_tools_params=self.standard_built_in_tools_params, ) callbacks = self.get_combined_callback_list( dynamic_success_callbacks=self.dynamic_async_success_callbacks, @@ -2076,18 +2078,18 @@ class Logging(LiteLLMLoggingBaseClass): ## STANDARDIZED LOGGING PAYLOAD - self.model_call_details["standard_logging_object"] = ( - get_standard_logging_object_payload( - kwargs=self.model_call_details, - init_response_obj={}, - start_time=start_time, - end_time=end_time, - logging_obj=self, - status="failure", - error_str=str(exception), - original_exception=exception, - standard_built_in_tools_params=self.standard_built_in_tools_params, - ) + self.model_call_details[ + "standard_logging_object" + ] = get_standard_logging_object_payload( + kwargs=self.model_call_details, + init_response_obj={}, + start_time=start_time, + end_time=end_time, + logging_obj=self, + status="failure", + error_str=str(exception), + original_exception=exception, + standard_built_in_tools_params=self.standard_built_in_tools_params, ) return start_time, end_time @@ -2861,9 +2863,9 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915 endpoint=arize_config.endpoint, ) - os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = ( - f"space_key={arize_config.space_key},api_key={arize_config.api_key}" - ) + os.environ[ + "OTEL_EXPORTER_OTLP_TRACES_HEADERS" + ] = f"space_key={arize_config.space_key},api_key={arize_config.api_key}" for callback in _in_memory_loggers: if ( isinstance(callback, ArizeLogger) @@ -2887,9 +2889,9 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915 # auth can be disabled on local deployments of arize phoenix if arize_phoenix_config.otlp_auth_headers is not None: - os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = ( - arize_phoenix_config.otlp_auth_headers - ) + os.environ[ + "OTEL_EXPORTER_OTLP_TRACES_HEADERS" + ] = arize_phoenix_config.otlp_auth_headers for callback in _in_memory_loggers: if ( @@ -2980,9 +2982,9 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915 exporter="otlp_http", endpoint="https://langtrace.ai/api/trace", ) - os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = ( - f"api_key={os.getenv('LANGTRACE_API_KEY')}" - ) + os.environ[ + "OTEL_EXPORTER_OTLP_TRACES_HEADERS" + ] = f"api_key={os.getenv('LANGTRACE_API_KEY')}" for callback in _in_memory_loggers: if ( isinstance(callback, OpenTelemetry) @@ -3526,10 +3528,10 @@ class StandardLoggingPayloadSetup: for key in StandardLoggingHiddenParams.__annotations__.keys(): if key in hidden_params: if key == "additional_headers": - clean_hidden_params["additional_headers"] = ( - StandardLoggingPayloadSetup.get_additional_headers( - hidden_params[key] - ) + clean_hidden_params[ + "additional_headers" + ] = StandardLoggingPayloadSetup.get_additional_headers( + hidden_params[key] ) else: clean_hidden_params[key] = hidden_params[key] # type: ignore @@ -3901,9 +3903,9 @@ def scrub_sensitive_keys_in_metadata(litellm_params: Optional[dict]): ): for k, v in metadata["user_api_key_metadata"].items(): if k == "logging": # prevent logging user logging keys - cleaned_user_api_key_metadata[k] = ( - "scrubbed_by_litellm_for_sensitive_keys" - ) + cleaned_user_api_key_metadata[ + k + ] = "scrubbed_by_litellm_for_sensitive_keys" else: cleaned_user_api_key_metadata[k] = v diff --git a/litellm/proxy/hooks/key_management_event_hooks.py b/litellm/proxy/hooks/key_management_event_hooks.py index 047dd976925..082b506ddd9 100644 --- a/litellm/proxy/hooks/key_management_event_hooks.py +++ b/litellm/proxy/hooks/key_management_event_hooks.py @@ -297,9 +297,10 @@ class KeyManagementEventHooks: @staticmethod async def _send_key_created_email(response: dict): - from enterprise.enterprise_callbacks.send_emails.base_email import ( + from litellm_enterprise.enterprise_callbacks.send_emails.base_email import ( BaseEmailLogger, ) + from litellm.proxy.proxy_server import general_settings, proxy_logging_obj from litellm.types.enterprise.enterprise_callbacks.send_emails import ( SendKeyCreatedEmailEvent, diff --git a/tests/litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py b/tests/litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py index c46912e3a8c..22d553e87c5 100644 --- a/tests/litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py +++ b/tests/litellm/enterprise/enterprise_callbacks/send_emails/test_base_email.py @@ -8,7 +8,10 @@ from fastapi.testclient import TestClient sys.path.insert(0, os.path.abspath("../../..")) -from enterprise.enterprise_callbacks.send_emails.base_email import BaseEmailLogger +from litellm_enterprise.enterprise_callbacks.send_emails.base_email import ( + BaseEmailLogger, +) + from litellm.proxy._types import Litellm_EntityType, WebhookEvent from litellm.types.enterprise.enterprise_callbacks.send_emails import ( SendKeyCreatedEmailEvent, diff --git a/tests/litellm/enterprise/enterprise_callbacks/send_emails/test_endpoints.py b/tests/litellm/enterprise/enterprise_callbacks/send_emails/test_endpoints.py index 2012b728b76..8169c457c24 100644 --- a/tests/litellm/enterprise/enterprise_callbacks/send_emails/test_endpoints.py +++ b/tests/litellm/enterprise/enterprise_callbacks/send_emails/test_endpoints.py @@ -9,7 +9,7 @@ from fastapi.testclient import TestClient sys.path.insert(0, os.path.abspath("../../..")) -from enterprise.enterprise_callbacks.send_emails.endpoints import ( +from litellm_enterprise.enterprise_callbacks.send_emails.endpoints import ( _get_email_settings, _save_email_settings, get_email_event_settings, @@ -17,6 +17,7 @@ from enterprise.enterprise_callbacks.send_emails.endpoints import ( router, update_event_settings, ) + from litellm.types.enterprise.enterprise_callbacks.send_emails import ( DefaultEmailSettings, EmailEvent, @@ -146,7 +147,7 @@ async def test_get_email_event_settings(mock_prisma_client, mock_user_api_key_au # Setup mocks with mock.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): with mock.patch( - "enterprise.enterprise_callbacks.send_emails.endpoints._get_email_settings", + "litellm_enterprise.enterprise_callbacks.send_emails.endpoints._get_email_settings", side_effect=mock_get_settings, ): # Call the endpoint function directly @@ -185,11 +186,11 @@ async def test_update_event_settings(mock_prisma_client, mock_user_api_key_auth) # Setup mocks with mock.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): with mock.patch( - "enterprise.enterprise_callbacks.send_emails.endpoints._get_email_settings", + "litellm_enterprise.enterprise_callbacks.send_emails.endpoints._get_email_settings", side_effect=mock_get_settings, ): with mock.patch( - "enterprise.enterprise_callbacks.send_emails.endpoints._save_email_settings", + "litellm_enterprise.enterprise_callbacks.send_emails.endpoints._save_email_settings", side_effect=mock_save_settings, ): # Create request with updated settings @@ -227,7 +228,7 @@ async def test_reset_event_settings(mock_prisma_client, mock_user_api_key_auth): # Setup mocks with mock.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client): with mock.patch( - "enterprise.enterprise_callbacks.send_emails.endpoints._save_email_settings", + "litellm_enterprise.enterprise_callbacks.send_emails.endpoints._save_email_settings", side_effect=mock_save_settings, ): # Call the endpoint function directly diff --git a/tests/litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py b/tests/litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py index 980fd49b948..dffa343d99f 100644 --- a/tests/litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py +++ b/tests/litellm/enterprise/enterprise_callbacks/send_emails/test_resend_email.py @@ -7,7 +7,9 @@ from httpx import Response sys.path.insert(0, os.path.abspath("../../..")) -from enterprise.enterprise_callbacks.send_emails.resend_email import ResendEmailLogger +from litellm_enterprise.enterprise_callbacks.send_emails.resend_email import ( + ResendEmailLogger, +) @pytest.fixture @@ -19,7 +21,7 @@ def mock_env_vars(): @pytest.fixture def mock_httpx_client(): with mock.patch( - "enterprise.enterprise_callbacks.send_emails.resend_email.get_async_httpx_client" + "litellm_enterprise.enterprise_callbacks.send_emails.resend_email.get_async_httpx_client" ) as mock_client: # Create a mock response mock_response = mock.AsyncMock(spec=Response) diff --git a/tests/logging_callback_tests/test_generic_api_callback.py b/tests/logging_callback_tests/test_generic_api_callback.py index 0bdec437fe9..c033f323ec8 100644 --- a/tests/logging_callback_tests/test_generic_api_callback.py +++ b/tests/logging_callback_tests/test_generic_api_callback.py @@ -28,7 +28,7 @@ from litellm.types.utils import ( ) verbose_logger.setLevel(logging.DEBUG) -from enterprise.enterprise_callbacks.generic_api_callback import GenericAPILogger +from litellm_enterprise.enterprise_callbacks.generic_api_callback import GenericAPILogger diff --git a/tests/logging_callback_tests/test_unit_tests_init_callbacks.py b/tests/logging_callback_tests/test_unit_tests_init_callbacks.py index f6edf71498f..71aece04076 100644 --- a/tests/logging_callback_tests/test_unit_tests_init_callbacks.py +++ b/tests/logging_callback_tests/test_unit_tests_init_callbacks.py @@ -42,8 +42,8 @@ from litellm.integrations.azure_storage.azure_storage import AzureBlobStorageLog from litellm.integrations.agentops import AgentOps from litellm.integrations.humanloop import HumanloopLogger from litellm.proxy.hooks.dynamic_rate_limiter import _PROXY_DynamicRateLimitHandler -from enterprise.enterprise_callbacks.generic_api_callback import GenericAPILogger -from enterprise.enterprise_callbacks.send_emails.resend_email import ResendEmailLogger +from litellm_enterprise.enterprise_callbacks.generic_api_callback import GenericAPILogger +from litellm_enterprise.enterprise_callbacks.send_emails.resend_email import ResendEmailLogger from unittest.mock import patch # clear prometheus collectors / registry