Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/semantic_guard/semantic_guard.py: 20%

95 statements  

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

1""" 

2Semantic Guard — embedding-based prompt injection detection. 

3 

4Uses semantic-router to match user prompts against known attack patterns 

5via embedding similarity. Smarter than regex (understands intent), lighter 

6than an LLM call (~20-50ms per request for embedding). 

7""" 

8 

9from typing import TYPE_CHECKING, Any, Final, Protocol 

10 

11from litellm._logging import verbose_logger 

12from litellm.integrations.custom_guardrail import ( 

13 CustomGuardrail, 

14 log_guardrail_information, 

15) 

16from litellm.proxy.guardrails.guardrail_hooks.semantic_guard.route_loader import ( 

17 SemanticGuardRouteLoader, 

18) 

19from litellm.types.guardrails import GuardrailEventHooks, Mode 

20from litellm.types.utils import CallTypes 

21 

22try: 

23 from fastapi.exceptions import HTTPException 

24except ImportError: 

25 HTTPException = None 

26 

27if TYPE_CHECKING: 27 ↛ 28line 27 didn't jump to line 28 because the condition on line 27 was never true

28 from semantic_router.routers import SemanticRouter 

29 

30 from litellm.caching import DualCache 

31 from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth 

32 from litellm.router import Router 

33 

34 

35class SemanticGuardrail(CustomGuardrail): 

36 """ 

37 Semantic matching guardrail that blocks requests matching known-bad patterns 

38 using embedding similarity via semantic-router. 

39 

40 Unlike regex, this understands intent: 

41 - "how to make a bomb?" -> may match harmful route (BLOCKED) 

42 - "tell me the spelling of bomb" -> does NOT match (ALLOWED) 

43 """ 

44 

45 def __init__( 

46 self, 

47 guardrail_name: str, 

48 llm_router: "Router", 

49 embedding_model: str, 

50 similarity_threshold: float, 

51 route_templates: list[str] | None = None, 

52 custom_routes_file: str | None = None, 

53 custom_routes: list[dict[str, object]] | None = None, 

54 on_flagged_action: str = "block", 

55 event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None, 

56 default_on: bool = False, 

57 **kwargs, 

58 ): 

59 super().__init__( 

60 guardrail_name=guardrail_name, 

61 supported_event_hooks=list(self.get_supported_event_hooks()), 

62 event_hook=event_hook or GuardrailEventHooks.pre_call, 

63 default_on=default_on, 

64 **kwargs, 

65 ) 

66 

67 self.guardrail_provider = "semantic_guard" 

68 self.embedding_model = embedding_model 

69 self.similarity_threshold = similarity_threshold 

70 self.on_flagged_action = on_flagged_action 

71 self.llm_router = llm_router 

72 

73 routes: Final = SemanticGuardRouteLoader.build_routes( 

74 route_templates=route_templates, 

75 custom_routes_file=custom_routes_file, 

76 custom_routes=custom_routes, 

77 global_threshold=similarity_threshold, 

78 ) 

79 

80 if not routes: 

81 raise ValueError("SemanticGuardrail: no routes configured. Provide route_templates or custom_routes.") 

82 

83 self.semantic_router: SemanticRouter = SemanticGuardRouteLoader.build_semantic_router( 

84 routes=routes, 

85 litellm_router=llm_router, 

86 embedding_model=embedding_model, 

87 global_threshold=similarity_threshold, 

88 ) 

89 

90 self.route_count = len(routes) 

91 verbose_logger.info( 

92 "SemanticGuardrail '%s' initialized with %s routes, embedding_model=%s, threshold=%s", 

93 guardrail_name, 

94 self.route_count, 

95 embedding_model, 

96 similarity_threshold, 

97 ) 

98 

99 @classmethod 

100 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: 

101 return [ 

102 GuardrailEventHooks.pre_call, 

103 GuardrailEventHooks.post_call, 

104 ] 

105 

106 @log_guardrail_information 

107 async def async_pre_call_hook( 

108 self, 

109 user_api_key_dict: "UserAPIKeyAuth", 

110 cache: "DualCache", 

111 data: dict, 

112 call_type: str, 

113 ): 

114 """Check user messages against semantic routes before LLM call.""" 

115 messages: Final = self.get_guardrails_messages_for_call_type(call_type=CallTypes(call_type), data=data) 

116 if not messages: 

117 return 

118 

119 user_text: Final = _extract_user_text(messages) 

120 if not user_text: 

121 return 

122 

