diff --git a/litellm/proxy/common_utils/registry_read_through.py b/litellm/proxy/common_utils/registry_read_through.py index e92803d3f51..460b348e188 100644 --- a/litellm/proxy/common_utils/registry_read_through.py +++ b/litellm/proxy/common_utils/registry_read_through.py @@ -22,6 +22,7 @@ if TYPE_CHECKING: from prisma.types import ( LiteLLM_AgentsTableInclude, LiteLLM_AgentsTableWhereUniqueInput, + LiteLLM_GuardrailsTableWhereInput, LiteLLM_ProxyModelTableWhereInput, ) @@ -132,20 +133,25 @@ async def _resync_model_deployments(model_name: str) -> bool: async def _resync_guardrails(guardrail_name: str) -> bool: from litellm.proxy import proxy_server from litellm.proxy.guardrails.guardrail_registry import ( + GUARDRAIL_RECONCILE_LOCK, IN_MEMORY_GUARDRAIL_HANDLER, - GuardrailRegistry, ) + from litellm.repositories.table_repositories import GuardrailsRepository + from litellm.types.guardrails import Guardrail if not _db_backed_registries_enabled("guardrails"): return False prisma_client: Final = proxy_server.prisma_client assert prisma_client is not None - row: Final = await GuardrailRegistry().get_guardrail_by_name_from_db( - guardrail_name=guardrail_name, prisma_client=prisma_client - ) + active_row_filter: Final[LiteLLM_GuardrailsTableWhereInput] = { + "guardrail_name": guardrail_name, + "status": "active", + } + row: Final = await GuardrailsRepository(prisma_client).table.find_first(where=active_row_filter) if row is None: return False - IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(guardrail=row) + async with GUARDRAIL_RECONCILE_LOCK: + IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db(guardrail=Guardrail(**dict(row))) return _initialized_guardrail(guardrail_name) is not None diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 5f7374581a2..f6d348c1045 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -1,5 +1,6 @@ # litellm/proxy/guardrails/guardrail_registry.py +import asyncio import importlib import os from collections.abc import Callable, Iterator, Mapping @@ -813,4 +814,6 @@ class InMemoryGuardrailHandler: # In Memory Guardrail Handler for LiteLLM Proxy ######################################################## IN_MEMORY_GUARDRAIL_HANDLER: Final = InMemoryGuardrailHandler() + +GUARDRAIL_RECONCILE_LOCK: Final = asyncio.Lock() ######################################################## diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index b0deb71f917..be1914ea50c 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -7090,38 +7090,40 @@ class ProxyConfig: async def _init_guardrails_in_db(self, prisma_client: PrismaClient): from litellm.proxy.guardrails.guardrail_registry import ( + GUARDRAIL_RECONCILE_LOCK, IN_MEMORY_GUARDRAIL_HANDLER, Guardrail, GuardrailRegistry, ) try: - guardrails_in_db: Final[list[Guardrail]] = await GuardrailRegistry.get_all_guardrails_from_db( - prisma_client=prisma_client - ) - verbose_proxy_logger.debug("guardrails from the DB %s", str(guardrails_in_db)) - db_guardrail_ids: Final[set] = set() - for guardrail in guardrails_in_db: - guardrail_id = guardrail.get("guardrail_id") - if guardrail_id: - db_guardrail_ids.add(guardrail_id) - try: - IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db( - guardrail=cast(Guardrail, guardrail), - ) - except Exception as e: # noqa: BLE001 # one unloadable row must not stop the remaining guardrails - verbose_proxy_logger.error( - "litellm.proxy.proxy_server.py::ProxyConfig:_init_guardrails_in_db - " - "skipping guardrail '%s' (ID: %s): %s: %s", - guardrail.get("guardrail_name"), - guardrail_id, - type(e).__name__, - e, - ) + async with GUARDRAIL_RECONCILE_LOCK: + guardrails_in_db: Final[list[Guardrail]] = await GuardrailRegistry.get_all_guardrails_from_db( + prisma_client=prisma_client + ) + verbose_proxy_logger.debug("guardrails from the DB %s", str(guardrails_in_db)) + db_guardrail_ids: Final[set] = set() + for guardrail in guardrails_in_db: + guardrail_id = guardrail.get("guardrail_id") + if guardrail_id: + db_guardrail_ids.add(guardrail_id) + try: + IN_MEMORY_GUARDRAIL_HANDLER.sync_guardrail_from_db( + guardrail=cast(Guardrail, guardrail), + ) + except Exception as e: # noqa: BLE001 # one unloadable row must not stop the remaining guardrails + verbose_proxy_logger.error( + "litellm.proxy.proxy_server.py::ProxyConfig:_init_guardrails_in_db - " + "skipping guardrail '%s' (ID: %s): %s: %s", + guardrail.get("guardrail_name"), + guardrail_id, + type(e).__name__, + e, + ) - # Drop in-memory DB-backed entries whose row was deleted on another - # pod. Config-loaded entries are never touched. - IN_MEMORY_GUARDRAIL_HANDLER.reconcile_db_guardrails(db_guardrail_ids=db_guardrail_ids) + # Drop in-memory DB-backed entries whose row was deleted on another + # pod. Config-loaded entries are never touched. + IN_MEMORY_GUARDRAIL_HANDLER.reconcile_db_guardrails(db_guardrail_ids=db_guardrail_ids) except Exception as e: verbose_proxy_logger.exception("litellm.proxy.proxy_server.py::ProxyConfig:_init_guardrails_in_db - %s", e) diff --git a/tests/test_litellm/proxy/common_utils/test_registry_read_through.py b/tests/test_litellm/proxy/common_utils/test_registry_read_through.py index 56c71e38ef5..f0fbdea4e85 100644 --- a/tests/test_litellm/proxy/common_utils/test_registry_read_through.py +++ b/tests/test_litellm/proxy/common_utils/test_registry_read_through.py @@ -260,6 +260,7 @@ class FakeGuardrailRow: "blocked_words": [{"keyword": "secret", "action": "BLOCK"}], }, "guardrail_info": {}, + "status": "active", }.items() ) @@ -277,7 +278,7 @@ async def test_get_guardrail_with_read_through_recovers_guardrail_created_on_sib guardrail_id: Final = "read-through-db-guardrail-id" guardrail_name: Final = "read-through-db-guardrail" prisma_client: Final = MagicMock() - prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock( + prisma_client.db.litellm_guardrailstable.find_first = AsyncMock( return_value=FakeGuardrailRow(guardrail_id, guardrail_name) ) prisma_client.db.litellm_guardrailstable.find_many = AsyncMock( @@ -290,8 +291,8 @@ async def test_get_guardrail_with_read_through_recovers_guardrail_created_on_sib guardrail: Final = await get_initialized_guardrail_with_read_through(guardrail_name=guardrail_name) assert guardrail is not None assert guardrail.guardrail_name == guardrail_name - prisma_client.db.litellm_guardrailstable.find_unique.assert_awaited_once_with( - where={"guardrail_name": guardrail_name} + prisma_client.db.litellm_guardrailstable.find_first.assert_awaited_once_with( + where={"guardrail_name": guardrail_name, "status": "active"} ) finally: IN_MEMORY_GUARDRAIL_HANDLER.delete_in_memory_guardrail(guardrail_id) @@ -307,13 +308,64 @@ async def test_get_guardrail_with_read_through_returns_none_for_unknown_guardrai ) prisma_client: Final = MagicMock() - prisma_client.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=None) + prisma_client.db.litellm_guardrailstable.find_first = AsyncMock(return_value=None) monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) monkeypatch.setattr(proxy_server, "store_model_in_db", True) assert await get_initialized_guardrail_with_read_through(guardrail_name="guardrail-nobody-created") is None +@pytest.mark.asyncio +async def test_resync_guardrails_never_loads_non_active_rows(monkeypatch): + from unittest.mock import AsyncMock, MagicMock + + import litellm.proxy.proxy_server as proxy_server + from litellm.proxy.common_utils.registry_read_through import _resync_guardrails + + pending_name: Final = "pending-review-guardrail" + prisma_client: Final = MagicMock() + prisma_client.db.litellm_guardrailstable.find_first = AsyncMock(return_value=None) + monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) + monkeypatch.setattr(proxy_server, "store_model_in_db", True) + + assert await _resync_guardrails(pending_name) is False + prisma_client.db.litellm_guardrailstable.find_first.assert_awaited_once_with( + where={"guardrail_name": pending_name, "status": "active"} + ) + + +@pytest.mark.asyncio +async def test_resync_guardrails_syncs_under_guardrail_reconcile_lock(monkeypatch): + from unittest.mock import AsyncMock, MagicMock + + import litellm.proxy.common_utils.registry_read_through as read_through_module + import litellm.proxy.proxy_server as proxy_server + from litellm.proxy.common_utils.registry_read_through import _resync_guardrails + from litellm.proxy.guardrails.guardrail_registry import ( + GUARDRAIL_RECONCILE_LOCK, + IN_MEMORY_GUARDRAIL_HANDLER, + ) + + guardrail_name: Final = "lock-scope-guardrail" + prisma_client: Final = MagicMock() + prisma_client.db.litellm_guardrailstable.find_first = AsyncMock( + return_value=FakeGuardrailRow("lock-scope-guardrail-id", guardrail_name) + ) + lock_states: list[bool] = [] + + def record_sync(guardrail) -> None: + lock_states.append(GUARDRAIL_RECONCILE_LOCK.locked()) + + monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) + monkeypatch.setattr(proxy_server, "store_model_in_db", True) + monkeypatch.setattr(IN_MEMORY_GUARDRAIL_HANDLER, "sync_guardrail_from_db", record_sync) + monkeypatch.setattr(read_through_module, "_initialized_guardrail", lambda guardrail_name: MagicMock()) + + assert await _resync_guardrails(guardrail_name) is True + assert lock_states == [True] + assert not GUARDRAIL_RECONCILE_LOCK.locked() + + @pytest.mark.asyncio async def test_resync_model_deployments_mutates_router_under_model_reconcile_lock(monkeypatch): from unittest.mock import AsyncMock, MagicMock diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index e7d7d4c9323..75aa716bb85 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -11126,6 +11126,34 @@ async def test_init_agents_in_db_rebuilds_registry_under_agent_reconcile_lock(mo assert not AGENT_RECONCILE_LOCK.locked() +@pytest.mark.asyncio +async def test_init_guardrails_in_db_snapshots_and_reconciles_under_guardrail_reconcile_lock(monkeypatch): + from litellm.proxy.guardrails.guardrail_registry import ( + GUARDRAIL_RECONCILE_LOCK, + IN_MEMORY_GUARDRAIL_HANDLER, + GuardrailRegistry, + ) + from litellm.proxy.proxy_server import ProxyConfig + + lock_states: list[bool] = [] + + async def fake_get_all_guardrails_from_db(prisma_client) -> list: + lock_states.append(GUARDRAIL_RECONCILE_LOCK.locked()) + return [] + + def fake_reconcile_db_guardrails(db_guardrail_ids) -> list: + lock_states.append(GUARDRAIL_RECONCILE_LOCK.locked()) + return [] + + monkeypatch.setattr(GuardrailRegistry, "get_all_guardrails_from_db", fake_get_all_guardrails_from_db) + monkeypatch.setattr(IN_MEMORY_GUARDRAIL_HANDLER, "reconcile_db_guardrails", fake_reconcile_db_guardrails) + + await ProxyConfig()._init_guardrails_in_db(prisma_client=MagicMock()) + + assert lock_states == [True, True] + assert not GUARDRAIL_RECONCILE_LOCK.locked() + + class TestEmbeddingsFailureHookRequestData: @pytest.mark.asyncio async def test_failure_hook_gets_post_setup_data_with_logging_obj(self):