mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(guardrails): never drop an Akto attachment the text check removed
- optional metadata (filename, title, format, media type) that isn't a string is ignored instead of failing the block - an attachment block that still can't be read blocks the request - search_result text is checked once, as a file, instead of also in the text check
This commit is contained in:
parent
72188daa71
commit
b9ff03a8c6
4 changed files with 97 additions and 19 deletions
|
|
@ -77,6 +77,7 @@ MCP_PATH: Final = "/mcp"
|
|||
MCP_TOOL_PREFIX: Final = "mcp"
|
||||
DEFAULT_BLOCK_REASON: Final = "Blocked by Akto Guardrails"
|
||||
UNMASKABLE_REASON: Final = "Content masked by Akto guardrail policy could not be applied"
|
||||
MALFORMED_ATTACHMENT_REASON: Final = "Attachment could not be read for the Akto guardrail check"
|
||||
UNREACHABLE_REASON: Final = "Akto guardrail service unreachable"
|
||||
BLOCKING_BEHAVIOURS: Final = frozenset(("block", ""))
|
||||
SESSION_ID_HEADER: Final = "x-akto-installer-akto_session_id"
|
||||
|
|
@ -684,6 +685,8 @@ class AktoGuardrail(CustomGuardrail):
|
|||
async def check_attachments(self, inputs: GenericGuardrailAPIInputs, request_data: Mapping[str, object]) -> None:
|
||||
"""Attachments can't be put back masked, so masking blocks."""
|
||||
found: Final = request_attachments(request_data)
|
||||
if found.malformed_count:
|
||||
raise self.blocked(MALFORMED_ATTACHMENT_REASON, streamed=False)
|
||||
if found.unsendable_count:
|
||||
verbose_proxy_logger.warning(
|
||||
"Akto: %d attachment(s) have no inline content or URL to check", found.unsendable_count
|
||||
|
|
|
|||
|
|
@ -18,14 +18,14 @@ from types import MappingProxyType
|
|||
from typing import Annotated, Final, Literal, TypeAlias, TypeVar
|
||||
from urllib.parse import unquote, unquote_to_bytes, urlparse
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError
|
||||
from pydantic import BaseModel, BeforeValidator, ConfigDict, Field, TypeAdapter, ValidationError
|
||||
|
||||
AttachmentType: TypeAlias = Literal["image", "audio", "file"]
|
||||
|
||||
_REMOTE_URI_SCHEMES: Final = ("http://", "https://")
|
||||
_URL_SAFE_TO_STANDARD: Final = str.maketrans("-_", "+/")
|
||||
_ATTACHMENT_BLOCK_TYPES: Final = frozenset(
|
||||
("image_url", "input_image", "input_audio", "file", "input_file", "image", "document", "video_url")
|
||||
("image_url", "input_image", "input_audio", "file", "input_file", "image", "document", "video_url", "search_result")
|
||||
)
|
||||
_OBJECT_MAPPING: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object])
|
||||
|
||||
|
|
@ -48,6 +48,15 @@ class Attachment:
|
|||
class RequestAttachments:
|
||||
attachments: tuple[Attachment, ...]
|
||||
unsendable_count: int
|
||||
malformed_count: int = 0
|
||||
|
||||
|
||||
def _text_or_none(value: object) -> object:
|
||||
return value if isinstance(value, str) else None
|
||||
|
||||
|
||||
# Optional metadata the provider ignores when malformed, so a bad value must not fail the whole block
|
||||
_Metadata: TypeAlias = Annotated[str | None, BeforeValidator(_text_or_none)]
|
||||
|
||||
|
||||
class _Model(BaseModel):
|
||||
|
|
@ -76,7 +85,7 @@ class _InputImageBlock(_Model):
|
|||
|
||||
class _InputAudio(_Model):
|
||||
data: str | None = None
|
||||
format: str | None = None
|
||||
format: _Metadata = None
|
||||
|
||||
|
||||
class _InputAudioBlock(_Model):
|
||||
|
|
@ -87,7 +96,7 @@ class _InputAudioBlock(_Model):
|
|||
class _FileData(_Model):
|
||||
file_data: str | None = None
|
||||
file_id: str | None = None
|
||||
filename: str | None = None
|
||||
filename: _Metadata = None
|
||||
|
||||
|
||||
class _FileBlock(_Model):
|
||||
|
|
@ -100,13 +109,13 @@ class _InputFileBlock(_Model):
|
|||
file_data: str | None = None
|
||||
file_url: str | None = None
|
||||
file_id: str | None = None
|
||||
filename: str | None = None
|
||||
filename: _Metadata = None
|
||||
|
||||
|
||||
class _Source(_Model):
|
||||
type: str | None = None
|
||||
type: _Metadata = None
|
||||
data: str | None = None
|
||||
media_type: str | None = None
|
||||
media_type: _Metadata = None
|
||||
url: str | None = None
|
||||
content: object = None
|
||||
|
||||
|
|
@ -124,12 +133,12 @@ class _ImageBlock(_Model):
|
|||
class _DocumentBlock(_Model):
|
||||
type: Literal["document"]
|
||||
source: _Source
|
||||
title: str | None = None
|
||||
title: _Metadata = None
|
||||
|
||||
|
||||
class _SearchResultBlock(_Model):
|
||||
type: Literal["search_result"]
|
||||
title: str | None = None
|
||||
title: _Metadata = None
|
||||
content: object = None
|
||||
|
||||
|
||||
|
|
@ -138,6 +147,10 @@ class _ToolResultBlock(_Model):
|
|||
content: object = None
|
||||
|
||||
|
||||
class _MalformedBlock(_Model):
|
||||
"""An attachment type that doesn't parse; it can't be checked, so it blocks."""
|
||||
|
||||
|
||||
class _Message(_Model):
|
||||
content: object = None
|
||||
output: object = None
|
||||
|
|
@ -158,6 +171,7 @@ _AttachmentBlock: TypeAlias = (
|
|||
_BLOCK_ADAPTER: Final[TypeAdapter[_AttachmentBlock]] = TypeAdapter(
|
||||
Annotated[_AttachmentBlock, Field(discriminator="type")]
|
||||
)
|
||||
_Block: TypeAlias = _AttachmentBlock | _MalformedBlock
|
||||
_TEXT_BLOCK_ADAPTER: Final[TypeAdapter[_TextBlock]] = TypeAdapter(_TextBlock)
|
||||
_MESSAGE_ADAPTER: Final[TypeAdapter[_Message]] = TypeAdapter(_Message)
|
||||
_ITEMS_ADAPTER: Final[TypeAdapter[list[object]]] = TypeAdapter(list[object])
|
||||
|
|
@ -171,17 +185,18 @@ _UNSENDABLE: Final[_Classified] = (None, True)
|
|||
def request_attachments(request_data: Mapping[str, object]) -> RequestAttachments:
|
||||
# Both, so a decoy "messages" can't hide attachments in a Responses API "input"
|
||||
containers: Final = (_parse(_ITEMS_ADAPTER, request_data.get(key)) or () for key in ("messages", "input"))
|
||||
blocks: Final = chain.from_iterable(_message_blocks(message) for message in chain.from_iterable(containers))
|
||||
blocks: Final = tuple(chain.from_iterable(_message_blocks(message) for message in chain.from_iterable(containers)))
|
||||
classified: Final = tuple(
|
||||
chain.from_iterable(_block_attachments(block, index) for index, block in enumerate(blocks))
|
||||
)
|
||||
return RequestAttachments(
|
||||
attachments=tuple(attachment for attachment, _ in classified if attachment is not None),
|
||||
unsendable_count=sum(1 for _, is_unsendable in classified if is_unsendable),
|
||||
malformed_count=sum(1 for block in blocks if isinstance(block, _MalformedBlock)),
|
||||
)
|
||||
|
||||
|
||||
def _message_blocks(message: object) -> tuple[_AttachmentBlock, ...]:
|
||||
def _message_blocks(message: object) -> tuple[_Block, ...]:
|
||||
parsed: Final = _parse(_MESSAGE_ADAPTER, message)
|
||||
top: Final = (_blocks(parsed.content) + _blocks(parsed.output)) if parsed else ()
|
||||
nested: Final = _nested_blocks(top)
|
||||
|
|
@ -189,11 +204,11 @@ def _message_blocks(message: object) -> tuple[_AttachmentBlock, ...]:
|
|||
return top + nested + _nested_blocks(nested)
|
||||
|
||||
|
||||
def _nested_blocks(blocks: tuple[_AttachmentBlock, ...]) -> tuple[_AttachmentBlock, ...]:
|
||||
def _nested_blocks(blocks: tuple[_Block, ...]) -> tuple[_Block, ...]:
|
||||
return tuple(chain.from_iterable(_blocks(_nested_content(block)) for block in blocks))
|
||||
|
||||
|
||||
def _nested_content(block: _AttachmentBlock) -> object:
|
||||
def _nested_content(block: _Block) -> object:
|
||||
match block:
|
||||
case _ToolResultBlock():
|
||||
return block.content
|
||||
|
|
@ -203,13 +218,21 @@ def _nested_content(block: _AttachmentBlock) -> object:
|
|||
return None
|
||||
|
||||
|
||||
def _blocks(content: object) -> tuple[_AttachmentBlock, ...]:
|
||||
def _blocks(content: object) -> tuple[_Block, ...]:
|
||||
items: Final = _parse(_ITEMS_ADAPTER, content)
|
||||
parsed: Final = (_parse(_BLOCK_ADAPTER, block) for block in items or ())
|
||||
parsed: Final = (_block(item) for item in items or ())
|
||||
return tuple(block for block in parsed if block is not None)
|
||||
|
||||
|
||||
def _block_attachments(block: _AttachmentBlock, index: int) -> tuple[_Classified, ...]:
|
||||
def _block(item: object) -> _Block | None:
|
||||
parsed: Final = _parse(_BLOCK_ADAPTER, item)
|
||||
if parsed is not None:
|
||||
return parsed
|
||||
block_type: Final = (_parse(_OBJECT_MAPPING, item) or {}).get("type")
|
||||
return _MalformedBlock() if block_type in _ATTACHMENT_BLOCK_TYPES else None
|
||||
|
||||
|
||||
def _block_attachments(block: _Block, index: int) -> tuple[_Classified, ...]:
|
||||
"""A file block can name several sources and providers differ on which they send, so all are checked."""
|
||||
match block:
|
||||
case _FileBlock():
|
||||
|
|
@ -238,7 +261,7 @@ def _from_file_id(file_id: str, name: str | None, index: int, kind: AttachmentTy
|
|||
return _from_uri(file_id, name, index, kind) if is_url else _UNSENDABLE
|
||||
|
||||
|
||||
def _classify_block(block: _AttachmentBlock, index: int) -> _Classified:
|
||||
def _classify_block(block: _Block, index: int) -> _Classified:
|
||||
match block:
|
||||
case _ImageURLBlock():
|
||||
return _from_uri(_url(block.image_url), None, index, "image")
|
||||
|
|
|
|||
|
|
@ -11,7 +11,11 @@ from fastapi import HTTPException
|
|||
|
||||
from litellm.exceptions import GuardrailRaisedException, Timeout
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.proxy.guardrails.guardrail_hooks.akto.akto import UNMASKABLE_REASON, AktoGuardrail
|
||||
from litellm.proxy.guardrails.guardrail_hooks.akto.akto import (
|
||||
MALFORMED_ATTACHMENT_REASON,
|
||||
UNMASKABLE_REASON,
|
||||
AktoGuardrail,
|
||||
)
|
||||
from litellm.proxy.guardrails.guardrail_registry import (
|
||||
guardrail_class_registry,
|
||||
guardrail_initializer_registry,
|
||||
|
|
@ -1666,6 +1670,21 @@ def test_the_client_ip_is_the_first_forwarded_hop(akto_pre_call):
|
|||
assert payload["ip"] == "10.0.0.1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_malformed_attachment_blocks_the_request(akto_pre_call):
|
||||
akto_pre_call.async_handler.post = AsyncMock(return_value=_mock_allowed_response())
|
||||
request_data = {
|
||||
"messages": [{"role": "user", "content": [{"type": "text", "text": "hi"}, {"type": "file", "file": "x"}]}]
|
||||
}
|
||||
|
||||
with pytest.raises(GuardrailRaisedException) as exc_info:
|
||||
await akto_pre_call.apply_guardrail(
|
||||
inputs=GenericGuardrailAPIInputs(texts=["hi"]), request_data=request_data, input_type="request"
|
||||
)
|
||||
|
||||
assert exc_info.value.message == MALFORMED_ATTACHMENT_REASON
|
||||
|
||||
|
||||
def test_the_proxy_recorded_ip_wins_over_a_client_forwarding_header(akto_pre_call):
|
||||
request_data = {
|
||||
"metadata": {"requester_ip_address": "203.0.113.7"},
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import base64
|
||||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -189,9 +190,40 @@ def test_request_attachments_names_files_by_their_type():
|
|||
Attachment("plain.png", "image", url="https://example.com/plain.png"),
|
||||
),
|
||||
unsendable_count=2,
|
||||
malformed_count=1,
|
||||
), "names get an extension from the media type; raw and line-wrapped base64 are sent; invalid base64 is not"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"block",
|
||||
[
|
||||
{"type": "file", "file": {"file_data": f"data:application/pdf;base64,{PDF_B64}", "filename": {}}},
|
||||
{"type": "input_file", "file_data": f"data:application/pdf;base64,{PDF_B64}", "filename": ["x"]},
|
||||
{"type": "document", "title": 7, "source": {"type": "base64", "media_type": None, "data": PDF_B64}},
|
||||
],
|
||||
)
|
||||
def test_bad_optional_metadata_does_not_hide_an_attachment(block):
|
||||
found = request_attachments({"messages": [{"role": "user", "content": [block]}]})
|
||||
|
||||
assert [attachment.content for attachment in found.attachments] == [PDF_B64]
|
||||
assert found.malformed_count == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"block",
|
||||
[
|
||||
{"type": "file", "file": "not a file block"},
|
||||
{"type": "input_audio"},
|
||||
{"type": "image_url", "image_url": {"url": 123}},
|
||||
{"type": "tool_result", "content": [{"type": "document", "source": "nope"}]},
|
||||
],
|
||||
)
|
||||
def test_an_attachment_that_cannot_be_read_is_counted_malformed(block):
|
||||
found = request_attachments({"messages": [{"role": "user", "content": [block]}]})
|
||||
|
||||
assert (found.attachments, found.malformed_count) == ((), 1), "it can't be checked, so it must not be dropped"
|
||||
|
||||
|
||||
def test_a_malformed_attachment_url_is_named_by_position():
|
||||
request_data = {
|
||||
"messages": [{"role": "user", "content": [{"type": "image_url", "image_url": "https://[::1/x.png"}]}]
|
||||
|
|
@ -348,12 +380,13 @@ def test_a_document_keeps_its_title_and_context_in_the_text_check():
|
|||
{"type": "tool_result", "content": [{"type": "search_result", "title": "results", "content": "secret"}]},
|
||||
],
|
||||
)
|
||||
def test_search_results_are_sent_as_text_files(block):
|
||||
def test_search_results_are_sent_as_text_files_and_kept_out_of_the_text_check(block):
|
||||
request_data = {"messages": [{"role": "user", "content": [block]}]}
|
||||
|
||||
assert request_attachments(request_data).attachments == (
|
||||
Attachment("results.txt", "file", content=base64.b64encode(b"secret").decode()),
|
||||
)
|
||||
assert "secret" not in json.dumps(without_attachment_content(request_data["messages"])), "checked once, as a file"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("video_url", [{"url": f"data:video/mp4;base64,{PNG_B64}"}, f"data:video/mp4;base64,{PNG_B64}"])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue