feat add microsfot purview logger

This commit is contained in:
Harshit28j 2026-03-12 01:45:41 +05:30
parent 08bb3808bc
commit 20ca676507
7 changed files with 1478 additions and 521 deletions

File diff suppressed because it is too large Load diff

View file

@ -454,5 +454,38 @@
}
},
"description": "SQS Queue (AWS) Logging Integration"
},
{
"id": "microsoft_purview",
"displayName": "Microsoft Purview",
"logo": "microsoft.svg",
"supports_key_team_logging": false,
"dynamic_params": {
"MICROSOFT_PURVIEW_TENANT_ID": {
"type": "text",
"ui_name": "Azure Tenant ID",
"description": "Azure Active Directory Tenant ID",
"required": true
},
"MICROSOFT_PURVIEW_CLIENT_ID": {
"type": "password",
"ui_name": "App Client ID",
"description": "App Registration Client ID with Content.Process.All permission",
"required": true
},
"MICROSOFT_PURVIEW_CLIENT_SECRET": {
"type": "password",
"ui_name": "App Client Secret",
"description": "App Registration Client Secret",
"required": true
},
"MICROSOFT_PURVIEW_APP_ID": {
"type": "text",
"ui_name": "App ID (GUID)",
"description": "Registered Application GUID in Purview policy location",
"required": false
}
},
"description": "Microsoft Purview AI Compliance & Audit Logging Integration"
}
]
]

View file

@ -0,0 +1,424 @@
"""
Microsoft Purview Integration - sends LLM prompts & responses to the Microsoft Graph
processContent API for compliance, DLP, and audit tracking.
Reference API: https://learn.microsoft.com/en-us/graph/api/userdatasecurityandgovernance-processcontent
"""
import asyncio
import os
import traceback
from collections import defaultdict
from typing import List, Optional, Dict, Any
from datetime import datetime, timezone
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.types.utils import StandardLoggingPayload
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
class MicrosoftPurviewLogger(CustomBatchLogger):
"""
Logger that sends LLM interactions to Microsoft Purview via the
Microsoft Graph processContent API for compliance and audit.
"""
def __init__(
self,
tenant_id: Optional[str] = None,
client_id: Optional[str] = None,
client_secret: Optional[str] = None,
app_name: Optional[str] = None,
app_version: Optional[str] = None,
app_id: Optional[str] = None,
default_user_id: Optional[str] = None,
graph_api_version: str = "v1.0",
log_prompts: bool = True,
log_responses: bool = True,
**kwargs,
):
"""
Initialize Microsoft Purview logger using the Graph API
"""
self.async_httpx_client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.LoggingCallback
)
self.tenant_id = (
tenant_id
or os.getenv("MICROSOFT_PURVIEW_TENANT_ID")
or os.getenv("AZURE_TENANT_ID")
)
self.client_id = (
client_id
or os.getenv("MICROSOFT_PURVIEW_CLIENT_ID")
or os.getenv("AZURE_CLIENT_ID")
)
self.client_secret = (
client_secret
or os.getenv("MICROSOFT_PURVIEW_CLIENT_SECRET")
or os.getenv("AZURE_CLIENT_SECRET")
)
self.app_name = app_name or os.getenv(
"MICROSOFT_PURVIEW_APP_NAME", "LiteLLM Proxy"
)
self.app_version = app_version or os.getenv(
"MICROSOFT_PURVIEW_APP_VERSION", getattr(litellm, "_version", "0.0.0")
)
self.app_id = app_id or os.getenv("MICROSOFT_PURVIEW_APP_ID")
self.default_user_id = default_user_id or os.getenv(
"MICROSOFT_PURVIEW_DEFAULT_USER_ID", "lite-llm-unknown-user"
)
self.graph_api_version = graph_api_version
# Boolean flags controls what is captured in logs
self.log_prompts = log_prompts
self.log_responses = log_responses
if not self.tenant_id:
raise ValueError(
"MICROSOFT_PURVIEW_TENANT_ID is required to use Microsoft Purview integration"
)
if not self.client_id:
raise ValueError(
"MICROSOFT_PURVIEW_CLIENT_ID is required to use Microsoft Purview integration"
)
if not self.client_secret:
raise ValueError(
"MICROSOFT_PURVIEW_CLIENT_SECRET is required to use Microsoft Purview integration"
)
# OAuth2 scope for Microsoft Graph
self.oauth_scope = "https://graph.microsoft.com/.default"
self.oauth_token: Optional[str] = None
self.oauth_token_expires_at: Optional[float] = None
self.flush_lock = asyncio.Lock()
super().__init__(**kwargs, flush_lock=self.flush_lock)
asyncio.create_task(self.periodic_flush())
self.log_queue: List[StandardLoggingPayload] = []
async def _get_oauth_token(self) -> str:
"""
Get OAuth2 Bearer token for Microsoft Graph
"""
import time
if (
self.oauth_token
and self.oauth_token_expires_at
and time.time() < self.oauth_token_expires_at - 60
): # Refresh 60 seconds before expiry
return self.oauth_token
assert self.tenant_id is not None, "tenant_id is required"
assert self.client_id is not None, "client_id is required"
assert self.client_secret is not None, "client_secret is required"
token_url = (
f"https://login.microsoftonline.com/{self.tenant_id}/oauth2/v2.0/token"
)
token_data = {
"client_id": self.client_id,
"client_secret": self.client_secret,
"scope": self.oauth_scope,
"grant_type": "client_credentials",
}
response = await self.async_httpx_client.post(
url=token_url,
data=token_data,
headers={"Content-Type": "application/x-www-form-urlencoded"},
)
if response.status_code != 200:
raise Exception(
f"Failed to get OAuth2 token: {response.status_code} - {response.text}"
)
token_response = response.json()
self.oauth_token = token_response.get("access_token")
expires_in = token_response.get("expires_in", 3600)
if not self.oauth_token:
raise Exception("OAuth2 token response did not contain access_token")
self.oauth_token_expires_at = time.time() + expires_in
return self.oauth_token
def _extract_user_id(self, payload: StandardLoggingPayload) -> str:
"""Get the user identity to map to the user-scoped Purview API"""
metadata = payload.get("metadata", {}) or {}
user_id = metadata.get("user_api_key_user_id")
if user_id:
return str(user_id)
end_user = payload.get("end_user")
if end_user:
return str(end_user)
# fallback if not available
return self.default_user_id
def _serialize_messages(self, messages: Any) -> str:
"""Serialize prompts to a string, limiting total size if necessary"""
if isinstance(messages, str):
text = messages
elif isinstance(messages, list):
try:
# Try to extract just text for readability in Purview
parts = []
for msg in messages:
if isinstance(msg, dict):
role = msg.get("role", "user")
content = msg.get("content", "")
if isinstance(content, str):
parts.append(f"[{role}]: {content}")
else:
parts.append(f"[{role}]: {safe_dumps(content)}")
else:
parts.append(str(msg))
text = "\n\n".join(parts)
except Exception:
text = safe_dumps(messages)
else:
text = safe_dumps(messages)
return text
def _extract_response_text(self, payload: StandardLoggingPayload) -> str:
"""Extract the model response"""
response = payload.get("response", {})
if not response:
return ""
if isinstance(response, str):
return response
try:
choices = response.get("choices", [])
if choices and len(choices) > 0:
message = choices[0].get("message", {})
content = message.get("content")
if content:
return str(content)
except Exception:
pass
return safe_dumps(response)
def _format_time(self, timestamp: Any) -> str:
"""Format timestamp to ISO 8601 strictly"""
try:
if timestamp:
dt = datetime.fromtimestamp(timestamp, tz=timezone.utc)
else:
dt = datetime.now(timezone.utc)
# Purview requires exactly yYYY-MM-DDThh:mm:ss format
return dt.strftime("%Y-%m-%dT%H:%M:%S")
except Exception:
return datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%S")
def _build_process_content_request(self, payload: StandardLoggingPayload) -> dict:
"""Assemble the Microsoft Graph API request body for processContent"""
entries = []
trace_id = payload.get("trace_id", "") or "purview-unknown-trace"
# 1. Add User Prompts
if self.log_prompts:
messages = payload.get("messages", [])
prompt_text = self._serialize_messages(messages)
if prompt_text:
entries.append(
{
"@odata.type": "microsoft.graph.processConversationMetadata",
"identifier": f"{trace_id}-prompt",
"content": {
"@odata.type": "microsoft.graph.textContent",
"data": prompt_text,
},
"name": "LLM Prompt",
"correlationId": trace_id,
"sequenceNumber": 0,
"isTruncated": False,
"createdDateTime": self._format_time(payload.get("startTime")),
"modifiedDateTime": self._format_time(payload.get("startTime")),
}
)
# 2. Add AI Response
if self.log_responses:
response_text = self._extract_response_text(payload)
if response_text:
entries.append(
{
"@odata.type": "microsoft.graph.processConversationMetadata",
"identifier": f"{trace_id}-response",
"content": {
"@odata.type": "microsoft.graph.textContent",
"data": response_text,
},
"name": "LLM Response",
"correlationId": trace_id,
"sequenceNumber": 1,
"isTruncated": False,
"createdDateTime": self._format_time(payload.get("endTime")),
"modifiedDateTime": self._format_time(payload.get("endTime")),
}
)
# If nothing to send based on configuration or empty payload
if not entries:
return {}
req_body = {
"contentToProcess": {
"contentEntries": entries,
"activityMetadata": {
"activity": "uploadText" # Or "downloadText". "uploadText" represents generating intent + receiving content
},
"integratedAppMetadata": {
"name": self.app_name,
"version": self.app_version,
},
}
}
# Add protectedAppMetadata if an app_id was specified (needed for full mapping in Purview)
if self.app_id:
req_body["contentToProcess"]["protectedAppMetadata"] = {
"name": self.app_name,
"version": self.app_version,
"applicationLocation": {
"@odata.type": "microsoft.graph.policyLocationApplication",
"value": self.app_id,
},
}
return req_body
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
"""Async log success events to Microsoft Purview API queue"""
try:
verbose_logger.debug(
"Microsoft Purview: Queueing success log for model %s",
kwargs.get("model"),
)
standard_logging_payload = kwargs.get("standard_logging_object", None)
if standard_logging_payload is None:
return
self.log_queue.append(standard_logging_payload)
if len(self.log_queue) >= self.batch_size:
await self.async_send_batch()
except Exception as e:
verbose_logger.exception(
f"Microsoft Purview Success Logging Error - {str(e)}\n{traceback.format_exc()}"
)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
"""Async log failure events to Microsoft Purview API queue"""
try:
verbose_logger.debug(
"Microsoft Purview: Queueing failure log for model %s",
kwargs.get("model"),
)
standard_logging_payload = kwargs.get("standard_logging_object", None)
if standard_logging_payload is None:
return
self.log_queue.append(standard_logging_payload)
if len(self.log_queue) >= self.batch_size:
await self.async_send_batch()
except Exception as e:
verbose_logger.exception(
f"Microsoft Purview Failure Logging Error - {str(e)}\n{traceback.format_exc()}"
)
async def async_send_batch(self):
"""
Sends the batch of logs to Microsoft Graph Process Content API
"""
try:
if not self.log_queue:
return
verbose_logger.debug(
"Microsoft Purview - about to flush %s events", len(self.log_queue)
)
# 1. Group payloads by user id (Since Graph API is per-user)
groups: Dict[str, list] = defaultdict(list)
for payload in self.log_queue:
user_id = self._extract_user_id(payload)
req_body = self._build_process_content_request(payload)
if req_body:
groups[user_id].append(req_body)
if not groups:
self.log_queue.clear()
return
# 2. Get OAuth2 Token
bearer_token = await self._get_oauth_token()
headers = {
"Authorization": f"Bearer {bearer_token}",
"Content-Type": "application/json",
}
# 3. Fire requests concurrently
# Although they belong to different users, we loop through. Graph API doesn't support batching processContent inside a single call currently.
# We process them in parallel
tasks = []
for user_id, requests in groups.items():
api_endpoint = f"https://graph.microsoft.com/{self.graph_api_version}/users/{user_id}/dataSecurityAndGovernance/processContent"
for req_body in requests:
tasks.append(
self.async_httpx_client.post(
url=api_endpoint, json=req_body, headers=headers
)
)
responses = await asyncio.gather(*tasks, return_exceptions=True)
for index, response in enumerate(responses):
if isinstance(response, Exception):
verbose_logger.error(
"Microsoft Purview Graph API encountered error: %s",
str(response),
)
elif response.status_code not in [200, 202, 204]:
verbose_logger.error(
"Microsoft Purview Graph API error: status_code=%s, response=%s",
response.status_code,
response.text,
)
verbose_logger.debug(
"Microsoft Purview: Flushed %s processContent calls", len(tasks)
)
except Exception as e:
verbose_logger.exception(
f"Microsoft Purview Error sending batch API - {str(e)}\n{traceback.format_exc()}"
)
finally:
self.log_queue.clear()

View file

@ -0,0 +1,61 @@
from typing import TypedDict, List, Optional
class TextContent(TypedDict):
odata_type: str # @odata.type -> "microsoft.graph.textContent"
data: str
class ProcessConversationMetadata(TypedDict):
odata_type: str # @odata.type -> "microsoft.graph.processConversationMetadata"
identifier: str
content: TextContent
name: str
correlationId: str
sequenceNumber: int
isTruncated: bool
createdDateTime: str
modifiedDateTime: str
class ActivityMetadata(TypedDict):
activity: str # e.g., "uploadText", "downloadText"
class OperatingSystemSpecifications(TypedDict):
operatingSystemPlatform: str
operatingSystemVersion: str
class DeviceMetadata(TypedDict, total=False):
deviceType: str
operatingSystemSpecifications: OperatingSystemSpecifications
ipAddress: str
class PolicyLocationApplication(TypedDict):
odata_type: str # @odata.type -> "microsoft.graph.policyLocationApplication"
value: str
class ProtectedApplicationMetadata(TypedDict, total=False):
name: str
version: str
applicationLocation: PolicyLocationApplication
class IntegratedApplicationMetadata(TypedDict):
name: str
version: str
class ProcessContentRequest(TypedDict):
contentEntries: List[ProcessConversationMetadata]
activityMetadata: ActivityMetadata
deviceMetadata: Optional[DeviceMetadata]
protectedAppMetadata: Optional[ProtectedApplicationMetadata]
integratedAppMetadata: IntegratedApplicationMetadata
class ProcessContentRequestBody(TypedDict):
contentToProcess: ProcessContentRequest

View file

@ -130,6 +130,7 @@ from ..integrations.argilla import ArgillaLogger
from ..integrations.arize.arize_phoenix import ArizePhoenixLogger
from ..integrations.athina import AthinaLogger
from ..integrations.azure_sentinel.azure_sentinel import AzureSentinelLogger
from ..integrations.microsoft_purview.microsoft_purview import MicrosoftPurviewLogger
from ..integrations.azure_storage.azure_storage import AzureBlobStorageLogger
from ..integrations.custom_prompt_management import CustomPromptManagement
from ..integrations.datadog.datadog import DataDogLogger
@ -352,9 +353,9 @@ class Logging(LiteLLMLoggingBaseClass):
)
self.function_id = function_id
self.streaming_chunks: List[Any] = [] # for generating complete stream response
self.sync_streaming_chunks: List[Any] = (
[]
) # for generating complete stream response
self.sync_streaming_chunks: List[
Any
] = [] # for generating complete stream response
self.log_raw_request_response = log_raw_request_response
# Initialize dynamic callbacks
@ -746,9 +747,9 @@ class Logging(LiteLLMLoggingBaseClass):
prompt_spec=prompt_spec,
dynamic_callback_params=dynamic_callback_params,
):
self.model_call_details["prompt_integration"] = (
logger.__class__.__name__
)
self.model_call_details[
"prompt_integration"
] = logger.__class__.__name__
return logger
except Exception:
# If check fails, continue to next logger
@ -816,9 +817,9 @@ class Logging(LiteLLMLoggingBaseClass):
if anthropic_cache_control_logger := AnthropicCacheControlHook.get_custom_logger_for_anthropic_cache_control_hook(
non_default_params
):
self.model_call_details["prompt_integration"] = (
anthropic_cache_control_logger.__class__.__name__
)
self.model_call_details[
"prompt_integration"
] = anthropic_cache_control_logger.__class__.__name__
return anthropic_cache_control_logger
#########################################################
@ -830,9 +831,9 @@ class Logging(LiteLLMLoggingBaseClass):
internal_usage_cache=None,
llm_router=None,
)
self.model_call_details["prompt_integration"] = (
vector_store_custom_logger.__class__.__name__
)
self.model_call_details[
"prompt_integration"
] = vector_store_custom_logger.__class__.__name__
# Add to global callbacks so post-call hooks are invoked
if (
vector_store_custom_logger
@ -892,9 +893,9 @@ class Logging(LiteLLMLoggingBaseClass):
model
): # if model name was changes pre-call, overwrite the initial model call name with the new one
self.model_call_details["model"] = model
self.model_call_details["litellm_params"]["api_base"] = (
self._get_masked_api_base(additional_args.get("api_base", ""))
)
self.model_call_details["litellm_params"][
"api_base"
] = self._get_masked_api_base(additional_args.get("api_base", ""))
def pre_call(self, input, api_key, model=None, additional_args={}): # noqa: PLR0915
# Log the exact input to the LLM API
@ -923,10 +924,10 @@ class Logging(LiteLLMLoggingBaseClass):
try:
# [Non-blocking Extra Debug Information in metadata]
if turn_off_message_logging is True:
_metadata["raw_request"] = (
"redacted by litellm. \
_metadata[
"raw_request"
] = "redacted by litellm. \
'litellm.turn_off_message_logging=True'"
)
else:
curl_command = self._get_request_curl_command(
api_base=additional_args.get("api_base", ""),
@ -937,34 +938,34 @@ class Logging(LiteLLMLoggingBaseClass):
_metadata["raw_request"] = str(curl_command)
# split up, so it's easier to parse in the UI
self.model_call_details["raw_request_typed_dict"] = (
RawRequestTypedDict(
raw_request_api_base=str(
additional_args.get("api_base") or ""
),
raw_request_body=self._get_raw_request_body(
additional_args.get("complete_input_dict", {})
),
# NOTE: setting ignore_sensitive_headers to True will cause
# the Authorization header to be leaked when calls to the health
# endpoint are made and fail.
raw_request_headers=self._get_masked_headers(
additional_args.get("headers", {}) or {},
),
error=None,
)
self.model_call_details[
"raw_request_typed_dict"
] = RawRequestTypedDict(
raw_request_api_base=str(
additional_args.get("api_base") or ""
),
raw_request_body=self._get_raw_request_body(
additional_args.get("complete_input_dict", {})
),
# NOTE: setting ignore_sensitive_headers to True will cause
# the Authorization header to be leaked when calls to the health
# endpoint are made and fail.
raw_request_headers=self._get_masked_headers(
additional_args.get("headers", {}) or {},
),
error=None,
)
except Exception as e:
self.model_call_details["raw_request_typed_dict"] = (
RawRequestTypedDict(
error=str(e),
)
self.model_call_details[
"raw_request_typed_dict"
] = RawRequestTypedDict(
error=str(e),
)
_metadata["raw_request"] = (
"Unable to Log \
_metadata[
"raw_request"
] = "Unable to Log \
raw request: {}".format(
str(e)
)
str(e)
)
if getattr(self, "logger_fn", None) and callable(self.logger_fn):
try:
@ -1265,13 +1266,13 @@ class Logging(LiteLLMLoggingBaseClass):
for callback in callbacks:
try:
if isinstance(callback, CustomLogger):
response: Optional[MCPPostCallResponseObject] = (
await callback.async_post_mcp_tool_call_hook(
kwargs=kwargs,
response_obj=post_mcp_tool_call_response_obj,
start_time=start_time,
end_time=end_time,
)
response: Optional[
MCPPostCallResponseObject
] = await callback.async_post_mcp_tool_call_hook(
kwargs=kwargs,
response_obj=post_mcp_tool_call_response_obj,
start_time=start_time,
end_time=end_time,
)
######################################################################
# if any of the callbacks modify the response, use the modified response
@ -1466,9 +1467,9 @@ class Logging(LiteLLMLoggingBaseClass):
verbose_logger.debug(
f"response_cost_failure_debug_information: {debug_info}"
)
self.model_call_details["response_cost_failure_debug_information"] = (
debug_info
)
self.model_call_details[
"response_cost_failure_debug_information"
] = debug_info
return None
try:
@ -1494,9 +1495,9 @@ class Logging(LiteLLMLoggingBaseClass):
verbose_logger.debug(
f"response_cost_failure_debug_information: {debug_info}"
)
self.model_call_details["response_cost_failure_debug_information"] = (
debug_info
)
self.model_call_details[
"response_cost_failure_debug_information"
] = debug_info
return None
@ -1652,9 +1653,9 @@ class Logging(LiteLLMLoggingBaseClass):
result=logging_result
)
self.model_call_details["standard_logging_object"] = (
self._build_standard_logging_payload(logging_result, start_time, end_time)
)
self.model_call_details[
"standard_logging_object"
] = self._build_standard_logging_payload(logging_result, start_time, end_time)
if (
standard_logging_payload := self.model_call_details.get(
@ -1732,9 +1733,9 @@ class Logging(LiteLLMLoggingBaseClass):
end_time = datetime.datetime.now()
if self.completion_start_time is None:
self.completion_start_time = end_time
self.model_call_details["completion_start_time"] = (
self.completion_start_time
)
self.model_call_details[
"completion_start_time"
] = self.completion_start_time
self.model_call_details["log_event_type"] = "successful_api_call"
self.model_call_details["end_time"] = end_time
@ -1771,10 +1772,10 @@ class Logging(LiteLLMLoggingBaseClass):
end_time=end_time,
)
elif isinstance(result, dict) or isinstance(result, list):
self.model_call_details["standard_logging_object"] = (
self._build_standard_logging_payload(
result, start_time, end_time
)
self.model_call_details[
"standard_logging_object"
] = self._build_standard_logging_payload(
result, start_time, end_time
)
if (
standard_logging_payload := self.model_call_details.get(
@ -1783,9 +1784,9 @@ class Logging(LiteLLMLoggingBaseClass):
) is not None:
emit_standard_logging_payload(standard_logging_payload)
elif standard_logging_object is not None:
self.model_call_details["standard_logging_object"] = (
standard_logging_object
)
self.model_call_details[
"standard_logging_object"
] = standard_logging_object
else:
self.model_call_details["response_cost"] = None
@ -1943,17 +1944,17 @@ class Logging(LiteLLMLoggingBaseClass):
verbose_logger.debug(
"Logging Details LiteLLM-Success Call streaming complete"
)
self.model_call_details["complete_streaming_response"] = (
complete_streaming_response
)
self.model_call_details["response_cost"] = (
self._response_cost_calculator(result=complete_streaming_response)
)
self.model_call_details[
"complete_streaming_response"
] = complete_streaming_response
self.model_call_details[
"response_cost"
] = self._response_cost_calculator(result=complete_streaming_response)
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details["standard_logging_object"] = (
self._build_standard_logging_payload(
complete_streaming_response, start_time, end_time
)
self.model_call_details[
"standard_logging_object"
] = self._build_standard_logging_payload(
complete_streaming_response, start_time, end_time
)
if (
standard_logging_payload := self.model_call_details.get(
@ -2287,10 +2288,10 @@ class Logging(LiteLLMLoggingBaseClass):
)
else:
if self.stream and complete_streaming_response:
self.model_call_details["complete_response"] = (
self.model_call_details.get(
"complete_streaming_response", {}
)
self.model_call_details[
"complete_response"
] = self.model_call_details.get(
"complete_streaming_response", {}
)
result = self.model_call_details["complete_response"]
openMeterLogger.log_success_event(
@ -2314,10 +2315,10 @@ class Logging(LiteLLMLoggingBaseClass):
)
else:
if self.stream and complete_streaming_response:
self.model_call_details["complete_response"] = (
self.model_call_details.get(
"complete_streaming_response", {}
)
self.model_call_details[
"complete_response"
] = self.model_call_details.get(
"complete_streaming_response", {}
)
result = self.model_call_details["complete_response"]
@ -2456,9 +2457,9 @@ class Logging(LiteLLMLoggingBaseClass):
if complete_streaming_response is not None:
print_verbose("Async success callbacks: Got a complete streaming response")
self.model_call_details["async_complete_streaming_response"] = (
complete_streaming_response
)
self.model_call_details[
"async_complete_streaming_response"
] = complete_streaming_response
try:
if self.model_call_details.get("cache_hit", False) is True:
@ -2469,10 +2470,10 @@ class Logging(LiteLLMLoggingBaseClass):
model_call_details=self.model_call_details
)
# base_model defaults to None if not set on model_info
self.model_call_details["response_cost"] = (
self._response_cost_calculator(
result=complete_streaming_response
)
self.model_call_details[
"response_cost"
] = self._response_cost_calculator(
result=complete_streaming_response
)
verbose_logger.debug(
@ -2485,10 +2486,10 @@ class Logging(LiteLLMLoggingBaseClass):
self.model_call_details["response_cost"] = None
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details["standard_logging_object"] = (
self._build_standard_logging_payload(
complete_streaming_response, start_time, end_time
)
self.model_call_details[
"standard_logging_object"
] = self._build_standard_logging_payload(
complete_streaming_response, start_time, end_time
)
# print standard logging payload
@ -2515,9 +2516,9 @@ class Logging(LiteLLMLoggingBaseClass):
# _success_handler_helper_fn
if self.model_call_details.get("standard_logging_object") is None:
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details["standard_logging_object"] = (
self._build_standard_logging_payload(result, start_time, end_time)
)
self.model_call_details[
"standard_logging_object"
] = self._build_standard_logging_payload(result, start_time, end_time)
# print standard logging payload
if (
@ -2760,18 +2761,18 @@ class Logging(LiteLLMLoggingBaseClass):
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details["standard_logging_object"] = (
get_standard_logging_object_payload(
kwargs=self.model_call_details,
init_response_obj={},
start_time=start_time,
end_time=end_time,
logging_obj=self,
status="failure",
error_str=str(exception),
original_exception=exception,
standard_built_in_tools_params=self.standard_built_in_tools_params,
)
self.model_call_details[
"standard_logging_object"
] = get_standard_logging_object_payload(
kwargs=self.model_call_details,
init_response_obj={},
start_time=start_time,
end_time=end_time,
logging_obj=self,
status="failure",
error_str=str(exception),
original_exception=exception,
standard_built_in_tools_params=self.standard_built_in_tools_params,
)
return start_time, end_time
@ -3678,6 +3679,14 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
_azure_sentinel_logger = AzureSentinelLogger()
_in_memory_loggers.append(_azure_sentinel_logger)
return _azure_sentinel_logger # type: ignore
elif logging_integration == "microsoft_purview":
for callback in _in_memory_loggers:
if isinstance(callback, MicrosoftPurviewLogger):
return callback # type: ignore
_microsoft_purview_logger = MicrosoftPurviewLogger()
_in_memory_loggers.append(_microsoft_purview_logger)
return _microsoft_purview_logger # type: ignore
elif logging_integration == "gcs_bucket":
for callback in _in_memory_loggers:
if isinstance(callback, GCSBucketLogger):
@ -3735,9 +3744,9 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
service_name=arize_config.project_name,
)
os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = (
f"space_id={arize_config.space_key or arize_config.space_id},api_key={arize_config.api_key}"
)
os.environ[
"OTEL_EXPORTER_OTLP_TRACES_HEADERS"
] = f"space_id={arize_config.space_key or arize_config.space_id},api_key={arize_config.api_key}"
for callback in _in_memory_loggers:
if (
isinstance(callback, ArizeLogger)
@ -3763,13 +3772,13 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
existing_attrs = os.environ.get("OTEL_RESOURCE_ATTRIBUTES", "")
# Add openinference.project.name attribute
if existing_attrs:
os.environ["OTEL_RESOURCE_ATTRIBUTES"] = (
f"{existing_attrs},openinference.project.name={arize_phoenix_config.project_name}"
)
os.environ[
"OTEL_RESOURCE_ATTRIBUTES"
] = f"{existing_attrs},openinference.project.name={arize_phoenix_config.project_name}"
else:
os.environ["OTEL_RESOURCE_ATTRIBUTES"] = (
f"openinference.project.name={arize_phoenix_config.project_name}"
)
os.environ[
"OTEL_RESOURCE_ATTRIBUTES"
] = f"openinference.project.name={arize_phoenix_config.project_name}"
# Set Phoenix project name from environment variable
phoenix_project_name = os.environ.get("PHOENIX_PROJECT_NAME", None)
@ -3777,19 +3786,19 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
existing_attrs = os.environ.get("OTEL_RESOURCE_ATTRIBUTES", "")
# Add openinference.project.name attribute
if existing_attrs:
os.environ["OTEL_RESOURCE_ATTRIBUTES"] = (
f"{existing_attrs},openinference.project.name={phoenix_project_name}"
)
os.environ[
"OTEL_RESOURCE_ATTRIBUTES"
] = f"{existing_attrs},openinference.project.name={phoenix_project_name}"
else:
os.environ["OTEL_RESOURCE_ATTRIBUTES"] = (
f"openinference.project.name={phoenix_project_name}"
)
os.environ[
"OTEL_RESOURCE_ATTRIBUTES"
] = f"openinference.project.name={phoenix_project_name}"
# auth can be disabled on local deployments of arize phoenix
if arize_phoenix_config.otlp_auth_headers is not None:
os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = (
arize_phoenix_config.otlp_auth_headers
)
os.environ[
"OTEL_EXPORTER_OTLP_TRACES_HEADERS"
] = arize_phoenix_config.otlp_auth_headers
for callback in _in_memory_loggers:
if (
@ -3965,9 +3974,9 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
exporter="otlp_http",
endpoint="https://langtrace.ai/api/trace",
)
os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = (
f"api_key={os.getenv('LANGTRACE_API_KEY')}"
)
os.environ[
"OTEL_EXPORTER_OTLP_TRACES_HEADERS"
] = f"api_key={os.getenv('LANGTRACE_API_KEY')}"
for callback in _in_memory_loggers:
if (
isinstance(callback, OpenTelemetry)
@ -4284,6 +4293,10 @@ def get_custom_logger_compatible_class( # noqa: PLR0915
for callback in _in_memory_loggers:
if isinstance(callback, AzureSentinelLogger):
return callback
elif logging_integration == "microsoft_purview":
for callback in _in_memory_loggers:
if isinstance(callback, MicrosoftPurviewLogger):
return callback
elif logging_integration == "gcs_bucket":
for callback in _in_memory_loggers:
if isinstance(callback, GCSBucketLogger):
@ -4881,10 +4894,10 @@ class StandardLoggingPayloadSetup:
for key in StandardLoggingHiddenParams.__annotations__.keys():
if key in hidden_params:
if key == "additional_headers":
clean_hidden_params["additional_headers"] = (
StandardLoggingPayloadSetup.get_additional_headers(
hidden_params[key]
)
clean_hidden_params[
"additional_headers"
] = StandardLoggingPayloadSetup.get_additional_headers(
hidden_params[key]
)
else:
clean_hidden_params[key] = hidden_params[key] # type: ignore
@ -5036,7 +5049,6 @@ class StandardLoggingPayloadSetup:
dynamic_litellm_session_id = litellm_params.get("litellm_session_id")
dynamic_litellm_trace_id = litellm_params.get("litellm_trace_id")
# Note: we recommend using `litellm_session_id` for session tracking
# `litellm_trace_id` is an internal litellm param
if dynamic_litellm_session_id:
@ -5507,9 +5519,9 @@ def scrub_sensitive_keys_in_metadata(litellm_params: Optional[dict]):
):
for k, v in metadata["user_api_key_metadata"].items():
if k == "logging": # prevent logging user logging keys
cleaned_user_api_key_metadata[k] = (
"scrubbed_by_litellm_for_sensitive_keys"
)
cleaned_user_api_key_metadata[
k
] = "scrubbed_by_litellm_for_sensitive_keys"
else:
cleaned_user_api_key_metadata[k] = v

View file

@ -0,0 +1,257 @@
import pytest
from unittest.mock import AsyncMock, MagicMock
from litellm.integrations.microsoft_purview.microsoft_purview import (
MicrosoftPurviewLogger,
)
@pytest.fixture
def valid_env_vars(monkeypatch):
monkeypatch.setenv("MICROSOFT_PURVIEW_TENANT_ID", "test-tenant")
monkeypatch.setenv("MICROSOFT_PURVIEW_CLIENT_ID", "test-client-id")
monkeypatch.setenv("MICROSOFT_PURVIEW_CLIENT_SECRET", "test-secret")
monkeypatch.setenv("MICROSOFT_PURVIEW_APP_NAME", "test-app")
monkeypatch.setenv("MICROSOFT_PURVIEW_APP_VERSION", "1.0.0")
monkeypatch.setenv("MICROSOFT_PURVIEW_APP_ID", "test-app-id")
@pytest.mark.asyncio
async def test_init_with_all_env_vars(valid_env_vars):
logger = MicrosoftPurviewLogger()
assert logger.tenant_id == "test-tenant"
assert logger.client_id == "test-client-id"
assert logger.client_secret == "test-secret"
assert logger.app_name == "test-app"
assert logger.app_version == "1.0.0"
assert logger.app_id == "test-app-id"
assert logger.oauth_scope == "https://graph.microsoft.com/.default"
def test_init_missing_tenant_id_raises(monkeypatch):
monkeypatch.delenv("MICROSOFT_PURVIEW_TENANT_ID", raising=False)
monkeypatch.delenv("AZURE_TENANT_ID", raising=False)
monkeypatch.setenv("MICROSOFT_PURVIEW_CLIENT_ID", "test-client-id")
monkeypatch.setenv("MICROSOFT_PURVIEW_CLIENT_SECRET", "test-secret")
with pytest.raises(
ValueError,
match="MICROSOFT_PURVIEW_TENANT_ID is required to use Microsoft Purview integration",
):
MicrosoftPurviewLogger()
def test_init_missing_client_id_raises(monkeypatch):
monkeypatch.setenv("MICROSOFT_PURVIEW_TENANT_ID", "test-tenant")
monkeypatch.delenv("MICROSOFT_PURVIEW_CLIENT_ID", raising=False)
monkeypatch.delenv("AZURE_CLIENT_ID", raising=False)
monkeypatch.setenv("MICROSOFT_PURVIEW_CLIENT_SECRET", "test-secret")
with pytest.raises(
ValueError,
match="MICROSOFT_PURVIEW_CLIENT_ID is required to use Microsoft Purview integration",
):
MicrosoftPurviewLogger()
def test_init_missing_client_secret_raises(monkeypatch):
monkeypatch.setenv("MICROSOFT_PURVIEW_TENANT_ID", "test-tenant")
monkeypatch.setenv("MICROSOFT_PURVIEW_CLIENT_ID", "test-client-id")
monkeypatch.delenv("MICROSOFT_PURVIEW_CLIENT_SECRET", raising=False)
monkeypatch.delenv("AZURE_CLIENT_SECRET", raising=False)
with pytest.raises(
ValueError,
match="MICROSOFT_PURVIEW_CLIENT_SECRET is required to use Microsoft Purview integration",
):
MicrosoftPurviewLogger()
@pytest.mark.asyncio
async def test_extract_user_id_from_metadata(valid_env_vars):
logger = MicrosoftPurviewLogger(default_user_id="default-user")
# Priority 1: metadata.user_api_key_user_id
payload1 = {"metadata": {"user_api_key_user_id": "test-user-1"}}
assert logger._extract_user_id(payload1) == "test-user-1"
# Priority 2: end_user
payload2 = {"end_user": "test-user-2"}
assert logger._extract_user_id(payload2) == "test-user-2"
# Priority 3: fallback
payload3 = {}
assert logger._extract_user_id(payload3) == "default-user"
@pytest.mark.asyncio
async def test_serialize_messages(valid_env_vars):
logger = MicrosoftPurviewLogger()
# List format
messages = [
{"role": "user", "content": "hello world"},
{"role": "assistant", "content": "hi"},
]
result = logger._serialize_messages(messages)
assert "[user]: hello world" in result
assert "[assistant]: hi" in result
# String format
assert logger._serialize_messages("hello world") == "hello world"
@pytest.mark.asyncio
async def test_extract_response_text(valid_env_vars):
logger = MicrosoftPurviewLogger()
# Standard format
payload = {
"response": {"choices": [{"message": {"content": "this is a response"}}]}
}
assert logger._extract_response_text(payload) == "this is a response"
# String format
payload_str = {"response": "just a string response"}
assert logger._extract_response_text(payload_str) == "just a string response"
@pytest.mark.asyncio
async def test_build_process_content_request(valid_env_vars):
logger = MicrosoftPurviewLogger()
payload = {
"trace_id": "test-trace-123",
"startTime": 1700000000,
"endTime": 1700000010,
"messages": [{"role": "user", "content": "What is 2+2?"}],
"response": {"choices": [{"message": {"content": "4"}}]},
}
req = logger._build_process_content_request(payload)
assert "contentToProcess" in req
content_to_process = req["contentToProcess"]
assert "contentEntries" in content_to_process
assert len(content_to_process["contentEntries"]) == 2
prompt_entry = content_to_process["contentEntries"][0]
assert prompt_entry["identifier"] == "test-trace-123-prompt"
assert prompt_entry["name"] == "LLM Prompt"
assert prompt_entry["content"]["data"] == "[user]: What is 2+2?"
assert prompt_entry["sequenceNumber"] == 0
assert prompt_entry["correlationId"] == "test-trace-123"
response_entry = content_to_process["contentEntries"][1]
assert response_entry["identifier"] == "test-trace-123-response"
assert response_entry["name"] == "LLM Response"
assert response_entry["content"]["data"] == "4"
assert response_entry["sequenceNumber"] == 1
assert response_entry["correlationId"] == "test-trace-123"
# Verify metadata
assert content_to_process["integratedAppMetadata"]["name"] == "test-app"
assert (
content_to_process["protectedAppMetadata"]["applicationLocation"]["value"]
== "test-app-id"
)
@pytest.mark.asyncio
async def test_async_log_success_event_queues(valid_env_vars):
logger = MicrosoftPurviewLogger(batch_size=5)
kwargs = {"model": "gpt-4", "standard_logging_object": {"trace_id": "1"}}
await logger.async_log_success_event(kwargs, None, None, None)
assert len(logger.log_queue) == 1
assert logger.log_queue[0]["trace_id"] == "1"
@pytest.mark.asyncio
async def test_async_log_failure_event_queues(valid_env_vars):
logger = MicrosoftPurviewLogger(batch_size=5)
kwargs = {"model": "gpt-4", "standard_logging_object": {"trace_id": "2"}}
await logger.async_log_failure_event(kwargs, None, None, None)
assert len(logger.log_queue) == 1
assert logger.log_queue[0]["trace_id"] == "2"
@pytest.mark.asyncio
async def test_async_send_batch_success(valid_env_vars, monkeypatch):
logger = MicrosoftPurviewLogger()
# Add dummy payload to queue
payload = {
"trace_id": "test-trace-123",
"metadata": {"user_api_key_user_id": "user-A"},
"messages": ["test message"],
}
logger.log_queue.append(payload)
# Mock token
async def mock_get_token():
return "fake-token"
logger._get_oauth_token = mock_get_token
# Mock the http client
mock_post = AsyncMock()
mock_response = MagicMock()
mock_response.status_code = 200
mock_post.return_value = mock_response
logger.async_httpx_client.post = mock_post
await logger.async_send_batch()
# Check that HTTP post was called correctly
assert mock_post.called
assert len(logger.log_queue) == 0
call_kwargs = mock_post.call_args[1]
assert "url" in call_kwargs
assert "users/user-A/dataSecurityAndGovernance/processContent" in call_kwargs["url"]
assert "Bearer fake-token" in call_kwargs["headers"]["Authorization"]
@pytest.mark.asyncio
async def test_oauth_token_caching(valid_env_vars):
logger = MicrosoftPurviewLogger()
import time
logger.oauth_token = "cached-token"
logger.oauth_token_expires_at = time.time() + 3600
mock_post = AsyncMock()
logger.async_httpx_client.post = mock_post
token = await logger._get_oauth_token()
assert token == "cached-token"
assert not mock_post.called
@pytest.mark.asyncio
async def test_oauth_token_refresh(valid_env_vars):
logger = MicrosoftPurviewLogger()
# Expired token
import time
logger.oauth_token = "expired-token"
logger.oauth_token_expires_at = time.time() - 3600
# Mock token response
mock_post = AsyncMock()
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = {"access_token": "new-token", "expires_in": 3600}
mock_post.return_value = mock_response
logger.async_httpx_client.post = mock_post
token = await logger._get_oauth_token()
assert token == "new-token"
assert logger.oauth_token == "new-token"
assert mock_post.called
assert logger.oauth_token_expires_at > time.time()