fix(guardrails): tolerate invalid stored logging-only scopes

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-29 11:34:18 +00:00
parent 748dcaa593
commit 48eae4287e
5 changed files with 191 additions and 12 deletions

View file

@ -32,7 +32,11 @@ from litellm.proxy.guardrails.guardrail_hooks.custom_code.sandbox import (
build_sandbox_globals,
compile_sandboxed,
)
from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry, _configured_event_hooks
from litellm.proxy.guardrails.guardrail_registry import (
GuardrailRegistry,
_configured_event_hooks,
parse_tolerant_litellm_params,
)
from litellm.proxy.guardrails.usage_endpoints import router as guardrails_usage_router
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
from litellm.repositories.prisma_protocols import TableActions
@ -1212,7 +1216,10 @@ async def patch_guardrail(
# Update litellm_params if default_on is provided or pii_entities_config is provided
existing_litellm_params: Final = _as_str_object_mapping(dict(existing_guardrail.get("litellm_params", {})))
current_litellm_params: Final = LitellmParams(**existing_litellm_params)
current_litellm_params: Final = parse_tolerant_litellm_params(
existing_litellm_params,
existing_guardrail.get("guardrail_name") or "Unknown",
)
requested_litellm_params: Final = (
request.litellm_params.model_dump(exclude_unset=True) if request.litellm_params is not None else {}
)
@ -1281,7 +1288,10 @@ async def patch_guardrail(
guardrail=Guardrail(
guardrail_id=guardrail_id,
guardrail_name=existing_guardrail.get("guardrail_name") or "",
litellm_params=LitellmParams(**existing_litellm_params),
litellm_params=parse_tolerant_litellm_params(
existing_litellm_params,
existing_guardrail.get("guardrail_name") or "Unknown",
),
guardrail_info=existing_guardrail.get(
"guardrail_info",
{}, # mutable-ok: Guardrail's own constructor takes a plain dict

View file

@ -499,6 +499,24 @@ def _configure_callback_scoping(
_apply_configured_bool_overrides(custom_guardrail_callback, litellm_params)
def parse_tolerant_litellm_params(
litellm_params_data: Mapping[str, object],
guardrail_name: str,
) -> LitellmParams:
try:
return LitellmParams(**litellm_params_data)
except ValidationError as validation_error:
if any(tuple(error["loc"]) != ("logging_only_scope",) for error in validation_error.errors()):
raise
verbose_proxy_logger.error(
"Guardrail %s: logging_only_scope=%r is not one of 'input', 'output' or 'both'. "
"Ignoring logging_only_scope; the guardrail keeps its configured mode.",
guardrail_name.replace("\r", "").replace("\n", ""),
str(litellm_params_data.get("logging_only_scope")).replace("\r", "").replace("\n", "")[:100],
)
return LitellmParams(**{**litellm_params_data, "logging_only_scope": None})
class InMemoryGuardrailHandler:
"""
Class that handles initializing guardrails and adding them to the CallbackManager
@ -557,7 +575,10 @@ class InMemoryGuardrailHandler:
verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data)
if isinstance(litellm_params_data, dict):
litellm_params = LitellmParams(**litellm_params_data)
if reject_invalid_logging_only_scope:
litellm_params = LitellmParams(**litellm_params_data)
else:
litellm_params = parse_tolerant_litellm_params(litellm_params_data, guardrail["guardrail_name"])
else:
litellm_params = litellm_params_data
@ -803,6 +824,7 @@ class InMemoryGuardrailHandler:
@staticmethod
def _normalize_litellm_params_for_comparison(
params: LitellmParams | Mapping[str, object] | None,
guardrail_name: str,
) -> Mapping[str, object] | None:
"""
Render litellm_params to a canonical dict so an in-memory LitellmParams and
@ -819,7 +841,7 @@ class InMemoryGuardrailHandler:
return params.model_dump()
if isinstance(params, dict):
try:
return LitellmParams(**params).model_dump()
return parse_tolerant_litellm_params(params, guardrail_name).model_dump()
except ValidationError as e:
verbose_proxy_logger.warning(
"Could not normalize guardrail litellm_params for comparison; treating the guardrail as changed. Error: %s",
@ -842,8 +864,12 @@ class InMemoryGuardrailHandler:
return True
# Compare litellm_params
existing_dict: Final = self._normalize_litellm_params_for_comparison(existing.get("litellm_params"))
new_dict: Final = self._normalize_litellm_params_for_comparison(new_guardrail.get("litellm_params"))
existing_dict: Final = self._normalize_litellm_params_for_comparison(
existing.get("litellm_params"), existing.get("guardrail_name", "Unknown")
)
new_dict: Final = self._normalize_litellm_params_for_comparison(
new_guardrail.get("litellm_params"), new_guardrail.get("guardrail_name", "Unknown")
)
# Compare and identify specific differences
changed_fields = {}

View file

@ -1393,8 +1393,9 @@ def test_logging_only_scope_observes_only_the_configured_direction_without_block
assert detail["requestsEvaluated"] == len(scanned_directions), detail
def test_logging_only_scope_without_logging_only_mode_is_ignored_at_load_and_keeps_blocking(
gateway: Gateway, tmp_path: Path
@pytest.mark.parametrize("logging_only_scope", ("input", "Input"))
def test_logging_only_scope_literal_or_mode_mismatch_is_ignored_at_load_and_keeps_blocking(
gateway: Gateway, tmp_path: Path, logging_only_scope: str
) -> None:
identity: Final = "guardrail" + uuid.uuid4().hex
prompt: Final = "synthetic invalid-scope prompt pineapple " + identity
@ -1433,7 +1434,7 @@ def test_logging_only_scope_without_logging_only_mode_is_ignored_at_load_and_kee
"litellm_params": {
"guardrail": "generic_guardrail_api",
"mode": "pre_call",
"logging_only_scope": "input",
"logging_only_scope": logging_only_scope,
"default_on": True,
"blocked_words": [{"keyword": "pineapple", "action": "BLOCK"}],
"api_base": policy.url,

View file

@ -1448,6 +1448,35 @@ async def test_patch_guardrail_tolerates_stored_unsupported_scope_on_unrelated_u
handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"])
@pytest.mark.asyncio
async def test_patch_guardrail_tolerates_invalid_stored_scope_on_unrelated_update(
mocker, monkeypatch, mock_guardrail_registry
):
handler, stored_guardrail = _setup_patch_scope_guardrail(
mocker,
monkeypatch,
mock_guardrail_registry,
_PatchScopeUnsupportedGuardrail,
"patch_scope_invalid_literal_test",
{"mode": "logging_only", "logging_only_scope": "sideways", "default_on": True},
)
request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(default_on=False))
try:
result = await patch_guardrail(
stored_guardrail["guardrail_id"],
request,
user_api_key_dict=MOCK_ADMIN_USER,
)
assert result["guardrail_id"] == stored_guardrail["guardrail_id"]
persisted_guardrail = mock_guardrail_registry.update_guardrail_in_db.call_args.kwargs["guardrail"]
assert persisted_guardrail["litellm_params"].logging_only_scope is None
assert persisted_guardrail["litellm_params"].default_on is False
finally:
handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"])
@pytest.mark.asyncio
async def test_patch_guardrail_rejects_explicit_unsupported_scope_and_rolls_back(
mocker, monkeypatch, mock_guardrail_registry
@ -1480,6 +1509,36 @@ async def test_patch_guardrail_rejects_explicit_unsupported_scope_and_rolls_back
handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"])
@pytest.mark.asyncio
async def test_patch_guardrail_rolls_back_invalid_stored_scope_tolerantly(mocker, monkeypatch, mock_guardrail_registry):
handler, stored_guardrail = _setup_patch_scope_guardrail(
mocker,
monkeypatch,
mock_guardrail_registry,
_PatchScopeUnsupportedGuardrail,
"patch_scope_invalid_rollback_test",
{"mode": "logging_only", "logging_only_scope": "sideways", "default_on": True},
)
request = PatchGuardrailRequest(litellm_params=BaseLitellmParams(logging_only_scope="output"))
try:
with pytest.raises(HTTPException) as exc_info:
await patch_guardrail(
stored_guardrail["guardrail_id"],
request,
user_api_key_dict=MOCK_ADMIN_USER,
)
assert exc_info.value.status_code == 422
assert mock_guardrail_registry.update_guardrail_in_db.call_count == 2
restored_guardrail = mock_guardrail_registry.update_guardrail_in_db.call_args_list[-1].kwargs["guardrail"]
assert restored_guardrail["litellm_params"].logging_only_scope is None
callback = handler.guardrail_id_to_custom_guardrail[stored_guardrail["guardrail_id"]]
assert callback.logging_only_scope is None
finally:
handler.delete_in_memory_guardrail(stored_guardrail["guardrail_id"])
@pytest.mark.parametrize(
"scenario,expected_result,expected_exception",
[

View file

@ -7,11 +7,12 @@ from pydantic import ValidationError
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy.guardrails.guardrail_registry import (
get_guardrail_initializer_from_hooks,
GuardrailRegistry,
InMemoryGuardrailHandler,
get_guardrail_initializer_from_hooks,
parse_tolerant_litellm_params,
)
from litellm.types.guardrails import GuardrailEventHooks, Guardrail, LitellmParams, LoggingOnlyScope, Mode
from litellm.types.guardrails import Guardrail, GuardrailEventHooks, LitellmParams, LoggingOnlyScope, Mode
from litellm.types.utils import GenericGuardrailAPIInputs
@ -475,6 +476,24 @@ def test_unnormalizable_db_params_register_as_changed_without_raising():
assert handler._has_guardrail_params_changed(gid, new) is True
def test_invalid_scope_literal_db_params_compare_equal_after_normalization():
handler = InMemoryGuardrailHandler()
raw = _db_litellm_params()
gid = "77777777-7777-7777-7777-777777777777"
handler.IN_MEMORY_GUARDRAILS[gid] = Guardrail(
guardrail_id=gid,
guardrail_name="cf",
litellm_params=LitellmParams(**{**raw, "logging_only_scope": None}),
)
new = Guardrail(
guardrail_id=gid,
guardrail_name="cf",
litellm_params={**raw, "logging_only_scope": "Input"},
)
assert handler._has_guardrail_params_changed(gid, new) is False
def _all_callback_lists():
import litellm
@ -960,6 +979,19 @@ class _LoggingOnlyScopeNativeGuardrail(_LoggingOnlyScopeSupportedGuardrail):
use_native_lifecycle_hooks: ClassVar[bool] = True
def _invalid_scope_content_filter_guardrail() -> Guardrail:
return Guardrail(
guardrail_id="invalid-scope-content-filter-test",
guardrail_name="invalid-scope-content-filter",
litellm_params={
"guardrail": "litellm_content_filter",
"mode": "pre_call",
"logging_only_scope": "Input",
"blocked_words": [{"keyword": "pineapple", "action": "BLOCK"}],
},
)
class TestLoggingOnlyScopeValidation:
def _initialize(
self,
@ -1094,6 +1126,57 @@ class TestLoggingOnlyScopeValidation:
with pytest.raises(ValidationError):
LitellmParams(guardrail="test", mode="logging_only", logging_only_scope="request")
def test_invalid_scope_literal_keeps_content_filter_registered_and_blocking(self) -> None:
import litellm
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.content_filter import (
ContentFilterGuardrail,
)
handler: Final = InMemoryGuardrailHandler()
callback_lists: Final = _all_callback_lists()
callback_snapshots: Final = [list(callback_list) for callback_list in callback_lists]
guardrail: Final = _invalid_scope_content_filter_guardrail()
try:
result: Final = handler.initialize_guardrail(guardrail=guardrail, source="config")
assert result is not None
callback: Final = handler.guardrail_id_to_custom_guardrail[result["guardrail_id"]]
assert isinstance(callback, ContentFilterGuardrail)
assert callback in litellm.callbacks
assert callback.logging_only_scope is None
assert callback.event_hook == GuardrailEventHooks.pre_call
assert callback._check_blocked_words("pineapple") is not None
finally:
handler.delete_in_memory_guardrail(guardrail["guardrail_id"])
for callback_list, snapshot in zip(callback_lists, callback_snapshots):
callback_list[:] = snapshot
def test_invalid_scope_literal_is_rejected_for_strict_initialization_without_callback_leakage(self) -> None:
handler: Final = InMemoryGuardrailHandler()
callback_lists: Final = _all_callback_lists()
callback_snapshots: Final = [list(callback_list) for callback_list in callback_lists]
with pytest.raises(ValueError):
handler.initialize_guardrail(
guardrail=_invalid_scope_content_filter_guardrail(),
source="config",
reject_invalid_logging_only_scope=True,
)
assert all(callback_list == snapshot for callback_list, snapshot in zip(callback_lists, callback_snapshots))
def test_invalid_scope_literal_does_not_tolerate_other_litellm_params_errors(self) -> None:
with pytest.raises(ValidationError):
parse_tolerant_litellm_params(
{
"guardrail": "litellm_content_filter",
"mode": "pre_call",
"logging_only_scope": "Input",
"default_on": "not-a-bool",
},
"invalid-scope-content-filter",
)
@pytest.mark.asyncio
async def test_update_guardrail_in_db_raises_when_row_missing():