mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
commit
f950274f93
46 changed files with 5735 additions and 358 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]],
|
||||
|
|
|
|||
167
litellm/proxy/prometheus_metrics_server.py
Normal file
167
litellm/proxy/prometheus_metrics_server.py
Normal 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()
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -21,6 +21,6 @@
|
|||
"limit": 117
|
||||
},
|
||||
"TQ008": {
|
||||
"limit": 10990
|
||||
"limit": 10992
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
167
tests/e2e/guardrails/test_tool_permission_guardrail_e2e.py
Normal file
167
tests/e2e/guardrails/test_tool_permission_guardrail_e2e.py
Normal 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]}"
|
||||
)
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
192
tests/e2e/logging/test_weave_log_e2e.py
Normal file
192
tests/e2e/logging/test_weave_log_e2e.py
Normal 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
|
||||
282
tests/e2e/logging/weave_reader.py
Normal file
282
tests/e2e/logging/weave_reader.py
Normal 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)
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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) ---------------------------------------------
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
259
tests/test_litellm/proxy/test_prometheus_metrics_server.py
Normal file
259
tests/test_litellm/proxy/test_prometheus_metrics_server.py
Normal 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
|
||||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)": [
|
||||
|
|
|
|||
|
|
@ -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({
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue