Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/aporia_ai/aporia_ai.py: 26%
92 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 AporiaAI for your LLM calls
4#
5# +-------------------------------------------------------------+
6# Thank you users! We ❤️ you! - Krrish & Ishaan
8import os
9import sys
11sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path
12import json
13import sys
14from typing import TYPE_CHECKING, Any, Final, Literal
16from fastapi import HTTPException
18from litellm._logging import verbose_proxy_logger
19from litellm.integrations.custom_guardrail import (
20 CustomGuardrail,
21 log_guardrail_information,
22)
23from litellm.litellm_core_utils.logging_utils import (
24 convert_litellm_response_object_to_str,
25)
26from litellm.llms.custom_httpx.http_handler import (
27 get_async_httpx_client,
28 httpxSpecialProvider,
29)
30from litellm.proxy._types import UserAPIKeyAuth
31from litellm.types.guardrails import GuardrailEventHooks
33GUARDRAIL_NAME: Final = "aporia"
35if TYPE_CHECKING: 35 ↛ 36line 35 didn't jump to line 36 because the condition on line 35 was never true
36 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
39class AporiaGuardrail(CustomGuardrail):
40 @classmethod
41 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
42 return [
43 GuardrailEventHooks.during_call,
44 GuardrailEventHooks.post_call,
45 ]
47 def __init__(self, api_key: str | None = None, api_base: str | None = None, **kwargs):
48 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
49 self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
50 self.aporia_api_key = api_key or os.environ["APORIO_API_KEY"]
51 self.aporia_api_base = api_base or os.environ["APORIO_API_BASE"]
52 super().__init__(**kwargs)
54 #### CALL HOOKS - proxy only ####
55 def transform_messages(self, messages: list[dict]) -> list[dict]:
56 supported_openai_roles: Final = ["system", "user", "assistant"]
57 default_role: Final = "other" # for unsupported roles - e.g. tool
58 new_messages: Final = []
59 for m in messages:
60 if m.get("role", "") in supported_openai_roles:
61 new_messages.append(m)
62 else:
63 new_messages.append(
64 {
65 "role": default_role,
66 **{key: value for key, value in m.items() if key != "role"},
67 }
68 )
70 return new_messages
72 async def prepare_aporia_request(self, new_messages: list[dict], response_string: str | None = None) -> dict:
73 data: Final[dict[str, Any]] = {}
74 if new_messages is not None:
75 data["messages"] = new_messages
76 if response_string is not None:
77 data["response"] = response_string
79 # Set validation target
80 if new_messages and response_string:
81 data["validation_target"] = "both"
82 elif new_messages:
83 data["validation_target"] = "prompt"
84 elif response_string:
85 data["validation_target"] = "response"
87 verbose_proxy_logger.debug("Aporia AI request: %s", data)
88 return data
90 async def make_aporia_api_request(
91 self,
92 request_data: dict,
93 new_messages: list[dict],
94 response_string: str | None = None,
95 ):
96 data: Final = await self.prepare_aporia_request(new_messages=new_messages, response_string=response_string)
98 data.update(self.get_guardrail_dynamic_request_body_params(request_data=request_data))
100 _json_data: Final = json.dumps(data)
102 """
103 export APORIO_API_KEY=<your key>
104 curl https://gr-prd-trial.aporia.com/some-id \
105 -X POST \
106 -H "X-APORIA-API-KEY: $APORIO_API_KEY" \
107 -H "Content-Type: application/json" \
108 -d '{
109 "messages": [
110 {
111 "role": "user",
112 "content": "This is a test prompt"
113 }
114 ],
115 }
116'
117 """
119 response: Final = await self.async_handler.post(
120 url=self.aporia_api_base + "/validate",
121 data=_json_data,
122 headers={
123 "X-APORIA-API-KEY": self.aporia_api_key,
124 "Content-Type": "application/json",
125 },
126 )
127 verbose_proxy_logger.debug("Aporia AI response: %s", response.text)
128 if response.status_code == 200:
129 # check if the response was flagged
130 _json_response: Final = response.json()
131 action: str = _json_response.get("action") # possible values are modify, passthrough, block, rephrase
132 if action == "block":
133 raise HTTPException(
134 status_code=400,
135 detail={
136 "error": "Violated guardrail policy",
137 "aporia_ai_response": _json_response,
138 },
139 )
141 @log_guardrail_information
142 async def async_post_call_success_hook(
143 self,
144 data: dict,
145 user_api_key_dict: UserAPIKeyAuth,
146 response,
147 ):
148 from litellm.proxy.common_utils.callback_utils import (
149 add_guardrail_to_applied_guardrails_header,
150 )
152 """
153 Use this for the post call moderation with Guardrails
154 """
155 event_type: Final[GuardrailEventHooks] = GuardrailEventHooks.post_call
156 if self.should_run_guardrail(data=data, event_type=event_type) is not True:
157 return
159 response_str: Final[str | None] = convert_litellm_response_object_to_str(response)
160 if response_str is not None:
161 await self.make_aporia_api_request(
162 request_data=data,
163 response_string=response_str,
164 new_messages=data.get("messages", []),
165 )
167 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
169 @log_guardrail_information
170 async def async_moderation_hook(
171 self,
172 data: dict,
173 user_api_key_dict: UserAPIKeyAuth,
174 call_type: Literal[
175 "completion",
176 "embeddings",
177 "image_generation",
178 "moderation",
179 "audio_transcription",
180 "responses",
181 "mcp_call",
182 "anthropic_messages",
183 ],
184 ):
185 from litellm.proxy.common_utils.callback_utils import (
186 add_guardrail_to_applied_guardrails_header,
187 )
188 from litellm.proxy.guardrails.guardrail_helpers import (
189 should_proceed_based_on_metadata,
190 )
192 event_type: Final[GuardrailEventHooks] = GuardrailEventHooks.during_call
193 if self.should_run_guardrail(data=data, event_type=event_type) is not True:
194 return
196 # old implementation - backwards compatibility
198 if (
199 await should_proceed_based_on_metadata(
200 data=data,
201 guardrail_name=GUARDRAIL_NAME,
202 )
203 is False
204 ):
205 return
207 new_messages: list[dict] | None = None
208 if "messages" in data and isinstance(data["messages"], list):
209 new_messages = self.transform_messages(messages=data["messages"])
211 if new_messages is not None:
212 await self.make_aporia_api_request(
213 request_data=data,
214 new_messages=new_messages,
215 )
216 add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=self.guardrail_name)
217 else:
218 verbose_proxy_logger.warning("Aporia AI: not running guardrail. No messages in data")
220 @staticmethod
221 def get_config_model() -> type["GuardrailConfigModel"] | None:
222 from litellm.types.proxy.guardrails.guardrail_hooks.aporia_ai import (
223 AporiaGuardrailConfigModel,
224 )
226 return AporiaGuardrailConfigModel