feat(guardrails): add Microsoft Purview DLP guardrail

This commit is contained in:
Sameer Kankute 2026-04-02 11:19:07 +05:30
parent 410ce761dc
commit c51b8fdc54
No known key found for this signature in database
5 changed files with 1296 additions and 3 deletions

View file

@ -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,
}

View file

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

View file

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

View file

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

View file

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