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
« 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
5import httpx
6from typing_extensions import ReadOnly, TypedDict, Unpack
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
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 )
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)
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]
60class _AnalyzePayload(TypedDict):
61 """Request body posted to the Vigil Guard analyze endpoint."""
63 text: ReadOnly[str]
64 source: ReadOnly[str]
65 mode: ReadOnly[str]
66 metadata: ReadOnly[Mapping[str, _MetadataValue]]
69class _AnalysisView(TypedDict):
70 """Typed read of the analyze endpoint's decoded JSON body."""
72 analysis: ReadOnly[Mapping[str, object]]
75class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object):
76 """Base-class constructor options this guardrail forwards untouched to CustomGuardrail."""
78 supported_event_hooks: ReadOnly[list[GuardrailEventHooks]]
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]: ...
92class VigilGuardMissingConfig(ValueError):
93 pass
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("/")
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
122 fallback: Final = (unreachable_fallback or "fail_closed").lower()
123 self.unreachable_fallback: _FallbackMode = "fail_open" if fallback == "fail_open" else "fail_closed"
125 self.timeout: httpx.Timeout = (
126 _DEFAULT_VIGIL_TIMEOUT if timeout is None else httpx.Timeout(timeout, connect=min(timeout, 5.0))
127 )
129 self.async_handler: _AsyncPostHandler = async_handler or get_async_httpx_client(
130 llm_provider=httpxSpecialProvider.GuardrailCallback,
131 )
133 forwarded: Final[_CustomGuardrailOptions] = {
134 "supported_event_hooks": list(self.get_supported_event_hooks()),
135 **kwargs,
136 }
138 super().__init__(**forwarded)
140 @staticmethod
141 def get_config_model() -> type["GuardrailConfigModel"] | None:
142 from litellm.types.proxy.guardrails.guardrail_hooks.vigil_guard import (
143 VigilGuardGuardrailConfigModel,
144 )
146 return VigilGuardGuardrailConfigModel
148 @classmethod
149 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
150 return [
151 GuardrailEventHooks.pre_call,
152 GuardrailEventHooks.post_call,
153 ]
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
169 source: Final = "user_input" if input_type == "request" else "model_output"
170 metadata: Final = self._collect_metadata(request_data, logging_obj)
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
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 )
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 )
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 )
222 if decision == "SANITIZED":
223 result_texts.append(self._resolve_sanitized_text(text, analysis))
224 else:
225 result_texts.append(text)
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)
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 )
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 )
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 )
270 return self._build_output(inputs, result_texts, result_tool_calls)
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
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
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
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
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"]
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
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 )
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"
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
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)
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
440 call_id: Final = self._extract_call_id(request_data, logging_obj)
441 if call_id:
442 collected["litellm_call_id"] = call_id
444 return collected
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
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