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

1""" 

2PromptGuard guardrail integration for LiteLLM. 

3 

4Calls the PromptGuard Guard API to scan messages for prompt 

5injection, PII, topic violations, and entity blocklist matches 

6before and after LLM calls. 

7""" 

8 

9import os 

10from typing import TYPE_CHECKING, Final, Literal, Optional, TypedDict 

11 

12from typing_extensions import ReadOnly, Unpack 

13from typing_extensions import TypedDict as ExtraItemsTypedDict 

14 

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 

28 

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 ) 

36 

37_DEFAULT_API_BASE: Final = "https://api.promptguard.co" 

38_GUARD_ENDPOINT: Final = "/api/v1/guard" 

39 

40 

41class PromptGuardGuardAPIResponse(TypedDict, total=False): 

42 """Body returned by the PromptGuard ``/api/v1/guard`` endpoint.""" 

43 

44 decision: ReadOnly[str] 

45 threat_type: ReadOnly[str] 

46 event_id: ReadOnly[str] 

47 confidence: ReadOnly[float] 

48 redacted_messages: ReadOnly[list[AllMessageValues]] 

49 

50 

51class PromptGuardHTTPView(TypedDict): 

52 """Typed read of the untyped JSON body returned by the httpx client.""" 

53 

54 guard_response: ReadOnly[PromptGuardGuardAPIResponse] 

55 

56 

57class _CustomGuardrailOptions(ExtraItemsTypedDict, total=False, extra_items=object): 

58 supported_event_hooks: ReadOnly[list[GuardrailEventHooks] | None] 

59 

60 

61class PromptGuardMissingCredentials(Exception): 

62 pass 

63 

64 

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 ) 

83 

84 self.api_base = (api_base or os.environ.get("PROMPTGUARD_API_BASE") or _DEFAULT_API_BASE).rstrip("/") 

85 

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 

95 

96 self.async_handler = get_async_httpx_client( 

97 llm_provider=httpxSpecialProvider.GuardrailCallback, 

98 ) 

99 

100 options: Final[_CustomGuardrailOptions] = { 

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

102 **kwargs, 

103 } 

104 

105 super().__init__(**options) 

106 

107 @staticmethod 

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

109 from litellm.types.proxy.guardrails.guardrail_hooks.promptguard import ( 

110 PromptGuardConfigModel, 

111 ) 

112 

113 return PromptGuardConfigModel 

114 

115 @classmethod 

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

117 return [ 

118 GuardrailEventHooks.pre_call, 

119 GuardrailEventHooks.post_call, 

120 ] 

121 

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

134 

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 

141 

142 direction: Final = "input" if input_type == "request" else "output" 

143 

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 

152 

153 endpoint: Final = f"{self.api_base}{_GUARD_ENDPOINT}" 

154 

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 ) 

162 

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 

184 

185 verbose_proxy_logger.debug( 

186 "PromptGuard: decision=%s threat=%s", 

187 result.get("decision"), 

188 result.get("threat_type"), 

189 ) 

190 

191 decision: Final = result.get("decision") or "allow" 

192 

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 ) 

202 

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 

214 

215 return inputs 

216 

217 @staticmethod 

218 def _extract_texts_from_messages(messages: list[AllMessageValues]) -> list[str]: 

219 """Extract text content from user-role messages only. 

220 

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