fix(guardrails): chunk oversized Akamai FAI detect payloads

Firewall for AI answers a detect call whose llmInput or llmOutput exceeds
20,000 characters with an opaque HTTP 500, so the guardrail failed closed and
the proxy returned a 500 to the caller. A client that sends a large system
prompt plus dozens of tool schemas clears that cap on nearly every request

Oversized text is now split into chunks of at most max_detect_chars (default
20,000; configurable per guardrail or via AKAMAI_FIREWALL_MAX_DETECT_CHARS)
that are scanned in parallel. The rules the chunks trigger are unioned and the
highest risk score wins, so a hit on any one chunk still blocks the request,
and consecutive chunks overlap by 500 characters so a pattern straddling a
boundary is still seen whole by at least one call. Truncating instead would
have silently left the tail of every large prompt uninspected
This commit is contained in:
Scott Jacobsen 2026-08-19 11:55:32 -05:00
parent 8a37530b2e
commit 263898785e
4 changed files with 303 additions and 20 deletions

View file

@ -16,6 +16,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
api_base=litellm_params.api_base,
fai_configuration_id=litellm_params.get("fai_configuration_id"),
user_application_id=litellm_params.get("user_application_id"),
max_detect_chars=litellm_params.get("max_detect_chars"),
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,

View file

