From d7f69d6ba165efb217a80abe9ea923cc780cfc4a Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 5 Oct 2026 22:06:53 -0700 Subject: [PATCH] refactor(logging): load enterprise alerting loggers lazily so import litellm skips proxy types (#44762) * refactor(logging): load enterprise alerting loggers lazily so import litellm skips proxy types Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(logging): cover lazy alerting dispatch and docs lookup Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): resolve litellm.proxy._types lazily on attribute access Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(imports): assert proxy types guard imports the checkout under test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: nate Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../enterprise_callbacks/callback_controls.py | 2 +- litellm/litellm_core_utils/litellm_logging.py | 84 ++++--- litellm/proxy/__init__.py | 9 + .../test_proxy_types_import.py | 218 +++++++++--------- .../test_callback_controls.py | 12 + .../test_litellm_logging.py | 167 +++++++++++++- 6 files changed, 347 insertions(+), 145 deletions(-) diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/callback_controls.py b/enterprise/litellm_enterprise/enterprise_callbacks/callback_controls.py index 8824f4c02de..bf99e1a72ca 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/callback_controls.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/callback_controls.py @@ -7,7 +7,6 @@ from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.llm_request_utils import ( get_proxy_server_request_headers, ) -from litellm.proxy._types import CommonProxyErrors from litellm.types.utils import StandardCallbackDynamicParams @@ -86,6 +85,7 @@ class EnterpriseCallbackControls: @staticmethod def _should_allow_dynamic_callback_disabling(): import litellm + from litellm.proxy._types import CommonProxyErrors from litellm.proxy.proxy_server import premium_user # Check if admin has disabled this feature diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index b2c22880f5a..15c53990d7c 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -3,6 +3,7 @@ # Logging function -> log the exact model details + what's being sent | Non-Blocking import copy import datetime +import functools import json import os import re @@ -14,7 +15,7 @@ from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence from datetime import datetime as dt_object from functools import lru_cache from types import MappingProxyType, TracebackType -from typing import TYPE_CHECKING, Any, Final, Literal, Union, cast +from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, Union, cast from httpx import Response from pydantic import BaseModel, JsonValue @@ -240,18 +241,6 @@ try: from litellm_enterprise.enterprise_callbacks.callback_controls import ( EnterpriseCallbackControls, ) - from litellm_enterprise.enterprise_callbacks.pagerduty.pagerduty import ( - PagerDutyAlerting, - ) - from litellm_enterprise.enterprise_callbacks.send_emails.resend_email import ( - ResendEmailLogger, - ) - from litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email import ( - SendGridEmailLogger, - ) - from litellm_enterprise.enterprise_callbacks.send_emails.smtp_email import ( - SMTPEmailLogger, - ) from litellm_enterprise.litellm_core_utils.litellm_logging import ( StandardLoggingPayloadSetup as EnterpriseStandardLoggingPayloadSetup, ) @@ -264,28 +253,41 @@ try: except Exception as e: verbose_logger.debug("[Non-Blocking] Unable to import GenericAPILogger - LiteLLM Enterprise Feature - %s", e) GenericAPILogger = CustomLogger - ResendEmailLogger = CustomLogger - SendGridEmailLogger = CustomLogger - SMTPEmailLogger = CustomLogger - PagerDutyAlerting = CustomLogger EnterpriseCallbackControls = None EnterpriseStandardLoggingPayloadSetupVAR = None + + +class _EnterpriseAlertingLoggers(NamedTuple): + pagerduty: type[CustomLogger] + resend_email: type[CustomLogger] + sendgrid_email: type[CustomLogger] + smtp_email: type[CustomLogger] + + +@functools.cache +def _enterprise_alerting_loggers() -> _EnterpriseAlertingLoggers: + try: + from litellm_enterprise.enterprise_callbacks.pagerduty.pagerduty import PagerDutyAlerting + from litellm_enterprise.enterprise_callbacks.send_emails.resend_email import ResendEmailLogger + from litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email import SendGridEmailLogger + from litellm_enterprise.enterprise_callbacks.send_emails.smtp_email import SMTPEmailLogger + except Exception as e: + verbose_logger.debug( + "[Non-Blocking] Unable to import enterprise alerting loggers - LiteLLM Enterprise Feature - %s", + e, + ) + return _EnterpriseAlertingLoggers(CustomLogger, CustomLogger, CustomLogger, CustomLogger) + return _EnterpriseAlertingLoggers(PagerDutyAlerting, ResendEmailLogger, SendGridEmailLogger, SMTPEmailLogger) + + if TYPE_CHECKING: from litellm.integrations.generic_api.generic_api_callback import ( GenericAPILogger as _GenericAPILoggerCls, ) _GENERIC_API_LOGGER_CLS: Final = _GenericAPILoggerCls - _RESEND_EMAIL_LOGGER_FACTORY: Final = CustomLogger - _SENDGRID_EMAIL_LOGGER_FACTORY: Final = CustomLogger - _SMTP_EMAIL_LOGGER_FACTORY: Final = CustomLogger - _PAGERDUTY_ALERTING_FACTORY: Final = CustomLogger else: _GENERIC_API_LOGGER_CLS: Final = GenericAPILogger - _RESEND_EMAIL_LOGGER_FACTORY: Final = ResendEmailLogger - _SENDGRID_EMAIL_LOGGER_FACTORY: Final = SendGridEmailLogger - _SMTP_EMAIL_LOGGER_FACTORY: Final = SMTPEmailLogger - _PAGERDUTY_ALERTING_FACTORY: Final = PagerDutyAlerting _in_memory_loggers: Final[list[CustomLogger]] = [] _STANDARD_LOGGING_METADATA_RESOLVED_KEYS: Final[frozenset[str]] = frozenset(("used_client_oauth_token",)) @@ -5062,10 +5064,11 @@ def _init_custom_logger_compatible_class( _in_memory_loggers.append(_otel_logger) return _otel_logger elif logging_integration == "pagerduty": + pagerduty_loggers: Final = _enterprise_alerting_loggers() for callback in _in_memory_loggers: - if isinstance(callback, PagerDutyAlerting): + if isinstance(callback, pagerduty_loggers.pagerduty): return callback - pagerduty_logger: Final = _PAGERDUTY_ALERTING_FACTORY(**custom_logger_init_args) + pagerduty_logger: Final = pagerduty_loggers.pagerduty(**custom_logger_init_args) _in_memory_loggers.append(pagerduty_logger) return pagerduty_logger elif logging_integration == "anthropic_cache_control_hook": @@ -5101,24 +5104,27 @@ def _init_custom_logger_compatible_class( _in_memory_loggers.append(generic_api_logger) return generic_api_logger elif logging_integration == "resend_email": + resend_email_loggers: Final = _enterprise_alerting_loggers() for callback in _in_memory_loggers: - if isinstance(callback, ResendEmailLogger): + if isinstance(callback, resend_email_loggers.resend_email): return callback - resend_email_logger: Final = _RESEND_EMAIL_LOGGER_FACTORY() + resend_email_logger: Final = resend_email_loggers.resend_email() _in_memory_loggers.append(resend_email_logger) return resend_email_logger elif logging_integration == "sendgrid_email": + sendgrid_email_loggers: Final = _enterprise_alerting_loggers() for callback in _in_memory_loggers: - if isinstance(callback, SendGridEmailLogger): + if isinstance(callback, sendgrid_email_loggers.sendgrid_email): return callback - sendgrid_email_logger: Final = _SENDGRID_EMAIL_LOGGER_FACTORY() + sendgrid_email_logger: Final = sendgrid_email_loggers.sendgrid_email() _in_memory_loggers.append(sendgrid_email_logger) return sendgrid_email_logger elif logging_integration == "smtp_email": + smtp_email_loggers: Final = _enterprise_alerting_loggers() for callback in _in_memory_loggers: - if isinstance(callback, SMTPEmailLogger): + if isinstance(callback, smtp_email_loggers.smtp_email): return callback - smtp_email_logger: Final = _SMTP_EMAIL_LOGGER_FACTORY() + smtp_email_logger: Final = smtp_email_loggers.smtp_email() _in_memory_loggers.append(smtp_email_logger) return smtp_email_logger elif logging_integration == "humanloop": @@ -5507,8 +5513,9 @@ def get_custom_logger_compatible_class( if isinstance(callback, MlflowLogger): return callback elif logging_integration == "pagerduty": + pagerduty_loggers: Final = _enterprise_alerting_loggers() for callback in _in_memory_loggers: - if isinstance(callback, PagerDutyAlerting): + if isinstance(callback, pagerduty_loggers.pagerduty): return callback elif logging_integration == "anthropic_cache_control_hook": for callback in _in_memory_loggers: @@ -5531,16 +5538,19 @@ def get_custom_logger_compatible_class( if isinstance(callback, _GENERIC_API_LOGGER_CLS): return callback elif logging_integration == "resend_email": + resend_email_loggers: Final = _enterprise_alerting_loggers() for callback in _in_memory_loggers: - if isinstance(callback, ResendEmailLogger): + if isinstance(callback, resend_email_loggers.resend_email): return callback elif logging_integration == "sendgrid_email": + sendgrid_email_loggers: Final = _enterprise_alerting_loggers() for callback in _in_memory_loggers: - if isinstance(callback, SendGridEmailLogger): + if isinstance(callback, sendgrid_email_loggers.sendgrid_email): return callback elif logging_integration == "smtp_email": + smtp_email_loggers: Final = _enterprise_alerting_loggers() for callback in _in_memory_loggers: - if isinstance(callback, SMTPEmailLogger): + if isinstance(callback, smtp_email_loggers.smtp_email): return callback elif logging_integration == "newrelic": from litellm.integrations.otel.logger import OpenTelemetryV2 diff --git a/litellm/proxy/__init__.py b/litellm/proxy/__init__.py index b6e690fd591..addeeadc9f8 100644 --- a/litellm/proxy/__init__.py +++ b/litellm/proxy/__init__.py @@ -1 +1,10 @@ +import importlib +from types import ModuleType + from . import * + + +def __getattr__(name: str) -> ModuleType: + if name == "_types": + return importlib.import_module("litellm.proxy._types") + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/tests/code_coverage_tests/test_proxy_types_import.py b/tests/code_coverage_tests/test_proxy_types_import.py index 9a7936198f3..9e493badec4 100644 --- a/tests/code_coverage_tests/test_proxy_types_import.py +++ b/tests/code_coverage_tests/test_proxy_types_import.py @@ -1,114 +1,122 @@ -import ast -import os +import subprocess import sys +from pathlib import Path +from typing import Final + +import pytest -def test_proxy_types_not_imported(): +_REPO_ROOT: Final = Path(__file__).resolve().parents[2] + + +def _run_import_check( + block_enterprise: bool, repo_root: Path +) -> subprocess.CompletedProcess[str]: + program: Final = f""" +import importlib.abc +import pathlib +import sys +import traceback + +repo_root = sys.argv[1] + +class ImportTracer(importlib.abc.MetaPathFinder): + def find_spec(self, name, path, target=None): + if {block_enterprise!r} and name.startswith("litellm_enterprise"): + raise ImportError("blocked litellm_enterprise imports") + if name == "litellm.proxy._types": + repository_frames = [ + frame for frame in traceback.extract_stack() + if repo_root in frame.filename + ] + print( + "FIRST _types importer chain:", + " | ".join( + f"{{frame.filename}}:{{frame.lineno}}" + for frame in repository_frames[-8:] + ), + ) + return None + +sys.meta_path.insert(0, ImportTracer()) +import litellm +if not pathlib.Path(litellm.__file__).resolve().is_relative_to(pathlib.Path(repo_root)): + raise SystemExit(f"imported litellm from {{litellm.__file__}}, expected under {{repo_root}}") +print("litellm.__file__:", litellm.__file__) +print("_types loaded:", "litellm.proxy._types" in sys.modules) """ - Test that proxy._types is not directly imported in litellm/__init__.py - by examining the source code using AST parsing. - """ - # Read the litellm/__init__.py file - # local_init_file = "../litellm/" - init_file_path = os.path.join("./litellm", "__init__.py") - if not os.path.exists(init_file_path): - raise Exception(f"Could not find {init_file_path}") - - with open(init_file_path, "r") as f: - content = f.read() - lines = content.splitlines() # Get lines for line number reporting - - try: - tree = ast.parse(content) - except SyntaxError as e: - raise Exception(f"Could not parse {init_file_path}: {e}") - - # Check for direct imports of proxy._types - found_imports = [] - - for node in ast.walk(tree): - if isinstance(node, ast.Import): - for alias in node.names: - if "proxy._types" in alias.name or "proxy/_types" in alias.name: - line_num = node.lineno - line_content = ( - lines[line_num - 1] if line_num <= len(lines) else "Unknown" - ) - import_statement = f"import {alias.name}" - found_imports.append( - { - "type": "import", - "line": line_num, - "content": line_content.strip(), - "statement": import_statement, - "module": alias.name, - } - ) - - elif isinstance(node, ast.ImportFrom): - if node.module and ( - "proxy._types" in node.module or "proxy/_types" in node.module - ): - line_num = node.lineno - line_content = ( - lines[line_num - 1] if line_num <= len(lines) else "Unknown" - ) - import_names = [alias.name for alias in node.names] - import_statement = ( - f"from {node.module} import {', '.join(import_names)}" - ) - found_imports.append( - { - "type": "from_import", - "line": line_num, - "content": line_content.strip(), - "statement": import_statement, - "module": node.module, - } - ) - - if found_imports: - print( - "āŒ BAD, this can import time to import litellm. Found direct imports of proxy._types in litellm/__init__.py:" - ) - print("=" * 80) - for imp in found_imports: - print(f"Line {imp['line']}: {imp['content']}") - print(f" Type: {imp['type']}") - print(f" Statement: {imp['statement']}") - print(f" Module: {imp['module']}") - print("-" * 80) - print("To fix this, please conditionally import this TYPE using TYPE_CHECKING") - - raise Exception( - f"Found {len(found_imports)} direct import(s) of proxy._types in litellm/__init__.py" - ) - - print("āœ“ No direct imports of proxy._types found in litellm/__init__.py") - return True - - -def main(): - """ - Main function to run the import test - """ - print("=" * 60) - print("Testing litellm import performance") - print( - "Checking that proxy._types is not directly imported from litellm/__init__.py" + return subprocess.run( + [sys.executable, "-I", "-c", program, str(repo_root.resolve())], + cwd=repo_root, + capture_output=True, + text=True, + check=False, ) - print("=" * 60) - try: - test_proxy_types_not_imported() - print("\n" + "=" * 60) - print( - "āœ“ Test passed! proxy._types is not directly imported from litellm/__init__.py" - ) - print("=" * 60) - except Exception as e: - print(f"\nāŒ Test failed: {e}") - print("=" * 60) + +def _run_proxy_types_attribute_access(repo_root: Path) -> subprocess.CompletedProcess[str]: + program: Final = """ +import pathlib +import sys +import litellm + +repo_root = sys.argv[1] +if not pathlib.Path(litellm.__file__).resolve().is_relative_to(pathlib.Path(repo_root)): + raise SystemExit(f"imported litellm from {litellm.__file__}, expected under {repo_root}") + +assert "litellm.proxy._types" not in sys.modules +print("_types loaded before attribute access: False") +user_api_key_auth = litellm.proxy._types.UserAPIKeyAuth +assert "litellm.proxy._types" in sys.modules +proxy_types = sys.modules["litellm.proxy._types"] +assert user_api_key_auth is proxy_types.UserAPIKeyAuth +print("_types loaded after attribute access: True") +print("UserAPIKeyAuth identity: True") + """ + return subprocess.run( + [sys.executable, "-I", "-c", program, str(repo_root.resolve())], + cwd=repo_root, + capture_output=True, + text=True, + check=False, + ) + + +@pytest.mark.parametrize("block_enterprise", (False, True)) +def test_import_litellm_does_not_load_proxy_types(block_enterprise: bool) -> None: + result: Final = _run_import_check(block_enterprise, _REPO_ROOT) + assert result.returncode == 0, result.stdout + result.stderr + assert "_types loaded: False" in result.stdout, result.stdout + result.stderr + + +def test_proxy_types_attribute_access_still_works() -> None: + result: Final = _run_proxy_types_attribute_access(_REPO_ROOT) + assert result.returncode == 0, result.stdout + result.stderr + assert "_types loaded before attribute access: False" in result.stdout + assert "_types loaded after attribute access: True" in result.stdout + assert "UserAPIKeyAuth identity: True" in result.stdout + + +def main(repo_root: Path = _REPO_ROOT) -> None: + results: Final = ( + ("enterprise blocked: False", _run_import_check(False, repo_root), "_types loaded: False"), + ("enterprise blocked: True", _run_import_check(True, repo_root), "_types loaded: False"), + ( + "proxy._types attribute access", + _run_proxy_types_attribute_access(repo_root), + "UserAPIKeyAuth identity: True", + ), + ) + for label, result, expected_output in results: + print(label) + print(result.stdout, end="") + if result.stderr: + print(result.stderr, file=sys.stderr, end="") + + if any( + result.returncode != 0 or expected_output not in result.stdout + for _, result, expected_output in results + ): sys.exit(1) diff --git a/tests/unit/enterprise/enterprise_callbacks/test_callback_controls.py b/tests/unit/enterprise/enterprise_callbacks/test_callback_controls.py index d67dc3cf6bc..063e9e0ba7f 100644 --- a/tests/unit/enterprise/enterprise_callbacks/test_callback_controls.py +++ b/tests/unit/enterprise/enterprise_callbacks/test_callback_controls.py @@ -412,3 +412,15 @@ class TestEnterpriseCallbackControls: "langfuse", litellm_params, standard_callback_dynamic_params ) assert result is True + + def test_non_premium_dynamic_callback_warning_uses_common_proxy_error(self, caplog): + from litellm.proxy._types import CommonProxyErrors + + caplog.set_level("WARNING") + with patch("litellm.allow_dynamic_callback_disabling", True): + with patch("litellm.proxy.proxy_server.premium_user", False): + assert ( + EnterpriseCallbackControls._should_allow_dynamic_callback_disabling() + is False + ) + assert CommonProxyErrors.not_premium_user.value in caplog.text diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index 9cdbd2a5d31..b42a6f99ee5 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -2,13 +2,15 @@ import asyncio import contextlib import copy import datetime +import importlib.abc import json import logging import os import sys import time -from collections.abc import Callable, Iterator, Mapping -from types import MappingProxyType +from collections.abc import Callable, Iterator, Mapping, Sequence +from importlib.machinery import ModuleSpec +from types import MappingProxyType, ModuleType from typing import Final, Literal from unittest.mock import AsyncMock, MagicMock, patch @@ -9373,3 +9375,164 @@ 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() + + +def test_enterprise_alerting_loggers_resolves_real_classes(): + from litellm_enterprise.enterprise_callbacks.pagerduty.pagerduty import PagerDutyAlerting + from litellm_enterprise.enterprise_callbacks.send_emails.resend_email import ResendEmailLogger + from litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email import SendGridEmailLogger + from litellm_enterprise.enterprise_callbacks.send_emails.smtp_email import SMTPEmailLogger + from litellm.litellm_core_utils import litellm_logging as logging_module + + logging_module._enterprise_alerting_loggers.cache_clear() + try: + loggers = logging_module._enterprise_alerting_loggers() + assert loggers.pagerduty is PagerDutyAlerting + assert loggers.resend_email is ResendEmailLogger + assert loggers.sendgrid_email is SendGridEmailLogger + assert loggers.smtp_email is SMTPEmailLogger + finally: + logging_module._enterprise_alerting_loggers.cache_clear() + + +def test_enterprise_alerting_loggers_falls_back_without_disabling_callback_controls( + monkeypatch, +): + from litellm_enterprise.enterprise_callbacks.callback_controls import ( + EnterpriseCallbackControls, + ) + from litellm.integrations.custom_logger import CustomLogger + from litellm.litellm_core_utils import litellm_logging as logging_module + + class BlockPagerDutyFinder(importlib.abc.MetaPathFinder): + def find_spec( + self, + fullname: str, + path: Sequence[str] | None = None, + target: ModuleType | None = None, + ) -> ModuleSpec | None: + if fullname == "litellm_enterprise.enterprise_callbacks.pagerduty.pagerduty": + raise ImportError("blocked pagerduty import") + return None + + module_prefixes: Final = ( + "litellm_enterprise.enterprise_callbacks.pagerduty", + "litellm_enterprise.enterprise_callbacks.send_emails", + ) + logging_module._enterprise_alerting_loggers.cache_clear() + for module_name in tuple(sys.modules): + if any( + module_name == prefix or module_name.startswith(f"{prefix}.") + for prefix in module_prefixes + ): + monkeypatch.delitem(sys.modules, module_name, raising=False) + monkeypatch.setattr(sys, "meta_path", [BlockPagerDutyFinder(), *sys.meta_path]) + + try: + loggers = logging_module._enterprise_alerting_loggers() + assert loggers.pagerduty is CustomLogger + assert loggers.resend_email is CustomLogger + assert loggers.sendgrid_email is CustomLogger + assert loggers.smtp_email is CustomLogger + assert logging_module.EnterpriseCallbackControls is EnterpriseCallbackControls + finally: + logging_module._enterprise_alerting_loggers.cache_clear() + + +@pytest.mark.parametrize("integration", ("smtp_email", "resend_email", "sendgrid_email")) +def test_init_email_alerting_logger_reuses_instance( + integration: Literal["smtp_email", "resend_email", "sendgrid_email"], + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.litellm_core_utils import litellm_logging as logging_module + from litellm_enterprise.enterprise_callbacks.send_emails.resend_email import ResendEmailLogger + from litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email import SendGridEmailLogger + from litellm_enterprise.enterprise_callbacks.send_emails.smtp_email import SMTPEmailLogger + + logger_class: Final = ( + SMTPEmailLogger + if integration == "smtp_email" + else ResendEmailLogger + if integration == "resend_email" + else SendGridEmailLogger + ) + monkeypatch.setenv("RESEND_API_KEY", "test-resend-key") + monkeypatch.setenv("SENDGRID_API_KEY", "test-sendgrid-key") + logging_module._in_memory_loggers.clear() + try: + logger = logging_module._init_custom_logger_compatible_class( + logging_integration=integration, + internal_usage_cache=None, + llm_router=None, + custom_logger_init_args={}, + ) + assert isinstance(logger, logger_class) + assert ( + logging_module._init_custom_logger_compatible_class( + logging_integration=integration, + internal_usage_cache=None, + llm_router=None, + custom_logger_init_args={}, + ) + is logger + ) + assert logging_module.get_custom_logger_compatible_class(integration) is logger + finally: + logging_module._in_memory_loggers.clear() + + +def test_init_pagerduty_logger_reuses_instance(monkeypatch): + from litellm.types.integrations.pagerduty import AlertingConfig + from litellm_enterprise.enterprise_callbacks.pagerduty.pagerduty import PagerDutyAlerting + from litellm.litellm_core_utils import litellm_logging as logging_module + + monkeypatch.setenv("PAGERDUTY_API_KEY", "test-pagerduty-key") + logging_module._in_memory_loggers.clear() + try: + logger = logging_module._init_custom_logger_compatible_class( + logging_integration="pagerduty", + internal_usage_cache=None, + llm_router=None, + custom_logger_init_args={ + "alerting_args": AlertingConfig( + failure_threshold=1, + failure_threshold_window_seconds=10, + ) + }, + ) + assert isinstance(logger, PagerDutyAlerting) + assert ( + logging_module._init_custom_logger_compatible_class( + logging_integration="pagerduty", + internal_usage_cache=None, + llm_router=None, + custom_logger_init_args={ + "alerting_args": AlertingConfig( + failure_threshold=1, + failure_threshold_window_seconds=10, + ) + }, + ) + is logger + ) + assert logging_module.get_custom_logger_compatible_class("pagerduty") is logger + finally: + logging_module._in_memory_loggers.clear() + + +@pytest.mark.parametrize("integration", ("smtp_email", "pagerduty", "resend_email", "sendgrid_email")) +def test_get_custom_logger_compatible_class_does_not_match_generic_api_logger( + integration: Literal["smtp_email", "pagerduty", "resend_email", "sendgrid_email"], + monkeypatch: pytest.MonkeyPatch, +): + from litellm.integrations.generic_api.generic_api_callback import GenericAPILogger + from litellm.litellm_core_utils import litellm_logging as logging_module + + monkeypatch.setenv("GENERIC_LOGGER_ENDPOINT", "https://generic-logger.test") + logging_module._in_memory_loggers.clear() + try: + with patch("asyncio.create_task", side_effect=lambda coro: coro.close()): + logging_module._in_memory_loggers.append(GenericAPILogger()) + assert logging_module.get_custom_logger_compatible_class(integration) is None + finally: + logging_module._in_memory_loggers.clear()