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

1"""Akto guardrail integration for LiteLLM proxy. 

2 

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. 

6 

7For monitor-only mode, enable only akto-ingest without akto-validate. 

8""" 

9 

10import asyncio 

11import json 

12import os 

13from datetime import datetime 

14from typing import TYPE_CHECKING, Final, Literal 

15 

16import httpx 

17from fastapi import HTTPException 

18from typing_extensions import NotRequired, ReadOnly, TypedDict, Unpack 

19 

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 

31 

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 

34 

35 

36class _CustomGuardrailKwargs(TypedDict): 

37 """Keyword arguments forwarded verbatim to CustomGuardrail.__init__.""" 

38 

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]]] 

55 

56 

57HTTP_PROXY_PATH: Final = "/api/http-proxy" 

58AKTO_CONNECTOR_NAME: Final = "litellm" 

59DEFAULT_GUARDRAIL_TIMEOUT: Final = 5 

60 

61 

62class AktoGuardrail(CustomGuardrail): 

63 """LiteLLM guardrail hook that validates and ingests LLM traffic via the Akto API.""" 

64 

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"} 

67 

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 ) 

74 

75 return AktoConfigModel 

76 

77 @classmethod 

78 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: 

79 return [ 

80 GuardrailEventHooks.pre_call, 

81 GuardrailEventHooks.post_call, 

82 ] 

83 

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. 

95 

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() 

108 

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.") 

112 

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.") 

116 

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") 

121 

122 init_kwargs: Final[_CustomGuardrailKwargs] = { 

123 **kwargs, 

124 "supported_event_hooks": list(self.get_supported_event_hooks()), 

125 } 

126 super().__init__(**init_kwargs) 

127 

128 verbose_proxy_logger.debug( 

129 "Akto guardrail initialized: base_url=%s fallback=%s", 

130 self.akto_base_url, 

131 self.unreachable_fallback, 

132 ) 

133 

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 

146 

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" 

155 

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 } 

162 

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 

172 

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 

186 

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} 

195 

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 [] 

206 

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"] 

212 

213 tool_calls: Final = inputs.get("tool_calls") 

214 if tool_calls: 

215 body["tool_calls"] = tool_calls 

216 

217 return body 

218 

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() 

228 

229 texts: Final = inputs.get("texts", []) 

230 if texts: 

231 return {"choices": [{"message": {"content": t, "role": "assistant"}} for t in texts]} 

232 return {} 

233 

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 

245 

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. 

255 

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) 

263 

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"} 

270 

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() 

279 

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 } 

306 

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 ) 

325 

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 ) 

360 

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 

374 

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 ) 

380 

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)) 

402 

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. 

412 

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 

422 

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 ) 

440 

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 ) 

471 

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) 

484 

485 return inputs