mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
115d4acdf0
commit
e27b8db1b9
2 changed files with 223 additions and 8 deletions
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue