From bde3f6ae4605d356ac298ce470396c1b89ec535d Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Fri, 4 Sep 2026 19:35:43 -0700 Subject: [PATCH 01/35] feat(ui): show guardrail usage units and cost on the Guardrails Monitor The overview table gains Usage Units and Cost columns plus a Guardrail Cost card, and the detail page gains a Usage & Cost section that breaks units and cost down by counter, team and key. Units the cost map could not price are called out next to the cost they are left out of. Both pages now read /guardrails/usage/* through $api.useQuery so the rows are typed from schema.d.ts; the hand-written PerformanceRow and the untyped fetch helpers are gone. fetchClient resolves fetch per request so integration tests that stub the global see typed-client calls too. Refs LIT-5652 --- .../_components/GuardrailDetail.test.tsx | 59 +++-- .../_components/GuardrailDetail.tsx | 12 +- .../GuardrailUsageBreakdown.test.tsx | 114 ++++++++++ .../_components/GuardrailUsageBreakdown.tsx | 159 ++++++++++++++ .../GuardrailsMonitorView.test.tsx | 20 +- .../_components/GuardrailsOverview.test.tsx | 204 +++++++++++++----- .../_components/GuardrailsOverview.tsx | 121 ++++++++--- .../page.integration.test.tsx | 27 ++- .../guardrails/useGuardrailsUsage.test.ts | 81 +++++++ .../hooks/guardrails/useGuardrailsUsage.ts | 38 ++++ .../GuardrailsMonitor/MetricCard.tsx | 2 +- .../components/GuardrailsMonitor/mockData.ts | 33 --- .../GuardrailsMonitor/usageUnits.test.ts | 55 +++++ .../GuardrailsMonitor/usageUnits.ts | 21 ++ .../src/components/networking.tsx | 57 ----- ui/litellm-dashboard/src/lib/http/api.ts | 12 +- 16 files changed, 789 insertions(+), 226 deletions(-) create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useGuardrailsUsage.test.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useGuardrailsUsage.ts create mode 100644 ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts create mode 100644 ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.ts diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.test.tsx index c7567aa80bb..bbcb8138d52 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.test.tsx @@ -2,12 +2,16 @@ import { describe, it, expect, vi, beforeEach } from "vitest"; import userEvent from "@testing-library/user-event"; import { render, screen, waitFor } from "@testing-library/react"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import type { GuardrailUsageDetail } from "@/app/(dashboard)/hooks/guardrails/useGuardrailsUsage"; import { GuardrailDetail } from "./GuardrailDetail"; -const mockGetGuardrailsUsageDetail = vi.fn(); +const mockUseGuardrailsUsageDetail = vi.fn(); +vi.mock("@/app/(dashboard)/hooks/guardrails/useGuardrailsUsage", () => ({ + useGuardrailsUsageDetail: (...args: unknown[]) => mockUseGuardrailsUsageDetail(...args), +})); + const mockGetGuardrailsUsageLogs = vi.fn(); vi.mock("@/components/networking", () => ({ - getGuardrailsUsageDetail: (...args: unknown[]) => mockGetGuardrailsUsageDetail(...args), getGuardrailsUsageLogs: (...args: unknown[]) => mockGetGuardrailsUsageLogs(...args), })); @@ -19,7 +23,8 @@ vi.mock("./EvaluationSettingsModal", () => ({ EvaluationSettingsModal: ({ open }: { open: boolean }) => (open ?
: null), })); -const detail = { +const detail: GuardrailUsageDetail = { + guardrail_id: "pii-detector", guardrail_name: "pii-detector", description: "Blocks personally identifiable information", status: "warning", @@ -29,12 +34,25 @@ const detail = { failRate: 20, avgScore: 0.4, avgLatency: 180, + trend: "stable", + time_series: [], + usage_units: { sensitiveInformationPolicyUnits: 4 }, + usage_units_daily: [], + usage_units_by_team: { "": { sensitiveInformationPolicyUnits: 4 } }, + usage_units_by_key: { "hash-1": { sensitiveInformationPolicyUnits: 4 } }, + cost: 0.0004, + cost_by_unit: { sensitiveInformationPolicyUnits: 0.0004 }, + cost_by_team: { "": 0.0004 }, + cost_by_key: { "hash-1": 0.0004 }, + untracked_usage_units: {}, }; +const loaded = (data: GuardrailUsageDetail | undefined) => ({ data, isLoading: false, error: null }); + const defaultProps = { guardrailId: "pii-detector", onBack: vi.fn(), - accessToken: "test-token", + accessToken: "test-token" as string | null, startDate: "2026-07-01", endDate: "2026-07-24", }; @@ -49,19 +67,19 @@ function renderDetail(props: Partial = {}) { describe("GuardrailDetail", () => { beforeEach(() => { vi.clearAllMocks(); - mockGetGuardrailsUsageDetail.mockResolvedValue(detail); + mockUseGuardrailsUsageDetail.mockReturnValue(loaded(detail)); mockGetGuardrailsUsageLogs.mockResolvedValue({ logs: [], total: 0 }); }); it("should show a busy indicator while the detail request is in flight", () => { - mockGetGuardrailsUsageDetail.mockReturnValue(new Promise(() => {})); + mockUseGuardrailsUsageDetail.mockReturnValue({ data: undefined, isLoading: true, error: null }); renderDetail(); expect(document.querySelector('[aria-busy="true"]')).toBeInTheDocument(); expect(screen.queryByText("pii-detector")).not.toBeInTheDocument(); }); it("should show an error message and a way back when the detail request fails", async () => { - mockGetGuardrailsUsageDetail.mockRejectedValue(new Error("boom")); + mockUseGuardrailsUsageDetail.mockReturnValue({ data: undefined, isLoading: false, error: new Error("boom") }); renderDetail(); expect(await screen.findByText("Failed to load guardrail details.")).toBeInTheDocument(); expect(screen.getByRole("button", { name: /back to overview/i })).toBeInTheDocument(); @@ -69,14 +87,13 @@ describe("GuardrailDetail", () => { it("should request the detail and the logs for the guardrail and date range", async () => { renderDetail(); - await waitFor(() => - expect(mockGetGuardrailsUsageDetail).toHaveBeenCalledWith( - "test-token", - "pii-detector", - "2026-07-01", - "2026-07-24", - ), - ); + expect(mockUseGuardrailsUsageDetail).toHaveBeenCalledWith({ + accessToken: "test-token", + guardrailId: "pii-detector", + startDate: "2026-07-01", + endDate: "2026-07-24", + }); + await waitFor(() => expect(mockGetGuardrailsUsageLogs).toHaveBeenCalled()); expect(mockGetGuardrailsUsageLogs).toHaveBeenCalledWith( "test-token", expect.objectContaining({ guardrailId: "pii-detector", startDate: "2026-07-01", endDate: "2026-07-24" }), @@ -100,11 +117,18 @@ describe("GuardrailDetail", () => { }); it("should show a placeholder when no latency has been recorded", async () => { - mockGetGuardrailsUsageDetail.mockResolvedValue({ ...detail, avgLatency: null }); + mockUseGuardrailsUsageDetail.mockReturnValue(loaded({ ...detail, avgLatency: null })); renderDetail(); expect(await screen.findByText("No data")).toBeInTheDocument(); }); + it("should show the usage and cost breakdown for the guardrail on the overview tab", async () => { + renderDetail(); + const section = await screen.findByRole("region", { name: "Usage and cost" }); + expect(section).toHaveTextContent("$0.0004"); + expect(section).toHaveTextContent("Sensitive Information Policy"); + }); + it("should call onBack when 'Back to Overview' is clicked", async () => { const user = userEvent.setup(); const onBack = vi.fn(); @@ -138,8 +162,9 @@ describe("GuardrailDetail", () => { }); it("should not request anything without an access token", () => { + mockUseGuardrailsUsageDetail.mockReturnValue(loaded(undefined)); renderDetail({ accessToken: null }); - expect(mockGetGuardrailsUsageDetail).not.toHaveBeenCalled(); + expect(mockUseGuardrailsUsageDetail).toHaveBeenCalledWith(expect.objectContaining({ accessToken: null })); expect(mockGetGuardrailsUsageLogs).not.toHaveBeenCalled(); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.tsx index 253477ffeac..1e82f1fee85 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.tsx @@ -1,13 +1,15 @@ import { useQuery } from "@tanstack/react-query"; import { ArrowLeft, Settings, Shield, TriangleAlert } from "lucide-react"; import React, { useMemo, useState } from "react"; -import { getGuardrailsUsageDetail, getGuardrailsUsageLogs } from "@/components/networking"; +import { getGuardrailsUsageLogs } from "@/components/networking"; +import { useGuardrailsUsageDetail } from "@/app/(dashboard)/hooks/guardrails/useGuardrailsUsage"; import { StatusBadge, type StatusTone } from "@/components/shared/table_cells/status_badge"; import { Badge } from "@/components/ui/badge"; import { Button } from "@/components/ui/button"; import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; import { EvaluationSettingsModal } from "./EvaluationSettingsModal"; +import { GuardrailUsageBreakdown } from "./GuardrailUsageBreakdown"; import { LogViewer } from "@/components/GuardrailsMonitor/LogViewer"; import { MetricCard } from "@/components/GuardrailsMonitor/MetricCard"; import type { LogEntry } from "@/components/GuardrailsMonitor/mockData"; @@ -36,11 +38,7 @@ export function GuardrailDetail({ guardrailId, onBack, accessToken = null, start data: detailData, isLoading: detailLoading, error: detailError, - } = useQuery({ - queryKey: ["guardrails-usage-detail", guardrailId, startDate, endDate], - queryFn: () => getGuardrailsUsageDetail(accessToken!, guardrailId, startDate, endDate), - enabled: !!accessToken && !!guardrailId, - }); + } = useGuardrailsUsageDetail({ accessToken, guardrailId, startDate, endDate }); const { data: logsData, isLoading: logsLoading } = useQuery({ queryKey: ["guardrails-usage-logs", guardrailId, logsPage, logsPageSize], queryFn: () => @@ -194,6 +192,8 @@ export function GuardrailDetail({ guardrailId, onBack, accessToken = null, start />
+ {detailData && } + {logViewer("all")} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx new file mode 100644 index 00000000000..3b24185f0b6 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx @@ -0,0 +1,114 @@ +import { render, screen, within } from "@testing-library/react"; +import { describe, expect, it } from "vitest"; +import type { GuardrailUsageDetail } from "@/app/(dashboard)/hooks/guardrails/useGuardrailsUsage"; +import { GuardrailUsageBreakdown } from "./GuardrailUsageBreakdown"; + +const detail: GuardrailUsageDetail = { + guardrail_id: "bedrock-pii-mask", + guardrail_name: "bedrock-pii-mask", + type: "pii", + provider: "Bedrock", + requestsEvaluated: 5, + failRate: 0, + avgScore: null, + avgLatency: 120, + status: "healthy", + trend: "stable", + description: null, + time_series: [], + usage_units: { contentPolicyUnits: 1000, sensitiveInformationPolicyUnits: 300, someFutureCounter: 7 }, + usage_units_daily: [], + usage_units_by_team: { + "team-a": { contentPolicyUnits: 900, sensitiveInformationPolicyUnits: 300 }, + "": { contentPolicyUnits: 100, someFutureCounter: 7 }, + }, + usage_units_by_key: { + "hash-1": { contentPolicyUnits: 1000, sensitiveInformationPolicyUnits: 300 }, + "hash-2": { someFutureCounter: 7 }, + }, + cost: 0.18, + cost_by_unit: { contentPolicyUnits: 0.15, sensitiveInformationPolicyUnits: 0.03, someFutureCounter: null }, + cost_by_team: { "team-a": 0.165, "": 0.015 }, + cost_by_key: { "hash-1": 0.18, "hash-2": null }, + untracked_usage_units: { someFutureCounter: 7 }, +}; + +const rowNamed = (name: string) => screen.getByRole("row", { name: new RegExp(name) }); + +describe("GuardrailUsageBreakdown", () => { + it("totals the units and the cost, and says how many units the cost leaves out", () => { + render(); + + const cost = screen.getByRole("group", { name: "Cost" }); + expect(cost).toHaveTextContent("$0.1800"); + expect(cost).toHaveTextContent("7 units unpriced"); + + const units = screen.getByRole("group", { name: "Usage Units" }); + expect(units).toHaveTextContent("1,307"); + expect(units).toHaveTextContent("3 counters"); + }); + + it("lists each counter with its units, cost and unpriced share", () => { + render(); + + const content = rowNamed("Content Policy"); + expect(within(content).getByText("1,000")).toBeInTheDocument(); + expect(within(content).getByText("$0.1500")).toBeInTheDocument(); + expect(within(content).getByText("—")).toBeInTheDocument(); + + const future = rowNamed("Some Future Counter"); + expect(within(future).getByText("7", { selector: ".text-warning" })).toBeInTheDocument(); + expect(within(future).getByText("—")).toBeInTheDocument(); + }); + + it("breaks units and cost down by team and by key, naming the rows without one", () => { + render(); + + expect(screen.getByRole("heading", { name: "By team" })).toBeInTheDocument(); + expect(screen.getByRole("heading", { name: "By key" })).toBeInTheDocument(); + const teamA = rowNamed("team-a"); + expect(within(teamA).getByText("1,200")).toBeInTheDocument(); + expect(within(teamA).getByText("$0.1650")).toBeInTheDocument(); + + const noTeam = rowNamed("No team"); + expect(within(noTeam).getByText("107")).toBeInTheDocument(); + expect(within(noTeam).getByText("$0.0150")).toBeInTheDocument(); + + const unpricedKey = rowNamed("hash-2"); + expect(within(unpricedKey).getByText("7")).toBeInTheDocument(); + expect(within(unpricedKey).getByText("—")).toBeInTheDocument(); + }); + + it("orders teams and keys by units, largest first", () => { + render(); + + const rows = screen.getAllByRole("row").map((row) => row.textContent ?? ""); + expect(rows.findIndex((text) => text.includes("team-a"))).toBeLessThan( + rows.findIndex((text) => text.includes("No team")), + ); + expect(rows.findIndex((text) => text.includes("hash-1"))).toBeLessThan( + rows.findIndex((text) => text.includes("hash-2")), + ); + }); + + it("says so when the window has no billable units instead of rendering empty tables", () => { + render( + , + ); + + expect(screen.getByText("No billable usage units were recorded in this period.")).toBeInTheDocument(); + expect(screen.queryByRole("table")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx new file mode 100644 index 00000000000..1eaab86c506 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx @@ -0,0 +1,159 @@ +import type { ColumnDef } from "@tanstack/react-table"; +import { CircleDollarSign } from "lucide-react"; +import React from "react"; +import type { GuardrailUsageDetail } from "@/app/(dashboard)/hooks/guardrails/useGuardrailsUsage"; +import { MetricCard } from "@/components/GuardrailsMonitor/MetricCard"; +import { counterLabel, formatCost, totalUnits, unpricedSummary } from "@/components/GuardrailsMonitor/usageUnits"; +import { DataTable } from "@/components/shared/DataTable"; +import { IdCell } from "@/components/shared/table_cells/id_cell"; +import { MoneyCell } from "@/components/shared/table_cells/money_cell"; + +interface CounterRow { + counter: string; + units: number; + cost: number | null; + unpriced: number; +} + +interface GroupRow { + id: string; + units: number; + cost: number | null; +} + +const counterRows = (detail: GuardrailUsageDetail): CounterRow[] => + Object.entries(detail.usage_units).map(([counter, units]) => ({ + counter, + units, + cost: detail.cost_by_unit[counter] ?? null, + unpriced: detail.untracked_usage_units[counter] ?? 0, + })); + +const groupRows = ( + unitsByGroup: GuardrailUsageDetail["usage_units_by_team"], + costByGroup: GuardrailUsageDetail["cost_by_team"], +): GroupRow[] => + Object.entries(unitsByGroup) + .map(([id, units]) => ({ id, units: totalUnits(units), cost: costByGroup[id] ?? null })) + .sort((a, b) => b.units - a.units); + +const counterColumns: ColumnDef[] = [ + { header: "Counter", accessorKey: "counter", cell: ({ row }) => counterLabel(row.original.counter) }, + { + header: "Units", + accessorKey: "units", + meta: { numeric: true }, + cell: ({ row }) => row.original.units.toLocaleString(), + }, + { + header: "Cost", + accessorKey: "cost", + meta: { numeric: true }, + cell: ({ row }) => , + }, + { + header: "Unpriced Units", + accessorKey: "unpriced", + meta: { numeric: true }, + cell: ({ row }) => + row.original.unpriced > 0 ? ( + {row.original.unpriced.toLocaleString()} + ) : ( + — + ), + }, +]; + +const groupColumns = (label: string, emptyLabel: string): ColumnDef[] => [ + { + header: label, + accessorKey: "id", + cell: ({ row }) => + row.original.id ? ( + + ) : ( + {emptyLabel} + ), + }, + { + header: "Units", + accessorKey: "units", + meta: { numeric: true }, + cell: ({ row }) => row.original.units.toLocaleString(), + }, + { + header: "Cost", + accessorKey: "cost", + meta: { numeric: true }, + cell: ({ row }) => , + }, +]; + +const teamColumns = groupColumns("Team", "No team"); +const keyColumns = groupColumns("Key", "No key"); + +const TableHeading = ({ title }: { title: string }) => ( +
{title}
+); + +export function GuardrailUsageBreakdown({ detail }: { detail: GuardrailUsageDetail }) { + const counters = counterRows(detail); + const unpriced = unpricedSummary(detail.untracked_usage_units); + + return ( +
+
+
Usage & Cost
+

+ Billable units the provider reported for this guardrail and what LiteLLM priced them at +

+
+ + {counters.length === 0 ? ( +

No billable usage units were recorded in this period.

+ ) : ( + <> +
+ } + subtitle={unpriced ?? undefined} + /> + +
+ + row.counter} + size="compact" + toolbar={() => } + /> + +
+ row.id || "no-team"} + size="compact" + toolbar={() => } + /> + row.id || "no-key"} + size="compact" + toolbar={() => } + /> +
+ + )} +
+ ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsMonitorView.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsMonitorView.test.tsx index 9f27daab6b3..83b3a8f5d58 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsMonitorView.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsMonitorView.test.tsx @@ -2,14 +2,15 @@ import { render, screen, waitFor } from "@testing-library/react"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { describe, expect, it, vi } from "vitest"; import GuardrailsMonitorView from "./GuardrailsMonitorView"; -import * as networking from "@/components/networking"; vi.mock("@/components/networking", () => ({ - getGuardrailsUsageOverview: vi.fn(), formatDate: vi.fn((d: Date) => d.toISOString().slice(0, 10)), })); -const mockGetGuardrailsUsageOverview = vi.mocked(networking.getGuardrailsUsageOverview); +const mockUseGuardrailsUsageOverview = vi.fn(); +vi.mock("@/app/(dashboard)/hooks/guardrails/useGuardrailsUsage", () => ({ + useGuardrailsUsageOverview: (...args: unknown[]) => mockUseGuardrailsUsageOverview(...args), +})); function wrapper({ children }: { children: React.ReactNode }) { const queryClient = new QueryClient({ @@ -22,23 +23,20 @@ function wrapper({ children }: { children: React.ReactNode }) { describe("GuardrailsMonitorView", () => { it("should render overview and fetch guardrails usage when accessToken is provided", async () => { - mockGetGuardrailsUsageOverview.mockResolvedValue({ - rows: [], - chart: [], - totalRequests: 0, - totalBlocked: 0, - passRate: 100, - }); + mockUseGuardrailsUsageOverview.mockReturnValue({ data: undefined, isLoading: true, error: null }); render(, { wrapper }); expect(await screen.findByRole("heading", { name: /Guardrails Monitor/i })).toBeInTheDocument(); await waitFor(() => { - expect(mockGetGuardrailsUsageOverview).toHaveBeenCalled(); + expect(mockUseGuardrailsUsageOverview).toHaveBeenCalledWith( + expect.objectContaining({ accessToken: "test-token", startDate: expect.any(String) }), + ); }); }); it("should render without crashing when accessToken is null", async () => { + mockUseGuardrailsUsageOverview.mockReturnValue({ data: undefined, isLoading: false, error: null }); render(, { wrapper }); expect(await screen.findByRole("heading", { name: /Guardrails Monitor/i })).toBeInTheDocument(); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx index c62505cc74f..3a56667b156 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx @@ -1,12 +1,15 @@ -import { render, screen, waitFor } from "@testing-library/react"; -import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { render, screen, waitFor, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; -import * as networking from "@/components/networking"; +import type { + GuardrailUsageOverview, + GuardrailUsageOverviewRow, +} from "@/app/(dashboard)/hooks/guardrails/useGuardrailsUsage"; import { GuardrailsOverview } from "./GuardrailsOverview"; -vi.mock("@/components/networking", () => ({ - getGuardrailsUsageOverview: vi.fn(), +const useGuardrailsUsageOverviewMock = vi.fn(); +vi.mock("@/app/(dashboard)/hooks/guardrails/useGuardrailsUsage", () => ({ + useGuardrailsUsageOverview: (...args: unknown[]) => useGuardrailsUsageOverviewMock(...args), })); vi.mock("./ScoreChart", () => ({ @@ -17,16 +20,63 @@ vi.mock("./EvaluationSettingsModal", () => ({ EvaluationSettingsModal: ({ open }: { open: boolean }) => (open ?
Evaluation settings modal
: null), })); -const mockGetGuardrailsUsageOverview = vi.mocked(networking.getGuardrailsUsageOverview); +const row = (overrides: Partial): GuardrailUsageOverviewRow => ({ + id: "guardrail", + name: "Guardrail", + type: "content_filter", + provider: "LiteLLM", + requestsEvaluated: 0, + failRate: 0, + avgScore: null, + avgLatency: null, + status: "healthy", + trend: "stable", + usageUnits: {}, + cost: null, + untrackedUsageUnits: {}, + ...overrides, +}); -function wrapper({ children }: { children: React.ReactNode }) { - const queryClient = new QueryClient({ - defaultOptions: { - queries: { retry: false }, - }, - }); - return {children}; -} +const overview: GuardrailUsageOverview = { + rows: [ + row({ + id: "guardrail-low", + name: "Low Failure Guardrail", + requestsEvaluated: 1200, + failRate: 2.5, + avgLatency: 45, + trend: "down", + }), + row({ + id: "guardrail-high", + name: "High Failure Guardrail", + provider: "Bedrock", + requestsEvaluated: 300, + failRate: 18, + status: "warning", + trend: "up", + usageUnits: { contentPolicyUnits: 1000, sensitiveInformationPolicyUnits: 250 }, + cost: 0.15, + untrackedUsageUnits: { sensitiveInformationPolicyUnits: 250 }, + }), + row({ + id: "guardrail-free", + name: "Free Bedrock Guardrail", + provider: "Bedrock", + requestsEvaluated: 10, + failRate: 0, + usageUnits: { contentPolicyUnits: 40 }, + cost: 0, + }), + ], + chart: [], + totalRequests: 1510, + totalBlocked: 84, + passRate: 94.4, + totalUsageUnits: { contentPolicyUnits: 1040, sensitiveInformationPolicyUnits: 250 }, + totalCost: 0.15, + totalUntrackedUsageUnits: { sensitiveInformationPolicyUnits: 250 }, +}; function renderOverview(onSelectGuardrail = vi.fn()) { return render( @@ -36,41 +86,24 @@ function renderOverview(onSelectGuardrail = vi.fn()) { endDate="2026-08-12" onSelectGuardrail={onSelectGuardrail} />, - { wrapper }, ); } +const rowNamed = (name: string) => screen.getByRole("row", { name: new RegExp(name) }); + describe("GuardrailsOverview", () => { beforeEach(() => { vi.clearAllMocks(); - mockGetGuardrailsUsageOverview.mockResolvedValue({ - rows: [ - { - id: "guardrail-low", - name: "Low Failure Guardrail", - type: "content_filter", - provider: "LiteLLM", - requestsEvaluated: 1200, - failRate: 2.5, - avgLatency: 45, - status: "healthy", - trend: "down", - }, - { - id: "guardrail-high", - name: "High Failure Guardrail", - type: "content_filter", - provider: "Bedrock", - requestsEvaluated: 300, - failRate: 18, - status: "warning", - trend: "up", - }, - ], - chart: [], - totalRequests: 1500, - totalBlocked: 84, - passRate: 94.4, + useGuardrailsUsageOverviewMock.mockReturnValue({ data: overview, isLoading: false, error: null }); + }); + + it("asks for the usage overview of the selected window", () => { + renderOverview(); + + expect(useGuardrailsUsageOverviewMock).toHaveBeenCalledWith({ + accessToken: "test-token", + startDate: "2026-08-01", + endDate: "2026-08-12", }); }); @@ -78,15 +111,7 @@ describe("GuardrailsOverview", () => { const onSelectGuardrail = vi.fn(); const user = userEvent.setup(); - render( - , - { wrapper }, - ); + renderOverview(onSelectGuardrail); expect(await screen.findByRole("columnheader", { name: "Guardrail" })).toBeInTheDocument(); expect(screen.getByRole("columnheader", { name: /Requests/ })).toBeInTheDocument(); @@ -105,6 +130,46 @@ describe("GuardrailsOverview", () => { expect(onSelectGuardrail).toHaveBeenCalledWith("guardrail-low"); }); + it("shows each guardrail's usage units and cost, marking the units cost leaves out", async () => { + renderOverview(); + + expect(await screen.findByRole("columnheader", { name: "Usage Units" })).toBeInTheDocument(); + expect(screen.getByRole("columnheader", { name: /Cost/ })).toBeInTheDocument(); + + const priced = rowNamed("High Failure Guardrail"); + expect(within(priced).getByText("1,250")).toBeInTheDocument(); + expect(within(priced).getByText("$0.1500")).toBeInTheDocument(); + expect(within(priced).getByLabelText("250 units unpriced")).toBeInTheDocument(); + + const free = rowNamed("Free Bedrock Guardrail"); + expect(within(free).getByText("40")).toBeInTheDocument(); + expect(within(free).getByText("$0.0000")).toBeInTheDocument(); + expect(within(free).queryByLabelText(/unpriced/)).not.toBeInTheDocument(); + + const unmetered = rowNamed("Low Failure Guardrail"); + expect(within(unmetered).getAllByText("—")).toHaveLength(2); + }); + + it("breaks the usage units down per counter on hover", async () => { + const user = userEvent.setup(); + renderOverview(); + + await user.hover(within(rowNamed("High Failure Guardrail")).getByText("1,250")); + + expect(await screen.findByText("Content Policy: 1,000")).toBeInTheDocument(); + expect(screen.getByText("Sensitive Information Policy: 250")).toBeInTheDocument(); + }); + + it("sorts by cost when its header is clicked", async () => { + const user = userEvent.setup(); + renderOverview(); + + await user.click(await screen.findByRole("button", { name: /Cost/ })); + + await waitFor(() => expect(screen.getAllByRole("row")[1]).toHaveTextContent("Low Failure Guardrail")); + expect(screen.getAllByRole("row")[3]).toHaveTextContent("High Failure Guardrail"); + }); + it("renders the page header and the export action", async () => { renderOverview(); @@ -117,15 +182,36 @@ describe("GuardrailsOverview", () => { it("renders every summary metric card", async () => { renderOverview(); - expect(await screen.findByText("1,500")).toBeInTheDocument(); + expect(await screen.findByText("1,510")).toBeInTheDocument(); expect(screen.getByText("Total Evaluations")).toBeInTheDocument(); expect(screen.getByText("Blocked Requests")).toBeInTheDocument(); expect(screen.getByText("84")).toBeInTheDocument(); expect(screen.getByText("Pass Rate")).toBeInTheDocument(); expect(screen.getByText("94.4%")).toBeInTheDocument(); - expect(screen.getByText("23ms")).toBeInTheDocument(); + expect(screen.getByText("15ms")).toBeInTheDocument(); expect(screen.getByText("Active Guardrails")).toBeInTheDocument(); - expect(screen.getByText("2")).toBeInTheDocument(); + expect(screen.getByText("3")).toBeInTheDocument(); + }); + + it("totals guardrail cost across the window and says how many units it leaves out", async () => { + renderOverview(); + + const card = await screen.findByRole("group", { name: "Guardrail Cost" }); + expect(card).toHaveTextContent("$0.1500"); + expect(card).toHaveTextContent("250 units unpriced"); + }); + + it("shows a dash for guardrail cost when nothing in the window was priced", async () => { + useGuardrailsUsageOverviewMock.mockReturnValue({ + data: { ...overview, totalCost: null, totalUntrackedUsageUnits: {} }, + isLoading: false, + error: null, + }); + renderOverview(); + + const card = await screen.findByRole("group", { name: "Guardrail Cost" }); + expect(card).toHaveTextContent("—"); + expect(card).not.toHaveTextContent("unpriced"); }); it("renders the table toolbar heading and its description", async () => { @@ -147,14 +233,18 @@ describe("GuardrailsOverview", () => { }); it("marks the overview busy while the usage request is in flight", async () => { - mockGetGuardrailsUsageOverview.mockReturnValue(new Promise(() => {})); + useGuardrailsUsageOverviewMock.mockReturnValue({ data: undefined, isLoading: true, error: null }); renderOverview(); await waitFor(() => expect(document.querySelector('[aria-busy="true"]')).toBeInTheDocument()); }); it("shows a failure message when the usage request rejects", async () => { - mockGetGuardrailsUsageOverview.mockRejectedValue(new Error("network down")); + useGuardrailsUsageOverviewMock.mockReturnValue({ + data: undefined, + isLoading: false, + error: new Error("network down"), + }); renderOverview(); expect(await screen.findByText("Failed to load data. Try again.")).toBeInTheDocument(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx index 5bc9eb16cee..67630fef13a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx @@ -1,10 +1,14 @@ -import { useQuery } from "@tanstack/react-query"; import type { ColumnDef, OnChangeFn, SortingState } from "@tanstack/react-table"; -import { Download, HeartPulse, Settings, TrendingUp, TriangleAlert } from "lucide-react"; +import { CircleDollarSign, Download, HeartPulse, Settings, TrendingUp, TriangleAlert } from "lucide-react"; import React, { useMemo, useState } from "react"; import { DataTable, DataTableSortHeader } from "@/components/shared/DataTable"; -import { getGuardrailsUsageOverview } from "@/components/networking"; -import { type PerformanceRow } from "@/components/GuardrailsMonitor/mockData"; +import { MoneyCell } from "@/components/shared/table_cells/money_cell"; +import { CellTooltip } from "@/components/shared/table_cells/cell_tooltip"; +import { + type GuardrailUsageOverviewRow, + useGuardrailsUsageOverview, +} from "@/app/(dashboard)/hooks/guardrails/useGuardrailsUsage"; +import { counterLabel, formatCost, totalUnits, unpricedSummary } from "@/components/GuardrailsMonitor/usageUnits"; import { Button } from "@/components/ui/button"; import { PageHeader } from "@/components/shared/PageHeader"; import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; @@ -20,7 +24,7 @@ interface GuardrailsOverviewProps { dateRangeControl?: React.ReactNode; } -type SortKey = "failRate" | "requestsEvaluated" | "avgLatency" | "falsePositiveRate" | "falseNegativeRate"; +type SortKey = "failRate" | "requestsEvaluated" | "avgLatency" | "cost"; const providerColors: Record = { Bedrock: "bg-warning/15 text-warning border-warning/20", @@ -30,14 +34,48 @@ const providerColors: Record = { Custom: "bg-muted text-muted-foreground border-border", }; -function computeMetricsFromRows(data: PerformanceRow[]) { - const totalRequests = data.reduce((sum, r) => sum + r.requestsEvaluated, 0); - const totalBlocked = data.reduce((sum, r) => sum + Math.round((r.requestsEvaluated * r.failRate) / 100), 0); - const passRate = totalRequests > 0 ? ((1 - totalBlocked / totalRequests) * 100).toFixed(1) : "0"; - const withLat = data.filter((r) => r.avgLatency != null); - const avgLatency = - withLat.length > 0 ? Math.round(withLat.reduce((sum, r) => sum + (r.avgLatency ?? 0), 0) / withLat.length) : 0; - return { totalRequests, totalBlocked, passRate, avgLatency, count: data.length }; +const EMPTY_METRICS = { + totalRequests: 0, + totalBlocked: 0, + passRate: "0", + avgLatency: 0, + count: 0, + totalCost: null as number | null, + unpriced: null as string | null, +}; + +function UsageUnitsCell({ units }: { units: GuardrailUsageOverviewRow["usageUnits"] }) { + const counters = Object.entries(units); + if (counters.length === 0) return —; + return ( + + {counters.map(([counter, n]) => ( +
  • + {counterLabel(counter)}: {n.toLocaleString()} +
  • + ))} + + } + trigger={{totalUnits(units).toLocaleString()}} + /> + ); +} + +function CostCell({ row }: { row: GuardrailUsageOverviewRow }) { + const unpriced = unpricedSummary(row.untrackedUsageUnits); + return ( + + {unpriced && ( + } + /> + )} + + + ); } export function GuardrailsOverview({ @@ -55,26 +93,22 @@ export function GuardrailsOverview({ data: guardrailsData, isLoading: guardrailsLoading, error: guardrailsError, - } = useQuery({ - queryKey: ["guardrails-usage-overview", startDate, endDate], - queryFn: () => getGuardrailsUsageOverview(accessToken!, startDate, endDate), - enabled: !!accessToken, - }); + } = useGuardrailsUsageOverview({ accessToken, startDate, endDate }); - const activeData: PerformanceRow[] = guardrailsData?.rows ?? []; + const activeData: GuardrailUsageOverviewRow[] = useMemo(() => guardrailsData?.rows ?? [], [guardrailsData]); const metrics = useMemo(() => { - if (guardrailsData) { - return { - totalRequests: guardrailsData.totalRequests ?? 0, - totalBlocked: guardrailsData.totalBlocked ?? 0, - passRate: String(guardrailsData.passRate ?? 0), - avgLatency: activeData.length - ? Math.round(activeData.reduce((s, r) => s + (r.avgLatency ?? 0), 0) / activeData.length) - : 0, - count: activeData.length, - }; - } - return computeMetricsFromRows(activeData); + if (!guardrailsData) return EMPTY_METRICS; + return { + totalRequests: guardrailsData.totalRequests, + totalBlocked: guardrailsData.totalBlocked, + passRate: String(guardrailsData.passRate), + avgLatency: activeData.length + ? Math.round(activeData.reduce((s, r) => s + (r.avgLatency ?? 0), 0) / activeData.length) + : 0, + count: activeData.length, + totalCost: guardrailsData.totalCost, + unpriced: unpricedSummary(guardrailsData.totalUntrackedUsageUnits), + }; }, [guardrailsData, activeData]); const chartData = guardrailsData?.chart; const sorted = useMemo(() => { @@ -88,7 +122,7 @@ export function GuardrailsOverview({ const isLoading = guardrailsLoading; const error = guardrailsError; - const columns: ColumnDef[] = [ + const columns: ColumnDef[] = [ { header: "Guardrail", accessorKey: "name", @@ -166,6 +200,20 @@ export function GuardrailsOverview({ ), }, + { + header: "Usage Units", + accessorKey: "usageUnits", + enableSorting: false, + meta: { numeric: true }, + cell: ({ row }) => , + }, + { + header: ({ column }) => , + accessorKey: "cost", + meta: { numeric: true }, + sortDescFirst: false, + cell: ({ row }) => , + }, { header: "Status", accessorKey: "status", @@ -187,7 +235,7 @@ export function GuardrailsOverview({ }, ]; - const sortableKeys: SortKey[] = ["failRate", "requestsEvaluated", "avgLatency"]; + const sortableKeys: SortKey[] = ["failRate", "requestsEvaluated", "avgLatency", "cost"]; const sorting = useMemo(() => [{ id: sortBy, desc: sortDir === "desc" }], [sortBy, sortDir]); const handleSortingChange: OnChangeFn = (updater) => { const nextSorting = typeof updater === "function" ? updater(sorting) : updater; @@ -236,6 +284,13 @@ export function GuardrailsOverview({ metrics.avgLatency > 150 ? "text-destructive" : metrics.avgLatency > 50 ? "text-warning" : "text-success" } /> + } + subtitle={metrics.unpriced ?? undefined} + /> diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/page.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/page.integration.test.tsx index fb521c0b8a3..ce8fda0ea69 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/page.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/page.integration.test.tsx @@ -11,7 +11,23 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ const fetchMock = vi.fn(); -const requestedUrls = () => fetchMock.mock.calls.map(([url]) => String(url)); +const requestUrl = (input: RequestInfo | URL) => (input instanceof Request ? input.url : String(input)); + +const requestedUrls = () => fetchMock.mock.calls.map(([input]) => requestUrl(input)); + +const emptyOverview = { + rows: [], + chart: [], + totalRequests: 0, + totalBlocked: 0, + passRate: 100, + totalUsageUnits: {}, + totalCost: null, + totalUntrackedUsageUnits: {}, +}; + +const jsonResponse = (body: unknown) => + new Response(JSON.stringify(body), { status: 200, headers: { "Content-Type": "application/json" } }); const renderAs = (userRole: string) => { useAuthorizedMock.mockReturnValue({ accessToken: "sk-test", userId: "u1", userRole }); @@ -25,12 +41,9 @@ describe("Guardrails Monitor page access by role", () => { beforeEach(() => { testQueryClient.clear(); vi.clearAllMocks(); - fetchMock.mockResolvedValue({ - ok: true, - status: 200, - statusText: "OK", - json: async () => ({ rows: [], chart: [], totalRequests: 0, totalBlocked: 0, passRate: 100 }), - }); + fetchMock.mockImplementation(async (input: RequestInfo | URL) => + jsonResponse(requestUrl(input).includes("/guardrails/usage/overview") ? emptyOverview : []), + ); vi.stubGlobal("fetch", fetchMock); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useGuardrailsUsage.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useGuardrailsUsage.test.ts new file mode 100644 index 00000000000..f0b2709484d --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useGuardrailsUsage.test.ts @@ -0,0 +1,81 @@ +import { renderHook } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { useGuardrailsUsageDetail, useGuardrailsUsageOverview } from "./useGuardrailsUsage"; + +const useQueryMock = vi.fn(); +vi.mock("@/lib/http/api", () => ({ + $api: { useQuery: (...args: unknown[]) => useQueryMock(...args) }, +})); + +const lastCall = () => { + const calls = useQueryMock.mock.calls; + return calls[calls.length - 1] as [string, string, unknown, { enabled: boolean }]; +}; + +describe("useGuardrailsUsageOverview", () => { + beforeEach(() => { + vi.clearAllMocks(); + useQueryMock.mockReturnValue({ data: undefined }); + }); + + it("queries GET /guardrails/usage/overview with the window as query params", () => { + renderHook(() => useGuardrailsUsageOverview({ accessToken: "sk", startDate: "2026-09-01", endDate: "2026-09-04" })); + + expect(lastCall().slice(0, 3)).toEqual([ + "get", + "/guardrails/usage/overview", + { params: { query: { start_date: "2026-09-01", end_date: "2026-09-04" } } }, + ]); + expect(lastCall()[3].enabled).toBe(true); + }); + + it("omits blank dates so the proxy applies its default window", () => { + renderHook(() => useGuardrailsUsageOverview({ accessToken: "sk", startDate: "", endDate: "" })); + + expect(lastCall()[2]).toEqual({ params: { query: { start_date: undefined, end_date: undefined } } }); + }); + + it("stays disabled without an access token", () => { + renderHook(() => useGuardrailsUsageOverview({ accessToken: null, startDate: "2026-09-01", endDate: "2026-09-04" })); + + expect(lastCall()[3].enabled).toBe(false); + }); +}); + +describe("useGuardrailsUsageDetail", () => { + beforeEach(() => { + vi.clearAllMocks(); + useQueryMock.mockReturnValue({ data: undefined }); + }); + + it("queries GET /guardrails/usage/detail/{guardrail_id} with the id as a path param", () => { + renderHook(() => + useGuardrailsUsageDetail({ + accessToken: "sk", + guardrailId: "bedrock-pii-mask", + startDate: "2026-09-01", + endDate: "2026-09-04", + }), + ); + + expect(lastCall().slice(0, 3)).toEqual([ + "get", + "/guardrails/usage/detail/{guardrail_id}", + { + params: { + path: { guardrail_id: "bedrock-pii-mask" }, + query: { start_date: "2026-09-01", end_date: "2026-09-04" }, + }, + }, + ]); + expect(lastCall()[3].enabled).toBe(true); + }); + + it("stays disabled without a guardrail id", () => { + renderHook(() => + useGuardrailsUsageDetail({ accessToken: "sk", guardrailId: "", startDate: "2026-09-01", endDate: "2026-09-04" }), + ); + + expect(lastCall()[3].enabled).toBe(false); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useGuardrailsUsage.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useGuardrailsUsage.ts new file mode 100644 index 00000000000..dc7f58fbc8f --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useGuardrailsUsage.ts @@ -0,0 +1,38 @@ +import { $api } from "@/lib/http/api"; +import type { components } from "@/lib/http/schema"; + +export type GuardrailUsageOverview = components["schemas"]["UsageOverviewResponse"]; +export type GuardrailUsageOverviewRow = components["schemas"]["UsageOverviewRow"]; +export type GuardrailUsageDetail = components["schemas"]["UsageDetailResponse"]; + +export interface GuardrailsUsageWindow { + accessToken: string | null; + startDate: string; + endDate: string; +} + +const dateQuery = (startDate: string, endDate: string) => ({ + start_date: startDate || undefined, + end_date: endDate || undefined, +}); + +export const useGuardrailsUsageOverview = ({ accessToken, startDate, endDate }: GuardrailsUsageWindow) => + $api.useQuery( + "get", + "/guardrails/usage/overview", + { params: { query: dateQuery(startDate, endDate) } }, + { enabled: Boolean(accessToken) }, + ); + +export const useGuardrailsUsageDetail = ({ + accessToken, + guardrailId, + startDate, + endDate, +}: GuardrailsUsageWindow & { guardrailId: string }) => + $api.useQuery( + "get", + "/guardrails/usage/detail/{guardrail_id}", + { params: { path: { guardrail_id: guardrailId }, query: dateQuery(startDate, endDate) } }, + { enabled: Boolean(accessToken && guardrailId) }, + ); diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/MetricCard.tsx b/ui/litellm-dashboard/src/components/GuardrailsMonitor/MetricCard.tsx index d5d249e4799..c0b5e0a50d1 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/MetricCard.tsx +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/MetricCard.tsx @@ -10,7 +10,7 @@ interface MetricCardProps { export function MetricCard({ label, value, valueColor = "text-foreground", icon, subtitle }: MetricCardProps) { return ( -
    +
    {label} {icon && {icon}} diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/mockData.ts b/ui/litellm-dashboard/src/components/GuardrailsMonitor/mockData.ts index 7d99ebe7c44..2b42f7907f1 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/mockData.ts +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/mockData.ts @@ -2,39 +2,6 @@ * Types for Guardrails Monitor dashboard (data from usage API). */ -export interface PerformanceRow { - id: string; - name: string; - type: string; - provider: string; - requestsEvaluated: number; - failRate: number; - avgScore?: number; - avgLatency?: number; - p95Latency?: number; - falsePositiveRate?: number; - falseNegativeRate?: number; - status: "healthy" | "warning" | "critical"; - trend: "up" | "down" | "stable"; -} - -export interface GuardrailDetailRecord { - name: string; - type: string; - provider: string; - requestsEvaluated: number; - failRate: number; - avgScore?: number; - avgLatency?: number; - p95Latency?: number; - falsePositiveRate?: number; - falsePositiveCount?: number; - falseNegativeRate?: number; - falseNegativeCount?: number; - status: string; - description: string; -} - export interface LogEntry { id: string; timestamp: string; diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts new file mode 100644 index 00000000000..29362cc3701 --- /dev/null +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts @@ -0,0 +1,55 @@ +import { describe, expect, it } from "vitest"; +import { counterLabel, formatCost, totalUnits, unpricedSummary } from "./usageUnits"; + +describe("formatCost", () => { + it("renders a dash when nothing was priced", () => { + expect(formatCost(null)).toBe("—"); + expect(formatCost(undefined)).toBe("—"); + }); + + it("keeps an explicit zero as a real price rather than a dash", () => { + expect(formatCost(0)).toBe("$0.0000"); + }); + + it("shows four decimals for the sub-cent amounts guardrail units cost", () => { + expect(formatCost(0.0003)).toBe("$0.0003"); + expect(formatCost(12.5)).toBe("$12.5000"); + }); + + it("flags amounts below the displayed precision instead of rounding them to zero", () => { + expect(formatCost(0.00001)).toBe("< $0.0001"); + }); +}); + +describe("totalUnits", () => { + it("sums every counter", () => { + expect(totalUnits({ contentPolicyUnits: 3, sensitiveInformationPolicyUnits: 4 })).toBe(7); + }); + + it("is zero for no counters", () => { + expect(totalUnits({})).toBe(0); + }); +}); + +describe("counterLabel", () => { + it("turns a Bedrock counter name into words without the Units suffix", () => { + expect(counterLabel("sensitiveInformationPolicyUnits")).toBe("Sensitive Information Policy"); + expect(counterLabel("contentPolicyUnits")).toBe("Content Policy"); + }); + + it("leaves a name it cannot split alone apart from capitalising it", () => { + expect(counterLabel("units")).toBe("Units"); + }); +}); + +describe("unpricedSummary", () => { + it("is null when every unit was priced", () => { + expect(unpricedSummary({})).toBeNull(); + expect(unpricedSummary({ contentPolicyUnits: 0 })).toBeNull(); + }); + + it("counts unpriced units across counters with a pluralised label", () => { + expect(unpricedSummary({ contentPolicyUnits: 1200, someFutureCounter: 34 })).toBe("1,234 units unpriced"); + expect(unpricedSummary({ someFutureCounter: 1 })).toBe("1 unit unpriced"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.ts b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.ts new file mode 100644 index 00000000000..f3a5e9d7140 --- /dev/null +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.ts @@ -0,0 +1,21 @@ +import { formatNumberWithCommas, getSpendString } from "@/utils/dataUtils"; + +export type UsageUnits = Readonly>; + +export const formatCost = (cost: number | null | undefined): string => { + if (cost == null) return "—"; + return cost === 0 ? `$${formatNumberWithCommas(0, 4)}` : getSpendString(cost, 4); +}; + +export const totalUnits = (units: UsageUnits): number => Object.values(units).reduce((sum, n) => sum + n, 0); + +export const counterLabel = (counter: string): string => + counter + .replace(/Units$/, "") + .replace(/([a-z0-9])([A-Z])/g, "$1 $2") + .replace(/^./, (c) => c.toUpperCase()); + +export const unpricedSummary = (untracked: UsageUnits): string | null => { + const total = totalUnits(untracked); + return total > 0 ? `${total.toLocaleString()} ${total === 1 ? "unit" : "units"} unpriced` : null; +}; diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index e06d457cfa9..ccf4bd748e5 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -3956,63 +3956,6 @@ export const rejectGuardrailSubmission = async ( }; // Guardrails / Policies usage (dashboard) -export const getGuardrailsUsageOverview = async (accessToken: string, startDate?: string, endDate?: string) => { - try { - let url = proxyBaseUrl ? `${proxyBaseUrl}/guardrails/usage/overview` : `/guardrails/usage/overview`; - const params = new URLSearchParams(); - if (startDate) params.append("start_date", startDate); - if (endDate) params.append("end_date", endDate); - if (params.toString()) url += `?${params.toString()}`; - const response = await fetch(url, { - method: "GET", - headers: { - [globalLitellmHeaderName]: `Bearer ${accessToken}`, - "Content-Type": "application/json", - }, - }); - if (!response.ok) { - const errorData = await response.json(); - throw new Error(deriveErrorMessage(errorData)); - } - return response.json(); - } catch (error) { - console.error("Failed to get guardrails usage overview:", error); - throw error; - } -}; - -export const getGuardrailsUsageDetail = async ( - accessToken: string, - guardrailId: string, - startDate?: string, - endDate?: string, -) => { - try { - let url = proxyBaseUrl - ? `${proxyBaseUrl}/guardrails/usage/detail/${encodeURIComponent(guardrailId)}` - : `/guardrails/usage/detail/${encodeURIComponent(guardrailId)}`; - const params = new URLSearchParams(); - if (startDate) params.append("start_date", startDate); - if (endDate) params.append("end_date", endDate); - if (params.toString()) url += `?${params.toString()}`; - const response = await fetch(url, { - method: "GET", - headers: { - [globalLitellmHeaderName]: `Bearer ${accessToken}`, - "Content-Type": "application/json", - }, - }); - if (!response.ok) { - const errorData = await response.json(); - throw new Error(deriveErrorMessage(errorData)); - } - return response.json(); - } catch (error) { - console.error("Failed to get guardrails usage detail:", error); - throw error; - } -}; - export const getGuardrailsUsageLogs = async ( accessToken: string, options: { diff --git a/ui/litellm-dashboard/src/lib/http/api.ts b/ui/litellm-dashboard/src/lib/http/api.ts index 508a27db78d..904e4e4f3ab 100644 --- a/ui/litellm-dashboard/src/lib/http/api.ts +++ b/ui/litellm-dashboard/src/lib/http/api.ts @@ -43,11 +43,15 @@ const middleware: Middleware = { * * The base URL is injected, not fixed at import: every request is built against * whatever registerBaseUrlGetter supplies at call time (a split-origin proxy or - * worker URL), falling back to the current origin. The middleware injects the - * auth header and maps non-2xx responses to ApiError so query functions can just - * read `.data`. + * worker URL), falling back to the current origin. `fetch` is looked up per + * request for the same reason, so a test that stubs the global sees these calls + * too. The middleware injects the auth header and maps non-2xx responses to + * ApiError so query functions can just read `.data`. */ -export const fetchClient = createFetchClient({ Request: BaseAwareRequest }); +export const fetchClient = createFetchClient({ + Request: BaseAwareRequest, + fetch: (request) => globalThis.fetch(request), +}); fetchClient.use(middleware); /** From 1d375d8ada91eb6f6aeceb8af8bc649a031a5012 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Fri, 4 Sep 2026 19:46:14 -0700 Subject: [PATCH 02/35] fix(guardrails): flag unpriced units per team and key, sort unknown cost last The detail endpoint now returns untracked_usage_units_by_team and untracked_usage_units_by_key next to the cost breakdowns, and the By team and By key tables show them in an Unpriced Units column, so a row that pairs its total units with a partial cost says how many units that cost leaves out. The overview comparator no longer treats a missing cost as zero: guardrails with no known cost sort last in both directions instead of mixing in with genuinely free ones. Refs LIT-5652 --- litellm/proxy/_lazy_openapi_snapshot.json | 24 ++++++++++- litellm/proxy/guardrails/usage_endpoints.py | 20 ++++++++-- .../proxy/guardrails/test_usage_endpoints.py | 8 ++++ .../_components/GuardrailDetail.test.tsx | 2 + .../GuardrailUsageBreakdown.test.tsx | 11 ++++- .../_components/GuardrailUsageBreakdown.tsx | 40 ++++++++++++------- .../_components/GuardrailsOverview.test.tsx | 16 ++++++-- .../_components/GuardrailsOverview.tsx | 9 +++-- ui/litellm-dashboard/src/lib/http/schema.d.ts | 12 ++++++ 9 files changed, 114 insertions(+), 28 deletions(-) diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index c24eea968f8..ddf6a59bea7 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -13218,6 +13218,26 @@ "title": "Untracked Usage Units", "type": "object" }, + "untracked_usage_units_by_key": { + "additionalProperties": { + "additionalProperties": { + "type": "integer" + }, + "type": "object" + }, + "title": "Untracked Usage Units By Key", + "type": "object" + }, + "untracked_usage_units_by_team": { + "additionalProperties": { + "additionalProperties": { + "type": "integer" + }, + "type": "object" + }, + "title": "Untracked Usage Units By Team", + "type": "object" + }, "usage_units": { "additionalProperties": { "type": "integer" @@ -13274,7 +13294,9 @@ "cost_by_unit", "cost_by_team", "cost_by_key", - "untracked_usage_units" + "untracked_usage_units", + "untracked_usage_units_by_team", + "untracked_usage_units_by_key" ], "title": "UsageDetailResponse", "type": "object" diff --git a/litellm/proxy/guardrails/usage_endpoints.py b/litellm/proxy/guardrails/usage_endpoints.py index 0390a2b5013..d3bd09f7d27 100644 --- a/litellm/proxy/guardrails/usage_endpoints.py +++ b/litellm/proxy/guardrails/usage_endpoints.py @@ -156,6 +156,14 @@ def _counter_name(row: "prisma_models.LiteLLM_DailyGuardrailUsageUnits") -> str: return row.usage_unit +def _team_of(row: "prisma_models.LiteLLM_DailyGuardrailUsageUnits") -> str: + return row.team_id + + +def _key_of(row: "prisma_models.LiteLLM_DailyGuardrailUsageUnits") -> str: + return row.api_key + + def _row_untracked_units(row: "prisma_models.LiteLLM_DailyGuardrailUsageUnits") -> int: """A row written before the cost column carries NULL cost and is untracked in full.""" return int(row.units) if row.cost is None else int(row.untracked_units) @@ -308,6 +316,8 @@ class UsageDetailResponse(BaseModel): cost_by_team: Mapping[str, float | None] cost_by_key: Mapping[str, float | None] untracked_usage_units: Mapping[str, int] + untracked_usage_units_by_team: Mapping[str, Mapping[str, int]] + untracked_usage_units_by_key: Mapping[str, Mapping[str, int]] class UsageLogEntry(BaseModel): @@ -705,13 +715,15 @@ async def guardrails_usage_detail( time_series=time_series, usage_units=_sum_counter_units(units_rows), usage_units_daily=units_daily, - usage_units_by_team=_by(units_rows, lambda r: r.team_id, _sum_counter_units), - usage_units_by_key=_by(units_rows, lambda r: r.api_key, _sum_counter_units), + usage_units_by_team=_by(units_rows, _team_of, _sum_counter_units), + usage_units_by_key=_by(units_rows, _key_of, _sum_counter_units), cost=_sum_tracked_cost(units_rows), cost_by_unit=_by(units_rows, _counter_name, _sum_tracked_cost), - cost_by_team=_by(units_rows, lambda r: r.team_id, _sum_tracked_cost), - cost_by_key=_by(units_rows, lambda r: r.api_key, _sum_tracked_cost), + cost_by_team=_by(units_rows, _team_of, _sum_tracked_cost), + cost_by_key=_by(units_rows, _key_of, _sum_tracked_cost), untracked_usage_units=_sum_untracked_units(units_rows), + untracked_usage_units_by_team=_by(units_rows, _team_of, _sum_untracked_units), + untracked_usage_units_by_key=_by(units_rows, _key_of, _sum_untracked_units), ) diff --git a/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py b/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py index ebb2be6edc2..1ff33c76035 100644 --- a/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py @@ -430,6 +430,13 @@ async def test_detail_breaks_cost_down_by_unit_day_team_and_key(): assert resp.cost_by_team.keys() == resp.usage_units_by_team.keys() assert resp.cost_by_key.keys() == resp.usage_units_by_key.keys() assert resp.untracked_usage_units == {"contentPolicyUnits": 50, "topicPolicyUnits": 10} + assert resp.untracked_usage_units_by_team == {"team-a": {"topicPolicyUnits": 10}, "": {"contentPolicyUnits": 50}} + assert resp.untracked_usage_units_by_key == { + "hash-1": {"topicPolicyUnits": 10}, + "hash-2": {"contentPolicyUnits": 50}, + } + assert resp.untracked_usage_units_by_team.keys() == resp.usage_units_by_team.keys() + assert resp.untracked_usage_units_by_key.keys() == resp.usage_units_by_key.keys() @pytest.mark.asyncio @@ -451,6 +458,7 @@ async def test_detail_degrades_units_to_empty_when_units_table_is_missing(): ) assert (resp.cost, resp.cost_by_unit, resp.cost_by_team, resp.cost_by_key) == (None, {}, {}, {}) assert resp.untracked_usage_units == {} + assert (resp.untracked_usage_units_by_team, resp.untracked_usage_units_by_key) == ({}, {}) # ---- logs ------------------------------------------------------------------- diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.test.tsx index bbcb8138d52..3d00e29245d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.test.tsx @@ -45,6 +45,8 @@ const detail: GuardrailUsageDetail = { cost_by_team: { "": 0.0004 }, cost_by_key: { "hash-1": 0.0004 }, untracked_usage_units: {}, + untracked_usage_units_by_team: {}, + untracked_usage_units_by_key: {}, }; const loaded = (data: GuardrailUsageDetail | undefined) => ({ data, isLoading: false, error: null }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx index 3b24185f0b6..7929b97b000 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx @@ -31,6 +31,8 @@ const detail: GuardrailUsageDetail = { cost_by_team: { "team-a": 0.165, "": 0.015 }, cost_by_key: { "hash-1": 0.18, "hash-2": null }, untracked_usage_units: { someFutureCounter: 7 }, + untracked_usage_units_by_team: { "team-a": {}, "": { someFutureCounter: 7 } }, + untracked_usage_units_by_key: { "hash-1": {}, "hash-2": { someFutureCounter: 7 } }, }; const rowNamed = (name: string) => screen.getByRole("row", { name: new RegExp(name) }); @@ -61,7 +63,7 @@ describe("GuardrailUsageBreakdown", () => { expect(within(future).getByText("—")).toBeInTheDocument(); }); - it("breaks units and cost down by team and by key, naming the rows without one", () => { + it("breaks units and cost down by team and by key, flagging the unpriced share of each row", () => { render(); expect(screen.getByRole("heading", { name: "By team" })).toBeInTheDocument(); @@ -69,14 +71,17 @@ describe("GuardrailUsageBreakdown", () => { const teamA = rowNamed("team-a"); expect(within(teamA).getByText("1,200")).toBeInTheDocument(); expect(within(teamA).getByText("$0.1650")).toBeInTheDocument(); + expect(within(teamA).getByText("—")).toBeInTheDocument(); + expect(within(teamA).queryByText("7")).not.toBeInTheDocument(); const noTeam = rowNamed("No team"); expect(within(noTeam).getByText("107")).toBeInTheDocument(); expect(within(noTeam).getByText("$0.0150")).toBeInTheDocument(); + expect(within(noTeam).getByText("7", { selector: ".text-warning" })).toBeInTheDocument(); const unpricedKey = rowNamed("hash-2"); - expect(within(unpricedKey).getByText("7")).toBeInTheDocument(); expect(within(unpricedKey).getByText("—")).toBeInTheDocument(); + expect(within(unpricedKey).getByText("7", { selector: ".text-warning" })).toBeInTheDocument(); }); it("orders teams and keys by units, largest first", () => { @@ -104,6 +109,8 @@ describe("GuardrailUsageBreakdown", () => { cost_by_team: {}, cost_by_key: {}, untracked_usage_units: {}, + untracked_usage_units_by_team: {}, + untracked_usage_units_by_key: {}, }} />, ); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx index 1eaab86c506..6e8b725aa2e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx @@ -19,6 +19,7 @@ interface GroupRow { id: string; units: number; cost: number | null; + unpriced: number; } const counterRows = (detail: GuardrailUsageDetail): CounterRow[] => @@ -32,11 +33,31 @@ const counterRows = (detail: GuardrailUsageDetail): CounterRow[] => const groupRows = ( unitsByGroup: GuardrailUsageDetail["usage_units_by_team"], costByGroup: GuardrailUsageDetail["cost_by_team"], + untrackedByGroup: GuardrailUsageDetail["untracked_usage_units_by_team"], ): GroupRow[] => Object.entries(unitsByGroup) - .map(([id, units]) => ({ id, units: totalUnits(units), cost: costByGroup[id] ?? null })) + .map(([id, units]) => ({ + id, + units: totalUnits(units), + cost: costByGroup[id] ?? null, + unpriced: totalUnits(untrackedByGroup[id] ?? {}), + })) .sort((a, b) => b.units - a.units); +const UnpricedUnitsCell = ({ unpriced }: { unpriced: number }) => + unpriced > 0 ? ( + {unpriced.toLocaleString()} + ) : ( + — + ); + +const unpricedColumn = (): ColumnDef => ({ + header: "Unpriced Units", + accessorKey: "unpriced", + meta: { numeric: true }, + cell: ({ row }) => , +}); + const counterColumns: ColumnDef[] = [ { header: "Counter", accessorKey: "counter", cell: ({ row }) => counterLabel(row.original.counter) }, { @@ -51,17 +72,7 @@ const counterColumns: ColumnDef[] = [ meta: { numeric: true }, cell: ({ row }) => , }, - { - header: "Unpriced Units", - accessorKey: "unpriced", - meta: { numeric: true }, - cell: ({ row }) => - row.original.unpriced > 0 ? ( - {row.original.unpriced.toLocaleString()} - ) : ( - — - ), - }, + unpricedColumn(), ]; const groupColumns = (label: string, emptyLabel: string): ColumnDef[] => [ @@ -87,6 +98,7 @@ const groupColumns = (label: string, emptyLabel: string): ColumnDef[] meta: { numeric: true }, cell: ({ row }) => , }, + unpricedColumn(), ]; const teamColumns = groupColumns("Team", "No team"); @@ -139,14 +151,14 @@ export function GuardrailUsageBreakdown({ detail }: { detail: GuardrailUsageDeta
    row.id || "no-team"} size="compact" toolbar={() => } /> row.id || "no-key"} size="compact" toolbar={() => } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx index 3a56667b156..c52645def70 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx @@ -160,14 +160,24 @@ describe("GuardrailsOverview", () => { expect(screen.getByText("Sensitive Information Policy: 250")).toBeInTheDocument(); }); - it("sorts by cost when its header is clicked", async () => { + it("sorts by cost when its header is clicked, keeping guardrails with no known cost last either way", async () => { const user = userEvent.setup(); renderOverview(); + const rowNames = () => + screen + .getAllByRole("row") + .slice(1) + .map((r) => r.textContent ?? ""); await user.click(await screen.findByRole("button", { name: /Cost/ })); + await waitFor(() => expect(rowNames()[0]).toContain("Free Bedrock Guardrail")); + expect(rowNames()[1]).toContain("High Failure Guardrail"); + expect(rowNames()[2]).toContain("Low Failure Guardrail"); - await waitFor(() => expect(screen.getAllByRole("row")[1]).toHaveTextContent("Low Failure Guardrail")); - expect(screen.getAllByRole("row")[3]).toHaveTextContent("High Failure Guardrail"); + await user.click(screen.getByRole("button", { name: /Cost/ })); + await waitFor(() => expect(rowNames()[0]).toContain("High Failure Guardrail")); + expect(rowNames()[1]).toContain("Free Bedrock Guardrail"); + expect(rowNames()[2]).toContain("Low Failure Guardrail"); }); it("renders the page header and the export action", async () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx index 67630fef13a..0bbda6015b4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx @@ -112,11 +112,12 @@ export function GuardrailsOverview({ }, [guardrailsData, activeData]); const chartData = guardrailsData?.chart; const sorted = useMemo(() => { + const mult = sortDir === "desc" ? -1 : 1; return [...activeData].sort((a, b) => { - const mult = sortDir === "desc" ? -1 : 1; - const aVal = a[sortBy] ?? 0; - const bVal = b[sortBy] ?? 0; - return (Number(aVal) - Number(bVal)) * mult; + const aVal = a[sortBy]; + const bVal = b[sortBy]; + if (aVal == null || bVal == null) return Number(aVal == null) - Number(bVal == null); + return (aVal - bVal) * mult; }); }, [activeData, sortBy, sortDir]); const isLoading = guardrailsLoading; diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 427e5deb555..c1f65299c52 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -38461,6 +38461,18 @@ export interface components { untracked_usage_units: { [key: string]: number; }; + /** Untracked Usage Units By Key */ + untracked_usage_units_by_key: { + [key: string]: { + [key: string]: number; + }; + }; + /** Untracked Usage Units By Team */ + untracked_usage_units_by_team: { + [key: string]: { + [key: string]: number; + }; + }; /** Usage Units */ usage_units: { [key: string]: number; From f73e6838000d7bca1af751fcbcf83fc59d663a9b Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Fri, 4 Sep 2026 20:26:13 -0700 Subject: [PATCH 03/35] chore(ui): drop the fetch lookup note from the fetchClient docblock --- ui/litellm-dashboard/src/lib/http/api.ts | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/ui/litellm-dashboard/src/lib/http/api.ts b/ui/litellm-dashboard/src/lib/http/api.ts index 904e4e4f3ab..9aa6bddf704 100644 --- a/ui/litellm-dashboard/src/lib/http/api.ts +++ b/ui/litellm-dashboard/src/lib/http/api.ts @@ -43,10 +43,9 @@ const middleware: Middleware = { * * The base URL is injected, not fixed at import: every request is built against * whatever registerBaseUrlGetter supplies at call time (a split-origin proxy or - * worker URL), falling back to the current origin. `fetch` is looked up per - * request for the same reason, so a test that stubs the global sees these calls - * too. The middleware injects the auth header and maps non-2xx responses to - * ApiError so query functions can just read `.data`. + * worker URL), falling back to the current origin. The middleware injects the + * auth header and maps non-2xx responses to ApiError so query functions can just + * read `.data`. */ export const fetchClient = createFetchClient({ Request: BaseAwareRequest, From 842623529048ceb836c746f8b99835260142a229 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 4 Sep 2026 21:25:46 -0700 Subject: [PATCH 04/35] feat(ocr): add Cohere Parse support for cohere and azure_ai --- litellm/llms/azure_ai/ocr/__init__.py | 2 + .../ocr/cohere_parse_transformation.py | 91 ++++++ litellm/llms/azure_ai/ocr/common_utils.py | 9 + litellm/llms/base_llm/ocr/transformation.py | 4 + litellm/llms/cohere/ocr/__init__.py | 3 + litellm/llms/cohere/ocr/transformation.py | 292 ++++++++++++++++++ ...odel_prices_and_context_window_backup.json | 19 ++ litellm/ocr/main.py | 2 + litellm/utils.py | 5 + model_prices_and_context_window.json | 19 ++ ...st_azure_ai_cohere_parse_transformation.py | 166 ++++++++++ .../llms/cohere/ocr/test_cohere_parse_cost.py | 58 ++++ .../ocr/test_cohere_parse_transformation.py | 217 +++++++++++++ .../ocr/test_ocr_native_format.py | 10 + 14 files changed, 897 insertions(+) create mode 100644 litellm/llms/azure_ai/ocr/cohere_parse_transformation.py create mode 100644 litellm/llms/cohere/ocr/__init__.py create mode 100644 litellm/llms/cohere/ocr/transformation.py create mode 100644 tests/test_litellm/llms/azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py create mode 100644 tests/test_litellm/llms/cohere/ocr/test_cohere_parse_cost.py create mode 100644 tests/test_litellm/llms/cohere/ocr/test_cohere_parse_transformation.py diff --git a/litellm/llms/azure_ai/ocr/__init__.py b/litellm/llms/azure_ai/ocr/__init__.py index ade1165b848..998d0570882 100644 --- a/litellm/llms/azure_ai/ocr/__init__.py +++ b/litellm/llms/azure_ai/ocr/__init__.py @@ -1,5 +1,6 @@ """Azure AI OCR module.""" +from .cohere_parse_transformation import AzureAICohereParseConfig from .common_utils import get_azure_ai_ocr_config from .document_intelligence.transformation import ( AzureDocumentIntelligenceOCRConfig, @@ -7,6 +8,7 @@ from .document_intelligence.transformation import ( from .transformation import AzureAIOCRConfig __all__ = [ + "AzureAICohereParseConfig", "AzureAIOCRConfig", "AzureDocumentIntelligenceOCRConfig", "get_azure_ai_ocr_config", diff --git a/litellm/llms/azure_ai/ocr/cohere_parse_transformation.py b/litellm/llms/azure_ai/ocr/cohere_parse_transformation.py new file mode 100644 index 00000000000..121f970c59b --- /dev/null +++ b/litellm/llms/azure_ai/ocr/cohere_parse_transformation.py @@ -0,0 +1,91 @@ +"""Cohere Parse served from Azure AI Foundry (`/providers/cohere/v2/parse`).""" + +from collections.abc import Mapping +from typing import Final + +import httpx + +from litellm.litellm_core_utils.prompt_templates.image_handling import ( + async_convert_url_to_base64, + convert_url_to_base64, +) +from litellm.llms.azure_ai.common_utils import get_azure_ai_auth_headers +from litellm.llms.cohere.ocr.transformation import COHERE_PARSE_PATH, CohereParseConfig +from litellm.secret_managers.main import get_secret_str + +AZURE_AI_API_KEY_ENV_VAR: Final = "AZURE_AI_API_KEY" +AZURE_AI_API_BASE_ENV_VAR: Final = "AZURE_AI_API_BASE" +AZURE_AI_COHERE_PROVIDER_PATH: Final = "/providers/cohere" +AZURE_AI_MODELS_PATH_SUFFIX: Final = "/models" + + +class AzureAICohereParseConfig(CohereParseConfig): + """Same request and response shape as Cohere Parse, behind Azure AI auth and URL layout. + + Foundry cannot fetch external URLs, so remote images are inlined as base64 data URIs. + """ + + def get_api_key_env_var(self) -> str | None: + return AZURE_AI_API_KEY_ENV_VAR + + def _llm_provider(self) -> str: + return "azure_ai" + + def validate_environment( + self, + headers: Mapping[str, str], + model: str, + api_key: str | None = None, + api_base: str | None = None, + litellm_params: Mapping[str, object] | None = None, + **kwargs: object, # kwargs-ok: BaseOCRConfig.validate_environment signature + ) -> dict[str, str]: # mutable-ok: BaseOCRConfig signature + resolved_base: Final = api_base or get_secret_str(AZURE_AI_API_BASE_ENV_VAR) + if resolved_base is None: + raise ValueError( + f"Missing Azure AI API Base - Set {AZURE_AI_API_BASE_ENV_VAR} environment variable " + "or pass api_base parameter" + ) + resolved_key: Final = api_key or get_secret_str(AZURE_AI_API_KEY_ENV_VAR) + return { # mutable-ok: BaseOCRConfig signature + **get_azure_ai_auth_headers(api_key=resolved_key, litellm_params=litellm_params), + "Content-Type": "application/json", + **headers, + } + + def get_complete_url( + self, + api_base: str | None, + model: str, + optional_params: Mapping[str, object], + litellm_params: Mapping[str, object] | None = None, + **kwargs: object, # kwargs-ok: BaseOCRConfig.get_complete_url signature + ) -> str: + resolved_base: Final = api_base or get_secret_str(AZURE_AI_API_BASE_ENV_VAR) + if resolved_base is None: + raise ValueError( + f"Missing Azure AI API Base - Set {AZURE_AI_API_BASE_ENV_VAR} environment variable " + "or pass api_base parameter" + ) + url: Final = httpx.URL(resolved_base) + if not url.is_absolute_url: + raise ValueError( + "Azure AI API Base must be an absolute URL including scheme (e.g. " + f"'https://.services.ai.azure.com'). Got api_base={resolved_base!r}." + ) + path: Final = url.path.rstrip("/") + if path.endswith(COHERE_PARSE_PATH): + return str(url.copy_with(path=path)) + if path.endswith(f"{AZURE_AI_COHERE_PROVIDER_PATH}/v2"): + return str(url.copy_with(path=f"{path}/parse")) + return str( + url.copy_with( + path=f"{path.removesuffix(AZURE_AI_MODELS_PATH_SUFFIX)}{AZURE_AI_COHERE_PROVIDER_PATH}{COHERE_PARSE_PATH}" + ) + ) + + def _resolve_image_url_sync(self, image_url: str) -> str: + return convert_url_to_base64(image_url) + + async def _resolve_image_url_async(self, image_url: str) -> str: + return await async_convert_url_to_base64(image_url) diff --git a/litellm/llms/azure_ai/ocr/common_utils.py b/litellm/llms/azure_ai/ocr/common_utils.py index ac1a1f5af0a..a4cd0c7a30b 100644 --- a/litellm/llms/azure_ai/ocr/common_utils.py +++ b/litellm/llms/azure_ai/ocr/common_utils.py @@ -24,6 +24,10 @@ def is_azure_document_intelligence_model(model: str) -> bool: return "doc-intelligence" in lowered or "documentintelligence" in lowered +def is_azure_cohere_parse_model(model: str) -> bool: + return "parse" in model.lower() + + def get_azure_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]: """ Determine which Azure AI OCR configuration to use based on the model name. @@ -46,6 +50,7 @@ def get_azure_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]: >>> get_azure_ai_ocr_config("azure_ai/pixtral-12b-2409") """ + from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig from litellm.llms.azure_ai.ocr.document_intelligence.transformation import ( AzureDocumentIntelligenceOCRConfig, ) @@ -56,6 +61,10 @@ def get_azure_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]: verbose_logger.debug("Routing %s to Azure Document Intelligence OCR config", model) return AzureDocumentIntelligenceOCRConfig() + if is_azure_cohere_parse_model(model): + verbose_logger.debug("Routing %s to Azure AI Cohere Parse config", model) + return AzureAICohereParseConfig() + # Default to Mistral-based OCR for other azure_ai models verbose_logger.debug("Routing %s to Azure AI (Mistral) OCR config", model) return AzureAIOCRConfig() diff --git a/litellm/llms/base_llm/ocr/transformation.py b/litellm/llms/base_llm/ocr/transformation.py index 75306cd572a..08ae077cb2f 100644 --- a/litellm/llms/base_llm/ocr/transformation.py +++ b/litellm/llms/base_llm/ocr/transformation.py @@ -142,6 +142,10 @@ class BaseOCRConfig: """ return None + def supports_rust_bridge(self) -> bool: + """Whether the Rust OCR bridge may serve this config when it is enabled for the provider.""" + return True + def map_ocr_params( self, non_default_params: dict, diff --git a/litellm/llms/cohere/ocr/__init__.py b/litellm/llms/cohere/ocr/__init__.py new file mode 100644 index 00000000000..7742c7e0035 --- /dev/null +++ b/litellm/llms/cohere/ocr/__init__.py @@ -0,0 +1,3 @@ +from litellm.llms.cohere.ocr.transformation import CohereParseConfig + +__all__ = ("CohereParseConfig",) diff --git a/litellm/llms/cohere/ocr/transformation.py b/litellm/llms/cohere/ocr/transformation.py new file mode 100644 index 00000000000..87454980aa4 --- /dev/null +++ b/litellm/llms/cohere/ocr/transformation.py @@ -0,0 +1,292 @@ +"""Cohere Parse (`POST /v2/parse`) exposed through LiteLLM's OCR interface.""" + +from collections.abc import Mapping, Sequence +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, Literal + +import httpx +from pydantic import BaseModel, ConfigDict, TypeAdapter +from typing_extensions import ReadOnly, TypedDict + +from litellm.exceptions import BadRequestError, UnsupportedParamsError +from litellm.llms.base_llm.ocr.transformation import ( + OCR_REQUEST_FORMAT_PARAM, + BaseOCRConfig, + DocumentType, + OCRPage, + OCRPageImage, + OCRRequestData, + OCRRequestFormat, + OCRResponse, + OCRUsageInfo, + parse_ocr_request_format, +) +from litellm.llms.cohere.common_utils import CohereError +from litellm.secret_managers.main import get_secret_str + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + +COHERE_API_KEY_ENV_VAR: Final = "COHERE_API_KEY" +COHERE_PARSE_API_BASE: Final = "https://api.cohere.com" +COHERE_PARSE_PATH: Final = "/v2/parse" +COHERE_PARSE_OUTPUT_FORMAT_PARAM: Final = "output_format" +COHERE_PARSE_OUTPUT_FORMATS: Final = ("markdown", "blocks") +COHERE_PARSE_DEFAULT_OUTPUT_FORMAT: Final = "markdown" +COHERE_PARSE_SUPPORTED_PARAMS: Final = (COHERE_PARSE_OUTPUT_FORMAT_PARAM, OCR_REQUEST_FORMAT_PARAM) +COHERE_PARSE_IMAGE_ONLY_MESSAGE: Final = ( + "Cohere Parse only accepts `image_url` documents (an image URL or a base64 image data URI); " + "`document_url` and PDF inputs are not supported." +) + +_NATIVE_RESPONSE_ADAPTER: Final = TypeAdapter(dict[str, object]) +_BOUNDING_BOX_ADAPTER: Final = TypeAdapter(Mapping[str, object]) + + +class _CohereParseDocument(TypedDict): + type: ReadOnly[Literal["image_url"]] + image_url: ReadOnly[str] + + +class _CohereParseRequestBody(TypedDict): + model: ReadOnly[str] + document: ReadOnly[_CohereParseDocument] + output_format: ReadOnly[str] + + +class _MarkdownPage(TypedDict): + index: ReadOnly[int] + markdown: ReadOnly[str] + images: ReadOnly[Sequence[OCRPageImage] | None] + + +class _BlocksPage(_MarkdownPage): + blocks: ReadOnly[Sequence[Mapping[str, object]]] + + +class _CohereParseMarkdown(BaseModel): + model_config = ConfigDict(frozen=True, extra="allow") + + content: str = "" + images: Sequence[Mapping[str, object]] | None = None + + +class _CohereParsePage(BaseModel): + model_config = ConfigDict(frozen=True, extra="allow") + + index: int | None = None + markdown: _CohereParseMarkdown | None = None + blocks: Sequence[Mapping[str, object]] | None = None + + +class _CohereParseBilledUnits(BaseModel): + model_config = ConfigDict(frozen=True, extra="allow") + + pages: int | None = None + + +class _CohereParseMeta(BaseModel): + model_config = ConfigDict(frozen=True, extra="allow") + + billed_units: _CohereParseBilledUnits | None = None + + +class _CohereParseResponse(BaseModel): + model_config = ConfigDict(frozen=True, extra="allow") + + pages: Sequence[_CohereParsePage] = () + meta: _CohereParseMeta | None = None + + +def _requested_format(optional_params: Mapping[str, object] | None) -> OCRRequestFormat: + if optional_params is None: + return "litellm" + return "native" if optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native" else "litellm" + + +def _page_image(image: Mapping[str, object]) -> OCRPageImage: + bounding_box: Final = image.get("bounding_box") + if not isinstance(bounding_box, Mapping): + return OCRPageImage.model_validate(image) + bbox: Final = _BOUNDING_BOX_ADAPTER.validate_python(bounding_box) + return OCRPageImage.model_validate(MappingProxyType({**image, "bbox": bbox})) + + +def _normalize_page(page: _CohereParsePage, position: int) -> OCRPage: + markdown: Final = page.markdown + images: Final = tuple(_page_image(image) for image in markdown.images) if markdown and markdown.images else None + normalized: Final[_MarkdownPage] = { + "index": page.index if page.index is not None else position, + "markdown": markdown.content if markdown else "", + "images": images, + } + if page.blocks is None: + return OCRPage.model_validate(normalized) + with_blocks: Final[_BlocksPage] = {**normalized, "blocks": page.blocks} + return OCRPage.model_validate(with_blocks) + + +def _billed_pages(parsed: _CohereParseResponse) -> int | None: + if parsed.meta is None or parsed.meta.billed_units is None: + return None + return parsed.meta.billed_units.pages + + +class CohereParseConfig(BaseOCRConfig): + """Cohere Parse, an image-only document understanding endpoint returning markdown or blocks.""" + + def get_supported_ocr_params(self, model: str) -> list[str]: # mutable-ok: BaseOCRConfig signature + return list(COHERE_PARSE_SUPPORTED_PARAMS) # mutable-ok: BaseOCRConfig signature + + def get_api_key_env_var(self) -> str | None: + return COHERE_API_KEY_ENV_VAR + + def supports_rust_bridge(self) -> bool: + return False + + def _llm_provider(self) -> str: + return "cohere" + + def map_ocr_params( + self, + non_default_params: Mapping[str, object], + optional_params: Mapping[str, object], + model: str, + ) -> dict[str, object]: # mutable-ok: BaseOCRConfig signature + output_format: Final = non_default_params.get(COHERE_PARSE_OUTPUT_FORMAT_PARAM) + if output_format is not None and output_format not in COHERE_PARSE_OUTPUT_FORMATS: + raise UnsupportedParamsError( + message=( + f"Invalid `{COHERE_PARSE_OUTPUT_FORMAT_PARAM}`: {output_format!r}. " + f"Expected one of {', '.join(COHERE_PARSE_OUTPUT_FORMATS)}." + ), + model=model, + llm_provider=self._llm_provider(), + ) + requested_format: Final = non_default_params.get(OCR_REQUEST_FORMAT_PARAM) + request_format: Final = parse_ocr_request_format(requested_format) if requested_format is not None else None + overrides: Final = tuple( + (key, value) + for key, value in ( + (COHERE_PARSE_OUTPUT_FORMAT_PARAM, output_format), + (OCR_REQUEST_FORMAT_PARAM, request_format), + ) + if value is not None + ) + return {**optional_params, **dict(overrides)} # mutable-ok: BaseOCRConfig signature + + def validate_environment( + self, + headers: Mapping[str, str], + model: str, + api_key: str | None = None, + api_base: str | None = None, + litellm_params: Mapping[str, object] | None = None, + **kwargs: object, # kwargs-ok: BaseOCRConfig.validate_environment signature + ) -> dict[str, str]: # mutable-ok: BaseOCRConfig signature + resolved_key: Final = api_key or get_secret_str(COHERE_API_KEY_ENV_VAR) + if resolved_key is None: + raise ValueError( + f"Missing {COHERE_API_KEY_ENV_VAR} - set it in the environment or pass api_key to " + "litellm.ocr()/litellm.aocr()" + ) + return { # mutable-ok: BaseOCRConfig signature + "Authorization": f"Bearer {resolved_key}", + "Content-Type": "application/json", + **headers, + } + + def get_complete_url( + self, + api_base: str | None, + model: str, + optional_params: Mapping[str, object], + litellm_params: Mapping[str, object] | None = None, + **kwargs: object, # kwargs-ok: BaseOCRConfig.get_complete_url signature + ) -> str: + url: Final = httpx.URL(api_base or COHERE_PARSE_API_BASE) + path: Final = url.path.rstrip("/") + if path.endswith(COHERE_PARSE_PATH): + return str(url.copy_with(path=path)) + if path.endswith("/v2"): + return str(url.copy_with(path=f"{path}/parse")) + return str(url.copy_with(path=f"{path}{COHERE_PARSE_PATH}")) + + def _image_url(self, document: DocumentType, model: str) -> str: + image_url: Final = document.get("image_url", "") + if document.get("type") != "image_url" or not image_url or image_url.startswith("data:application/pdf"): + raise BadRequestError( + message=COHERE_PARSE_IMAGE_ONLY_MESSAGE, + model=model, + llm_provider=self._llm_provider(), + ) + return image_url + + def _resolve_image_url_sync(self, image_url: str) -> str: + return image_url + + async def _resolve_image_url_async(self, image_url: str) -> str: + return image_url + + def _build_request(self, model: str, image_url: str, optional_params: Mapping[str, object]) -> OCRRequestData: + body: Final[_CohereParseRequestBody] = { + "model": model, + "document": {"type": "image_url", "image_url": image_url}, + "output_format": str( + optional_params.get(COHERE_PARSE_OUTPUT_FORMAT_PARAM, COHERE_PARSE_DEFAULT_OUTPUT_FORMAT) + ), + } + return OCRRequestData(data=dict(body), files=None) # mutable-ok: OCRRequestData.data is a dict + + def transform_ocr_request( + self, + model: str, + document: DocumentType, + optional_params: Mapping[str, object], + headers: Mapping[str, str], + **kwargs: object, # kwargs-ok: BaseOCRConfig.transform_ocr_request signature + ) -> OCRRequestData: + image_url: Final = self._resolve_image_url_sync(self._image_url(document, model)) + return self._build_request(model=model, image_url=image_url, optional_params=optional_params) + + async def async_transform_ocr_request( + self, + model: str, + document: DocumentType, + optional_params: Mapping[str, object], + headers: Mapping[str, str], + **kwargs: object, # kwargs-ok: BaseOCRConfig.async_transform_ocr_request signature + ) -> OCRRequestData: + image_url: Final = await self._resolve_image_url_async(self._image_url(document, model)) + return self._build_request(model=model, image_url=image_url, optional_params=optional_params) + + def transform_ocr_response( + self, + model: str, + raw_response: httpx.Response, + logging_obj: "LiteLLMLoggingObj", + optional_params: Mapping[str, object] | None = None, + **kwargs: object, # kwargs-ok: BaseOCRConfig.transform_ocr_response signature + ) -> OCRResponse: + native: Final = _NATIVE_RESPONSE_ADAPTER.validate_python(raw_response.json()) + parsed: Final = _CohereParseResponse.model_validate(native) + pages: Final = [ # mutable-ok: OCRResponse.pages is a list + _normalize_page(page, position) for position, page in enumerate(parsed.pages) + ] + billed_pages: Final = _billed_pages(parsed) + response: Final = OCRResponse( + pages=pages, + model=model, + usage_info=OCRUsageInfo(pages_processed=billed_pages if billed_pages is not None else len(pages)), + ) + if _requested_format(optional_params) == "native": + response.set_provider_native_response(native) + return response + + def get_error_class( + self, + error_message: str, + status_code: int, + headers: Mapping[str, str], + ) -> Exception: + return CohereError(status_code=status_code, message=error_message) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 2459ed940e0..4273ec54472 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -10243,6 +10243,16 @@ ], "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/mistral/" }, + "azure_ai/Cohere-parse-v5": { + "deprecation_date": "2026-12-15", + "litellm_provider": "azure_ai", + "mode": "ocr", + "ocr_cost_per_page": 0.0015, + "source": "https://cohere.com/blog/parse", + "supported_endpoints": [ + "/v1/ocr" + ] + }, "azure_ai/doc-intelligence/prebuilt-read": { "litellm_provider": "azure_ai", "ocr_cost_per_page": 0.0015, @@ -14116,6 +14126,15 @@ "output_vector_size": 1536, "supports_embedding_image_input": true }, + "cohere/parse-v5.0": { + "litellm_provider": "cohere", + "mode": "ocr", + "ocr_cost_per_page": 0.0015, + "source": "https://cohere.com/blog/parse", + "supported_endpoints": [ + "/v1/ocr" + ] + }, "cohere.rerank-v3-5:0": { "input_cost_per_query": 0.002, "input_cost_per_token": 0.0, diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index b260ec6e06f..6c68971f8d5 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -191,6 +191,8 @@ def _prepare_ocr_request( def _rust_ocr_supported(prepared_request: _PreparedOCRRequest) -> bool: if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native": return False + if not prepared_request.provider_config.supports_rust_bridge(): + return False return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS diff --git a/litellm/utils.py b/litellm/utils.py index 9d20d32d147..52c1859b525 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -9294,6 +9294,11 @@ class ProviderConfigManager: return get_vertex_ai_ocr_config(model=model) + if provider == litellm.LlmProviders.COHERE: + from litellm.llms.cohere.ocr.transformation import CohereParseConfig + + return CohereParseConfig() + if provider == litellm.LlmProviders.REDUCTO: from litellm.llms.reducto.ocr.transformation import ( ReductoParseLegacyConfig, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 2459ed940e0..4273ec54472 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -10243,6 +10243,16 @@ ], "source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/mistral/" }, + "azure_ai/Cohere-parse-v5": { + "deprecation_date": "2026-12-15", + "litellm_provider": "azure_ai", + "mode": "ocr", + "ocr_cost_per_page": 0.0015, + "source": "https://cohere.com/blog/parse", + "supported_endpoints": [ + "/v1/ocr" + ] + }, "azure_ai/doc-intelligence/prebuilt-read": { "litellm_provider": "azure_ai", "ocr_cost_per_page": 0.0015, @@ -14116,6 +14126,15 @@ "output_vector_size": 1536, "supports_embedding_image_input": true }, + "cohere/parse-v5.0": { + "litellm_provider": "cohere", + "mode": "ocr", + "ocr_cost_per_page": 0.0015, + "source": "https://cohere.com/blog/parse", + "supported_endpoints": [ + "/v1/ocr" + ] + }, "cohere.rerank-v3-5:0": { "input_cost_per_query": 0.002, "input_cost_per_token": 0.0, diff --git a/tests/test_litellm/llms/azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py b/tests/test_litellm/llms/azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py new file mode 100644 index 00000000000..01e1f59184c --- /dev/null +++ b/tests/test_litellm/llms/azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py @@ -0,0 +1,166 @@ +import base64 +import json + +import pytest + +import litellm +from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig +from litellm.llms.azure_ai.ocr.common_utils import get_azure_ai_ocr_config +from litellm.llms.azure_ai.ocr.document_intelligence.transformation import AzureDocumentIntelligenceOCRConfig +from litellm.llms.azure_ai.ocr.transformation import AzureAIOCRConfig + +MODEL = "azure_ai/Cohere-parse-v5" +API_BASE = "https://resource.services.ai.azure.com" +PARSE_URL = f"{API_BASE}/providers/cohere/v2/parse" +IMAGE_URL = "https://example.com/receipt.png" +PNG_BYTES = base64.b64decode( + "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg==" +) +PNG_DATA_URI = f"data:image/png;base64,{base64.b64encode(PNG_BYTES).decode()}" + + +def _parse_response() -> dict: + return { + "id": "882bf973-9dfa-4d02-9d30-709247008efd", + "pages": [{"index": 0, "type": "markdown", "markdown": {"content": "# Receipt\n\nTotal Due: $4.00"}}], + "meta": {"api_version": {"version": "2"}, "billed_units": {"pages": 1}}, + } + + +@pytest.fixture() +def disable_aiohttp_transport(monkeypatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + yield + litellm.in_memory_llm_clients_cache.flush_cache() + + +@pytest.mark.parametrize( + "model, expected_config", + [ + ("Cohere-parse-v5", AzureAICohereParseConfig), + ("cohere-parse-v5", AzureAICohereParseConfig), + ("parse-v5", AzureAICohereParseConfig), + ("mistral-ocr-4-0", AzureAIOCRConfig), + ("mistral-document-ai-2512", AzureAIOCRConfig), + ("doc-intelligence/prebuilt-read", AzureDocumentIntelligenceOCRConfig), + ], +) +def test_azure_ai_ocr_routing(model: str, expected_config: type) -> None: + assert type(get_azure_ai_ocr_config(model)) is expected_config + + +@pytest.mark.parametrize( + "api_base, expected_url", + [ + (API_BASE, PARSE_URL), + (f"{API_BASE}/", PARSE_URL), + (f"{API_BASE}/models", PARSE_URL), + (f"{API_BASE}/providers/cohere/v2", PARSE_URL), + (f"{API_BASE}/providers/cohere/v2/parse", PARSE_URL), + ], +) +def test_get_complete_url_targets_the_cohere_provider_route(api_base: str, expected_url: str) -> None: + url = AzureAICohereParseConfig().get_complete_url(api_base=api_base, model="Cohere-parse-v5", optional_params={}) + + assert url == expected_url + + +def test_get_complete_url_falls_back_to_env_api_base(monkeypatch) -> None: + monkeypatch.setenv("AZURE_AI_API_BASE", API_BASE) + + url = AzureAICohereParseConfig().get_complete_url(api_base=None, model="Cohere-parse-v5", optional_params={}) + + assert url == PARSE_URL + + +def test_get_complete_url_requires_api_base(monkeypatch) -> None: + monkeypatch.delenv("AZURE_AI_API_BASE", raising=False) + + with pytest.raises(ValueError, match="AZURE_AI_API_BASE"): + AzureAICohereParseConfig().get_complete_url(api_base=None, model="Cohere-parse-v5", optional_params={}) + + +def test_get_complete_url_rejects_relative_api_base() -> None: + with pytest.raises(ValueError, match="absolute URL"): + AzureAICohereParseConfig().get_complete_url( + api_base="resource.services.ai.azure.com", model="Cohere-parse-v5", optional_params={} + ) + + +def test_validate_environment_requires_api_base(monkeypatch) -> None: + monkeypatch.delenv("AZURE_AI_API_BASE", raising=False) + + with pytest.raises(ValueError, match="AZURE_AI_API_BASE"): + AzureAICohereParseConfig().validate_environment(headers={}, model="Cohere-parse-v5", api_key="key") + + +@pytest.mark.asyncio +async def test_aocr_inlines_remote_image_and_posts_to_foundry(disable_aiohttp_transport, respx_mock): + respx_mock.get(IMAGE_URL).respond(content=PNG_BYTES, headers={"Content-Type": "image/png"}) + route = respx_mock.post(PARSE_URL).respond(json=_parse_response()) + + response = await litellm.aocr( + model=MODEL, + document={"type": "image_url", "image_url": IMAGE_URL}, + api_base=API_BASE, + api_key="azure-key", + ) + + request = route.calls.last.request + assert request.headers["Authorization"] == "Bearer azure-key" + assert json.loads(request.content) == { + "model": "Cohere-parse-v5", + "document": {"type": "image_url", "image_url": PNG_DATA_URI}, + "output_format": "markdown", + } + assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00" + assert response.usage_info.pages_processed == 1 + + +@pytest.mark.asyncio +async def test_aocr_passes_data_uri_through_without_fetching(disable_aiohttp_transport, respx_mock): + route = respx_mock.post(PARSE_URL).respond(json=_parse_response()) + + await litellm.aocr( + model=MODEL, + document={"type": "image_url", "image_url": PNG_DATA_URI}, + api_base=API_BASE, + api_key="azure-key", + output_format="blocks", + ) + + body = json.loads(route.calls.last.request.content) + assert body["document"]["image_url"] == PNG_DATA_URI + assert body["output_format"] == "blocks" + + +def test_ocr_sync_inlines_remote_image(respx_mock): + respx_mock.get(IMAGE_URL).respond(content=PNG_BYTES, headers={"Content-Type": "image/png"}) + route = respx_mock.post(PARSE_URL).respond(json=_parse_response()) + + response = litellm.ocr( + model=MODEL, + document={"type": "image_url", "image_url": IMAGE_URL}, + api_base=API_BASE, + api_key="azure-key", + ) + + assert json.loads(route.calls.last.request.content)["document"]["image_url"] == PNG_DATA_URI + assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00" + + +@pytest.mark.asyncio +async def test_aocr_rejects_pdf_before_calling_foundry(disable_aiohttp_transport, respx_mock): + route = respx_mock.post(PARSE_URL).respond(json=_parse_response()) + + with pytest.raises(litellm.BadRequestError, match="only accepts `image_url` documents") as exc_info: + await litellm.aocr( + model=MODEL, + document={"type": "document_url", "document_url": "https://example.com/doc.pdf"}, + api_base=API_BASE, + api_key="azure-key", + ) + + assert exc_info.value.llm_provider == "azure_ai" + assert not route.called diff --git a/tests/test_litellm/llms/cohere/ocr/test_cohere_parse_cost.py b/tests/test_litellm/llms/cohere/ocr/test_cohere_parse_cost.py new file mode 100644 index 00000000000..dfa3c7a056e --- /dev/null +++ b/tests/test_litellm/llms/cohere/ocr/test_cohere_parse_cost.py @@ -0,0 +1,58 @@ +import json +from pathlib import Path + +import pytest + +import litellm +from litellm.cost_calculator import completion_cost +from litellm.llms.base_llm.ocr.transformation import OCRPage, OCRResponse, OCRUsageInfo + +COST_PER_PAGE = 0.0015 +REPO_ROOT = Path(__file__).parents[5] +COST_MAPS = [ + REPO_ROOT / "model_prices_and_context_window.json", + REPO_ROOT / "litellm" / "model_prices_and_context_window_backup.json", +] +MODELS = [("cohere/parse-v5.0", "cohere"), ("azure_ai/Cohere-parse-v5", "azure_ai")] + + +def _ocr_response(model: str, pages_processed: int) -> OCRResponse: + return OCRResponse( + pages=[OCRPage(index=i, markdown=f"page {i}") for i in range(pages_processed)], + model=model, + usage_info=OCRUsageInfo(pages_processed=pages_processed), + ) + + +@pytest.mark.parametrize("cost_map_path", COST_MAPS, ids=lambda path: path.name) +@pytest.mark.parametrize("model, provider", MODELS) +def test_pricing_entry(cost_map_path: Path, model: str, provider: str) -> None: + with open(cost_map_path) as f: + info = json.load(f).get(model) + + assert info is not None, f"{model} missing from {cost_map_path.name}" + assert info["litellm_provider"] == provider + assert info["mode"] == "ocr" + assert info["supported_endpoints"] == ["/v1/ocr"] + assert info["ocr_cost_per_page"] == COST_PER_PAGE + + +@pytest.mark.parametrize("model, provider", MODELS) +def test_model_info_resolves_ocr_mode_and_price(local_model_cost_map, model: str, provider: str) -> None: + info = litellm.get_model_info(model=model, custom_llm_provider=provider) + + assert info["mode"] == "ocr" + assert info["ocr_cost_per_page"] == COST_PER_PAGE + + +@pytest.mark.parametrize("model, provider", MODELS) +@pytest.mark.parametrize("pages_processed", [1, 3]) +def test_cost_scales_with_billed_pages(local_model_cost_map, model: str, provider: str, pages_processed: int) -> None: + cost = completion_cost( + completion_response=_ocr_response(model.split("/", 1)[1], pages_processed), + model=model, + custom_llm_provider=provider, + call_type="ocr", + ) + + assert cost == pytest.approx(COST_PER_PAGE * pages_processed) diff --git a/tests/test_litellm/llms/cohere/ocr/test_cohere_parse_transformation.py b/tests/test_litellm/llms/cohere/ocr/test_cohere_parse_transformation.py new file mode 100644 index 00000000000..64f2bb383c2 --- /dev/null +++ b/tests/test_litellm/llms/cohere/ocr/test_cohere_parse_transformation.py @@ -0,0 +1,217 @@ +import json + +import pytest + +import litellm + +PARSE_URL = "https://api.cohere.com/v2/parse" +MODEL = "cohere/parse-v5.0" +IMAGE_DOCUMENT = {"type": "image_url", "image_url": "https://example.com/receipt.png"} +BOUNDING_BOX = {"top_left_x": 0, "top_left_y": 0, "bottom_right_x": 32, "bottom_right_y": 32} + + +def _markdown_response(billed_pages: int | None = 2) -> dict: + return { + "id": "272900cc-04c0-4da2-a505-2cea58d231bf", + "pages": [ + { + "index": 0, + "type": "markdown", + "markdown": { + "content": "# Receipt\n\nTotal Due: $4.00", + "images": [ + { + "id": "img-0", + "description": "A parking receipt", + "category": "other", + "bounding_box": BOUNDING_BOX, + "bounding_box_normalized": { + "top_left_x": 0, + "top_left_y": 0, + "bottom_right_x": 1, + "bottom_right_y": 1, + }, + } + ], + }, + }, + {"index": 1, "type": "markdown", "markdown": {"content": "Page two"}}, + ], + **( + {"meta": {"api_version": {"version": "2"}, "billed_units": {"pages": billed_pages}}} if billed_pages else {} + ), + } + + +def _blocks_response() -> dict: + return { + "id": "94474f83-e30d-4763-b4bc-52af6e12c4f7", + "pages": [ + { + "index": 0, + "type": "blocks", + "blocks": [{"type": "text", "text": "Total Due: $4.00"}], + } + ], + "meta": {"api_version": {"version": "2"}, "billed_units": {"pages": 1}}, + } + + +@pytest.fixture() +def disable_aiohttp_transport(monkeypatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + yield + litellm.in_memory_llm_clients_cache.flush_cache() + + +@pytest.mark.asyncio +async def test_aocr_sends_markdown_parse_request_and_normalizes_pages(disable_aiohttp_transport, respx_mock): + route = respx_mock.post(PARSE_URL).respond(json=_markdown_response()) + + response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key") + + request = route.calls.last.request + assert request.headers["Authorization"] == "Bearer test-key" + assert json.loads(request.content) == { + "model": "parse-v5.0", + "document": IMAGE_DOCUMENT, + "output_format": "markdown", + } + assert response.object == "ocr" + assert [page.index for page in response.pages] == [0, 1] + assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00" + assert response.pages[1].markdown == "Page two" + assert response.pages[1].images is None + image = response.pages[0].images[0] + assert image.bbox == BOUNDING_BOX + assert image.model_extra["description"] == "A parking receipt" + assert image.model_extra["bounding_box_normalized"]["bottom_right_x"] == 1 + assert response.usage_info.pages_processed == 2 + assert response.get_provider_native_response() is None + + +@pytest.mark.asyncio +async def test_aocr_usage_prefers_billed_units_over_page_count(disable_aiohttp_transport, respx_mock): + respx_mock.post(PARSE_URL).respond(json=_markdown_response(billed_pages=3)) + + response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key") + + assert response.usage_info.pages_processed == 3 + + +@pytest.mark.asyncio +async def test_aocr_usage_falls_back_to_page_count_without_meta(disable_aiohttp_transport, respx_mock): + respx_mock.post(PARSE_URL).respond(json=_markdown_response(billed_pages=None)) + + response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key") + + assert response.usage_info.pages_processed == 2 + + +@pytest.mark.asyncio +async def test_aocr_blocks_output_format_forwards_param_and_keeps_blocks(disable_aiohttp_transport, respx_mock): + route = respx_mock.post(PARSE_URL).respond(json=_blocks_response()) + + response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", output_format="blocks") + + assert json.loads(route.calls.last.request.content)["output_format"] == "blocks" + assert response.pages[0].markdown == "" + assert response.pages[0].model_extra["blocks"] == [{"type": "text", "text": "Total Due: $4.00"}] + assert response.usage_info.pages_processed == 1 + + +@pytest.mark.asyncio +async def test_aocr_native_format_carries_provider_payload(disable_aiohttp_transport, respx_mock): + payload = _markdown_response() + route = respx_mock.post(PARSE_URL).respond(json=payload) + + response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", req_format="native") + + assert "req_format" not in json.loads(route.calls.last.request.content) + assert response.get_provider_native_response() == payload + assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00" + + +@pytest.mark.asyncio +async def test_aocr_rejects_unknown_output_format_before_calling_provider(disable_aiohttp_transport, respx_mock): + route = respx_mock.post(PARSE_URL).respond(json=_markdown_response()) + + with pytest.raises(litellm.BadRequestError, match="Invalid `output_format`: 'html'") as exc_info: + await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", output_format="html") + + assert exc_info.value.status_code == 400 + assert not route.called + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "document", + [ + {"type": "document_url", "document_url": "https://example.com/doc.pdf"}, + {"type": "image_url", "image_url": "data:application/pdf;base64,JVBERi0="}, + {"type": "image_url", "image_url": ""}, + ], +) +async def test_aocr_rejects_non_image_documents_before_calling_provider( + disable_aiohttp_transport, respx_mock, document +): + route = respx_mock.post(PARSE_URL).respond(json=_markdown_response()) + + with pytest.raises(litellm.BadRequestError, match="only accepts `image_url` documents") as exc_info: + await litellm.aocr(model=MODEL, document=document, api_key="test-key") + + assert exc_info.value.status_code == 400 + assert not route.called + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "api_base, expected_url", + [ + ("https://gateway.example.com", "https://gateway.example.com/v2/parse"), + ("https://gateway.example.com/cohere/", "https://gateway.example.com/cohere/v2/parse"), + ("https://gateway.example.com/v2", "https://gateway.example.com/v2/parse"), + ("https://gateway.example.com/v2/parse", "https://gateway.example.com/v2/parse"), + ], +) +async def test_aocr_posts_to_api_base_variants(disable_aiohttp_transport, respx_mock, api_base, expected_url): + route = respx_mock.post(expected_url).respond(json=_markdown_response()) + + await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", api_base=api_base) + + assert route.called + + +@pytest.mark.asyncio +async def test_aocr_surfaces_provider_error_with_its_status_and_message(disable_aiohttp_transport, respx_mock): + respx_mock.post(PARSE_URL).respond( + status_code=400, json={"id": "83b0d95e", "message": "output_format must be `blocks` or `markdown`"} + ) + + with pytest.raises(litellm.BadRequestError, match="output_format must be") as exc_info: + await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key") + + assert exc_info.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_aocr_reads_api_key_from_environment(disable_aiohttp_transport, respx_mock, monkeypatch): + monkeypatch.setenv("COHERE_API_KEY", "env-key") + route = respx_mock.post(PARSE_URL).respond(json=_markdown_response()) + + await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT) + + assert route.calls.last.request.headers["Authorization"] == "Bearer env-key" + + +@pytest.mark.asyncio +async def test_aocr_without_api_key_names_the_env_var(disable_aiohttp_transport, respx_mock, monkeypatch): + monkeypatch.delenv("COHERE_API_KEY", raising=False) + monkeypatch.setattr(litellm, "cohere_key", None) + route = respx_mock.post(PARSE_URL).respond(json=_markdown_response()) + + with pytest.raises(Exception, match="Missing COHERE_API_KEY"): + await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT) + + assert not route.called diff --git a/tests/test_litellm/ocr/test_ocr_native_format.py b/tests/test_litellm/ocr/test_ocr_native_format.py index 463213a2071..249fbda713e 100644 --- a/tests/test_litellm/ocr/test_ocr_native_format.py +++ b/tests/test_litellm/ocr/test_ocr_native_format.py @@ -4,11 +4,14 @@ providers that don't support a native response must reject it, and the Rust bridge (which only returns the normalized shape) must not serve native requests. """ +import dataclasses from unittest.mock import MagicMock import pytest import litellm +from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig +from litellm.llms.cohere.ocr.transformation import CohereParseConfig from litellm.ocr.main import _PreparedOCRRequest, _rust_ocr_supported DOCUMENT = {"type": "document_url", "document_url": "https://example.com/doc.pdf"} @@ -39,6 +42,13 @@ def test_rust_ocr_skipped_for_native_format(): assert _rust_ocr_supported(_prepared({"req_format": "native"})) is False +@pytest.mark.parametrize("provider_config", [CohereParseConfig(), AzureAICohereParseConfig()]) +def test_rust_ocr_skipped_for_configs_without_bridge_support(provider_config): + prepared = dataclasses.replace(_prepared({}), provider_config=provider_config) + + assert _rust_ocr_supported(prepared) is False + + @pytest.mark.asyncio async def test_native_format_rejected_for_provider_without_support_as_bad_request(): with pytest.raises(litellm.BadRequestError, match="not supported for provider") as exc_info: From d3a179f98871abdacd3b50d041aa61205cfec564 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 4 Sep 2026 22:07:26 -0700 Subject: [PATCH 05/35] fix(azure_ai): route only cohere parse deployment names to Cohere Parse --- litellm/llms/azure_ai/ocr/common_utils.py | 3 ++- .../azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py | 4 +++- 2 files changed, 5 insertions(+), 2 deletions(-) diff --git a/litellm/llms/azure_ai/ocr/common_utils.py b/litellm/llms/azure_ai/ocr/common_utils.py index a4cd0c7a30b..2ca2ad9ec2f 100644 --- a/litellm/llms/azure_ai/ocr/common_utils.py +++ b/litellm/llms/azure_ai/ocr/common_utils.py @@ -25,7 +25,8 @@ def is_azure_document_intelligence_model(model: str) -> bool: def is_azure_cohere_parse_model(model: str) -> bool: - return "parse" in model.lower() + lowered: Final = model.lower() + return "cohere" in lowered and "parse" in lowered def get_azure_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]: diff --git a/tests/test_litellm/llms/azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py b/tests/test_litellm/llms/azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py index 01e1f59184c..8c4dd25aa77 100644 --- a/tests/test_litellm/llms/azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py +++ b/tests/test_litellm/llms/azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py @@ -40,7 +40,9 @@ def disable_aiohttp_transport(monkeypatch): [ ("Cohere-parse-v5", AzureAICohereParseConfig), ("cohere-parse-v5", AzureAICohereParseConfig), - ("parse-v5", AzureAICohereParseConfig), + ("cohere/parse-v5", AzureAICohereParseConfig), + ("invoice-parser", AzureAIOCRConfig), + ("parse-v5", AzureAIOCRConfig), ("mistral-ocr-4-0", AzureAIOCRConfig), ("mistral-document-ai-2512", AzureAIOCRConfig), ("doc-intelligence/prebuilt-read", AzureDocumentIntelligenceOCRConfig), From 004a8201167eb0bd85d52e86473adb9b9b8d1ee7 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 4 Sep 2026 22:52:12 -0700 Subject: [PATCH 06/35] fix(ocr): send each provider a health-check document it accepts Health checks probed every OCR deployment with a PDF, which Cohere Parse rejects, so /health, background health checks, and the UI Test Connection button marked Cohere Parse deployments unhealthy. BaseOCRConfig gains a get_health_check_document hook (PDF by default) that CohereParseConfig overrides with a 1x1 PNG data URI. cohere also gains ocr in the provider endpoint matrix --- .../health_check_helpers.py | 18 +++++++----- litellm/llms/base_llm/ocr/transformation.py | 8 +++++ litellm/llms/cohere/ocr/transformation.py | 9 ++++++ .../provider_endpoints_support_backup.json | 1 + provider_endpoints_support.json | 1 + .../test_health_check_helpers.py | 29 +++++++++++++++++++ ...st_azure_ai_cohere_parse_transformation.py | 16 ++++++++++ .../ocr/test_cohere_parse_transformation.py | 12 ++++++++ 8 files changed, 87 insertions(+), 7 deletions(-) diff --git a/litellm/litellm_core_utils/health_check_helpers.py b/litellm/litellm_core_utils/health_check_helpers.py index c745bbea5c4..9f8878d36f0 100644 --- a/litellm/litellm_core_utils/health_check_helpers.py +++ b/litellm/litellm_core_utils/health_check_helpers.py @@ -6,14 +6,13 @@ import base64 from collections.abc import Awaitable, Callable from typing import TYPE_CHECKING, Final, Literal -from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS +from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, DocumentType +from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS, LlmProviders if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging from litellm.types.utils import ImageResponse -# Minimal PDF for health checks - base64 encoded 1-page PDF with just "test" -TEST_PDF_URL = "data:application/pdf;base64,JVBERi0xLjQKJeLjz9MKMyAwIG9iago8PC9UeXBlIC9QYWdlCi9QYXJlbnQgMSAwIFIKL01lZGlhQm94IFswIDAgNjEyIDc5Ml0KL0NvbnRlbnRzIDQgMCBSCi9SZXNvdXJjZXMgPDwvRm9udCA8PC9GMSAyIDAgUj4+Pj4+PgplbmRvYmoKNCAwIG9iago8PC9MZW5ndGggNDQ+PgpzdHJlYW0KQlQKL0YxIDI0IFRmCjEwMCA3MDAgVGQKKHRlc3QpIFRqCkVUCmVuZHN0cmVhbQplbmRvYmoKMiAwIG9iago8PC9UeXBlIC9Gb250Ci9TdWJ0eXBlIC9UeXBlMQovQmFzZUZvbnQgL0hlbHZldGljYT4+CmVuZG9iagoxIDAgb2JqCjw8L1R5cGUgL1BhZ2VzCi9LaWRzIFszIDAgUl0KL0NvdW50IDE+PgplbmRvYmoKNSAwIG9iago8PC9UeXBlIC9DYXRhbG9nCi9QYWdlcyAxIDAgUj4+CmVuZG9iagp0cmFpbGVyCjw8L1NpemUgNgovUm9vdCA1IDAgUj4+CnN0YXJ0eHJlZgozMjQKJSVFT0Y=" # Minimal image for health checks - base64 encoded 512x512 blue circle on a white background PNG TEST_IMAGE_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAAgAAAAIACAIAAAB7GkOtAAAJk0lEQVR42u3VQREAIRADwVWCOmTjBVzwSLorCri6nbkAVBpPACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAAAgAAAIAZdY+HgEBgIRr/meeGgGA8EMvDAgAuPh6gACAi68HCAA4+mKAAICjLwYIALj7SoAAgLuvBAgA7r4pAQKAu29KgADg7psSIAA4/SYDCADuvikBAoDTbzKAAOD0mwwgADj9JgMIAE6/yQACgNNvMoAA4PSbDCAAOP0mAwgATr/JAAKA668BIAA4/TKAAOD0mwwgALj+pgEIAE6/yQACgOtvGoAA4PSbDCAAuP6mAQgATr/JAAKA628agADg9JsMIAC4/qYBCACuv2kAAoDrbxqAAOD0mwwgALj+pgEIAK6/aQACgOtvGoAA4PqbBiAArr+ZBiAATr+ZDCAArr+ZBiAArr+ZBiAArr+ZBiAArr+ZBggArr+ZBggArr+ZBggArr+ZBggArr+ZBggArr+ZBggArr+ZBggAAmAmAAKA62+mAQKA62+mAQKA62/m1xYAXH/TAAQA1980AAHA9TcNQAAEwEwAEADX30wDEADX30wDEADX30wDEADX30wDEAABMBMABMD1N9MABMD1N9MABEAAzAQAAXD9zTQAAXD9zTRAABAAMwEQAFx/Mw0QAFx/Mw0QAATATAAEANffTAMEANffTAMEAAEwEwABwPU30wABQADMBEAAXH8z0wABcP3NTAMEQADMTAAEwPU3Mw0QAAEwMwEQANffTAMQAAEwEwAEwPU30wAEQADMBAABcP3NNAABEAAzAUAAXH8zDUAABMBMABAA199MAwQAATATAAHA9TfTAAFAAMwEQABw/c00QAAQADMBEAAEwEwABMD1NzMNEAABMDMBEADX38w0QAAEwMwEQAAEwMwEQABcfzPTAAEQADMTAAEQADMTAAFw/c1MAwRAAMxMAARAAMxMAATA9TczDRAAATATAARAAMwEAAFw/c00AAEQADMBQAAEwEwAEADX30wDBAABMBMAAUAAzARAABAAMwEQAFx/Mw0QAAEwMwEQAAEwMwEQAAEwMwEQANffzDRAAATAzARAAATAzARAAATAzARAAATAzARAAFx/M9MAARAAMxMAARAAMxMAARAAMxMAARAAMxMAAXD9zUwDBEAAzEwABEAAzAQAARAAMwFAAATATAAQAAEwEwABQADMBEAAEAAzARAAXH8zDRAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAA19/MNEAANMDM9UcABMBMABAAATATAAHwBAJgJgACgACYCYAAIABmAiAA+IvMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAANMDMXH8BEAAzEwABEAAzEwABEAAzEwABEAAzEwABEAAzEwABEAAzAUAABMBMABAADTBz/REAATATAAFAAMwEQAAQADMBEAAEwEwABAANMHP9BUAAzEwABEAAzEwABEAAzEwABEAAzEwABEADzMz1FwABMDMBEAABMDMBEAABMDMBEAANMDPXXwAEwMwEQAAEwMwEQAAEwMwEQAA0wMxcfwEQADMTAAEQADMBQAA0wMz1RwAEwEwAEAABMBMABEADzFx/AUAAzARAABAAMwEQADTAzPUXAATATAAEAAEwEwABQAPMXH8BEAAzEwABEAAzEwAB0AAzc/0FQADMTAAEQAPMzPUXAAEwMwEQAAEwMwEQAA0wM9dfAATAzARAADTAzFx/ARAAMwFAADTAzPVHAATATAAQAA0wc/0RAAEwEwAEQAPMXH8EQADMBAAB0AAz118AEAAzARAANMDM9RcABMBMAAQADTBz/QUAATATAAFAA8xcfwFAA8xcfwFAAMwEQADQADPXXwAEwMwEQAA0wMxcfwHQADNz/QVAAMxMAARAA8xcfwRAA8xcfwRAAMwEAAHQADPXHwHQADPXHwEQADMBQAA0wMz1RwA0wMz1RwAEwEwAEAANMHP9BQANMHP9BQANMHP9BQANMHP9BQABMBMAAUADzFx/AUADzFx/AUADzFx/AUADzFx/AUADzFx/AUADzPVHABAAEwAEAA0w1x8BQAPM9UcA0ABz/REANMBcfwQADTDXHwHQADPXHwHQADPXHwHQADPXHwHQADPXHwHQADPXHwGQATOnHwHQADPXHwHQADPXXwDQADPXXwDQADPXXwDQADPXXwCQAXP6EQA0wFx/BAANMNcfAUADzPVHAJABc/oRADTAXH8EABkwpx8BQAPM9UcAkAFz+hEANMBcfwQAGTCnHwFAA8z1RwCQAXP6EQBkwJx+BAANMNcfAUAGzOlHAJABc/oRAGTA6QcBQAacfhAAZMDpRwBABpx+BABkwOlHAEAGnH4EAJTA3UcAQAacfgQAlMDdRwBACdx9BACUwN1HAEAJ3H0EAJTA3UcAQAwcfQQAmmLgsyIA0NIDHw4BgJYe+DQIAISHwVMjAJDQDI+AAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACACAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAIAABNHpialFcmLajuAAAAAElFTkSuQmCC" @@ -29,6 +28,14 @@ def get_image_file_for_health_check() -> bytes: return base64.b64decode(TEST_IMAGE_BASE64) +def _ocr_health_check_document(model: str, custom_llm_provider: str) -> DocumentType: + from litellm.utils import ProviderConfigManager + + provider: Final = next((known for known in LlmProviders if known.value == custom_llm_provider), None) + config: Final = ProviderConfigManager.get_provider_ocr_config(model=model, provider=provider) if provider else None + return (config or BaseOCRConfig()).get_health_check_document() + + class HealthCheckHelpers: @staticmethod async def ahealth_check_wildcard_models( @@ -247,9 +254,6 @@ class HealthCheckHelpers: ), "ocr": lambda: litellm.aocr( **_filter_model_params(model_params=model_params), - document={ - "type": "document_url", - "document_url": TEST_PDF_URL, - }, + document=_ocr_health_check_document(model=model, custom_llm_provider=custom_llm_provider), ), } diff --git a/litellm/llms/base_llm/ocr/transformation.py b/litellm/llms/base_llm/ocr/transformation.py index 08ae077cb2f..8111f9a194a 100644 --- a/litellm/llms/base_llm/ocr/transformation.py +++ b/litellm/llms/base_llm/ocr/transformation.py @@ -33,6 +33,8 @@ OCR_REQUEST_FORMAT_HEADER: Final = "x-req-format" PROVIDER_NATIVE_RESPONSE_KEY: Final = "provider_native_response" +HEALTH_CHECK_PDF_DATA_URI: Final = "data:application/pdf;base64,JVBERi0xLjQKJeLjz9MKMyAwIG9iago8PC9UeXBlIC9QYWdlCi9QYXJlbnQgMSAwIFIKL01lZGlhQm94IFswIDAgNjEyIDc5Ml0KL0NvbnRlbnRzIDQgMCBSCi9SZXNvdXJjZXMgPDwvRm9udCA8PC9GMSAyIDAgUj4+Pj4+PgplbmRvYmoKNCAwIG9iago8PC9MZW5ndGggNDQ+PgpzdHJlYW0KQlQKL0YxIDI0IFRmCjEwMCA3MDAgVGQKKHRlc3QpIFRqCkVUCmVuZHN0cmVhbQplbmRvYmoKMiAwIG9iago8PC9UeXBlIC9Gb250Ci9TdWJ0eXBlIC9UeXBlMQovQmFzZUZvbnQgL0hlbHZldGljYT4+CmVuZG9iagoxIDAgb2JqCjw8L1R5cGUgL1BhZ2VzCi9LaWRzIFszIDAgUl0KL0NvdW50IDE+PgplbmRvYmoKNSAwIG9iago8PC9UeXBlIC9DYXRhbG9nCi9QYWdlcyAxIDAgUj4+CmVuZG9iagp0cmFpbGVyCjw8L1NpemUgNgovUm9vdCA1IDAgUj4+CnN0YXJ0eHJlZgozMjQKJSVFT0Y=" + def parse_ocr_request_format(value: object) -> OCRRequestFormat: if value == "litellm": @@ -146,6 +148,12 @@ class BaseOCRConfig: """Whether the Rust OCR bridge may serve this config when it is enabled for the provider.""" return True + def get_health_check_document(self) -> DocumentType: + return { # mutable-ok: litellm.aocr rejects any document that is not a dict + "type": "document_url", + "document_url": HEALTH_CHECK_PDF_DATA_URI, + } + def map_ocr_params( self, non_default_params: dict, diff --git a/litellm/llms/cohere/ocr/transformation.py b/litellm/llms/cohere/ocr/transformation.py index 87454980aa4..dd15d5360a6 100644 --- a/litellm/llms/cohere/ocr/transformation.py +++ b/litellm/llms/cohere/ocr/transformation.py @@ -34,6 +34,9 @@ COHERE_PARSE_OUTPUT_FORMAT_PARAM: Final = "output_format" COHERE_PARSE_OUTPUT_FORMATS: Final = ("markdown", "blocks") COHERE_PARSE_DEFAULT_OUTPUT_FORMAT: Final = "markdown" COHERE_PARSE_SUPPORTED_PARAMS: Final = (COHERE_PARSE_OUTPUT_FORMAT_PARAM, OCR_REQUEST_FORMAT_PARAM) +COHERE_PARSE_HEALTH_CHECK_IMAGE_DATA_URI: Final = ( + "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGP4//8/AAX+Av4N70a4AAAAAElFTkSuQmCC" +) COHERE_PARSE_IMAGE_ONLY_MESSAGE: Final = ( "Cohere Parse only accepts `image_url` documents (an image URL or a base64 image data URI); " "`document_url` and PDF inputs are not supported." @@ -144,6 +147,12 @@ class CohereParseConfig(BaseOCRConfig): def supports_rust_bridge(self) -> bool: return False + def get_health_check_document(self) -> DocumentType: + return { # mutable-ok: litellm.aocr rejects any document that is not a dict + "type": "image_url", + "image_url": COHERE_PARSE_HEALTH_CHECK_IMAGE_DATA_URI, + } + def _llm_provider(self) -> str: return "cohere" diff --git a/litellm/provider_endpoints_support_backup.json b/litellm/provider_endpoints_support_backup.json index 9d6b1e18f59..dbeaccdda2d 100644 --- a/litellm/provider_endpoints_support_backup.json +++ b/litellm/provider_endpoints_support_backup.json @@ -559,6 +559,7 @@ "moderations": false, "batches": false, "rerank": true, + "ocr": true, "a2a": true, "interactions": true } diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 41ed8e1d975..c71f4a82a4a 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -594,6 +594,7 @@ "moderations": false, "batches": false, "rerank": true, + "ocr": true, "a2a": true, "interactions": true } diff --git a/tests/test_litellm/litellm_core_utils/test_health_check_helpers.py b/tests/test_litellm/litellm_core_utils/test_health_check_helpers.py index ee2a31beff7..89b377af3a0 100644 --- a/tests/test_litellm/litellm_core_utils/test_health_check_helpers.py +++ b/tests/test_litellm/litellm_core_utils/test_health_check_helpers.py @@ -453,3 +453,32 @@ async def test_realtime_health_check_uses_model_level_vertex_params(): "Authorization": "Bearer model-level-token", "x-goog-user-project": "model-level-project", } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "model, custom_llm_provider, expected_document_type, expected_uri_prefix", + [ + ("mistral/mistral-ocr-latest", "mistral", "document_url", "data:application/pdf;base64,"), + ("azure_ai/mistral-document-ai-2512", "azure_ai", "document_url", "data:application/pdf;base64,"), + ("cohere/parse-v5.0", "cohere", "image_url", "data:image/png;base64,"), + ("azure_ai/Cohere-parse-v5", "azure_ai", "image_url", "data:image/png;base64,"), + ], +) +async def test_ocr_health_check_sends_the_document_kind_the_provider_config_accepts( + model, custom_llm_provider, expected_document_type, expected_uri_prefix +): + handlers = HealthCheckHelpers.get_mode_handlers( + model=model, + custom_llm_provider=custom_llm_provider, + model_params={"model": model, "api_key": "sk-test"}, + ) + + with patch( # test-quality-ok: the public health-check path has no dependency injection seam + "litellm.aocr", new_callable=AsyncMock, return_value={} + ) as mock_aocr: + await handlers["ocr"]() + + document = mock_aocr.call_args.kwargs["document"] + assert document["type"] == expected_document_type + assert document[expected_document_type].startswith(expected_uri_prefix) diff --git a/tests/test_litellm/llms/azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py b/tests/test_litellm/llms/azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py index 8c4dd25aa77..3f98e9b6a2d 100644 --- a/tests/test_litellm/llms/azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py +++ b/tests/test_litellm/llms/azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py @@ -166,3 +166,19 @@ async def test_aocr_rejects_pdf_before_calling_foundry(disable_aiohttp_transport assert exc_info.value.llm_provider == "azure_ai" assert not route.called + + +@pytest.mark.asyncio +async def test_ahealth_check_ocr_sends_an_image_to_the_foundry_cohere_parse_deployment( + disable_aiohttp_transport, respx_mock +): + route = respx_mock.post(PARSE_URL).respond(json=_parse_response()) + + result = await litellm.ahealth_check( + model_params={"model": MODEL, "api_base": API_BASE, "api_key": "test-key"}, mode="ocr" + ) + + document = json.loads(route.calls.last.request.content)["document"] + assert document["type"] == "image_url" + assert document["image_url"].startswith("data:image/png;base64,") + assert "error" not in result diff --git a/tests/test_litellm/llms/cohere/ocr/test_cohere_parse_transformation.py b/tests/test_litellm/llms/cohere/ocr/test_cohere_parse_transformation.py index 64f2bb383c2..cb9af56f5e0 100644 --- a/tests/test_litellm/llms/cohere/ocr/test_cohere_parse_transformation.py +++ b/tests/test_litellm/llms/cohere/ocr/test_cohere_parse_transformation.py @@ -215,3 +215,15 @@ async def test_aocr_without_api_key_names_the_env_var(disable_aiohttp_transport, await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT) assert not route.called + + +@pytest.mark.asyncio +async def test_ahealth_check_ocr_sends_an_image_cohere_parse_accepts(disable_aiohttp_transport, respx_mock): + route = respx_mock.post(PARSE_URL).respond(json=_markdown_response()) + + result = await litellm.ahealth_check(model_params={"model": MODEL, "api_key": "test-key"}, mode="ocr") + + document = json.loads(route.calls.last.request.content)["document"] + assert document["type"] == "image_url" + assert document["image_url"].startswith("data:image/png;base64,") + assert "error" not in result From 5a22edb6c3e223f1fecd08eeb966b4ce65b7b3ce Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Sat, 5 Sep 2026 12:07:48 -0700 Subject: [PATCH 07/35] feat(ui): explain how guardrail usage and cost are calculated Adds a "How is this calculated?" hover to the Guardrail Cost card on the overview and to the Cost and Usage Units cards on the detail page. The overview hint lists each guardrail's cost and the total; the detail cost hint shows units x per-unit price per counter with unpriced units called out, and the units hint shows the per-counter sum. Also moves the Status column to the front of the overview table. Refs LIT-5652 --- .../GuardrailUsageBreakdown.test.tsx | 32 +++++++++ .../_components/GuardrailUsageBreakdown.tsx | 31 +++++++- .../_components/GuardrailsOverview.test.tsx | 14 ++++ .../_components/GuardrailsOverview.tsx | 67 ++++++++++++----- .../GuardrailsMonitor/MetricCard.tsx | 25 ++++++- .../GuardrailsMonitor/usageUnits.test.ts | 71 ++++++++++++++++++- .../GuardrailsMonitor/usageUnits.ts | 32 +++++++++ 7 files changed, 250 insertions(+), 22 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx index 7929b97b000..ba90ca8e6ff 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx @@ -1,4 +1,5 @@ import { render, screen, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; import { describe, expect, it } from "vitest"; import type { GuardrailUsageDetail } from "@/app/(dashboard)/hooks/guardrails/useGuardrailsUsage"; import { GuardrailUsageBreakdown } from "./GuardrailUsageBreakdown"; @@ -84,6 +85,37 @@ describe("GuardrailUsageBreakdown", () => { expect(within(unpricedKey).getByText("7", { selector: ".text-warning" })).toBeInTheDocument(); }); + it("explains the cost math per counter on hover", async () => { + const user = userEvent.setup(); + render(); + + await user.hover( + within(screen.getByRole("group", { name: "Cost" })).getByRole("button", { name: /How is this calculated/ }), + ); + + expect(await screen.findByText("Content Policy: 1,000 × $0.00015 = $0.1500")).toBeInTheDocument(); + expect(screen.getByText("Sensitive Information Policy: 300 × $0.0001 = $0.0300")).toBeInTheDocument(); + expect(screen.getByText("Some Future Counter: 7 units with no known price, left out")).toBeInTheDocument(); + expect(screen.getByText("Total: $0.1800")).toBeInTheDocument(); + }); + + it("explains the units sum on hover", async () => { + const user = userEvent.setup(); + render(); + + await user.hover( + within(screen.getByRole("group", { name: "Usage Units" })).getByRole("button", { + name: /How is this calculated/, + }), + ); + + expect( + await screen.findByText( + "Content Policy 1,000 + Sensitive Information Policy 300 + Some Future Counter 7 = 1,307", + ), + ).toBeInTheDocument(); + }); + it("orders teams and keys by units, largest first", () => { render(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx index 6e8b725aa2e..01dd8f79ce4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx @@ -3,7 +3,14 @@ import { CircleDollarSign } from "lucide-react"; import React from "react"; import type { GuardrailUsageDetail } from "@/app/(dashboard)/hooks/guardrails/useGuardrailsUsage"; import { MetricCard } from "@/components/GuardrailsMonitor/MetricCard"; -import { counterLabel, formatCost, totalUnits, unpricedSummary } from "@/components/GuardrailsMonitor/usageUnits"; +import { + counterLabel, + counterMathLine, + formatCost, + totalUnits, + unitsSumLine, + unpricedSummary, +} from "@/components/GuardrailsMonitor/usageUnits"; import { DataTable } from "@/components/shared/DataTable"; import { IdCell } from "@/components/shared/table_cells/id_cell"; import { MoneyCell } from "@/components/shared/table_cells/money_cell"; @@ -104,6 +111,26 @@ const groupColumns = (label: string, emptyLabel: string): ColumnDef[] const teamColumns = groupColumns("Team", "No team"); const keyColumns = groupColumns("Key", "No key"); +const CostMath = ({ counters, total }: { counters: CounterRow[]; total: number | null }) => ( +
    + {counters.map((row) => ( +
    {counterMathLine(row)}
    + ))} +
    Total: {formatCost(total)}
    +
    Per-unit prices come from the bedrock/guardrails entry in the cost map.
    +
    +); + +const UnitsMath = ({ units }: { units: GuardrailUsageDetail["usage_units"] }) => ( +
    +
    {unitsSumLine(units)}
    +
    + Bedrock reports one unit per 1,000 characters of the message for each policy the guardrail has on, on every call, + blocked or not. +
    +
    +); + const TableHeading = ({ title }: { title: string }) => (
    {title}
    ); @@ -132,11 +159,13 @@ export function GuardrailUsageBreakdown({ detail }: { detail: GuardrailUsageDeta valueColor={detail.cost != null ? "text-foreground" : "text-muted-foreground"} icon={} subtitle={unpriced ?? undefined} + hint={} /> } />
    diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx index c52645def70..14e070afc0d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx @@ -211,6 +211,20 @@ describe("GuardrailsOverview", () => { expect(card).toHaveTextContent("250 units unpriced"); }); + it("explains the guardrail cost total on hover", async () => { + const user = userEvent.setup(); + renderOverview(); + + const card = await screen.findByRole("group", { name: "Guardrail Cost" }); + await user.hover(within(card).getByRole("button", { name: /How is this calculated/ })); + + expect(await screen.findByText("High Failure Guardrail: $0.1500")).toBeInTheDocument(); + expect(screen.getByText("Free Bedrock Guardrail: $0.0000")).toBeInTheDocument(); + expect(screen.queryByText(/Low Failure Guardrail: /)).not.toBeInTheDocument(); + expect(screen.getByText("Total: $0.1500")).toBeInTheDocument(); + expect(screen.getByText(/250 units unpriced had no known price and are left out/)).toBeInTheDocument(); + }); + it("shows a dash for guardrail cost when nothing in the window was priced", async () => { useGuardrailsUsageOverviewMock.mockReturnValue({ data: { ...overview, totalCost: null, totalUntrackedUsageUnits: {} }, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx index 0bbda6015b4..3df7058baba 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx @@ -63,6 +63,34 @@ function UsageUnitsCell({ units }: { units: GuardrailUsageOverviewRow["usageUnit ); } +function TotalCostMath({ + rows, + total, + unpriced, +}: { + rows: GuardrailUsageOverviewRow[]; + total: number | null; + unpriced: string | null; +}) { + return ( +
    + {rows + .filter((row) => row.cost != null) + .map((row) => ( +
    + {row.name}: {formatCost(row.cost)} +
    + ))} +
    Total: {formatCost(total)}
    +
    + {`Each guardrail's cost is its units per policy × that policy's per-unit price from the cost map, added up${ + unpriced ? `; ${unpriced} had no known price and are left out` : "" + }. Open a guardrail for its per-policy math.`} +
    +
    + ); +} + function CostCell({ row }: { row: GuardrailUsageOverviewRow }) { const unpriced = unpricedSummary(row.untrackedUsageUnits); return ( @@ -124,6 +152,25 @@ export function GuardrailsOverview({ const error = guardrailsError; const columns: ColumnDef[] = [ + { + header: "Status", + accessorKey: "status", + enableSorting: false, + cell: ({ row }) => ( + + + {row.original.status} + + ), + }, { header: "Guardrail", accessorKey: "name", @@ -215,25 +262,6 @@ export function GuardrailsOverview({ sortDescFirst: false, cell: ({ row }) => , }, - { - header: "Status", - accessorKey: "status", - enableSorting: false, - cell: ({ row }) => ( - - - {row.original.status} - - ), - }, ]; const sortableKeys: SortKey[] = ["failRate", "requestsEvaluated", "avgLatency", "cost"]; @@ -291,6 +319,7 @@ export function GuardrailsOverview({ valueColor={metrics.totalCost != null ? "text-foreground" : "text-muted-foreground"} icon={} subtitle={metrics.unpriced ?? undefined} + hint={} />
    diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/MetricCard.tsx b/ui/litellm-dashboard/src/components/GuardrailsMonitor/MetricCard.tsx index c0b5e0a50d1..1805dc797e4 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/MetricCard.tsx +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/MetricCard.tsx @@ -1,4 +1,6 @@ +import { CircleHelp } from "lucide-react"; import React, { type ReactNode } from "react"; +import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; interface MetricCardProps { label: string; @@ -6,9 +8,10 @@ interface MetricCardProps { valueColor?: string; icon?: ReactNode; subtitle?: string; + hint?: ReactNode; } -export function MetricCard({ label, value, valueColor = "text-foreground", icon, subtitle }: MetricCardProps) { +export function MetricCard({ label, value, valueColor = "text-foreground", icon, subtitle, hint }: MetricCardProps) { return (
    @@ -17,6 +20,26 @@ export function MetricCard({ label, value, valueColor = "text-foreground", icon,
    {value}
    {subtitle &&

    {subtitle}

    } + {hint && ( + + + + + How is this calculated? + + } + /> + + {hint} + + + + )}
    ); } diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts index 29362cc3701..560010bd852 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts @@ -1,5 +1,14 @@ import { describe, expect, it } from "vitest"; -import { counterLabel, formatCost, totalUnits, unpricedSummary } from "./usageUnits"; +import { + counterLabel, + counterMathLine, + formatCost, + formatUnitPrice, + totalUnits, + unitPrice, + unitsSumLine, + unpricedSummary, +} from "./usageUnits"; describe("formatCost", () => { it("renders a dash when nothing was priced", () => { @@ -53,3 +62,63 @@ describe("unpricedSummary", () => { expect(unpricedSummary({ someFutureCounter: 1 })).toBe("1 unit unpriced"); }); }); + +describe("unitPrice", () => { + it("backs the per-unit price out of the priced share only", () => { + expect(unitPrice({ counter: "contentPolicyUnits", units: 1200, unpriced: 200, cost: 0.15 })).toBeCloseTo( + 0.00015, + 10, + ); + }); + + it("is null when nothing was priced", () => { + expect(unitPrice({ counter: "someFutureCounter", units: 7, unpriced: 7, cost: null })).toBeNull(); + expect(unitPrice({ counter: "someFutureCounter", units: 7, unpriced: 7, cost: 0 })).toBeNull(); + }); +}); + +describe("formatUnitPrice", () => { + it("keeps the significant decimals and drops trailing zeros", () => { + expect(formatUnitPrice(0.0001)).toBe("$0.0001"); + expect(formatUnitPrice(0.00015)).toBe("$0.00015"); + expect(formatUnitPrice(0)).toBe("$0"); + expect(formatUnitPrice(1)).toBe("$1"); + }); +}); + +describe("counterMathLine", () => { + it("shows units × price = cost for a fully priced counter", () => { + expect(counterMathLine({ counter: "contentPolicyUnits", units: 1000, unpriced: 0, cost: 0.15 })).toBe( + "Content Policy: 1,000 × $0.00015 = $0.1500", + ); + }); + + it("prices only the priced share and calls out the rest", () => { + expect(counterMathLine({ counter: "sensitiveInformationPolicyUnits", units: 8, unpriced: 2, cost: 0.0006 })).toBe( + "Sensitive Information Policy: 6 × $0.0001 = $0.0006 (2 unpriced left out)", + ); + }); + + it("says so when a counter has no known price at all", () => { + expect(counterMathLine({ counter: "someFutureCounter", units: 7, unpriced: 7, cost: null })).toBe( + "Some Future Counter: 7 units with no known price, left out", + ); + expect(counterMathLine({ counter: "someFutureCounter", units: 1, unpriced: 1, cost: null })).toBe( + "Some Future Counter: 1 unit with no known price, left out", + ); + }); + + it("shows a free counter as × $0", () => { + expect(counterMathLine({ counter: "wordPolicyUnits", units: 2, unpriced: 0, cost: 0 })).toBe( + "Word Policy: 2 × $0 = $0.0000", + ); + }); +}); + +describe("unitsSumLine", () => { + it("adds the counters up in order", () => { + expect(unitsSumLine({ contentPolicyUnits: 2, topicPolicyUnits: 2, wordPolicyUnits: 1200 })).toBe( + "Content Policy 2 + Topic Policy 2 + Word Policy 1,200 = 1,204", + ); + }); +}); diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.ts b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.ts index f3a5e9d7140..05e046aaace 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.ts +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.ts @@ -19,3 +19,35 @@ export const unpricedSummary = (untracked: UsageUnits): string | null => { const total = totalUnits(untracked); return total > 0 ? `${total.toLocaleString()} ${total === 1 ? "unit" : "units"} unpriced` : null; }; + +export interface CounterMath { + readonly counter: string; + readonly units: number; + readonly unpriced: number; + readonly cost: number | null; +} + +export const pricedUnits = ({ units, unpriced }: Pick): number => + Math.max(units - unpriced, 0); + +export const unitPrice = (row: CounterMath): number | null => { + const priced = pricedUnits(row); + return row.cost != null && priced > 0 ? row.cost / priced : null; +}; + +export const formatUnitPrice = (price: number): string => `$${price.toFixed(6).replace(/\.?0+$/, "")}`; + +export const counterMathLine = (row: CounterMath): string => { + const label = counterLabel(row.counter); + const price = unitPrice(row); + if (price == null) { + return `${label}: ${row.units.toLocaleString()} ${row.units === 1 ? "unit" : "units"} with no known price, left out`; + } + const line = `${label}: ${pricedUnits(row).toLocaleString()} × ${formatUnitPrice(price)} = ${formatCost(row.cost)}`; + return row.unpriced > 0 ? `${line} (${row.unpriced.toLocaleString()} unpriced left out)` : line; +}; + +export const unitsSumLine = (units: UsageUnits): string => + `${Object.entries(units) + .map(([counter, n]) => `${counterLabel(counter)} ${n.toLocaleString()}`) + .join(" + ")} = ${totalUnits(units).toLocaleString()}`; From e66ba0533fe43a626b8a157fae59f507894aa13b Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Sat, 5 Sep 2026 12:23:07 -0700 Subject: [PATCH 08/35] fix(ui): keep guardrail cost hints provider neutral and link to a pricing request The hint copy described Bedrock's unit semantics and cost map entry even though any provider's units reach this view, so it now explains the math in provider-neutral terms. When units have no known price, the hint says so and links to a prefilled GitHub feature request (provider and counter names filled in) so the reader can ask for pricing. Per-unit prices below $0.000001 now read "< $0.000001" instead of "$0". Refs LIT-5652 --- .../GuardrailUsageBreakdown.test.tsx | 28 +++++++++++++++++++ .../_components/GuardrailUsageBreakdown.tsx | 15 +++++----- .../_components/GuardrailsOverview.test.tsx | 6 +++- .../_components/GuardrailsOverview.tsx | 26 ++++++++++------- .../GuardrailsMonitor/UnpricedNote.tsx | 21 ++++++++++++++ .../GuardrailsMonitor/usageUnits.test.ts | 22 +++++++++++++++ .../GuardrailsMonitor/usageUnits.ts | 15 +++++++++- 7 files changed, 113 insertions(+), 20 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/GuardrailsMonitor/UnpricedNote.tsx diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx index ba90ca8e6ff..3a0a4c38ecb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx @@ -97,6 +97,34 @@ describe("GuardrailUsageBreakdown", () => { expect(screen.getByText("Sensitive Information Policy: 300 × $0.0001 = $0.0300")).toBeInTheDocument(); expect(screen.getByText("Some Future Counter: 7 units with no known price, left out")).toBeInTheDocument(); expect(screen.getByText("Total: $0.1800")).toBeInTheDocument(); + expect(screen.getByText(/7 units with no known price are left out of the cost/)).toBeInTheDocument(); + const issueLink = screen.getByRole("link", { name: "Request pricing on GitHub" }); + expect(issueLink).toHaveAttribute("target", "_blank"); + const issueUrl = new URL(issueLink.getAttribute("href") ?? ""); + expect(issueUrl.searchParams.get("title")).toBe("[Feature]: add Bedrock guardrail pricing to the cost map"); + expect(issueUrl.searchParams.get("the-feature")).toContain("someFutureCounter"); + }); + + it("does not ask for pricing when every unit was priced", async () => { + const user = userEvent.setup(); + render( + , + ); + + await user.hover( + within(screen.getByRole("group", { name: "Cost" })).getByRole("button", { name: /How is this calculated/ }), + ); + + expect(await screen.findByText("Total: $0.1500")).toBeInTheDocument(); + expect(screen.queryByRole("link", { name: "Request pricing on GitHub" })).not.toBeInTheDocument(); }); it("explains the units sum on hover", async () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx index 01dd8f79ce4..e67425adb10 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx @@ -3,6 +3,7 @@ import { CircleDollarSign } from "lucide-react"; import React from "react"; import type { GuardrailUsageDetail } from "@/app/(dashboard)/hooks/guardrails/useGuardrailsUsage"; import { MetricCard } from "@/components/GuardrailsMonitor/MetricCard"; +import { UnpricedNote } from "@/components/GuardrailsMonitor/UnpricedNote"; import { counterLabel, counterMathLine, @@ -111,23 +112,21 @@ const groupColumns = (label: string, emptyLabel: string): ColumnDef[] const teamColumns = groupColumns("Team", "No team"); const keyColumns = groupColumns("Key", "No key"); -const CostMath = ({ counters, total }: { counters: CounterRow[]; total: number | null }) => ( +const CostMath = ({ counters, detail }: { counters: CounterRow[]; detail: GuardrailUsageDetail }) => (
    {counters.map((row) => (
    {counterMathLine(row)}
    ))} -
    Total: {formatCost(total)}
    -
    Per-unit prices come from the bedrock/guardrails entry in the cost map.
    +
    Total: {formatCost(detail.cost)}
    +
    Each counter is its priced units × the per-unit price LiteLLM has for it in the cost map.
    +
    ); const UnitsMath = ({ units }: { units: GuardrailUsageDetail["usage_units"] }) => (
    {unitsSumLine(units)}
    -
    - Bedrock reports one unit per 1,000 characters of the message for each policy the guardrail has on, on every call, - blocked or not. -
    +
    Units are the billable counters the provider reported for this guardrail, added up over every call.
    ); @@ -159,7 +158,7 @@ export function GuardrailUsageBreakdown({ detail }: { detail: GuardrailUsageDeta valueColor={detail.cost != null ? "text-foreground" : "text-muted-foreground"} icon={} subtitle={unpriced ?? undefined} - hint={} + hint={} /> { expect(screen.getByText("Free Bedrock Guardrail: $0.0000")).toBeInTheDocument(); expect(screen.queryByText(/Low Failure Guardrail: /)).not.toBeInTheDocument(); expect(screen.getByText("Total: $0.1500")).toBeInTheDocument(); - expect(screen.getByText(/250 units unpriced had no known price and are left out/)).toBeInTheDocument(); + expect(screen.getByText(/250 units with no known price are left out of the cost/)).toBeInTheDocument(); + const issueLink = screen.getByRole("link", { name: "Request pricing on GitHub" }); + const issueUrl = new URL(issueLink.getAttribute("href") ?? ""); + expect(issueUrl.searchParams.get("template")).toBe("feature_request.yml"); + expect(issueUrl.searchParams.get("the-feature")).toContain("sensitiveInformationPolicyUnits"); }); it("shows a dash for guardrail cost when nothing in the window was priced", async () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx index 3df7058baba..33be85f3c81 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx @@ -8,7 +8,14 @@ import { type GuardrailUsageOverviewRow, useGuardrailsUsageOverview, } from "@/app/(dashboard)/hooks/guardrails/useGuardrailsUsage"; -import { counterLabel, formatCost, totalUnits, unpricedSummary } from "@/components/GuardrailsMonitor/usageUnits"; +import { UnpricedNote } from "@/components/GuardrailsMonitor/UnpricedNote"; +import { + counterLabel, + formatCost, + totalUnits, + unpricedSummary, + type UsageUnits, +} from "@/components/GuardrailsMonitor/usageUnits"; import { Button } from "@/components/ui/button"; import { PageHeader } from "@/components/shared/PageHeader"; import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; @@ -41,7 +48,7 @@ const EMPTY_METRICS = { avgLatency: 0, count: 0, totalCost: null as number | null, - unpriced: null as string | null, + untracked: {} as UsageUnits, }; function UsageUnitsCell({ units }: { units: GuardrailUsageOverviewRow["usageUnits"] }) { @@ -66,11 +73,11 @@ function UsageUnitsCell({ units }: { units: GuardrailUsageOverviewRow["usageUnit function TotalCostMath({ rows, total, - unpriced, + untracked, }: { rows: GuardrailUsageOverviewRow[]; total: number | null; - unpriced: string | null; + untracked: UsageUnits; }) { return (
    @@ -83,10 +90,9 @@ function TotalCostMath({ ))}
    Total: {formatCost(total)}
    - {`Each guardrail's cost is its units per policy × that policy's per-unit price from the cost map, added up${ - unpriced ? `; ${unpriced} had no known price and are left out` : "" - }. Open a guardrail for its per-policy math.`} + {`Each guardrail's cost is its units per counter × that counter's per-unit price from the cost map, added up. Open a guardrail for its per-counter math.`}
    +
    ); } @@ -135,7 +141,7 @@ export function GuardrailsOverview({ : 0, count: activeData.length, totalCost: guardrailsData.totalCost, - unpriced: unpricedSummary(guardrailsData.totalUntrackedUsageUnits), + untracked: guardrailsData.totalUntrackedUsageUnits, }; }, [guardrailsData, activeData]); const chartData = guardrailsData?.chart; @@ -318,8 +324,8 @@ export function GuardrailsOverview({ value={formatCost(metrics.totalCost)} valueColor={metrics.totalCost != null ? "text-foreground" : "text-muted-foreground"} icon={} - subtitle={metrics.unpriced ?? undefined} - hint={} + subtitle={unpricedSummary(metrics.untracked) ?? undefined} + hint={} />
    diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/UnpricedNote.tsx b/ui/litellm-dashboard/src/components/GuardrailsMonitor/UnpricedNote.tsx new file mode 100644 index 00000000000..174124d18aa --- /dev/null +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/UnpricedNote.tsx @@ -0,0 +1,21 @@ +import React from "react"; +import { pricingIssueUrl, totalUnits, type UsageUnits } from "./usageUnits"; + +export function UnpricedNote({ unpriced, provider }: { unpriced: UsageUnits; provider?: string }) { + const total = totalUnits(unpriced); + if (total === 0) return null; + const [noun, verb] = total === 1 ? ["unit", "is"] : ["units", "are"]; + return ( +
    + {`${total.toLocaleString()} ${noun} with no known price ${verb} left out of the cost. `} + + Request pricing on GitHub + +
    + ); +} diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts index 560010bd852..9b45eefaf54 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts @@ -4,6 +4,7 @@ import { counterMathLine, formatCost, formatUnitPrice, + pricingIssueUrl, totalUnits, unitPrice, unitsSumLine, @@ -84,6 +85,10 @@ describe("formatUnitPrice", () => { expect(formatUnitPrice(0)).toBe("$0"); expect(formatUnitPrice(1)).toBe("$1"); }); + + it("never shows a positive price as free", () => { + expect(formatUnitPrice(0.0000002)).toBe("< $0.000001"); + }); }); describe("counterMathLine", () => { @@ -122,3 +127,20 @@ describe("unitsSumLine", () => { ); }); }); + +describe("pricingIssueUrl", () => { + it("prefills the feature request with the provider and the unpriced counters", () => { + const url = new URL(pricingIssueUrl({ text_records: 5, someFutureCounter: 7 }, "azure/prompt_shield")); + + expect(url.origin + url.pathname).toBe("https://github.com/BerriAI/litellm/issues/new"); + expect(url.searchParams.get("template")).toBe("feature_request.yml"); + expect(url.searchParams.get("title")).toBe("[Feature]: add azure/prompt_shield guardrail pricing to the cost map"); + expect(url.searchParams.get("the-feature")).toContain("text_records, someFutureCounter"); + }); + + it("stays generic when no provider is known", () => { + const url = new URL(pricingIssueUrl({ text_records: 5 })); + + expect(url.searchParams.get("title")).toBe("[Feature]: add guardrail pricing to the cost map"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.ts b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.ts index 05e046aaace..f914a3e9698 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.ts +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.ts @@ -35,7 +35,10 @@ export const unitPrice = (row: CounterMath): number | null => { return row.cost != null && priced > 0 ? row.cost / priced : null; }; -export const formatUnitPrice = (price: number): string => `$${price.toFixed(6).replace(/\.?0+$/, "")}`; +export const formatUnitPrice = (price: number): string => { + const fixed = price.toFixed(6).replace(/\.?0+$/, ""); + return price > 0 && Number(fixed) === 0 ? "< $0.000001" : `$${fixed}`; +}; export const counterMathLine = (row: CounterMath): string => { const label = counterLabel(row.counter); @@ -51,3 +54,13 @@ export const unitsSumLine = (units: UsageUnits): string => `${Object.entries(units) .map(([counter, n]) => `${counterLabel(counter)} ${n.toLocaleString()}`) .join(" + ")} = ${totalUnits(units).toLocaleString()}`; + +export const pricingIssueUrl = (unpriced: UsageUnits, provider?: string): string => { + const subject = provider ? `${provider} guardrail` : "guardrail"; + const params = new URLSearchParams({ + template: "feature_request.yml", + title: `[Feature]: add ${subject} pricing to the cost map`, + "the-feature": `LiteLLM has no price for these ${subject} usage units, so the Guardrails Monitor leaves them out of the cost: ${Object.keys(unpriced).join(", ")}`, + }); + return `https://github.com/BerriAI/litellm/issues/new?${params.toString()}`; +}; From def734923f683055f04fc5e46d9ef78ecf912381 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Sat, 5 Sep 2026 12:55:30 -0700 Subject: [PATCH 09/35] test(e2e): prove Vertex context caching on the first cold call and on the spend row --- .../e2e/llm_translation/test_cache_control.py | 76 +++++++++++++++++-- tests/e2e/models.py | 1 + 2 files changed, 71 insertions(+), 6 deletions(-) diff --git a/tests/e2e/llm_translation/test_cache_control.py b/tests/e2e/llm_translation/test_cache_control.py index 0d224061381..3ad98bc6072 100644 --- a/tests/e2e/llm_translation/test_cache_control.py +++ b/tests/e2e/llm_translation/test_cache_control.py @@ -8,8 +8,12 @@ Each case asserts the feature actually happened, not just a 200. Coverage matrix cache-read usage tokens > 0. service_tier is out of scope for Bedrock; AWS Bedrock does not expose an OpenAI-style request service tier, so that cell is intentionally not covered here. -- Vertex (gemini-2.5-flash): prompt caching via ``cache_control`` context - caching; the second identical call must report cached prompt tokens > 0. +- Vertex (gemini-2.5-flash): explicit context caching via ``cache_control`` + with a 5-minute ttl. litellm builds the Vertex cache before the generate + call, so a never-seen prefix must come back cached on its very first call + (Gemini's implicit caching cannot hit a cold prefix), the cached count must + cover the marked block, and the spend row must be billed below the uncached + price of the prompt. - Anthropic (claude-haiku-4-5, direct): the same ``cache_control`` prefix over the OpenAI-compatible route; the second call must report cache-read tokens > 0. - OpenAI (gpt-5.6): automatic prompt caching needs no request marker, so the @@ -26,7 +30,8 @@ built from the typed content blocks shared in ``endpoints_client.py``. from __future__ import annotations import time -from collections.abc import Callable +from collections.abc import Callable, Iterator +from typing import Final import pytest from pydantic import BaseModel @@ -45,6 +50,10 @@ BEDROCK_MODEL = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" VERTEX_MODEL = "vertex_ai/gemini-2.5-flash" ANTHROPIC_MODEL = "anthropic/claude-haiku-4-5-20251001" OPENAI_MODEL = "openai/gpt-5.6" +VERTEX_CACHE_TTL: Final = "300s" +VERTEX_COLD_CALL_ATTEMPTS: Final = 3 +VERTEX_MINIMUM_CACHED_TOKENS: Final = 1024 +CACHED_SHARE_OF_PROMPT: Final = 0.9 class CacheChatBody(BaseModel): @@ -77,14 +86,14 @@ def _cached_read_tokens(usage: Usage | None) -> int: def _cache_chat( - client: PassthroughClient, key: str, model: str, prefix: str + client: PassthroughClient, key: str, model: str, prefix: str, ttl: str | None = None ) -> Result[ChatResponse]: body = CacheChatBody( model=model, messages=[ RichMessage( role="system", - content=[TextBlock(text=prefix, cache_control=CacheControl())], + content=[TextBlock(text=prefix, cache_control=CacheControl(ttl=ttl))], ), RichMessage(role="user", content=[TextBlock(text="Reply with one word.")]), ], @@ -138,6 +147,58 @@ def _assert_cache_read_on_second_call( ) +def _cold_cache_calls(send: Callable[[str], Result[ChatResponse]]) -> Iterator[ChatResponse]: + for _ in range(VERTEX_COLD_CALL_ATTEMPTS): + yield unwrap(send(_cacheable_prefix())) + + +def _first_cold_call_reads_cache(model: str, send: Callable[[str], Result[ChatResponse]]) -> ChatResponse: + completion: Final = next( + ( + candidate + for candidate in _cold_cache_calls(send) + if _cached_read_tokens(candidate.usage) >= VERTEX_MINIMUM_CACHED_TOKENS + ), + None, + ) + assert completion is not None, ( + f"{model}: {VERTEX_COLD_CALL_ATTEMPTS} never-seen prompts marked with cache_control all reported fewer " + f"than {VERTEX_MINIMUM_CACHED_TOKENS} cached tokens on their first call; explicit context caching did " + "not engage" + ) + assert completion.choices, f"{model}: cached call returned no choices: {completion}" + usage: Final = completion.usage + cached: Final = _cached_read_tokens(usage) + assert usage and usage.prompt_tokens and cached >= CACHED_SHARE_OF_PROMPT * usage.prompt_tokens, ( + f"{model}: only {cached} of {usage.prompt_tokens if usage else None} prompt tokens were served from the " + "cache; the cache_control block was not cached whole" + ) + return completion + + +def _input_rate(client: PassthroughClient, model: str) -> float: + entry: Final = next((row for row in client.proxy.model_info() if row.model_name == model), None) + assert entry and entry.model_info.input_cost_per_token, f"/model/info resolved no input rate for {model}" + return entry.model_info.input_cost_per_token + + +def _assert_billed_below_uncached_prompt(client: PassthroughClient, model: str, completion: ChatResponse) -> None: + assert completion.id, f"{model}: cached completion carried no id to find its spend row by" + usage: Final = completion.usage + assert usage and usage.prompt_tokens, f"{model}: cached completion carried no prompt_tokens: {usage}" + rows: Final = client.proxy.poll_logs_for_request_id(completion.id, predicate=lambda rs: (rs[0].spend or 0) > 0) + assert rows, f"{model}: no costed /spend/logs row for request {completion.id}" + row: Final = rows[0] + assert row.prompt_tokens == usage.prompt_tokens, ( + f"{model}: spend row prompt_tokens {row.prompt_tokens} != response prompt_tokens {usage.prompt_tokens}" + ) + uncached_prompt_cost: Final = usage.prompt_tokens * _input_rate(client, model) + assert row.spend is not None and row.spend < uncached_prompt_cost, ( + f"{model}: spend {row.spend} is not below the uncached price of the prompt alone ({uncached_prompt_cost} for " + f"{usage.prompt_tokens} tokens); cache-read pricing was not applied" + ) + + class TestCacheControl: @pytest.mark.covers( "llm.chat_completions.bedrock_converse.prompt_cache_5m.nonstream.works", @@ -174,7 +235,10 @@ class TestCacheControl: ) resources.defer(lambda: client.proxy.delete_model(model_id)) key = resources.key() - _assert_cache_read_on_second_call(model, lambda prefix: _cache_chat(client, key, model, prefix)) + completion = _first_cold_call_reads_cache( + model, lambda prefix: _cache_chat(client, key, model, prefix, ttl=VERTEX_CACHE_TTL) + ) + _assert_billed_below_uncached_prompt(client, model, completion) @pytest.mark.covers( "llm.chat_completions.anthropic.prompt_cache_5m.nonstream.works", diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 5de49ead3ed..016a9de56b8 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -181,6 +181,7 @@ class ChatMessage(BaseModel): class CacheControl(BaseModel): type: str = "ephemeral" + ttl: str | None = None class TextBlock(BaseModel): From 29b93b57aaa0cef7660da3636df377422a2cd704 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Sat, 5 Sep 2026 12:58:18 -0700 Subject: [PATCH 10/35] feat(ui): show the guardrail cost math in a popover table MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The "How is this calculated?" hover was a plain-text tooltip. It is now a popover (opens on hover or click) with a title, the formula, a table of one row per counter or guardrail (units, × price, = cost, with unpriced units called out under the row) and a total row, so the math reads as a worked sum instead of a sentence. Refs LIT-5652 --- .../GuardrailUsageBreakdown.test.tsx | 55 ++++++++++----- .../_components/GuardrailUsageBreakdown.tsx | 26 +++---- .../_components/GuardrailsOverview.test.tsx | 25 ++++--- .../_components/GuardrailsOverview.tsx | 25 ++++--- .../GuardrailsMonitor/CalcPopover.tsx | 67 +++++++++++++++++++ .../GuardrailsMonitor/MetricCard.tsx | 23 +------ .../GuardrailsMonitor/UnpricedNote.tsx | 4 +- .../GuardrailsMonitor/usageUnits.test.ts | 55 +++++++++------ .../GuardrailsMonitor/usageUnits.ts | 30 ++++++--- 9 files changed, 204 insertions(+), 106 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/GuardrailsMonitor/CalcPopover.tsx diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx index 3a0a4c38ecb..db7855ab0d5 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.test.tsx @@ -85,20 +85,33 @@ describe("GuardrailUsageBreakdown", () => { expect(within(unpricedKey).getByText("7", { selector: ".text-warning" })).toBeInTheDocument(); }); - it("explains the cost math per counter on hover", async () => { + const cellsOf = (dialog: HTMLElement): string[][] => + within(dialog) + .getAllByRole("row") + .map((row) => + within(row) + .getAllByRole("cell") + .map((cell) => cell.textContent ?? ""), + ); + + it("lays the cost math out per counter as units × price = cost", async () => { const user = userEvent.setup(); render(); - await user.hover( + await user.click( within(screen.getByRole("group", { name: "Cost" })).getByRole("button", { name: /How is this calculated/ }), ); - expect(await screen.findByText("Content Policy: 1,000 × $0.00015 = $0.1500")).toBeInTheDocument(); - expect(screen.getByText("Sensitive Information Policy: 300 × $0.0001 = $0.0300")).toBeInTheDocument(); - expect(screen.getByText("Some Future Counter: 7 units with no known price, left out")).toBeInTheDocument(); - expect(screen.getByText("Total: $0.1800")).toBeInTheDocument(); - expect(screen.getByText(/7 units with no known price are left out of the cost/)).toBeInTheDocument(); - const issueLink = screen.getByRole("link", { name: "Request pricing on GitHub" }); + const dialog = await screen.findByRole("dialog", { name: "How this cost is calculated" }); + expect(cellsOf(dialog)).toEqual([ + ["Content Policy", "1,000", "× $0.00015", "= $0.1500"], + ["Sensitive Information Policy", "300", "× $0.0001", "= $0.0300"], + ["Some Future Counter", "7", "× —", "= —"], + ["no known price, left out"], + ["Total", "$0.1800"], + ]); + expect(within(dialog).getByText(/7 units with no known price are left out of the cost/)).toBeInTheDocument(); + const issueLink = within(dialog).getByRole("link", { name: "Request pricing on GitHub" }); expect(issueLink).toHaveAttribute("target", "_blank"); const issueUrl = new URL(issueLink.getAttribute("href") ?? ""); expect(issueUrl.searchParams.get("title")).toBe("[Feature]: add Bedrock guardrail pricing to the cost map"); @@ -119,29 +132,35 @@ describe("GuardrailUsageBreakdown", () => { />, ); - await user.hover( + await user.click( within(screen.getByRole("group", { name: "Cost" })).getByRole("button", { name: /How is this calculated/ }), ); - expect(await screen.findByText("Total: $0.1500")).toBeInTheDocument(); - expect(screen.queryByRole("link", { name: "Request pricing on GitHub" })).not.toBeInTheDocument(); + const dialog = await screen.findByRole("dialog", { name: "How this cost is calculated" }); + expect(cellsOf(dialog)).toEqual([ + ["Content Policy", "1,000", "× $0.00015", "= $0.1500"], + ["Total", "$0.1500"], + ]); + expect(within(dialog).queryByRole("link", { name: "Request pricing on GitHub" })).not.toBeInTheDocument(); }); - it("explains the units sum on hover", async () => { + it("lays the units sum out per counter", async () => { const user = userEvent.setup(); render(); - await user.hover( + await user.click( within(screen.getByRole("group", { name: "Usage Units" })).getByRole("button", { name: /How is this calculated/, }), ); - expect( - await screen.findByText( - "Content Policy 1,000 + Sensitive Information Policy 300 + Some Future Counter 7 = 1,307", - ), - ).toBeInTheDocument(); + const dialog = await screen.findByRole("dialog", { name: "How usage units add up" }); + expect(cellsOf(dialog)).toEqual([ + ["Content Policy", "1,000"], + ["Sensitive Information Policy", "300"], + ["Some Future Counter", "7"], + ["Total", "1,307"], + ]); }); it("orders teams and keys by units, largest first", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx index e67425adb10..27d9ba5162f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailUsageBreakdown.tsx @@ -2,14 +2,15 @@ import type { ColumnDef } from "@tanstack/react-table"; import { CircleDollarSign } from "lucide-react"; import React from "react"; import type { GuardrailUsageDetail } from "@/app/(dashboard)/hooks/guardrails/useGuardrailsUsage"; +import { CalcPopover, MathTable } from "@/components/GuardrailsMonitor/CalcPopover"; import { MetricCard } from "@/components/GuardrailsMonitor/MetricCard"; import { UnpricedNote } from "@/components/GuardrailsMonitor/UnpricedNote"; import { counterLabel, - counterMathLine, + counterMathRow, formatCost, totalUnits, - unitsSumLine, + unitsMathRows, unpricedSummary, } from "@/components/GuardrailsMonitor/usageUnits"; import { DataTable } from "@/components/shared/DataTable"; @@ -113,21 +114,20 @@ const teamColumns = groupColumns("Team", "No team"); const keyColumns = groupColumns("Key", "No key"); const CostMath = ({ counters, detail }: { counters: CounterRow[]; detail: GuardrailUsageDetail }) => ( -
    - {counters.map((row) => ( -
    {counterMathLine(row)}
    - ))} -
    Total: {formatCost(detail.cost)}
    -
    Each counter is its priced units × the per-unit price LiteLLM has for it in the cost map.
    + + +

    Per-unit prices come from the cost map LiteLLM ships with.

    -
    + ); const UnitsMath = ({ units }: { units: GuardrailUsageDetail["usage_units"] }) => ( -
    -
    {unitsSumLine(units)}
    -
    Units are the billable counters the provider reported for this guardrail, added up over every call.
    -
    + + +

    + Units are the billable counters the provider reported for this guardrail, added up over every call. +

    +
    ); const TableHeading = ({ title }: { title: string }) => ( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx index 16361bdec27..959ed8b172e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx @@ -211,19 +211,28 @@ describe("GuardrailsOverview", () => { expect(card).toHaveTextContent("250 units unpriced"); }); - it("explains the guardrail cost total on hover", async () => { + it("lays the guardrail cost total out per guardrail", async () => { const user = userEvent.setup(); renderOverview(); const card = await screen.findByRole("group", { name: "Guardrail Cost" }); - await user.hover(within(card).getByRole("button", { name: /How is this calculated/ })); + await user.click(within(card).getByRole("button", { name: /How is this calculated/ })); - expect(await screen.findByText("High Failure Guardrail: $0.1500")).toBeInTheDocument(); - expect(screen.getByText("Free Bedrock Guardrail: $0.0000")).toBeInTheDocument(); - expect(screen.queryByText(/Low Failure Guardrail: /)).not.toBeInTheDocument(); - expect(screen.getByText("Total: $0.1500")).toBeInTheDocument(); - expect(screen.getByText(/250 units with no known price are left out of the cost/)).toBeInTheDocument(); - const issueLink = screen.getByRole("link", { name: "Request pricing on GitHub" }); + const dialog = await screen.findByRole("dialog", { name: "How this cost is calculated" }); + const cells = within(dialog) + .getAllByRole("row") + .map((row) => + within(row) + .getAllByRole("cell") + .map((cell) => cell.textContent ?? ""), + ); + expect(cells).toEqual([ + ["High Failure Guardrail", "$0.1500"], + ["Free Bedrock Guardrail", "$0.0000"], + ["Total", "$0.1500"], + ]); + expect(within(dialog).getByText(/250 units with no known price are left out of the cost/)).toBeInTheDocument(); + const issueLink = within(dialog).getByRole("link", { name: "Request pricing on GitHub" }); const issueUrl = new URL(issueLink.getAttribute("href") ?? ""); expect(issueUrl.searchParams.get("template")).toBe("feature_request.yml"); expect(issueUrl.searchParams.get("the-feature")).toContain("sensitiveInformationPolicyUnits"); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx index 33be85f3c81..468e6967d81 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx @@ -8,6 +8,7 @@ import { type GuardrailUsageOverviewRow, useGuardrailsUsageOverview, } from "@/app/(dashboard)/hooks/guardrails/useGuardrailsUsage"; +import { CalcPopover, MathTable } from "@/components/GuardrailsMonitor/CalcPopover"; import { UnpricedNote } from "@/components/GuardrailsMonitor/UnpricedNote"; import { counterLabel, @@ -80,20 +81,18 @@ function TotalCostMath({ untracked: UsageUnits; }) { return ( -
    - {rows - .filter((row) => row.cost != null) - .map((row) => ( -
    - {row.name}: {formatCost(row.cost)} -
    - ))} -
    Total: {formatCost(total)}
    -
    - {`Each guardrail's cost is its units per counter × that counter's per-unit price from the cost map, added up. Open a guardrail for its per-counter math.`} -
    + + row.cost != null) + .map((row) => ({ label: row.name, parts: [formatCost(row.cost)], note: null }))} + total={formatCost(total)} + /> +

    + {`Each guardrail's cost is its units per counter × that counter's per-unit price from the cost map. Open a guardrail for its per-counter math.`} +

    -
    + ); } diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/CalcPopover.tsx b/ui/litellm-dashboard/src/components/GuardrailsMonitor/CalcPopover.tsx new file mode 100644 index 00000000000..686992a6fb4 --- /dev/null +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/CalcPopover.tsx @@ -0,0 +1,67 @@ +import { CircleHelp } from "lucide-react"; +import React, { type ReactNode } from "react"; +import { Popover, PopoverContent, PopoverTitle, PopoverTrigger } from "@/components/ui/popover"; +import type { MathRow } from "./usageUnits"; + +export function CalcPopover({ title, formula, children }: { title: string; formula: string; children: ReactNode }) { + return ( + + + } + > + + How is this calculated? + + + {title} + {formula} + {children} + + + ); +} + +export function MathTable({ rows, total }: { rows: readonly MathRow[]; total: string }) { + const width = 1 + Math.max(...rows.map((row) => row.parts.length), 1); + return ( + + + {rows.map((row) => ( + + + + {row.parts.map((part, i) => ( + + ))} + + {row.note && ( + + + + )} + + ))} + + + + + + + +
    {row.label} + {part} +
    + {row.note} +
    + Total + {total}
    + ); +} diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/MetricCard.tsx b/ui/litellm-dashboard/src/components/GuardrailsMonitor/MetricCard.tsx index 1805dc797e4..008dc279f13 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/MetricCard.tsx +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/MetricCard.tsx @@ -1,6 +1,4 @@ -import { CircleHelp } from "lucide-react"; import React, { type ReactNode } from "react"; -import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; interface MetricCardProps { label: string; @@ -20,26 +18,7 @@ export function MetricCard({ label, value, valueColor = "text-foreground", icon,
    {value}
    {subtitle &&

    {subtitle}

    } - {hint && ( - - - - - How is this calculated? - - } - /> - - {hint} - - - - )} + {hint} ); } diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/UnpricedNote.tsx b/ui/litellm-dashboard/src/components/GuardrailsMonitor/UnpricedNote.tsx index 174124d18aa..b43d9d18841 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/UnpricedNote.tsx +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/UnpricedNote.tsx @@ -6,7 +6,7 @@ export function UnpricedNote({ unpriced, provider }: { unpriced: UsageUnits; pro if (total === 0) return null; const [noun, verb] = total === 1 ? ["unit", "is"] : ["units", "are"]; return ( -
    +

    {`${total.toLocaleString()} ${noun} with no known price ${verb} left out of the cost. `} Request pricing on GitHub -

    +

    ); } diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts index 9b45eefaf54..8e2baaaa3c5 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts @@ -1,13 +1,13 @@ import { describe, expect, it } from "vitest"; import { counterLabel, - counterMathLine, + counterMathRow, formatCost, formatUnitPrice, pricingIssueUrl, totalUnits, unitPrice, - unitsSumLine, + unitsMathRows, unpricedSummary, } from "./usageUnits"; @@ -91,40 +91,51 @@ describe("formatUnitPrice", () => { }); }); -describe("counterMathLine", () => { +describe("counterMathRow", () => { it("shows units × price = cost for a fully priced counter", () => { - expect(counterMathLine({ counter: "contentPolicyUnits", units: 1000, unpriced: 0, cost: 0.15 })).toBe( - "Content Policy: 1,000 × $0.00015 = $0.1500", - ); + expect(counterMathRow({ counter: "contentPolicyUnits", units: 1000, unpriced: 0, cost: 0.15 })).toEqual({ + label: "Content Policy", + parts: ["1,000", "× $0.00015", "= $0.1500"], + note: null, + }); }); it("prices only the priced share and calls out the rest", () => { - expect(counterMathLine({ counter: "sensitiveInformationPolicyUnits", units: 8, unpriced: 2, cost: 0.0006 })).toBe( - "Sensitive Information Policy: 6 × $0.0001 = $0.0006 (2 unpriced left out)", + expect(counterMathRow({ counter: "sensitiveInformationPolicyUnits", units: 8, unpriced: 2, cost: 0.0006 })).toEqual( + { + label: "Sensitive Information Policy", + parts: ["6", "× $0.0001", "= $0.0006"], + note: "2 unpriced units left out", + }, ); + expect( + counterMathRow({ counter: "sensitiveInformationPolicyUnits", units: 8, unpriced: 1, cost: 0.0007 }).note, + ).toBe("1 unpriced unit left out"); }); it("says so when a counter has no known price at all", () => { - expect(counterMathLine({ counter: "someFutureCounter", units: 7, unpriced: 7, cost: null })).toBe( - "Some Future Counter: 7 units with no known price, left out", - ); - expect(counterMathLine({ counter: "someFutureCounter", units: 1, unpriced: 1, cost: null })).toBe( - "Some Future Counter: 1 unit with no known price, left out", - ); + expect(counterMathRow({ counter: "someFutureCounter", units: 7, unpriced: 7, cost: null })).toEqual({ + label: "Some Future Counter", + parts: ["7", "× —", "= —"], + note: "no known price, left out", + }); }); it("shows a free counter as × $0", () => { - expect(counterMathLine({ counter: "wordPolicyUnits", units: 2, unpriced: 0, cost: 0 })).toBe( - "Word Policy: 2 × $0 = $0.0000", - ); + expect(counterMathRow({ counter: "wordPolicyUnits", units: 2, unpriced: 0, cost: 0 }).parts).toEqual([ + "2", + "× $0", + "= $0.0000", + ]); }); }); -describe("unitsSumLine", () => { - it("adds the counters up in order", () => { - expect(unitsSumLine({ contentPolicyUnits: 2, topicPolicyUnits: 2, wordPolicyUnits: 1200 })).toBe( - "Content Policy 2 + Topic Policy 2 + Word Policy 1,200 = 1,204", - ); +describe("unitsMathRows", () => { + it("lists the counters in order with their counts", () => { + expect(unitsMathRows({ contentPolicyUnits: 2, wordPolicyUnits: 1200 })).toEqual([ + { label: "Content Policy", parts: ["2"], note: null }, + { label: "Word Policy", parts: ["1,200"], note: null }, + ]); }); }); diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.ts b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.ts index f914a3e9698..c47442de200 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.ts +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.ts @@ -40,20 +40,34 @@ export const formatUnitPrice = (price: number): string => { return price > 0 && Number(fixed) === 0 ? "< $0.000001" : `$${fixed}`; }; -export const counterMathLine = (row: CounterMath): string => { +export interface MathRow { + readonly label: string; + readonly parts: readonly string[]; + readonly note: string | null; +} + +export const counterMathRow = (row: CounterMath): MathRow => { const label = counterLabel(row.counter); const price = unitPrice(row); if (price == null) { - return `${label}: ${row.units.toLocaleString()} ${row.units === 1 ? "unit" : "units"} with no known price, left out`; + return { label, parts: [row.units.toLocaleString(), "× —", "= —"], note: "no known price, left out" }; } - const line = `${label}: ${pricedUnits(row).toLocaleString()} × ${formatUnitPrice(price)} = ${formatCost(row.cost)}`; - return row.unpriced > 0 ? `${line} (${row.unpriced.toLocaleString()} unpriced left out)` : line; + return { + label, + parts: [pricedUnits(row).toLocaleString(), `× ${formatUnitPrice(price)}`, `= ${formatCost(row.cost)}`], + note: + row.unpriced > 0 + ? `${row.unpriced.toLocaleString()} unpriced ${row.unpriced === 1 ? "unit" : "units"} left out` + : null, + }; }; -export const unitsSumLine = (units: UsageUnits): string => - `${Object.entries(units) - .map(([counter, n]) => `${counterLabel(counter)} ${n.toLocaleString()}`) - .join(" + ")} = ${totalUnits(units).toLocaleString()}`; +export const unitsMathRows = (units: UsageUnits): readonly MathRow[] => + Object.entries(units).map(([counter, n]) => ({ + label: counterLabel(counter), + parts: [n.toLocaleString()], + note: null, + })); export const pricingIssueUrl = (unpriced: UsageUnits, provider?: string): string => { const subject = provider ? `${provider} guardrail` : "guardrail"; From b98f8ee2c5be9bbaa1b54297f7cfdcb57ec4e615 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Sat, 5 Sep 2026 13:00:24 -0700 Subject: [PATCH 11/35] test(e2e): retry a fresh prefix when Vertex rejects the cache create on its minimum-token check --- tests/e2e/llm_translation/test_cache_control.py | 16 +++++++++++----- 1 file changed, 11 insertions(+), 5 deletions(-) diff --git a/tests/e2e/llm_translation/test_cache_control.py b/tests/e2e/llm_translation/test_cache_control.py index 3ad98bc6072..4e11ad6cec5 100644 --- a/tests/e2e/llm_translation/test_cache_control.py +++ b/tests/e2e/llm_translation/test_cache_control.py @@ -37,7 +37,7 @@ import pytest from pydantic import BaseModel from e2e_config import unique_marker -from e2e_http import Result, unwrap +from e2e_http import Result, UnknownApiError, unwrap from endpoints_client import CacheControl, RichMessage, TextBlock from lifecycle import ResourceManager from models import ChatBody, ChatMessage, ChatResponse, LiteLLMParamsBody, Usage @@ -54,6 +54,7 @@ VERTEX_CACHE_TTL: Final = "300s" VERTEX_COLD_CALL_ATTEMPTS: Final = 3 VERTEX_MINIMUM_CACHED_TOKENS: Final = 1024 CACHED_SHARE_OF_PROMPT: Final = 0.9 +VERTEX_CACHE_REJECTION_MARKER: Final = "minimum token count to start explicit caching" class CacheChatBody(BaseModel): @@ -149,7 +150,12 @@ def _assert_cache_read_on_second_call( def _cold_cache_calls(send: Callable[[str], Result[ChatResponse]]) -> Iterator[ChatResponse]: for _ in range(VERTEX_COLD_CALL_ATTEMPTS): - yield unwrap(send(_cacheable_prefix())) + result = send(_cacheable_prefix()) + match result: + case UnknownApiError(status_code=400, body=body) if VERTEX_CACHE_REJECTION_MARKER in body: + continue + case _: + yield unwrap(result) def _first_cold_call_reads_cache(model: str, send: Callable[[str], Result[ChatResponse]]) -> ChatResponse: @@ -162,9 +168,9 @@ def _first_cold_call_reads_cache(model: str, send: Callable[[str], Result[ChatRe None, ) assert completion is not None, ( - f"{model}: {VERTEX_COLD_CALL_ATTEMPTS} never-seen prompts marked with cache_control all reported fewer " - f"than {VERTEX_MINIMUM_CACHED_TOKENS} cached tokens on their first call; explicit context caching did " - "not engage" + f"{model}: {VERTEX_COLD_CALL_ATTEMPTS} never-seen prompts marked with cache_control were each either " + f"rejected by Vertex's minimum-token check or served with fewer than {VERTEX_MINIMUM_CACHED_TOKENS} " + "cached tokens on their first call; explicit context caching did not engage" ) assert completion.choices, f"{model}: cached call returned no choices: {completion}" usage: Final = completion.usage From 1c0172b477cfbc52c1a74ef55397a6531a999ced Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Sat, 5 Sep 2026 13:04:33 -0700 Subject: [PATCH 12/35] style(e2e): annotate the cold-call locals as Final --- tests/e2e/llm_translation/test_cache_control.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/e2e/llm_translation/test_cache_control.py b/tests/e2e/llm_translation/test_cache_control.py index 4e11ad6cec5..7530485ad79 100644 --- a/tests/e2e/llm_translation/test_cache_control.py +++ b/tests/e2e/llm_translation/test_cache_control.py @@ -150,7 +150,7 @@ def _assert_cache_read_on_second_call( def _cold_cache_calls(send: Callable[[str], Result[ChatResponse]]) -> Iterator[ChatResponse]: for _ in range(VERTEX_COLD_CALL_ATTEMPTS): - result = send(_cacheable_prefix()) + result: Final = send(_cacheable_prefix()) match result: case UnknownApiError(status_code=400, body=body) if VERTEX_CACHE_REJECTION_MARKER in body: continue @@ -241,7 +241,7 @@ class TestCacheControl: ) resources.defer(lambda: client.proxy.delete_model(model_id)) key = resources.key() - completion = _first_cold_call_reads_cache( + completion: Final = _first_cold_call_reads_cache( model, lambda prefix: _cache_chat(client, key, model, prefix, ttl=VERTEX_CACHE_TTL) ) _assert_billed_below_uncached_prompt(client, model, completion) From b56e4f80a4e97b60c47b4f177ff22b170ec9fa30 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Sat, 5 Sep 2026 13:05:03 -0700 Subject: [PATCH 13/35] test(e2e): make each cold cache call a single-assignment helper so its result stays Final --- .../e2e/llm_translation/test_cache_control.py | 21 +++++++++---------- 1 file changed, 10 insertions(+), 11 deletions(-) diff --git a/tests/e2e/llm_translation/test_cache_control.py b/tests/e2e/llm_translation/test_cache_control.py index 7530485ad79..a18e03c982b 100644 --- a/tests/e2e/llm_translation/test_cache_control.py +++ b/tests/e2e/llm_translation/test_cache_control.py @@ -30,7 +30,7 @@ built from the typed content blocks shared in ``endpoints_client.py``. from __future__ import annotations import time -from collections.abc import Callable, Iterator +from collections.abc import Callable from typing import Final import pytest @@ -148,22 +148,21 @@ def _assert_cache_read_on_second_call( ) -def _cold_cache_calls(send: Callable[[str], Result[ChatResponse]]) -> Iterator[ChatResponse]: - for _ in range(VERTEX_COLD_CALL_ATTEMPTS): - result: Final = send(_cacheable_prefix()) - match result: - case UnknownApiError(status_code=400, body=body) if VERTEX_CACHE_REJECTION_MARKER in body: - continue - case _: - yield unwrap(result) +def _cold_cache_call(send: Callable[[str], Result[ChatResponse]]) -> ChatResponse | None: + result: Final = send(_cacheable_prefix()) + match result: + case UnknownApiError(status_code=400, body=body) if VERTEX_CACHE_REJECTION_MARKER in body: + return None + case _: + return unwrap(result) def _first_cold_call_reads_cache(model: str, send: Callable[[str], Result[ChatResponse]]) -> ChatResponse: completion: Final = next( ( candidate - for candidate in _cold_cache_calls(send) - if _cached_read_tokens(candidate.usage) >= VERTEX_MINIMUM_CACHED_TOKENS + for candidate in (_cold_cache_call(send) for _ in range(VERTEX_COLD_CALL_ATTEMPTS)) + if candidate is not None and _cached_read_tokens(candidate.usage) >= VERTEX_MINIMUM_CACHED_TOKENS ), None, ) From 5df0e12e0f2628ed847c8110759a589f3fa1c138 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 5 Sep 2026 13:08:03 -0700 Subject: [PATCH 14/35] feat(guardrails): add non-blocking flag() verdict to custom code guardrails (#39728) Custom code guardrails could only allow(), block(reason) or modify(). This adds flag(reason, metadata={}) which lets the request or response through unchanged and records a guardrail_flagged entry carrying the guardrail name, configured mode, evaluated input_type (request or response), reason and structured metadata. The new status is threaded through the request-level guardrail_status aggregation, the Guardrails Monitor rollup (flagged_count), Request Logs (action=flagged, most severe phase wins when a guardrail runs pre and post call) and the Request Logs detail view in the dashboard, which now renders FLAGGED with warning styling instead of falling into FAILED. Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/litellm_core_utils/litellm_logging.py | 2 + .../custom_code/custom_code_guardrail.py | 27 ++++++ .../guardrail_hooks/custom_code/primitives.py | 28 +++++- litellm/proxy/guardrails/usage_endpoints.py | 18 ++-- litellm/proxy/guardrails/usage_tracking.py | 6 +- litellm/types/utils.py | 5 +- .../test_litellm_logging.py | 18 +++- .../guardrails/test_custom_code_security.py | 65 +++++++++++++ .../proxy/guardrails/test_usage_endpoints.py | 97 +++++++++++++++++++ .../proxy/guardrails/test_usage_tracking.py | 21 ++++ .../custom_code/CustomCodeModal.tsx | 1 + .../GuardrailViewer/GuardrailViewer.test.tsx | 16 +++ .../GuardrailViewer/GuardrailViewer.tsx | 95 ++++++++++++------ .../LogDetailContent.test.tsx | 16 ++- .../LogDetailsDrawer/LogDetailContent.tsx | 31 +++--- 15 files changed, 386 insertions(+), 60 deletions(-) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 83e0b4d84f1..c31c4323157 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -5881,6 +5881,7 @@ def _get_status_fields( # Mapping for legacy guardrail status values to new GuardrailStatus values GUARDRAIL_STATUS_MAP: Final[dict[str, GuardrailStatus]] = { "success": "success", + "guardrail_flagged": "guardrail_flagged", "blocked": "guardrail_intervened", # legacy "guardrail_intervened": "guardrail_intervened", # direct "failure": "guardrail_failed_to_respond", # legacy @@ -5902,6 +5903,7 @@ def _get_status_fields( GUARDRAIL_STATUS_SEVERITY: Final[tuple[GuardrailStatus, ...]] = ( "not_run", "success", + "guardrail_flagged", "guardrail_failed_to_respond", "guardrail_intervened", ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py index 830dec8d80d..d5ef1e949b8 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py @@ -36,6 +36,7 @@ Example: block when response rejects the user (input_type response only): import asyncio import threading +import time from collections.abc import Callable, Mapping from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast @@ -93,6 +94,7 @@ class CustomCodeGuardrail(CustomGuardrail): that returns one of: - allow() - let the request/response through - block(reason) - reject with a message + - flag(reason) - let it through but log a non-blocking violation - modify(texts=...) - transform the content Example: @@ -227,6 +229,7 @@ class CustomCodeGuardrail(CustomGuardrail): raise CustomCodeExecutionError(f"Custom code guardrail not compiled: {self._compile_error}") raise CustomCodeExecutionError("Custom code guardrail not compiled") + start_time: Final = time.time() try: # Prepare inputs dict for the function @@ -245,6 +248,7 @@ class CustomCodeGuardrail(CustomGuardrail): inputs=inputs, request_data=request_data, input_type=input_type, + start_time=start_time, ) except HTTPException: @@ -290,6 +294,7 @@ class CustomCodeGuardrail(CustomGuardrail): inputs: GenericGuardrailAPIInputs, request_data: dict[str, object], input_type: Literal["request", "response"], + start_time: float, ) -> GenericGuardrailAPIInputs: """ Process the result from the custom code function. @@ -299,6 +304,7 @@ class CustomCodeGuardrail(CustomGuardrail): inputs: The original inputs request_data: The request data input_type: "request" or "response" + start_time: Unix timestamp of when the guardrail started running, used for the flagged log entry Returns: GenericGuardrailAPIInputs - possibly modified @@ -348,6 +354,27 @@ class CustomCodeGuardrail(CustomGuardrail): }, ) + elif action == "flag": + flag_reason: Final = result.get("reason", "Flagged by custom code guardrail") + verbose_proxy_logger.info( + "Custom code guardrail '%s': Flagging %s - %s", self.guardrail_name, input_type, flag_reason + ) + end_time: Final = time.time() + self.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response={ # mutable-ok: logging helper requires a dict + "action": "flag", + "reason": flag_reason, + "input_type": input_type, + "metadata": result.get("metadata") or {}, # mutable-ok: logging helper requires a dict + }, + request_data=request_data, + guardrail_status="guardrail_flagged", + start_time=start_time, + end_time=end_time, + duration=end_time - start_time, + ) + return inputs + elif action == "modify": verbose_proxy_logger.debug("Custom code guardrail '%s': Modifying %s", self.guardrail_name, input_type) diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py index 24801aa2df1..d5dbfaeb84b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py @@ -8,7 +8,7 @@ and provide safe, sandboxed functionality for common guardrail operations. import json import re from collections.abc import Mapping, Sequence -from typing import Final +from typing import Final, Literal from urllib.parse import urlparse import httpx @@ -51,6 +51,31 @@ def block(reason: str, detection_info: Mapping[str, object] | None = None) -> di return result +class FlagResult(TypedDict): + action: ReadOnly[Literal["flag"]] + reason: ReadOnly[str] + metadata: ReadOnly[Mapping[str, object]] + + +def flag(reason: str, metadata: Mapping[str, object] | None = None) -> FlagResult: + """ + Let the request/response proceed unchanged but record a non-blocking violation. + + Args: + reason: Human-readable reason for flagging + metadata: Optional structured metadata stored alongside the reason + + Returns: + Dict indicating the request should be flagged but allowed + """ + result: Final[FlagResult] = { + "action": "flag", + "reason": reason, + "metadata": metadata if metadata is not None else {}, + } + return result + + def modify( texts: Sequence[str] | None = None, images: Sequence[object] | None = None, @@ -787,6 +812,7 @@ def get_custom_code_primitives() -> dict[str, object]: # Result types "allow": allow, "block": block, + "flag": flag, "modify": modify, # Regex "regex_match": regex_match, diff --git a/litellm/proxy/guardrails/usage_endpoints.py b/litellm/proxy/guardrails/usage_endpoints.py index 62145b9ede9..014ba3d1472 100644 --- a/litellm/proxy/guardrails/usage_endpoints.py +++ b/litellm/proxy/guardrails/usage_endpoints.py @@ -17,6 +17,7 @@ from typing_extensions import NotRequired, ReadOnly, TypedDict from litellm._logging import verbose_proxy_logger from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.guardrails.usage_tracking import guardrail_status_to_action from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import ( DailyGuardrailMetricsRepository, @@ -41,6 +42,7 @@ if TYPE_CHECKING: router: Final = APIRouter() _EMPTY_UNITS: Final[Mapping[str, int]] = MappingProxyType({}) +_ACTION_SEVERITY: Final[Mapping[str, int]] = MappingProxyType({"passed": 0, "flagged": 1, "blocked": 2}) _T = TypeVar("_T") @@ -759,21 +761,17 @@ def _usage_log_entry_from_row( except Exception: meta = {} guardrail_info_list: Final[Sequence[_GuardrailRunInfo]] = (meta or {}).get("guardrail_information") or [] - entry_for_guardrail: _GuardrailRunInfo | None = None - for gi in guardrail_info_list: - if (gi.get("guardrail_id") or gi.get("guardrail_name")) == r.guardrail_id: - entry_for_guardrail = gi - break + entry_for_guardrail: Final[_GuardrailRunInfo | None] = max( + (gi for gi in guardrail_info_list if (gi.get("guardrail_id") or gi.get("guardrail_name")) == r.guardrail_id), + key=lambda gi: _ACTION_SEVERITY[guardrail_status_to_action(gi.get("guardrail_status"))], + default=None, + ) action_val = "passed" score_val = None latency_val = None reason_val = None if entry_for_guardrail: - st: Final = (entry_for_guardrail.get("guardrail_status") or "").lower() - if "intervened" in st or "block" in st: - action_val = "blocked" - elif "fail" in st or "error" in st: - action_val = "flagged" + action_val = guardrail_status_to_action(entry_for_guardrail.get("guardrail_status")) duration: Final = entry_for_guardrail.get("duration") if duration is not None: latency_val = round(float(duration) * 1000, 0) diff --git a/litellm/proxy/guardrails/usage_tracking.py b/litellm/proxy/guardrails/usage_tracking.py index a20ad3935e5..df967058cf0 100644 --- a/litellm/proxy/guardrails/usage_tracking.py +++ b/litellm/proxy/guardrails/usage_tracking.py @@ -190,14 +190,14 @@ async def _upsert_rows_with_retry( return await _upsert_rows_with_retry(retryable, upsert_row, label, sleep, retries_left - 1) -def _guardrail_status_to_action(status: str | None) -> str: +def guardrail_status_to_action(status: str | None) -> str: """Map StandardLogging guardrail_status to blocked/passed/flagged.""" if not status: return "passed" s: Final = (status or "").lower() if "intervened" in s or "block" in s: return "blocked" - if "fail" in s or "error" in s: + if "flagged" in s or "fail" in s or "error" in s: return "flagged" return "passed" @@ -367,7 +367,7 @@ async def process_spend_logs_guardrail_usage( continue key = _MetricsKey(guardrail_id, date_key) daily_guardrail[key]["requests_evaluated"] += 1 - action = _guardrail_status_to_action(entry.get("guardrail_status")) + action = guardrail_status_to_action(entry.get("guardrail_status")) if action == "passed": daily_guardrail[key]["passed_count"] += 1 elif action == "blocked": diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 2118fe77aad..61c2fc8c5a5 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3078,7 +3078,9 @@ class GuardrailMode(TypedDict, total=False): default: str | list[str] | None -GuardrailStatus = Literal["success", "guardrail_intervened", "guardrail_failed_to_respond", "not_run"] +GuardrailStatus = Literal[ + "success", "guardrail_flagged", "guardrail_intervened", "guardrail_failed_to_respond", "not_run" +] # Fields on a guardrail record whose values can quote the caller's prompt: the payload sent to the # guardrail, the provider response that echoes it back, and the two first-party hooks that inline @@ -3320,6 +3322,7 @@ class StandardLoggingPayloadStatusFields(TypedDict, total=False): """ Status of guardrail execution: - 'success': Guardrail ran and allowed content through + - 'guardrail_flagged': Guardrail allowed content through but recorded a non-blocking violation - 'guardrail_intervened': Guardrail blocked or modified content - 'guardrail_failed_to_respond': Guardrail had technical failure - 'not_run': No guardrail was run diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 1991170707d..16a99713a06 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -16,7 +16,10 @@ from litellm._logging import session_id_var, trace_id_var from litellm.constants import SENTRY_DENYLIST, SENTRY_PII_DENYLIST from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging -from litellm.litellm_core_utils.litellm_logging import set_callbacks +from litellm.litellm_core_utils.litellm_logging import ( + _get_status_fields, + set_callbacks, +) from litellm.types.utils import ModelResponse, TextCompletionResponse @@ -6441,3 +6444,16 @@ def test_passthrough_embeddings_result_swapped_for_callbacks(): assert isinstance(swapped_result, EmbeddingResponse) assert swapped_result.data[0]["embedding"] == [0.1, 0.2, 0.3] + + +def test_get_status_fields_ranks_guardrail_flagged_between_success_and_intervened(): + """LIT-6894: a non-blocking flagged verdict must outrank success in the + request-level guardrail_status but never mask an intervention.""" + flagged = {"guardrail_status": "guardrail_flagged"} + + assert _get_status_fields( + "success", [{"guardrail_status": "success"}, flagged], None + )["guardrail_status"] == "guardrail_flagged" + assert _get_status_fields( + "success", [flagged, {"guardrail_status": "guardrail_intervened"}], None + )["guardrail_status"] == "guardrail_intervened" diff --git a/tests/test_litellm/proxy/guardrails/test_custom_code_security.py b/tests/test_litellm/proxy/guardrails/test_custom_code_security.py index f93ecfc3010..7971cf62c9a 100644 --- a/tests/test_litellm/proxy/guardrails/test_custom_code_security.py +++ b/tests/test_litellm/proxy/guardrails/test_custom_code_security.py @@ -197,6 +197,71 @@ async def test_custom_code_post_call_block_raises_http_400(): } +FLAG_CODE = ( + "def apply_guardrail(inputs, request_data, input_type):\n" + ' return flag("audit hit", metadata={"category": "topic"})\n' +) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("input_type", ["request", "response"]) +async def test_custom_code_flag_passes_content_through_and_records_flagged_entry(input_type): + """LIT-6894: flag() must not raise, must return the content unchanged and must log + exactly one guardrail_flagged entry (the decorator must not add a second "success").""" + guardrail = CustomCodeGuardrail(custom_code=FLAG_CODE, guardrail_name="t", event_hook=["pre_call", "post_call"]) + request_data = {"model": "test-model", "litellm_metadata": {}} + + result = await guardrail.apply_guardrail( + inputs={"texts": ["hello"]}, + request_data=request_data, + input_type=input_type, + ) + + assert result == {"texts": ["hello"]} + entries = request_data["litellm_metadata"]["standard_logging_guardrail_information"] + assert len(entries) == 1 + entry = entries[0] + assert entry["guardrail_status"] == "guardrail_flagged" + assert entry["guardrail_name"] == "t" + assert entry["guardrail_mode"] == ["pre_call", "post_call"] + assert entry["guardrail_response"] == { + "action": "flag", + "reason": "audit hit", + "input_type": input_type, + "metadata": {"category": "topic"}, + } + assert entry["duration"] is not None and entry["duration"] >= 0 + + +@pytest.mark.asyncio +async def test_custom_code_flag_default_reason_and_empty_metadata(): + code = "def apply_guardrail(inputs, request_data, input_type):\n return flag('just a note')\n" + guardrail = _compile(code) + request_data = {"model": "m", "litellm_metadata": {}} + + await guardrail.apply_guardrail(inputs={"texts": ["x"]}, request_data=request_data, input_type="request") + + entry = request_data["litellm_metadata"]["standard_logging_guardrail_information"][0] + assert entry["guardrail_response"] == { + "action": "flag", + "reason": "just a note", + "input_type": "request", + "metadata": {}, + } + + +@pytest.mark.asyncio +async def test_custom_code_allow_still_records_success_not_flagged(): + code = "def apply_guardrail(inputs, request_data, input_type):\n return allow()\n" + guardrail = _compile(code) + request_data = {"model": "m", "litellm_metadata": {}} + + await guardrail.apply_guardrail(inputs={"texts": ["x"]}, request_data=request_data, input_type="request") + + entries = request_data["litellm_metadata"]["standard_logging_guardrail_information"] + assert [e["guardrail_status"] for e in entries] == ["success"] + + def test_typical_sync_guardrail_still_works(): code = ( "def apply_guardrail(inputs, request_data, input_type):\n" diff --git a/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py b/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py index ebb2be6edc2..4e5a7ad4b2b 100644 --- a/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py +++ b/tests/test_litellm/proxy/guardrails/test_usage_endpoints.py @@ -477,6 +477,103 @@ async def test_logs_resolves_config_guardrail_logical_name(): assert where["guardrail_id"] == {"in": ["yaml-uuid", "yaml-pii"]} +def _index_row(request_id: str, guardrail_id: str = "cc-flag") -> Any: + r = MagicMock(spec=["request_id", "guardrail_id", "policy_id", "start_time"]) + r.request_id = request_id + r.guardrail_id = guardrail_id + return r + + +def _spend_log(request_id: str, *guardrail_statuses: str, guardrail_id: str = "cc-flag") -> Any: + sl = MagicMock(spec=["request_id", "metadata", "startTime", "model", "messages", "response"]) + sl.request_id = request_id + sl.startTime = datetime(2026, 4, 25, 12, 0) + sl.model = "gpt-4o-mini" + sl.messages = [{"role": "user", "content": "hi"}] + sl.response = "ok" + sl.metadata = { + "guardrail_information": [ + { + "guardrail_name": guardrail_id, + "guardrail_status": status, + "guardrail_response": ( + {"action": "flag", "reason": "audit hit"} if status == "guardrail_flagged" else "allow" + ), + "duration": 0.002, + } + for status in guardrail_statuses + ] + } + return sl + + +@pytest.mark.asyncio +async def test_logs_reports_flagged_action_for_guardrail_flagged_status(): + """LIT-6894: Request Logs surface a custom code flag() verdict as flagged with its reason.""" + prisma = _prisma(index_find_many=[_index_row("r-flag"), _index_row("r-pass"), _index_row("r-block")]) + prisma.db.litellm_spendlogs.find_many = AsyncMock( + return_value=[ + _spend_log("r-flag", "guardrail_flagged"), + _spend_log("r-pass", "success"), + _spend_log("r-block", "guardrail_intervened"), + ] + ) + p1, p2 = _patches(prisma, _config_handler()) + with p1, p2: + resp = await guardrails_usage_logs( + guardrail_id="cc-flag", + policy_id=None, + page=1, + page_size=50, + action=None, + start_date=START, + end_date=END, + user_api_key_dict=ADMIN, + ) + flagged_only = await guardrails_usage_logs( + guardrail_id="cc-flag", + policy_id=None, + page=1, + page_size=50, + action="flagged", + start_date=START, + end_date=END, + user_api_key_dict=ADMIN, + ) + assert [(log.id, log.action) for log in resp.logs] == [ + ("r-flag", "flagged"), + ("r-pass", "passed"), + ("r-block", "blocked"), + ] + assert resp.logs[0].reason == "{'action': 'flag', 'reason': 'audit hit'}" + assert [log.id for log in flagged_only.logs] == ["r-flag"] + + +@pytest.mark.asyncio +async def test_logs_reports_post_call_flag_when_pre_call_allowed(): + """LIT-6894: a guardrail on mode [pre_call, post_call] that allows the request but flags the response + shows as flagged, not hidden behind the pre_call allow entry.""" + prisma = _prisma(index_find_many=[_index_row("r-post-flag")]) + prisma.db.litellm_spendlogs.find_many = AsyncMock( + return_value=[_spend_log("r-post-flag", "success", "guardrail_flagged")] + ) + p1, p2 = _patches(prisma, _config_handler()) + with p1, p2: + resp = await guardrails_usage_logs( + guardrail_id="cc-flag", + policy_id=None, + page=1, + page_size=50, + action=None, + start_date=START, + end_date=END, + user_api_key_dict=ADMIN, + ) + assert [(log.id, log.action, log.reason) for log in resp.logs] == [ + ("r-post-flag", "flagged", "{'action': 'flag', 'reason': 'audit hit'}") + ] + + # ---- date window cap (LIT-5762) --------------------------------------------- diff --git a/tests/test_litellm/proxy/guardrails/test_usage_tracking.py b/tests/test_litellm/proxy/guardrails/test_usage_tracking.py index ae360b281cb..110de7dbe70 100644 --- a/tests/test_litellm/proxy/guardrails/test_usage_tracking.py +++ b/tests/test_litellm/proxy/guardrails/test_usage_tracking.py @@ -105,6 +105,27 @@ async def test_usage_units_rolled_up_by_guardrail_team_key_and_date(): } +@pytest.mark.asyncio +async def test_flagged_status_counts_as_flagged_not_passed_or_blocked(): + """LIT-6894: a custom code flag() verdict lands in flagged_count on the Monitor rollup.""" + prisma = _prisma() + logs = [ + _payload("r1", guardrail_status="success"), + _payload("r2", guardrail_status="guardrail_flagged"), + _payload("r3", guardrail_status="guardrail_intervened"), + ] + + await process_spend_logs_guardrail_usage(prisma, logs) + + create = prisma.db.litellm_dailyguardrailmetrics.upsert.call_args.kwargs["data"]["create"] + assert (create["requests_evaluated"], create["passed_count"], create["flagged_count"], create["blocked_count"]) == ( + 3, + 1, + 1, + 1, + ) + + def _fake_sleep() -> tuple[AsyncMock, list[float]]: delays: list[float] = [] sleep = AsyncMock(side_effect=lambda delay: delays.append(delay)) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx index a69824f32d3..05a48598859 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx @@ -112,6 +112,7 @@ const PRIMITIVES = { "Return Values": [ { name: "allow()", desc: "Let request/response through" }, { name: "block(reason)", desc: "Reject with message" }, + { name: "flag(reason, metadata={})", desc: "Let through, record a non-blocking violation" }, { name: "modify(texts=[], images=[], tool_calls=[])", desc: "Transform content" }, ], "HTTP Requests (async)": [ diff --git a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.test.tsx b/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.test.tsx index b5e04c72440..aabac50a661 100644 --- a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.test.tsx @@ -33,6 +33,22 @@ describe("GuardrailViewer", () => { expect(screen.getByText("1235ms")).toBeInTheDocument(); }); + it("renders guardrail_flagged as FLAGGED (warning), not FAILED", () => { + const data = makeGuardrailInformation({ + guardrail_name: "cc-flag", + guardrail_status: "guardrail_flagged", + guardrail_provider: "custom_code", + }); + renderWithProviders(); + + expect(screen.getByText(/0 Passed/)).toBeInTheDocument(); + expect(screen.getByText(/1 Flagged/)).toBeInTheDocument(); + const badges = screen.getAllByText("FLAGGED"); + expect(badges.length).toBeGreaterThan(0); + expect(badges[0]).toHaveClass("text-warning"); + expect(screen.queryByText("FAILED")).not.toBeInTheDocument(); + }); + it("calculates and displays masked entity totals", async () => { const user = userEvent.setup(); const data = makeGuardrailInformation({ diff --git a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx b/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx index 863f4117510..271b8f6ce05 100644 --- a/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx @@ -133,8 +133,27 @@ const getTotalMasked = (entry: GuardrailInformation): number => { ); }; -const isEntrySuccess = (entry: GuardrailInformation): boolean => { - return (entry.guardrail_status ?? "").toLowerCase() === "success"; +type EntryOutcome = "passed" | "flagged" | "failed"; + +const getEntryOutcome = (entry: GuardrailInformation): EntryOutcome => { + const status = (entry.guardrail_status ?? "").toLowerCase(); + if (status === "success") return "passed"; + if (status === "guardrail_flagged") return "flagged"; + return "failed"; +}; + +const isEntrySuccess = (entry: GuardrailInformation): boolean => getEntryOutcome(entry) === "passed"; + +const OUTCOME_LABEL: Record = { + passed: "PASSED", + flagged: "FLAGGED", + failed: "FAILED", +}; + +const OUTCOME_BADGE_CLASS: Record = { + passed: "bg-success/15 text-success border border-success/20", + flagged: "bg-warning/15 text-warning border border-warning/20", + failed: "bg-destructive/15 text-destructive border border-destructive/20", }; const getRiskColor = (score: number): string => { @@ -202,6 +221,19 @@ const FailCircleIcon = ({ className }: { className?: string }) => ( ); +const FlagCircleIcon = ({ className }: { className?: string }) => ( + + + + +); + +const OutcomeIcon = ({ outcome }: { outcome: EntryOutcome }) => { + if (outcome === "passed") return ; + if (outcome === "flagged") return ; + return ; +}; + const PlayCircleIcon = () => ( @@ -318,8 +350,7 @@ interface TimelineEntry { type: "request" | "guardrail" | "llm" | "response"; label: string; offsetMs: number; - status?: string; - isSuccess?: boolean; + outcome?: EntryOutcome; } const RequestLifecycle = ({ entries }: { entries: GuardrailInformation[] }) => { @@ -348,8 +379,7 @@ const RequestLifecycle = ({ entries }: { entries: GuardrailInformation[] }) => { type: "guardrail", label: `Pre-call guardrail: ${getDisplayName(e)}`, offsetMs, - status: isEntrySuccess(e) ? "PASSED" : "FAILED", - isSuccess: isEntrySuccess(e), + outcome: getEntryOutcome(e), }); } @@ -372,8 +402,7 @@ const RequestLifecycle = ({ entries }: { entries: GuardrailInformation[] }) => { type: "guardrail", label: `During-call guardrail: ${getDisplayName(e)}`, offsetMs, - status: isEntrySuccess(e) ? "PASSED" : "FAILED", - isSuccess: isEntrySuccess(e), + outcome: getEntryOutcome(e), }); } @@ -384,8 +413,7 @@ const RequestLifecycle = ({ entries }: { entries: GuardrailInformation[] }) => { type: "guardrail", label: `Post-call guardrail: ${getDisplayName(e)}`, offsetMs, - status: isEntrySuccess(e) ? "PASSED" : "FAILED", - isSuccess: isEntrySuccess(e), + outcome: getEntryOutcome(e), }); } @@ -410,10 +438,8 @@ const RequestLifecycle = ({ entries }: { entries: GuardrailInformation[] }) => { ) : item.type === "llm" ? ( - ) : item.isSuccess ? ( - ) : ( - + )} {idx < timeline.length - 1 &&
    } @@ -425,13 +451,11 @@ const RequestLifecycle = ({ entries }: { entries: GuardrailInformation[] }) => { {item.label} - {item.status && ( + {item.outcome && ( - {item.status} + {OUTCOME_LABEL[item.outcome]} )} T+{item.offsetMs}ms @@ -455,7 +479,7 @@ const formatGuardrailCost = (cost: number): string => { const EvaluationCard = ({ entry }: { entry: GuardrailInformation }) => { const [expanded, setExpanded] = useState(false); - const success = isEntrySuccess(entry); + const outcome = getEntryOutcome(entry); const totalMasked = getTotalMasked(entry); const displayName = getDisplayName(entry); const durationStr = formatDurationMs(entry.duration); @@ -490,7 +514,9 @@ const EvaluationCard = ({ entry }: { entry: GuardrailInformation }) => { onClick={() => setExpanded(!expanded)} > {/* Status icon */} -
    {success ? : }
    +
    + +
    {/* Name + badges */}
    @@ -501,13 +527,9 @@ const EvaluationCard = ({ entry }: { entry: GuardrailInformation }) => { - {success ? "PASSED" : "FAILED"} + {OUTCOME_LABEL[outcome]} {matchCountStr && ( @@ -528,7 +550,7 @@ const EvaluationCard = ({ entry }: { entry: GuardrailInformation }) => { )} - {riskScore != null && success && ( + {riskScore != null && outcome === "passed" && ( getEntryOutcome(e) === "flagged").length; const allPassed = passedCount === guardrailEntries.length; + const headerOutcome: EntryOutcome = allPassed + ? "passed" + : passedCount + flaggedCount === guardrailEntries.length + ? "flagged" + : "failed"; const totalOverheadMs = useMemo(() => { return Math.round(guardrailEntries.reduce((sum, e) => sum + (e.duration ?? 0), 0) * 1000); @@ -709,11 +737,7 @@ const GuardrailViewer = ({ data, accessToken, logEntry }: GuardrailViewerProps) | {allPassed ? ( @@ -728,6 +752,13 @@ const GuardrailViewer = ({ data, accessToken, logEntry }: GuardrailViewerProps) ) : null} {passedCount} Passed + {flaggedCount > 0 && ( + + {flaggedCount} Flagged + + )}
    diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.test.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.test.tsx index a679dc49427..2e9bce5048f 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.test.tsx @@ -1,7 +1,7 @@ import { render, screen, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { describe, expect, it, vi } from "vitest"; -import { LogDetailContent } from "./LogDetailContent"; +import { GuardrailJumpLink, LogDetailContent } from "./LogDetailContent"; import type { LogEntry } from "../columns"; vi.mock("../GuardrailViewer/GuardrailViewer", () => ({ @@ -489,3 +489,17 @@ describe("LogDetailContent", () => { expect(within(descriptions).getByText("-")).toBeInTheDocument(); }); }); + +describe("GuardrailJumpLink", () => { + it.each([ + [["success", "success"], "text-success", "\u2713"], + [["success", "guardrail_flagged"], "text-warning", "\u26A0"], + [["guardrail_flagged", "guardrail_intervened"], "text-destructive", "\u2717"], + ])("styles %j as %s", (statuses, expectedClass, glyph) => { + render( ({ guardrail_status: s }))} />); + + const pill = screen.getByText(/2 guardrails evaluated/); + expect(pill).toHaveClass(expectedClass); + expect(pill).toHaveTextContent(glyph); + }); +}); diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.tsx index 4c5c7b7b43f..052f1ec8802 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.tsx @@ -635,11 +635,24 @@ function RequestResponseSection({ ); } +const GUARDRAIL_JUMP_LINK_STYLE = { + passed: { className: "border border-success/20 bg-success/10 text-success", glyph: "\u2713" }, + flagged: { className: "border border-warning/20 bg-warning/10 text-warning", glyph: "\u26A0" }, + failed: { className: "border border-destructive/20 bg-destructive/10 text-destructive", glyph: "\u2717" }, +} as const; + +const isPassedStatus = (status: unknown) => status === "pass" || status === "passed" || status === "success"; +const isFlaggedStatus = (status: unknown) => status === "flagged" || status === "guardrail_flagged"; + +const guardrailJumpLinkOutcome = (statuses: unknown[]): keyof typeof GUARDRAIL_JUMP_LINK_STYLE => { + if (statuses.every(isPassedStatus)) return "passed"; + if (statuses.every((s) => isPassedStatus(s) || isFlaggedStatus(s))) return "flagged"; + return "failed"; +}; + export function GuardrailJumpLink({ guardrailEntries }: { guardrailEntries: any[] }) { - const allPassed = guardrailEntries.every((e) => { - const status = e?.guardrail_status || e?.status; - return status === "pass" || status === "passed" || status === "success"; - }); + const outcome = guardrailJumpLinkOutcome(guardrailEntries.map((e) => e?.guardrail_status || e?.status)); + const { className, glyph } = GUARDRAIL_JUMP_LINK_STYLE[outcome]; const handleClick = () => { const el = document.getElementById("guardrail-section"); @@ -650,11 +663,7 @@ export function GuardrailJumpLink({ guardrailEntries }: { guardrailEntries: any[
    - {allPassed ? "\u2713" : "\u2717"} {guardrailEntries.length} guardrail{guardrailEntries.length !== 1 ? "s" : ""}{" "} - evaluated + {glyph} {guardrailEntries.length} guardrail + {guardrailEntries.length !== 1 ? "s" : ""} evaluated {"\u2193"}
    From d6bc8fe289271457b733190f8abd9c261599b009 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Sat, 5 Sep 2026 13:23:38 -0700 Subject: [PATCH 15/35] test(e2e): ask the streamed /v1/messages pin for a reply long enough to span several deltas Anthropic now returns the 64-token 'count to 20' reply in one to three content_block_delta events, measured directly against api.anthropic.com and through proxies at 7672399 and 49a1145 alike, so the incrementality assertion (at least two deltas) failed in litellm-e2e builds 125, 130 and the 278 rerun with no proxy change behind it. A 'count to 100' reply at max_tokens 400 arrived in five to fifty deltas across every measured run --- tests/e2e/llm_translation/test_messages_e2e.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/e2e/llm_translation/test_messages_e2e.py b/tests/e2e/llm_translation/test_messages_e2e.py index e0bedd72eac..a5f36a8cbdd 100644 --- a/tests/e2e/llm_translation/test_messages_e2e.py +++ b/tests/e2e/llm_translation/test_messages_e2e.py @@ -169,9 +169,9 @@ class TestAnthropicMessages: key, AnthropicMessagesBody( model=model, - max_tokens=64, + max_tokens=400, stream=True, - messages=[ChatMessage(role="user", content="Count from 1 to 20, one number per line.")], + messages=[ChatMessage(role="user", content="Count from 1 to 100, one number per line.")], ), ) require_successful_call(result) From 80839bb33c318851af40b125c54c262bbe5dc90f Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 5 Sep 2026 13:26:09 -0700 Subject: [PATCH 16/35] feat(proxy): serve Prometheus /metrics from a separate process via --prometheus_metrics_port (#39889) * feat(proxy): serve Prometheus /metrics from a separate process via --prometheus_metrics_port Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(proxy): ruff format prometheus_metrics_server Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): fail fast when the separate metrics server cannot start and force the multiproc dir whenever it is enabled - wait for the child's /health before starting uvicorn; raise a ClickException if it exits first (port in use) - create PROMETHEUS_MULTIPROC_DIR whenever --prometheus_metrics_port is set, so DB-configured prometheus callbacks work - honour lowercase prometheus_multiproc_dir; validate the port before spawning - cover main() entry point, readiness, bind failure and wildcard-host probing in tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): pin metrics-server readiness to the child pid so another service on the port cannot pass the health check Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(proxy): probe metrics-server readiness through the shared HTTPHandler instead of bare httpx.get Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(proxy): serve only /metrics on the prometheus metrics port Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): validate metrics server CLI args with pydantic instead of typing.cast Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): satisfy metrics server lint gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- basedpyright-code-budget.json | 4 +- litellm/proxy/prometheus_metrics_server.py | 167 +++++++++++ litellm/proxy/proxy_cli.py | 90 ++++-- .../proxy/test_prometheus_cleanup.py | 40 +++ .../proxy/test_prometheus_metrics_server.py | 259 ++++++++++++++++++ tests/test_litellm/proxy/test_proxy_cli.py | 118 ++++++++ type-discipline-budget.json | 4 +- 7 files changed, 649 insertions(+), 33 deletions(-) create mode 100644 litellm/proxy/prometheus_metrics_server.py create mode 100644 tests/test_litellm/proxy/test_prometheus_metrics_server.py diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index b876cf2d69a..0b0a61192e6 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -105,13 +105,13 @@ "limit": 109 }, "reportUnknownMemberType": { - "limit": 38283 + "limit": 38271 }, "reportUnknownParameterType": { "limit": 19584 }, "reportUnknownVariableType": { - "limit": 29829 + "limit": 29814 }, "reportUnnecessaryCast": { "limit": 110 diff --git a/litellm/proxy/prometheus_metrics_server.py b/litellm/proxy/prometheus_metrics_server.py new file mode 100644 index 00000000000..4a9651d62e1 --- /dev/null +++ b/litellm/proxy/prometheus_metrics_server.py @@ -0,0 +1,167 @@ +"""Serve Prometheus `/metrics` from its own process so a scrape never runs on an inference worker. + +Workers write their samples to `PROMETHEUS_MULTIPROC_DIR`; this process reads them back with a +``MultiProcessCollector`` and serves the aggregated output on a separate port. The proxy CLI starts +it with ``--prometheus_metrics_port``. It can also run as a sidecar sharing the same directory: +``python -m litellm.proxy.prometheus_metrics_server --host 0.0.0.0 --port 4001``. +""" + +from __future__ import annotations + +import argparse +import atexit +import os +import subprocess +import sys +import threading +import time +from collections.abc import Sequence +from contextlib import closing +from types import MappingProxyType +from typing import Final + +import httpx +from fastapi import FastAPI +from prometheus_client import CollectorRegistry, multiprocess +from pydantic import BaseModel, ConfigDict +from starlette.types import ASGIApp, Message, Receive, Scope, Send + +from litellm.integrations.prometheus_metrics_endpoint import make_metrics_asgi_app +from litellm.llms.custom_httpx.http_handler import HTTPHandler + +METRICS_PATH: Final = "/metrics" +PID_HEADER: Final = "x-litellm-metrics-pid" +_PARENT_POLL_INTERVAL_SECONDS: Final = 1.0 +_STARTUP_TIMEOUT_SECONDS: Final = 30.0 +_STARTUP_POLL_INTERVAL_SECONDS: Final = 0.1 +_STARTUP_PROBE_TIMEOUT_SECONDS: Final = 1.0 +_WILDCARD_TO_LOOPBACK: Final = MappingProxyType({"0.0.0.0": "127.0.0.1", "::": "::1"}) + + +class _CliArgs(BaseModel): + model_config = ConfigDict(frozen=True) + + host: str + port: int + multiproc_dir: str | None + + +class MetricsServerStartupError(RuntimeError): + """The metrics process died or never answered on its port before the proxy started serving.""" + + +def _add_pid_header(app: ASGIApp) -> ASGIApp: + async def app_with_pid(scope: Scope, receive: Receive, send: Send) -> None: + async def send_with_pid(message: Message) -> None: + if message["type"] == "http.response.start": + await send( + { + **message, + "headers": [ + *message["headers"], + (PID_HEADER.encode(), str(os.getpid()).encode()), + ], + } + ) + return + await send(message) + + await app(scope, receive, send_with_pid) + + return app_with_pid + + +def build_metrics_app(multiproc_dir: str) -> FastAPI: + registry: Final = CollectorRegistry() + multiprocess.MultiProcessCollector(registry, path=multiproc_dir) + app: Final = FastAPI(title="LiteLLM Prometheus metrics", docs_url=None, redoc_url=None, openapi_url=None) + app.mount(METRICS_PATH, _add_pid_header(make_metrics_asgi_app(registry))) + + return app + + +def _exit_when_parent_dies(parent_pid: int) -> None: + def watch() -> None: + while os.getppid() == parent_pid: + time.sleep(_PARENT_POLL_INTERVAL_SECONDS) + os._exit(0) + + threading.Thread(target=watch, name="litellm-metrics-parent-watchdog", daemon=True).start() + + +def run_metrics_server(host: str, port: int, multiproc_dir: str) -> None: + import uvicorn + + _exit_when_parent_dies(os.getppid()) + uvicorn.run(build_metrics_app(multiproc_dir), host=host, port=port, log_level="warning", access_log=False) + + +def metrics_url(host: str, port: int) -> str: + probe_host: Final = _WILDCARD_TO_LOOPBACK.get(host, host) + netloc: Final = f"[{probe_host}]" if ":" in probe_host else probe_host + return f"http://{netloc}:{port}{METRICS_PATH}" + + +def _answered_by(http: HTTPHandler, url: str, pid: int) -> bool: + """True only when the metrics response comes from our child, not from whatever else holds the port.""" + try: + response: Final = http.get(url) # pyright: ignore[reportUnknownMemberType] # HTTPHandler.get exposes untyped optional mappings + return response.status_code == 200 and response.headers.get(PID_HEADER) == str(pid) + except httpx.TransportError: + return False + + +def _wait_until_serving(process: subprocess.Popen[bytes], host: str, port: int) -> None: + url: Final = metrics_url(host, port) + deadline: Final = time.monotonic() + _STARTUP_TIMEOUT_SECONDS + with closing(HTTPHandler(timeout=_STARTUP_PROBE_TIMEOUT_SECONDS)) as http: + while time.monotonic() < deadline: + if (returncode := process.poll()) is not None: + raise MetricsServerStartupError( + f"Prometheus metrics server exited with code {returncode} before serving {host}:{port}; " + "is the port already in use?" + ) + if _answered_by(http, url, process.pid): + return + time.sleep(_STARTUP_POLL_INTERVAL_SECONDS) + process.terminate() + raise MetricsServerStartupError( + f"Prometheus metrics server did not answer {url} within {_STARTUP_TIMEOUT_SECONDS:.0f}s" + ) + + +def start_metrics_server_process(host: str, port: int, multiproc_dir: str) -> subprocess.Popen[bytes]: + """Spawn the metrics server next to the proxy and block until it answers on its port.""" + process: Final = subprocess.Popen( + ( + sys.executable, + "-m", + "litellm.proxy.prometheus_metrics_server", + "--host", + host, + "--port", + str(port), + "--multiproc_dir", + multiproc_dir, + ) + ) + atexit.register(process.terminate) + _wait_until_serving(process, host, port) + return process + + +def main(argv: Sequence[str] | None = None) -> None: + parser: Final = argparse.ArgumentParser( + description="Serve LiteLLM Prometheus metrics from PROMETHEUS_MULTIPROC_DIR" + ) + parser.add_argument("--host", default="0.0.0.0") + parser.add_argument("--port", type=int, required=True) + parser.add_argument("--multiproc_dir", default=os.environ.get("PROMETHEUS_MULTIPROC_DIR")) + args: Final = _CliArgs.model_validate(vars(parser.parse_args(argv))) + if not args.multiproc_dir: + parser.error("--multiproc_dir or PROMETHEUS_MULTIPROC_DIR is required") + run_metrics_server(host=args.host, port=args.port, multiproc_dir=args.multiproc_dir) + + +if __name__ == "__main__": + main() diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index e780beb4410..e245367b1b4 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -7,7 +7,7 @@ import re import subprocess import sys import urllib.parse as urlparse -from collections.abc import Iterable +from collections.abc import Iterable, Mapping, Sequence from pathlib import Path from typing import TYPE_CHECKING, Any, Final @@ -610,48 +610,49 @@ class ProxyInitializationHelpers: return None # Let uvicorn choose the default loop on Windows return "uvloop" + @staticmethod + def _prometheus_callback_configured(litellm_settings: Mapping[str, object] | None) -> bool: + if litellm_settings is None: + return False + configured: Final = tuple( + litellm_settings.get(key) for key in ("callbacks", "success_callback", "failure_callback") + ) + return any( + setting == "prometheus" + if isinstance(setting, str) + else isinstance(setting, Sequence) and "prometheus" in setting + for setting in configured + ) + @staticmethod def _maybe_setup_prometheus_multiproc_dir( num_workers: int, litellm_settings: dict | None, - ) -> None: + prometheus_metrics_port: int | None = None, + ) -> str | None: """ - Auto-create PROMETHEUS_MULTIPROC_DIR when running with multiple workers - and prometheus is configured as a callback. + Auto-create PROMETHEUS_MULTIPROC_DIR when another process needs to read the samples: extra workers + with prometheus configured as a callback in config.yaml, or the separate metrics server (always, since + callbacks may also be enabled from the DB after startup). """ import tempfile - if num_workers <= 1 or litellm_settings is None: - return - - # Check if prometheus is in any callback list - # Each setting can be a list or a single string; normalize to list - callbacks = litellm_settings.get("callbacks") or [] - success_callbacks = litellm_settings.get("success_callback") or [] - failure_callbacks = litellm_settings.get("failure_callback") or [] - if isinstance(callbacks, str): - callbacks = [callbacks] - if isinstance(success_callbacks, str): - success_callbacks = [success_callbacks] - if isinstance(failure_callbacks, str): - failure_callbacks = [failure_callbacks] - all_callbacks: Final = callbacks + success_callbacks + failure_callbacks - if "prometheus" not in all_callbacks: - return + if prometheus_metrics_port is None and ( + num_workers <= 1 or not ProxyInitializationHelpers._prometheus_callback_configured(litellm_settings) + ): + return None from litellm.proxy.prometheus_cleanup import wipe_directory - multiproc_dir = os.environ.get("PROMETHEUS_MULTIPROC_DIR") or os.environ.get("prometheus_multiproc_dir") - - auto_created: Final = not multiproc_dir - if not multiproc_dir: - multiproc_dir = os.path.join(tempfile.gettempdir(), "litellm_prometheus_multiproc") - os.environ["PROMETHEUS_MULTIPROC_DIR"] = multiproc_dir + configured_dir: Final = os.environ.get("PROMETHEUS_MULTIPROC_DIR") or os.environ.get("prometheus_multiproc_dir") + multiproc_dir: Final = configured_dir or os.path.join(tempfile.gettempdir(), "litellm_prometheus_multiproc") + os.environ["PROMETHEUS_MULTIPROC_DIR"] = multiproc_dir os.makedirs(multiproc_dir, exist_ok=True) wipe_directory(multiproc_dir) - action: Final = "Auto-created" if auto_created else "Using existing" + action: Final = "Using existing" if configured_dir else "Auto-created" print(f"LiteLLM: {action} PROMETHEUS_MULTIPROC_DIR={multiproc_dir}") + return multiproc_dir @click.command() @@ -930,6 +931,19 @@ class ProxyInitializationHelpers: default=False, help="Enable uvicorn hot reload (dev only). Also reloads when the --config YAML file changes. Incompatible with --num_workers>1, --run_gunicorn, and --run_hypercorn.", ) +@click.option( + "--prometheus_metrics_port", + default=None, + type=click.IntRange(min=1, max=65535), + help=( + "Serve Prometheus /metrics from a separate process on this port (bound to --host) so scraping and " + "multi-worker aggregation never run on an inference worker's event loop. Samples appear once the " + "`prometheus` callback is enabled (config.yaml or DB). /metrics stays mounted on the main port as well; " + "the separate port has no virtual-key auth, so keep it off public ingress. Startup fails if the metrics " + "server cannot bind." + ), + envvar="PROMETHEUS_METRICS_PORT", +) def run_server( cli_args, host, @@ -980,6 +994,7 @@ def run_server( enforce_prisma_migration_check: bool, use_v2_migration_resolver: bool, reload: bool, + prometheus_metrics_port: int | None, ): if cli_args: if cli_args == ("xai-oauth", "login"): @@ -1364,6 +1379,8 @@ def run_server( ) if port == 4000 and ProxyInitializationHelpers._is_port_in_use(port): port = random.randint(1024, 49152) + if prometheus_metrics_port == port: + raise click.UsageError("--prometheus_metrics_port must differ from --port") import litellm @@ -1374,9 +1391,10 @@ def run_server( from litellm.proxy.proxy_server import app # Auto-create PROMETHEUS_MULTIPROC_DIR for multi-worker setups - ProxyInitializationHelpers._maybe_setup_prometheus_multiproc_dir( + prometheus_multiproc_dir: Final = ProxyInitializationHelpers._maybe_setup_prometheus_multiproc_dir( num_workers=num_workers, litellm_settings=litellm_settings if config else None, + prometheus_metrics_port=prometheus_metrics_port, ) # Skip server startup if requested (after all setup is done) @@ -1384,6 +1402,20 @@ def run_server( print("LiteLLM: Setup complete. Skipping server startup as requested.") return + if prometheus_metrics_port is not None and prometheus_multiproc_dir is not None: + from litellm.proxy.prometheus_metrics_server import MetricsServerStartupError, start_metrics_server_process + + try: + metrics_process: Final = start_metrics_server_process( + host=host, port=prometheus_metrics_port, multiproc_dir=prometheus_multiproc_dir + ) + except MetricsServerStartupError as error: + raise click.ClickException(str(error)) from error + print( + f"\033[1;32mLiteLLM: Serving Prometheus metrics on {host}:{prometheus_metrics_port}/metrics " + f"(pid {metrics_process.pid})\033[0m" + ) + running_uvicorn: Final = run_gunicorn is False and run_hypercorn is False uvicorn_args: Final = ProxyInitializationHelpers._get_default_unvicorn_init_args( host=host, diff --git a/tests/test_litellm/proxy/test_prometheus_cleanup.py b/tests/test_litellm/proxy/test_prometheus_cleanup.py index ca5476d6af9..93b9b694c2c 100644 --- a/tests/test_litellm/proxy/test_prometheus_cleanup.py +++ b/tests/test_litellm/proxy/test_prometheus_cleanup.py @@ -131,3 +131,43 @@ class TestMaybeSetupPrometheusMultiprocDir: # Cleanup os.environ.pop("PROMETHEUS_MULTIPROC_DIR", None) + + @pytest.mark.parametrize( + "litellm_settings", + [ + {"callbacks": ["prometheus"]}, + {"callbacks": ["langfuse"]}, + None, + ], + ) + def test_separate_metrics_port_forces_dir_for_single_worker(self, litellm_settings): + """The separate metrics process reads the samples, so one worker still needs the shared dir, even when + prometheus is not in config.yaml (callbacks can be turned on from the DB after startup).""" + with patch.dict(os.environ, {}, clear=False): + os.environ.pop("PROMETHEUS_MULTIPROC_DIR", None) + os.environ.pop("prometheus_multiproc_dir", None) + + result_dir = ProxyInitializationHelpers._maybe_setup_prometheus_multiproc_dir( + num_workers=1, + litellm_settings=litellm_settings, + prometheus_metrics_port=4001, + ) + + assert result_dir is not None + assert os.environ.get("PROMETHEUS_MULTIPROC_DIR") == result_dir + assert os.path.isdir(result_dir) + + os.environ.pop("PROMETHEUS_MULTIPROC_DIR", None) + + def test_lowercase_env_var_is_reused_and_exported_uppercase(self, tmp_path): + """prometheus_client honours both spellings; the metrics server only reads the uppercase one.""" + with patch.dict(os.environ, {"prometheus_multiproc_dir": str(tmp_path)}, clear=False): + os.environ.pop("PROMETHEUS_MULTIPROC_DIR", None) + + result_dir = ProxyInitializationHelpers._maybe_setup_prometheus_multiproc_dir( + num_workers=4, + litellm_settings={"callbacks": "prometheus"}, + ) + + assert result_dir == str(tmp_path) + assert os.environ["PROMETHEUS_MULTIPROC_DIR"] == str(tmp_path) diff --git a/tests/test_litellm/proxy/test_prometheus_metrics_server.py b/tests/test_litellm/proxy/test_prometheus_metrics_server.py new file mode 100644 index 00000000000..fc1fa381fa4 --- /dev/null +++ b/tests/test_litellm/proxy/test_prometheus_metrics_server.py @@ -0,0 +1,259 @@ +"""The separate metrics server must aggregate PROMETHEUS_MULTIPROC_DIR, expose only /metrics, and follow its +parent's lifetime. + +Everything here runs on loopback against a child of this test process; no LLM keys or external network. +""" + +from __future__ import annotations + +import os +import socket +import subprocess +import sys +import threading +import time +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from pathlib import Path +from typing import Final +from unittest.mock import patch + +import httpx +import pytest +from fastapi.testclient import TestClient +from prometheus_client import values + +from litellm.proxy.prometheus_metrics_server import ( + PID_HEADER, + MetricsServerStartupError, + build_metrics_app, + main, + metrics_url, + start_metrics_server_process, +) + +_STARTUP_TIMEOUT_SECONDS: Final = 60.0 +_SHUTDOWN_TIMEOUT_SECONDS: Final = 15.0 + + +def _write_worker_sample(pid: int, value: float) -> None: + """Write one counter sample into PROMETHEUS_MULTIPROC_DIR the way a proxy worker would.""" + counter: Final = values.MultiProcessValue(process_identifier=lambda: pid)( + "counter", + "litellm_requests_metric_total", + "litellm_requests_metric_total", + ("model",), + ("gpt-5",), + "Total number of LLM calls", + ) + counter.inc(value) + + +def _free_port() -> int: + with socket.socket() as sock: + sock.bind(("127.0.0.1", 0)) + return sock.getsockname()[1] + + +def _wait_for_metrics(port: int, pid: int) -> httpx.Response: + deadline: Final = time.monotonic() + _STARTUP_TIMEOUT_SECONDS + while time.monotonic() < deadline: + try: + response: Final = httpx.get(f"http://127.0.0.1:{port}/metrics", follow_redirects=True, timeout=1.0) + if response.status_code == 200 and response.headers.get(PID_HEADER) == str(pid): + return response + except httpx.TransportError: + pass + time.sleep(0.2) + raise AssertionError(f"metrics server on port {port} never served metrics") + + +def _wait_until_down(port: int) -> None: + deadline: Final = time.monotonic() + _SHUTDOWN_TIMEOUT_SECONDS + while time.monotonic() < deadline: + try: + httpx.get(f"http://127.0.0.1:{port}/metrics", timeout=1.0) + except httpx.TransportError: + return + time.sleep(0.2) + raise AssertionError(f"metrics server on port {port} kept running after its parent died") + + +def test_metrics_app_aggregates_multiproc_dir_and_reports_pid(tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("PROMETHEUS_MULTIPROC_DIR", str(tmp_path)) + _write_worker_sample(pid=1001, value=2) + _write_worker_sample(pid=1002, value=3) + other_dir: Final = tmp_path / "other" + other_dir.mkdir() + + client: Final = TestClient(build_metrics_app(str(tmp_path))) + metrics: Final = client.get("/metrics") + assert metrics.status_code == 200 + assert metrics.headers[PID_HEADER] == str(os.getpid()) + assert 'litellm_requests_metric_total{model="gpt-5"} 5.0' in metrics.text + + assert client.get("/health").status_code == 404 + + empty: Final = TestClient(build_metrics_app(str(other_dir))).get("/metrics") + assert empty.status_code == 200 + assert "litellm_requests_metric_total" not in empty.text + + +@pytest.mark.parametrize( + ("host", "expected"), + ( + ("0.0.0.0", "http://127.0.0.1:4001/metrics"), + ("::", "http://[::1]:4001/metrics"), + ("10.1.2.3", "http://10.1.2.3:4001/metrics"), + ("metrics.internal", "http://metrics.internal:4001/metrics"), + ), +) +def test_metrics_url_probes_loopback_for_wildcard_binds(host: str, expected: str): + assert metrics_url(host, 4001) == expected + + +def test_main_serves_the_app_for_the_given_dir_with_uvicorn(tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("PROMETHEUS_MULTIPROC_DIR", str(tmp_path)) + _write_worker_sample(pid=2001, value=6) + with patch("uvicorn.run") as run: + main(["--host", "10.1.2.3", "--port", "4001", "--multiproc_dir", str(tmp_path)]) + + run.assert_called_once() + assert run.call_args.kwargs["host"] == "10.1.2.3" + assert run.call_args.kwargs["port"] == 4001 + client: Final = TestClient(run.call_args.args[0]) + metrics: Final = client.get("/metrics") + assert metrics.status_code == 200 + assert metrics.headers[PID_HEADER] == str(os.getpid()) + assert 'litellm_requests_metric_total{model="gpt-5"} 6.0' in client.get("/metrics").text + + +def test_main_falls_back_to_env_multiproc_dir(tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("PROMETHEUS_MULTIPROC_DIR", str(tmp_path)) + with patch("uvicorn.run") as run: + main(["--port", "4001"]) + + (app,), served_on = run.call_args + assert served_on["host"] == "0.0.0.0" + metrics: Final = TestClient(app).get("/metrics") + assert metrics.status_code == 200 + assert metrics.headers[PID_HEADER] == str(os.getpid()) + + +def test_main_rejects_missing_multiproc_dir(monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("PROMETHEUS_MULTIPROC_DIR", raising=False) + with patch("uvicorn.run") as run, pytest.raises(SystemExit) as exit_info: + main(["--port", "4001"]) + + assert exit_info.value.code == 2 + run.assert_not_called() + + +def test_start_metrics_server_process_returns_only_once_child_serves(tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("PROMETHEUS_MULTIPROC_DIR", str(tmp_path)) + _write_worker_sample(pid=3001, value=4) + port: Final = _free_port() + with patch("atexit.register") as register: + process: Final = start_metrics_server_process(host="127.0.0.1", port=port, multiproc_dir=str(tmp_path)) + try: + register.assert_called_once_with(process.terminate) + assert process.poll() is None + startup_metrics: Final = httpx.get(f"http://127.0.0.1:{port}/metrics", follow_redirects=True, timeout=5.0) + assert startup_metrics.status_code == 200 + assert startup_metrics.headers[PID_HEADER] == str(process.pid) + metrics: Final = httpx.get(f"http://127.0.0.1:{port}/metrics", follow_redirects=True, timeout=10.0) + assert 'litellm_requests_metric_total{model="gpt-5"} 4.0' in metrics.text + finally: + process.kill() + process.wait(timeout=10) + + +def test_start_metrics_server_process_fails_when_port_is_taken(tmp_path: Path): + with socket.socket() as occupied: + occupied.bind(("127.0.0.1", 0)) + occupied.listen() + port: Final = occupied.getsockname()[1] + with ( + patch("atexit.register"), + pytest.raises( + MetricsServerStartupError, match=rf"exited with code [1-9]\d* before serving 127.0.0.1:{port}" + ), + ): + start_metrics_server_process(host="127.0.0.1", port=port, multiproc_dir=str(tmp_path)) + + +class _ImpostorMetrics(BaseHTTPRequestHandler): + """An unrelated service already on the port that answers /metrics with 200 and plausible metrics.""" + + def do_GET(self) -> None: + body: Final = b"# HELP impostor_metric A plausible metric\n# TYPE impostor_metric counter\nimpostor_metric 1\n" + self.send_response(200) + self.send_header("Content-Type", "text/plain") + self.end_headers() + self.wfile.write(body) + + def log_message(self, format: str, *args: object) -> None: + return + + +def test_start_metrics_server_process_rejects_metrics_from_another_service_on_the_port(tmp_path: Path): + with ThreadingHTTPServer(("127.0.0.1", 0), _ImpostorMetrics) as impostor: + threading.Thread(target=impostor.serve_forever, daemon=True).start() + port: Final = impostor.server_address[1] + impostor_response: Final = httpx.get(f"http://127.0.0.1:{port}/metrics") + assert impostor_response.status_code == 200 + assert "# HELP impostor_metric" in impostor_response.text + assert PID_HEADER not in impostor_response.headers + with ( + patch("atexit.register"), + pytest.raises(MetricsServerStartupError, match=rf"exited with code [1-9]\d* before serving 127.0.0.1:{port}"), + ): + start_metrics_server_process(host="127.0.0.1", port=port, multiproc_dir=str(tmp_path)) + impostor.shutdown() + + +def test_metrics_server_process_serves_and_exits_with_parent(tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("PROMETHEUS_MULTIPROC_DIR", str(tmp_path)) + _write_worker_sample(pid=2001, value=7) + port: Final = _free_port() + server_argv: Final = ( + sys.executable, + "-m", + "litellm.proxy.prometheus_metrics_server", + "--host", + "127.0.0.1", + "--port", + str(port), + "--multiproc_dir", + str(tmp_path), + ) + parent: Final = subprocess.Popen( + ( + sys.executable, + "-c", + "import subprocess, sys, time; p = subprocess.Popen(sys.argv[1:]); print(p.pid, flush=True); time.sleep(600)", + *server_argv, + ), + stdout=subprocess.PIPE, + text=True, + ) + assert parent.stdout is not None + server_pid: Final = int(parent.stdout.readline()) + try: + metrics: Final = _wait_for_metrics(port, server_pid) + assert metrics.status_code == 200 + assert metrics.headers[PID_HEADER] == str(server_pid) + + scrape: Final = httpx.get(f"http://127.0.0.1:{port}/metrics", follow_redirects=True, timeout=10.0) + assert scrape.status_code == 200 + assert scrape.headers[PID_HEADER] == str(server_pid) + assert 'litellm_requests_metric_total{model="gpt-5"} 7.0' in scrape.text + + parent.kill() + parent.wait(timeout=10) + _wait_until_down(port) + finally: + parent.kill() + try: + os.kill(server_pid, 9) + except ProcessLookupError: + pass diff --git a/tests/test_litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py index 9256706d340..0c20d5e0ff0 100644 --- a/tests/test_litellm/proxy/test_proxy_cli.py +++ b/tests/test_litellm/proxy/test_proxy_cli.py @@ -662,6 +662,124 @@ class TestProxyInitializationHelpers: assert "Invalid value for '--limit_concurrency'" in result.output mock_uvicorn_run.assert_not_called() + @patch("uvicorn.run") + @patch("httpx.HTTPTransport.handle_request") + @patch("atexit.register") + @patch("subprocess.Popen") + @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") # test-quality-ok: run_server always wires the DB; same isolation as the sibling CLI tests above + @patch( # test-quality-ok: run_server always wires the DB; same isolation as the sibling CLI tests above + "litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False + ) + def test_prometheus_metrics_port_starts_separate_metrics_process( + self, + mock_should_update, + mock_setup_db, + mock_popen, + mock_atexit_register, + mock_handle_request, + mock_uvicorn_run, + tmp_path, + ): + """--prometheus_metrics_port must spawn `python -m litellm.proxy.prometheus_metrics_server` on --host + with the shared multiproc dir, wait for its /metrics response, and only then start uvicorn. It must stay off by + default, refuse to share --port, and abort the proxy when the child dies before serving.""" + import httpx + from click.testing import CliRunner + + from litellm.proxy.proxy_cli import run_server + + runner = CliRunner() + mock_popen.return_value = MagicMock(pid=4242, **{"poll.return_value": None}) + probed_urls: list[str] = [] + + def child_metrics(request: httpx.Request) -> httpx.Response: + probed_urls.append(str(request.url)) + return httpx.Response(200, headers={"x-litellm-metrics-pid": "4242"}, content=b"") + + mock_handle_request.side_effect = child_metrics + mock_proxy_module = MagicMock( + app=MagicMock(), + ProxyConfig=MagicMock(), + KeyManagementSettings=MagicMock(), + save_worker_config=MagicMock(), + ) + clean_env = { + k: v + for k, v in os.environ.items() + if k not in ("DATABASE_URL", "DIRECT_URL", "PROMETHEUS_METRICS_PORT") + } + clean_env["PROMETHEUS_MULTIPROC_DIR"] = str(tmp_path) + with ( + patch.dict(os.environ, clean_env, clear=True), + patch.dict( + "sys.modules", + { + "proxy_server": mock_proxy_module, + "litellm.proxy.proxy_server": mock_proxy_module, + }, + ), + patch( # test-quality-ok: same isolation as the sibling CLI tests above + "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + ) as mock_get_args, + ): + mock_get_args.side_effect = lambda *a, **k: { + "app": "litellm.proxy.proxy_server:app", + "host": "localhost", + "port": 8000, + } + + result = runner.invoke( + run_server, + ["--local", "--host", "127.0.0.1", "--port", "4000", "--prometheus_metrics_port", "4001"], + ) + assert ( + result.exit_code == 0 + ), f"exit_code={result.exit_code}, output={result.output}" + mock_popen.assert_called_once() + spawned = list(mock_popen.call_args.args[0]) + assert spawned[1:3] == ["-m", "litellm.proxy.prometheus_metrics_server"] + assert spawned[3:] == ["--host", "127.0.0.1", "--port", "4001", "--multiproc_dir", str(tmp_path)] + assert probed_urls == ["http://127.0.0.1:4001/metrics"] + assert "Serving Prometheus metrics on 127.0.0.1:4001/metrics (pid 4242)" in result.output + mock_uvicorn_run.assert_called_once() + + mock_popen.reset_mock() + mock_uvicorn_run.reset_mock() + mock_popen.return_value = MagicMock(pid=4243, **{"poll.return_value": 1}) + result = runner.invoke( + run_server, + ["--local", "--port", "4000", "--prometheus_metrics_port", "4001"], + ) + assert result.exit_code == 1, f"exit_code={result.exit_code}, output={result.output}" + assert "Prometheus metrics server exited with code 1 before serving 0.0.0.0:4001" in result.output + mock_uvicorn_run.assert_not_called() + + mock_popen.reset_mock() + mock_uvicorn_run.reset_mock() + result = runner.invoke(run_server, ["--local"]) + assert ( + result.exit_code == 0 + ), f"exit_code={result.exit_code}, output={result.output}" + mock_popen.assert_not_called() + mock_uvicorn_run.assert_called_once() + + mock_uvicorn_run.reset_mock() + result = runner.invoke( + run_server, + ["--local", "--port", "4000", "--prometheus_metrics_port", "4000"], + ) + assert result.exit_code == 2 + assert "--prometheus_metrics_port must differ from --port" in result.output + mock_popen.assert_not_called() + mock_uvicorn_run.assert_not_called() + + result = runner.invoke( + run_server, ["--local", "--prometheus_metrics_port", "0"] + ) + assert result.exit_code == 2 + assert "Invalid value for '--prometheus_metrics_port'" in result.output + mock_popen.assert_not_called() + @pytest.mark.parametrize( "timeout_config,expected_timeout", [ diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 225134b4e2b..e35e470c979 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -3,7 +3,7 @@ "limit": 22180 }, "LIT002": { - "limit": 26745 + "limit": 26729 }, "LIT003": { "limit": 261 @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16462 + "limit": 16430 }, "LIT011": { "limit": 5506 From 9832d6e4a6e3cc832e4425f759243f97926e6d46 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Sat, 5 Sep 2026 13:51:25 -0700 Subject: [PATCH 17/35] fix(mcp): scan and mask MCP tool call arguments in unified guardrails (#35142) * fix(mcp): scan and mask MCP tool call arguments in unified guardrails A guardrail configured with mode pre_mcp_call was handed only a synthetic tool definition (name plus an empty parameters schema), so it never saw the argument values it was configured to inspect, and any rewrite it returned was discarded. Detection could not fire and masking could not take effect, while the applied-guardrails metadata still reported the guardrail as having run. Pass every string leaf of the tool call arguments as texts, and fold the guardrail's rewritten leaves back into modified_arguments, which is the channel the MCP call path reads to decide what to send upstream. The leaf walk reuses the json_string_leaves / with_json_string_leaves helpers the tool result path already uses, so both directions share one bounded traversal. Two guardrails running concurrently under run_in_parallel scan the same payload snapshot, so each returns a full replacement derived from the original leaf. Rewrites of the same leaf to different values are rejected rather than silently losing one redaction; a leaf that already holds this guardrail's own replacement is convergent and still masks, which is what the bundled content filter does when it rewrites the arguments itself as well as through texts. * fix(mcp): annotate guardrail argument rewrites Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(tests): isolate MCP guardrail callback state Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore: ratchet LIT010 budget after merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(tests): remove duplicate Bedrock hook parameter Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(mcp): fail closed when guardrail rewrites cannot be mapped to MCP arguments Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(tests): patch the guardrail translation mappings cache where staging now keeps it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../guardrail_translation/handler.py | 137 ++++-- .../test_mcp_guardrail_handler.py | 429 +++++++++++++++++- type-discipline-budget.json | 2 +- 3 files changed, 530 insertions(+), 38 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/guardrail_translation/handler.py b/litellm/proxy/_experimental/mcp_server/guardrail_translation/handler.py index 4918229c2b8..c0235077ecd 100644 --- a/litellm/proxy/_experimental/mcp_server/guardrail_translation/handler.py +++ b/litellm/proxy/_experimental/mcp_server/guardrail_translation/handler.py @@ -1,16 +1,19 @@ """ MCP Guardrail Handler for Unified Guardrails. -Converts an MCP call_tool (name + arguments) into a single OpenAI-compatible -tool_call and passes it to apply_guardrail. Works with the synthetic payload -from ProxyLogging._convert_mcp_to_llm_format. +Converts an MCP call_tool (name + arguments) into the OpenAI-compatible shape +apply_guardrail expects: the tool as a single-entry ``tools`` definition, and +every string leaf of the call arguments as ``texts`` so text guardrails can +detect and mask sensitive values in the payload. Works with the synthetic +request from ProxyLogging._convert_mcp_to_llm_format. Note: For MCP tool definitions (schema) -> OpenAI tools=[], see litellm.experimental_mcp_client.tools.transform_mcp_tool_to_openai_tool when you have a full MCP Tool from list_tools. Here we only have the call -payload (name + arguments) so we just build the tool_call. +payload (name + arguments) so we just build the tool definition. """ +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Final from fastapi import HTTPException @@ -20,6 +23,8 @@ from litellm._logging import verbose_proxy_logger from litellm.experimental_mcp_client.tools import transform_mcp_tool_to_openai_tool from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation from litellm.proxy._experimental.mcp_server.utils import ( + MAX_STRUCTURED_CONTENT_SCAN_DEPTH, + JSONLeafPath, json_string_leaves, json_unrewritable_labels, mcp_content_item_text, @@ -42,6 +47,72 @@ if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +def _blocked(reason: str) -> HTTPException: + return HTTPException(status_code=400, detail={"error": f"Content blocked: {reason}"}) + + +def _too_deeply_nested() -> HTTPException: + return _blocked( + f"MCP tool call arguments exceed the maximum nesting depth of {MAX_STRUCTURED_CONTENT_SCAN_DEPTH} " + "and cannot be scanned by the configured guardrail" + ) + + +def _argument_replacements( + argument_leaves: tuple[tuple[JSONLeafPath, str], ...], + masked_texts: Sequence[str] | None, +) -> Mapping[JSONLeafPath, str]: + """Positionally pair the guardrail's returned texts with the leaves they came from. + + Only leaves the guardrail actually rewrote are returned, so a guardrail that + detects nothing leaves the outbound tool call byte-identical. A guardrail that + returns the wrong number of texts fails closed, because a positional write-back + would scramble the arguments rather than mask them. + """ + if masked_texts is not None and len(masked_texts) != len(argument_leaves): + raise _blocked( + f"guardrail returned {len(masked_texts)} texts for {len(argument_leaves)} MCP tool call argument strings, " + "so the redaction cannot be mapped back to the arguments" + ) + return {path: masked for (path, original), masked in zip(argument_leaves, masked_texts or ()) if masked != original} + + +def _conflicting_rewrite_paths( + scanned_leaves: tuple[tuple[JSONLeafPath, str], ...], + current_leaves: tuple[tuple[JSONLeafPath, str], ...], + replacements: Mapping[JSONLeafPath, str], +) -> tuple[JSONLeafPath, ...]: + """Paths another guardrail already rewrote differently from what this one wants. + + Guardrails opted into ``run_in_parallel`` all scan the same payload snapshot, so + each one returns a full replacement string derived from the *original* leaf. Two + of them rewriting one leaf to different values cannot be merged: writing either + result discards the other guardrail's redaction. A leaf still holding the text + this guardrail was handed, or already holding this guardrail's own replacement, + is safe to write; the latter is how a guardrail that masks the arguments itself + as well as through ``texts`` gets there first. Anything else fails closed, + including a payload reshaped so the leaves no longer line up, because the + write-back is positional and would land a redaction on the wrong value. + """ + if tuple(path for path, _ in scanned_leaves) != tuple(path for path, _ in current_leaves): + return tuple(replacements) + return tuple( + path + for (path, scanned), (_, current) in zip(scanned_leaves, current_leaves) + if path in replacements and current not in (scanned, replacements[path]) + ) + + +def _conflicting_rewrite(paths: tuple[JSONLeafPath, ...]) -> HTTPException: + return _blocked( + "two guardrails running concurrently rewrote the same MCP tool call " + f"argument{'s' if len(paths) > 1 else ''} " + f"({', '.join('.'.join(str(part) for part in path) for path in paths)}); " + "their redactions cannot be merged. Remove run_in_parallel from one of them so they " + "run in sequence." + ) + + class MCPGuardrailTranslationHandler(BaseTranslation): """Guardrail translation handler for MCP tool calls (passes a single tool_call to guardrail).""" @@ -52,10 +123,8 @@ class MCPGuardrailTranslationHandler(BaseTranslation): litellm_logging_obj: "LiteLLMLoggingObj | None" = None, ) -> dict[str, Any]: mcp_tool_name: Final = data.get("mcp_tool_name") or data.get("name") - mcp_arguments = data.get("mcp_arguments") or data.get("arguments") + mcp_arguments: Final[object] = data.get("mcp_arguments") or data.get("arguments") mcp_tool_description: Final = data.get("mcp_tool_description") or data.get("description") - if mcp_arguments is None or not isinstance(mcp_arguments, dict): - mcp_arguments = {} if not mcp_tool_name: verbose_proxy_logger.debug("MCP Guardrail: mcp_tool_name missing") @@ -84,16 +153,37 @@ class MCPGuardrailTranslationHandler(BaseTranslation): strict=fn.get("strict", False) or False, # Default to False if None ), } + argument_leaves: Final = json_string_leaves(mcp_arguments) + if argument_leaves is None: + raise _too_deeply_nested() inputs: Final[GenericGuardrailAPIInputs] = GenericGuardrailAPIInputs( tools=[tool_def], + texts=[text for _, text in argument_leaves], ) - await guardrail_to_apply.apply_guardrail( + guarded: Final = await guardrail_to_apply.apply_guardrail( inputs=inputs, request_data=data, input_type="request", logging_obj=litellm_logging_obj, ) + replacements: Final = _argument_replacements( + argument_leaves=argument_leaves, + masked_texts=guarded.get("texts") if guarded else None, + ) + if not replacements: + return data + + current_arguments: Final[object] = data.get("mcp_arguments") or data.get("arguments") + current_leaves: Final = json_string_leaves(current_arguments) + if current_leaves is None: + raise _too_deeply_nested() + conflicting: Final = _conflicting_rewrite_paths(argument_leaves, current_leaves, replacements) + if conflicting: + raise _conflicting_rewrite(conflicting) + masked_arguments: Final = with_json_string_leaves(current_arguments, replacements) + data["mcp_arguments"] = masked_arguments # rebind-ok: preserve the mask for the outbound MCP call + data["modified_arguments"] = masked_arguments # rebind-ok: expose the applied mask to the caller return data async def process_output_response( @@ -131,14 +221,8 @@ class MCPGuardrailTranslationHandler(BaseTranslation): structured_leaves: Final = json_string_leaves(structured) if structured is not None else () structured_labels: Final = json_unrewritable_labels(structured) if structured is not None else () if structured_leaves is None or structured_labels is None: - raise HTTPException( - status_code=400, - detail={ - "error": ( - "Content blocked: MCP tool result structuredContent is nested too deeply to be scanned " - "by the configured guardrail" - ) - }, + raise _blocked( + "MCP tool result structuredContent is nested too deeply to be scanned by the configured guardrail" ) if not text_blocks and not structured_leaves and not structured_labels: @@ -158,12 +242,10 @@ class MCPGuardrailTranslationHandler(BaseTranslation): if masked_texts is None: return response if len(masked_texts) != len(originals): - verbose_proxy_logger.warning( - "MCP Guardrail: guardrail returned %d texts for %d tool result texts; leaving the result unmasked", - len(masked_texts), - len(originals), + raise _blocked( + f"guardrail returned {len(masked_texts)} texts for {len(originals)} MCP tool result texts, " + "so the redaction cannot be mapped back to the result" ) - return response split: Final = len(text_blocks) if content is not None: @@ -173,15 +255,10 @@ class MCPGuardrailTranslationHandler(BaseTranslation): label_start: Final = split + len(structured_leaves) if any(masked != original for original, masked in zip(structured_labels, masked_texts[label_start:])): - raise HTTPException( - status_code=400, - detail={ - "error": ( - "Content blocked: MCP tool result matched a masking rule on a non-rewritable field " - "(a structuredContent key or numeric value), which cannot be redacted without changing " - "the payload contract" - ) - }, + raise _blocked( + "MCP tool result matched a masking rule on a non-rewritable field " + "(a structuredContent key or numeric value), which cannot be redacted without changing " + "the payload contract" ) structured_replacements: Final = { diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/guardrail_translation/test_mcp_guardrail_handler.py b/tests/test_litellm/proxy/_experimental/mcp_server/guardrail_translation/test_mcp_guardrail_handler.py index 2e286a237c4..28959054195 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/guardrail_translation/test_mcp_guardrail_handler.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/guardrail_translation/test_mcp_guardrail_handler.py @@ -1,13 +1,22 @@ """Tests for the MCP guardrail translation handler.""" +import asyncio + import pytest +from fastapi import HTTPException from mcp.types import CallToolResult, ImageContent, TextContent +import litellm +import litellm.llms as litellm_llms +from litellm.caching.caching import DualCache from litellm.exceptions import BlockedPiiEntityError from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.proxy._experimental.mcp_server.guardrail_translation.handler import ( MCPGuardrailTranslationHandler, ) +from litellm.proxy._experimental.mcp_server.utils import MAX_STRUCTURED_CONTENT_SCAN_DEPTH +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.utils import ProxyLogging from litellm.types.utils import GenericGuardrailAPIInputs @@ -24,12 +33,11 @@ class MockGuardrail(CustomGuardrail): self.call_count += 1 self.last_inputs = inputs self.last_request_data = request_data - return None # Guardrail doesn't modify for MCP tools @pytest.mark.asyncio async def test_process_input_messages_updates_content(): - """Handler should pass tool definition to guardrail when mcp_tool_name is present.""" + """Handler should pass the tool definition and the argument strings to the guardrail.""" handler = MCPGuardrailTranslationHandler() guardrail = MockGuardrail() @@ -45,7 +53,7 @@ async def test_process_input_messages_updates_content(): assert result == data # Guardrail was called assert guardrail.call_count == 1 - # Guardrail received tools (not texts) with tool definition + # Guardrail received tools with the tool definition assert guardrail.last_inputs is not None tools = guardrail.last_inputs.get("tools", []) assert len(tools) == 1 @@ -85,6 +93,412 @@ async def test_process_input_messages_handles_minimal_data(): assert tools[0]["function"]["name"] == "simple_tool" +class ArgumentMaskingGuardrail(CustomGuardrail): + """Unified guardrail that rewrites every text it is handed, like presidio does.""" + + def __init__( + self, + secret: str = "jane.doe@example.com", + replacement: str = "", + texts_override: list[str] | None = None, + **kwargs, + ): + kwargs.setdefault("guardrail_name", "argument-masking-mcp-guardrail") + super().__init__(**kwargs) + self.secret = secret + self.replacement = replacement + self.texts_override = texts_override + self.seen_texts: list[str] | None = None + + def _mask(self, text: str) -> str: + return text.replace(self.secret, self.replacement) + + async def apply_guardrail(self, inputs, request_data, input_type, **kwargs): + self.seen_texts = list(inputs.get("texts") or []) + if self.texts_override is not None: + inputs["texts"] = self.texts_override + else: + inputs["texts"] = [self._mask(text) for text in self.seen_texts] + return inputs + + +@pytest.fixture +def restore_callbacks(monkeypatch): + """Restore the process-wide state driving pre_call_hook through unified_guardrail. + + litellm.llms memoizes the guardrail translation mappings in a module global, and + ProxyLogging caches callback capabilities keyed on id()s of litellm.callbacks, + so leaving either populated leaks into unrelated tests in the same worker. + """ + monkeypatch.setattr(litellm, "callbacks", litellm.callbacks) + monkeypatch.setattr( + litellm_llms, + "endpoint_guardrail_translation_mappings", + litellm_llms.endpoint_guardrail_translation_mappings, + ) + yield + ProxyLogging._callback_capabilities_cache.clear() + + +@pytest.mark.asyncio +async def test_argument_strings_are_handed_to_the_guardrail(): + """A guardrail must see the argument values, not just the tool definition. + + Without this the guardrail is handed a name and an empty schema, so no + sensitive-data detection can ever fire on an MCP tool call. + """ + handler = MCPGuardrailTranslationHandler() + guardrail = MockGuardrail() + + data = { + "mcp_tool_name": "search", + "mcp_arguments": {"query": "contact jane.doe@example.com about the invoice"}, + } + + await handler.process_input_messages(data, guardrail) + + assert guardrail.last_inputs is not None + assert guardrail.last_inputs.get("texts") == ["contact jane.doe@example.com about the invoice"] + + +@pytest.mark.asyncio +async def test_masked_arguments_are_written_back_for_the_call_path(): + """A mask only takes effect once it lands in modified_arguments.""" + handler = MCPGuardrailTranslationHandler() + guardrail = ArgumentMaskingGuardrail() + + data = { + "mcp_tool_name": "search", + "mcp_arguments": {"query": "contact jane.doe@example.com about the invoice"}, + } + + result = await handler.process_input_messages(data, guardrail) + + masked = {"query": "contact about the invoice"} + assert result["modified_arguments"] == masked + assert result["mcp_arguments"] == masked + + +@pytest.mark.asyncio +async def test_nested_arguments_keep_their_shape_when_masked(): + """Masking rewrites string leaves in place and preserves non-string values.""" + handler = MCPGuardrailTranslationHandler() + guardrail = ArgumentMaskingGuardrail() + + arguments = { + "recipients": ["jane.doe@example.com", "ops@example.net"], + "envelope": {"reply_to": "jane.doe@example.com", "retries": 3, "urgent": True, "cc": None}, + "count": 2, + } + data = {"mcp_tool_name": "send_email", "mcp_arguments": arguments} + + result = await handler.process_input_messages(data, guardrail) + + assert guardrail.seen_texts == [ + "jane.doe@example.com", + "ops@example.net", + "jane.doe@example.com", + ] + assert result["modified_arguments"] == { + "recipients": ["", "ops@example.net"], + "envelope": {"reply_to": "", "retries": 3, "urgent": True, "cc": None}, + "count": 2, + } + + +@pytest.mark.asyncio +async def test_clean_arguments_are_not_overridden(): + """A guardrail that changes nothing must not set modified_arguments.""" + handler = MCPGuardrailTranslationHandler() + guardrail = ArgumentMaskingGuardrail() + + data = {"mcp_tool_name": "search", "mcp_arguments": {"query": "quarterly revenue"}} + + result = await handler.process_input_messages(data, guardrail) + + assert "modified_arguments" not in result + assert result["mcp_arguments"] == {"query": "quarterly revenue"} + + +@pytest.mark.asyncio +async def test_guardrail_returning_wrong_text_count_blocks_the_call(): + """Write-back is positional, so a length mismatch must block the call.""" + handler = MCPGuardrailTranslationHandler() + guardrail = ArgumentMaskingGuardrail(texts_override=["only", "two", "texts"]) + + arguments = {"query": "contact jane.doe@example.com about the invoice"} + data = {"mcp_tool_name": "search", "mcp_arguments": arguments} + + with pytest.raises(HTTPException) as exc_info: + await handler.process_input_messages(data, guardrail) + + assert exc_info.value.status_code == 400 + assert "modified_arguments" not in data + + +@pytest.mark.asyncio +async def test_deeply_nested_arguments_are_blocked_rather_than_skipped(): + """Arguments too deep to walk must block instead of passing unscanned.""" + handler = MCPGuardrailTranslationHandler() + guardrail = ArgumentMaskingGuardrail() + + nested: dict = {"leaf": "jane.doe@example.com"} + for _ in range(MAX_STRUCTURED_CONTENT_SCAN_DEPTH + 1): + nested = {"next": nested} + + data = {"mcp_tool_name": "search", "mcp_arguments": nested} + + with pytest.raises(HTTPException) as exc_info: + await handler.process_input_messages(data, guardrail) + + assert exc_info.value.status_code == 400 + + +class SelfWritingMaskingGuardrail(ArgumentMaskingGuardrail): + """Masks through ``texts`` and writes the masked arguments itself. + + The shape the bundled content filter guardrail already has: it rewrites + ``request_data["mcp_arguments"]`` from inside ``apply_guardrail`` as well as + returning masked texts. + """ + + async def apply_guardrail(self, inputs, request_data, input_type, **kwargs): + returned = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + arguments = request_data.get("mcp_arguments") or {} + masked = {key: self._mask(value) if isinstance(value, str) else value for key, value in arguments.items()} + request_data["mcp_arguments"] = masked + request_data["modified_arguments"] = masked + return returned + + +@pytest.mark.asyncio +async def test_guardrail_that_masks_the_arguments_itself_is_not_treated_as_a_conflict(): + """Converging on the same replacement is not an unmergeable rewrite. + + A guardrail that both returns masked texts and rewrites the arguments in + request_data must still mask, not be rejected as if a second guardrail had + clobbered the leaf. + """ + handler = MCPGuardrailTranslationHandler() + guardrail = SelfWritingMaskingGuardrail() + + data = { + "mcp_tool_name": "search", + "mcp_arguments": {"query": "contact jane.doe@example.com about the invoice"}, + } + + result = await handler.process_input_messages(data, guardrail) + + assert result["modified_arguments"] == {"query": "contact about the invoice"} + + +class ReshapingGuardrail(ArgumentMaskingGuardrail): + """Masks through ``texts`` while moving the secret to a different path.""" + + def __init__(self, reshaped: dict, **kwargs): + super().__init__(**kwargs) + self.reshaped = reshaped + + async def apply_guardrail(self, inputs, request_data, input_type, **kwargs): + returned = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + request_data["mcp_arguments"] = self.reshaped + return returned + + +@pytest.mark.asyncio +async def test_arguments_reshaped_under_the_guardrail_fail_closed(): + """A payload that no longer lines up leaf for leaf must block, not be written blind. + + Write-back pairs masked texts to leaves positionally, so a tree another guardrail + reshaped would take the redaction on the wrong value. + """ + handler = MCPGuardrailTranslationHandler() + guardrail = ReshapingGuardrail({"query": "contact jane.doe@example.com", "note": "added"}) + + data = { + "mcp_tool_name": "search", + "mcp_arguments": {"query": "contact jane.doe@example.com about the invoice"}, + } + + with pytest.raises(HTTPException) as exc_info: + await handler.process_input_messages(data, guardrail) + + assert exc_info.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_arguments_shortened_under_the_guardrail_fail_closed(): + handler = MCPGuardrailTranslationHandler() + guardrail = ReshapingGuardrail({"padding": "jane.doe@example.com"}) + + data = { + "mcp_tool_name": "search", + "mcp_arguments": {"padding": "x", "secret": "jane.doe@example.com"}, + } + + with pytest.raises(HTTPException) as exc_info: + await handler.process_input_messages(data, guardrail) + + assert exc_info.value.status_code == 400 + assert "jane.doe@example.com" not in str(data.get("modified_arguments")) + + +@pytest.mark.asyncio +async def test_a_renamed_argument_key_blocks_rather_than_dropping_the_mask(): + """The leak this closes: same text, new path, so the write-back would find nothing. + + Matching purely on position would see an unchanged value and write the mask to a + path that no longer exists, shipping the secret while reporting a clean scan. + """ + handler = MCPGuardrailTranslationHandler() + guardrail = ReshapingGuardrail({"renamed": "jane.doe@example.com", "other": "kept"}) + + data = { + "mcp_tool_name": "search", + "mcp_arguments": {"query": "jane.doe@example.com", "other": "kept"}, + } + + with pytest.raises(HTTPException) as exc_info: + await handler.process_input_messages(data, guardrail) + + assert exc_info.value.status_code == 400 + assert "jane.doe@example.com" not in str(data.get("modified_arguments")) + + +@pytest.mark.parametrize("run_in_parallel", [False, True]) +@pytest.mark.asyncio +async def test_masked_arguments_reach_the_outbound_mcp_call(restore_callbacks, monkeypatch, run_in_parallel): + """End to end over the real MCP pre-call path, not just the handler. + + Drives the same sequence mcp_server_manager.call_tool uses: + synthetic payload -> pre_call_hook -> arguments sent upstream. + + Covers run_in_parallel both ways: that path shares one payload snapshot and + discards whatever a guardrail returns, so the mask has to land on the caller's + dict rather than on a copy of it. + """ + guardrail = ArgumentMaskingGuardrail( + event_hook="pre_mcp_call", + default_on=True, + run_in_parallel=run_in_parallel, + ) + monkeypatch.setattr(litellm, "callbacks", [guardrail]) + + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + arguments = {"query": "contact jane.doe@example.com about the invoice"} + pre_hook_kwargs = { + "name": "search", + "arguments": arguments, + "server_name": "test-server", + "user_api_key_auth": UserAPIKeyAuth(api_key="sk-test", user_id="test-user"), + } + + request_obj = proxy_logging_obj._create_mcp_request_object_from_kwargs(pre_hook_kwargs) + synthetic_data = proxy_logging_obj._convert_mcp_to_llm_format(request_obj, pre_hook_kwargs) + + modified_data = await proxy_logging_obj.pre_call_hook( + user_api_key_dict=pre_hook_kwargs["user_api_key_auth"], + data=synthetic_data, + call_type="call_mcp_tool", + ) + modified_kwargs = proxy_logging_obj._convert_mcp_hook_response_to_kwargs(modified_data, pre_hook_kwargs) + + assert modified_kwargs["arguments"] == {"query": "contact about the invoice"} + + +class SlowSubstitutionGuardrail(CustomGuardrail): + """Rewrites one substring, after a delay, so two instances genuinely interleave.""" + + def __init__(self, needle: str, replacement: str, delay: float, **kwargs): + super().__init__(**kwargs) + self.needle = needle + self.replacement = replacement + self.delay = delay + + async def apply_guardrail(self, inputs, request_data, input_type, **kwargs): + await asyncio.sleep(self.delay) + inputs["texts"] = [text.replace(self.needle, self.replacement) for text in (inputs.get("texts") or [])] + return inputs + + +def _two_interleaving_maskers(run_in_parallel: bool): + return [ + SlowSubstitutionGuardrail( + "jane.doe@example.com", + "", + 0.02, + guardrail_name="mask-email", + event_hook="pre_mcp_call", + default_on=True, + run_in_parallel=run_in_parallel, + ), + SlowSubstitutionGuardrail( + "415-555-0132", + "", + 0.04, + guardrail_name="mask-phone", + event_hook="pre_mcp_call", + default_on=True, + run_in_parallel=run_in_parallel, + ), + ] + + +async def _arguments_sent_upstream(arguments: dict): + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + pre_hook_kwargs = { + "name": "search", + "arguments": arguments, + "server_name": "test-server", + "user_api_key_auth": UserAPIKeyAuth(api_key="sk-test", user_id="test-user"), + } + request_obj = proxy_logging_obj._create_mcp_request_object_from_kwargs(pre_hook_kwargs) + modified_data = await proxy_logging_obj.pre_call_hook( + user_api_key_dict=pre_hook_kwargs["user_api_key_auth"], + data=proxy_logging_obj._convert_mcp_to_llm_format(request_obj, pre_hook_kwargs), + call_type="call_mcp_tool", + ) + return proxy_logging_obj._convert_mcp_hook_response_to_kwargs(modified_data, pre_hook_kwargs)["arguments"] + + +@pytest.mark.asyncio +async def test_two_sequential_guardrails_both_masks_survive(restore_callbacks, monkeypatch): + """The recommended config: each guardrail sees the previous one's output.""" + monkeypatch.setattr(litellm, "callbacks", _two_interleaving_maskers(run_in_parallel=False)) + + sent = await _arguments_sent_upstream({"note": "mail jane.doe@example.com or call 415-555-0132"}) + + assert sent == {"note": "mail or call "} + + +@pytest.mark.asyncio +async def test_two_parallel_guardrails_on_separate_arguments_both_masks_survive(restore_callbacks, monkeypatch): + """Concurrent rewrites of different leaves compose; neither is lost.""" + monkeypatch.setattr(litellm, "callbacks", _two_interleaving_maskers(run_in_parallel=True)) + + sent = await _arguments_sent_upstream({"email": "jane.doe@example.com", "phone": "415-555-0132"}) + + assert sent == {"email": "", "phone": ""} + + +@pytest.mark.asyncio +async def test_two_parallel_guardrails_on_one_argument_block_instead_of_losing_a_mask(restore_callbacks, monkeypatch): + """Unmergeable concurrent rewrites must fail closed, not ship one redaction. + + Both guardrails derive a full replacement string from the same snapshot, so + writing either result would silently discard the other's redaction and leak + the value it was configured to mask. + """ + monkeypatch.setattr(litellm, "callbacks", _two_interleaving_maskers(run_in_parallel=True)) + original = "mail jane.doe@example.com or call 415-555-0132" + + with pytest.raises(HTTPException) as exc_info: + await _arguments_sent_upstream({"note": original}) + + assert exc_info.value.status_code == 400 + assert "note" in str(exc_info.value.detail) + + class MaskingGuardrail(CustomGuardrail): """Guardrail that rewrites every scanned text, recording what it saw.""" @@ -190,8 +604,8 @@ async def test_process_output_response_handles_result_without_content(): @pytest.mark.asyncio -async def test_process_output_response_leaves_result_unmasked_on_text_count_mismatch(): - """A guardrail returning the wrong number of texts must not shuffle content.""" +async def test_process_output_response_blocks_on_text_count_mismatch(): + """A guardrail returning the wrong number of texts must block the result.""" handler = MCPGuardrailTranslationHandler() guardrail = MaskingGuardrail(masked_texts=[""]) result = CallToolResult( @@ -202,9 +616,10 @@ async def test_process_output_response_leaves_result_unmasked_on_text_count_mism isError=False, ) - returned = await handler.process_output_response(response=result, guardrail_to_apply=guardrail) + with pytest.raises(HTTPException) as exc_info: + await handler.process_output_response(response=result, guardrail_to_apply=guardrail) - assert [item.text for item in returned.content] == ["jane@example.com", "415-555-0132"] + assert exc_info.value.status_code == 400 class SubstitutingGuardrail(CustomGuardrail): diff --git a/type-discipline-budget.json b/type-discipline-budget.json index e35e470c979..e7186dfe186 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16430 + "limit": 16426 }, "LIT011": { "limit": 5506 From a46a076b2abd46b88f65d6d21d7afd9c052bb826 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 5 Sep 2026 14:22:26 -0700 Subject: [PATCH 18/35] fix(proxy): reject ambiguous name or alias keys in mcp_tool_permissions on write (#39947) Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../organization_endpoints.py | 6 + .../object_permission_utils.py | 71 ++++++++++- .../test_internal_user_endpoints.py | 2 + .../test_key_management_endpoints.py | 1 + .../test_organization_endpoints.py | 31 +++++ .../test_team_endpoints.py | 1 + .../test_object_permission_utils.py | 119 ++++++++++++++++++ 7 files changed, 229 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index 0af9f816318..1e711b036d2 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -46,6 +46,7 @@ from litellm.proxy.management_endpoints.common_utils import ( from litellm.proxy.management_helpers.object_permission_utils import ( handle_update_object_permission_common, prepare_object_permission_upsert, + reject_ambiguous_mcp_tool_permission_keys, ) from litellm.proxy.management_helpers.utils import ( get_new_internal_user_defaults, @@ -606,6 +607,11 @@ async def _set_object_permission( return None if data.object_permission is not None: + await reject_ambiguous_mcp_tool_permission_keys( + new_mcp_tool_permissions=data.object_permission.mcp_tool_permissions, + existing_mcp_tool_permissions=None, + prisma_client=prisma_client, + ) created_object_permission: Final = await _table(ObjectPermissionRepository(prisma_client)).create( data=data.object_permission.model_dump(exclude_none=True), ) diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index a2fbf80422c..daab38d3662 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -5,10 +5,13 @@ organizations, teams, and keys. import json from collections.abc import Mapping, Sequence +from collections.abc import Set as AbstractSet from dataclasses import dataclass +from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Optional from fastapi import HTTPException, status +from pydantic import TypeAdapter from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid @@ -103,6 +106,11 @@ async def prepare_object_permission_upsert( if existing_object_permission is not None else {} ) + await reject_ambiguous_mcp_tool_permission_keys( + new_mcp_tool_permissions=new_object_permission.get("mcp_tool_permissions"), + existing_mcp_tool_permissions=existing_fields.get("mcp_tool_permissions"), + prisma_client=prisma_client, + ) merged: Final[dict[str, object]] = { **existing_fields, **new_object_permission, @@ -194,6 +202,12 @@ async def _set_object_permission( k: v for k, v in permission_data.items() if v is not None and k != "object_permission_id" } + await reject_ambiguous_mcp_tool_permission_keys( + new_mcp_tool_permissions=clean_data.get("mcp_tool_permissions"), + existing_mcp_tool_permissions=None, + prisma_client=prisma_client, + ) + # Serialize mcp_tool_permissions to JSON string for GraphQL compatibility if "mcp_tool_permissions" in clean_data: clean_data["mcp_tool_permissions"] = safe_dumps(clean_data["mcp_tool_permissions"]) @@ -226,7 +240,7 @@ def _mcp_server_identifier_matches(server: Any, identifier: str) -> bool: async def _get_db_mcp_servers_by_identifiers( - identifiers: set[str], + identifiers: AbstractSet[str], prisma_client: PrismaClient | None, ) -> "Sequence[prisma_models.LiteLLM_MCPServerTable]": if prisma_client is None or not identifiers: @@ -245,7 +259,7 @@ async def _get_db_mcp_servers_by_identifiers( async def _resolve_mcp_server_identifiers_to_ids( - identifiers: set[str], + identifiers: AbstractSet[str], prisma_client: PrismaClient | None, ) -> dict[str, set[str]]: """ @@ -286,6 +300,59 @@ async def _resolve_mcp_server_identifiers_to_ids( return resolved +_MCP_TOOL_PERMISSIONS_ADAPTER: Final = TypeAdapter(dict[str, list[str] | None]) + + +def _mcp_tool_permission_entries(raw: object) -> Mapping[str, frozenset[str]]: + parsed: Final[Mapping[str, Sequence[str] | None]] = ( + _MCP_TOOL_PERMISSIONS_ADAPTER.validate_json(raw) + if isinstance(raw, str) + else _MCP_TOOL_PERMISSIONS_ADAPTER.validate_python(raw) + if isinstance(raw, Mapping) + else MappingProxyType({}) + ) + return MappingProxyType({identifier: frozenset(tools or ()) for identifier, tools in parsed.items()}) + + +async def reject_ambiguous_mcp_tool_permission_keys( + new_mcp_tool_permissions: object, + existing_mcp_tool_permissions: object, + prisma_client: PrismaClient | None, +) -> None: + """ + A name or alias shared by several MCP servers cannot key ``mcp_tool_permissions``: + the read path unions the entry into every match, so no edit can narrow one of + those servers without also changing the other. An exact server_id is never + ambiguous, even when another server uses that string as its alias. Entries the + row already stores with the same tool list are left alone, so unrelated edits + to such an entity still succeed. + + Raises HTTPException(400) naming the colliding servers. + """ + requested: Final = _mcp_tool_permission_entries(new_mcp_tool_permissions) + stored: Final = _mcp_tool_permission_entries(existing_mcp_tool_permissions) + resolved: Final = await _resolve_mcp_server_identifiers_to_ids( + identifiers=frozenset(identifier for identifier, tools in requested.items() if stored.get(identifier) != tools), + prisma_client=prisma_client, + ) + collisions: Final = "; ".join( + f"'{identifier}' matches MCP servers {sorted(server_ids)}" + for identifier, server_ids in sorted(resolved.items()) + if identifier not in server_ids and len(server_ids) > 1 + ) + if not collisions: + return + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={ # mutable-ok: HTTPException.detail has no immutable form; same shape as the sibling errors here + "error": ( + f"Ambiguous mcp_tool_permissions key: {collisions}. " + "Key tool permissions by server_id when servers share a name or alias." + ) + }, + ) + + def _drop_stale_object_permission_mcp_servers( object_permission: ObjectPermissionDict, identifier_to_server_ids: dict[str, set[str]], diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index 0d3ea5863a2..d1d669cae38 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -3917,6 +3917,7 @@ def _object_permission_mocks(mocker, existing_object_permission_id=None): mock_prisma_client.db.litellm_objectpermissiontable.upsert = mocker.AsyncMock( return_value=SimpleNamespace(object_permission_id="perm-new") ) + mock_prisma_client.db.litellm_mcpservertable.find_many = mocker.AsyncMock(return_value=[]) mock_prisma_client.update_data = mocker.AsyncMock( return_value={"user_id": "target-user"} ) @@ -4146,6 +4147,7 @@ async def test_new_user_persists_the_requested_mcp_entitlement(mocker): mock_prisma_client.db.litellm_objectpermissiontable.create = mocker.AsyncMock( return_value=SimpleNamespace(object_permission_id="perm-created") ) + mock_prisma_client.db.litellm_mcpservertable.find_many = mocker.AsyncMock(return_value=[]) mock_prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock( return_value=None ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 47571497f74..8766b1a1868 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -963,6 +963,7 @@ async def test_key_generation_with_mcp_tool_permissions(monkeypatch): mock_prisma_client.db = MagicMock() mock_prisma_client.db.litellm_objectpermissiontable = MagicMock() mock_prisma_client.db.litellm_objectpermissiontable.create = mock_create + mock_prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[]) async def _insert_data_side_effect(*args, **kwargs): table_name = kwargs.get("table_name") diff --git a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py index da68492e3d7..4d13e054e46 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py @@ -1183,6 +1183,37 @@ async def test_find_member_if_email_missing_row_raises_documented_400(): } +@pytest.mark.asyncio +async def test_new_organization_rejects_shared_alias_tool_permission_key(): + """/organization/new creates its permission row through its own helper, so the + ambiguous mcp_tool_permissions key check (LIT-4982) has to run there too.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionBase, NewOrganizationRequest + from litellm.proxy.management_endpoints.organization_endpoints import ( + _set_object_permission, + ) + + prisma_client = MagicMock() + prisma_client.db.litellm_mcpservertable.find_many = AsyncMock( + return_value=[ + MagicMock(server_id="wiki-a-id", alias="wiki", server_name="wiki_a"), + MagicMock(server_id="wiki-b-id", alias="wiki", server_name="wiki_b"), + ] + ) + prisma_client.db.litellm_objectpermissiontable.create = AsyncMock() + data = NewOrganizationRequest( + organization_alias="org", + object_permission=LiteLLM_ObjectPermissionBase(mcp_tool_permissions={"wiki": ["ask_question"]}), + ) + + with pytest.raises(HTTPException) as exc_info: + await _set_object_permission(data=data, prisma_client=prisma_client) + + assert exc_info.value.status_code == 400 + assert "wiki-a-id" in str(exc_info.value.detail) + assert "wiki-b-id" in str(exc_info.value.detail) + prisma_client.db.litellm_objectpermissiontable.create.assert_not_called() + + def test_v2_update_organization_is_in_openapi_schema(): """PATCH /v2/organization/{organization_id} is documented in the generated OpenAPI spec.""" from fastapi import FastAPI diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 46678c8ff6a..051e6bed4fd 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -651,6 +651,7 @@ async def test_new_team_with_mcp_tool_permissions(mock_db_client, mock_admin_aut mock_db_client.db.litellm_objectpermissiontable = MagicMock() mock_db_client.db.litellm_objectpermissiontable.create = mock_obj_perm_create + mock_db_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[]) # Mock model table mock_db_client.db.litellm_modeltable = MagicMock() diff --git a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py index f2b6b799271..d7ebb1f60bf 100644 --- a/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_object_permission_utils.py @@ -17,6 +17,7 @@ from litellm.proxy.management_helpers.object_permission_utils import ( _resolve_team_allowed_mcp_servers, _set_object_permission, enforce_all_proxy_mcp_servers_grant_is_admin_only, + prepare_object_permission_upsert, validate_key_mcp_servers_against_team, validate_key_search_tools_against_team, validate_key_vector_stores_against_team, @@ -41,6 +42,7 @@ async def test_set_object_permission(): mock_prisma_client.db.litellm_objectpermissiontable.create = AsyncMock( return_value=mock_created_permission ) + mock_prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[]) # Test data with object_permission data_json = { @@ -1349,6 +1351,123 @@ async def test_validate_key_update_sentinels_do_not_grandfather(monkeypatch): assert exc_info.value.status_code == 403 +# ---- Tests for rejecting ambiguous mcp_tool_permissions keys on write (LIT-4982) ---- + + +_SHARED_ALIAS_DB_SERVERS = ( + _make_mock_mcp_server("wiki-a-id", alias="wiki", server_name="wiki_a"), + _make_mock_mcp_server("wiki-b-id", alias="wiki", server_name="wiki_b"), + _make_mock_mcp_server("gh-a-id", alias="gh_a", server_name="github"), + _make_mock_mcp_server("gh-b-id", alias="gh_b", server_name="github"), + _make_mock_mcp_server("solo-id", alias="solo", server_name="Solo Server"), + _make_mock_mcp_server("shadow-id", alias="solo-id", server_name="shadow"), +) + + +def _make_ambiguity_prisma(existing_tool_permissions=None): + """Mock prisma client whose MCP server table holds _SHARED_ALIAS_DB_SERVERS and whose + object permission row (if any) stores the given mcp_tool_permissions JSON string.""" + mock_prisma = MagicMock() + mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=list(_SHARED_ALIAS_DB_SERVERS)) + mock_prisma.db.litellm_objectpermissiontable.create = AsyncMock( + return_value=MagicMock(object_permission_id="perm-id") + ) + existing_row = None + if existing_tool_permissions is not None: + existing_row = MagicMock() + existing_row.model_dump.return_value = { + "object_permission_id": "perm-id", + "mcp_tool_permissions": json.dumps(existing_tool_permissions), + } + mock_prisma.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=existing_row) + return mock_prisma + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "identifier, colliding_ids", + [("wiki", ("wiki-a-id", "wiki-b-id")), ("github", ("gh-a-id", "gh-b-id"))], +) +async def test_set_object_permission_rejects_shared_alias_or_name_tool_permission_key(identifier, colliding_ids): + """An alias or server_name two servers share cannot key mcp_tool_permissions on + create: the write is rejected with 400 naming both servers and nothing is persisted.""" + mock_prisma = _make_ambiguity_prisma() + data_json = {"object_permission": {"mcp_tool_permissions": {identifier: ["read_wiki_structure"]}}} + + with pytest.raises(HTTPException) as exc_info: + await _set_object_permission(data_json=data_json, prisma_client=mock_prisma) + + assert exc_info.value.status_code == 400 + assert all(server_id in str(exc_info.value.detail) for server_id in colliding_ids) + mock_prisma.db.litellm_objectpermissiontable.create.assert_not_called() + + +@pytest.mark.asyncio +async def test_prepare_object_permission_upsert_rejects_shared_alias_tool_permission_key(): + """The update seam shared by key/team/org/user/customer/agent rejects a new + shared-alias key when the existing row does not already hold it.""" + mock_prisma = _make_ambiguity_prisma(existing_tool_permissions={"solo-id": ["tool1"]}) + + with pytest.raises(HTTPException) as exc_info: + await prepare_object_permission_upsert( + new_object_permission={"mcp_tool_permissions": {"wiki": ["ask_question"]}}, + existing_object_permission_id="perm-id", + prisma_client=mock_prisma, + ) + + assert exc_info.value.status_code == 400 + assert "'wiki'" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_unambiguous_tool_permission_keys_persist_verbatim(): + """Exact ids (even when another server uses that id string as its alias), + unique aliases, and an id plus alias pointing at one server all still write.""" + mock_prisma = _make_ambiguity_prisma() + tool_permissions = { + "wiki-a-id": ["ask_question"], + "wiki-b-id": ["read_wiki_structure"], + "solo-id": ["tool1"], + "solo": ["tool2"], + "Solo Server": ["tool3"], + } + + upsert = await prepare_object_permission_upsert( + new_object_permission={"mcp_tool_permissions": dict(tool_permissions)}, + existing_object_permission_id=None, + prisma_client=mock_prisma, + ) + + assert json.loads(upsert.record["mcp_tool_permissions"]) == tool_permissions + + +@pytest.mark.asyncio +async def test_stored_ambiguous_tool_permission_key_is_grandfathered_until_changed(): + """A shared-alias entry already on the row may be re-sent unchanged so unrelated + edits succeed, but changing its tool list is rejected.""" + mock_prisma = _make_ambiguity_prisma(existing_tool_permissions={"wiki": ["read_wiki_structure"]}) + + upsert = await prepare_object_permission_upsert( + new_object_permission={ + "mcp_tool_permissions": {"wiki": ["read_wiki_structure"], "solo-id": ["tool1"]}, + }, + existing_object_permission_id="perm-id", + prisma_client=mock_prisma, + ) + assert json.loads(upsert.record["mcp_tool_permissions"]) == { + "wiki": ["read_wiki_structure"], + "solo-id": ["tool1"], + } + + with pytest.raises(HTTPException) as exc_info: + await prepare_object_permission_upsert( + new_object_permission={"mcp_tool_permissions": {"wiki": ["ask_question"]}}, + existing_object_permission_id="perm-id", + prisma_client=mock_prisma, + ) + assert exc_info.value.status_code == 400 + + def test_object_permission_dict_mirrors_pydantic_model(): """ObjectPermissionDict must stay field-for-field aligned with LiteLLM_ObjectPermissionBase. If a new field is added to the Pydantic From b99d8ac38ee021b4a58b1fcfd4f4fbc1cf0c5b62 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Sat, 5 Sep 2026 14:27:08 -0700 Subject: [PATCH 19/35] refactor(ui): keep guardrail usage code under the inline-object-arg lint budget The staging merge pushed local/no-large-inline-object-arg to 567 against a 554 ceiling, and 15 of those hits came from this branch. useGuardrailsUsageDetail now takes the guardrail id positionally with the date window as its second argument, the usageUnits tests build CounterMath rows through a positional helper, and the overview fixture spreads a base row inside the array instead of calling a factory Claude-Session: https://claude.ai/code/session_01EX13mWex6RaBo9PYnkAtFW --- .../_components/GuardrailDetail.test.tsx | 8 ++-- .../_components/GuardrailDetail.tsx | 2 +- .../GuardrailsMonitorView.test.tsx | 3 +- .../_components/GuardrailsOverview.test.tsx | 20 +++++---- .../guardrails/useGuardrailsUsage.test.ts | 5 +-- .../hooks/guardrails/useGuardrailsUsage.ts | 10 ++--- .../GuardrailsMonitor/usageUnits.test.ts | 43 +++++++++---------- 7 files changed, 46 insertions(+), 45 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.test.tsx index 3d00e29245d..9f2a4c42228 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.test.tsx @@ -89,9 +89,8 @@ describe("GuardrailDetail", () => { it("should request the detail and the logs for the guardrail and date range", async () => { renderDetail(); - expect(mockUseGuardrailsUsageDetail).toHaveBeenCalledWith({ + expect(mockUseGuardrailsUsageDetail).toHaveBeenCalledWith("pii-detector", { accessToken: "test-token", - guardrailId: "pii-detector", startDate: "2026-07-01", endDate: "2026-07-24", }); @@ -166,7 +165,10 @@ describe("GuardrailDetail", () => { it("should not request anything without an access token", () => { mockUseGuardrailsUsageDetail.mockReturnValue(loaded(undefined)); renderDetail({ accessToken: null }); - expect(mockUseGuardrailsUsageDetail).toHaveBeenCalledWith(expect.objectContaining({ accessToken: null })); + expect(mockUseGuardrailsUsageDetail).toHaveBeenCalledWith( + "pii-detector", + expect.objectContaining({ accessToken: null }), + ); expect(mockGetGuardrailsUsageLogs).not.toHaveBeenCalled(); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.tsx index 1e82f1fee85..81c39258f67 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.tsx @@ -38,7 +38,7 @@ export function GuardrailDetail({ guardrailId, onBack, accessToken = null, start data: detailData, isLoading: detailLoading, error: detailError, - } = useGuardrailsUsageDetail({ accessToken, guardrailId, startDate, endDate }); + } = useGuardrailsUsageDetail(guardrailId, { accessToken, startDate, endDate }); const { data: logsData, isLoading: logsLoading } = useQuery({ queryKey: ["guardrails-usage-logs", guardrailId, logsPage, logsPageSize], queryFn: () => diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsMonitorView.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsMonitorView.test.tsx index e86b6fa53b6..df106fc38c3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsMonitorView.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsMonitorView.test.tsx @@ -107,7 +107,8 @@ describe("GuardrailsMonitorView", () => { expect(await screen.findByRole("heading", { name: "PII Guard" })).toBeInTheDocument(); expect(mockUseGuardrailsUsageDetail).toHaveBeenCalledWith( - expect.objectContaining({ accessToken: "test-token", guardrailId: "gr-pii", startDate: expect.any(String) }), + "gr-pii", + expect.objectContaining({ accessToken: "test-token", startDate: expect.any(String) }), ); expect(screen.queryByRole("heading", { name: /Guardrails Monitor/i })).not.toBeInTheDocument(); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx index 959ed8b172e..ead5fc9f845 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.test.tsx @@ -20,7 +20,7 @@ vi.mock("./EvaluationSettingsModal", () => ({ EvaluationSettingsModal: ({ open }: { open: boolean }) => (open ?
    Evaluation settings modal
    : null), })); -const row = (overrides: Partial): GuardrailUsageOverviewRow => ({ +const baseRow: GuardrailUsageOverviewRow = { id: "guardrail", name: "Guardrail", type: "content_filter", @@ -34,20 +34,21 @@ const row = (overrides: Partial): GuardrailUsageOverv usageUnits: {}, cost: null, untrackedUsageUnits: {}, - ...overrides, -}); +}; const overview: GuardrailUsageOverview = { rows: [ - row({ + { + ...baseRow, id: "guardrail-low", name: "Low Failure Guardrail", requestsEvaluated: 1200, failRate: 2.5, avgLatency: 45, trend: "down", - }), - row({ + }, + { + ...baseRow, id: "guardrail-high", name: "High Failure Guardrail", provider: "Bedrock", @@ -58,8 +59,9 @@ const overview: GuardrailUsageOverview = { usageUnits: { contentPolicyUnits: 1000, sensitiveInformationPolicyUnits: 250 }, cost: 0.15, untrackedUsageUnits: { sensitiveInformationPolicyUnits: 250 }, - }), - row({ + }, + { + ...baseRow, id: "guardrail-free", name: "Free Bedrock Guardrail", provider: "Bedrock", @@ -67,7 +69,7 @@ const overview: GuardrailUsageOverview = { failRate: 0, usageUnits: { contentPolicyUnits: 40 }, cost: 0, - }), + }, ], chart: [], totalRequests: 1510, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useGuardrailsUsage.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useGuardrailsUsage.test.ts index f0b2709484d..70fc874fb50 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useGuardrailsUsage.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useGuardrailsUsage.test.ts @@ -50,9 +50,8 @@ describe("useGuardrailsUsageDetail", () => { it("queries GET /guardrails/usage/detail/{guardrail_id} with the id as a path param", () => { renderHook(() => - useGuardrailsUsageDetail({ + useGuardrailsUsageDetail("bedrock-pii-mask", { accessToken: "sk", - guardrailId: "bedrock-pii-mask", startDate: "2026-09-01", endDate: "2026-09-04", }), @@ -73,7 +72,7 @@ describe("useGuardrailsUsageDetail", () => { it("stays disabled without a guardrail id", () => { renderHook(() => - useGuardrailsUsageDetail({ accessToken: "sk", guardrailId: "", startDate: "2026-09-01", endDate: "2026-09-04" }), + useGuardrailsUsageDetail("", { accessToken: "sk", startDate: "2026-09-01", endDate: "2026-09-04" }), ); expect(lastCall()[3].enabled).toBe(false); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useGuardrailsUsage.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useGuardrailsUsage.ts index dc7f58fbc8f..5569bdc8beb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useGuardrailsUsage.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/guardrails/useGuardrailsUsage.ts @@ -24,12 +24,10 @@ export const useGuardrailsUsageOverview = ({ accessToken, startDate, endDate }: { enabled: Boolean(accessToken) }, ); -export const useGuardrailsUsageDetail = ({ - accessToken, - guardrailId, - startDate, - endDate, -}: GuardrailsUsageWindow & { guardrailId: string }) => +export const useGuardrailsUsageDetail = ( + guardrailId: string, + { accessToken, startDate, endDate }: GuardrailsUsageWindow, +) => $api.useQuery( "get", "/guardrails/usage/detail/{guardrail_id}", diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts index 8e2baaaa3c5..dd5b564b49a 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/usageUnits.test.ts @@ -9,8 +9,16 @@ import { unitPrice, unitsMathRows, unpricedSummary, + type CounterMath, } from "./usageUnits"; +const counterOf = (counter: string, units: number, unpriced: number, cost: number | null): CounterMath => ({ + counter, + units, + unpriced, + cost, +}); + describe("formatCost", () => { it("renders a dash when nothing was priced", () => { expect(formatCost(null)).toBe("—"); @@ -66,15 +74,12 @@ describe("unpricedSummary", () => { describe("unitPrice", () => { it("backs the per-unit price out of the priced share only", () => { - expect(unitPrice({ counter: "contentPolicyUnits", units: 1200, unpriced: 200, cost: 0.15 })).toBeCloseTo( - 0.00015, - 10, - ); + expect(unitPrice(counterOf("contentPolicyUnits", 1200, 200, 0.15))).toBeCloseTo(0.00015, 10); }); it("is null when nothing was priced", () => { - expect(unitPrice({ counter: "someFutureCounter", units: 7, unpriced: 7, cost: null })).toBeNull(); - expect(unitPrice({ counter: "someFutureCounter", units: 7, unpriced: 7, cost: 0 })).toBeNull(); + expect(unitPrice(counterOf("someFutureCounter", 7, 7, null))).toBeNull(); + expect(unitPrice(counterOf("someFutureCounter", 7, 7, 0))).toBeNull(); }); }); @@ -93,7 +98,7 @@ describe("formatUnitPrice", () => { describe("counterMathRow", () => { it("shows units × price = cost for a fully priced counter", () => { - expect(counterMathRow({ counter: "contentPolicyUnits", units: 1000, unpriced: 0, cost: 0.15 })).toEqual({ + expect(counterMathRow(counterOf("contentPolicyUnits", 1000, 0, 0.15))).toEqual({ label: "Content Policy", parts: ["1,000", "× $0.00015", "= $0.1500"], note: null, @@ -101,20 +106,18 @@ describe("counterMathRow", () => { }); it("prices only the priced share and calls out the rest", () => { - expect(counterMathRow({ counter: "sensitiveInformationPolicyUnits", units: 8, unpriced: 2, cost: 0.0006 })).toEqual( - { - label: "Sensitive Information Policy", - parts: ["6", "× $0.0001", "= $0.0006"], - note: "2 unpriced units left out", - }, + expect(counterMathRow(counterOf("sensitiveInformationPolicyUnits", 8, 2, 0.0006))).toEqual({ + label: "Sensitive Information Policy", + parts: ["6", "× $0.0001", "= $0.0006"], + note: "2 unpriced units left out", + }); + expect(counterMathRow(counterOf("sensitiveInformationPolicyUnits", 8, 1, 0.0007)).note).toBe( + "1 unpriced unit left out", ); - expect( - counterMathRow({ counter: "sensitiveInformationPolicyUnits", units: 8, unpriced: 1, cost: 0.0007 }).note, - ).toBe("1 unpriced unit left out"); }); it("says so when a counter has no known price at all", () => { - expect(counterMathRow({ counter: "someFutureCounter", units: 7, unpriced: 7, cost: null })).toEqual({ + expect(counterMathRow(counterOf("someFutureCounter", 7, 7, null))).toEqual({ label: "Some Future Counter", parts: ["7", "× —", "= —"], note: "no known price, left out", @@ -122,11 +125,7 @@ describe("counterMathRow", () => { }); it("shows a free counter as × $0", () => { - expect(counterMathRow({ counter: "wordPolicyUnits", units: 2, unpriced: 0, cost: 0 }).parts).toEqual([ - "2", - "× $0", - "= $0.0000", - ]); + expect(counterMathRow(counterOf("wordPolicyUnits", 2, 0, 0)).parts).toEqual(["2", "× $0", "= $0.0000"]); }); }); From d56affa81413ab0c04c3e3a94aae3286f3b2f8c7 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Sat, 5 Sep 2026 14:44:21 -0700 Subject: [PATCH 20/35] test(e2e): judge /v1/messages streaming on the clock, not on the provider's delta count The Anthropic and Together AI /v1/messages streaming tests required at least two content_block_delta events. How many deltas a reply is split into is the provider's choice, and Haiku answers a short count in one or two, so the assertion failed on provider variance with no change in the proxy: four of the day's full runs on the PR e2e gate went red on it on 2026-09-05. The harness now stamps when each SSE event reached the client (StreamingResponse.stream_event_arrivals, index-aligned with stream_events, with the clock injectable so the reader has a unit test). Both tests ask for a reply long enough to take seconds to generate and require the first content delta to land at least STREAM_MIN_LEAD_SECONDS before message_stop. A relayed stream shows a lead of about two seconds. A proxy that buffered the response delivers every event in one burst and fails every time, which a whole-response buffering relay in front of a live proxy confirmed. The event-grammar assertions are unchanged. Replay hands the proxy its recorded chunks back to back, so timing says nothing there. The assertion is gated on provider_paces_stream() and replay proves the grammar only, which tests/e2e/CLAUDE.md now says. --- tests/e2e/CLAUDE.md | 2 +- tests/e2e/e2e_config.py | 8 ++ tests/e2e/e2e_http.py | 99 ++++++++++++------- .../e2e/llm_translation/test_messages_e2e.py | 38 ++++--- .../llm_translation/test_together_ai_e2e.py | 29 ++++-- tests/e2e/test_e2e_http.py | 57 ++++++++++- 6 files changed, 173 insertions(+), 60 deletions(-) diff --git a/tests/e2e/CLAUDE.md b/tests/e2e/CLAUDE.md index 14b2b4e3299..89c04208d65 100644 --- a/tests/e2e/CLAUDE.md +++ b/tests/e2e/CLAUDE.md @@ -77,7 +77,7 @@ Mark live tests with `@pytest.mark.e2e` (on the class or the module). Pure cover The seam is `provider_edge.py`: `start_provider_edge` boots an in-process HTTP server (one shared instance per pytest process, `e2e_config.provider_edge_base` is the accessor) that mounts each supported provider under a path prefix (`EDGE_MOUNTS`: `/openai` -> `https://api.openai.com`, `/anthropic` -> `https://api.anthropic.com`). A test participates by registering its deployment with `api_base=provider_edge_base("openai")` plus the provider's path suffix; `quota_management/spend_tracking/test_provider_edge_spend_e2e.py` is the reference. In live mode the accessor returns None and the deployment defaults to the real provider, so an edge-wired test runs in all three modes unchanged. Non-wired tests hit their providers live in every mode. The edge binds `E2E_PROVIDER_EDGE_BIND_HOST` (default 127.0.0.1) and advertises `E2E_PROVIDER_EDGE_ADVERTISE_HOST` in the api_base it hands out, for proxies running in containers -A bundle (default `tests/e2e/.fixtures`, override with `E2E_FIXTURE_DIR`) is a directory: `manifest.json` carries the record timestamp, harness git version, and format version, and each test gets a subdirectory holding one JSON file per provider call in call order (`0000-post-openai-v1-chat-completions.json`). Request headers are never stored (provider credentials never touch disk), non-JSON request bodies store a canonicalized sha256 digest instead of the bytes, `multipart/form-data` bodies store their ordinary fields plus a JSON list of the uploaded parts' `[field, filename, content-type]` triples and a digest of their content, so the per-request random boundary and the envelope never reach the key, and responses store status, filtered headers, and the verbatim body base64-encoded, which is part of why bundles are gitignored. Responses come in two shapes told apart by a `kind` tag: an ordinary one holding a single base64 body, and, for a response the provider streamed (`content-type: text/event-stream`), one holding its transfer chunks in order plus why the stream ended early if it did, so replay reproduces the split points the provider chose instead of one coalesced body. `fixture_bundle.py` owns the format, and `BUNDLE_FORMAT_VERSION` is checked on load, so a bundle recorded under older rules is refused by name rather than partially read. Record serves the proxy the same filtered stored response replay will serve later, chunk for chunk on a stream, so the two modes are byte-identical from the proxy's side of the socket +A bundle (default `tests/e2e/.fixtures`, override with `E2E_FIXTURE_DIR`) is a directory: `manifest.json` carries the record timestamp, harness git version, and format version, and each test gets a subdirectory holding one JSON file per provider call in call order (`0000-post-openai-v1-chat-completions.json`). Request headers are never stored (provider credentials never touch disk), non-JSON request bodies store a canonicalized sha256 digest instead of the bytes, `multipart/form-data` bodies store their ordinary fields plus a JSON list of the uploaded parts' `[field, filename, content-type]` triples and a digest of their content, so the per-request random boundary and the envelope never reach the key, and responses store status, filtered headers, and the verbatim body base64-encoded, which is part of why bundles are gitignored. Responses come in two shapes told apart by a `kind` tag: an ordinary one holding a single base64 body, and, for a response the provider streamed (`content-type: text/event-stream`), one holding its transfer chunks in order plus why the stream ended early if it did, so replay reproduces the split points the provider chose instead of one coalesced body. `fixture_bundle.py` owns the format, and `BUNDLE_FORMAT_VERSION` is checked on load, so a bundle recorded under older rules is refused by name rather than partially read. Record serves the proxy the same filtered stored response replay will serve later, chunk for chunk on a stream, so the two modes are byte-identical from the proxy's side of the socket. Replay does not reproduce the provider's inter-chunk timing (chunks go out as fast as the socket takes them), so a test that judges streaming on the clock, such as the `stream_event_arrivals` lead between the first content delta and `message_stop`, gates that assertion on `provider_paces_stream()` and proves only the event grammar in replay Multipart identity is the fiddly corner, and the rules exist because each one had a collision behind it. A part counts as an upload when it carries a filename or declares its own content type, and everything else is an ordinary field. Field names get a `name[n]` suffix on repeats, with a literal `[` doubled first, so a form that repeats `purpose` never keys the same as one that literally sends `purpose[1]`. A field whose name reads as a credential is stored as ``, which stays key-preserving because the key is recomputed from the stored request rather than saved alongside it, so the live request carrying the real value still matches its redacted fixture. A field value that is not UTF-8 is stored as a base64 sha256 digest, base64 and not hex because the canonicalizer rewrites any 64-character hex run to `` and would fold every binary value onto one key. The uploaded parts contribute a JSON list rather than a `field:filename` string, so a separator inside a filename cannot impersonate a field boundary, and their byte length is stored for a reader's benefit but deliberately left out of the key, since the canonicalizer absorbs timestamp and id drift inside a file that changes its length diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index 21c5a338dc3..09d17b9a0db 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -10,6 +10,7 @@ import os import time import uuid from pathlib import Path +from typing import Final from dotenv import load_dotenv @@ -192,6 +193,13 @@ def provider_edge_base(mount: str) -> str | None: ) +STREAM_MIN_LEAD_SECONDS: Final = 1.0 + + +def provider_paces_stream() -> bool: + return parse_fixture_mode(FIXTURE_MODE_RAW) != "replay" + + def unique_marker() -> str: """A short unique token per call/run, so concurrent runs and the shared response cache never collide on prompts, tags, or customer ids. In record diff --git a/tests/e2e/e2e_http.py b/tests/e2e/e2e_http.py index bc76eb3ea7a..3a4e69cfb43 100644 --- a/tests/e2e/e2e_http.py +++ b/tests/e2e/e2e_http.py @@ -16,9 +16,9 @@ requests itself imports. from __future__ import annotations import time -from collections.abc import Callable +from collections.abc import Callable, Mapping from dataclasses import dataclass -from typing import Generator, Generic, Iterator, Literal, NewType, Protocol, TypeVar, cast +from typing import Final, Generator, Generic, Iterator, Literal, NewType, Protocol, TypeVar, cast import pytest import requests @@ -132,7 +132,12 @@ class StreamingResponse(BaseModel): non-streaming `application/json`), the response headers (lowercased names, e.g. the x-ratelimit-* pacing headers and retry-after on a 429), and the body. SpendLogs.request_id is the completion body id, not call_id. Used by passthrough - and streaming, where one validated JSON model does not fit.""" + and streaming, where one validated JSON model does not fit. + + ``stream_event_arrivals`` is index-aligned with ``stream_events`` and holds the + seconds after the request was sent at which each event reached the client, so a + test can tell a relayed stream (events spread over the provider's generation time) + from a buffered one (every event in one burst at the end).""" status_code: int call_id: str | None = None # x-litellm-call-id header @@ -142,6 +147,7 @@ class StreamingResponse(BaseModel): body: str chunks: int = 0 # streamed events (0 for non-streaming) stream_events: list[str] = [] + stream_event_arrivals: list[float] = [] # First in-stream error event, if any. A streamed call commits its HTTP 200 # before the upstream completes, so upstream failures (e.g. insufficient # quota) arrive as SSE error events inside an otherwise-successful response; @@ -184,7 +190,20 @@ class BinaryStream(BaseModel): return "chunked" in (self.transfer_encoding or "") -def _hdr(resp: requests.Response, name: str) -> str | None: +class SseResponse(Protocol): + @property + def status_code(self) -> int: ... + + @property + def headers(self) -> Mapping[str, str]: ... + + @property + def text(self) -> str: ... + + def iter_lines(self) -> Iterator[bytes]: ... + + +def _hdr(resp: SseResponse, name: str) -> str | None: value = resp.headers.get(name) return value if isinstance(value, str) else None @@ -457,7 +476,7 @@ def probe( return ProbeResult(status_code=resp.status_code, body=resp.text) -def _parse_response_cost(resp: requests.Response) -> float | None: +def _parse_response_cost(resp: SseResponse) -> float | None: raw = _hdr(resp, "x-litellm-response-cost") if raw is None or raw == "": return None @@ -467,11 +486,26 @@ def _parse_response_cost(resp: requests.Response) -> float | None: return None -def _streaming_outcome(resp: requests.Response, stream: bool) -> StreamingResponse: - call_id = _hdr(resp, "x-litellm-call-id") - response_cost = _parse_response_cost(resp) - content_type = _hdr(resp, "content-type") - headers = {name.lower(): value for name, value in resp.headers.items()} +_SSE_DATA_PREFIX: Final = b"data: " +_SSE_DONE: Final = "[DONE]" + + +def _is_stream_error_line(line: bytes) -> bool: + return ( + line.startswith(b"event: error") + or b'"type":"error"' in line + or b'"type": "error"' in line + or line.startswith(b'data: {"error"') + ) + + +def streaming_outcome( + resp: SseResponse, stream: bool, *, sent_at: float, clock: Callable[[], float] = time.monotonic +) -> StreamingResponse: + call_id: Final = _hdr(resp, "x-litellm-call-id") + response_cost: Final = _parse_response_cost(resp) + content_type: Final = _hdr(resp, "content-type") + headers: Final = {name.lower(): value for name, value in resp.headers.items()} if not stream or not (200 <= resp.status_code < 300): return StreamingResponse( status_code=resp.status_code, @@ -481,29 +515,13 @@ def _streaming_outcome(resp: requests.Response, stream: bool) -> StreamingRespon headers=headers, body=resp.text, ) - lines = cast("Iterator[bytes]", resp.iter_lines()) - chunks = 0 - stream_error: str | None = None - stream_events: list[str] = [] - stream_done = False - for line in lines: - if not line: - continue - chunks += 1 - decoded_line = line.decode(errors="replace") - if decoded_line.startswith("data: "): - payload = decoded_line.removeprefix("data: ") - if payload == "[DONE]": - stream_done = True - else: - stream_events.append(payload) - if stream_error is None and ( - line.startswith(b"event: error") - or b'"type":"error"' in line - or b'"type": "error"' in line - or line.startswith(b'data: {"error"') - ): - stream_error = line.decode(errors="replace")[:300] + stamped: Final = tuple((line, clock() - sent_at) for line in resp.iter_lines() if line) + payloads: Final = tuple( + (line.removeprefix(_SSE_DATA_PREFIX).decode(errors="replace"), arrived) + for line, arrived in stamped + if line.startswith(_SSE_DATA_PREFIX) + ) + events: Final = tuple((payload, arrived) for payload, arrived in payloads if payload != _SSE_DONE) return StreamingResponse( status_code=resp.status_code, call_id=call_id, @@ -511,10 +529,14 @@ def _streaming_outcome(resp: requests.Response, stream: bool) -> StreamingRespon content_type=content_type, headers=headers, body="", - chunks=chunks, - stream_events=stream_events, - stream_done=stream_done, - stream_error=stream_error, + chunks=len(stamped), + stream_events=[payload for payload, _ in events], + stream_event_arrivals=[arrived for _, arrived in events], + stream_done=any(payload == _SSE_DONE for payload, _ in payloads), + stream_error=next( + (line.decode(errors="replace")[:300] for line, _ in stamped if _is_stream_error_line(line)), + None, + ), ) @@ -531,6 +553,7 @@ def send( x-litellm-call-id header. For native/passthrough bodies and for calls judged by status rather than a typed JSON model (e.g. a budget block is a non-2xx). With ``stream=True`` the SSE body is consumed and its events counted instead.""" + sent_at: Final = time.monotonic() try: resp = request_with_retry( lambda: requests.post( @@ -544,7 +567,7 @@ def send( ) except requests.RequestException as exc: return StreamingResponse(status_code=-1, body=str(exc)) - return _streaming_outcome(resp, stream) + return streaming_outcome(resp, stream, sent_at=sent_at) def stream( diff --git a/tests/e2e/llm_translation/test_messages_e2e.py b/tests/e2e/llm_translation/test_messages_e2e.py index a5f36a8cbdd..1227b1c7119 100644 --- a/tests/e2e/llm_translation/test_messages_e2e.py +++ b/tests/e2e/llm_translation/test_messages_e2e.py @@ -9,7 +9,12 @@ litellm-regression-tests/tests/test_inference_endpoints.py. from __future__ import annotations import pytest -from e2e_config import provider_edge_base, unique_marker +from e2e_config import ( + STREAM_MIN_LEAD_SECONDS, + provider_edge_base, + provider_paces_stream, + unique_marker, +) from e2e_http import assert_client_error, require_successful_call, unwrap from endpoints_client import EndpointsClient, MessagesResult from lifecycle import ResourceManager @@ -159,19 +164,24 @@ class TestAnthropicMessages: """Edge-wired like its non-streaming siblings, so record and replay both carry the streamed response. - Asserts the shape of the event sequence, not just that deltas and a stop - appeared somewhere in it: the answer arrives across several deltas, and the - usage event sits between the last of them and ``message_stop``. A replay that - coalesced the response into one buffered body could not satisfy either.""" + Asserts what the proxy controls. The event grammar arrives intact: the usage + event sits between the last content delta and ``message_stop``. And the relay + is incremental, judged on the clock rather than by counting deltas: how many + deltas a reply is split into is the provider's choice (Haiku often sends a + 20-line count as one), so a count threshold flaked on provider variance. A + reply that takes seconds to generate must reach the client with its first + delta well before ``message_stop``; a proxy that buffered would deliver every + event in one burst. Replay hands the proxy its recorded chunks back to back, + so only live and record runs can judge the timing.""" model, key = self._register(endpoints_client, resources) result = endpoints_client.proxy.messages_stream( key, AnthropicMessagesBody( model=model, - max_tokens=400, + max_tokens=800, stream=True, - messages=[ChatMessage(role="user", content="Count from 1 to 100, one number per line.")], + messages=[ChatMessage(role="user", content="Count from 1 to 200, one number per line.")], ), ) require_successful_call(result) @@ -186,10 +196,7 @@ class TestAnthropicMessages: delta_positions = [ index for index, event in enumerate(events) if event.type == "content_block_delta" ] - assert len(delta_positions) >= 2, ( - f"stream carried {len(delta_positions)} content deltas, so it was not " - f"incremental: {types}" - ) + assert delta_positions, f"stream carried no content deltas: {types}" text = "".join( event.delta.text for event in events @@ -209,6 +216,15 @@ class TestAnthropicMessages: f"usage did not land between the last content delta and message_stop: {types}" ) + first_delta_at = result.stream_event_arrivals[delta_positions[0]] + stop_at = result.stream_event_arrivals[stop_position] + if provider_paces_stream(): + assert stop_at - first_delta_at >= STREAM_MIN_LEAD_SECONDS, ( + f"first content delta reached the client {first_delta_at:.2f}s after the request " + f"and message_stop {stop_at:.2f}s after it; a relayed stream shows the first delta " + f"at least {STREAM_MIN_LEAD_SECONDS}s before the end, so the response was buffered" + ) + @pytest.mark.covers("llm.messages.anthropic.tool_use.nonstream.works") def test_messages_tool_use( self, endpoints_client: EndpointsClient, resources: ResourceManager diff --git a/tests/e2e/llm_translation/test_together_ai_e2e.py b/tests/e2e/llm_translation/test_together_ai_e2e.py index 2c8a7a3aa20..26daf9040e0 100644 --- a/tests/e2e/llm_translation/test_together_ai_e2e.py +++ b/tests/e2e/llm_translation/test_together_ai_e2e.py @@ -23,7 +23,7 @@ from datetime import date from typing import Final import pytest -from e2e_config import unique_marker +from e2e_config import STREAM_MIN_LEAD_SECONDS, provider_paces_stream, unique_marker from e2e_http import StreamingResponse, require_successful_call, unwrap from lifecycle import ResourceManager from models import ( @@ -81,7 +81,7 @@ PERSON_RESPONSE_FORMAT: dict[str, object] = { } WEATHER_PROMPT = "What is the weather in Paris? Use the tool." WEATHER_REPORT = "Paris: 22 degrees Celsius, clear skies, wind from the northwest at 9 km/h" -COUNTING_PROMPT = "Count from 1 to 20, one number per line." +COUNTING_PROMPT = "Count from 1 to 200, one number per line." WEATHER_TOOL = ChatTool( function=ChatToolFunction( @@ -753,7 +753,7 @@ class TestTogetherMessages: key, AnthropicMessagesBody( model=model, - max_tokens=512, + max_tokens=2048, stream=True, messages=[ChatMessage(role="user", content=COUNTING_PROMPT)], ), @@ -763,11 +763,24 @@ class TestTogetherMessages: assert not result.stream_error, f"stream errored: {result.stream_error}" events = [_MessagesStreamEvent.model_validate_json(event) for event in result.stream_events] types = [event.type for event in events] - text_deltas = [ + delta_positions = [ + index for index, event in enumerate(events) if event.type == "content_block_delta" + ] + assert delta_positions, f"stream carried no content deltas: {types}" + text = "".join( event.delta.text for event in events - if event.type == "content_block_delta" and event.delta is not None and event.delta.text - ] - assert len(text_deltas) >= 2, f"stream was not incremental: {types}" - assert "20" in "".join(text_deltas), f"streamed text lost the answer: {text_deltas}" + if event.type == "content_block_delta" and event.delta is not None + ) + assert "200" in text, f"streamed text lost the answer: {text[:300]!r}" assert "message_stop" in types, f"stream never reached message_stop: {types}" + + stop_position = types.index("message_stop") + first_delta_at = result.stream_event_arrivals[delta_positions[0]] + stop_at = result.stream_event_arrivals[stop_position] + if provider_paces_stream(): + assert stop_at - first_delta_at >= STREAM_MIN_LEAD_SECONDS, ( + f"first content delta reached the client {first_delta_at:.2f}s after the request " + f"and message_stop {stop_at:.2f}s after it; a relayed stream shows the first delta " + f"at least {STREAM_MIN_LEAD_SECONDS}s before the end, so the response was buffered" + ) diff --git a/tests/e2e/test_e2e_http.py b/tests/e2e/test_e2e_http.py index 007801a797a..dc92e31fc5c 100644 --- a/tests/e2e/test_e2e_http.py +++ b/tests/e2e/test_e2e_http.py @@ -12,12 +12,13 @@ monkeypatches anything. from __future__ import annotations -from collections.abc import Callable, Sequence +from collections.abc import Callable, Iterator, Mapping, Sequence from dataclasses import dataclass, field +from types import MappingProxyType import pytest -from e2e_http import RETRY_ATTEMPTS, TRANSIENT_STATUSES, request_with_retry +from e2e_http import RETRY_ATTEMPTS, TRANSIENT_STATUSES, request_with_retry, streaming_outcome @dataclass @@ -80,3 +81,55 @@ class TestTransientRetryPolicy: assert result is responses[RETRY_ATTEMPTS - 1] assert sleep.delays == [0.5, 1.0] assert [r.close_calls for r in responses] == [1, 1, 0, 0] + + +@dataclass(frozen=True, slots=True) +class FakeSseResponse: + lines: Sequence[bytes] + status_code: int = 200 + headers: Mapping[str, str] = MappingProxyType({"content-type": "text/event-stream"}) + text: str = "" + + def iter_lines(self) -> Iterator[bytes]: + return iter(self.lines) + + +def _ticking_clock(start: float, step: float) -> Callable[[], float]: + ticks = iter(range(10_000)) + return lambda: start + step * next(ticks) + + +class TestStreamEventArrivals: + def test_each_event_is_stamped_at_the_moment_its_line_arrives(self) -> None: + resp = FakeSseResponse( + lines=( + b"event: message_start", + b'data: {"type":"message_start"}', + b"", + b"event: ping", + b'data: {"type":"ping"}', + b"event: content_block_delta", + b'data: {"type":"content_block_delta"}', + b"data: [DONE]", + ) + ) + + result = streaming_outcome(resp, True, sent_at=100.0, clock=_ticking_clock(start=100.0, step=0.5)) + + assert result.stream_events == [ + '{"type":"message_start"}', + '{"type":"ping"}', + '{"type":"content_block_delta"}', + ] + assert result.stream_event_arrivals == [0.5, 1.5, 2.5] + assert result.stream_done + assert result.chunks == 7 + + def test_a_non_streaming_outcome_carries_no_arrivals(self) -> None: + resp = FakeSseResponse(lines=(), status_code=400, text="bad request") + + result = streaming_outcome(resp, True, sent_at=0.0, clock=_ticking_clock(start=0.0, step=1.0)) + + assert result.stream_events == [] + assert result.stream_event_arrivals == [] + assert result.body == "bad request" From 0cb759772caccee55fae6ca37dfdd05cf47c1da8 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Sat, 5 Sep 2026 14:49:30 -0700 Subject: [PATCH 21/35] fix(ui): show indirectly granted and name-keyed MCP servers in the tool matrix (#35154) * fix(ui): show indirectly granted and name-keyed MCP servers in the tool matrix The MCP tool permission editor was fed the direct server list only, so a server a principal reaches through an access group or a toolset never appeared in the matrix. That single blind spot produced two opposite bugs depending on how a save handler filtered mcp_tool_permissions: filtering by the selected servers deletes an indirect server's allowlist, and because a missing entry means "no restriction from this level", the principal silently gains every tool on it; not filtering leaves a stale entry that keeps a removed access group's server reachable, since a server named under mcp_tool_permissions is entitled on purpose. The editor now resolves the selected access groups and toolsets to their servers and renders them alongside the direct ones, badged with where the grant comes from, so an admin can see and clear an inherited server's tools like any other. Resolution reuses the data the selector already loads: access groups resolve from each server's mcp_access_groups, toolsets from the toolset's own tool list. When that data cannot be loaded the editor says so instead of rendering an empty list, because an absent inherited server reads as "there are none". Servers named only by an mcp_tool_permissions key are listed too, which is what makes a leftover entry visible; the opt-out sentinel still renders nothing, since it short-circuits the backend resolver to zero servers. Opening the editor no longer applies the delete-blocked-by-default allowlist to an inherited server. Writing an entry for one would narrow a grant the admin never touched just by opening the form; direct servers keep that default. Both components also matched on server_id alone, while the backend accepts a server id, name or alias interchangeably. A grant or allowlist written by API or config with a name rendered as a selected server with no tools under it, which reads as "this server has no tools". Matching now covers all three identifiers, and an edit writes back to the key the entry already uses rather than forking a second id-keyed entry. The same mismatch could also put one server under several keys at once, its id and its name for instance. The backend unions every key's list, so reading one key understated what was in force and writing one key left the others granting. The resolver now reports, per server, the key an edit keeps, the equivalent keys it supersedes, and the union those keys allow; the card renders the union and every write goes through one function that writes the kept key and drops the superseded ones. A key that also names a DIFFERENT server, which happens when two servers share a name, is never dropped, because dropping it would strip the neighbouring server's restriction; the card names such a key and says its tools stay allowed until the servers no longer share the name, so an admin is told rather than left to infer it from an edit that bounces back. A third divergence from the backend sat in the same matching. The backend resolves an identifier with exact-id precedence: a string that is a registry server id names that server and stops, and only a string that is no server's id falls back to name and alias, which can name several. Matching all three fields at once meant a server merely named after another server's id joined the matrix as if it had been selected, and because it landed there as a directly selected server it also received the delete-blocked default write on open. Since an mcp_tool_permissions key is itself a grant source, saving then handed out a server nobody granted, with no admin gesture involved. Identifier resolution now mirrors the backend's precedence, and a key is read as this server's only when it resolves back to it, so an entry that belongs to the id's owner is neither read into this server's allowlist nor overwritten by an edit made against it. A toolset grant was also invisible to the tool matrix. The backend unions a toolset's tools with whatever mcp_tool_permissions allows, so a toolset-only grant restricts the server to that toolset's tools; the editor read the map alone, found no entry and rendered every tool on the server as allowed. Deselecting one from that state wrote all the others as a permission entry, and the union turned a revocation into a grant of every tool the toolset never included. The resolved entry now carries the toolset's tools, so the matrix opens on what is actually in force, the delete-blocked default is withheld from a server a toolset restricts, and a write keeps out the tools only the toolset accounts for so a grant that ends with the toolset does not become a standing one. Those tools cannot be revoked from this screen at all, since the backend unions them in; they render allowed and locked and the card says which of them a toolset holds open and where to go to revoke them. That guard originally covered only the keys an edit supersedes, on the assumption that the key it keeps names one server. It does not when a shared key is a server's only entry: it then becomes the key an edit writes, and writing it moves the other server's allowlist too, which is the widening the guard exists to prevent. The key an edit writes is now the first one naming this server and no other, falling back to the server's own id, so a shared key is never written through and an edit against one card cannot reach the server behind the other. Both cards say the shared key holds tools open, since neither can revoke them. No owner's save handler changes here. With the full effective set now available to the editor, the key and team handlers can filter against it instead of guessing, which makes the internal-user surface's unfiltered save redundant Resolves LIT-4963 Resolves LIT-4958 * chore: drop tsbuildinfo churn from merge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): satisfy dashboard lint budgets Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): keep MCP tool allowlists for indirect grants Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): keep standing MCP grants on team save Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(ui): format TeamInfo and hoist inline object args Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): keep MCP tool allowlists for team servers granted indirectly (#35153) * fix(ui): filter team MCP tool allowlists against the effective server set Saving a team filtered mcp_tool_permissions down to the directly selected servers. A server reached through an access group or a toolset is never in that list, so any save dropped its entry, including a save that only changed the team alias. Because the resolver unions tool-permission keys into the entitled server set and treats a missing entry as "no restriction from this level", the team kept the server and lost the tool allowlist on it Filtering on the direct list alone cannot get this right in either direction. Keeping every entry a level did not directly select leaves a removed access group's server reachable through its own stale entry, which breaks revocation. Dropping on deselection alone widens a server that an access group still supplies The save handler now resolves the effective server set with resolveEffectiveMcpServers and keeps an entry only when something other than the entry itself still grants that server: a direct selection, a selected access group, or a selected toolset. Unified access group ids are added when that selection is untouched, since the loaded server list is then still accurate When the server or toolset list cannot be resolved, every entry is kept and the admin is told the allowlists were saved unchanged. Pruning on incomplete knowledge is the direction that silently widens, so it only happens when the editor can show the server became unreachable. A failed lookup and a changed access group selection are separate cases in a tagged union, so the notice names what actually happened instead of describing the intentional one as a failure, and both hooks gate the filter symmetrically so a save fired before toolsets settle cannot resolve against an empty toolset list Resolves LIT-4961 * fix(ui): resolve team MCP grants from access group metadata and refuse unsafe saves Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): resolve team access group grants from team info when the access group list is role-gated Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): match every selected access group by id instead of by count Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): reload team access group grants at save time Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): keep frontend lint budget within limit Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(ui): cover a standing allowlist no group grant covers at load or save Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(ui): keep MCP grant inputs in named variables for the lint budget Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): guard MCP default write on toolset load, keep create toolsets, fix flat view Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../_components/add_agent_form.test.tsx | 45 ++ .../agents/_components/add_agent_form.tsx | 3 + .../src/components/Teams.test.tsx | 26 + ui/litellm-dashboard/src/components/Teams.tsx | 8 +- .../MCPToolPermissions.test.tsx | 659 +++++++++++++++++- .../MCPToolPermissions.tsx | 204 ++++-- .../effectiveMcpServers.test.ts | 524 ++++++++++++++ .../effectiveMcpServers.ts | 216 ++++++ .../mcp_tools/McpCrudPermissionPanel.tsx | 19 +- .../organisms/create_key_button.tsx | 9 +- .../permissions/MCPServerPermissions.test.tsx | 75 +- .../permissions/MCPServerPermissions.tsx | 25 +- .../src/components/team/TeamInfo.test.tsx | 519 +++++++++++++- .../src/components/team/TeamInfo.tsx | 179 ++++- .../components/templates/key_edit_view.tsx | 11 +- 15 files changed, 2440 insertions(+), 82 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.test.ts create mode 100644 ui/litellm-dashboard/src/components/mcp_server_management/effectiveMcpServers.ts diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx index 4e637344e2c..ccac244f019 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx @@ -24,6 +24,30 @@ vi.mock("./agent_form_fields", () => ({ default: () =>
    , })); +vi.mock("@/components/mcp_server_management/MCPServerSelector", () => ({ + default: ({ + onChange, + }: { + onChange: (selection: { servers: string[]; accessGroups: string[]; toolsets: string[] }) => void; + }) => ( + + ), +})); + +vi.mock("@/components/mcp_server_management/MCPToolPermissions", () => ({ + default: () => null, +})); + +vi.mock("@/components/common_components/team_dropdown", () => ({ + default: () => null, +})); + const a2aInfo: AgentCreateInfo = { agent_type: "a2a", agent_type_display_name: "A2A Agent", @@ -97,4 +121,25 @@ describe("AddAgentForm logos", () => { expect(warnSpy).toHaveBeenCalledTimes(2); warnSpy.mockRestore(); }); + + it("includes selected MCP toolsets in the create payload", async () => { + const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never }); + vi.mocked(networking.createAgentCall).mockResolvedValue({ + agent_id: "agent-1", + agent_name: "Test Agent", + } as never); + vi.mocked(networking.keyListCall).mockResolvedValue({ keys: [] }); + + renderForm(); + await user.click(screen.getByRole("button", { name: "Next →" })); + await user.click(screen.getByTestId("select-mcp-toolset")); + await user.click(screen.getByRole("button", { name: "Next →" })); + await user.click(screen.getByRole("button", { name: "Next →" })); + await user.click(screen.getByText(/Skip for now/)); + await user.click(screen.getByRole("button", { name: "Create Agent →" })); + + await vi.waitFor(() => expect(networking.createAgentCall).toHaveBeenCalled()); + const [, payload] = vi.mocked(networking.createAgentCall).mock.calls[0]; + expect(payload.object_permission).toEqual({ mcp_toolsets: ["ts-1"] }); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx index 51445ec1bdb..108bae977e1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx @@ -361,6 +361,7 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok const objectPermission: Record = { ...(mcpServersAndGroups.servers?.length ? { mcp_servers: mcpServersAndGroups.servers } : {}), ...(mcpServersAndGroups.accessGroups?.length ? { mcp_access_groups: mcpServersAndGroups.accessGroups } : {}), + ...(mcpServersAndGroups.toolsets?.length ? { mcp_toolsets: mcpServersAndGroups.toolsets } : {}), ...(Object.keys(toolPermissions).length ? { mcp_tool_permissions: toolPermissions } : {}), ...(entitlementModels.length ? { models: entitlementModels } : {}), ...(entitlementAgents.length ? { agents: entitlementAgents } : {}), @@ -520,6 +521,8 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok ) => form.setValue("mcp_tool_permissions", toolPerms)} /> diff --git a/ui/litellm-dashboard/src/components/Teams.test.tsx b/ui/litellm-dashboard/src/components/Teams.test.tsx index 2bceb00aae1..c7df35197e6 100644 --- a/ui/litellm-dashboard/src/components/Teams.test.tsx +++ b/ui/litellm-dashboard/src/components/Teams.test.tsx @@ -17,6 +17,22 @@ import { import Teams from "./Teams"; import { chooseSelectOption } from "../../tests/test-utils"; +vi.mock("./mcp_server_management/MCPServerSelector", () => ({ + default: ({ + onChange, + }: { + onChange: (selection: { servers: string[]; accessGroups: string[]; toolsets: string[] }) => void; + }) => ( + + ), +})); + const can = vi.fn(); vi.mock("@/app/(dashboard)/hooks/useCan", () => ({ default: (...args: unknown[]) => can(...args), @@ -1343,6 +1359,16 @@ describe("Teams - the exact bytes the create call sends", () => { }); }); + it("includes selected MCP toolsets in the create object permission", async () => { + await openCreateModal(); + await openSection("MCP Settings", /Allowed MCP Servers/); + fireEvent.click(screen.getByTestId("select-mcp-toolset")); + + const payload = await submit(); + + expect(payload.object_permission).toStrictEqual({ mcp_toolsets: ["ts-1"] }); + }); + it.each([ ["MCP Settings", /Allowed MCP Servers/, ["allowed_mcp_servers_and_groups", "mcp_tool_permissions"]], ["Agent Settings", /Allowed Agents/, ["allowed_agents_and_groups"]], diff --git a/ui/litellm-dashboard/src/components/Teams.tsx b/ui/litellm-dashboard/src/components/Teams.tsx index 4be91f22339..c4060163c78 100644 --- a/ui/litellm-dashboard/src/components/Teams.tsx +++ b/ui/litellm-dashboard/src/components/Teams.tsx @@ -443,6 +443,7 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser (formValues.allowed_mcp_servers_and_groups && (formValues.allowed_mcp_servers_and_groups.servers?.length > 0 || formValues.allowed_mcp_servers_and_groups.accessGroups?.length > 0 || + formValues.allowed_mcp_servers_and_groups.toolsets?.length > 0 || formValues.allowed_mcp_servers_and_groups.toolPermissions)) ) { if (!formValues.object_permission) { @@ -453,13 +454,16 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser delete formValues.allowed_vector_store_ids; } if (formValues.allowed_mcp_servers_and_groups) { - const { servers, accessGroups } = formValues.allowed_mcp_servers_and_groups; + const { servers, accessGroups, toolsets } = formValues.allowed_mcp_servers_and_groups; if (servers && servers.length > 0) { formValues.object_permission.mcp_servers = servers; } if (accessGroups && accessGroups.length > 0) { formValues.object_permission.mcp_access_groups = accessGroups; } + if (toolsets && toolsets.length > 0) { + formValues.object_permission.mcp_toolsets = toolsets; + } delete formValues.allowed_mcp_servers_and_groups; } @@ -1086,6 +1090,8 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser form.setValue("mcp_tool_permissions", toolPerms)} /> diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.test.tsx b/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.test.tsx index 91f1a45f858..69a761b4723 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.test.tsx @@ -2,9 +2,11 @@ import { useState } from "react"; import { describe, it, expect, vi, beforeEach } from "vitest"; import { screen, waitFor } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; -import { renderWithProviders } from "../../../tests/test-utils"; +import { renderWithProviders, testQueryClient } from "../../../tests/test-utils"; import MCPToolPermissions from "./MCPToolPermissions"; import * as networking from "../networking"; +import { NO_MCP_SERVERS_SENTINEL } from "../mcp_tools/constants"; +import type { MCPToolset } from "../mcp_tools/types"; vi.mock("../networking"); @@ -15,6 +17,8 @@ describe("MCPToolPermissions", () => { beforeEach(() => { vi.clearAllMocks(); + testQueryClient.clear(); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([]); }); it("should update tool permissions when user selects a tool", async () => { @@ -71,9 +75,10 @@ describe("MCPToolPermissions", () => { await userEvent.click(screen.getByRole("checkbox", { name: "read_wiki_structure" })); // Verify onChange was called with read_wiki_structure removed - expect(mockOnChange).toHaveBeenCalledWith({ + const expectedToolPermissions = { [mockServerId]: ["read_wiki_contents", "ask_question"], - }); + }; + expect(mockOnChange).toHaveBeenCalledWith(expectedToolPermissions); // Verify API calls // Note: useMCPServers uses useAuthorized() internally, which returns "123" from global mock @@ -184,6 +189,654 @@ describe("MCPToolPermissions", () => { }); }); + describe("servers reached indirectly", () => { + const groupServer = { + server_id: "srv-group-1", + server_name: "Group Server", + alias: "Group Server", + mcp_access_groups: ["production-group"], + }; + const groupTools = [ + { name: "list_issues", description: "List issues" }, + { name: "delete_issue", description: "Delete an issue" }, + ]; + + it("renders the tool matrix for a server granted only through an access group", async () => { + vi.mocked(networking.fetchMCPServers).mockResolvedValue([groupServer]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: groupTools, error: false }); + + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + expect(await screen.findByText("Group Server")).toBeInTheDocument(); + expect(await screen.findByText("list_issues")).toBeInTheDocument(); + expect(screen.getByText("delete_issue")).toBeInTheDocument(); + expect(networking.listMCPTools).toHaveBeenCalledWith(mockAccessToken, groupServer.server_id); + }); + + it("shows every tool selected in flat view for an unrestricted access-group server", async () => { + vi.mocked(networking.fetchMCPServers).mockResolvedValue([groupServer]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: groupTools, error: false }); + + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + await screen.findByText("Group Server"); + await userEvent.click(screen.getByText("Flat List")); + + const [listIssues, deleteIssue] = screen.getAllByRole("checkbox"); + expect(listIssues).toBeChecked(); + expect(deleteIssue).toBeChecked(); + + await userEvent.click(listIssues); + expect(mockOnChange).toHaveBeenCalledWith({ [groupServer.server_id]: ["delete_issue"] }); + }); + + it("marks an access-group server as inherited and leaves a directly selected one unmarked", async () => { + const directServer = { server_id: "srv-direct-1", server_name: "Direct Server", alias: "Direct Server" }; + vi.mocked(networking.fetchMCPServers).mockResolvedValue([directServer, groupServer]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: groupTools, error: false }); + + renderWithProviders( + , + ); + + expect(await screen.findByText("Direct Server")).toBeInTheDocument(); + expect(await screen.findByText("Group Server")).toBeInTheDocument(); + expect(screen.getByText("Via access group: production-group")).toBeInTheDocument(); + expect(screen.queryAllByText(/^Via /)).toHaveLength(1); + }); + + it("renders a toolset server as inherited from that toolset", async () => { + const toolsetServer = { server_id: "srv-toolset-1", server_name: "Toolset Server", alias: "Toolset Server" }; + vi.mocked(networking.fetchMCPServers).mockResolvedValue([toolsetServer]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([ + { + toolset_id: "ts-1", + toolset_name: "Support Toolset", + tools: [{ server_id: toolsetServer.server_id, tool_name: "list_issues" }], + }, + ]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: groupTools, error: false }); + + renderWithProviders( + , + ); + + expect(await screen.findByText("Toolset Server")).toBeInTheDocument(); + expect(screen.getByText("Via toolset: Support Toolset")).toBeInTheDocument(); + }); + + // The backend adds a toolset's tools to whatever mcp_tool_permissions holds, so showing the + // server as unrestricted would invite a deselection that grants every other tool on it. + it("shows a toolset's own tools as the allowed set and locks them", async () => { + const toolsetServer = { server_id: "srv-toolset-1", server_name: "Toolset Server", alias: "Toolset Server" }; + vi.mocked(networking.fetchMCPServers).mockResolvedValue([toolsetServer]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([ + { + toolset_id: "ts-1", + toolset_name: "Support Toolset", + tools: [{ server_id: toolsetServer.server_id, tool_name: "list_issues" }], + }, + ]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: groupTools, error: false }); + + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + expect(await screen.findByText("list_issues")).toBeInTheDocument(); + expect( + screen.getByText( + "list_issues is granted by a selected toolset, so it stays allowed here; edit the toolset to revoke it", + ), + ).toBeInTheDocument(); + + await userEvent.click(screen.getByText("Flat List")); + const [listIssues, deleteIssue] = screen.getAllByRole("checkbox"); + expect(listIssues).toBeChecked(); + expect(listIssues).toBeDisabled(); + expect(deleteIssue).not.toBeChecked(); + + await userEvent.click(listIssues); + expect(mockOnChange).not.toHaveBeenCalled(); + }); + + it("ignores a click on a locked tool in the risk-group view", async () => { + const toolsetServer = { server_id: "srv-toolset-1", server_name: "Toolset Server", alias: "Toolset Server" }; + vi.mocked(networking.fetchMCPServers).mockResolvedValue([toolsetServer]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([ + { + toolset_id: "ts-1", + toolset_name: "Support Toolset", + tools: [{ server_id: toolsetServer.server_id, tool_name: "list_issues" }], + }, + ]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: groupTools, error: false }); + + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + await userEvent.click(await screen.findByText("list_issues")); + expect(mockOnChange).not.toHaveBeenCalled(); + }); + + // Turning a risk group off must not drop a tool the entry grants in its own right, which the + // toolset happens to grant too: that tool outlives the toolset and the admin did not clear it. + it("keeps a locked tool the entry also grants when its risk group is turned off", async () => { + const toolsetServer = { server_id: "srv-toolset-1", server_name: "Toolset Server", alias: "Toolset Server" }; + vi.mocked(networking.fetchMCPServers).mockResolvedValue([toolsetServer]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([ + { + toolset_id: "ts-1", + toolset_name: "Support Toolset", + tools: [{ server_id: toolsetServer.server_id, tool_name: "list_issues" }], + }, + ]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: groupTools, error: false }); + + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + expect(await screen.findByText("list_issues")).toBeInTheDocument(); + // First checkbox is the header toggle of the group holding list_issues. + await userEvent.click(screen.getAllByRole("checkbox")[0]); + + expect(mockOnChange).toHaveBeenCalledWith({ [toolsetServer.server_id]: ["list_issues"] }); + }); + + // Copying the toolset's tools into the entry would outlive the toolset, so a write keeps only + // what this level grants on its own. + it("leaves a toolset's tools out of the entry a Select All writes", async () => { + const toolsetServer = { server_id: "srv-toolset-1", server_name: "Toolset Server", alias: "Toolset Server" }; + vi.mocked(networking.fetchMCPServers).mockResolvedValue([toolsetServer]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([ + { + toolset_id: "ts-1", + toolset_name: "Support Toolset", + tools: [{ server_id: toolsetServer.server_id, tool_name: "list_issues" }], + }, + ]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: groupTools, error: false }); + + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + expect(await screen.findByText("list_issues")).toBeInTheDocument(); + await userEvent.click(screen.getByText("Select All")); + + expect(mockOnChange).toHaveBeenCalledWith({ [toolsetServer.server_id]: ["delete_issue"] }); + }); + + // The default narrows an unrestricted server; against a toolset-restricted one it would widen + // the grant to every non-delete tool the server exposes. + it("does not write the delete-blocked default for a directly selected server a toolset restricts", async () => { + const directServer = { server_id: "srv-direct-1", server_name: "Direct Server", alias: "Direct Server" }; + vi.mocked(networking.fetchMCPServers).mockResolvedValue([directServer]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([ + { + toolset_id: "ts-1", + toolset_name: "Support Toolset", + tools: [{ server_id: directServer.server_id, tool_name: "list_issues" }], + }, + ]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: groupTools, error: false }); + + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + expect(await screen.findByText("list_issues")).toBeInTheDocument(); + expect(mockOnChange).not.toHaveBeenCalled(); + }); + + it("waits for toolsets before writing the delete-blocked default", async () => { + const directServer = { server_id: "srv-direct-1", server_name: "Direct Server", alias: "Direct Server" }; + let resolveToolsets: (toolsets: MCPToolset[]) => void = () => {}; + const pendingToolsets = new Promise((resolve) => { + resolveToolsets = resolve; + }); + vi.mocked(networking.fetchMCPServers).mockResolvedValue([directServer]); + vi.mocked(networking.fetchMCPToolsets).mockReturnValue(pendingToolsets); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: groupTools, error: false }); + + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + await screen.findByText("Direct Server"); + expect(mockOnChange).not.toHaveBeenCalled(); + + resolveToolsets([ + { + toolset_id: "ts-1", + toolset_name: "Support Toolset", + tools: [{ server_id: directServer.server_id, tool_name: "list_issues" }], + }, + ]); + + await screen.findByText("list_issues"); + expect(screen.getByRole("checkbox", { name: "list_issues" })).toHaveAttribute("aria-disabled", "true"); + expect(mockOnChange).not.toHaveBeenCalled(); + }); + + // The backend resolves a selection that is a registry id to that server alone. Rendering the + // server merely named after it would fire the default write against a server nobody granted, + // and a tool-permission entry is itself a grant. + it.each([ + { label: "id owner first", idOwnerFirst: true }, + { label: "name twin first", idOwnerFirst: false }, + ])("does not offer a server merely named after a selected id ($label)", async ({ idOwnerFirst }) => { + const idOwner = { server_id: "srv-collide", server_name: "Payments", alias: "Payments" }; + const nameTwin = { server_id: "srv-twin", server_name: "srv-collide", alias: "srv-collide" }; + vi.mocked(networking.fetchMCPServers).mockResolvedValue(idOwnerFirst ? [idOwner, nameTwin] : [nameTwin, idOwner]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: groupTools, error: false }); + + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + expect(await screen.findByText("Payments")).toBeInTheDocument(); + expect(screen.queryByText("srv-collide")).not.toBeInTheDocument(); + await waitFor(() => { + expect(mockOnChange).toHaveBeenCalledWith({ "srv-collide": ["list_issues"] }); + }); + expect(mockOnChange.mock.calls.every(([written]) => !Object.hasOwn(written, "srv-twin"))).toBe(true); + expect(networking.listMCPTools).not.toHaveBeenCalledWith(mockAccessToken, "srv-twin"); + }); + + it("does not write a default allowlist for an inherited server", async () => { + vi.mocked(networking.fetchMCPServers).mockResolvedValue([groupServer]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: groupTools, error: false }); + + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + expect(await screen.findByText("list_issues")).toBeInTheDocument(); + expect(mockOnChange).not.toHaveBeenCalled(); + }); + + it("keeps blocking delete tools by default for a directly selected server", async () => { + const directServer = { server_id: "srv-direct-1", server_name: "Direct Server", alias: "Direct Server" }; + vi.mocked(networking.fetchMCPServers).mockResolvedValue([directServer]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: groupTools, error: false }); + + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + await waitFor(() => { + expect(mockOnChange).toHaveBeenCalledWith({ [directServer.server_id]: ["list_issues"] }); + }); + }); + + it("shows a server that only a stale tool-permission entry still entitles", async () => { + vi.mocked(networking.fetchMCPServers).mockResolvedValue([groupServer]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: groupTools, error: false }); + + renderWithProviders( + , + ); + + expect(await screen.findByText("Group Server")).toBeInTheDocument(); + expect(screen.getByText("Via tool permissions")).toBeInTheDocument(); + }); + + it("shows nothing for a principal blocked from every MCP server", async () => { + vi.mocked(networking.fetchMCPServers).mockResolvedValue([groupServer]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: groupTools, error: false }); + + const { container } = renderWithProviders( + , + ); + + expect(container).toBeEmptyDOMElement(); + expect(networking.listMCPTools).not.toHaveBeenCalled(); + }); + + it("warns instead of showing no inherited servers when the server list cannot be loaded", async () => { + vi.mocked(networking.fetchMCPServers).mockRejectedValue(new Error("boom")); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([]); + + renderWithProviders( + , + ); + + expect(await screen.findByText("Unable to load MCP servers")).toBeInTheDocument(); + }); + + it("warns when the selected toolsets cannot be resolved to servers", async () => { + vi.mocked(networking.fetchMCPServers).mockResolvedValue([]); + vi.mocked(networking.fetchMCPToolsets).mockRejectedValue(new Error("boom")); + + renderWithProviders( + , + ); + + expect(await screen.findByText("Unable to load toolsets")).toBeInTheDocument(); + }); + }); + + describe("grants keyed by server name", () => { + const namedServer = { + server_id: "1f4bd6c1-0000-4000-8000-000000000001", + server_name: "github_mcp", + alias: "GitHub", + }; + const namedTools = [ + { name: "list_issues", description: "List issues" }, + { name: "delete_issue", description: "Delete an issue" }, + ]; + + it("renders the tool matrix for a grant that names the server instead of its id", async () => { + vi.mocked(networking.fetchMCPServers).mockResolvedValue([namedServer]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: namedTools, error: false }); + + renderWithProviders( + , + ); + + expect(await screen.findByText("github_mcp")).toBeInTheDocument(); + expect(await screen.findByText("list_issues")).toBeInTheDocument(); + expect(screen.getByText("delete_issue")).toBeInTheDocument(); + expect(networking.listMCPTools).toHaveBeenCalledWith(mockAccessToken, namedServer.server_id); + }); + + it("writes an edit back to the name key instead of adding a second id-keyed entry", async () => { + vi.mocked(networking.fetchMCPServers).mockResolvedValue([namedServer]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: namedTools, error: false }); + + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + expect(await screen.findByText("list_issues")).toBeInTheDocument(); + await userEvent.click(screen.getByRole("button", { name: "Deselect All" })); + + expect(mockOnChange).toHaveBeenCalledWith({ github_mcp: [] }); + }); + }); + + describe("a server named by several equivalent keys", () => { + const namedServer = { + server_id: "1f4bd6c1-0000-4000-8000-000000000001", + server_name: "github_mcp", + alias: "GitHub", + mcp_access_groups: ["production-group"], + }; + const namedTools = [ + { name: "list_issues", description: "List issues" }, + { name: "create_issue", description: "Open an issue" }, + { name: "delete_issue", description: "Delete an issue" }, + ]; + + const renderWithBothKeys = (onChange: () => void) => + renderWithProviders( + , + ); + + beforeEach(() => { + vi.mocked(networking.fetchMCPServers).mockResolvedValue([namedServer]); + vi.mocked(networking.fetchMCPToolsets).mockResolvedValue([]); + vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: namedTools, error: false }); + }); + + it("renders one card showing the union both keys grant", async () => { + renderWithBothKeys(vi.fn()); + + expect(await screen.findByText("github_mcp")).toBeInTheDocument(); + expect(screen.getAllByText("github_mcp")).toHaveLength(1); + expect(await screen.findByText("list_issues")).toBeInTheDocument(); + + // Flat view keeps checkbox order identical to the fetched tool order. + await userEvent.click(screen.getByText("Flat List")); + const [listIssues, createIssue, deleteIssue] = screen.getAllByRole("checkbox"); + expect(listIssues).toBeChecked(); + expect(createIssue).toBeChecked(); + expect(deleteIssue).not.toBeChecked(); + }); + + it("removes a deselected tool from every equivalent key, leaving one entry for the server", async () => { + const mockOnChange = vi.fn(); + renderWithBothKeys(mockOnChange); + + expect(await screen.findByText("list_issues")).toBeInTheDocument(); + await userEvent.click(screen.getByText("Flat List")); + await userEvent.click(screen.getAllByRole("checkbox")[0]); + + const written = mockOnChange.mock.calls.at(-1)?.[0] as Record; + expect(Object.keys(written)).toEqual([namedServer.server_id]); + expect(written[namedServer.server_id]).not.toContain("list_issues"); + expect(written[namedServer.server_id]).toContain("create_issue"); + }); + + // Both catalog orders, because a name resolves to two servers here and a first-match + // implementation is only wrong in one of them. + it.each([ + { label: "edited server first", editedFirst: true }, + { label: "twin first", editedFirst: false }, + ])( + "says on the card when a key names another server too, since its tools cannot be revoked here ($label)", + async ({ editedFirst }) => { + const twin = { server_id: "1f4bd6c1-0000-4000-8000-000000000002", server_name: "github_mcp", alias: "Twin" }; + vi.mocked(networking.fetchMCPServers).mockResolvedValue( + editedFirst ? [namedServer, twin] : [twin, namedServer], + ); + + renderWithProviders( + , + ); + + // Both cards say it: the shared key grants on either server and neither card can revoke it, + // so an admin looking at either one has to be told the same thing. + expect( + await screen.findAllByText( + 'Also granted by "github_mcp", which names another server too. Those tools stay allowed here until the servers no longer share that name', + ), + ).toHaveLength(2); + }, + ); + + // The shared key is the twin's only entry, so it would otherwise be the key an edit writes, + // and writing it would move the allowlist of the server the admin is not looking at. + it.each([ + { label: "edited server first", editedFirst: true }, + { label: "twin first", editedFirst: false }, + ])("edits the twin through its own id rather than the shared key ($label)", async ({ editedFirst }) => { + const twin = { server_id: "1f4bd6c1-0000-4000-8000-000000000002", server_name: "github_mcp", alias: "Twin" }; + vi.mocked(networking.fetchMCPServers).mockResolvedValue(editedFirst ? [namedServer, twin] : [twin, namedServer]); + + const mockOnChange = vi.fn(); + renderWithProviders( + , + ); + + // The directly selected twin is the first card; both share the display name "github_mcp". + expect(await screen.findAllByText("list_issues")).toHaveLength(2); + await userEvent.click(screen.getAllByText("Select All")[0]); + + const written = mockOnChange.mock.calls.at(-1)?.[0] as Record; + expect(written["github_mcp"]).toEqual(["list_issues"]); + expect(written[twin.server_id]).toEqual(["list_issues", "create_issue", "delete_issue"]); + }); + + it("says nothing about shared names when every key names one server", async () => { + renderWithBothKeys(vi.fn()); + + expect(await screen.findByText("github_mcp")).toBeInTheDocument(); + expect(screen.queryByText(/names another server too/)).not.toBeInTheDocument(); + }); + + it("badges the server once, by its strongest grant, when a key and a group both name it", async () => { + renderWithProviders( + , + ); + + expect(await screen.findByText("github_mcp")).toBeInTheDocument(); + expect(screen.getByText("Via access group: production-group")).toBeInTheDocument(); + expect(screen.queryByText("Via tool permissions")).not.toBeInTheDocument(); + expect(screen.queryAllByText(/^Via /)).toHaveLength(1); + }); + }); + describe("risk-group (CRUD) view", () => { const crudTools = [ { name: "list_documents", description: "List every document" }, diff --git a/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.tsx b/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.tsx index 9d7c8cd452b..c866e9cc011 100644 --- a/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.tsx +++ b/ui/litellm-dashboard/src/components/mcp_server_management/MCPToolPermissions.tsx @@ -1,28 +1,62 @@ import React, { useEffect, useRef, useState, useMemo } from "react"; import { listMCPTools } from "../networking"; -import { MCPTool, MCPServer } from "../mcp_tools/types"; +import { MCPTool } from "../mcp_tools/types"; import { RadioGroup, RadioGroupItem } from "@/components/ui/radio-group"; import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; import { useMCPServers } from "../../app/(dashboard)/hooks/mcpServers/useMCPServers"; +import { useMCPToolsets } from "../../app/(dashboard)/hooks/mcpServers/useMCPToolsets"; import McpCrudPermissionPanel from "../mcp_tools/McpCrudPermissionPanel"; import { classifyToolOp } from "../../utils/mcpToolCrudClassification"; +import { NO_MCP_SERVERS_SENTINEL } from "../mcp_tools/constants"; +import { + EffectiveMcpServer, + McpGrantSource, + applyToolPermissionWrite, + mcpAllowedToolsFor, + resolveEffectiveMcpServers, +} from "./effectiveMcpServers"; interface MCPToolPermissionsProps { accessToken: string; - selectedServers: string[]; + selectedServers: readonly string[]; + selectedAccessGroups?: readonly string[]; + selectedToolsets?: readonly string[]; toolPermissions: Record; onChange: (toolPermissions: Record) => void; disabled?: boolean; } +const NO_SELECTION: readonly string[] = []; + +interface InheritedBadge { + readonly label: string; + readonly className: string; +} + +const inheritedBadgeFor = (source: McpGrantSource): InheritedBadge | null => { + switch (source.kind) { + case "direct": + return null; + case "accessGroup": + return { label: `Via access group: ${source.name}`, className: "text-green-700 bg-green-50 border-green-200" }; + case "toolset": + return { label: `Via toolset: ${source.name}`, className: "text-purple-700 bg-purple-50 border-purple-200" }; + case "toolPermission": + return { label: "Via tool permissions", className: "text-amber-700 bg-amber-50 border-amber-200" }; + } +}; + const MCPToolPermissions: React.FC = ({ accessToken, selectedServers, + selectedAccessGroups = NO_SELECTION, + selectedToolsets = NO_SELECTION, toolPermissions, onChange, disabled = false, }) => { - const { data: allServers = [] } = useMCPServers(); + const { data: allServers = [], isError: serversFailed, isLoading: serversLoading } = useMCPServers(); + const { data: toolsets = [], isError: toolsetsFailed, isLoading: toolsetsLoading } = useMCPToolsets(); const [serverTools, setServerTools] = useState>({}); const [loadingTools, setLoadingTools] = useState>({}); const [toolErrors, setToolErrors] = useState>({}); @@ -36,15 +70,25 @@ const MCPToolPermissions: React.FC = ({ toolPermissionsRef.current = toolPermissions; }, [toolPermissions]); - // Filter servers based on selectedServers - const servers = useMemo(() => { - if (selectedServers.length === 0) return []; - return allServers.filter((server: MCPServer) => selectedServers.includes(server.server_id)); - }, [allServers, selectedServers]); + // Every server this permission level reaches, not just the directly selected ones: a server + // reached through an access group or a toolset needs its allowlist visible and editable too. + const effectiveMcpInput = { + allServers, + selectedServers, + selectedAccessGroups, + selectedToolsets, + toolsets, + toolPermissions, + }; + const servers = useMemo( + () => resolveEffectiveMcpServers(effectiveMcpInput), + [allServers, selectedServers, selectedAccessGroups, selectedToolsets, toolsets, toolPermissions], + ); // Fetch tools for a specific server; applies delete-blocked-by-default for new servers. // `token` is passed explicitly so the closure never captures a stale accessToken. - const fetchToolsForServer = async (serverId: string, token: string) => { + const fetchToolsForServer = async (entry: EffectiveMcpServer, token: string) => { + const serverId = entry.server.server_id; setLoadingTools((prev) => ({ ...prev, [serverId]: true })); setToolErrors((prev) => ({ ...prev, [serverId]: "" })); @@ -58,14 +102,18 @@ const MCPToolPermissions: React.FC = ({ const fetchedTools: MCPTool[] = response.tools || []; setServerTools((prev) => ({ ...prev, [serverId]: fetchedTools })); - // For servers that have no permissions stored yet, block delete tools by default. + // Default only unrestricted direct servers to non-delete tools. // Read latest permissions from the ref to avoid clobbering concurrent results. const latestPermissions = toolPermissionsRef.current; - if (!latestPermissions[serverId] && fetchedTools.length > 0) { + const isDirect = entry.source.kind === "direct"; + const unrestricted = + mcpAllowedToolsFor(entry.server, latestPermissions, allServers) === undefined && + entry.toolsetTools === undefined; + if (isDirect && unrestricted && (selectedToolsets.length === 0 || !toolsetsFailed) && fetchedTools.length > 0) { const nonDeleteTools = fetchedTools .filter((t) => classifyToolOp(t.name, t.description || "") !== "delete") .map((t) => t.name); - onChange({ ...latestPermissions, [serverId]: nonDeleteTools }); + onChange(applyToolPermissionWrite({ toolPermissions: latestPermissions, entry, allowed: nonDeleteTools })); } } } catch (err) { @@ -79,58 +127,124 @@ const MCPToolPermissions: React.FC = ({ // Auto-fetch tools when servers or accessToken change useEffect(() => { - servers.forEach((server) => { - if (!serverTools[server.server_id] && !loadingTools[server.server_id]) { - fetchToolsForServer(server.server_id, accessToken); + if (toolsetsLoading) return; + servers.forEach((entry) => { + const serverId = entry.server.server_id; + if (!serverTools[serverId] && !loadingTools[serverId]) { + fetchToolsForServer(entry, accessToken); } }); // fetchToolsForServer is defined in this render scope but receives `accessToken` // as an explicit argument, so it is safe to omit from deps here. // eslint-disable-next-line react-hooks/exhaustive-deps - }, [servers, accessToken]); + }, [servers, accessToken, toolsetsLoading]); - const handleCrudPanelChange = (serverId: string, allowed: string[]) => { - onChange({ ...toolPermissions, [serverId]: allowed }); + // Every write goes through here so an edit is authoritative for the SERVER, not for one of the + // equivalent keys that may name it. + const writeAllowedTools = (entry: EffectiveMcpServer, allowed: string[]) => { + onChange(applyToolPermissionWrite({ toolPermissions, entry, allowed })); }; - const handleSelectAll = (serverId: string) => { - const tools = serverTools[serverId] || []; - onChange({ ...toolPermissions, [serverId]: tools.map((t) => t.name) }); + const handleSelectAll = (entry: EffectiveMcpServer) => { + const tools = serverTools[entry.server.server_id] || []; + writeAllowedTools( + entry, + tools.map((t) => t.name), + ); }; - const handleDeselectAll = (serverId: string) => { - onChange({ ...toolPermissions, [serverId]: [] }); - }; + // The opt-out sentinel short-circuits the backend resolver to zero servers, so nothing stored + // here is in force and showing a tool matrix would claim otherwise. + if (selectedServers.includes(NO_MCP_SERVERS_SENTINEL)) { + return null; + } - if (selectedServers.length === 0) { + const selectionSizes = [ + selectedServers.length, + selectedAccessGroups.length, + selectedToolsets.length, + Object.keys(toolPermissions).length, + ]; + if (!selectionSizes.some((size) => size > 0)) { return null; } return (
    - {servers.map((server) => { - const serverName = server.server_name || server.alias || server.server_id; - const tools = serverTools[server.server_id] || []; - const selectedTools = toolPermissions[server.server_id] || []; - const isLoading = loadingTools[server.server_id]; - const error = toolErrors[server.server_id]; - const viewMode = viewModes[server.server_id] ?? "crud"; + {serversFailed && ( +
    +

    Unable to load MCP servers

    +

    + This list is incomplete; servers granted directly or through an access group may be missing. Reload before + changing tool permissions +

    +
    + )} + + {toolsetsFailed && selectedToolsets.length > 0 && ( +
    +

    Unable to load toolsets

    +

    + Servers reached through the selected toolsets are not listed below +

    +
    + )} + + {serversLoading && ( +
    + +

    Loading MCP servers...

    +
    + )} + + {servers.map((entry) => { + const server = entry.server; + const serverId = server.server_id; + const serverName = server.server_name || server.alias || serverId; + const tools = serverTools[serverId] || []; + const selectedTools = entry.allowedTools ?? tools.map((t) => t.name); + const isLoading = loadingTools[serverId]; + const error = toolErrors[serverId]; + const viewMode = viewModes[serverId] ?? "crud"; + const inherited = inheritedBadgeFor(entry.source); + // The backend adds a toolset's tools to whatever this map allows, so these stay on however + // the boxes are ticked. Locking them is what keeps the matrix an honest picture of the grant. + const toolsetTools = entry.toolsetTools ?? []; return ( -
    +
    {/* Header */}
    -

    {serverName}

    +
    +

    {serverName}

    + {inherited && ( + + {inherited.label} + + )} +
    {server.description &&

    {server.description}

    } + {entry.ambiguousKeys.length > 0 && ( +

    + {`Also granted by ${entry.ambiguousKeys.map((key) => `"${key}"`).join(", ")}, which names another server too. Those tools stay allowed here until the servers no longer share that name`} +

    + )} + {toolsetTools.length > 0 && ( +

    + {toolsetTools.length === 1 + ? `${toolsetTools[0]} is granted by a selected toolset, so it stays allowed here; edit the toolset to revoke it` + : `${toolsetTools.join(", ")} are granted by a selected toolset, so they stay allowed here; edit the toolset to revoke them`} +

    + )}
    {!disabled && tools.length > 0 && ( - setViewModes((prev) => ({ ...prev, [server.server_id]: next as "crud" | "flat" })) - } + onValueChange={(next) => setViewModes((prev) => ({ ...prev, [serverId]: next as "crud" | "flat" }))} className="flex w-auto items-center gap-4" >