Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py: 20%

150 statements  

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

1#!/usr/bin/env python3 

2""" 

3Azure Prompt Shield Native Guardrail Integrationfor LiteLLM 

4""" 

5 

6import math 

7from collections.abc import Mapping, MutableMapping 

8from contextvars import ContextVar 

9from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, NoReturn, cast 

10 

11from fastapi import HTTPException 

12 

13from litellm._logging import verbose_proxy_logger 

14from litellm.integrations.custom_guardrail import ( 

15 CustomGuardrail, 

16 log_guardrail_information, 

17) 

18from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import ( 

19 AZURE_PROMPT_SHIELD_TEXT_RECORD_UNIT, 

20 azure_prompt_shield_guardrail_cost, 

21) 

22from litellm.secret_managers.main import get_secret_str 

23from litellm.types.guardrails import GuardrailEventHooks 

24from litellm.types.utils import ( 

25 CallTypesLiteral, 

26 GenericGuardrailAPIInputs, 

27 GuardrailTracingDetail, 

28) 

29 

30from .base import AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH, AzureGuardrailBase 

31 

32if TYPE_CHECKING: 32 ↛ 33line 32 didn't jump to line 33 because the condition on line 32 was never true

33 from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj 

34 from litellm.proxy._types import UserAPIKeyAuth 

35 from litellm.types.guardrails import LitellmParams 

36 from litellm.types.llms.openai import AllMessageValues 

37 from litellm.types.proxy.guardrails.guardrail_hooks.azure.azure_prompt_shield import ( 

38 AzurePromptShieldGuardrailResponse, 

39 ) 

40 from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel 

41 

42 

43# Per-invocation billing counters. A ContextVar rather than request metadata: the 

44# decorator can swap out ``request_data``, metadata is client-forgeable, and 

45# concurrent guardrails run in separate tasks with their own context copy. 

46_billing_usage_stash: Final[ContextVar[dict[str, int] | None]] = ContextVar( # mutable-ok: task-local stash 

47 "azure_prompt_shield_billing_usage", default=None 

48) 

49 

50 

51def _resolved_secret_value(value: object) -> object: 

52 """Resolve ``os.environ/<VAR>`` references the way guardrail api_key/api_base 

53 are resolved; any other value passes through unchanged. A reference that 

54 resolves to nothing raises instead of silently disabling pricing, so an 

55 intended-paid deployment fails fast rather than starting in usage-only mode.""" 

56 if isinstance(value, str) and value.startswith("os.environ/"): 

57 resolved: Final = get_secret_str(value) 

58 if resolved is None or not resolved.strip(): 

59 raise ValueError(f"Azure Prompt Shield: {value!r} resolves to an unset or blank environment variable") 

60 return resolved 

61 return value 

62 

63 

64def _updated_param(litellm_params: "LitellmParams | dict", key: str) -> object: # mutable-ok: DB dict 

65 """Read one param from a Mapping or a pydantic object, including pydantic 

66 extras (cost_tier / price_per_1000_text_records live there), which the base 

67 class ``vars()`` loop never sees.""" 

68 if isinstance(litellm_params, Mapping): 

69 return litellm_params.get(key) 

70 return getattr(litellm_params, key, None) 

71 

72 

73def _resolved_cost_tier(raw: object) -> str | None: 

74 """Normalize the configured cost_tier to 'free' / 'paid' / None.""" 

75 value: Final = _resolved_secret_value(raw) 

76 if value is None or (isinstance(value, str) and not value.strip()): 

77 return None 

78 tier: Final = str(value).strip().lower() 

79 if tier not in ("free", "paid"): 

80 raise ValueError(f"Azure Prompt Shield: cost_tier must be 'free' or 'paid', got {value!r}") 

81 return tier 

82 

83 

84def _resolved_price(raw: object, cost_tier: str | None) -> float | None: 

85 """Normalize price_per_1000_text_records and validate it against the tier. 

86 

87 A 'paid' tier requires a positive price so a misconfigured deployment fails at 

88 startup instead of silently reporting a wrong cost; an omitted price with no 

89 tier means usage-only tracking (no cost estimate).""" 

90 value: Final = _resolved_secret_value(raw) 

91 price: Final = _price_from_value(value) 

92 if cost_tier == "paid" and (price is None or price <= 0): 

93 raise ValueError("Azure Prompt Shield: cost_tier 'paid' requires a positive price_per_1000_text_records") 

94 return price 

95 

96 

97def _price_from_value(value: object) -> float | None: 

98 """Parse a resolved price value into a float; None for an unset/blank value.""" 

99 if value is None or (isinstance(value, str) and not value.strip()): 

100 return None 

101 if isinstance(value, bool) or not isinstance(value, (int, float, str)): 

102 raise TypeError(f"Azure Prompt Shield: price_per_1000_text_records must be a number, got {value!r}") 

103 try: 

104 price: Final = float(value) 

105 except ValueError as e: 

106 raise ValueError(f"Azure Prompt Shield: price_per_1000_text_records must be a number, got {value!r}") from e 

107 if not math.isfinite(price) or price < 0: 

108 raise ValueError( 

109 f"Azure Prompt Shield: price_per_1000_text_records must be a finite, non-negative number, got {value!r}" 

110 ) 

111 return price 

112 

113 

114class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrail): 

115 """ 

116 LiteLLM Built-in Guardrail for Azure Content Safety Guardrail (Prompt Shield). 

117 

118 This guardrail scans prompts and responses using the Azure Prompt Shield API to detect 

119 malicious content, injection attempts, and policy violations. 

120 

121 Configuration: 

122 guardrail_name: Name of the guardrail instance 

123 api_key: Azure Prompt Shield API key 

124 api_base: Azure Prompt Shield API endpoint 

125 default_on: Whether to enable by default 

126 """ 

127 

128 use_native_lifecycle_hooks: ClassVar[bool] = True 

129 

130 def __init__( 

131 self, 

132 guardrail_name: str, 

133 api_key: str, 

134 api_base: str, 

135 **kwargs, 

136 ): 

137 """Initialize Azure Prompt Shield guardrail handler.""" 

138 # AzureGuardrailBase.__init__ stores api_key, api_base, api_version, 

139 # async_handler and forwards the rest to CustomGuardrail. 

140 super().__init__( 

141 api_key=api_key, 

142 api_base=api_base, 

143 guardrail_name=guardrail_name, 

144 supported_event_hooks=list(self.get_supported_event_hooks()), 

145 **kwargs, 

146 ) 

147 

148 # Plain (non-Final) attributes: ``update_in_memory_litellm_params`` 

149 # re-resolves them when the guardrail is updated in place. 

150 self.cost_tier: str | None = _resolved_cost_tier(kwargs.get("cost_tier")) 

151 self.price_per_1000_text_records: float | None = _resolved_price( 

152 kwargs.get("price_per_1000_text_records"), self.cost_tier 

153 ) 

154 

155 verbose_proxy_logger.debug("Initialized Azure Prompt Shield Guardrail: %s", guardrail_name) 

156 

157 async def async_make_request( 

158 self, 

159 user_prompt: str, 

160 usage_accumulator: MutableMapping[str, int], # mutable-ok: callee-filled accumulator 

161 ) -> "AzurePromptShieldGuardrailResponse": 

162 """ 

163 Make a request to the Azure Prompt Shield API. 

164 

165 Long prompts are automatically split at word boundaries into chunks 

166 that respect the Azure Content Safety 10 000-character limit. Each 

167 chunk is analysed independently; an attack in *any* chunk raises 

168 an HTTPException immediately. 

169 

170 ``usage_accumulator`` collects billable usage per SUBMITTED chunk: 

171 ``requests`` (Azure API calls), ``input_characters``, and 

172 ``text_records`` (ceil(chunk_chars / 1000), Azure's billing unit). 

173 A chunk that triggers an intervention was still submitted and billed, 

174 so it is counted before the block is raised; chunks after it are 

175 never submitted and never counted. 

176 """ 

177 from litellm.types.proxy.guardrails.guardrail_hooks.azure.azure_prompt_shield import ( 

178 AzurePromptShieldGuardrailRequestBody, 

179 AzurePromptShieldGuardrailResponse, 

180 ) 

181 

182 from .base import AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH 

183 

184 chunks: Final = self.split_text_by_words(user_prompt, AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH) 

185 

186 last_response: AzurePromptShieldGuardrailResponse | None = None 

187 

188 for chunk in chunks: 

189 request_body = AzurePromptShieldGuardrailRequestBody(documents=[], userPrompt=chunk) 

190 response_json = await self._post_to_content_safety("text:shieldPrompt", cast(dict, request_body)) 

191 

192 last_response = cast(AzurePromptShieldGuardrailResponse, response_json) 

193 

194 usage_accumulator["requests"] = usage_accumulator.get("requests", 0) + 1 

195 usage_accumulator["input_characters"] = usage_accumulator.get("input_characters", 0) + len(chunk) 

196 usage_accumulator[AZURE_PROMPT_SHIELD_TEXT_RECORD_UNIT] = usage_accumulator.get( 

197 AZURE_PROMPT_SHIELD_TEXT_RECORD_UNIT, 0 

198 ) + math.ceil(len(chunk) / AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH) 

199 

200 if last_response["userPromptAnalysis"].get("attackDetected"): 

201 verbose_proxy_logger.warning( 

202 "Azure Prompt Shield: Attack detected in chunk of length %d", 

203 len(chunk), 

204 ) 

205 raise HTTPException( 

206 status_code=400, 

207 detail={ 

208 "error": "Violated Azure Prompt Shield guardrail policy", 

209 "detection_message": f"Attack detected: {last_response['userPromptAnalysis']}", 

210 }, 

211 ) 

212 

213 # chunks is always non-empty (split_text_by_words guarantees ≥1 element) 

214 assert last_response is not None 

215 return last_response 

216 

217 @log_guardrail_information 

218 async def apply_guardrail( 

219 self, 

220 inputs: GenericGuardrailAPIInputs, 

221 request_data: dict, 

222 input_type: Literal["request", "response"], 

223 logging_obj: "LiteLLMLoggingObj | None" = None, 

224 ) -> GenericGuardrailAPIInputs: 

225 _billing_usage_stash.set(None) 

226 usage: Final[dict[str, int]] = {} # mutable-ok: per-invocation billing accumulator 

227 try: 

228 for text in inputs.get("texts") or (): 

229 if text: 

230 await self.async_make_request(user_prompt=text, usage_accumulator=usage) 

231 finally: 

232 self._record_billing_usage(usage) 

233 return inputs 

234 

235 @log_guardrail_information 

236 async def async_pre_call_hook( 

237 self, 

238 user_api_key_dict: "UserAPIKeyAuth", 

239 cache: Any, 

240 data: dict[str, Any], 

241 call_type: CallTypesLiteral, 

242 ) -> dict[str, Any] | None: 

243 """ 

244 Pre-call hook to scan user prompts before sending to LLM. 

245 

246 Raises HTTPException if content should be blocked. 

247 """ 

248 _billing_usage_stash.set(None) 

249 verbose_proxy_logger.debug( 

250 "Azure Prompt Shield: Running pre-call prompt scan, on call_type: %s", 

251 call_type, 

252 ) 

253 new_messages: Final[list[AllMessageValues] | None] = data.get("messages") 

254 if new_messages is None: 

255 verbose_proxy_logger.warning("Azure Prompt Shield: not running guardrail. No messages in data") 

256 return data 

257 user_prompt: Final = self.get_user_prompt(new_messages) 

258 

259 if user_prompt: 

260 verbose_proxy_logger.debug("Azure Prompt Shield: User prompt: %s", user_prompt) 

261 usage: Final[dict[str, int]] = {} # mutable-ok: per-invocation billing accumulator 

262 try: 

263 await self.async_make_request( 

264 user_prompt=user_prompt, 

265 usage_accumulator=usage, 

266 ) 

267 finally: 

268 self._record_billing_usage(usage) 

269 else: 

270 verbose_proxy_logger.warning("Azure Prompt Shield: No user prompt found") 

271 return None 

272 

273 def update_in_memory_litellm_params(self, litellm_params: "LitellmParams | dict") -> None: # mutable-ok: DB dict 

274 """Apply updated params in place, re-resolving billing and credentials. 

275 

276 Pricing is read via ``_updated_param`` (the values are pydantic extras, and 

277 the immediate PUT sync hands this method the raw DB dict). Pricing and any 

278 ``os.environ/`` credential references are validated and resolved BEFORE any 

279 state is mutated, so an invalid update leaves the running guardrail 

280 untouched and a raw reference never overwrites a resolved credential. 

281 """ 

282 cost_tier: Final = _resolved_cost_tier(_updated_param(litellm_params, "cost_tier")) 

283 price: Final = _resolved_price(_updated_param(litellm_params, "price_per_1000_text_records"), cost_tier) 

284 resolved_credentials: dict[str, object] = {} # mutable-ok: staged before mutation 

285 for cred_key in ("api_key", "api_base"): 

286 cred_value = _updated_param(litellm_params, cred_key) 

287 if isinstance(cred_value, str) and cred_value.startswith("os.environ/"): 

288 resolved_credentials[cred_key] = _resolved_secret_value(cred_value) 

289 if isinstance(litellm_params, Mapping): 

290 for key, value in litellm_params.items(): 

291 setattr(self, key, resolved_credentials.get(key, value)) 

292 else: 

293 super().update_in_memory_litellm_params(litellm_params) 

294 for cred_key, cred_value in resolved_credentials.items(): 

295 setattr(self, cred_key, cred_value) 

296 self.cost_tier = cost_tier 

297 self.price_per_1000_text_records = price 

298 

299 def _record_billing_usage(self, usage: Mapping[str, int]) -> None: 

300 """Stash this invocation's usage counters for the ``_process_*`` call the 

301 decorator runs next in the same asyncio task; overwrites any leftover.""" 

302 _billing_usage_stash.set(dict(usage) if usage else None) # mutable-ok: fresh snapshot, popped by _process_* 

303 

304 def _pop_billing_tracing_detail(self) -> GuardrailTracingDetail | None: 

305 """Build the billing tracing detail from the stashed usage counters, priced 

306 with the configured tier/price. ``guardrail_cost_in_spend=False`` keeps the 

307 estimated cost out of ``response_cost`` and budget enforcement: Azure 

308 guardrail cost is reported on logs, OTEL spans, and the UI, never billed 

309 against team/user/key budgets (LIT-5917).""" 

310 usage: Final = _billing_usage_stash.get() 

311 _billing_usage_stash.set(None) 

312 if not usage: 

313 return None 

314 cost: Final = azure_prompt_shield_guardrail_cost( 

315 usage_units=usage, 

316 cost_tier=self.cost_tier, 

317 price_per_1000_text_records=self.price_per_1000_text_records, 

318 ) 

319 if cost is None: 

320 return GuardrailTracingDetail(guardrail_usage=usage) 

321 return GuardrailTracingDetail( 

322 guardrail_usage=usage, 

323 guardrail_cost=cost, 

324 guardrail_cost_in_spend=False, 

325 ) 

326 

327 def _process_response( 

328 self, 

329 response: dict | None, # mutable-ok: matches CustomGuardrail._process_response signature 

330 request_data: dict, # mutable-ok: matches CustomGuardrail._process_response signature 

331 start_time: float | None = None, 

332 end_time: float | None = None, 

333 duration: float | None = None, 

334 event_type: GuardrailEventHooks | None = None, 

335 original_inputs: dict | None = None, # mutable-ok: matches CustomGuardrail._process_response signature 

336 ) -> dict | None: # mutable-ok: matches CustomGuardrail._process_response return 

337 """Override to attach the Azure billing tracing detail (usage counters and 

338 estimated cost) and the ``azure`` provider label to the recorded guardrail 

339 information. Follows the OpenAI moderation override pattern 

340 (openai/moderations.py).""" 

341 guardrail_response: Final = self._summarize_guardrail_response( 

342 response=response, 

343 original_inputs=original_inputs, 

344 event_type=event_type, 

345 ) 

346 self.add_standard_logging_guardrail_information_to_request_data( 

347 guardrail_json_response=guardrail_response, 

348 request_data=request_data, 

349 guardrail_status="success", 

350 duration=duration, 

351 start_time=start_time, 

352 end_time=end_time, 

353 event_type=event_type, 

354 guardrail_provider="azure", 

355 tracing_detail=self._pop_billing_tracing_detail(), 

356 ) 

357 return response 

358 

359 def _process_error( 

360 self, 

361 e: Exception, 

362 request_data: dict, # mutable-ok: matches CustomGuardrail._process_error signature 

363 start_time: float | None = None, 

364 end_time: float | None = None, 

365 duration: float | None = None, 

366 event_type: GuardrailEventHooks | None = None, 

367 ) -> NoReturn: 

368 """Override to attach the Azure billing tracing detail to the blocked/error 

369 guardrail record; a chunk that triggered an intervention was still submitted 

370 to (and billed by) Azure, so its usage is recorded on this path too.""" 

371 guardrail_status: Final = ( 

372 "guardrail_intervened" if self._is_guardrail_intervention(e) else "guardrail_failed_to_respond" 

373 ) 

374 self.add_standard_logging_guardrail_information_to_request_data( 

375 guardrail_json_response=e, 

376 request_data=request_data, 

377 guardrail_status=guardrail_status, 

378 duration=duration, 

379 start_time=start_time, 

380 end_time=end_time, 

381 event_type=event_type, 

382 guardrail_provider="azure", 

383 tracing_detail=self._pop_billing_tracing_detail(), 

384 ) 

385 raise e 

386 

387 @staticmethod 

388 def get_config_model() -> type["GuardrailConfigModel"] | None: 

389 """ 

390 Get the config model for the Azure Prompt Shield guardrail. 

391 """ 

392 from litellm.types.proxy.guardrails.guardrail_hooks.azure.azure_prompt_shield import ( 

393 AzurePromptShieldGuardrailConfigModel, 

394 ) 

395 

396 return AzurePromptShieldGuardrailConfigModel 

397 

398 @classmethod 

399 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: 

400 return [ 

401 GuardrailEventHooks.pre_call, 

402 GuardrailEventHooks.during_call, 

403 ]