mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
parent
a76ba8c01e
commit
abadad3020
5 changed files with 2585 additions and 582 deletions
|
|
@ -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
334
litellm/proxy/guardrails/guardrail_hooks/akto/attachments.py
Normal file
334
litellm/proxy/guardrails/guardrail_hooks/akto/attachments.py
Normal 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
|
||||
|
|
@ -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
Loading…
Add table
Reference in a new issue