add ingestion for blocked responses in AktoGuardrail

This commit is contained in:
rzeta-10 2026-03-23 23:54:03 +05:30
parent a0a6cd3940
commit 61fa049f7b
2 changed files with 94 additions and 18 deletions

View file

@ -4,6 +4,7 @@ Mode:
- pre_call: Validates request against Akto guardrails, blocks if flagged.
"""
import asyncio
import json
import os
from datetime import datetime
@ -64,6 +65,7 @@ class AktoGuardrail(CustomGuardrail):
self.async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback,
)
self.background_tasks: set = set()
self.akto_base_url = (
akto_base_url or os.environ.get("AKTO_GUARDRAIL_API_BASE", "")
@ -91,6 +93,30 @@ class AktoGuardrail(CustomGuardrail):
]
super().__init__(**kwargs)
# ── Helpers ──
def schedule(self, coro) -> None:
"""Schedule a fire-and-forget background task"""
task = asyncio.create_task(coro)
self.background_tasks.add(task)
task.add_done_callback(self.background_tasks.discard)
async def ingest_blocked_request(self, payload: dict) -> None:
"""Fire-and-forget: ingest a blocked request to Akto for audit."""
try:
await self.async_handler.post(
url=f"{self.akto_base_url}{HTTP_PROXY_PATH}",
data=json.dumps(payload),
params={"akto_connector": AKTO_CONNECTOR_NAME, "ingest_data": "true"},
headers={
"content-type": "application/json",
"Authorization": self.akto_api_key,
},
timeout=self.guardrail_timeout,
)
except Exception as e:
verbose_proxy_logger.error("Akto blocked-request ingest error: %s", e)
# ── Payload builders ──
@staticmethod
@ -338,6 +364,16 @@ class AktoGuardrail(CustomGuardrail):
)
if not allowed:
blocked_payload = self.build_akto_payload(
inputs, request_data, status_code=403
)
blocked_payload["responsePayload"] = json.dumps(
{"x-blocked-by": "Akto Proxy", "reason": reason}
)
blocked_payload["responseHeaders"] = json.dumps(
{"content-type": "application/json"}
)
self.schedule(self.ingest_blocked_request(blocked_payload))
detail = (
f"Blocked by Akto Guardrails: {reason}"
if reason

View file

@ -1,3 +1,4 @@
import asyncio
import json
import os
from unittest.mock import AsyncMock, MagicMock, patch
@ -7,7 +8,10 @@ import pytest
from starlette.exceptions import HTTPException
from litellm.types.utils import GenericGuardrailAPIInputs
from litellm.proxy.guardrails.guardrail_registry import guardrail_initializer_registry, guardrail_class_registry
from litellm.proxy.guardrails.guardrail_registry import (
guardrail_initializer_registry,
guardrail_class_registry,
)
from litellm.proxy.guardrails.guardrail_hooks.akto.akto import AktoGuardrail
@ -70,14 +74,18 @@ def sample_request_data() -> dict:
def _mock_allowed_response():
mock = MagicMock(spec=httpx.Response)
mock.status_code = 200
mock.json.return_value = {"data": {"guardrailsResult": {"Allowed": True, "Reason": ""}}}
mock.json.return_value = {
"data": {"guardrailsResult": {"Allowed": True, "Reason": ""}}
}
return mock
def _mock_blocked_response(reason="Prompt injection detected"):
mock = MagicMock(spec=httpx.Response)
mock.status_code = 200
mock.json.return_value = {"data": {"guardrailsResult": {"Allowed": False, "Reason": reason}}}
mock.json.return_value = {
"data": {"guardrailsResult": {"Allowed": False, "Reason": reason}}
}
return mock
@ -173,8 +181,12 @@ def test_build_akto_payload(akto_validate, sample_inputs, sample_request_data):
assert len(payload["time"]) >= 13
def test_build_akto_payload_with_response(akto_validate, sample_inputs, sample_request_data):
payload = akto_validate.build_akto_payload(sample_inputs, sample_request_data, include_response=True)
def test_build_akto_payload_with_response(
akto_validate, sample_inputs, sample_request_data
):
payload = akto_validate.build_akto_payload(
sample_inputs, sample_request_data, include_response=True
)
resp_body = json.loads(payload["responsePayload"])
assert "choices" in resp_body
@ -193,7 +205,6 @@ def test_build_akto_payload_custom_ids(sample_request_data):
assert payload["akto_vxlan_id"] == "7"
# ---------------------------------------------------------------------------
# Guardrail response handling
# ---------------------------------------------------------------------------
@ -202,7 +213,9 @@ def test_build_akto_payload_custom_ids(sample_request_data):
def test_handle_guardrail_response_allowed():
mock_resp = MagicMock(spec=httpx.Response)
mock_resp.status_code = 200
mock_resp.json.return_value = {"data": {"guardrailsResult": {"Allowed": True, "Reason": ""}}}
mock_resp.json.return_value = {
"data": {"guardrailsResult": {"Allowed": True, "Reason": ""}}
}
allowed, reason = AktoGuardrail.handle_guardrail_response(mock_resp)
assert allowed is True
assert reason == ""
@ -211,7 +224,9 @@ def test_handle_guardrail_response_allowed():
def test_handle_guardrail_response_blocked():
mock_resp = MagicMock(spec=httpx.Response)
mock_resp.status_code = 200
mock_resp.json.return_value = {"data": {"guardrailsResult": {"Allowed": False, "Reason": "PII detected"}}}
mock_resp.json.return_value = {
"data": {"guardrailsResult": {"Allowed": False, "Reason": "PII detected"}}
}
allowed, reason = AktoGuardrail.handle_guardrail_response(mock_resp)
assert allowed is False
assert reason == "PII detected"
@ -251,7 +266,6 @@ def test_handle_guardrail_response_non_dict():
assert allowed is True
def test_handle_guardrail_response_non_json_body():
mock_resp = MagicMock(spec=httpx.Response)
mock_resp.status_code = 200
@ -290,7 +304,9 @@ async def test_pre_call_allowed(akto_validate, sample_inputs, sample_request_dat
@pytest.mark.asyncio
async def test_pre_call_blocked(akto_validate, sample_inputs, sample_request_data):
akto_validate.async_handler.post = AsyncMock(return_value=_mock_blocked_response("PII"))
akto_validate.async_handler.post = AsyncMock(
side_effect=[_mock_blocked_response("PII"), MagicMock(status_code=200)]
)
with pytest.raises(HTTPException) as exc_info:
await akto_validate.apply_guardrail(
@ -299,9 +315,19 @@ async def test_pre_call_blocked(akto_validate, sample_inputs, sample_request_dat
input_type="request",
)
await asyncio.gather(*akto_validate.background_tasks)
assert exc_info.value.status_code == 403
assert "PII" in exc_info.value.detail
akto_validate.async_handler.post.assert_called_once()
assert akto_validate.async_handler.post.call_count == 2
# Second call is the blocked-request ingestion
ingest_call = akto_validate.async_handler.post.call_args_list[1].kwargs
assert ingest_call["params"].get("ingest_data") == "true"
assert "guardrails" not in ingest_call["params"]
ingest_payload = json.loads(ingest_call["data"])
assert ingest_payload["statusCode"] == "403"
assert json.loads(ingest_payload["responsePayload"])["x-blocked-by"] == "Akto Proxy"
# ---------------------------------------------------------------------------
@ -310,7 +336,9 @@ async def test_pre_call_blocked(akto_validate, sample_inputs, sample_request_dat
@pytest.mark.asyncio
async def test_validate_response_noop(akto_validate, sample_inputs, sample_request_data):
async def test_validate_response_noop(
akto_validate, sample_inputs, sample_request_data
):
akto_validate.async_handler.post = AsyncMock()
result = await akto_validate.apply_guardrail(
@ -337,10 +365,14 @@ async def test_fail_open_on_unreachable():
guardrail_name="fail-open-test",
event_hook="pre_call",
)
g.async_handler.post = AsyncMock(side_effect=httpx.ConnectError("Connection refused"))
g.async_handler.post = AsyncMock(
side_effect=httpx.ConnectError("Connection refused")
)
inputs = GenericGuardrailAPIInputs(texts=["test"], model="gpt-4")
result = await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request")
result = await g.apply_guardrail(
inputs=inputs, request_data={}, input_type="request"
)
assert result.get("texts") == ["test"]
@ -354,7 +386,9 @@ async def test_fail_closed_on_unreachable():
guardrail_name="fail-closed-test",
event_hook="pre_call",
)
g.async_handler.post = AsyncMock(side_effect=httpx.ConnectError("Connection refused"))
g.async_handler.post = AsyncMock(
side_effect=httpx.ConnectError("Connection refused")
)
inputs = GenericGuardrailAPIInputs(texts=["test"], model="gpt-4")
with pytest.raises(HTTPException) as exc_info:
@ -384,7 +418,9 @@ async def test_fail_open_on_http_error():
)
inputs = GenericGuardrailAPIInputs(texts=["test"], model="gpt-4")
result = await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request")
result = await g.apply_guardrail(
inputs=inputs, request_data={}, input_type="request"
)
assert result.get("texts") == ["test"]
@ -419,7 +455,9 @@ async def test_fail_closed_on_http_error():
def test_extract_request_path_from_metadata():
path = AktoGuardrail.extract_request_path({"metadata": {"user_api_key_request_route": "/v1/embeddings"}})
path = AktoGuardrail.extract_request_path(
{"metadata": {"user_api_key_request_route": "/v1/embeddings"}}
)
assert path == "/v1/embeddings"
@ -435,7 +473,9 @@ def test_extract_request_path_non_dict_metadata():
def test_resolve_metadata_value():
assert (
AktoGuardrail.resolve_metadata_value({"metadata": {"user_api_key_user_id": "u1"}}, "user_api_key_user_id")
AktoGuardrail.resolve_metadata_value(
{"metadata": {"user_api_key_user_id": "u1"}}, "user_api_key_user_id"
)
== "u1"
)
assert (