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

347 statements  

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

1from __future__ import annotations 

2 

3from collections.abc import AsyncGenerator 

4from datetime import datetime 

5from typing import Final, Literal, TypeGuard 

6 

7from fastapi import HTTPException 

8from httpx import HTTPError 

9from httpx import Response as HttpxResponse 

10from pydantic import BaseModel, TypeAdapter, ValidationError 

11 

12import litellm 

13from litellm._logging import verbose_proxy_logger 

14from litellm.integrations.custom_guardrail import CustomGuardrail 

15from litellm.llms.custom_httpx.http_handler import ( 

16 get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] 

17 httpxSpecialProvider, 

18) 

19from litellm.proxy._types import UserAPIKeyAuth 

20from litellm.proxy.common_utils.callback_utils import ( 

21 add_guardrail_to_applied_guardrails_header, # pyright: ignore[reportUnknownVariableType] 

22) 

23from litellm.proxy.guardrails._content_utils import build_inspection_messages 

24from litellm.secret_managers.main import get_secret_str 

25from litellm.types.guardrails import GuardrailEventHooks, Mode 

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

27from litellm.types.proxy.guardrails.guardrail_hooks.repelloai import ( 

28 RepelloAIAnalyzeResponse, 

29) 

30from litellm.types.utils import ( 

31 CallTypesLiteral, 

32 GuardrailStatus, 

33 LLMResponseTypes, 

34 ModelResponse, 

35 ModelResponseStream, 

36) 

37 

38DEFAULT_REPELLOAI_API_BASE: Final = "https://argusapi.repello.ai/sdk/v1" 

39DEFAULT_REPELLOAI_TIMEOUT: Final = 30.0 

40BLOCKED_VERDICT: Final = "blocked" 

41FLAGGED_VERDICT: Final = "flagged" 

42PASSED_VERDICT: Final = "passed" 

43 

44# Argus returns these for a permanently broken guardrail (bad key, unknown 

45# asset_id, malformed payload), not a transient outage. They must always 

46# block, never honour fail_open. 

47CONFIG_ERROR_STATUS_CODES: Final = frozenset({400, 401, 403, 404, 422}) 

48_SCHEMA_SCALAR_KEYS: Final = frozenset(("name", "description", "title", "const", "default")) 

49_SCHEMA_LIST_KEYS: Final = frozenset(("enum", "examples")) 

50_SCHEMA_EXTRACTED_KEYS: Final = _SCHEMA_SCALAR_KEYS | _SCHEMA_LIST_KEYS 

51 

52 

53class RepelloAIGuardrailMissingSecrets(Exception): 

54 pass 

55 

56 

57def _is_object_dict(value: object) -> TypeGuard[dict[str, object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip 

58 return isinstance(value, dict) 

59 

60 

61def _is_object_list(value: object) -> TypeGuard[list[object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip 

62 return isinstance(value, list) 

63 

64 

65class RepelloAIGuardrail(CustomGuardrail): 

66 @classmethod 

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

68 return [ 

69 GuardrailEventHooks.pre_call, 

70 GuardrailEventHooks.post_call, 

71 ] 

72 

73 @staticmethod 

74 def _get_field(obj: object, key: str) -> object: 

75 if _is_object_dict(obj): 

76 return obj.get(key) 

77 return getattr(obj, key, None) 

78 

79 @classmethod 

80 def _extract_tool_call_args_from_message(cls, message: object) -> list[str]: 

81 args: Final[list[str]] = [] 

82 

83 tool_calls: Final = cls._get_field(message, "tool_calls") 

84 if _is_object_list(tool_calls): 

85 for tool_call in tool_calls: 

86 function = cls._get_field(tool_call, "function") 

87 arguments = cls._get_field(function, "arguments") 

88 if isinstance(arguments, str) and arguments.strip(): 

89 args.append(arguments) 

90 

91 function_call: Final = cls._get_field(message, "function_call") 

92 arguments = cls._get_field(function_call, "arguments") 

93 if isinstance(arguments, str) and arguments.strip(): 

94 args.append(arguments) 

95 

96 return args 

97 

98 @staticmethod 

99 def _iter_schema_text(node: object) -> list[str]: 

100 texts: Final[list[str]] = [] 

101 stack: Final[list[object]] = [node] 

102 

103 while stack: 

104 current = stack.pop() 

105 if _is_object_dict(current): 

106 for key in _SCHEMA_SCALAR_KEYS: 

107 value = current.get(key) 

108 if isinstance(value, str) and value: 

109 texts.append(value) 

110 for key in _SCHEMA_LIST_KEYS: 

111 items = current.get(key) 

112 if _is_object_list(items): 

113 for item in items: 

114 if isinstance(item, str) and item: 

115 texts.append(item) 

116 remaining: list[object] = [v for k, v in current.items() if k not in _SCHEMA_EXTRACTED_KEYS] 

117 stack.extend(reversed(remaining)) 

118 elif _is_object_list(current): 

119 stack.extend(reversed(current)) 

120 

121 return texts 

122 

123 @classmethod 

124 def _extract_tool_definition_text(cls, data: dict[str, object]) -> list[str]: 

125 texts: Final[list[str]] = [] 

126 

127 tools: Final = data.get("tools") 

128 for tool in tools if _is_object_list(tools) else []: 

129 if not _is_object_dict(tool): 

130 continue 

131 function = tool.get("function") 

132 if _is_object_dict(function): 

133 texts.extend(cls._iter_schema_text(function)) 

134 

135 functions: Final = data.get("functions") 

136 for function in functions if _is_object_list(functions) else []: 

137 if _is_object_dict(function): 

138 texts.extend(cls._iter_schema_text(function)) 

139 

140 return texts 

141 

142 def __init__( 

143 self, 

144 api_key: str | None = None, 

145 api_base: str | None = None, 

146 asset_id: str | None = None, 

147 unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed", 

148 guardrail_name: str | None = None, 

149 event_hook: (GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None) = None, 

150 default_on: bool = False, 

151 ): 

152 self.repelloai_api_key = api_key or get_secret_str("ARGUS_API_KEY") or get_secret_str("REPELLOAI_API_KEY") or "" 

153 if not self.repelloai_api_key: 

154 raise RepelloAIGuardrailMissingSecrets( 

155 "Couldn't get Repello API key. Set `ARGUS_API_KEY` in the environment " 

156 "or pass `api_key` to the guardrail in the config file." 

157 ) 

158 

159 self.asset_id = asset_id 

160 if not self.asset_id: 

161 raise ValueError( 

162 "Repello guardrail requires an `asset_id`. Create an asset in the Repello " 

163 "dashboard and set `asset_id` on the guardrail in the config file." 

164 ) 

165 

166 self.api_base = api_base or get_secret_str("REPELLOAI_API_BASE") or DEFAULT_REPELLOAI_API_BASE 

167 self.unreachable_fallback: Literal["fail_closed", "fail_open"] = ( 

168 "fail_open" if unreachable_fallback == "fail_open" else "fail_closed" 

169 ) 

170 self.async_handler = get_async_httpx_client( 

171 llm_provider=httpxSpecialProvider.GuardrailCallback, 

172 params={"timeout": DEFAULT_REPELLOAI_TIMEOUT}, 

173 ) 

174 super().__init__( # pyright: ignore[reportUnknownMemberType] 

175 guardrail_name=guardrail_name, 

176 event_hook=event_hook, 

177 default_on=default_on, 

178 supported_event_hooks=list(self.get_supported_event_hooks()), 

179 ) 

180 

181 async def _call_analyze( 

182 self, 

183 text: str, 

184 stage: Literal["prompt", "response"], 

185 request_data: dict[str, object], 

186 event_type: GuardrailEventHooks, 

187 ) -> RepelloAIAnalyzeResponse | None: 

188 endpoint: Final = f"{self.api_base}/analyze/{stage}" 

189 request: Final[dict[str, object]] = { 

190 "asset_id": self.asset_id or "", 

191 "scan_data": {stage: text}, 

192 } 

193 

194 status: GuardrailStatus = "success" 

195 guardrail_json_response: str | dict[str, object] | list[dict[str, object]] = "" 

196 start_time: Final[datetime] = datetime.now() 

197 repelloai_response: RepelloAIAnalyzeResponse | None = None 

198 try: 

199 verbose_proxy_logger.debug("RepelloAI Argus request: %s", request) 

200 response: Final[HttpxResponse] = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post is untyped 

201 url=endpoint, 

202 headers={"X-API-Key": self.repelloai_api_key}, 

203 json=request, 

204 ) 

205 self._raise_for_config_error(response) 

206 response.raise_for_status() 

207 try: 

208 repelloai_response = TypeAdapter(RepelloAIAnalyzeResponse).validate_json(response.text) 

209 except ValidationError as e: 

210 raise HTTPException( 

211 status_code=500, 

212 detail={ 

213 "error": "RepelloAI Argus guardrail returned invalid JSON", 

214 "status_code": response.status_code, 

215 }, 

216 ) from e 

217 verbose_proxy_logger.debug("RepelloAI Argus response: %s", repelloai_response) 

218 if self._verdict_blocks(repelloai_response): 

219 status = "guardrail_intervened" 

220 return repelloai_response 

221 except HTTPException as e: 

222 status = "guardrail_failed_to_respond" 

223 guardrail_json_response = str(e.detail) if not isinstance(e.detail, (dict, list)) else e.detail 

224 raise 

225 except HTTPError as e: 

226 status = "guardrail_failed_to_respond" 

227 guardrail_json_response = str(e) 

228 return self._handle_unreachable(e) 

229 except Exception as e: 

230 status = "guardrail_failed_to_respond" 

231 guardrail_json_response = str(e) 

232 raise HTTPException(status_code=500, detail={"error": "RepelloAI Argus guardrail failed"}) from e 

233 finally: 

234 end_time: Final = datetime.now() 

235 if repelloai_response is not None: 

236 guardrail_json_response = dict(repelloai_response) 

237 self.add_standard_logging_guardrail_information_to_request_data( # pyright: ignore[reportUnknownMemberType] 

238 guardrail_json_response=guardrail_json_response, 

239 guardrail_status=status, 

240 request_data=request_data, 

241 start_time=start_time.timestamp(), 

242 end_time=end_time.timestamp(), 

243 duration=(end_time - start_time).total_seconds(), 

244 masked_entity_count={}, 

245 event_type=event_type, 

246 ) 

247 

248 @staticmethod 

249 def _raise_for_config_error(response: HttpxResponse) -> None: 

250 if response.status_code in CONFIG_ERROR_STATUS_CODES: 

251 raise HTTPException( 

252 status_code=500, 

253 detail={ 

254 "error": "RepelloAI Argus guardrail is misconfigured", 

255 "status_code": response.status_code, 

256 }, 

257 ) 

258 

259 def _verdict_blocks(self, repelloai_response: RepelloAIAnalyzeResponse | None) -> bool: 

260 if repelloai_response is None: 

261 return False 

262 verdict: Final = repelloai_response.get("verdict") 

263 if verdict == BLOCKED_VERDICT: 

264 return True 

265 if verdict in (PASSED_VERDICT, FLAGGED_VERDICT): 

266 return False 

267 verbose_proxy_logger.warning( 

268 "RepelloAI Argus returned an unrecognized verdict (%s) - blocking.", 

269 verdict, 

270 ) 

271 return True 

272 

273 def _handle_unreachable(self, error: Exception) -> RepelloAIAnalyzeResponse | None: 

274 verbose_proxy_logger.warning("RepelloAI Argus unreachable: %s", str(error)) 

275 if self.unreachable_fallback == "fail_closed": 

276 raise HTTPException( 

277 status_code=500, 

278 detail={"error": "RepelloAI Argus guardrail unreachable"}, 

279 ) 

280 return None 

281 

282 def _raise_if_blocked(self, repelloai_response: RepelloAIAnalyzeResponse | None) -> None: 

283 if repelloai_response is None: 

284 return 

285 if self._verdict_blocks(repelloai_response): 

286 raise HTTPException( 

287 status_code=400, 

288 detail=self._format_blocked_detail(repelloai_response), 

289 ) 

290 self._log_flagged_verdict(repelloai_response) 

291 

292 @classmethod 

293 def _format_blocked_detail(cls, repelloai_response: RepelloAIAnalyzeResponse) -> str: 

294 policies: Final = repelloai_response.get("policies_violated") 

295 if not isinstance(policies, list) or not policies: 

296 return "Blocked by RepelloAI Argus guardrail." 

297 

298 formatted_policies: Final[list[str]] = [] 

299 for policy in policies: 

300 policy_name = policy.get("policy_name") or "unknown_policy" 

301 details: list[str] = [] 

302 action_taken = policy.get("action_taken") 

303 if action_taken: 

304 details.append(f"action: {action_taken}") 

305 policy_details = policy.get("details") 

306 if isinstance(policy_details, dict): 

307 score = policy_details.get("score") 

308 if score is not None: 

309 details.append(f"score: {score}") 

310 suffix = f" ({', '.join(details)})" if details else "" 

311 formatted_policies.append(f"{policy_name}{suffix}") 

312 

313 if not formatted_policies: 

314 return "Blocked by RepelloAI Argus guardrail." 

315 return f"Blocked by RepelloAI Argus guardrail. Policies violated: {'; '.join(formatted_policies)}." 

316 

317 @staticmethod 

318 def _log_flagged_verdict(repelloai_response: RepelloAIAnalyzeResponse) -> None: 

319 if repelloai_response.get("verdict") == FLAGGED_VERDICT: 

320 verbose_proxy_logger.warning( 

321 "RepelloAI Argus flagged content (allowed): %s", 

322 repelloai_response.get("policies_violated"), 

323 ) 

324 

325 @staticmethod 

326 def _extract_prompt_message_text(data: dict[str, object]) -> list[str]: 

327 messages: Final = build_inspection_messages(data) 

328 return [content for message in messages if isinstance(content := message.get("content"), str) and content] 

329 

330 @staticmethod 

331 def _extract_input_text_parts(content: object) -> list[str]: 

332 if not _is_object_list(content): 

333 return [] 

334 return [ 

335 text 

336 for part in content 

337 if _is_object_dict(part) and part.get("type") == "input_text" 

338 if isinstance(text := part.get("text"), str) and text 

339 ] 

340 

341 @staticmethod 

342 def _extract_prompt_field_text(data: dict[str, object]) -> list[str]: 

343 prompt: Final = data.get("prompt") 

344 if isinstance(prompt, str) and prompt: 

345 return [prompt] 

346 if _is_object_list(prompt): 

347 return [item for item in prompt if isinstance(item, str) and item] 

348 return [] 

349 

350 @classmethod 

351 def _extract_prompt_text(cls, data: dict[str, object]) -> str | None: 

352 texts: Final = cls._extract_prompt_message_text(data) 

353 texts.extend(cls._extract_prompt_field_text(data)) 

354 

355 instructions: Final = data.get("instructions") 

356 if isinstance(instructions, str) and instructions: 

357 texts.append(instructions) 

358 

359 raw_messages: Final = data.get("messages") 

360 if _is_object_list(raw_messages): 

361 for message in raw_messages: 

362 texts.extend(cls._extract_tool_call_args_from_message(message)) 

363 

364 raw_input: Final = data.get("input") 

365 if _is_object_list(raw_input): 

366 for item in raw_input: 

367 if _is_object_dict(item): 

368 if "role" not in item: 

369 continue 

370 texts.extend(cls._extract_tool_call_args_from_message(item)) 

371 texts.extend(cls._extract_input_text_parts(item.get("content"))) 

372 

373 texts.extend(cls._extract_tool_definition_text(data)) 

374 return "\n".join(text for text in texts if text) if texts else None 

375 

376 async def async_pre_call_hook( 

377 self, 

378 user_api_key_dict: UserAPIKeyAuth, 

379 cache: litellm.DualCache, 

380 data: dict[str, object], 

381 call_type: CallTypesLiteral, 

382 ) -> Exception | str | dict[str, object] | None: 

383 verbose_proxy_logger.debug("RepelloAI Argus: pre_call_hook") 

384 

385 event_type: Final = GuardrailEventHooks.pre_call 

386 if ( 

387 self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType] 

388 data=data, event_type=event_type 

389 ) 

390 is not True 

391 ): 

392 return data 

393 

394 text: Final = self._extract_prompt_text(data) 

395 if not text: 

396 verbose_proxy_logger.warning("RepelloAI Argus: no inspectable prompt text in data - skipping.") 

397 return data 

398 

399 repelloai_response: Final = await self._call_analyze( 

400 text=text, 

401 stage="prompt", 

402 request_data=data, 

403 event_type=event_type, 

404 ) 

405 self._raise_if_blocked(repelloai_response) 

406 

407 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name) 

408 return data 

409 

410 async def async_post_call_success_hook( 

411 self, 

412 data: dict[str, object], 

413 user_api_key_dict: UserAPIKeyAuth, 

414 response: LLMResponseTypes, 

415 ): 

416 verbose_proxy_logger.debug("RepelloAI Argus: post_call_success_hook") 

417 

418 event_type: Final = GuardrailEventHooks.post_call 

419 if ( 

420 self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType] 

421 data=data, event_type=event_type 

422 ) 

423 is not True 

424 ): 

425 return response 

426 

427 text: Final = self._extract_response_text(response) 

428 if not text: 

429 verbose_proxy_logger.warning("RepelloAI Argus: no inspectable response text - skipping.") 

430 return response 

431 

432 repelloai_response: Final = await self._call_analyze( 

433 text=text, 

434 stage="response", 

435 request_data=data, 

436 event_type=event_type, 

437 ) 

438 self._raise_if_blocked(repelloai_response) 

439 

440 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name) 

441 return response 

442 

443 async def async_post_call_streaming_iterator_hook( 

444 self, 

445 user_api_key_dict: UserAPIKeyAuth, 

446 response: AsyncGenerator[ModelResponseStream, None], 

447 request_data: dict[str, object], 

448 ) -> AsyncGenerator[ModelResponseStream, None]: 

