diff --git a/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py b/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py index 8d336c9d4f8..4e3738a6dab 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py +++ b/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/akto/akto_attachments.py b/litellm/proxy/guardrails/guardrail_hooks/akto/akto_attachments.py index b8b285dd07a..98a008e86fe 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/akto/akto_attachments.py +++ b/litellm/proxy/guardrails/guardrail_hooks/akto/akto_attachments.py @@ -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 diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto.py b/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto.py index a0875a6f5ef..58a84513dcc 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/akto/test_akto.py @@ -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):