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
« 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
8import json
9import os
10from typing import TYPE_CHECKING, Final, Literal, TypedDict
12from fastapi import HTTPException
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
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
30class GuardrailsAIResponse(TypedDict):
31 callId: str
32 rawLlmOutput: str
33 validatedOutput: str
34 validationPassed: bool
37class InferenceData(TypedDict):
38 name: str
39 shape: list[int]
40 data: list
41 datatype: str
44class GuardrailsAIResponsePreCall(TypedDict):
45 modelname: str
46 modelversion: str
47 outputs: list[InferenceData]
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)
69 async def make_guardrails_ai_api_request(self, llm_output: str, request_data: dict) -> GuardrailsAIResponse:
70 from httpx import URL
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
96 async def make_guardrails_ai_api_request_pre_call_request(self, text_input: str, request_data: dict) -> str:
97 from httpx import URL
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"
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 )
131 _json_response: Final = GuardrailsAIResponsePreCall(**response.json())
132 response = _json_response.get("outputs", [])[0].get("data", [])[0]
133 return response
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 )
141 # Only process completion-related call types
142 if call_type not in ["completion", "acompletion"]:
143 return data
145 if "messages" not in data: # invalid request
146 return data
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)
160 return data
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)
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)
188 return kwargs, result
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
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 )
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
210 if not isinstance(response, litellm.ModelResponse):
211 return
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)
217 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
219 return
221 @staticmethod
222 def get_config_model() -> type["GuardrailConfigModel"] | None:
223 from litellm.types.proxy.guardrails.guardrail_hooks.guardrails_ai import (
224 GuardrailsAIGuardrailConfigModel,
225 )
227 return GuardrailsAIGuardrailConfigModel
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 ]