From 8a3534246113db35d9e95296dcb69cdeb0ca1419 Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 1 Oct 2026 05:47:10 +0000 Subject: [PATCH] fix(guardrails): tolerate invalid stored params on reads Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/guardrails/guardrail_endpoints.py | 19 ++- .../proxy/guardrails/guardrail_registry.py | 14 ++- .../test_logging_only_scope_config.py | 68 +++++++++++ .../guardrails/test_guardrail_endpoints.py | 113 +++++++++++++++++- 4 files changed, 204 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index 8c7212ae376..113ddfe8f06 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -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 ) diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 01ded884500..a2cbde95c36 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -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: diff --git a/tests/integration/observability/test_logging_only_scope_config.py b/tests/integration/observability/test_logging_only_scope_config.py index 9ef53da04b4..c3b82caf8d2 100644 --- a/tests/integration/observability/test_logging_only_scope_config.py +++ b/tests/integration/observability/test_logging_only_scope_config.py @@ -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"), ( diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index d9d306602f2..310c4f5a386 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -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"""