diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 5ba71a7fa12..13bfc675a01 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -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 = { diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index 9a4860a7977..9667156d3a8 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -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. diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py index 45955270d5d..08b099cc9ef 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py @@ -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