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

1import threading 

2import time 

3import uuid 

4from collections import OrderedDict 

5from collections.abc import Mapping, Sequence 

6from typing import TYPE_CHECKING, Any, Final 

7 

8from typing_extensions import NotRequired, TypedDict 

9 

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) 

19 

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 

23 

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" 

27 

28# Protection scope cache TTL in seconds (1 hour, per Microsoft recommendation). 

29SCOPE_CACHE_TTL_SECONDS: Final = 3600.0 

30 

31 

32class GraphTokenResponse(TypedDict): 

33 access_token: str 

34 expires_in: NotRequired[int] 

35 

36 

37class PurviewGuardrailBase: 

38 """ 

39 Base class for Microsoft Purview guardrails. 

40 

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

45 

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) 

58 

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 

65 

66 # Token cache: (access_token, expires_at_epoch) 

67 self._token_cache: tuple[str, float] | None = None 

68 

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

82 

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

87 

88 # ------------------------------------------------------------------ 

89 # OAuth2 token management 

90 # ------------------------------------------------------------------ 

91 

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] 

98 

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 

121 

122 # ------------------------------------------------------------------ 

123 # Graph API helpers 

124 # ------------------------------------------------------------------ 

125 

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. 

133 

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) 

144 

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 

152 

153 # ------------------------------------------------------------------ 

154 # Protection scopes 

155 # ------------------------------------------------------------------ 

156 

157 async def _compute_protection_scopes(self, user_id: str) -> tuple[str, Mapping[str, object]]: 

158 """Call protectionScopes/compute and cache with ETag. 

159 

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

165 

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] 

171 

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 } 

182 

183 response_json, response_headers = await self._graph_post(url, body) 

184 etag: Final = response_headers.get("etag", response_headers.get("ETag", "")) 

185 

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 

199 

200 # ------------------------------------------------------------------ 

201 # Process content 

202 # ------------------------------------------------------------------ 

203 

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. 

213 

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 } 

255 

256 extra_headers: Final[dict[str, str]] = {} 

257 if etag: 

258 extra_headers["If-None-Match"] = etag 

259 

260 response_json, _ = await self._graph_post(url, body, extra_headers) 

261 

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) 

266 

267 return response_json 

268 

269 # ------------------------------------------------------------------ 

270 # User ID resolution 

271 # ------------------------------------------------------------------ 

272 

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. 

275 

276 Returns the strongest available identity walking down four sources, in 

277 decreasing trust order: 

278 

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 

283 

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 

292 

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) 

295 

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) 

303 

304 uid = metadata.get(self.user_id_field) 

305 if uid: 

306 return str(uid) 

307 

308 return None 

309 

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 {} 

318 

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. 

321 

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``. 

327 

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. 

331 

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) 

337 

338 return None 

339 

340 def _resolve_user_id_from_logging_kwargs(self, kwargs: Mapping[str, object]) -> str | None: 

341 """Trusted-identity-only resolver for logging-only hooks. 

342 

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 

357 

358 # ------------------------------------------------------------------ 

359 # Policy action evaluation 

360 # ------------------------------------------------------------------ 

361 

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

368 

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 

374 

375 # ------------------------------------------------------------------ 

376 # Prompt text for DLP 

377 # ------------------------------------------------------------------ 

378 

379 @staticmethod 

380 def is_token_id_prompt(prompt: str | Sequence[object] | None) -> bool: 

381 """Return True if ``prompt`` carries OpenAI completions token ids. 

382 

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 

396 

397 @staticmethod 

398 def completion_prompt_to_str(prompt: str | Sequence[object] | None) -> str | None: 

399 """Normalize OpenAI ``/v1/completions`` ``prompt`` for text DLP. 

400 

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 

421 

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. 

425 

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

432 

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) 

445 

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) 

458 

459 return args 

460 

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. 

463 

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. 

469 

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