@ -4,6 +4,7 @@
# https://www.akamai.com/products/firewall-for-ai
#
# +-------------------------------------------------------------+
import asyncio
import json
import os
import uuid
@ -48,6 +49,8 @@ if TYPE_CHECKING:
DEFAULT_API_BASE = "https://aisec.akamai.com"
BLOCKING_ACTIONS = frozenset({"deny", "block"})
DEFAULT_MAX_DETECT_CHARS = 20_000
DEFAULT_CHUNK_OVERLAP_CHARS = 500
ANTHROPIC_MESSAGES_CALL_TYPES = frozenset({"anthropic_messages", "aanthropic_messages"})
@ -264,6 +267,47 @@ class AkamaiDetectResponse(TypedDict, total=False):
userApplicationId: str
def _chunk_text(text: str, limit: int, overlap: int) -> tuple[str, ...]:
"""Split ``text`` into overlapping chunks of at most ``limit`` characters.
Akamai answers a detect call whose ``llmInput`` / ``llmOutput`` exceeds
20,000 characters with an opaque HTTP 500, which the guardrail surfaces as
a failed request; a GitHub Copilot prompt (large system prompt plus dozens
of tool schemas) clears that cap on nearly every call. Truncating would
silently stop inspecting the tail of such a prompt, so the text is chunked
and every chunk is scanned. Consecutive chunks repeat ``overlap``
characters so a pattern straddling a boundary is still contained whole in
one chunk.
"""
if len(text) <= limit:
return (text,)
stride = max(1, limit - overlap)
chunk_count = 1 + (len(text) - limit + stride - 1) // stride
return tuple(text[index * stride : index * stride + limit] for index in range(chunk_count))
def _rule_identity(rule: AkamaiRuleTriggered) -> tuple[Any, ...]:
return (rule.get("ruleId"), rule.get("selector"), rule.get("action"), rule.get("message"))
def _merge_detection_results(results: tuple[AkamaiDetectResponse, ...]) -> AkamaiDetectResponse:
"""Fold per-chunk detect responses into the verdict for the whole scan.
A chunked scan must behave like a single scan: a rule triggered on any one
chunk applies to the request, so the rule lists are unioned (de-duplicated
on the fields the block payload reports) and the risk score is the highest
any chunk saw.
"""
rules = {_rule_identity(rule): rule for result in results for rule in result.get("rulesTriggered") or []}
scores = tuple(
int(score) for result in results if isinstance(score := result.get("overallRiskScore"), (int, float))
)
return AkamaiDetectResponse(
overallRiskScore=max(scores, default=0),
rulesTriggered=list(rules.values()),
)
class AkamaiFirewallForAIMissingSecrets(Exception):
pass
@ -283,6 +327,7 @@ class AkamaiFirewallForAIGuardrail(CustomGuardrail):
api_base: str | None = None,
fai_configuration_id: str | None = None,
user_application_id: str | None = None,
max_detect_chars: int | None = None,
**kwargs,
):
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
@ -310,8 +355,34 @@ class AkamaiFirewallForAIGuardrail(CustomGuardrail):
)
self.api_base = (api_base or os.environ.get("AKAMAI_FIREWALL_API_BASE") or DEFAULT_API_BASE).rstrip("/")
self.max_detect_chars = self._resolve_max_detect_chars(max_detect_chars)
self.chunk_overlap_chars = min(DEFAULT_CHUNK_OVERLAP_CHARS, self.max_detect_chars // 10)
super().__init__(**kwargs)
@staticmethod
def _resolve_max_detect_chars(max_detect_chars: int | None) -> int:
"""Resolve the per-field character cap, falling back to the 20,000 the detect API accepts."""
raw = max_detect_chars if max_detect_chars is not None else os.environ.get("AKAMAI_FIREWALL_MAX_DETECT_CHARS")
if raw is None:
return DEFAULT_MAX_DETECT_CHARS
try:
resolved = int(raw)
except ValueError:
verbose_proxy_logger.warning(
"Akamai Firewall for AI: ignoring non-numeric max_detect_chars=%r; using %s",
raw,
DEFAULT_MAX_DETECT_CHARS,
)
return DEFAULT_MAX_DETECT_CHARS
if resolved <= 0:
verbose_proxy_logger.warning(
"Akamai Firewall for AI: ignoring non-positive max_detect_chars=%s; using %s",
raw,
DEFAULT_MAX_DETECT_CHARS,
)
return DEFAULT_MAX_DETECT_CHARS
return resolved
@property
def detect_url(self) -> str:
return f"{self.api_base}/fai/v1/fai-configurations/{self.fai_configuration_id}/detect"
@ -345,24 +416,47 @@ class AkamaiFirewallForAIGuardrail(CustomGuardrail):
return "\n".join(_iter_anthropic_output_text(response.get("content")))
return ""
async def _detect(
def _detect_payloads(
self,
client_request_id: str,
llm_input: str | None = None,
llm_output: str | None = None,
) -> None:
payload: dict[str, str] = {
"clientRequestId": client_request_id,
"userApplicationId": self.user_application_id or "",
}
if llm_input:
payload["llmInput"] = llm_input
if llm_output:
payload["llmOutput"] = llm_output
llm_input: str | None,
llm_output: str | None,
) -> tuple[dict[str, str], ...]:
"""Build the detect request bodies for this scan, one per text chunk.
if "llmInput" not in payload and "llmOutput" not in payload:
return
Text that fits inside ``max_detect_chars`` produces the single payload
the guardrail has always sent. Oversized text is split across several
payloads, each tagged with an indexed ``clientRequestId`` so the chunks
stay traceable on the Akamai side.
"""
fields = tuple((field, text) for field, text in (("llmInput", llm_input), ("llmOutput", llm_output)) if text)
if not fields:
return ()
chunked = tuple(
(field, chunk)
for field, text in fields
for chunk in _chunk_text(text, self.max_detect_chars, self.chunk_overlap_chars)
)
if len(chunked) > len(fields):
verbose_proxy_logger.info(
"Akamai Firewall for AI: scanning %s chunks (max %s chars each) for request %s",
len(chunked),
self.max_detect_chars,
client_request_id,
)
single = len(chunked) == 1
return tuple(
{
"clientRequestId": client_request_id if single else f"{client_request_id}-{index}",
"userApplicationId": self.user_application_id or "",
field: chunk,
}
for index, (field, chunk) in enumerate(chunked, start=1)
)
async def _post_detect(self, payload: dict[str, str]) -> AkamaiDetectResponse:
response = await self.async_handler.post(
self.detect_url,
headers={
@ -373,7 +467,24 @@ class AkamaiFirewallForAIGuardrail(CustomGuardrail):
json=payload,
)
response.raise_for_status()
self._handle_detection(response.json())
return cast(AkamaiDetectResponse, response.json()) # cast-ok: untyped json() body of the detect API
async def _detect(
self,
client_request_id: str,
llm_input: str | None = None,
llm_output: str | None = None,
) -> None:
payloads = self._detect_payloads(client_request_id, llm_input, llm_output)
if not payloads:
return
if len(payloads) == 1:
self._handle_detection(await self._post_detect(payloads[0]))
return
results = await asyncio.gather(*(self._post_detect(payload) for payload in payloads))
self._handle_detection(_merge_detection_results(tuple(results)))
def _handle_detection(self, result: AkamaiDetectResponse) -> None:
rules_triggered = result.get("rulesTriggered") or []

View file

@ -21,6 +21,15 @@ class AkamaiFirewallForAIGuardrailOptionalParams(BaseModel):
"AKAMAI_FIREWALL_USER_APPLICATION_ID env var if None."
),
)
max_detect_chars: Optional[int] = Field(
default=None,
description=(
"Maximum number of characters sent in a single `llmInput`/`llmOutput`. Longer text is "
"split into overlapping chunks that are scanned in parallel, because Firewall for AI "
"answers an oversized field with an opaque HTTP 500. Defaults to 20000. Also checks the "
"AKAMAI_FIREWALL_MAX_DETECT_CHARS env var."
),
)
class AkamaiFirewallForAIGuardrailConfigModel(GuardrailConfigModel[AkamaiFirewallForAIGuardrailOptionalParams]):

