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

98 statements  

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

1# +-------------------------------------------------------------+ 

2# 

3# Use GuardrailsAI for your LLM calls 

4# 

5# +-------------------------------------------------------------+ 

6# Thank you for using Litellm! - Krrish & Ishaan 

7 

8import json 

9import os 

10from typing import TYPE_CHECKING, Final, Literal, TypedDict 

11 

12from fastapi import HTTPException 

13 

14import litellm 

15from litellm._logging import verbose_proxy_logger 

16from litellm.integrations.custom_guardrail import ( 

17 CustomGuardrail, 

18 log_guardrail_information, 

19) 

20from litellm.litellm_core_utils.prompt_templates.common_utils import ( 

21 get_content_from_model_response, 

22) 

23from litellm.proxy._types import UserAPIKeyAuth 

24from litellm.types.guardrails import GuardrailEventHooks 

25 

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

27 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel 

28 

29 

30class GuardrailsAIResponse(TypedDict): 

31 callId: str 

32 rawLlmOutput: str 

33 validatedOutput: str 

34 validationPassed: bool 

35 

36 

37class InferenceData(TypedDict): 

38 name: str 

39 shape: list[int] 

40 data: list 

41 datatype: str 

42 

43 

44class GuardrailsAIResponsePreCall(TypedDict): 

45 modelname: str 

46 modelversion: str 

47 outputs: list[InferenceData] 

48 

49 

50class GuardrailsAI(CustomGuardrail): 

51 def __init__( 

52 self, 

53 guard_name: str, 

54 api_base: str | None = None, 

55 guardrails_ai_api_input_format: Literal["inputs", "llmOutput"] = "llmOutput", 

56 **kwargs, 

57 ): 

58 if guard_name is None: 

59 raise Exception( 

60 "GuardrailsAIException - Please pass the Guardrails AI guard name via 'litellm_params::guard_name'" 

61 ) 

62 # store kwargs as optional_params 

63 self.guardrails_ai_api_base = api_base or os.getenv("GUARDRAILS_AI_API_BASE") or "http://0.0.0.0:8000" 

64 self.guardrails_ai_guard_name = guard_name 

65 self.optional_params = kwargs 

66 self.guardrails_ai_api_input_format = guardrails_ai_api_input_format 

67 super().__init__(supported_event_hooks=list(self.get_supported_event_hooks()), **kwargs) 

68 

69 async def make_guardrails_ai_api_request(self, llm_output: str, request_data: dict) -> GuardrailsAIResponse: 

70 from httpx import URL 

71 

72 data: Final = { 

73 "llmOutput": llm_output, 

74 **self.get_guardrail_dynamic_request_body_params(request_data=request_data), 

75 } 

76 _json_data: Final = json.dumps(data) 

77 response: Final = await litellm.module_level_aclient.post( 

78 url=str(URL(self.guardrails_ai_api_base).join(f"guards/{self.guardrails_ai_guard_name}/validate")), 

79 data=_json_data, 

80 headers={ 

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

82 }, 

83 ) 

84 verbose_proxy_logger.debug("guardrails_ai response: %s", response) 

85 _json_response: Final = GuardrailsAIResponse(**response.json()) 

86 if _json_response.get("validationPassed") is False: 

87 raise HTTPException( 

88 status_code=400, 

89 detail={ 

90 "error": "Violated guardrail policy", 

91 "guardrails_ai_response": _json_response, 

92 }, 

93 ) 

94 return _json_response 

95 

96 async def make_guardrails_ai_api_request_pre_call_request(self, text_input: str, request_data: dict) -> str: 

97 from httpx import URL 

98 

99 # This branch of code does not work with current version of GuardrailsAI API (as of July 2025), and it is unclear if it ever worked. 

100 # Use guardrails_ai_api_input_format: "llmOutput" config line for all guardrails (which is the default anyway) 

101 # We can still use the "pre_call" mode to validate the inputs even if the API input format is technicallt "llmOutput" 

102 

103 data: Final = { 

104 "inputs": [ 

105 { 

106 "name": "text", 

107 "shape": [1], 

108 "data": [text_input], 

109 "datatype": "BYTES", # not sure what this should be, but Guardrail's response sets BYTES for text response - https://github.com/guardrails-ai/detect_pii/blob/e4719a95a26f6caacb78d46ebb4768317032bee5/app.py#L40C31-L40C36 

110 } 

111 ] 

112 } 

113 _json_data: Final = json.dumps(data) 

114 response = await litellm.module_level_aclient.post( 

115 url=str(URL(self.guardrails_ai_api_base).join(f"guards/{self.guardrails_ai_guard_name}/validate")), 

116 data=_json_data, 

117 headers={ 

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

119 }, 

120 ) 

121 verbose_proxy_logger.debug("guardrails_ai response: %s", response) 

122 if response.status_code == 400: 

123 raise HTTPException( 

124 status_code=400, 

125 detail={ 

126 "error": "Violated guardrail policy", 

127 "guardrails_ai_response": response.json(), 

128 }, 

129 ) 

130 

131 _json_response: Final = GuardrailsAIResponsePreCall(**response.json()) 

132 response = _json_response.get("outputs", [])[0].get("data", [])[0] 

133 return response 

134 

135 async def process_input(self, data: dict, call_type: str) -> dict: 

136 from litellm.litellm_core_utils.prompt_templates.common_utils import ( 

137 get_last_user_message, 

138 set_last_user_message, 

139 ) 

140 

141 # Only process completion-related call types 

142 if call_type not in ["completion", "acompletion"]: 

143 return data 

144 

145 if "messages" not in data: # invalid request 

146 return data 

147 

148 text: Final = get_last_user_message(data["messages"]) 

149 if text is None: 

150 return data 

151 if self.guardrails_ai_api_input_format == "inputs": 

152 updated_text = await self.make_guardrails_ai_api_request_pre_call_request( 

153 text_input=text, request_data=data 

154 ) 

155 else: 

156 _result: Final = await self.make_guardrails_ai_api_request(llm_output=text, request_data=data) 

157 updated_text = _result.get("validatedOutput") or _result.get("rawLlmOutput") or text 

158 data["messages"] = set_last_user_message(data["messages"], updated_text) 

159 

160 return data 

161 

162 @log_guardrail_information 

163 async def async_pre_call_hook( 

164 self, 

165 user_api_key_dict: UserAPIKeyAuth, 

166 cache: litellm.DualCache, 

167 data: dict, 

168 call_type: Literal[ 

169 "completion", 

170 "text_completion", 

171 "embeddings", 

172 "image_generation", 

173 "moderation", 

174 "audio_transcription", 

175 "pass_through_endpoint", 

176 "rerank", 

177 "mcp_call", 

178 ], 

179 ) -> ( 

180 Exception | str | dict | None 

181 ): # raise exception if invalid, return a str for the user to receive - if rejected, or return a modified dictionary for passing into litellm 

182 return await self.process_input(data=data, call_type=call_type) 

183 

184 async def async_logging_hook(self, kwargs: dict, result: object, call_type: str) -> tuple[dict, object]: 

185 if call_type == "acompletion" or call_type == "completion": 

186 kwargs = await self.process_input(data=kwargs, call_type=call_type) 

187 

188 return kwargs, result 

189 

190 @log_guardrail_information 

191 async def async_post_call_success_hook( 

192 self, 

193 data: dict, 

194 user_api_key_dict: UserAPIKeyAuth, 

195 response, 

196 ): 

197 """ 

198 Runs on response from LLM API call 

199 

200 It can be used to reject a response 

201 """ 

202 from litellm.proxy.common_utils.callback_utils import ( 

203 add_guardrail_to_applied_guardrails_header, 

204 ) 

205 

206 event_type: Final[GuardrailEventHooks] = GuardrailEventHooks.post_call 

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

208 return 

209 

210 if not isinstance(response, litellm.ModelResponse): 

211 return 

212 

213 response_str: Final[str] = get_content_from_model_response(response) 

214 if response_str is not None and len(response_str) > 0: 

215 await self.make_guardrails_ai_api_request(llm_output=response_str, request_data=data) 

216 

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

218 

219 return 

220 

221 @staticmethod 

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

223 from litellm.types.proxy.guardrails.guardrail_hooks.guardrails_ai import ( 

224 GuardrailsAIGuardrailConfigModel, 

225 ) 

226 

227 return GuardrailsAIGuardrailConfigModel 

228 

229 @classmethod 

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

231 return [ 

232 GuardrailEventHooks.post_call, 

233 GuardrailEventHooks.pre_call, 

234 GuardrailEventHooks.logging_only, 

235 ]