Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/promptguard/promptguard.py: 26%
110 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
1"""
2PromptGuard guardrail integration for LiteLLM.
4Calls the PromptGuard Guard API to scan messages for prompt
5injection, PII, topic violations, and entity blocklist matches
6before and after LLM calls.
7"""
9import os
10from typing import TYPE_CHECKING, Final, Literal, Optional, TypedDict
12from typing_extensions import ReadOnly, Unpack
13from typing_extensions import TypedDict as ExtraItemsTypedDict
15from litellm._logging import verbose_proxy_logger
16from litellm.exceptions import GuardrailRaisedException
17from litellm.integrations.custom_guardrail import (
18 CustomGuardrail,
19 log_guardrail_information,
20)
21from litellm.llms.custom_httpx.http_handler import (
22 get_async_httpx_client,
23 httpxSpecialProvider,
24)
25from litellm.types.guardrails import GuardrailEventHooks
26from litellm.types.llms.openai import AllMessageValues
27from litellm.types.utils import GenericGuardrailAPIInputs
29if TYPE_CHECKING: 29 ↛ 30line 29 didn't jump to line 30 because the condition on line 29 was never true
30 from litellm.litellm_core_utils.litellm_logging import (
31 Logging as LiteLLMLoggingObj,
32 )
33 from litellm.types.proxy.guardrails.guardrail_hooks.base import (
34 GuardrailConfigModel,
35 )
37_DEFAULT_API_BASE: Final = "https://api.promptguard.co"
38_GUARD_ENDPOINT: Final = "/api/v1/guard"
41class PromptGuardGuardAPIResponse(TypedDict, total=False):
42 """Body returned by the PromptGuard ``/api/v1/guard`` endpoint."""
44 decision: ReadOnly[str]
45 threat_type: ReadOnly[str]
46 event_id: ReadOnly[str]
47 confidence: ReadOnly[float]
48 redacted_messages: ReadOnly[list[AllMessageValues]]
51class PromptGuardHTTPView(TypedDict):
52 """Typed read of the untyped JSON body returned by the httpx client."""
54 guard_response: ReadOnly[PromptGuardGuardAPIResponse]
57class _CustomGuardrailOptions(ExtraItemsTypedDict, total=False, extra_items=object):
58 supported_event_hooks: ReadOnly[list[GuardrailEventHooks] | None]
61class PromptGuardMissingCredentials(Exception):
62 pass
65class PromptGuardGuardrail(CustomGuardrail):
66 def __init__(
67 self,
68 api_key: str | None = None,
69 api_base: str | None = None,
70 block_on_error: bool | None = None,
71 **kwargs: Unpack[_CustomGuardrailOptions],
72 ) -> None:
73 self.api_key = api_key or os.environ.get(
74 "PROMPTGUARD_API_KEY",
75 )
76 if not self.api_key:
77 raise PromptGuardMissingCredentials(
78 "PromptGuard API key is required. "
79 "Set PROMPTGUARD_API_KEY in the "
80 "environment or pass api_key in "
81 "the guardrail config."
82 )
84 self.api_base = (api_base or os.environ.get("PROMPTGUARD_API_BASE") or _DEFAULT_API_BASE).rstrip("/")
86 if block_on_error is None:
87 env: Final = os.environ.get("PROMPTGUARD_BLOCK_ON_ERROR", "true")
88 self.block_on_error = env.lower() in (
89 "true",
90 "1",
91 "yes",
92 )
93 else:
94 self.block_on_error = block_on_error
96 self.async_handler = get_async_httpx_client(
97 llm_provider=httpxSpecialProvider.GuardrailCallback,
98 )
100 options: Final[_CustomGuardrailOptions] = {
101 "supported_event_hooks": list(self.get_supported_event_hooks()),
102 **kwargs,
103 }
105 super().__init__(**options)
107 @staticmethod
108 def get_config_model() -> type["GuardrailConfigModel"] | None:
109 from litellm.types.proxy.guardrails.guardrail_hooks.promptguard import (
110 PromptGuardConfigModel,
111 )
113 return PromptGuardConfigModel
115 @classmethod
116 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
117 return [
118 GuardrailEventHooks.pre_call,
119 GuardrailEventHooks.post_call,
120 ]
122 @log_guardrail_information
123 async def apply_guardrail(
124 self,
125 inputs: GenericGuardrailAPIInputs,
126 request_data: dict[str, object],
127 input_type: Literal["request", "response"],
128 logging_obj: Optional["LiteLLMLoggingObj"] = None,
129 ) -> GenericGuardrailAPIInputs:
130 texts: Final = inputs.get("texts", [])
131 images: Final = inputs.get("images", [])
132 structured_messages: Final = inputs.get("structured_messages", [])
133 model: Final = inputs.get("model")
135 if structured_messages:
136 messages = list(structured_messages)
137 elif texts:
138 messages = [{"role": "user", "content": text} for text in texts]
139 else:
140 return inputs
142 direction: Final = "input" if input_type == "request" else "output"
144 payload: Final[dict[str, object]] = {
145 "messages": messages,
146 "direction": direction,
147 }
148 if model:
149 payload["model"] = model
150 if images:
151 payload["images"] = images
153 endpoint: Final = f"{self.api_base}{_GUARD_ENDPOINT}"
155 verbose_proxy_logger.debug(
156 "PromptGuard: %s direction=%s msgs=%d imgs=%d",
157 endpoint,
158 direction,
159 len(messages),
160 len(images),
161 )
163 try:
164 response: Final = await self.async_handler.post(
165 url=endpoint,
166 headers={
167 "X-API-Key": self.api_key,
168 "Content-Type": "application/json",
169 },
170 json=payload,
171 timeout=10.0,
172 )
173 response.raise_for_status()
174 view: Final[PromptGuardHTTPView] = {"guard_response": response.json()}
175 result: Final = view["guard_response"]
176 except Exception as exc:
177 verbose_proxy_logger.error("PromptGuard API error: %s", str(exc))
178 if self.block_on_error:
179 raise GuardrailRaisedException(
180 guardrail_name=self.guardrail_name,
181 message=f"PromptGuard API unreachable (block_on_error=True): {exc}",
182 ) from exc
183 return inputs
185 verbose_proxy_logger.debug(
186 "PromptGuard: decision=%s threat=%s",
187 result.get("decision"),
188 result.get("threat_type"),
189 )
191 decision: Final = result.get("decision") or "allow"
193 if decision == "block":
194 threat_type: Final = result.get("threat_type", "unknown")
195 event_id: Final = result.get("event_id", "")
196 confidence: Final = result.get("confidence", 0.0)
197 raise GuardrailRaisedException(
198 guardrail_name=self.guardrail_name,
199 message=(f"Blocked by PromptGuard: {threat_type} (confidence={confidence}, event_id={event_id})"),
200 blocked_content=True,
201 )
203 if decision == "redact":
204 redacted: Final = result.get("redacted_messages")
205 if redacted:
206 if structured_messages:
207 inputs["structured_messages"] = redacted
208 if "texts" in inputs:
209 extracted: Final = self._extract_texts_from_messages(
210 redacted,
211 )
212 if extracted:
213 inputs["texts"] = extracted
215 return inputs
217 @staticmethod
218 def _extract_texts_from_messages(messages: list[AllMessageValues]) -> list[str]:
219 """Extract text content from user-role messages only.
221 Only user messages are extracted to avoid injecting system or
222 assistant content into the ``texts`` list, which should mirror
223 the original user-provided input.
224 """
225 texts: Final[list[str]] = []
226 for message in messages:
227 if message.get("role") != "user":
228 continue
229 content = message.get("content")
230 if isinstance(content, str):
231 texts.append(content)
232 elif isinstance(content, list):
233 for item in content:
234 if isinstance(item, dict) and item.get("type") == "text":
235 text = item.get("text")
236 if text:
237 texts.append(text)
238 return texts