diff --git a/litellm/integrations/compression_interception/handler.py b/litellm/integrations/compression_interception/handler.py index c6ae7d9e82..8899089500 100644 --- a/litellm/integrations/compression_interception/handler.py +++ b/litellm/integrations/compression_interception/handler.py @@ -72,8 +72,13 @@ class CompressionInterceptionLogger(CustomLogger): compression_params: CompressionInterceptionConfig = {} if "compression_interception_params" in litellm_settings: compression_params = litellm_settings["compression_interception_params"] - elif "compression_interception" in callback_specific_params: - compression_params = callback_specific_params["compression_interception"] + elif "compression_interception" in callback_specific_params and isinstance( + callback_specific_params["compression_interception"], dict + ): + compression_params = cast( + CompressionInterceptionConfig, + callback_specific_params["compression_interception"], + ) return CompressionInterceptionLogger.from_config_yaml(compression_params) async def async_pre_call_deployment_hook( diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 37528e7dcd..79f9b16bba 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -1339,8 +1339,13 @@ class WebSearchInterceptionLogger(CustomLogger): websearch_params: WebSearchInterceptionConfig = {} if "websearch_interception_params" in litellm_settings: websearch_params = litellm_settings["websearch_interception_params"] - elif "websearch_interception" in callback_specific_params: - websearch_params = callback_specific_params["websearch_interception"] + elif "websearch_interception" in callback_specific_params and isinstance( + callback_specific_params["websearch_interception"], dict + ): + websearch_params = cast( + WebSearchInterceptionConfig, + callback_specific_params["websearch_interception"], + ) # Use classmethod to initialize from config return WebSearchInterceptionLogger.from_config_yaml(websearch_params) diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index a65e737f24..c630294c1e 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -40,8 +40,10 @@ def initialize_callbacks_on_proxy( # noqa: PLR0915 premium_user: bool, config_file_path: str, litellm_settings: dict, - callback_specific_params: dict = {}, + callback_specific_params: Optional[dict] = None, ): + if not isinstance(callback_specific_params, dict): + callback_specific_params = {} from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.logging_callback_manager import ( LoggingCallbackManager, @@ -166,7 +168,12 @@ def initialize_callbacks_on_proxy( # noqa: PLR0915 ) init_params = {} - if "lakera_prompt_injection" in callback_specific_params: + if ( + "lakera_prompt_injection" in callback_specific_params + and isinstance( + callback_specific_params["lakera_prompt_injection"], dict + ) + ): init_params = callback_specific_params["lakera_prompt_injection"] lakera_moderations_object = lakeraAI_Moderation(**init_params) imported_list.append(lakera_moderations_object) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ba23175c10..96c9cd1e8f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -4079,6 +4079,7 @@ class ProxyConfig: premium_user=premium_user, config_file_path=config_file_path, litellm_settings=litellm_settings, + callback_specific_params=callback_settings, ) elif key == "model_group_settings": diff --git a/tests/test_litellm/integrations/compression_interception/test_compression_interception_handler.py b/tests/test_litellm/integrations/compression_interception/test_compression_interception_handler.py index 56e5a94cd4..ffa81abf86 100644 --- a/tests/test_litellm/integrations/compression_interception/test_compression_interception_handler.py +++ b/tests/test_litellm/integrations/compression_interception/test_compression_interception_handler.py @@ -32,6 +32,40 @@ def test_initialize_from_proxy_config(): assert logger.compression_target == 789 +def test_initialize_from_proxy_config_ignores_non_dict_callback_specific_params(): + """Regression (#29590): a non-dict value under + callback_settings.compression_interception must not crash initialization. + + Forwarding callback_settings as callback_specific_params activates this + branch; without the isinstance(dict) guard a non-dict value reached + from_config_yaml(...).get(...) and raised AttributeError at proxy startup. + The value is ignored and the logger falls back to defaults. + """ + logger = CompressionInterceptionLogger.initialize_from_proxy_config( + litellm_settings={}, + callback_specific_params={"compression_interception": True}, + ) + + assert logger.enabled is True + assert logger.compression_trigger == 200_000 + + +def test_initialize_from_proxy_config_honors_dict_callback_specific_params(): + """A valid dict under callback_settings.compression_interception is applied.""" + logger = CompressionInterceptionLogger.initialize_from_proxy_config( + litellm_settings={}, + callback_specific_params={ + "compression_interception": { + "enabled": False, + "compression_trigger": 12345, + } + }, + ) + + assert logger.enabled is False + assert logger.compression_trigger == 12345 + + @pytest.mark.asyncio async def test_pre_call_hook_compresses_messages_and_injects_tool(monkeypatch): """Test pre-call hook compresses and stores per-call cache.""" diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py index 1095126511..c2a502b34e 100644 --- a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py @@ -34,6 +34,35 @@ def test_initialize_from_proxy_config(): assert logger.search_tool_name == "my-search" +def test_initialize_from_proxy_config_ignores_non_dict_callback_specific_params(): + """Regression (#29590): a non-dict value under + callback_settings.websearch_interception must not crash initialization. + + Forwarding callback_settings as callback_specific_params activates this + branch; without the isinstance(dict) guard a non-dict value reached + from_config_yaml(...).get(...) and raised AttributeError at proxy startup. + The value is ignored and the logger falls back to defaults. + """ + logger = WebSearchInterceptionLogger.initialize_from_proxy_config( + litellm_settings={}, + callback_specific_params={"websearch_interception": True}, + ) + + assert logger.search_tool_name is None + + +def test_initialize_from_proxy_config_honors_dict_callback_specific_params(): + """A valid dict under callback_settings.websearch_interception is applied.""" + logger = WebSearchInterceptionLogger.initialize_from_proxy_config( + litellm_settings={}, + callback_specific_params={ + "websearch_interception": {"search_tool_name": "ws-tool"} + }, + ) + + assert logger.search_tool_name == "ws-tool" + + @pytest.mark.asyncio async def test_async_should_run_agentic_loop(): """Test that agentic loop is NOT triggered for wrong provider or missing WebSearch tool""" diff --git a/tests/test_litellm/proxy/common_utils/test_callback_utils.py b/tests/test_litellm/proxy/common_utils/test_callback_utils.py index d328d68dcd..36ff3f3c39 100644 --- a/tests/test_litellm/proxy/common_utils/test_callback_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_callback_utils.py @@ -1,7 +1,9 @@ import copy import sys import os -from types import SimpleNamespace +from types import ModuleType, SimpleNamespace + +import pytest sys.path.insert( 0, os.path.abspath("../../..") @@ -309,3 +311,98 @@ def test_encrypt_callback_vars_only_encrypts_credential_fields(monkeypatch): assert cv["langfuse_host"] == "https://cloud.langfuse.com" assert cv["langsmith_project"] == "my-proj" assert cv["langsmith_base_url"] == "https://smith.example" + + +def test_initialize_callbacks_on_proxy_lakera_ignores_non_dict_callback_settings( + monkeypatch, +): + """Regression: a non-dict value under callback_settings.lakera_prompt_injection + must not crash initialize_callbacks_on_proxy. + + Forwarding callback_settings as callback_specific_params (so callbacks like + DatadogCostManagementLogger receive their init params) exposes the lakera + branch, which previously did lakeraAI_Moderation(**callback_specific_params[ + "lakera_prompt_injection"]) with no isinstance(dict) guard. For a config like + {"lakera_prompt_injection": "x"} that is `**"x"` -> TypeError: argument after + ** must be a mapping, not str. The branch now guards on isinstance(dict), + matching the presidio / datadog_cost_management branches. + """ + captured = {} + + class _DummyLakera: + def __init__(self, **kwargs): + captured["kwargs"] = kwargs + + # Inject a fake lakera_ai module so the branch's + # `from ...lakera_ai import lakeraAI_Moderation` resolves to our stub without + # importing the real module (which imports proxy_server symbols not present + # under the stubbed proxy_server below). + fake_lakera = ModuleType("litellm.proxy.guardrails.guardrail_hooks.lakera_ai") + fake_lakera.lakeraAI_Moderation = _DummyLakera + monkeypatch.setitem( + sys.modules, + "litellm.proxy.guardrails.guardrail_hooks.lakera_ai", + fake_lakera, + ) + monkeypatch.setitem( + sys.modules, + "litellm.proxy.proxy_server", + SimpleNamespace(prisma_client=None), + ) + + original_callbacks = ( + list(litellm.callbacks) if isinstance(litellm.callbacks, list) else [] + ) + litellm.callbacks = [] + try: + # A non-dict value must be ignored (init_params stays {}), not **-unpacked. + initialize_callbacks_on_proxy( + value=["lakera_prompt_injection"], + premium_user=False, + config_file_path=".", + litellm_settings={}, + callback_specific_params={"lakera_prompt_injection": "any-string"}, + ) + assert captured["kwargs"] == {} + assert any(isinstance(c, _DummyLakera) for c in litellm.callbacks) + finally: + litellm.callbacks = original_callbacks + + +@pytest.mark.parametrize("bad_root", [None, True]) +def test_initialize_callbacks_on_proxy_non_dict_callback_specific_params_root( + monkeypatch, bad_root +): + """Regression: a blank `callback_settings:` key in YAML loads as None (and + `callback_settings: true` as a bool); load_config forwards that value + verbatim as callback_specific_params. Membership tests like + `"compression_interception" in callback_specific_params` then raise + TypeError and abort proxy startup. A non-dict root must be normalized to {} + so the callback initializes with its defaults. + """ + monkeypatch.setitem( + sys.modules, + "litellm.proxy.proxy_server", + SimpleNamespace(prisma_client=None), + ) + from litellm.integrations.compression_interception.handler import ( + CompressionInterceptionLogger, + ) + + original_callbacks = ( + list(litellm.callbacks) if isinstance(litellm.callbacks, list) else [] + ) + litellm.callbacks = [] + try: + initialize_callbacks_on_proxy( + value=["compression_interception"], + premium_user=False, + config_file_path=".", + litellm_settings={}, + callback_specific_params=bad_root, + ) + assert any( + isinstance(c, CompressionInterceptionLogger) for c in litellm.callbacks + ) + finally: + litellm.callbacks = original_callbacks diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 164538a275..677d358428 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -601,6 +601,98 @@ async def test_ProxyConfig_load_config_missing_file_raises(monkeypatch): await pc.load_config(router=None, config_file_path="/no/file.yaml") +@pytest.mark.asyncio +async def test_ProxyConfig_load_config_forwards_callback_specific_params( + tmp_path, monkeypatch +): + """Regression: callback_settings from config must be forwarded to + initialize_callbacks_on_proxy as callback_specific_params. + + Callbacks like DatadogCostManagementLogger read their init params (e.g. + cost_tag_keys) from callback_specific_params[]. If the + argument is dropped at the call site, they silently initialize with empty + params and the configured allowlist never takes effect. + """ + f = tmp_path / "c.yaml" + f.write_text( + "model_list: []\n" + "general_settings: {}\n" + "callback_settings:\n" + " datadog_cost_management:\n" + " cost_tag_keys:\n" + " - capability\n" + " - platform\n" + " - ai_product\n" + "litellm_settings:\n" + ' callbacks: ["datadog_cost_management"]\n' + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) + + captured = {} + + def _fake_initialize_callbacks_on_proxy(**kwargs): + captured.update(kwargs) + + monkeypatch.setattr( + "litellm.proxy.proxy_server.initialize_callbacks_on_proxy", + _fake_initialize_callbacks_on_proxy, + ) + + pc = ProxyConfig() + await pc.load_config(router=None, config_file_path=str(f)) + + # The callbacks branch must forward the loaded callback_settings. + assert captured.get("callback_specific_params") == { + "datadog_cost_management": { + "cost_tag_keys": ["capability", "platform", "ai_product"] + } + } + + +@pytest.mark.asyncio +async def test_ProxyConfig_load_config_blank_callback_settings_does_not_crash( + tmp_path, monkeypatch +): + """Regression: `callback_settings:` with no body loads as None because + dict.get() only falls back to the default when the key is absent. The None + was forwarded verbatim to initialize_callbacks_on_proxy, where the first + `"" in callback_specific_params` membership test raised + TypeError: argument of type 'NoneType' is not iterable, aborting startup. + Startup must succeed and the callback must initialize with its defaults. + """ + f = tmp_path / "c.yaml" + f.write_text( + "model_list: []\n" + "general_settings: {}\n" + "callback_settings:\n" + "litellm_settings:\n" + ' callbacks: ["compression_interception"]\n' + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) + + from litellm.integrations.compression_interception.handler import ( + CompressionInterceptionLogger, + ) + + original_callbacks = ( + list(litellm.callbacks) if isinstance(litellm.callbacks, list) else [] + ) + litellm.callbacks = [] + try: + pc = ProxyConfig() + await pc.load_config(router=None, config_file_path=str(f)) + + assert any( + isinstance(c, CompressionInterceptionLogger) for c in litellm.callbacks + ) + finally: + litellm.callbacks = original_callbacks + + # --------------------------------------------------------------------------- # ProxyConfig._init_non_llm_configs # ---------------------------------------------------------------------------