mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
feat(guardrails): wire apply_guardrail into proxy logging callbacks (#28970)
* feat(guardrails): wire apply_guardrail into proxy logging callbacks Route /apply_guardrail through pre/post proxy hooks and LiteLLM success/failure handlers so Langfuse and OTEL integrations receive input/output on guardrail-only requests. Co-authored-by: Cursor <cursoragent@cursor.com> * fix(guardrails): fix Greptile review comments on apply_guardrail logging Co-authored-by: Cursor <cursoragent@cursor.com> * fix(apply_guardrail): preserve original exception and capture modified response - Capture return value from post_call_success_hook so callback-modified responses propagate to the caller. - Wrap success/failure logging calls in defensive try/except so logging infrastructure failures don't replace the user-visible response or mask the original guardrail exception. Co-authored-by: Yassin Kortam <yassin@berri.ai> * Fix mypy * fix(apply_guardrail): isolate failure logging and use post-hook response for logging - Split async_failure_handler and post_call_failure_hook into independent try/except blocks so a callback bug in one does not silently skip the other. - Build response_for_logging inside _emit_guardrail_success_logs after post_call_success_hook runs, so logged data matches the response the caller actually receives when the hook modifies the response. Co-authored-by: Yassin Kortam <yassin@berri.ai> * fix(apply_guardrail): fix black formatting and update tests for fastapi_request param - Run black on guardrail_endpoints.py to fix CI formatting check - Add _mock_proxy_logging() helper to enterprise guardrail tests to patch proxy-server globals imported at call time - Pass fastapi_request=Mock() in all direct apply_guardrail test calls to match updated function signature Co-authored-by: Cursor <cursoragent@cursor.com> * fix(guardrails): use transformed exception from post_call_failure_hook in apply_guardrail Co-authored-by: Yassin Kortam <yassin@berri.ai> * fix(guardrails): isolate sync/async logging handlers in apply_guardrail Separate each logging handler call into its own try/except so a failure in the async handler does not silently skip the sync handler submission (and vice versa). Matches the docstring's defensive intent. Co-authored-by: Yassin Kortam <yassin@berri.ai> * fix(apply_guardrail): guard transformed_exception with isinstance check Co-authored-by: Cursor <cursoragent@cursor.com> * test(guardrails): mock proxy globals in not_found test and share apply_guardrail logging fixture - Add proxy-server global mocks to test_apply_guardrail_not_found so the failure-path post_call_failure_hook call doesn't touch the real proxy logging singleton. - Extract the duplicated _mock_proxy_logging context manager out of the two enterprise apply_guardrail test files into a shared conftest fixture so the helper stays in one place. * fix(guardrails): use update_messages to keep logging obj in sync Co-authored-by: Yassin Kortam <yassin@berri.ai> --------- Co-authored-by: Cursor <cursoragent@cursor.com> Co-authored-by: Yassin Kortam <yassin@berri.ai> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
95015de733
commit
5bd59b33e6
6 changed files with 362 additions and 47 deletions
|
|
@ -839,6 +839,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"aget_run",
|
||||
"acancel_run",
|
||||
"adelete_run",
|
||||
"apply_guardrail",
|
||||
],
|
||||
version: Optional[str] = None,
|
||||
user_model: Optional[str] = None,
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ from datetime import datetime, timezone
|
|||
from typing import Any, Dict, List, Literal, Optional, Type, TypeVar, Union, cast
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.proxy.common_utils.path_utils import safe_join
|
||||
|
|
@ -2187,9 +2187,97 @@ async def test_custom_code_guardrail(
|
|||
)
|
||||
|
||||
|
||||
def _resolve_guardrail_input_type(
|
||||
active_guardrail: CustomGuardrail, input_type: str
|
||||
) -> Literal["request", "response"]:
|
||||
"""Return the effective input_type, auto-upgrading to 'response' for post_call guardrails."""
|
||||
if input_type == "request":
|
||||
hook = getattr(active_guardrail, "event_hook", None)
|
||||
if hook == GuardrailEventHooks.post_call or hook == "post_call":
|
||||
return "response"
|
||||
return "response" if input_type == "response" else "request"
|
||||
|
||||
|
||||
def _patch_logging_obj_for_guardrail(
|
||||
litellm_logging_obj: Any, request: ApplyGuardrailRequest
|
||||
) -> None:
|
||||
"""Configure the logging object so Langfuse/OTEL extract input and output correctly."""
|
||||
litellm_logging_obj.call_type = "pass_through_endpoint"
|
||||
litellm_logging_obj.model_call_details["call_type"] = "pass_through_endpoint"
|
||||
litellm_logging_obj.update_messages(
|
||||
request.messages
|
||||
if request.messages
|
||||
else [{"role": "user", "content": request.text}]
|
||||
)
|
||||
|
||||
|
||||
async def _emit_guardrail_success_logs(
|
||||
proxy_logging_obj: Any,
|
||||
litellm_logging_obj: Any,
|
||||
data: dict,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: ApplyGuardrailResponse,
|
||||
start_time: datetime,
|
||||
) -> ApplyGuardrailResponse:
|
||||
"""Fire proxy and LiteLLM success hooks after a successful guardrail run.
|
||||
|
||||
Each hook is wrapped defensively so a callback failure never prevents the
|
||||
caller from receiving the guardrail response. Returns the (possibly
|
||||
hook-modified) response.
|
||||
"""
|
||||
from litellm.litellm_core_utils.thread_pool_executor import (
|
||||
executor as thread_pool_executor,
|
||||
)
|
||||
|
||||
try:
|
||||
modified = await proxy_logging_obj.post_call_success_hook(
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=response,
|
||||
)
|
||||
if isinstance(modified, ApplyGuardrailResponse):
|
||||
response = modified
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception("apply_guardrail: post_call_success_hook failed")
|
||||
|
||||
# Build the logging payload after post_call_success_hook so that logged
|
||||
# data matches what the caller actually receives if the hook modified
|
||||
# the response.
|
||||
response_for_logging = {"response": response.model_dump(exclude_none=True)}
|
||||
|
||||
if litellm_logging_obj is not None:
|
||||
end_time = datetime.now(timezone.utc)
|
||||
try:
|
||||
await litellm_logging_obj.async_success_handler(
|
||||
result=response_for_logging,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
cache_hit=False,
|
||||
)
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception(
|
||||
"apply_guardrail: async_success_handler failed"
|
||||
)
|
||||
try:
|
||||
thread_pool_executor.submit(
|
||||
litellm_logging_obj.success_handler,
|
||||
response_for_logging,
|
||||
start_time,
|
||||
end_time,
|
||||
False,
|
||||
)
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception(
|
||||
"apply_guardrail: success_handler submit failed"
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
@router.post("/guardrails/apply_guardrail", response_model=ApplyGuardrailResponse)
|
||||
@router.post("/apply_guardrail", response_model=ApplyGuardrailResponse)
|
||||
async def apply_guardrail(
|
||||
fastapi_request: Request,
|
||||
request: ApplyGuardrailRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
|
|
@ -2198,8 +2286,29 @@ async def apply_guardrail(
|
|||
|
||||
This endpoint allows testing guardrails by applying them to custom text inputs.
|
||||
"""
|
||||
import traceback
|
||||
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.litellm_core_utils.thread_pool_executor import (
|
||||
executor as thread_pool_executor,
|
||||
)
|
||||
from litellm.proxy.proxy_server import (
|
||||
general_settings,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
version,
|
||||
)
|
||||
from litellm.proxy.utils import handle_exception_on_proxy
|
||||
|
||||
data: dict = {
|
||||
"guardrail_name": request.guardrail_name,
|
||||
"input": [request.text],
|
||||
"messages": request.messages or [],
|
||||
"metadata": {"route": "/apply_guardrail"},
|
||||
}
|
||||
litellm_logging_obj = None
|
||||
start_time = datetime.now(timezone.utc)
|
||||
|
||||
try:
|
||||
active_guardrail: Optional[CustomGuardrail] = (
|
||||
GUARDRAIL_REGISTRY.get_initialized_guardrail_callback(
|
||||
|
|
@ -2212,23 +2321,25 @@ async def apply_guardrail(
|
|||
detail=f"Guardrail '{request.guardrail_name}' not found. Please ensure the guardrail is configured in your LiteLLM proxy.",
|
||||
)
|
||||
|
||||
request_data: dict = {}
|
||||
if request.messages:
|
||||
request_data["messages"] = request.messages
|
||||
request_processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
data, litellm_logging_obj = (
|
||||
await request_processor.common_processing_pre_call_logic(
|
||||
request=fastapi_request,
|
||||
general_settings=general_settings,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
version=version,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
proxy_config=proxy_config,
|
||||
route_type="apply_guardrail",
|
||||
)
|
||||
)
|
||||
|
||||
# Auto-detect input_type: if the caller didn't specify "response" but the
|
||||
# guardrail only runs post_call (e.g. LLM-as-a-judge), use "response" so
|
||||
# the test actually exercises the guardrail logic.
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
if litellm_logging_obj is not None:
|
||||
_patch_logging_obj_for_guardrail(litellm_logging_obj, request)
|
||||
|
||||
resolved_input_type = request.input_type
|
||||
if resolved_input_type == "request":
|
||||
hook = getattr(active_guardrail, "event_hook", None)
|
||||
if hook == GuardrailEventHooks.post_call or hook == "post_call":
|
||||
resolved_input_type = "response"
|
||||
|
||||
_input_type: Literal["request", "response"] = (
|
||||
"response" if resolved_input_type == "response" else "request"
|
||||
request_data: dict = {"messages": request.messages} if request.messages else {}
|
||||
_input_type = _resolve_guardrail_input_type(
|
||||
active_guardrail, request.input_type
|
||||
)
|
||||
guardrailed_inputs = await active_guardrail.apply_guardrail(
|
||||
inputs={"texts": [request.text]},
|
||||
|
|
@ -2236,13 +2347,55 @@ async def apply_guardrail(
|
|||
input_type=_input_type,
|
||||
)
|
||||
response_text = guardrailed_inputs.get("texts", [])
|
||||
|
||||
return ApplyGuardrailResponse(
|
||||
response = ApplyGuardrailResponse(
|
||||
response_text=response_text[0] if response_text else request.text
|
||||
)
|
||||
except Exception as e:
|
||||
if litellm_logging_obj is not None and not isinstance(e, HTTPException):
|
||||
try:
|
||||
await litellm_logging_obj.async_failure_handler(
|
||||
exception=e,
|
||||
traceback_exception=traceback.format_exc(),
|
||||
)
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception(
|
||||
"apply_guardrail: async_failure_handler failed"
|
||||
)
|
||||
try:
|
||||
thread_pool_executor.submit(
|
||||
litellm_logging_obj.failure_handler,
|
||||
e,
|
||||
traceback.format_exc(),
|
||||
)
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception(
|
||||
"apply_guardrail: failure_handler submit failed"
|
||||
)
|
||||
try:
|
||||
transformed_exception = await proxy_logging_obj.post_call_failure_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
original_exception=e,
|
||||
request_data=data,
|
||||
)
|
||||
if isinstance(transformed_exception, Exception):
|
||||
e = transformed_exception
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception(
|
||||
"apply_guardrail: post_call_failure_hook failed"
|
||||
)
|
||||
raise handle_exception_on_proxy(e)
|
||||
|
||||
# Success logging outside except so a hook error never triggers failure handlers.
|
||||
response = await _emit_guardrail_success_logs(
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=response,
|
||||
start_time=start_time,
|
||||
)
|
||||
return response
|
||||
|
||||
|
||||
# Usage (dashboard) endpoints: overview, detail, logs
|
||||
router.include_router(guardrails_usage_router)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,42 @@
|
|||
"""Shared fixtures for guardrail apply_guardrail tests."""
|
||||
|
||||
from contextlib import contextmanager
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _mock_proxy_logging():
|
||||
"""Patch the proxy-server globals that apply_guardrail imports at call time."""
|
||||
mock_proxy_logging = MagicMock()
|
||||
mock_proxy_logging.post_call_success_hook = AsyncMock(return_value=None)
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock(return_value=None)
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.async_success_handler = AsyncMock(return_value=None)
|
||||
mock_logging_obj.async_failure_handler = AsyncMock(return_value=None)
|
||||
mock_logging_obj.success_handler = MagicMock(return_value=None)
|
||||
mock_logging_obj.failure_handler = MagicMock(return_value=None)
|
||||
mock_logging_obj.model_call_details = {}
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing"
|
||||
) as mock_proc_cls,
|
||||
patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging),
|
||||
patch("litellm.proxy.proxy_server.general_settings", {}),
|
||||
patch("litellm.proxy.proxy_server.proxy_config", MagicMock()),
|
||||
patch("litellm.proxy.proxy_server.version", "0.0.0"),
|
||||
):
|
||||
mock_proc = MagicMock()
|
||||
mock_proc.common_processing_pre_call_logic = AsyncMock(
|
||||
return_value=({}, mock_logging_obj)
|
||||
)
|
||||
mock_proc_cls.return_value = mock_proc
|
||||
yield mock_proxy_logging
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_proxy_logging_ctx():
|
||||
"""Return the proxy-logging context manager factory for use as `with ctx():`."""
|
||||
return _mock_proxy_logging
|
||||
|
|
@ -18,14 +18,19 @@ from litellm.types.guardrails import ApplyGuardrailRequest, ApplyGuardrailRespon
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_endpoint_returns_correct_response():
|
||||
async def test_apply_guardrail_endpoint_returns_correct_response(
|
||||
mock_proxy_logging_ctx,
|
||||
):
|
||||
"""Test that apply_guardrail endpoint returns ApplyGuardrailResponse object"""
|
||||
from litellm.proxy.guardrails.guardrail_endpoints import apply_guardrail
|
||||
|
||||
# Mock the guardrail registry
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
|
||||
) as mock_registry:
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
|
||||
) as mock_registry,
|
||||
mock_proxy_logging_ctx(),
|
||||
):
|
||||
# Create a mock guardrail
|
||||
mock_guardrail = Mock(spec=CustomGuardrail)
|
||||
# Apply guardrail returns GenericGuardrailAPIInputs (dict with texts key)
|
||||
|
|
@ -49,7 +54,9 @@ async def test_apply_guardrail_endpoint_returns_correct_response():
|
|||
|
||||
# Call the endpoint
|
||||
response = await apply_guardrail(
|
||||
request=request, user_api_key_dict=user_api_key_dict
|
||||
fastapi_request=Mock(),
|
||||
request=request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify the response is of the correct type
|
||||
|
|
@ -65,15 +72,18 @@ async def test_apply_guardrail_endpoint_returns_correct_response():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_endpoint_guardrail_not_found():
|
||||
async def test_apply_guardrail_endpoint_guardrail_not_found(mock_proxy_logging_ctx):
|
||||
"""Test that apply_guardrail endpoint raises exception when guardrail not found"""
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.guardrails.guardrail_endpoints import apply_guardrail
|
||||
|
||||
# Mock the guardrail registry to return None
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
|
||||
) as mock_registry:
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
|
||||
) as mock_registry,
|
||||
mock_proxy_logging_ctx(),
|
||||
):
|
||||
mock_registry.get_initialized_guardrail_callback.return_value = None
|
||||
|
||||
# Create the request
|
||||
|
|
@ -86,26 +96,35 @@ async def test_apply_guardrail_endpoint_guardrail_not_found():
|
|||
|
||||
# Verify exception is raised
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await apply_guardrail(request=request, user_api_key_dict=user_api_key_dict)
|
||||
await apply_guardrail(
|
||||
fastapi_request=Mock(),
|
||||
request=request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
assert "non-existent-guardrail" in exc_info.value.message
|
||||
assert "not found" in exc_info.value.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_endpoint_with_presidio_guardrail():
|
||||
async def test_apply_guardrail_endpoint_with_presidio_guardrail(mock_proxy_logging_ctx):
|
||||
"""Test apply_guardrail endpoint with a Presidio-like guardrail"""
|
||||
from litellm.proxy.guardrails.guardrail_endpoints import apply_guardrail
|
||||
|
||||
# Mock the guardrail registry
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
|
||||
) as mock_registry:
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
|
||||
) as mock_registry,
|
||||
mock_proxy_logging_ctx(),
|
||||
):
|
||||
# Create a mock guardrail that simulates Presidio behavior
|
||||
mock_guardrail = Mock(spec=CustomGuardrail)
|
||||
# Simulate masking PII entities - returns GenericGuardrailAPIInputs (dict with texts key)
|
||||
mock_guardrail.apply_guardrail = AsyncMock(
|
||||
return_value={"texts": ["My name is [PERSON] and my email is [EMAIL_ADDRESS]"]}
|
||||
return_value={
|
||||
"texts": ["My name is [PERSON] and my email is [EMAIL_ADDRESS]"]
|
||||
}
|
||||
)
|
||||
|
||||
# Configure the registry to return our mock guardrail
|
||||
|
|
@ -124,7 +143,9 @@ async def test_apply_guardrail_endpoint_with_presidio_guardrail():
|
|||
|
||||
# Call the endpoint
|
||||
response = await apply_guardrail(
|
||||
request=request, user_api_key_dict=user_api_key_dict
|
||||
fastapi_request=Mock(),
|
||||
request=request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify the response is of the correct type
|
||||
|
|
@ -138,14 +159,17 @@ async def test_apply_guardrail_endpoint_with_presidio_guardrail():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_endpoint_without_optional_params():
|
||||
async def test_apply_guardrail_endpoint_without_optional_params(mock_proxy_logging_ctx):
|
||||
"""Test apply_guardrail endpoint without optional language and entities parameters"""
|
||||
from litellm.proxy.guardrails.guardrail_endpoints import apply_guardrail
|
||||
|
||||
# Mock the guardrail registry
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
|
||||
) as mock_registry:
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
|
||||
) as mock_registry,
|
||||
mock_proxy_logging_ctx(),
|
||||
):
|
||||
# Create a mock guardrail
|
||||
mock_guardrail = Mock(spec=CustomGuardrail)
|
||||
# Returns GenericGuardrailAPIInputs (dict with texts key)
|
||||
|
|
@ -166,7 +190,9 @@ async def test_apply_guardrail_endpoint_without_optional_params():
|
|||
|
||||
# Call the endpoint
|
||||
response = await apply_guardrail(
|
||||
request=request, user_api_key_dict=user_api_key_dict
|
||||
fastapi_request=Mock(),
|
||||
request=request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify the response is of the correct type
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ Test the Bedrock guardrail apply_guardrail functionality
|
|||
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, patch
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -153,7 +153,7 @@ async def test_bedrock_apply_guardrail_api_failure():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_apply_guardrail_endpoint_integration():
|
||||
async def test_bedrock_apply_guardrail_endpoint_integration(mock_proxy_logging_ctx):
|
||||
"""Test the full endpoint integration with Bedrock guardrail"""
|
||||
from litellm.proxy.guardrails.guardrail_endpoints import apply_guardrail
|
||||
|
||||
|
|
@ -165,9 +165,12 @@ async def test_bedrock_apply_guardrail_endpoint_integration():
|
|||
)
|
||||
|
||||
# Mock the guardrail registry
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
|
||||
) as mock_registry:
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
|
||||
) as mock_registry,
|
||||
mock_proxy_logging_ctx(),
|
||||
):
|
||||
# Mock the make_bedrock_api_request method
|
||||
with patch.object(
|
||||
guardrail, "make_bedrock_api_request", new_callable=AsyncMock
|
||||
|
|
@ -194,7 +197,9 @@ async def test_bedrock_apply_guardrail_endpoint_integration():
|
|||
|
||||
# Call the endpoint
|
||||
response = await apply_guardrail(
|
||||
request=request, user_api_key_dict=user_api_key_dict
|
||||
fastapi_request=Mock(),
|
||||
request=request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify the response
|
||||
|
|
|
|||
|
|
@ -1149,6 +1149,13 @@ async def test_apply_guardrail_not_found(mocker):
|
|||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_registry
|
||||
)
|
||||
|
||||
mock_proxy_logging = mocker.Mock()
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mocker.patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging)
|
||||
mocker.patch("litellm.proxy.proxy_server.general_settings", {})
|
||||
mocker.patch("litellm.proxy.proxy_server.proxy_config", mocker.Mock())
|
||||
mocker.patch("litellm.proxy.proxy_server.version", "test")
|
||||
|
||||
# Create request
|
||||
request = ApplyGuardrailRequest(
|
||||
guardrail_name="non-existent-guardrail", text="Test input text"
|
||||
|
|
@ -1159,7 +1166,11 @@ async def test_apply_guardrail_not_found(mocker):
|
|||
|
||||
# Call endpoint and expect ProxyException
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await apply_guardrail(request=request, user_api_key_dict=mock_user_auth)
|
||||
await apply_guardrail(
|
||||
fastapi_request=mocker.Mock(),
|
||||
request=request,
|
||||
user_api_key_dict=mock_user_auth,
|
||||
)
|
||||
|
||||
# Verify error details
|
||||
assert str(exc_info.value.code) == "404"
|
||||
|
|
@ -1186,6 +1197,25 @@ async def test_apply_guardrail_execution_error(mocker):
|
|||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_registry
|
||||
)
|
||||
|
||||
mock_logging_obj = mocker.Mock()
|
||||
mock_logging_obj.async_failure_handler = AsyncMock()
|
||||
mock_logging_obj.model_call_details = {}
|
||||
mock_processor = mocker.Mock()
|
||||
mock_processor.common_processing_pre_call_logic = AsyncMock(
|
||||
return_value=({"guardrail_name": "test-guardrail"}, mock_logging_obj)
|
||||
)
|
||||
mocker.patch(
|
||||
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing",
|
||||
return_value=mock_processor,
|
||||
)
|
||||
mock_proxy_logging = mocker.Mock()
|
||||
mock_proxy_logging.post_call_failure_hook = AsyncMock()
|
||||
mocker.patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging)
|
||||
mocker.patch("litellm.proxy.proxy_server.general_settings", {})
|
||||
mocker.patch("litellm.proxy.proxy_server.proxy_config", mocker.Mock())
|
||||
mocker.patch("litellm.proxy.proxy_server.version", "test")
|
||||
mocker.patch("litellm.litellm_core_utils.thread_pool_executor.executor")
|
||||
|
||||
# Create request
|
||||
request = ApplyGuardrailRequest(
|
||||
guardrail_name="test-guardrail", text="Test input text with forbidden content"
|
||||
|
|
@ -1196,12 +1226,70 @@ async def test_apply_guardrail_execution_error(mocker):
|
|||
|
||||
# Call endpoint and expect ProxyException
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await apply_guardrail(request=request, user_api_key_dict=mock_user_auth)
|
||||
await apply_guardrail(
|
||||
fastapi_request=mocker.Mock(),
|
||||
request=request,
|
||||
user_api_key_dict=mock_user_auth,
|
||||
)
|
||||
|
||||
# Verify error is properly handled
|
||||
assert "Bedrock guardrail failed" in str(exc_info.value.message)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_invokes_logging_pipeline(mocker):
|
||||
mock_guardrail = mocker.Mock()
|
||||
mock_guardrail.apply_guardrail = AsyncMock(return_value={"texts": ["masked"]})
|
||||
|
||||
mock_registry = mocker.Mock()
|
||||
mock_registry.get_initialized_guardrail_callback.return_value = mock_guardrail
|
||||
mocker.patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_registry
|
||||
)
|
||||
|
||||
mock_logging_obj = mocker.Mock()
|
||||
mock_logging_obj.async_success_handler = AsyncMock()
|
||||
mock_logging_obj.model_call_details = {}
|
||||
mock_processor = mocker.Mock()
|
||||
mock_processor.common_processing_pre_call_logic = AsyncMock(
|
||||
return_value=({"guardrail_name": "test-guardrail"}, mock_logging_obj)
|
||||
)
|
||||
mocker.patch(
|
||||
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing",
|
||||
return_value=mock_processor,
|
||||
)
|
||||
|
||||
mock_proxy_logging = mocker.Mock()
|
||||
mock_proxy_logging.post_call_success_hook = AsyncMock()
|
||||
mocker.patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging)
|
||||
mocker.patch("litellm.proxy.proxy_server.general_settings", {})
|
||||
mocker.patch("litellm.proxy.proxy_server.proxy_config", mocker.Mock())
|
||||
mocker.patch("litellm.proxy.proxy_server.version", "test")
|
||||
mock_executor = mocker.Mock()
|
||||
mocker.patch(
|
||||
"litellm.litellm_core_utils.thread_pool_executor.executor", mock_executor
|
||||
)
|
||||
|
||||
request = ApplyGuardrailRequest(
|
||||
guardrail_name="test-guardrail", text="hello@example.com"
|
||||
)
|
||||
response = await apply_guardrail(
|
||||
fastapi_request=mocker.Mock(),
|
||||
request=request,
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
)
|
||||
|
||||
assert response.response_text == "masked"
|
||||
mock_processor.common_processing_pre_call_logic.assert_awaited_once()
|
||||
mock_proxy_logging.post_call_success_hook.assert_awaited_once()
|
||||
mock_logging_obj.async_success_handler.assert_awaited_once()
|
||||
assert mock_logging_obj.call_type == "pass_through_endpoint"
|
||||
mock_executor.submit.assert_called_once()
|
||||
assert mock_logging_obj.async_success_handler.await_args.kwargs["result"] == {
|
||||
"response": {"response_text": "masked"}
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_guardrail_info_endpoint_config_guardrail(mocker):
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue