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

1import json 

2import os 

3import time 

4from collections.abc import Mapping, Sequence 

5from typing import TYPE_CHECKING, Annotated, Final, Literal, NamedTuple, Optional, cast 

6 

7from fastapi import HTTPException 

8from pydantic import BaseModel, ConfigDict, Field, ValidationError 

9from typing_extensions import override 

10 

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 

34 

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 

38 

39 

40class CrowdStrikeAIDRGuardrailMissingSecrets(Exception): 

41 """Custom exception for missing CrowdStrike AIDR secrets.""" 

42 

43 

44class _TextContentPart(BaseModel): 

45 model_config = ConfigDict(extra="forbid") 

46 

47 type: Literal["text"] = "text" 

48 text: str 

49 

50 

51class _ImageUrl(BaseModel): 

52 url: str 

53 

54 

55class _ImageUrlContentPart(BaseModel): 

56 model_config = ConfigDict(extra="forbid") 

57 

58 type: Literal["image_url"] = "image_url" 

59 image_url: _ImageUrl 

60 

61 

62_ContentPart = Annotated[_TextContentPart | _ImageUrlContentPart, Field(discriminator="type")] 

63 

64 

65class _Message(BaseModel): 

66 role: str 

67 content: str | list[_ContentPart] | None = None 

68 

69 

70class _GuardInput(BaseModel): 

71 messages: list[_Message] 

72 tools: Sequence[OpenAIChatCompletionToolParam] | None = None 

73 

74 

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

84 

85 

86class _GuardChatCompletionsResponse(BaseModel): 

87 result: _GuardChatCompletionsResult | None = None 

88 

89 

90class _FilteredMessages(NamedTuple): 

91 """Subset of a conversation selected for guardrail analysis.""" 

92 

93 messages: list[AllMessageValues] 

94 """Messages subset.""" 

95 indices: tuple[int, ...] 

96 """Positions of the subset's messages in the original list.""" 

97 

98 

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

104 

105 

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 

118 

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

126 

127 # Any other types are not recognized by the CrowdStrike AIDR API. 

128 

129 return parts 

130 

131 

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

139 

140 

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

148 

149 

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 

158 

159 

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 ) 

173 

174 

175def _messages_since_last_assistant( 

176 messages: Sequence[AllMessageValues], 

177) -> _FilteredMessages: 

178 if not messages: 

179 return _FilteredMessages([], ()) 

180 

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) 

184 

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 

190 

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) 

195 

196 

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

213 

214 

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 

226 

227 

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] 

242 

243 

244class CrowdStrikeAIDRHandler(CustomGuardrail): 

245 """ 

246 CrowdStrike AIDR AI Guardrail handler to interact with the CrowdStrike AIDR 

247 AI Guard service. 

248 """ 

249 

250 @classmethod 

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

252 return [ 

253 GuardrailEventHooks.pre_call, 

254 GuardrailEventHooks.post_call, 

255 ] 

256 

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. 

272 

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 ) 

296 

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 ) 

302 

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 ) 

308 

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 ) 

315 

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 

321 

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

326 

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. 

335 

336 Args: 

337 payload (dict): The request payload. 

338 hook_name (str): Name of the hook calling this function (for logging). 

339 

340 Raises: 

341 HTTPException: If the CrowdStrike AIDR API returns a 'blocked: true' response. 

342 Exception: For other API call failures. 

343 

344 Returns: 

345 The parsed `result` body of the API response. 

346 """ 

347 endpoint: Final = f"{self.api_base}/v1/guard_chat_completions" 

348 

349 headers: Final = { 

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

351 "Content-Type": "application/json", 

352 } 

353 

354 verbose_proxy_logger.debug( 

355 "CrowdStrike AIDR Guardrail (%s): Calling endpoint %s with payload: %s", hook_name, endpoint, payload 

356 ) 

357 

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

361 

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 

365 

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 ) 

377 

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 ) 

397 

398 return result 

399 

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

405 

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 

420 

421 if tools: 

422 guard_input.tools = tools 

423 

424 return _GuardInputWithIndices(guard_input, indices) 

425 

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 ) 

432 

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] 

436 

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

464 

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) 

468 

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) 

486 

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) 

497 

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

503 

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

520 

521 ai_guard_payload: Final[dict[str, object]] = { 

522 "guard_input": guard_input.model_dump(mode="json"), 

523 "event_type": event_type, 

524 } 

525 

526 model: Final = inputs.get("model") 

527 if model: 

528 ai_guard_payload["model"] = model 

529 

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 

535 

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 

541 

542 result: Final = await self._call_or_fail_open(ai_guard_payload, hook_name, request_data) 

543 

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) 

546 

547 if not result.transformed or result.guard_output is None: 

548 return inputs 

549 

550 guard_output: Final = result.guard_output 

551 

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

556 

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 

569 

570 return result_inputs 

571 

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 ) 

578 

579 return CrowdStrikeAIDRGuardrailConfigModel