mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(model_armor): let a streamed de-identify match mask instead of blocking
A de-identify template reports MATCH_FOUND for every redaction it makes. The streaming block check omitted allow_sanitization, so with mask_response_content enabled that match read as a refusal and the client got a 400 where the non-streaming sibling returned the redacted text. Pass the flag through, as the non-streaming hook already does, and stamp the logged status from the same decision so the spend row agrees with what the client received. Also drop Any from the chat-completion assembler's parameter; stream_chunk_builder takes a bare list, so list[object] carries the mutability requirement without erasing the element type.
This commit is contained in:
parent
8bd62b4366
commit
c3e7224ee2
2 changed files with 119 additions and 7 deletions
|
|
@ -941,7 +941,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
|
||||
@staticmethod
|
||||
def _assemble_chat_completion_stream(
|
||||
all_chunks: list[Any], # mutable-ok: stream_chunk_builder only accepts a mutable list
|
||||
all_chunks: list[object], # mutable-ok: stream_chunk_builder only accepts a mutable list
|
||||
) -> ModelResponse | TextCompletionResponse | None:
|
||||
"""Assemble chat-completion chunks, returning ``None`` when they cannot be assembled."""
|
||||
from litellm.main import stream_chunk_builder
|
||||
|
|
@ -1065,14 +1065,19 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
|
|||
if isinstance(request_data, dict):
|
||||
_, metadata = get_or_create_metadata_bucket(request_data)
|
||||
metadata["_model_armor_response"] = self._build_logging_response(armor_response)
|
||||
metadata["_model_armor_status"] = "blocked" if self._should_block_content(armor_response) else "success"
|
||||
metadata["_model_armor_status"] = (
|
||||
"blocked"
|
||||
if self._should_block_content(armor_response, allow_sanitization=self.mask_response_content)
|
||||
else "success"
|
||||
)
|
||||
|
||||
# Add guardrail to applied_guardrails BEFORE potential blocking
|
||||
# This ensures guardrail is recorded even when it blocks the request
|
||||
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
|
||||
|
||||
# Check if blocked
|
||||
if self._should_block_content(armor_response):
|
||||
# Check if blocked. Mirrors the non-streaming sibling: with masking on, a de-identify
|
||||
# match is a redaction to apply below, not a refusal
|
||||
if self._should_block_content(armor_response, allow_sanitization=self.mask_response_content):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=self._build_block_error_detail(
|
||||
|
|
|
|||
|
|
@ -3814,14 +3814,54 @@ _MODEL_ARMOR_BLOCK = {
|
|||
}
|
||||
}
|
||||
|
||||
# The streaming hook checks _should_block_content without allow_sanitization, so a
|
||||
# deidentifyResult MATCH_FOUND blocks rather than masks. The root-level sanitizedText
|
||||
# fallback in _get_sanitized_content is the shape that reaches the masking branch.
|
||||
# The root-level sanitizedText fallback in _get_sanitized_content, i.e. a rewrite that trips no
|
||||
# named filter
|
||||
_MODEL_ARMOR_SANITIZED = {
|
||||
"sanitizedText": "my card is [REDACTED]",
|
||||
"sanitizationResult": {"filterMatchState": "NO_MATCH_FOUND"},
|
||||
}
|
||||
|
||||
# The shape a real de-identify template returns: the SDP filter both matches and hands back the
|
||||
# rewritten text, so whether it blocks or masks is decided by allow_sanitization alone
|
||||
_MODEL_ARMOR_DEIDENTIFIED = {
|
||||
"sanitizationResult": {
|
||||
"filterMatchState": "MATCH_FOUND",
|
||||
"filterResults": {
|
||||
"sdp": {
|
||||
"sdpFilterResult": {
|
||||
"deidentifyResult": {
|
||||
"matchState": "MATCH_FOUND",
|
||||
"data": {"text": "my card is [REDACTED]"},
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def _chat_completion_chunks():
|
||||
"""The chat-completions surface: typed ModelResponseStream chunks."""
|
||||
return (
|
||||
litellm.types.utils.ModelResponseStream(
|
||||
choices=[
|
||||
litellm.types.utils.StreamingChoices(
|
||||
index=0,
|
||||
delta=litellm.types.utils.Delta(content="my card is 4111-1111-1111-1111"),
|
||||
)
|
||||
]
|
||||
),
|
||||
litellm.types.utils.ModelResponseStream(
|
||||
choices=[
|
||||
litellm.types.utils.StreamingChoices(
|
||||
index=0,
|
||||
delta=litellm.types.utils.Delta(content=""),
|
||||
finish_reason="stop",
|
||||
)
|
||||
]
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _surface_guardrail(**kwargs):
|
||||
guardrail = ModelArmorGuardrail(
|
||||
|
|
@ -4391,3 +4431,70 @@ async def test_streaming_hook_does_not_forward_typed_chunks_that_end_with_an_err
|
|||
body = b"".join(item if isinstance(item, bytes) else str(item).encode() for item in delivered)
|
||||
assert b"4111-1111-1111-1111" not in body
|
||||
assert b"could not be assembled for scanning" in body
|
||||
|
||||
|
||||
def _delivered_bytes(delivered):
|
||||
return b"".join(
|
||||
item
|
||||
if isinstance(item, bytes)
|
||||
else item.encode()
|
||||
if isinstance(item, str)
|
||||
else str(item.model_dump() if hasattr(item, "model_dump") else item).encode()
|
||||
for item in delivered
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("chunks, case", [(None, "chat_completions"), (_ANTHROPIC_SSE_CHUNKS, "anthropic_sse")])
|
||||
async def test_streaming_deidentify_match_masks_when_masking_is_enabled(chunks, case):
|
||||
"""A de-identify template reports MATCH_FOUND for every redaction it makes, so reading that
|
||||
match as a refusal makes mask_response_content unusable on a stream: the client gets an error
|
||||
where its non-streaming sibling gets redacted text. The block check has to allow sanitization
|
||||
exactly as the non-streaming hook does."""
|
||||
guardrail = _surface_guardrail(mask_response_content=True)
|
||||
post = _armor_post_mock(_MODEL_ARMOR_DEIDENTIFIED)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", post):
|
||||
delivered = await _drain_surface_hook(
|
||||
guardrail, _chat_completion_chunks() if chunks is None else chunks
|
||||
)
|
||||
|
||||
body = _delivered_bytes(delivered)
|
||||
assert b"[REDACTED]" in body, case
|
||||
assert b"4111-1111-1111-1111" not in body, case
|
||||
assert b"blocked by Model Armor" not in body, case
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("chunks, case", [(None, "chat_completions"), (_ANTHROPIC_SSE_CHUNKS, "anthropic_sse")])
|
||||
async def test_streaming_deidentify_match_still_blocks_when_masking_is_disabled(chunks, case):
|
||||
"""Without mask_response_content there is nowhere to put the rewritten text, so the same
|
||||
de-identify match must still end the stream rather than release the original."""
|
||||
guardrail = _surface_guardrail()
|
||||
post = _armor_post_mock(_MODEL_ARMOR_DEIDENTIFIED)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", post):
|
||||
delivered = await _drain_surface_hook(
|
||||
guardrail, _chat_completion_chunks() if chunks is None else chunks
|
||||
)
|
||||
|
||||
body = _delivered_bytes(delivered)
|
||||
assert b"Streaming response blocked by Model Armor" in body, case
|
||||
assert b"4111-1111-1111-1111" not in body, case
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streaming_deidentify_match_logs_masked_run_as_success_not_blocked():
|
||||
"""The status stamped on request metadata feeds the spend log, so it has to agree with what
|
||||
the client actually received: a masked stream is a success, not a block."""
|
||||
guardrail = _surface_guardrail(mask_response_content=True)
|
||||
request_data = {
|
||||
"model": "claude-haiku",
|
||||
"messages": [{"role": "user", "content": "show me a card"}],
|
||||
"metadata": {"guardrails": ["model-armor-test"]},
|
||||
}
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", _armor_post_mock(_MODEL_ARMOR_DEIDENTIFIED)):
|
||||
await _drain_surface_hook(guardrail, _chat_completion_chunks(), request_data=request_data)
|
||||
|
||||
assert request_data["metadata"]["_model_armor_status"] == "success"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue