From 48eae4287eb49154f4091cf6bfd79d8bdb3d657b Mon Sep 17 00:00:00 2001 From: yucheng Date: Tue, 29 Sep 2026 11:34:18 +0000 Subject: [PATCH] fix(guardrails): tolerate invalid stored logging-only scopes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/guardrails/guardrail_endpoints.py | 16 +++- .../proxy/guardrails/guardrail_registry.py | 34 +++++++- .../observability/test_guardrail_effects.py | 7 +- .../guardrails/test_guardrail_endpoints.py | 59 +++++++++++++ .../guardrails/test_guardrail_registry.py | 87 ++++++++++++++++++- 5 files changed, 191 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index c3ed2fd8cc9..01fb2bb47c6 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index 404468a387d..e4f00cdf80b 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -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 = {} diff --git a/tests/integration/observability/test_guardrail_effects.py b/tests/integration/observability/test_guardrail_effects.py index cf182188e86..a1e0c10ce5f 100644 --- a/tests/integration/observability/test_guardrail_effects.py +++ b/tests/integration/observability/test_guardrail_effects.py @@ -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, diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 12310f5d67c..a09b746d4db 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -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", [ diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py index 56cfdd120f2..be01d8ee89f 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_registry.py @@ -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():