mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(proxy): preserve database setting types during startup
This commit is contained in:
parent
8196a38dc0
commit
08fd5f86c8
3 changed files with 78 additions and 12 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue