fix(guardrails): inspect tool-call output and enforce streaming for Akamai FAI

Post-call output extraction only read message.content, so a
tool-call-only response (content is None with populated tool-call
arguments) produced empty output text, short-circuited the detect
request and was released to the caller uninspected. Reuse the shared
get_content_from_model_response helper so tool-call names and arguments
and the legacy function_call are inspected too.

Streaming responses were only handled by async_post_call_success_hook,
which the proxy runs through the audit-only deferred stream path after
the content has already been delivered, so a block was logged rather
than enforced. Add async_post_call_streaming_iterator_hook to buffer,
assemble via stream_chunk_builder and inspect the completion before any
bytes reach the client, emitting an SSE error event when Akamai blocks.
This commit is contained in:
Scott Jacobsen 2026-07-27 16:07:08 -05:00
parent 115d4acdf0
commit e27b8db1b9
2 changed files with 223 additions and 8 deletions

View file

@ -4,11 +4,13 @@
# https://www.akamai.com/products/firewall-for-ai
#
# +-------------------------------------------------------------+
import json
import os
import uuid
from typing import (
TYPE_CHECKING,
Any,
AsyncGenerator,
TypedDict,
)
@ -29,7 +31,6 @@ from litellm.proxy.guardrails._content_utils import iter_message_text
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import (
CallTypesLiteral,
Choices,
EmbeddingResponse,
ImageResponse,
ModelResponse,
@ -120,14 +121,13 @@ class AkamaiFirewallForAIGuardrail(CustomGuardrail):
@staticmethod
def _output_text(response: ModelResponse | Any) -> str:
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_content_from_model_response,
)
if not isinstance(response, ModelResponse):
return ""
fragments = [
choice.message.content
for choice in response.choices
if isinstance(choice, Choices) and isinstance(choice.message.content, str) and choice.message.content
]
return "\n".join(fragments)
return get_content_from_model_response(response)
async def _detect(
self,
@ -234,6 +234,53 @@ class AkamaiFirewallForAIGuardrail(CustomGuardrail):
await self._detect(client_request_id=self._client_request_id(data), llm_output=self._output_text(response))
return response
async def async_post_call_streaming_iterator_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
response: Any,
request_data: dict,
) -> AsyncGenerator[Any, None]:
if self.should_run_guardrail(data=request_data, event_type=GuardrailEventHooks.post_call) is not True:
async for chunk in response:
yield chunk
return
from litellm.main import stream_chunk_builder
chunks = [chunk async for chunk in response]
if not chunks:
return
assembled = stream_chunk_builder(chunks=chunks)
if not isinstance(assembled, ModelResponse):
for chunk in chunks:
yield chunk
return
try:
await self._detect(
client_request_id=self._client_request_id(request_data),
llm_output=self._output_text(assembled),
)
except HTTPException as exc:
error_obj = dict(exc.detail) if isinstance(exc.detail, dict) else {"message": str(exc.detail)}
error_obj["code"] = exc.status_code
yield f"data: {json.dumps({'error': error_obj})}\n\n"
return
except Exception as exc:
verbose_proxy_logger.exception("Akamai Firewall for AI: streaming output scan failed: %s", exc)
error_obj = {
"message": "Akamai Firewall for AI scan failed; response withheld",
"type": "guardrail_scan_error",
"code": 500,
"guardrail": self.guardrail_name,
}
yield f"data: {json.dumps({'error': error_obj})}\n\n"
return
for chunk in chunks:
yield chunk
@staticmethod
def get_config_model() -> type["GuardrailConfigModel"] | None:
from litellm.types.proxy.guardrails.guardrail_hooks.akamai_firewall_for_ai import (

View file

@ -12,7 +12,17 @@ from litellm.proxy.guardrails.guardrail_hooks.akamai_firewall_for_ai.akamai_fire
AkamaiFirewallForAIMissingSecrets,
)
from litellm.proxy.proxy_server import UserAPIKeyAuth
from litellm.types.utils import Choices, Message, ModelResponse
from litellm.types.utils import (
ChatCompletionDeltaToolCall,
ChatCompletionMessageToolCall,
Choices,
Delta,
Function,
Message,
ModelResponse,
ModelResponseStream,
StreamingChoices,
)
sys.path.insert(0, os.path.abspath("../.."))
import litellm
@ -213,3 +223,161 @@ async def test_no_api_call_when_no_text():
)
assert result == data
mock_post.assert_not_called()
def _tool_call_response() -> ModelResponse:
"""A completion whose only output lives in tool-call arguments (content is None)."""
return ModelResponse(
choices=[
Choices(
index=0,
message=Message(
role="assistant",
content=None,
tool_calls=[
ChatCompletionMessageToolCall(
id="call_1",
type="function",
function=Function(name="exfiltrate", arguments='{"secret": "AKIA-super-secret"}'),
)
],
),
)
]
)
async def _aiter(chunks):
for chunk in chunks:
yield chunk
@pytest.mark.asyncio
async def test_output_hook_inspects_tool_call_arguments():
"""Regression: tool-call arguments (content=None) must be sent to Akamai and blocked.
Before the fix ``_output_text`` only read ``message.content``, so a
tool-call-only response produced empty output text, ``_detect`` short
circuited, no request was made and the payload was released uninspected.
"""
guardrail = _init("post_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(BLOCK_BODY)),
) as mock_post:
with pytest.raises(HTTPException):
await guardrail.async_post_call_success_hook(
data=data, user_api_key_dict=UserAPIKeyAuth(), response=_tool_call_response()
)
body = mock_post.call_args.kwargs["json"]
assert "AKIA-super-secret" in body["llmOutput"]
assert "exfiltrate" in body["llmOutput"]
@pytest.mark.asyncio
async def test_streaming_hook_blocks_before_delivery():
"""Regression: a blocking verdict on a streamed response must withhold the content.
Guardrails that only override ``async_post_call_success_hook`` are run by
the deferred stream path after the bytes are already delivered, so the
block is not enforced. The streaming iterator hook must buffer, inspect
and emit an SSE error instead of the original chunks.
"""
guardrail = _init("post_call")
request_data = {"litellm_call_id": "req-1", "guardrails": ["akamai-guard"]}
chunks = [
ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(role="assistant", content="here is a "))]),
ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(content="secret"))]),
]
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new=AsyncMock(return_value=_response(BLOCK_BODY)),
) as mock_post:
yielded = [
chunk
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(), response=_aiter(chunks), request_data=request_data
)
]
assert mock_post.call_args.kwargs["json"]["llmOutput"] == "here is a secret"
# none of the original model chunks are delivered
assert all(not isinstance(chunk, ModelResponseStream) for chunk in yielded)
# a single SSE error event carrying the Akamai block is emitted instead
assert len(yielded) == 1 and isinstance(yielded[0], str)
assert "Blocked by Akamai Firewall for AI" in yielded[0]
@pytest.mark.asyncio
async def test_streaming_hook_inspects_tool_call_arguments():
"""Tool-call arguments streamed as deltas must be assembled, inspected and blocked."""
guardrail = _init("post_call")
request_data = {"litellm_call_id": "req-1", "guardrails": ["akamai-guard"]}
chunks = [
ModelResponseStream(
choices=[
StreamingChoices(
index=0,
delta=Delta(
role="assistant",
content=None,
tool_calls=[
ChatCompletionDeltaToolCall(
index=0, id="call_1", type="function", function=Function(name="exfiltrate", arguments='{"secret":')
)
],
),
)
]
),
ModelResponseStream(
choices=[
StreamingChoices(
index=0,
delta=Delta(
tool_calls=[
ChatCompletionDeltaToolCall(index=0, function=Function(name=None, arguments=' "AKIA-super-secret"}'))
]
),
)
]
),
]
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new=AsyncMock(return_value=_response(BLOCK_BODY)),
) as mock_post:
yielded = [
chunk
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(), response=_aiter(chunks), request_data=request_data
)
]
assert "AKIA-super-secret" in mock_post.call_args.kwargs["json"]["llmOutput"]
assert len(yielded) == 1 and "Blocked by Akamai Firewall for AI" in yielded[0]
@pytest.mark.asyncio
async def test_streaming_hook_passes_through_when_clean():
"""A clean verdict yields the original chunks unchanged after inspection."""
guardrail = _init("post_call")
request_data = {"litellm_call_id": "req-1", "guardrails": ["akamai-guard"]}
chunks = [
ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(role="assistant", content="all "))]),
ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(content="clear"))]),
]
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new=AsyncMock(return_value=_response(CLEAN_BODY)),
) as mock_post:
yielded = [
chunk
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(), response=_aiter(chunks), request_data=request_data
)
]
assert mock_post.call_args.kwargs["json"]["llmOutput"] == "all clear"
assert yielded == chunks