Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_add_milvus_grpc_transport

# Conflicts:
#	basedpyright-code-budget.json
#	type-discipline-budget.json
This commit is contained in:
mateo-berri 2026-09-05 14:33:05 -07:00
commit f950274f93
46 changed files with 5735 additions and 358 deletions

View file

@ -24,8 +24,7 @@ jobs:
runs-on: ubuntu-latest
timeout-minutes: 5
permissions:
issues: write # PR comments use the issues API
pull-requests: read # Current-head validation rejects stale workflow runs
pull-requests: write
steps:
- name: Link release wheel report on PR

View file

@ -99,19 +99,19 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 44133
"limit": 44134
},
"reportUnknownLambdaType": {
"limit": 109
},
"reportUnknownMemberType": {
"limit": 38267
"limit": 38255
},
"reportUnknownParameterType": {
"limit": 19582
},
"reportUnknownVariableType": {
"limit": 29826
"limit": 29814
},
"reportUnnecessaryCast": {
"limit": 110

View file

@ -13,6 +13,7 @@ warnings.filterwarnings("ignore", message=".*`ReadOnly` qualifier.*")
### INIT VARIABLES #########################
import threading
import os
import sys
# Load .env before any other litellm imports so env vars (e.g. LITELLM_UI_SESSION_DURATION) are available
import dotenv as _dotenv
@ -45,8 +46,6 @@ from typing import (
TYPE_CHECKING,
Union,
)
from litellm.types.integrations.datadog import DatadogInitParams
from litellm.types.integrations.newrelic import NewRelicInitParams
from litellm._logging import (
set_verbose,
_turn_on_debug,
@ -95,8 +94,7 @@ from litellm.constants import (
DEFAULT_SOFT_BUDGET,
DEFAULT_ALLOWED_FAILS,
)
import httpx
# httpx is lazy-loaded via __getattr__
# register_async_client_cleanup is lazy-loaded and called on first access
litellm_mode = os.getenv("LITELLM_MODE", "DEV") # "PRODUCTION", "DEV"
@ -364,8 +362,6 @@ guardrail_name_config_map: Dict[str, GuardrailItem] = {}
include_cost_in_streaming_usage: bool = False
reasoning_auto_summary: bool = False
### PROMPTS ####
from litellm.types.prompts.init_prompts import PromptSpec
prompt_name_config_map: Dict[str, PromptSpec] = {}
##################
@ -1271,206 +1267,203 @@ openai_video_generation_models = ["sora-2"]
# get_llm_provider is lazy-loaded via __getattr__
# remove_index_from_tool_calls is lazy-loaded via __getattr__
# Import KeyManagementSettings here (before utils import) because _key_management_settings
# is accessed during import time in secret_managers/main.py (via dd_tracing -> datadog -> _service_logger -> utils)
from litellm.types.secret_managers.main import KeyManagementSettings
# SDK symbols previously imported eagerly here are lazy-loaded via __getattr__
# (_SDK_SYMBOLS_IMPORT_MAP in _lazy_imports_registry.py); mirrored under TYPE_CHECKING
# so static type checkers still see them
if TYPE_CHECKING:
_key_management_settings: KeyManagementSettings
_key_management_settings: KeyManagementSettings = KeyManagementSettings()
from .utils import client
# client must be imported immediately as it's used as a decorator at function definition time
from .utils import client
from .llms.custom_llm import CustomLLM
from .llms.anthropic.common_utils import AnthropicModelInfo
from .llms.ai21.chat.transformation import AI21ChatConfig, AI21ChatConfig as AI21Config
from .llms.deprecated_providers.palm import (
PalmConfig,
) # here to prevent breaking changes
from .llms.deprecated_providers.aleph_alpha import AlephAlphaConfig
from .llms.gemini.common_utils import GeminiModelInfo
# Note: Most other utils imports are lazy-loaded via __getattr__ to avoid loading utils.py
# (which imports tiktoken) at import time
from .llms.vertex_ai.vertex_embeddings.transformation import (
VertexAITextEmbeddingConfig,
)
from .llms.custom_llm import CustomLLM
from .llms.anthropic.common_utils import AnthropicModelInfo
from .llms.ai21.chat.transformation import AI21ChatConfig, AI21ChatConfig as AI21Config
from .llms.deprecated_providers.palm import (
PalmConfig,
) # here to prevent breaking changes
from .llms.deprecated_providers.aleph_alpha import AlephAlphaConfig
from .llms.gemini.common_utils import GeminiModelInfo
vertexAITextEmbeddingConfig = VertexAITextEmbeddingConfig()
from .llms.bedrock.embed.amazon_titan_v2_transformation import (
AmazonTitanV2Config,
)
from .llms.topaz.common_utils import TopazModelInfo
from .llms.vertex_ai.vertex_embeddings.transformation import (
VertexAITextEmbeddingConfig,
)
# OpenAIOSeriesConfig is lazy loaded - openaiOSeriesConfig will be created on first access
# OpenAIGPTConfig, OpenAIGPT5Config, etc. are lazy loaded - instances will be created on first access
from .llms.xai.common_utils import XAIModelInfo
vertexAITextEmbeddingConfig = VertexAITextEmbeddingConfig()
# PublicAI now uses JSON-based configuration (see litellm/llms/openai_like/providers.json)
# All remaining configs are now lazy loaded - see _lazy_imports_registry.py
# Import LlmProviders here (before main import) because it's imported during import time
# in multiple places including openai.py (via main import)
from .llms.bedrock.embed.amazon_titan_v2_transformation import (
AmazonTitanV2Config,
)
from .llms.topaz.common_utils import TopazModelInfo
## Lazy loading this is not straightforward, will leave it here for now.
from .main import *
from .compression import compress
# OpenAIOSeriesConfig is lazy loaded - openaiOSeriesConfig will be created on first access
# OpenAIGPTConfig, OpenAIGPT5Config, etc. are lazy loaded - instances will be created on first access
from .llms.xai.common_utils import XAIModelInfo
# Skills API
from .skills.main import (
create_skill,
acreate_skill,
list_skills,
alist_skills,
get_skill,
aget_skill,
delete_skill,
adelete_skill,
)
from .evals.main import (
create_eval,
acreate_eval,
list_evals,
alist_evals,
get_eval,
aget_eval,
delete_eval,
adelete_eval,
cancel_eval,
acancel_eval,
create_run,
acreate_run,
list_runs,
alist_runs,
get_run,
aget_run,
delete_run,
adelete_run,
cancel_run,
acancel_run,
)
from .integrations import *
from .llms.custom_httpx.async_client_cleanup import close_litellm_async_clients
from .exceptions import (
AuthenticationError,
InvalidRequestError,
BadRequestError,
ImageFetchError,
NotFoundError,
PermissionDeniedError,
RateLimitError,
RateLimitErrorCategory,
RateLimitType,
ServiceUnavailableError,
BadGatewayError,
OpenAIError,
ContextWindowExceededError,
ContentPolicyViolationError,
BudgetExceededError,
APIError,
Timeout,
APIConnectionError,
UnsupportedParamsError,
APIResponseValidationError,
UnprocessableEntityError,
InternalServerError,
JSONSchemaValidationError,
LITELLM_EXCEPTION_TYPES,
MockException,
)
from .budget_manager import BudgetManager
from .proxy.proxy_cli import run_server
from .router import Router
from .assistants.main import *
from .batches.main import *
from .images.main import *
from .videos.main import *
from .batch_completion.main import *
from .rerank_api.main import *
from .llms.anthropic.experimental_pass_through.messages.handler import *
from .responses.main import *
# PublicAI now uses JSON-based configuration (see litellm/llms/openai_like/providers.json)
# All remaining configs are now lazy loaded - see _lazy_imports_registry.py
# Interactions API is available as litellm.interactions module
# Usage: litellm.interactions.create(), litellm.interactions.get(), etc.
from . import interactions
from .interactions.agents.main import (
acreate as acreate_agent,
create as create_agent,
alist as alist_agents,
list as list_agents,
aget as aget_agent,
get as get_agent,
adelete as adelete_agent,
delete as delete_agent,
alist_versions as alist_agent_versions,
list_versions as list_agent_versions,
)
from .skills.main import (
create_skill,
acreate_skill,
list_skills,
alist_skills,
get_skill,
aget_skill,
delete_skill,
adelete_skill,
)
from .containers.main import *
from .ocr.main import *
from .rust_bridge import rust
from .rag.main import *
from .sandbox.main import *
from .search.main import *
from .realtime_api.main import (
_arealtime,
acreate_realtime_client_secret,
acreate_realtime_transcription_session,
arealtime_calls,
)
from .responses.main import _aresponses_websocket
from .fine_tuning.main import *
from .files.main import *
from .vector_store_files.main import (
acreate as avector_store_file_create,
adelete as avector_store_file_delete,
alist as avector_store_file_list,
aretrieve as avector_store_file_retrieve,
aretrieve_content as avector_store_file_content,
aupdate as avector_store_file_update,
create as vector_store_file_create,
delete as vector_store_file_delete,
list as vector_store_file_list,
retrieve as vector_store_file_retrieve,
retrieve_content as vector_store_file_content,
update as vector_store_file_update,
)
from .scheduler import *
# Import LlmProviders here (before main import) because it's imported during import time
# in multiple places including openai.py (via main import)
from litellm.types.utils import LlmProviders
### ADAPTERS ###
import litellm.anthropic_interface as anthropic
## Lazy loading this is not straightforward, will leave it here for now.
from .main import *
from .compression import compress
### Vector Store Registry ###
# Skills API
from .skills.main import (
create_skill,
acreate_skill,
list_skills,
alist_skills,
get_skill,
aget_skill,
delete_skill,
adelete_skill,
)
from .evals.main import (
create_eval,
acreate_eval,
list_evals,
alist_evals,
get_eval,
aget_eval,
delete_eval,
adelete_eval,
cancel_eval,
acancel_eval,
create_run,
acreate_run,
list_runs,
alist_runs,
get_run,
aget_run,
delete_run,
adelete_run,
cancel_run,
acancel_run,
)
from .integrations import *
from .llms.custom_httpx.async_client_cleanup import close_litellm_async_clients
from .exceptions import (
AuthenticationError,
InvalidRequestError,
BadRequestError,
ImageFetchError,
NotFoundError,
PermissionDeniedError,
RateLimitError,
RateLimitErrorCategory,
RateLimitType,
ServiceUnavailableError,
BadGatewayError,
OpenAIError,
ContextWindowExceededError,
ContentPolicyViolationError,
BudgetExceededError,
APIError,
Timeout,
APIConnectionError,
UnsupportedParamsError,
APIResponseValidationError,
UnprocessableEntityError,
InternalServerError,
JSONSchemaValidationError,
LITELLM_EXCEPTION_TYPES,
MockException,
)
from .budget_manager import BudgetManager
from .proxy.proxy_cli import run_server
from .router import Router
from .assistants.main import *
from .batches.main import *
from .images.main import *
from .videos.main import *
from .batch_completion.main import *
from .rerank_api.main import *
from .llms.anthropic.experimental_pass_through.messages.handler import *
from .responses.main import *
### RAG ###
from . import rag
# Interactions API is available as litellm.interactions module
# Usage: litellm.interactions.create(), litellm.interactions.get(), etc.
from . import interactions
from .interactions.agents.main import (
acreate as acreate_agent,
create as create_agent,
alist as alist_agents,
list as list_agents,
aget as aget_agent,
get as get_agent,
adelete as adelete_agent,
delete as delete_agent,
alist_versions as alist_agent_versions,
list_versions as list_agent_versions,
)
from .skills.main import (
create_skill,
acreate_skill,
list_skills,
alist_skills,
get_skill,
aget_skill,
delete_skill,
adelete_skill,
)
from .containers.main import *
from .ocr.main import *
from .rust_bridge import rust
from .rag.main import *
from .sandbox.main import *
from .search.main import *
from .realtime_api.main import (
_arealtime,
acreate_realtime_client_secret,
acreate_realtime_transcription_session,
arealtime_calls,
)
from .responses.main import _aresponses_websocket
from .fine_tuning.main import *
from .files.main import *
from .vector_store_files.main import (
acreate as avector_store_file_create,
adelete as avector_store_file_delete,
alist as avector_store_file_list,
aretrieve as avector_store_file_retrieve,
aretrieve_content as avector_store_file_content,
aupdate as avector_store_file_update,
create as vector_store_file_create,
delete as vector_store_file_delete,
list as vector_store_file_list,
retrieve as vector_store_file_retrieve,
retrieve_content as vector_store_file_content,
update as vector_store_file_update,
)
from .scheduler import *
### CUSTOM LLMs ###
### CLI UTILITIES ###
from litellm.litellm_core_utils.cli_token_utils import get_litellm_gateway_api_key
### PASSTHROUGH ###
from .passthrough import allm_passthrough_route, llm_passthrough_route
from .google_genai import agenerate_content
### ADAPTERS ###
from .types.adapter import AdapterItem
import litellm.anthropic_interface as anthropic
adapters: List[AdapterItem] = []
### Vector Store Registry ###
from .vector_stores.vector_store_registry import (
VectorStoreRegistry,
VectorStoreIndexRegistry,
)
vector_store_registry: Optional[VectorStoreRegistry] = None
vector_store_index_registry: Optional[VectorStoreIndexRegistry] = None
### RAG ###
from . import rag
### CUSTOM LLMs ###
from .types.llms.custom_llm import CustomLLMItem
custom_provider_map: List[CustomLLMItem] = []
_custom_providers: List[str] = [] # internal helper util, used to track names of custom providers
disable_hf_tokenizer_download: Optional[bool] = (
@ -1478,13 +1471,6 @@ disable_hf_tokenizer_download: Optional[bool] = (
)
global_disable_no_log_param: bool = False
### CLI UTILITIES ###
from litellm.litellm_core_utils.cli_token_utils import get_litellm_gateway_api_key
### PASSTHROUGH ###
from .passthrough import allm_passthrough_route, llm_passthrough_route
from .google_genai import agenerate_content
### GLOBAL CONFIG ###
global_bitbucket_config: Optional[Dict[str, Any]] = None
@ -1508,10 +1494,21 @@ def set_global_gitlab_config(config: Dict[str, Any]) -> None:
# Lazy loading system for heavy modules to reduce initial import time and memory usage
if TYPE_CHECKING:
import httpx
from litellm.types.utils import ModelInfo as _ModelInfoType
from litellm.types.utils import PriorityReservationSettings
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.caching.caching import Cache
from litellm.types.adapter import AdapterItem
from litellm.types.integrations.datadog import DatadogInitParams
from litellm.types.integrations.newrelic import NewRelicInitParams
from litellm.types.llms.custom_llm import CustomLLMItem
from litellm.types.prompts.init_prompts import PromptSpec
from litellm.vector_stores.vector_store_registry import (
VectorStoreIndexRegistry,
VectorStoreRegistry,
)
# Type stubs for lazy-loaded configs to help mypy
from .llms.bedrock.chat.converse_transformation import (
@ -2187,16 +2184,6 @@ if TYPE_CHECKING:
# Track if async client cleanup has been registered (for lazy loading)
_async_client_cleanup_registered = False
# Eager loading for backwards compatibility with VCR and other HTTP recording tools
# When LITELLM_DISABLE_LAZY_LOADING is set, lazy-loaded attributes are loaded at import time
# For now, this only affects encoding (tiktoken) as it was the only reported issue
# See: https://github.com/BerriAI/litellm/issues/18659
# This ensures encoding is initialized before VCR starts recording HTTP requests
if os.getenv("LITELLM_DISABLE_LAZY_LOADING", "").lower() in ("1", "true", "yes", "on"):
# Load encoding at import time (pre-#18070 behavior)
# This ensures encoding is initialized before VCR starts recording
from .main import encoding
def __getattr__(name: str) -> Any:
"""Lazy import handler with cached registry for improved performance."""
@ -2276,6 +2263,8 @@ def __getattr__(name: str) -> Any:
"openAIGPT5Config": "OpenAIGPT5Config",
"nvidiaNimConfig": "NvidiaNimConfig",
"nvidiaNimEmbeddingConfig": "NvidiaNimEmbeddingConfig",
"vertexAITextEmbeddingConfig": "VertexAITextEmbeddingConfig",
"_key_management_settings": "KeyManagementSettings",
}
if name in _config_instances:
from ._lazy_imports import get_litellm_globals
@ -2393,7 +2382,30 @@ def __getattr__(name: str) -> Any:
return locals()[name]
from ._lazy_imports import lazy_import_litellm_submodule
submodule: Final = lazy_import_litellm_submodule(name)
if submodule is not None:
return submodule
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
from ._lazy_imports import LiteLLMModule
from ._lazy_imports_registry import STAR_IMPORT_PUBLIC_NAMES
sys.modules[__name__].__class__ = LiteLLMModule
__all__ = list(STAR_IMPORT_PUBLIC_NAMES) # mutable-ok: star imports require __all__ to be a list of str
# ALL_LITELLM_RESPONSE_TYPES is lazy-loaded via __getattr__ to avoid loading utils at import time
# Eager loading for backwards compatibility with VCR and other HTTP recording tools
# When LITELLM_DISABLE_LAZY_LOADING is set, lazy-loaded attributes are loaded at import time
# For now, this only affects encoding (tiktoken) as it was the only reported issue
# See: https://github.com/BerriAI/litellm/issues/18659
# This ensures encoding is initialized before VCR starts recording HTTP requests
# This block stays at the bottom so __getattr__ can resolve attributes main.py needs during its import
if os.getenv("LITELLM_DISABLE_LAZY_LOADING", "").lower() in ("1", "true", "yes", "on"):
from .main import encoding

View file

@ -16,9 +16,10 @@ until they're actually needed.
"""
import importlib
import importlib.util
import sys
from collections.abc import Callable, Mapping
from types import ModuleType
from types import MappingProxyType, ModuleType
from typing import TYPE_CHECKING, Any, Final, cast
from typing_extensions import ReadOnly, TypedDict
@ -34,6 +35,8 @@ from ._lazy_imports_registry import (
_LITELLM_LOGGING_IMPORT_MAP,
_LLM_CONFIGS_IMPORT_MAP,
_LLM_PROVIDER_LOGIC_IMPORT_MAP,
_SDK_MODULE_ALIASES,
_SDK_SYMBOLS_IMPORT_MAP,
_TOKEN_COUNTER_IMPORT_MAP,
_TYPES_IMPORT_MAP,
_TYPES_UTILS_IMPORT_MAP,
@ -78,7 +81,10 @@ def _get_utils_globals() -> dict[str, object]:
This is where we cache imported attributes so we don't import them twice.
When you do `litellm.utils.some_function`, it gets stored in this dictionary.
"""
return sys.modules["litellm.utils"].__dict__
cached: Final = sys.modules.get("litellm.utils")
if cached is not None:
return cached.__dict__
return importlib.import_module("litellm.utils").__dict__
def _get_module_level_client_timeout(litellm_globals: Mapping[str, Any]) -> "float | httpx.Timeout | None":
@ -214,6 +220,10 @@ def _get_lazy_import_registry() -> dict[str, Callable[[str], object]]:
_LAZY_IMPORT_REGISTRY[name] = _lazy_import_llm_provider_logic
for name in UTILS_MODULE_NAMES:
_LAZY_IMPORT_REGISTRY[name] = _lazy_import_utils_module
for name in _SDK_SYMBOLS_IMPORT_MAP:
_LAZY_IMPORT_REGISTRY.setdefault(name, _lazy_import_sdk_symbols)
for name in _SDK_MODULE_ALIASES:
_LAZY_IMPORT_REGISTRY.setdefault(name, _lazy_import_sdk_module_alias)
return _LAZY_IMPORT_REGISTRY
@ -229,7 +239,7 @@ def _module_attribute(module: ModuleType, attr_name: str) -> object:
return attribute["value"]
def _generic_lazy_import(name: str, import_map: dict[str, tuple[str, str]], category: str) -> object:
def _generic_lazy_import(name: str, import_map: Mapping[str, tuple[str, str]], category: str) -> object:
"""
Generic function that handles lazy importing for most attributes.
@ -350,6 +360,86 @@ def _lazy_import_llm_provider_logic(name: str) -> object:
return _generic_lazy_import(name, _LLM_PROVIDER_LOGIC_IMPORT_MAP, "LLM provider logic")
def _lazy_import_sdk_symbols(name: str) -> object:
"""Handler for SDK symbols previously imported eagerly at the bottom of litellm/__init__.py"""
return _generic_lazy_import(name, _SDK_SYMBOLS_IMPORT_MAP, "SDK symbols")
def _lazy_import_sdk_module_alias(name: str) -> object:
"""Handler for litellm attributes that bind a module (e.g. litellm.anthropic)"""
_globals: Final = get_litellm_globals()
if name in _globals:
return _globals[name]
module: Final = importlib.import_module(_SDK_MODULE_ALIASES[name])
_globals[name] = module # rebind-ok: caches the resolved module alias on the package
return module
_SHADOWABLE_SDK_FUNCTIONS: Final = MappingProxyType(
{
"batch_completion": ("litellm.batch_completion.main", "batch_completion"),
"ocr": ("litellm.ocr.main", "ocr"),
"responses": ("litellm.responses.main", "responses"),
"search": ("litellm.search.main", "search"),
}
)
def _shadowable_function_property(name: str) -> property:
"""Property keeping litellm.<name> bound to the SDK function even after the import
machinery binds the identically named litellm.<name> subpackage onto the litellm module."""
module_path, attr_name = _SHADOWABLE_SDK_FUNCTIONS[name]
def _get(module: ModuleType) -> object:
stored: Final = module.__dict__.get(name)
if stored is not None and not (isinstance(stored, ModuleType) and stored.__name__ == f"litellm.{name}"):
return stored
value: Final = _module_attribute(importlib.import_module(module_path), attr_name)
module.__dict__[name] = value # rebind-ok: caches the resolved function on the litellm module
return value
def _set(module: ModuleType, value: object) -> None:
module.__dict__[name] = value # rebind-ok: property setter must store assignments on the module
return property(_get, _set)
class LiteLLMModule(ModuleType):
"""Module type installed on the litellm package so function names shadowed by
same-named subpackages (litellm.responses, ...) keep resolving to the functions."""
batch_completion = _shadowable_function_property("batch_completion")
ocr = _shadowable_function_property("ocr")
responses = _shadowable_function_property("responses")
search = _shadowable_function_property("search")
def lazy_import_submodule(package: str, name: str) -> "ModuleType | None":
"""Resolve <package>.<name> as a submodule (e.g. litellm.utils) when no other handler matches"""
if name.startswith("__") or not name.isidentifier():
return None
qualified_name: Final = f"{package}.{name}"
try:
spec: Final = importlib.util.find_spec(qualified_name)
except ModuleNotFoundError:
return None
if spec is None:
return None
try:
module: Final = importlib.import_module(qualified_name)
except ModuleNotFoundError as exc:
if exc.name == qualified_name:
return None
raise
sys.modules[package].__dict__[name] = module # rebind-ok: caches the resolved submodule on the package
return module
def lazy_import_litellm_submodule(name: str) -> "ModuleType | None":
"""Resolve litellm.<name> as a submodule (e.g. litellm.utils) when no other handler matches"""
return lazy_import_submodule("litellm", name)
def _lazy_import_utils_module(name: str) -> object:
"""
Handler for utils module lazy imports.

File diff suppressed because it is too large Load diff

View file

@ -5881,6 +5881,7 @@ def _get_status_fields(
# Mapping for legacy guardrail status values to new GuardrailStatus values
GUARDRAIL_STATUS_MAP: Final[dict[str, GuardrailStatus]] = {
"success": "success",
"guardrail_flagged": "guardrail_flagged",
"blocked": "guardrail_intervened", # legacy
"guardrail_intervened": "guardrail_intervened", # direct
"failure": "guardrail_failed_to_respond", # legacy
@ -5902,6 +5903,7 @@ def _get_status_fields(
GUARDRAIL_STATUS_SEVERITY: Final[tuple[GuardrailStatus, ...]] = (
"not_run",
"success",
"guardrail_flagged",
"guardrail_failed_to_respond",
"guardrail_intervened",
)

View file

@ -1 +1,11 @@
from . import *
from types import ModuleType
from typing import Final
def __getattr__(name: str) -> ModuleType:
from litellm._lazy_imports import lazy_import_submodule
submodule: Final = lazy_import_submodule(__name__, name)
if submodule is not None:
return submodule
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")

View file

@ -1,16 +1,19 @@
"""
MCP Guardrail Handler for Unified Guardrails.
Converts an MCP call_tool (name + arguments) into a single OpenAI-compatible
tool_call and passes it to apply_guardrail. Works with the synthetic payload
from ProxyLogging._convert_mcp_to_llm_format.
Converts an MCP call_tool (name + arguments) into the OpenAI-compatible shape
apply_guardrail expects: the tool as a single-entry ``tools`` definition, and
every string leaf of the call arguments as ``texts`` so text guardrails can
detect and mask sensitive values in the payload. Works with the synthetic
request from ProxyLogging._convert_mcp_to_llm_format.
Note: For MCP tool definitions (schema) -> OpenAI tools=[], see
litellm.experimental_mcp_client.tools.transform_mcp_tool_to_openai_tool
when you have a full MCP Tool from list_tools. Here we only have the call
payload (name + arguments) so we just build the tool_call.
payload (name + arguments) so we just build the tool definition.
"""
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final
from fastapi import HTTPException
@ -20,6 +23,8 @@ from litellm._logging import verbose_proxy_logger
from litellm.experimental_mcp_client.tools import transform_mcp_tool_to_openai_tool
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.proxy._experimental.mcp_server.utils import (
MAX_STRUCTURED_CONTENT_SCAN_DEPTH,
JSONLeafPath,
json_string_leaves,
json_unrewritable_labels,
mcp_content_item_text,
@ -42,6 +47,72 @@ if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
def _blocked(reason: str) -> HTTPException:
return HTTPException(status_code=400, detail={"error": f"Content blocked: {reason}"})
def _too_deeply_nested() -> HTTPException:
return _blocked(
f"MCP tool call arguments exceed the maximum nesting depth of {MAX_STRUCTURED_CONTENT_SCAN_DEPTH} "
"and cannot be scanned by the configured guardrail"
)
def _argument_replacements(
argument_leaves: tuple[tuple[JSONLeafPath, str], ...],
masked_texts: Sequence[str] | None,
) -> Mapping[JSONLeafPath, str]:
"""Positionally pair the guardrail's returned texts with the leaves they came from.
Only leaves the guardrail actually rewrote are returned, so a guardrail that
detects nothing leaves the outbound tool call byte-identical. A guardrail that
returns the wrong number of texts fails closed, because a positional write-back
would scramble the arguments rather than mask them.
"""
if masked_texts is not None and len(masked_texts) != len(argument_leaves):
raise _blocked(
f"guardrail returned {len(masked_texts)} texts for {len(argument_leaves)} MCP tool call argument strings, "
"so the redaction cannot be mapped back to the arguments"
)
return {path: masked for (path, original), masked in zip(argument_leaves, masked_texts or ()) if masked != original}
def _conflicting_rewrite_paths(
scanned_leaves: tuple[tuple[JSONLeafPath, str], ...],
current_leaves: tuple[tuple[JSONLeafPath, str], ...],
replacements: Mapping[JSONLeafPath, str],
) -> tuple[JSONLeafPath, ...]:
"""Paths another guardrail already rewrote differently from what this one wants.
Guardrails opted into ``run_in_parallel`` all scan the same payload snapshot, so
each one returns a full replacement string derived from the *original* leaf. Two
of them rewriting one leaf to different values cannot be merged: writing either
result discards the other guardrail's redaction. A leaf still holding the text
this guardrail was handed, or already holding this guardrail's own replacement,
is safe to write; the latter is how a guardrail that masks the arguments itself
as well as through ``texts`` gets there first. Anything else fails closed,
including a payload reshaped so the leaves no longer line up, because the
write-back is positional and would land a redaction on the wrong value.
"""
if tuple(path for path, _ in scanned_leaves) != tuple(path for path, _ in current_leaves):
return tuple(replacements)
return tuple(
path
for (path, scanned), (_, current) in zip(scanned_leaves, current_leaves)
if path in replacements and current not in (scanned, replacements[path])
)
def _conflicting_rewrite(paths: tuple[JSONLeafPath, ...]) -> HTTPException:
return _blocked(
"two guardrails running concurrently rewrote the same MCP tool call "
f"argument{'s' if len(paths) > 1 else ''} "
f"({', '.join('.'.join(str(part) for part in path) for path in paths)}); "
"their redactions cannot be merged. Remove run_in_parallel from one of them so they "
"run in sequence."
)
class MCPGuardrailTranslationHandler(BaseTranslation):
"""Guardrail translation handler for MCP tool calls (passes a single tool_call to guardrail)."""
@ -52,10 +123,8 @@ class MCPGuardrailTranslationHandler(BaseTranslation):
litellm_logging_obj: "LiteLLMLoggingObj | None" = None,
) -> dict[str, Any]:
mcp_tool_name: Final = data.get("mcp_tool_name") or data.get("name")
mcp_arguments = data.get("mcp_arguments") or data.get("arguments")
mcp_arguments: Final[object] = data.get("mcp_arguments") or data.get("arguments")
mcp_tool_description: Final = data.get("mcp_tool_description") or data.get("description")
if mcp_arguments is None or not isinstance(mcp_arguments, dict):
mcp_arguments = {}
if not mcp_tool_name:
verbose_proxy_logger.debug("MCP Guardrail: mcp_tool_name missing")
@ -84,16 +153,37 @@ class MCPGuardrailTranslationHandler(BaseTranslation):
strict=fn.get("strict", False) or False, # Default to False if None
),
}
argument_leaves: Final = json_string_leaves(mcp_arguments)
if argument_leaves is None:
raise _too_deeply_nested()
inputs: Final[GenericGuardrailAPIInputs] = GenericGuardrailAPIInputs(
tools=[tool_def],
texts=[text for _, text in argument_leaves],
)
await guardrail_to_apply.apply_guardrail(
guarded: Final = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,
)
replacements: Final = _argument_replacements(
argument_leaves=argument_leaves,
masked_texts=guarded.get("texts") if guarded else None,
)
if not replacements:
return data
current_arguments: Final[object] = data.get("mcp_arguments") or data.get("arguments")
current_leaves: Final = json_string_leaves(current_arguments)
if current_leaves is None:
raise _too_deeply_nested()
conflicting: Final = _conflicting_rewrite_paths(argument_leaves, current_leaves, replacements)
if conflicting:
raise _conflicting_rewrite(conflicting)
masked_arguments: Final = with_json_string_leaves(current_arguments, replacements)
data["mcp_arguments"] = masked_arguments # rebind-ok: preserve the mask for the outbound MCP call
data["modified_arguments"] = masked_arguments # rebind-ok: expose the applied mask to the caller
return data
async def process_output_response(
@ -131,14 +221,8 @@ class MCPGuardrailTranslationHandler(BaseTranslation):
structured_leaves: Final = json_string_leaves(structured) if structured is not None else ()
structured_labels: Final = json_unrewritable_labels(structured) if structured is not None else ()
if structured_leaves is None or structured_labels is None:
raise HTTPException(
status_code=400,
detail={
"error": (
"Content blocked: MCP tool result structuredContent is nested too deeply to be scanned "
"by the configured guardrail"
)
},
raise _blocked(
"MCP tool result structuredContent is nested too deeply to be scanned by the configured guardrail"
)
if not text_blocks and not structured_leaves and not structured_labels:
@ -158,12 +242,10 @@ class MCPGuardrailTranslationHandler(BaseTranslation):
if masked_texts is None:
return response
if len(masked_texts) != len(originals):
verbose_proxy_logger.warning(
"MCP Guardrail: guardrail returned %d texts for %d tool result texts; leaving the result unmasked",
len(masked_texts),
len(originals),
raise _blocked(
f"guardrail returned {len(masked_texts)} texts for {len(originals)} MCP tool result texts, "
"so the redaction cannot be mapped back to the result"
)
return response
split: Final = len(text_blocks)
if content is not None:
@ -173,15 +255,10 @@ class MCPGuardrailTranslationHandler(BaseTranslation):
label_start: Final = split + len(structured_leaves)
if any(masked != original for original, masked in zip(structured_labels, masked_texts[label_start:])):
raise HTTPException(
status_code=400,
detail={
"error": (
"Content blocked: MCP tool result matched a masking rule on a non-rewritable field "
"(a structuredContent key or numeric value), which cannot be redacted without changing "
"the payload contract"
)
},
raise _blocked(
"MCP tool result matched a masking rule on a non-rewritable field "
"(a structuredContent key or numeric value), which cannot be redacted without changing "
"the payload contract"
)
structured_replacements: Final = {

View file

@ -36,6 +36,7 @@ Example: block when response rejects the user (input_type response only):
import asyncio
import threading
import time
from collections.abc import Callable, Mapping
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast
@ -93,6 +94,7 @@ class CustomCodeGuardrail(CustomGuardrail):
that returns one of:
- allow() - let the request/response through
- block(reason) - reject with a message
- flag(reason) - let it through but log a non-blocking violation
- modify(texts=...) - transform the content
Example:
@ -227,6 +229,7 @@ class CustomCodeGuardrail(CustomGuardrail):
raise CustomCodeExecutionError(f"Custom code guardrail not compiled: {self._compile_error}")
raise CustomCodeExecutionError("Custom code guardrail not compiled")
start_time: Final = time.time()
try:
# Prepare inputs dict for the function
@ -245,6 +248,7 @@ class CustomCodeGuardrail(CustomGuardrail):
inputs=inputs,
request_data=request_data,
input_type=input_type,
start_time=start_time,
)
except HTTPException:
@ -290,6 +294,7 @@ class CustomCodeGuardrail(CustomGuardrail):
inputs: GenericGuardrailAPIInputs,
request_data: dict[str, object],
input_type: Literal["request", "response"],
start_time: float,
) -> GenericGuardrailAPIInputs:
"""
Process the result from the custom code function.
@ -299,6 +304,7 @@ class CustomCodeGuardrail(CustomGuardrail):
inputs: The original inputs
request_data: The request data
input_type: "request" or "response"
start_time: Unix timestamp of when the guardrail started running, used for the flagged log entry
Returns:
GenericGuardrailAPIInputs - possibly modified
@ -348,6 +354,27 @@ class CustomCodeGuardrail(CustomGuardrail):
},
)
elif action == "flag":
flag_reason: Final = result.get("reason", "Flagged by custom code guardrail")
verbose_proxy_logger.info(
"Custom code guardrail '%s': Flagging %s - %s", self.guardrail_name, input_type, flag_reason
)
end_time: Final = time.time()
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response={ # mutable-ok: logging helper requires a dict
"action": "flag",
"reason": flag_reason,
"input_type": input_type,
"metadata": result.get("metadata") or {}, # mutable-ok: logging helper requires a dict
},
request_data=request_data,
guardrail_status="guardrail_flagged",
start_time=start_time,
end_time=end_time,
duration=end_time - start_time,
)
return inputs
elif action == "modify":
verbose_proxy_logger.debug("Custom code guardrail '%s': Modifying %s", self.guardrail_name, input_type)

View file

@ -8,7 +8,7 @@ and provide safe, sandboxed functionality for common guardrail operations.
import json
import re
from collections.abc import Mapping, Sequence
from typing import Final
from typing import Final, Literal
from urllib.parse import urlparse
import httpx
@ -51,6 +51,31 @@ def block(reason: str, detection_info: Mapping[str, object] | None = None) -> di
return result
class FlagResult(TypedDict):
action: ReadOnly[Literal["flag"]]
reason: ReadOnly[str]
metadata: ReadOnly[Mapping[str, object]]
def flag(reason: str, metadata: Mapping[str, object] | None = None) -> FlagResult:
"""
Let the request/response proceed unchanged but record a non-blocking violation.
Args:
reason: Human-readable reason for flagging
metadata: Optional structured metadata stored alongside the reason
Returns:
Dict indicating the request should be flagged but allowed
"""
result: Final[FlagResult] = {
"action": "flag",
"reason": reason,
"metadata": metadata if metadata is not None else {},
}
return result
def modify(
texts: Sequence[str] | None = None,
images: Sequence[object] | None = None,
@ -787,6 +812,7 @@ def get_custom_code_primitives() -> dict[str, object]:
# Result types
"allow": allow,
"block": block,
"flag": flag,
"modify": modify,
# Regex
"regex_match": regex_match,

View file

@ -17,6 +17,7 @@ from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.guardrails.usage_tracking import guardrail_status_to_action
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import (
DailyGuardrailMetricsRepository,
@ -41,6 +42,7 @@ if TYPE_CHECKING:
router: Final = APIRouter()
_EMPTY_UNITS: Final[Mapping[str, int]] = MappingProxyType({})
_ACTION_SEVERITY: Final[Mapping[str, int]] = MappingProxyType({"passed": 0, "flagged": 1, "blocked": 2})
_T = TypeVar("_T")
@ -759,21 +761,17 @@ def _usage_log_entry_from_row(
except Exception:
meta = {}
guardrail_info_list: Final[Sequence[_GuardrailRunInfo]] = (meta or {}).get("guardrail_information") or []
entry_for_guardrail: _GuardrailRunInfo | None = None
for gi in guardrail_info_list:
if (gi.get("guardrail_id") or gi.get("guardrail_name")) == r.guardrail_id:
entry_for_guardrail = gi
break
entry_for_guardrail: Final[_GuardrailRunInfo | None] = max(
(gi for gi in guardrail_info_list if (gi.get("guardrail_id") or gi.get("guardrail_name")) == r.guardrail_id),
key=lambda gi: _ACTION_SEVERITY[guardrail_status_to_action(gi.get("guardrail_status"))],
default=None,
)
action_val = "passed"
score_val = None
latency_val = None
reason_val = None
if entry_for_guardrail:
st: Final = (entry_for_guardrail.get("guardrail_status") or "").lower()
if "intervened" in st or "block" in st:
action_val = "blocked"
elif "fail" in st or "error" in st:
action_val = "flagged"
action_val = guardrail_status_to_action(entry_for_guardrail.get("guardrail_status"))
duration: Final = entry_for_guardrail.get("duration")
if duration is not None:
latency_val = round(float(duration) * 1000, 0)

View file

@ -190,14 +190,14 @@ async def _upsert_rows_with_retry(
return await _upsert_rows_with_retry(retryable, upsert_row, label, sleep, retries_left - 1)
def _guardrail_status_to_action(status: str | None) -> str:
def guardrail_status_to_action(status: str | None) -> str:
"""Map StandardLogging guardrail_status to blocked/passed/flagged."""
if not status:
return "passed"
s: Final = (status or "").lower()
if "intervened" in s or "block" in s:
return "blocked"
if "fail" in s or "error" in s:
if "flagged" in s or "fail" in s or "error" in s:
return "flagged"
return "passed"
@ -367,7 +367,7 @@ async def process_spend_logs_guardrail_usage(
continue
key = _MetricsKey(guardrail_id, date_key)
daily_guardrail[key]["requests_evaluated"] += 1
action = _guardrail_status_to_action(entry.get("guardrail_status"))
action = guardrail_status_to_action(entry.get("guardrail_status"))
if action == "passed":
daily_guardrail[key]["passed_count"] += 1
elif action == "blocked":

View file

@ -46,6 +46,7 @@ from litellm.proxy.management_endpoints.common_utils import (
from litellm.proxy.management_helpers.object_permission_utils import (
handle_update_object_permission_common,
prepare_object_permission_upsert,
reject_ambiguous_mcp_tool_permission_keys,
)
from litellm.proxy.management_helpers.utils import (
get_new_internal_user_defaults,
@ -606,6 +607,11 @@ async def _set_object_permission(
return None
if data.object_permission is not None:
await reject_ambiguous_mcp_tool_permission_keys(
new_mcp_tool_permissions=data.object_permission.mcp_tool_permissions,
existing_mcp_tool_permissions=None,
prisma_client=prisma_client,
)
created_object_permission: Final = await _table(ObjectPermissionRepository(prisma_client)).create(
data=data.object_permission.model_dump(exclude_none=True),
)

View file

@ -5,10 +5,13 @@ organizations, teams, and keys.
import json
from collections.abc import Mapping, Sequence
from collections.abc import Set as AbstractSet
from dataclasses import dataclass
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Optional
from fastapi import HTTPException, status
from pydantic import TypeAdapter
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
@ -103,6 +106,11 @@ async def prepare_object_permission_upsert(
if existing_object_permission is not None
else {}
)
await reject_ambiguous_mcp_tool_permission_keys(
new_mcp_tool_permissions=new_object_permission.get("mcp_tool_permissions"),
existing_mcp_tool_permissions=existing_fields.get("mcp_tool_permissions"),
prisma_client=prisma_client,
)
merged: Final[dict[str, object]] = {
**existing_fields,
**new_object_permission,
@ -194,6 +202,12 @@ async def _set_object_permission(
k: v for k, v in permission_data.items() if v is not None and k != "object_permission_id"
}
await reject_ambiguous_mcp_tool_permission_keys(
new_mcp_tool_permissions=clean_data.get("mcp_tool_permissions"),
existing_mcp_tool_permissions=None,
prisma_client=prisma_client,
)
# Serialize mcp_tool_permissions to JSON string for GraphQL compatibility
if "mcp_tool_permissions" in clean_data:
clean_data["mcp_tool_permissions"] = safe_dumps(clean_data["mcp_tool_permissions"])
@ -226,7 +240,7 @@ def _mcp_server_identifier_matches(server: Any, identifier: str) -> bool:
async def _get_db_mcp_servers_by_identifiers(
identifiers: set[str],
identifiers: AbstractSet[str],
prisma_client: PrismaClient | None,
) -> "Sequence[prisma_models.LiteLLM_MCPServerTable]":
if prisma_client is None or not identifiers:
@ -245,7 +259,7 @@ async def _get_db_mcp_servers_by_identifiers(
async def _resolve_mcp_server_identifiers_to_ids(
identifiers: set[str],
identifiers: AbstractSet[str],
prisma_client: PrismaClient | None,
) -> dict[str, set[str]]:
"""
@ -286,6 +300,59 @@ async def _resolve_mcp_server_identifiers_to_ids(
return resolved
_MCP_TOOL_PERMISSIONS_ADAPTER: Final = TypeAdapter(dict[str, list[str] | None])
def _mcp_tool_permission_entries(raw: object) -> Mapping[str, frozenset[str]]:
parsed: Final[Mapping[str, Sequence[str] | None]] = (
_MCP_TOOL_PERMISSIONS_ADAPTER.validate_json(raw)
if isinstance(raw, str)
else _MCP_TOOL_PERMISSIONS_ADAPTER.validate_python(raw)
if isinstance(raw, Mapping)
else MappingProxyType({})
)
return MappingProxyType({identifier: frozenset(tools or ()) for identifier, tools in parsed.items()})
async def reject_ambiguous_mcp_tool_permission_keys(
new_mcp_tool_permissions: object,
existing_mcp_tool_permissions: object,
prisma_client: PrismaClient | None,
) -> None:
"""
A name or alias shared by several MCP servers cannot key ``mcp_tool_permissions``:
the read path unions the entry into every match, so no edit can narrow one of
those servers without also changing the other. An exact server_id is never
ambiguous, even when another server uses that string as its alias. Entries the
row already stores with the same tool list are left alone, so unrelated edits
to such an entity still succeed.
Raises HTTPException(400) naming the colliding servers.
"""
requested: Final = _mcp_tool_permission_entries(new_mcp_tool_permissions)
stored: Final = _mcp_tool_permission_entries(existing_mcp_tool_permissions)
resolved: Final = await _resolve_mcp_server_identifiers_to_ids(
identifiers=frozenset(identifier for identifier, tools in requested.items() if stored.get(identifier) != tools),
prisma_client=prisma_client,
)
collisions: Final = "; ".join(
f"'{identifier}' matches MCP servers {sorted(server_ids)}"
for identifier, server_ids in sorted(resolved.items())
if identifier not in server_ids and len(server_ids) > 1
)
if not collisions:
return
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={ # mutable-ok: HTTPException.detail has no immutable form; same shape as the sibling errors here
"error": (
f"Ambiguous mcp_tool_permissions key: {collisions}. "
"Key tool permissions by server_id when servers share a name or alias."
)
},
)
def _drop_stale_object_permission_mcp_servers(
object_permission: ObjectPermissionDict,
identifier_to_server_ids: dict[str, set[str]],

View file

@ -0,0 +1,167 @@
"""Serve Prometheus `/metrics` from its own process so a scrape never runs on an inference worker.
Workers write their samples to `PROMETHEUS_MULTIPROC_DIR`; this process reads them back with a
``MultiProcessCollector`` and serves the aggregated output on a separate port. The proxy CLI starts
it with ``--prometheus_metrics_port``. It can also run as a sidecar sharing the same directory:
``python -m litellm.proxy.prometheus_metrics_server --host 0.0.0.0 --port 4001``.
"""
from __future__ import annotations
import argparse
import atexit
import os
import subprocess
import sys
import threading
import time
from collections.abc import Sequence
from contextlib import closing
from types import MappingProxyType
from typing import Final
import httpx
from fastapi import FastAPI
from prometheus_client import CollectorRegistry, multiprocess
from pydantic import BaseModel, ConfigDict
from starlette.types import ASGIApp, Message, Receive, Scope, Send
from litellm.integrations.prometheus_metrics_endpoint import make_metrics_asgi_app
from litellm.llms.custom_httpx.http_handler import HTTPHandler
METRICS_PATH: Final = "/metrics"
PID_HEADER: Final = "x-litellm-metrics-pid"
_PARENT_POLL_INTERVAL_SECONDS: Final = 1.0
_STARTUP_TIMEOUT_SECONDS: Final = 30.0
_STARTUP_POLL_INTERVAL_SECONDS: Final = 0.1
_STARTUP_PROBE_TIMEOUT_SECONDS: Final = 1.0
_WILDCARD_TO_LOOPBACK: Final = MappingProxyType({"0.0.0.0": "127.0.0.1", "::": "::1"})
class _CliArgs(BaseModel):
model_config = ConfigDict(frozen=True)
host: str
port: int
multiproc_dir: str | None
class MetricsServerStartupError(RuntimeError):
"""The metrics process died or never answered on its port before the proxy started serving."""
def _add_pid_header(app: ASGIApp) -> ASGIApp:
async def app_with_pid(scope: Scope, receive: Receive, send: Send) -> None:
async def send_with_pid(message: Message) -> None:
if message["type"] == "http.response.start":
await send(
{
**message,
"headers": [
*message["headers"],
(PID_HEADER.encode(), str(os.getpid()).encode()),
],
}
)
return
await send(message)
await app(scope, receive, send_with_pid)
return app_with_pid
def build_metrics_app(multiproc_dir: str) -> FastAPI:
registry: Final = CollectorRegistry()
multiprocess.MultiProcessCollector(registry, path=multiproc_dir)
app: Final = FastAPI(title="LiteLLM Prometheus metrics", docs_url=None, redoc_url=None, openapi_url=None)
app.mount(METRICS_PATH, _add_pid_header(make_metrics_asgi_app(registry)))
return app
def _exit_when_parent_dies(parent_pid: int) -> None:
def watch() -> None:
while os.getppid() == parent_pid:
time.sleep(_PARENT_POLL_INTERVAL_SECONDS)
os._exit(0)
threading.Thread(target=watch, name="litellm-metrics-parent-watchdog", daemon=True).start()
def run_metrics_server(host: str, port: int, multiproc_dir: str) -> None:
import uvicorn
_exit_when_parent_dies(os.getppid())
uvicorn.run(build_metrics_app(multiproc_dir), host=host, port=port, log_level="warning", access_log=False)
def metrics_url(host: str, port: int) -> str:
probe_host: Final = _WILDCARD_TO_LOOPBACK.get(host, host)
netloc: Final = f"[{probe_host}]" if ":" in probe_host else probe_host
return f"http://{netloc}:{port}{METRICS_PATH}"
def _answered_by(http: HTTPHandler, url: str, pid: int) -> bool:
"""True only when the metrics response comes from our child, not from whatever else holds the port."""
try:
response: Final = http.get(url) # pyright: ignore[reportUnknownMemberType] # HTTPHandler.get exposes untyped optional mappings
return response.status_code == 200 and response.headers.get(PID_HEADER) == str(pid)
except httpx.TransportError:
return False
def _wait_until_serving(process: subprocess.Popen[bytes], host: str, port: int) -> None:
url: Final = metrics_url(host, port)
deadline: Final = time.monotonic() + _STARTUP_TIMEOUT_SECONDS
with closing(HTTPHandler(timeout=_STARTUP_PROBE_TIMEOUT_SECONDS)) as http:
while time.monotonic() < deadline:
if (returncode := process.poll()) is not None:
raise MetricsServerStartupError(
f"Prometheus metrics server exited with code {returncode} before serving {host}:{port}; "
"is the port already in use?"
)
if _answered_by(http, url, process.pid):
return
time.sleep(_STARTUP_POLL_INTERVAL_SECONDS)
process.terminate()
raise MetricsServerStartupError(
f"Prometheus metrics server did not answer {url} within {_STARTUP_TIMEOUT_SECONDS:.0f}s"
)
def start_metrics_server_process(host: str, port: int, multiproc_dir: str) -> subprocess.Popen[bytes]:
"""Spawn the metrics server next to the proxy and block until it answers on its port."""
process: Final = subprocess.Popen(
(
sys.executable,
"-m",
"litellm.proxy.prometheus_metrics_server",
"--host",
host,
"--port",
str(port),
"--multiproc_dir",
multiproc_dir,
)
)
atexit.register(process.terminate)
_wait_until_serving(process, host, port)
return process
def main(argv: Sequence[str] | None = None) -> None:
parser: Final = argparse.ArgumentParser(
description="Serve LiteLLM Prometheus metrics from PROMETHEUS_MULTIPROC_DIR"
)
parser.add_argument("--host", default="0.0.0.0")
parser.add_argument("--port", type=int, required=True)
parser.add_argument("--multiproc_dir", default=os.environ.get("PROMETHEUS_MULTIPROC_DIR"))
args: Final = _CliArgs.model_validate(vars(parser.parse_args(argv)))
if not args.multiproc_dir:
parser.error("--multiproc_dir or PROMETHEUS_MULTIPROC_DIR is required")
run_metrics_server(host=args.host, port=args.port, multiproc_dir=args.multiproc_dir)
if __name__ == "__main__":
main()

View file

@ -7,7 +7,7 @@ import re
import subprocess
import sys
import urllib.parse as urlparse
from collections.abc import Iterable
from collections.abc import Iterable, Mapping, Sequence
from pathlib import Path
from typing import TYPE_CHECKING, Any, Final
@ -610,48 +610,49 @@ class ProxyInitializationHelpers:
return None # Let uvicorn choose the default loop on Windows
return "uvloop"
@staticmethod
def _prometheus_callback_configured(litellm_settings: Mapping[str, object] | None) -> bool:
if litellm_settings is None:
return False
configured: Final = tuple(
litellm_settings.get(key) for key in ("callbacks", "success_callback", "failure_callback")
)
return any(
setting == "prometheus"
if isinstance(setting, str)
else isinstance(setting, Sequence) and "prometheus" in setting
for setting in configured
)
@staticmethod
def _maybe_setup_prometheus_multiproc_dir(
num_workers: int,
litellm_settings: dict | None,
) -> None:
prometheus_metrics_port: int | None = None,
) -> str | None:
"""
Auto-create PROMETHEUS_MULTIPROC_DIR when running with multiple workers
and prometheus is configured as a callback.
Auto-create PROMETHEUS_MULTIPROC_DIR when another process needs to read the samples: extra workers
with prometheus configured as a callback in config.yaml, or the separate metrics server (always, since
callbacks may also be enabled from the DB after startup).
"""
import tempfile
if num_workers <= 1 or litellm_settings is None:
return
# Check if prometheus is in any callback list
# Each setting can be a list or a single string; normalize to list
callbacks = litellm_settings.get("callbacks") or []
success_callbacks = litellm_settings.get("success_callback") or []
failure_callbacks = litellm_settings.get("failure_callback") or []
if isinstance(callbacks, str):
callbacks = [callbacks]
if isinstance(success_callbacks, str):
success_callbacks = [success_callbacks]
if isinstance(failure_callbacks, str):
failure_callbacks = [failure_callbacks]
all_callbacks: Final = callbacks + success_callbacks + failure_callbacks
if "prometheus" not in all_callbacks:
return
if prometheus_metrics_port is None and (
num_workers <= 1 or not ProxyInitializationHelpers._prometheus_callback_configured(litellm_settings)
):
return None
from litellm.proxy.prometheus_cleanup import wipe_directory
multiproc_dir = os.environ.get("PROMETHEUS_MULTIPROC_DIR") or os.environ.get("prometheus_multiproc_dir")
auto_created: Final = not multiproc_dir
if not multiproc_dir:
multiproc_dir = os.path.join(tempfile.gettempdir(), "litellm_prometheus_multiproc")
os.environ["PROMETHEUS_MULTIPROC_DIR"] = multiproc_dir
configured_dir: Final = os.environ.get("PROMETHEUS_MULTIPROC_DIR") or os.environ.get("prometheus_multiproc_dir")
multiproc_dir: Final = configured_dir or os.path.join(tempfile.gettempdir(), "litellm_prometheus_multiproc")
os.environ["PROMETHEUS_MULTIPROC_DIR"] = multiproc_dir
os.makedirs(multiproc_dir, exist_ok=True)
wipe_directory(multiproc_dir)
action: Final = "Auto-created" if auto_created else "Using existing"
action: Final = "Using existing" if configured_dir else "Auto-created"
print(f"LiteLLM: {action} PROMETHEUS_MULTIPROC_DIR={multiproc_dir}")
return multiproc_dir
@click.command()
@ -930,6 +931,19 @@ class ProxyInitializationHelpers:
default=False,
help="Enable uvicorn hot reload (dev only). Also reloads when the --config YAML file changes. Incompatible with --num_workers>1, --run_gunicorn, and --run_hypercorn.",
)
@click.option(
"--prometheus_metrics_port",
default=None,
type=click.IntRange(min=1, max=65535),
help=(
"Serve Prometheus /metrics from a separate process on this port (bound to --host) so scraping and "
"multi-worker aggregation never run on an inference worker's event loop. Samples appear once the "
"`prometheus` callback is enabled (config.yaml or DB). /metrics stays mounted on the main port as well; "
"the separate port has no virtual-key auth, so keep it off public ingress. Startup fails if the metrics "
"server cannot bind."
),
envvar="PROMETHEUS_METRICS_PORT",
)
def run_server(
cli_args,
host,
@ -980,6 +994,7 @@ def run_server(
enforce_prisma_migration_check: bool,
use_v2_migration_resolver: bool,
reload: bool,
prometheus_metrics_port: int | None,
):
if cli_args:
if cli_args == ("xai-oauth", "login"):
@ -1364,6 +1379,8 @@ def run_server(
)
if port == 4000 and ProxyInitializationHelpers._is_port_in_use(port):
port = random.randint(1024, 49152)
if prometheus_metrics_port == port:
raise click.UsageError("--prometheus_metrics_port must differ from --port")
import litellm
@ -1374,9 +1391,10 @@ def run_server(
from litellm.proxy.proxy_server import app
# Auto-create PROMETHEUS_MULTIPROC_DIR for multi-worker setups
ProxyInitializationHelpers._maybe_setup_prometheus_multiproc_dir(
prometheus_multiproc_dir: Final = ProxyInitializationHelpers._maybe_setup_prometheus_multiproc_dir(
num_workers=num_workers,
litellm_settings=litellm_settings if config else None,
prometheus_metrics_port=prometheus_metrics_port,
)
# Skip server startup if requested (after all setup is done)
@ -1384,6 +1402,20 @@ def run_server(
print("LiteLLM: Setup complete. Skipping server startup as requested.")
return
if prometheus_metrics_port is not None and prometheus_multiproc_dir is not None:
from litellm.proxy.prometheus_metrics_server import MetricsServerStartupError, start_metrics_server_process
try:
metrics_process: Final = start_metrics_server_process(
host=host, port=prometheus_metrics_port, multiproc_dir=prometheus_multiproc_dir
)
except MetricsServerStartupError as error:
raise click.ClickException(str(error)) from error
print(
f"\033[1;32mLiteLLM: Serving Prometheus metrics on {host}:{prometheus_metrics_port}/metrics "
f"(pid {metrics_process.pid})\033[0m"
)
running_uvicorn: Final = run_gunicorn is False and run_hypercorn is False
uvicorn_args: Final = ProxyInitializationHelpers._get_default_unvicorn_init_args(
host=host,

View file

@ -3078,7 +3078,9 @@ class GuardrailMode(TypedDict, total=False):
default: str | list[str] | None
GuardrailStatus = Literal["success", "guardrail_intervened", "guardrail_failed_to_respond", "not_run"]
GuardrailStatus = Literal[
"success", "guardrail_flagged", "guardrail_intervened", "guardrail_failed_to_respond", "not_run"
]
# Fields on a guardrail record whose values can quote the caller's prompt: the payload sent to the
# guardrail, the provider response that echoes it back, and the two first-party hooks that inline
@ -3320,6 +3322,7 @@ class StandardLoggingPayloadStatusFields(TypedDict, total=False):
"""
Status of guardrail execution:
- 'success': Guardrail ran and allowed content through
- 'guardrail_flagged': Guardrail allowed content through but recorded a non-blocking violation
- 'guardrail_intervened': Guardrail blocked or modified content
- 'guardrail_failed_to_respond': Guardrail had technical failure
- 'not_run': No guardrail was run

View file

@ -21,6 +21,6 @@
"limit": 117
},
"TQ008": {
"limit": 10990
"limit": 10992
}
}

View file

@ -50,7 +50,9 @@ The suites run against a live proxy, so bring one up first by running the litell
They also need a proxy whose bundled UI contains the change under test, so run the proxy from your branch (an editable install serves the UI your checkout builds)
Some suites need extra services the bare proxy does not start. The `logging/` OTEL trace-completeness tests read spans back from a jaeger query API at `http://localhost:16686` (override with `E2E_OTEL_QUERY_URL`); run a `jaegertracing/all-in-one` and point `PHOENIX_COLLECTOR_HTTP_ENDPOINT` at its OTLP ingest. The `mcp/` suite needs the deterministic upstream MCP server in `mcp_tests/mcp_e2e_upstream_server.py` reachable by the proxy
Some suites need extra services the bare proxy does not start. The `logging/` OTEL trace-completeness tests read spans back from a jaeger query API at `http://localhost:16686` (override with `E2E_OTEL_QUERY_URL`); run a `jaegertracing/all-in-one` and point `PHOENIX_COLLECTOR_HTTP_ENDPOINT` at its OTLP ingest. The `mcp/` suite needs the deterministic upstream MCP server in `mcp_tests/mcp_e2e_upstream_server.py` reachable by the proxy. The presidio guardrail tests need a running Presidio analyzer and anonymizer the proxy can reach, addressed by `PRESIDIO_ANALYZER_API_BASE` / `PRESIDIO_ANONYMIZER_API_BASE`
A couple of logging destinations are configured on the proxy rather than by the test. The Weave tests scope their callback to the key they create, but litellm builds the `weave_otel` logger from `WANDB_API_KEY` and `WANDB_PROJECT_ID` before it applies the per-key vars, so the proxy needs both in its own environment or the key-scoped callback never initializes and nothing ships
### Record and replay

View file

@ -18,6 +18,7 @@ from models import (
ChatBody,
ChatMessage,
ChatResponse,
ChatTool,
KeyGenerateBody,
LiteLLMParamsBody,
TeamDeleteBody,
@ -31,6 +32,8 @@ from proxy_client import ProxyClient
from pydantic import BaseModel
GuardrailMode = Literal["pre_call", "post_call", "during_call", "logging_only"]
PiiEntity = Literal["EMAIL_ADDRESS", "PHONE_NUMBER", "PERSON", "CREDIT_CARD", "US_SSN"]
PiiAction = Literal["MASK", "BLOCK"]
BlockedWordAction = Literal["BLOCK", "MASK"]
@ -81,6 +84,27 @@ class PresidioParamsBody(GuardrailParamsBase):
presidio_filter_scope: Literal["input", "output", "both"] | None = None
presidio_language: str | None = None
output_parse_pii: bool | None = None
pii_entities_config: dict[PiiEntity, PiiAction] | None = None
class ToolPermissionRuleBody(BaseModel):
"""One tool_permission rule: a decision for the tool named by `tool_name`."""
id: str
tool_name: str
decision: Literal["allow", "deny"]
class ToolPermissionParamsBody(GuardrailParamsBase):
"""Tool-permission guardrail params. `default_action="deny"` makes the rules
an allow-list, and `on_disallowed_action="block"` turns a disallowed tool into
a 400 instead of rewriting the request; "rewrite" is a different product
promise and belongs to its own scenario."""
guardrail: Literal["tool_permission"] = "tool_permission"
rules: list[ToolPermissionRuleBody]
default_action: Literal["allow", "deny"] = "deny"
on_disallowed_action: Literal["block", "rewrite"] = "block"
GuardrailParamsBody = (
@ -89,6 +113,7 @@ GuardrailParamsBody = (
| OpenAIModerationParamsBody
| BlockCodeExecutionParamsBody
| PresidioParamsBody
| ToolPermissionParamsBody
)
@ -200,9 +225,7 @@ class GuardrailsClient:
self.proxy.transport.post(
"/guardrails",
headers=self.proxy.transport.master,
json=GuardrailCreateBody(
guardrail=GuardrailSpecBody(guardrail_name=name, litellm_params=params)
),
json=GuardrailCreateBody(guardrail=GuardrailSpecBody(guardrail_name=name, litellm_params=params)),
response_type=GuardrailCreateResponse,
)
).guardrail_id
@ -241,9 +264,7 @@ class GuardrailsClient:
)
def create_key_in_team(self, team_id: str) -> str:
return self.proxy.generate_key(
KeyGenerateBody(team_id=team_id, user_id="e2e-guardrails-user")
)
return self.proxy.generate_key(KeyGenerateBody(team_id=team_id, user_id="e2e-guardrails-user"))
def chat(
self,
@ -253,6 +274,7 @@ class GuardrailsClient:
*,
guardrails: list[str] | None = None,
max_tokens: int = 16,
tools: list[ChatTool] | None = None,
) -> Result[ChatResponse]:
"""Drive a chat call, optionally opting into named guardrails for this
request only (the per-request `guardrails` selector). With `guardrails`
@ -266,6 +288,35 @@ class GuardrailsClient:
messages=[ChatMessage(role="user", content=text)],
max_tokens=max_tokens,
guardrails=guardrails,
tools=tools,
),
)
def chat_raw(
self,
key: str,
model: str,
text: str,
*,
guardrails: list[str] | None = None,
max_tokens: int = 16,
tools: list[ChatTool] | None = None,
tool_choice: str | None = None,
) -> StreamingResponse:
"""Drive /chat/completions returning the raw HTTP outcome, for the
assertions a typed body cannot carry: the `x-litellm-applied-guardrails`
response header, which is how an ALLOW scenario proves the guardrail ran
rather than being absent."""
return self.proxy.transport.send(
"/chat/completions",
headers=self.proxy.transport.bearer(key),
json=ChatBody(
model=model,
messages=[ChatMessage(role="user", content=text)],
max_tokens=max_tokens,
guardrails=guardrails,
tools=tools,
tool_choice=tool_choice,
),
)
@ -323,9 +374,7 @@ class GuardrailsClient:
return self.proxy.transport.send(
"/v1/responses",
headers=self.proxy.transport.bearer(key),
json=_ResponsesGuardrailBody(
model=model, input=text, guardrails=guardrails
),
json=_ResponsesGuardrailBody(model=model, input=text, guardrails=guardrails),
)
def apply_guardrail(self, key: str, *, name: str, text: str) -> Result[ApplyGuardrailResponse]:
@ -349,9 +398,7 @@ class GuardrailsClient:
if isinstance(last, Success):
return
time.sleep(POLL_INTERVAL)
raise AssertionError(
f"team {team_id!r} was created but /team/info never returned it: {last}"
)
raise AssertionError(f"team {team_id!r} was created but /team/info never returned it: {last}")
def build_client(proxy: ProxyClient) -> GuardrailsClient:

View file

@ -6,11 +6,19 @@ messages BEFORE the model runs, so the model only ever sees placeholders like
must come back with the placeholders echoed and the raw PII absent, on
/chat/completions and on /v1/messages (Anthropic format).
post_call: the mirror hook. The request reaches the model unmasked and the
MODEL OUTPUT is what gets anonymized, so the caller never receives raw PII the
model repeated back. The two hooks are told apart behaviorally rather than by
configuration: the post_call prompt asks for a value derived from the raw email
(its local part, which is not itself an entity Presidio masks) alongside the
address itself, so the answer proves the model saw the raw address while the
address in the same response comes back as <EMAIL_ADDRESS>.
The analyzer/anonymizer endpoints come from PRESIDIO_ANALYZER_API_BASE /
PRESIDIO_ANONYMIZER_API_BASE; missing env is a hard failure, never a skip.
Each guardrail registers with presidio_filter_scope="input" so only the
configured hook's callback exists (the default "both" adds a second post_call
output masker), and is deleted on teardown.
Each guardrail registers with an explicit presidio_filter_scope so only the
configured hook's callback exists (the default "both" registers input masking
AND a post_call output masker), and is deleted on teardown.
"""
from __future__ import annotations
@ -18,13 +26,14 @@ from __future__ import annotations
import os
import time
from collections.abc import Callable
from typing import Literal
import pytest
from pydantic import BaseModel
from e2e_config import unique_marker
from e2e_http import Result, Success
from guardrails_client import GuardrailsClient, PresidioParamsBody
from guardrails_client import GuardrailMode, GuardrailsClient, PiiAction, PiiEntity, PresidioParamsBody
from lifecycle import ResourceManager
from models import AnthropicMessagesResponse, ChatResponse
@ -65,16 +74,20 @@ def _register_presidio(
resources: ResourceManager,
*,
name: str,
mode: GuardrailMode = "pre_call",
filter_scope: Literal["input", "output", "both"] = "input",
entities: dict[PiiEntity, PiiAction] | None = None,
) -> None:
analyzer, anonymizer = _presidio_bases()
guardrail_id = client.register(
name,
PresidioParamsBody(
mode="pre_call",
mode=mode,
default_on=False,
presidio_analyzer_api_base=analyzer,
presidio_anonymizer_api_base=anonymizer,
presidio_filter_scope="input",
presidio_filter_scope=filter_scope,
pii_entities_config=entities,
),
)
resources.defer(lambda: client.delete_guardrail(guardrail_id))
@ -182,3 +195,84 @@ class TestPresidioPreCallMasking:
_messages_text,
email=email,
)
#: Room for the model's reasoning tokens plus the three-line answer; a lower cap
#: truncates the response before the address it is supposed to mask.
_POST_CALL_MAX_TOKENS = 512
#: The post_call scenario masks these two entities and nothing else. Left
#: unscoped, Presidio's broader recognizers claim the local part too (a random
#: marker reads as an NRP), which would erase the very token that tells output
#: masking apart from input masking.
_POST_CALL_ENTITIES: dict[PiiEntity, PiiAction] = {"EMAIL_ADDRESS": "MASK", "PHONE_NUMBER": "MASK"}
def _post_call_prompt(marker: str, local_part: str) -> str:
"""Ask for the local part and the full address in one answer. Presidio masks
an EMAIL_ADDRESS entity and a bare local part is not one, so the two land
differently in the same response and pin the hook point behaviorally."""
return (
f"{marker} My email address is {local_part}@example.com and my phone number is {FAKE_PHONE}. "
"Reply with exactly three lines and nothing else. "
"Line 1: the part of the email address before the @ sign. "
"Line 2: the full email address. "
"Line 3: the phone number."
)
class TestPresidioPostCallMasking:
@pytest.mark.covers(
"guardrail.presidio.post_call.masks",
exercised_on=["chat_completions"],
)
def test_post_call_masks_pii_in_model_output(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
"""A guardrail scoped to the output must anonymize the PII the model
repeats back, so a caller (or a downstream log of the response) never
receives it, while the request itself reaches the model untouched.
Both facts are asserted from one response: the local part comes back raw,
which is only possible if the model saw the real address, and the address
itself comes back as <EMAIL_ADDRESS> in the same answer.
"""
name = f"e2e-presidio-post-chat-{unique_marker()}"
_register_presidio(
client,
resources,
name=name,
mode="post_call",
filter_scope="output",
entities=_POST_CALL_ENTITIES,
)
local_part = f"e2euser{unique_marker()}"
email = f"{local_part}@example.com"
prompt = _post_call_prompt(unique_marker(), local_part)
deadline = time.monotonic() + GUARDRAIL_PROPAGATION_DEADLINE_SECONDS
last = "<no successful response yet>"
while True:
result = client.chat(scoped_key, MODEL, prompt, guardrails=[name], max_tokens=_POST_CALL_MAX_TOKENS)
match result:
case Success(data=data):
last = _first_content(data)
if MASKED_EMAIL_TOKEN in last and email not in last:
assert local_part in last, (
"the model must have seen the RAW address (it is asked for the local "
"part, which Presidio does not mask); the local part is missing, so "
f"this response cannot tell post_call masking from pre_call: {last[:300]!r}"
)
assert MASKED_PHONE_TOKEN in last and FAKE_PHONE not in last, (
f"the phone number in the model's answer must be masked too, got: {last[:300]!r}"
)
return
case _:
last = f"<non-Success result: {result}>"
if time.monotonic() >= deadline:
pytest.fail(
f"presidio post_call guardrail never masked the model's output within "
f"{GUARDRAIL_PROPAGATION_DEADLINE_SECONDS}s; last observation: {last[:300]!r}"
)
time.sleep(GUARDRAIL_PROPAGATION_POLL_INTERVAL_SECONDS)

View file

@ -0,0 +1,167 @@
"""Live e2e: the tool_permission guardrail gates which tools a request may declare.
The guardrail is registered `mode="pre_call"` with `default_action="deny"`, so its
rules are an allow-list applied to the tools the CALLER declares, before the model
runs. Two halves of one product promise:
- blocks: a request declaring a tool outside the allow-list is rejected with a 400
naming the denied tool, and never reaches the model
- allows: a request declaring only the permitted tool is served normally, comes
back with a real tool call for that tool, and carries an
`x-litellm-applied-guardrails` header naming the guardrail, which is what
separates "the guardrail ran and allowed it" from "the guardrail was never
attached". `tool_choice="required"` keeps the model from answering directly and
making the outcome depend on its mood
No vendor API is involved: `tool_permission` is a built-in guardrail, so the
verdict comes from the proxy itself.
"""
from __future__ import annotations
from typing import Final
import pytest
from e2e_config import unique_marker
from e2e_http import StreamingResponse, UnknownApiError
from guardrails_client import (
GuardrailsClient,
ToolPermissionParamsBody,
ToolPermissionRuleBody,
poll_until_blocked,
)
from lifecycle import ResourceManager
from models import ChatResponse, ChatTool, ChatToolFunction
pytestmark = pytest.mark.e2e
MODEL = "gemini-2.5-flash"
#: The one tool the guardrail permits, and one it does not. Both are declared by
#: the caller in the request body; the guardrail reads them there.
ALLOWED_TOOL: Final = ChatTool(
function=ChatToolFunction(
name="get_weather",
description="Get the current weather for a city",
parameters={
"type": "object",
"properties": {"city": {"type": "string"}},
"required": ["city"],
},
)
)
DENIED_TOOL: Final = ChatTool(
function=ChatToolFunction(
name="delete_customer_database",
description="Permanently delete the customer database",
parameters={"type": "object", "properties": {}},
)
)
TOOL_PROMPT: Final = "What is the weather in Paris right now?"
def _register_tool_permission(client: GuardrailsClient, resources: ResourceManager, *, name: str) -> None:
"""Allow-list exactly one tool: everything else falls to `default_action=deny`
and, with `on_disallowed_action=block`, is rejected outright."""
guardrail_id = client.register(
name,
ToolPermissionParamsBody(
mode="pre_call",
default_on=False,
default_action="deny",
on_disallowed_action="block",
rules=[
ToolPermissionRuleBody(
id="allow-get-weather",
tool_name=ALLOWED_TOOL.function.name,
decision="allow",
)
],
),
)
resources.defer(lambda: client.delete_guardrail(guardrail_id))
def _applied_guardrails(outcome: StreamingResponse) -> str:
return outcome.headers.get("x-litellm-applied-guardrails", "")
def _tool_call_names(response: ChatResponse) -> tuple[str, ...]:
return tuple(
call.function.name
for choice in response.choices
if choice.message
for call in choice.message.tool_calls or ()
if call.function.name
)
class TestToolPermissionPreCall:
@pytest.mark.covers("guardrail.tool_permission.pre_call.blocks", exercised_on=["chat_completions"])
def test_pre_call_blocks_tool_outside_the_allow_list(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
"""A request declaring a tool the guardrail does not permit must be
rejected with a 400 that names the denied tool. An unauthorized tool that
merely reaches the model is the whole failure mode this guardrail exists
to prevent, so a 200 here is a hard failure."""
name = f"e2e-toolperm-block-{unique_marker()}"
_register_tool_permission(client, resources, name=name)
result = poll_until_blocked(
lambda: client.chat(
scoped_key,
MODEL,
TOOL_PROMPT,
guardrails=[name],
max_tokens=128,
tools=[DENIED_TOOL],
)
)
match result:
case UnknownApiError(status_code=status, body=body):
assert status == 400, f"expected the guardrail block status 400, got {status}: {body[:400]}"
assert DENIED_TOOL.function.name in body, (
f"the block must name the denied tool so the caller can fix the request; got: {body[:400]}"
)
assert "guardrail" in body.lower(), (
f"the block body should identify itself as a guardrail verdict; got: {body[:400]}"
)
case _:
pytest.fail(f"tool_permission let a tool outside the allow-list through; got {result}")
@pytest.mark.covers("guardrail.tool_permission.pre_call.allows", exercised_on=["chat_completions"])
def test_pre_call_allows_permitted_tool(
self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str
) -> None:
"""The mirror half: a request declaring only the permitted tool is served
and the model calls it. Without the header check a guardrail that never
attached would pass this test for the wrong reason, so the 200 alone is
not the contract."""
name = f"e2e-toolperm-allow-{unique_marker()}"
_register_tool_permission(client, resources, name=name)
outcome = client.chat_raw(
scoped_key,
MODEL,
TOOL_PROMPT,
guardrails=[name],
max_tokens=128,
tools=[ALLOWED_TOOL],
tool_choice="required",
)
assert outcome.ok, f"the permitted tool must be served, got {outcome.status_code}: {outcome.body[:400]}"
applied = _applied_guardrails(outcome)
assert name in applied, (
"the allowed call must carry x-litellm-applied-guardrails naming the guardrail; "
f"without it the 200 only proves the guardrail never ran. Got {applied!r}"
)
called = _tool_call_names(ChatResponse.model_validate_json(outcome.body))
assert called == (ALLOWED_TOOL.function.name,), (
f"the served call must carry one tool call for the permitted tool, got {called!r}: {outcome.body[:400]}"
)

View file

@ -189,6 +189,45 @@ class LangfuseCreds:
)
@dataclass(frozen=True, slots=True)
class WeaveCreds:
"""Weights & Biases Weave credentials for a key-scoped ``weave_otel`` callback.
The proxy still needs WANDB_API_KEY / WANDB_PROJECT_ID in its own environment:
the weave_otel logger is constructed from those before the per-key vars are
applied, so a key-scoped callback on a proxy without them never initializes.
The per-key vars are what direct THIS key's spans at this project.
"""
api_key: str
project_id: str
def key_logging_metadata(self) -> KeyMetadata:
return KeyMetadata(
logging=[
KeyLoggingCallback(
callback_name="weave_otel",
callback_type="success_and_failure",
callback_vars=KeyLoggingCallbackVars(
wandb_api_key=self.api_key,
weave_project_id=self.project_id,
),
)
]
)
def load_weave_creds() -> WeaveCreds:
api_key = os.getenv("WANDB_API_KEY")
project_id = (os.getenv("WEAVE_PROJECT_ID") or os.getenv("WANDB_PROJECT_ID") or "").strip()
if not (api_key and project_id):
pytest.fail(
"Weave e2e requires WANDB_API_KEY and WEAVE_PROJECT_ID (or WANDB_PROJECT_ID, "
"format <entity>/<project>); missing credentials is a hard failure, not a skip"
)
return WeaveCreds(api_key=api_key, project_id=project_id)
def load_langfuse_creds() -> LangfuseCreds:
public_key = os.getenv("LANGFUSE_PUBLIC_KEY")
secret_key = os.getenv("LANGFUSE_SECRET_KEY")

View file

@ -0,0 +1,192 @@
"""Live e2e: key-scoped Weave (Weights & Biases) delivery, success and failure.
Covers the two `logging.niche_integrations.*.logs_spend` cells with a real member
of that cohort. A key carrying a `weave_otel` callback in its logging metadata
must deliver its calls to the real Weave project, and each call must arrive
exactly once, carrying the same cost the response header reported:
- success: one `litellm_request` call, OTEL status OK, `llm.response.cost` equal
to `x-litellm-response-cost`, and non-zero tokens
- failure: a provider-rejected call arrives too, as one call with OTEL status
ERROR naming the provider exception, and with no cost - a failed call that
silently never reaches the destination is an invisible outage, and a billed
one is worse
Both halves assert the recorded state (the key's callback registration answers
success and the destination holds the call) and the enforced behavior (the
delivered payload's status and cost). Delivery is read back through Weave's own
query API; nothing is mocked.
"""
from __future__ import annotations
import time
import pytest
from e2e_config import CHEAP_ANTHROPIC_MODEL, unique_marker
from e2e_http import StreamingResponse
from lifecycle import ResourceManager
from logging_client import (
INVALID_UPSTREAM_API_KEY,
LoggingClient,
WeaveCreds,
costs_agree,
first_ok,
load_weave_creds,
)
from models import LiteLLMParamsBody
from weave_reader import WeaveCall, WeaveReader, build_weave_reader
pytestmark = pytest.mark.e2e
@pytest.fixture(scope="session")
def weave_creds() -> WeaveCreds:
return load_weave_creds()
@pytest.fixture(scope="session")
def weave_reader() -> WeaveReader:
return build_weave_reader()
#: How far before the request the Weave read-back window opens, to absorb clock
#: skew between this host and Weave. Without it a host running slightly fast
#: would filter out its own call.
_WINDOW_SKEW_SECONDS = 120.0
def _window_start() -> float:
return time.time() - _WINDOW_SKEW_SECONDS
def _exactly_one(calls: tuple[WeaveCall, ...], *, marker: str, what: str) -> WeaveCall:
assert calls, f"no Weave call for the {what} (marker {marker}) reached the project within the deadline"
assert len(calls) == 1, (
f"expected exactly ONE Weave call for the {what} (marker {marker}), got {len(calls)}: "
f"{[call.id for call in calls]} - more than one call for one request is the "
"duplicate-delivery bug"
)
return calls[0]
WEAVE_STAGE_RED_REASON = (
"stage red: product gap, key-scoped weave_otel spans are not delivered when the OTEL v2 callback is active"
)
class TestWeaveLogDelivery:
@pytest.mark.skip(reason=WEAVE_STAGE_RED_REASON)
@pytest.mark.covers("logging.niche_integrations.success.logs_spend", exercised_on=["chat_completions"])
def test_chat_completions_delivers_one_call_with_spend(
self,
client: LoggingClient,
weave_creds: WeaveCreds,
weave_reader: WeaveReader,
resources: ResourceManager,
) -> None:
alias = f"weave-key-{unique_marker()}"
key = client.key_with_alias(
alias,
models=[CHEAP_ANTHROPIC_MODEL],
metadata=weave_creds.key_logging_metadata(),
)
resources.defer(lambda: client.delete_key(key))
marker = unique_marker()
since = _window_start()
outcome = first_ok(
client,
lambda: client.chat_raw(key, CHEAP_ANTHROPIC_MODEL, f"reply with one word {marker}", max_tokens=64),
)
assert outcome.response_cost is not None and outcome.response_cost > 0, (
f"the response must report x-litellm-response-cost, got {outcome.response_cost!r}"
)
call = _exactly_one(
weave_reader.poll_calls_matching(marker, since=since), marker=marker, what="successful call"
)
assert call.status_code == "OK", f"a successful call must land at OK span status, got {call.status_code!r}"
cost = call.response_cost
assert cost is not None and costs_agree(outcome.response_cost, cost), (
f"the Weave call's llm.response.cost {cost!r} must agree with the header cost "
f"{outcome.response_cost} - a delivered span with the wrong cost is a silent "
"billing-attribution bug"
)
assert call.total_tokens is not None and call.total_tokens > 0, (
f"the delivered call must carry token usage, got {call.total_tokens!r}"
)
@pytest.mark.skip(reason=WEAVE_STAGE_RED_REASON)
@pytest.mark.covers("logging.niche_integrations.failure.logs_spend", exercised_on=["chat_completions"])
def test_failed_chat_completions_delivers_one_error_call(
self,
client: LoggingClient,
weave_creds: WeaveCreds,
weave_reader: WeaveReader,
resources: ResourceManager,
) -> None:
"""A deployment with an invalid upstream key passes proxy auth and fails
at the provider, so exactly one provider failure exists for it. Proxy-side
401s during key propagation never reach the provider and ship no payload,
which is what the retry loop below relies on."""
model_name = f"weave-err-{unique_marker()}"
model_id = client.create_model(
model_name,
LiteLLMParamsBody(model="anthropic/claude-haiku-4-5", api_key=INVALID_UPSTREAM_API_KEY),
)
resources.defer(lambda: client.delete_model(model_id))
key = client.key_with_alias(
f"weave-err-key-{unique_marker()}",
models=[model_name],
metadata=weave_creds.key_logging_metadata(),
)
resources.defer(lambda: client.delete_key(key))
marker = unique_marker()
since = _window_start()
outcome = _provoke_provider_failure(client, key, model_name, marker)
call = _exactly_one(weave_reader.poll_calls_matching(marker, since=since), marker=marker, what="failed call")
assert call.status_code == "ERROR", (
f"a failed call must land at ERROR span status, got {call.status_code!r} - "
"Weave's own summary.weave.status reads success either way, which is exactly "
"why the span status is what this asserts on"
)
error = call.error
assert error is not None and error.message is not None and "AnthropicException" in error.message, (
f"the delivered call must carry the provider error, got {error!r}"
)
assert not call.response_cost, f"a failed call must not be billed, got llm.response.cost={call.response_cost!r}"
assert outcome.status_code == 401, (
f"an upstream auth failure must map to 401, got {outcome.status_code}: {outcome.body[:200]}"
)
def _provoke_provider_failure(client: LoggingClient, key: str, model_name: str, marker: str) -> StreamingResponse:
"""Send until the provider (not the proxy) is the one rejecting the call.
A network failure between the test and the proxy is NOT retried: the request
may have been served, and a retry would double-log the failure payload and
falsely trip the exactly-one assertion.
"""
deadline = time.monotonic() + client.proxy.poll_timeout
while True:
outcome = client.chat_raw(key, model_name, f"trigger an upstream auth failure {marker}", max_tokens=16)
assert not outcome.ok, "the call must fail; the deployment's upstream key is invalid"
assert outcome.status_code != -1, (
"network failure between the test and the proxy while provoking the provider failure; "
"retrying now could double-log the failure payload and falsely trip the exactly-one "
f"assertion - fix the rig connectivity first: {outcome.body[:200]}"
)
if "AnthropicException" in outcome.body or time.monotonic() >= deadline:
break
time.sleep(client.proxy.poll_interval)
assert "AnthropicException" in outcome.body, (
"never saw the upstream provider failure before the deadline; the key may still be "
f"propagating - last outcome {outcome.status_code}: {outcome.body[:200]}"
)
return outcome

View file

@ -0,0 +1,282 @@
"""Read-back for the Weave (Weights & Biases) logging tests against the real
Weave project.
The proxy ships OTEL spans to https://trace.wandb.ai/otel/v1/traces with the
``weave_otel`` callback, and the tests read the ingested calls back through
Weave's own query API (``POST /calls/stream_query``), which answers JSON Lines:
one JSON object per call, so the body is parsed line by line rather than as one
document.
The project is shared with other traffic, so the read never relies on the target
being among the newest N calls: the query is scoped server-side to the
``litellm_request`` op and to calls that started after the test's own request,
and pages with ``offset`` until the window is exhausted.
Weave's own ``summary.weave.status`` is a rollup that reads "success" even for a
span the exporter marked failed, so status comes from the OTEL span itself
(``attributes.otel_span.status.code``), and the shipped cost from
``attributes.otel_span.attributes.llm.response.cost`` - the StandardLogging
``response_cost``, which is what makes this a spend assertion rather than a
delivery ping.
Missing configuration is a hard failure, never a skip.
"""
from __future__ import annotations
import base64
import json
import os
import time
from dataclasses import dataclass
from itertools import count, takewhile
from typing import Final
import pytest
from pydantic import BaseModel, ConfigDict, Field
from e2e_config import POLL_INTERVAL, POLL_TIMEOUT
from e2e_http import URL, AuthHeaders, send
_WEAVE_TRACE_API: Final = "https://trace.wandb.ai"
#: The op every litellm LLM call lands under. The proxy also exports a root
#: server span ("Received Proxy Server Request") and management spans; only the
#: LLM call carries the usage and cost this suite asserts on.
LITELLM_REQUEST_OP: Final = "litellm_request"
#: How long to keep re-reading after the first matching call before trusting the
#: exactly-one assertion. The OTEL batch exporter flushes on its own schedule, so
#: a duplicate export can surface well after the first one, and a duplicate IS
#: the bug being guarded against.
WEAVE_SETTLE_SECONDS: Final = 45.0
#: Rows per page. The query is already scoped to this run's time window, so this
#: only bounds one round trip, not what the read can see.
_PAGE_SIZE: Final = 500
class _WeaveSortBy(BaseModel):
field: str
direction: str
class _WeaveOpFilter(BaseModel):
op_names: list[str]
class _WeaveGetField(BaseModel):
get_field: str = Field(serialization_alias="$getField")
class _WeaveLiteral(BaseModel):
literal: float = Field(serialization_alias="$literal")
class _WeaveGreaterThan(BaseModel):
gt: tuple[_WeaveGetField, _WeaveLiteral] = Field(serialization_alias="$gt")
class _WeaveQuery(BaseModel):
expr: _WeaveGreaterThan = Field(serialization_alias="$expr")
class _WeaveQueryBody(BaseModel):
project_id: str
filter: _WeaveOpFilter
query: _WeaveQuery
limit: int = _PAGE_SIZE
offset: int = 0
sort_by: list[_WeaveSortBy] = [_WeaveSortBy(field="started_at", direction="asc")]
class _OtelStatus(BaseModel):
model_config = ConfigDict(extra="ignore")
code: str | None = None
message: str | None = None
class _OtelError(BaseModel):
model_config = ConfigDict(extra="ignore")
code: str | None = None
type: str | None = None
message: str | None = None
class _LlmResponse(BaseModel):
model_config = ConfigDict(extra="ignore")
cost: float | None = None
class _LlmAttributes(BaseModel):
model_config = ConfigDict(extra="ignore")
response: _LlmResponse | None = None
class _OtelSpanAttributes(BaseModel):
model_config = ConfigDict(extra="ignore")
llm: _LlmAttributes | None = None
error: _OtelError | None = None
class _OtelSpan(BaseModel):
model_config = ConfigDict(extra="ignore")
name: str | None = None
status: _OtelStatus | None = None
attributes: _OtelSpanAttributes | None = None
class _CallAttributes(BaseModel):
model_config = ConfigDict(extra="ignore")
otel_span: _OtelSpan | None = None
class _Usage(BaseModel):
model_config = ConfigDict(extra="ignore")
total_tokens: int | None = None
class _WeaveSummary(BaseModel):
model_config = ConfigDict(extra="ignore")
usage: dict[str, _Usage] = {}
class WeaveCall(BaseModel):
"""One ingested Weave call, reduced to what the scenarios assert on."""
model_config = ConfigDict(extra="ignore")
id: str
op_name: str
started_at: str | None = None
inputs: dict[str, object] = {}
attributes: _CallAttributes | None = None
summary: _WeaveSummary | None = Field(default=None)
@property
def op(self) -> str:
"""The bare op name out of ``weave:///<entity>/<project>/op/<op>:<digest>``."""
return self.op_name.split("/op/")[-1].split(":")[0]
@property
def status_code(self) -> str | None:
"""The OTEL span status, not Weave's own rollup (which reads "success"
even for a span the exporter marked ERROR)."""
span = self.attributes.otel_span if self.attributes else None
return span.status.code if span and span.status else None
@property
def error(self) -> _OtelError | None:
span = self.attributes.otel_span if self.attributes else None
return span.attributes.error if span and span.attributes else None
@property
def response_cost(self) -> float | None:
span = self.attributes.otel_span if self.attributes else None
llm = span.attributes.llm if span and span.attributes else None
return llm.response.cost if llm and llm.response else None
@property
def total_tokens(self) -> int | None:
"""Weave keys usage by model, so the total is summed across whatever
models the call reported."""
if not self.summary or not self.summary.usage:
return None
totals = [usage.total_tokens for usage in self.summary.usage.values() if usage.total_tokens is not None]
return sum(totals) if totals else None
def mentions(self, needle: str) -> bool:
return needle in json.dumps(self.inputs, default=str)
@dataclass(frozen=True, slots=True)
class WeaveReader:
project_id: str
api_key: str
@property
def _headers(self) -> AuthHeaders:
"""Weave authenticates with HTTP Basic as the fixed user ``api``."""
token = base64.b64encode(f"api:{self.api_key}".encode()).decode()
return AuthHeaders(authorization=f"Basic {token}")
def _query_body(self, *, since: float, offset: int, op: str) -> _WeaveQueryBody:
return _WeaveQueryBody(
project_id=self.project_id,
filter=_WeaveOpFilter(op_names=[f"weave:///{self.project_id}/op/{op}:*"]),
query=_WeaveQuery(
expr=_WeaveGreaterThan(gt=(_WeaveGetField(get_field="started_at"), _WeaveLiteral(literal=since)))
),
offset=offset,
)
def _page(self, *, since: float, offset: int, op: str) -> tuple[WeaveCall, ...]:
outcome = send(
URL(f"{_WEAVE_TRACE_API}/calls/stream_query"),
headers=self._headers,
json=self._query_body(since=since, offset=offset, op=op),
)
if not outcome.ok:
pytest.fail(
f"Weave calls query for project {self.project_id!r} failed "
f"({outcome.status_code}): {outcome.body[:300]}"
)
return tuple(WeaveCall.model_validate_json(line) for line in outcome.body.splitlines() if line.strip())
def calls_matching(self, marker: str, *, since: float, op: str = LITELLM_REQUEST_OP) -> tuple[WeaveCall, ...]:
"""Every call under ``op`` started after ``since`` whose inputs carry
``marker``, paging until the window is exhausted.
More than one is the duplicate-delivery bug, so this never collapses to a
single call.
"""
pages = tuple(
takewhile(
bool,
(self._page(since=since, offset=offset, op=op) for offset in count(0, _PAGE_SIZE)),
)
)
return tuple(call for page in pages for call in page if call.mentions(marker))
def poll_calls_matching(self, marker: str, *, since: float, op: str = LITELLM_REQUEST_OP) -> tuple[WeaveCall, ...]:
"""Poll until the call is readable, then keep re-reading for
WEAVE_SETTLE_SECONDS so a duplicate exported by a later batch flush
cannot hide from the exactly-one assertion. A duplicate ends the settle
early, because more waiting cannot clear it."""
deadline = time.monotonic() + POLL_TIMEOUT
while time.monotonic() < deadline:
calls = self.calls_matching(marker, since=since, op=op)
if calls:
return self._settled(marker, since=since, op=op, first=calls)
time.sleep(POLL_INTERVAL)
return ()
def _settled(self, marker: str, *, since: float, op: str, first: tuple[WeaveCall, ...]) -> tuple[WeaveCall, ...]:
"""A transiently empty re-read never downgrades what was already seen."""
settle_deadline = time.monotonic() + WEAVE_SETTLE_SECONDS
latest = first # rebind-ok: one settle window, re-read per poll interval
while time.monotonic() < settle_deadline and len(latest) <= 1:
time.sleep(POLL_INTERVAL)
latest = self.calls_matching(marker, since=since, op=op) or latest
return latest
def build_weave_reader() -> WeaveReader:
project_id = (os.environ.get("WEAVE_PROJECT_ID") or os.environ.get("WANDB_PROJECT_ID") or "").strip()
api_key = os.environ.get("WANDB_API_KEY", "").strip()
if not project_id or not api_key:
pytest.fail(
"Weave e2e requires WANDB_API_KEY and WEAVE_PROJECT_ID (or WANDB_PROJECT_ID, "
"format <entity>/<project>): the test reads the proxy's weave_otel delivery "
"back from the real Weave project; missing credentials is a hard failure, not a skip"
)
return WeaveReader(project_id=project_id, api_key=api_key)

View file

@ -35,6 +35,8 @@ class KeyLoggingCallbackVars(BaseModel):
langfuse_public_key: str | None = None
langfuse_secret_key: str | None = None
langfuse_host: str | None = None
wandb_api_key: str | None = None
weave_project_id: str | None = None
class KeyLoggingCallback(BaseModel):

View file

@ -16,7 +16,10 @@ from litellm._logging import session_id_var, trace_id_var
from litellm.constants import SENTRY_DENYLIST, SENTRY_PII_DENYLIST
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging
from litellm.litellm_core_utils.litellm_logging import set_callbacks
from litellm.litellm_core_utils.litellm_logging import (
_get_status_fields,
set_callbacks,
)
from litellm.types.utils import ModelResponse, TextCompletionResponse
@ -6441,3 +6444,16 @@ def test_passthrough_embeddings_result_swapped_for_callbacks():
assert isinstance(swapped_result, EmbeddingResponse)
assert swapped_result.data[0]["embedding"] == [0.1, 0.2, 0.3]
def test_get_status_fields_ranks_guardrail_flagged_between_success_and_intervened():
"""LIT-6894: a non-blocking flagged verdict must outrank success in the
request-level guardrail_status but never mask an intervention."""
flagged = {"guardrail_status": "guardrail_flagged"}
assert _get_status_fields(
"success", [{"guardrail_status": "success"}, flagged], None
)["guardrail_status"] == "guardrail_flagged"
assert _get_status_fields(
"success", [flagged, {"guardrail_status": "guardrail_intervened"}], None
)["guardrail_status"] == "guardrail_intervened"

View file

@ -1,13 +1,22 @@
"""Tests for the MCP guardrail translation handler."""
import asyncio
import pytest
from fastapi import HTTPException
from mcp.types import CallToolResult, ImageContent, TextContent
import litellm
import litellm.llms as litellm_llms
from litellm.caching.caching import DualCache
from litellm.exceptions import BlockedPiiEntityError
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._experimental.mcp_server.guardrail_translation.handler import (
MCPGuardrailTranslationHandler,
)
from litellm.proxy._experimental.mcp_server.utils import MAX_STRUCTURED_CONTENT_SCAN_DEPTH
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
from litellm.types.utils import GenericGuardrailAPIInputs
@ -24,12 +33,11 @@ class MockGuardrail(CustomGuardrail):
self.call_count += 1
self.last_inputs = inputs
self.last_request_data = request_data
return None # Guardrail doesn't modify for MCP tools
@pytest.mark.asyncio
async def test_process_input_messages_updates_content():
"""Handler should pass tool definition to guardrail when mcp_tool_name is present."""
"""Handler should pass the tool definition and the argument strings to the guardrail."""
handler = MCPGuardrailTranslationHandler()
guardrail = MockGuardrail()
@ -45,7 +53,7 @@ async def test_process_input_messages_updates_content():
assert result == data
# Guardrail was called
assert guardrail.call_count == 1
# Guardrail received tools (not texts) with tool definition
# Guardrail received tools with the tool definition
assert guardrail.last_inputs is not None
tools = guardrail.last_inputs.get("tools", [])
assert len(tools) == 1
@ -85,6 +93,412 @@ async def test_process_input_messages_handles_minimal_data():
assert tools[0]["function"]["name"] == "simple_tool"
class ArgumentMaskingGuardrail(CustomGuardrail):
"""Unified guardrail that rewrites every text it is handed, like presidio does."""
def __init__(
self,
secret: str = "jane.doe@example.com",
replacement: str = "<EMAIL_ADDRESS>",
texts_override: list[str] | None = None,
**kwargs,
):
kwargs.setdefault("guardrail_name", "argument-masking-mcp-guardrail")
super().__init__(**kwargs)
self.secret = secret
self.replacement = replacement
self.texts_override = texts_override
self.seen_texts: list[str] | None = None
def _mask(self, text: str) -> str:
return text.replace(self.secret, self.replacement)
async def apply_guardrail(self, inputs, request_data, input_type, **kwargs):
self.seen_texts = list(inputs.get("texts") or [])
if self.texts_override is not None:
inputs["texts"] = self.texts_override
else:
inputs["texts"] = [self._mask(text) for text in self.seen_texts]
return inputs
@pytest.fixture
def restore_callbacks(monkeypatch):
"""Restore the process-wide state driving pre_call_hook through unified_guardrail.
litellm.llms memoizes the guardrail translation mappings in a module global, and
ProxyLogging caches callback capabilities keyed on id()s of litellm.callbacks,
so leaving either populated leaks into unrelated tests in the same worker.
"""
monkeypatch.setattr(litellm, "callbacks", litellm.callbacks)
monkeypatch.setattr(
litellm_llms,
"endpoint_guardrail_translation_mappings",
litellm_llms.endpoint_guardrail_translation_mappings,
)
yield
ProxyLogging._callback_capabilities_cache.clear()
@pytest.mark.asyncio
async def test_argument_strings_are_handed_to_the_guardrail():
"""A guardrail must see the argument values, not just the tool definition.
Without this the guardrail is handed a name and an empty schema, so no
sensitive-data detection can ever fire on an MCP tool call.
"""
handler = MCPGuardrailTranslationHandler()
guardrail = MockGuardrail()
data = {
"mcp_tool_name": "search",
"mcp_arguments": {"query": "contact jane.doe@example.com about the invoice"},
}
await handler.process_input_messages(data, guardrail)
assert guardrail.last_inputs is not None
assert guardrail.last_inputs.get("texts") == ["contact jane.doe@example.com about the invoice"]
@pytest.mark.asyncio
async def test_masked_arguments_are_written_back_for_the_call_path():
"""A mask only takes effect once it lands in modified_arguments."""
handler = MCPGuardrailTranslationHandler()
guardrail = ArgumentMaskingGuardrail()
data = {
"mcp_tool_name": "search",
"mcp_arguments": {"query": "contact jane.doe@example.com about the invoice"},
}
result = await handler.process_input_messages(data, guardrail)
masked = {"query": "contact <EMAIL_ADDRESS> about the invoice"}
assert result["modified_arguments"] == masked
assert result["mcp_arguments"] == masked
@pytest.mark.asyncio
async def test_nested_arguments_keep_their_shape_when_masked():
"""Masking rewrites string leaves in place and preserves non-string values."""
handler = MCPGuardrailTranslationHandler()
guardrail = ArgumentMaskingGuardrail()
arguments = {
"recipients": ["jane.doe@example.com", "ops@example.net"],
"envelope": {"reply_to": "jane.doe@example.com", "retries": 3, "urgent": True, "cc": None},
"count": 2,
}
data = {"mcp_tool_name": "send_email", "mcp_arguments": arguments}
result = await handler.process_input_messages(data, guardrail)
assert guardrail.seen_texts == [
"jane.doe@example.com",
"ops@example.net",
"jane.doe@example.com",
]
assert result["modified_arguments"] == {
"recipients": ["<EMAIL_ADDRESS>", "ops@example.net"],
"envelope": {"reply_to": "<EMAIL_ADDRESS>", "retries": 3, "urgent": True, "cc": None},
"count": 2,
}
@pytest.mark.asyncio
async def test_clean_arguments_are_not_overridden():
"""A guardrail that changes nothing must not set modified_arguments."""
handler = MCPGuardrailTranslationHandler()
guardrail = ArgumentMaskingGuardrail()
data = {"mcp_tool_name": "search", "mcp_arguments": {"query": "quarterly revenue"}}
result = await handler.process_input_messages(data, guardrail)
assert "modified_arguments" not in result
assert result["mcp_arguments"] == {"query": "quarterly revenue"}
@pytest.mark.asyncio
async def test_guardrail_returning_wrong_text_count_blocks_the_call():
"""Write-back is positional, so a length mismatch must block the call."""
handler = MCPGuardrailTranslationHandler()
guardrail = ArgumentMaskingGuardrail(texts_override=["only", "two", "texts"])
arguments = {"query": "contact jane.doe@example.com about the invoice"}
data = {"mcp_tool_name": "search", "mcp_arguments": arguments}
with pytest.raises(HTTPException) as exc_info:
await handler.process_input_messages(data, guardrail)
assert exc_info.value.status_code == 400
assert "modified_arguments" not in data
@pytest.mark.asyncio
async def test_deeply_nested_arguments_are_blocked_rather_than_skipped():
"""Arguments too deep to walk must block instead of passing unscanned."""
handler = MCPGuardrailTranslationHandler()
guardrail = ArgumentMaskingGuardrail()
nested: dict = {"leaf": "jane.doe@example.com"}
for _ in range(MAX_STRUCTURED_CONTENT_SCAN_DEPTH + 1):
nested = {"next": nested}
data = {"mcp_tool_name": "search", "mcp_arguments": nested}
with pytest.raises(HTTPException) as exc_info:
await handler.process_input_messages(data, guardrail)
assert exc_info.value.status_code == 400
class SelfWritingMaskingGuardrail(ArgumentMaskingGuardrail):
"""Masks through ``texts`` and writes the masked arguments itself.
The shape the bundled content filter guardrail already has: it rewrites
``request_data["mcp_arguments"]`` from inside ``apply_guardrail`` as well as
returning masked texts.
"""
async def apply_guardrail(self, inputs, request_data, input_type, **kwargs):
returned = await super().apply_guardrail(inputs, request_data, input_type, **kwargs)
arguments = request_data.get("mcp_arguments") or {}
masked = {key: self._mask(value) if isinstance(value, str) else value for key, value in arguments.items()}
request_data["mcp_arguments"] = masked
request_data["modified_arguments"] = masked
return returned
@pytest.mark.asyncio
async def test_guardrail_that_masks_the_arguments_itself_is_not_treated_as_a_conflict():
"""Converging on the same replacement is not an unmergeable rewrite.
A guardrail that both returns masked texts and rewrites the arguments in
request_data must still mask, not be rejected as if a second guardrail had
clobbered the leaf.
"""
handler = MCPGuardrailTranslationHandler()
guardrail = SelfWritingMaskingGuardrail()
data = {
"mcp_tool_name": "search",
"mcp_arguments": {"query": "contact jane.doe@example.com about the invoice"},
}
result = await handler.process_input_messages(data, guardrail)
assert result["modified_arguments"] == {"query": "contact <EMAIL_ADDRESS> about the invoice"}
class ReshapingGuardrail(ArgumentMaskingGuardrail):
"""Masks through ``texts`` while moving the secret to a different path."""
def __init__(self, reshaped: dict, **kwargs):
super().__init__(**kwargs)
self.reshaped = reshaped
async def apply_guardrail(self, inputs, request_data, input_type, **kwargs):
returned = await super().apply_guardrail(inputs, request_data, input_type, **kwargs)
request_data["mcp_arguments"] = self.reshaped
return returned
@pytest.mark.asyncio
async def test_arguments_reshaped_under_the_guardrail_fail_closed():
"""A payload that no longer lines up leaf for leaf must block, not be written blind.
Write-back pairs masked texts to leaves positionally, so a tree another guardrail
reshaped would take the redaction on the wrong value.
"""
handler = MCPGuardrailTranslationHandler()
guardrail = ReshapingGuardrail({"query": "contact jane.doe@example.com", "note": "added"})
data = {
"mcp_tool_name": "search",
"mcp_arguments": {"query": "contact jane.doe@example.com about the invoice"},
}
with pytest.raises(HTTPException) as exc_info:
await handler.process_input_messages(data, guardrail)
assert exc_info.value.status_code == 400
@pytest.mark.asyncio
async def test_arguments_shortened_under_the_guardrail_fail_closed():
handler = MCPGuardrailTranslationHandler()
guardrail = ReshapingGuardrail({"padding": "jane.doe@example.com"})
data = {
"mcp_tool_name": "search",
"mcp_arguments": {"padding": "x", "secret": "jane.doe@example.com"},
}
with pytest.raises(HTTPException) as exc_info:
await handler.process_input_messages(data, guardrail)
assert exc_info.value.status_code == 400
assert "jane.doe@example.com" not in str(data.get("modified_arguments"))
@pytest.mark.asyncio
async def test_a_renamed_argument_key_blocks_rather_than_dropping_the_mask():
"""The leak this closes: same text, new path, so the write-back would find nothing.
Matching purely on position would see an unchanged value and write the mask to a
path that no longer exists, shipping the secret while reporting a clean scan.
"""
handler = MCPGuardrailTranslationHandler()
guardrail = ReshapingGuardrail({"renamed": "jane.doe@example.com", "other": "kept"})
data = {
"mcp_tool_name": "search",
"mcp_arguments": {"query": "jane.doe@example.com", "other": "kept"},
}
with pytest.raises(HTTPException) as exc_info:
await handler.process_input_messages(data, guardrail)
assert exc_info.value.status_code == 400
assert "jane.doe@example.com" not in str(data.get("modified_arguments"))
@pytest.mark.parametrize("run_in_parallel", [False, True])
@pytest.mark.asyncio
async def test_masked_arguments_reach_the_outbound_mcp_call(restore_callbacks, monkeypatch, run_in_parallel):
"""End to end over the real MCP pre-call path, not just the handler.
Drives the same sequence mcp_server_manager.call_tool uses:
synthetic payload -> pre_call_hook -> arguments sent upstream.
Covers run_in_parallel both ways: that path shares one payload snapshot and
discards whatever a guardrail returns, so the mask has to land on the caller's
dict rather than on a copy of it.
"""
guardrail = ArgumentMaskingGuardrail(
event_hook="pre_mcp_call",
default_on=True,
run_in_parallel=run_in_parallel,
)
monkeypatch.setattr(litellm, "callbacks", [guardrail])
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
arguments = {"query": "contact jane.doe@example.com about the invoice"}
pre_hook_kwargs = {
"name": "search",
"arguments": arguments,
"server_name": "test-server",
"user_api_key_auth": UserAPIKeyAuth(api_key="sk-test", user_id="test-user"),
}
request_obj = proxy_logging_obj._create_mcp_request_object_from_kwargs(pre_hook_kwargs)
synthetic_data = proxy_logging_obj._convert_mcp_to_llm_format(request_obj, pre_hook_kwargs)
modified_data = await proxy_logging_obj.pre_call_hook(
user_api_key_dict=pre_hook_kwargs["user_api_key_auth"],
data=synthetic_data,
call_type="call_mcp_tool",
)
modified_kwargs = proxy_logging_obj._convert_mcp_hook_response_to_kwargs(modified_data, pre_hook_kwargs)
assert modified_kwargs["arguments"] == {"query": "contact <EMAIL_ADDRESS> about the invoice"}
class SlowSubstitutionGuardrail(CustomGuardrail):
"""Rewrites one substring, after a delay, so two instances genuinely interleave."""
def __init__(self, needle: str, replacement: str, delay: float, **kwargs):
super().__init__(**kwargs)
self.needle = needle
self.replacement = replacement
self.delay = delay
async def apply_guardrail(self, inputs, request_data, input_type, **kwargs):
await asyncio.sleep(self.delay)
inputs["texts"] = [text.replace(self.needle, self.replacement) for text in (inputs.get("texts") or [])]
return inputs
def _two_interleaving_maskers(run_in_parallel: bool):
return [
SlowSubstitutionGuardrail(
"jane.doe@example.com",
"<EMAIL_ADDRESS>",
0.02,
guardrail_name="mask-email",
event_hook="pre_mcp_call",
default_on=True,
run_in_parallel=run_in_parallel,
),
SlowSubstitutionGuardrail(
"415-555-0132",
"<PHONE_NUMBER>",
0.04,
guardrail_name="mask-phone",
event_hook="pre_mcp_call",
default_on=True,
run_in_parallel=run_in_parallel,
),
]
async def _arguments_sent_upstream(arguments: dict):
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
pre_hook_kwargs = {
"name": "search",
"arguments": arguments,
"server_name": "test-server",
"user_api_key_auth": UserAPIKeyAuth(api_key="sk-test", user_id="test-user"),
}
request_obj = proxy_logging_obj._create_mcp_request_object_from_kwargs(pre_hook_kwargs)
modified_data = await proxy_logging_obj.pre_call_hook(
user_api_key_dict=pre_hook_kwargs["user_api_key_auth"],
data=proxy_logging_obj._convert_mcp_to_llm_format(request_obj, pre_hook_kwargs),
call_type="call_mcp_tool",
)
return proxy_logging_obj._convert_mcp_hook_response_to_kwargs(modified_data, pre_hook_kwargs)["arguments"]
@pytest.mark.asyncio
async def test_two_sequential_guardrails_both_masks_survive(restore_callbacks, monkeypatch):
"""The recommended config: each guardrail sees the previous one's output."""
monkeypatch.setattr(litellm, "callbacks", _two_interleaving_maskers(run_in_parallel=False))
sent = await _arguments_sent_upstream({"note": "mail jane.doe@example.com or call 415-555-0132"})
assert sent == {"note": "mail <EMAIL_ADDRESS> or call <PHONE_NUMBER>"}
@pytest.mark.asyncio
async def test_two_parallel_guardrails_on_separate_arguments_both_masks_survive(restore_callbacks, monkeypatch):
"""Concurrent rewrites of different leaves compose; neither is lost."""
monkeypatch.setattr(litellm, "callbacks", _two_interleaving_maskers(run_in_parallel=True))
sent = await _arguments_sent_upstream({"email": "jane.doe@example.com", "phone": "415-555-0132"})
assert sent == {"email": "<EMAIL_ADDRESS>", "phone": "<PHONE_NUMBER>"}
@pytest.mark.asyncio
async def test_two_parallel_guardrails_on_one_argument_block_instead_of_losing_a_mask(restore_callbacks, monkeypatch):
"""Unmergeable concurrent rewrites must fail closed, not ship one redaction.
Both guardrails derive a full replacement string from the same snapshot, so
writing either result would silently discard the other's redaction and leak
the value it was configured to mask.
"""
monkeypatch.setattr(litellm, "callbacks", _two_interleaving_maskers(run_in_parallel=True))
original = "mail jane.doe@example.com or call 415-555-0132"
with pytest.raises(HTTPException) as exc_info:
await _arguments_sent_upstream({"note": original})
assert exc_info.value.status_code == 400
assert "note" in str(exc_info.value.detail)
class MaskingGuardrail(CustomGuardrail):
"""Guardrail that rewrites every scanned text, recording what it saw."""
@ -190,8 +604,8 @@ async def test_process_output_response_handles_result_without_content():
@pytest.mark.asyncio
async def test_process_output_response_leaves_result_unmasked_on_text_count_mismatch():
"""A guardrail returning the wrong number of texts must not shuffle content."""
async def test_process_output_response_blocks_on_text_count_mismatch():
"""A guardrail returning the wrong number of texts must block the result."""
handler = MCPGuardrailTranslationHandler()
guardrail = MaskingGuardrail(masked_texts=["<EMAIL_ADDRESS>"])
result = CallToolResult(
@ -202,9 +616,10 @@ async def test_process_output_response_leaves_result_unmasked_on_text_count_mism
isError=False,
)
returned = await handler.process_output_response(response=result, guardrail_to_apply=guardrail)
with pytest.raises(HTTPException) as exc_info:
await handler.process_output_response(response=result, guardrail_to_apply=guardrail)
assert [item.text for item in returned.content] == ["jane@example.com", "415-555-0132"]
assert exc_info.value.status_code == 400
class SubstitutingGuardrail(CustomGuardrail):

View file

@ -197,6 +197,71 @@ async def test_custom_code_post_call_block_raises_http_400():
}
FLAG_CODE = (
"def apply_guardrail(inputs, request_data, input_type):\n"
' return flag("audit hit", metadata={"category": "topic"})\n'
)
@pytest.mark.asyncio
@pytest.mark.parametrize("input_type", ["request", "response"])
async def test_custom_code_flag_passes_content_through_and_records_flagged_entry(input_type):
"""LIT-6894: flag() must not raise, must return the content unchanged and must log
exactly one guardrail_flagged entry (the decorator must not add a second "success")."""
guardrail = CustomCodeGuardrail(custom_code=FLAG_CODE, guardrail_name="t", event_hook=["pre_call", "post_call"])
request_data = {"model": "test-model", "litellm_metadata": {}}
result = await guardrail.apply_guardrail(
inputs={"texts": ["hello"]},
request_data=request_data,
input_type=input_type,
)
assert result == {"texts": ["hello"]}
entries = request_data["litellm_metadata"]["standard_logging_guardrail_information"]
assert len(entries) == 1
entry = entries[0]
assert entry["guardrail_status"] == "guardrail_flagged"
assert entry["guardrail_name"] == "t"
assert entry["guardrail_mode"] == ["pre_call", "post_call"]
assert entry["guardrail_response"] == {
"action": "flag",
"reason": "audit hit",
"input_type": input_type,
"metadata": {"category": "topic"},
}
assert entry["duration"] is not None and entry["duration"] >= 0
@pytest.mark.asyncio
async def test_custom_code_flag_default_reason_and_empty_metadata():
code = "def apply_guardrail(inputs, request_data, input_type):\n return flag('just a note')\n"
guardrail = _compile(code)
request_data = {"model": "m", "litellm_metadata": {}}
await guardrail.apply_guardrail(inputs={"texts": ["x"]}, request_data=request_data, input_type="request")
entry = request_data["litellm_metadata"]["standard_logging_guardrail_information"][0]
assert entry["guardrail_response"] == {
"action": "flag",
"reason": "just a note",
"input_type": "request",
"metadata": {},
}
@pytest.mark.asyncio
async def test_custom_code_allow_still_records_success_not_flagged():
code = "def apply_guardrail(inputs, request_data, input_type):\n return allow()\n"
guardrail = _compile(code)
request_data = {"model": "m", "litellm_metadata": {}}
await guardrail.apply_guardrail(inputs={"texts": ["x"]}, request_data=request_data, input_type="request")
entries = request_data["litellm_metadata"]["standard_logging_guardrail_information"]
assert [e["guardrail_status"] for e in entries] == ["success"]
def test_typical_sync_guardrail_still_works():
code = (
"def apply_guardrail(inputs, request_data, input_type):\n"

View file

@ -477,6 +477,103 @@ async def test_logs_resolves_config_guardrail_logical_name():
assert where["guardrail_id"] == {"in": ["yaml-uuid", "yaml-pii"]}
def _index_row(request_id: str, guardrail_id: str = "cc-flag") -> Any:
r = MagicMock(spec=["request_id", "guardrail_id", "policy_id", "start_time"])
r.request_id = request_id
r.guardrail_id = guardrail_id
return r
def _spend_log(request_id: str, *guardrail_statuses: str, guardrail_id: str = "cc-flag") -> Any:
sl = MagicMock(spec=["request_id", "metadata", "startTime", "model", "messages", "response"])
sl.request_id = request_id
sl.startTime = datetime(2026, 4, 25, 12, 0)
sl.model = "gpt-4o-mini"
sl.messages = [{"role": "user", "content": "hi"}]
sl.response = "ok"
sl.metadata = {
"guardrail_information": [
{
"guardrail_name": guardrail_id,
"guardrail_status": status,
"guardrail_response": (
{"action": "flag", "reason": "audit hit"} if status == "guardrail_flagged" else "allow"
),
"duration": 0.002,
}
for status in guardrail_statuses
]
}
return sl
@pytest.mark.asyncio
async def test_logs_reports_flagged_action_for_guardrail_flagged_status():
"""LIT-6894: Request Logs surface a custom code flag() verdict as flagged with its reason."""
prisma = _prisma(index_find_many=[_index_row("r-flag"), _index_row("r-pass"), _index_row("r-block")])
prisma.db.litellm_spendlogs.find_many = AsyncMock(
return_value=[
_spend_log("r-flag", "guardrail_flagged"),
_spend_log("r-pass", "success"),
_spend_log("r-block", "guardrail_intervened"),
]
)
p1, p2 = _patches(prisma, _config_handler())
with p1, p2:
resp = await guardrails_usage_logs(
guardrail_id="cc-flag",
policy_id=None,
page=1,
page_size=50,
action=None,
start_date=START,
end_date=END,
user_api_key_dict=ADMIN,
)
flagged_only = await guardrails_usage_logs(
guardrail_id="cc-flag",
policy_id=None,
page=1,
page_size=50,
action="flagged",
start_date=START,
end_date=END,
user_api_key_dict=ADMIN,
)
assert [(log.id, log.action) for log in resp.logs] == [
("r-flag", "flagged"),
("r-pass", "passed"),
("r-block", "blocked"),
]
assert resp.logs[0].reason == "{'action': 'flag', 'reason': 'audit hit'}"
assert [log.id for log in flagged_only.logs] == ["r-flag"]
@pytest.mark.asyncio
async def test_logs_reports_post_call_flag_when_pre_call_allowed():
"""LIT-6894: a guardrail on mode [pre_call, post_call] that allows the request but flags the response
shows as flagged, not hidden behind the pre_call allow entry."""
prisma = _prisma(index_find_many=[_index_row("r-post-flag")])
prisma.db.litellm_spendlogs.find_many = AsyncMock(
return_value=[_spend_log("r-post-flag", "success", "guardrail_flagged")]
)
p1, p2 = _patches(prisma, _config_handler())
with p1, p2:
resp = await guardrails_usage_logs(
guardrail_id="cc-flag",
policy_id=None,
page=1,
page_size=50,
action=None,
start_date=START,
end_date=END,
user_api_key_dict=ADMIN,
)
assert [(log.id, log.action, log.reason) for log in resp.logs] == [
("r-post-flag", "flagged", "{'action': 'flag', 'reason': 'audit hit'}")
]
# ---- date window cap (LIT-5762) ---------------------------------------------

View file

@ -105,6 +105,27 @@ async def test_usage_units_rolled_up_by_guardrail_team_key_and_date():
}
@pytest.mark.asyncio
async def test_flagged_status_counts_as_flagged_not_passed_or_blocked():
"""LIT-6894: a custom code flag() verdict lands in flagged_count on the Monitor rollup."""
prisma = _prisma()
logs = [
_payload("r1", guardrail_status="success"),
_payload("r2", guardrail_status="guardrail_flagged"),
_payload("r3", guardrail_status="guardrail_intervened"),
]
await process_spend_logs_guardrail_usage(prisma, logs)
create = prisma.db.litellm_dailyguardrailmetrics.upsert.call_args.kwargs["data"]["create"]
assert (create["requests_evaluated"], create["passed_count"], create["flagged_count"], create["blocked_count"]) == (
3,
1,
1,
1,
)
def _fake_sleep() -> tuple[AsyncMock, list[float]]:
delays: list[float] = []
sleep = AsyncMock(side_effect=lambda delay: delays.append(delay))

View file

@ -3917,6 +3917,7 @@ def _object_permission_mocks(mocker, existing_object_permission_id=None):
mock_prisma_client.db.litellm_objectpermissiontable.upsert = mocker.AsyncMock(
return_value=SimpleNamespace(object_permission_id="perm-new")
)
mock_prisma_client.db.litellm_mcpservertable.find_many = mocker.AsyncMock(return_value=[])
mock_prisma_client.update_data = mocker.AsyncMock(
return_value={"user_id": "target-user"}
)
@ -4146,6 +4147,7 @@ async def test_new_user_persists_the_requested_mcp_entitlement(mocker):
mock_prisma_client.db.litellm_objectpermissiontable.create = mocker.AsyncMock(
return_value=SimpleNamespace(object_permission_id="perm-created")
)
mock_prisma_client.db.litellm_mcpservertable.find_many = mocker.AsyncMock(return_value=[])
mock_prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(
return_value=None
)

View file

@ -963,6 +963,7 @@ async def test_key_generation_with_mcp_tool_permissions(monkeypatch):
mock_prisma_client.db = MagicMock()
mock_prisma_client.db.litellm_objectpermissiontable = MagicMock()
mock_prisma_client.db.litellm_objectpermissiontable.create = mock_create
mock_prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
async def _insert_data_side_effect(*args, **kwargs):
table_name = kwargs.get("table_name")

View file

@ -1183,6 +1183,37 @@ async def test_find_member_if_email_missing_row_raises_documented_400():
}
@pytest.mark.asyncio
async def test_new_organization_rejects_shared_alias_tool_permission_key():
"""/organization/new creates its permission row through its own helper, so the
ambiguous mcp_tool_permissions key check (LIT-4982) has to run there too."""
from litellm.proxy._types import LiteLLM_ObjectPermissionBase, NewOrganizationRequest
from litellm.proxy.management_endpoints.organization_endpoints import (
_set_object_permission,
)
prisma_client = MagicMock()
prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(
return_value=[
MagicMock(server_id="wiki-a-id", alias="wiki", server_name="wiki_a"),
MagicMock(server_id="wiki-b-id", alias="wiki", server_name="wiki_b"),
]
)
prisma_client.db.litellm_objectpermissiontable.create = AsyncMock()
data = NewOrganizationRequest(
organization_alias="org",
object_permission=LiteLLM_ObjectPermissionBase(mcp_tool_permissions={"wiki": ["ask_question"]}),
)
with pytest.raises(HTTPException) as exc_info:
await _set_object_permission(data=data, prisma_client=prisma_client)
assert exc_info.value.status_code == 400
assert "wiki-a-id" in str(exc_info.value.detail)
assert "wiki-b-id" in str(exc_info.value.detail)
prisma_client.db.litellm_objectpermissiontable.create.assert_not_called()
def test_v2_update_organization_is_in_openapi_schema():
"""PATCH /v2/organization/{organization_id} is documented in the generated OpenAPI spec."""
from fastapi import FastAPI

View file

@ -651,6 +651,7 @@ async def test_new_team_with_mcp_tool_permissions(mock_db_client, mock_admin_aut
mock_db_client.db.litellm_objectpermissiontable = MagicMock()
mock_db_client.db.litellm_objectpermissiontable.create = mock_obj_perm_create
mock_db_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
# Mock model table
mock_db_client.db.litellm_modeltable = MagicMock()

View file

@ -17,6 +17,7 @@ from litellm.proxy.management_helpers.object_permission_utils import (
_resolve_team_allowed_mcp_servers,
_set_object_permission,
enforce_all_proxy_mcp_servers_grant_is_admin_only,
prepare_object_permission_upsert,
validate_key_mcp_servers_against_team,
validate_key_search_tools_against_team,
validate_key_vector_stores_against_team,
@ -41,6 +42,7 @@ async def test_set_object_permission():
mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock(
return_value=mock_created_permission
)
mock_prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
# Test data with object_permission
data_json = {
@ -1349,6 +1351,123 @@ async def test_validate_key_update_sentinels_do_not_grandfather(monkeypatch):
assert exc_info.value.status_code == 403
# ---- Tests for rejecting ambiguous mcp_tool_permissions keys on write (LIT-4982) ----
_SHARED_ALIAS_DB_SERVERS = (
_make_mock_mcp_server("wiki-a-id", alias="wiki", server_name="wiki_a"),
_make_mock_mcp_server("wiki-b-id", alias="wiki", server_name="wiki_b"),
_make_mock_mcp_server("gh-a-id", alias="gh_a", server_name="github"),
_make_mock_mcp_server("gh-b-id", alias="gh_b", server_name="github"),
_make_mock_mcp_server("solo-id", alias="solo", server_name="Solo Server"),
_make_mock_mcp_server("shadow-id", alias="solo-id", server_name="shadow"),
)
def _make_ambiguity_prisma(existing_tool_permissions=None):
"""Mock prisma client whose MCP server table holds _SHARED_ALIAS_DB_SERVERS and whose
object permission row (if any) stores the given mcp_tool_permissions JSON string."""
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=list(_SHARED_ALIAS_DB_SERVERS))
mock_prisma.db.litellm_objectpermissiontable.create = AsyncMock(
return_value=MagicMock(object_permission_id="perm-id")
)
existing_row = None
if existing_tool_permissions is not None:
existing_row = MagicMock()
existing_row.model_dump.return_value = {
"object_permission_id": "perm-id",
"mcp_tool_permissions": json.dumps(existing_tool_permissions),
}
mock_prisma.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=existing_row)
return mock_prisma
@pytest.mark.asyncio
@pytest.mark.parametrize(
"identifier, colliding_ids",
[("wiki", ("wiki-a-id", "wiki-b-id")), ("github", ("gh-a-id", "gh-b-id"))],
)
async def test_set_object_permission_rejects_shared_alias_or_name_tool_permission_key(identifier, colliding_ids):
"""An alias or server_name two servers share cannot key mcp_tool_permissions on
create: the write is rejected with 400 naming both servers and nothing is persisted."""
mock_prisma = _make_ambiguity_prisma()
data_json = {"object_permission": {"mcp_tool_permissions": {identifier: ["read_wiki_structure"]}}}
with pytest.raises(HTTPException) as exc_info:
await _set_object_permission(data_json=data_json, prisma_client=mock_prisma)
assert exc_info.value.status_code == 400
assert all(server_id in str(exc_info.value.detail) for server_id in colliding_ids)
mock_prisma.db.litellm_objectpermissiontable.create.assert_not_called()
@pytest.mark.asyncio
async def test_prepare_object_permission_upsert_rejects_shared_alias_tool_permission_key():
"""The update seam shared by key/team/org/user/customer/agent rejects a new
shared-alias key when the existing row does not already hold it."""
mock_prisma = _make_ambiguity_prisma(existing_tool_permissions={"solo-id": ["tool1"]})
with pytest.raises(HTTPException) as exc_info:
await prepare_object_permission_upsert(
new_object_permission={"mcp_tool_permissions": {"wiki": ["ask_question"]}},
existing_object_permission_id="perm-id",
prisma_client=mock_prisma,
)
assert exc_info.value.status_code == 400
assert "'wiki'" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_unambiguous_tool_permission_keys_persist_verbatim():
"""Exact ids (even when another server uses that id string as its alias),
unique aliases, and an id plus alias pointing at one server all still write."""
mock_prisma = _make_ambiguity_prisma()
tool_permissions = {
"wiki-a-id": ["ask_question"],
"wiki-b-id": ["read_wiki_structure"],
"solo-id": ["tool1"],
"solo": ["tool2"],
"Solo Server": ["tool3"],
}
upsert = await prepare_object_permission_upsert(
new_object_permission={"mcp_tool_permissions": dict(tool_permissions)},
existing_object_permission_id=None,
prisma_client=mock_prisma,
)
assert json.loads(upsert.record["mcp_tool_permissions"]) == tool_permissions
@pytest.mark.asyncio
async def test_stored_ambiguous_tool_permission_key_is_grandfathered_until_changed():
"""A shared-alias entry already on the row may be re-sent unchanged so unrelated
edits succeed, but changing its tool list is rejected."""
mock_prisma = _make_ambiguity_prisma(existing_tool_permissions={"wiki": ["read_wiki_structure"]})
upsert = await prepare_object_permission_upsert(
new_object_permission={
"mcp_tool_permissions": {"wiki": ["read_wiki_structure"], "solo-id": ["tool1"]},
},
existing_object_permission_id="perm-id",
prisma_client=mock_prisma,
)
assert json.loads(upsert.record["mcp_tool_permissions"]) == {
"wiki": ["read_wiki_structure"],
"solo-id": ["tool1"],
}
with pytest.raises(HTTPException) as exc_info:
await prepare_object_permission_upsert(
new_object_permission={"mcp_tool_permissions": {"wiki": ["ask_question"]}},
existing_object_permission_id="perm-id",
prisma_client=mock_prisma,
)
assert exc_info.value.status_code == 400
def test_object_permission_dict_mirrors_pydantic_model():
"""ObjectPermissionDict must stay field-for-field aligned with
LiteLLM_ObjectPermissionBase. If a new field is added to the Pydantic

View file

@ -131,3 +131,43 @@ class TestMaybeSetupPrometheusMultiprocDir:
# Cleanup
os.environ.pop("PROMETHEUS_MULTIPROC_DIR", None)
@pytest.mark.parametrize(
"litellm_settings",
[
{"callbacks": ["prometheus"]},
{"callbacks": ["langfuse"]},
None,
],
)
def test_separate_metrics_port_forces_dir_for_single_worker(self, litellm_settings):
"""The separate metrics process reads the samples, so one worker still needs the shared dir, even when
prometheus is not in config.yaml (callbacks can be turned on from the DB after startup)."""
with patch.dict(os.environ, {}, clear=False):
os.environ.pop("PROMETHEUS_MULTIPROC_DIR", None)
os.environ.pop("prometheus_multiproc_dir", None)
result_dir = ProxyInitializationHelpers._maybe_setup_prometheus_multiproc_dir(
num_workers=1,
litellm_settings=litellm_settings,
prometheus_metrics_port=4001,
)
assert result_dir is not None
assert os.environ.get("PROMETHEUS_MULTIPROC_DIR") == result_dir
assert os.path.isdir(result_dir)
os.environ.pop("PROMETHEUS_MULTIPROC_DIR", None)
def test_lowercase_env_var_is_reused_and_exported_uppercase(self, tmp_path):
"""prometheus_client honours both spellings; the metrics server only reads the uppercase one."""
with patch.dict(os.environ, {"prometheus_multiproc_dir": str(tmp_path)}, clear=False):
os.environ.pop("PROMETHEUS_MULTIPROC_DIR", None)
result_dir = ProxyInitializationHelpers._maybe_setup_prometheus_multiproc_dir(
num_workers=4,
litellm_settings={"callbacks": "prometheus"},
)
assert result_dir == str(tmp_path)
assert os.environ["PROMETHEUS_MULTIPROC_DIR"] == str(tmp_path)

View file

@ -0,0 +1,259 @@
"""The separate metrics server must aggregate PROMETHEUS_MULTIPROC_DIR, expose only /metrics, and follow its
parent's lifetime.
Everything here runs on loopback against a child of this test process; no LLM keys or external network.
"""
from __future__ import annotations
import os
import socket
import subprocess
import sys
import threading
import time
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from typing import Final
from unittest.mock import patch
import httpx
import pytest
from fastapi.testclient import TestClient
from prometheus_client import values
from litellm.proxy.prometheus_metrics_server import (
PID_HEADER,
MetricsServerStartupError,
build_metrics_app,
main,
metrics_url,
start_metrics_server_process,
)
_STARTUP_TIMEOUT_SECONDS: Final = 60.0
_SHUTDOWN_TIMEOUT_SECONDS: Final = 15.0
def _write_worker_sample(pid: int, value: float) -> None:
"""Write one counter sample into PROMETHEUS_MULTIPROC_DIR the way a proxy worker would."""
counter: Final = values.MultiProcessValue(process_identifier=lambda: pid)(
"counter",
"litellm_requests_metric_total",
"litellm_requests_metric_total",
("model",),
("gpt-5",),
"Total number of LLM calls",
)
counter.inc(value)
def _free_port() -> int:
with socket.socket() as sock:
sock.bind(("127.0.0.1", 0))
return sock.getsockname()[1]
def _wait_for_metrics(port: int, pid: int) -> httpx.Response:
deadline: Final = time.monotonic() + _STARTUP_TIMEOUT_SECONDS
while time.monotonic() < deadline:
try:
response: Final = httpx.get(f"http://127.0.0.1:{port}/metrics", follow_redirects=True, timeout=1.0)
if response.status_code == 200 and response.headers.get(PID_HEADER) == str(pid):
return response
except httpx.TransportError:
pass
time.sleep(0.2)
raise AssertionError(f"metrics server on port {port} never served metrics")
def _wait_until_down(port: int) -> None:
deadline: Final = time.monotonic() + _SHUTDOWN_TIMEOUT_SECONDS
while time.monotonic() < deadline:
try:
httpx.get(f"http://127.0.0.1:{port}/metrics", timeout=1.0)
except httpx.TransportError:
return
time.sleep(0.2)
raise AssertionError(f"metrics server on port {port} kept running after its parent died")
def test_metrics_app_aggregates_multiproc_dir_and_reports_pid(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("PROMETHEUS_MULTIPROC_DIR", str(tmp_path))
_write_worker_sample(pid=1001, value=2)
_write_worker_sample(pid=1002, value=3)
other_dir: Final = tmp_path / "other"
other_dir.mkdir()
client: Final = TestClient(build_metrics_app(str(tmp_path)))
metrics: Final = client.get("/metrics")
assert metrics.status_code == 200
assert metrics.headers[PID_HEADER] == str(os.getpid())
assert 'litellm_requests_metric_total{model="gpt-5"} 5.0' in metrics.text
assert client.get("/health").status_code == 404
empty: Final = TestClient(build_metrics_app(str(other_dir))).get("/metrics")
assert empty.status_code == 200
assert "litellm_requests_metric_total" not in empty.text
@pytest.mark.parametrize(
("host", "expected"),
(
("0.0.0.0", "http://127.0.0.1:4001/metrics"),
("::", "http://[::1]:4001/metrics"),
("10.1.2.3", "http://10.1.2.3:4001/metrics"),
("metrics.internal", "http://metrics.internal:4001/metrics"),
),
)
def test_metrics_url_probes_loopback_for_wildcard_binds(host: str, expected: str):
assert metrics_url(host, 4001) == expected
def test_main_serves_the_app_for_the_given_dir_with_uvicorn(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("PROMETHEUS_MULTIPROC_DIR", str(tmp_path))
_write_worker_sample(pid=2001, value=6)
with patch("uvicorn.run") as run:
main(["--host", "10.1.2.3", "--port", "4001", "--multiproc_dir", str(tmp_path)])
run.assert_called_once()
assert run.call_args.kwargs["host"] == "10.1.2.3"
assert run.call_args.kwargs["port"] == 4001
client: Final = TestClient(run.call_args.args[0])
metrics: Final = client.get("/metrics")
assert metrics.status_code == 200
assert metrics.headers[PID_HEADER] == str(os.getpid())
assert 'litellm_requests_metric_total{model="gpt-5"} 6.0' in client.get("/metrics").text
def test_main_falls_back_to_env_multiproc_dir(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("PROMETHEUS_MULTIPROC_DIR", str(tmp_path))
with patch("uvicorn.run") as run:
main(["--port", "4001"])
(app,), served_on = run.call_args
assert served_on["host"] == "0.0.0.0"
metrics: Final = TestClient(app).get("/metrics")
assert metrics.status_code == 200
assert metrics.headers[PID_HEADER] == str(os.getpid())
def test_main_rejects_missing_multiproc_dir(monkeypatch: pytest.MonkeyPatch):
monkeypatch.delenv("PROMETHEUS_MULTIPROC_DIR", raising=False)
with patch("uvicorn.run") as run, pytest.raises(SystemExit) as exit_info:
main(["--port", "4001"])
assert exit_info.value.code == 2
run.assert_not_called()
def test_start_metrics_server_process_returns_only_once_child_serves(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("PROMETHEUS_MULTIPROC_DIR", str(tmp_path))
_write_worker_sample(pid=3001, value=4)
port: Final = _free_port()
with patch("atexit.register") as register:
process: Final = start_metrics_server_process(host="127.0.0.1", port=port, multiproc_dir=str(tmp_path))
try:
register.assert_called_once_with(process.terminate)
assert process.poll() is None
startup_metrics: Final = httpx.get(f"http://127.0.0.1:{port}/metrics", follow_redirects=True, timeout=5.0)
assert startup_metrics.status_code == 200
assert startup_metrics.headers[PID_HEADER] == str(process.pid)
metrics: Final = httpx.get(f"http://127.0.0.1:{port}/metrics", follow_redirects=True, timeout=10.0)
assert 'litellm_requests_metric_total{model="gpt-5"} 4.0' in metrics.text
finally:
process.kill()
process.wait(timeout=10)
def test_start_metrics_server_process_fails_when_port_is_taken(tmp_path: Path):
with socket.socket() as occupied:
occupied.bind(("127.0.0.1", 0))
occupied.listen()
port: Final = occupied.getsockname()[1]
with (
patch("atexit.register"),
pytest.raises(
MetricsServerStartupError, match=rf"exited with code [1-9]\d* before serving 127.0.0.1:{port}"
),
):
start_metrics_server_process(host="127.0.0.1", port=port, multiproc_dir=str(tmp_path))
class _ImpostorMetrics(BaseHTTPRequestHandler):
"""An unrelated service already on the port that answers /metrics with 200 and plausible metrics."""
def do_GET(self) -> None:
body: Final = b"# HELP impostor_metric A plausible metric\n# TYPE impostor_metric counter\nimpostor_metric 1\n"
self.send_response(200)
self.send_header("Content-Type", "text/plain")
self.end_headers()
self.wfile.write(body)
def log_message(self, format: str, *args: object) -> None:
return
def test_start_metrics_server_process_rejects_metrics_from_another_service_on_the_port(tmp_path: Path):
with ThreadingHTTPServer(("127.0.0.1", 0), _ImpostorMetrics) as impostor:
threading.Thread(target=impostor.serve_forever, daemon=True).start()
port: Final = impostor.server_address[1]
impostor_response: Final = httpx.get(f"http://127.0.0.1:{port}/metrics")
assert impostor_response.status_code == 200
assert "# HELP impostor_metric" in impostor_response.text
assert PID_HEADER not in impostor_response.headers
with (
patch("atexit.register"),
pytest.raises(MetricsServerStartupError, match=rf"exited with code [1-9]\d* before serving 127.0.0.1:{port}"),
):
start_metrics_server_process(host="127.0.0.1", port=port, multiproc_dir=str(tmp_path))
impostor.shutdown()
def test_metrics_server_process_serves_and_exits_with_parent(tmp_path: Path, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("PROMETHEUS_MULTIPROC_DIR", str(tmp_path))
_write_worker_sample(pid=2001, value=7)
port: Final = _free_port()
server_argv: Final = (
sys.executable,
"-m",
"litellm.proxy.prometheus_metrics_server",
"--host",
"127.0.0.1",
"--port",
str(port),
"--multiproc_dir",
str(tmp_path),
)
parent: Final = subprocess.Popen(
(
sys.executable,
"-c",
"import subprocess, sys, time; p = subprocess.Popen(sys.argv[1:]); print(p.pid, flush=True); time.sleep(600)",
*server_argv,
),
stdout=subprocess.PIPE,
text=True,
)
assert parent.stdout is not None
server_pid: Final = int(parent.stdout.readline())
try:
metrics: Final = _wait_for_metrics(port, server_pid)
assert metrics.status_code == 200
assert metrics.headers[PID_HEADER] == str(server_pid)
scrape: Final = httpx.get(f"http://127.0.0.1:{port}/metrics", follow_redirects=True, timeout=10.0)
assert scrape.status_code == 200
assert scrape.headers[PID_HEADER] == str(server_pid)
assert 'litellm_requests_metric_total{model="gpt-5"} 7.0' in scrape.text
parent.kill()
parent.wait(timeout=10)
_wait_until_down(port)
finally:
parent.kill()
try:
os.kill(server_pid, 9)
except ProcessLookupError:
pass

View file

@ -662,6 +662,124 @@ class TestProxyInitializationHelpers:
assert "Invalid value for '--limit_concurrency'" in result.output
mock_uvicorn_run.assert_not_called()
@patch("uvicorn.run")
@patch("httpx.HTTPTransport.handle_request")
@patch("atexit.register")
@patch("subprocess.Popen")
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") # test-quality-ok: run_server always wires the DB; same isolation as the sibling CLI tests above
@patch( # test-quality-ok: run_server always wires the DB; same isolation as the sibling CLI tests above
"litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False
)
def test_prometheus_metrics_port_starts_separate_metrics_process(
self,
mock_should_update,
mock_setup_db,
mock_popen,
mock_atexit_register,
mock_handle_request,
mock_uvicorn_run,
tmp_path,
):
"""--prometheus_metrics_port must spawn `python -m litellm.proxy.prometheus_metrics_server` on --host
with the shared multiproc dir, wait for its /metrics response, and only then start uvicorn. It must stay off by
default, refuse to share --port, and abort the proxy when the child dies before serving."""
import httpx
from click.testing import CliRunner
from litellm.proxy.proxy_cli import run_server
runner = CliRunner()
mock_popen.return_value = MagicMock(pid=4242, **{"poll.return_value": None})
probed_urls: list[str] = []
def child_metrics(request: httpx.Request) -> httpx.Response:
probed_urls.append(str(request.url))
return httpx.Response(200, headers={"x-litellm-metrics-pid": "4242"}, content=b"")
mock_handle_request.side_effect = child_metrics
mock_proxy_module = MagicMock(
app=MagicMock(),
ProxyConfig=MagicMock(),
KeyManagementSettings=MagicMock(),
save_worker_config=MagicMock(),
)
clean_env = {
k: v
for k, v in os.environ.items()
if k not in ("DATABASE_URL", "DIRECT_URL", "PROMETHEUS_METRICS_PORT")
}
clean_env["PROMETHEUS_MULTIPROC_DIR"] = str(tmp_path)
with (
patch.dict(os.environ, clean_env, clear=True),
patch.dict(
"sys.modules",
{
"proxy_server": mock_proxy_module,
"litellm.proxy.proxy_server": mock_proxy_module,
},
),
patch( # test-quality-ok: same isolation as the sibling CLI tests above
"litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args"
) as mock_get_args,
):
mock_get_args.side_effect = lambda *a, **k: {
"app": "litellm.proxy.proxy_server:app",
"host": "localhost",
"port": 8000,
}
result = runner.invoke(
run_server,
["--local", "--host", "127.0.0.1", "--port", "4000", "--prometheus_metrics_port", "4001"],
)
assert (
result.exit_code == 0
), f"exit_code={result.exit_code}, output={result.output}"
mock_popen.assert_called_once()
spawned = list(mock_popen.call_args.args[0])
assert spawned[1:3] == ["-m", "litellm.proxy.prometheus_metrics_server"]
assert spawned[3:] == ["--host", "127.0.0.1", "--port", "4001", "--multiproc_dir", str(tmp_path)]
assert probed_urls == ["http://127.0.0.1:4001/metrics"]
assert "Serving Prometheus metrics on 127.0.0.1:4001/metrics (pid 4242)" in result.output
mock_uvicorn_run.assert_called_once()
mock_popen.reset_mock()
mock_uvicorn_run.reset_mock()
mock_popen.return_value = MagicMock(pid=4243, **{"poll.return_value": 1})
result = runner.invoke(
run_server,
["--local", "--port", "4000", "--prometheus_metrics_port", "4001"],
)
assert result.exit_code == 1, f"exit_code={result.exit_code}, output={result.output}"
assert "Prometheus metrics server exited with code 1 before serving 0.0.0.0:4001" in result.output
mock_uvicorn_run.assert_not_called()
mock_popen.reset_mock()
mock_uvicorn_run.reset_mock()
result = runner.invoke(run_server, ["--local"])
assert (
result.exit_code == 0
), f"exit_code={result.exit_code}, output={result.output}"
mock_popen.assert_not_called()
mock_uvicorn_run.assert_called_once()
mock_uvicorn_run.reset_mock()
result = runner.invoke(
run_server,
["--local", "--port", "4000", "--prometheus_metrics_port", "4000"],
)
assert result.exit_code == 2
assert "--prometheus_metrics_port must differ from --port" in result.output
mock_popen.assert_not_called()
mock_uvicorn_run.assert_not_called()
result = runner.invoke(
run_server, ["--local", "--prometheus_metrics_port", "0"]
)
assert result.exit_code == 2
assert "Invalid value for '--prometheus_metrics_port'" in result.output
mock_popen.assert_not_called()
@pytest.mark.parametrize(
"timeout_config,expected_timeout",
[

View file

@ -1,5 +1,8 @@
"""Simple tests for lazy import functionality."""
import importlib
import json
import subprocess
import sys
import pytest
@ -7,6 +10,10 @@ import pytest
import litellm
from litellm._lazy_imports import (
_SDK_MODULE_ALIASES,
_SDK_SYMBOLS_IMPORT_MAP,
lazy_import_litellm_submodule,
_lazy_import_sdk_symbols,
COST_CALCULATOR_NAMES,
LITELLM_LOGGING_NAMES,
UTILS_NAMES,
@ -346,3 +353,83 @@ def test_utils_module_lazy_imports():
assert name in utils_globals
_verify_only_requested_name_imported_in_utils(name, UTILS_MODULE_NAMES)
def test_sdk_symbols_lazy_imports():
"""Every symbol previously imported eagerly in litellm/__init__.py resolves to the source module attribute."""
for name, (module_path, attr_name) in _SDK_SYMBOLS_IMPORT_MAP.items():
resolved = getattr(litellm, name)
expected = getattr(importlib.import_module(module_path), attr_name)
assert resolved is expected, f"litellm.{name} is not {module_path}.{attr_name}"
def test_sdk_module_aliases():
"""Module-valued attributes (litellm.anthropic, litellm.httpx, ...) resolve to the aliased modules."""
for name, module_path in _SDK_MODULE_ALIASES.items():
assert getattr(litellm, name) is importlib.import_module(module_path)
def test_litellm_submodule_fallback():
"""litellm.<submodule> attribute access resolves real submodules and returns None for unknown names."""
assert lazy_import_litellm_submodule("budget_manager") is importlib.import_module("litellm.budget_manager")
assert litellm.utils is importlib.import_module("litellm.utils")
assert lazy_import_litellm_submodule("not_a_real_submodule") is None
with pytest.raises(AttributeError):
_ = litellm.not_a_real_attribute
def test_missing_attribute_stays_attribute_error_when_find_spec_lies(monkeypatch):
"""getattr(litellm, name, default) must not leak ModuleNotFoundError when find_spec is patched to always succeed."""
monkeypatch.setattr(importlib.util, "find_spec", lambda name: object())
assert getattr(litellm, "not_a_real_submodule", None) is None
with pytest.raises(AttributeError):
_ = litellm.not_a_real_attribute
def test_proxy_private_submodule_resolves_in_fresh_process():
"""litellm.proxy._types resolves without an eager proxy import (used by documentation checks)."""
code = "import litellm\nprint(litellm.proxy._types.__name__)\n"
result = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True)
assert result.returncode == 0, result.stderr
assert result.stdout.strip() == "litellm.proxy._types"
def test_lazy_instances_are_singletons():
"""Lazily created instances are cached, so repeated access returns the same object."""
assert litellm._key_management_settings is litellm._key_management_settings
assert litellm.vertexAITextEmbeddingConfig is litellm.vertexAITextEmbeddingConfig
from litellm.types.secret_managers.main import KeyManagementSettings
assert isinstance(litellm._key_management_settings, KeyManagementSettings)
def test_star_import_exports_public_api():
"""`from litellm import *` keeps exporting the full public surface despite lazy loading."""
code = (
"from litellm import *\n"
"import litellm\n"
"missing = [n for n in litellm.__all__ if n not in dir()]\n"
"assert not missing, missing[:20]\n"
"assert callable(completion) and callable(Router)\n"
)
result = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True)
assert result.returncode == 0, result.stderr
@pytest.mark.skipif(sys.platform != "linux", reason="reads /proc for RSS")
def test_import_litellm_stays_lightweight():
"""`import litellm` must not pull in the SDK/proxy heavyweights or blow up RSS (LIT-6607)."""
code = (
"import json, re, sys\n"
"import litellm\n"
"heavy = [m for m in ('litellm.main', 'litellm.utils', 'litellm.router', 'litellm.proxy.proxy_cli',\n"
" 'tiktoken', 'fastapi', 'grpc', 'boto3') if m in sys.modules]\n"
"with open('/proc/self/status') as f:\n"
" rss_kb = int(re.search(r'VmRSS:\\s+(\\d+) kB', f.read()).group(1))\n"
"print(json.dumps({'total': len(sys.modules), 'heavy': heavy, 'rss_mb': rss_kb / 1024}))\n"
)
result = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True, check=True)
stats = json.loads(result.stdout)
assert stats["heavy"] == [], f"heavy modules imported eagerly: {stats['heavy']}"
assert stats["total"] < 800, f"import litellm loaded {stats['total']} modules"
assert stats["rss_mb"] < 75, f"import litellm used {stats['rss_mb']:.1f} MB RSS"

View file

@ -3,7 +3,7 @@
"limit": 22171
},
"LIT002": {
"limit": 26745
"limit": 26728
},
"LIT003": {
"limit": 261
@ -27,7 +27,7 @@
"limit": 0
},
"LIT010": {
"limit": 16451
"limit": 16415
},
"LIT011": {
"limit": 5506

View file

@ -112,6 +112,7 @@ const PRIMITIVES = {
"Return Values": [
{ name: "allow()", desc: "Let request/response through" },
{ name: "block(reason)", desc: "Reject with message" },
{ name: "flag(reason, metadata={})", desc: "Let through, record a non-blocking violation" },
{ name: "modify(texts=[], images=[], tool_calls=[])", desc: "Transform content" },
],
"HTTP Requests (async)": [

View file

@ -33,6 +33,22 @@ describe("GuardrailViewer", () => {
expect(screen.getByText("1235ms")).toBeInTheDocument();
});
it("renders guardrail_flagged as FLAGGED (warning), not FAILED", () => {
const data = makeGuardrailInformation({
guardrail_name: "cc-flag",
guardrail_status: "guardrail_flagged",
guardrail_provider: "custom_code",
});
renderWithProviders(<GuardrailViewer data={data} />);
expect(screen.getByText(/0 Passed/)).toBeInTheDocument();
expect(screen.getByText(/1 Flagged/)).toBeInTheDocument();
const badges = screen.getAllByText("FLAGGED");
expect(badges.length).toBeGreaterThan(0);
expect(badges[0]).toHaveClass("text-warning");
expect(screen.queryByText("FAILED")).not.toBeInTheDocument();
});
it("calculates and displays masked entity totals", async () => {
const user = userEvent.setup();
const data = makeGuardrailInformation({

View file

@ -133,8 +133,27 @@ const getTotalMasked = (entry: GuardrailInformation): number => {
);
};
const isEntrySuccess = (entry: GuardrailInformation): boolean => {
return (entry.guardrail_status ?? "").toLowerCase() === "success";
type EntryOutcome = "passed" | "flagged" | "failed";
const getEntryOutcome = (entry: GuardrailInformation): EntryOutcome => {
const status = (entry.guardrail_status ?? "").toLowerCase();
if (status === "success") return "passed";
if (status === "guardrail_flagged") return "flagged";
return "failed";
};
const isEntrySuccess = (entry: GuardrailInformation): boolean => getEntryOutcome(entry) === "passed";
const OUTCOME_LABEL: Record<EntryOutcome, string> = {
passed: "PASSED",
flagged: "FLAGGED",
failed: "FAILED",
};
const OUTCOME_BADGE_CLASS: Record<EntryOutcome, string> = {
passed: "bg-success/15 text-success border border-success/20",
flagged: "bg-warning/15 text-warning border border-warning/20",
failed: "bg-destructive/15 text-destructive border border-destructive/20",
};
const getRiskColor = (score: number): string => {
@ -202,6 +221,19 @@ const FailCircleIcon = ({ className }: { className?: string }) => (
</svg>
);
const FlagCircleIcon = ({ className }: { className?: string }) => (
<svg width="22" height="22" viewBox="0 0 22 22" fill="none" className={className}>
<circle cx="11" cy="11" r="10" stroke="#D97706" strokeWidth="1.5" fill="#FFFBEB" />
<path d="M11 6.5v5M11 14.5v.5" stroke="#D97706" strokeWidth="1.5" strokeLinecap="round" />
</svg>
);
const OutcomeIcon = ({ outcome }: { outcome: EntryOutcome }) => {
if (outcome === "passed") return <CheckCircleIcon />;
if (outcome === "flagged") return <FlagCircleIcon />;
return <FailCircleIcon />;
};
const PlayCircleIcon = () => (
<svg width="22" height="22" viewBox="0 0 22 22" fill="none">
<circle cx="11" cy="11" r="10" stroke="#3B82F6" strokeWidth="1.5" fill="#EFF6FF" />
@ -318,8 +350,7 @@ interface TimelineEntry {
type: "request" | "guardrail" | "llm" | "response";
label: string;
offsetMs: number;
status?: string;
isSuccess?: boolean;
outcome?: EntryOutcome;
}
const RequestLifecycle = ({ entries }: { entries: GuardrailInformation[] }) => {
@ -348,8 +379,7 @@ const RequestLifecycle = ({ entries }: { entries: GuardrailInformation[] }) => {
type: "guardrail",
label: `Pre-call guardrail: ${getDisplayName(e)}`,
offsetMs,
status: isEntrySuccess(e) ? "PASSED" : "FAILED",
isSuccess: isEntrySuccess(e),
outcome: getEntryOutcome(e),
});
}
@ -372,8 +402,7 @@ const RequestLifecycle = ({ entries }: { entries: GuardrailInformation[] }) => {
type: "guardrail",
label: `During-call guardrail: ${getDisplayName(e)}`,
offsetMs,
status: isEntrySuccess(e) ? "PASSED" : "FAILED",
isSuccess: isEntrySuccess(e),
outcome: getEntryOutcome(e),
});
}
@ -384,8 +413,7 @@ const RequestLifecycle = ({ entries }: { entries: GuardrailInformation[] }) => {
type: "guardrail",
label: `Post-call guardrail: ${getDisplayName(e)}`,
offsetMs,
status: isEntrySuccess(e) ? "PASSED" : "FAILED",
isSuccess: isEntrySuccess(e),
outcome: getEntryOutcome(e),
});
}
@ -410,10 +438,8 @@ const RequestLifecycle = ({ entries }: { entries: GuardrailInformation[] }) => {
<GrayDotIcon />
) : item.type === "llm" ? (
<PlayCircleIcon />
) : item.isSuccess ? (
<CheckCircleIcon />
) : (
<FailCircleIcon />
<OutcomeIcon outcome={item.outcome ?? "failed"} />
)}
</div>
{idx < timeline.length - 1 && <div className="w-0.5 bg-border grow" style={{ minHeight: "24px" }} />}
@ -425,13 +451,11 @@ const RequestLifecycle = ({ entries }: { entries: GuardrailInformation[] }) => {
<span className={`text-sm ${item.type === "llm" ? "text-info font-medium" : "text-foreground"}`}>
{item.label}
</span>
{item.status && (
{item.outcome && (
<span
className={`px-1.5 py-0.5 rounded text-[10px] font-bold uppercase ${
item.isSuccess ? "bg-success/15 text-success" : "bg-destructive/15 text-destructive"
}`}
className={`px-1.5 py-0.5 rounded text-[10px] font-bold uppercase ${OUTCOME_BADGE_CLASS[item.outcome]}`}
>
{item.status}
{OUTCOME_LABEL[item.outcome]}
</span>
)}
<span className="text-xs text-muted-foreground font-mono ml-auto shrink-0">T+{item.offsetMs}ms</span>
@ -455,7 +479,7 @@ const formatGuardrailCost = (cost: number): string => {
const EvaluationCard = ({ entry }: { entry: GuardrailInformation }) => {
const [expanded, setExpanded] = useState(false);
const success = isEntrySuccess(entry);
const outcome = getEntryOutcome(entry);
const totalMasked = getTotalMasked(entry);
const displayName = getDisplayName(entry);
const durationStr = formatDurationMs(entry.duration);
@ -490,7 +514,9 @@ const EvaluationCard = ({ entry }: { entry: GuardrailInformation }) => {
onClick={() => setExpanded(!expanded)}
>
{/* Status icon */}
<div className="shrink-0">{success ? <CheckCircleIcon /> : <FailCircleIcon />}</div>
<div className="shrink-0">
<OutcomeIcon outcome={outcome} />
</div>
{/* Name + badges */}
<div className="flex items-center gap-2 flex-wrap flex-1 min-w-0">
@ -501,13 +527,9 @@ const EvaluationCard = ({ entry }: { entry: GuardrailInformation }) => {
</span>
<span
className={`px-2 py-0.5 rounded text-[11px] font-semibold uppercase shrink-0 ${
success
? "bg-success/15 text-success border border-success/20"
: "bg-destructive/15 text-destructive border border-destructive/20"
}`}
className={`px-2 py-0.5 rounded text-[11px] font-semibold uppercase shrink-0 ${OUTCOME_BADGE_CLASS[outcome]}`}
>
{success ? "PASSED" : "FAILED"}
{OUTCOME_LABEL[outcome]}
</span>
{matchCountStr && (
@ -528,7 +550,7 @@ const EvaluationCard = ({ entry }: { entry: GuardrailInformation }) => {
</span>
)}
{riskScore != null && success && (
{riskScore != null && outcome === "passed" && (
<TooltipProvider>
<Tooltip>
<TooltipTrigger
@ -673,7 +695,13 @@ const GuardrailViewer = ({ data, accessToken, logEntry }: GuardrailViewerProps)
}, [data]);
const passedCount = guardrailEntries.filter(isEntrySuccess).length;
const flaggedCount = guardrailEntries.filter((e) => getEntryOutcome(e) === "flagged").length;
const allPassed = passedCount === guardrailEntries.length;
const headerOutcome: EntryOutcome = allPassed
? "passed"
: passedCount + flaggedCount === guardrailEntries.length
? "flagged"
: "failed";
const totalOverheadMs = useMemo(() => {
return Math.round(guardrailEntries.reduce((sum, e) => sum + (e.duration ?? 0), 0) * 1000);
@ -709,11 +737,7 @@ const GuardrailViewer = ({ data, accessToken, logEntry }: GuardrailViewerProps)
</span>
<span className="text-muted-foreground">|</span>
<span
className={`inline-flex items-center gap-1 px-2 py-0.5 rounded-full text-xs font-semibold ${
allPassed
? "bg-success/10 text-success border border-success/20"
: "bg-destructive/10 text-destructive border border-destructive/20"
}`}
className={`inline-flex items-center gap-1 px-2 py-0.5 rounded-full text-xs font-semibold ${OUTCOME_BADGE_CLASS[headerOutcome]}`}
>
{allPassed ? (
<svg width="12" height="12" viewBox="0 0 12 12" fill="none">
@ -728,6 +752,13 @@ const GuardrailViewer = ({ data, accessToken, logEntry }: GuardrailViewerProps)
) : null}
{passedCount} Passed
</span>
{flaggedCount > 0 && (
<span
className={`inline-flex items-center px-2 py-0.5 rounded-full text-xs font-semibold ${OUTCOME_BADGE_CLASS.flagged}`}
>
{flaggedCount} Flagged
</span>
)}
</div>
</div>
</div>

View file

@ -1,7 +1,7 @@
import { render, screen, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { describe, expect, it, vi } from "vitest";
import { LogDetailContent } from "./LogDetailContent";
import { GuardrailJumpLink, LogDetailContent } from "./LogDetailContent";
import type { LogEntry } from "../columns";
vi.mock("../GuardrailViewer/GuardrailViewer", () => ({
@ -489,3 +489,17 @@ describe("LogDetailContent", () => {
expect(within(descriptions).getByText("-")).toBeInTheDocument();
});
});
describe("GuardrailJumpLink", () => {
it.each([
[["success", "success"], "text-success", "\u2713"],
[["success", "guardrail_flagged"], "text-warning", "\u26A0"],
[["guardrail_flagged", "guardrail_intervened"], "text-destructive", "\u2717"],
])("styles %j as %s", (statuses, expectedClass, glyph) => {
render(<GuardrailJumpLink guardrailEntries={statuses.map((s) => ({ guardrail_status: s }))} />);
const pill = screen.getByText(/2 guardrails evaluated/);
expect(pill).toHaveClass(expectedClass);
expect(pill).toHaveTextContent(glyph);
});
});

View file

@ -635,11 +635,24 @@ function RequestResponseSection({
);
}
const GUARDRAIL_JUMP_LINK_STYLE = {
passed: { className: "border border-success/20 bg-success/10 text-success", glyph: "\u2713" },
flagged: { className: "border border-warning/20 bg-warning/10 text-warning", glyph: "\u26A0" },
failed: { className: "border border-destructive/20 bg-destructive/10 text-destructive", glyph: "\u2717" },
} as const;
const isPassedStatus = (status: unknown) => status === "pass" || status === "passed" || status === "success";
const isFlaggedStatus = (status: unknown) => status === "flagged" || status === "guardrail_flagged";
const guardrailJumpLinkOutcome = (statuses: unknown[]): keyof typeof GUARDRAIL_JUMP_LINK_STYLE => {
if (statuses.every(isPassedStatus)) return "passed";
if (statuses.every((s) => isPassedStatus(s) || isFlaggedStatus(s))) return "flagged";
return "failed";
};
export function GuardrailJumpLink({ guardrailEntries }: { guardrailEntries: any[] }) {
const allPassed = guardrailEntries.every((e) => {
const status = e?.guardrail_status || e?.status;
return status === "pass" || status === "passed" || status === "success";
});
const outcome = guardrailJumpLinkOutcome(guardrailEntries.map((e) => e?.guardrail_status || e?.status));
const { className, glyph } = GUARDRAIL_JUMP_LINK_STYLE[outcome];
const handleClick = () => {
const el = document.getElementById("guardrail-section");
@ -650,11 +663,7 @@ export function GuardrailJumpLink({ guardrailEntries }: { guardrailEntries: any[
<div style={{ textAlign: "left", marginBottom: 12 }}>
<div
onClick={handleClick}
className={
allPassed
? "border border-success/20 bg-success/10 text-success"
: "border border-destructive/20 bg-destructive/10 text-destructive"
}
className={className}
style={{
display: "inline-flex",
alignItems: "center",
@ -666,8 +675,8 @@ export function GuardrailJumpLink({ guardrailEntries }: { guardrailEntries: any[
fontWeight: 500,
}}
>
{allPassed ? "\u2713" : "\u2717"} {guardrailEntries.length} guardrail{guardrailEntries.length !== 1 ? "s" : ""}{" "}
evaluated
{glyph} {guardrailEntries.length} guardrail
{guardrailEntries.length !== 1 ? "s" : ""} evaluated
<span style={{ fontSize: 11, opacity: 0.7 }}>{"\u2193"}</span>
</div>
</div>