diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 140bcefc1a9..441b3b836a1 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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 ## diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index 2b20876ba22..422bdc13780 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -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 diff --git a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py index ca224726361..defea08594f 100644 --- a/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_guardrail_endpoints.py @@ -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.""" diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useRegisterGuardrail.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useRegisterGuardrail.ts new file mode 100644 index 00000000000..3135e8326fc --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useRegisterGuardrail.ts @@ -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; + guardrail_info?: Record; + 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 => { + 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({ + mutationFn: async (params) => { + if (!accessToken) { + throw new Error("Access token is required"); + } + return registerGuardrail(accessToken, params); + }, + onSuccess: () => { + queryClient.invalidateQueries({ queryKey: guardrailKeys.all }); + }, + }); +}; diff --git a/ui/litellm-dashboard/src/components/guardrails.test.tsx b/ui/litellm-dashboard/src/components/guardrails.test.tsx index 8cafc18eb9a..99c2474347e 100644 --- a/ui/litellm-dashboard/src/components/guardrails.test.tsx +++ b/ui/litellm-dashboard/src/components/guardrails.test.tsx @@ -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: () =>
Mock Guardrail Test Playground
, })); +vi.mock("./guardrails/TeamGuardrailsTab", () => ({ + TeamGuardrailsTab: () =>
Mock Team Guardrails Tab
, +})); + vi.mock("@/utils/roles", () => ({ isAdminRole: vi.fn((role: string) => role === "admin"), })); @@ -99,6 +103,8 @@ describe("GuardrailsPanel", () => { it("should render the component", async () => { render(); 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(); }); }); diff --git a/ui/litellm-dashboard/src/components/guardrails.tsx b/ui/litellm-dashboard/src/components/guardrails.tsx index 56bba2724d0..fee9d02d3ae 100644 --- a/ui/litellm-dashboard/src/components/guardrails.tsx +++ b/ui/litellm-dashboard/src/components/guardrails.tsx @@ -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 = ({ accessToken, userRole const [guardrailToDelete, setGuardrailToDelete] = useState(null); const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); const [selectedGuardrailId, setSelectedGuardrailId] = useState(null); - const [activeTab, setActiveTab] = useState(0); - const isAdmin = userRole ? isAdminRole(userRole) : false; const fetchGuardrails = async () => { @@ -135,122 +132,130 @@ const GuardrailsPanel: React.FC = ({ accessToken, userRole return (
- - - Guardrail Garden - Guardrails - Test Playground - Submitted Guardrails - - - - {/* Guardrail Garden Tab */} - - - - - {/* Existing Guardrails Tab */} - -
- , - label: "Add Provider Guardrail", - onClick: handleAddGuardrail, - }, - { - key: "custom_code", - icon: , - label: "Create Custom Code Guardrail", - onClick: handleAddCustomCodeGuardrail, - }, - ], - }} - trigger={["click"]} - disabled={!accessToken} - > - - -
- - {selectedGuardrailId ? ( - setSelectedGuardrailId(null)} - accessToken={accessToken} - isAdmin={isAdmin} - /> - ) : ( - setSelectedGuardrailId(id)} - /> - )} - - - - - - + ), }, - ]} - onCancel={handleDeleteCancel} - onOk={handleDeleteConfirm} - confirmLoading={isDeleting} - /> -
+ { + key: "guardrails", + label: "Guardrails", + children: ( + <> +
+ , + label: "Add Provider Guardrail", + onClick: handleAddGuardrail, + }, + { + key: "custom_code", + icon: , + label: "Create Custom Code Guardrail", + onClick: handleAddCustomCodeGuardrail, + }, + ], + }} + trigger={["click"]} + disabled={!accessToken} + > + + +
- {/* Test Playground Tab */} - - setActiveTab(0)} - /> - + {selectedGuardrailId ? ( + setSelectedGuardrailId(null)} + accessToken={accessToken} + isAdmin={isAdmin} + /> + ) : ( + setSelectedGuardrailId(id)} + /> + )} - {/* Team Guardrails Tab */} - - - -
-
+ + + + + + + ), + }, + { + key: "playground", + label: "Test Playground", + disabled: !accessToken, + children: ( + {}} + /> + ), + }, + ] + : []), + { + key: "submitted", + label: "Submitted Guardrails", + children: , + }, + ]} + />
); }; diff --git a/ui/litellm-dashboard/src/components/guardrails/TeamGuardrailsTab.tsx b/ui/litellm-dashboard/src/components/guardrails/TeamGuardrailsTab.tsx index 8fbbae56124..b03ac92ba2b 100644 --- a/ui/litellm-dashboard/src/components/guardrails/TeamGuardrailsTab.tsx +++ b/ui/litellm-dashboard/src/components/guardrails/TeamGuardrailsTab.tsx @@ -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(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) {