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
« 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
13from pydantic import BaseModel, TypeAdapter, ValidationError
14from typing_extensions import NotRequired, ReadOnly, TypedDict
15from websockets.asyncio.client import ClientConnection, connect
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)
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
46class AimGuardrailMissingSecrets(Exception):
47 pass
50class AimRequiredAction(TypedDict):
51 """The ``required_action`` block of an Aim ``/fw/v1/analyze`` response."""
53 action_type: ReadOnly[NotRequired[str]]
54 detection_message: ReadOnly[str]
57class AimAnalysisResult(TypedDict):
58 """The ``analysis_result`` block of an Aim ``/fw/v1/analyze`` response."""
60 policy_drill_down: ReadOnly[Mapping[str, object]]
63class AimRedactedMessage(TypedDict):
64 """One entry of Aim's ``redacted_chat.all_redacted_messages``."""
66 role: ReadOnly[str]
67 content: ReadOnly[str]
70class AimRedactedChat(TypedDict):
71 """The ``redacted_chat`` block of an Aim ``/fw/v1/analyze`` response."""
73 all_redacted_messages: ReadOnly[Sequence[AimRedactedMessage]]
76_REDACTED_CHAT_ADAPTER: Final = TypeAdapter(AimRedactedChat)
79class AimAnalyzeResponse(TypedDict):
80 """Body returned by Aim's ``POST /fw/v1/analyze``."""
82 required_action: ReadOnly[AimRequiredAction]
83 analysis_result: ReadOnly[AimAnalysisResult]
84 redacted_chat: ReadOnly[NotRequired[AimRedactedChat]]
87class AimOutputGuardrailResult(TypedDict, total=False):
88 """Outcome of inspecting one model completion with Aim."""
90 detection_message: ReadOnly[str]
91 redacted_output: ReadOnly[str]
94class AimStreamMessage(TypedDict, total=False):
95 """One frame of Aim's ``/fw/v1/analyze/stream`` websocket protocol."""
97 verified_chunk: ReadOnly[Mapping[str, object]]
98 done: ReadOnly[bool]
99 blocking_message: ReadOnly[str]
102AimStreamChunk: TypeAlias = BaseModel | Mapping[str, object] | str | bytes
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 ]
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)
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)
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
168 await self.call_aim_guardrail(data, hook="moderation", key_alias=user_api_key_dict.key_alias)
169 return data
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
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)]
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 )
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")
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
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)
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}
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}
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 )
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
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
439 raise StreamingCallbackError(blocking_message)
440 verbose_proxy_logger.error("Unknown message received from AIM: %s", result)
441 return
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}))
456 @staticmethod
457 def get_config_model() -> type["GuardrailConfigModel"] | None:
458 from litellm.types.proxy.guardrails.guardrail_hooks.aim import (
459 AimGuardrailConfigModel,
460 )
462 return AimGuardrailConfigModel