mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(guardrails): ignore a client-sent response in Akto output checks
- a "response" field sent in the request body is no longer scanned in place of the model's reply - MCP tool calls in the reply are still checked in that case - split match or-patterns so CodeQL can follow the bound names - cover a masked payload that is not JSON and drop unused test imports
This commit is contained in:
parent
2573fa0107
commit
0c70b60f65
3 changed files with 71 additions and 11 deletions
|
|
@ -382,11 +382,17 @@ class AktoGuardrail(CustomGuardrail):
|
|||
}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def model_response(request_data: Mapping[str, object]) -> object:
|
||||
"""Translators keep a "response" already in the request, so one the client sent isn't the model's."""
|
||||
client_body: Final = as_mapping(as_mapping(request_data.get("proxy_server_request")).get("body"))
|
||||
return None if "response" in client_body else request_data.get("response")
|
||||
|
||||
@staticmethod
|
||||
def build_response_body(
|
||||
inputs: GenericGuardrailAPIInputs, request_data: Mapping[str, object]
|
||||
) -> Mapping[str, object]:
|
||||
model_response: Final = request_data.get("response")
|
||||
model_response: Final = AktoGuardrail.model_response(request_data)
|
||||
if isinstance(model_response, BaseModel):
|
||||
return model_response.model_dump()
|
||||
response_mapping: Final = as_mapping(model_response)
|
||||
|
|
@ -784,7 +790,13 @@ class AktoGuardrail(CustomGuardrail):
|
|||
# Only the complete response is under "response"; mid-stream checks get "responses"
|
||||
complete_response: Final = request_data.get("response")
|
||||
streamed: Final = bool(request_data.get("stream"))
|
||||
tool_calls: Final = self.response_mcp_tool_calls(complete_response) if complete_response is not None else ()
|
||||
model_response: Final = self.model_response(request_data)
|
||||
tool_call_source: Final = (
|
||||
model_response
|
||||
if model_response is not None
|
||||
else {"choices": [{"message": {"tool_calls": list(inputs.get("tool_calls") or ())}}]}
|
||||
)
|
||||
tool_calls: Final = self.response_mcp_tool_calls(tool_call_source) if complete_response is not None else ()
|
||||
return await self.settle(
|
||||
self.check_and_record(
|
||||
inputs,
|
||||
|
|
|
|||
|
|
@ -184,7 +184,9 @@ def _nested_blocks(blocks: tuple[_AttachmentBlock, ...]) -> tuple[_AttachmentBlo
|
|||
|
||||
def _nested_content(block: _AttachmentBlock) -> object:
|
||||
match block:
|
||||
case _ToolResultBlock(content=content) | _DocumentBlock(source=_Source(type="content", content=content)):
|
||||
case _ToolResultBlock(content=content):
|
||||
return content
|
||||
case _DocumentBlock(source=_Source(type="content", content=content)):
|
||||
return content
|
||||
case _:
|
||||
return None
|
||||
|
|
@ -198,11 +200,15 @@ def _blocks(content: object) -> tuple[_AttachmentBlock, ...]:
|
|||
|
||||
def _classify_block(block: _AttachmentBlock, index: int) -> _Classified:
|
||||
match block:
|
||||
case _ImageURLBlock(image_url=_ImageURL(url=url)) | _InputImageBlock(image_url=url):
|
||||
case _ImageURLBlock(image_url=_ImageURL(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)):
|
||||
case _InputImageBlock(image_url=url):
|
||||
return _from_uri(url, None, index, "image")
|
||||
case _VideoURLBlock(video_url=_ImageURL(url=url)):
|
||||
return _from_uri(url, None, index, "file")
|
||||
case _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
|
||||
|
|
|
|||
|
|
@ -12,12 +12,6 @@ from fastapi import HTTPException
|
|||
from litellm.exceptions import GuardrailRaisedException, Timeout
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.proxy.guardrails.guardrail_hooks.akto.akto import UNMASKABLE_REASON, AktoGuardrail
|
||||
from litellm.proxy.guardrails.guardrail_hooks.akto.akto_attachments import (
|
||||
Attachment,
|
||||
RequestAttachments,
|
||||
request_attachments,
|
||||
without_attachment_content,
|
||||
)
|
||||
from litellm.proxy.guardrails.guardrail_registry import (
|
||||
guardrail_class_registry,
|
||||
guardrail_initializer_registry,
|
||||
|
|
@ -733,6 +727,21 @@ async def test_pre_call_forwards_akto_masked_prompt(akto_pre_call):
|
|||
assert result["texts"] == ["be brief", "card XXXX"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_blocks_a_masked_payload_that_is_not_json(akto_pre_call):
|
||||
result = {"Allowed": True, "Modified": True, "ModifiedPayload": "card XXXX", "behaviour": "alert"}
|
||||
akto_pre_call.async_handler.post = AsyncMock(return_value=_response({"data": {"guardrailsResult": result}}))
|
||||
|
||||
with pytest.raises(GuardrailRaisedException) as exc_info:
|
||||
await akto_pre_call.apply_guardrail(
|
||||
inputs=GenericGuardrailAPIInputs(texts=[f"card {CARD}"]),
|
||||
request_data={"messages": [{"role": "user", "content": f"card {CARD}"}]},
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert exc_info.value.message == UNMASKABLE_REASON
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_blocks_masking_that_also_hits_a_tool_description(akto_pre_call):
|
||||
akto_pre_call.async_handler.post = _masking_akto("requestPayload", CARD)
|
||||
|
|
@ -1524,6 +1533,39 @@ async def test_a_dict_model_response_is_recorded_as_sent(akto_post_call, sample_
|
|||
assert json.loads(json.loads(payload["responsePayload"])["body"]) == tool_use_only
|
||||
|
||||
|
||||
def _with_client_response(request_data):
|
||||
fake = {"choices": [{"message": {"role": "assistant", "content": "ok"}}]}
|
||||
return {**request_data, "response": fake, "proxy_server_request": {"body": {"response": fake}}}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_response_sent_by_the_client_is_not_scanned_in_place_of_the_reply(akto_post_call, sample_request_data):
|
||||
akto_post_call.async_handler.post = AsyncMock(return_value=_mock_blocked_response("PII Policy violated"))
|
||||
|
||||
with pytest.raises(GuardrailRaisedException):
|
||||
await akto_post_call.apply_guardrail(
|
||||
inputs=GenericGuardrailAPIInputs(texts=[f"card {CARD}"]),
|
||||
request_data=_with_client_response(sample_request_data),
|
||||
input_type="response",
|
||||
)
|
||||
|
||||
[(_, payload)] = _calls(akto_post_call)
|
||||
assert f"card {CARD}" in payload["responsePayload"], "the model's reply is scanned, not the client's"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_tool_calls_are_checked_when_the_client_sends_a_response(akto_post_call, sample_request_data):
|
||||
akto_post_call.async_handler.post = AsyncMock(return_value=_mock_allowed_response())
|
||||
|
||||
await akto_post_call.apply_guardrail(
|
||||
inputs=GenericGuardrailAPIInputs(texts=[], tool_calls=[MCP_TOOL_CALL]),
|
||||
request_data=_with_client_response(sample_request_data),
|
||||
input_type="response",
|
||||
)
|
||||
|
||||
assert "/mcp" in [payload["path"] for _, payload in _calls(akto_post_call)]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_one_text_masked_two_ways_blocks(akto_pre_call):
|
||||
def respond(**kwargs):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue