mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Merge 409b402553 into 1c61c2606e
This commit is contained in:
commit
48f7ca1b97
7 changed files with 835 additions and 125 deletions
|
|
@ -1,5 +1,9 @@
|
|||
import importlib
|
||||
import importlib.util
|
||||
import os
|
||||
from collections.abc import Iterable, Iterator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -80,93 +84,181 @@ def get_cost_for_web_search_request(custom_llm_provider: str, usage: "Usage", mo
|
|||
return None
|
||||
|
||||
|
||||
def discover_guardrail_translation_mappings() -> dict[CallTypes, type["BaseTranslation"]]:
|
||||
_GUARDRAIL_TRANSLATION_PACKAGE: Final = "guardrail_translation"
|
||||
_MCP_GUARDRAIL_TRANSLATION_MODULE: Final = "litellm.proxy._experimental.mcp_server.guardrail_translation"
|
||||
_NO_MAPPINGS: Final[Mapping[CallTypes, type["BaseTranslation"]]] = MappingProxyType({})
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class GuardrailTranslationDiscovery:
|
||||
"""
|
||||
The outcome of one scan for guardrail translation handlers.
|
||||
|
||||
unavailable maps each package that failed to import to the reason, which is what tells a complete
|
||||
result apart from one that is missing handlers and therefore has to be retried.
|
||||
"""
|
||||
|
||||
mappings: Mapping[CallTypes, type["BaseTranslation"]]
|
||||
unavailable: Mapping[str, str]
|
||||
|
||||
|
||||
def _bundled_guardrail_translation_modules() -> Iterator[str]:
|
||||
"""Yield the import path of every guardrail_translation package shipped under litellm/llms."""
|
||||
llms_dir: Final = os.path.dirname(__file__)
|
||||
for root, dirs, files in os.walk(llms_dir):
|
||||
dirs[:] = tuple(d for d in dirs if not d.startswith("__") and d != "base_llm")
|
||||
if os.path.basename(root) == _GUARDRAIL_TRANSLATION_PACKAGE and "__init__.py" in files:
|
||||
yield "litellm." + os.path.relpath(root, os.path.dirname(llms_dir)).replace(os.sep, ".")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _UnavailablePackage:
|
||||
"""
|
||||
Why one guardrail_translation package could not be imported.
|
||||
|
||||
missing_dependency is set when the package asked for a module this install does not have at all, which is
|
||||
the one failure that says the package is absent rather than momentarily unimportable.
|
||||
"""
|
||||
|
||||
reason: str
|
||||
missing_dependency: bool
|
||||
|
||||
|
||||
def _is_absent(module_name: str) -> bool:
|
||||
"""Whether this install has no module of that name, as opposed to one that is present but failed to import."""
|
||||
root: Final = module_name.partition(".")[0]
|
||||
try:
|
||||
return importlib.util.find_spec(root) is None
|
||||
except (ImportError, ValueError):
|
||||
return False
|
||||
|
||||
|
||||
def _import_guardrail_translations(
|
||||
module_path: str,
|
||||
) -> Mapping[CallTypes, type["BaseTranslation"]] | _UnavailablePackage:
|
||||
"""Import one guardrail_translation package, reporting why that failed instead of raising."""
|
||||
try:
|
||||
module: Final = importlib.import_module(module_path)
|
||||
except Exception as e: # noqa: BLE001 # a package failing at import time for any reason is unavailable, not fatal
|
||||
return _UnavailablePackage(
|
||||
reason=f"{type(e).__name__}: {e}",
|
||||
missing_dependency=isinstance(e, ModuleNotFoundError) and _is_absent(e.name or ""),
|
||||
)
|
||||
mappings: Final = getattr(module, "guardrail_translation_mappings", None)
|
||||
if not isinstance(mappings, dict):
|
||||
return _NO_MAPPINGS
|
||||
declared: Final[Mapping[CallTypes, type[BaseTranslation]]] = mappings
|
||||
return declared
|
||||
|
||||
|
||||
def _guardrail_translations_from(
|
||||
module_path: str,
|
||||
) -> Mapping[CallTypes, type["BaseTranslation"]] | _UnavailablePackage:
|
||||
"""
|
||||
Import one package's handlers, tolerating an install that does not ship the optional MCP server.
|
||||
|
||||
litellm ships every package under llms, so a failure there is a gap to retry rather than a fact about the
|
||||
install. The MCP package instead arrives with the proxy extra, and a dependency this install does not have
|
||||
at all means it serves no MCP endpoints for a guardrail to scan, so there is nothing to retry or report. A
|
||||
dependency that is installed and still fails to import is a broken install, which is reported and retried.
|
||||
"""
|
||||
result: Final = _import_guardrail_translations(module_path)
|
||||
if (
|
||||
module_path == _MCP_GUARDRAIL_TRANSLATION_MODULE
|
||||
and isinstance(result, _UnavailablePackage)
|
||||
and result.missing_dependency
|
||||
):
|
||||
verbose_logger.debug("%s is not installed: %s", module_path, result.reason)
|
||||
return _NO_MAPPINGS
|
||||
return result
|
||||
|
||||
|
||||
def _guardrail_translation_modules() -> Iterator[str]:
|
||||
"""Yield every module that can declare guardrail translation handlers, the optional MCP one last."""
|
||||
yield from _bundled_guardrail_translation_modules()
|
||||
yield _MCP_GUARDRAIL_TRANSLATION_MODULE
|
||||
|
||||
|
||||
def _discover(
|
||||
module_paths: Iterable[str], already_found: Mapping[CallTypes, type["BaseTranslation"]]
|
||||
) -> GuardrailTranslationDiscovery:
|
||||
imported: Final = tuple((module_path, _guardrail_translations_from(module_path)) for module_path in module_paths)
|
||||
found: Final = (
|
||||
already_found,
|
||||
*(result for _, result in imported if not isinstance(result, _UnavailablePackage)),
|
||||
)
|
||||
return GuardrailTranslationDiscovery(
|
||||
mappings=MappingProxyType(
|
||||
{call_type: handler for mappings in found for call_type, handler in mappings.items()}
|
||||
),
|
||||
unavailable=MappingProxyType(
|
||||
{module_path: result.reason for module_path, result in imported if isinstance(result, _UnavailablePackage)}
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def discover_guardrail_translations() -> GuardrailTranslationDiscovery:
|
||||
"""
|
||||
Scan the llms tree, plus the optional MCP package, for guardrail translation handlers.
|
||||
|
||||
Returns:
|
||||
GuardrailTranslationDiscovery: the handlers found, and the packages that failed to import
|
||||
"""
|
||||
return _discover(_guardrail_translation_modules(), already_found=_NO_MAPPINGS)
|
||||
|
||||
|
||||
def discover_guardrail_translation_mappings() -> Mapping[CallTypes, type["BaseTranslation"]]:
|
||||
"""
|
||||
Discover guardrail translation mappings by scanning the llms directory structure.
|
||||
|
||||
Scans for modules with guardrail_translation_mappings dictionaries and aggregates them.
|
||||
|
||||
Returns:
|
||||
Dict[CallTypes, Type[BaseTranslation]]: A dictionary mapping call types to their translation handler classes
|
||||
Mapping[CallTypes, Type[BaseTranslation]]: the call types that have a translation handler class
|
||||
"""
|
||||
discovered_mappings: Final[dict[CallTypes, type[BaseTranslation]]] = {}
|
||||
return discover_guardrail_translations().mappings
|
||||
|
||||
try:
|
||||
# Get the path to the llms directory
|
||||
current_dir: Final = os.path.dirname(__file__)
|
||||
llms_dir: Final = current_dir
|
||||
|
||||
if not os.path.exists(llms_dir):
|
||||
verbose_logger.debug("llms directory not found")
|
||||
return discovered_mappings
|
||||
|
||||
# Recursively scan for guardrail_translation directories
|
||||
for root, dirs, files in os.walk(llms_dir):
|
||||
# Skip __pycache__ and base_llm directories
|
||||
dirs[:] = [d for d in dirs if not d.startswith("__") and d != "base_llm"]
|
||||
|
||||
# Check if this is a guardrail_translation directory with __init__.py
|
||||
if os.path.basename(root) == "guardrail_translation" and "__init__.py" in files:
|
||||
# Build the module path relative to litellm
|
||||
rel_path = os.path.relpath(root, os.path.dirname(llms_dir))
|
||||
module_path = "litellm." + rel_path.replace(os.sep, ".")
|
||||
|
||||
try:
|
||||
# Import the module
|
||||
verbose_logger.debug("Discovering guardrail translations in: %s", module_path)
|
||||
|
||||
module = importlib.import_module(module_path)
|
||||
|
||||
# Check for guardrail_translation_mappings dictionary
|
||||
if hasattr(module, "guardrail_translation_mappings"):
|
||||
mappings = getattr(module, "guardrail_translation_mappings")
|
||||
if isinstance(mappings, dict):
|
||||
discovered_mappings.update(mappings)
|
||||
verbose_logger.debug(
|
||||
"Found guardrail_translation_mappings in %s: %s", module_path, list(mappings.keys())
|
||||
)
|
||||
|
||||
except ImportError as e:
|
||||
verbose_logger.error("Could not import %s: %s", module_path, e)
|
||||
continue
|
||||
except Exception as e:
|
||||
verbose_logger.error("Error processing %s: %s", module_path, e)
|
||||
continue
|
||||
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.guardrail_translation import (
|
||||
guardrail_translation_mappings as mcp_guardrail_translation_mappings,
|
||||
)
|
||||
|
||||
discovered_mappings.update(mcp_guardrail_translation_mappings)
|
||||
verbose_logger.debug(
|
||||
"Loaded MCP guardrail translation mappings: %s",
|
||||
list(mcp_guardrail_translation_mappings.keys()),
|
||||
)
|
||||
except ImportError:
|
||||
verbose_logger.debug("MCP guardrail translation mappings not available; skipping")
|
||||
|
||||
verbose_logger.debug(
|
||||
"Discovered %s guardrail translation mappings: %s",
|
||||
len(discovered_mappings),
|
||||
list(discovered_mappings.keys()),
|
||||
def _announce_discovery(previous: GuardrailTranslationDiscovery | None, current: GuardrailTranslationDiscovery) -> None:
|
||||
if previous is None and not current.unavailable:
|
||||
return
|
||||
if previous is None:
|
||||
verbose_logger.error(
|
||||
"Could not import guardrail translation handlers from %s; guardrails cannot run for their call types "
|
||||
"until the import succeeds, which every lookup retries. %s",
|
||||
", ".join(current.unavailable),
|
||||
"; ".join(f"{module_path}: {reason}" for module_path, reason in current.unavailable.items()),
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.error("Error discovering guardrail translation mappings: %s", e)
|
||||
|
||||
return discovered_mappings
|
||||
return
|
||||
recovered: Final = tuple(
|
||||
module_path for module_path in previous.unavailable if module_path not in current.unavailable
|
||||
)
|
||||
if not recovered:
|
||||
return
|
||||
verbose_logger.warning("Guardrail translation handlers from %s are available again.", ", ".join(recovered))
|
||||
|
||||
|
||||
# Cache the discovered mappings
|
||||
endpoint_guardrail_translation_mappings: dict[CallTypes, type["BaseTranslation"]] | None = None
|
||||
guardrail_translation_discovery: GuardrailTranslationDiscovery | None = None
|
||||
|
||||
|
||||
def load_guardrail_translation_mappings():
|
||||
global endpoint_guardrail_translation_mappings
|
||||
if endpoint_guardrail_translation_mappings is None:
|
||||
endpoint_guardrail_translation_mappings = discover_guardrail_translation_mappings()
|
||||
return endpoint_guardrail_translation_mappings
|
||||
def load_guardrail_translation_mappings() -> Mapping[CallTypes, type["BaseTranslation"]]:
|
||||
"""
|
||||
Return the guardrail translation handlers, retrying any bundled package that could not be imported last time.
|
||||
|
||||
Serving an incomplete scan as if it were complete would silently strip the missing call types off every
|
||||
guardrail for the rest of the process, so the packages that failed are imported again on each lookup, and
|
||||
only the part that succeeded is kept.
|
||||
"""
|
||||
global guardrail_translation_discovery
|
||||
cached: Final = guardrail_translation_discovery
|
||||
if cached is not None and not cached.unavailable:
|
||||
return cached.mappings
|
||||
discovery: Final = (
|
||||
discover_guardrail_translations()
|
||||
if cached is None
|
||||
else _discover(cached.unavailable, already_found=cached.mappings)
|
||||
)
|
||||
_announce_discovery(previous=cached, current=discovery)
|
||||
guardrail_translation_discovery = discovery
|
||||
return discovery.mappings
|
||||
|
||||
|
||||
def get_guardrail_translation_mapping(call_type: CallTypes) -> type["BaseTranslation"]:
|
||||
|
|
@ -182,18 +274,10 @@ def get_guardrail_translation_mapping(call_type: CallTypes) -> type["BaseTransla
|
|||
Raises:
|
||||
ValueError: If no translation mapping exists for the given call type
|
||||
"""
|
||||
global endpoint_guardrail_translation_mappings
|
||||
|
||||
# Lazy load the mappings on first access
|
||||
if endpoint_guardrail_translation_mappings is None:
|
||||
endpoint_guardrail_translation_mappings = discover_guardrail_translation_mappings()
|
||||
|
||||
# Get the translation handler class for the call type
|
||||
if call_type not in endpoint_guardrail_translation_mappings:
|
||||
mappings: Final = load_guardrail_translation_mappings()
|
||||
if call_type not in mappings:
|
||||
raise ValueError(
|
||||
f"No guardrail translation mapping found for call_type: {call_type}. "
|
||||
f"Available mappings: {list(endpoint_guardrail_translation_mappings.keys())}"
|
||||
f"Available mappings: {list(mappings.keys())}"
|
||||
)
|
||||
|
||||
# Return the handler class directly
|
||||
return endpoint_guardrail_translation_mappings[call_type]
|
||||
return mappings[call_type]
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ Unified Guardrail, leveraging LiteLLM's /applyGuardrail endpoint
|
|||
import copy
|
||||
import json
|
||||
from collections.abc import AsyncGenerator, AsyncIterable, Awaitable, Callable, Mapping, Sequence
|
||||
from functools import lru_cache
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -163,6 +164,72 @@ def _ensure_litellm_metadata(data: dict, user_api_key_dict: UserAPIKeyAuth) -> N
|
|||
data["litellm_metadata"] = user_metadata
|
||||
|
||||
|
||||
_UNSCANNED_WARNING_KEYS: Final = 4096
|
||||
|
||||
|
||||
def _resolved_call_type(call_type: str | None) -> CallTypes | None:
|
||||
"""Return the CallTypes member a route's call type names, or None when the enum has no member for it."""
|
||||
try:
|
||||
return CallTypes(call_type)
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
|
||||
@lru_cache(maxsize=_UNSCANNED_WARNING_KEYS)
|
||||
def _warn_left_unscanned_once(
|
||||
guardrail_name: str | None,
|
||||
request_route: str | None,
|
||||
call_type: str | None,
|
||||
consequence: str,
|
||||
) -> None:
|
||||
if _resolved_call_type(call_type) is not None:
|
||||
verbose_proxy_logger.warning(
|
||||
"Guardrail '%s' selected for route '%s' but call type '%s' has no guardrail translation handler; %s. "
|
||||
"Add a guardrail translation handler for that call type.",
|
||||
guardrail_name,
|
||||
request_route,
|
||||
call_type,
|
||||
consequence,
|
||||
)
|
||||
return
|
||||
unscannable, remedy = (
|
||||
(
|
||||
f"call type '{call_type}' is not one litellm can scan",
|
||||
"Map the route to a CallTypes member in API_ROUTE_TO_CALL_TYPES",
|
||||
)
|
||||
if call_type
|
||||
else ("its call type could not be resolved", "Add the route to API_ROUTE_TO_CALL_TYPES")
|
||||
)
|
||||
verbose_proxy_logger.warning(
|
||||
"Guardrail '%s' selected for route '%s' but %s, so no guardrail can run on that route; %s. %s.",
|
||||
guardrail_name,
|
||||
request_route,
|
||||
unscannable,
|
||||
consequence,
|
||||
remedy,
|
||||
)
|
||||
|
||||
|
||||
def _warn_left_unscanned(
|
||||
guardrail_to_apply: "CustomGuardrail",
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
call_type: str | None,
|
||||
consequence: str,
|
||||
) -> None:
|
||||
"""
|
||||
Say that a selected guardrail could not scan this request, once per route and reason.
|
||||
|
||||
Most proxy routes have no translation handler and never will, so a line per request would bury the
|
||||
outage it is meant to surface, and repeating it adds nothing an operator can act on twice.
|
||||
"""
|
||||
_warn_left_unscanned_once(
|
||||
guardrail_name=guardrail_to_apply.guardrail_name,
|
||||
request_route=user_api_key_dict.request_route,
|
||||
call_type=call_type,
|
||||
consequence=consequence,
|
||||
)
|
||||
|
||||
|
||||
class UnifiedLLMGuardrails(CustomLogger):
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -213,14 +280,17 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
return data
|
||||
|
||||
mappings: Final = load_guardrail_translation_mappings()
|
||||
resolved_call_type: Final = _resolved_call_type(call_type)
|
||||
if resolved_call_type is None or resolved_call_type not in mappings:
|
||||
_warn_left_unscanned(
|
||||
guardrail_to_apply=guardrail_to_apply,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
consequence="skipping pre-call scanning",
|
||||
)
|
||||
return data
|
||||
|
||||
try:
|
||||
if CallTypes(call_type) not in mappings:
|
||||
return data
|
||||
except ValueError:
|
||||
return data # handle unmapped call types
|
||||
|
||||
endpoint_translation: Final = _as_endpoint_translation(mappings[CallTypes(call_type)]())
|
||||
endpoint_translation: Final = _as_endpoint_translation(mappings[resolved_call_type]())
|
||||
|
||||
_ensure_litellm_metadata(data, user_api_key_dict)
|
||||
|
||||
|
|
@ -263,10 +333,17 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
return data
|
||||
|
||||
mappings: Final = load_guardrail_translation_mappings()
|
||||
if call_type is not None and CallTypes(call_type) not in mappings:
|
||||
resolved_call_type: Final = _resolved_call_type(call_type)
|
||||
if resolved_call_type is None or resolved_call_type not in mappings:
|
||||
_warn_left_unscanned(
|
||||
guardrail_to_apply=guardrail_to_apply,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
consequence="skipping during-call scanning",
|
||||
)
|
||||
return data
|
||||
|
||||
endpoint_translation: Final = _as_endpoint_translation(mappings[CallTypes(call_type)]())
|
||||
endpoint_translation: Final = _as_endpoint_translation(mappings[resolved_call_type]())
|
||||
|
||||
_ensure_litellm_metadata(data, user_api_key_dict)
|
||||
|
||||
|
|
@ -311,7 +388,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
if user_api_key_dict.request_route is not None:
|
||||
call_types: Final = get_call_types_for_route(user_api_key_dict.request_route)
|
||||
if call_types is not None and len(call_types) > 0:
|
||||
call_type = call_types[0]
|
||||
call_type = call_types[0].value
|
||||
if call_type is None:
|
||||
call_type = _infer_call_type(call_type=None, completion_response=response)
|
||||
|
||||
|
|
@ -327,28 +404,18 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
):
|
||||
call_type = logging_call_type
|
||||
|
||||
if call_type is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"Guardrail '%s' selected for route '%s' but its call type could not be resolved; "
|
||||
"skipping post-call scanning. Add the route to API_ROUTE_TO_CALL_TYPES.",
|
||||
guardrail_to_apply.guardrail_name,
|
||||
user_api_key_dict.request_route,
|
||||
)
|
||||
return response
|
||||
|
||||
mappings: Final = load_guardrail_translation_mappings()
|
||||
|
||||
if CallTypes(call_type) not in mappings:
|
||||
verbose_proxy_logger.warning(
|
||||
"Guardrail '%s' selected for route '%s' but call type '%s' has no guardrail translation handler; "
|
||||
"skipping post-call scanning.",
|
||||
guardrail_to_apply.guardrail_name,
|
||||
user_api_key_dict.request_route,
|
||||
call_type,
|
||||
resolved_call_type: Final = _resolved_call_type(call_type)
|
||||
if resolved_call_type is None or resolved_call_type not in mappings:
|
||||
_warn_left_unscanned(
|
||||
guardrail_to_apply=guardrail_to_apply,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
consequence="skipping post-call scanning",
|
||||
)
|
||||
return response
|
||||
|
||||
endpoint_translation: Final = _as_endpoint_translation(mappings[CallTypes(call_type)]())
|
||||
endpoint_translation: Final = _as_endpoint_translation(mappings[resolved_call_type]())
|
||||
|
||||
try:
|
||||
response = await endpoint_translation.process_output_response(
|
||||
|
|
@ -450,7 +517,7 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
has no in-stream error frame) the exception is re-raised so the proxy
|
||||
can report it with a real HTTP status.
|
||||
"""
|
||||
if call_type is not None and CallTypes(call_type) in A2A_CALL_TYPES:
|
||||
if _resolved_call_type(call_type) in A2A_CALL_TYPES:
|
||||
yield _a2a_jsonrpc_error_chunk(exc, _get_a2a_request_id(responses_so_far, request_data))
|
||||
return
|
||||
if stream_started and endpoint_translation is not None:
|
||||
|
|
@ -1046,7 +1113,13 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
call_type = _infer_call_type(call_type=None, completion_response=item)
|
||||
|
||||
# If call type not supported, just pass through all chunks
|
||||
if call_type is None or CallTypes(call_type) not in mappings:
|
||||
if _resolved_call_type(call_type) not in mappings:
|
||||
_warn_left_unscanned(
|
||||
guardrail_to_apply=guardrail_to_apply,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
consequence="streaming this response to the client unscanned",
|
||||
)
|
||||
yield item
|
||||
async for remaining_item in response:
|
||||
yield remaining_item
|
||||
|
|
@ -1152,14 +1225,15 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
yield item
|
||||
|
||||
# Stream has ended - do final processing with all collected chunks
|
||||
if call_type is not None and CallTypes(call_type) in mappings:
|
||||
final_call_type: Final = _resolved_call_type(call_type)
|
||||
if final_call_type is not None and final_call_type in mappings:
|
||||
verbose_proxy_logger.debug(
|
||||
"Processing final streaming response with all %s chunks for guardrail %s",
|
||||
len(responses_so_far),
|
||||
guardrail_to_apply.guardrail_name,
|
||||
)
|
||||
|
||||
endpoint_translation = mappings[CallTypes(call_type)]()
|
||||
endpoint_translation = mappings[final_call_type]()
|
||||
|
||||
# When buffering, snapshot the original chunks before moderation.
|
||||
# A shallow copy suffices: end-of-stream
|
||||
|
|
|
|||
210
tests/test_litellm/llms/test_guardrail_translation_discovery.py
Normal file
210
tests/test_litellm/llms/test_guardrail_translation_discovery.py
Normal file
|
|
@ -0,0 +1,210 @@
|
|||
import importlib.abc
|
||||
import importlib.util
|
||||
import logging
|
||||
import sys
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from types import ModuleType
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm.llms as llms_package
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.guardrail_translation.handler import (
|
||||
MCPGuardrailTranslationHandler,
|
||||
)
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
OPENAI_CHAT_TRANSLATION_MODULE = "litellm.llms.openai.chat.guardrail_translation"
|
||||
MCP_TRANSLATION_MODULE = "litellm.proxy._experimental.mcp_server.guardrail_translation"
|
||||
|
||||
|
||||
class RaisingLoader(importlib.abc.MetaPathFinder, importlib.abc.Loader):
|
||||
"""Serves one module path, and raises the given error when Python executes it."""
|
||||
|
||||
def __init__(self, module_path: str, error: BaseException) -> None:
|
||||
self.module_path = module_path
|
||||
self.error = error
|
||||
|
||||
def find_spec(self, fullname: str, path=None, target=None):
|
||||
if fullname != self.module_path:
|
||||
return None
|
||||
return importlib.util.spec_from_loader(fullname, self)
|
||||
|
||||
def create_module(self, spec) -> ModuleType | None:
|
||||
return None
|
||||
|
||||
def exec_module(self, module: ModuleType) -> None:
|
||||
raise self.error
|
||||
|
||||
|
||||
@contextmanager
|
||||
def raising_on_import(module_path: str, error: BaseException) -> Iterator[None]:
|
||||
with pytest.MonkeyPatch.context() as mp:
|
||||
mp.delitem(sys.modules, module_path, raising=False)
|
||||
mp.setattr(sys, "meta_path", [RaisingLoader(module_path, error), *sys.meta_path])
|
||||
yield
|
||||
|
||||
|
||||
@contextmanager
|
||||
def unimportable(module_path: str) -> Iterator[None]:
|
||||
with pytest.MonkeyPatch.context() as mp:
|
||||
mp.setitem(sys.modules, module_path, None)
|
||||
yield
|
||||
|
||||
|
||||
@contextmanager
|
||||
def capturing(caplog: pytest.LogCaptureFixture, level: int) -> Iterator[None]:
|
||||
with pytest.MonkeyPatch.context() as mp:
|
||||
mp.setattr(verbose_logger, "propagate", True)
|
||||
caplog.set_level(level, logger=verbose_logger.name)
|
||||
yield
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_guardrail_translation_discovery():
|
||||
saved = llms_package.guardrail_translation_discovery
|
||||
llms_package.guardrail_translation_discovery = None
|
||||
yield
|
||||
llms_package.guardrail_translation_discovery = saved
|
||||
|
||||
|
||||
def test_discovery_reports_the_handler_package_it_could_not_import():
|
||||
with unimportable(OPENAI_CHAT_TRANSLATION_MODULE):
|
||||
discovery = llms_package.discover_guardrail_translations()
|
||||
|
||||
assert tuple(discovery.unavailable) == (OPENAI_CHAT_TRANSLATION_MODULE,)
|
||||
assert "None in sys.modules" in discovery.unavailable[OPENAI_CHAT_TRANSLATION_MODULE]
|
||||
assert CallTypes.acompletion not in discovery.mappings
|
||||
assert CallTypes.completion not in discovery.mappings
|
||||
assert CallTypes.aembedding in discovery.mappings
|
||||
|
||||
|
||||
def test_complete_discovery_reports_nothing_unavailable():
|
||||
discovery = llms_package.discover_guardrail_translations()
|
||||
|
||||
assert not discovery.unavailable
|
||||
assert CallTypes.acompletion in discovery.mappings
|
||||
|
||||
|
||||
def test_the_next_lookup_retries_a_package_that_failed_to_import():
|
||||
with unimportable(OPENAI_CHAT_TRANSLATION_MODULE):
|
||||
partial = llms_package.load_guardrail_translation_mappings()
|
||||
|
||||
assert CallTypes.acompletion not in partial
|
||||
assert CallTypes.aembedding in partial
|
||||
|
||||
recovered = llms_package.load_guardrail_translation_mappings()
|
||||
|
||||
assert CallTypes.acompletion in recovered
|
||||
assert CallTypes.completion in recovered
|
||||
assert CallTypes.aembedding in recovered
|
||||
assert not llms_package.guardrail_translation_discovery.unavailable
|
||||
|
||||
|
||||
def test_a_complete_discovery_is_not_scanned_again():
|
||||
healthy = llms_package.load_guardrail_translation_mappings()
|
||||
|
||||
assert CallTypes.acompletion in healthy
|
||||
|
||||
with unimportable(OPENAI_CHAT_TRANSLATION_MODULE):
|
||||
after_the_package_breaks = llms_package.load_guardrail_translation_mappings()
|
||||
|
||||
assert CallTypes.acompletion in after_the_package_breaks
|
||||
assert not llms_package.guardrail_translation_discovery.unavailable
|
||||
|
||||
|
||||
def test_a_package_that_keeps_failing_is_reported_once_and_its_recovery_announced(caplog):
|
||||
with capturing(caplog, logging.INFO), unimportable(OPENAI_CHAT_TRANSLATION_MODULE):
|
||||
for _ in range(3):
|
||||
llms_package.load_guardrail_translation_mappings()
|
||||
|
||||
errors = [record for record in caplog.records if record.levelno == logging.ERROR]
|
||||
assert len(errors) == 1, [record.getMessage() for record in errors]
|
||||
assert OPENAI_CHAT_TRANSLATION_MODULE in errors[0].getMessage()
|
||||
assert "None in sys.modules" in errors[0].getMessage()
|
||||
|
||||
caplog.clear()
|
||||
with capturing(caplog, logging.INFO):
|
||||
llms_package.load_guardrail_translation_mappings()
|
||||
llms_package.load_guardrail_translation_mappings()
|
||||
|
||||
recoveries = [record for record in caplog.records if "available again" in record.getMessage()]
|
||||
assert len(recoveries) == 1
|
||||
assert OPENAI_CHAT_TRANSLATION_MODULE in recoveries[0].getMessage()
|
||||
assert recoveries[0].levelno >= logging.WARNING
|
||||
assert not [record for record in caplog.records if record.levelno == logging.ERROR]
|
||||
|
||||
|
||||
def test_lookup_recovers_after_a_failed_discovery():
|
||||
with unimportable(OPENAI_CHAT_TRANSLATION_MODULE):
|
||||
with pytest.raises(ValueError, match="acompletion"):
|
||||
llms_package.get_guardrail_translation_mapping(CallTypes.acompletion)
|
||||
|
||||
assert llms_package.get_guardrail_translation_mapping(CallTypes.acompletion) is not None
|
||||
|
||||
|
||||
def test_an_mcp_package_that_fails_to_import_is_reported_and_retried():
|
||||
with raising_on_import(MCP_TRANSLATION_MODULE, AttributeError("module 'mcp.types' has no attribute 'ToolCall'")):
|
||||
partial = llms_package.load_guardrail_translation_mappings()
|
||||
|
||||
assert CallTypes.call_mcp_tool not in partial
|
||||
assert CallTypes.acompletion in partial
|
||||
assert tuple(llms_package.guardrail_translation_discovery.unavailable) == (MCP_TRANSLATION_MODULE,)
|
||||
assert "AttributeError" in llms_package.guardrail_translation_discovery.unavailable[MCP_TRANSLATION_MODULE]
|
||||
|
||||
recovered = llms_package.load_guardrail_translation_mappings()
|
||||
|
||||
assert CallTypes.call_mcp_tool in recovered
|
||||
assert CallTypes.acompletion in recovered
|
||||
assert not llms_package.guardrail_translation_discovery.unavailable
|
||||
|
||||
|
||||
def test_an_install_without_the_mcp_server_is_discovered_once_and_quietly(caplog):
|
||||
absent = ModuleNotFoundError("No module named 'mcp'", name="mcp")
|
||||
|
||||
with capturing(caplog, logging.DEBUG), unimportable("mcp"), raising_on_import(MCP_TRANSLATION_MODULE, absent):
|
||||
first = llms_package.load_guardrail_translation_mappings()
|
||||
second = llms_package.load_guardrail_translation_mappings()
|
||||
|
||||
assert first is second
|
||||
assert CallTypes.call_mcp_tool not in first
|
||||
assert CallTypes.acompletion in first
|
||||
assert not llms_package.guardrail_translation_discovery.unavailable
|
||||
assert not [record for record in caplog.records if record.levelno >= logging.WARNING]
|
||||
|
||||
|
||||
def test_a_broken_mcp_dependency_is_reported_and_retried(caplog):
|
||||
"""An mcp the install has but cannot import is a broken install, not a lean one, so it must be loud."""
|
||||
broken = ModuleNotFoundError("No module named 'mcp.types'", name="mcp.types")
|
||||
|
||||
with capturing(caplog, logging.DEBUG), raising_on_import(MCP_TRANSLATION_MODULE, broken):
|
||||
partial = llms_package.load_guardrail_translation_mappings()
|
||||
|
||||
assert CallTypes.call_mcp_tool not in partial
|
||||
assert CallTypes.acompletion in partial
|
||||
assert tuple(llms_package.guardrail_translation_discovery.unavailable) == (MCP_TRANSLATION_MODULE,)
|
||||
errors = [record for record in caplog.records if record.levelno >= logging.ERROR]
|
||||
assert len(errors) == 1, [record.getMessage() for record in caplog.records]
|
||||
assert "mcp.types" in errors[0].getMessage()
|
||||
|
||||
recovered = llms_package.load_guardrail_translation_mappings()
|
||||
|
||||
assert CallTypes.call_mcp_tool in recovered
|
||||
assert not llms_package.guardrail_translation_discovery.unavailable
|
||||
|
||||
|
||||
class _StandInMCPHandler:
|
||||
"""A bundled package's handler that claims the call type the MCP package owns."""
|
||||
|
||||
|
||||
def test_the_mcp_package_wins_a_handler_collision_with_a_bundled_package(monkeypatch):
|
||||
"""MCP handlers are the specialised ones, so a bundled package declaring the same call type must not shadow them."""
|
||||
colliding = ModuleType("litellm.llms.colliding_stub.guardrail_translation")
|
||||
colliding.guardrail_translation_mappings = {CallTypes.call_mcp_tool: _StandInMCPHandler}
|
||||
monkeypatch.setitem(sys.modules, colliding.__name__, colliding)
|
||||
monkeypatch.setattr(llms_package, "_bundled_guardrail_translation_modules", lambda: iter((colliding.__name__,)))
|
||||
|
||||
discovery = llms_package.discover_guardrail_translations()
|
||||
|
||||
assert discovery.mappings[CallTypes.call_mcp_tool] is MCPGuardrailTranslationHandler
|
||||
|
|
@ -133,8 +133,8 @@ def restore_callbacks(monkeypatch):
|
|||
monkeypatch.setattr(litellm, "callbacks", litellm.callbacks)
|
||||
monkeypatch.setattr(
|
||||
litellm_llms,
|
||||
"endpoint_guardrail_translation_mappings",
|
||||
litellm_llms.endpoint_guardrail_translation_mappings,
|
||||
"guardrail_translation_discovery",
|
||||
litellm_llms.guardrail_translation_discovery,
|
||||
)
|
||||
yield
|
||||
ProxyLogging._callback_capabilities_cache.clear()
|
||||
|
|
|
|||
13
tests/test_litellm/proxy/guardrails/conftest.py
Normal file
13
tests/test_litellm/proxy/guardrails/conftest.py
Normal file
|
|
@ -0,0 +1,13 @@
|
|||
import pytest
|
||||
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail import (
|
||||
unified_guardrail as unified_module,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _forget_unscanned_warnings():
|
||||
"""The unscanned warning fires once per route and reason, so every guardrail test starts with nothing remembered."""
|
||||
unified_module._warn_left_unscanned_once.cache_clear()
|
||||
yield
|
||||
unified_module._warn_left_unscanned_once.cache_clear()
|
||||
|
|
@ -1,14 +1,24 @@
|
|||
import pytest
|
||||
from unittest.mock import MagicMock, patch
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm.llms as llms_package
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_hooks.openai.moderations import (
|
||||
OpenAIModerationGuardrail,
|
||||
)
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
from litellm.types.utils import ModelResponseStream, ModelResponse
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.utils import (
|
||||
CallTypes,
|
||||
Delta,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
StreamingChoices,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -516,3 +526,74 @@ async def test_openai_moderation_streaming_sampled_when_end_of_stream_only_disab
|
|||
f"because chunk 6 already scanned the full text), "
|
||||
f"got {patched_make_request.await_count}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def reset_guardrail_translation_discovery():
|
||||
saved = llms_package.guardrail_translation_discovery
|
||||
llms_package.guardrail_translation_discovery = None
|
||||
yield llms_package
|
||||
llms_package.guardrail_translation_discovery = saved
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_moderation_still_runs_after_a_failed_translation_discovery(
|
||||
reset_guardrail_translation_discovery,
|
||||
):
|
||||
"""
|
||||
A guardrail translation discovery that could not import the chat handler must not silently
|
||||
disable moderation for the rest of the process.
|
||||
"""
|
||||
with pytest.MonkeyPatch.context() as poison:
|
||||
poison.setitem(sys.modules, "litellm.llms.openai.chat.guardrail_translation", None)
|
||||
poisoned = llms_package.load_guardrail_translation_mappings()
|
||||
assert CallTypes.acompletion not in poisoned
|
||||
|
||||
with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}):
|
||||
openai_guardrail = OpenAIModerationGuardrail(
|
||||
guardrail_name="test-openai-moderation",
|
||||
event_hook="post_call",
|
||||
)
|
||||
unified_guardrail = UnifiedLLMGuardrails()
|
||||
|
||||
mock_mod_response = MagicMock()
|
||||
mock_mod_response.results = []
|
||||
|
||||
async def mock_stream():
|
||||
chunks_data = ["Hello", " ", "world", "!", " Goodbye"]
|
||||
for i, content in enumerate(chunks_data):
|
||||
yield ModelResponseStream(
|
||||
model="gpt-4",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(content=content, role="assistant"),
|
||||
finish_reason=(
|
||||
"stop" if i == len(chunks_data) - 1 else None
|
||||
),
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
openai_guardrail, "async_make_request", return_value=mock_mod_response
|
||||
) as patched_make_request:
|
||||
chunks_received = 0
|
||||
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
api_key="test", request_route="/chat/completions"
|
||||
),
|
||||
response=mock_stream(),
|
||||
request_data={
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"guardrail_to_apply": openai_guardrail,
|
||||
"metadata": {"guardrails": ["test-openai-moderation"]},
|
||||
},
|
||||
):
|
||||
chunks_received += 1
|
||||
|
||||
assert chunks_received == 5
|
||||
assert patched_make_request.await_count > 0, (
|
||||
"Moderation never ran: the failed discovery was cached and the streaming hook "
|
||||
"passed every chunk through unscanned"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
"""Tests for unified guardrail."""
|
||||
|
||||
import logging
|
||||
from contextlib import contextmanager
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -2289,3 +2290,250 @@ class TestTranslationMappingsAreReadLive:
|
|||
for name, value in vars(unified_module).items()
|
||||
if isinstance(value, dict) and CallTypes.aocr in value
|
||||
]
|
||||
|
||||
|
||||
class TestUnscannedStreamIsAnnounced:
|
||||
"""The streaming hook must never forward a whole response unscanned without saying so."""
|
||||
|
||||
@staticmethod
|
||||
async def _drive(caplog, monkeypatch, request_route, mappings, response_chunks):
|
||||
_patch_translation_mappings(monkeypatch, mappings)
|
||||
|
||||
async def stream():
|
||||
for chunk in response_chunks:
|
||||
yield chunk
|
||||
|
||||
caplog.set_level(logging.WARNING, logger="LiteLLM Proxy")
|
||||
unified_module.verbose_proxy_logger.addHandler(caplog.handler)
|
||||
try:
|
||||
chunks = [
|
||||
chunk
|
||||
async for chunk in UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
api_key="test", request_route=request_route
|
||||
),
|
||||
response=stream(),
|
||||
request_data={
|
||||
"guardrail_to_apply": RecordingGuardrail(),
|
||||
"metadata": {"guardrails": ["recording-guardrail"]},
|
||||
},
|
||||
)
|
||||
]
|
||||
finally:
|
||||
unified_module.verbose_proxy_logger.removeHandler(caplog.handler)
|
||||
|
||||
return chunks, [
|
||||
record.getMessage()
|
||||
for record in caplog.records
|
||||
if record.levelno >= logging.WARNING
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_warns_when_the_call_type_has_no_translation_handler(
|
||||
self, caplog, monkeypatch
|
||||
):
|
||||
chunks, warnings = await self._drive(
|
||||
caplog,
|
||||
monkeypatch,
|
||||
request_route="/chat/completions",
|
||||
mappings={CallTypes.aembedding: _NoopTranslation},
|
||||
response_chunks=[
|
||||
ModelResponseStream(
|
||||
choices=[StreamingChoices(index=0, delta=Delta(content=content))]
|
||||
)
|
||||
for content in ("a", "b", "c")
|
||||
],
|
||||
)
|
||||
|
||||
assert len(chunks) == 3
|
||||
assert any(
|
||||
"no guardrail translation handler" in message
|
||||
and "Add a guardrail translation handler for that call type." in message
|
||||
and "recording-guardrail" in message
|
||||
and "/chat/completions" in message
|
||||
for message in warnings
|
||||
), warnings
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_warns_when_the_call_type_cannot_be_resolved(self, caplog, monkeypatch):
|
||||
chunks, warnings = await self._drive(
|
||||
caplog,
|
||||
monkeypatch,
|
||||
request_route="/v1/not-a-mapped-route",
|
||||
mappings=load_guardrail_translation_mappings(),
|
||||
response_chunks=[{"event": "delta", "text": content} for content in ("a", "b", "c")],
|
||||
)
|
||||
|
||||
assert len(chunks) == 3
|
||||
assert any(
|
||||
"call type could not be resolved" in message
|
||||
and "Add the route to API_ROUTE_TO_CALL_TYPES." in message
|
||||
and "recording-guardrail" in message
|
||||
for message in warnings
|
||||
), warnings
|
||||
|
||||
|
||||
class TestUnscannedRequestIsAnnounced:
|
||||
"""A request hook that cannot scan must say so instead of passing the request through in silence."""
|
||||
|
||||
@staticmethod
|
||||
def _request(guardrail):
|
||||
return {
|
||||
"guardrail_to_apply": guardrail,
|
||||
"model": "gpt-4o",
|
||||
"messages": [{"role": "user", "content": "hello world"}],
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
@contextmanager
|
||||
def _capturing(caplog):
|
||||
caplog.set_level(logging.WARNING, logger="LiteLLM Proxy")
|
||||
unified_module.verbose_proxy_logger.addHandler(caplog.handler)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
unified_module.verbose_proxy_logger.removeHandler(caplog.handler)
|
||||
|
||||
@staticmethod
|
||||
def _warnings(caplog):
|
||||
return [record.getMessage() for record in caplog.records if record.levelno >= logging.WARNING]
|
||||
|
||||
@staticmethod
|
||||
def _distinct_warnings(caplog):
|
||||
"""caplog holds every record twice here, once through the handler above and once through propagation."""
|
||||
by_record = {id(record): record for record in caplog.records if record.levelno >= logging.WARNING}
|
||||
return [record.getMessage() for record in by_record.values()]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_warns_when_the_call_type_has_no_translation_handler(self, caplog, monkeypatch):
|
||||
_patch_translation_mappings(monkeypatch, {CallTypes.aembedding: _NoopTranslation})
|
||||
guardrail = RecordingGuardrail()
|
||||
data = self._request(guardrail)
|
||||
|
||||
with self._capturing(caplog):
|
||||
returned = await UnifiedLLMGuardrails().async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/moderations"),
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type=CallTypes.acompletion.value,
|
||||
)
|
||||
|
||||
assert guardrail.apply_calls == []
|
||||
assert returned["messages"] == [{"role": "user", "content": "hello world"}]
|
||||
assert any(
|
||||
"no guardrail translation handler" in message
|
||||
and "skipping pre-call scanning" in message
|
||||
and "recording-guardrail" in message
|
||||
and "/v1/moderations" in message
|
||||
and "acompletion" in message
|
||||
for message in self._warnings(caplog)
|
||||
), self._warnings(caplog)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_during_call_warns_when_the_call_type_has_no_translation_handler(self, caplog, monkeypatch):
|
||||
_patch_translation_mappings(monkeypatch, {CallTypes.aembedding: _NoopTranslation})
|
||||
guardrail = RecordingGuardrail()
|
||||
data = self._request(guardrail)
|
||||
|
||||
with self._capturing(caplog):
|
||||
returned = await UnifiedLLMGuardrails().async_moderation_hook(
|
||||
data=data,
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/moderations"),
|
||||
call_type=CallTypes.acompletion.value,
|
||||
)
|
||||
|
||||
assert guardrail.apply_calls == []
|
||||
assert returned["messages"] == [{"role": "user", "content": "hello world"}]
|
||||
assert any(
|
||||
"skipping during-call scanning" in message and "recording-guardrail" in message
|
||||
for message in self._warnings(caplog)
|
||||
), self._warnings(caplog)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_warns_instead_of_raising_on_a_call_type_outside_the_enum(self, caplog, monkeypatch):
|
||||
_patch_translation_mappings(monkeypatch, load_guardrail_translation_mappings())
|
||||
guardrail = RecordingGuardrail()
|
||||
data = self._request(guardrail)
|
||||
|
||||
with self._capturing(caplog):
|
||||
returned = await UnifiedLLMGuardrails().async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/moderations"),
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type="not_a_call_type",
|
||||
)
|
||||
|
||||
assert guardrail.apply_calls == []
|
||||
assert returned["messages"] == [{"role": "user", "content": "hello world"}]
|
||||
assert any(
|
||||
"call type 'not_a_call_type' is not one litellm can scan" in message
|
||||
and "Map the route to a CallTypes member in API_ROUTE_TO_CALL_TYPES." in message
|
||||
and "skipping pre-call scanning" in message
|
||||
for message in self._warnings(caplog)
|
||||
), self._warnings(caplog)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_during_call_warns_instead_of_raising_on_a_call_type_outside_the_enum(self, caplog, monkeypatch):
|
||||
_patch_translation_mappings(monkeypatch, load_guardrail_translation_mappings())
|
||||
guardrail = RecordingGuardrail()
|
||||
data = self._request(guardrail)
|
||||
|
||||
with self._capturing(caplog):
|
||||
returned = await UnifiedLLMGuardrails().async_moderation_hook(
|
||||
data=data,
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/moderations"),
|
||||
call_type="not_a_call_type",
|
||||
)
|
||||
|
||||
assert guardrail.apply_calls == []
|
||||
assert returned["messages"] == [{"role": "user", "content": "hello world"}]
|
||||
assert any(
|
||||
"call type 'not_a_call_type' is not one litellm can scan" in message
|
||||
and "Map the route to a CallTypes member in API_ROUTE_TO_CALL_TYPES." in message
|
||||
and "skipping during-call scanning" in message
|
||||
for message in self._warnings(caplog)
|
||||
), self._warnings(caplog)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_route_with_no_handler_warns_once_instead_of_once_per_request(self, caplog, monkeypatch):
|
||||
_patch_translation_mappings(monkeypatch, {CallTypes.aembedding: _NoopTranslation})
|
||||
guardrail = RecordingGuardrail()
|
||||
|
||||
with self._capturing(caplog):
|
||||
for _ in range(5):
|
||||
await UnifiedLLMGuardrails().async_moderation_hook(
|
||||
data=self._request(guardrail),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/responses/resp_123"),
|
||||
call_type="aget_responses",
|
||||
)
|
||||
await UnifiedLLMGuardrails().async_moderation_hook(
|
||||
data=self._request(guardrail),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/moderations"),
|
||||
call_type="aget_responses",
|
||||
)
|
||||
|
||||
unscanned = [
|
||||
message for message in self._distinct_warnings(caplog) if "skipping during-call scanning" in message
|
||||
]
|
||||
assert len(unscanned) == 2, unscanned
|
||||
assert "/v1/responses/resp_123" in unscanned[0]
|
||||
assert "aget_responses" in unscanned[0]
|
||||
assert "/v1/moderations" in unscanned[1]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_names_the_call_type_the_route_maps_to(self, caplog, monkeypatch):
|
||||
_patch_translation_mappings(monkeypatch, {CallTypes.aembedding: _NoopTranslation})
|
||||
guardrail = RecordingGuardrail()
|
||||
|
||||
with self._capturing(caplog):
|
||||
await UnifiedLLMGuardrails().async_post_call_success_hook(
|
||||
data={"guardrail_to_apply": guardrail, "model": "gpt-4o"},
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/responses"),
|
||||
response=litellm.ModelResponse(),
|
||||
)
|
||||
|
||||
assert guardrail.apply_calls == []
|
||||
unscanned = [message for message in self._warnings(caplog) if "skipping post-call scanning" in message]
|
||||
assert unscanned, self._warnings(caplog)
|
||||
assert "call type 'aresponses'" in unscanned[0], unscanned[0]
|
||||
assert "CallTypes." not in unscanned[0], unscanned[0]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue