diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 74b4d145c32..8e7920fc507 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6964,8 +6964,7 @@ class ProxyConfig: return current_config elif param_name == "litellm_settings" and isinstance(db_param_value, dict): for key, value in db_param_value.items(): - if key in LITELLM_SETTINGS_SAFE_DB_OVERRIDES: # params that are safe to override with db values - setattr(litellm, key, value) + self._apply_safe_litellm_setting_override(key, value) # If param doesn't exist in config, add it if param_name not in current_config: @@ -7261,15 +7260,19 @@ class ProxyConfig: if not isinstance(litellm_settings, dict): return for key, value in litellm_settings.items(): - if key not in LITELLM_SETTINGS_SAFE_DB_OVERRIDES: - continue - if key == "prometheus_emit_input_sequence_length_label": - if isinstance(value, bool): - setattr(litellm, key, value) - elif isinstance(value, str) and (normalized_value := str_to_bool(value)) is not None: - setattr(litellm, key, normalized_value) - continue - setattr(litellm, key, value) + self._apply_safe_litellm_setting_override(key, value) + + @staticmethod + def _apply_safe_litellm_setting_override(key: str, value: object) -> None: + if key not in LITELLM_SETTINGS_SAFE_DB_OVERRIDES: + return + if key == "prometheus_emit_input_sequence_length_label": + if isinstance(value, bool): + setattr(litellm, key, value) + elif isinstance(value, str) and (normalized_value := str_to_bool(value)) is not None: + setattr(litellm, key, normalized_value) + return + setattr(litellm, key, value) async def _init_semantic_filter_settings_in_db(self, prisma_client: PrismaClient): """ @@ -9022,7 +9025,7 @@ class ProxyStartupEvent: @staticmethod async def resolve_store_model_in_db(prisma_client: PrismaClient | None, configured: bool) -> bool: - if get_secret_bool("STORE_MODEL_IN_DB", configured) or configured: + if (get_secret_bool("STORE_MODEL_IN_DB", configured) or configured) is True: return True if prisma_client is None: return False diff --git a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py index 759d91571a6..97d49bfb908 100644 --- a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py +++ b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py @@ -774,6 +774,51 @@ async def test_resolve_store_model_in_db_uses_config_or_db( assert prisma_client.db.litellm_config.find_first.await_count == (0 if configured else 1) +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("configured", "environment", "db_value", "expected"), + ( + ("false", None, False, False), + ("true", None, False, False), + ("false", "false", False, False), + ("false", "true", False, True), + (True, "false", False, True), + (False, "false", True, True), + ("false", None, True, True), + ), +) +async def test_resolve_store_model_in_db_preserves_legacy_config_and_env_precedence( + monkeypatch: pytest.MonkeyPatch, + configured: bool | str, + environment: str | None, + db_value: bool, + expected: bool, +): + monkeypatch.setattr(ps, "store_model_in_db", configured) + if environment is None: + monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False) + else: + monkeypatch.setenv("STORE_MODEL_IN_DB", environment) + prisma_client: Final = MagicMock() + prisma_client.db.litellm_config.find_first = AsyncMock( + return_value=MagicMock(param_value={"store_model_in_db": db_value}) + ) + + assert await ProxyStartupEvent.resolve_store_model_in_db(prisma_client, ps.store_model_in_db) is expected + assert prisma_client.db.litellm_config.find_first.await_count == ( + 0 if configured is True or environment == "true" else 1 + ) + + +@pytest.mark.asyncio +async def test_resolve_store_model_in_db_continues_after_database_failure(monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False) + prisma_client: Final = MagicMock() + prisma_client.db.litellm_config.find_first = AsyncMock(side_effect=RuntimeError("database unavailable")) + + assert await ProxyStartupEvent.resolve_store_model_in_db(prisma_client, configured=False) is False + + @pytest.mark.asyncio @pytest.mark.parametrize("param_value", ("legacy", ["store_model_in_db"], 7)) async def test_resolve_store_model_in_db_ignores_non_mapping_row(monkeypatch: pytest.MonkeyPatch, param_value: object): 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 56b3f9cebbf..9bbe8fcc12d 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -3133,6 +3133,24 @@ def test_ProxyConfig__update_config_fields_merges_dict(): assert out == {"general_settings": {"a": 1, "b": 3, "c": 4, "d": 5}} +@pytest.mark.parametrize( + ("db_value", "expected"), + ((True, True), (False, False), ("true", True), ("false", False), ("invalid", False)), +) +def test_update_config_fields_normalizes_input_sequence_length_flag( + monkeypatch: pytest.MonkeyPatch, db_value: object, expected: bool +): + monkeypatch.setattr(litellm, "prometheus_emit_input_sequence_length_label", False) + + ProxyConfig()._update_config_fields( + current_config={}, + param_name="litellm_settings", + db_param_value={"prometheus_emit_input_sequence_length_label": db_value}, + ) + + assert litellm.prometheus_emit_input_sequence_length_label is expected + + def test_ProxyConfig__update_config_fields_invalid_param_raises(): pc = ProxyConfig() with pytest.raises(TypeError):