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
« 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"""
6import math
7from collections.abc import Mapping, MutableMapping
8from contextvars import ContextVar
9from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, NoReturn, cast
11from fastapi import HTTPException
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)
30from .base import AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH, AzureGuardrailBase
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
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)
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
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)
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
84def _resolved_price(raw: object, cost_tier: str | None) -> float | None:
85 """Normalize price_per_1000_text_records and validate it against the tier.
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
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
114class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrail):
115 """
116 LiteLLM Built-in Guardrail for Azure Content Safety Guardrail (Prompt Shield).
118 This guardrail scans prompts and responses using the Azure Prompt Shield API to detect
119 malicious content, injection attempts, and policy violations.
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 """
128 use_native_lifecycle_hooks: ClassVar[bool] = True
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 )
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 )
155 verbose_proxy_logger.debug("Initialized Azure Prompt Shield Guardrail: %s", guardrail_name)
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.
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.
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 )
182 from .base import AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH
184 chunks: Final = self.split_text_by_words(user_prompt, AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH)
186 last_response: AzurePromptShieldGuardrailResponse | None = None
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))
192 last_response = cast(AzurePromptShieldGuardrailResponse, response_json)
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)
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 )
213 # chunks is always non-empty (split_text_by_words guarantees ≥1 element)
214 assert last_response is not None
215 return last_response
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
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.
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)
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
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.
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
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_*
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 )
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
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
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 )
396 return AzurePromptShieldGuardrailConfigModel
398 @classmethod
399 def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
400 return [
401 GuardrailEventHooks.pre_call,
402 GuardrailEventHooks.during_call,
403 ]