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

247 statements  

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

1from collections.abc import Awaitable, Mapping, Sequence 

2from json import JSONDecodeError 

3from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypeAlias, cast 

4 

5import httpx 

6from typing_extensions import ReadOnly, TypedDict, Unpack 

7 

8from litellm._logging import verbose_proxy_logger 

9from litellm.exceptions import GuardrailRaisedException 

10from litellm.exceptions import Timeout as LiteLLMTimeout 

11from litellm.integrations.custom_guardrail import ( 

12 CustomGuardrail, 

13 log_guardrail_information, 

14) 

15from litellm.llms.custom_httpx.http_handler import ( 

16 get_async_httpx_client, 

17 httpxSpecialProvider, 

18) 

19from litellm.secret_managers.main import get_secret_str 

20from litellm.types.guardrails import GuardrailEventHooks 

21from litellm.types.llms.openai import ChatCompletionToolCallChunk 

22from litellm.types.utils import ChatCompletionMessageToolCall, GenericGuardrailAPIInputs 

23 

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

25 from litellm.litellm_core_utils.litellm_logging import ( 

26 Logging as LiteLLMLoggingObj, 

27 ) 

28 from litellm.types.proxy.guardrails.guardrail_hooks.base import ( 

29 GuardrailConfigModel, 

30 ) 

31 

32 

33_ANALYZE_ENDPOINT: Final = "/v1/guard/analyze" 

34_DEFAULT_VIGIL_TIMEOUT: Final = httpx.Timeout(10.0, connect=5.0) 

35_BLOCK_REASON_MAX_CHARS: Final = 500 

36_METADATA_STRING_MAX_CHARS: Final = 500 

37_METADATA_ARRAY_MAX_ITEMS: Final = 10 

38_VALID_DECISIONS: Final = ("ALLOWED", "SANITIZED", "BLOCKED") 

39_TRANSIENT_STATUS_CODES: Final = frozenset({429, 502, 503, 504}) 

40_METADATA_ALLOWLIST: Final = ( 

41 "model", 

42 "model_group", 

43 "provider", 

44 "region", 

45 "deployment", 

46 "user", 

47 "user_id", 

48 "session_id", 

49 "conversation_id", 

50 "request_id", 

51 "tenant_id", 

52 "org_id", 

53) 

54 

55_FallbackMode: TypeAlias = Literal["fail_closed", "fail_open"] 

56_MetadataValue: TypeAlias = str | int | float | Sequence[str | int | float] 

57_ToolCalls: TypeAlias = list[ChatCompletionToolCallChunk] | list[ChatCompletionMessageToolCall] 

58 

59 

60class _AnalyzePayload(TypedDict): 

61 """Request body posted to the Vigil Guard analyze endpoint.""" 

62 

63 text: ReadOnly[str] 

64 source: ReadOnly[str] 

65 mode: ReadOnly[str] 

66 metadata: ReadOnly[Mapping[str, _MetadataValue]] 

67 

68 

69class _AnalysisView(TypedDict): 

70 """Typed read of the analyze endpoint's decoded JSON body.""" 

71 

72 analysis: ReadOnly[Mapping[str, object]] 

73 

74 

75class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object): 

76 """Base-class constructor options this guardrail forwards untouched to CustomGuardrail.""" 

77 

78 supported_event_hooks: ReadOnly[list[GuardrailEventHooks]] 

79 

80 

81class _AsyncPostHandler(Protocol): 

82 def post( 82 ↛ exitline 82 didn't return from function 'post' because

83 self, 

84 *, 

85 url: str, 

86 headers: dict[str, str], 

87 json: _AnalyzePayload, 

88 timeout: httpx.Timeout, 

89 ) -> Awaitable[httpx.Response]: ... 

90 

91 

92class VigilGuardMissingConfig(ValueError): 

93 pass 

94 

95 

96class VigilGuardGuardrail(CustomGuardrail): 

97 def __init__( 

98 self, 

99 api_base: str | None = None, 

100 api_key: str | None = None, 

101 unreachable_fallback: str | None = None, 

102 timeout: float | None = None, 

103 async_handler: _AsyncPostHandler | None = None, 

104 **kwargs: Unpack[_CustomGuardrailOptions], 

105 ) -> None: 

106 resolved_base: Final = api_base or get_secret_str("VIGIL_GUARD_URL") 

107 if not resolved_base: 

108 raise VigilGuardMissingConfig( 

109 "Vigil Guard api_base is required. Set api_base in the guardrail " 

110 "config or the VIGIL_GUARD_URL environment variable." 

111 ) 

112 self.api_base = resolved_base.rstrip("/") 

113 

114 resolved_key: Final = api_key or get_secret_str("VIGIL_GUARD_API_KEY") 

115 if not resolved_key: 

116 raise VigilGuardMissingConfig( 

117 "Vigil Guard api_key is required. Set api_key in the guardrail " 

118 "config or the VIGIL_GUARD_API_KEY environment variable." 

119 ) 

120 self.api_key = resolved_key 

121 

122 fallback: Final = (unreachable_fallback or "fail_closed").lower() 

123 self.unreachable_fallback: _FallbackMode = "fail_open" if fallback == "fail_open" else "fail_closed" 

124 

125 self.timeout: httpx.Timeout = ( 

126 _DEFAULT_VIGIL_TIMEOUT if timeout is None else httpx.Timeout(timeout, connect=min(timeout, 5.0)) 

127 ) 

128 

129 self.async_handler: _AsyncPostHandler = async_handler or get_async_httpx_client( 

130 llm_provider=httpxSpecialProvider.GuardrailCallback, 

131 ) 

132 

133 forwarded: Final[_CustomGuardrailOptions] = { 

134 "supported_event_hooks": list(self.get_supported_event_hooks()), 

135 **kwargs, 

136 } 

137 

138 super().__init__(**forwarded) 

139 

140 @staticmethod 

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

142 from litellm.types.proxy.guardrails.guardrail_hooks.vigil_guard import ( 

143 VigilGuardGuardrailConfigModel, 

144 ) 

145 

146 return VigilGuardGuardrailConfigModel 

147 

148 @classmethod 

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

150 return [ 

151 GuardrailEventHooks.pre_call, 

152 GuardrailEventHooks.post_call, 

153 ] 

154 

155 @log_guardrail_information 

156 async def apply_guardrail( 

157 self, 

158 inputs: GenericGuardrailAPIInputs, 

159 request_data: dict, 

160 input_type: Literal["request", "response"], 

161 logging_obj: Optional["LiteLLMLoggingObj"] = None, 

162 ) -> GenericGuardrailAPIInputs: 

163 texts: Final = inputs.get("texts") or [] 

164 has_text: Final = any(isinstance(text, str) and text.strip() for text in texts) 

165 tool_call_args: Final = self._tool_call_arguments(inputs.get("tool_calls")) if input_type == "response" else [] 

166 if not has_text and not tool_call_args: 

167 return inputs 

168 

169 source: Final = "user_input" if input_type == "request" else "model_output" 

170 metadata: Final = self._collect_metadata(request_data, logging_obj) 

171 

172 result_texts: Final[list[str]] = [] 

173 for index, text in enumerate(texts): 

174 if not isinstance(text, str) or not text.strip(): 

175 result_texts.append(text) 

176 continue 

177 

178 try: 

179 analysis = await self._analyze(text=text, source=source, metadata=metadata) 

180 except ( 

181 httpx.HTTPError, 

182 LiteLLMTimeout, 

183 JSONDecodeError, 

184 OSError, 

185 ) as exc: 

186 return self._handle_backend_failure( 

187 exc, 

188 inputs, 

189 source, 

190 result_texts + list(texts[index:]), 

191 inputs.get("tool_calls"), 

192 ) 

193 

194 decision = analysis.get("decision") if isinstance(analysis, dict) else None 

195 if decision not in _VALID_DECISIONS: 

196 verbose_proxy_logger.error( 

197 "Vigil Guard unrecognized decision for guardrail_name=%s source=%s: %r", 

198 self.guardrail_name, 

199 source, 

200 decision, 

201 ) 

202 if self.unreachable_fallback == "fail_open": 

203 return self._build_output( 

204 inputs, 

205 result_texts + list(texts[index:]), 

206 inputs.get("tool_calls"), 

207 ) 

208 raise GuardrailRaisedException( 

209 guardrail_name=self.guardrail_name, 

210 message="Vigil Guard returned an unrecognized decision.", 

211 should_wrap_with_default_message=False, 

212 ) 

213 

214 if decision == "BLOCKED": 

215 raise GuardrailRaisedException( 

216 guardrail_name=self.guardrail_name, 

217 message=self._build_block_reason(analysis), 

218 should_wrap_with_default_message=False, 

219 blocked_content=True, 

220 ) 

221 

222 if decision == "SANITIZED": 

223 result_texts.append(self._resolve_sanitized_text(text, analysis)) 

224 else: 

225 result_texts.append(text) 

226 

227 result_tool_calls = inputs.get("tool_calls") 

228 for tc_index, arguments in tool_call_args: 

229 try: 

230 analysis = await self._analyze(text=arguments, source=source, metadata=metadata) 

231 except ( 

232 httpx.HTTPError, 

233 LiteLLMTimeout, 

234 JSONDecodeError, 

235 OSError, 

236 ) as exc: 

237 return self._handle_backend_failure(exc, inputs, source, result_texts, result_tool_calls) 

238 

239 decision = analysis.get("decision") if isinstance(analysis, dict) else None 

240 if decision not in _VALID_DECISIONS: 

241 verbose_proxy_logger.error( 

242 "Vigil Guard unrecognized decision for guardrail_name=%s source=%s: %r", 

243 self.guardrail_name, 

244 source, 

245 decision, 

246 ) 

247 if self.unreachable_fallback == "fail_open": 

248 return self._build_output(inputs, result_texts, result_tool_calls) 

249 raise GuardrailRaisedException( 

250 guardrail_name=self.guardrail_name, 

251 message="Vigil Guard returned an unrecognized decision.", 

252 should_wrap_with_default_message=False, 

253 ) 

254 

255 if decision == "BLOCKED": 

256 raise GuardrailRaisedException( 

257 guardrail_name=self.guardrail_name, 

258 message=self._build_block_reason(analysis), 

259 should_wrap_with_default_message=False, 

260 blocked_content=True, 

261 ) 

262 

263 if decision == "SANITIZED": 

264 result_tool_calls = self._set_tool_call_arguments( 

265 result_tool_calls, 

266 tc_index, 

267 self._resolve_sanitized_text(arguments, analysis), 

268 ) 

269 

270 return self._build_output(inputs, result_texts, result_tool_calls) 

271 

272 def _handle_backend_failure( 

273 self, 

274 exc: Exception, 

275 inputs: GenericGuardrailAPIInputs, 

276 source: str, 

277 final_texts: list[str], 

278 final_tool_calls: _ToolCalls | None, 

279 ) -> GenericGuardrailAPIInputs: 

280 if self.unreachable_fallback == "fail_open": 

281 verbose_proxy_logger.error( 

282 "Vigil Guard backend failure with fail_open; allowing request " 

283 "unscanned. guardrail_name=%s source=%s error=%s", 

284 self.guardrail_name, 

285 source, 

286 str(exc), 

287 ) 

288 return self._build_output(inputs, final_texts, final_tool_calls) 

289 verbose_proxy_logger.error( 

290 "Vigil Guard backend failure with fail_closed; blocking request. guardrail_name=%s source=%s error=%s", 

291 self.guardrail_name, 

292 source, 

293 str(exc), 

294 ) 

