Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py: 14%
200 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# Use Generic Guardrail API for your LLM calls
4#
5# +-------------------------------------------------------------+
6# Thank you users! We ❤️ you! - Krrish & Ishaan
8import fnmatch
9import os
10from collections.abc import Mapping, Sequence
11from typing import TYPE_CHECKING, Any, Final, Literal, Optional
13import httpx
15from litellm._logging import verbose_proxy_logger
16from litellm._version import version as litellm_version
17from litellm.exceptions import GuardrailRaisedException, Timeout
18from litellm.integrations.custom_guardrail import (
19 CustomGuardrail,
20 log_guardrail_information,
21)
22from litellm.llms.custom_httpx.http_handler import (
23 get_async_httpx_client,
24 httpxSpecialProvider,
25)
26from litellm.types.guardrails import GuardrailEventHooks
27from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
28from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
29 GenericGuardrailAPIMetadata,
30 GenericGuardrailAPIRequest,
31 GenericGuardrailAPIResponse,
32 GuardrailToolParam,
33)
34from litellm.types.utils import GenericGuardrailAPIInputs
36if TYPE_CHECKING: 36 ↛ 37line 36 didn't jump to line 37 because the condition on line 36 was never true
37 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
38 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
40GUARDRAIL_NAME: Final = "generic_guardrail_api"
42# Headers whose values are forwarded as-is (case-insensitive). Glob patterns supported (e.g. x-stainless-*, x-litellm*).
43_HEADER_VALUE_ALLOWLIST: Final = frozenset(
44 {
45 "host",
46 "accept-encoding",
47 "connection",
48 "accept",
49 "content-type",
50 "user-agent",
51 "x-stainless-*",
52 "x-litellm-*",
53 "content-length",
54 }
55)
57# Placeholder for headers that exist but are not on the allowlist (we don't expose their value).
58_HEADER_PRESENT_PLACEHOLDER: Final = "[present]"
61def _header_value_allowed(
62 header_name: str,
63 extra_allowlist: set[str] | None = None,
64) -> bool:
65 """Return True if this header's value may be forwarded (allowlist, including globs and extra_headers)."""
66 lower: Final = header_name.lower()
67 if lower in _HEADER_VALUE_ALLOWLIST:
68 return True
69 for pattern in _HEADER_VALUE_ALLOWLIST:
70 if "*" in pattern and fnmatch.fnmatch(lower, pattern):
71 return True
72 if extra_allowlist and lower in extra_allowlist:
73 return True
74 return False
77def _sanitize_inbound_headers(
78 headers: object,
79 extra_allowlist: set[str] | None = None,
80) -> dict[str, str] | None:
81 """
82 Sanitize inbound headers before passing them to a 3rd party guardrail service.
84 - Allowlist: default allowlist + extra_allowlist (from litellm_params.extra_headers); only these have values forwarded.
85 - All other headers are included with value "[present]" so the guardrail knows the header existed.
86 - Coerces values to str (for JSON serialization).
87 """
88 if not headers or not isinstance(headers, dict):
89 return None
91 sanitized: Final[dict[str, str]] = {}
92 for k, v in headers.items():
93 if k is None:
94 continue
95 key = str(k)
96 if _header_value_allowed(key, extra_allowlist=extra_allowlist):
97 try:
98 sanitized[key] = str(v)
99 except Exception:
100 continue
101 else:
102 sanitized[key] = _HEADER_PRESENT_PLACEHOLDER
104 return sanitized or None
107def _extract_inbound_headers(
108 request_data: dict,
109 logging_obj: Optional["LiteLLMLoggingObj"],
110 extra_allowlist: set[str] | None = None,
111) -> dict[str, str] | None:
112 """
113 Extract inbound headers from available request context.
115 We try multiple locations to support different call paths:
116 - proxy endpoints: request_data["proxy_server_request"]["headers"]
117 - if the guardrail is passed the proxy_server_request object directly
118 - metadata headers captured in litellm_pre_call_utils
119 - response hooks: fallback to logging_obj.model_call_details
120 """
121 # 1) Most common path (proxy): full request context in proxy_server_request
122 headers = request_data.get("proxy_server_request", {}).get("headers")
123 if headers:
124 return _sanitize_inbound_headers(headers, extra_allowlist=extra_allowlist)
126 # 2) Some guardrails pass proxy_server_request as request_data itself
127 headers = request_data.get("headers")
128 if headers:
129 return _sanitize_inbound_headers(headers, extra_allowlist=extra_allowlist)
131 # 3) Pre-call: headers stored in request metadata
132 metadata_headers: Final = (request_data.get("metadata") or {}).get("headers")
133 if metadata_headers:
134 return _sanitize_inbound_headers(metadata_headers, extra_allowlist=extra_allowlist)
136 litellm_metadata_headers: Final = (request_data.get("litellm_metadata") or {}).get("headers")
137 if litellm_metadata_headers:
138 return _sanitize_inbound_headers(litellm_metadata_headers, extra_allowlist=extra_allowlist)
140 # 4) Post-call: headers not present on response; fallback to logging object
141 if logging_obj and getattr(logging_obj, "model_call_details", None):
142 try:
143 details: Final = logging_obj.model_call_details or {}
144 headers = details.get("litellm_params", {}).get("metadata", {}).get("headers", None)
145 if headers:
146 return _sanitize_inbound_headers(headers, extra_allowlist=extra_allowlist)
147 except Exception:
148 pass
150 return None
153def _structured_rows_to_write_back(
154 original_rows: Sequence[AllMessageValues] | None,
155 shown_rows: Sequence[AllMessageValues] | None,
156 returned_rows: Sequence[AllMessageValues],
157) -> tuple[AllMessageValues, ...] | None:
158 """The request model drops row keys its message types do not declare, so a
159 row the server echoes back verbatim is restored to the original row object.
160 A server that echoes every row back unchanged has not rewritten anything
161 per row, so its answer is read from texts, as it was before rows could be
162 returned at all."""
163 if original_rows is None or shown_rows is None or len(returned_rows) != len(original_rows):
164 return tuple(returned_rows)
165 if all(returned == shown for shown, returned in zip(shown_rows, returned_rows)):
166 return None
167 return tuple(
168 original if returned == shown else returned
169 for original, shown, returned in zip(original_rows, shown_rows, returned_rows)
170 )
173class GenericGuardrailAPI(CustomGuardrail):
174 """
175 Generic Guardrail API integration for LiteLLM.
177 This integration allows you to use any guardrail API that follows the
178 LiteLLM Basic Guardrail API spec without needing to write custom integration code.
180 The API should accept a POST request with:
181 {
182 "text": str,
183 "request_body": dict,
184 "additional_provider_specific_params": dict
185 }
187 And return:
188 {
189 "action": "BLOCKED" | "NONE" | "GUARDRAIL_INTERVENED",
190 "blocked_reason": str (optional, only if action is BLOCKED),
191 "text": str (optional, modified text if action is GUARDRAIL_INTERVENED)
192 }
193 """
195 def __init__(
196 self,
197 headers: dict[str, Any] | None = None,
198 api_base: str | None = None,
199 api_key: str | None = None,
200 additional_provider_specific_params: Mapping[str, object] | None = None,
201 unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
202 fail_on_error: bool | None = True,
203 extra_headers: list | None = None,
204 streaming_end_of_stream_only: bool | None = None,
205 streaming_sampling_rate: int | None = None,
206 streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = None,
207 **kwargs,
208 ):
209 self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
210 self.headers = headers or {}
211 self.extra_headers = extra_headers or []
213 # If api_key is provided, add it as x-api-key header
214 if api_key:
215 self.headers["x-api-key"] = api_key
217 base_url = api_base or os.environ.get("GENERIC_GUARDRAIL_API_BASE")
219 if not base_url:
220 raise ValueError(
221 "api_base is required for Generic Guardrail API. "
222 "Set GENERIC_GUARDRAIL_API_BASE environment variable or pass it in litellm_params"
223 )
225 # Append the endpoint path if not already present
226 if not base_url.endswith("/beta/litellm_basic_guardrail_api"):
227 base_url = base_url.rstrip("/")
228 self.api_base = f"{base_url}/beta/litellm_basic_guardrail_api"
229 else:
230 self.api_base = base_url
232 self.additional_provider_specific_params = additional_provider_specific_params or {}
234 self.unreachable_fallback: Literal["fail_closed", "fail_open"] = unreachable_fallback
236 self.fail_on_error: bool = True if fail_on_error is None else fail_on_error
238 # Read by UnifiedLLMGuardrails.async_post_call_streaming_iterator_hook
239 # via getattr(guardrail_to_apply, "streaming_*", default).
240 self.streaming_end_of_stream_only: bool = (
241 False if streaming_end_of_stream_only is None else streaming_end_of_stream_only
242 )
243 if streaming_sampling_rate is not None and streaming_sampling_rate < 1:
244 raise ValueError(f"streaming_sampling_rate must be >= 1 (got {streaming_sampling_rate})")
245 self.streaming_sampling_rate: int = 5 if streaming_sampling_rate is None else streaming_sampling_rate
247 # Read by UnifiedLLMGuardrails.async_post_call_streaming_iterator_hook.
248 # "block_only" (default) drops text rewrites on the streaming path;
249 # "incremental_diff" emits them as synthetic deltas.
250 self.streaming_transform_mode: Literal["block_only", "incremental_diff"] = (
251 "block_only" if streaming_transform_mode is None else streaming_transform_mode
252 )
254 # Set supported event hooks
255 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
257 super().__init__(**kwargs)
259 verbose_proxy_logger.debug("Generic Guardrail API initialized with api_base: %s", self.api_base)
261 def _extract_user_api_key_metadata(self, request_data: dict) -> GenericGuardrailAPIMetadata:
262 """
263 Extract user API key metadata from request_data.
265 Args:
266 request_data: Request data dictionary that may contain:
267 - metadata (for input requests) with user_api_key_* fields
268 - litellm_metadata (for output responses) with user_api_key_* fields
270 Returns:
271 GenericGuardrailAPIMetadata with extracted user information
272 """
273 result_metadata: Final = GenericGuardrailAPIMetadata()
275 # Get the source of metadata - try both locations
276 # 1. For output responses: litellm_metadata (set by handlers with prefixed keys)
277 # 2. For input requests: metadata (already present in request_data with prefixed keys)
278 litellm_metadata: Final = request_data.get("litellm_metadata", {})
279 top_level_metadata: Final = request_data.get("metadata", {})
281 # Merge both sources, preferring litellm_metadata if both exist
282 metadata_dict: Final = {**top_level_metadata, **litellm_metadata}
284 if not metadata_dict:
285 return result_metadata
287 # Dynamically iterate through GenericGuardrailAPIMetadata fields
288 # and extract matching fields from the source metadata
289 # Fields in metadata are already prefixed with 'user_api_key_'
290 for field_name in GenericGuardrailAPIMetadata.__annotations__:
291 value = metadata_dict.get(field_name)
292 if value is not None:
293 result_metadata[field_name] = value
295 # handle user_api_key_token = user_api_key_hash
296 if metadata_dict.get("user_api_key_token") is not None:
297 result_metadata["user_api_key_hash"] = metadata_dict.get("user_api_key_token")
299 verbose_proxy_logger.debug(
300 "Generic Guardrail API: Extracted user metadata: %s",
301 {k: v for k, v in result_metadata.items() if v is not None},
302 )
304 return result_metadata
306 def _fail_open_passthrough(
307 self,
308 *,
309 inputs: GenericGuardrailAPIInputs,
310 input_type: Literal["request", "response"],
311 logging_obj: Optional["LiteLLMLoggingObj"],
312 error: Exception,
313 http_status_code: int | None = None,
314 ) -> GenericGuardrailAPIInputs:
315 status_suffix: Final = f" http_status_code={http_status_code}" if http_status_code else ""
316 verbose_proxy_logger.critical(
317 "Generic Guardrail API error (fail-open). Proceeding without guardrail.%s "
318 "guardrail_name=%s api_base=%s input_type=%s litellm_call_id=%s litellm_trace_id=%s",
319 status_suffix,
320 getattr(self, "guardrail_name", None),
321 getattr(self, "api_base", None),
322 input_type,
323 getattr(logging_obj, "litellm_call_id", None) if logging_obj else None,
324 getattr(logging_obj, "litellm_trace_id", None) if logging_obj else None,
325 exc_info=error,
326 )
327 # Keep flow going - treat as action=NONE (no modifications)
328 return_inputs: Final[GenericGuardrailAPIInputs] = {}
329 return_inputs.update(inputs)
330 return return_inputs
332 def _build_request_headers(self) -> dict:
333 """Build HTTP headers for the guardrail API request."""
334 headers: Final = {"Content-Type": "application/json"}
335 if self.headers:
336 headers.update(self.headers)
337 return headers
339 def _build_guardrail_return_inputs(
340 self,
341 *,
342 texts: list,
343 images: list[str] | None,
344 tools: list[ChatCompletionToolParam] | None,
345 structured_messages: Sequence[AllMessageValues] | None,
346 shown_messages: Sequence[AllMessageValues] | None,
347 guardrail_response: GenericGuardrailAPIResponse,
348 ) -> GenericGuardrailAPIInputs:
349 # Action is NONE or no modifications needed
350 return_inputs: Final = GenericGuardrailAPIInputs(texts=texts)
351 if guardrail_response.texts:
352 return_inputs["texts"] = guardrail_response.texts
353 if guardrail_response.images:
354 return_inputs["images"] = guardrail_response.images
355 elif images:
356 return_inputs["images"] = images
357 if guardrail_response.tools:
358 return_inputs["tools"] = guardrail_response.tools
359 elif tools:
360 return_inputs["tools"] = tools
361 rows_to_write_back: Final = (
362 _structured_rows_to_write_back(structured_messages, shown_messages, guardrail_response.structured_messages)
363 if guardrail_response.structured_messages
364 else None
365 )
366 if rows_to_write_back is not None:
367 return_inputs["structured_messages"] = list(rows_to_write_back) # mutable-ok: guardrail inputs take a list
368 if guardrail_response.stream_holdback_chars is not None:
369 return_inputs["stream_holdback_chars"] = guardrail_response.stream_holdback_chars
370 return return_inputs
372 def _handle_guardrail_request_error(
373 self,
374 error: Exception,
375 inputs: GenericGuardrailAPIInputs,
376 input_type: Literal["request", "response"],
377 logging_obj: Optional["LiteLLMLoggingObj"],
378 is_unreachable: bool = True,
379 ) -> GenericGuardrailAPIInputs:
380 unreachable_fail_open: Final = is_unreachable and self.unreachable_fallback == "fail_open"
381 if unreachable_fail_open or not self.fail_on_error:
382 http_status_code: Final = getattr(getattr(error, "response", None), "status_code", None)
383 return self._fail_open_passthrough(
384 inputs=inputs,
385 input_type=input_type,
386 logging_obj=logging_obj,
387 error=error,
388 **({"http_status_code": http_status_code} if http_status_code else {}),
389 )
390 verbose_proxy_logger.error("Generic Guardrail API: failed to make request: %s", str(error))
391 raise Exception(f"Generic Guardrail API failed: {error}")
393 @log_guardrail_information
394 async def apply_guardrail(
395 self,
396 inputs: GenericGuardrailAPIInputs,
397 request_data: dict,
398 input_type: Literal["request", "response"],
399 logging_obj: Optional["LiteLLMLoggingObj"] = None,
400 ) -> GenericGuardrailAPIInputs:
401 """
402 Apply the Generic Guardrail API to the given inputs.
404 This is the main method that gets called by the framework.
406 Args:
407 inputs: Dictionary containing:
408 - texts: List of texts to check
409 - images: Optional list of images to check
410 - tool_calls: Optional list of tool calls to check
411 request_data: Request data dictionary containing user_api_key_dict and other metadata
412 input_type: Whether this is a "request" or "response" guardrail
413 logging_obj: Optional logging object for tracking the guardrail execution
415 Returns:
416 Tuple of (processed texts, processed images)
418 Raises:
419 Exception: If the guardrail blocks the request
420 """
421 verbose_proxy_logger.debug("Generic Guardrail API: Applying guardrail to text")
423 # Extract texts and images from inputs
424 texts: Final = inputs.get("texts", [])
425 images: Final = inputs.get("images")
426 tools: Final = inputs.get("tools")
427 structured_messages: Final = inputs.get("structured_messages")
428 tool_calls: Final = inputs.get("tool_calls")
429 model: Final = inputs.get("model")
431 # Use provided request_data or create an empty dict
432 if request_data is None:
433 request_data = {}
435 request_body: Final = request_data.get("body") or {}
437 # Merge additional provider specific params from config and dynamic params
438 additional_params: Final = {**self.additional_provider_specific_params}
440 # Get dynamic params from request if available
441 dynamic_params: Final = self.get_guardrail_dynamic_request_body_params(request_body)
442 if dynamic_params:
443 additional_params.update(dynamic_params)
445 # Extract user API key metadata
446 user_metadata: Final = self._extract_user_api_key_metadata(request_data)
447 extra_allowlist = {h.lower() for h in self.extra_headers if isinstance(h, str)} if self.extra_headers else None
448 inbound_headers: Final = _extract_inbound_headers(
449 request_data=request_data,
450 logging_obj=logging_obj,
451 extra_allowlist=extra_allowlist,
452 )
454 try:
455 # Create request payload
456 guardrail_request: Final = GenericGuardrailAPIRequest(
457 litellm_call_id=logging_obj.litellm_call_id if logging_obj else None,
458 litellm_trace_id=logging_obj.litellm_trace_id if logging_obj else None,
459 texts=texts,
460 request_data=user_metadata,
461 request_headers=inbound_headers,
462 litellm_version=litellm_version,
463 images=images,
464 tools=([GuardrailToolParam.model_validate(t) for t in tools] if tools else None),
465 structured_messages=structured_messages,
466 tool_calls=tool_calls,
467 additional_provider_specific_params=additional_params,
468 input_type=input_type,
469 model=model,
470 )
472 headers: Final = self._build_request_headers()
474 # Make the API request
475 # Use mode="json" to ensure all iterables are converted to lists
476 response: Final = await self.async_handler.post(
477 url=self.api_base,
478 json=guardrail_request.model_dump(mode="json"),
479 headers=headers,
480 )
482 response.raise_for_status()
483 response_json: Final = response.json()
485 verbose_proxy_logger.debug("Generic Guardrail API response: %s", response_json)
487 guardrail_response: Final = GenericGuardrailAPIResponse.from_dict(response_json)
489 # Handle the response
490 if guardrail_response.action == "BLOCKED":
491 # Block the request
492 error_message: Final = guardrail_response.blocked_reason or "Content violates policy"
493 verbose_proxy_logger.warning("Generic Guardrail API blocked request: %s", error_message)
494 raise GuardrailRaisedException(
495 guardrail_name=GUARDRAIL_NAME,
496 message=error_message,
497 should_wrap_with_default_message=False,
498 blocked_content=True,
499 )
501 return self._build_guardrail_return_inputs(
502 texts=texts,
503 images=images,
504 tools=tools,
505 structured_messages=structured_messages,
506 shown_messages=guardrail_request.structured_messages,
507 guardrail_response=guardrail_response,
508 )
510 except GuardrailRaisedException:
511 raise
512 except Timeout as e:
513 return self._handle_guardrail_request_error(e, inputs, input_type, logging_obj)
514 except httpx.HTTPStatusError as e:
515 status_code: Final = getattr(getattr(e, "response", None), "status_code", None)
516 is_unreachable: Final = status_code in (502, 503, 504)
517 return self._handle_guardrail_request_error(
518 e, inputs, input_type, logging_obj, is_unreachable=is_unreachable
519 )
520 except httpx.RequestError as e:
521 return self._handle_guardrail_request_error(e, inputs, input_type, logging_obj)
522 except Exception as e:
523 return self._handle_guardrail_request_error(e, inputs, input_type, logging_obj, is_unreachable=False)
525 @staticmethod
526 def get_config_model() -> type["GuardrailConfigModel"] | None:
527 from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
528 GenericGuardrailAPIConfigModel,
529 )
531 return GenericGuardrailAPIConfigModel
533 @classmethod
534 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
535 return [
536 GuardrailEventHooks.pre_call,
537 GuardrailEventHooks.post_call,
538 GuardrailEventHooks.during_call,
539 ]