mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
refactor(guardrails): retry guardrail rotation by bounded recursion instead of a rebound cursor
- each attempt re-reads the row and recurses with attempts_left - 1, so no loop variable is rebound - cover the give-up path after GUARDRAIL_ROTATION_ATTEMPTS writes
This commit is contained in:
parent
79fcef6256
commit
27ac33a1a6
3 changed files with 48 additions and 19 deletions
|
|
@ -141,27 +141,33 @@ def guardrail_from_db_row(row: Iterable[tuple[str, object]]) -> Guardrail:
|
|||
|
||||
|
||||
async def _rotate_guardrail_row(
|
||||
prisma_client: PrismaClient, row: "prisma_models.LiteLLM_GuardrailsTable", encryption_key: str
|
||||
prisma_client: PrismaClient,
|
||||
row: "prisma_models.LiteLLM_GuardrailsTable | None",
|
||||
encryption_key: str,
|
||||
attempts_left: int = GUARDRAIL_ROTATION_ATTEMPTS,
|
||||
) -> int:
|
||||
current: prisma_models.LiteLLM_GuardrailsTable | None = row
|
||||
for _ in range(GUARDRAIL_ROTATION_ATTEMPTS):
|
||||
if current is None or not isinstance(current.litellm_params, Mapping):
|
||||
return 0
|
||||
rotated_params: dict[str, object] = encrypt_guardrail_litellm_params(
|
||||
decrypt_guardrail_litellm_params(current.litellm_params), new_encryption_key=encryption_key
|
||||
)
|
||||
if rotated_params == current.litellm_params:
|
||||
return 0
|
||||
if await _guardrail_table(prisma_client).update_many(
|
||||
where={"guardrail_id": current.guardrail_id, "updated_at": current.updated_at},
|
||||
data={"litellm_params": safe_dumps(rotated_params)},
|
||||
):
|
||||
return 1
|
||||
current = await _guardrail_table(prisma_client).find_unique(where={"guardrail_id": row.guardrail_id})
|
||||
verbose_proxy_logger.warning(
|
||||
"Guardrail %s kept changing during master key rotation; its secrets were not re-encrypted", row.guardrail_id
|
||||
"""Re-encrypt one row's params under encryption_key with a compare-and-set on updated_at.
|
||||
A row edited since it was read is re-read and retried, up to attempts_left writes. Returns 1 when rewritten."""
|
||||
if row is None or not isinstance(row.litellm_params, Mapping):
|
||||
return 0
|
||||
rotated_params: Final = encrypt_guardrail_litellm_params(
|
||||
decrypt_guardrail_litellm_params(row.litellm_params), new_encryption_key=encryption_key
|
||||
)
|
||||
return 0
|
||||
if rotated_params == row.litellm_params:
|
||||
return 0
|
||||
if await _guardrail_table(prisma_client).update_many(
|
||||
where={"guardrail_id": row.guardrail_id, "updated_at": row.updated_at},
|
||||
data={"litellm_params": safe_dumps(rotated_params)},
|
||||
):
|
||||
return 1
|
||||
if attempts_left <= 1:
|
||||
verbose_proxy_logger.warning(
|
||||
"Guardrail %s kept changing during master key rotation; its secrets were not re-encrypted",
|
||||
row.guardrail_id,
|
||||
)
|
||||
return 0
|
||||
latest_row: Final = await _guardrail_table(prisma_client).find_unique(where={"guardrail_id": row.guardrail_id})
|
||||
return await _rotate_guardrail_row(prisma_client, latest_row, encryption_key, attempts_left - 1)
|
||||
|
||||
|
||||
guardrail_initializer_registry: Final = {
|
||||
|
|
|
|||
|
|
@ -38,6 +38,7 @@ IGNORE_FUNCTIONS = [
|
|||
"_mask_sequence", # max depth set.
|
||||
"_encrypted_param", # max depth set.
|
||||
"_decrypted_param", # max depth set.
|
||||
"_rotate_guardrail_row", # bounded by attempts_left.
|
||||
"_delete_nested_value_custom", # max depth set (bounded by number of path segments).
|
||||
"filter_exceptions_from_params", # max depth set (default 20) to prevent infinite recursion.
|
||||
"__getattr__", # lazy loading pattern in litellm/__init__.py with proper caching to prevent infinite recursion.
|
||||
|
|
|
|||
|
|
@ -1293,3 +1293,25 @@ async def test_rotate_guardrail_params_retries_a_row_edited_during_rotation(monk
|
|||
assert decrypt_guardrail_litellm_params(_stored_params(prisma_client.db.litellm_guardrailstable.update_many)) == {
|
||||
"api_key": "edited-key"
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rotate_guardrail_params_gives_up_on_a_row_that_keeps_changing(monkeypatch):
|
||||
from litellm.constants import GUARDRAIL_ROTATION_ATTEMPTS
|
||||
from litellm.proxy.guardrails.guardrail_registry import encrypt_guardrail_litellm_params
|
||||
|
||||
monkeypatch.delenv("LITELLM_SALT_KEY", raising=False)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-old-master")
|
||||
row = _Row(guardrail_id="g-1", updated_at="t1", litellm_params=encrypt_guardrail_litellm_params({"api_key": "k"}))
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_guardrailstable.find_many = AsyncMock(return_value=[row])
|
||||
prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row)
|
||||
prisma_client.db.litellm_guardrailstable.update_many = AsyncMock(return_value=0)
|
||||
|
||||
rows_updated = await GuardrailRegistry.rotate_guardrail_params_master_key(
|
||||
prisma_client=prisma_client, new_master_key="sk-new-master"
|
||||
)
|
||||
|
||||
assert rows_updated == 0
|
||||
assert prisma_client.db.litellm_guardrailstable.update_many.await_count == GUARDRAIL_ROTATION_ATTEMPTS
|
||||
assert prisma_client.db.litellm_guardrailstable.find_unique.await_count == GUARDRAIL_ROTATION_ATTEMPTS - 1
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue