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
« 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.
10import asyncio
11from concurrent.futures import ThreadPoolExecutor
12from difflib import SequenceMatcher
13from typing import Final, Literal
15from fastapi import HTTPException
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
32HEURISTICS_EXECUTOR: Final = ThreadPoolExecutor(
33 max_workers=PROMPT_INJECTION_HEURISTICS_MAX_THREADS, thread_name_prefix="prompt-injection-heuristics"
34)
37class _OPTIONAL_PromptInjectionDetection(CustomLogger):
38 enforces_request_content: bool = True
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
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 ]
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)
85 if litellm.set_verbose is True:
86 print(print_statement) # noqa: T201
88 def update_environment(self, router: Router | None = None):
89 self.llm_router = router
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 )
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 )
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
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 )
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()
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)
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]
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
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)
179 is_prompt_attack = False
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)
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 )
202 return data
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 )
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}")
234 if self.prompt_injection_params is None:
235 return None
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
242 prompt_injection_system_prompt: Final = getattr(
243 self.prompt_injection_params,
244 "llm_api_system_prompt",
245 prompt_injection_detection_default_pt(),
246 )
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 )
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
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 )
280 return is_prompt_attack