123 route_choice: Final = _get_top_route_choice(self.semantic_router(text=user_text)) 

124 if route_choice is not None and route_choice.name: 

125 _handle_match( 

126 guardrail=self, 

127 route_name=route_choice.name, 

128 similarity_score=getattr(route_choice, "similarity_score", None), 

129 user_text=user_text, 

130 data=data, 

131 ) 

132 

133 return 

134 

135 @log_guardrail_information 

136 async def async_post_call_success_hook( 

137 self, 

138 data: dict, 

139 user_api_key_dict: "UserAPIKeyAuth", 

140 response, 

141 ): 

142 """Optionally check LLM response for attack patterns.""" 

143 response_text: Final = _extract_response_text(response) 

144 if not response_text: 

145 return response 

146 

147 route_choice: Final = _get_top_route_choice(self.semantic_router(text=response_text)) 

148 if route_choice is not None and route_choice.name: 

149 _handle_match( 

150 guardrail=self, 

151 route_name=route_choice.name, 

152 similarity_score=getattr(route_choice, "similarity_score", None), 

153 user_text=response_text, 

154 data=data, 

155 ) 

156 

157 return response 

158 

159 

160class _RouteChoice(Protocol): 

161 """The semantic-router match this guardrail reads: the route that fired, if any.""" 

162 

163 @property 

164 def name(self) -> str | None: ... 164 ↛ exitline 164 didn't return from function 'name' because

165 

166 

167def _get_top_route_choice(result: _RouteChoice | list[_RouteChoice] | None) -> _RouteChoice | None: 

168 """Extract the top RouteChoice from SemanticRouter result. 

169 

170 SemanticRouter.__call__ can return RouteChoice or List[RouteChoice]. 

171 """ 

172 if result is None: 

173 return None 

174 if isinstance(result, list): 

175 return result[0] if result else None 

176 return result 

177 

178 

179def _extract_user_text(messages: list) -> str: 

180 """Extract the latest user message text.""" 

181 for msg in reversed(messages): 

182 if isinstance(msg, dict) and msg.get("role") == "user": 

183 content = msg.get("content", "") 

184 if isinstance(content, str): 

185 return content 

186 if isinstance(content, list): 

187 return " ".join(block.get("text", "") if isinstance(block, dict) else str(block) for block in content) 

188 return "" 

189 

190 

191def _extract_response_text(response: Any) -> str: 

192 """Extract text from every LLM response choice.""" 

193 if hasattr(response, "choices") and response.choices: 

194 text_parts: Final[list[str]] = [] 

195 for choice in response.choices: 

196 if hasattr(choice, "message") and choice.message: 

197 text = _content_to_text(choice.message.content) 

198 if text: 

199 text_parts.append(text) 

200 return "\n".join(text_parts) 

201 return "" 

202 

203 

204def _content_to_text(content: object) -> str: 

205 if isinstance(content, str): 

206 return content 

207 if isinstance(content, list): 

208 text_parts: Final = [ 

209 block.get("text") for block in content if isinstance(block, dict) and isinstance(block.get("text"), str) 

210 ] 

211 return " ".join(part for part in text_parts if part) 

212 return "" 

213 

214 

215def _handle_match( 

216 guardrail: SemanticGuardrail, 

217 route_name: str, 

218 similarity_score: float | None, 

219 user_text: str, 

220 data: dict, 

221) -> None: 

222 """Block or passthrough based on config.""" 

223 violation_msg = f"Request blocked by semantic guardrail '{guardrail.guardrail_name}'. Matched route: {route_name}" 

224 

225 detection_info: Final = { 

226 "route_name": route_name, 

227 "similarity_score": similarity_score, 

228 "guardrail": guardrail.guardrail_name, 

229 } 

230 

231 verbose_logger.warning( 

232 "SemanticGuard match: route=%s, score=%s, action=%s", route_name, similarity_score, guardrail.on_flagged_action 

233 ) 

234 

235 if guardrail.on_flagged_action == "passthrough": 

236 guardrail.raise_passthrough_exception( 

237 violation_message=violation_msg, 

238 request_data=data, 

239 detection_info=detection_info, 

240 ) 

241 else: 

242 raise HTTPException( # pyright: ignore[reportOptionalCall] # fastapi is installed wherever this proxy hook runs 

243 status_code=400, 

244 detail={ 

245 "error": violation_msg, 

246 "route": route_name, 

247 "similarity_score": similarity_score, 

248 "type": "semantic_guard_violation", 

249 }, 

250 )