mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge pull request #25038 from BerriAI/litellm_feat-add-guardrail
feat: allow adding team guardrails from the UI
This commit is contained in:
commit
8ecbf757b2
8 changed files with 608 additions and 136 deletions
|
|
@ -666,6 +666,9 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/invitation/delete",
|
||||
# Team guardrail submission - requires team-scoped key; endpoint enforces team_id
|
||||
"/guardrails/register",
|
||||
# Team guardrail submissions - endpoint scopes results to caller's teams (non-admin)
|
||||
"/guardrails/submissions",
|
||||
"/guardrails/submissions/{guardrail_id}",
|
||||
] # routes that manage their own allowed/disallowed logic
|
||||
|
||||
## Org Admin Routes ##
|
||||
|
|
|
|||
|
|
@ -542,6 +542,7 @@ class RegisterGuardrailRequest(BaseModel):
|
|||
str, Any
|
||||
] # guardrail, mode, api_base required; api_key, headers, etc. optional
|
||||
guardrail_info: Optional[Dict[str, Any]] = None
|
||||
team_id: Optional[str] = None
|
||||
|
||||
def get_litellm_params_dict(self) -> Dict[str, Any]:
|
||||
return dict(self.litellm_params)
|
||||
|
|
@ -603,12 +604,24 @@ async def register_guardrail(
|
|||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Prisma client not initialized")
|
||||
|
||||
if not user_api_key_dict.team_id:
|
||||
# Resolve team_id: prefer request body, fall back to API key's team
|
||||
team_id = request.team_id or user_api_key_dict.team_id
|
||||
if not team_id:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Registration requires an API key associated with a team. Use a team-scoped key.",
|
||||
detail="team_id is required. Provide it in the request body or use a team-scoped API key.",
|
||||
)
|
||||
|
||||
# Validate team membership for non-admin users when team differs from key
|
||||
is_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
if not is_admin and team_id != user_api_key_dict.team_id:
|
||||
user_team_ids = await _get_user_team_ids(user_api_key_dict)
|
||||
if team_id not in user_team_ids:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"You are not a member of team {team_id!r}",
|
||||
)
|
||||
|
||||
params = request.get_litellm_params_dict()
|
||||
if params.get("guardrail") != GENERIC_GUARDRAIL_API:
|
||||
raise HTTPException(
|
||||
|
|
@ -673,7 +686,7 @@ async def register_guardrail(
|
|||
"litellm_params": litellm_params_str,
|
||||
"guardrail_info": guardrail_info_str,
|
||||
"status": "pending_review",
|
||||
"team_id": user_api_key_dict.team_id,
|
||||
"team_id": team_id,
|
||||
"submitted_at": now,
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
|
|
@ -703,6 +716,30 @@ def _parse_json_field(value: Any) -> Optional[Dict[str, Any]]:
|
|||
return None
|
||||
|
||||
|
||||
async def _get_user_team_ids(user_api_key_dict: UserAPIKeyAuth) -> List[str]:
|
||||
"""Return the list of team_ids the caller belongs to (empty list if none)."""
|
||||
from litellm.proxy.auth.auth_checks import get_user_object
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if not user_api_key_dict.user_id or prisma_client is None:
|
||||
return []
|
||||
user_obj = await get_user_object(
|
||||
user_id=user_api_key_dict.user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
parent_otel_span=user_api_key_dict.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if user_obj is None or not user_obj.teams:
|
||||
return []
|
||||
return [t for t in user_obj.teams if t]
|
||||
|
||||
|
||||
def _row_to_submission_item(row: Any) -> GuardrailSubmissionItem:
|
||||
guardrail_info = _parse_json_field(row.guardrail_info) or {}
|
||||
team_guardrail = row.team_id is not None
|
||||
|
|
@ -735,27 +772,49 @@ async def list_guardrail_submissions(
|
|||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
List team guardrail submissions (admin only). Returns only guardrails with a team_id.
|
||||
List team guardrail submissions. Returns only guardrails with a team_id.
|
||||
|
||||
Admins see all submissions. Non-admin users see submissions for teams they are
|
||||
a member of.
|
||||
|
||||
Status values: pending_review (team-registered, awaiting approval), active (approved), rejected.
|
||||
|
||||
Optional filters:
|
||||
- status: pending_review | active | rejected
|
||||
- team_id: filter by specific team
|
||||
- team_id: filter by specific team (non-admins must be a member of that team)
|
||||
- search: name/description
|
||||
"""
|
||||
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")
|
||||
|
||||
is_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
visible_team_ids: Optional[List[str]] = None
|
||||
if not is_admin:
|
||||
visible_team_ids = await _get_user_team_ids(user_api_key_dict)
|
||||
if team_id is not None and team_id not in visible_team_ids:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"You are not a member of team {team_id!r}",
|
||||
)
|
||||
|
||||
try:
|
||||
# Single query: fetch all team guardrails (team_id is not null)
|
||||
where_clause: Dict[str, Any] = {"team_id": {"not": None}}
|
||||
if visible_team_ids is not None:
|
||||
if not visible_team_ids:
|
||||
# Non-admin with no team memberships: nothing visible.
|
||||
return ListGuardrailSubmissionsResponse(
|
||||
submissions=[],
|
||||
summary=GuardrailSubmissionSummary(
|
||||
total=0, pending_review=0, active=0, rejected=0
|
||||
),
|
||||
)
|
||||
where_clause["team_id"] = {"in": visible_team_ids}
|
||||
|
||||
# Single query: fetch team guardrails visible to the caller
|
||||
all_team_rows = await prisma_client.db.litellm_guardrailstable.find_many(
|
||||
where={"team_id": {"not": None}},
|
||||
where=where_clause,
|
||||
order={"created_at": "desc"},
|
||||
)
|
||||
|
||||
|
|
@ -816,15 +875,14 @@ async def get_guardrail_submission(
|
|||
guardrail_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""Get a single guardrail submission by id (admin only)."""
|
||||
"""Get a single guardrail submission by id. Non-admins may only access submissions for teams they belong to."""
|
||||
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")
|
||||
|
||||
is_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
|
||||
try:
|
||||
row = await prisma_client.db.litellm_guardrailstable.find_unique(
|
||||
where={"guardrail_id": guardrail_id}
|
||||
|
|
@ -833,6 +891,13 @@ async def get_guardrail_submission(
|
|||
raise HTTPException(
|
||||
status_code=404, detail="Guardrail submission not found"
|
||||
)
|
||||
if not is_admin:
|
||||
visible_team_ids = await _get_user_team_ids(user_api_key_dict)
|
||||
if row.team_id is None or row.team_id not in visible_team_ids:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail="You are not a member of the team that owns this submission",
|
||||
)
|
||||
return _row_to_submission_item(row)
|
||||
except HTTPException:
|
||||
raise
|
||||
|
|
|
|||
|
|
@ -1220,6 +1220,59 @@ async def test_register_guardrail_requires_team_id(mocker):
|
|||
assert "team" in exc_info.value.detail.lower()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_guardrail_non_admin_cross_team_allowed(mocker):
|
||||
"""Non-admin may register for a team in their user.teams list even if the key's team_id differs."""
|
||||
mock_prisma = mocker.Mock()
|
||||
mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=None)
|
||||
created = mocker.Mock(
|
||||
guardrail_id="g1",
|
||||
guardrail_name=MOCK_REGISTER_REQUEST.guardrail_name,
|
||||
status="pending_review",
|
||||
submitted_at=datetime.now(),
|
||||
)
|
||||
mock_prisma.db.litellm_guardrailstable.create = AsyncMock(return_value=created)
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
mocker.patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints._get_user_team_ids",
|
||||
AsyncMock(return_value=["team-alpha", "team-beta"]),
|
||||
)
|
||||
req = RegisterGuardrailRequest(
|
||||
guardrail_name=MOCK_REGISTER_REQUEST.guardrail_name,
|
||||
team_id="team-beta",
|
||||
litellm_params=MOCK_REGISTER_REQUEST.litellm_params,
|
||||
)
|
||||
user = UserAPIKeyAuth(
|
||||
user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER, team_id="team-alpha"
|
||||
)
|
||||
|
||||
result = await register_guardrail(req, user)
|
||||
|
||||
assert result.guardrail_id == "g1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_guardrail_non_admin_cross_team_forbidden(mocker):
|
||||
"""Non-admin gets 403 when registering for a team they are not a member of."""
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mocker.Mock())
|
||||
mocker.patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints._get_user_team_ids",
|
||||
AsyncMock(return_value=["team-alpha"]),
|
||||
)
|
||||
req = RegisterGuardrailRequest(
|
||||
guardrail_name=MOCK_REGISTER_REQUEST.guardrail_name,
|
||||
team_id="team-other",
|
||||
litellm_params=MOCK_REGISTER_REQUEST.litellm_params,
|
||||
)
|
||||
user = UserAPIKeyAuth(
|
||||
user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER, team_id="team-alpha"
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await register_guardrail(req, user)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_guardrail_duplicate_name(mocker):
|
||||
"""Register returns 400 when guardrail_name already exists."""
|
||||
|
|
@ -1237,13 +1290,82 @@ async def test_register_guardrail_duplicate_name(mocker):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_guardrail_submissions_requires_admin(mocker):
|
||||
"""List submissions returns 403 when user is not admin."""
|
||||
async def test_list_guardrail_submissions_non_admin_scoped_to_own_teams(mocker):
|
||||
"""Non-admin callers see only submissions for teams they belong to."""
|
||||
mock_prisma = mocker.Mock()
|
||||
own_team_row = mocker.Mock(
|
||||
guardrail_id="mine",
|
||||
guardrail_name="mine-guard",
|
||||
status="pending_review",
|
||||
team_id="team-mine",
|
||||
litellm_params={},
|
||||
guardrail_info={},
|
||||
submitted_at=None,
|
||||
reviewed_at=None,
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
)
|
||||
find_many = AsyncMock(return_value=[own_team_row])
|
||||
mock_prisma.db.litellm_guardrailstable.find_many = find_many
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
mocker.patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints._get_user_team_ids",
|
||||
AsyncMock(return_value=["team-mine"]),
|
||||
)
|
||||
user = UserAPIKeyAuth(
|
||||
user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
result = await list_guardrail_submissions(user_api_key_dict=user)
|
||||
|
||||
# DB query scoped to visible teams
|
||||
where_clause = find_many.call_args.kwargs["where"]
|
||||
assert where_clause["team_id"] == {"in": ["team-mine"]}
|
||||
assert len(result.submissions) == 1
|
||||
assert result.submissions[0].team_id == "team-mine"
|
||||
# Summary counts reflect only visible teams
|
||||
assert result.summary.total == 1
|
||||
assert result.summary.pending_review == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_guardrail_submissions_non_admin_no_teams(mocker):
|
||||
"""Non-admin caller with no team memberships gets an empty list (not 403)."""
|
||||
mock_prisma = mocker.Mock()
|
||||
find_many = AsyncMock(return_value=[])
|
||||
mock_prisma.db.litellm_guardrailstable.find_many = find_many
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
mocker.patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints._get_user_team_ids",
|
||||
AsyncMock(return_value=[]),
|
||||
)
|
||||
user = UserAPIKeyAuth(
|
||||
user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
result = await list_guardrail_submissions(user_api_key_dict=user)
|
||||
|
||||
assert result.submissions == []
|
||||
assert result.summary.total == 0
|
||||
assert find_many.call_count == 0 # no DB query when user has no teams
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_guardrail_submissions_non_admin_team_filter_forbidden(mocker):
|
||||
"""Non-admin caller filtering by a team they're not in gets 403."""
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mocker.Mock())
|
||||
user = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER)
|
||||
mocker.patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints._get_user_team_ids",
|
||||
AsyncMock(return_value=["team-mine"]),
|
||||
)
|
||||
user = UserAPIKeyAuth(
|
||||
user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await list_guardrail_submissions(user_api_key_dict=user)
|
||||
await list_guardrail_submissions(
|
||||
team_id="team-other", user_api_key_dict=user
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
|
||||
|
|
@ -1354,6 +1476,69 @@ async def test_get_guardrail_submission_not_found(mocker):
|
|||
assert exc_info.value.status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_guardrail_submission_non_admin_own_team(mocker):
|
||||
"""Non-admin caller can fetch a submission belonging to one of their teams."""
|
||||
mock_prisma = mocker.Mock()
|
||||
row = mocker.Mock(
|
||||
guardrail_id="sub-1",
|
||||
guardrail_name="team-guard",
|
||||
status="pending_review",
|
||||
team_id="team-mine",
|
||||
litellm_params={},
|
||||
guardrail_info={},
|
||||
submitted_at=None,
|
||||
reviewed_at=None,
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
)
|
||||
mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row)
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
mocker.patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints._get_user_team_ids",
|
||||
AsyncMock(return_value=["team-mine"]),
|
||||
)
|
||||
user = UserAPIKeyAuth(
|
||||
user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
result = await get_guardrail_submission("sub-1", user)
|
||||
|
||||
assert result.guardrail_id == "sub-1"
|
||||
assert result.team_id == "team-mine"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_guardrail_submission_non_admin_other_team_forbidden(mocker):
|
||||
"""Non-admin caller gets 403 when fetching a submission for a team they're not in."""
|
||||
mock_prisma = mocker.Mock()
|
||||
row = mocker.Mock(
|
||||
guardrail_id="sub-1",
|
||||
guardrail_name="team-guard",
|
||||
status="pending_review",
|
||||
team_id="team-other",
|
||||
litellm_params={},
|
||||
guardrail_info={},
|
||||
submitted_at=None,
|
||||
reviewed_at=None,
|
||||
created_at=datetime.now(),
|
||||
updated_at=datetime.now(),
|
||||
)
|
||||
mock_prisma.db.litellm_guardrailstable.find_unique = AsyncMock(return_value=row)
|
||||
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
mocker.patch(
|
||||
"litellm.proxy.guardrails.guardrail_endpoints._get_user_team_ids",
|
||||
AsyncMock(return_value=["team-mine"]),
|
||||
)
|
||||
user = UserAPIKeyAuth(
|
||||
user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await get_guardrail_submission("sub-1", user)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_approve_guardrail_submission_success(mocker):
|
||||
"""Approve sets status to active and initializes guardrail in memory."""
|
||||
|
|
|
|||
|
|
@ -0,0 +1,74 @@
|
|||
import { useMutation, useQueryClient } from "@tanstack/react-query";
|
||||
import {
|
||||
getProxyBaseUrl,
|
||||
getGlobalLitellmHeaderName,
|
||||
deriveErrorMessage,
|
||||
handleError,
|
||||
} from "@/components/networking";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import { createQueryKeys } from "../common/queryKeysFactory";
|
||||
|
||||
// ── Types ────────────────────────────────────────────────────────────────────
|
||||
|
||||
export interface RegisterGuardrailParams {
|
||||
guardrail_name: string;
|
||||
litellm_params: Record<string, unknown>;
|
||||
guardrail_info?: Record<string, unknown>;
|
||||
team_id?: string;
|
||||
}
|
||||
|
||||
export interface RegisterGuardrailResponse {
|
||||
guardrail_id: string;
|
||||
guardrail_name: string;
|
||||
status: string;
|
||||
submitted_at?: string | null;
|
||||
}
|
||||
|
||||
// ── Fetch function ───────────────────────────────────────────────────────────
|
||||
|
||||
const registerGuardrail = async (
|
||||
accessToken: string,
|
||||
params: RegisterGuardrailParams,
|
||||
): Promise<RegisterGuardrailResponse> => {
|
||||
const baseUrl = getProxyBaseUrl();
|
||||
const url = `${baseUrl}/guardrails/register`;
|
||||
|
||||
const response = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
[getGlobalLitellmHeaderName()]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
body: JSON.stringify(params),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorData = await response.json().catch(() => ({}));
|
||||
const errorMessage = deriveErrorMessage(errorData);
|
||||
handleError(errorMessage);
|
||||
throw new Error(errorMessage);
|
||||
}
|
||||
|
||||
return response.json();
|
||||
};
|
||||
|
||||
// ── Hook ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
const guardrailKeys = createQueryKeys("guardrails");
|
||||
|
||||
export const useRegisterGuardrail = () => {
|
||||
const { accessToken } = useAuthorized();
|
||||
const queryClient = useQueryClient();
|
||||
|
||||
return useMutation<RegisterGuardrailResponse, Error, RegisterGuardrailParams>({
|
||||
mutationFn: async (params) => {
|
||||
if (!accessToken) {
|
||||
throw new Error("Access token is required");
|
||||
}
|
||||
return registerGuardrail(accessToken, params);
|
||||
},
|
||||
onSuccess: () => {
|
||||
queryClient.invalidateQueries({ queryKey: guardrailKeys.all });
|
||||
},
|
||||
});
|
||||
};
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
import { render, screen } from "@testing-library/react";
|
||||
import { render, screen, fireEvent } from "@testing-library/react";
|
||||
import { beforeAll, beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import GuardrailsPanel from "./guardrails";
|
||||
import { getGuardrailsList } from "./networking";
|
||||
|
|
@ -40,6 +40,10 @@ vi.mock("./guardrails/GuardrailTestPlayground", () => ({
|
|||
default: () => <div>Mock Guardrail Test Playground</div>,
|
||||
}));
|
||||
|
||||
vi.mock("./guardrails/TeamGuardrailsTab", () => ({
|
||||
TeamGuardrailsTab: () => <div>Mock Team Guardrails Tab</div>,
|
||||
}));
|
||||
|
||||
vi.mock("@/utils/roles", () => ({
|
||||
isAdminRole: vi.fn((role: string) => role === "admin"),
|
||||
}));
|
||||
|
|
@ -99,6 +103,8 @@ describe("GuardrailsPanel", () => {
|
|||
it("should render the component", async () => {
|
||||
render(<GuardrailsPanel {...defaultProps} />);
|
||||
expect(screen.getByText("Guardrails")).toBeInTheDocument();
|
||||
// Activate the Guardrails tab so its content (including the Add button) is rendered
|
||||
fireEvent.click(screen.getByText("Guardrails"));
|
||||
expect(screen.getByText("+ Add New Guardrail")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
import React, { useState, useEffect } from "react";
|
||||
import { Button, TabGroup, TabList, Tab, TabPanels, TabPanel } from "@tremor/react";
|
||||
import { Dropdown } from "antd";
|
||||
import { Button, Dropdown, Tabs } from "antd";
|
||||
import { DownOutlined, PlusOutlined, CodeOutlined } from "@ant-design/icons";
|
||||
import { getGuardrailsList, deleteGuardrailCall } from "./networking";
|
||||
import AddGuardrailForm from "./guardrails/add_guardrail_form";
|
||||
|
|
@ -48,8 +47,6 @@ const GuardrailsPanel: React.FC<GuardrailsPanelProps> = ({ accessToken, userRole
|
|||
const [guardrailToDelete, setGuardrailToDelete] = useState<Guardrail | null>(null);
|
||||
const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false);
|
||||
const [selectedGuardrailId, setSelectedGuardrailId] = useState<string | null>(null);
|
||||
const [activeTab, setActiveTab] = useState<number>(0);
|
||||
|
||||
const isAdmin = userRole ? isAdminRole(userRole) : false;
|
||||
|
||||
const fetchGuardrails = async () => {
|
||||
|
|
@ -135,122 +132,130 @@ const GuardrailsPanel: React.FC<GuardrailsPanelProps> = ({ accessToken, userRole
|
|||
|
||||
return (
|
||||
<div className="w-full mx-auto flex-auto overflow-y-auto m-8 p-2">
|
||||
<TabGroup index={activeTab} onIndexChange={setActiveTab}>
|
||||
<TabList className="mb-4">
|
||||
<Tab>Guardrail Garden</Tab>
|
||||
<Tab>Guardrails</Tab>
|
||||
<Tab disabled={!accessToken || guardrailsList.length === 0}>Test Playground</Tab>
|
||||
<Tab>Submitted Guardrails</Tab>
|
||||
</TabList>
|
||||
|
||||
<TabPanels>
|
||||
{/* Guardrail Garden Tab */}
|
||||
<TabPanel>
|
||||
<GuardrailGarden
|
||||
accessToken={accessToken}
|
||||
onGuardrailCreated={handleSuccess}
|
||||
/>
|
||||
</TabPanel>
|
||||
|
||||
{/* Existing Guardrails Tab */}
|
||||
<TabPanel>
|
||||
<div className="flex justify-between items-center mb-4">
|
||||
<Dropdown
|
||||
menu={{
|
||||
items: [
|
||||
{
|
||||
key: "provider",
|
||||
icon: <PlusOutlined />,
|
||||
label: "Add Provider Guardrail",
|
||||
onClick: handleAddGuardrail,
|
||||
},
|
||||
{
|
||||
key: "custom_code",
|
||||
icon: <CodeOutlined />,
|
||||
label: "Create Custom Code Guardrail",
|
||||
onClick: handleAddCustomCodeGuardrail,
|
||||
},
|
||||
],
|
||||
}}
|
||||
trigger={["click"]}
|
||||
disabled={!accessToken}
|
||||
>
|
||||
<Button disabled={!accessToken}>
|
||||
+ Add New Guardrail <DownOutlined className="ml-2" />
|
||||
</Button>
|
||||
</Dropdown>
|
||||
</div>
|
||||
|
||||
{selectedGuardrailId ? (
|
||||
<GuardrailInfoView
|
||||
guardrailId={selectedGuardrailId}
|
||||
onClose={() => setSelectedGuardrailId(null)}
|
||||
accessToken={accessToken}
|
||||
isAdmin={isAdmin}
|
||||
/>
|
||||
) : (
|
||||
<GuardrailTable
|
||||
guardrailsList={guardrailsList}
|
||||
isLoading={isLoading}
|
||||
onDeleteClick={handleDeleteClick}
|
||||
accessToken={accessToken}
|
||||
onGuardrailUpdated={fetchGuardrails}
|
||||
isAdmin={isAdmin}
|
||||
onGuardrailClick={(id) => setSelectedGuardrailId(id)}
|
||||
/>
|
||||
)}
|
||||
|
||||
<AddGuardrailForm
|
||||
visible={isAddModalVisible}
|
||||
onClose={handleCloseModal}
|
||||
accessToken={accessToken}
|
||||
onSuccess={handleSuccess}
|
||||
/>
|
||||
|
||||
<CustomCodeModal
|
||||
visible={isCustomCodeModalVisible}
|
||||
onClose={handleCloseCustomCodeModal}
|
||||
accessToken={accessToken}
|
||||
onSuccess={handleSuccess}
|
||||
/>
|
||||
|
||||
<DeleteResourceModal
|
||||
isOpen={isDeleteModalOpen}
|
||||
title="Delete Guardrail"
|
||||
message={`Are you sure you want to delete guardrail: ${guardrailToDelete?.guardrail_name}? This action cannot be undone.`}
|
||||
resourceInformationTitle="Guardrail Information"
|
||||
resourceInformation={[
|
||||
{ label: "Name", value: guardrailToDelete?.guardrail_name },
|
||||
{ label: "ID", value: guardrailToDelete?.guardrail_id, code: true },
|
||||
{ label: "Provider", value: providerDisplayName },
|
||||
{ label: "Mode", value: guardrailToDelete?.litellm_params.mode },
|
||||
<Tabs
|
||||
defaultActiveKey="submitted"
|
||||
items={[
|
||||
...(isAdmin
|
||||
? [
|
||||
{
|
||||
label: "Default On",
|
||||
value: guardrailToDelete?.litellm_params.default_on ? "Yes" : "No",
|
||||
key: "garden",
|
||||
label: "Guardrail Garden",
|
||||
children: (
|
||||
<GuardrailGarden
|
||||
accessToken={accessToken}
|
||||
onGuardrailCreated={handleSuccess}
|
||||
/>
|
||||
),
|
||||
},
|
||||
]}
|
||||
onCancel={handleDeleteCancel}
|
||||
onOk={handleDeleteConfirm}
|
||||
confirmLoading={isDeleting}
|
||||
/>
|
||||
</TabPanel>
|
||||
{
|
||||
key: "guardrails",
|
||||
label: "Guardrails",
|
||||
children: (
|
||||
<>
|
||||
<div className="flex justify-between items-center mb-4">
|
||||
<Dropdown
|
||||
menu={{
|
||||
items: [
|
||||
{
|
||||
key: "provider",
|
||||
icon: <PlusOutlined />,
|
||||
label: "Add Provider Guardrail",
|
||||
onClick: handleAddGuardrail,
|
||||
},
|
||||
{
|
||||
key: "custom_code",
|
||||
icon: <CodeOutlined />,
|
||||
label: "Create Custom Code Guardrail",
|
||||
onClick: handleAddCustomCodeGuardrail,
|
||||
},
|
||||
],
|
||||
}}
|
||||
trigger={["click"]}
|
||||
disabled={!accessToken}
|
||||
>
|
||||
<Button disabled={!accessToken}>
|
||||
+ Add New Guardrail <DownOutlined className="ml-2" />
|
||||
</Button>
|
||||
</Dropdown>
|
||||
</div>
|
||||
|
||||
{/* Test Playground Tab */}
|
||||
<TabPanel>
|
||||
<GuardrailTestPlayground
|
||||
guardrailsList={guardrailsList}
|
||||
isLoading={isLoading}
|
||||
accessToken={accessToken}
|
||||
onClose={() => setActiveTab(0)}
|
||||
/>
|
||||
</TabPanel>
|
||||
{selectedGuardrailId ? (
|
||||
<GuardrailInfoView
|
||||
guardrailId={selectedGuardrailId}
|
||||
onClose={() => setSelectedGuardrailId(null)}
|
||||
accessToken={accessToken}
|
||||
isAdmin={isAdmin}
|
||||
/>
|
||||
) : (
|
||||
<GuardrailTable
|
||||
guardrailsList={guardrailsList}
|
||||
isLoading={isLoading}
|
||||
onDeleteClick={handleDeleteClick}
|
||||
accessToken={accessToken}
|
||||
onGuardrailUpdated={fetchGuardrails}
|
||||
isAdmin={isAdmin}
|
||||
onGuardrailClick={(id) => setSelectedGuardrailId(id)}
|
||||
/>
|
||||
)}
|
||||
|
||||
{/* Team Guardrails Tab */}
|
||||
<TabPanel>
|
||||
<TeamGuardrailsTab accessToken={accessToken} />
|
||||
</TabPanel>
|
||||
</TabPanels>
|
||||
</TabGroup>
|
||||
<AddGuardrailForm
|
||||
visible={isAddModalVisible}
|
||||
onClose={handleCloseModal}
|
||||
accessToken={accessToken}
|
||||
onSuccess={handleSuccess}
|
||||
/>
|
||||
|
||||
<CustomCodeModal
|
||||
visible={isCustomCodeModalVisible}
|
||||
onClose={handleCloseCustomCodeModal}
|
||||
accessToken={accessToken}
|
||||
onSuccess={handleSuccess}
|
||||
/>
|
||||
|
||||
<DeleteResourceModal
|
||||
isOpen={isDeleteModalOpen}
|
||||
title="Delete Guardrail"
|
||||
message={`Are you sure you want to delete guardrail: ${guardrailToDelete?.guardrail_name}? This action cannot be undone.`}
|
||||
resourceInformationTitle="Guardrail Information"
|
||||
resourceInformation={[
|
||||
{ label: "Name", value: guardrailToDelete?.guardrail_name },
|
||||
{ label: "ID", value: guardrailToDelete?.guardrail_id, code: true },
|
||||
{ label: "Provider", value: providerDisplayName },
|
||||
{ label: "Mode", value: guardrailToDelete?.litellm_params.mode },
|
||||
{
|
||||
label: "Default On",
|
||||
value: guardrailToDelete?.litellm_params.default_on ? "Yes" : "No",
|
||||
},
|
||||
]}
|
||||
onCancel={handleDeleteCancel}
|
||||
onOk={handleDeleteConfirm}
|
||||
confirmLoading={isDeleting}
|
||||
/>
|
||||
</>
|
||||
),
|
||||
},
|
||||
{
|
||||
key: "playground",
|
||||
label: "Test Playground",
|
||||
disabled: !accessToken,
|
||||
children: (
|
||||
<GuardrailTestPlayground
|
||||
guardrailsList={guardrailsList}
|
||||
isLoading={isLoading}
|
||||
accessToken={accessToken}
|
||||
onClose={() => {}}
|
||||
/>
|
||||
),
|
||||
},
|
||||
]
|
||||
: []),
|
||||
{
|
||||
key: "submitted",
|
||||
label: "Submitted Guardrails",
|
||||
children: <TeamGuardrailsTab accessToken={accessToken} />,
|
||||
},
|
||||
]}
|
||||
/>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ import {
|
|||
AlertCircleIcon,
|
||||
InfoIcon,
|
||||
} from "lucide-react";
|
||||
import { Modal, Form, Input, Select } from "antd";
|
||||
import {
|
||||
listGuardrailSubmissions,
|
||||
approveGuardrailSubmission,
|
||||
|
|
@ -22,6 +23,8 @@ import {
|
|||
type GuardrailSubmissionItem,
|
||||
} from "@/components/networking";
|
||||
import NotificationsManager from "@/components/molecules/notifications_manager";
|
||||
import TeamDropdown from "@/components/common_components/team_dropdown";
|
||||
import { useRegisterGuardrail } from "@/app/(dashboard)/hooks/guardrails/useRegisterGuardrail";
|
||||
|
||||
type GuardrailStatus = "active" | "pending" | "rejected";
|
||||
|
||||
|
|
@ -820,6 +823,9 @@ export function TeamGuardrailsTab({ accessToken }: TeamGuardrailsTabProps) {
|
|||
const [isLoading, setIsLoading] = useState(true);
|
||||
const [error, setError] = useState<string | null>(null);
|
||||
const [searchDebounced, setSearchDebounced] = useState("");
|
||||
const [isSubmitModalOpen, setIsSubmitModalOpen] = useState(false);
|
||||
const [submitForm] = Form.useForm();
|
||||
const registerGuardrail = useRegisterGuardrail();
|
||||
|
||||
useEffect(() => {
|
||||
const t = setTimeout(() => setSearchDebounced(search), 300);
|
||||
|
|
@ -1006,6 +1012,7 @@ export function TeamGuardrailsTab({ accessToken }: TeamGuardrailsTabProps) {
|
|||
</select>
|
||||
<button
|
||||
type="button"
|
||||
onClick={() => setIsSubmitModalOpen(true)}
|
||||
className="ml-auto flex items-center gap-2 bg-blue-500 hover:bg-blue-600 text-white text-sm font-medium px-4 py-2 rounded-md transition-colors"
|
||||
>
|
||||
<PlusIcon className="h-4 w-4" />
|
||||
|
|
@ -1076,6 +1083,134 @@ export function TeamGuardrailsTab({ accessToken }: TeamGuardrailsTabProps) {
|
|||
onCancel={() => setConfirmAction(null)}
|
||||
/>
|
||||
)}
|
||||
|
||||
<Modal
|
||||
title="Submit Guardrail for Review"
|
||||
open={isSubmitModalOpen}
|
||||
onCancel={() => {
|
||||
setIsSubmitModalOpen(false);
|
||||
submitForm.resetFields();
|
||||
}}
|
||||
onOk={() => submitForm.submit()}
|
||||
okText="Submit for Review"
|
||||
>
|
||||
<div className="rounded-md bg-blue-50 border border-blue-200 px-4 py-3 text-sm text-blue-800 mb-4">
|
||||
Your guardrail will be sent for admin review before it becomes active.
|
||||
</div>
|
||||
<Form
|
||||
form={submitForm}
|
||||
layout="vertical"
|
||||
initialValues={{ mode: "pre_call" }}
|
||||
onFinish={async (values) => {
|
||||
const litellm_params: Record<string, unknown> = {
|
||||
...(values.extra_litellm_params ? JSON.parse(values.extra_litellm_params) : {}),
|
||||
guardrail: "generic_guardrail_api",
|
||||
mode: values.mode,
|
||||
api_base: values.api_base,
|
||||
};
|
||||
try {
|
||||
await registerGuardrail.mutateAsync({
|
||||
team_id: values.team_id,
|
||||
guardrail_name: values.guardrail_name,
|
||||
litellm_params,
|
||||
guardrail_info: values.guardrail_info ? JSON.parse(values.guardrail_info) : undefined,
|
||||
});
|
||||
NotificationsManager.success("Guardrail submitted for review");
|
||||
setIsSubmitModalOpen(false);
|
||||
submitForm.resetFields();
|
||||
fetchSubmissions();
|
||||
} catch {
|
||||
// error already handled by networking layer
|
||||
}
|
||||
}}
|
||||
>
|
||||
<Form.Item
|
||||
label="Team"
|
||||
name="team_id"
|
||||
rules={[{ required: true, message: "Select a team" }]}
|
||||
>
|
||||
<TeamDropdown />
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
label="Guardrail Name"
|
||||
name="guardrail_name"
|
||||
rules={[{ required: true, message: "Enter a guardrail name" }]}
|
||||
>
|
||||
<Input placeholder="e.g. pii-detection" />
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
label="Mode"
|
||||
name="mode"
|
||||
rules={[{ required: true, message: "Select a mode" }]}
|
||||
>
|
||||
<Select>
|
||||
<Select.Option value="pre_call">Pre Call</Select.Option>
|
||||
<Select.Option value="post_call">Post Call</Select.Option>
|
||||
<Select.Option value="during_call">During Call</Select.Option>
|
||||
</Select>
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
label="API Base URL"
|
||||
name="api_base"
|
||||
rules={[
|
||||
{ required: true, message: "Enter the API base URL" },
|
||||
{ type: "url", message: "Must be a valid URL" },
|
||||
]}
|
||||
>
|
||||
<Input placeholder="https://your-guardrail-api.com/v1/check" className="font-mono" />
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
label="Additional litellm_params (optional)"
|
||||
name="extra_litellm_params"
|
||||
tooltip="JSON object merged into litellm_params. e.g. forward_api_key, headers, model, unreachable_fallback"
|
||||
rules={[
|
||||
{
|
||||
validator: (_, value) => {
|
||||
if (!value) return Promise.resolve();
|
||||
try {
|
||||
const parsed = JSON.parse(value);
|
||||
if (typeof parsed !== "object" || Array.isArray(parsed)) {
|
||||
return Promise.reject("Must be a JSON object");
|
||||
}
|
||||
return Promise.resolve();
|
||||
} catch {
|
||||
return Promise.reject("Invalid JSON");
|
||||
}
|
||||
},
|
||||
},
|
||||
]}
|
||||
>
|
||||
<Input.TextArea
|
||||
rows={3}
|
||||
className="font-mono text-xs"
|
||||
placeholder='{"forward_api_key": true, "headers": {"X-Custom": "value"}}'
|
||||
/>
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
label="Guardrail Info (optional)"
|
||||
name="guardrail_info"
|
||||
rules={[
|
||||
{
|
||||
validator: (_, value) => {
|
||||
if (!value) return Promise.resolve();
|
||||
try {
|
||||
JSON.parse(value);
|
||||
return Promise.resolve();
|
||||
} catch {
|
||||
return Promise.reject("Invalid JSON");
|
||||
}
|
||||
},
|
||||
},
|
||||
]}
|
||||
>
|
||||
<Input.TextArea
|
||||
rows={3}
|
||||
className="font-mono text-xs"
|
||||
placeholder='{"description": "Detects PII in requests"}'
|
||||
/>
|
||||
</Form.Item>
|
||||
</Form>
|
||||
</Modal>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -137,7 +137,6 @@ const menuGroups: MenuGroup[] = [
|
|||
page: "guardrails",
|
||||
label: "Guardrails",
|
||||
icon: <SafetyOutlined />,
|
||||
roles: all_admin_roles,
|
||||
},
|
||||
{
|
||||
key: "policies",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue