mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
748dcaa593
commit
48eae4287e
5 changed files with 191 additions and 12 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 = {}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue