mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: gate guardrail read-through to active rows and serialize it with the reload reconcile
This commit is contained in:
parent
e9c01da233
commit
f8a23aab09
5 changed files with 125 additions and 34 deletions
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
########################################################
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue