Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_initializers.py: 17%

81 statements  

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

1# litellm/proxy/guardrails/guardrail_initializers.py 

2from typing import Any, Final 

3 

4import litellm 

5from litellm.integrations.custom_guardrail import CustomGuardrail 

6from litellm.proxy._types import CommonProxyErrors 

7from litellm.types.guardrails import * 

8 

9 

10def initialize_bedrock(litellm_params: LitellmParams, guardrail: Guardrail): 

11 from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import ( 

12 BedrockGuardrail, 

13 ) 

14 

15 streaming_params: Final = BedrockGuardrailStreamingParams.from_extras(litellm_params.model_extra) 

16 _bedrock_callback: Final = BedrockGuardrail( 

17 guardrail_name=guardrail.get("guardrail_name", ""), 

18 event_hook=litellm_params.mode, 

19 guardrailIdentifier=litellm_params.guardrailIdentifier, 

20 guardrailVersion=litellm_params.guardrailVersion, 

21 checks=litellm_params.checks, 

22 content_filter_threshold=litellm_params.content_filter_threshold, 

23 prompt_attack_threshold=litellm_params.prompt_attack_threshold, 

24 pii_confidence_threshold=litellm_params.pii_confidence_threshold, 

25 chunk_budget_chars=litellm_params.chunk_budget_chars, 

26 contextual_grounding_from_messages=litellm_params.contextual_grounding_from_messages, 

27 default_on=litellm_params.default_on, 

28 disable_exception_on_block=litellm_params.disable_exception_on_block, 

29 mask_request_content=litellm_params.mask_request_content, 

30 mask_response_content=litellm_params.mask_response_content, 

31 aws_region_name=litellm_params.aws_region_name, 

32 aws_access_key_id=litellm_params.aws_access_key_id, 

33 aws_secret_access_key=litellm_params.aws_secret_access_key, 

34 aws_session_token=litellm_params.aws_session_token, 

35 aws_session_name=litellm_params.aws_session_name, 

36 aws_profile_name=litellm_params.aws_profile_name, 

37 aws_role_name=litellm_params.aws_role_name, 

38 aws_web_identity_token=litellm_params.aws_web_identity_token, 

39 aws_sts_endpoint=litellm_params.aws_sts_endpoint, 

40 aws_external_id=litellm_params.aws_external_id, 

41 aws_bedrock_runtime_endpoint=litellm_params.aws_bedrock_runtime_endpoint, 

42 experimental_use_latest_role_message_only=litellm_params.experimental_use_latest_role_message_only, 

43 only_scan_new_messages=litellm_params.only_scan_new_messages or False, 

44 streaming_buffer_until_moderated=streaming_params.streaming_buffer_until_moderated, 

45 streaming_sampling_rate=streaming_params.streaming_sampling_rate, 

46 streaming_end_of_stream_only=streaming_params.streaming_end_of_stream_only, 

47 streaming_buffer_release_on_scan=streaming_params.streaming_buffer_release_on_scan, 

48 ) 

49 litellm.logging_callback_manager.add_litellm_callback(_bedrock_callback) 

50 return _bedrock_callback 

51 

52 

53def initialize_lakera(litellm_params: LitellmParams, guardrail: Guardrail): 

54 from litellm.proxy.guardrails.guardrail_hooks.lakera_ai import lakeraAI_Moderation 

55 

56 _lakera_callback: Final = lakeraAI_Moderation( 

57 api_base=litellm_params.api_base, 

58 api_key=litellm_params.api_key, 

59 guardrail_name=guardrail.get("guardrail_name", ""), 

60 event_hook=litellm_params.mode, 

61 category_thresholds=litellm_params.category_thresholds, 

62 default_on=litellm_params.default_on, 

63 ) 

64 litellm.logging_callback_manager.add_litellm_callback(_lakera_callback) 

65 return _lakera_callback 

66 

67 

68def initialize_lakera_v2(litellm_params: LitellmParams, guardrail: Guardrail): 

69 from litellm.proxy.guardrails.guardrail_hooks.lakera_ai_v2 import LakeraAIGuardrail 

70 

71 _lakera_v2_callback: Final = LakeraAIGuardrail( 

72 api_base=litellm_params.api_base, 

73 api_key=litellm_params.api_key, 

74 guardrail_name=guardrail.get("guardrail_name", ""), 

75 event_hook=litellm_params.mode, 

76 default_on=litellm_params.default_on, 

77 project_id=litellm_params.project_id, 

78 payload=litellm_params.payload, 

79 breakdown=litellm_params.breakdown, 

80 metadata=litellm_params.metadata, 

81 dev_info=litellm_params.dev_info, 

82 on_flagged=litellm_params.on_flagged, 

83 skip_system_message_in_guardrail=litellm_params.skip_system_message_in_guardrail, 

84 skip_tool_message_in_guardrail=litellm_params.skip_tool_message_in_guardrail, 

85 advisory_system_message=litellm_params.advisory_system_message, 

86 ) 

87 litellm.logging_callback_manager.add_litellm_callback(_lakera_v2_callback) 

88 return _lakera_v2_callback 

89 

90 

91_MCP_EVENT_HOOKS: Final = frozenset( 

92 { 

93 GuardrailEventHooks.pre_mcp_call.value, 

94 GuardrailEventHooks.during_mcp_call.value, 

95 GuardrailEventHooks.post_mcp_call.value, 

96 } 

97) 

98 

99 

100def _configured_event_hooks(mode: str | list[str] | Mode) -> tuple[str, ...]: 

101 if isinstance(mode, str): 

102 return (mode,) 

103 if isinstance(mode, list): 

104 return tuple(mode) 

105 return tuple( 

106 hook 

107 for value in (*mode.tags.values(), mode.default) 

108 if value is not None 

109 for hook in ((value,) if isinstance(value, str) else value) 

110 ) 

111 

112 

113def _is_mcp_only_mode(mode: str | list[str] | Mode) -> bool: 

114 hooks: Final = _configured_event_hooks(mode) 

115 return bool(hooks) and all(hook in _MCP_EVENT_HOOKS for hook in hooks) 

116 

117 

118def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) -> tuple[CustomGuardrail, ...]: 

119 from litellm.proxy.guardrails.guardrail_hooks.presidio import ( 

120 _OPTIONAL_PresidioPIIMasking, 

121 ) 

122 

123 explicit_filter_scope: Final = litellm_params.presidio_filter_scope 

124 filter_scope: Final = explicit_filter_scope or ("input" if _is_mcp_only_mode(litellm_params.mode) else "both") 

125 run_input: Final = filter_scope in ("input", "both") 

126 run_output: Final = filter_scope in ("output", "both") 

127 

128 def _make_presidio_callback(**overrides) -> CustomGuardrail: 

129 params: Final = dict( 

130 guardrail_name=guardrail.get("guardrail_name", ""), 

131 event_hook=litellm_params.mode, 

132 output_parse_pii=litellm_params.output_parse_pii, 

133 presidio_ad_hoc_recognizers=litellm_params.presidio_ad_hoc_recognizers, 

134 mock_redacted_text=litellm_params.mock_redacted_text, 

135 default_on=litellm_params.default_on, 

136 pii_entities_config=litellm_params.pii_entities_config, 

137 presidio_score_thresholds=litellm_params.presidio_score_thresholds, 

138 presidio_analyzer_api_base=litellm_params.presidio_analyzer_api_base, 

139 presidio_anonymizer_api_base=litellm_params.presidio_anonymizer_api_base, 

140 presidio_language=litellm_params.presidio_language, 

141 presidio_entities_deny_list=litellm_params.presidio_entities_deny_list, 

142 apply_to_output=False, 

143 ) 

144 params.update(overrides) 

145 # Passed outside the heterogeneous params dict so the argument keeps 

146 # its precise int | None type. 

147 callback: Final = _OPTIONAL_PresidioPIIMasking( 

148 presidio_analyze_chunk_size_bytes=litellm_params.presidio_analyze_chunk_size_bytes, 

149 **params, 

150 ) 

151 litellm.logging_callback_manager.add_litellm_callback(callback) 

152 return callback 

153 

154 input_callback: Final = _make_presidio_callback() if run_input else None 

155 unmask_output_callback: Final = ( 

156 _make_presidio_callback( 

157 output_parse_pii=True, 

158 event_hook=GuardrailEventHooks.post_call.value, 

159 ) 

160 if run_input and litellm_params.output_parse_pii 

161 else None 

162 ) 

163 mask_output_callback: Final = ( 

164 _make_presidio_callback( 

165 apply_to_output=True, 

166 event_hook=GuardrailEventHooks.post_call.value, 

167 output_parse_pii=False, 

168 mask_response_content=True, 

169 ) 

170 if run_output 

171 else None 

172 ) 

173 return tuple( 

174 callback for callback in (input_callback, unmask_output_callback, mask_output_callback) if callback is not None 

175 ) 

176 

177 

178def initialize_hide_secrets(litellm_params: LitellmParams, guardrail: Guardrail): 

179 try: 

180 from litellm_enterprise.enterprise_callbacks.secret_detection import ( 

181 _ENTERPRISE_SecretDetection, 

182 ) 

183 except ImportError: 

184 raise Exception("Trying to use Secret Detection" + CommonProxyErrors.missing_enterprise_package.value) 

185 

186 _secret_detection_object: Final = _ENTERPRISE_SecretDetection( 

187 detect_secrets_config=litellm_params.detect_secrets_config, 

188 event_hook=litellm_params.mode, 

189 guardrail_name=guardrail.get("guardrail_name", ""), 

190 default_on=litellm_params.default_on, 

191 ) 

192 litellm.logging_callback_manager.add_litellm_callback(_secret_detection_object) 

193 return _secret_detection_object 

194 

195 

196def initialize_tool_permission(litellm_params: LitellmParams, guardrail: Guardrail): 

197 from litellm.proxy.guardrails.guardrail_hooks.tool_permission import ( 

198 ToolPermissionGuardrail, 

199 ) 

200 

201 rules: list[dict[str, Any]] | None = None 

202 if litellm_params.rules: 

203 rules = [] 

204 for rule in litellm_params.rules: 

205 if hasattr(rule, "model_dump"): 

206 rules.append(rule.model_dump()) 

207 else: 

208 rules.append(dict(rule)) 

209 

210 _tool_permission_callback: Final = ToolPermissionGuardrail( 

211 guardrail_name=guardrail.get("guardrail_name", ""), 

212 event_hook=litellm_params.mode, 

213 rules=rules, 

214 default_action=getattr(litellm_params, "default_action", "deny"), 

215 on_disallowed_action=getattr(litellm_params, "on_disallowed_action", "block"), 

216 default_on=litellm_params.default_on, 

217 violation_message_template=litellm_params.violation_message_template, 

218 ) 

219 litellm.logging_callback_manager.add_litellm_callback(_tool_permission_callback) 

220 return _tool_permission_callback 

221 

222 

223def initialize_lasso( 

224 litellm_params: LitellmParams, 

225 guardrail: Guardrail, 

226): 

227 from litellm.proxy.guardrails.guardrail_hooks.lasso import LassoGuardrail 

228 

229 _lasso_callback: Final = LassoGuardrail( 

230 guardrail_name=guardrail.get("guardrail_name", ""), 

231 lasso_api_key=litellm_params.api_key, 

232 api_base=litellm_params.api_base, 

233 user_id=litellm_params.lasso_user_id, 

234 conversation_id=litellm_params.lasso_conversation_id, 

235 mask=litellm_params.mask, 

236 event_hook=litellm_params.mode, 

237 default_on=litellm_params.default_on, 

238 ) 

239 litellm.logging_callback_manager.add_litellm_callback(_lasso_callback) 

240 

241 return _lasso_callback 

242 

243 

244def initialize_panw_prisma_airs(litellm_params, guardrail): 

245 from litellm.proxy.guardrails.guardrail_hooks.panw_prisma_airs import ( 

246 PanwPrismaAirsHandler, 

247 ) 

248 

249 if not litellm_params.api_key: 

250 raise ValueError("PANW Prisma AIRS: api_key is required") 

251 if not litellm_params.profile_name: 

252 raise ValueError("PANW Prisma AIRS: profile_name is required") 

253 

254 _panw_callback: Final = PanwPrismaAirsHandler( 

255 guardrail_name=guardrail.get("guardrail_name", "panw_prisma_airs"), # Use .get() with default 

256 api_key=litellm_params.api_key, 

257 api_base=litellm_params.api_base or "https://service.api.aisecurity.paloaltonetworks.com/v1/scan/sync/request", 

258 profile_name=litellm_params.profile_name, 

259 default_on=litellm_params.default_on, 

260 mask_on_block=getattr(litellm_params, "mask_on_block", False), 

261 mask_request_content=getattr(litellm_params, "mask_request_content", False), 

262 mask_response_content=getattr(litellm_params, "mask_response_content", False), 

263 app_name=getattr(litellm_params, "app_name", None), 

264 fallback_on_error=getattr(litellm_params, "fallback_on_error", "block"), 

265 # `timeout` is now declared on BaseLitellmParams (Optional[float] = None), 

266 # so the attribute always exists. The Pydantic validator on LitellmParams 

267 # coerces strings to float, but None still means "use handler default" — 

268 # guard against float(None) here. 

269 timeout=( 

270 float(getattr(litellm_params, "timeout", None)) 

271 if getattr(litellm_params, "timeout", None) is not None 

272 else 10.0 

273 ), 

274 violation_message_template=litellm_params.violation_message_template, 

275 ) 

276 litellm.logging_callback_manager.add_litellm_callback(_panw_callback) 

277 

278 return _panw_callback