refactor akto guardrails to use only pre_call mode, update docs and tests

This commit is contained in:
rzeta-10 2026-03-23 12:41:54 +05:30
parent c89496f378
commit da6a6b3cdf
5 changed files with 70 additions and 310 deletions

View file

@ -3,15 +3,9 @@
## Overview
[Akto](https://www.akto.io/) provides API security guardrails and data ingestion for LLM traffic.
Akto now uses a **two-entry guardrail pattern** in LiteLLM:
- `akto-validate` (`pre_call`) for request validation
- `akto-ingest` (`post_call`) for request/response ingestion
The Akto guardrail uses `pre_call` mode — it validates requests before the LLM call and blocks if flagged.
There is no `on_flagged` setting anymore.
Use these as two separate guardrails in `config.yaml`:
- `guardrail_name: "akto-validate"`
- `guardrail_name: "akto-ingest"`
For non-blocking traffic monitoring/ingestion, use the Akto logging integration (`success_callback: ["akto"]`).
## 1. Get Your Akto Credentials
@ -21,14 +15,6 @@ Set up the Akto Guardrail API Service and grab:
## 2. Configure in `config.yaml`
### Block + Ingest (recommended)
Use both entries below. This gives you:
- pre-call block decision
- post-call ingestion for allowed traffic
Keep these as two separate entries (`akto-validate` and `akto-ingest`).
```yaml
guardrails:
- guardrail_name: "akto-validate"
@ -40,31 +26,6 @@ guardrails:
default_on: true
unreachable_fallback: fail_closed # optional: fail_open | fail_closed (default: fail_closed)
guardrail_timeout: 5 # optional, default: 5
akto_account_id: "1000000" # optional, env fallback: AKTO_ACCOUNT_ID
akto_vxlan_id: "0" # optional, env fallback: AKTO_VXLAN_ID
- guardrail_name: "akto-ingest"
litellm_params:
guardrail: akto
mode: post_call
akto_base_url: os.environ/AKTO_GUARDRAIL_API_BASE
akto_api_key: os.environ/AKTO_API_KEY
default_on: true
```
### Monitor-only mode
If you only want logging/ingestion and no blocking, keep only `akto-ingest`.
```yaml
guardrails:
- guardrail_name: "akto-ingest"
litellm_params:
guardrail: akto
mode: post_call
akto_base_url: os.environ/AKTO_GUARDRAIL_API_BASE
akto_api_key: os.environ/AKTO_API_KEY
default_on: true
```
## 3. Test It
@ -96,29 +57,13 @@ If a request gets blocked:
## 4. How It Works
**Block + Ingest mode:**
```
Request → LiteLLM → Akto guardrail check
→ Allowed → forward to LLM → ingest response
→ Blocked → ingest blocked marker → 403 error
Request → LiteLLM → Akto guardrail check (pre_call, awaited)
→ Allowed → LLM call → response
→ Blocked → 403 error
```
**Monitor-only mode:**
```
Request → LiteLLM → forward to LLM → get response
→ Send to Akto (guardrails + ingest) → log only
```
## 5. Event behavior
| Entry | LiteLLM hook | Akto call behavior |
|------|---|---|
| `akto-validate` | `pre_call` | Awaited call with `guardrails=true`, `ingest_data=false` |
| `akto-ingest` | `post_call` | Fire-and-forget call with `guardrails=true`, `ingest_data=true` |
When blocked in `pre_call`, LiteLLM sends one fire-and-forget ingest payload with blocked metadata and returns `403`.
## 6. Parameters
## 5. Parameters
| Parameter | Env Variable | Default | Description |
|-----------|-------------|---------|-------------|
@ -130,7 +75,7 @@ When blocked in `pre_call`, LiteLLM sends one fire-and-forget ingest payload wit
| `guardrail_timeout` | — | `5` | Timeout in seconds |
| `default_on` | — | `true` (recommended) | Enables the guardrail entry by default |
## 7. Error Handling
## 6. Error Handling
| Scenario | `fail_closed` (default) | `fail_open` |
|----------|------------------------|-------------|

View file

@ -15,8 +15,6 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
_akto_callback = AktoGuardrail(
akto_base_url=getattr(litellm_params, "akto_base_url", None),
akto_api_key=getattr(litellm_params, "akto_api_key", None),
akto_account_id=getattr(litellm_params, "akto_account_id", None),
akto_vxlan_id=getattr(litellm_params, "akto_vxlan_id", None),
unreachable_fallback=getattr(
litellm_params, "unreachable_fallback", "fail_closed"
),

View file

@ -1,13 +1,9 @@
"""Akto guardrail integration for LiteLLM proxy.
Uses a two-config-entry pattern:
- akto-validate (pre_call): Checks request against Akto guardrails, blocks if flagged.
- akto-ingest (post_call): Sends request+response to Akto for data ingestion.
For monitor-only mode, enable only akto-ingest without akto-validate.
Mode:
- pre_call: Validates request against Akto guardrails, blocks if flagged.
"""
import asyncio
import json
import os
from datetime import datetime
@ -38,10 +34,7 @@ DEFAULT_GUARDRAIL_TIMEOUT = 5
class AktoGuardrail(CustomGuardrail):
"""LiteLLM guardrail hook that validates and ingests LLM traffic via the Akto API."""
# Maps event_hook to the input_type it should handle; mismatches are no-ops
HOOK_TO_INPUT = {"pre_call": "request", "post_call": "response"}
"""Validates LLM requests against Akto guardrails."""
@staticmethod
def get_config_model() -> Type["GuardrailConfigModel"]:
@ -56,8 +49,6 @@ class AktoGuardrail(CustomGuardrail):
self,
akto_base_url: Optional[str] = None,
akto_api_key: Optional[str] = None,
akto_account_id: Optional[str] = None,
akto_vxlan_id: Optional[str] = None,
unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
guardrail_timeout: Optional[int] = None,
**kwargs: Any,
@ -75,7 +66,6 @@ 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", "")
@ -95,22 +85,15 @@ class AktoGuardrail(CustomGuardrail):
"fail_closed", "fail_open"
] = unreachable_fallback
self.guardrail_timeout = guardrail_timeout or DEFAULT_GUARDRAIL_TIMEOUT
self.akto_account_id = akto_account_id or os.environ.get(
"AKTO_ACCOUNT_ID", "1000000"
)
self.akto_vxlan_id = akto_vxlan_id or os.environ.get("AKTO_VXLAN_ID", "0")
self.akto_account_id = os.environ.get("AKTO_ACCOUNT_ID", "1000000")
self.akto_vxlan_id = os.environ.get("AKTO_VXLAN_ID", "0")
kwargs["supported_event_hooks"] = [
GuardrailEventHooks.pre_call,
GuardrailEventHooks.post_call,
]
super().__init__(**kwargs)
verbose_proxy_logger.debug(
"Akto guardrail initialized: base_url=%s fallback=%s",
self.akto_base_url,
self.unreachable_fallback,
)
# ── Payload builders ──
@staticmethod
def resolve_metadata_value(request_data: Optional[dict], key: str) -> Optional[str]:
@ -305,25 +288,23 @@ class AktoGuardrail(CustomGuardrail):
"contextSource": "AGENTIC",
}
async def send_request(
self,
*,
guardrails: bool,
ingest_data: bool,
payload: dict,
) -> httpx.Response:
"""Send an HTTP POST to the Akto API endpoint."""
endpoint = f"{self.akto_base_url}{HTTP_PROXY_PATH}"
params = self.build_query_params(guardrails=guardrails, ingest_data=ingest_data)
headers = self.prepare_headers()
# ── HTTP ──
async def send_to_akto(self, payload: dict) -> httpx.Response:
"""POST payload to Akto guardrail API for validation."""
return await self.async_handler.post(
url=endpoint,
data=json.dumps(payload),
params=params,
headers=headers,
params={"akto_connector": AKTO_CONNECTOR_NAME, "guardrails": "true"},
headers={
"content-type": "application/json",
"Authorization": self.akto_api_key,
},
timeout=self.guardrail_timeout,
)
# ── Response parsing ──
@staticmethod
def handle_guardrail_response(response: httpx.Response) -> Tuple[bool, str]:
"""Parse the Akto guardrail response. Returns (allowed, reason)."""
@ -409,84 +390,30 @@ class AktoGuardrail(CustomGuardrail):
input_type: Literal["request", "response"],
logging_obj=None,
) -> GenericGuardrailAPIInputs:
"""Main entry point called by LiteLLM's guardrail framework.
Pre_call (input_type="request"):
- Awaits guardrail check. If blocked, fires off ingest with 403 marker and raises.
Post_call (input_type="response"):
- Fire-and-forget combined guardrail + ingest call.
"""
# Skip if this hook doesn't handle the current input_type
expected = self.HOOK_TO_INPUT.get(str(self.event_hook))
if expected and expected != input_type:
"""Pre_call: validate request against Akto guardrails, block if flagged."""
if input_type != "request":
return inputs
if input_type == "request":
# Pre_call: awaited guardrail check (no ingestion)
payload = self.build_akto_payload(
inputs, request_data, include_response=False
payload = self.build_akto_payload(request_data)
try:
response = await self.send_to_akto(payload)
allowed, reason = self.parse_guardrail_response(response)
except HTTPException:
raise
except (httpx.RequestError, httpx.HTTPStatusError) as e:
if self.unreachable_fallback == "fail_open":
verbose_proxy_logger.critical("Akto unreachable (fail-open): %s", e)
return inputs
raise HTTPException(
status_code=503, detail="Akto guardrail service unreachable"
)
try:
response = await self.send_request(
guardrails=True,
ingest_data=False,
payload=payload,
)
allowed, reason = self.handle_guardrail_response(response)
except HTTPException:
raise
except (httpx.RequestError, httpx.HTTPStatusError) as e:
return self.handle_unreachable(
inputs=inputs,
error=e,
)
if not allowed:
# Build a blocked marker payload with 403 status and reason
blocked_payload = self.build_akto_payload(
inputs,
request_data,
include_response=False,
status_code=403,
)
blocked_payload["responsePayload"] = json.dumps(
{
"body": json.dumps(
{"x-blocked-by": "Akto Proxy", "reason": reason}
),
}
)
blocked_payload["responseHeaders"] = json.dumps(
{"content-type": "application/json"},
)
# Fire-and-forget ingest of the blocked request, then raise 403
task = asyncio.create_task(
self.fire_and_forget_request(
guardrails=False,
ingest_data=True,
payload=blocked_payload,
)
)
self.background_tasks.add(task)
task.add_done_callback(self.background_tasks.discard)
raise HTTPException(
status_code=403,
detail=reason or "Blocked by Akto Guardrails",
)
elif input_type == "response":
# Post_call: fire-and-forget combined guardrail + ingest
payload = self.build_akto_payload(
inputs, request_data, include_response=True
if not allowed:
detail = (
f"Blocked by Akto Guardrails: {reason}"
if reason
else "Blocked by Akto Guardrails"
)
task = asyncio.create_task(
self.fire_and_forget_request(
guardrails=True,
ingest_data=True,
payload=payload,
)
)
self.background_tasks.add(task)
task.add_done_callback(self.background_tasks.discard)
raise HTTPException(status_code=403, detail=detail)
return inputs

View file

@ -9,9 +9,8 @@ class AktoConfigModel(GuardrailConfigModel):
"""
Config for the Akto guardrail.
Use two separate config entries to control behaviour:
akto-validate (mode: pre_call) -> check guardrails, block if flagged
akto-ingest (mode: post_call) -> ingest request+response data
Mode:
pre_call -> validate request, block if flagged
"""
akto_base_url: Optional[str] = Field(
@ -30,16 +29,6 @@ class AktoConfigModel(GuardrailConfigModel):
description="API key for Akto. Env: AKTO_API_KEY.",
)
akto_account_id: Optional[str] = Field(
default=None,
description="Akto account ID for multi-tenant deployments. Env: AKTO_ACCOUNT_ID. Default: '1000000'.",
)
akto_vxlan_id: Optional[str] = Field(
default=None,
description="Akto VXLAN ID. Env: AKTO_VXLAN_ID. Default: '0'.",
)
unreachable_fallback: Literal["fail_closed", "fail_open"] = Field(
default="fail_closed",
description="What to do when Akto is unreachable. 'fail_open' = allow, 'fail_closed' = block.",

View file

@ -1,4 +1,3 @@
import asyncio
import json
import os
from unittest.mock import AsyncMock, MagicMock, patch
@ -43,18 +42,6 @@ def akto_validate():
)
@pytest.fixture
def akto_ingest():
"""AktoGuardrail configured for post_call (akto-ingest)."""
return AktoGuardrail(
akto_base_url="http://localhost:9090",
akto_api_key="test-token",
unreachable_fallback="fail_open",
guardrail_name="test-akto-ingest",
event_hook="post_call",
)
@pytest.fixture
def sample_inputs() -> GenericGuardrailAPIInputs:
return GenericGuardrailAPIInputs(
@ -131,8 +118,8 @@ def test_init_from_env():
"AKTO_VXLAN_ID": "42",
},
):
g = AktoGuardrail(guardrail_name="env-test", event_hook="post_call")
assert g.akto_base_url == "http://env-host:9090"
g = AktoGuardrail(guardrail_name="t", event_hook="pre_call")
assert g.akto_base_url == "http://env:9090"
assert g.akto_api_key == "env-token"
assert g.guardrail_timeout == 5
assert g.akto_account_id == "2000000"
@ -152,29 +139,11 @@ def test_init_defaults():
assert g.akto_vxlan_id == "0"
def test_background_tasks_per_instance():
a = AktoGuardrail(
akto_base_url="http://localhost:9090",
akto_api_key="test-token",
guardrail_name="instance-a",
event_hook="pre_call",
)
b = AktoGuardrail(
akto_base_url="http://localhost:9090",
akto_api_key="test-token",
guardrail_name="instance-b",
event_hook="post_call",
)
assert a.background_tasks is not b.background_tasks
# ── Payload ──
# ---------------------------------------------------------------------------
# Payload format tests
# ---------------------------------------------------------------------------
def test_build_akto_payload_format(akto_validate, sample_inputs, sample_request_data):
payload = akto_validate.build_akto_payload(sample_inputs, sample_request_data, include_response=False)
def test_build_akto_payload(akto_validate, sample_request_data):
payload = akto_validate.build_akto_payload(sample_request_data)
assert payload["path"] == "/v1/chat/completions"
assert payload["method"] == "POST"
@ -209,18 +178,17 @@ def test_build_akto_payload_with_response(akto_validate, sample_inputs, sample_r
assert "choices" in resp_body
def test_build_akto_payload_custom_account_ids(sample_inputs, sample_request_data):
g = AktoGuardrail(
akto_base_url="http://localhost:9090",
akto_api_key="test-token",
akto_account_id="9999",
akto_vxlan_id="7",
guardrail_name="custom-ids-test",
event_hook="pre_call",
)
payload = g.build_akto_payload(sample_inputs, sample_request_data, include_response=False)
assert payload["akto_account_id"] == "9999"
assert payload["akto_vxlan_id"] == "7"
def test_build_akto_payload_custom_ids(sample_request_data):
with patch.dict(os.environ, {"AKTO_ACCOUNT_ID": "9999", "AKTO_VXLAN_ID": "7"}):
g = AktoGuardrail(
akto_base_url="http://x",
akto_api_key="tok",
guardrail_name="t",
event_hook="pre_call",
)
payload = g.build_akto_payload(sample_request_data)
assert payload["akto_account_id"] == "9999"
assert payload["akto_vxlan_id"] == "7"
def test_build_query_params():
@ -331,24 +299,17 @@ async def test_pre_call_allowed(akto_validate, sample_inputs, sample_request_dat
assert result == sample_inputs
akto_validate.async_handler.post.assert_called_once()
call_params = akto_validate.async_handler.post.call_args.kwargs["params"]
assert call_params.get("guardrails") == "true"
assert "ingest_data" not in call_params
params = akto_validate.async_handler.post.call_args.kwargs["params"]
assert params.get("guardrails") == "true"
assert "ingest_data" not in params
# ---------------------------------------------------------------------------
# Pre-call (akto-validate) — blocked
# ---------------------------------------------------------------------------
# ── Pre-call: blocked ──
@pytest.mark.asyncio
async def test_pre_call_blocked(akto_validate, sample_inputs, sample_request_data):
akto_validate.async_handler.post = AsyncMock(
side_effect=[
_mock_blocked_response("PII detected"),
_mock_allowed_response(),
]
)
akto_validate.async_handler.post = AsyncMock(return_value=_mock_blocked("PII"))
with pytest.raises(HTTPException) as exc_info:
await akto_validate.apply_guardrail(
@ -357,25 +318,9 @@ async def test_pre_call_blocked(akto_validate, sample_inputs, sample_request_dat
input_type="request",
)
await asyncio.sleep(0)
await asyncio.sleep(0)
assert exc_info.value.status_code == 403
assert akto_validate.async_handler.post.call_count == 2
first_call_params = akto_validate.async_handler.post.call_args_list[0].kwargs["params"]
assert first_call_params.get("guardrails") == "true"
second_call_params = akto_validate.async_handler.post.call_args_list[1].kwargs["params"]
assert second_call_params.get("ingest_data") == "true"
assert "guardrails" not in second_call_params
second_payload = json.loads(akto_validate.async_handler.post.call_args_list[1].kwargs["data"])
assert second_payload["statusCode"] == "403"
resp_body = json.loads(second_payload["responsePayload"])
inner = json.loads(resp_body["body"])
assert inner["x-blocked-by"] == "Akto Proxy"
assert inner["reason"] == "PII detected"
assert exc.value.status_code == 403
assert "PII" in exc.value.detail
akto_validate.async_handler.post.assert_called_once()
# ---------------------------------------------------------------------------
@ -397,50 +342,6 @@ async def test_validate_response_noop(akto_validate, sample_inputs, sample_reque
akto_validate.async_handler.post.assert_not_called()
# ---------------------------------------------------------------------------
# Post-call (akto-ingest) — combined guardrail + ingest
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_post_call_combined(akto_ingest, sample_inputs, sample_request_data):
akto_ingest.async_handler.post = AsyncMock(return_value=_mock_allowed_response())
result = await akto_ingest.apply_guardrail(
inputs=sample_inputs,
request_data=sample_request_data,
input_type="response",
)
await asyncio.sleep(0)
await asyncio.sleep(0)
assert result == sample_inputs
akto_ingest.async_handler.post.assert_called_once()
call_params = akto_ingest.async_handler.post.call_args.kwargs["params"]
assert call_params.get("guardrails") == "true"
assert call_params.get("ingest_data") == "true"
# ---------------------------------------------------------------------------
# Post-call (akto-ingest) — request input is no-op
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_ingest_request_noop(akto_ingest, sample_inputs, sample_request_data):
akto_ingest.async_handler.post = AsyncMock()
result = await akto_ingest.apply_guardrail(
inputs=sample_inputs,
request_data=sample_request_data,
input_type="request",
)
assert result == sample_inputs
akto_ingest.async_handler.post.assert_not_called()
# ---------------------------------------------------------------------------
# Fail-open / fail-closed
# ---------------------------------------------------------------------------