mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
2056b73256
commit
8a35342461
4 changed files with 204 additions and 10 deletions
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
(
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue