mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
* add _aguardrail_helper for LB * add _aguardrail_helper on router.py * test_proxy_logging_pre_call_hook_load_balancing * add _execute_guardrail_with_load_balancing * add LB TEsting * docs guard lb * fix linting * fix lint
134 lines
4.5 KiB
Python
134 lines
4.5 KiB
Python
from typing import Any, Dict, List, Literal, Optional, Union
|
|
|
|
import litellm
|
|
from litellm._logging import verbose_proxy_logger
|
|
from litellm.caching.caching import DualCache
|
|
from litellm.integrations.custom_guardrail import CustomGuardrail
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
from litellm.proxy.guardrails.guardrail_helpers import should_proceed_based_on_metadata
|
|
from litellm.types.utils import CallTypesLiteral
|
|
|
|
# Global counter for tracking which guardrail was called (for load balancing tests)
|
|
guardrail_lb_call_count: Dict[str, int] = {"A": 0, "B": 0}
|
|
|
|
|
|
class GuardrailForLBTestingA(CustomGuardrail):
|
|
"""Guardrail A for load balancing testing."""
|
|
|
|
async def async_pre_call_hook(
|
|
self,
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
cache: DualCache,
|
|
data: dict,
|
|
call_type: CallTypesLiteral,
|
|
) -> Optional[Union[Exception, str, dict]]:
|
|
guardrail_lb_call_count["A"] += 1
|
|
verbose_proxy_logger.info(
|
|
f"GuardrailForLBTestingA called. Total A calls: {guardrail_lb_call_count['A']}"
|
|
)
|
|
return data
|
|
|
|
|
|
class GuardrailForLBTestingB(CustomGuardrail):
|
|
"""Guardrail B for load balancing testing."""
|
|
|
|
async def async_pre_call_hook(
|
|
self,
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
cache: DualCache,
|
|
data: dict,
|
|
call_type: CallTypesLiteral,
|
|
) -> Optional[Union[Exception, str, dict]]:
|
|
guardrail_lb_call_count["B"] += 1
|
|
verbose_proxy_logger.info(
|
|
f"GuardrailForLBTestingB called. Total B calls: {guardrail_lb_call_count['B']}"
|
|
)
|
|
return data
|
|
|
|
|
|
class myCustomGuardrail(CustomGuardrail):
|
|
def __init__(
|
|
self,
|
|
**kwargs,
|
|
):
|
|
# store kwargs as optional_params
|
|
self.optional_params = kwargs
|
|
|
|
super().__init__(**kwargs)
|
|
|
|
async def async_pre_call_hook(
|
|
self,
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
cache: DualCache,
|
|
data: dict,
|
|
call_type: CallTypesLiteral,
|
|
) -> Optional[Union[Exception, str, dict]]:
|
|
"""
|
|
Runs before the LLM API call
|
|
Runs on only Input
|
|
Use this if you want to MODIFY the input
|
|
"""
|
|
|
|
# In this guardrail, if a user inputs `litellm` we will mask it and then send it to the LLM
|
|
_messages = data.get("messages")
|
|
if _messages:
|
|
for message in _messages:
|
|
_content = message.get("content")
|
|
if isinstance(_content, str):
|
|
if "litellm" in _content.lower():
|
|
_content = _content.replace("litellm", "********")
|
|
message["content"] = _content
|
|
|
|
verbose_proxy_logger.debug(
|
|
"async_pre_call_hook: Message after masking %s", _messages
|
|
)
|
|
|
|
return data
|
|
|
|
async def async_moderation_hook(
|
|
self,
|
|
data: dict,
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
call_type: CallTypesLiteral,
|
|
):
|
|
"""
|
|
Runs in parallel to LLM API call
|
|
Runs on only Input
|
|
|
|
This can NOT modify the input, only used to reject or accept a call before going to LLM API
|
|
"""
|
|
|
|
# this works the same as async_pre_call_hook, but just runs in parallel as the LLM API Call
|
|
# In this guardrail, if a user inputs `litellm` we will mask it.
|
|
_messages = data.get("messages")
|
|
if _messages:
|
|
for message in _messages:
|
|
_content = message.get("content")
|
|
if isinstance(_content, str):
|
|
if "litellm" in _content.lower():
|
|
raise ValueError("Guardrail failed words - `litellm` detected")
|
|
|
|
async def async_post_call_success_hook(
|
|
self,
|
|
data: dict,
|
|
user_api_key_dict: UserAPIKeyAuth,
|
|
response,
|
|
):
|
|
"""
|
|
Runs on response from LLM API call
|
|
|
|
It can be used to reject a response
|
|
|
|
If a response contains the word "coffee" -> we will raise an exception
|
|
"""
|
|
verbose_proxy_logger.debug("async_pre_call_hook response: %s", response)
|
|
if isinstance(response, litellm.ModelResponse):
|
|
for choice in response.choices:
|
|
if isinstance(choice, litellm.Choices):
|
|
verbose_proxy_logger.debug("async_pre_call_hook choice: %s", choice)
|
|
if (
|
|
choice.message.content
|
|
and isinstance(choice.message.content, str)
|
|
and "coffee" in choice.message.content
|
|
):
|
|
raise ValueError("Guardrail failed Coffee Detected")
|