Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/lucid-burnell-e0be3a

This commit is contained in:
Yuneng Jiang 2026-05-22 18:10:10 -07:00
commit a7feddba43
No known key found for this signature in database
13 changed files with 4280 additions and 68 deletions

View file

@ -2541,7 +2541,6 @@ jobs:
paths:
- litellm-docker-database.tar.zst
test_bad_database_url:
machine:
image: ubuntu-2204:2024.04.1

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -5,6 +5,9 @@ from typing import Any, Dict, List, Literal, Optional, Union
from pydantic import BaseModel, ConfigDict, Field, field_validator
from typing_extensions import Required, TypedDict
from litellm.types.proxy.guardrails.guardrail_hooks.akto import (
AktoConfigModel,
)
from litellm.types.proxy.guardrails.guardrail_hooks.block_code_execution import (
BlockCodeExecutionGuardrailConfigModel,
)
@ -17,9 +20,6 @@ from litellm.types.proxy.guardrails.guardrail_hooks.grayswan import (
from litellm.types.proxy.guardrails.guardrail_hooks.ibm import (
IBMGuardrailsBaseConfigModel,
)
from litellm.types.proxy.guardrails.guardrail_hooks.akto import (
AktoConfigModel,
)
from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import (
ContentFilterCategoryConfig,
)
@ -93,6 +93,7 @@ class SupportedGuardrailIntegrations(Enum):
GENERIC_GUARDRAIL_API = "generic_guardrail_api"
QUALIFIRE = "qualifire"
CUSTOM_CODE = "custom_code"
MICROSOFT_PURVIEW = "microsoft_purview"
SEMANTIC_GUARD = "semantic_guard"
MCP_END_USER_PERMISSION = "mcp_end_user_permission"
BLOCK_CODE_EXECUTION = "block_code_execution"

View file

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

View file

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