Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py: 23%
226 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""Akto guardrail integration for LiteLLM proxy.
3Uses a two-config-entry pattern:
4 - akto-validate (pre_call): Checks request against Akto guardrails, blocks if flagged.
5 - akto-ingest (post_call): Sends request+response to Akto for data ingestion.
7For monitor-only mode, enable only akto-ingest without akto-validate.
8"""
10import asyncio
11import json
12import os
13from datetime import datetime
14from typing import TYPE_CHECKING, Final, Literal
16import httpx
17from fastapi import HTTPException
18from typing_extensions import NotRequired, ReadOnly, TypedDict, Unpack
20from litellm._logging import verbose_proxy_logger
21from litellm.integrations.custom_guardrail import (
22 CustomGuardrail,
23 log_guardrail_information,
24)
25from litellm.llms.custom_httpx.http_handler import (
26 get_async_httpx_client,
27 httpxSpecialProvider,
28)
29from litellm.types.guardrails import GuardrailEventHooks, Mode
30from litellm.types.utils import GenericGuardrailAPIInputs
32if TYPE_CHECKING: 32 ↛ 33line 32 didn't jump to line 33 because the condition on line 32 was never true
33 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
36class _CustomGuardrailKwargs(TypedDict):
37 """Keyword arguments forwarded verbatim to CustomGuardrail.__init__."""
39 guardrail_name: NotRequired[ReadOnly[str | None]]
40 event_hook: NotRequired[ReadOnly[GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None]]
41 default_on: NotRequired[ReadOnly[bool]]
42 mask_request_content: NotRequired[ReadOnly[bool]]
43 mask_response_content: NotRequired[ReadOnly[bool]]
44 violation_message_template: NotRequired[ReadOnly[str | None]]
45 end_session_after_n_fails: NotRequired[ReadOnly[int | None]]
46 on_violation: NotRequired[ReadOnly[str | None]]
47 realtime_violation_message: NotRequired[ReadOnly[str | None]]
48 on_sensitive_data: NotRequired[ReadOnly[str | None]]
49 sensitive_data_route_to_model: NotRequired[ReadOnly[str | None]]
50 sticky_session_routing: NotRequired[ReadOnly[bool]]
51 run_in_parallel: NotRequired[ReadOnly[bool]]
52 scan_raw_request: NotRequired[ReadOnly[bool]]
53 only_scan_new_messages: NotRequired[ReadOnly[bool]]
54 supported_event_hooks: NotRequired[ReadOnly[list[GuardrailEventHooks]]]
57HTTP_PROXY_PATH: Final = "/api/http-proxy"
58AKTO_CONNECTOR_NAME: Final = "litellm"
59DEFAULT_GUARDRAIL_TIMEOUT: Final = 5
62class AktoGuardrail(CustomGuardrail):
63 """LiteLLM guardrail hook that validates and ingests LLM traffic via the Akto API."""
65 # Maps event_hook to the input_type it should handle; mismatches are no-ops
66 HOOK_TO_INPUT = {"pre_call": "request", "post_call": "response"}
68 @staticmethod
69 def get_config_model() -> type["GuardrailConfigModel"]:
70 """Return the Pydantic config model for YAML-based initialization."""
71 from litellm.types.proxy.guardrails.guardrail_hooks.akto import (
72 AktoConfigModel,
73 )
75 return AktoConfigModel
77 @classmethod
78 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
79 return [
80 GuardrailEventHooks.pre_call,
81 GuardrailEventHooks.post_call,
82 ]
84 def __init__(
85 self,
86 akto_base_url: str | None = None,
87 akto_api_key: str | None = None,
88 akto_account_id: str | None = None,
89 akto_vxlan_id: str | None = None,
90 unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
91 guardrail_timeout: int | None = None,
92 **kwargs: Unpack[_CustomGuardrailKwargs],
93 ) -> None:
94 """Initialize the Akto guardrail.
96 Args:
97 akto_base_url: Akto API base URL. Falls back to AKTO_GUARDRAIL_API_BASE env var.
98 akto_api_key: Akto API key. Falls back to AKTO_API_KEY env var.
99 akto_account_id: Akto account ID. Falls back to AKTO_ACCOUNT_ID env var, then "1000000".
100 akto_vxlan_id: Akto VXLAN ID. Falls back to AKTO_VXLAN_ID env var, then "0".
101 unreachable_fallback: Behavior when Akto is unreachable — block or allow.
102 guardrail_timeout: HTTP timeout in seconds for Akto API calls.
103 """
104 self.async_handler = get_async_httpx_client(
105 llm_provider=httpxSpecialProvider.GuardrailCallback,
106 )
107 self.background_tasks: set = set()
109 self.akto_base_url = (akto_base_url or os.environ.get("AKTO_GUARDRAIL_API_BASE", "")).rstrip("/")
110 if not self.akto_base_url:
111 raise ValueError("akto_base_url is required. Set AKTO_GUARDRAIL_API_BASE or pass it in litellm_params.")
113 self.akto_api_key = akto_api_key or os.environ.get("AKTO_API_KEY", "")
114 if not self.akto_api_key:
115 raise ValueError("akto_api_key is required. Set AKTO_API_KEY or pass it in litellm_params.")
117 self.unreachable_fallback: Literal["fail_closed", "fail_open"] = unreachable_fallback
118 self.guardrail_timeout = guardrail_timeout or DEFAULT_GUARDRAIL_TIMEOUT
119 self.akto_account_id = akto_account_id or os.environ.get("AKTO_ACCOUNT_ID", "1000000")
120 self.akto_vxlan_id = akto_vxlan_id or os.environ.get("AKTO_VXLAN_ID", "0")
122 init_kwargs: Final[_CustomGuardrailKwargs] = {
123 **kwargs,
124 "supported_event_hooks": list(self.get_supported_event_hooks()),
125 }
126 super().__init__(**init_kwargs)
128 verbose_proxy_logger.debug(
129 "Akto guardrail initialized: base_url=%s fallback=%s",
130 self.akto_base_url,
131 self.unreachable_fallback,
132 )
134 @staticmethod
135 def resolve_metadata_value(request_data: dict | None, key: str) -> str | None:
136 """Look up a metadata value from litellm_metadata or metadata dicts."""
137 if request_data is None:
138 return None
139 for dict_key in ("litellm_metadata", "metadata"):
140 container = request_data.get(dict_key) or {}
141 if isinstance(container, dict) and container:
142 value = container.get(key)
143 if value is not None:
144 return str(value).strip()
145 return None
147 @staticmethod
148 def extract_request_path(request_data: dict) -> str:
149 """Extract the API route from request metadata, defaulting to /v1/chat/completions."""
150 metadata = request_data.get("metadata") or {}
151 if not isinstance(metadata, dict):
152 metadata = {}
153 route: Final = metadata.get("user_api_key_request_route")
154 return route if route else "/v1/chat/completions"
156 def prepare_headers(self) -> dict[str, str]:
157 """Build HTTP headers for the Akto API call."""
158 return {
159 "content-type": "application/json",
160 "Authorization": self.akto_api_key,
161 }
163 @staticmethod
164 def build_query_params(*, guardrails: bool, ingest_data: bool) -> dict[str, str]:
165 """Build query params that control Akto backend behavior (guardrail check and/or data ingestion)."""
166 params: Final[dict[str, str]] = {"akto_connector": AKTO_CONNECTOR_NAME}
167 if guardrails:
168 params["guardrails"] = "true"
169 if ingest_data:
170 params["ingest_data"] = "true"
171 return params
173 @staticmethod
174 def build_request_headers(request_data: dict) -> dict[str, str]:
175 """Build the requestHeaders field from proxy request headers."""
176 headers: Final[dict[str, str]] = {"content-type": "application/json"}
177 proxy_req: Final = request_data.get("proxy_server_request", {})
178 if not isinstance(proxy_req, dict):
179 return headers
180 proxy_req_headers: Final = proxy_req.get("headers")
181 if isinstance(proxy_req_headers, dict):
182 for key, val in proxy_req_headers.items():
183 if key and val:
184 headers[str(key).lower()] = str(val)
185 return headers
187 @staticmethod
188 def build_request_body(
189 inputs: GenericGuardrailAPIInputs,
190 request_data: dict | None = None,
191 ) -> dict[str, object]:
192 """Build the LLM request body from guardrail inputs (messages, model, tools)."""
193 model: Final = inputs.get("model", "") or ""
194 body: Final[dict[str, object]] = {"model": model}
196 structured: Final = inputs.get("structured_messages")
197 if structured:
198 body["messages"] = structured
199 elif request_data is not None and request_data.get("messages"):
200 body["messages"] = request_data["messages"]
201 if request_data.get("model"):
202 body["model"] = request_data["model"]
203 else:
204 texts: Final = inputs.get("texts", [])
205 body["messages"] = [{"role": "user", "content": t} for t in texts] if texts else []
207 tools: Final = inputs.get("tools")
208 if tools:
209 body["tools"] = tools
210 elif request_data is not None and request_data.get("tools"):
211 body["tools"] = request_data["tools"]
213 tool_calls: Final = inputs.get("tool_calls")
214 if tool_calls:
215 body["tool_calls"] = tool_calls
217 return body
219 @staticmethod
220 def build_response_body(
221 inputs: GenericGuardrailAPIInputs,
222 request_data: dict | None = None,
223 ) -> dict[str, object]:
224 """Build the LLM response body, preferring the actual model response if available."""
225 model_response: Final = request_data.get("response") if request_data else None
226 if model_response is not None and hasattr(model_response, "model_dump"):
227 return model_response.model_dump()
229 texts: Final = inputs.get("texts", [])
230 if texts:
231 return {"choices": [{"message": {"content": t, "role": "assistant"}} for t in texts]}
232 return {}
234 @staticmethod
235 def build_tag_metadata(request_data: dict) -> dict[str, str]:
236 """Build tag/metadata dict with user_id and team_id for Akto tracking."""
237 tag: Final[dict[str, str]] = {"gen-ai": "Gen AI"}
238 user_id: Final = AktoGuardrail.resolve_metadata_value(request_data, "user_api_key_user_id")
239 team_id: Final = AktoGuardrail.resolve_metadata_value(request_data, "user_api_key_team_id")
240 if user_id:
241 tag["user_id"] = user_id
242 if team_id:
243 tag["team_id"] = team_id
244 return tag
246 def build_akto_payload(
247 self,
248 inputs: GenericGuardrailAPIInputs,
249 request_data: dict,
250 *,
251 status_code: int = 200,
252 include_response: bool = False,
253 ) -> dict[str, object]:
254 """Build the flat MIRRORING payload sent to Akto's HTTP proxy endpoint.
256 All body fields use double-encoding: json.dumps({"body": json.dumps(actual_body)})
257 to match the canonical CLI hook format.
258 """
259 request_path: Final = self.extract_request_path(request_data)
260 request_headers: Final = self.build_request_headers(request_data)
261 request_body: Final = self.build_request_body(inputs, request_data)
262 tag: Final = self.build_tag_metadata(request_data)
264 response_payload = json.dumps({}) # Empty body wrapper when no response yet
265 response_headers: dict[str, str] = {}
266 if include_response:
267 response_body: Final = self.build_response_body(inputs, request_data)
268 response_payload = json.dumps({"body": json.dumps(response_body)}) # Double-encoded
269 response_headers = {"content-type": "application/json"}
271 # Extract client IP from proxy headers
272 ip = ""
273 proxy_req: Final = request_data.get("proxy_server_request", {})
274 proxy_headers: Final = proxy_req.get("headers", {}) if isinstance(proxy_req, dict) else {}
275 if isinstance(proxy_headers, dict):
276 ip = proxy_headers.get("x-forwarded-for") or proxy_headers.get("x-real-ip") or ""
277 if "," in ip:
278 ip = ip.split(",")[0].strip()
280 return {
281 "path": request_path,
282 "requestHeaders": json.dumps(request_headers),
283 "responseHeaders": json.dumps(response_headers),
284 "method": "POST",
285 "requestPayload": json.dumps({"body": json.dumps(request_body)}), # Double-encoded
286 "responsePayload": response_payload,
287 "ip": ip,
288 "destIp": "127.0.0.1",
289 "time": str(int(datetime.now().timestamp() * 1000)),
290 "statusCode": str(status_code),
291 "type": "HTTP/1.1",
292 "status": str(status_code),
293 "akto_account_id": self.akto_account_id,
294 "akto_vxlan_id": self.akto_vxlan_id,
295 "is_pending": "false",
296 "source": "MIRRORING",
297 "direction": None,
298 "process_id": None,
299 "socket_id": None,
300 "daemonset_id": None,
301 "enabled_graph": None,
302 "tag": json.dumps(tag),
303 "metadata": json.dumps(tag),
304 "contextSource": "AGENTIC",
305 }
307 async def send_request(
308 self,
309 *,
310 guardrails: bool,
311 ingest_data: bool,
312 payload: dict,
313 ) -> httpx.Response:
314 """Send an HTTP POST to the Akto API endpoint."""
315 endpoint: Final = f"{self.akto_base_url}{HTTP_PROXY_PATH}"
316 params: Final = self.build_query_params(guardrails=guardrails, ingest_data=ingest_data)
317 headers: Final = self.prepare_headers()
318 return await self.async_handler.post(
319 url=endpoint,
320 data=json.dumps(payload),
321 params=params,
322 headers=headers,
323 timeout=self.guardrail_timeout,
324 )
326 @staticmethod
327 def handle_guardrail_response(response: httpx.Response) -> tuple[bool, str]:
328 """Parse the Akto guardrail response. Returns (allowed, reason)."""
329 if response.status_code != 200:
330 verbose_proxy_logger.error("Akto returned HTTP %d", response.status_code)
331 raise httpx.HTTPStatusError(
332 f"Akto returned unexpected status {response.status_code}",
333 request=response.request,
334 response=response,
335 )
336 try:
337 result: Final = response.json()
338 except (json.JSONDecodeError, ValueError) as e:
339 response_text: Final = getattr(response, "text", "")
340 verbose_proxy_logger.error(
341 "Akto returned non-JSON body for status 200: %r",
342 response_text[:200],
343 )
344 raise httpx.RequestError(
345 "Akto returned non-JSON body",
346 request=response.request,
347 ) from e
348 if not isinstance(result, dict):
349 return True, ""
350 data: Final = result.get("data") or {}
351 if not isinstance(data, dict):
352 return True, ""
353 guardrails_result: Final = data.get("guardrailsResult") or {}
354 if not isinstance(guardrails_result, dict):
355 return True, ""
356 return (
357 bool(guardrails_result.get("Allowed", True)),
358 str(guardrails_result.get("Reason", "")),
359 )
361 def handle_unreachable(
362 self,
363 inputs: GenericGuardrailAPIInputs,
364 error: Exception,
365 ) -> GenericGuardrailAPIInputs:
366 """Handle Akto being unreachable based on fail_open/fail_closed config."""
367 if self.unreachable_fallback == "fail_open":
368 verbose_proxy_logger.critical(
369 "Akto unreachable (fail-open): %s",
370 str(error),
371 exc_info=error,
372 )
373 return inputs
375 verbose_proxy_logger.error("Akto unreachable (fail-closed): %s", str(error))
376 raise HTTPException(
377 status_code=503,
378 detail="Akto guardrail service unreachable",
379 )
381 async def fire_and_forget_request(
382 self,
383 *,
384 guardrails: bool,
385 ingest_data: bool,
386 payload: dict,
387 ) -> None:
388 """Send a request without awaiting it in the caller. Errors are logged, not raised."""
389 try:
390 response: Final = await self.send_request(
391 guardrails=guardrails,
392 ingest_data=ingest_data,
393 payload=payload,
394 )
395 if response.status_code != 200:
396 verbose_proxy_logger.error(
397 "Akto fire-and-forget returned HTTP %d",
398 response.status_code,
399 )
400 except Exception as e:
401 verbose_proxy_logger.error("Akto fire-and-forget error: %s", str(e))
403 @log_guardrail_information
404 async def apply_guardrail(
405 self,
406 inputs: GenericGuardrailAPIInputs,
407 request_data: dict,
408 input_type: Literal["request", "response"],
409 logging_obj=None,
410 ) -> GenericGuardrailAPIInputs:
411 """Main entry point called by LiteLLM's guardrail framework.
413 Pre_call (input_type="request"):
414 - Awaits guardrail check. If blocked, fires off ingest with 403 marker and raises.
415 Post_call (input_type="response"):
416 - Fire-and-forget combined guardrail + ingest call.
417 """
418 # Skip if this hook doesn't handle the current input_type
419 expected: Final = self.HOOK_TO_INPUT.get(str(self.event_hook))
420 if expected and expected != input_type:
421 return inputs
423 if input_type == "request":
424 # Pre_call: awaited guardrail check (no ingestion)
425 payload = self.build_akto_payload(inputs, request_data, include_response=False)
426 try:
427 response: Final = await self.send_request(
428 guardrails=True,
429 ingest_data=False,
430 payload=payload,
431 )
432 allowed, reason = self.handle_guardrail_response(response)
433 except HTTPException:
434 raise
435 except (httpx.RequestError, httpx.HTTPStatusError) as e:
436 return self.handle_unreachable(
437 inputs=inputs,
438 error=e,
439 )
441 if not allowed:
442 # Build a blocked marker payload with 403 status and reason
443 blocked_payload: Final = self.build_akto_payload(
444 inputs,
445 request_data,
446 include_response=False,
447 status_code=403,
448 )
449 blocked_payload["responsePayload"] = json.dumps(
450 {
451 "body": json.dumps({"x-blocked-by": "Akto Proxy", "reason": reason}),
452 }
453 )
454 blocked_payload["responseHeaders"] = json.dumps(
455 {"content-type": "application/json"},
456 )
457 # Fire-and-forget ingest of the blocked request, then raise 403
458 task = asyncio.create_task(
459 self.fire_and_forget_request(
460 guardrails=False,
461 ingest_data=True,
462 payload=blocked_payload,
463 )
464 )
465 self.background_tasks.add(task)
466 task.add_done_callback(self.background_tasks.discard)
467 raise HTTPException(
468 status_code=403,
469 detail=reason or "Blocked by Akto Guardrails",
470 )
472 elif input_type == "response":
473 # Post_call: fire-and-forget combined guardrail + ingest
474 payload = self.build_akto_payload(inputs, request_data, include_response=True)
475 task = asyncio.create_task(
476 self.fire_and_forget_request(
477 guardrails=True,
478 ingest_data=True,
479 payload=payload,
480 )
481 )
482 self.background_tasks.add(task)
483 task.add_done_callback(self.background_tasks.discard)
485 return inputs