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

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

2# 

3# Use AporiaAI for your LLM calls 

4# 

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

6# Thank you users! We ❤️ you! - Krrish & Ishaan 

7 

8import os 

9import sys 

10 

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 

15 

16from fastapi import HTTPException 

17 

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 

32 

33GUARDRAIL_NAME: Final = "aporia" 

34 

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 

37 

38 

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 ] 

46 

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) 

53 

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 ) 

69 

70 return new_messages 

71 

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 

78 

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" 

86 

87 verbose_proxy_logger.debug("Aporia AI request: %s", data) 

88 return data 

89 

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) 

97 

98 data.update(self.get_guardrail_dynamic_request_body_params(request_data=request_data)) 

99 

100 _json_data: Final = json.dumps(data) 

101 

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

118 

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 ) 

140 

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 ) 

151 

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 

158 

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 ) 

166 

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

168 

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 ) 

191 

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 

195 

196 # old implementation - backwards compatibility 

197 

198 if ( 

199 await should_proceed_based_on_metadata( 

200 data=data, 

201 guardrail_name=GUARDRAIL_NAME, 

202 ) 

203 is False 

204 ): 

205 return 

206 

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

210 

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

219 

220 @staticmethod 

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

222 from litellm.types.proxy.guardrails.guardrail_hooks.aporia_ai import ( 

223 AporiaGuardrailConfigModel, 

224 ) 

225 

226 return AporiaGuardrailConfigModel