Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py: 15%
202 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
1import threading
2import time
3import uuid
4from collections import OrderedDict
5from collections.abc import Mapping, Sequence
6from typing import TYPE_CHECKING, Any, Final
8from typing_extensions import NotRequired, TypedDict
10from litellm._logging import verbose_proxy_logger
11from litellm.litellm_core_utils.prompt_templates.common_utils import (
12 convert_content_list_to_str,
13)
14from litellm.litellm_core_utils.url_utils import encode_url_path_segment
15from litellm.llms.custom_httpx.http_handler import (
16 get_async_httpx_client,
17 httpxSpecialProvider,
18)
20if TYPE_CHECKING: 20 ↛ 21line 20 didn't jump to line 21 because the condition on line 20 was never true
21 from litellm.proxy._types import UserAPIKeyAuth
22 from litellm.types.llms.openai import AllMessageValues
24GRAPH_API_BASE: Final = "https://graph.microsoft.com/v1.0"
25TOKEN_ENDPOINT_TEMPLATE: Final = "https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token"
26GRAPH_SCOPE: Final = "https://graph.microsoft.com/.default"
28# Protection scope cache TTL in seconds (1 hour, per Microsoft recommendation).
29SCOPE_CACHE_TTL_SECONDS: Final = 3600.0
32class GraphTokenResponse(TypedDict):
33 access_token: str
34 expires_in: NotRequired[int]
37class PurviewGuardrailBase:
38 """
39 Base class for Microsoft Purview guardrails.
41 Manages OAuth2 client-credentials token acquisition, protection scope
42 computation with ETag caching, and authenticated POST calls to the
43 Microsoft Graph API.
44 """
46 def __init__(
47 self,
48 tenant_id: str,
49 client_id: str,
50 client_secret: str,
51 purview_app_name: str = "LiteLLM",
52 user_id_field: str = "user_id",
53 **kwargs: object,
54 ) -> None:
55 # Forward remaining kwargs to the next class in the MRO
56 # (typically CustomGuardrail).
57 super().__init__(**kwargs)
59 self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
60 self.tenant_id = tenant_id
61 self.client_id = client_id
62 self.client_secret = client_secret
63 self.purview_app_name = purview_app_name
64 self.user_id_field = user_id_field
66 # Token cache: (access_token, expires_at_epoch)
67 self._token_cache: tuple[str, float] | None = None
69 # Protection scope cache: user_id -> (etag, scope_response, fetched_at)
70 # Capped at 1000 entries (LRU eviction) to avoid unbounded growth.
71 self._scope_cache: OrderedDict[str, tuple[str, Mapping[str, object], float]] = OrderedDict()
72 self._scope_cache_maxsize = 1000
73 # Use a threading.Lock (not asyncio.Lock) because this lock is acquired
74 # from both the proxy's main asyncio event loop and from short-lived
75 # event loops created by the logging_hook thread fallback. In Python
76 # 3.10+ an asyncio.Lock is bound to the first event loop that acquires
77 # it and raises RuntimeError from any other loop, which would silently
78 # break audit logging via the thread fallback. All critical sections
79 # below are pure in-memory dict ops with no awaits, so a synchronous
80 # lock is both correct and sufficient.
81 self._cache_lock = threading.Lock()
83 @staticmethod
84 def _encode_graph_user_id(user_id: str) -> str:
85 """Percent-encode Entra user id for Graph ``/users/{id}/...`` path segments."""
86 return encode_url_path_segment(user_id, field_name="user_id")
88 # ------------------------------------------------------------------
89 # OAuth2 token management
90 # ------------------------------------------------------------------
92 async def _get_access_token(self) -> str:
93 """Acquire or return cached OAuth2 token via client_credentials grant."""
94 now: Final = time.time()
95 with self._cache_lock:
96 if self._token_cache and self._token_cache[1] > now + 60:
97 return self._token_cache[0]
99 url: Final = TOKEN_ENDPOINT_TEMPLATE.format(tenant_id=self.tenant_id)
100 data: Final = {
101 "grant_type": "client_credentials",
102 "client_id": self.client_id,
103 "client_secret": self.client_secret,
104 "scope": GRAPH_SCOPE,
105 }
106 response: Final = await self.async_handler.post(
107 url=url,
108 data=data,
109 headers={"Content-Type": "application/x-www-form-urlencoded"},
110 )
111 response.raise_for_status()
112 token_data: Final[GraphTokenResponse] = response.json()
113 access_token: Final = token_data["access_token"]
114 expires_in: Final = int(token_data.get("expires_in", 3599))
115 # Recompute ``now`` after the await so the expiry reflects when the
116 # token was actually received, not when the request started.
117 with self._cache_lock:
118 self._token_cache = (access_token, time.time() + expires_in)
119 verbose_proxy_logger.debug("Purview: acquired new OAuth2 token (expires_in=%ds)", expires_in)
120 return access_token
122 # ------------------------------------------------------------------
123 # Graph API helpers
124 # ------------------------------------------------------------------
126 async def _graph_post(
127 self,
128 url: str,
129 json_body: dict[str, object],
130 extra_headers: Mapping[str, str] | None = None,
131 ) -> tuple[dict[str, object], dict[str, str]]:
132 """POST to Graph API with bearer auth.
134 Returns:
135 Tuple of (response_json, response_headers).
136 """
137 token: Final = await self._get_access_token()
138 headers: Final = {
139 "Authorization": f"Bearer {token}",
140 "Content-Type": "application/json",
141 }
142 if extra_headers:
143 headers.update(extra_headers)
145 verbose_proxy_logger.debug("Purview Graph POST %s", url)
146 response: Final = await self.async_handler.post(url=url, headers=headers, json=json_body)
147 response.raise_for_status()
148 response_json: Final[dict[str, object]] = response.json()
149 response_headers: Final = dict(response.headers)
150 verbose_proxy_logger.debug("Purview Graph response: %s", response_json)
151 return response_json, response_headers
153 # ------------------------------------------------------------------
154 # Protection scopes
155 # ------------------------------------------------------------------
157 async def _compute_protection_scopes(self, user_id: str) -> tuple[str, Mapping[str, object]]:
158 """Call protectionScopes/compute and cache with ETag.
160 Returns:
161 Tuple of (etag, scope_response).
162 """
163 encoded_user_id: Final = self._encode_graph_user_id(user_id)
164 now: Final = time.time()
166 with self._cache_lock:
167 cached: Final = self._scope_cache.get(user_id)
168 if cached and (now - cached[2]) < SCOPE_CACHE_TTL_SECONDS:
169 self._scope_cache.move_to_end(user_id)
170 return cached[0], cached[1]
172 url: Final = f"{GRAPH_API_BASE}/users/{encoded_user_id}/dataSecurityAndGovernance/protectionScopes/compute"
173 body: Final[dict[str, object]] = {
174 "activities": "uploadText,downloadText",
175 "locations": [
176 {
177 "@odata.type": "microsoft.graph.policyLocationApplication",
178 "value": self.client_id,
179 }
180 ],
181 }
183 response_json, response_headers = await self._graph_post(url, body)
184 etag: Final = response_headers.get("etag", response_headers.get("ETag", ""))
186 # Recompute ``now`` after the await so the TTL reflects when the
187 # scope response was actually received, not when the request started.
188 fetched_at: Final = time.time()
189 with self._cache_lock:
190 self._scope_cache[user_id] = (etag, response_json, fetched_at)
191 # Move refreshed entry to the end so it is treated as most-recently-used.
192 # OrderedDict.__setitem__ preserves existing insertion order for known
193 # keys, so an explicit move_to_end() call is required.
194 self._scope_cache.move_to_end(user_id)
195 # Evict least-recently-used entry when cache exceeds max size.
196 while len(self._scope_cache) > self._scope_cache_maxsize:
197 self._scope_cache.popitem(last=False)
198 return etag, response_json
200 # ------------------------------------------------------------------
201 # Process content
202 # ------------------------------------------------------------------
204 async def _process_content(
205 self,
206 user_id: str,
207 text: str,
208 activity: str,
209 etag: str,
210 correlation_id: str | None = None,
211 ) -> dict[str, object]:
212 """Call processContent for DLP policy evaluation.
214 Args:
215 user_id: Entra object ID of the user.
216 text: The content to evaluate.
217 activity: ``"uploadText"`` for prompts, ``"downloadText"`` for responses.
218 etag: Cached ETag from protectionScopes/compute.
219 correlation_id: Optional conversation/thread ID.
220 """
221 encoded_user_id: Final = self._encode_graph_user_id(user_id)
222 url: Final = f"{GRAPH_API_BASE}/users/{encoded_user_id}/dataSecurityAndGovernance/processContent"
223 body: Final[dict[str, object]] = {
224 "contentToProcess": {
225 "contentEntries": [
226 {
227 "@odata.type": "microsoft.graph.processConversationMetadata",
228 "identifier": str(uuid.uuid4()),
229 "content": {
230 "@odata.type": "microsoft.graph.textContent",
231 "data": text,
232 },
233 "name": f"{self.purview_app_name} message",
234 "correlationId": correlation_id or str(uuid.uuid4()),
235 "sequenceNumber": 0,
236 "isTruncated": False,
237 }
238 ],
239 "activityMetadata": {"activity": activity},
240 "deviceMetadata": {},
241 "protectedAppMetadata": {
242 "name": self.purview_app_name,
243 "version": "1.0",
244 "applicationLocation": {
245 "@odata.type": "microsoft.graph.policyLocationApplication",
246 "value": self.client_id,
247 },
248 },
249 "integratedAppMetadata": {
250 "name": self.purview_app_name,
251 "version": "1.0",
252 },
253 }
254 }
256 extra_headers: Final[dict[str, str]] = {}
257 if etag:
258 extra_headers["If-None-Match"] = etag
260 response_json, _ = await self._graph_post(url, body, extra_headers)
262 # If policies changed, invalidate scope cache so next call re-fetches.
263 if response_json.get("protectionScopeState") == "modified":
264 with self._cache_lock:
265 self._scope_cache.pop(user_id, None)
267 return response_json
269 # ------------------------------------------------------------------
270 # User ID resolution
271 # ------------------------------------------------------------------
273 def _resolve_user_id(self, data: Mapping[str, object], user_api_key_dict: "UserAPIKeyAuth") -> str | None:
274 """Resolve the Entra user object ID from request data or auth context.
276 Returns the strongest available identity walking down four sources, in
277 decreasing trust order:
279 1. ``user_api_key_dict.user_id`` — LiteLLM key / JWT-bound user
280 2. ``user_api_key_dict.end_user_id`` — request-derived
281 3. ``metadata["user_api_key_user_id"]`` — proxy-injected from the key
282 4. ``metadata[user_id_field]`` — caller-supplied
284 Used only by blocking-mode resolution to disambiguate "no identity at
285 all" from "caller supplied an untrusted identity" for the error
286 message. Neither blocking nor audit DLP feeds the untrusted
287 fallbacks (2, 4) into Purview itself.
288 """
289 trusted: Final = self._resolve_trusted_user_id(data, user_api_key_dict)
290 if trusted:
291 return trusted
293 if hasattr(user_api_key_dict, "end_user_id") and user_api_key_dict.end_user_id:
294 return str(user_api_key_dict.end_user_id)
296 metadata_value: Final[object] = data.get("metadata") or data.get("litellm_metadata") or {}
297 if not isinstance(metadata_value, Mapping):
298 return None
299 metadata: Final[Mapping[str, object]] = metadata_value
300 uid = metadata.get("user_api_key_user_id")
301 if uid:
302 return str(uid)
304 uid = metadata.get(self.user_id_field)
305 if uid:
306 return str(uid)
308 return None
310 @staticmethod
311 def _logging_kwargs_metadata(kwargs: Mapping[str, object]) -> Mapping[str, object]:
312 """Metadata dict from ``model_call_details`` / logging kwargs."""
313 litellm_params: Final[object] = kwargs.get("litellm_params") or {}
314 if not isinstance(litellm_params, dict):
315 return {}
316 md: Final = litellm_params.get("metadata")
317 return md if isinstance(md, dict) else {}
319 def _resolve_trusted_user_id(self, data: Mapping[str, object], user_api_key_dict: "UserAPIKeyAuth") -> str | None:
320 """Resolve user ID from API-key/JWT-bound identity for blocking DLP.
322 Uses only ``UserAPIKeyAuth.user_id`` (bound on the LiteLLM key or JWT).
323 Intentionally omits ``UserAPIKeyAuth.end_user_id`` because the proxy sets
324 it from caller-controlled request fields (``user``, ``metadata.user_id``,
325 ``safety_identifier``, custom headers, etc.) via
326 ``get_end_user_id_from_request_body``.
328 Also omits ``metadata[user_id_field]`` and
329 ``metadata["user_api_key_user_id"]`` for the same impersonation risk when
330 the key has no bound user.
332 Returns ``None`` when no authenticated identity is available. Blocking
333 hooks must fail closed rather than skip the DLP check.
334 """
335 if hasattr(user_api_key_dict, "user_id") and user_api_key_dict.user_id:
336 return str(user_api_key_dict.user_id)
338 return None
340 def _resolve_user_id_from_logging_kwargs(self, kwargs: Mapping[str, object]) -> str | None:
341 """Trusted-identity-only resolver for logging-only hooks.
343 Uses only the proxy-injected ``user_api_key_user_id`` (populated from
344 the API-key/JWT-bound ``UserAPIKeyAuth.user_id`` after the proxy
345 strips every caller-supplied ``user_api_key_*`` key from the request
346 metadata). Caller-influenceable sources (``user_api_key_end_user_id``,
347 ``metadata[user_id_field]``) are not used here so a caller cannot
348 cause Purview audit records to be written under a victim's identity.
349 Returns ``None`` when no trusted identity is available so the audit
350 is skipped rather than misattributed.
351 """
352 md: Final = self._logging_kwargs_metadata(kwargs)
353 uid: Final = md.get("user_api_key_user_id") or kwargs.get("user_api_key_user_id")
354 if uid:
355 return str(uid)
356 return None
358 # ------------------------------------------------------------------
359 # Policy action evaluation
360 # ------------------------------------------------------------------
362 @staticmethod
363 def _should_block(response: dict[str, Any]) -> bool:
364 """Return True if any policyAction requires blocking."""
365 for action in response.get("policyActions", []):
366 odata_type = action.get("@odata.type", "")
367 action_field = action.get("action", "")
369 if "restrictAccessAction" in odata_type or action_field == "restrictAccess":
370 restriction = action.get("restrictionAction", "")
371 if restriction == "block":
372 return True
373 return False
375 # ------------------------------------------------------------------
376 # Prompt text for DLP
377 # ------------------------------------------------------------------
379 @staticmethod
380 def is_token_id_prompt(prompt: str | Sequence[object] | None) -> bool:
381 """Return True if ``prompt`` carries OpenAI completions token ids.
383 Covers every list shape that ``completion_prompt_to_str`` cannot decode
384 for Purview, including flat ``list[int]`` (single token-id prompt),
385 ``list[list[int]]`` (multi-prompt token-id batches), and mixed lists
386 that include any token-id sub-array.
387 """
388 if not isinstance(prompt, list) or not prompt:
389 return False
390 for x in prompt:
391 if isinstance(x, int):
392 return True
393 if isinstance(x, list) and x and any(isinstance(y, int) for y in x):
394 return True
395 return False
397 @staticmethod
398 def completion_prompt_to_str(prompt: str | Sequence[object] | None) -> str | None:
399 """Normalize OpenAI ``/v1/completions`` ``prompt`` for text DLP.
401 Supports string prompts and list-of-string prompts. List-of-token-id prompts
402 are skipped (no plaintext for Purview to evaluate).
403 """
404 if prompt is None:
405 return None
406 if isinstance(prompt, str):
407 stripped: Final = prompt.strip()
408 return stripped or None
409 if isinstance(prompt, list) and prompt:
410 if all(isinstance(x, str) for x in prompt):
411 joined = "\n".join(s.strip() for s in prompt if isinstance(s, str))
412 return joined.strip() or None
413 if all(isinstance(x, int) for x in prompt):
414 verbose_proxy_logger.debug("Purview DLP: completions prompt is token ids only; skipping text scan")
415 return None
416 str_parts: Final = [x for x in prompt if isinstance(x, str)]
417 if str_parts:
418 joined = "\n".join(s.strip() for s in str_parts)
419 return joined.strip() or None
420 return None
422 @staticmethod
423 def _extract_tool_call_args_from_message(message: object) -> list[str]:
424 """Return plaintext arguments strings from tool_calls and function_call fields.
426 Covers both the request path (assistant messages in chat histories that
427 carry tool_calls / function_call) and the response path (model-generated
428 tool calls returned in a ModelResponse). Both dict-style and object-style
429 representations are handled.
430 """
431 args: Final[list[str]] = []
433 # tool_calls: [{"function": {"arguments": "..."}}]
434 tool_calls: Final[Sequence[object] | None] = (
435 message.get("tool_calls") if isinstance(message, dict) else getattr(message, "tool_calls", None)
436 )
437 if tool_calls:
438 for tc in tool_calls:
439 fn = tc.get("function") if isinstance(tc, dict) else getattr(tc, "function", None)
440 if fn is None:
441 continue
442 arguments = fn.get("arguments") if isinstance(fn, dict) else getattr(fn, "arguments", None)
443 if isinstance(arguments, str) and arguments.strip():
444 args.append(arguments)
446 # Legacy function_call: {"arguments": "..."}
447 function_call: Final = (
448 message.get("function_call") if isinstance(message, dict) else getattr(message, "function_call", None)
449 )
450 if function_call is not None:
451 arguments = (
452 function_call.get("arguments")
453 if isinstance(function_call, dict)
454 else getattr(function_call, "arguments", None)
455 )
456 if isinstance(arguments, str) and arguments.strip():
457 args.append(arguments)
459 return args
461 def get_prompt_text_for_dlp(self, messages: list["AllMessageValues"]) -> str | None:
462 """Concatenate text from every chat message (all roles) for pre-call DLP.
464 Evaluates the same payload the model receives, not only the trailing user
465 turn. Each message is separated by ``\\n\\n`` so that tokens at message
466 boundaries are not merged (e.g., ``"end of msg1\\n\\nstart of msg2"``
467 rather than ``"end of msg1start of msg2"``), which preserves DLP pattern
468 detection accuracy across message boundaries.
470 Tool-call arguments (``tool_calls[].function.arguments`` and
471 ``function_call.arguments``) are included alongside message content so
472 that sensitive data hidden in function arguments is not bypassed.
473 """
474 if not messages:
475 return None
476 parts: Final[list[str]] = []
477 for msg in messages:
478 segments: list[str] = []
479 content = convert_content_list_to_str(message=msg).strip()
480 if content:
481 segments.append(content)
482 segments.extend(self._extract_tool_call_args_from_message(msg))
483 combined = "\n".join(segments)
484 if combined.strip():
485 parts.append(combined.strip())
486 text: Final = "\n\n".join(parts)
487 return text or None