fix(proxy): preserve database setting types during startup

This commit is contained in:
Yucheng He 2026-09-07 12:35:52 -07:00
parent 8196a38dc0
commit 08fd5f86c8
3 changed files with 78 additions and 12 deletions

View file

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

View file

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

View file

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