Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/noma/noma_v2.py: 24%
172 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# Noma Security V2 Guardrail Integration for LiteLLM
4#
5# +-------------------------------------------------------------+
7import enum
8import json
9import os
10from collections.abc import Callable, Mapping
11from datetime import datetime
12from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, cast
13from urllib.parse import urlparse
15from litellm._logging import verbose_proxy_logger
16from litellm.integrations.custom_guardrail import (
17 CustomGuardrail,
18 log_guardrail_information,
19)
20from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
21from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
22from litellm.llms.custom_httpx.http_handler import (
23 get_async_httpx_client,
24 httpxSpecialProvider,
25)
26from litellm.proxy.guardrails.guardrail_hooks.noma.noma import NomaBlockedMessage
27from litellm.types.guardrail_base_init import GuardrailBaseInitKwargs
28from litellm.types.guardrails import GuardrailEventHooks
29from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailStatus
31if TYPE_CHECKING: 31 ↛ 32line 31 didn't jump to line 32 because the condition on line 31 was never true
32 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
33 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
36_DEFAULT_API_BASE: Final = "https://api.noma.security/"
37_AIDR_SCAN_ENDPOINT: Final = "/litellm/guardrail"
38_INTERVENED_INPUT_FIELDS: Final = ("texts", "images", "tools", "tool_calls")
39_DEFAULT_API_BASE_HOSTNAME: Final = urlparse(_DEFAULT_API_BASE).hostname
41_GuardrailJsonResponse: TypeAlias = Exception | str | dict[str, object]
43_KEYS_DUPLICATING_SCAN_INPUTS: Final = ("messages", "input")
44_LOGGING_KEYS_DUPLICATING_SCAN_INPUTS: Final = _KEYS_DUPLICATING_SCAN_INPUTS + (
45 "additional_args",
46 "standard_logging_object",
47 "original_response",
48)
51class _Action(str, enum.Enum):
52 BLOCKED = "BLOCKED"
53 NONE = "NONE"
54 GUARDRAIL_INTERVENED = "GUARDRAIL_INTERVENED"
57class NomaV2Guardrail(CustomGuardrail):
58 def __init__(
59 self,
60 api_key: str | None = None,
61 api_base: str | None = None,
62 application_id: str | None = None,
63 monitor_mode: bool | None = None,
64 block_failures: bool | None = None,
65 **kwargs: Any,
66 ) -> None:
67 self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
69 self.api_key = api_key or os.environ.get("NOMA_API_KEY")
70 self.api_base = (api_base or os.environ.get("NOMA_API_BASE") or _DEFAULT_API_BASE).rstrip("/")
71 self.application_id = application_id or os.environ.get("NOMA_APPLICATION_ID")
72 if monitor_mode is None:
73 self.monitor_mode = os.environ.get("NOMA_MONITOR_MODE", "false").lower() == "true"
74 else:
75 self.monitor_mode = monitor_mode
77 if block_failures is None:
78 self.block_failures = os.environ.get("NOMA_BLOCK_FAILURES", "true").lower() == "true"
79 else:
80 self.block_failures = block_failures
82 if self._requires_api_key(api_base=self.api_base) and not self.api_key:
83 raise ValueError("Noma v2 guardrail requires api_key when using Noma SaaS endpoint")
85 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
87 base_kwargs: Final[GuardrailBaseInitKwargs] = kwargs
88 super().__init__(**base_kwargs)
90 @staticmethod
91 def get_config_model() -> type["GuardrailConfigModel"] | None:
92 from litellm.types.proxy.guardrails.guardrail_hooks.noma import (
93 NomaV2GuardrailConfigModel,
94 )
96 return NomaV2GuardrailConfigModel
98 @classmethod
99 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
100 return [
101 GuardrailEventHooks.pre_call,
102 GuardrailEventHooks.during_call,
103 GuardrailEventHooks.post_call,
104 GuardrailEventHooks.pre_mcp_call,
105 GuardrailEventHooks.during_mcp_call,
106 ]
108 def _get_authorization_header(self) -> str:
109 if not self.api_key:
110 return ""
111 return f"Bearer {self.api_key}"
113 @staticmethod
114 def _requires_api_key(api_base: str) -> bool:
115 parsed: Final = urlparse(api_base)
116 return parsed.hostname == _DEFAULT_API_BASE_HOSTNAME
118 @staticmethod
119 def _get_non_empty_str(value: object) -> str | None:
120 if not isinstance(value, str):
121 return None
122 stripped: Final = value.strip()
123 return stripped or None
125 def _resolve_action_from_response(
126 self,
127 response_json: Mapping[str, object],
128 ) -> _Action:
129 action: Final = response_json.get("action")
130 if isinstance(action, str):
131 try:
132 return _Action(action)
133 except ValueError:
134 pass
136 raise ValueError("Noma v2 response missing valid action")
138 def _build_scan_payload(
139 self,
140 inputs: GenericGuardrailAPIInputs,
141 request_data: dict,
142 input_type: Literal["request", "response"],
143 logging_obj: Optional["LiteLLMLoggingObj"],
144 application_id: str | None,
145 ) -> dict:
146 payload_request_data: Final = self._sanitize_payload_for_transport(
147 {key: value for key, value in request_data.items() if key not in _KEYS_DUPLICATING_SCAN_INPUTS}
148 )
149 if logging_obj is not None:
150 model_call_details: Final = getattr(logging_obj, "model_call_details", None)
151 payload_request_data["litellm_logging_obj"] = (
152 {
153 key: value
154 for key, value in model_call_details.items()
155 if key not in _LOGGING_KEYS_DUPLICATING_SCAN_INPUTS
156 }
157 if isinstance(model_call_details, dict)
158 else model_call_details
159 )
161 payload: Final[dict[str, object]] = {
162 "inputs": inputs,
163 "request_data": payload_request_data,
164 "input_type": input_type,
165 "monitor_mode": self.monitor_mode,
166 }
167 if application_id:
168 payload["application_id"] = application_id
169 return payload
171 @staticmethod
172 def _sanitize_payload_for_transport(payload: dict) -> dict:
173 def _default(obj: object) -> object:
174 model_dump: Final[Callable[[], Mapping[str, object]] | None] = getattr(obj, "model_dump", None)
175 if model_dump is not None:
176 try:
177 return model_dump()
178 except Exception:
179 pass
180 return str(obj)
182 try:
183 json_str = json.dumps(payload, default=_default)
184 except (ValueError, TypeError):
185 json_str = safe_dumps(payload)
187 safe_payload: Final[object] = safe_json_loads(json_str, default={})
188 if safe_payload == {} and payload:
189 verbose_proxy_logger.warning(
190 "Noma v2 guardrail: payload serialization failed, falling back to empty payload"
191 )
193 if isinstance(safe_payload, dict):
194 return safe_payload
196 verbose_proxy_logger.warning(
197 "Noma v2 guardrail: payload sanitization produced non-dict output (type=%s), falling back to empty payload",
198 type(safe_payload).__name__,
199 )
200 return {}
202 async def _call_noma_scan(
203 self,
204 payload: dict,
205 ) -> dict[str, object]:
206 headers: Final[dict[str, str]] = {"Content-Type": "application/json"}
207 authorization_header: Final = self._get_authorization_header()
208 if authorization_header:
209 headers["Authorization"] = authorization_header
211 endpoint: Final = f"{self.api_base}{_AIDR_SCAN_ENDPOINT}"
212 sanitized_payload: Final = self._sanitize_payload_for_transport(payload)
213 response: Final = await self.async_handler.post(
214 url=endpoint,
215 headers=headers,
216 json=sanitized_payload,
217 )
218 verbose_proxy_logger.debug(
219 "Noma v2 AIDR response: status_code=%s body=%s",
220 response.status_code,
221 response.text,
222 )
223 response.raise_for_status()
224 response_json: Final[dict[str, object]] = response.json()
225 verbose_proxy_logger.debug(
226 "Noma v2 AIDR response parsed: %s",
227 json.dumps(response_json, default=str),
228 )
229 return response_json
231 def _add_guardrail_observability(
232 self,
233 request_data: dict,
234 start_time: datetime,
235 guardrail_status: GuardrailStatus,
236 guardrail_json_response: _GuardrailJsonResponse,
237 ) -> None:
238 end_time: Final = datetime.now()
239 duration: Final = (end_time - start_time).total_seconds()
240 self.add_standard_logging_guardrail_information_to_request_data(
241 guardrail_provider="noma_v2",
242 guardrail_json_response=guardrail_json_response,
243 request_data=request_data,
244 guardrail_status=guardrail_status,
245 start_time=start_time.timestamp(),
246 end_time=end_time.timestamp(),
247 duration=duration,
248 )
250 def _apply_action(
251 self,
252 inputs: GenericGuardrailAPIInputs,
253 response_json: dict,
254 action: _Action,
255 ) -> GenericGuardrailAPIInputs:
256 if action == _Action.BLOCKED:
257 raise NomaBlockedMessage(response_json)
259 if action == _Action.GUARDRAIL_INTERVENED:
260 updated_inputs: Final = cast(GenericGuardrailAPIInputs, dict(inputs))
261 for field in _INTERVENED_INPUT_FIELDS:
262 value = response_json.get(field)
263 if isinstance(value, list):
264 updated_inputs[field] = value
265 return updated_inputs
267 return inputs
269 @log_guardrail_information
270 async def apply_guardrail(
271 self,
272 inputs: GenericGuardrailAPIInputs,
273 request_data: dict,
274 input_type: Literal["request", "response"],
275 logging_obj: Optional["LiteLLMLoggingObj"] = None,
276 ) -> GenericGuardrailAPIInputs:
277 start_time: Final = datetime.now()
278 guardrail_status: GuardrailStatus = "success"
279 guardrail_json_response: _GuardrailJsonResponse = {}
280 dynamic_params = self.get_guardrail_dynamic_request_body_params(request_data)
281 if not isinstance(dynamic_params, dict):
282 dynamic_params = {}
283 response_json: dict[str, object] | None = None
285 # Per-request dynamic params can override configured application context.
286 application_id = self._get_non_empty_str(dynamic_params.get("application_id"))
288 if application_id is None:
289 application_id = self._get_non_empty_str(self.application_id)
291 # Fall back to API key alias for per-key traceability in Noma dashboard
292 # (ports v1 fallback from PR #16832).
293 if application_id is None:
294 application_id = self._get_non_empty_str(
295 request_data.get("litellm_metadata", {}).get("user_api_key_alias")
296 ) or self._get_non_empty_str(request_data.get("metadata", {}).get("user_api_key_alias"))
298 try:
299 payload: Final = self._build_scan_payload(
300 inputs=inputs,
301 request_data=request_data,
302 input_type=input_type,
303 logging_obj=logging_obj,
304 application_id=application_id,
305 )
307 response_json = await self._call_noma_scan(payload=payload)
308 if self.monitor_mode:
309 action = _Action.NONE
310 else:
311 action = self._resolve_action_from_response(response_json=response_json)
312 guardrail_json_response = response_json
313 verbose_proxy_logger.debug(
314 "Noma v2 guardrail decision: input_type=%s action=%s",
315 input_type,
316 action.value,
317 )
318 processed_inputs: Final = self._apply_action(
319 inputs=inputs,
320 response_json=response_json,
321 action=action,
322 )
324 guardrail_status = "success" if action == _Action.NONE else "guardrail_intervened"
325 return processed_inputs
327 except NomaBlockedMessage as e:
328 guardrail_status = "guardrail_intervened"
329 blocked_detail: Final[dict[str, object]] = {"error": "blocked"}
330 guardrail_json_response = (
331 response_json if isinstance(response_json, dict) else getattr(e, "detail", blocked_detail)
332 )
333 raise
334 except Exception as e:
335 guardrail_status = "guardrail_failed_to_respond"
336 guardrail_json_response = str(e)
337 verbose_proxy_logger.error("Noma v2 guardrail failed: %s", str(e))
338 if self.block_failures:
339 raise
340 return inputs
341 finally:
342 self._add_guardrail_observability(
343 request_data=request_data,
344 start_time=start_time,
345 guardrail_status=guardrail_status,
346 guardrail_json_response=guardrail_json_response,
347 )