mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
add ingestion for blocked responses in AktoGuardrail
This commit is contained in:
parent
a0a6cd3940
commit
61fa049f7b
2 changed files with 94 additions and 18 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue