fix: gate guardrail read-through to active rows and serialize it with the reload reconcile

This commit is contained in:
mateo-berri 2026-08-19 01:46:33 -07:00
parent e9c01da233
commit f8a23aab09
5 changed files with 125 additions and 34 deletions

View file

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

View file

@ -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()
########################################################

View file

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

View file

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

View file

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