Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py: 27%

198 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1# +-------------------------------------------------------------+ 

2# 

3# Use Aim Security Guardrails for your LLM calls 

4# https://www.aim.security/ 

5# 

6# +-------------------------------------------------------------+ 

7import asyncio 

8import json 

9import os 

10from collections.abc import AsyncGenerator, AsyncIterator, Mapping, Sequence 

11from typing import TYPE_CHECKING, Final, TypeAlias 

12 

13from pydantic import BaseModel, TypeAdapter, ValidationError 

14from typing_extensions import NotRequired, ReadOnly, TypedDict 

15from websockets.asyncio.client import ClientConnection, connect 

16 

17from litellm import DualCache 

18from litellm._logging import verbose_proxy_logger 

19from litellm._version import version as litellm_version 

20from litellm.integrations.custom_guardrail import CustomGuardrail 

21from litellm.llms.custom_httpx.http_handler import ( 

22 get_async_httpx_client, 

23 httpxSpecialProvider, 

24) 

25from litellm.proxy._types import ProxyException, UserAPIKeyAuth 

26from litellm.proxy.guardrails._content_utils import ( 

27 apply_redacted_messages_back, 

28 build_inspection_messages, 

29 has_non_string_content, 

30 is_non_conversational_call_type, 

31 is_string_batch_input, 

32) 

33from litellm.types.guardrails import GuardrailEventHooks 

34from litellm.types.utils import ( 

35 CallTypesLiteral, 

36 Choices, 

37 LLMResponseTypes, 

38 ModelResponse, 

39 ModelResponseStream, 

40) 

41 

42if TYPE_CHECKING: 42 ↛ 43line 42 didn't jump to line 43 because the condition on line 42 was never true

43 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel 

44 

45 

46class AimGuardrailMissingSecrets(Exception): 

47 pass 

48 

49 

50class AimRequiredAction(TypedDict): 

51 """The ``required_action`` block of an Aim ``/fw/v1/analyze`` response.""" 

52 

53 action_type: ReadOnly[NotRequired[str]] 

54 detection_message: ReadOnly[str] 

55 

56 

57class AimAnalysisResult(TypedDict): 

58 """The ``analysis_result`` block of an Aim ``/fw/v1/analyze`` response.""" 

59 

60 policy_drill_down: ReadOnly[Mapping[str, object]] 

61 

62 

63class AimRedactedMessage(TypedDict): 

64 """One entry of Aim's ``redacted_chat.all_redacted_messages``.""" 

65 

66 role: ReadOnly[str] 

67 content: ReadOnly[str] 

68 

69 

70class AimRedactedChat(TypedDict): 

71 """The ``redacted_chat`` block of an Aim ``/fw/v1/analyze`` response.""" 

72 

73 all_redacted_messages: ReadOnly[Sequence[AimRedactedMessage]] 

74 

75 

76_REDACTED_CHAT_ADAPTER: Final = TypeAdapter(AimRedactedChat) 

77 

78 

79class AimAnalyzeResponse(TypedDict): 

80 """Body returned by Aim's ``POST /fw/v1/analyze``.""" 

81 

82 required_action: ReadOnly[AimRequiredAction] 

83 analysis_result: ReadOnly[AimAnalysisResult] 

84 redacted_chat: ReadOnly[NotRequired[AimRedactedChat]] 

85 

86 

87class AimOutputGuardrailResult(TypedDict, total=False): 

88 """Outcome of inspecting one model completion with Aim.""" 

89 

90 detection_message: ReadOnly[str] 

91 redacted_output: ReadOnly[str] 

92 

93 

94class AimStreamMessage(TypedDict, total=False): 

95 """One frame of Aim's ``/fw/v1/analyze/stream`` websocket protocol.""" 

96 

97 verified_chunk: ReadOnly[Mapping[str, object]] 

98 done: ReadOnly[bool] 

99 blocking_message: ReadOnly[str] 

100 

101 

102AimStreamChunk: TypeAlias = BaseModel | Mapping[str, object] | str | bytes 

103 

104 

105class AimGuardrail(CustomGuardrail): 

106 @classmethod 

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

108 return [ 

109 GuardrailEventHooks.pre_call, 

110 GuardrailEventHooks.during_call, 

111 GuardrailEventHooks.post_call, 

112 ] 

113 

114 def __init__( 

115 self, 

116 api_key: str | None = None, 

117 api_base: str | None = None, 

118 inspect_embeddings: bool | None = None, 

119 **kwargs, 

120 ): 

121 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) 

122 self.inspect_embeddings: Final = inspect_embeddings is True 

123 ssl_verify: Final = kwargs.pop("ssl_verify", None) 

124 self.async_handler = get_async_httpx_client( 

125 llm_provider=httpxSpecialProvider.GuardrailCallback, 

126 params={"ssl_verify": ssl_verify} if ssl_verify is not None else None, 

127 ) 

128 self.api_key = api_key or os.environ.get("AIM_API_KEY") 

129 if not self.api_key: 

130 msg: Final = ( 

131 "Couldn't get Aim api key, either set the `AIM_API_KEY` in the environment or " 

132 "pass it as a parameter to the guardrail in the config file" 

133 ) 

134 raise AimGuardrailMissingSecrets(msg) 

135 self.api_base = api_base or os.environ.get("AIM_API_BASE") or "https://api.aim.security" 

136 self.ws_api_base = self.api_base.replace("http://", "ws://").replace("https://", "wss://") 

137 self.dlp_entities: list[dict] = [] 

138 self._max_dlp_entities = 100 

139 super().__init__(**kwargs) 

140 

141 async def async_pre_call_hook( 

142 self, 

143 user_api_key_dict: UserAPIKeyAuth, 

144 cache: DualCache, 

145 data: dict, 

146 call_type: CallTypesLiteral, 

147 ) -> Exception | str | dict | None: 

148 verbose_proxy_logger.debug("Inside AIM Pre-Call Hook") 

149 # /embeddings carries ``input`` — documents being indexed, not a prompt — which 

150 # the flatten lifts into synthetic chat messages. A verdict on that text then 

151 # blocks or silently rewrites a request that was never a conversation. 

152 if is_non_conversational_call_type(call_type) and not self.inspect_embeddings: 

153 verbose_proxy_logger.debug("Aim: skipping non-conversational call type %s", call_type) 

154 return data 

155 return await self.call_aim_guardrail(data, hook="pre_call", key_alias=user_api_key_dict.key_alias) 

156 

157 async def async_moderation_hook( 

158 self, 

159 data: dict, 

160 user_api_key_dict: UserAPIKeyAuth, 

161 call_type: CallTypesLiteral, 

162 ) -> Exception | str | dict | None: 

163 verbose_proxy_logger.debug("Inside AIM Moderation Hook") 

164 if is_non_conversational_call_type(call_type) and not self.inspect_embeddings: 

165 verbose_proxy_logger.debug("Aim: skipping non-conversational call type %s", call_type) 

166 return data 

167 

168 await self.call_aim_guardrail(data, hook="moderation", key_alias=user_api_key_dict.key_alias) 

169 return data 

170 

171 async def call_aim_guardrail(self, data: dict, hook: str, key_alias: str | None) -> dict: 

172 user_email: Final = data.get("metadata", {}).get("headers", {}).get("x-aim-user-email") 

173 call_id: Final = data.get("litellm_call_id") 

174 headers: Final = self._build_aim_headers( 

175 hook=hook, 

176 key_alias=key_alias, 

177 user_email=user_email, 

178 litellm_call_id=call_id, 

179 ) 

180 response: Final = await self.async_handler.post( 

181 f"{self.api_base}/fw/v1/analyze", 

182 headers=headers, 

183 json={"messages": self._build_aim_inspection_messages(data)}, 

184 ) 

185 response.raise_for_status() 

186 res: Final[AimAnalyzeResponse] = response.json() 

187 required_action: Final = res.get("required_action") 

188 action_type: Final = required_action and required_action.get("action_type", None) 

189 if action_type is None: 

190 verbose_proxy_logger.debug("Aim: No required action specified") 

191 return data 

192 if action_type == "monitor_action": 

193 verbose_proxy_logger.info("Aim: monitor action") 

194 elif action_type == "block_action": 

195 self._handle_block_action(res["analysis_result"], required_action) 

196 elif action_type == "anonymize_action": 

197 return self._anonymize_request(res, data) 

198 else: 

199 verbose_proxy_logger.error("Aim: %s action", action_type) 

200 return data 

201 

202 @staticmethod 

203 def _build_aim_inspection_messages(data: dict) -> list[dict[str, str]]: 

204 """AIM validates against the OpenAI chat schema. Bare ``role: "tool"`` 

205 without ``tool_call_id`` and bare ``role: "function"`` without ``name`` 

206 are rejected; the flatten drops those fields, so any role outside 

207 ``{system, user, assistant}`` collapses to ``user`` for the AIM POST.""" 

208 safe_roles: Final = {"system", "user", "assistant"} 

209 return [{**m, "role": "user"} if m["role"] not in safe_roles else m for m in build_inspection_messages(data)] 

210 

211 @staticmethod 

212 def _rejection(message: str, *, openai_code: str | None = None) -> ProxyException: 

213 return ProxyException( 

214 message=message, 

215 type="invalid_request_error", 

216 param=None, 

217 code=400, 

218 openai_code=openai_code, 

219 ) 

220 

221 def _handle_block_action(self, analysis_result: AimAnalysisResult, required_action: AimRequiredAction) -> None: 

222 detection_message: Final = required_action.get("detection_message", None) 

223 verbose_proxy_logger.info( 

224 "Aim: Violation detected enabled policies: {policies}".format( 

225 policies=list(analysis_result["policy_drill_down"].keys()), 

226 ), 

227 ) 

228 raise self._rejection(detection_message, openai_code="content_policy_violation") 

229 

230 def _anonymize_request(self, res: AimAnalyzeResponse, data: dict) -> dict: 

231 verbose_proxy_logger.info("Aim: anonymize action") 

232 redacted_chat: Final = res.get("redacted_chat") 

233 if not redacted_chat: 

234 return data 

235 # Aim returns text-only redacted messages. Overwriting 

236 # ``data["messages"]`` with that would silently strip image/audio 

237 # parts from a multimodal request — degrade to block so the 

238 # multimodal payload is never silently rewritten. 

239 if has_non_string_content(data) and not is_string_batch_input(data): 

240 raise self._rejection( 

241 "Aim: anonymize action requested for multimodal input " 

242 "but mask-in-place would drop non-text parts. Send the " 

243 "request with plain string content to use anonymize, " 

244 "or rely on block-mode policies." 

245 ) 

246 try: 

247 redacted_chat_model: Final = _REDACTED_CHAT_ADAPTER.validate_python(redacted_chat) 

248 except ValidationError: 

249 raise self._rejection( 

250 "Aim: anonymize action returned malformed redacted messages, " 

251 "so the request cannot be rewritten without forwarding unredacted text." 

252 ) from None 

253 redacted_messages: Final = list(redacted_chat_model["all_redacted_messages"]) 

254 if len(redacted_messages) != len(build_inspection_messages(data)): 

255 raise self._rejection( 

256 "Aim: anonymize action returned a redacted batch of a different " 

257 "size than the inspected input, so the request cannot be " 

258 "rewritten without forwarding unredacted text." 

259 ) 

260 # Write back to ``messages`` AND ``input``. The Responses-API 

261 # backend reads ``input``; writing only to ``messages`` would let 

262 # unredacted text reach the LLM for ``/v1/responses`` calls. 

263 if not apply_redacted_messages_back(data, redacted_messages): 

264 raise self._rejection( 

265 "Aim: anonymize action returned a redacted batch of a different " 

266 "size than the inspected input, so the request cannot be " 

267 "rewritten without forwarding unredacted text." 

268 ) 

269 return data 

270 

271 async def call_aim_guardrail_on_output( 

272 self, request_data: dict, output: str, hook: str, key_alias: str | None 

273 ) -> AimOutputGuardrailResult | None: 

