fix(guardrails): tolerate invalid stored params on reads

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-10-01 05:47:10 +00:00
parent 2056b73256
commit 8a35342461
4 changed files with 204 additions and 10 deletions

View file

@ -79,7 +79,6 @@ if TYPE_CHECKING:
router: Final = APIRouter()
GUARDRAIL_REGISTRY: Final = GuardrailRegistry()
_BASE_LITELLM_PARAMS_ADAPTER: Final[TypeAdapter[BaseLitellmParams]] = TypeAdapter(BaseLitellmParams)
def _as_str_object_mapping(mapping: Mapping[str, object]) -> Mapping[str, object]:
@ -271,7 +270,11 @@ async def list_guardrails_v2(
number_of_asterisks=4,
)
masked_litellm_params = (
_BASE_LITELLM_PARAMS_ADAPTER.validate_python(masked_litellm_params_dict)
parse_tolerant_litellm_params(
masked_litellm_params_dict,
guardrail.get("guardrail_name") or "Unknown",
params_model=BaseLitellmParams,
)
if masked_litellm_params_dict
else None
)
@ -314,7 +317,11 @@ async def list_guardrails_v2(
number_of_asterisks=4,
)
masked_in_memory_litellm_params_typed = (
_BASE_LITELLM_PARAMS_ADAPTER.validate_python(masked_in_memory_litellm_params)
parse_tolerant_litellm_params(
masked_in_memory_litellm_params,
guardrail.get("guardrail_name") or "Unknown",
params_model=BaseLitellmParams,
)
if masked_in_memory_litellm_params
else None
)
@ -1395,7 +1402,11 @@ async def get_guardrail_info(guardrail_id: str):
number_of_asterisks=4,
)
masked_litellm_params = (
_BASE_LITELLM_PARAMS_ADAPTER.validate_python(masked_litellm_params_dict)
parse_tolerant_litellm_params(
masked_litellm_params_dict,
result.get("guardrail_name") or "Unknown",
params_model=BaseLitellmParams,
)
if masked_litellm_params_dict
else None
)

View file

@ -7,9 +7,9 @@ from collections.abc import Callable, Iterator, Mapping, Sequence
from datetime import datetime, timezone
from itertools import chain, count
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, TypeAlias, cast
from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol, TypeAlias, TypeVar, cast
from pydantic import ValidationError
from pydantic import BaseModel, ValidationError
import litellm
from litellm import Router
@ -500,12 +500,16 @@ def _configure_callback_scoping(
_apply_configured_bool_overrides(custom_guardrail_callback, litellm_params)
_ParamsT = TypeVar("_ParamsT", bound=BaseModel)
def parse_tolerant_litellm_params(
litellm_params_data: Mapping[str, object],
guardrail_name: str,
) -> LitellmParams:
params_model: type[_ParamsT] = LitellmParams,
) -> _ParamsT:
try:
return LitellmParams(**litellm_params_data)
return params_model.model_validate(litellm_params_data)
except ValidationError as validation_error:
if any(tuple(error["loc"]) != ("logging_only_scope",) for error in validation_error.errors()):
raise
@ -515,7 +519,7 @@ def parse_tolerant_litellm_params(
guardrail_name.replace("\r", "").replace("\n", ""),
str(litellm_params_data.get("logging_only_scope")).replace("\r", "").replace("\n", "")[:100],
)
return LitellmParams(**MappingProxyType({**litellm_params_data, "logging_only_scope": None}))
return params_model.model_validate(MappingProxyType({**litellm_params_data, "logging_only_scope": None}))
class InMemoryGuardrailHandler:

View file

@ -49,6 +49,74 @@ from integration._support.wire import Reply, Request
from pydantic import JsonValue
def test_G7_database_invalid_scope_reads_preserve_pre_call_blocking(gateway: Gateway, tmp_path: Path) -> None:
identity: Final = f"logging-scope-g7-{uuid.uuid4().hex}"
blocked_word: Final = f"pineapple{uuid.uuid4().hex[:8]}"
prompt: Final = f"synthetic request containing {blocked_word}"
reply: Final = f"synthetic upstream response {identity}"
scenario_id: Final = f"phase12-g7-{uuid.uuid4().hex}"
guardrail_id: Final = str(uuid.uuid5(uuid.NAMESPACE_URL, identity))
upstream_handle: Final = register_scenario(scenario_id, _provider_response("chat", scenario_id, reply, False))
try:
with gateway.scenario() as scenario:
model: Final = scenario.model(
model="openai/gpt-4o-mini",
api_base=f"{upstream_handle.api_base()}/v1",
api_key="synthetic-provider-key",
)
write_rows(
'INSERT INTO "LiteLLM_GuardrailsTable" '
"(guardrail_id, guardrail_name, litellm_params, guardrail_info, updated_at) "
"VALUES (%s, %s, %s::jsonb, %s::jsonb, NOW())",
(
guardrail_id,
identity,
json.dumps(
{
"guardrail": "litellm_content_filter",
"mode": "pre_call",
"logging_only_scope": "Input",
"default_on": True,
"blocked_words": [{"keyword": blocked_word, "action": "BLOCK"}],
}
),
"{}",
),
)
try:
config: Final = _empty_proxy_configuration(tmp_path, identity)
with owned_proxy(gateway, tmp_path, {}, config=config, workers=1) as candidate:
listing_response: Final = candidate.request("GET", "/v2/guardrails/list")
assert listing_response.status_code == 200, listing_response.text
listing: Final = JSON_OBJECT.validate_json(listing_response.content)
rows: Final = tuple(object_value(row) for row in listing["guardrails"])
stored_guardrail: Final = next(row for row in rows if row.get("guardrail_id") == guardrail_id)
stored_params: Final = object_value(stored_guardrail["litellm_params"])
expected_scope: Final = "Input" if _is_base_audit_leg() else None
assert stored_params.get("logging_only_scope") == expected_scope, stored_guardrail
assert stored_params["mode"] == "pre_call", stored_guardrail
assert stored_params["guardrail"] == "litellm_content_filter", stored_guardrail
info_response: Final = candidate.request("GET", f"/guardrails/{guardrail_id}/info")
assert info_response.status_code == 200, info_response.text
info: Final = JSON_OBJECT.validate_json(info_response.content)
info_params: Final = object_value(info["litellm_params"])
assert info_params.get("logging_only_scope") == expected_scope, info
blocked_response: Final = candidate.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": prompt}]},
)
assert blocked_response.status_code == 400, blocked_response.text
assert _drain_upstream(gateway.upstream_url) == ()
finally:
_delete_database_guardrail(identity)
finally:
delete_scenario(upstream_handle)
@pytest.mark.parametrize(
("row_id", "scope", "blocked_side"),
(

View file

@ -2,7 +2,8 @@ import json
import time
from collections.abc import Callable
from datetime import datetime
from typing import Dict, List, Optional
from types import MappingProxyType
from typing import Dict, Final, List, Optional
from unittest.mock import AsyncMock
import pytest
@ -67,6 +68,46 @@ MOCK_DB_GUARDRAIL = {
"updated_at": datetime.now(),
}
_INVALID_SCOPE_LITELLM_PARAMS: Final = MappingProxyType(
{
"guardrail": "litellm_content_filter",
"mode": "pre_call",
"logging_only_scope": "Input",
"blocked_words": [{"keyword": "synthetic blocked phrase", "action": "BLOCK"}],
}
)
_INVALID_SCOPE_DB_GUARDRAIL: Final = MappingProxyType(
{
"guardrail_id": "invalid-scope-db-guardrail",
"guardrail_name": "Invalid scope DB guardrail",
"litellm_params": _INVALID_SCOPE_LITELLM_PARAMS,
"guardrail_info": MappingProxyType({}),
}
)
_INVALID_SCOPE_IN_MEMORY_GUARDRAIL: Final = MappingProxyType(
{
"guardrail_id": "invalid-scope-in-memory-guardrail",
"guardrail_name": "Invalid scope in-memory guardrail",
"litellm_params": _INVALID_SCOPE_LITELLM_PARAMS,
"guardrail_info": MappingProxyType({}),
}
)
_VALID_SCOPE_DB_GUARDRAIL: Final = MappingProxyType(
{
"guardrail_id": "valid-scope-db-guardrail",
"guardrail_name": "Valid scope DB guardrail",
"litellm_params": MappingProxyType(
{
"guardrail": "litellm_content_filter",
"mode": "pre_call",
"logging_only_scope": "output",
"blocked_words": [{"keyword": "synthetic blocked phrase", "action": "BLOCK"}],
}
),
"guardrail_info": MappingProxyType({}),
}
)
MOCK_CONFIG_GUARDRAIL = {
"guardrail_id": "test-config-guardrail",
"guardrail_name": "Test Config Guardrail",
@ -229,6 +270,54 @@ async def test_list_guardrails_v2_with_db_and_config(mocker, mock_prisma_client,
assert isinstance(config_guardrail.litellm_params, BaseLitellmParams)
@pytest.mark.asyncio
async def test_list_guardrails_v2_normalizes_invalid_scope_and_keeps_other_db_rows(
mocker, mock_prisma_client, mock_in_memory_handler
):
mock_prisma_client.db.litellm_guardrailstable.find_many.return_value = (
_INVALID_SCOPE_DB_GUARDRAIL,
_VALID_SCOPE_DB_GUARDRAIL,
)
mock_in_memory_handler.list_in_memory_guardrails.return_value = ()
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mocker.patch(
"litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER",
mock_in_memory_handler,
)
response: Final = await list_guardrails_v2(user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN))
invalid_scope_row: Final = next(
guardrail for guardrail in response.guardrails if guardrail.guardrail_id == "invalid-scope-db-guardrail"
)
assert invalid_scope_row.litellm_params is not None
assert invalid_scope_row.litellm_params.logging_only_scope is None
assert invalid_scope_row.litellm_params.mode == "pre_call"
assert invalid_scope_row.litellm_params.guardrail == "litellm_content_filter"
assert any(guardrail.guardrail_id == "valid-scope-db-guardrail" for guardrail in response.guardrails)
@pytest.mark.asyncio
async def test_list_guardrails_v2_normalizes_invalid_scope_in_memory(
mocker, mock_prisma_client, mock_in_memory_handler
):
mock_prisma_client.db.litellm_guardrailstable.find_many.return_value = ()
mock_in_memory_handler.list_in_memory_guardrails.return_value = (_INVALID_SCOPE_IN_MEMORY_GUARDRAIL,)
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mocker.patch(
"litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER",
mock_in_memory_handler,
)
response: Final = await list_guardrails_v2(user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN))
assert len(response.guardrails) == 1
assert response.guardrails[0].litellm_params is not None
assert response.guardrails[0].litellm_params.logging_only_scope is None
assert response.guardrails[0].litellm_params.mode == "pre_call"
assert response.guardrails[0].litellm_params.guardrail == "litellm_content_filter"
@pytest.mark.asyncio
async def test_list_guardrails_v2_skips_stale_db_backed_in_memory_entries(mocker):
"""
@ -493,6 +582,28 @@ async def test_get_guardrail_info_from_db(mocker, mock_prisma_client):
assert response.guardrail_info == {"description": "Test guardrail from DB"}
@pytest.mark.asyncio
async def test_get_guardrail_info_normalizes_invalid_scope_from_db(
mocker, mock_guardrail_registry, mock_in_memory_handler
):
mock_guardrail_registry.get_guardrail_by_id_from_db.return_value = _INVALID_SCOPE_DB_GUARDRAIL
mock_in_memory_handler.get_guardrail_by_id.return_value = None
mocker.patch("litellm.proxy.proxy_server.prisma_client", mocker.Mock())
mocker.patch("litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_guardrail_registry)
mocker.patch(
"litellm.proxy.guardrails.guardrail_registry.IN_MEMORY_GUARDRAIL_HANDLER",
mock_in_memory_handler,
)
response: Final = await get_guardrail_info("invalid-scope-db-guardrail")
assert response.guardrail_id == "invalid-scope-db-guardrail"
assert response.litellm_params is not None
assert response.litellm_params.logging_only_scope is None
assert response.litellm_params.mode == "pre_call"
assert response.litellm_params.guardrail == "litellm_content_filter"
@pytest.mark.asyncio
async def test_get_guardrail_info_from_config(mocker, mock_prisma_client, mock_in_memory_handler):
"""Test getting guardrail info from config when not found in DB"""