449 from litellm import main as litellm_main 

450 

451 event_type: Final = GuardrailEventHooks.post_call 

452 if ( 

453 self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType] 

454 data=request_data, event_type=event_type 

455 ) 

456 is not True 

457 ): 

458 async for chunk in response: 

459 yield chunk 

460 return 

461 

462 chunks: Final[list[ModelResponseStream]] = [] 

463 async for chunk in response: 

464 chunks.append(chunk) 

465 

466 assembled = litellm_main.stream_chunk_builder( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] 

467 chunks=chunks 

468 ) 

469 text: Final = self._extract_response_text(assembled) if isinstance(assembled, ModelResponse) else None 

470 if text: 

471 repelloai_response: Final = await self._call_analyze( 

472 text=text, 

473 stage="response", 

474 request_data=request_data, 

475 event_type=event_type, 

476 ) 

477 if repelloai_response is not None: 

478 self._log_flagged_verdict(repelloai_response) 

479 if self._verdict_blocks(repelloai_response): 

480 from litellm.proxy.proxy_server import StreamingCallbackError 

481 

482 raise StreamingCallbackError("Blocked by RepelloAI Argus guardrail") 

483 add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=self.guardrail_name) 

484 else: 

485 verbose_proxy_logger.warning( 

486 "RepelloAI Argus: no inspectable text in streamed response; skipping scan. " 

487 "guardrail=%s assembled_type=%s", 

488 self.guardrail_name, 

489 type(assembled).__name__, 

490 ) 

