Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/hooks/prompt_injection_detection.py: 16%

110 statements  

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

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

2# 

3# Prompt Injection Detection 

4# 

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

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

7## Reject a call if it contains a prompt injection attack. 

8 

9 

10import asyncio 

11from concurrent.futures import ThreadPoolExecutor 

12from difflib import SequenceMatcher 

13from typing import Final, Literal 

14 

15from fastapi import HTTPException 

16 

17import litellm 

18from litellm._logging import verbose_proxy_logger 

19from litellm.caching.caching import DualCache 

20from litellm.constants import ( 

21 DEFAULT_PROMPT_INJECTION_SIMILARITY_THRESHOLD, 

22 PROMPT_INJECTION_HEURISTICS_MAX_THREADS, 

23) 

24from litellm.integrations.custom_logger import CustomLogger 

25from litellm.litellm_core_utils.prompt_templates.factory import ( 

26 prompt_injection_detection_default_pt, 

27) 

28from litellm.proxy._types import LiteLLMPromptInjectionParams, UserAPIKeyAuth 

29from litellm.router import Router 

30from litellm.utils import get_formatted_prompt 

31 

32HEURISTICS_EXECUTOR: Final = ThreadPoolExecutor( 

33 max_workers=PROMPT_INJECTION_HEURISTICS_MAX_THREADS, thread_name_prefix="prompt-injection-heuristics" 

34) 

35 

36 

37class _OPTIONAL_PromptInjectionDetection(CustomLogger): 

38 enforces_request_content: bool = True 

39 

40 # Class variables or attributes 

41 def __init__( 

42 self, 

43 prompt_injection_params: LiteLLMPromptInjectionParams | None = None, 

44 ): 

45 self.prompt_injection_params = prompt_injection_params 

46 self.llm_router: Router | None = None 

47 

48 self.verbs = [ 

49 "Ignore", 

50 "Disregard", 

51 "Skip", 

52 "Forget", 

53 "Neglect", 

54 "Overlook", 

55 "Omit", 

56 "Bypass", 

57 "Pay no attention to", 

58 "Do not follow", 

59 "Do not obey", 

60 ] 

61 self.adjectives = [ 

62 "", 

63 "prior", 

64 "previous", 

65 "preceding", 

66 "above", 

67 "foregoing", 

68 "earlier", 

69 "initial", 

70 ] 

71 self.prepositions = [ 

72 "", 

73 "and start over", 

74 "and start anew", 

75 "and begin afresh", 

76 "and start from scratch", 

77 ] 

78 

79 def print_verbose(self, print_statement, level: Literal["INFO", "DEBUG"] = "DEBUG"): 

80 if level == "INFO": 

81 verbose_proxy_logger.info(print_statement) 

82 elif level == "DEBUG": 

83 verbose_proxy_logger.debug(print_statement) 

84 

85 if litellm.set_verbose is True: 

86 print(print_statement) # noqa: T201 

87 

88 def update_environment(self, router: Router | None = None): 

89 self.llm_router = router 

90 

91 if self.prompt_injection_params is not None and self.prompt_injection_params.llm_api_check is True: 

92 if self.llm_router is None: 

93 raise Exception( 

94 "PromptInjectionDetection: Model List not set. Required for Prompt Injection detection." 

95 ) 

96 

97 self.print_verbose( 

98 f"model_names: {self.llm_router.model_names}; self.prompt_injection_params.llm_api_name: {self.prompt_injection_params.llm_api_name}" 

99 ) 

100 if ( 

101 self.prompt_injection_params.llm_api_name is None 

102 or self.prompt_injection_params.llm_api_name not in self.llm_router.model_names 

103 ): 

104 raise Exception( 

105 "PromptInjectionDetection: Invalid LLM API Name. LLM API Name must be a 'model_name' in 'model_list'." 

106 ) 

107 

108 def generate_injection_keywords(self) -> list[str]: 

109 combinations: Final = [] 

110 for verb in self.verbs: 

111 for adj in self.adjectives: 

112 for prep in self.prepositions: 

113 phrase = " ".join(filter(None, [verb, adj, prep])).strip() 

114 if len(phrase.split()) > 2: # additional check to ensure more than 2 words 

115 combinations.append(phrase.lower()) 

116 return combinations 

117 

118 async def check_user_input_similarity_off_loop(self, user_input: str) -> bool: 

119 return await asyncio.get_running_loop().run_in_executor( 

120 HEURISTICS_EXECUTOR, self.check_user_input_similarity, user_input 

121 ) 

122 

123 def check_user_input_similarity( 

124 self, 

125 user_input: str, 

126 similarity_threshold: float = DEFAULT_PROMPT_INJECTION_SIMILARITY_THRESHOLD, 

127 ) -> bool: 

128 user_input_lower: Final = user_input.lower() 

129 keywords: Final = self.generate_injection_keywords() 

130 

131 for keyword in keywords: 

132 # Calculate the length of the keyword to extract substrings of the same length from user input 

133 keyword_length = len(keyword) 

134 

135 for i in range(len(user_input_lower) - keyword_length + 1): 

136 # Extract a substring of the same length as the keyword 

137 substring = user_input_lower[i : i + keyword_length] 

138 

139 # Calculate similarity 

140 match_ratio = SequenceMatcher(None, substring, keyword).ratio() 

141 if match_ratio > similarity_threshold: 

142 self.print_verbose( 

143 print_statement=f"Rejected user input - {user_input}. {match_ratio} similar to {keyword}", 

144 level="INFO", 

145 ) 

146 return True # Found a highly similar substring 

147 return False # No substring crossed the threshold 

148 

149 async def async_pre_call_hook( 

150 self, 

151 user_api_key_dict: UserAPIKeyAuth, 

152 cache: DualCache, 

153 data: dict, 

154 call_type: str, # "completion", "embeddings", "image_generation", "moderation" 

155 ): 

156 try: 

157 """ 

158 - check if user id part of call 

159 - check if user id part of blocked list 

160 """ 

161 self.print_verbose("Inside Prompt Injection Detection Pre-Call Hook") 

162 try: 

163 assert call_type in [ 

164 "acompletion", 

165 "completion", 

166 "text_completion", 

167 "embeddings", 

168 "image_generation", 

169 "moderation", 

170 "audio_transcription", 

171 ] 

172 except Exception: 

173 self.print_verbose( 

174 f"Call Type - {call_type}, not in accepted list - ['completion','embeddings','image_generation','moderation','audio_transcription']" 

175 ) 

176 return data 

177 formatted_prompt: Final = get_formatted_prompt(data=data, call_type=call_type) 

178 

179 is_prompt_attack = False 

180 

181 if self.prompt_injection_params is not None: 

182 # 1. check if heuristics check turned on 

183 if self.prompt_injection_params.heuristics_check is True: 

184 is_prompt_attack = await self.check_user_input_similarity_off_loop(formatted_prompt) 

185 if is_prompt_attack is True: 

186 raise HTTPException( 

187 status_code=400, 

188 detail={"error": "Rejected message. This is a prompt injection attack."}, 

189 ) 

190 # 2. check if vector db similarity check turned on [TODO] Not Implemented yet 

191 if self.prompt_injection_params.vector_db_check is True: 

192 pass 

193 else: 

194 is_prompt_attack = await self.check_user_input_similarity_off_loop(formatted_prompt) 

195 

196 if is_prompt_attack is True: 

197 raise HTTPException( 

198 status_code=400, 

199 detail={"error": "Rejected message. This is a prompt injection attack."}, 

200 ) 

201 

202 return data 

203 

204 except HTTPException as e: 

205 if ( 

206 e.status_code == 400 

207 and isinstance(e.detail, dict) 

208 and "error" in e.detail 

209 and self.prompt_injection_params is not None 

210 and self.prompt_injection_params.reject_as_response 

211 ): 

212 return e.detail.get("error") 

213 raise e 

214 except Exception as e: 

215 verbose_proxy_logger.exception( 

216 "litellm.proxy.hooks.prompt_injection_detection.py::async_pre_call_hook(): Exception occured - %s", e 

217 ) 

218 

219 async def async_moderation_hook( 

220 self, 

221 data: dict, 

222 user_api_key_dict: UserAPIKeyAuth, 

223 call_type: Literal[ 

224 "acompletion", 

225 "completion", 

226 "embeddings", 

227 "image_generation", 

228 "moderation", 

229 "audio_transcription", 

230 ], 

231 ) -> bool | None: 

232 self.print_verbose(f"IN ASYNC MODERATION HOOK - self.prompt_injection_params = {self.prompt_injection_params}") 

233 

234 if self.prompt_injection_params is None: 

235 return None 

236 

237 formatted_prompt: Final = get_formatted_prompt(data=data, call_type=call_type) 

238 if not formatted_prompt: 

239 return None 

240 is_prompt_attack = False 

241 

242 prompt_injection_system_prompt: Final = getattr( 

243 self.prompt_injection_params, 

244 "llm_api_system_prompt", 

245 prompt_injection_detection_default_pt(), 

246 ) 

247 

248 # 3. check if llm api check turned on 

249 if ( 

250 self.prompt_injection_params.llm_api_check is True 

251 and self.prompt_injection_params.llm_api_name is not None 

252 and self.llm_router is not None 

253 ): 

254 # make a call to the llm api 

255 response: Final = await self.llm_router.acompletion( 

256 model=self.prompt_injection_params.llm_api_name, 

257 messages=[ 

258 { 

259 "role": "system", 

260 "content": prompt_injection_system_prompt, 

261 }, 

262 {"role": "user", "content": formatted_prompt}, 

263 ], 

264 ) 

265 

266 self.print_verbose(f"Received LLM Moderation response: {response}") 

267 self.print_verbose(f"llm_api_fail_call_string: {self.prompt_injection_params.llm_api_fail_call_string}") 

268 if isinstance(response, litellm.ModelResponse) and isinstance(response.choices[0], litellm.Choices): 

269 fail_call_string: Final = self.prompt_injection_params.llm_api_fail_call_string 

270 content: Final = response.choices[0].message.content 

271 if fail_call_string is not None and content is not None and fail_call_string in content: 

272 is_prompt_attack = True 

273 

274 if is_prompt_attack is True: 

275 raise HTTPException( 

276 status_code=400, 

277 detail={"error": "Rejected message. This is a prompt injection attack."}, 

278 ) 

279 

280 return is_prompt_attack