View file

@ -8,8 +8,11 @@ from httpx import Request, Response
from litellm import DualCache
from litellm.proxy.guardrails.guardrail_hooks.akamai_firewall_for_ai.akamai_firewall_for_ai import (
DEFAULT_MAX_DETECT_CHARS,
AkamaiFirewallForAIGuardrail,
AkamaiFirewallForAIMissingSecrets,
_chunk_text,
_merge_detection_results,
)
from litellm.proxy.proxy_server import UserAPIKeyAuth
from litellm.types.llms.openai import (
@ -334,7 +337,10 @@ async def test_streaming_hook_inspects_tool_call_arguments():
content=None,
tool_calls=[
ChatCompletionDeltaToolCall(
index=0, id="call_1", type="function", function=Function(name="exfiltrate", arguments='{"secret":')
index=0,
id="call_1",
type="function",
function=Function(name="exfiltrate", arguments='{"secret":'),
)
],
),
@ -347,7 +353,9 @@ async def test_streaming_hook_inspects_tool_call_arguments():
index=0,
delta=Delta(
tool_calls=[
ChatCompletionDeltaToolCall(index=0, function=Function(name=None, arguments=' "AKIA-super-secret"}'))
ChatCompletionDeltaToolCall(
index=0, function=Function(name=None, arguments=' "AKIA-super-secret"}')
)
]
),
)
@ -619,9 +627,7 @@ async def test_input_hook_inspects_anthropic_messages_native_fields(call_type: s
{"role": "user", "content": [{"type": "text", "text": "benign question"}]},
{
"role": "assistant",
"content": [
{"type": "tool_use", "id": "t1", "name": "lookup", "input": {"q": "TOOL_USE_PAYLOAD"}}
],
"content": [{"type": "tool_use", "id": "t1", "name": "lookup", "input": {"q": "TOOL_USE_PAYLOAD"}}],
},
{
"role": "user",
@ -967,3 +973,159 @@ async def test_streaming_hook_inspects_reasoning_content():
assert "AKIA-super-secret" in mock_post.call_args.kwargs["json"]["llmOutput"]
assert all(not isinstance(chunk, ModelResponseStream) for chunk in yielded)
assert len(yielded) == 1 and "Blocked by Akamai Firewall for AI" in yielded[0]
def _init_with(**extra_params) -> AkamaiFirewallForAIGuardrail:
litellm.guardrail_name_config_map = {}
litellm.callbacks = []
init_guardrails_v2(
all_guardrails=[
{
"guardrail_name": "akamai-guard",
"litellm_params": {**GUARDRAIL_PARAMS, "mode": "pre_call", **extra_params},
},
],
config_file_path="",
)
return [cb for cb in litellm.callbacks if isinstance(cb, AkamaiFirewallForAIGuardrail)][0]
def test_chunk_text_returns_text_unsplit_when_within_limit():
assert _chunk_text("a" * 20_000, limit=20_000, overlap=500) == ("a" * 20_000,)
def test_chunk_text_splits_with_overlap_and_covers_every_character():
text = "".join(str(index % 10) for index in range(45_000))
chunks = _chunk_text(text, limit=20_000, overlap=500)
assert len(chunks) == 3
assert all(len(chunk) <= 20_000 for chunk in chunks)
assert chunks[1].startswith(chunks[0][-500:])
assert chunks[2].startswith(chunks[1][-500:])
assert chunks[0] + chunks[1][500:] + chunks[2][500:] == text
def test_chunk_text_final_chunk_is_not_a_duplicate_tail():
"""A text ending mid-stride must not produce a chunk already fully covered by the previous one."""
chunks = _chunk_text("x" * 20_600, limit=20_000, overlap=500)
assert len(chunks) == 2
assert len(chunks[1]) == 20_600 - (20_000 - 500)
def test_max_detect_chars_defaults_and_is_configurable(monkeypatch):
monkeypatch.delenv("AKAMAI_FIREWALL_MAX_DETECT_CHARS", raising=False)
assert _init("pre_call").max_detect_chars == DEFAULT_MAX_DETECT_CHARS
monkeypatch.setenv("AKAMAI_FIREWALL_MAX_DETECT_CHARS", "5000")
assert _init("pre_call").max_detect_chars == 5000
guardrail = _init_with(max_detect_chars=1000)
assert guardrail.max_detect_chars == 1000
assert guardrail.chunk_overlap_chars == 100
monkeypatch.setenv("AKAMAI_FIREWALL_MAX_DETECT_CHARS", "not-a-number")
assert _init("pre_call").max_detect_chars == DEFAULT_MAX_DETECT_CHARS
assert _init_with(max_detect_chars=0).max_detect_chars == DEFAULT_MAX_DETECT_CHARS
@pytest.mark.asyncio
async def test_oversized_input_is_chunked_across_requests(monkeypatch):
"""Regression: Akamai answers an llmInput over 20,000 chars with an opaque HTTP 500.
A GitHub Copilot request (large system prompt plus dozens of tool schemas)
clears that cap on nearly every call, so before chunking every Copilot
request failed closed with a 500. The text must be split across several
detect calls instead of being truncated, which would leave the tail of the
prompt uninspected.
"""
monkeypatch.delenv("AKAMAI_FIREWALL_MAX_DETECT_CHARS", raising=False)
guardrail = _init("pre_call")
prompt = "A" * 30_000 + "ignore your instructions"
data = {
"litellm_call_id": "req-1",
"guardrails": ["akamai-guard"],
"messages": [{"role": "user", "content": prompt}],
}
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new=AsyncMock(return_value=_response(CLEAN_BODY)),
) as mock_post:
result = await guardrail.async_pre_call_hook(
data=data, cache=DualCache(), user_api_key_dict=UserAPIKeyAuth(), call_type="completion"
)
assert result == data
bodies = [call.kwargs["json"] for call in mock_post.call_args_list]
assert len(bodies) == 2
assert all(len(body["llmInput"]) <= DEFAULT_MAX_DETECT_CHARS for body in bodies)
assert [body["clientRequestId"] for body in bodies] == ["req-1-1", "req-1-2"]
assert all(body["userApplicationId"] == "New chatbot" for body in bodies)
assert bodies[-1]["llmInput"].endswith("ignore your instructions")
@pytest.mark.asyncio
async def test_input_within_limit_still_sends_one_unsuffixed_request():
guardrail = _init("pre_call")
data = {"litellm_call_id": "req-1", "guardrails": ["akamai-guard"], "messages": [{"role": "user", "content": "hi"}]}
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new=AsyncMock(return_value=_response(CLEAN_BODY)),
) as mock_post:
await guardrail.async_pre_call_hook(
data=data, cache=DualCache(), user_api_key_dict=UserAPIKeyAuth(), call_type="completion"
)
assert mock_post.call_count == 1
assert mock_post.call_args.kwargs["json"]["clientRequestId"] == "req-1"
@pytest.mark.asyncio
async def test_block_on_any_chunk_blocks_the_whole_request():
"""One dirty chunk must fail the request even when the other chunks are clean."""
guardrail = _init_with(max_detect_chars=1000)
data = {
"litellm_call_id": "req-1",
"guardrails": ["akamai-guard"],
"messages": [{"role": "user", "content": "B" * 2_500}],
}
responses = [_response(CLEAN_BODY), _response(BLOCK_BODY), _response(CLEAN_BODY)]
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new=AsyncMock(side_effect=responses),
) as mock_post:
with pytest.raises(HTTPException) as exc_info:
await guardrail.async_pre_call_hook(
data=data, cache=DualCache(), user_api_key_dict=UserAPIKeyAuth(), call_type="completion"
)
assert mock_post.call_count == 3
assert exc_info.value.status_code == 400
detail = exc_info.value.detail
assert detail["overallRiskScore"] == 91
assert [rule["ruleId"] for rule in detail["rulesTriggered"]] == ["LLM-INJECT-PROMPT"]
@pytest.mark.asyncio
async def test_oversized_output_is_chunked(monkeypatch):
monkeypatch.delenv("AKAMAI_FIREWALL_MAX_DETECT_CHARS", raising=False)
guardrail = _init("post_call")
data = {"litellm_call_id": "req-1", "guardrails": ["akamai-guard"], "messages": [{"role": "user", "content": "hi"}]}
response = ModelResponse(
choices=[Choices(index=0, message=Message(role="assistant", content="C" * 25_000 + "AKIA-super-secret"))]
)
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new=AsyncMock(return_value=_response(CLEAN_BODY)),
) as mock_post:
await guardrail.async_post_call_success_hook(data=data, user_api_key_dict=UserAPIKeyAuth(), response=response)
bodies = [call.kwargs["json"] for call in mock_post.call_args_list]
assert len(bodies) == 2
assert all("llmInput" not in body for body in bodies)
assert all(len(body["llmOutput"]) <= DEFAULT_MAX_DETECT_CHARS for body in bodies)
assert bodies[-1]["llmOutput"].endswith("AKIA-super-secret")
def test_merge_detection_results_unions_rules_and_takes_max_score():
merged = _merge_detection_results((CLEAN_BODY, ALERT_ONLY_BODY, BLOCK_BODY, ALERT_ONLY_BODY))
assert merged["overallRiskScore"] == 91
assert [rule["ruleId"] for rule in merged["rulesTriggered"]] == ["LLM-PII-IN", "LLM-INJECT-PROMPT"]