491 

492 for chunk in chunks: 

493 yield chunk 

494 

495 @staticmethod 

496 def _extract_response_text(response: object) -> str | None: 

497 if _is_object_dict(response): 

498 response_dict = response 

499 elif isinstance(response, ModelResponse): 

500 response_dict = ( 

501 response.model_dump() # pyright: ignore[reportUnknownMemberType] 

502 ) 

503 else: 

504 output_text: Final = getattr(response, "output_text", None) 

505 if isinstance(output_text, str) and output_text: 

506 return output_text 

507 response_dict = {} 

508 

509 text: Final = RepelloAIGuardrail._extract_chat_completion_text(response_dict) 

510 if text: 

511 return text 

512 return RepelloAIGuardrail._extract_responses_api_text(response_dict) 

513 

514 @classmethod 

515 def _extract_chat_completion_text(cls, response_dict: dict[str, object]) -> str | None: 

516 choices: Final = response_dict.get("choices") 

517 if not _is_object_list(choices): 

518 return None 

519 parts: Final[list[str]] = [] 

520 for choice in choices: 

521 if not _is_object_dict(choice): 

522 continue 

523 message = choice.get("message") 

524 if _is_object_dict(message): 

525 content = message.get("content") 

526 if isinstance(content, str) and content: 

527 parts.append(content) 

528 parts.extend(cls._extract_tool_call_args_from_message(message)) 

529 text = choice.get("text") 

530 if isinstance(text, str) and text: 

531 parts.append(text) 

532 return "\n".join(parts) if parts else None 

533 

534 @staticmethod 

535 def _extract_responses_api_text(response_dict: dict[str, object]) -> str | None: 

536 output: Final = response_dict.get("output") 

537 if not _is_object_list(output): 

538 return None 

539 texts: Final[list[str]] = [] 

540 for output_item in output: 

541 if not _is_object_dict(output_item): 

542 continue 

543 item_type = output_item.get("type") 

544 if item_type == "function_call": 

545 arguments = output_item.get("arguments") 

546 if isinstance(arguments, str) and arguments: 

547 texts.append(arguments) 

548 continue 

549 if item_type != "message": 

550 continue 

551 content = output_item.get("content") 

552 if not _is_object_list(content): 

553 continue 

554 for content_item in content: 

555 if not _is_object_dict(content_item): 

556 continue 

557 if content_item.get("type") not in ("output_text", "text"): 

558 continue 

559 text = content_item.get("text") 

560 if isinstance(text, str) and text: 

561 texts.append(text) 

562 return "".join(texts) if texts else None 

563 

564 @staticmethod 

565 def get_config_model() -> type[GuardrailConfigModel[BaseModel]] | None: 

566 from litellm.types.proxy.guardrails.guardrail_hooks.repelloai import ( 

567 RepelloAIGuardrailConfigModel, 

568 ) 

569 

570 return RepelloAIGuardrailConfigModel