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
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1"""
2Semantic Guard — embedding-based prompt injection detection.
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"""
9from typing import TYPE_CHECKING, Any, Final, Protocol
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
22try:
23 from fastapi.exceptions import HTTPException
24except ImportError:
25 HTTPException = None
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
30 from litellm.caching import DualCache
31 from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
32 from litellm.router import Router
35class SemanticGuardrail(CustomGuardrail):
36 """
37 Semantic matching guardrail that blocks requests matching known-bad patterns
38 using embedding similarity via semantic-router.
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 """
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 )
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
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 )
80 if not routes:
81 raise ValueError("SemanticGuardrail: no routes configured. Provide route_templates or custom_routes.")
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 )
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 )
99 @classmethod
100 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
101 return [
102 GuardrailEventHooks.pre_call,
103 GuardrailEventHooks.post_call,
104 ]
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
119 user_text: Final = _extract_user_text(messages)
120 if not user_text:
121 return
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 )
133 return
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
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 )
157 return response
160class _RouteChoice(Protocol):
161 """The semantic-router match this guardrail reads: the route that fired, if any."""
163 @property
164 def name(self) -> str | None: ... 164 ↛ exitline 164 didn't return from function 'name' because
167def _get_top_route_choice(result: _RouteChoice | list[_RouteChoice] | None) -> _RouteChoice | None:
168 """Extract the top RouteChoice from SemanticRouter result.
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
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 ""
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 ""
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 ""
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}"
225 detection_info: Final = {
226 "route_name": route_name,
227 "similarity_score": similarity_score,
228 "guardrail": guardrail.guardrail_name,
229 }
231 verbose_logger.warning(
232 "SemanticGuard match: route=%s, score=%s, action=%s", route_name, similarity_score, guardrail.on_flagged_action
233 )
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 )