mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(guardrails): fail closed on non-text request masks + bound Alice mask reconstruction
Addresses two review findings on the request-side join. Non-text masks are no longer discarded: when a MASK redacts a detection-only piece (tool-call args or a tool / function description), reconstruction recovers it but those pieces cannot be spliced back into the wire format, so forwarding the original unredacted value would leak it. _scan_request now compares the recovered detection-only pieces against the originals and fails closed (block) when they differ, instead of slicing them off and forwarding the originals. Reconstruction is now bounded: difflib.SequenceMatcher(autojunk=False) is O(n*m) worst case and runs synchronously on the event loop, and the scan budget allows up to a million characters, so a large repetitive MASK-triggering prompt could wedge the loop. reconstruct() now fails closed (returns None -> block) when the joined document exceeds RECONSTRUCT_MAX_CHARS (two chunks' worth), which keeps the worst-case alignment sub-second while still covering ordinary multi-message chats. Non-MASK requests of any size are unaffected.
This commit is contained in:
parent
8bceb610c8
commit
e7b0c9232c
5 changed files with 106 additions and 36 deletions
|
|
@ -358,8 +358,21 @@ class WonderFenceGuardrail(CustomGuardrail):
|
|||
recovered = reconstruct(pieces, verdict.masked_text or "")
|
||||
if recovered is None:
|
||||
logger.warning(
|
||||
"Alice WonderFence (apply_guardrail request): MASK reconstruction failed "
|
||||
"(a joiner or part boundary landed inside a masked span); failing closed. guardrail=%s correlation_id=%s",
|
||||
"Alice WonderFence (apply_guardrail request): MASK reconstruction unavailable "
|
||||
"(document too large, or a joiner / part boundary landed inside a masked span); "
|
||||
"failing closed. guardrail=%s correlation_id=%s",
|
||||
self.guardrail_name,
|
||||
correlation_id,
|
||||
)
|
||||
raise WonderFenceBlockedError(block_detail([verdict], self.guardrail_name, self.block_message))
|
||||
if recovered[n_text:] != pieces[n_text:]:
|
||||
# The mask redacted a detection-only piece (tool-call args or a
|
||||
# tool / function description). Those are not maskable in place
|
||||
# (the joined form is not the wire format), so we cannot forward
|
||||
# the original unredacted value; fail closed rather than leak it.
|
||||
logger.warning(
|
||||
"Alice WonderFence (apply_guardrail request): MASK landed in a detection-only "
|
||||
"piece (tool-call args / tool or function description); failing closed. guardrail=%s correlation_id=%s",
|
||||
self.guardrail_name,
|
||||
correlation_id,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ import litellm
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
from .chunked_evaluation import SegmentVerdict
|
||||
from .chunked_evaluation import MAX_PROMPT_CHARS, SegmentVerdict
|
||||
from .credentials import get_metadata
|
||||
from .exceptions import WonderFenceBlockedError, WonderFenceScanBudgetExceeded
|
||||
|
||||
|
|
@ -18,6 +18,14 @@ logger = verbose_proxy_logger.getChild("alice_wonderfence")
|
|||
|
||||
JOINER = "\n"
|
||||
|
||||
# Upper bound on the document ``reconstruct`` will align. ``SequenceMatcher`` is
|
||||
# O(n*m) worst case and runs synchronously on the event loop, so a large
|
||||
# repetitive MASK-triggering prompt could otherwise wedge it. MASK on a document
|
||||
# larger than this fails closed (block) rather than run the quadratic alignment;
|
||||
# non-MASK requests of any size are unaffected. Two chunks' worth keeps the
|
||||
# worst case sub-second while still covering ordinary multi-message chats.
|
||||
RECONSTRUCT_MAX_CHARS = 2 * MAX_PROMPT_CHARS
|
||||
|
||||
|
||||
def build_analysis_context(
|
||||
request_data: dict,
|
||||
|
|
@ -195,18 +203,21 @@ def reconstruct(parts: list[str], masked: str) -> list[str] | None:
|
|||
document. We align original-vs-masked with ``difflib.SequenceMatcher`` (no
|
||||
sentinel injected) and map each part's char range through the alignment.
|
||||
|
||||
Fails closed (returns ``None``) when the structure is not recoverable: every
|
||||
``JOINER`` between parts must survive the mask as an unmodified ``\\n`` (a
|
||||
mask spanning a joiner would merge parts), and no part boundary may land
|
||||
inside a changed block. Returns one masked string per input part, in order;
|
||||
``[]`` for no parts. Assumes masking is span substitution that preserves the
|
||||
non-masked characters; if the service reflows whitespace the joiner-survival
|
||||
check trips and we fail closed rather than misassign.
|
||||
Fails closed (returns ``None``) when the structure is not recoverable: the
|
||||
document exceeds ``RECONSTRUCT_MAX_CHARS`` (bounds the quadratic alignment
|
||||
cost); any ``JOINER`` between parts does not survive the mask as an
|
||||
unmodified ``\\n`` (a mask spanning a joiner would merge parts); or a part
|
||||
boundary lands inside a changed block. Returns one masked string per input
|
||||
part, in order; ``[]`` for no parts. Assumes masking is span substitution
|
||||
that preserves the non-masked characters; if the service reflows whitespace
|
||||
the joiner-survival check trips and we fail closed rather than misassign.
|
||||
"""
|
||||
if not parts:
|
||||
return []
|
||||
|
||||
original = JOINER.join(parts)
|
||||
if len(original) > RECONSTRUCT_MAX_CHARS or len(masked) > RECONSTRUCT_MAX_CHARS:
|
||||
return None
|
||||
starts = [0, *accumulate(len(p) + len(JOINER) for p in parts)][: len(parts)]
|
||||
ranges = [(s, s + len(p)) for s, p in zip(starts, parts)]
|
||||
joiners = [end for (_s, end) in ranges[:-1]]
|
||||
|
|
|
|||
|
|
@ -498,3 +498,37 @@ async def test_apply_guardrail_cap_is_not_bypassed_by_fail_open(make_guardrail,
|
|||
assert exc.value.status_code == 400
|
||||
assert exc.value.detail["limit"] == "max_scan_chars"
|
||||
client.evaluate_prompt.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_mask_on_oversized_document_fails_closed(make_guardrail, make_request_data):
|
||||
"""A MASK on a document too large to reconstruct within the bounded
|
||||
alignment cost fails closed (block) rather than run the quadratic
|
||||
SequenceMatcher on the event loop or forward unmasked content."""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.processing import (
|
||||
RECONSTRUCT_MAX_CHARS,
|
||||
)
|
||||
|
||||
# Keep the char cap high enough to reach scanning, but exceed the
|
||||
# reconstruction bound so MASK cannot be applied.
|
||||
guardrail, client = make_guardrail(max_scan_chars=RECONSTRUCT_MAX_CHARS * 2)
|
||||
guardrail._client_cache["default-api-key"] = client
|
||||
big = "a " * RECONSTRUCT_MAX_CHARS # > RECONSTRUCT_MAX_CHARS chars
|
||||
|
||||
def evaluate(prompt, **kwargs):
|
||||
r = Mock()
|
||||
r.action = "MASK"
|
||||
r.action_text = "[REDACTED]"
|
||||
r.detections = []
|
||||
r.correlation_id = None
|
||||
return r
|
||||
|
||||
client.evaluate_prompt.side_effect = evaluate
|
||||
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": [big]},
|
||||
request_data=make_request_data(),
|
||||
input_type="request",
|
||||
)
|
||||
assert exc.value.status_code == 400
|
||||
|
|
|
|||
|
|
@ -44,11 +44,11 @@ async def test_apply_guardrail_blocks_on_tool_call_arguments(guardrail_and_clien
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_tool_call_args_are_detection_only(guardrail_and_client, make_request_data):
|
||||
"""On the request side, tool-call args are rendered into the joined document
|
||||
as detection-only pieces: they can BLOCK/DETECT but a MASK is never spliced
|
||||
back into the arguments string (the joined form is not the wire format).
|
||||
Message text still masks; the args survive untouched."""
|
||||
async def test_apply_guardrail_request_tool_call_args_mask_fails_closed(guardrail_and_client, make_request_data):
|
||||
"""On the request side, tool-call args are detection-only pieces in the join.
|
||||
A MASK that redacts an arg cannot be spliced back into the wire-format
|
||||
arguments string, so forwarding the original unredacted value would leak it;
|
||||
the request fails closed (block) instead."""
|
||||
guardrail, client = guardrail_and_client
|
||||
|
||||
def evaluate(prompt, **kwargs):
|
||||
|
|
@ -65,14 +65,13 @@ async def test_apply_guardrail_request_tool_call_args_are_detection_only(guardra
|
|||
"texts": ["benign"],
|
||||
"tool_calls": [_tool_call('{"body": "secret value"}')],
|
||||
}
|
||||
out = await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=make_request_data(),
|
||||
input_type="request",
|
||||
)
|
||||
assert out["tool_calls"][0]["function"]["arguments"] == '{"body": "secret value"}'
|
||||
assert out["texts"] == ["benign"]
|
||||
client.evaluate_prompt.assert_awaited_once()
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await guardrail.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=make_request_data(),
|
||||
input_type="request",
|
||||
)
|
||||
assert exc.value.status_code == 400
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -211,9 +210,10 @@ async def test_apply_guardrail_blocks_on_tool_parameter_description(guardrail_an
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_tool_definitions_are_detection_only(guardrail_and_client, make_request_data):
|
||||
"""Tool definitions are scanned detection-only (they can BLOCK/DETECT) but a
|
||||
MASK is never written back into the schema; the description survives."""
|
||||
async def test_apply_guardrail_tool_definition_mask_fails_closed(guardrail_and_client, make_request_data):
|
||||
"""Tool definitions are detection-only; a MASK that would redact a
|
||||
description cannot be spliced back into the schema, so the request fails
|
||||
closed rather than forward the original unredacted description."""
|
||||
guardrail, client = guardrail_and_client
|
||||
|
||||
def evaluate(prompt, **kwargs):
|
||||
|
|
@ -230,8 +230,9 @@ async def test_apply_guardrail_tool_definitions_are_detection_only(guardrail_and
|
|||
"texts": ["hi"],
|
||||
"tools": [_tool_def(description="contains secret stuff")],
|
||||
}
|
||||
out = await guardrail.apply_guardrail(inputs=inputs, request_data=make_request_data(), input_type="request")
|
||||
assert out["tools"][0]["function"]["description"] == "contains secret stuff"
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await guardrail.apply_guardrail(inputs=inputs, request_data=make_request_data(), input_type="request")
|
||||
assert exc.value.status_code == 400
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -366,9 +367,9 @@ async def test_apply_guardrail_legacy_function_detect_does_not_mutate(guardrail_
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_legacy_function_definitions_are_detection_only(guardrail_and_client, make_request_data):
|
||||
"""Legacy functions[] descriptions are scanned detection-only; a MASK is
|
||||
never spliced back into request_data['functions']."""
|
||||
async def test_apply_guardrail_legacy_function_definition_mask_fails_closed(guardrail_and_client, make_request_data):
|
||||
"""Legacy functions[] descriptions are detection-only; a MASK that would
|
||||
redact one fails closed rather than forward the original unredacted value."""
|
||||
guardrail, client = guardrail_and_client
|
||||
|
||||
def evaluate(prompt, **kwargs):
|
||||
|
|
@ -382,9 +383,11 @@ async def test_apply_guardrail_legacy_function_definitions_are_detection_only(gu
|
|||
client.evaluate_prompt.side_effect = evaluate
|
||||
|
||||
request_data = make_request_data(functions=[_legacy_function(description="contains secret stuff")])
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["hi"]},
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await guardrail.apply_guardrail(
|
||||
inputs={"texts": ["hi"]},
|
||||
request_data=request_data,
|
||||
input_type="request",
|
||||
)
|
||||
assert exc.value.status_code == 400
|
||||
assert request_data["functions"][0]["description"] == "contains secret stuff"
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.exceptions impor
|
|||
)
|
||||
from litellm.proxy.guardrails.guardrail_hooks.alice_wonderfence.processing import (
|
||||
JOINER,
|
||||
RECONSTRUCT_MAX_CHARS,
|
||||
apply_response_verdicts,
|
||||
check_scan_budget,
|
||||
function_definition_segments,
|
||||
|
|
@ -63,6 +64,14 @@ def test_reconstruct_empty_parts_is_empty_list():
|
|||
assert reconstruct([], "") == []
|
||||
|
||||
|
||||
def test_reconstruct_fails_closed_when_document_exceeds_bound():
|
||||
"""Reconstruction is bounded to avoid the quadratic SequenceMatcher cost
|
||||
blocking the event loop; an over-bound document fails closed (None) instead
|
||||
of running the alignment."""
|
||||
big = "a" * (RECONSTRUCT_MAX_CHARS + 1)
|
||||
assert reconstruct([big], big) is None
|
||||
|
||||
|
||||
# --------------- check_scan_budget (total-work cap) ---------------
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue