Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py: 24%
296 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
1import json
2import os
3import time
4from collections.abc import Mapping, Sequence
5from typing import TYPE_CHECKING, Annotated, Final, Literal, NamedTuple, Optional, cast
7from fastapi import HTTPException
8from pydantic import BaseModel, ConfigDict, Field, ValidationError
9from typing_extensions import override
11from litellm._logging import verbose_proxy_logger
12from litellm.integrations.custom_guardrail import (
13 CustomGuardrail,
14 log_guardrail_information,
15)
16from litellm.llms.base_llm.guardrail_translation.utils import (
17 effective_skip_system_message_for_guardrail,
18 effective_skip_tool_message_for_guardrail,
19)
20from litellm.llms.custom_httpx.http_handler import (
21 AsyncHTTPHandler,
22 get_async_httpx_client,
23 httpxSpecialProvider,
24)
25from litellm.proxy.common_utils.callback_utils import (
26 add_guardrail_to_applied_guardrails_header,
27)
28from litellm.types.guardrails import GuardrailEventHooks, LitellmParams
29from litellm.types.llms.openai import AllMessageValues, OpenAIChatCompletionToolParam
30from litellm.types.proxy.guardrails.guardrail_hooks.crowdstrike_aidr import (
31 CrowdStrikeAIDRGuardrailConfigModelOptionalParams,
32)
33from litellm.types.utils import GenericGuardrailAPIInputs
35if TYPE_CHECKING: 35 ↛ 36line 35 didn't jump to line 36 because the condition on line 35 was never true
36 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
37 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
40class CrowdStrikeAIDRGuardrailMissingSecrets(Exception):
41 """Custom exception for missing CrowdStrike AIDR secrets."""
44class _TextContentPart(BaseModel):
45 model_config = ConfigDict(extra="forbid")
47 type: Literal["text"] = "text"
48 text: str
51class _ImageUrl(BaseModel):
52 url: str
55class _ImageUrlContentPart(BaseModel):
56 model_config = ConfigDict(extra="forbid")
58 type: Literal["image_url"] = "image_url"
59 image_url: _ImageUrl
62_ContentPart = Annotated[_TextContentPart | _ImageUrlContentPart, Field(discriminator="type")]
65class _Message(BaseModel):
66 role: str
67 content: str | list[_ContentPart] | None = None
70class _GuardInput(BaseModel):
71 messages: list[_Message]
72 tools: Sequence[OpenAIChatCompletionToolParam] | None = None
75class _GuardChatCompletionsResult(BaseModel):
76 guard_output: _GuardInput | None = None
77 """Updated structured prompt."""
78 blocked: bool | None = None
79 """Whether or not the prompt triggered a block detection."""
80 transformed: bool | None = None
81 """Whether or not the original input was transformed."""
82 detectors: dict[str, object] | None = None
83 """Result of the policy analyzing and input prompt."""
86class _GuardChatCompletionsResponse(BaseModel):
87 result: _GuardChatCompletionsResult | None = None
90class _FilteredMessages(NamedTuple):
91 """Subset of a conversation selected for guardrail analysis."""
93 messages: list[AllMessageValues]
94 """Messages subset."""
95 indices: tuple[int, ...]
96 """Positions of the subset's messages in the original list."""
99class _GuardInputWithIndices(NamedTuple):
100 guard_input: _GuardInput
101 """Guard API payload."""
102 sent_indices: tuple[int, ...]
103 """Positions of the guard input's messages in the original list."""
106def _normalize_content(raw: object) -> str | list[_ContentPart] | None:
107 if raw is None:
108 return None
109 if isinstance(raw, str):
110 return raw
111 if not isinstance(raw, list):
112 return json.dumps(raw)
113 parts: Final[list[_ContentPart]] = []
114 for block in raw:
115 if not isinstance(block, dict):
116 parts.append(_TextContentPart(text=json.dumps(block)))
117 continue
119 t = block.get("type")
120 if t == "text" and isinstance(block.get("text"), str):
121 parts.append(_TextContentPart(text=cast(str, block["text"])))
122 elif t == "image_url":
123 iu = block.get("image_url")
124 url = iu if isinstance(iu, str) else str((iu or {}).get("url", ""))
125 parts.append(_ImageUrlContentPart(image_url=_ImageUrl(url=url)))
127 # Any other types are not recognized by the CrowdStrike AIDR API.
129 return parts
132def _extract_text_from_content(content: object) -> str:
133 if isinstance(content, str):
134 return content
135 if isinstance(content, list):
136 parts = [item.get("text", "") for item in content if isinstance(item, dict) and item.get("type") == "text"]
137 return "\n".join(parts)
138 return ""
141def _extract_text_from_message(message: _Message) -> str:
142 content: Final = message.content
143 if isinstance(content, str):
144 return content
145 if content is None:
146 return ""
147 return "\n".join(part.text for part in content if isinstance(part, _TextContentPart))
150def _merge_metadata_bags(request_data: Mapping[str, object]) -> Mapping[str, object] | None:
151 merged: Final[dict[str, object]] = {}
152 present = False
153 for bag in (request_data.get("metadata"), request_data.get("litellm_metadata")):
154 if isinstance(bag, Mapping):
155 present = True
156 merged.update(bag)
157 return merged if present else None
160def streaming_params_from_litellm_params(
161 litellm_params: LitellmParams,
162) -> CrowdStrikeAIDRGuardrailConfigModelOptionalParams:
163 extras: Final[Mapping[str, object]] = litellm_params.model_extra or {}
164 nested: Final = litellm_params.optional_params
165 optional_params: Final[Mapping[str, object]] = {} if nested is None else nested.model_dump()
166 return CrowdStrikeAIDRGuardrailConfigModelOptionalParams.model_validate(
167 {
168 name: value
169 for name in CrowdStrikeAIDRGuardrailConfigModelOptionalParams.model_fields
170 if (value := optional_params.get(name, extras.get(name))) is not None
171 }
172 )
175def _messages_since_last_assistant(
176 messages: Sequence[AllMessageValues],
177) -> _FilteredMessages:
178 if not messages:
179 return _FilteredMessages([], ())
181 if messages[-1]["role"] == "assistant":
182 indices = tuple(i for i, m in enumerate(messages) if m["role"] == "system") + (len(messages) - 1,)
183 return _FilteredMessages([messages[i] for i in indices], indices)
185 last_assistant_idx = -1
186 for i in range(len(messages) - 1, -1, -1):
187 if messages[i]["role"] == "assistant":
188 last_assistant_idx = i
189 break
191 system_indices: Final = tuple(i for i in range(last_assistant_idx + 1) if messages[i]["role"] == "system")
192 tail_indices: Final = tuple(range(last_assistant_idx + 1, len(messages)))
193 indices = system_indices + tail_indices
194 return _FilteredMessages([messages[i] for i in indices], indices)
197def _merge_request_transforms(
198 guard_output: _GuardInput,
199 structured_messages: list[AllMessageValues] | None,
200 texts: list[str],
201 sent_indices: tuple[int, ...],
202) -> list[str]:
203 returned_texts: Final = [_extract_text_from_message(msg) for msg in guard_output.messages]
204 original_texts: Final = (
205 [_extract_text_from_content(m.get("content")) for m in structured_messages] if structured_messages else texts
206 )
207 replacements: Final = {
208 idx: returned_texts[pos]
209 for pos, idx in enumerate(sent_indices)
210 if pos < len(returned_texts) and idx < len(original_texts)
211 }
212 return [replacements.get(idx, original) for idx, original in enumerate(original_texts)]
215def _apply_message_redaction(original: AllMessageValues, redacted: _Message) -> AllMessageValues:
216 content: Final = original.get("content")
217 if isinstance(content, str):
218 return cast(AllMessageValues, {**original, "content": _extract_text_from_message(redacted)})
219 if isinstance(content, list) and _extract_text_from_content(content):
220 redacted_content: Final = redacted.content
221 new_content: Final = (
222 [part.model_dump() for part in redacted_content] if isinstance(redacted_content, list) else redacted_content
223 )
224 return cast(AllMessageValues, {**original, "content": new_content})
225 return original
228def _redacted_messages(
229 processed_messages: list[AllMessageValues],
230 guard_output: _GuardInput,
231 sent_indices: tuple[int, ...],
232 full_messages: list[AllMessageValues],
233) -> list[AllMessageValues] | None:
234 redactions: Final = {
235 id(processed_messages[idx]): _apply_message_redaction(processed_messages[idx], guard_output.messages[pos])
236 for pos, idx in enumerate(sent_indices)
237 if pos < len(guard_output.messages) and idx < len(processed_messages)
238 }
239 if not redactions.keys() <= {id(message) for message in full_messages}:
240 return None
241 return [redactions.get(id(message), message) for message in full_messages]
244class CrowdStrikeAIDRHandler(CustomGuardrail):
245 """
246 CrowdStrike AIDR AI Guardrail handler to interact with the CrowdStrike AIDR
247 AI Guard service.
248 """
250 @classmethod
251 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
252 return [
253 GuardrailEventHooks.pre_call,
254 GuardrailEventHooks.post_call,
255 ]
257 def __init__(
258 self,
259 guardrail_name: str,
260 api_key: str | None = None,
261 api_base: str | None = None,
262 fail_on_error: bool | None = True,
263 streaming_buffer_until_moderated: bool | None = None,
264 streaming_buffer_release_on_scan: bool | None = None,
265 streaming_end_of_stream_only: bool | None = None,
266 streaming_sampling_rate: int | None = None,
267 async_handler: AsyncHTTPHandler | None = None,
268 **kwargs,
269 ) -> None:
270 """
271 Initializes the CrowdStrikeAIDRHandler.
273 Args:
274 guardrail_name (str): The name of the guardrail instance.
275 api_key (str | None): The CrowdStrike AIDR API key. Reads from CS_AIDR_TOKEN env var if None.
276 api_base (str | None): The CrowdStrike AIDR API base URL. Reads from CS_AIDR_BASE_URL env var if None.
277 streaming_end_of_stream_only (bool | None): Scan streamed output once at end of stream instead of
278 every streaming_sampling_rate chunks. Defaults to False.
279 streaming_sampling_rate (int | None): Scan the accumulated streamed output every Nth chunk. Defaults to 5.
280 async_handler (AsyncHTTPHandler | None): HTTP client to call AI Guard with. Defaults to the shared
281 guardrail-callback client.
282 **kwargs: Additional arguments passed to the CustomGuardrail base class.
283 """
284 self.async_handler = async_handler or get_async_httpx_client(
285 llm_provider=httpxSpecialProvider.GuardrailCallback
286 )
287 self.fail_on_error = True if fail_on_error is None else fail_on_error
288 self._set_streaming_params(
289 CrowdStrikeAIDRGuardrailConfigModelOptionalParams(
290 streaming_end_of_stream_only=streaming_end_of_stream_only,
291 streaming_sampling_rate=streaming_sampling_rate,
292 streaming_buffer_until_moderated=streaming_buffer_until_moderated,
293 streaming_buffer_release_on_scan=streaming_buffer_release_on_scan,
294 )
295 )
297 self.api_key = api_key or os.environ.get("CS_AIDR_TOKEN")
298 if not self.api_key:
299 raise CrowdStrikeAIDRGuardrailMissingSecrets(
300 "CrowdStrike AIDR API Key not found. Set CS_AIDR_TOKEN environment variable or pass it in litellm_params."
301 )
303 self.api_base = api_base or os.environ.get("CS_AIDR_BASE_URL")
304 if not self.api_base:
305 raise CrowdStrikeAIDRGuardrailMissingSecrets(
306 "CrowdStrike AIDR API base URL is required. Set CS_AIDR_BASE_URL environment variable or pass it in litellm_params."
307 )
309 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
310 # Pass relevant kwargs to the parent class
311 super().__init__(guardrail_name=guardrail_name, **kwargs)
312 verbose_proxy_logger.debug(
313 "Initialized CrowdStrike AIDR Guardrail: name=%s, api_base=%s", guardrail_name, self.api_base
314 )
316 def _set_streaming_params(self, streaming_params: CrowdStrikeAIDRGuardrailConfigModelOptionalParams) -> None:
317 self.streaming_buffer_until_moderated: bool = streaming_params.streaming_buffer_until_moderated or False
318 self.streaming_buffer_release_on_scan: bool = streaming_params.streaming_buffer_release_on_scan or False
319 self.streaming_end_of_stream_only: bool = streaming_params.streaming_end_of_stream_only or False
320 self.streaming_sampling_rate: int = streaming_params.streaming_sampling_rate or 5
322 @override
323 def update_in_memory_litellm_params(self, litellm_params: LitellmParams) -> None:
324 super().update_in_memory_litellm_params(litellm_params)
325 self._set_streaming_params(streaming_params_from_litellm_params(litellm_params))
327 async def _call_crowdstrike_aidr_guard(
328 self, payload: dict[str, object], hook_name: str
329 ) -> _GuardChatCompletionsResult:
330 """
331 Makes the API call to the CrowdStrike AIDR AI Guard endpoint.
332 The function itself will raise an error if a response should be blocked,
333 but otherwise will return a list of redacted messages that the caller
334 should act on.
336 Args:
337 payload (dict): The request payload.
338 hook_name (str): Name of the hook calling this function (for logging).
340 Raises:
341 HTTPException: If the CrowdStrike AIDR API returns a 'blocked: true' response.
342 Exception: For other API call failures.
344 Returns:
345 The parsed `result` body of the API response.
346 """
347 endpoint: Final = f"{self.api_base}/v1/guard_chat_completions"
349 headers: Final = {
350 "Authorization": f"Bearer {self.api_key}",
351 "Content-Type": "application/json",
352 }
354 verbose_proxy_logger.debug(
355 "CrowdStrike AIDR Guardrail (%s): Calling endpoint %s with payload: %s", hook_name, endpoint, payload
356 )
358 response: Final = await self.async_handler.post(url=endpoint, json=payload, headers=headers)
359 assert response is not None
360 response.raise_for_status()
362 response_body: Final[object] = response.json()
363 raw_result: Final[object] = response_body.get("result") if isinstance(response_body, dict) else None
364 blocked_signal: Final[object] = raw_result.get("blocked") if isinstance(raw_result, dict) else None
366 if blocked_signal:
367 verbose_proxy_logger.warning(
368 "CrowdStrike AIDR Guardrail (%s): Request blocked. Verdict: %s", hook_name, blocked_signal
369 )
370 raise HTTPException(
371 status_code=400, # Bad Request, indicating violation
372 detail={
373 "error": "Violated CrowdStrike AIDR guardrail policy",
374 "guardrail_name": self.guardrail_name,
375 },
376 )
378 try:
379 result: Final = (
380 _GuardChatCompletionsResponse.model_validate(response_body).result or _GuardChatCompletionsResult()
381 )
382 except ValidationError as validation_error:
383 transformed_signal: Final[object] = raw_result.get("transformed") if isinstance(raw_result, dict) else None
384 if transformed_signal:
385 raise HTTPException(
386 status_code=500,
387 detail={ # mutable-ok: one-shot HTTPException detail payload, never mutated after construction
388 "error": "CrowdStrike AIDR returned a transformed response litellm could not parse; "
389 "failing closed instead of dropping the delivered redactions",
390 "guardrail_name": self.guardrail_name,
391 },
392 ) from validation_error
393 raise
394 verbose_proxy_logger.debug(
395 "CrowdStrike AIDR Guardrail (%s): Request passed. Response: %s", hook_name, result.detectors
396 )
398 return result
400 def _build_guard_input_for_request(self, inputs: GenericGuardrailAPIInputs) -> _GuardInputWithIndices | None:
401 guard_input: Final = _GuardInput(messages=[], tools=[])
402 structured_messages: Final = inputs.get("structured_messages")
403 texts: Final = inputs.get("texts", [])
404 tools: Final = inputs.get("tools")
406 if structured_messages:
407 filtered: Final = _messages_since_last_assistant(structured_messages)
408 for message in filtered.messages:
409 content = _normalize_content(message.get("content"))
410 if content is None or len(content) == 0:
411 content = ""
412 guard_input.messages.append(_Message(role=message["role"], content=content))
413 indices = filtered.indices
414 elif texts:
415 guard_input.messages = [_Message(role="user", content=text) for text in texts]
416 indices = tuple(range(len(texts)))
417 else:
418 verbose_proxy_logger.warning("CrowdStrike AIDR Guardrail: No messages or texts provided for input request")
419 return None
421 if tools:
422 guard_input.tools = tools
424 return _GuardInputWithIndices(guard_input, indices)
426 def _build_guard_input_for_response(self, inputs: GenericGuardrailAPIInputs) -> _GuardInput:
427 output_texts: Final[list[str]] = inputs.get("texts", [])
428 return _GuardInput(
429 messages=[_Message(role="assistant", content=text) for text in output_texts],
430 tools=inputs.get("tools", []),
431 )
433 def _extract_transformed_texts(self, guard_output: _GuardInput, num_assistant_messages: int) -> list[str]:
434 tail: Final = guard_output.messages[-num_assistant_messages:] if num_assistant_messages > 0 else []
435 return [_extract_text_from_message(msg) for msg in tail]
437 async def _call_or_fail_open(
438 self, payload: dict[str, object], hook_name: str, request_data: dict[str, object]
439 ) -> _GuardChatCompletionsResult:
440 start_time: Final = time.time()
441 try:
442 return await self._call_crowdstrike_aidr_guard(payload, hook_name)
443 except HTTPException:
444 raise
445 except Exception as error:
446 if self.fail_on_error:
447 raise
448 verbose_proxy_logger.error(
449 "CrowdStrike AIDR Guardrail failed open | hook_name: %s error: %s",
450 hook_name,
451 error,
452 exc_info=True,
453 )
454 end_time: Final = time.time()
455 self.add_standard_logging_guardrail_information_to_request_data(
456 guardrail_json_response=error,
457 request_data=request_data,
458 guardrail_status="guardrail_failed_to_respond",
459 start_time=start_time,
460 end_time=end_time,
461 duration=end_time - start_time,
462 )
463 return _GuardChatCompletionsResult()
465 @override
466 def structured_messages_cover_full_request(self) -> bool:
467 return effective_skip_system_message_for_guardrail(self) or effective_skip_tool_message_for_guardrail(self)
469 def _writeback_messages(
470 self,
471 structured_messages: list[AllMessageValues],
472 guard_output: _GuardInput,
473 sent_indices: tuple[int, ...],
474 request_data: dict[str, object],
475 ) -> list[AllMessageValues] | None:
476 if effective_skip_system_message_for_guardrail(self) or effective_skip_tool_message_for_guardrail(self):
477 request_messages: Final = request_data.get("messages")
478 full_messages = (
479 cast("list[AllMessageValues]", request_messages)
480 if isinstance(request_messages, list)
481 else structured_messages
482 )
483 else:
484 full_messages = structured_messages
485 return _redacted_messages(structured_messages, guard_output, sent_indices, full_messages)
487 @log_guardrail_information
488 @override
489 async def apply_guardrail(
490 self,
491 inputs: GenericGuardrailAPIInputs,
492 request_data: dict,
493 input_type: Literal["request", "response"],
494 logging_obj: Optional["LiteLLMLoggingObj"] = None,
495 ) -> GenericGuardrailAPIInputs:
496 verbose_proxy_logger.debug("CrowdStrike AIDR Guardrail: Applying guardrail to %s", input_type)
498 # Extract inputs
499 texts: Final = inputs.get("texts", [])
500 structured_messages: Final = inputs.get("structured_messages")
501 tools: Final = inputs.get("tools")
502 tool_calls: Final = inputs.get("tool_calls")
504 # Build guard_input based on input_type
505 sent_indices: tuple[int, ...] = ()
506 if input_type == "request":
507 request_result: Final = self._build_guard_input_for_request(inputs)
508 if request_result is None:
509 return inputs
510 guard_input = request_result.guard_input
511 sent_indices = request_result.sent_indices
512 event_type = "input"
513 hook_name = "apply_guardrail (request)"
514 else:
515 guard_input = self._build_guard_input_for_response(inputs)
516 if len(guard_input.messages) == 0:
517 return inputs
518 event_type = "output"
519 hook_name = "apply_guardrail (response)"
521 ai_guard_payload: Final[dict[str, object]] = {
522 "guard_input": guard_input.model_dump(mode="json"),
523 "event_type": event_type,
524 }
526 model: Final = inputs.get("model")
527 if model:
528 ai_guard_payload["model"] = model
530 metadata: Final = _merge_metadata_bags(request_data)
531 if metadata is not None:
532 user_id: Final = metadata.get("user_api_key_user_id")
533 if user_id:
534 ai_guard_payload["user_id"] = user_id
536 extra_info: Final[dict[str, object]] = {}
537 user_email: Final = metadata.get("user_api_key_user_email")
538 if user_email:
539 extra_info["user_name"] = user_email
540 ai_guard_payload["extra_info"] = extra_info
542 result: Final = await self._call_or_fail_open(ai_guard_payload, hook_name, request_data)
544 if "body" in request_data or "messages" in request_data:
545 add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name)
547 if not result.transformed or result.guard_output is None:
548 return inputs
550 guard_output: Final = result.guard_output
552 if input_type == "request":
553 transformed_texts = _merge_request_transforms(guard_output, structured_messages, texts, sent_indices)
554 else:
555 transformed_texts = self._extract_transformed_texts(guard_output, len(texts))
557 result_inputs: Final[GenericGuardrailAPIInputs] = {"texts": transformed_texts}
558 if tools:
559 result_inputs["tools"] = tools
560 if tool_calls:
561 result_inputs["tool_calls"] = tool_calls
562 if structured_messages:
563 rebuilt: Final = (
564 self._writeback_messages(structured_messages, guard_output, sent_indices, request_data)
565 if input_type == "request"
566 else None
567 )
568 result_inputs["structured_messages"] = rebuilt if rebuilt is not None else structured_messages
570 return result_inputs
572 @override
573 @staticmethod
574 def get_config_model() -> type["GuardrailConfigModel"] | None:
575 from litellm.types.proxy.guardrails.guardrail_hooks.crowdstrike_aidr import (
576 CrowdStrikeAIDRGuardrailConfigModel,
577 )
579 return CrowdStrikeAIDRGuardrailConfigModel