diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 4609d57ff15..ae559f30857 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -26,6 +26,7 @@ from collections.abc import ( MutableMapping, Sequence, ) +from dataclasses import dataclass from datetime import datetime, timedelta, timezone from itertools import chain from types import MappingProxyType, UnionType @@ -5104,6 +5105,61 @@ def _log_ignored_cost_map_copy(model_id: str, fields: tuple[str, ...]) -> None: ) +_DB_CONFIG_CALLBACK_EVENT_TYPES: Final[Mapping[str, tuple[Literal["success", "failure"], ...]]] = MappingProxyType( + {"success_callback": ("success",), "failure_callback": ("failure",), "callbacks": ("success", "failure")} +) +_CALLBACK_LIST_NAMES: Final = ( + "callbacks", + "success_callback", + "failure_callback", + "_async_success_callback", + "_async_failure_callback", +) + + +def _callback_list_entries(list_name: str) -> tuple[tuple[str, object], ...]: + return tuple((list_name, entry) for entry in getattr(litellm, list_name)) + + +def _registered_callback_entries() -> tuple[tuple[str, object], ...]: + return tuple(chain.from_iterable(_callback_list_entries(list_name) for list_name in _CALLBACK_LIST_NAMES)) + + +def _entries_missing_from( + entries: tuple[tuple[str, object], ...], others: tuple[tuple[str, object], ...] +) -> tuple[tuple[str, object], ...]: + other_ids: Final = frozenset((list_name, id(entry)) for list_name, entry in others) + return tuple((list_name, entry) for list_name, entry in entries if (list_name, id(entry)) not in other_ids) + + +@dataclass(frozen=True, slots=True) +class _DbCallbackRegistration: + added: tuple[tuple[str, object], ...] = () + displaced: tuple[tuple[str, object], ...] = () + + +def _restore_displaced_callback(list_name: str, entry: object) -> None: + callback_list: Final = getattr(litellm, list_name) + if entry not in callback_list: + callback_list.append(entry) + + +def _unregister_db_callback(registration: _DbCallbackRegistration) -> None: + for list_name, entry in registration.added: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + getattr(litellm, list_name), entry, require_self=False + ) + for list_name, entry in registration.displaced: + _restore_displaced_callback(list_name, entry) + + +def _configured_db_callbacks(litellm_settings: Mapping[str, object], setting_key: str) -> tuple[tuple[str, str], ...]: + callbacks: Final = litellm_settings.get(setting_key) + if not isinstance(callbacks, list): + return () + return tuple((setting_key, callback) for callback in callbacks if isinstance(callback, str)) + + class ProxyConfig: """ Abstraction class on top of config loading/updating logic. Gives us one place to control all config updating logic. @@ -5136,6 +5192,7 @@ class ProxyConfig: self.litellm_settings: Final[SettingsStore] = SettingsStore("litellm_settings") self.environment_variables: Final[SettingsStore] = SettingsStore("environment_variables") self._warned_shadowed_keys: frozenset[tuple[Section, str]] = frozenset() + self._db_config_callback_entries: Mapping[tuple[str, str], _DbCallbackRegistration] = MappingProxyType({}) self._settings_stores: Final[Mapping[Section, SettingsStore]] = MappingProxyType( { "general_settings": self.settings, @@ -7160,7 +7217,7 @@ class ProxyConfig: still_desired_ids: frozenset[str] | None = None # Load config separately so a timeout here doesn't block model loading - config_data: dict = {} + config_data: dict | None = None search_tools = None try: config_data = await proxy_config.get_config() @@ -7219,8 +7276,8 @@ class ProxyConfig: if llm_router is not None: llm_model_list = llm_router.get_model_list() - # check if user set any callbacks in Config Table - self._add_callbacks_from_db_config(config_data) + if config_data is not None: + self._add_callbacks_from_db_config(config_data) # router settings await self._add_router_settings_from_db_config(llm_router=llm_router, prisma_client=prisma_client) @@ -7253,37 +7310,34 @@ class ProxyConfig: litellm.logging_callback_manager.add_litellm_callback(callback) def _add_callbacks_from_db_config(self, config_data: dict) -> None: - """ - Adds callbacks from DB config to litellm - """ litellm_settings: Final = config_data.get("litellm_settings", {}) or {} - success_callbacks: Final = litellm_settings.get("success_callback", None) - failure_callbacks: Final = litellm_settings.get("failure_callback", None) - callbacks: Final = litellm_settings.get("callbacks", None) + configured: Final = tuple( + chain.from_iterable( + _configured_db_callbacks(litellm_settings, setting_key) + for setting_key in _DB_CONFIG_CALLBACK_EVENT_TYPES + ) + ) + still_configured: Final = frozenset(configured) + for key, registration in self._db_config_callback_entries.items(): + if key not in still_configured: + _unregister_db_callback(registration) + self._db_config_callback_entries = MappingProxyType( + {key: self._register_db_config_callback(*key) for key in dict.fromkeys(configured)} + ) - if success_callbacks is not None and isinstance(success_callbacks, list): - for success_callback in success_callbacks: - self._add_callback_from_db_to_in_memory_litellm_callbacks( - callback=success_callback, - event_types=["success"], - existing_callbacks=litellm.success_callback, - ) - - if failure_callbacks is not None and isinstance(failure_callbacks, list): - for failure_callback in failure_callbacks: - self._add_callback_from_db_to_in_memory_litellm_callbacks( - callback=failure_callback, - event_types=["failure"], - existing_callbacks=litellm.failure_callback, - ) - - if callbacks is not None and isinstance(callbacks, list): - for callback in callbacks: - self._add_callback_from_db_to_in_memory_litellm_callbacks( - callback=callback, - event_types=["success", "failure"], - existing_callbacks=litellm.callbacks, - ) + def _register_db_config_callback(self, setting_key: str, callback: str) -> _DbCallbackRegistration: + previous: Final = self._db_config_callback_entries.get((setting_key, callback), _DbCallbackRegistration()) + before: Final = _registered_callback_entries() + self._add_callback_from_db_to_in_memory_litellm_callbacks( + callback=callback, + event_types=list(_DB_CONFIG_CALLBACK_EVENT_TYPES[setting_key]), + existing_callbacks=getattr(litellm, setting_key), + ) + after: Final = _registered_callback_entries() + return _DbCallbackRegistration( + added=previous.added + _entries_missing_from(after, before), + displaced=previous.displaced + _entries_missing_from(before, after), + ) def _encrypt_env_variables(self, environment_variables: dict, new_encryption_key: str | None = None) -> dict: """ diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index d5b57fcc309..df05cf0987e 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -11997,6 +11997,100 @@ async def test_db_stored_datadog_redaction_settings_apply_before_logger_init(mon assert litellm.turn_off_message_logging is True +def _reset_runtime_callbacks(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.litellm_core_utils import litellm_logging + + for list_name in ("success_callback", "_async_success_callback", "failure_callback", "_async_failure_callback"): + monkeypatch.setattr(litellm, list_name, []) + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm_logging, "_in_memory_loggers", []) + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk-test") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-test") + monkeypatch.setenv("HUMANLOOP_API_KEY", "test-key") + + +def _runtime_callback_names() -> frozenset[str]: + manager = litellm.logging_callback_manager + return frozenset(manager._get_callback_string(callback) for callback in manager._get_all_callbacks()) + + +@pytest.mark.parametrize("setting_key", ["success_callback", "failure_callback", "callbacks"]) +@pytest.mark.parametrize("callback_name", ["langfuse_otel", "helicone"]) +def test_db_config_sync_unregisters_a_callback_the_stored_config_no_longer_lists( + monkeypatch: pytest.MonkeyPatch, setting_key: str, callback_name: str +): + import litellm.proxy.proxy_server as ps + + _reset_runtime_callbacks(monkeypatch) + pc = ps.ProxyConfig() + + for _ in range(2): + pc._add_callbacks_from_db_config({"litellm_settings": {setting_key: [callback_name]}}) + assert callback_name in _runtime_callback_names() + + pc._add_callbacks_from_db_config({"litellm_settings": {setting_key: []}}) + assert callback_name not in _runtime_callback_names() + + +def test_db_config_sync_keeps_callbacks_it_did_not_register(monkeypatch: pytest.MonkeyPatch): + import litellm.proxy.proxy_server as ps + from litellm.utils import _add_custom_logger_callback_to_specific_event + + _reset_runtime_callbacks(monkeypatch) + _add_custom_logger_callback_to_specific_event("langfuse_otel", "success") + litellm.logging_callback_manager.add_litellm_success_callback("helicone") + pc = ps.ProxyConfig() + + pc._add_callbacks_from_db_config( + {"litellm_settings": {"success_callback": ["langfuse_otel", "helicone", "humanloop", "supabase"]}} + ) + assert {"humanloop", "supabase"} <= _runtime_callback_names() + + pc._add_callbacks_from_db_config({"litellm_settings": {"success_callback": []}}) + remaining: Final = _runtime_callback_names() + assert {"langfuse_otel", "helicone"} <= remaining + assert not {"humanloop", "supabase"} & remaining + + +def test_db_config_sync_restores_a_code_callback_it_replaced(monkeypatch: pytest.MonkeyPatch): + import litellm.proxy.proxy_server as ps + + _reset_runtime_callbacks(monkeypatch) + litellm.logging_callback_manager.add_litellm_success_callback("langfuse_otel") + pc = ps.ProxyConfig() + + pc._add_callbacks_from_db_config({"litellm_settings": {"success_callback": ["langfuse_otel"]}}) + assert "langfuse_otel" not in litellm.success_callback + assert "langfuse_otel" in _runtime_callback_names() + + pc._add_callbacks_from_db_config({"litellm_settings": {"success_callback": []}}) + assert litellm.success_callback == ["langfuse_otel"] + + +@pytest.mark.asyncio +async def test_failed_config_load_keeps_callbacks_the_stored_config_registered(monkeypatch: pytest.MonkeyPatch): + import litellm.proxy.proxy_server as ps + + _reset_runtime_callbacks(monkeypatch) + pc = ps.ProxyConfig() + monkeypatch.setattr(ps, "proxy_config", pc) + monkeypatch.setattr(ps, "llm_router", None) + monkeypatch.setattr(ps, "master_key", "sk-1234") + monkeypatch.setattr( + pc, "get_config", AsyncMock(return_value={"litellm_settings": {"success_callback": ["helicone"]}}) + ) + await pc._update_llm_router(new_models=[], proxy_logging_obj=MagicMock()) + assert "helicone" in _runtime_callback_names() + + monkeypatch.setattr(pc, "get_config", AsyncMock(side_effect=TimeoutError("config read timed out"))) + await pc._update_llm_router(new_models=[], proxy_logging_obj=MagicMock()) + assert "helicone" in _runtime_callback_names() + + monkeypatch.setattr(pc, "get_config", AsyncMock(return_value={"litellm_settings": {"success_callback": []}})) + await pc._update_llm_router(new_models=[], proxy_logging_obj=MagicMock()) + assert "helicone" not in _runtime_callback_names() + + @pytest.mark.parametrize( "field_name", [