From b9ff03a8c6f7fec349887fd303cba93d738f1166 Mon Sep 17 00:00:00 2001 From: Rohan Date: Sat, 3 Oct 2026 20:17:10 +0530 Subject: [PATCH] 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 --- .../guardrails/guardrail_hooks/akto/akto.py | 3 + .../guardrail_hooks/akto/akto_attachments.py | 57 +++++++++++++------ .../guardrail_hooks/akto/test_akto.py | 21 ++++++- .../akto/test_akto_attachments.py | 35 +++++++++++- 4 files changed, 97 insertions(+), 19 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py b/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py index bbdea4cff14..aeabfe239c0 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py +++ b/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/akto/akto_attachments.py b/litellm/proxy/guardrails/guardrail_hooks/akto/akto_attachments.py index 801b0db4db7..caee65d312d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/akto/akto_attachments.py +++ b/litellm/proxy/guardrails/guardrail_hooks/akto/akto_attachments.py @@ -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") diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto.py b/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto.py index ef1ad7d3b3a..db15aae0fa2 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto.py @@ -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"}, diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto_attachments.py b/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto_attachments.py index a0ad0c4be95..428b0b3c245 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto_attachments.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto_attachments.py @@ -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}"])