From 8d4b5ed571512379e6d794175e9db36133a7f851 Mon Sep 17 00:00:00 2001 From: Harshit28j Date: Mon, 16 Mar 2026 16:38:53 +0530 Subject: [PATCH] fix: req changes on tests --- .../proxy/guardrails/guardrail_endpoints.py | 27 ++++++++++++++++--- .../guardrails/test_guardrail_endpoints.py | 16 +++++++++-- 2 files changed, 38 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index c08a7070f90..5757e42de48 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -462,6 +462,7 @@ async def update_guardrail( f"Immediate sync: Failed to update '{guardrail_name}' (ID: {guardrail_id}) in memory: {update_error}" ) # Rollback: restore previous guardrail data in DB + rollback_ok = False try: await GUARDRAIL_REGISTRY.update_guardrail_in_db( guardrail_id=guardrail_id, @@ -473,14 +474,20 @@ async def update_guardrail( guardrail_id=guardrail_id, guardrail=cast(Guardrail, existing_guardrail), ) + rollback_ok = True except Exception: verbose_proxy_logger.error( "Failed to rollback guardrail DB entry after update failure" ) + status = ( + "rolled back" + if rollback_ok + else "rollback also failed, DB/memory may be inconsistent" + ) raise HTTPException( status_code=422, - detail=f"Guardrail update failed, rolled back: {update_error}", + detail=f"Guardrail update failed, {status}: {update_error}", ) return result @@ -557,19 +564,26 @@ async def delete_guardrail( f"Immediate sync: Failed to remove guardrail '{guardrail_name}' (ID: {guardrail_id}) from memory: {delete_error}" ) # Rollback: re-create the DB entry with the ORIGINAL guardrail_id + rollback_ok = False try: await GUARDRAIL_REGISTRY.add_guardrail_to_db( guardrail=cast(Guardrail, existing_guardrail), prisma_client=prisma_client, guardrail_id=guardrail_id, ) + rollback_ok = True except Exception: verbose_proxy_logger.error( "Failed to rollback guardrail DB deletion after memory removal failure" ) + status = ( + "rolled back" + if rollback_ok + else "rollback also failed, DB/memory may be inconsistent" + ) raise HTTPException( status_code=422, - detail=f"Guardrail delete failed, rolled back: {delete_error}", + detail=f"Guardrail delete failed, {status}: {delete_error}", ) return result @@ -1155,6 +1169,7 @@ async def patch_guardrail( f"Immediate sync: Failed to update '{guardrail_name}' (ID: {guardrail_id}) in memory: {update_error}" ) # Rollback: restore previous guardrail data in DB + rollback_ok = False try: await GUARDRAIL_REGISTRY.update_guardrail_in_db( guardrail_id=guardrail_id, @@ -1165,14 +1180,20 @@ async def patch_guardrail( IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db( guardrail=cast(Guardrail, existing_guardrail), ) + rollback_ok = True except Exception: verbose_proxy_logger.error( "Failed to rollback guardrail DB entry after patch failure" ) + status = ( + "rolled back" + if rollback_ok + else "rollback also failed, DB/memory may be inconsistent" + ) raise HTTPException( status_code=422, - detail=f"Guardrail patch failed, rolled back: {update_error}", + detail=f"Guardrail patch failed, {status}: {update_error}", ) return result diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index f46f62b9b11..c61c7401238 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -837,7 +837,11 @@ async def test_update_guardrail_endpoint( assert "Prisma client not initialized" in str(exc_info.value.detail) elif scenario == "success_sync_fails": assert exc_info.value.status_code == 422 - assert "rolled back" in str(exc_info.value.detail) + assert "Guardrail update failed" in str(exc_info.value.detail) + # Verify rollback was attempted: update_in_db called twice (initial + rollback) + assert mock_guardrail_registry.update_guardrail_in_db.call_count == 2 + # Verify rollback attempted in-memory re-init with old config + assert mock_in_memory_handler.update_in_memory_guardrail.call_count == 2 else: result = await update_guardrail( "test-guardrail-id", MOCK_UPDATE_REQUEST, user_api_key_dict=MOCK_ADMIN_USER @@ -942,7 +946,11 @@ async def test_patch_guardrail_endpoint( if scenario == "success_sync_fails": assert exc_info.value.status_code == 422 - assert "rolled back" in str(exc_info.value.detail) + assert "Guardrail patch failed" in str(exc_info.value.detail) + # Verify rollback was attempted: update_in_db called twice (initial + rollback) + assert mock_guardrail_registry.update_guardrail_in_db.call_count == 2 + # Verify rollback attempted in-memory re-sync with old config + assert mock_in_memory_handler.sync_guardrail_from_db.call_count == 2 else: result = await patch_guardrail( "test-guardrail-id", MOCK_PATCH_REQUEST, user_api_key_dict=MOCK_ADMIN_USER @@ -1017,6 +1025,10 @@ async def test_delete_guardrail_endpoint( if scenario == "success_sync_fails": assert exc_info.value.status_code == 422 assert "rolled back" in str(exc_info.value.detail) + # Verify rollback was attempted: add_guardrail_to_db called with original ID + mock_guardrail_registry.add_guardrail_to_db.assert_called_once() + call_kwargs = mock_guardrail_registry.add_guardrail_to_db.call_args + assert call_kwargs.kwargs.get("guardrail_id") == expected_result else: result = await delete_guardrail( guardrail_id=expected_result, user_api_key_dict=MOCK_ADMIN_USER