mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat add microsfot purview logger
This commit is contained in:
parent
08bb3808bc
commit
20ca676507
7 changed files with 1478 additions and 521 deletions
File diff suppressed because it is too large
Load diff
|
|
@ -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"
|
||||
}
|
||||
]
|
||||
]
|
||||
0
litellm/integrations/microsoft_purview/__init__.py
Normal file
0
litellm/integrations/microsoft_purview/__init__.py
Normal file
424
litellm/integrations/microsoft_purview/microsoft_purview.py
Normal file
424
litellm/integrations/microsoft_purview/microsoft_purview.py
Normal 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()
|
||||
61
litellm/integrations/microsoft_purview/types.py
Normal file
61
litellm/integrations/microsoft_purview/types.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
257
tests/litellm/integrations/test_microsoft_purview.py
Normal file
257
tests/litellm/integrations/test_microsoft_purview.py
Normal 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()
|
||||
Loading…
Add table
Reference in a new issue