diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index b38b5bb1f6a..819b3a26447 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -7,12 +7,13 @@ import inspect import json import os from datetime import datetime, timezone -from typing import Any, Dict, List, Optional, Type, TypeVar, Union, cast +from typing import Any, Dict, List, Literal, Optional, Type, TypeVar, Union, cast from urllib.parse import urlparse from fastapi import APIRouter, Depends, HTTPException from pydantic import BaseModel +from litellm.exceptions import GuardrailRaisedException from litellm.proxy.common_utils.path_utils import safe_join from litellm._logging import verbose_proxy_logger @@ -21,6 +22,9 @@ from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.generic_guardrail_api import ( + GenericGuardrailAPI, +) from litellm.proxy.guardrails.guardrail_hooks.custom_code.sandbox import ( build_sandbox_globals, compile_sandboxed, @@ -47,6 +51,7 @@ from litellm.types.guardrails import ( SupportedGuardrailIntegrations, ToolPermissionGuardrailConfigModel, ) +from litellm.types.utils import GenericGuardrailAPIInputs #### GUARDRAILS ENDPOINTS #### @@ -593,6 +598,12 @@ class ListGuardrailSubmissionsResponse(BaseModel): summary: GuardrailSubmissionSummary +class TestGuardrailSubmissionResponse(BaseModel): + success: bool + action: str + message: str + + @router.post( "/guardrails/register", tags=["Guardrails"], @@ -770,6 +781,112 @@ def _row_to_submission_item(row: Any) -> GuardrailSubmissionItem: ) +def _get_guardrail_test_input_type( + mode: Any, +) -> Literal["request", "response"]: + if isinstance(mode, list): + modes = mode + elif mode is None: + modes = [] + else: + modes = [mode] + + normalized_modes = [] + for current_mode in modes: + if current_mode is None: + continue + normalized_modes.append( + str(getattr(current_mode, "value", current_mode)).lower() + ) + + request_modes = {"pre_call", "during_call", "pre_mcp_call"} + response_modes = {"post_call", "post_mcp_call"} + + if any(current_mode in request_modes for current_mode in normalized_modes): + return "request" + if any(current_mode in response_modes for current_mode in normalized_modes): + return "response" + return "request" + + +async def _test_generic_guardrail_submission( + *, + guardrail_name: str, + litellm_params: Dict[str, Any], +) -> TestGuardrailSubmissionResponse: + guardrail_client = GenericGuardrailAPI( + api_base=litellm_params.get("api_base"), + api_key=litellm_params.get("api_key"), + headers=litellm_params.get("headers"), + additional_provider_specific_params=litellm_params.get( + "additional_provider_specific_params", {} + ), + # Approval should verify the endpoint is reachable regardless of + # the runtime fail-open/fail-closed behavior configured later. + unreachable_fallback="fail_closed", + extra_headers=litellm_params.get("extra_headers"), + guardrail_name=guardrail_name, + event_hook=litellm_params.get("mode"), + default_on=bool(litellm_params.get("default_on", False)), + ) + + probe_text = "LiteLLM guardrail connectivity check" + model = str(litellm_params.get("model") or "test-model") + input_type = _get_guardrail_test_input_type(litellm_params.get("mode")) + test_inputs: GenericGuardrailAPIInputs = { + "texts": [probe_text], + "model": model, + } + request_data = { + "model": model, + "body": { + "model": model, + "messages": [{"role": "user", "content": probe_text}], + }, + "metadata": {}, + } + + try: + result = await guardrail_client.apply_guardrail( + inputs=test_inputs, + request_data=request_data, + input_type=input_type, + ) + except GuardrailRaisedException as e: + return TestGuardrailSubmissionResponse( + success=True, + action="BLOCKED", + message=e.message or "Guardrail endpoint responded and blocked the probe.", + ) + except Exception as e: + verbose_proxy_logger.warning( + "Guardrail submission endpoint test failed for %s: %s", + guardrail_name, + e, + ) + raise HTTPException( + status_code=502, + detail=f"Guardrail endpoint check failed: {str(e)}", + ) + + action = ( + "GUARDRAIL_INTERVENED" + if result.get("texts") != test_inputs.get("texts") + or result.get("images") != test_inputs.get("images") + or result.get("tools") != test_inputs.get("tools") + else "NONE" + ) + message = "Guardrail endpoint responded successfully." + if action == "GUARDRAIL_INTERVENED": + message = "Guardrail endpoint responded successfully and modified the probe." + + return TestGuardrailSubmissionResponse( + success=True, + action=action, + message=message, + ) + + @router.get( "/guardrails/submissions", tags=["Guardrails"], @@ -916,6 +1033,61 @@ async def get_guardrail_submission( raise HTTPException(status_code=500, detail=str(e)) +@router.post( + "/guardrails/submissions/{guardrail_id}/test", + tags=["Guardrails"], + response_model=TestGuardrailSubmissionResponse, +) +async def probe_guardrail_submission_endpoint( + guardrail_id: str, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """Send a probe request to a submitted team guardrail endpoint.""" + from litellm.proxy.proxy_server import prisma_client + + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException(status_code=403, detail="Admin access required") + + if prisma_client is None: + raise HTTPException(status_code=500, detail="Prisma client not initialized") + + try: + row = await prisma_client.db.litellm_guardrailstable.find_unique( + where={"guardrail_id": guardrail_id} + ) + if row is None: + raise HTTPException( + status_code=404, detail="Guardrail submission not found" + ) + if row.team_id is None: + raise HTTPException( + status_code=400, + detail="Only team-submitted guardrails can be tested from this endpoint", + ) + + litellm_params = _parse_json_field(row.litellm_params) + if not litellm_params: + raise HTTPException( + status_code=500, + detail="Guardrail litellm_params is missing or invalid", + ) + if litellm_params.get("guardrail") != GENERIC_GUARDRAIL_API: + raise HTTPException( + status_code=400, + detail="Only generic_guardrail_api submissions support endpoint testing", + ) + + return await _test_generic_guardrail_submission( + guardrail_name=row.guardrail_name, + litellm_params=litellm_params, + ) + except HTTPException: + raise + except Exception as e: + verbose_proxy_logger.exception("Error testing guardrail submission: %s", e) + raise HTTPException(status_code=500, detail=str(e)) + + @router.post( "/guardrails/submissions/{guardrail_id}/approve", tags=["Guardrails"], @@ -948,12 +1120,6 @@ async def approve_guardrail_submission( detail=f"Guardrail is not pending review (status={row.status})", ) - now = datetime.now(timezone.utc) - await prisma_client.db.litellm_guardrailstable.update( - where={"guardrail_id": guardrail_id}, - data={"status": "active", "reviewed_at": now, "updated_at": now}, - ) - litellm_params = _parse_json_field(row.litellm_params) guardrail_info = _parse_json_field(row.guardrail_info) if not litellm_params: @@ -961,6 +1127,12 @@ async def approve_guardrail_submission( status_code=500, detail="Guardrail litellm_params is missing or invalid", ) + now = datetime.now(timezone.utc) + await prisma_client.db.litellm_guardrailstable.update( + where={"guardrail_id": guardrail_id}, + data={"status": "active", "reviewed_at": now, "updated_at": now}, + ) + guardrail_dict = { "guardrail_id": row.guardrail_id, "guardrail_name": row.guardrail_name, diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index 0e49e24496f..9f04926e2ae 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -28,6 +28,7 @@ from litellm.proxy.guardrails.guardrail_endpoints import ( list_guardrail_submissions, list_guardrails_v2, patch_guardrail, + probe_guardrail_submission_endpoint, register_guardrail, reject_guardrail_submission, update_guardrail, @@ -1692,6 +1693,80 @@ async def test_approve_guardrail_submission_not_pending(mocker): assert exc_info.value.status_code == 400 +@pytest.mark.asyncio +async def test_test_guardrail_submission_endpoint_success(mocker): + """Testing a submitted guardrail delegates to the generic endpoint probe.""" + mock_prisma = mocker.Mock() + row = mocker.Mock( + guardrail_id="sub-1", + guardrail_name="my-guard", + status="pending_review", + team_id="team-1", + litellm_params={ + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "api_base": "https://g.com", + }, + ) + mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row) + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) + expected = mocker.Mock( + success=True, + action="NONE", + message="Guardrail endpoint responded successfully.", + ) + test_submission = AsyncMock(return_value=expected) + mocker.patch( + "litellm.proxy.guardrails.guardrail_endpoints._test_generic_guardrail_submission", + test_submission, + ) + user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + + result = await probe_guardrail_submission_endpoint("sub-1", user) + + assert result is expected + test_submission.assert_awaited_once_with( + guardrail_name="my-guard", + litellm_params=row.litellm_params, + ) + + +@pytest.mark.asyncio +async def test_approve_guardrail_submission_ignores_endpoint_probe_failures(mocker): + """Approve should remain independent from the manual endpoint probe.""" + mock_prisma = mocker.Mock() + row = mocker.Mock( + guardrail_id="approve-me", + guardrail_name="my-guard", + status="pending_review", + litellm_params={ + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "api_base": "https://g.com", + }, + guardrail_info={}, + ) + mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row) + mock_prisma.db.litellm_guardrailstable.update = AsyncMock() + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma) + test_submission = AsyncMock( + side_effect=HTTPException( + status_code=502, detail="Guardrail endpoint check failed: timeout" + ) + ) + mocker.patch( + "litellm.proxy.guardrails.guardrail_endpoints._test_generic_guardrail_submission", + test_submission, + ) + user = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + + result = await approve_guardrail_submission("approve-me", user) + + assert result["status"] == "active" + mock_prisma.db.litellm_guardrailstable.update.assert_called_once() + test_submission.assert_not_awaited() + + @pytest.mark.asyncio async def test_reject_guardrail_submission_success(mocker): """Reject sets status to rejected.""" diff --git a/ui/litellm-dashboard/src/components/guardrails/TeamGuardrailsTab.tsx b/ui/litellm-dashboard/src/components/guardrails/TeamGuardrailsTab.tsx index b03ac92ba2b..dc5435cd71c 100644 --- a/ui/litellm-dashboard/src/components/guardrails/TeamGuardrailsTab.tsx +++ b/ui/litellm-dashboard/src/components/guardrails/TeamGuardrailsTab.tsx @@ -19,6 +19,7 @@ import { listGuardrailSubmissions, approveGuardrailSubmission, rejectGuardrailSubmission, + testGuardrailSubmission, updateGuardrailCall, type GuardrailSubmissionItem, } from "@/components/networking"; @@ -394,6 +395,8 @@ type DetailPanelProps = { onClose: () => void; onApprove: () => void; onReject: () => void; + onTestEndpoint: () => Promise; + isTestingEndpoint: boolean; onToggleForwardKey: () => void; onUpdateCustomHeaders: ( customHeaders: { key: string; value: string }[] @@ -406,6 +409,8 @@ function DetailPanel({ onClose, onApprove, onReject, + onTestEndpoint, + isTestingEndpoint, onToggleForwardKey, onUpdateCustomHeaders, onUpdateExtraHeaders, @@ -704,10 +709,12 @@ function DetailPanel({
{g.status === "pending" && (
@@ -824,6 +831,7 @@ export function TeamGuardrailsTab({ accessToken }: TeamGuardrailsTabProps) { const [error, setError] = useState(null); const [searchDebounced, setSearchDebounced] = useState(""); const [isSubmitModalOpen, setIsSubmitModalOpen] = useState(false); + const [testingEndpointId, setTestingEndpointId] = useState(null); const [submitForm] = Form.useForm(); const registerGuardrail = useRegisterGuardrail(); @@ -943,8 +951,10 @@ export function TeamGuardrailsTab({ accessToken }: TeamGuardrailsTabProps) { if (selectedId === id) setSelectedId(null); await fetchSubmissions(); NotificationsManager.success("Guardrail approved"); - } catch { - NotificationsManager.fromBackend("Failed to approve guardrail"); + } catch (err) { + NotificationsManager.fromBackend( + err instanceof Error ? err.message : "Failed to approve guardrail" + ); } } @@ -961,6 +971,25 @@ export function TeamGuardrailsTab({ accessToken }: TeamGuardrailsTabProps) { } } + async function handleTestEndpoint(id: string) { + if (!accessToken) return; + setTestingEndpointId(id); + try { + const result = await testGuardrailSubmission(accessToken, id); + NotificationsManager.success( + result.action === "NONE" + ? result.message + : `${result.message} (${result.action})` + ); + } catch (err) { + NotificationsManager.fromBackend( + err instanceof Error ? err.message : "Failed to reach guardrail endpoint" + ); + } finally { + setTestingEndpointId((current) => (current === id ? null : current)); + } + } + function toggleHeaders(id: string) { setExpandedHeaders((prev) => { const next = new Set(prev); @@ -1060,6 +1089,8 @@ export function TeamGuardrailsTab({ accessToken }: TeamGuardrailsTabProps) { onReject={() => setConfirmAction({ id: selected.id, action: "reject" }) } + onTestEndpoint={() => handleTestEndpoint(selected.id)} + isTestingEndpoint={testingEndpointId === selected.id} onToggleForwardKey={() => toggleForwardKey(selected.id)} onUpdateCustomHeaders={(customHeaders) => updateCustomHeaders(selected.id, customHeaders) diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index fb0e70842db..916fe4462b6 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -5061,6 +5061,12 @@ interface ListGuardrailSubmissionsResponse { summary: GuardrailSubmissionSummary; } +export interface TestGuardrailSubmissionResponse { + success: boolean; + action: string; + message: string; +} + export const listGuardrailSubmissions = async ( accessToken: string, params?: { status?: string; team_id?: string; team_guardrail?: boolean; search?: string } @@ -5088,6 +5094,29 @@ export const listGuardrailSubmissions = async ( return response.json(); }; +export const testGuardrailSubmission = async ( + accessToken: string, + guardrailId: string +): Promise => { + const url = proxyBaseUrl + ? `${proxyBaseUrl}/guardrails/submissions/${encodeURIComponent(guardrailId)}/test` + : `/guardrails/submissions/${encodeURIComponent(guardrailId)}/test`; + const response = await fetch(url, { + method: "POST", + headers: { + [globalLitellmHeaderName]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + if (!response.ok) { + const errorData = await response.json().catch(() => ({})); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + return response.json(); +}; + export const approveGuardrailSubmission = async ( accessToken: string, guardrailId: string