274 user_email: Final = request_data.get("metadata", {}).get("headers", {}).get("x-aim-user-email") 

275 call_id: Final = request_data.get("litellm_call_id") 

276 response: Final = await self.async_handler.post( 

277 f"{self.api_base}/fw/v1/analyze", 

278 headers=self._build_aim_headers( 

279 hook=hook, 

280 key_alias=key_alias, 

281 user_email=user_email, 

282 litellm_call_id=call_id, 

283 ), 

284 json={ 

285 "messages": self._build_aim_inspection_messages(request_data) 

286 + [{"role": "assistant", "content": output}] 

287 }, 

288 ) 

289 response.raise_for_status() 

290 res: Final[AimAnalyzeResponse] = response.json() 

291 required_action: Final = res.get("required_action") 

292 action_type: Final = required_action and required_action.get("action_type", None) 

293 if action_type and action_type == "block_action": 

294 return self._handle_block_action_on_output(res["analysis_result"], required_action) 

295 redacted_chat: Final = res.get("redacted_chat", None) 

296 

297 if action_type != "anonymize_action": 

298 return {"redacted_output": output} 

299 try: 

300 redacted_chat_model: Final = _REDACTED_CHAT_ADAPTER.validate_python(redacted_chat) 

301 except ValidationError: 

302 raise self._rejection( 

303 "Aim: anonymize action returned malformed redacted output, " 

304 "so the response cannot be rewritten without forwarding unredacted text." 

305 ) from None 

306 redacted_messages: Final = redacted_chat_model["all_redacted_messages"] 

307 inspected_messages: Final = self._build_aim_inspection_messages(request_data) 

308 if len(redacted_messages) != len(inspected_messages) + 1: 

309 raise self._rejection( 

310 "Aim: anonymize action returned an invalid redacted output count, " 

311 "so the response cannot be rewritten without forwarding unredacted text." 

312 ) 

313 redacted_output: Final = redacted_messages[-1]["content"] 

314 if not redacted_output: 

315 raise self._rejection( 

316 "Aim: anonymize action returned empty redacted output, " 

317 "so the response cannot be rewritten without forwarding unredacted text." 

318 ) 

319 return {"redacted_output": redacted_output} 

320 

321 def _handle_block_action_on_output( 

322 self, analysis_result: AimAnalysisResult, required_action: AimRequiredAction 

323 ) -> AimOutputGuardrailResult | None: 

324 detection_message: Final = required_action.get("detection_message", None) 

325 verbose_proxy_logger.info( 

326 "Aim: detected: {detected}, enabled policies: {policies}".format( 

327 detected=True, 

328 policies=list(analysis_result["policy_drill_down"].keys()), 

329 ), 

330 ) 

331 return {"detection_message": detection_message} 

332 

333 def _build_aim_headers( 

334 self, 

335 *, 

336 hook: str, 

337 key_alias: str | None, 

338 user_email: str | None, 

339 litellm_call_id: str | None, 

340 ): 

341 """ 

342 A helper function to build the http headers that are required by AIM guardrails. 

343 """ 

344 return ( 

345 { 

346 "Authorization": f"Bearer {self.api_key}", 

347 # Used by Aim to apply only the guardrails that should be applied in a specific request phase. 

348 "x-aim-litellm-hook": hook, 

349 # Used by Aim to track LiteLLM version and provide backward compatibility. 

350 "x-aim-litellm-version": litellm_version, 

351 } 

352 # Used by Aim to track together single call input and output 

353 | ({"x-aim-call-id": litellm_call_id} if litellm_call_id else {}) 

354 # Used by Aim to track guardrails violations by user. 

355 | ({"x-aim-user-email": user_email} if user_email else {}) 

356 | ( 

357 { 

358 # Used by Aim apply only the guardrails that are associated with the key alias. 

359 "x-aim-gateway-key-alias": key_alias, 

360 } 

361 if key_alias 

362 else {} 

363 ) 

364 ) 

365 

