feat(guardrails): extend Akto guardrail to responses, MCP tools, attachments and masking

- post_call now waits for Akto and blocks or masks the reply instead of only logging it
- pre_mcp_call and post_mcp_call check MCP tool arguments and tool results
- attached images, audio and files are sent to Akto's file check
- streamed replies are checked every streaming_sampling_rate chunks
- context_source routes traffic to Akto's endpoint or agentic policies
- tags carry user email, team alias and key alias for attribution
This commit is contained in:
Rohan 2026-10-03 11:58:10 +05:30
parent a76ba8c01e
commit abadad3020
5 changed files with 2585 additions and 582 deletions

View file

@ -2,7 +2,7 @@ from typing import TYPE_CHECKING, Final
from litellm.types.guardrails import SupportedGuardrailIntegrations
from .akto import AktoGuardrail
from .akto import AktoGuardrail, streaming_sampling_rate_from
if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
@ -12,12 +12,16 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
import litellm
_akto_callback: Final = AktoGuardrail(
akto_base_url=getattr(litellm_params, "akto_base_url", None),
akto_api_key=getattr(litellm_params, "akto_api_key", None),
akto_account_id=getattr(litellm_params, "akto_account_id", None),
akto_vxlan_id=getattr(litellm_params, "akto_vxlan_id", None),
unreachable_fallback=getattr(litellm_params, "unreachable_fallback", "fail_closed"),
guardrail_timeout=getattr(litellm_params, "guardrail_timeout", None),
akto_base_url=litellm_params.akto_base_url,
akto_api_key=litellm_params.akto_api_key,
akto_account_id=litellm_params.akto_account_id,
akto_vxlan_id=litellm_params.akto_vxlan_id,
context_source=litellm_params.context_source,
akto_metadata=litellm_params.akto_metadata,
streaming_sampling_rate=streaming_sampling_rate_from(litellm_params),
guardrail_timeout=litellm_params.guardrail_timeout,
file_guardrail_timeout=litellm_params.file_guardrail_timeout,
unreachable_fallback=litellm_params.unreachable_fallback,
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,334 @@
"""Attachment blocks sent to Akto's file guardrail, including those inside ``tool_result`` blocks:
OpenAI chat ``image_url``, ``input_audio``, ``file``, ``video_url``
Anthropic ``image``, ``document``
Responses API ``input_image``, ``input_file``
A block with neither inline bytes nor a URL (an OpenAI ``file_id``) is unsendable.
"""
import base64
import binascii
import mimetypes
import posixpath
from collections.abc import Mapping
from dataclasses import dataclass
from itertools import chain
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
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")
)
_OBJECT_MAPPING: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object])
_T: Final = TypeVar("_T")
@dataclass(frozen=True, slots=True)
class Attachment:
filename: str
type: AttachmentType
content: str | None = None
url: str | None = None
def as_payload(self) -> Mapping[str, str]:
fields: Final = (("filename", self.filename), ("type", self.type), ("content", self.content), ("url", self.url))
return MappingProxyType({key: value for key, value in fields if value is not None})
@dataclass(frozen=True, slots=True)
class RequestAttachments:
attachments: tuple[Attachment, ...]
unsendable_count: int
class _Model(BaseModel):
model_config = ConfigDict(extra="ignore")
class _ImageURL(_Model):
url: str | None = None
class _ImageURLBlock(_Model):
type: Literal["image_url"]
image_url: _ImageURL | str
class _VideoURLBlock(_Model):
type: Literal["video_url"]
video_url: _ImageURL | str
class _InputImageBlock(_Model):
type: Literal["input_image"]
image_url: str | None = None
class _InputAudio(_Model):
data: str | None = None
format: str | None = None
class _InputAudioBlock(_Model):
type: Literal["input_audio"]
input_audio: _InputAudio
class _FileData(_Model):
file_data: str | None = None
filename: str | None = None
class _FileBlock(_Model):
type: Literal["file"]
file: _FileData
class _InputFileBlock(_Model):
type: Literal["input_file"]
file_data: str | None = None
file_url: str | None = None
filename: str | None = None
class _Source(_Model):
type: str | None = None
data: str | None = None
media_type: str | None = None
url: str | None = None
content: object = None
class _TextBlock(_Model):
type: Literal["text"]
text: str
class _ImageBlock(_Model):
type: Literal["image"]
source: _Source
class _DocumentBlock(_Model):
type: Literal["document"]
source: _Source
title: str | None = None
class _ToolResultBlock(_Model):
type: Literal["tool_result"]
content: object = None
class _Message(_Model):
content: object = None
output: object = None
_AttachmentBlock: TypeAlias = (
_ImageURLBlock
| _VideoURLBlock
| _InputImageBlock
| _InputAudioBlock
| _FileBlock
| _InputFileBlock
| _ImageBlock
| _DocumentBlock
| _ToolResultBlock
)
_BLOCK_ADAPTER: Final[TypeAdapter[_AttachmentBlock]] = TypeAdapter(
Annotated[_AttachmentBlock, Field(discriminator="type")]
)
_TEXT_BLOCK_ADAPTER: Final[TypeAdapter[_TextBlock]] = TypeAdapter(_TextBlock)
_MESSAGE_ADAPTER: Final[TypeAdapter[_Message]] = TypeAdapter(_Message)
_ITEMS_ADAPTER: Final[TypeAdapter[list[object]]] = TypeAdapter(list[object])
# (attachment, is_unsendable); (None, False) is a block that isn't an attachment
_Classified: TypeAlias = tuple[Attachment | None, bool]
_NOT_AN_ATTACHMENT: Final[_Classified] = (None, False)
_UNSENDABLE: Final[_Classified] = (None, True)
def request_attachments(request_data: Mapping[str, object]) -> RequestAttachments:
messages: Final = _parse(_ITEMS_ADAPTER, request_data.get("messages")) or _parse(
_ITEMS_ADAPTER, request_data.get("input")
)
blocks: Final = chain.from_iterable(_message_blocks(message) for message in messages or ())
classified: Final = tuple(_classify_block(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),
)
def _message_blocks(message: object) -> tuple[_AttachmentBlock, ...]:
parsed: Final = _parse(_MESSAGE_ADAPTER, message)
top: Final = (_blocks(parsed.content) + _blocks(parsed.output)) if parsed else ()
nested: Final = _nested_blocks(top)
# tool_result -> document -> image is the deepest the APIs nest
return top + nested + _nested_blocks(nested)
def _nested_blocks(blocks: tuple[_AttachmentBlock, ...]) -> tuple[_AttachmentBlock, ...]:
return tuple(chain.from_iterable(_blocks(_nested_content(block)) for block in blocks))
def _nested_content(block: _AttachmentBlock) -> object:
match block:
case _ToolResultBlock(content=content) | _DocumentBlock(source=_Source(type="content", content=content)):
return content
case _:
return None
def _blocks(content: object) -> tuple[_AttachmentBlock, ...]:
items: Final = _parse(_ITEMS_ADAPTER, content)
parsed: Final = (_parse(_BLOCK_ADAPTER, block) for block in items or ())
return tuple(block for block in parsed if block is not None)
def _classify_block(block: _AttachmentBlock, index: int) -> _Classified:
match block:
case _ImageURLBlock(image_url=_ImageURL(url=url)) | _InputImageBlock(image_url=url):
return _from_uri(url, None, index, "image")
case _ImageURLBlock(image_url=str(url)):
return _from_uri(url, None, index, "image")
case _VideoURLBlock(video_url=_ImageURL(url=url)) | _VideoURLBlock(video_url=str(url)):
return _from_uri(url, None, index, "file")
case _InputAudioBlock(input_audio=_InputAudio(data=str(data), format=audio_format)):
name: Final = f"attachment-{index}.{audio_format}" if audio_format else None
return _from_base64(data, name, index, "audio", None)
case _InputAudioBlock():
return _UNSENDABLE
case _FileBlock(file=file):
return _from_uri(file.file_data, file.filename, index, "file")
case _InputFileBlock():
return _from_uri(block.file_data or block.file_url, block.filename, index, "file")
case _ImageBlock(source=source):
return _from_source(source, None, index, "image")
case _DocumentBlock(source=source, title=title):
return _from_source(source, title, index, "file")
case _:
return _NOT_AN_ATTACHMENT
def _from_uri(raw_uri: str | None, name: str | None, index: int, kind: AttachmentType) -> _Classified:
uri: Final = (raw_uri or "").strip()
if not uri:
return _UNSENDABLE
if uri.lower().startswith(_REMOTE_URI_SCHEMES):
return Attachment(_filename(name, index, url=uri), kind, url=uri), False
media_type, data = _parse_data_uri(uri)
return _from_base64(data, name, index, kind, media_type)
def _from_source(source: _Source, name: str | None, index: int, kind: AttachmentType) -> _Classified:
"""base64, plain text, text blocks or a URL; a file_id has nothing to send."""
match source:
case _Source(type="base64", data=str(data)):
return _from_base64(data, name, index, kind, source.media_type)
case _Source(type="text", data=str(data)):
content: Final = base64.b64encode(data.encode(errors="surrogatepass")).decode()
return Attachment(_filename(name, index, source.media_type or "text/plain"), kind, content=content), False
case _Source(type="content", content=text_blocks) if text := _joined_text(text_blocks):
encoded: Final = base64.b64encode(text.encode(errors="surrogatepass")).decode()
return Attachment(_filename(name, index, "text/plain"), kind, content=encoded), False
case _Source(type="url", url=str(url)) if url:
return Attachment(_filename(name, index, url=url), kind, url=url), False
case _:
return _UNSENDABLE
def _joined_text(content: object) -> str:
if isinstance(content, str):
return content
blocks: Final = (_parse(_TEXT_BLOCK_ADAPTER, item) for item in _parse(_ITEMS_ADAPTER, content) or ())
return "\n".join(block.text for block in blocks if block is not None)
def _from_base64(data: str, name: str | None, index: int, kind: AttachmentType, media_type: str | None) -> _Classified:
content: Final = _standard_base64(data)
if content is None:
return _UNSENDABLE
return Attachment(_filename(name, index, media_type), kind, content=content), False
def _standard_base64(data: str) -> str | None:
"""Padded standard base64, accepting line breaks, missing padding and URL-safe characters."""
compact: Final = "".join(data.split()).translate(_URL_SAFE_TO_STANDARD)
padded: Final = compact + "=" * (-len(compact) % 4)
return padded if compact and _is_base64(padded) else None
def _parse_data_uri(uri: str) -> tuple[str | None, str]:
"""(media type, base64 data); a plain data URI's text is encoded, anything else is taken as raw base64."""
if uri[:5].lower() != "data:" or "," not in uri:
return None, uri
header, data = uri[5:].split(",", 1)
params: Final = header.split(";")
encoded: Final = params[-1].strip().lower() == "base64"
return params[0], data if encoded else base64.b64encode(
unquote_to_bytes(data.encode(errors="surrogatepass"))
).decode()
def _filename(name: str | None, index: int, media_type: str | None = None, url: str | None = None) -> str:
"""The client's name, else the URL's, with an extension from the media type when it has none."""
stem: Final = posixpath.basename((name or "").strip()) or _url_basename(url) or f"attachment-{index}"
extension: Final = mimetypes.guess_extension(media_type.split(";")[0].strip()) if media_type else None
return stem if posixpath.splitext(stem)[1] or not extension else f"{stem}{extension}"
def _url_basename(url: str | None) -> str:
try:
return posixpath.basename(unquote(urlparse(url or "").path))
except ValueError:
return ""
def _is_base64(data: str) -> bool:
try:
base64.b64decode(data, validate=True)
except (binascii.Error, ValueError):
return False
return True
def without_attachment_content(messages: object) -> object:
items: Final = _parse(_ITEMS_ADAPTER, messages)
return (
messages
if items is None
else tuple(_without_content(_without_content(message, "content"), "output") for message in items)
)
def _without_content(value: object, key: str) -> object:
mapping: Final = _parse(_OBJECT_MAPPING, value)
blocks: Final = _parse(_ITEMS_ADAPTER, mapping.get(key)) if mapping else None
if mapping is None or blocks is None:
return value
return {**mapping, key: tuple(_block_without_content(block) for block in blocks)}
def _block_without_content(block: object) -> object:
block_type: Final = (_parse(_OBJECT_MAPPING, block) or {}).get("type")
if block_type in _ATTACHMENT_BLOCK_TYPES:
return {"type": block_type}
return _without_content(block, "content") if block_type == "tool_result" else block
def _parse(adapter: TypeAdapter[_T], value: object) -> _T | None:
try:
return adapter.validate_python(value)
except ValidationError:
return None

View file

@ -1,17 +1,28 @@
from typing import Literal
from pydantic import Field
from pydantic import BaseModel, Field
from .base import GuardrailConfigModel
class AktoConfigModel(GuardrailConfigModel):
"""
Config for the Akto guardrail.
class AktoGuardrailConfigModelOptionalParams(BaseModel):
streaming_sampling_rate: int | None = Field(
default=None,
ge=1,
description=(
"Check the streamed response every Nth chunk; the stream pauses at that chunk until Akto replies. "
"1 checks every chunk. Default: 5."
),
)
Use two separate config entries to control behaviour:
akto-validate (mode: pre_call) -> check guardrails, block if flagged
akto-ingest (mode: post_call) -> ingest request+response data
class AktoConfigModel(GuardrailConfigModel[AktoGuardrailConfigModelOptionalParams]):
"""
Config for the Akto guardrail. Each mode checks the traffic with Akto, then blocks or masks it:
pre_call -> LLM request
post_call -> LLM response
pre_mcp_call -> MCP tool call
post_mcp_call -> MCP tool result
"""
akto_base_url: str | None = Field(
@ -40,16 +51,39 @@ class AktoConfigModel(GuardrailConfigModel):
description="Akto VXLAN ID. Env: AKTO_VXLAN_ID. Default: '0'.",
)
unreachable_fallback: Literal["fail_closed", "fail_open"] = Field(
default="fail_closed",
description="What to do when Akto is unreachable. 'fail_open' = allow, 'fail_closed' = block.",
context_source: Literal["ENDPOINT", "AGENTIC"] | None = Field(
default=None,
description="Akto context the traffic belongs to: 'ENDPOINT' (Atlas) or 'AGENTIC' (Argus). Default: ENDPOINT.",
)
akto_metadata: dict | None = Field( # mutable-ok: UI type derivation maps dict to "object"
default=None,
description=(
"JSON object sent to Akto. 'policy_name': comma-separated Akto policies to enforce (empty enforces all). "
'Example: {"policy_name": "PII Strict, Secrets"}.'
),
)
guardrail_timeout: int | None = Field(
default=None,
ge=1,
description="HTTP timeout in seconds. Default: 5.",
)
file_guardrail_timeout: int | None = Field(
default=None,
ge=1,
description="HTTP timeout in seconds for checking attached files. Default: 10.",
)
unreachable_fallback: Literal["fail_closed", "fail_open"] = Field(
default="fail_closed",
description=(
"What to do when Akto is unreachable, times out or errors. 'fail_closed' = block (default), "
"'fail_open' = allow."
),
)
@staticmethod
def ui_friendly_name() -> str:
return "Akto"

File diff suppressed because it is too large Load diff