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:
Yucheng He 2026-09-28 17:14:48 -07:00
parent 79fcef6256
commit 27ac33a1a6
3 changed files with 48 additions and 19 deletions

View file

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

View file

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

View file

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