366 async def async_post_call_success_hook( 

367 self, 

368 data: dict, 

369 user_api_key_dict: UserAPIKeyAuth, 

370 response: LLMResponseTypes, 

371 ) -> LLMResponseTypes: 

372 if not (isinstance(response, ModelResponse) and response.choices): 

373 return response 

374 # Inspect every choice — when ``n>1`` the additional completions 

375 # used to bypass Aim entirely because the hook only inspected 

376 # ``choices[0]``. Run inspections concurrently so multi-completion 

377 # responses don't pay an n× latency penalty. 

378 choices_to_inspect: Final = [c for c in response.choices if isinstance(c, Choices)] 

379 if not choices_to_inspect: 

380 return response 

381 # ``return_exceptions=True`` lets every inspection finish even if 

382 # one fails — without it, the first exception would propagate and 

383 # leave the remaining tasks running in the background. 

384 results: Final = await asyncio.gather( 

385 *( 

386 self.call_aim_guardrail_on_output( 

387 data, 

388 choice.message.content or "", 

389 hook="output", 

390 key_alias=user_api_key_dict.key_alias, 

391 ) 

392 for choice in choices_to_inspect 

393 ), 

394 return_exceptions=True, 

395 ) 

396 for choice, aim_output_guardrail_result in zip(choices_to_inspect, results): 

397 if isinstance(aim_output_guardrail_result, BaseException): 

398 raise aim_output_guardrail_result 

399 if aim_output_guardrail_result and ( 

400 detection_message := aim_output_guardrail_result.get("detection_message") 

401 ): 

402 raise self._rejection( 

403 detection_message, 

404 openai_code="content_policy_violation", 

405 ) 

406 if aim_output_guardrail_result and aim_output_guardrail_result.get("redacted_output"): 

407 choice.message.content = aim_output_guardrail_result.get("redacted_output") 

408 return response 

409 

410 async def async_post_call_streaming_iterator_hook( 

411 self, 

412 user_api_key_dict: UserAPIKeyAuth, 

413 response: AsyncIterator[AimStreamChunk], 

414 request_data: dict, 

415 ) -> AsyncGenerator[ModelResponseStream, None]: 

416 user_email: Final = request_data.get("metadata", {}).get("headers", {}).get("x-aim-user-email") 

417 call_id: Final = request_data.get("litellm_call_id") 

418 async with connect( 

419 f"{self.ws_api_base}/fw/v1/analyze/stream", 

420 additional_headers=self._build_aim_headers( 

421 hook="output", 

422 key_alias=user_api_key_dict.key_alias, 

423 user_email=user_email, 

424 litellm_call_id=call_id, 

425 ), 

426 ) as websocket: 

427 sender: Final = asyncio.create_task(self.forward_the_stream_to_aim(websocket, response)) 

428 while True: 

429 result: AimStreamMessage = json.loads(await websocket.recv()) 

430 if verified_chunk := result.get("verified_chunk"): 

431 yield ModelResponseStream.model_validate(verified_chunk) 

432 else: 

433 sender.cancel() 

434 if result.get("done"): 

435 return 

436 if blocking_message := result.get("blocking_message"): 

437 from litellm.proxy.proxy_server import StreamingCallbackError 

438 

439 raise StreamingCallbackError(blocking_message) 

440 verbose_proxy_logger.error("Unknown message received from AIM: %s", result) 

441 return 

442 

443 async def forward_the_stream_to_aim( 

444 self, 

445 websocket: ClientConnection, 

446 response_iter: AsyncIterator[AimStreamChunk], 

447 ) -> None: 

448 async for chunk in response_iter: 

449 if isinstance(chunk, BaseModel): 

450 chunk = chunk.model_dump_json() 

451 if isinstance(chunk, dict): 

452 chunk = json.dumps(chunk) 

453 await websocket.send(chunk) 

454 await websocket.send(json.dumps({"done": True})) 

455 

456 @staticmethod 

457 def get_config_model() -> type["GuardrailConfigModel"] | None: 

458 from litellm.types.proxy.guardrails.guardrail_hooks.aim import ( 

459 AimGuardrailConfigModel, 

460 ) 

461 

462 return AimGuardrailConfigModel