mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat(guardrails): add Microsoft Purview DLP guardrail
This commit is contained in:
parent
410ce761dc
commit
c51b8fdc54
5 changed files with 1296 additions and 3 deletions
|
|
@ -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,
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
Loading…
Add table
Reference in a new issue