295 raise GuardrailRaisedException( 

296 guardrail_name=self.guardrail_name, 

297 message="Vigil Guard backend unreachable; request blocked by fail_closed policy.", 

298 should_wrap_with_default_message=False, 

299 ) from exc 

300 

301 @staticmethod 

302 def _build_output( 

303 inputs: GenericGuardrailAPIInputs, 

304 final_texts: list[str], 

305 final_tool_calls: Any, 

306 ) -> GenericGuardrailAPIInputs: 

307 # When nothing was changed, return the input shape verbatim so the guardrail 

308 # logs "allow" rather than "mask". When a text or a tool-call argument was 

309 # changed (sanitized), return only the remap-relevant keys and drop 

310 # structured_messages so a stale, unsanitized payload cannot reach the model. 

311 texts_changed: Final = final_texts != (inputs.get("texts") or []) 

312 tool_calls_changed: Final = final_tool_calls != inputs.get("tool_calls") 

313 if not texts_changed and not tool_calls_changed: 

314 return cast(GenericGuardrailAPIInputs, dict(inputs)) 

315 guardrailed: Final[GenericGuardrailAPIInputs] = {"texts": final_texts} 

316 if "images" in inputs: 

317 guardrailed["images"] = inputs["images"] 

318 if "tools" in inputs: 

319 guardrailed["tools"] = inputs["tools"] 

320 if tool_calls_changed: 

321 guardrailed["tool_calls"] = final_tool_calls 

322 return guardrailed 

323 

324 @staticmethod 

325 def _tool_call_arguments(tool_calls: Sequence[object] | None) -> list[tuple[int, str]]: 

326 pairs: Final[list[tuple[int, str]]] = [] 

327 if isinstance(tool_calls, list): 

328 for index, tool_call in enumerate(tool_calls): 

329 function = tool_call.get("function") if isinstance(tool_call, dict) else None 

330 arguments = function.get("arguments") if isinstance(function, dict) else None 

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

332 pairs.append((index, arguments)) 

333 return pairs 

334 

335 @staticmethod 

336 def _set_tool_call_arguments(tool_calls: Any, index: int, arguments: str) -> list[Any]: 

337 updated: Final = list(tool_calls) 

338 tool_call: Final = dict(updated[index]) 

339 function: Final = dict(tool_call.get("function") or {}) 

340 function["arguments"] = arguments 

341 tool_call["function"] = function 

342 updated[index] = tool_call 

343 return updated 

344 

345 async def _analyze(self, text: str, source: str, metadata: Mapping[str, _MetadataValue]) -> Mapping[str, object]: 

346 payload: Final[_AnalyzePayload] = { 

347 "text": text, 

348 "source": source, 

349 "mode": "full", 

350 "metadata": metadata, 

351 } 

352 endpoint: Final = f"{self.api_base}{_ANALYZE_ENDPOINT}" 

353 headers: Final = { 

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

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

356 } 

357 response: Final = await self._post_with_retry(endpoint, headers, payload) 

358 decoded: Final[_AnalysisView] = {"analysis": response.json()} 

359 return decoded["analysis"] 

360 

361 async def _post_with_retry( 

362 self, endpoint: str, headers: dict[str, str], payload: _AnalyzePayload 

363 ) -> httpx.Response: 

364 for attempt in range(2): 

365 try: 

366 response = await self.async_handler.post( 

367 url=endpoint, 

368 headers=headers, 

369 json=payload, 

370 timeout=self.timeout, 

371 ) 

372 response.raise_for_status() 

373 return response 

374 except Exception as exc: 

375 if attempt == 0 and self._is_transient(exc): 

376 verbose_proxy_logger.debug( 

377 "Vigil Guard transient failure; retrying once: %s", 

378 type(exc).__name__, 

379 ) 

380 continue 

381 raise 

382 raise AssertionError("unreachable") # pragma: no cover 

383 

384 @staticmethod 

385 def _is_transient(exc: Exception) -> bool: 

386 if isinstance(exc, httpx.HTTPStatusError): 

387 return exc.response.status_code in _TRANSIENT_STATUS_CODES 

388 return isinstance( 

389 exc, 

390 ( 

391 httpx.ConnectError, 

392 httpx.ConnectTimeout, 

393 httpx.ReadTimeout, 

394 httpx.RemoteProtocolError, 

395 LiteLLMTimeout, 

396 ), 

397 ) 

398 

399 @staticmethod 

400 def _build_block_reason(analysis: Mapping[str, object]) -> str: 

401 for key in ("blockMessage", "decisionReason"): 

402 value = analysis.get(key) 

403 if isinstance(value, str) and value.strip(): 

404 return value.strip()[:_BLOCK_REASON_MAX_CHARS] 

405 categories: Final = analysis.get("categories") 

406 if isinstance(categories, list): 

407 names: Final = [c for c in categories if isinstance(c, str) and c.strip()] 

408 if names: 

409 return ", ".join(names)[:_BLOCK_REASON_MAX_CHARS] 

410 return "Blocked by policy" 

411 

412 @staticmethod 

413 def _resolve_sanitized_text(original: str, analysis: Mapping[str, object]) -> str: 

414 for key in ("sanitizedText", "outputText"): 

415 value = analysis.get(key) 

416 if isinstance(value, str): 

417 return value 

418 return original 

419 

420 def _collect_metadata( 

421 self, request_data: dict, logging_obj: Optional["LiteLLMLoggingObj"] 

422 ) -> Mapping[str, _MetadataValue]: 

423 sources: Final[list[dict]] = [] 

424 if isinstance(request_data, dict): 

425 sources.append(request_data) 

426 for nested_key in ("metadata", "litellm_metadata"): 

427 nested = request_data.get(nested_key) 

428 if isinstance(nested, dict): 

429 sources.append(nested) 

430 

431 collected: Final[dict[str, _MetadataValue]] = {} 

432 for field in _METADATA_ALLOWLIST: 

433 for source in sources: 

434 if field in source and source[field] is not None: 

435 clamped = self._clamp_metadata_value(source[field]) 

436 if clamped is not None: 

437 collected[field] = clamped 

438 break 

439 

440 call_id: Final = self._extract_call_id(request_data, logging_obj) 

441 if call_id: 

442 collected["litellm_call_id"] = call_id 

443 

444 return collected 

445 

446 @staticmethod 

447 def _clamp_metadata_value(value: object) -> _MetadataValue | None: 

448 if isinstance(value, bool): 

449 return None 

450 if isinstance(value, str): 

451 return value[:_METADATA_STRING_MAX_CHARS] 

452 if isinstance(value, (int, float)): 

453 return value 

454 if isinstance(value, list): 

455 clamped: Final[list[str | int | float]] = [] 

456 for item in value[:_METADATA_ARRAY_MAX_ITEMS]: 

457 if isinstance(item, bool): 

458 continue 

459 if isinstance(item, str): 

460 clamped.append(item[:_METADATA_STRING_MAX_CHARS]) 

461 elif isinstance(item, (int, float)): 

462 clamped.append(item) 

463 return clamped or None 

464 return None 

465 

466 @staticmethod 

467 def _extract_call_id(request_data: dict, logging_obj: Optional["LiteLLMLoggingObj"]) -> str | None: 

468 if logging_obj is not None: 

469 call_id = getattr(logging_obj, "litellm_call_id", None) 

470 if isinstance(call_id, str) and call_id: 

471 return call_id 

472 if isinstance(request_data, dict): 

473 call_id = request_data.get("litellm_call_id") 

474 if isinstance(call_id, str) and call_id: 

475 return call_id 

476 metadata: Final = request_data.get("metadata") 

477 if isinstance(metadata, dict): 

478 nested: Final = metadata.get("litellm_call_id") 

479 if isinstance(nested, str) and nested: 

480 return nested 

481 return None