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

1# +-------------------------------------------------------------+ 

2# 

3# Use Generic Guardrail API for your LLM calls 

4# 

5# +-------------------------------------------------------------+ 

6# Thank you users! We ❤️ you! - Krrish & Ishaan 

7 

8import fnmatch 

9import os 

10from collections.abc import Mapping, Sequence 

11from typing import TYPE_CHECKING, Any, Final, Literal, Optional 

12 

13import httpx 

14 

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 

35 

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 

39 

40GUARDRAIL_NAME: Final = "generic_guardrail_api" 

41 

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) 

56 

57# Placeholder for headers that exist but are not on the allowlist (we don't expose their value). 

58_HEADER_PRESENT_PLACEHOLDER: Final = "[present]" 

59 

60 

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 

75 

76 

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. 

83 

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 

90 

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 

103 

104 return sanitized or None 

105 

106 

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. 

114 

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) 

125 

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) 

130 

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) 

135 

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) 

139 

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 

149 

150 return None 

151 

152 

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 ) 

171 

172 

173class GenericGuardrailAPI(CustomGuardrail): 

174 """ 

175 Generic Guardrail API integration for LiteLLM. 

176 

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. 

179 

180 The API should accept a POST request with: 

181 { 

182 "text": str, 

183 "request_body": dict, 

184 "additional_provider_specific_params": dict 

185 } 

186 

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 """ 

194 

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 [] 

212 

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 

216 

217 base_url = api_base or os.environ.get("GENERIC_GUARDRAIL_API_BASE") 

218 

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 ) 

224 

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 

231 

232 self.additional_provider_specific_params = additional_provider_specific_params or {} 

233 

234 self.unreachable_fallback: Literal["fail_closed", "fail_open"] = unreachable_fallback 

235 

236 self.fail_on_error: bool = True if fail_on_error is None else fail_on_error 

237 

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 

246 

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 ) 

253 

254 # Set supported event hooks 

255 kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) 

256 

257 super().__init__(**kwargs) 

258 

259 verbose_proxy_logger.debug("Generic Guardrail API initialized with api_base: %s", self.api_base) 

260 

261 def _extract_user_api_key_metadata(self, request_data: dict) -> GenericGuardrailAPIMetadata: 

262 """ 

263 Extract user API key metadata from request_data. 

264 

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 

269 

270 Returns: 

271 GenericGuardrailAPIMetadata with extracted user information 

272 """ 

273 result_metadata: Final = GenericGuardrailAPIMetadata() 

274 

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", {}) 

280 

281 # Merge both sources, preferring litellm_metadata if both exist 

282 metadata_dict: Final = {**top_level_metadata, **litellm_metadata} 

283 

284 if not metadata_dict: 

285 return result_metadata 

286 

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 

294 

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") 

298 

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 ) 

303 

304 return result_metadata 

305 

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 

331 

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 

338 

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 

371 

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}") 

392 

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. 

403 

404 This is the main method that gets called by the framework. 

405 

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 

414 

415 Returns: 

416 Tuple of (processed texts, processed images) 

417 

418 Raises: 

419 Exception: If the guardrail blocks the request 

420 """ 

421 verbose_proxy_logger.debug("Generic Guardrail API: Applying guardrail to text") 

422 

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") 

430 

431 # Use provided request_data or create an empty dict 

432 if request_data is None: 

433 request_data = {} 

434 

435 request_body: Final = request_data.get("body") or {} 

436 

437 # Merge additional provider specific params from config and dynamic params 

438 additional_params: Final = {**self.additional_provider_specific_params} 

439 

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) 

444 

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 ) 

453 

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 ) 

471 

472 headers: Final = self._build_request_headers() 

473 

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 ) 

481 

482 response.raise_for_status() 

483 response_json: Final = response.json() 

484 

485 verbose_proxy_logger.debug("Generic Guardrail API response: %s", response_json) 

486 

487 guardrail_response: Final = GenericGuardrailAPIResponse.from_dict(response_json) 

488 

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 ) 

500 

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 ) 

509 

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) 

524 

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 ) 

530 

531 return GenericGuardrailAPIConfigModel 

532 

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 ]