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

126 statements  

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

1# litellm/proxy/guardrails/guardrail_hooks/pangea.py 

2import os 

3from typing import TYPE_CHECKING, Final 

4 

5from fastapi import HTTPException 

6 

7from litellm._logging import verbose_proxy_logger 

8from litellm.caching.dual_cache import DualCache 

9from litellm.integrations.custom_guardrail import ( 

10 CustomGuardrail, 

11 log_guardrail_information, 

12) 

13from litellm.llms.custom_httpx.http_handler import ( 

14 get_async_httpx_client, 

15 httpxSpecialProvider, 

16) 

17from litellm.proxy._types import UserAPIKeyAuth 

18from litellm.proxy.common_utils.callback_utils import ( 

19 add_guardrail_to_applied_guardrails_header, 

20) 

21from litellm.types.guardrails import GuardrailEventHooks 

22from litellm.types.utils import ( 

23 Choices, 

24 LLMResponseTypes, 

25 ModelResponse, 

26 TextCompletionResponse, 

27) 

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.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel 

31 

32 

33class PangeaGuardrailMissingSecrets(Exception): 

34 """Custom exception for missing Pangea secrets.""" 

35 

36 

37class _TextCompletionRequest: 

38 def __init__(self, body: dict[str, object]) -> None: 

39 self.body = body 

40 

41 def get_messages(self) -> list[dict]: 

42 return [{"role": "user", "content": self.body["prompt"]}] 

43 

44 # This mutates the original dict, but we'll still return it anyways 

45 def update_original_body(self, prompt_messages: list[dict]) -> dict[str, object]: 

46 assert len(prompt_messages) == 1 

47 self.body["prompt"] = prompt_messages[0]["content"] 

48 return self.body 

49 

50 

51class PangeaHandler(CustomGuardrail): 

52 """ 

53 Pangea AI Guardrail handler to interact with the Pangea AI Guard service. 

54 

55 This class implements the necessary hooks to call the Pangea AI Guard API 

56 for input and output scanning based on the configured recipe. 

57 """ 

58 

59 def __init__( 

60 self, 

61 guardrail_name: str, 

62 pangea_input_recipe: str | None = None, 

63 pangea_output_recipe: str | None = None, 

64 api_key: str | None = None, 

65 api_base: str | None = None, 

66 **kwargs, 

67 ): 

68 """ 

69 Initializes the PangeaHandler. 

70 

71 Args: 

72 guardrail_name (str): The name of the guardrail instance. 

73 pangea_recipe (str): The Pangea recipe key to use for scanning. 

74 api_key (Optional[str]): The Pangea API key. Reads from PANGEA_API_KEY env var if None. 

75 api_base (Optional[str]): The Pangea API base URL. Reads from PANGEA_API_BASE env var or uses default if None. 

76 **kwargs: Additional arguments passed to the CustomGuardrail base class. 

77 """ 

78 self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) 

79 self.api_key = api_key or os.environ.get("PANGEA_API_KEY") 

80 if not self.api_key: 

81 raise PangeaGuardrailMissingSecrets( 

82 "Pangea API Key not found. Set PANGEA_API_KEY environment variable or pass it in litellm_params." 

83 ) 

84 

85 # Default Pangea base URL if not provided 

86 self.api_base = api_base or os.environ.get("PANGEA_API_BASE") or "https://ai-guard.aws.us.pangea.cloud" 

87 self.pangea_input_recipe = pangea_input_recipe 

88 self.pangea_output_recipe = pangea_output_recipe 

89 

90 # Pass relevant kwargs to the parent class 

91 super().__init__( 

92 guardrail_name=guardrail_name, 

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

94 **kwargs, 

95 ) 

96 verbose_proxy_logger.debug( 

97 "Initialized Pangea Guardrail: name=%s, recipe=%s, api_base=%s", 

98 guardrail_name, 

99 pangea_input_recipe, 

100 self.api_base, 

101 ) 

102 

103 async def _call_pangea_ai_guard(self, api: str, payload: dict, hook_name: str) -> dict: 

104 """ 

105 Makes the API call to the Pangea AI Guard endpoint. 

106 The function itself will raise an error in the case that a response 

107 should be blocked, but will return a list of redacted messages that the caller 

108 should act on. 

109 

110 Args: 

111 api (str): Which API to use (text/guard or v1beta/guard) 

112 payload (dict): The request payload. 

113 request_data (dict): Original request data (used for logging/headers). 

114 hook_name (str): Name of the hook calling this function (for logging). 

115 

116 Raises: 

117 HTTPException: If the Pangea API returns a 'blocked: true' response. 

118 Exception: For other API call failures. 

119 

120 Returns: 

121 list[dict]: The original response body 

122 """ 

123 endpoint: Final = f"{self.api_base}/{api}" 

124 

125 headers: Final = { 

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

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

128 } 

129 

130 verbose_proxy_logger.debug( 

131 "Pangea Guardrail (%s): Calling endpoint %s with payload: %s", hook_name, endpoint, payload 

132 ) 

133 

134 response: Final = await self.async_handler.post(url=endpoint, json=payload, headers=headers) 

135 response.raise_for_status() 

136 

137 result: Final = response.json() 

138 

139 if result.get("result", {}).get("blocked"): 

140 verbose_proxy_logger.warning("Pangea Guardrail (%s): Request blocked. Response: %s", hook_name, result) 

141 raise HTTPException( 

142 status_code=400, # Bad Request, indicating violation 

143 detail={ 

144 "error": "Violated Pangea guardrail policy", 

145 "guardrail_name": self.guardrail_name, 

146 }, 

147 ) 

148 verbose_proxy_logger.debug( 

149 "Pangea Guardrail (%s): Request passed. Response: %s", hook_name, result.get("result", {}).get("detectors") 

150 ) 

151 

152 return result 

153 

154 async def _async_pre_call_hook( 

155 self, 

156 user_api_key_dict: UserAPIKeyAuth, 

157 cache: DualCache, 

158 data: dict, 

159 call_type: str, 

160 ): 

161 transformer = None 

162 messages: object = None 

163 if call_type == "text_completion" or call_type == "atext_completion": 

164 transformer = _TextCompletionRequest(data) 

165 messages = transformer.get_messages() 

166 else: 

167 messages = data.get("messages") 

168 

169 ai_guard_payload: Final = { 

170 "debug": False, 

171 "input": {"messages": messages, "tools": data.get("tools")}, 

172 "event_type": "input", 

173 } 

174 if self.pangea_input_recipe: 

175 ai_guard_payload["recipe"] = self.pangea_input_recipe 

176 

177 ai_guard_response = await self._call_pangea_ai_guard("v1beta/guard", ai_guard_payload, "async_pre_call_hook") 

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

179 

180 if not ai_guard_response.get("result", {}).get("transformed"): 

181 return 

182 

183 output: Final = ai_guard_response.get("result", {}).get("output", {}) 

184 if call_type == "text_completion" or call_type == "atext_completion": 

185 data = transformer.update_original_body(output["messages"]) 

186 else: 

187 data["messages"] = output["messages"] 

188 return data 

189 

190 @log_guardrail_information 

191 async def async_pre_call_hook( 

192 self, 

193 user_api_key_dict: UserAPIKeyAuth, 

194 cache: DualCache, 

195 data: dict, 

196 call_type: str, 

197 ): 

198 event_type: Final = GuardrailEventHooks.pre_call 

199 if self.should_run_guardrail(data=data, event_type=event_type) is not True: 

200 verbose_proxy_logger.debug( 

201 "Pangea Guardrail (async_pre_call_hook): Guardrail is disabled %s.", self.guardrail_name 

202 ) 

203 return data 

204 

205 try: 

206 return await self._async_pre_call_hook(user_api_key_dict, cache, data, call_type) 

207 except HTTPException: 

208 raise 

209 except Exception as e: 

210 raise HTTPException( 

211 status_code=500, 

212 detail={ 

213 "error": "Error in Pangea Guardrail", 

214 "guardrail_name": self.guardrail_name, 

215 "exceptions": str(e), 

216 }, 

217 ) from e 

218 

219 async def _async_post_call_success_hook( 

220 self, 

221 data: dict, 

222 user_api_key_dict: UserAPIKeyAuth, 

223 # This union isn't actually correct -- it can get other response types depending on the API called 

224 response: LLMResponseTypes, 

225 ): 

226 if isinstance(response, TextCompletionResponse): 

227 # Assume the earlier call type as well 

228 input_messages = _TextCompletionRequest(data).get_messages() 

229 elif isinstance(response, ModelResponse): 

230 messages: Final = data.get("messages") 

231 if messages is None: 

232 return # No messages to check 

233 input_messages = messages 

234 else: 

235 return 

236 

237 if choices := response.get("choices"): 

238 if isinstance(choices, list): 

239 serialized_choices: Final = [] 

240 for c in choices: 

241 if isinstance(c, Choices): 

242 try: 

243 serialized_choices.append(c.model_dump()) 

244 except Exception: 

245 serialized_choices.append(c.dict()) 

246 else: 

247 serialized_choices.append(c) 

248 choices = serialized_choices 

249 

250 ai_guard_payload: Final = { 

251 "debug": False, 

252 "input": { 

253 "messages": input_messages, 

254 "tools": data.get("tools"), 

255 "choices": choices, 

256 }, 

257 "event_type": "output", 

258 } 

259 

260 if self.pangea_output_recipe: 

261 ai_guard_payload["recipe"] = self.pangea_output_recipe 

262 

263 ai_guard_response = await self._call_pangea_ai_guard("v1beta/guard", ai_guard_payload, "async_pre_call_hook") 

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

265 

266 if not ai_guard_response.get("result", {}).get("transformed"): 

267 return 

268 

269 output: Final = ai_guard_response.get("result", {}).get("output", {}) 

270 response.choices = output["choices"] 

271 return response 

272 

273 @log_guardrail_information 

274 async def async_post_call_success_hook( 

275 self, 

276 data: dict, 

277 user_api_key_dict: UserAPIKeyAuth, 

278 # This union isn't actually correct -- it can get other response types depending on the API called 

279 response: LLMResponseTypes, 

280 ): 

281 """ 

282 Guardrail hook run after a successful LLM call (scans output). 

283 

284 Args: 

285 data (dict): The original request data. 

286 user_api_key_dict (UserAPIKeyAuth): User API key details. 

287 response (LLMResponseTypes): The response object from the LLM call. 

288 """ 

289 event_type: Final = GuardrailEventHooks.post_call 

290 if self.should_run_guardrail(data=data, event_type=event_type) is not True: 

291 verbose_proxy_logger.debug( 

292 "Pangea Guardrail (async_pre_call_hook): Guardrail is disabled %s.", self.guardrail_name 

293 ) 

294 return data 

295 try: 

296 return await self._async_post_call_success_hook(data, user_api_key_dict, response) 

297 except HTTPException: 

298 raise 

299 except Exception as e: 

300 raise HTTPException( 

301 status_code=500, 

302 detail={ 

303 "error": "Error in Pangea Guardrail", 

304 "guardrail_name": self.guardrail_name, 

305 "exceptions": str(e), 

306 }, 

307 ) from e 

308 

309 @staticmethod 

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

311 from litellm.types.proxy.guardrails.guardrail_hooks.pangea import ( 

312 PangeaGuardrailConfigModel, 

313 ) 

314 

315 return PangeaGuardrailConfigModel 

316 

317 @classmethod 

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

319 return [ 

320 GuardrailEventHooks.pre_call, 

321 GuardrailEventHooks.post_call, 

322 ]