diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/__init__.py new file mode 100644 index 00000000000..273089d4b4a --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/__init__.py @@ -0,0 +1,60 @@ +from typing import TYPE_CHECKING + +from litellm.types.guardrails import SupportedGuardrailIntegrations + +from .purview_dlp import MicrosoftPurviewDLPGuardrail + +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + + +def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"): + import litellm + from litellm.types.guardrails import GuardrailEventHooks + + tenant_id = litellm_params.get("tenant_id") + client_id = litellm_params.get("client_id") + + # client_secret can be passed via the standard api_key field or as + # a dedicated client_secret parameter. + client_secret = litellm_params.api_key or litellm_params.get("client_secret") + + if not tenant_id: + raise ValueError("Microsoft Purview: tenant_id is required") + if not client_id: + raise ValueError("Microsoft Purview: client_id is required") + if not client_secret: + raise ValueError("Microsoft Purview: client_secret (or api_key) is required") + + guardrail_name = guardrail.get("guardrail_name") + if not guardrail_name: + raise ValueError("Microsoft Purview: guardrail_name is required") + + mode = litellm_params.mode + logging_only = False + if isinstance(mode, str) and mode == GuardrailEventHooks.logging_only.value: + logging_only = True + + purview_guardrail = MicrosoftPurviewDLPGuardrail( + guardrail_name=guardrail_name, + tenant_id=str(tenant_id), + client_id=str(client_id), + client_secret=str(client_secret), + purview_app_name=str(litellm_params.get("purview_app_name") or "LiteLLM"), + user_id_field=str(litellm_params.get("user_id_field") or "user_id"), + logging_only=logging_only, + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, + ) + + litellm.logging_callback_manager.add_litellm_callback(purview_guardrail) + return purview_guardrail + + +guardrail_initializer_registry = { + SupportedGuardrailIntegrations.MICROSOFT_PURVIEW.value: initialize_guardrail, +} + +guardrail_class_registry = { + SupportedGuardrailIntegrations.MICROSOFT_PURVIEW.value: MicrosoftPurviewDLPGuardrail, +} diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py new file mode 100644 index 00000000000..c0b2021728a --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/base.py @@ -0,0 +1,303 @@ +import time +import uuid +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple + +from litellm._logging import verbose_proxy_logger +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) + +if TYPE_CHECKING: + from litellm.types.llms.openai import AllMessageValues + +GRAPH_API_BASE = "https://graph.microsoft.com/v1.0" +TOKEN_ENDPOINT_TEMPLATE = ( + "https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token" +) +GRAPH_SCOPE = "https://graph.microsoft.com/.default" + +# Protection scope cache TTL in seconds (1 hour, per Microsoft recommendation). +SCOPE_CACHE_TTL_SECONDS = 3600.0 + + +class PurviewGuardrailBase: + """ + Base class for Microsoft Purview guardrails. + + Manages OAuth2 client-credentials token acquisition, protection scope + computation with ETag caching, and authenticated POST calls to the + Microsoft Graph API. + """ + + def __init__( + self, + tenant_id: str, + client_id: str, + client_secret: str, + purview_app_name: str = "LiteLLM", + user_id_field: str = "user_id", + **kwargs: Any, + ): + # Forward remaining kwargs to the next class in the MRO + # (typically CustomGuardrail). + super().__init__(**kwargs) + + self.async_handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback + ) + self.tenant_id = tenant_id + self.client_id = client_id + self.client_secret = client_secret + self.purview_app_name = purview_app_name + self.user_id_field = user_id_field + + # Token cache: (access_token, expires_at_epoch) + self._token_cache: Optional[Tuple[str, float]] = None + + # Protection scope cache: user_id -> (etag, scope_response, fetched_at) + self._scope_cache: Dict[str, Tuple[str, Dict, float]] = {} + + # ------------------------------------------------------------------ + # OAuth2 token management + # ------------------------------------------------------------------ + + async def _get_access_token(self) -> str: + """Acquire or return cached OAuth2 token via client_credentials grant.""" + now = time.time() + if self._token_cache and self._token_cache[1] > now + 60: + return self._token_cache[0] + + url = TOKEN_ENDPOINT_TEMPLATE.format(tenant_id=self.tenant_id) + data = { + "grant_type": "client_credentials", + "client_id": self.client_id, + "client_secret": self.client_secret, + "scope": GRAPH_SCOPE, + } + response = await self.async_handler.post( + url=url, + data=data, + headers={"Content-Type": "application/x-www-form-urlencoded"}, + ) + token_data = response.json() + access_token = token_data["access_token"] + expires_in = int(token_data.get("expires_in", 3599)) + self._token_cache = (access_token, now + expires_in) + verbose_proxy_logger.debug( + "Purview: acquired new OAuth2 token (expires_in=%ds)", expires_in + ) + return access_token + + # ------------------------------------------------------------------ + # Graph API helpers + # ------------------------------------------------------------------ + + async def _graph_post( + self, + url: str, + json_body: Dict[str, Any], + extra_headers: Optional[Dict[str, str]] = None, + ) -> Tuple[Dict[str, Any], Dict[str, str]]: + """POST to Graph API with bearer auth. + + Returns: + Tuple of (response_json, response_headers). + """ + token = await self._get_access_token() + headers = { + "Authorization": f"Bearer {token}", + "Content-Type": "application/json", + } + if extra_headers: + headers.update(extra_headers) + + verbose_proxy_logger.debug("Purview Graph POST %s", url) + response = await self.async_handler.post( + url=url, headers=headers, json=json_body + ) + response_json: Dict[str, Any] = response.json() + response_headers = dict(response.headers) + verbose_proxy_logger.debug("Purview Graph response: %s", response_json) + return response_json, response_headers + + # ------------------------------------------------------------------ + # Protection scopes + # ------------------------------------------------------------------ + + async def _compute_protection_scopes( + self, user_id: str + ) -> Tuple[str, Dict[str, Any]]: + """Call protectionScopes/compute and cache with ETag. + + Returns: + Tuple of (etag, scope_response). + """ + cached = self._scope_cache.get(user_id) + now = time.time() + + if cached and (now - cached[2]) < SCOPE_CACHE_TTL_SECONDS: + return cached[0], cached[1] + + url = ( + f"{GRAPH_API_BASE}/users/{user_id}" + "/dataSecurityAndGovernance/protectionScopes/compute" + ) + body: Dict[str, Any] = { + "activities": "uploadText,downloadText", + "locations": [ + { + "@odata.type": "microsoft.graph.policyLocationApplication", + "value": self.client_id, + } + ], + } + + response_json, response_headers = await self._graph_post(url, body) + etag = response_headers.get("etag", response_headers.get("ETag", "")) + + self._scope_cache[user_id] = (etag, response_json, now) + return etag, response_json + + # ------------------------------------------------------------------ + # Process content + # ------------------------------------------------------------------ + + async def _process_content( + self, + user_id: str, + text: str, + activity: str, + etag: str, + correlation_id: Optional[str] = None, + ) -> Dict[str, Any]: + """Call processContent for DLP policy evaluation. + + Args: + user_id: Entra object ID of the user. + text: The content to evaluate. + activity: ``"uploadText"`` for prompts, ``"downloadText"`` for responses. + etag: Cached ETag from protectionScopes/compute. + correlation_id: Optional conversation/thread ID. + """ + url = ( + f"{GRAPH_API_BASE}/users/{user_id}" + "/dataSecurityAndGovernance/processContent" + ) + body: Dict[str, Any] = { + "contentToProcess": { + "contentEntries": [ + { + "@odata.type": "microsoft.graph.processConversationMetadata", + "identifier": str(uuid.uuid4()), + "content": { + "@odata.type": "microsoft.graph.textContent", + "data": text, + }, + "name": f"{self.purview_app_name} message", + "correlationId": correlation_id or str(uuid.uuid4()), + "sequenceNumber": 0, + "isTruncated": False, + } + ], + "activityMetadata": {"activity": activity}, + "deviceMetadata": {}, + "protectedAppMetadata": { + "name": self.purview_app_name, + "version": "1.0", + "applicationLocation": { + "@odata.type": "microsoft.graph.policyLocationApplication", + "value": self.client_id, + }, + }, + "integratedAppMetadata": { + "name": self.purview_app_name, + "version": "1.0", + }, + } + } + + extra_headers: Dict[str, str] = {} + if etag: + extra_headers["If-None-Match"] = etag + + response_json, _ = await self._graph_post(url, body, extra_headers) + + # If policies changed, invalidate scope cache so next call re-fetches. + if response_json.get("protectionScopeState") == "modified": + self._scope_cache.pop(user_id, None) + + return response_json + + # ------------------------------------------------------------------ + # User ID resolution + # ------------------------------------------------------------------ + + def _resolve_user_id( + self, data: Dict[str, Any], user_api_key_dict: Any + ) -> Optional[str]: + """Resolve the Entra user object ID from request data or auth context. + + Resolution order: + 1. ``metadata[user_id_field]`` (explicit per-request mapping) + 2. ``user_api_key_dict.user_id`` + 3. ``user_api_key_dict.end_user_id`` + """ + metadata = data.get("metadata") or data.get("litellm_metadata") or {} + uid = metadata.get(self.user_id_field) + if uid: + return str(uid) + if hasattr(user_api_key_dict, "user_id") and user_api_key_dict.user_id: + return str(user_api_key_dict.user_id) + if hasattr(user_api_key_dict, "end_user_id") and user_api_key_dict.end_user_id: + return str(user_api_key_dict.end_user_id) + return None + + # ------------------------------------------------------------------ + # Policy action evaluation + # ------------------------------------------------------------------ + + @staticmethod + def _should_block(response: Dict[str, Any]) -> bool: + """Return True if any policyAction requires blocking.""" + for action in response.get("policyActions", []): + odata_type = action.get("@odata.type", "") + action_field = action.get("action", "") + + if "restrictAccessAction" in odata_type or action_field == "restrictAccess": + restriction = action.get("restrictionAction", "") + if restriction == "block": + return True + return False + + # ------------------------------------------------------------------ + # User prompt extraction (same pattern as AzureGuardrailBase) + # ------------------------------------------------------------------ + + def get_user_prompt(self, messages: List["AllMessageValues"]) -> Optional[str]: + """Get the last consecutive block of user messages as a single string.""" + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + convert_content_list_to_str, + ) + + if not messages: + return None + + user_messages = [] + for message in reversed(messages): + if message.get("role") == "user": + user_messages.append(message) + else: + break + + if not user_messages: + return None + + user_messages.reverse() + user_prompt = "" + for message in user_messages: + text_content = convert_content_list_to_str(message) + user_prompt += text_content + "\n" + + result = user_prompt.strip() + return result if result else None diff --git a/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py new file mode 100644 index 00000000000..f6fcae77ce2 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/microsoft_purview/purview_dlp.py @@ -0,0 +1,302 @@ +""" +Microsoft Purview DLP Guardrail for LiteLLM. + +Supports three modes: +- pre_call: Block sensitive data in prompts before they reach the LLM. +- post_call: Block sensitive data in LLM responses. +- logging_only: Log interactions to Purview for audit/compliance without blocking. +""" + +import asyncio +from datetime import datetime +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Type, Union + +from fastapi import HTTPException + +from litellm._logging import verbose_proxy_logger +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) +from litellm.types.guardrails import GuardrailEventHooks + +from .base import PurviewGuardrailBase + +if TYPE_CHECKING: + from litellm.proxy._types import UserAPIKeyAuth + from litellm.types.proxy.guardrails.guardrail_hooks.base import ( + GuardrailConfigModel, + ) + from litellm.types.utils import ( + CallTypesLiteral, + EmbeddingResponse, + ImageResponse, + ModelResponse, + ) + + +class MicrosoftPurviewDLPGuardrail(PurviewGuardrailBase, CustomGuardrail): + """ + Microsoft Purview DLP guardrail. + + Evaluates prompts and responses against Microsoft Purview DLP policies + via the Microsoft Graph ``processContent`` API. + """ + + def __init__( + self, + guardrail_name: str, + tenant_id: str, + client_id: str, + client_secret: str, + purview_app_name: str = "LiteLLM", + user_id_field: str = "user_id", + logging_only: bool = False, + **kwargs: Any, + ): + if logging_only: + kwargs["event_hook"] = GuardrailEventHooks.logging_only + + supported_event_hooks = [ + GuardrailEventHooks.pre_call, + GuardrailEventHooks.post_call, + GuardrailEventHooks.logging_only, + ] + + super().__init__( + tenant_id=tenant_id, + client_id=client_id, + client_secret=client_secret, + purview_app_name=purview_app_name, + user_id_field=user_id_field, + guardrail_name=guardrail_name, + supported_event_hooks=supported_event_hooks, + **kwargs, + ) + self._logging_only = logging_only + self.guardrail_provider = "microsoft_purview" + verbose_proxy_logger.info( + "Initialized Microsoft Purview DLP Guardrail: %s (logging_only=%s)", + guardrail_name, + logging_only, + ) + + @staticmethod + def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: + return None # Config model can be added later for UI support + + # ------------------------------------------------------------------ + # Core DLP check + # ------------------------------------------------------------------ + + async def _check_content( + self, + user_id: str, + text: str, + activity: str, + request_data: Dict[str, Any], + block_on_violation: bool = True, + ) -> Dict[str, Any]: + """Evaluate content against Purview DLP policies. + + Args: + user_id: Entra object ID. + text: Content to evaluate. + activity: ``"uploadText"`` or ``"downloadText"``. + request_data: Original request dict (used for logging metadata). + block_on_violation: If False, log only — do not raise. + + Returns: + The processContent response dict. + """ + start_time = datetime.now() + status = "success" + response: Dict[str, Any] = {} + + try: + etag, _ = await self._compute_protection_scopes(user_id) + response = await self._process_content( + user_id=user_id, + text=text, + activity=activity, + etag=etag, + ) + + if self._should_block(response): + status = "guardrail_intervened" + except Exception: + status = "guardrail_failed_to_respond" + raise + finally: + end_time = datetime.now() + self.add_standard_logging_guardrail_information_to_request_data( + guardrail_provider=self.guardrail_provider, + guardrail_json_response=response, + request_data=request_data, + guardrail_status=status, + start_time=start_time.timestamp(), + end_time=end_time.timestamp(), + duration=(end_time - start_time).total_seconds(), + ) + + if block_on_violation and status == "guardrail_intervened": + raise HTTPException( + status_code=400, + detail={ + "error": "Microsoft Purview DLP: Content blocked by policy", + "activity": activity, + }, + ) + + return response + + # ------------------------------------------------------------------ + # Pre-call hook — DLP on prompts + # ------------------------------------------------------------------ + + @log_guardrail_information + async def async_pre_call_hook( + self, + user_api_key_dict: "UserAPIKeyAuth", + cache: Any, + data: Dict[str, Any], + call_type: "CallTypesLiteral", + ) -> Optional[Dict[str, Any]]: + """Check user prompt against Purview DLP policies before LLM call.""" + user_id = self._resolve_user_id(data, user_api_key_dict) + if not user_id: + verbose_proxy_logger.warning( + "Purview DLP: No user_id found, skipping pre-call check" + ) + return data + + messages: Optional[List] = data.get("messages") + if not messages: + return data + + user_prompt = self.get_user_prompt(messages) + if user_prompt: + await self._check_content( + user_id=user_id, + text=user_prompt, + activity="uploadText", + request_data=data, + block_on_violation=True, + ) + return None + + # ------------------------------------------------------------------ + # Post-call hook — DLP on responses + # ------------------------------------------------------------------ + + async def async_post_call_success_hook( + self, + data: dict, + user_api_key_dict: "UserAPIKeyAuth", + response: Union[Any, "ModelResponse", "EmbeddingResponse", "ImageResponse"], + ) -> Any: + """Check LLM response against Purview DLP policies.""" + from litellm.types.utils import Choices, ModelResponse + + user_id = self._resolve_user_id(data, user_api_key_dict) + if not user_id: + verbose_proxy_logger.warning( + "Purview DLP: No user_id found, skipping post-call check" + ) + return response + + if ( + isinstance(response, ModelResponse) + and response.choices + and isinstance(response.choices[0], Choices) + ): + content = response.choices[0].message.content or "" + if content: + await self._check_content( + user_id=user_id, + text=content, + activity="downloadText", + request_data=data, + block_on_violation=True, + ) + return response + + # ------------------------------------------------------------------ + # Logging-only hook — audit without blocking + # ------------------------------------------------------------------ + + def logging_hook( + self, kwargs: dict, result: Any, call_type: str + ) -> Tuple[dict, Any]: + """Sync wrapper for async_logging_hook (follows Presidio pattern).""" + from concurrent.futures import ThreadPoolExecutor + + def run_in_new_loop() -> Tuple[dict, Any]: + new_loop = asyncio.new_event_loop() + try: + asyncio.set_event_loop(new_loop) + return new_loop.run_until_complete( + self.async_logging_hook( + kwargs=kwargs, result=result, call_type=call_type + ) + ) + finally: + new_loop.close() + asyncio.set_event_loop(None) + + try: + _ = asyncio.get_running_loop() + with ThreadPoolExecutor(max_workers=1) as executor: + future = executor.submit(run_in_new_loop) + return future.result() + except RuntimeError: + return run_in_new_loop() + + async def async_logging_hook( + self, kwargs: dict, result: Any, call_type: str + ) -> Tuple[dict, Any]: + """Send both prompt and response to Purview for audit logging. + + Errors are logged but never raised — this mode is non-blocking. + """ + try: + metadata = kwargs.get("metadata") or kwargs.get("litellm_metadata") or {} + user_id = metadata.get(self.user_id_field) or kwargs.get( + "user_api_key_user_id" + ) + + if not user_id: + verbose_proxy_logger.debug("Purview audit: no user_id, skipping") + return kwargs, result + + # Log prompt (uploadText) + messages = kwargs.get("messages") + if messages: + user_prompt = self.get_user_prompt(messages) + if user_prompt: + await self._check_content( + user_id=user_id, + text=user_prompt, + activity="uploadText", + request_data=kwargs, + block_on_violation=False, + ) + + # Log response (downloadText) + from litellm.types.utils import Choices, ModelResponse + + if isinstance(result, ModelResponse) and result.choices: + if isinstance(result.choices[0], Choices): + content = result.choices[0].message.content or "" + if content: + await self._check_content( + user_id=user_id, + text=content, + activity="downloadText", + request_data=kwargs, + block_on_violation=False, + ) + except Exception as e: + verbose_proxy_logger.error("Purview audit logging error: %s", e) + + return kwargs, result diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 751113400d3..98a833ccb67 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -5,6 +5,9 @@ from typing import Any, Dict, List, Literal, Optional, Union from pydantic import BaseModel, ConfigDict, Field, field_validator from typing_extensions import Required, TypedDict +from litellm.types.proxy.guardrails.guardrail_hooks.akto import ( + AktoConfigModel, +) from litellm.types.proxy.guardrails.guardrail_hooks.block_code_execution import ( BlockCodeExecutionGuardrailConfigModel, ) @@ -17,9 +20,6 @@ from litellm.types.proxy.guardrails.guardrail_hooks.grayswan import ( from litellm.types.proxy.guardrails.guardrail_hooks.ibm import ( IBMGuardrailsBaseConfigModel, ) -from litellm.types.proxy.guardrails.guardrail_hooks.akto import ( - AktoConfigModel, -) from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import ( ContentFilterCategoryConfig, ) @@ -93,6 +93,7 @@ class SupportedGuardrailIntegrations(Enum): GENERIC_GUARDRAIL_API = "generic_guardrail_api" QUALIFIRE = "qualifire" CUSTOM_CODE = "custom_code" + MICROSOFT_PURVIEW = "microsoft_purview" SEMANTIC_GUARD = "semantic_guard" MCP_END_USER_PERMISSION = "mcp_end_user_permission" BLOCK_CODE_EXECUTION = "block_code_execution" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py new file mode 100644 index 00000000000..c4ed94b75f5 --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_microsoft_purview.py @@ -0,0 +1,627 @@ +"""Unit tests for the Microsoft Purview DLP guardrail.""" + +import time +from unittest.mock import AsyncMock, Mock, patch + +import pytest +from fastapi import HTTPException + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.guardrails.guardrail_hooks.microsoft_purview.base import ( + PurviewGuardrailBase, +) +from litellm.proxy.guardrails.guardrail_hooks.microsoft_purview.purview_dlp import ( + MicrosoftPurviewDLPGuardrail, +) + + +def _make_guardrail(**kwargs) -> MicrosoftPurviewDLPGuardrail: + """Helper to construct a guardrail with test defaults.""" + defaults = { + "guardrail_name": "test-purview", + "tenant_id": "test-tenant-id", + "client_id": "test-client-id", + "client_secret": "test-client-secret", + } + defaults.update(kwargs) + return MicrosoftPurviewDLPGuardrail(**defaults) + + +def _mock_token_response(): + """Mock a successful OAuth2 token response.""" + resp = Mock() + resp.json.return_value = { + "access_token": "mock-access-token", + "expires_in": 3600, + } + return resp + + +def _mock_graph_response(policy_actions=None, protection_scope_state="unchanged"): + """Mock a processContent Graph API response.""" + resp = Mock() + body = { + "protectionScopeState": protection_scope_state, + "policyActions": policy_actions or [], + "processingErrors": [], + } + resp.json.return_value = body + resp.headers = {"ETag": "test-etag-123"} + return resp + + +def _mock_scope_response(): + """Mock a protectionScopes/compute Graph API response.""" + resp = Mock() + resp.json.return_value = { + "value": [ + { + "activities": "uploadText,downloadText", + "executionMode": "evaluateInline", + "policyActions": [], + } + ] + } + resp.headers = {"ETag": "scope-etag-123"} + return resp + + +# --------------------------------------------------------------- +# _should_block +# --------------------------------------------------------------- + + +class TestShouldBlock: + def test_empty_policy_actions(self): + assert PurviewGuardrailBase._should_block({"policyActions": []}) is False + + def test_no_policy_actions_key(self): + assert PurviewGuardrailBase._should_block({}) is False + + def test_restrict_access_block(self): + response = { + "policyActions": [ + { + "@odata.type": "#microsoft.graph.restrictAccessAction", + "action": "restrictAccess", + "restrictionAction": "block", + } + ] + } + assert PurviewGuardrailBase._should_block(response) is True + + def test_restrict_access_non_block(self): + response = { + "policyActions": [ + { + "@odata.type": "#microsoft.graph.restrictAccessAction", + "action": "restrictAccess", + "restrictionAction": "warn", + } + ] + } + assert PurviewGuardrailBase._should_block(response) is False + + def test_non_restrict_action(self): + response = { + "policyActions": [ + { + "@odata.type": "#microsoft.graph.auditAction", + "action": "audit", + } + ] + } + assert PurviewGuardrailBase._should_block(response) is False + + def test_multiple_actions_one_blocks(self): + response = { + "policyActions": [ + {"action": "audit"}, + { + "@odata.type": "#microsoft.graph.restrictAccessAction", + "action": "restrictAccess", + "restrictionAction": "block", + }, + ] + } + assert PurviewGuardrailBase._should_block(response) is True + + +# --------------------------------------------------------------- +# User ID resolution +# --------------------------------------------------------------- + + +class TestResolveUserId: + def test_from_metadata(self): + guardrail = _make_guardrail() + data = {"metadata": {"user_id": "entra-user-123"}} + assert guardrail._resolve_user_id(data, Mock()) == "entra-user-123" + + def test_custom_field(self): + guardrail = _make_guardrail(user_id_field="entra_id") + data = {"metadata": {"entra_id": "custom-user-456"}} + assert guardrail._resolve_user_id(data, Mock()) == "custom-user-456" + + def test_from_user_api_key_dict_user_id(self): + guardrail = _make_guardrail() + auth = UserAPIKeyAuth(api_key="test", user_id="key-user-789") + assert guardrail._resolve_user_id({}, auth) == "key-user-789" + + def test_from_end_user_id(self): + guardrail = _make_guardrail() + auth = Mock() + auth.user_id = None + auth.end_user_id = "end-user-101" + assert guardrail._resolve_user_id({}, auth) == "end-user-101" + + def test_none_when_missing(self): + guardrail = _make_guardrail() + auth = Mock() + auth.user_id = None + auth.end_user_id = None + assert guardrail._resolve_user_id({}, auth) is None + + +# --------------------------------------------------------------- +# Pre-call hook +# --------------------------------------------------------------- + + +class TestPreCallHook: + @pytest.mark.asyncio + async def test_pre_call_allow(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data={"messages": [{"role": "user", "content": "Hello, how are you?"}]}, + call_type="completion", + ) + + mock_check.assert_called_once() + assert mock_check.call_args.kwargs["activity"] == "uploadText" + assert mock_check.call_args.kwargs["block_on_violation"] is True + + @pytest.mark.asyncio + async def test_pre_call_block(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.side_effect = HTTPException( + status_code=400, + detail={"error": "Microsoft Purview DLP: Content blocked by policy"}, + ) + + with pytest.raises(HTTPException) as exc_info: + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth( + api_key="test", user_id="user-123" + ), + cache=None, + data={ + "messages": [ + { + "role": "user", + "content": "SSN: 123-45-6789", + } + ] + }, + call_type="completion", + ) + + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_pre_call_no_user_id_skips(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + result = await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + cache=None, + data={"messages": [{"role": "user", "content": "Hello"}]}, + call_type="completion", + ) + + mock_check.assert_not_called() + # Returns data when skipping + assert result is not None + + @pytest.mark.asyncio + async def test_pre_call_no_messages_skips(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test", user_id="user-123"), + cache=None, + data={}, + call_type="completion", + ) + + mock_check.assert_not_called() + + +# --------------------------------------------------------------- +# Post-call hook +# --------------------------------------------------------------- + + +class TestPostCallHook: + @pytest.mark.asyncio + async def test_post_call_allow(self): + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + response = ModelResponse( + choices=[ + Choices( + index=0, message=Message(content="Safe response", role="assistant") + ) + ], + ) + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.return_value = {"policyActions": []} + + result = await guardrail.async_post_call_success_hook( + data={"metadata": {"user_id": "user-123"}}, + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=response, + ) + + mock_check.assert_called_once() + assert mock_check.call_args.kwargs["activity"] == "downloadText" + assert result is response + + @pytest.mark.asyncio + async def test_post_call_block(self): + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + response = ModelResponse( + choices=[ + Choices( + index=0, + message=Message( + content="Credit card: 4532-6677-8521-3500", + role="assistant", + ), + ) + ], + ) + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + mock_check.side_effect = HTTPException( + status_code=400, + detail={"error": "Microsoft Purview DLP: Content blocked by policy"}, + ) + + with pytest.raises(HTTPException) as exc_info: + await guardrail.async_post_call_success_hook( + data={"metadata": {"user_id": "user-123"}}, + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=response, + ) + + assert exc_info.value.status_code == 400 + + @pytest.mark.asyncio + async def test_post_call_no_user_id_skips(self): + from litellm.types.utils import Choices, Message, ModelResponse + + guardrail = _make_guardrail() + response = ModelResponse( + choices=[ + Choices(index=0, message=Message(content="Response", role="assistant")) + ], + ) + + with patch.object( + guardrail, "_check_content", new_callable=AsyncMock + ) as mock_check: + result = await guardrail.async_post_call_success_hook( + data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="test"), + response=response, + ) + + mock_check.assert_not_called() + assert result is response + + +# --------------------------------------------------------------- +# _check_content — integration-level +# --------------------------------------------------------------- + + +class TestCheckContent: + @pytest.mark.asyncio + async def test_check_content_allow(self): + guardrail = _make_guardrail() + + with ( + patch.object( + guardrail, + "_compute_protection_scopes", + new_callable=AsyncMock, + return_value=("etag-1", {}), + ), + patch.object( + guardrail, + "_process_content", + new_callable=AsyncMock, + return_value={ + "protectionScopeState": "unchanged", + "policyActions": [], + }, + ), + ): + result = await guardrail._check_content( + user_id="user-1", + text="Hello world", + activity="uploadText", + request_data={}, + block_on_violation=True, + ) + + assert result["policyActions"] == [] + + @pytest.mark.asyncio + async def test_check_content_block(self): + guardrail = _make_guardrail() + + with ( + patch.object( + guardrail, + "_compute_protection_scopes", + new_callable=AsyncMock, + return_value=("etag-1", {}), + ), + patch.object( + guardrail, + "_process_content", + new_callable=AsyncMock, + return_value={ + "protectionScopeState": "unchanged", + "policyActions": [ + { + "@odata.type": "#microsoft.graph.restrictAccessAction", + "action": "restrictAccess", + "restrictionAction": "block", + } + ], + }, + ), + ): + with pytest.raises(HTTPException) as exc_info: + await guardrail._check_content( + user_id="user-1", + text="SSN: 123-45-6789", + activity="uploadText", + request_data={}, + block_on_violation=True, + ) + + assert exc_info.value.status_code == 400 + assert "blocked by policy" in str(exc_info.value.detail) + + @pytest.mark.asyncio + async def test_check_content_logging_only_no_block(self): + """In logging_only mode, violations should NOT raise.""" + guardrail = _make_guardrail(logging_only=True) + + with ( + patch.object( + guardrail, + "_compute_protection_scopes", + new_callable=AsyncMock, + return_value=("etag-1", {}), + ), + patch.object( + guardrail, + "_process_content", + new_callable=AsyncMock, + return_value={ + "protectionScopeState": "unchanged", + "policyActions": [ + { + "@odata.type": "#microsoft.graph.restrictAccessAction", + "action": "restrictAccess", + "restrictionAction": "block", + } + ], + }, + ), + ): + # Should NOT raise even though violation detected + result = await guardrail._check_content( + user_id="user-1", + text="SSN: 123-45-6789", + activity="uploadText", + request_data={}, + block_on_violation=False, + ) + + assert len(result["policyActions"]) == 1 + + +# --------------------------------------------------------------- +# Token caching +# --------------------------------------------------------------- + + +class TestTokenCaching: + @pytest.mark.asyncio + async def test_token_cached(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail.async_handler, "post", return_value=_mock_token_response() + ) as mock_post: + token1 = await guardrail._get_access_token() + token2 = await guardrail._get_access_token() + + assert token1 == "mock-access-token" + assert token2 == "mock-access-token" + # Should only call the token endpoint once (cached) + assert mock_post.call_count == 1 + + @pytest.mark.asyncio + async def test_token_refreshed_on_expiry(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail.async_handler, "post", return_value=_mock_token_response() + ) as mock_post: + await guardrail._get_access_token() + + # Expire the token + guardrail._token_cache = ("old-token", time.time() - 10) + + await guardrail._get_access_token() + + # Should have called token endpoint twice + assert mock_post.call_count == 2 + + +# --------------------------------------------------------------- +# Protection scope caching +# --------------------------------------------------------------- + + +class TestScopeCaching: + @pytest.mark.asyncio + async def test_scope_cached(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_graph_post", new_callable=AsyncMock + ) as mock_post: + mock_post.return_value = ( + { + "value": [ + {"activities": "uploadText", "executionMode": "evaluateInline"} + ] + }, + {"ETag": "scope-etag"}, + ) + + etag1, _ = await guardrail._compute_protection_scopes("user-1") + etag2, _ = await guardrail._compute_protection_scopes("user-1") + + assert etag1 == "scope-etag" + assert etag2 == "scope-etag" + assert mock_post.call_count == 1 + + @pytest.mark.asyncio + async def test_scope_invalidated_on_modified(self): + guardrail = _make_guardrail() + + with patch.object( + guardrail, "_graph_post", new_callable=AsyncMock + ) as mock_post: + # First call: compute scopes + mock_post.return_value = ( + {"value": []}, + {"ETag": "etag-1"}, + ) + await guardrail._compute_protection_scopes("user-1") + + # processContent returns modified + mock_post.return_value = ( + {"protectionScopeState": "modified", "policyActions": []}, + {}, + ) + await guardrail._process_content("user-1", "text", "uploadText", "etag-1") + + # Scope cache should be invalidated + assert "user-1" not in guardrail._scope_cache + + +# --------------------------------------------------------------- +# Initializer validation +# --------------------------------------------------------------- + + +class TestInitializerValidation: + def test_missing_tenant_id(self): + from litellm.proxy.guardrails.guardrail_hooks.microsoft_purview import ( + initialize_guardrail, + ) + + litellm_params = Mock() + litellm_params.get.return_value = None + litellm_params.api_key = "secret" + litellm_params.mode = "pre_call" + + with pytest.raises(ValueError, match="tenant_id is required"): + initialize_guardrail(litellm_params, {"guardrail_name": "test"}) + + def test_missing_client_id(self): + from litellm.proxy.guardrails.guardrail_hooks.microsoft_purview import ( + initialize_guardrail, + ) + + litellm_params = Mock() + litellm_params.get.side_effect = lambda k, *a: ( + "test-tenant" if k == "tenant_id" else None + ) + litellm_params.api_key = "secret" + litellm_params.mode = "pre_call" + + with pytest.raises(ValueError, match="client_id is required"): + initialize_guardrail(litellm_params, {"guardrail_name": "test"}) + + def test_missing_client_secret(self): + from litellm.proxy.guardrails.guardrail_hooks.microsoft_purview import ( + initialize_guardrail, + ) + + litellm_params = Mock() + litellm_params.get.side_effect = lambda k, *a: { + "tenant_id": "test-tenant", + "client_id": "test-client", + }.get(k) + litellm_params.api_key = None + litellm_params.mode = "pre_call" + + with pytest.raises(ValueError, match="client_secret"): + initialize_guardrail(litellm_params, {"guardrail_name": "test"}) + + +# --------------------------------------------------------------- +# Auto-discovery registration +# --------------------------------------------------------------- + + +class TestRegistration: + def test_registry_contains_microsoft_purview(self): + from litellm.proxy.guardrails.guardrail_hooks.microsoft_purview import ( + guardrail_class_registry, + guardrail_initializer_registry, + ) + + assert "microsoft_purview" in guardrail_initializer_registry + assert "microsoft_purview" in guardrail_class_registry + assert ( + guardrail_class_registry["microsoft_purview"] + is MicrosoftPurviewDLPGuardrail + )