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 <nate@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-05 22:06:53 -07:00 • committed by GitHub
parent 0b633aa9c8
commit d7f69d6ba1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 347 additions and 145 deletions

View file

@ -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

View file

@ -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

View file

@ -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}")

View file

@ -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)

View file

@ -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

View file

@ -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()