mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
0b633aa9c8
commit
d7f69d6ba1
6 changed files with 347 additions and 145 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue