mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/lucid-burnell-e0be3a
This commit is contained in:
commit
a7feddba43
13 changed files with 4280 additions and 68 deletions
|
|
@ -2541,7 +2541,6 @@ jobs:
|
|||
paths:
|
||||
- litellm-docker-database.tar.zst
|
||||
|
||||
|
||||
test_bad_database_url:
|
||||
machine:
|
||||
image: ubuntu-2204:2024.04.1
|
||||
|
|
|
|||
|
|
@ -1204,12 +1204,8 @@ def get_last_user_message(messages: List[AllMessageValues]) -> Optional[str]:
|
|||
{"role": "assistant", "content": "I'm good, thank you!"},
|
||||
{"role": "user", "content": "What is the weather in Tokyo?"},
|
||||
]
|
||||
get_user_prompt(messages) -> "What is the weather in Tokyo?"
|
||||
get_last_user_message(messages) -> "What is the weather in Tokyo?"
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
)
|
||||
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ import hashlib
|
|||
import json
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
from typing import Any, Callable, Dict, List, Literal, Optional, Set, Tuple, Union, cast
|
||||
from urllib.parse import urlparse
|
||||
|
||||
|
|
@ -250,6 +251,10 @@ class MCPServerManager:
|
|||
}
|
||||
"""
|
||||
self._upstream_initialize_instructions_by_server_id: Dict[str, str] = {}
|
||||
# Per-server monotonic timestamp of last upstream prefetch attempt (success,
|
||||
# empty result, or failure). Used to throttle re-probes for servers that do
|
||||
# not return instructions, and to apply a short cooldown after failures.
|
||||
self._upstream_initialize_instructions_probed_at: Dict[str, float] = {}
|
||||
|
||||
def _remember_upstream_initialize_instructions(
|
||||
self, server: MCPServer, client: MCPClient
|
||||
|
|
@ -260,6 +265,80 @@ class MCPServerManager:
|
|||
raw
|
||||
).strip()
|
||||
|
||||
async def _ensure_upstream_initialize_instructions_cached(
|
||||
self, server: MCPServer
|
||||
) -> None:
|
||||
"""
|
||||
Open one upstream session and cache InitializeResult.instructions if missing.
|
||||
|
||||
No-op when:
|
||||
- YAML/DB instructions are set on the server record,
|
||||
- server is OpenAPI (spec_path),
|
||||
- non-empty upstream instructions are already cached,
|
||||
- auth preconditions match health_check_server's skip rules
|
||||
(per-user auth / missing static auth token),
|
||||
- a prior probe attempt for this server is within
|
||||
MCP_HEALTH_CHECK_TIMEOUT seconds (the probe is a health-check-shaped
|
||||
op and already uses this knob for its inner call timeout; reusing it
|
||||
as the cooldown avoids reconnecting on every gateway initialize when
|
||||
upstream returns empty or fails).
|
||||
"""
|
||||
if server.spec_path:
|
||||
return
|
||||
if server.instructions and server.instructions.strip():
|
||||
return
|
||||
if self._upstream_initialize_instructions_by_server_id.get(server.server_id):
|
||||
return
|
||||
if server.requires_per_user_auth:
|
||||
return
|
||||
if (
|
||||
server.auth_type
|
||||
and server.auth_type != MCPAuth.none
|
||||
and server.auth_type != MCPAuth.aws_sigv4
|
||||
and not server.authentication_token
|
||||
):
|
||||
return
|
||||
|
||||
last_probed_at = self._upstream_initialize_instructions_probed_at.get(
|
||||
server.server_id
|
||||
)
|
||||
if (
|
||||
last_probed_at is not None
|
||||
and (time.monotonic() - last_probed_at) < MCP_HEALTH_CHECK_TIMEOUT
|
||||
):
|
||||
return
|
||||
|
||||
# Record the attempt up-front so that a failure / empty response does not
|
||||
# cause every subsequent initialize request to re-open the upstream session.
|
||||
self._upstream_initialize_instructions_probed_at[server.server_id] = (
|
||||
time.monotonic()
|
||||
)
|
||||
|
||||
try:
|
||||
extra_headers: Optional[Dict[str, str]] = (
|
||||
dict(server.static_headers) if server.static_headers else None
|
||||
)
|
||||
client = await self._create_mcp_client(
|
||||
server=server,
|
||||
mcp_auth_header=None,
|
||||
extra_headers=extra_headers,
|
||||
stdio_env=None,
|
||||
)
|
||||
|
||||
async def _noop(_session):
|
||||
return "ok"
|
||||
|
||||
await asyncio.wait_for(
|
||||
client.run_with_session(_noop), timeout=MCP_HEALTH_CHECK_TIMEOUT
|
||||
)
|
||||
self._remember_upstream_initialize_instructions(server, client)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
"Upstream initialize instructions prefetch failed for %s: %s",
|
||||
server.name,
|
||||
e,
|
||||
)
|
||||
|
||||
def get_registry(self) -> Dict[str, MCPServer]:
|
||||
"""
|
||||
Get the registered MCP Servers from the registry and union with the config MCP Servers
|
||||
|
|
@ -280,6 +359,7 @@ class MCPServerManager:
|
|||
"""
|
||||
verbose_logger.debug("Loading MCP Servers from config-----")
|
||||
self._upstream_initialize_instructions_by_server_id.clear()
|
||||
self._upstream_initialize_instructions_probed_at.clear()
|
||||
|
||||
# Track which aliases have been used to ensure only first occurrence is used
|
||||
used_aliases = set()
|
||||
|
|
@ -3141,6 +3221,7 @@ class MCPServerManager:
|
|||
|
||||
verbose_logger.debug("Loading MCP servers from database into registry...")
|
||||
self._upstream_initialize_instructions_by_server_id.clear()
|
||||
self._upstream_initialize_instructions_probed_at.clear()
|
||||
|
||||
# perform authz check to filter the mcp servers user has access to
|
||||
prisma_client = get_prisma_client_or_throw(
|
||||
|
|
|
|||
|
|
@ -1165,7 +1165,7 @@ if MCP_AVAILABLE:
|
|||
def _merge_gateway_initialize_instructions(
|
||||
allowed_mcp_servers: List[MCPServer],
|
||||
) -> Optional[str]:
|
||||
"""YAML/DB override, else in-memory upstream text from list_tools / health_check / call_tool."""
|
||||
"""YAML/DB override, else upstream text (prefetch on init, or list_tools / health_check / call_tool cache)."""
|
||||
if not allowed_mcp_servers:
|
||||
return None
|
||||
|
||||
|
|
@ -1206,6 +1206,20 @@ if MCP_AVAILABLE:
|
|||
mcp_servers=mcp_servers,
|
||||
client_ip=client_ip,
|
||||
)
|
||||
if allowed:
|
||||
# return_exceptions=True: a per-server probe failure (incl. CancelledError
|
||||
# bubbled from anyio task group teardown on connection refused) must not
|
||||
# cancel sibling probes or 500 the gateway initialize request.
|
||||
await asyncio.gather(
|
||||
*[
|
||||
global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(
|
||||
s
|
||||
)
|
||||
for s in allowed
|
||||
if s is not None
|
||||
],
|
||||
return_exceptions=True,
|
||||
)
|
||||
merged = _merge_gateway_initialize_instructions(allowed_mcp_servers=allowed)
|
||||
tok = _mcp_gateway_initialize_instructions.set(merged)
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -2,6 +2,9 @@ import re
|
|||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_last_user_message,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
|
|
@ -134,32 +137,4 @@ class AzureGuardrailBase:
|
|||
]
|
||||
get_user_prompt(messages) -> "What is the weather in Tokyo?"
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
)
|
||||
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
# Iterate from the end to find the last consecutive block of user messages
|
||||
user_messages = []
|
||||
for message in reversed(messages):
|
||||
if message.get("role") == "user":
|
||||
user_messages.append(message)
|
||||
else:
|
||||
# Stop when we hit a non-user message
|
||||
break
|
||||
|
||||
if not user_messages:
|
||||
return None
|
||||
|
||||
# Reverse to get the messages in chronological order
|
||||
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
|
||||
return get_last_user_message(messages)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,57 @@
|
|||
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
|
||||
|
||||
tenant_id = getattr(litellm_params, "tenant_id", None)
|
||||
client_id = getattr(litellm_params, "client_id", None)
|
||||
|
||||
# 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 getattr(
|
||||
litellm_params, "client_secret", None
|
||||
)
|
||||
|
||||
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")
|
||||
|
||||
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(
|
||||
getattr(litellm_params, "purview_app_name", None) or "LiteLLM"
|
||||
),
|
||||
user_id_field=str(getattr(litellm_params, "user_id_field", None) or "user_id"),
|
||||
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,515 @@
|
|||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from collections import OrderedDict
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
)
|
||||
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)
|
||||
# Capped at 1000 entries (LRU eviction) to avoid unbounded growth.
|
||||
self._scope_cache: OrderedDict[str, Tuple[str, Dict[str, Any], float]] = (
|
||||
OrderedDict()
|
||||
)
|
||||
self._scope_cache_maxsize = 1000
|
||||
# Use a threading.Lock (not asyncio.Lock) because this lock is acquired
|
||||
# from both the proxy's main asyncio event loop and from short-lived
|
||||
# event loops created by the logging_hook thread fallback. In Python
|
||||
# 3.10+ an asyncio.Lock is bound to the first event loop that acquires
|
||||
# it and raises RuntimeError from any other loop, which would silently
|
||||
# break audit logging via the thread fallback. All critical sections
|
||||
# below are pure in-memory dict ops with no awaits, so a synchronous
|
||||
# lock is both correct and sufficient.
|
||||
self._cache_lock = threading.Lock()
|
||||
|
||||
@staticmethod
|
||||
def _encode_graph_user_id(user_id: str) -> str:
|
||||
"""Percent-encode Entra user id for Graph ``/users/{id}/...`` path segments."""
|
||||
return encode_url_path_segment(user_id, field_name="user_id")
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# OAuth2 token management
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _get_access_token(self) -> str:
|
||||
"""Acquire or return cached OAuth2 token via client_credentials grant."""
|
||||
now = time.time()
|
||||
with self._cache_lock:
|
||||
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"},
|
||||
)
|
||||
response.raise_for_status()
|
||||
token_data = response.json()
|
||||
access_token = token_data["access_token"]
|
||||
expires_in = int(token_data.get("expires_in", 3599))
|
||||
# Recompute ``now`` after the await so the expiry reflects when the
|
||||
# token was actually received, not when the request started.
|
||||
with self._cache_lock:
|
||||
self._token_cache = (access_token, time.time() + 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.raise_for_status()
|
||||
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).
|
||||
"""
|
||||
encoded_user_id = self._encode_graph_user_id(user_id)
|
||||
now = time.time()
|
||||
|
||||
with self._cache_lock:
|
||||
cached = self._scope_cache.get(user_id)
|
||||
if cached and (now - cached[2]) < SCOPE_CACHE_TTL_SECONDS:
|
||||
self._scope_cache.move_to_end(user_id)
|
||||
return cached[0], cached[1]
|
||||
|
||||
url = (
|
||||
f"{GRAPH_API_BASE}/users/{encoded_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", ""))
|
||||
|
||||
# Recompute ``now`` after the await so the TTL reflects when the
|
||||
# scope response was actually received, not when the request started.
|
||||
fetched_at = time.time()
|
||||
with self._cache_lock:
|
||||
self._scope_cache[user_id] = (etag, response_json, fetched_at)
|
||||
# Move refreshed entry to the end so it is treated as most-recently-used.
|
||||
# OrderedDict.__setitem__ preserves existing insertion order for known
|
||||
# keys, so an explicit move_to_end() call is required.
|
||||
self._scope_cache.move_to_end(user_id)
|
||||
# Evict least-recently-used entry when cache exceeds max size.
|
||||
while len(self._scope_cache) > self._scope_cache_maxsize:
|
||||
self._scope_cache.popitem(last=False)
|
||||
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.
|
||||
"""
|
||||
encoded_user_id = self._encode_graph_user_id(user_id)
|
||||
url = (
|
||||
f"{GRAPH_API_BASE}/users/{encoded_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":
|
||||
with self._cache_lock:
|
||||
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.
|
||||
|
||||
Returns the strongest available identity walking down four sources, in
|
||||
decreasing trust order:
|
||||
|
||||
1. ``user_api_key_dict.user_id`` — LiteLLM key / JWT-bound user
|
||||
2. ``user_api_key_dict.end_user_id`` — request-derived
|
||||
3. ``metadata["user_api_key_user_id"]`` — proxy-injected from the key
|
||||
4. ``metadata[user_id_field]`` — caller-supplied
|
||||
|
||||
Used only by blocking-mode resolution to disambiguate "no identity at
|
||||
all" from "caller supplied an untrusted identity" for the error
|
||||
message. Neither blocking nor audit DLP feeds the untrusted
|
||||
fallbacks (2, 4) into Purview itself.
|
||||
"""
|
||||
trusted = self._resolve_trusted_user_id(data, user_api_key_dict)
|
||||
if trusted:
|
||||
return trusted
|
||||
|
||||
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)
|
||||
|
||||
metadata = data.get("metadata") or data.get("litellm_metadata") or {}
|
||||
uid = metadata.get("user_api_key_user_id")
|
||||
if uid:
|
||||
return str(uid)
|
||||
|
||||
uid = metadata.get(self.user_id_field)
|
||||
if uid:
|
||||
return str(uid)
|
||||
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _logging_kwargs_metadata(kwargs: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Metadata dict from ``model_call_details`` / logging kwargs."""
|
||||
litellm_params = kwargs.get("litellm_params") or {}
|
||||
if not isinstance(litellm_params, dict):
|
||||
return {}
|
||||
md = litellm_params.get("metadata")
|
||||
return md if isinstance(md, dict) else {}
|
||||
|
||||
def _resolve_trusted_user_id(
|
||||
self, data: Dict[str, Any], user_api_key_dict: Any
|
||||
) -> Optional[str]:
|
||||
"""Resolve user ID from API-key/JWT-bound identity for blocking DLP.
|
||||
|
||||
Uses only ``UserAPIKeyAuth.user_id`` (bound on the LiteLLM key or JWT).
|
||||
Intentionally omits ``UserAPIKeyAuth.end_user_id`` because the proxy sets
|
||||
it from caller-controlled request fields (``user``, ``metadata.user_id``,
|
||||
``safety_identifier``, custom headers, etc.) via
|
||||
``get_end_user_id_from_request_body``.
|
||||
|
||||
Also omits ``metadata[user_id_field]`` and
|
||||
``metadata["user_api_key_user_id"]`` for the same impersonation risk when
|
||||
the key has no bound user.
|
||||
|
||||
Returns ``None`` when no authenticated identity is available. Blocking
|
||||
hooks must fail closed rather than skip the DLP check.
|
||||
"""
|
||||
if hasattr(user_api_key_dict, "user_id") and user_api_key_dict.user_id:
|
||||
return str(user_api_key_dict.user_id)
|
||||
|
||||
return None
|
||||
|
||||
def _resolve_user_id_from_logging_kwargs(
|
||||
self, kwargs: Dict[str, Any]
|
||||
) -> Optional[str]:
|
||||
"""Trusted-identity-only resolver for logging-only hooks.
|
||||
|
||||
Uses only the proxy-injected ``user_api_key_user_id`` (populated from
|
||||
the API-key/JWT-bound ``UserAPIKeyAuth.user_id`` after the proxy
|
||||
strips every caller-supplied ``user_api_key_*`` key from the request
|
||||
metadata). Caller-influenceable sources (``user_api_key_end_user_id``,
|
||||
``metadata[user_id_field]``) are not used here so a caller cannot
|
||||
cause Purview audit records to be written under a victim's identity.
|
||||
Returns ``None`` when no trusted identity is available so the audit
|
||||
is skipped rather than misattributed.
|
||||
"""
|
||||
md = self._logging_kwargs_metadata(kwargs)
|
||||
uid = md.get("user_api_key_user_id") or kwargs.get("user_api_key_user_id")
|
||||
if uid:
|
||||
return str(uid)
|
||||
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
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Prompt text for DLP
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def is_token_id_prompt(prompt: Any) -> bool:
|
||||
"""Return True if ``prompt`` carries OpenAI completions token ids.
|
||||
|
||||
Covers every list shape that ``completion_prompt_to_str`` cannot decode
|
||||
for Purview, including flat ``list[int]`` (single token-id prompt),
|
||||
``list[list[int]]`` (multi-prompt token-id batches), and mixed lists
|
||||
that include any token-id sub-array.
|
||||
"""
|
||||
if not isinstance(prompt, list) or not prompt:
|
||||
return False
|
||||
for x in prompt:
|
||||
if isinstance(x, int):
|
||||
return True
|
||||
if isinstance(x, list) and x and any(isinstance(y, int) for y in x):
|
||||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def completion_prompt_to_str(prompt: Any) -> Optional[str]:
|
||||
"""Normalize OpenAI ``/v1/completions`` ``prompt`` for text DLP.
|
||||
|
||||
Supports string prompts and list-of-string prompts. List-of-token-id prompts
|
||||
are skipped (no plaintext for Purview to evaluate).
|
||||
"""
|
||||
if prompt is None:
|
||||
return None
|
||||
if isinstance(prompt, str):
|
||||
stripped = prompt.strip()
|
||||
return stripped or None
|
||||
if isinstance(prompt, list) and prompt:
|
||||
if all(isinstance(x, str) for x in prompt):
|
||||
joined = "\n".join(s.strip() for s in prompt if isinstance(s, str))
|
||||
return joined.strip() or None
|
||||
if all(isinstance(x, int) for x in prompt):
|
||||
verbose_proxy_logger.debug(
|
||||
"Purview DLP: completions prompt is token ids only; skipping text scan"
|
||||
)
|
||||
return None
|
||||
str_parts = [x for x in prompt if isinstance(x, str)]
|
||||
if str_parts:
|
||||
joined = "\n".join(s.strip() for s in str_parts)
|
||||
return joined.strip() or None
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _extract_tool_call_args_from_message(message: Any) -> List[str]:
|
||||
"""Return plaintext arguments strings from tool_calls and function_call fields.
|
||||
|
||||
Covers both the request path (assistant messages in chat histories that
|
||||
carry tool_calls / function_call) and the response path (model-generated
|
||||
tool calls returned in a ModelResponse). Both dict-style and object-style
|
||||
representations are handled.
|
||||
"""
|
||||
args: List[str] = []
|
||||
|
||||
# tool_calls: [{"function": {"arguments": "..."}}]
|
||||
tool_calls = (
|
||||
message.get("tool_calls")
|
||||
if isinstance(message, dict)
|
||||
else getattr(message, "tool_calls", None)
|
||||
)
|
||||
if tool_calls:
|
||||
for tc in tool_calls:
|
||||
fn = (
|
||||
tc.get("function")
|
||||
if isinstance(tc, dict)
|
||||
else getattr(tc, "function", None)
|
||||
)
|
||||
if fn is None:
|
||||
continue
|
||||
arguments = (
|
||||
fn.get("arguments")
|
||||
if isinstance(fn, dict)
|
||||
else getattr(fn, "arguments", None)
|
||||
)
|
||||
if isinstance(arguments, str) and arguments.strip():
|
||||
args.append(arguments)
|
||||
|
||||
# Legacy function_call: {"arguments": "..."}
|
||||
function_call = (
|
||||
message.get("function_call")
|
||||
if isinstance(message, dict)
|
||||
else getattr(message, "function_call", None)
|
||||
)
|
||||
if function_call is not None:
|
||||
arguments = (
|
||||
function_call.get("arguments")
|
||||
if isinstance(function_call, dict)
|
||||
else getattr(function_call, "arguments", None)
|
||||
)
|
||||
if isinstance(arguments, str) and arguments.strip():
|
||||
args.append(arguments)
|
||||
|
||||
return args
|
||||
|
||||
def get_prompt_text_for_dlp(
|
||||
self, messages: List["AllMessageValues"]
|
||||
) -> Optional[str]:
|
||||
"""Concatenate text from every chat message (all roles) for pre-call DLP.
|
||||
|
||||
Evaluates the same payload the model receives, not only the trailing user
|
||||
turn. Each message is separated by ``\\n\\n`` so that tokens at message
|
||||
boundaries are not merged (e.g., ``"end of msg1\\n\\nstart of msg2"``
|
||||
rather than ``"end of msg1start of msg2"``), which preserves DLP pattern
|
||||
detection accuracy across message boundaries.
|
||||
|
||||
Tool-call arguments (``tool_calls[].function.arguments`` and
|
||||
``function_call.arguments``) are included alongside message content so
|
||||
that sensitive data hidden in function arguments is not bypassed.
|
||||
"""
|
||||
if not messages:
|
||||
return None
|
||||
parts: List[str] = []
|
||||
for msg in messages:
|
||||
segments: List[str] = []
|
||||
content = convert_content_list_to_str(message=msg).strip()
|
||||
if content:
|
||||
segments.append(content)
|
||||
segments.extend(self._extract_tool_call_args_from_message(msg))
|
||||
combined = "\n".join(segments)
|
||||
if combined.strip():
|
||||
parts.append(combined.strip())
|
||||
text = "\n\n".join(parts)
|
||||
return text or None
|
||||
|
|
@ -0,0 +1,734 @@
|
|||
"""
|
||||
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
|
||||
import threading
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
AsyncGenerator,
|
||||
Dict,
|
||||
List,
|
||||
Optional,
|
||||
Tuple,
|
||||
Type,
|
||||
Union,
|
||||
cast,
|
||||
)
|
||||
|
||||
import httpx
|
||||
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 litellm.types.utils import (
|
||||
Choices,
|
||||
GuardrailStatus,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
ResponsesAPIResponse,
|
||||
TextChoices,
|
||||
TextCompletionResponse,
|
||||
)
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
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",
|
||||
**kwargs: Any,
|
||||
):
|
||||
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.guardrail_provider = "microsoft_purview"
|
||||
verbose_proxy_logger.info(
|
||||
"Initialized Microsoft Purview DLP Guardrail: %s",
|
||||
guardrail_name,
|
||||
)
|
||||
|
||||
@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: GuardrailStatus = "success"
|
||||
response: Dict[str, Any] = {}
|
||||
|
||||
try:
|
||||
etag, _ = await self._compute_protection_scopes(user_id)
|
||||
correlation_id = request_data.get("litellm_call_id") or str(uuid.uuid4())
|
||||
response = await self._process_content(
|
||||
user_id=user_id,
|
||||
text=text,
|
||||
activity=activity,
|
||||
etag=etag,
|
||||
correlation_id=correlation_id,
|
||||
)
|
||||
|
||||
if self._should_block(response):
|
||||
status = "guardrail_intervened"
|
||||
except HTTPException:
|
||||
status = "guardrail_failed_to_respond"
|
||||
raise
|
||||
except httpx.HTTPStatusError as exc:
|
||||
# Preserve the upstream Graph API status code (e.g. 429, 503) so
|
||||
# callers can distinguish a transient infrastructure error from a
|
||||
# DLP policy block (signaled separately as HTTP 400 below) and can
|
||||
# implement retry-after handling on rate limits. 401/403 upstream
|
||||
# responses indicate a proxy-side credential / consent problem the
|
||||
# caller can do nothing about, so they are mapped to 502.
|
||||
status = "guardrail_failed_to_respond"
|
||||
if block_on_violation:
|
||||
upstream_status = exc.response.status_code
|
||||
client_status = (
|
||||
502 if upstream_status in (401, 403) else upstream_status
|
||||
)
|
||||
headers: Optional[Dict[str, str]] = None
|
||||
retry_after = exc.response.headers.get("retry-after")
|
||||
if retry_after:
|
||||
headers = {"Retry-After": retry_after}
|
||||
raise HTTPException(
|
||||
status_code=client_status,
|
||||
detail={
|
||||
"error": "Microsoft Purview DLP: upstream policy evaluation failed",
|
||||
"activity": activity,
|
||||
"upstream_status": upstream_status,
|
||||
"exception": str(exc),
|
||||
},
|
||||
headers=headers,
|
||||
) from exc
|
||||
verbose_proxy_logger.warning(
|
||||
"Purview DLP: API/network error in logging-only mode (not re-raised): %s",
|
||||
exc,
|
||||
)
|
||||
except Exception as exc:
|
||||
status = "guardrail_failed_to_respond"
|
||||
if block_on_violation:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Microsoft Purview DLP: upstream policy evaluation failed",
|
||||
"activity": activity,
|
||||
"exception": str(exc),
|
||||
},
|
||||
) from exc
|
||||
verbose_proxy_logger.warning(
|
||||
"Purview DLP: API/network error in logging-only mode (not re-raised): %s",
|
||||
exc,
|
||||
)
|
||||
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
|
||||
|
||||
@staticmethod
|
||||
def _extract_responses_api_function_call_args(result: Any) -> List[str]:
|
||||
"""Return tool-call argument strings from a ``ResponsesAPIResponse.output``.
|
||||
|
||||
``ResponsesAPIResponse.output_text`` only aggregates ``output_text``
|
||||
content blocks and ignores ``function_call`` items. Model-generated
|
||||
tool-call arguments can themselves contain sensitive data, so we
|
||||
extract them explicitly to keep DLP coverage consistent with the
|
||||
chat (``ModelResponse``) path.
|
||||
"""
|
||||
args: List[str] = []
|
||||
output = getattr(result, "output", None)
|
||||
if not output:
|
||||
return args
|
||||
for item in output:
|
||||
if isinstance(item, dict):
|
||||
item_type = item.get("type")
|
||||
arguments = item.get("arguments")
|
||||
else:
|
||||
item_type = getattr(item, "type", None)
|
||||
arguments = getattr(item, "arguments", None)
|
||||
if item_type == "function_call" and isinstance(arguments, str):
|
||||
if arguments.strip():
|
||||
args.append(arguments)
|
||||
return args
|
||||
|
||||
def _completion_response_text_parts(self, result: Any) -> List[str]:
|
||||
"""Collect non-empty text segments from chat, text completions, or responses API.
|
||||
|
||||
Includes assistant message content *and* model-generated tool-call
|
||||
arguments so that sensitive data returned inside function calls is not
|
||||
missed by the DLP scan.
|
||||
"""
|
||||
parts: List[str] = []
|
||||
if isinstance(result, TextCompletionResponse) and result.choices:
|
||||
for text_choice in result.choices:
|
||||
if not isinstance(text_choice, TextChoices):
|
||||
continue
|
||||
raw = text_choice.get("text")
|
||||
if isinstance(raw, str) and raw.strip():
|
||||
parts.append(raw)
|
||||
elif isinstance(result, ResponsesAPIResponse):
|
||||
text = result.output_text
|
||||
if text and text.strip():
|
||||
parts.append(text)
|
||||
# Include tool-call arguments from ``function_call`` output items
|
||||
# (``output_text`` ignores them).
|
||||
parts.extend(self._extract_responses_api_function_call_args(result))
|
||||
elif isinstance(result, ModelResponse) and result.choices:
|
||||
for chat_choice in result.choices:
|
||||
if not isinstance(chat_choice, Choices):
|
||||
continue
|
||||
msg = chat_choice.message
|
||||
if msg is None:
|
||||
continue
|
||||
raw = (
|
||||
msg.get("content")
|
||||
if isinstance(msg, dict)
|
||||
else getattr(msg, "content", None)
|
||||
)
|
||||
if isinstance(raw, str) and raw.strip():
|
||||
parts.append(raw)
|
||||
# Include tool-call arguments returned by the model
|
||||
parts.extend(self._extract_tool_call_args_from_message(msg))
|
||||
return parts
|
||||
|
||||
def _assemble_responses_api_from_chunks(
|
||||
self, chunks: List[Any]
|
||||
) -> Tuple[bool, Optional[ResponsesAPIResponse]]:
|
||||
"""Extract the final ``ResponsesAPIResponse`` from a buffered Responses API stream.
|
||||
|
||||
Returns a ``(is_responses_api_stream, assembled)`` tuple so the caller
|
||||
can distinguish "not a Responses API stream" (fall through to
|
||||
``stream_chunk_builder``) from "Responses API stream but no final
|
||||
response event was received" (fail closed with an accurate error).
|
||||
When the stream is a Responses API stream the latest event carrying a
|
||||
``ResponsesAPIResponse`` body is returned (``response.completed``, or
|
||||
``response.failed`` / ``response.incomplete`` as fallbacks).
|
||||
"""
|
||||
looks_like_responses_api = False
|
||||
final: Optional[ResponsesAPIResponse] = None
|
||||
for chunk in chunks:
|
||||
event_type = getattr(chunk, "type", None)
|
||||
if isinstance(event_type, str) and event_type.startswith("response."):
|
||||
looks_like_responses_api = True
|
||||
candidate = getattr(chunk, "response", None)
|
||||
if isinstance(candidate, ResponsesAPIResponse):
|
||||
final = candidate
|
||||
return looks_like_responses_api, final
|
||||
|
||||
def _responses_api_input_to_str(
|
||||
self, data: Dict[str, Any], raise_on_failure: bool = False
|
||||
) -> Optional[str]:
|
||||
"""Extract DLP-scannable text from a Responses API request ``input`` field.
|
||||
|
||||
``input`` may be a plain string or a list of input items (messages). In
|
||||
the latter case the items are converted to chat messages via the standard
|
||||
LiteLLM transformation and then concatenated by ``get_prompt_text_for_dlp``.
|
||||
|
||||
When ``raise_on_failure`` is True (blocking mode), a transformation error
|
||||
raises ``HTTPException`` so the request is fail-closed. In logging-only
|
||||
mode the error is swallowed and ``None`` is returned so audit attempts on
|
||||
the response side can still run.
|
||||
"""
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
|
||||
input_data = data.get("input")
|
||||
if input_data is None and not data.get("instructions"):
|
||||
return None
|
||||
try:
|
||||
# Always transform via messages so ``instructions`` become a system message
|
||||
# (string ``input`` alone would skip instructions and bypass DLP).
|
||||
messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
|
||||
input=input_data if input_data is not None else "",
|
||||
responses_api_request=data,
|
||||
)
|
||||
return self.get_prompt_text_for_dlp(cast(List[Any], messages))
|
||||
except Exception:
|
||||
verbose_proxy_logger.warning(
|
||||
"Purview DLP: failed to transform responses API input",
|
||||
exc_info=True,
|
||||
)
|
||||
if raise_on_failure:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": (
|
||||
"Microsoft Purview DLP: Responses API input could "
|
||||
"not be transformed for DLP scanning in blocking mode"
|
||||
),
|
||||
},
|
||||
)
|
||||
return None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Identity resolution for blocking modes
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _resolve_user_id_for_blocking(
|
||||
self,
|
||||
data: Dict[str, Any],
|
||||
user_api_key_dict: Any,
|
||||
) -> str:
|
||||
"""Resolve user ID for blocking (pre_call / post_call) DLP hooks.
|
||||
|
||||
Uses only trusted proxy-authenticated sources (``_resolve_trusted_user_id``).
|
||||
Caller-supplied ``UserAPIKeyAuth.end_user_id`` (from request ``user``,
|
||||
``metadata.user_id``, ``safety_identifier``, etc.) and
|
||||
``metadata[user_id_field]`` are rejected (fail closed) because they can
|
||||
impersonate another Entra user's Purview policy.
|
||||
|
||||
Raises ``HTTPException`` when no API-key-bound ``user_id`` exists or when
|
||||
only caller-influenceable identity fields are available (fail closed).
|
||||
"""
|
||||
trusted_id = self._resolve_trusted_user_id(data, user_api_key_dict)
|
||||
if trusted_id:
|
||||
return trusted_id
|
||||
|
||||
if self._resolve_user_id(data, user_api_key_dict):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": (
|
||||
"Microsoft Purview DLP: No proxy-authenticated user identity; "
|
||||
"bind user_id to the API key (caller-supplied metadata cannot "
|
||||
"be used for blocking DLP)"
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": (
|
||||
"Microsoft Purview DLP: No proxy-authenticated user identity; "
|
||||
"bind user_id to the API key for blocking DLP"
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# 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_for_blocking(data, user_api_key_dict)
|
||||
|
||||
prompt_text: Optional[str] = None
|
||||
if call_type in ("responses", "aresponses"):
|
||||
# Route Responses API calls to the responses-specific extractor
|
||||
# before the generic ``messages`` branch. This mirrors
|
||||
# ``async_logging_hook`` and ensures ``instructions`` (system
|
||||
# prompt) content is included in the DLP scan, and prevents a
|
||||
# crafted ``messages`` key in the request from being scanned in
|
||||
# place of the actual ``input``.
|
||||
prompt_text = self._responses_api_input_to_str(data, raise_on_failure=True)
|
||||
elif call_type in ("text_completion", "atext_completion"):
|
||||
raw_prompt = data.get("prompt")
|
||||
# Reject every token-id prompt shape Purview cannot evaluate —
|
||||
# flat ``list[int]`` (single prompt), ``list[list[int]]`` (multi-prompt
|
||||
# batches), and mixed lists that include any token-id sub-array.
|
||||
# Empty/whitespace-only strings also yield ``prompt_text is None`` but
|
||||
# contain no sensitive data and pass through harmlessly below.
|
||||
if self.is_token_id_prompt(raw_prompt):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": (
|
||||
"Microsoft Purview DLP: Token-id completion prompts "
|
||||
"cannot be scanned for DLP in blocking mode"
|
||||
),
|
||||
},
|
||||
)
|
||||
prompt_text = self.completion_prompt_to_str(raw_prompt)
|
||||
else:
|
||||
messages: Optional[List] = data.get("messages")
|
||||
if messages:
|
||||
prompt_text = self.get_prompt_text_for_dlp(cast(List[Any], messages))
|
||||
|
||||
if not prompt_text:
|
||||
return data
|
||||
|
||||
await self._check_content(
|
||||
user_id=user_id,
|
||||
text=prompt_text,
|
||||
activity="uploadText",
|
||||
request_data=data,
|
||||
block_on_violation=True,
|
||||
)
|
||||
return data
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Post-call hook — DLP on responses
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@log_guardrail_information
|
||||
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 (non-streaming only).
|
||||
|
||||
Streaming responses are handled by ``async_post_call_streaming_iterator_hook``
|
||||
which buffers all chunks before scanning. The proxy automatically skips
|
||||
this hook for requests that have a streaming iterator hook defined.
|
||||
"""
|
||||
user_id = self._resolve_user_id_for_blocking(data, user_api_key_dict)
|
||||
|
||||
parts = self._completion_response_text_parts(response)
|
||||
|
||||
if parts:
|
||||
combined = "\n\n---\n\n".join(parts)
|
||||
await self._check_content(
|
||||
user_id=user_id,
|
||||
text=combined,
|
||||
activity="downloadText",
|
||||
request_data=data,
|
||||
block_on_violation=True,
|
||||
)
|
||||
return response
|
||||
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: "UserAPIKeyAuth",
|
||||
response: Any,
|
||||
request_data: dict,
|
||||
) -> AsyncGenerator[ModelResponseStream, None]:
|
||||
"""Check streaming LLM responses against Purview DLP policies.
|
||||
|
||||
All chunks are buffered before the DLP scan so that no content is
|
||||
delivered to the client if a policy violation is detected. After a
|
||||
clean scan the assembled response is re-yielded chunk-by-chunk via a
|
||||
``MockResponseIterator`` so the caller receives normal streaming output.
|
||||
|
||||
The proxy automatically skips ``async_post_call_success_hook`` for
|
||||
guardrails that define this method, preventing duplicate scans.
|
||||
"""
|
||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||
from litellm.main import stream_chunk_builder
|
||||
|
||||
# Resolve user ID up-front so identity failures don't waste work
|
||||
# buffering and assembling the stream.
|
||||
user_id = self._resolve_user_id_for_blocking(request_data, user_api_key_dict)
|
||||
|
||||
# Buffer the entire stream before any DLP scan.
|
||||
all_chunks: List[ModelResponseStream] = []
|
||||
async for chunk in response:
|
||||
all_chunks.append(chunk)
|
||||
|
||||
# Responses API streams emit typed events (e.g. ``response.completed``)
|
||||
# whose final event carries the full ``ResponsesAPIResponse`` — these
|
||||
# are not understood by ``stream_chunk_builder`` (which is built for
|
||||
# chat/text-completion deltas). Detect and scan them via the same
|
||||
# ``_completion_response_text_parts`` path used by non-streaming.
|
||||
(
|
||||
is_responses_api_stream,
|
||||
responses_api_assembled,
|
||||
) = self._assemble_responses_api_from_chunks(all_chunks)
|
||||
if is_responses_api_stream:
|
||||
if responses_api_assembled is None:
|
||||
# Fail closed: Responses API events were seen but no final
|
||||
# ``response.completed`` / ``response.failed`` /
|
||||
# ``response.incomplete`` event carrying a ``ResponsesAPIResponse``
|
||||
# body was received, so we cannot scan the content.
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": (
|
||||
"Microsoft Purview DLP: Incomplete Responses API "
|
||||
"stream — no final response event received for "
|
||||
"DLP scanning; blocking response."
|
||||
),
|
||||
},
|
||||
)
|
||||
parts = self._completion_response_text_parts(responses_api_assembled)
|
||||
if parts:
|
||||
combined = "\n\n---\n\n".join(parts)
|
||||
await self._check_content(
|
||||
user_id=user_id,
|
||||
text=combined,
|
||||
activity="downloadText",
|
||||
request_data=request_data,
|
||||
block_on_violation=True,
|
||||
)
|
||||
for chunk in all_chunks:
|
||||
yield chunk
|
||||
return
|
||||
|
||||
assembled_response = stream_chunk_builder(chunks=all_chunks)
|
||||
|
||||
if assembled_response is None and all_chunks:
|
||||
# Fail closed: stream_chunk_builder dropped all chunks, so we cannot
|
||||
# scan the content. Refuse to release the buffered chunks.
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": (
|
||||
"Microsoft Purview DLP: Unable to assemble streamed "
|
||||
"response for scanning; blocking response."
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
if isinstance(
|
||||
assembled_response, (TextCompletionResponse, ResponsesAPIResponse)
|
||||
):
|
||||
parts = self._completion_response_text_parts(assembled_response)
|
||||
if parts:
|
||||
combined = "\n\n---\n\n".join(parts)
|
||||
await self._check_content(
|
||||
user_id=user_id,
|
||||
text=combined,
|
||||
activity="downloadText",
|
||||
request_data=request_data,
|
||||
block_on_violation=True,
|
||||
)
|
||||
for chunk in all_chunks:
|
||||
yield chunk
|
||||
return
|
||||
|
||||
if not isinstance(assembled_response, ModelResponse):
|
||||
# Non-content response (e.g. embeddings) — pass through unchanged.
|
||||
for chunk in all_chunks:
|
||||
yield chunk
|
||||
return
|
||||
|
||||
parts = self._completion_response_text_parts(assembled_response)
|
||||
if parts:
|
||||
combined = "\n\n---\n\n".join(parts)
|
||||
# Raises HTTPException(400) on violation — no chunks are yielded.
|
||||
await self._check_content(
|
||||
user_id=user_id,
|
||||
text=combined,
|
||||
activity="downloadText",
|
||||
request_data=request_data,
|
||||
block_on_violation=True,
|
||||
)
|
||||
|
||||
# DLP passed — re-yield chunks from the assembled chat response.
|
||||
mock_response = MockResponseIterator(model_response=assembled_response)
|
||||
async for chunk in mock_response:
|
||||
yield chunk
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Logging-only hook — audit without blocking
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def logging_hook(
|
||||
self, kwargs: dict, result: Any, call_type: str
|
||||
) -> Tuple[dict, Any]:
|
||||
"""Fire-and-forget async audit logging; returns original (kwargs, result) immediately.
|
||||
|
||||
In the proxy's async success path, litellm independently calls both
|
||||
``logging_hook`` (sync) and ``async_logging_hook`` (async) for every
|
||||
``CustomGuardrail`` callback. To avoid making two complete sets of
|
||||
Purview API calls per request, this sync hook is a no-op whenever an
|
||||
event loop is running — the framework's async path will invoke
|
||||
``async_logging_hook`` directly.
|
||||
|
||||
For genuine sync-only call paths (no running event loop, so the async
|
||||
success handler will not fire either), schedule ``async_logging_hook``
|
||||
on a short-lived background daemon thread so audit logging still runs
|
||||
without blocking the caller on two Graph API round-trips.
|
||||
"""
|
||||
|
||||
try:
|
||||
asyncio.get_running_loop()
|
||||
# Async context — let the framework's async success handler invoke
|
||||
# async_logging_hook to avoid duplicate Purview API calls. Log so
|
||||
# the deferral is observable if the framework ever stops dispatching
|
||||
# async_logging_hook on a given code path (otherwise audit silently
|
||||
# drops).
|
||||
verbose_proxy_logger.debug(
|
||||
"Purview audit: deferring to async_logging_hook (running event loop detected)"
|
||||
)
|
||||
return kwargs, result
|
||||
except RuntimeError:
|
||||
pass
|
||||
|
||||
async def _log_safe() -> None:
|
||||
try:
|
||||
await self.async_logging_hook(
|
||||
kwargs=kwargs, result=result, call_type=call_type
|
||||
)
|
||||
except Exception as exc:
|
||||
verbose_proxy_logger.error(
|
||||
"Purview audit background logging error: %s", exc
|
||||
)
|
||||
|
||||
def _run_in_new_loop() -> None:
|
||||
new_loop = asyncio.new_event_loop()
|
||||
try:
|
||||
asyncio.set_event_loop(new_loop)
|
||||
new_loop.run_until_complete(_log_safe())
|
||||
finally:
|
||||
new_loop.close()
|
||||
asyncio.set_event_loop(None)
|
||||
|
||||
thread = threading.Thread(target=_run_in_new_loop, daemon=True)
|
||||
thread.start()
|
||||
|
||||
return kwargs, result
|
||||
|
||||
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.
|
||||
Each audit call (prompt and response) is wrapped in its own try/except
|
||||
so a failure on the first does not prevent the second from running.
|
||||
"""
|
||||
user_id = self._resolve_user_id_from_logging_kwargs(kwargs)
|
||||
if not user_id:
|
||||
verbose_proxy_logger.debug("Purview audit: no user_id, skipping")
|
||||
return kwargs, result
|
||||
|
||||
# Log prompt (uploadText)
|
||||
try:
|
||||
prompt_text: Optional[str] = None
|
||||
if call_type in ("responses", "aresponses"):
|
||||
# Responses API: route to the responses-specific extractor
|
||||
# before the generic ``messages`` branch. litellm's logging
|
||||
# pipeline stores the raw responses ``input`` (a string or a
|
||||
# list of input items) under ``model_call_details["messages"]``
|
||||
# via ``function_setup``, which is NOT the chat message format
|
||||
# ``get_prompt_text_for_dlp`` expects. Use the original
|
||||
# ``input`` / ``instructions`` keys that ``pre_call`` and
|
||||
# ``update_environment_variables`` persist on the call details.
|
||||
prompt_text = self._responses_api_input_to_str(kwargs)
|
||||
elif call_type in ("text_completion", "atext_completion"):
|
||||
prompt_text = self.completion_prompt_to_str(kwargs.get("prompt"))
|
||||
else:
|
||||
messages = kwargs.get("messages")
|
||||
if messages:
|
||||
prompt_text = self.get_prompt_text_for_dlp(
|
||||
cast(List[Any], messages)
|
||||
)
|
||||
|
||||
if prompt_text:
|
||||
await self._check_content(
|
||||
user_id=user_id,
|
||||
text=prompt_text,
|
||||
activity="uploadText",
|
||||
request_data=kwargs,
|
||||
block_on_violation=False,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Purview audit logging error (prompt): %s", e)
|
||||
|
||||
# Log response (downloadText) — runs regardless of prompt audit outcome
|
||||
try:
|
||||
parts = self._completion_response_text_parts(result)
|
||||
if parts:
|
||||
combined = "\n\n---\n\n".join(parts)
|
||||
await self._check_content(
|
||||
user_id=user_id,
|
||||
text=combined,
|
||||
activity="downloadText",
|
||||
request_data=kwargs,
|
||||
block_on_violation=False,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Purview audit logging error (response): %s", e)
|
||||
|
||||
return kwargs, result
|
||||
|
|
@ -1,5 +1,9 @@
|
|||
from typing import TYPE_CHECKING, List, Optional
|
||||
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_last_user_message,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
|
@ -21,32 +25,4 @@ class OpenAIGuardrailBase:
|
|||
]
|
||||
get_user_prompt(messages) -> "What is the weather in Tokyo?"
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
)
|
||||
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
# Iterate from the end to find the last consecutive block of user messages
|
||||
user_messages = []
|
||||
for message in reversed(messages):
|
||||
if message.get("role") == "user":
|
||||
user_messages.append(message)
|
||||
else:
|
||||
# Stop when we hit a non-user message
|
||||
break
|
||||
|
||||
if not user_messages:
|
||||
return None
|
||||
|
||||
# Reverse to get the messages in chronological order
|
||||
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
|
||||
return get_last_user_message(messages)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -993,6 +993,11 @@ def test_vertex_ai_stream(provider):
|
|||
|
||||
except litellm.RateLimitError as e:
|
||||
pass
|
||||
except litellm.exceptions.MidStreamFallbackError as e:
|
||||
# Streaming 429s are wrapped in MidStreamFallbackError so the
|
||||
# Router can fall back; treat as a transient rate-limit pass.
|
||||
if not isinstance(e.original_exception, litellm.RateLimitError):
|
||||
pytest.fail(f"Error occurred: {e}")
|
||||
except Exception as e:
|
||||
pytest.fail(f"Error occurred: {e}")
|
||||
|
||||
|
|
|
|||
|
|
@ -3036,6 +3036,206 @@ class TestMergeGatewayInitializeInstructions:
|
|||
)
|
||||
|
||||
|
||||
class TestEnsureUpstreamInitializeInstructionsCached:
|
||||
@pytest.mark.asyncio
|
||||
async def test_skips_when_yaml_instructions_set(self):
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
server = _make_instruction_server(
|
||||
server_id="yaml-only", instructions="from yaml"
|
||||
)
|
||||
with patch.object(
|
||||
global_mcp_server_manager, "_create_mcp_client", AsyncMock()
|
||||
) as mock_create:
|
||||
await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(
|
||||
server
|
||||
)
|
||||
mock_create.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skips_when_already_cached(self):
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
server = _make_instruction_server(server_id="cached-only", instructions=None)
|
||||
global_mcp_server_manager._upstream_initialize_instructions_by_server_id[
|
||||
"cached-only"
|
||||
] = "warm"
|
||||
try:
|
||||
with patch.object(
|
||||
global_mcp_server_manager, "_create_mcp_client", AsyncMock()
|
||||
) as mock_create:
|
||||
await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(
|
||||
server
|
||||
)
|
||||
mock_create.assert_not_awaited()
|
||||
finally:
|
||||
global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop(
|
||||
"cached-only", None
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skips_when_spec_path_set(self):
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
server = _make_instruction_server(
|
||||
server_id="openapi-spec", spec_path="/openapi.json", url=None
|
||||
)
|
||||
with patch.object(
|
||||
global_mcp_server_manager, "_create_mcp_client", AsyncMock()
|
||||
) as mock_create:
|
||||
await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(
|
||||
server
|
||||
)
|
||||
mock_create.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runs_upstream_session_and_caches(self):
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
server = _make_instruction_server(server_id="cold-server", instructions=None)
|
||||
fake_client = MagicMock()
|
||||
fake_client.run_with_session = AsyncMock(return_value="ok")
|
||||
fake_client._last_initialize_instructions = " upstream says hi "
|
||||
|
||||
with patch.object(
|
||||
global_mcp_server_manager,
|
||||
"_create_mcp_client",
|
||||
AsyncMock(return_value=fake_client),
|
||||
):
|
||||
try:
|
||||
await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(
|
||||
server
|
||||
)
|
||||
assert (
|
||||
global_mcp_server_manager._upstream_initialize_instructions_by_server_id[
|
||||
"cold-server"
|
||||
]
|
||||
== "upstream says hi"
|
||||
)
|
||||
finally:
|
||||
global_mcp_server_manager._upstream_initialize_instructions_by_server_id.pop(
|
||||
"cold-server", None
|
||||
)
|
||||
global_mcp_server_manager._upstream_initialize_instructions_probed_at.pop(
|
||||
"cold-server", None
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cooldown_after_empty_upstream_response(self):
|
||||
"""Upstream returns no instructions → next call within cooldown must not reconnect."""
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
server = _make_instruction_server(server_id="empty-server", instructions=None)
|
||||
fake_client = MagicMock()
|
||||
fake_client.run_with_session = AsyncMock(return_value="ok")
|
||||
fake_client._last_initialize_instructions = None # upstream sent nothing
|
||||
|
||||
create = AsyncMock(return_value=fake_client)
|
||||
with patch.object(global_mcp_server_manager, "_create_mcp_client", create):
|
||||
try:
|
||||
await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(
|
||||
server
|
||||
)
|
||||
await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(
|
||||
server
|
||||
)
|
||||
assert create.await_count == 1, (
|
||||
"Second probe within cooldown must not reconnect to upstream"
|
||||
)
|
||||
assert (
|
||||
"empty-server"
|
||||
not in global_mcp_server_manager._upstream_initialize_instructions_by_server_id
|
||||
)
|
||||
assert (
|
||||
"empty-server"
|
||||
in global_mcp_server_manager._upstream_initialize_instructions_probed_at
|
||||
)
|
||||
finally:
|
||||
global_mcp_server_manager._upstream_initialize_instructions_probed_at.pop(
|
||||
"empty-server", None
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cooldown_after_upstream_failure(self):
|
||||
"""run_with_session raises → cooldown applies, no immediate retry."""
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
server = _make_instruction_server(server_id="boom-server", instructions=None)
|
||||
fake_client = MagicMock()
|
||||
fake_client.run_with_session = AsyncMock(side_effect=RuntimeError("upstream down"))
|
||||
fake_client._last_initialize_instructions = None
|
||||
|
||||
create = AsyncMock(return_value=fake_client)
|
||||
with patch.object(global_mcp_server_manager, "_create_mcp_client", create):
|
||||
try:
|
||||
await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(
|
||||
server
|
||||
)
|
||||
await global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(
|
||||
server
|
||||
)
|
||||
assert create.await_count == 1, (
|
||||
"Second probe within cooldown must not reconnect after failure"
|
||||
)
|
||||
assert (
|
||||
"boom-server"
|
||||
not in global_mcp_server_manager._upstream_initialize_instructions_by_server_id
|
||||
)
|
||||
assert (
|
||||
"boom-server"
|
||||
in global_mcp_server_manager._upstream_initialize_instructions_probed_at
|
||||
)
|
||||
finally:
|
||||
global_mcp_server_manager._upstream_initialize_instructions_probed_at.pop(
|
||||
"boom-server", None
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reload_resets_probe_cooldown(self):
|
||||
"""load_servers_from_config clears the negative-cache map so reloads re-probe."""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
global_mcp_server_manager._upstream_initialize_instructions_probed_at[
|
||||
"reload-target"
|
||||
] = 1.0
|
||||
try:
|
||||
await global_mcp_server_manager.load_servers_from_config({})
|
||||
assert (
|
||||
"reload-target"
|
||||
not in global_mcp_server_manager._upstream_initialize_instructions_probed_at
|
||||
)
|
||||
finally:
|
||||
global_mcp_server_manager._upstream_initialize_instructions_probed_at.pop(
|
||||
"reload-target", None
|
||||
)
|
||||
|
||||
|
||||
class TestGatewayCreateInitializationOptions:
|
||||
"""Tests for the patched server.create_initialization_options via ContextVar."""
|
||||
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
Loading…
Add table
Reference in a new issue