diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index a5865e71c2c..6782208458e 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -839,6 +839,7 @@ class ProxyBaseLLMRequestProcessing: "aget_run", "acancel_run", "adelete_run", + "apply_guardrail", ], version: Optional[str] = None, user_model: Optional[str] = None, diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index e55f3b6e16b..e0e4bdcf4a4 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -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) diff --git a/tests/enterprise/litellm_enterprise/proxy/guardrails/conftest.py b/tests/enterprise/litellm_enterprise/proxy/guardrails/conftest.py new file mode 100644 index 00000000000..4dd5c3d88ca --- /dev/null +++ b/tests/enterprise/litellm_enterprise/proxy/guardrails/conftest.py @@ -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 diff --git a/tests/enterprise/litellm_enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py b/tests/enterprise/litellm_enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py index 0d27df50d15..e5074c44210 100644 --- a/tests/enterprise/litellm_enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py +++ b/tests/enterprise/litellm_enterprise/proxy/guardrails/test_apply_guardrail_endpoint.py @@ -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 diff --git a/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py b/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py index dff444168c2..d1caf398540 100644 --- a/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py +++ b/tests/enterprise/litellm_enterprise/proxy/guardrails/test_bedrock_apply_guardrail.py @@ -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 diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 0d7becd3e2e..ce8f0802ae1 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -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): """