From 3e287b43a04e2a2a41cd97d70d7b6a910e29268a Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Mon, 10 Aug 2026 14:54:16 -0700 Subject: [PATCH 01/10] docs(terraform): describe the provider release as automatic The runbook still read as a fully manual flow: dispatch the publish workflow by hand, then approve a second gate in the mirror repo. Neither is true now. project-releaser checks the provider changelog on every release except adhoc, nightly included, and dispatches the publish itself when the topmost released heading has moved ahead of the mirror's tags, so cutting the version heading is what ships the provider. The mirror's own release workflow no longer gates, leaving one approval in project-releaser. --- terraform/provider/RELEASING.md | 10 +++++++--- 1 file changed, 7 insertions(+), 3 deletions(-) diff --git a/terraform/provider/RELEASING.md b/terraform/provider/RELEASING.md index 1dc296f29b8..7b359047e2f 100644 --- a/terraform/provider/RELEASING.md +++ b/terraform/provider/RELEASING.md @@ -106,19 +106,23 @@ Before creating a release: 4. **Land the changes in BerriAI/litellm** - Open a PR to `BerriAI/litellm` updating `terraform/provider/CHANGELOG.md` (and any source changes) and merge it. Note the merge commit SHA; the release workflow takes it as `git_ref` + Open a PR to `BerriAI/litellm` updating `terraform/provider/CHANGELOG.md` (and any source changes) and merge it ### 2. Mirror and Tag via project-releaser The provider source lives at `terraform/provider/` in `BerriAI/litellm`; `BerriAI/terraform-provider-litellm` is a thin release mirror. Do not commit or tag the mirror directly +Normally there is nothing to do here. `BerriAI/project-releaser`'s release pipeline runs the same check on every release except `adhoc`, nightly included: it reads the topmost released heading in `terraform/provider/CHANGELOG.md`, probes the mirror for `v`, and dispatches `Publish Terraform provider` only when the changelog has moved ahead of what the mirror carries. Cutting the version heading in step 1 is therefore what releases the provider, and the next release picks it up, so the wait is a day rather than a week + +Dispatch by hand only for an out-of-band release, or to recover a run that failed: + 1. Go to `BerriAI/project-releaser` > **Actions** > `Publish Terraform provider` 2. Click **Run workflow**: - `git_ref`: full 40-char commit SHA from `BerriAI/litellm` to release from - `provider_version`: the new version without the `v` prefix (e.g. `0.3.0`) - `dry_run`: optional; validates without pushing -3. The workflow rsyncs `terraform/provider/` into the mirror repo, commits, and pushes tag `v` -4. The tag push triggers the mirror's `Release` workflow (goreleaser), which is gated by the `production-release` environment approval + +Automatic or manual, the run waits on the `production-release` approval in `project-releaser`, then rsyncs `terraform/provider/` into the mirror repo, commits, and pushes tag `v`. That approval is the only one in the flow. The tag push triggers the mirror's `Release` workflow (goreleaser), which runs unattended **Important**: - Tags must follow the format: `v..` (e.g., `v0.1.2`, `v1.0.0`) From 8f0644e63ff48cae494a0c60378f6fafbff21c0d Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Mon, 10 Aug 2026 15:00:45 -0700 Subject: [PATCH 02/10] fix(ui): scope Virtual Keys and Logs team lists to the caller The Virtual Keys table and the Logs page team filter both asked for every team on the proxy, which /v2/team/list and /team/list reject with a 401 for any role below proxy admin or org admin. Both endpoints answer the same request with the caller's own teams when it carries a user_id, so send one. Only the two unscoped call sites change. The remaining callers either already role-branch or render on surfaces gated to roles the endpoints answer broadly, and scoping those would shrink the list they see: a proxy admin scoped to their own id gets nothing back, and an org admin scoped on /team/list loses the org teams they administer but do not belong to. The shared helper reads the display-form session role rather than all_admin_roles, which mixes display labels with raw role names and so does not match the "Org Admin" value the dashboard actually holds. --- .../(dashboard)/hooks/teams/useTeams.test.ts | 73 +++++++++++++++++++ .../app/(dashboard)/hooks/teams/useTeams.ts | 19 +++-- .../key_team_helpers/filter_helpers.test.ts | 43 ++++++++++- .../key_team_helpers/filter_helpers.ts | 10 ++- .../view_logs/log_filter_logic.test.tsx | 37 ++++++++++ .../components/view_logs/log_filter_logic.tsx | 7 +- ui/litellm-dashboard/src/utils/roles.test.ts | 40 ++++++++++ ui/litellm-dashboard/src/utils/roles.ts | 9 +++ 8 files changed, 225 insertions(+), 13 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.test.ts index 66dfc43cebb..8980c772c9b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.test.ts @@ -820,6 +820,18 @@ describe("useAllTeams", () => { }); const requestedPage = (url: string) => new URLSearchParams(url.split("?")[1]).get("page"); + const requestedUserId = (url: string) => new URLSearchParams(url.split("?")[1]).get("user_id"); + const asRole = (userRole: string, userId = "test-user-id") => + mockUseAuthorized.mockReturnValue({ + accessToken: "test-access-token", + userId, + userRole, + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); it("paginates /v2/team/list to completion and concatenates every page", async () => { fetchMock.mockImplementation((url: string) => @@ -892,4 +904,65 @@ describe("useAllTeams", () => { await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(2)); }); + + it("scopes the request to the caller for an internal user and returns their teams", async () => { + asRole("Internal User", "member-7"); + fetchMock.mockResolvedValue(pageResponse(mockTeams, 1, 1)); + + const { result } = renderHook(() => useAllTeams(), { wrapper }); + + await waitFor(() => expect(result.current.isSuccess).toBe(true)); + + // A scoped call that comes back empty is the failure this guards against: the + // 401 disappears but the page still shows no teams. + expect(result.current.data).toEqual(mockTeams); + expect(result.current.data?.length).toBeGreaterThan(0); + expect(requestedUserId(fetchMock.mock.calls[0][0] as string)).toBe("member-7"); + }); + + it("carries user_id on every page of a scoped multi-page result", async () => { + asRole("Internal Viewer", "member-7"); + fetchMock.mockImplementation((url: string) => + Promise.resolve( + requestedPage(url) === "1" ? pageResponse([mockTeams[0]], 1, 2) : pageResponse([mockTeams[1]], 2, 2), + ), + ); + + const { result } = renderHook(() => useAllTeams(), { wrapper }); + + await waitFor(() => expect(result.current.isSuccess).toBe(true)); + + expect(fetchMock).toHaveBeenCalledTimes(2); + const scopes = fetchMock.mock.calls.map((call) => requestedUserId(call[0] as string)); + expect(scopes).toEqual(["member-7", "member-7"]); + }); + + it.each(["Admin", "Admin Viewer", "Org Admin"])( + "sends no user_id for %s so the broad list is left intact", + async (userRole) => { + asRole(userRole); + fetchMock.mockResolvedValue(pageResponse(mockTeams, 1, 1)); + + const { result } = renderHook(() => useAllTeams(), { wrapper }); + + await waitFor(() => expect(result.current.isSuccess).toBe(true)); + + expect(requestedUserId(fetchMock.mock.calls[0][0] as string)).toBeNull(); + }, + ); + + it("refetches when the scope changes even though the access token has not", async () => { + asRole("Internal User", "member-7"); + fetchMock.mockResolvedValue(pageResponse(mockTeams, 1, 1)); + + const { result, rerender } = renderHook(() => useAllTeams(), { wrapper }); + await waitFor(() => expect(result.current.isSuccess).toBe(true)); + expect(fetchMock).toHaveBeenCalledTimes(1); + + asRole("Internal User", "member-8"); + rerender(); + + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(2)); + expect(requestedUserId(fetchMock.mock.calls[1][0] as string)).toBe("member-8"); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.ts index 4061026b94d..e209a1d7273 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.ts @@ -5,6 +5,7 @@ import { fetchTeams } from "@/app/(dashboard)/networking"; import { createQueryKeys } from "@/app/(dashboard)/hooks/common/queryKeysFactory"; import { teamInfoCall } from "@/components/networking"; import { getProxyBaseUrl, getGlobalLitellmHeaderName, deriveErrorMessage, handleError } from "@/components/networking"; +import { teamListScopeUserId } from "@/utils/roles"; export interface TeamsResponse { teams: Team[]; @@ -116,24 +117,30 @@ export const useTeams = (): UseQueryResult => { const ALL_TEAMS_PAGE_SIZE = 100; -const fetchAllTeamsPaged = async (accessToken: string): Promise => { - const firstPage: TeamsResponse = await teamListCall(accessToken, 1, ALL_TEAMS_PAGE_SIZE); +const fetchAllTeamsPaged = async (accessToken: string, userID: string | null): Promise => { + const firstPage: TeamsResponse = await teamListCall(accessToken, 1, ALL_TEAMS_PAGE_SIZE, { userID }); const totalPages = firstPage.total_pages ?? 1; if (totalPages <= 1) return firstPage.teams; const remainingPages: TeamsResponse[] = await Promise.all( - Array.from({ length: totalPages - 1 }, (_, i) => teamListCall(accessToken, i + 2, ALL_TEAMS_PAGE_SIZE)), + Array.from({ length: totalPages - 1 }, (_, i) => teamListCall(accessToken, i + 2, ALL_TEAMS_PAGE_SIZE, { userID })), ); return [firstPage, ...remainingPages].flatMap((page) => page.teams); }; export const useAllTeams = (): UseQueryResult => { - const { accessToken } = useAuthorized(); + const { accessToken, userId, userRole } = useAuthorized(); + const scopedUserID = teamListScopeUserId(userRole, userId); return useQuery({ queryKey: teamKeys.list({ - filters: { scope: "all", pageSize: ALL_TEAMS_PAGE_SIZE, accessToken: accessToken ?? "" }, + filters: { + scope: "all", + pageSize: ALL_TEAMS_PAGE_SIZE, + accessToken: accessToken ?? "", + userID: scopedUserID ?? "", + }, }), - queryFn: async () => await fetchAllTeamsPaged(accessToken!), + queryFn: async () => await fetchAllTeamsPaged(accessToken!, scopedUserID), enabled: Boolean(accessToken), staleTime: 30000, }); diff --git a/ui/litellm-dashboard/src/components/key_team_helpers/filter_helpers.test.ts b/ui/litellm-dashboard/src/components/key_team_helpers/filter_helpers.test.ts index 637325ad98d..15c45153026 100644 --- a/ui/litellm-dashboard/src/components/key_team_helpers/filter_helpers.test.ts +++ b/ui/litellm-dashboard/src/components/key_team_helpers/filter_helpers.test.ts @@ -1,11 +1,12 @@ -import { describe, expect, it, vi } from "vitest"; -import { fetchTeamFilterOptions } from "./filter_helpers"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { fetchAllTeams, fetchTeamFilterOptions } from "./filter_helpers"; const mockKeyListCall = vi.fn(); +const mockTeamListCall = vi.fn(); vi.mock("@/components/networking", () => ({ keyListCall: (...args: unknown[]) => mockKeyListCall(...args), - teamListCall: vi.fn(), + teamListCall: (...args: unknown[]) => mockTeamListCall(...args), organizationListCall: vi.fn(), })); @@ -78,3 +79,39 @@ describe("fetchTeamFilterOptions", () => { expect(result).toEqual({ keyAliases: [], organizationIds: [], userIds: [] }); }); }); + +describe("fetchAllTeams", () => { + beforeEach(() => { + mockTeamListCall.mockReset(); + }); + + it("forwards the scoping user id to /team/list and returns the rows it answers with", async () => { + mockTeamListCall.mockResolvedValue([{ team_id: "team-a" }, { team_id: "team-b" }]); + + const teams = await fetchAllTeams("tok-123", null, "member-7"); + + expect(mockTeamListCall).toHaveBeenCalledWith("tok-123", null, "member-7"); + expect(teams.map((team) => team.team_id)).toEqual(["team-a", "team-b"]); + }); + + it("sends no user id when the caller is entitled to the broad list", async () => { + mockTeamListCall.mockResolvedValue([]); + + await fetchAllTeams("tok-123"); + + expect(mockTeamListCall).toHaveBeenCalledWith("tok-123", null, null); + }); + + it("keeps the organization filter independent of the scoping user id", async () => { + mockTeamListCall.mockResolvedValue([]); + + await fetchAllTeams("tok-123", "org-1", "member-7"); + + expect(mockTeamListCall).toHaveBeenCalledWith("tok-123", "org-1", "member-7"); + }); + + it("returns an empty list without calling the endpoint when there is no access token", async () => { + expect(await fetchAllTeams(null, null, "member-7")).toEqual([]); + expect(mockTeamListCall).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/key_team_helpers/filter_helpers.ts b/ui/litellm-dashboard/src/components/key_team_helpers/filter_helpers.ts index fb701b4656b..7eef4d3a8b3 100644 --- a/ui/litellm-dashboard/src/components/key_team_helpers/filter_helpers.ts +++ b/ui/litellm-dashboard/src/components/key_team_helpers/filter_helpers.ts @@ -114,9 +114,15 @@ export const fetchTeamFilterOptions = async ( * Fetches all teams across all pages * @param accessToken The access token for API authentication * @param organizationId Optional organization ID to filter teams + * @param userID Scopes the list to that user's teams. Required for roles the endpoint + * does not grant a broad list to; see `teamListScopeUserId` * @returns Array of all teams */ -export const fetchAllTeams = async (accessToken: string | null, organizationId?: string | null): Promise => { +export const fetchAllTeams = async ( + accessToken: string | null, + organizationId?: string | null, + userID?: string | null, +): Promise => { if (!accessToken) return []; try { @@ -125,7 +131,7 @@ export const fetchAllTeams = async (accessToken: string | null, organizationId?: let hasMorePages = true; while (hasMorePages) { - const response = await teamListCall(accessToken, organizationId || null, null); + const response = await teamListCall(accessToken, organizationId || null, userID ?? null); // Add teams from this page allTeams = [...allTeams, ...response]; diff --git a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx index 080bf6b380a..45ef1d017ac 100644 --- a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx @@ -26,6 +26,8 @@ vi.mock("@/components/key_team_helpers/filter_helpers", () => ({ })); import { uiSpendLogsCall } from "../networking"; +import { fetchAllTeams } from "@/components/key_team_helpers/filter_helpers"; +import type { Team } from "../key_team_helpers/key_list"; const emptyResponse: PaginatedResponse = { data: [], @@ -198,6 +200,41 @@ describe("useLogFilterLogic", () => { }); }); + describe("team filter list scope", () => { + const callerTeams = [{ team_id: "team-a" }, { team_id: "team-b" }] as Team[]; + + it("scopes /team/list to an internal user and still surfaces their teams", async () => { + vi.mocked(fetchAllTeams).mockResolvedValue(callerTeams); + + const { result } = renderFilterHook({ userRole: "Internal User", userID: "member-7" }); + + await waitFor(() => expect(fetchAllTeams).toHaveBeenCalled()); + expect(fetchAllTeams).toHaveBeenCalledWith("test-token", null, "member-7"); + // Without the scope the request 401s and the filter falls back to an empty + // list, so the rows matter as much as the argument. + await waitFor(() => expect(result.current.allTeams).toEqual(callerTeams)); + }); + + it("scopes /team/list for an internal viewer", async () => { + vi.mocked(fetchAllTeams).mockResolvedValue(callerTeams); + + renderFilterHook({ userRole: "Internal Viewer", userID: "member-7" }); + + await waitFor(() => expect(fetchAllTeams).toHaveBeenCalledWith("test-token", null, "member-7")); + }); + + it.each(["Admin", "Admin Viewer", "Org Admin"])( + "leaves /team/list unscoped for %s so the broad list survives", + async (userRole) => { + vi.mocked(fetchAllTeams).mockResolvedValue(callerTeams); + + renderFilterHook({ userRole, userID: "member-7" }); + + await waitFor(() => expect(fetchAllTeams).toHaveBeenCalledWith("test-token", null, null)); + }, + ); + }); + it("returns an empty payload and does not crash when the call fails", async () => { vi.mocked(uiSpendLogsCall).mockRejectedValue(new Error("boom")); const { result } = renderFilterHook(); diff --git a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx index 8244066cadb..e1089c6a16c 100644 --- a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx @@ -4,6 +4,7 @@ import type { ColumnFiltersState, PaginationState, SortingState } from "@tanstac import { uiSpendLogsCall } from "../networking"; import { Team } from "../key_team_helpers/key_list"; import { fetchAllTeams } from "../../components/key_team_helpers/filter_helpers"; +import { teamListScopeUserId } from "../../utils/roles"; import { defaultPageSize } from "../constants"; import { LOGS_SORT_FIELD_MAP, type LogEntry, type LogsSortField } from "./columns"; @@ -194,11 +195,13 @@ export function useLogFilterLogic({ total_pages: 0, }; + const teamListUserID = teamListScopeUserId(userRole, userID); + const allTeamsQueryOptions: UseQueryOptions = { - queryKey: ["allTeamsForLogFilters", accessToken], + queryKey: ["allTeamsForLogFilters", accessToken, teamListUserID], queryFn: async () => { if (!accessToken) return []; - const teamsData = await fetchAllTeams(accessToken); + const teamsData = await fetchAllTeams(accessToken, null, teamListUserID); return teamsData || []; }, enabled: !!accessToken, diff --git a/ui/litellm-dashboard/src/utils/roles.test.ts b/ui/litellm-dashboard/src/utils/roles.test.ts index 83f633bc299..209353d3e3d 100644 --- a/ui/litellm-dashboard/src/utils/roles.test.ts +++ b/ui/litellm-dashboard/src/utils/roles.test.ts @@ -1,5 +1,6 @@ import { describe, it, expect } from "vitest"; import { + all_admin_roles, effectiveSessionRole, isAdminRole, isProxyAdminRole, @@ -8,6 +9,7 @@ import { isViewOnlySessionRole, rolesAllowedToViewWriteScopedPages, rolesWithWriteAccess, + teamListScopeUserId, } from "./roles"; import { Team } from "@/components/networking"; @@ -236,4 +238,42 @@ describe("roles", () => { expect(isViewOnlySessionRole("proxy_admin_viewer")).toBe(true); }); }); + + describe("teamListScopeUserId", () => { + const SESSION_USER_ID = "user-1"; + + // The truth table is driven through effectiveSessionRole rather than hand-written + // labels, so it keeps holding if the raw -> display mapping ever moves. + it.each(["proxy_admin", "proxy_admin_viewer", "org_admin"])( + "leaves %s unscoped so the endpoint keeps returning its broad list", + (rawRole) => { + expect(teamListScopeUserId(effectiveSessionRole(rawRole), SESSION_USER_ID)).toBeNull(); + }, + ); + + it.each(["internal_user", "internal_user_viewer", "internal_viewer", "app_user"])( + "scopes %s to its own user id, which is what the endpoint authorizes on", + (rawRole) => { + expect(teamListScopeUserId(effectiveSessionRole(rawRole), SESSION_USER_ID)).toBe(SESSION_USER_ID); + }, + ); + + it("also accepts the Admin Viewer label that formatUserRole emits", () => { + expect(teamListScopeUserId("Admin Viewer", SESSION_USER_ID)).toBeNull(); + }); + + it("scopes an unknown or absent role rather than assuming a broad list", () => { + expect(teamListScopeUserId(null, SESSION_USER_ID)).toBe(SESSION_USER_ID); + expect(teamListScopeUserId("Undefined Role", SESSION_USER_ID)).toBe(SESSION_USER_ID); + }); + + it("keeps Org Admin broad even though all_admin_roles carries only the raw org_admin", () => { + // all_admin_roles mixes display labels with raw role names, so isAdminRole is + // false for the value useAuthorized actually supplies for an org admin. Relying + // on it here would scope org admins down to their direct memberships. + expect(all_admin_roles).not.toContain(effectiveSessionRole("org_admin")); + expect(isAdminRole(effectiveSessionRole("org_admin"))).toBe(false); + expect(teamListScopeUserId(effectiveSessionRole("org_admin"), SESSION_USER_ID)).toBeNull(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/utils/roles.ts b/ui/litellm-dashboard/src/utils/roles.ts index 8d226313f78..17a0ab11824 100644 --- a/ui/litellm-dashboard/src/utils/roles.ts +++ b/ui/litellm-dashboard/src/utils/roles.ts @@ -77,3 +77,12 @@ export const effectiveSessionRole = (rawUserRole?: string): string => { export const isViewOnlySessionRole = (rawUserRole?: string): boolean => viewOnlyRawRoles.includes(rawUserRole?.toLowerCase() ?? ""); + +// Session roles (the value `useAuthorized().userRole` supplies) that /team/list and +// /v2/team/list already answer with a broad list: proxy-wide for admins, org-wide for +// org admins. Sending a user_id for those narrows the response to direct memberships, +// so only the roles the endpoints would otherwise reject carry one. +const sessionRolesWithBroadTeamList: string[] = ["Admin", "Admin Viewer", "Org Admin"]; + +export const teamListScopeUserId = (userRole: string | null, userId: string | null): string | null => + sessionRolesWithBroadTeamList.includes(userRole ?? "") ? null : userId; From 5096fc79274216211ceb45b806411218b95706b4 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Mon, 10 Aug 2026 15:10:08 -0700 Subject: [PATCH 03/10] fix(ui): gate the Old Usage page behind a proxy-admin capability The Old Usage nav entry carried no role restriction, so every role saw it and the page immediately fired eight /global/spend/* requests that the proxy withholds from non-admins, producing a wall of 401s. Gate the nav entry, the page, and both of its mount effects behind a single viewGlobalSpend capability scoped to proxy_admin and proxy_admin_viewer, matching what the backend actually serves. Also drop the session JWT that adminspendByProvider put in the /global/spend/provider query string; the handler never read it. --- .../old-usage/_components/usage.test.tsx | 80 ++++++++++++++++--- .../old-usage/_components/usage.tsx | 34 ++++++-- .../src/components/leftnav.test.tsx | 25 ++++++ .../src/components/leftnav.tsx | 8 +- .../src/components/networking.tsx | 2 - .../src/utils/capabilities.test.ts | 46 +++++++++++ .../src/utils/capabilities.ts | 5 +- 7 files changed, 180 insertions(+), 20 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.test.tsx index e3db50b7300..4cf210c6e38 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.test.tsx @@ -1,6 +1,6 @@ import React from "react"; import { describe, it, expect, vi, beforeEach } from "vitest"; -import { screen, waitFor, within } from "@testing-library/react"; +import { act, screen, waitFor, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { renderWithProviders } from "../../../../../tests/test-utils"; import UsagePage from "./usage"; @@ -49,6 +49,17 @@ const renderUsage = (overrides: Partial> />, ); +// Mount fires two effects whose requests sit behind a promise chain +// (proxy settings, then the spend query). "proves the flush window is wide +// enough" below keeps this honest: it asserts the same flush surfaces those +// requests for an admin, so a denied role's silence means the gate held. +const flushPendingRequests = async () => { + await act(async () => { + await new Promise((resolve) => setTimeout(resolve, 0)); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); +}; + beforeEach(() => { vi.clearAllMocks(); networking.getProxyUISettings.mockResolvedValue(UNLIMITED_SETTINGS); @@ -185,18 +196,67 @@ describe("old usage page", () => { }); }); - describe("as a non-admin", () => { - it("renders only the All Up tab and skips admin-only queries", async () => { - renderUsage({ userRole: "Internal User" }); + // Every role below is served 401 on /global/spend/* by the proxy. Org admins + // and team admins reach the UI as "Internal User" — `org_admin` is an + // organization membership role, never a top-level user_role. + describe.each(["Internal User", "Internal Viewer", "internal_user", "internal_user_viewer", "Org Admin"])( + "as %s", + (userRole) => { + it("shows the admin-only notice instead of the usage dashboard", async () => { + renderUsage({ userRole }); + + expect(await screen.findByText(/Proxy-wide usage is only available to admin users/i)).toBeInTheDocument(); + expect(screen.queryByRole("tab", { name: "All Up" })).not.toBeInTheDocument(); + }); + + it("fires no /global/spend or /global/activity request", async () => { + renderUsage({ userRole }); + + await screen.findByText(/Proxy-wide usage is only available to admin users/i); + await flushPendingRequests(); + + expect(networking.getProxyUISettings).not.toHaveBeenCalled(); + expect(networking.adminSpendLogsCall).not.toHaveBeenCalled(); + expect(networking.adminTopKeysCall).not.toHaveBeenCalled(); + expect(networking.adminTopModelsCall).not.toHaveBeenCalled(); + expect(networking.adminTopEndUsersCall).not.toHaveBeenCalled(); + expect(networking.teamSpendLogsCall).not.toHaveBeenCalled(); + expect(networking.tagsSpendLogsCall).not.toHaveBeenCalled(); + expect(networking.allTagNamesCall).not.toHaveBeenCalled(); + expect(networking.adminspendByProvider).not.toHaveBeenCalled(); + expect(networking.adminGlobalActivity).not.toHaveBeenCalled(); + expect(networking.adminGlobalActivityPerModel).not.toHaveBeenCalled(); + }); + }, + ); + + describe("the admin-only gate", () => { + it("proves the flush window is wide enough to catch a leaked request", async () => { + renderUsage({ userRole: "Admin" }); + + await flushPendingRequests(); + + expect(networking.getProxyUISettings).toHaveBeenCalled(); + expect(networking.adminSpendLogsCall).toHaveBeenCalled(); + expect(networking.tagsSpendLogsCall).toHaveBeenCalled(); + expect(networking.adminGlobalActivity).toHaveBeenCalled(); + }); + + it("still lets an admin through, so the notice is a real gate and not a dead branch", async () => { + renderUsage({ userRole: "Admin" }); expect(await screen.findByRole("tab", { name: "All Up" })).toBeInTheDocument(); - expect(screen.queryByRole("tab", { name: "Team Based Usage" })).not.toBeInTheDocument(); - expect(screen.queryByRole("tab", { name: "Customer Usage" })).not.toBeInTheDocument(); - expect(screen.queryByRole("tab", { name: "Tag Based Usage" })).not.toBeInTheDocument(); - + expect(screen.queryByText(/Proxy-wide usage is only available to admin users/i)).not.toBeInTheDocument(); await waitFor(() => expect(networking.adminSpendLogsCall).toHaveBeenCalled()); - expect(networking.teamSpendLogsCall).not.toHaveBeenCalled(); - expect(networking.adminTopEndUsersCall).not.toHaveBeenCalled(); + }); + + it("does not put the session token in the provider spend query", async () => { + renderUsage({ userRole: "Admin", token: "session-jwt-value" }); + + await waitFor(() => expect(networking.adminspendByProvider).toHaveBeenCalled()); + const callArgs = networking.adminspendByProvider.mock.calls[0]; + expect(callArgs).not.toContain("session-jwt-value"); + expect(callArgs[0]).toBe("sk-test"); }); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.tsx index 3d55f9bb698..5b2f8547822 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.tsx @@ -37,6 +37,7 @@ import { } from "@/components/networking"; import TopKeyView from "@/components/UsagePage/components/EntityUsage/TopKeyView"; import { MoneyCell } from "@/components/shared/table_cells"; +import { hasCapability } from "@/utils/capabilities"; import { formatNumberWithCommas } from "@/utils/dataUtils"; interface UsagePageProps { @@ -90,6 +91,7 @@ const TeamSpendBarList: React.FC<{ data: TeamSpendTotal[] }> = ({ data }) => { }; const UsagePage: React.FC = ({ accessToken, token, userRole, userID, keys, premiumUser }) => { + const canViewGlobalSpend = hasCapability(userRole, "viewGlobalSpend"); const currentDate = new Date(); const [keySpendData, setKeySpendData] = useState([]); const [topKeys, setTopKeys] = useState([]); @@ -155,8 +157,11 @@ const UsagePage: React.FC = ({ accessToken, token, userRole, use }; useEffect(() => { + if (!canViewGlobalSpend) { + return; + } updateTagSpendData(dateValue.from, dateValue.to); - }, [dateValue, selectedTags]); + }, [canViewGlobalSpend, dateValue, selectedTags]); const updateEndUserData = async ( startTime: Date | undefined, @@ -319,10 +324,7 @@ const UsagePage: React.FC = ({ accessToken, token, userRole, use const fetchProviderSpend = () => fetchAndSetData( - () => - accessToken && token - ? adminspendByProvider(accessToken, token, startTime, endTime) - : Promise.reject("No access token or token"), + () => (accessToken ? adminspendByProvider(accessToken, startTime, endTime) : Promise.reject("No access token")), setSpendByProvider, "Error fetching provider spend", ); @@ -467,6 +469,9 @@ const UsagePage: React.FC = ({ accessToken, token, userRole, use useEffect(() => { const initlizeUsageData = async () => { + if (!canViewGlobalSpend) { + return; + } if (accessToken && token && userRole && userID) { const proxy_settings: ProxySettings | undefined = await fetchProxySettings(); if (proxy_settings) { @@ -493,7 +498,24 @@ const UsagePage: React.FC = ({ accessToken, token, userRole, use }; initlizeUsageData(); - }, [accessToken, token, userRole, userID, startTime, endTime]); + }, [canViewGlobalSpend, accessToken, token, userRole, userID, startTime, endTime]); + + if (!canViewGlobalSpend) { + return ( +
+ + + Usage + + +

+ Proxy-wide usage is only available to admin users. Your own usage is on the Usage page. +

+
+
+
+ ); + } if (proxySettings?.DISABLE_EXPENSIVE_DB_QUERIES) { return ( diff --git a/ui/litellm-dashboard/src/components/leftnav.test.tsx b/ui/litellm-dashboard/src/components/leftnav.test.tsx index f795076ff03..04aae64c768 100644 --- a/ui/litellm-dashboard/src/components/leftnav.test.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.test.tsx @@ -6,6 +6,7 @@ import Sidebar, { menuGroups, getBreadcrumb } from "./leftnav"; vi.mock("../utils/roles", () => { return { all_admin_roles: ["admin", "admin_viewer"], + old_admin_roles: ["admin", "admin_viewer"], internalUserRoles: ["internal"], rolesWithWriteAccess: ["admin", "internal"], rolesAllowedToViewWriteScopedPages: ["admin", "internal", "admin_viewer"], @@ -266,6 +267,30 @@ describe("Sidebar (leftnav)", () => { }); expect(screen.queryByText("Prompts")).not.toBeInTheDocument(); }); + + it("should hide Old Usage from internal users while keeping other Experimental children", async () => { + mockUseAuthorized.mockReturnValue(internalAuth); + renderWithProviders(); + + act(() => { + fireEvent.click(screen.getByText("Experimental")); + }); + await waitFor(() => { + expect(screen.getByText("API Playground")).toBeInTheDocument(); + }); + expect(screen.queryByText("Old Usage")).not.toBeInTheDocument(); + }); + + it("should show Old Usage to admins", async () => { + renderWithProviders(); + + act(() => { + fireEvent.click(screen.getByText("Experimental")); + }); + await waitFor(() => { + expect(screen.getByText("Old Usage")).toBeInTheDocument(); + }); + }); }); it("should show Organizations tab for organization admins", () => { diff --git a/ui/litellm-dashboard/src/components/leftnav.tsx b/ui/litellm-dashboard/src/components/leftnav.tsx index 12d124b059e..805428f8cd8 100644 --- a/ui/litellm-dashboard/src/components/leftnav.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.tsx @@ -289,7 +289,13 @@ const menuGroups: MenuGroup[] = [ icon: , roles: all_admin_roles, }, - { key: "4", page: "usage", label: "Old Usage", icon: }, + { + key: "4", + page: "usage", + label: "Old Usage", + icon: , + roles: rolesWithCapability("viewGlobalSpend"), + }, ], }, ], diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 17a5ca37990..25c87560cbd 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -2098,7 +2098,6 @@ export const adminTopEndUsersCall = async ( export const adminspendByProvider = async ( accessToken: string, - keyToken: string | null, startTime: string | undefined, endTime: string | undefined, ) => { @@ -2107,7 +2106,6 @@ export const adminspendByProvider = async ( accessToken, query: { ...(startTime && endTime ? { start_date: startTime, end_date: endTime } : {}), - ...(keyToken ? { api_key: keyToken } : {}), }, }); return data; diff --git a/ui/litellm-dashboard/src/utils/capabilities.test.ts b/ui/litellm-dashboard/src/utils/capabilities.test.ts index 858658309ec..ab3358b80ac 100644 --- a/ui/litellm-dashboard/src/utils/capabilities.test.ts +++ b/ui/litellm-dashboard/src/utils/capabilities.test.ts @@ -1,6 +1,7 @@ import { describe, expect, it } from "vitest"; import { hasCapability, rolesWithCapability } from "./capabilities"; +import { effectiveSessionRole } from "./roles"; describe("hasCapability", () => { it.each(["Admin", "Admin Viewer", "proxy_admin", "proxy_admin_viewer"])( @@ -67,6 +68,51 @@ describe.each(["viewAuditLogs", "viewDeletedTeams"] as const)("hasCapability - % ); }); +// Backend truth table for the `/global/spend/*` routes the Old Usage page calls +// (verified against a live proxy): only proxy_admin and proxy_admin_viewer are +// served. Org admins and team admins carry `internal_user` as their top-level +// user_role, so `effectiveSessionRole` renders them "Internal User" — an org +// admin never reaches the UI as "Org Admin" or `org_admin`. +describe("hasCapability - viewGlobalSpend", () => { + it.each(["Admin", "Admin Viewer", "proxy_admin", "proxy_admin_viewer"])("should grant it to %s", (role) => { + expect(hasCapability(role, "viewGlobalSpend")).toBe(true); + }); + + it.each([ + "Internal User", + "Internal Viewer", + "internal_user", + "internal_user_viewer", + "Org Admin", + "org_admin", + "App User", + "Unknown Role", + "", + null, + undefined, + ])("should deny it to %s", (role) => { + expect(hasCapability(role, "viewGlobalSpend")).toBe(false); + }); + + it("should deny it to every role an org admin or team admin can present at runtime", () => { + const orgAdminSessionRole = effectiveSessionRole("internal_user"); + const teamAdminSessionRole = effectiveSessionRole("internal_user"); + + expect(orgAdminSessionRole).toBe("Internal User"); + expect(hasCapability(orgAdminSessionRole, "viewGlobalSpend")).toBe(false); + expect(hasCapability(teamAdminSessionRole, "viewGlobalSpend")).toBe(false); + }); + + it.each([ + ["proxy_admin", true], + ["proxy_admin_viewer", true], + ["internal_user", false], + ["internal_user_viewer", false], + ] as const)("should match the backend for a %s session", (rawRole, expected) => { + expect(hasCapability(effectiveSessionRole(rawRole), "viewGlobalSpend")).toBe(expected); + }); +}); + describe("rolesWithCapability", () => { it("should return a copy so callers cannot mutate the capability map", () => { const roles = rolesWithCapability("viewToolPolicies"); diff --git a/ui/litellm-dashboard/src/utils/capabilities.ts b/ui/litellm-dashboard/src/utils/capabilities.ts index 8171ef9a512..014bf8530f4 100644 --- a/ui/litellm-dashboard/src/utils/capabilities.ts +++ b/ui/litellm-dashboard/src/utils/capabilities.ts @@ -1,4 +1,6 @@ -import { all_admin_roles } from "./roles"; +import { all_admin_roles, old_admin_roles } from "./roles"; + +const proxyAdminOnlyRoles = [...old_admin_roles, "proxy_admin", "proxy_admin_viewer"]; const CAPABILITY_ROLES = { viewToolPolicies: all_admin_roles, @@ -6,6 +8,7 @@ const CAPABILITY_ROLES = { viewDeletedTeams: all_admin_roles, viewPolicies: all_admin_roles, viewPrompts: all_admin_roles, + viewGlobalSpend: proxyAdminOnlyRoles, } as const satisfies Record; export type Capability = keyof typeof CAPABILITY_ROLES; From 729ec315e2f9041697db0ddb253a9b50aae821be Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Mon, 10 Aug 2026 15:19:39 -0700 Subject: [PATCH 04/10] refactor(ui): make illegal DataTable prop combinations unrepresentable DataTable accepted any mix of its 40-odd props and rejected the incoherent combinations at runtime, from a validator that threw during the first render. A caller only found out it had wired server sorting without a `sorting` prop when the page blew up in front of them. Split the public prop type into mode-keyed unions instead, so the compiler rejects those combinations at the call site. `validateDataTableConfig` and `DataTableConfigError` go away; the component body reads an unchanged flat `DataTableResolvedProps`, which every union member is assignable to, so there is no narrowing inside it. All 44 existing call sites typecheck against the new union unchanged, which `next build` covers. That build only typechecks the app module graph, so the prop type itself needed a gate of its own: `npm run test:types` runs vitest's typecheck mode over `*.test-d.tsx`, and the unit workflow now runs it. The four guards deleted from `DataTable.test.tsx` come back there as compile-time assertions, and loosening the union back to the flat shape fails all five. --- .github/workflows/test-litellm-ui-unit.yml | 5 ++ ui/litellm-dashboard/package.json | 1 + .../shared/DataTable/DataTable.test-d.tsx | 67 ++++++++++++++++ .../shared/DataTable/DataTable.test.tsx | 41 ---------- .../components/shared/DataTable/DataTable.tsx | 68 ++++------------ .../DataTable/DataTableRowSelection.test.tsx | 14 +--- .../src/components/shared/DataTable/index.ts | 3 +- .../src/components/shared/DataTable/types.ts | 80 ++++++++++++++++++- ui/litellm-dashboard/vitest.config.ts | 5 ++ 9 files changed, 175 insertions(+), 109 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test-d.tsx diff --git a/.github/workflows/test-litellm-ui-unit.yml b/.github/workflows/test-litellm-ui-unit.yml index 8f2199017d9..69cbc082d98 100644 --- a/.github/workflows/test-litellm-ui-unit.yml +++ b/.github/workflows/test-litellm-ui-unit.yml @@ -42,6 +42,11 @@ jobs: - name: Install dependencies run: npm ci + - name: Run UI type tests (Vitest) + env: + CI: "true" + run: npm run test:types + - name: Run UI unit tests (Vitest) env: CI: "true" diff --git a/ui/litellm-dashboard/package.json b/ui/litellm-dashboard/package.json index 3953b41a2c2..62bf5fff4b4 100644 --- a/ui/litellm-dashboard/package.json +++ b/ui/litellm-dashboard/package.json @@ -10,6 +10,7 @@ "lint": "eslint .", "test": "vitest", "test:dot": "vitest --reporter=dot", + "test:types": "vitest --run --typecheck.only", "test:watch": "vitest -w", "test:coverage": "vitest run --coverage", "format": "prettier --write .", diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test-d.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test-d.tsx new file mode 100644 index 00000000000..7bbd4f918cd --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test-d.tsx @@ -0,0 +1,67 @@ +import type { ColumnDef, PaginationState, RowSelectionState, SortingState } from "@tanstack/react-table"; + +import { DataTable } from "./DataTable"; + +interface Row { + id: string; + name: string; +} + +const data: Row[] = []; +const columns: ColumnDef[] = []; +const sorting: SortingState = [{ id: "name", desc: false }]; +const pagination: PaginationState = { pageIndex: 0, pageSize: 10 }; +const rowSelection: RowSelectionState = { r1: true }; +const noop = () => {}; + +export const uncontrolled = ; + +export const controlled = ( + +); + +export const clientSortingWithServerPagination = ( + +); + +// @ts-expect-error sortingMode="server" requires `sorting` and `onSortingChange` +export const serverSortingWithoutState = ; + +// @ts-expect-error paginationMode="server" requires `pagination`, `onPaginationChange` and `rowCount` +export const serverPaginationWithoutState = ; + +// @ts-expect-error filterMode="server" requires `columnFilters` and `onColumnFiltersChange` +export const serverFilteringWithoutState = ; + +export const bothSortingSources = ( + // @ts-expect-error `defaultSorting` seeds uncontrolled sorting, so it cannot pair with a controlled `sorting` + +); + +export const selectionWithoutHandler = ( + // @ts-expect-error a controlled `rowSelection` needs `onRowSelectionChange` or selection changes are dropped + +); diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx index 8555a10c326..3afdd2849ae 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx @@ -337,14 +337,6 @@ describe("DataTable filtering", () => { ); expect(names()).toEqual(["Charlie", "Alice", "Bob"]); }); - - it("throws when server filtering is missing required props", () => { - const spy = vi.spyOn(console, "error").mockImplementation(() => {}); - expect(() => render()).toThrow( - /filterMode='server'/, - ); - spy.mockRestore(); - }); }); describe("DataTable loading", () => { @@ -664,36 +656,3 @@ describe("DataTable layout", () => { expect(container.querySelector("thead")?.className).not.toContain("bg-background"); }); }); - -describe("DataTable misconfiguration guards", () => { - it("throws when server sorting is missing required props", () => { - const spy = vi.spyOn(console, "error").mockImplementation(() => {}); - expect(() => render()).toThrow( - /sortingMode='server'/, - ); - spy.mockRestore(); - }); - - it("throws when server pagination is missing required props", () => { - const spy = vi.spyOn(console, "error").mockImplementation(() => {}); - expect(() => render()).toThrow( - /paginationMode='server'/, - ); - spy.mockRestore(); - }); - - it("throws when both defaultSorting and sorting are provided", () => { - const spy = vi.spyOn(console, "error").mockImplementation(() => {}); - expect(() => - render( - , - ), - ).toThrow(/defaultSorting/); - spy.mockRestore(); - }); -}); diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx index 8cd0e25dfc4..2f47887d01f 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx @@ -42,7 +42,15 @@ import { cn } from "@/lib/cva.config"; import "./columnMeta"; import { DataTablePagination, DEFAULT_PAGE_SIZE_OPTIONS } from "./DataTablePagination"; -import type { ColumnPinnedSide, DataTableProps, DataTableSize, FilterMode, PaginationMode, SortingMode } from "./types"; +import type { + ColumnPinnedSide, + DataTableProps, + DataTableResolvedProps, + DataTableSize, + FilterMode, + PaginationMode, + SortingMode, +} from "./types"; const INTERACTIVE_SELECTOR = "button, a, input, select, textarea, [role=checkbox], [data-row-click-exempt]"; @@ -64,47 +72,6 @@ const FILL_CLASSES = { const NO_FILL_CLASSES = { outer: "", frame: "", body: "", header: "" } as const; -export class DataTableConfigError extends Error { - constructor(messages: readonly string[]) { - super(`DataTable misconfiguration:\n- ${messages.join("\n- ")}`); - this.name = "DataTableConfigError"; - } -} - -export function validateDataTableConfig( - props: DataTableProps, -): readonly string[] { - const serverSortingIncomplete = - props.sortingMode === "server" && (props.sorting === undefined || props.onSortingChange === undefined); - - const serverPaginationPropsMissing = - props.pagination === undefined || props.onPaginationChange === undefined || props.rowCount === undefined; - const serverPaginationIncomplete = props.paginationMode === "server" && serverPaginationPropsMissing; - - const serverFilteringIncomplete = - props.filterMode === "server" && (props.columnFilters === undefined || props.onColumnFiltersChange === undefined); - - const bothSortingSources = props.defaultSorting !== undefined && props.sorting !== undefined; - const bothFilterSources = props.defaultColumnFilters !== undefined && props.columnFilters !== undefined; - - const controlledSelectionIncomplete = props.rowSelection !== undefined && props.onRowSelectionChange === undefined; - - return [ - serverSortingIncomplete ? "sortingMode='server' requires both `sorting` and `onSortingChange`." : null, - serverPaginationIncomplete - ? "paginationMode='server' requires `pagination`, `onPaginationChange`, and `rowCount`." - : null, - serverFilteringIncomplete ? "filterMode='server' requires both `columnFilters` and `onColumnFiltersChange`." : null, - bothSortingSources ? "Provide either `defaultSorting` (uncontrolled) or `sorting` (controlled), not both." : null, - bothFilterSources - ? "Provide either `defaultColumnFilters` (uncontrolled) or `columnFilters` (controlled), not both." - : null, - controlledSelectionIncomplete - ? "Controlled `rowSelection` requires `onRowSelectionChange`; without it selection changes are dropped." - : null, - ].filter((message): message is string => message !== null); -} - function columnDefId(column: ColumnDef): string | undefined { if ("id" in column && typeof column.id === "string") { return column.id; @@ -442,7 +409,9 @@ function useControllable( return { value: internal, onChange: setInternal }; } -function useDataTableInstance(props: DataTableProps): Table { +function useDataTableInstance( + props: DataTableResolvedProps, +): Table { const { data, columns, @@ -532,14 +501,7 @@ function useDataTableInstance(props: DataTablePro } export function DataTable(props: DataTableProps) { - // Validate once at construction so a misconfig surfaces immediately instead of on every render. - useState(() => { - const errors = validateDataTableConfig(props); - if (errors.length > 0) { - throw new DataTableConfigError(errors); - } - return null; - }); + const resolved: DataTableResolvedProps = props; const { isLoading = false, @@ -559,9 +521,9 @@ export function DataTable(props: DataTableProps { await user.click(rowBox("m1")); expect(selectedCount()).toHaveTextContent("1"); }); - - it("rejects controlled rowSelection without onRowSelectionChange", () => { - const errors = validateDataTableConfig({ data, columns, rowSelection: { m1: true } }); - - expect(errors).toContain( - "Controlled `rowSelection` requires `onRowSelectionChange`; without it selection changes are dropped.", - ); - }); - - it("does not complain when selection is left uncontrolled", () => { - expect(validateDataTableConfig({ data, columns })).toHaveLength(0); - }); }); diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/index.ts b/ui/litellm-dashboard/src/components/shared/DataTable/index.ts index 62ddd1b0742..39a887ba948 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/index.ts +++ b/ui/litellm-dashboard/src/components/shared/DataTable/index.ts @@ -1,6 +1,6 @@ import "./columnMeta"; -export { DataTable, DataTableConfigError, validateDataTableConfig } from "./DataTable"; +export { DataTable } from "./DataTable"; export { DataTableFilterDrawer, DataTableFilterField, type FilterDraft } from "./DataTableFilterDrawer"; export { DataTablePagination, DEFAULT_PAGE_SIZE_OPTIONS } from "./DataTablePagination"; export { createSelectionColumn } from "./DataTableSelectionColumn"; @@ -17,6 +17,7 @@ export type { ColumnPinnedSide, ColumnResizeMode, DataTableProps, + DataTableResolvedProps, DataTableSize, FilterMode, PaginationMode, diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/types.ts b/ui/litellm-dashboard/src/components/shared/DataTable/types.ts index dd578f4df45..58f0d2c2539 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/types.ts +++ b/ui/litellm-dashboard/src/components/shared/DataTable/types.ts @@ -21,7 +21,11 @@ export type DataTableSize = "compact" | "default"; export type ColumnPinnedSide = "left" | "right"; export type DataTableSkeletonShape = "text" | "twoLine" | "badge" | "chips" | "meter"; -export interface DataTableProps { +/** + * The flat shape the component reads internally. Every member of the public + * `DataTableProps` union is assignable to it, so the component body needs no narrowing. + */ +export interface DataTableResolvedProps { data: TData[]; columns: ColumnDef[]; getRowId?: (row: TData, index: number, parent?: Row) => string; @@ -81,3 +85,77 @@ export interface DataTableProps { paginationSlot?: (table: Table) => React.ReactNode; footer?: (table: Table) => React.ReactNode; } + +type DataTableBaseProps = Omit< + DataTableResolvedProps, + | "sortingMode" + | "sorting" + | "onSortingChange" + | "defaultSorting" + | "paginationMode" + | "pagination" + | "onPaginationChange" + | "rowCount" + | "filterMode" + | "columnFilters" + | "onColumnFiltersChange" + | "defaultColumnFilters" + | "rowSelection" + | "onRowSelectionChange" +>; + +type SortingProps = + | { + sorting: SortingState; + onSortingChange: OnChangeFn; + sortingMode?: SortingMode; + defaultSorting?: never; + } + | { + sortingMode?: Exclude; + sorting?: never; + onSortingChange?: never; + defaultSorting?: SortingState; + }; + +type PaginationProps = + | { + paginationMode: "server"; + pagination: PaginationState; + onPaginationChange: OnChangeFn; + rowCount: number; + } + | { + paginationMode?: Exclude; + pagination?: PaginationState; + onPaginationChange?: OnChangeFn; + rowCount?: number; + }; + +type FilterProps = + | { + columnFilters: ColumnFiltersState; + onColumnFiltersChange: OnChangeFn; + filterMode?: FilterMode; + defaultColumnFilters?: never; + } + | { + filterMode?: Exclude; + columnFilters?: never; + onColumnFiltersChange?: never; + defaultColumnFilters?: ColumnFiltersState; + }; + +type RowSelectionProps = + | { rowSelection: RowSelectionState; onRowSelectionChange: OnChangeFn } + | { rowSelection?: never; onRowSelectionChange?: OnChangeFn }; + +/** + * Public prop type. The mode-keyed unions make the combinations + * `validateDataTableConfig` used to reject at runtime unrepresentable instead. + */ +export type DataTableProps = DataTableBaseProps & + SortingProps & + PaginationProps & + FilterProps & + RowSelectionProps; diff --git a/ui/litellm-dashboard/vitest.config.ts b/ui/litellm-dashboard/vitest.config.ts index 19af41ab8ed..469f7fa3520 100644 --- a/ui/litellm-dashboard/vitest.config.ts +++ b/ui/litellm-dashboard/vitest.config.ts @@ -29,6 +29,7 @@ const config: ViteUserConfig = { exclude: [ "**/*.d.ts", "**/*.test.*", + "**/*.test-d.*", "**/*.spec.*", "tests/**", @@ -45,6 +46,10 @@ const config: ViteUserConfig = { }, exclude: ["node_modules/**"], include: ["src/**/*.test.ts", "src/**/*.test.tsx", "tests/**/*.test.ts", "tests/**/*.test.tsx"], + typecheck: { + include: ["src/**/*.test-d.ts", "src/**/*.test-d.tsx"], + ignoreSourceErrors: true, + }, }, resolve: { alias: { From 3ced0e433ad107593fbfca53d0fad7a681587439 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Mon, 10 Aug 2026 15:21:14 -0700 Subject: [PATCH 05/10] refactor(ui): drop a doc comment naming the deleted DataTable validator --- ui/litellm-dashboard/src/components/shared/DataTable/types.ts | 4 ---- 1 file changed, 4 deletions(-) diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/types.ts b/ui/litellm-dashboard/src/components/shared/DataTable/types.ts index 58f0d2c2539..8c4d7cbb161 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/types.ts +++ b/ui/litellm-dashboard/src/components/shared/DataTable/types.ts @@ -150,10 +150,6 @@ type RowSelectionProps = | { rowSelection: RowSelectionState; onRowSelectionChange: OnChangeFn } | { rowSelection?: never; onRowSelectionChange?: OnChangeFn }; -/** - * Public prop type. The mode-keyed unions make the combinations - * `validateDataTableConfig` used to reject at runtime unrepresentable instead. - */ export type DataTableProps = DataTableBaseProps & SortingProps & PaginationProps & From e7450b11ba562e265650e5543a049270d7ee06f7 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Mon, 10 Aug 2026 15:38:59 -0700 Subject: [PATCH 06/10] test(ui): drop redundant commentary from the team-list scoping tests The removed comments restated the test names and the assertions directly below them. The reasoning they carried is already recorded in the commit that introduced the fix and in the pull request body. --- .../src/app/(dashboard)/hooks/teams/useTeams.test.ts | 2 -- .../src/components/view_logs/log_filter_logic.test.tsx | 2 -- ui/litellm-dashboard/src/utils/roles.test.ts | 5 ----- 3 files changed, 9 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.test.ts index 8980c772c9b..fa3f15124cf 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.test.ts @@ -913,8 +913,6 @@ describe("useAllTeams", () => { await waitFor(() => expect(result.current.isSuccess).toBe(true)); - // A scoped call that comes back empty is the failure this guards against: the - // 401 disappears but the page still shows no teams. expect(result.current.data).toEqual(mockTeams); expect(result.current.data?.length).toBeGreaterThan(0); expect(requestedUserId(fetchMock.mock.calls[0][0] as string)).toBe("member-7"); diff --git a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx index 45ef1d017ac..17d26dc00f3 100644 --- a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx @@ -210,8 +210,6 @@ describe("useLogFilterLogic", () => { await waitFor(() => expect(fetchAllTeams).toHaveBeenCalled()); expect(fetchAllTeams).toHaveBeenCalledWith("test-token", null, "member-7"); - // Without the scope the request 401s and the filter falls back to an empty - // list, so the rows matter as much as the argument. await waitFor(() => expect(result.current.allTeams).toEqual(callerTeams)); }); diff --git a/ui/litellm-dashboard/src/utils/roles.test.ts b/ui/litellm-dashboard/src/utils/roles.test.ts index 209353d3e3d..6430b8277f1 100644 --- a/ui/litellm-dashboard/src/utils/roles.test.ts +++ b/ui/litellm-dashboard/src/utils/roles.test.ts @@ -242,8 +242,6 @@ describe("roles", () => { describe("teamListScopeUserId", () => { const SESSION_USER_ID = "user-1"; - // The truth table is driven through effectiveSessionRole rather than hand-written - // labels, so it keeps holding if the raw -> display mapping ever moves. it.each(["proxy_admin", "proxy_admin_viewer", "org_admin"])( "leaves %s unscoped so the endpoint keeps returning its broad list", (rawRole) => { @@ -268,9 +266,6 @@ describe("roles", () => { }); it("keeps Org Admin broad even though all_admin_roles carries only the raw org_admin", () => { - // all_admin_roles mixes display labels with raw role names, so isAdminRole is - // false for the value useAuthorized actually supplies for an org admin. Relying - // on it here would scope org admins down to their direct memberships. expect(all_admin_roles).not.toContain(effectiveSessionRole("org_admin")); expect(isAdminRole(effectiveSessionRole("org_admin"))).toBe(false); expect(teamListScopeUserId(effectiveSessionRole("org_admin"), SESSION_USER_ID)).toBeNull(); From dc69f6e4a23014182e51b1bca382b0c4545970e1 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Mon, 10 Aug 2026 15:39:31 -0700 Subject: [PATCH 07/10] test(ui): trim rationale comments in the Old Usage gate tests Drop the duplicated org_admin note and shorten the flush-window note to the one line that keeps the liveness test from looking redundant. --- .../app/(dashboard)/old-usage/_components/usage.test.tsx | 9 ++------- ui/litellm-dashboard/src/utils/capabilities.test.ts | 5 ----- 2 files changed, 2 insertions(+), 12 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.test.tsx index 4cf210c6e38..0e4455ee912 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.test.tsx @@ -49,10 +49,7 @@ const renderUsage = (overrides: Partial> />, ); -// Mount fires two effects whose requests sit behind a promise chain -// (proxy settings, then the spend query). "proves the flush window is wide -// enough" below keeps this honest: it asserts the same flush surfaces those -// requests for an admin, so a denied role's silence means the gate held. +// Width of this window is guarded by "proves the flush window is wide enough". const flushPendingRequests = async () => { await act(async () => { await new Promise((resolve) => setTimeout(resolve, 0)); @@ -196,9 +193,7 @@ describe("old usage page", () => { }); }); - // Every role below is served 401 on /global/spend/* by the proxy. Org admins - // and team admins reach the UI as "Internal User" — `org_admin` is an - // organization membership role, never a top-level user_role. + // org_admin is an organization membership role; those users reach the UI as "Internal User". describe.each(["Internal User", "Internal Viewer", "internal_user", "internal_user_viewer", "Org Admin"])( "as %s", (userRole) => { diff --git a/ui/litellm-dashboard/src/utils/capabilities.test.ts b/ui/litellm-dashboard/src/utils/capabilities.test.ts index 56eaecd7a5f..3492223c5e8 100644 --- a/ui/litellm-dashboard/src/utils/capabilities.test.ts +++ b/ui/litellm-dashboard/src/utils/capabilities.test.ts @@ -38,11 +38,6 @@ describe("hasCapability", () => { }); }); -// Backend truth table for the `/global/spend/*` routes the Old Usage page calls -// (verified against a live proxy): only proxy_admin and proxy_admin_viewer are -// served. Org admins and team admins carry `internal_user` as their top-level -// user_role, so `effectiveSessionRole` renders them "Internal User" — an org -// admin never reaches the UI as "Org Admin" or `org_admin`. describe("hasCapability - viewGlobalSpend", () => { it.each(ADMIN_ROLES)("should grant it to %s", (role) => { expect(hasCapability(role, "viewGlobalSpend")).toBe(true); From 4f1c92b9751af9cbdf21d19f4ca70f2742d4d781 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Mon, 10 Aug 2026 15:42:09 -0700 Subject: [PATCH 08/10] fix(ui): register type-test files as knip entry points knip derives its entry points from vitest's `test.include`, which does not cover `test.typecheck.include`, so the new `*.test-d.tsx` file read as an unused file and failed the lint job. Declare the glob as an entry point. Also drops the doc comment on `DataTableResolvedProps`; the rationale for the resolved/public split belongs in the commit that introduced it. --- ui/litellm-dashboard/knip.json | 2 +- ui/litellm-dashboard/src/components/shared/DataTable/types.ts | 4 ---- 2 files changed, 1 insertion(+), 5 deletions(-) diff --git a/ui/litellm-dashboard/knip.json b/ui/litellm-dashboard/knip.json index 48b39e8122d..f6cd8ace112 100644 --- a/ui/litellm-dashboard/knip.json +++ b/ui/litellm-dashboard/knip.json @@ -1,6 +1,6 @@ { "$schema": "https://unpkg.com/knip@5/schema.json", - "entry": ["scripts/**/*.{ts,mjs}", "src/components/ui/**/*.{ts,tsx}"], + "entry": ["scripts/**/*.{ts,mjs}", "src/components/ui/**/*.{ts,tsx}", "src/**/*.test-d.{ts,tsx}"], "project": ["src/**/*.{ts,tsx}", "tests/**/*.{ts,tsx}", "scripts/**/*.{ts,mjs}"], "ignore": ["src/lib/http/schema.d.ts"], "ignoreDependencies": [ diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/types.ts b/ui/litellm-dashboard/src/components/shared/DataTable/types.ts index 8c4d7cbb161..c767a0a64c0 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/types.ts +++ b/ui/litellm-dashboard/src/components/shared/DataTable/types.ts @@ -21,10 +21,6 @@ export type DataTableSize = "compact" | "default"; export type ColumnPinnedSide = "left" | "right"; export type DataTableSkeletonShape = "text" | "twoLine" | "badge" | "chips" | "meter"; -/** - * The flat shape the component reads internally. Every member of the public - * `DataTableProps` union is assignable to it, so the component body needs no narrowing. - */ export interface DataTableResolvedProps { data: TData[]; columns: ColumnDef[]; From 5c1623888ec5e0fb37fffc771f1e5e381082e705 Mon Sep 17 00:00:00 2001 From: yucheng-berri Date: Mon, 10 Aug 2026 16:37:12 -0700 Subject: [PATCH 09/10] fix(arize): trace MCP tool calls instead of crashing on CallToolResult (#36453) * fix(arize): stop MCP CallToolResult from aborting span attribute setting `call_mcp_tool` logs the MCP SDK's `CallToolResult`, a Pydantic model with no `.get`. `_coerce_response_obj_for_attrs` left it untouched and `_set_request_attributes` then raised AttributeError, which aborted the rest of the attribute block, so MCP tool spans lost their invocation params, input messages, and outputs. Dump Pydantic models that lack `.get` to a dict, and guard the response id/model reads the same way `_set_response_attributes` already does so any other uncoercible response object degrades instead of crashing. * feat(arize): render MCP tool calls as OpenInference TOOL spans `call_mcp_tool` spans carry neither `messages` nor `choices`, so every generic extraction path left Input and Output blank and the span showed only provider/model metadata. Emit `tool.name` from `metadata.mcp_tool_call_metadata`, `input.value` from the tool arguments, and `output.value` from the `CallToolResult` content (text parts when present, JSON otherwise). Arguments and results are user content, so the input/output emit is gated on the same `should_redact_message_logging` check the passthrough normalizer uses. Reuse `_to_plain_dict` for the Pydantic coercion instead of the local BaseModel branch added in the previous commit. * fix(arize): annotate the new MCP helper parameters The strict-rule gate flagged three new ANN001 violations. Type the payload as StandardLoggingPayload | None and the coerced response as object, which the isinstance guards already narrow. * fix(arize): annotate the MCP helper against the type-discipline gate LIT001 bans mutable collections in annotations, so the kwargs parameter becomes Mapping[str, object]. should_redact_message_logging still declares a dict it only ever reads, and widening it would cascade into core_helpers, so the call carries a scoped ignore instead. Narrow the payload by None rather than isinstance now that it is typed, and annotate the values read out of the untyped logging payload. * fix(arize): record empty MCP arguments and results instead of dropping them Zero-argument tools record arguments={} and successful calls can return content=[]; both were skipped by truthiness, leaving the generic placeholder on Input and nothing on Output. Read structuredContent when content yields no text, and cover the list_mcp_tools response shape. * fix(arize): keep media parts in mixed MCP results A result mixing text and media returned the text alone, so Arize showed text/plain and dropped the image or resource parts. --------- Co-authored-by: Sean Lee --- litellm/integrations/arize/_utils.py | 79 ++++- .../integrations/arize/test_arize_utils.py | 323 ++++++++++++++++++ 2 files changed, 399 insertions(+), 3 deletions(-) diff --git a/litellm/integrations/arize/_utils.py b/litellm/integrations/arize/_utils.py index 8c494794858..e7e1ab538d5 100644 --- a/litellm/integrations/arize/_utils.py +++ b/litellm/integrations/arize/_utils.py @@ -1,4 +1,5 @@ import json +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final from typing_extensions import override @@ -12,7 +13,7 @@ from litellm.litellm_core_utils.redact_messages import ( should_redact_message_logging, ) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps -from litellm.types.utils import StandardLoggingPayload +from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall, StandardLoggingPayload if TYPE_CHECKING: from opentelemetry.trace import Span @@ -22,6 +23,7 @@ from litellm.integrations._types.open_inference import ( ImageAttributes, MessageAttributes, MessageContentAttributes, + OpenInferenceMimeTypeValues, OpenInferenceSpanKindValues, SpanAttributes, ToolCallAttributes, @@ -480,6 +482,7 @@ def set_attributes(span: "Span", kwargs, response_obj, attributes: type[BaseLLMO response_obj_for_attrs, slp, ) + _safe_emit("mcp tool attrs", _maybe_set_mcp_tool_attrs, span, kwargs, slp, response_obj_for_attrs) def _sanitize_optional_params(optional_params: dict | None) -> dict: @@ -538,9 +541,12 @@ def _set_request_attributes( if optional_params.get("user"): safe_set_attribute(span, "llm.user", optional_params.get("user")) - if response_obj and response_obj.get("id"): + if not hasattr(response_obj, "get"): + return + + if response_obj.get("id"): safe_set_attribute(span, "llm.response.id", response_obj.get("id")) - if response_obj and response_obj.get("model"): + if response_obj.get("model"): safe_set_attribute(span, "llm.response.model", response_obj.get("model")) @@ -588,6 +594,8 @@ def _coerce_response_obj_for_attrs(response_obj): - dicts and Pydantic models that already expose `.get` are returned unchanged (preserves all current behavior, including the Responses API flow which relies on Pydantic attribute access). + - Pydantic models without `.get` (e.g. the MCP SDK's `CallToolResult`, + logged for `call_mcp_tool` spans) are dumped to a dict. - `httpx.Response` and other text-only responses (passthrough routes) are JSON-decoded so the standard extraction paths can read fields like `id`, `model`, and `usage`. On failure the original object is returned @@ -595,6 +603,9 @@ def _coerce_response_obj_for_attrs(response_obj): """ if response_obj is None or hasattr(response_obj, "get"): return response_obj + dumped: Final = _to_plain_dict(response_obj) + if isinstance(dumped, dict): + return dumped text: Final = getattr(response_obj, "text", None) if isinstance(text, str) and text: try: @@ -1058,3 +1069,65 @@ def _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs): except Exception: return None return None + + +def _maybe_set_mcp_tool_attrs( + span: "Span", + kwargs: Mapping[str, object], + standard_logging_payload: StandardLoggingPayload | None, + coerced_response_obj: object, +) -> None: + """Render `call_mcp_tool` spans as OpenInference TOOL spans. + + MCP tool calls carry neither `messages` nor `choices`, so the generic + extraction paths leave Input/Output blank. The tool name and arguments live + in `metadata.mcp_tool_call_metadata`; the result is an MCP `CallToolResult` + whose `content` is a list of typed parts. + """ + if standard_logging_payload is None: + return + if standard_logging_payload.get("call_type") != CallTypes.call_mcp_tool.value: + return + + metadata: Final = standard_logging_payload.get("metadata") + mcp_meta: Final[StandardLoggingMCPToolCall | None] = metadata.get("mcp_tool_call_metadata") if metadata else None + if mcp_meta is None: + return + + tool_name: Final = mcp_meta.get("name") or mcp_meta.get("namespaced_tool_name") + if tool_name: + safe_set_attribute(span, SpanAttributes.TOOL_NAME, tool_name) + + if should_redact_message_logging(kwargs): # pyright: ignore[reportArgumentType] # reads, never mutates + return + + arguments: Final[object] = mcp_meta.get("arguments") + if arguments is not None: + safe_set_attribute(span, SpanAttributes.INPUT_VALUE, safe_dumps(arguments)) + safe_set_attribute(span, SpanAttributes.INPUT_MIME_TYPE, OpenInferenceMimeTypeValues.JSON.value) + + _set_mcp_tool_output(span, coerced_response_obj) + + +def _has_only_text_parts(content: object) -> bool: + return not isinstance(content, list) or all(_coerce_text([part]) is not None for part in content) + + +def _set_mcp_tool_output(span: "Span", coerced_response_obj: object) -> None: + if not isinstance(coerced_response_obj, Mapping): + return + + content: Final[object] = coerced_response_obj.get("content") + text: Final[str | None] = _coerce_text(content) + if text and _has_only_text_parts(content): + safe_set_attribute(span, SpanAttributes.OUTPUT_VALUE, text) + safe_set_attribute(span, SpanAttributes.OUTPUT_MIME_TYPE, OpenInferenceMimeTypeValues.TEXT.value) + return + + structured: Final[object] = coerced_response_obj.get("structuredContent") + payload: Final[object] = content if content else structured if structured is not None else content + if payload is None: + return + + safe_set_attribute(span, SpanAttributes.OUTPUT_VALUE, safe_dumps(payload)) + safe_set_attribute(span, SpanAttributes.OUTPUT_MIME_TYPE, OpenInferenceMimeTypeValues.JSON.value) diff --git a/tests/test_litellm/integrations/arize/test_arize_utils.py b/tests/test_litellm/integrations/arize/test_arize_utils.py index 83c3351319a..b02fe35cad0 100644 --- a/tests/test_litellm/integrations/arize/test_arize_utils.py +++ b/tests/test_litellm/integrations/arize/test_arize_utils.py @@ -1193,3 +1193,326 @@ def test_arize_coerce_response_obj_returns_original_on_bad_json(): obj = BadJson() assert _coerce_response_obj_for_attrs(obj) is obj + + +def test_arize_mcp_call_tool_result_does_not_break_attribute_setting(): + """`call_mcp_tool` logs the MCP SDK's `CallToolResult`, a Pydantic model + with no `.get`. It used to raise inside `_set_request_attributes`, aborting + the whole attribute block (input messages, invocation params, outputs).""" + from unittest.mock import MagicMock + + from mcp.types import CallToolResult, TextContent + + span = MagicMock() + kwargs = { + "model": "MCP: get_weather", + "standard_logging_object": { + "model_parameters": {"user": "u-1"}, + "metadata": {}, + "call_type": "call_mcp_tool", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "mcp"}, + } + response_obj = CallToolResult( + content=[TextContent(type="text", text="sunny, 21C")], isError=False + ) + + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + + span.record_exception.assert_not_called() + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.OPENINFERENCE_SPAN_KIND] == "TOOL" + assert written["llm.request.type"] == "call_mcp_tool" + # Emitted after the old crash point, so absent before the fix. + assert written[SpanAttributes.LLM_INVOCATION_PARAMETERS] == '{"user": "u-1"}' + assert written[SpanAttributes.USER_ID] == "u-1" + + +def test_arize_coerce_response_obj_dumps_pydantic_without_get(): + from mcp.types import CallToolResult, TextContent + + from litellm.integrations.arize._utils import _coerce_response_obj_for_attrs + + result = CallToolResult(content=[TextContent(type="text", text="hi")], isError=False) + coerced = _coerce_response_obj_for_attrs(result) + + assert isinstance(coerced, dict) + assert coerced["isError"] is False + assert coerced["content"][0]["text"] == "hi" + + +def test_arize_request_attributes_survive_uncoercible_response_obj(): + """A response object that is neither dict-like nor coercible (binary + passthrough body, SDK object) must not abort attribute setting.""" + from unittest.mock import MagicMock + + span = MagicMock() + kwargs = { + "model": "gpt-4o", + "standard_logging_object": { + "model_parameters": {}, + "metadata": {}, + "call_type": "completion", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + } + + class Opaque: + pass + + ArizeLogger.set_arize_attributes(span, kwargs, Opaque()) + + span.record_exception.assert_not_called() + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written["llm.provider"] == "openai" + + +def _mcp_kwargs(mcp_tool_call_metadata=None, **overrides): + return { + "model": "MCP: get_weather", + "standard_logging_object": { + "model_parameters": {}, + "metadata": { + "mcp_tool_call_metadata": mcp_tool_call_metadata + or { + "name": "get_weather", + "arguments": {"city": "Seoul"}, + "namespaced_tool_name": "weather-mcp/get_weather", + } + }, + "call_type": "call_mcp_tool", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "mcp"}, + **overrides, + } + + +def test_arize_mcp_tool_span_renders_name_input_and_output(): + """`call_mcp_tool` spans have no messages/choices, so Input and Output came + out blank. Render them from mcp_tool_call_metadata + CallToolResult.""" + from unittest.mock import MagicMock + + from mcp.types import CallToolResult, TextContent + + span = MagicMock() + response_obj = CallToolResult( + content=[TextContent(type="text", text="sunny, 21C")], isError=False + ) + + ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.TOOL_NAME] == "get_weather" + assert written[SpanAttributes.INPUT_VALUE] == '{"city": "Seoul"}' + assert written[SpanAttributes.INPUT_MIME_TYPE] == "application/json" + assert written[SpanAttributes.OUTPUT_VALUE] == "sunny, 21C" + assert written[SpanAttributes.OUTPUT_MIME_TYPE] == "text/plain" + + +def test_arize_mcp_tool_span_serializes_non_text_content(): + """Image/resource results have no text part, so fall back to JSON.""" + from unittest.mock import MagicMock + + from mcp.types import CallToolResult, ImageContent + + span = MagicMock() + response_obj = CallToolResult( + content=[ImageContent(type="image", data="Zm9v", mimeType="image/png")], + isError=False, + ) + + ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.OUTPUT_MIME_TYPE] == "application/json" + assert "image/png" in written[SpanAttributes.OUTPUT_VALUE] + + +def test_arize_mcp_tool_span_respects_message_redaction(): + """Tool arguments and results are user content. With redaction on, only the + tool name may reach the span.""" + from unittest.mock import MagicMock + + from mcp.types import CallToolResult, TextContent + + span = MagicMock() + response_obj = CallToolResult( + content=[TextContent(type="text", text="SSN 123-45-6789")], isError=False + ) + + ArizeLogger.set_arize_attributes( + span, + _mcp_kwargs(standard_callback_dynamic_params={"turn_off_message_logging": True}), + response_obj, + ) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.TOOL_NAME] == "get_weather" + assert SpanAttributes.INPUT_VALUE not in written + assert SpanAttributes.OUTPUT_VALUE not in written + + +def test_arize_non_mcp_span_gets_no_tool_name(): + """The MCP emitter must not fire on ordinary completions.""" + from unittest.mock import MagicMock + + from litellm.types.utils import Choices, ModelResponse + + span = MagicMock() + kwargs = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "standard_logging_object": { + "model_parameters": {}, + "metadata": {"mcp_tool_call_metadata": {"name": "get_weather"}}, + "call_type": "completion", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + } + response_obj = ModelResponse( + choices=[Choices(message={"role": "assistant", "content": "hello"})], + model="gpt-4o", + id="r-1", + ) + + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert SpanAttributes.TOOL_NAME not in written + assert written[SpanAttributes.OUTPUT_VALUE] == "hello" + + +def test_arize_mcp_tool_span_renders_empty_arguments(): + from unittest.mock import MagicMock + + from mcp.types import CallToolResult, TextContent + + span = MagicMock() + kwargs = _mcp_kwargs(mcp_tool_call_metadata={"name": "ping", "arguments": {}}) + response_obj = CallToolResult(content=[TextContent(type="text", text="pong")], isError=False) + + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.INPUT_VALUE] == "{}" + assert written[SpanAttributes.INPUT_MIME_TYPE] == "application/json" + + +def test_arize_mcp_tool_span_renders_empty_content(): + from unittest.mock import MagicMock + + from mcp.types import CallToolResult + + span = MagicMock() + response_obj = CallToolResult(content=[], isError=False) + + ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.OUTPUT_VALUE] == "[]" + assert written[SpanAttributes.OUTPUT_MIME_TYPE] == "application/json" + + +def test_arize_mcp_tool_span_falls_back_to_structured_content(): + from unittest.mock import MagicMock + + from mcp.types import CallToolResult + + span = MagicMock() + response_obj = CallToolResult(content=[], structuredContent={"temp_c": 21}, isError=False) + + ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.OUTPUT_VALUE] == '{"temp_c": 21}' + assert written[SpanAttributes.OUTPUT_MIME_TYPE] == "application/json" + + +def test_arize_list_mcp_tools_response_does_not_break_attribute_setting(): + from unittest.mock import MagicMock + + span = MagicMock() + kwargs = { + "model": "MCP: list_tools", + "messages": [{"role": "user", "content": "list"}], + "standard_logging_object": { + "model_parameters": {}, + "metadata": {}, + "call_type": "list_mcp_tools", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "mcp"}, + } + + ArizeLogger.set_arize_attributes(span, kwargs, [{"name": "get_weather"}]) + + span.record_exception.assert_not_called() + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written["llm.input_messages.0.message.content"] == "list" + + +def test_arize_mcp_tool_span_serializes_mixed_text_and_media(): + from unittest.mock import MagicMock + + from mcp.types import CallToolResult, ImageContent, TextContent + + span = MagicMock() + response_obj = CallToolResult( + content=[ + TextContent(type="text", text="see image"), + ImageContent(type="image", data="Zm9v", mimeType="image/png"), + ], + isError=False, + ) + + ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.OUTPUT_MIME_TYPE] == "application/json" + assert "see image" in written[SpanAttributes.OUTPUT_VALUE] + assert "image/png" in written[SpanAttributes.OUTPUT_VALUE] + + +def test_arize_mcp_tool_span_without_response_object_keeps_name_and_input(): + from unittest.mock import MagicMock + + span = MagicMock() + + ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), None) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.TOOL_NAME] == "get_weather" + assert written[SpanAttributes.INPUT_VALUE] == '{"city": "Seoul"}' + assert SpanAttributes.OUTPUT_VALUE not in written + + +def test_arize_mcp_tool_span_without_content_emits_no_output(): + from unittest.mock import MagicMock + + span = MagicMock() + + ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), {"isError": False}) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.TOOL_NAME] == "get_weather" + assert SpanAttributes.OUTPUT_VALUE not in written + + +def test_arize_mcp_emitter_is_inert_without_a_standard_logging_object(): + from unittest.mock import MagicMock + + span = MagicMock() + kwargs = { + "model": "MCP: get_weather", + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "mcp"}, + } + + ArizeLogger.set_arize_attributes(span, kwargs, None) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert SpanAttributes.TOOL_NAME not in written From 363d56f91788d0ea34cb02ae54c7081e06ec6607 Mon Sep 17 00:00:00 2001 From: Deepanshu Lulla Date: Mon, 10 Aug 2026 19:47:17 -0400 Subject: [PATCH 10/10] feat(proxy): add per-deployment keepalive_seconds SSE heartbeat to prevent load-balancer timeout on long streams (#34423) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat(proxy): add per-deployment keepalive_seconds SSE heartbeat for long-running streams Adds _iter_with_keepalive, _keepalive_from_deployment_config, and _resolve_keepalive_seconds helpers to proxy_server.py. When enabled (keepalive_seconds > 0 in request body or deployment litellm_params), async_data_generator emits ': ping\n\n' SSE comment frames every N seconds during idle upstream intervals, preventing load-balancer idle-timeout drops on long chain-of-thought reasoning streams. The hot path (keepalive_seconds absent or 0) is a plain async-for with no per-chunk Task wrapping — zero overhead. Includes 8 new unit tests covering sentinel emission, hot-path pass-through, early-close cleanup, priority resolution, deployment-config lookup, and end-to-end heartbeat emission through async_data_generator. Registers keepalive_seconds in all_litellm_params (types/utils.py) so the parameter is not stripped from request bodies. Adds the field to LiteLLMParamsTypedDict and GenericLiteLLMParams (types/router.py) so deployment YAML config is parsed and validated. Co-Authored-By: Claude Sonnet 4.6 * fix(proxy): narrow BaseException to CancelledError to fix BLE001 strict lint gate * fix: use explicit None check instead of truthiness in keepalive_seconds extraction `float(raw or 0)` would treat any falsy value (including the integer 0) as absent and substitute 0.0 before float() saw it. Replace with `float(raw) if raw is not None else 0.0` so a caller-supplied zero is correctly passed through to the `value <= 0` guard that disables keepalive, rather than being silently overwritten. * fix(proxy): don't guess a deployment's keepalive_seconds when model_id is missing When a streaming response lacks _hidden_params.model_id, the fallback that looks up keepalive_seconds by model_name previously returned the first configured deployment's value, which could apply the wrong interval (or override an explicit disable) when multiple deployments share the same model_name with different keepalive_seconds settings. Only resolve the fallback when every deployment agrees; otherwise leave it unset. * fix(proxy): also treat an unset keepalive_seconds as disagreement in the fallback The model_name fallback for keepalive_seconds only compared configured values, filtering out deployments that leave the field unset entirely. That meant a deployment with no keepalive_seconds configured could still inherit a sibling deployment's interval when model_id is unavailable. Compare the raw per-deployment value (including None for unset) so an unconfigured deployment never silently adopts another's heartbeat. * fix(proxy): deployment-level keepalive_seconds: 0 is a hard disable clients can't override Previously an authenticated client's request-level keepalive_seconds always took precedence over the deployment default, including when a deployment operator explicitly set keepalive_seconds: 0 to disable heartbeats. That let any client re-enable heartbeats for a deployment the operator opted out of, using them to keep an idle-looking stream alive past a load balancer's idle timeout and hold a parallel-request slot open longer than intended. Treat an explicit deployment-level 0 as authoritative: resolve the deployment's configured value first, and short-circuit to disabled before ever looking at the request body if the deployment hard-disabled it. * fix(proxy): a stale (unresolvable) model_id must not fall through to model_name guessing A populated _hidden_params.model_id names the specific deployment that served a stream. If that ID no longer resolves (e.g. a deployment removed by a config reload mid-stream), the resolver was falling through to the model_name-based fallback, letting a currently-live sibling deployment's keepalive_seconds silently apply to a stream it never served. Return None once a populated model_id fails to resolve, rather than degrading to a guess. * fix(proxy): keepalive_seconds is operator-only by default; require deployment opt-in for client override A security review flagged that a client's request-level keepalive_seconds could unilaterally enable heartbeats for any deployment, even one that never configured keepalive_seconds at all, letting an authenticated client defeat load-balancer idle timeouts and hold a parallel-request slot open for longer than the deployment operator ever intended, with no way for the operator to prevent it short of explicitly setting keepalive_seconds: 0. Add allow_client_keepalive_override (default False) to LiteLLMParamsTypedDict and GenericLiteLLMParams. _resolve_keepalive_seconds now ignores the request body's keepalive_seconds entirely unless the resolved deployment explicitly grants override permission; only the deployment's own configured value (or disabled, if unset) applies otherwise. An explicit deployment-level 0 still takes priority over everything, including a grant of override permission. * fix(proxy): register allow_client_keepalive_override in all_litellm_params Caught during live proxy verification against the real Anthropic API: allow_client_keepalive_override was added to LiteLLMParamsTypedDict and GenericLiteLLMParams but never registered in all_litellm_params, so it leaked straight through into the provider request body as an unrecognized field. Anthropic rejected every call on a deployment that had this field configured with a 400 ("Extra inputs are not permitted"), regardless of its value. Register it alongside keepalive_seconds so it's stripped before reaching the provider, matching what keepalive_seconds already does. * feat(proxy): support keepalive_seconds via x-litellm-keepalive-seconds header Some clients (e.g. the Vercel AI SDK) can set custom headers more easily than extra JSON body fields. Add x-litellm-keepalive-seconds, following the existing x-litellm-timeout/x-litellm-stream-timeout/x-litellm-num-retries convention in LiteLLMProxyRequestSetup: the header merges into the same data["keepalive_seconds"] field the request body already populates, so it goes through the exact same _resolve_keepalive_seconds precedence and the allow_client_keepalive_override gate -- a header can't enable heartbeats for a deployment that hasn't opted in any more than the body field can. Verified live against the real Anthropic API: the header produces real heartbeats on an opt-in deployment (88 pings over a genuine long-reasoning stall) and is silently ignored on a deployment without override permission (0 pings), matching the existing body-field behavior exactly. * chore: rebase onto litellm_internal_staging, drop unrelated credential_migration.py reformat, fix budget-ratchet drift Rebased onto the current litellm_internal_staging (merge-base was 5 days stale). Dropped the now-redundant schema.d.ts-only regen commit entirely (the new base's own schema.d.ts already supersedes it) and regenerated schema.d.ts fresh against the new base. Reverted litellm/proxy/management_endpoints/credential_migration.py to exactly match litellm_internal_staging: it was a pure reformat with no semantic change, unrelated to this PR, flagged by review as unnecessary noise in an encryption-migration file. Fixed two lint-budget-ratchet failures caused by the base's ceilings tightening since this branch last synced (other merged work lowered ANN401/LIT001 budgets; this code was previously under budget and didn't change): - _iter_with_keepalive's aiter param: Any -> AsyncIterator[Any], a real narrowing (it's always the result of .__aiter__()). - _keepalive_from_deployment_config/_resolve_keepalive_seconds's request_data param: dict[str, Any] -> Mapping[str, Any], matching the existing read-only-dict convention already used elsewhere in this file (_apply_ssrf_general_settings, _build_redis_usage_cache, etc.) for params that are only ever read, never mutated. - response/raw params: dropped the explicit `Any` annotation to match async_data_generator's own (deliberately unannotated) `response` param, its actual caller. - litellm_pre_call_utils.py's new headers param: dict -> Mapping[str, str], same read-only-dict rationale. * fix(proxy): freeze the transient collections in the keepalive helpers _iter_with_keepalive and _keepalive_from_deployment_config built a set literal for asyncio.wait, a set comprehension for the per-deployment config-agreement check, and two dict-literal fallbacks, all flagged by the LIT002 mutable-collection-construction gate. Switched to a tuple for asyncio.wait, a frozenset-wrapped generator plus next(iter(...)) for the config check, and a shared MappingProxyType({}) empty mapping for the fallbacks. * fix(proxy): trust metadata.model_info.id over the stale model group after a router fallback Greptile P1: when a streaming request falls back from model group A to group B and the response's _hidden_params carries no model_id, _keepalive_from_deployment_config fell straight through to guessing via request_data["model"], which still names the pre-fallback group A since the fallback handler mutates its own local **kwargs copy, not this dict. request_data[metadata|litellm_metadata]["model_info"]["id"], by contrast, is mutated on this same dict by Router._update_kwargs_with_deployment on every attempt including fallbacks (the same source ProxyLogging._build_litellm_call_info uses for logging), so check it before falling through to the model-name guess. Added two regression tests that fail on the prior code (assert get_model_list is never called once metadata.model_info.id resolves) and pass with the fix. * Revert "fix(proxy): trust metadata.model_info.id over the stale model group after a router fallback" This reverts commit d7790678645695b20f25880315238b49c31a9143. * fix(proxy): satisfy the new LIT010/ANN001 gates in the keepalive helpers litellm_internal_staging picked up a LIT010 (every local/module variable must be declared Final unless it's genuinely rebound) and tightened ANN001 (missing parameter annotations) since this branch last synced. Annotated every single-assignment local and module constant with Final, suppressed pending's loop-carried reassignment with # rebind-ok, and typed the previously-bare response/raw parameters as object with isinstance narrowing at their use sites instead of cast (LIT006 discourages cast; validate into a concrete type instead). Also swapped the hand-rolled getattr(response, "_hidden_params", None) + isinstance(hidden, dict) check for the existing get_hidden_params_dict() helper already used for this exact purpose elsewhere in this file and in common_request_processing.py. * fix(proxy): re-resolve keepalive_seconds per chunk to track mid-stream fallback Greptile P1: the router can perform a mid-stream fallback to a different deployment partway through a stream (MidStreamFallbackError in router.py), and Router._apply_fallback_hidden_params_to_item merges the fallback deployment's hidden params onto every subsequent chunk. But _resolve_keepalive_seconds was only ever called once, before iteration started, against the pre-fallback response wrapper, so a stream that fell back to a deployment with a different (or disabled) keepalive policy kept using the original deployment's interval for the rest of the stream. _iter_with_keepalive now takes a resolve_keepalive_seconds(item) callback and re-resolves after every real chunk using that chunk's own _hidden_params (which do carry the fallback deployment's identity), rather than trusting the value picked before iteration began. Updated the three existing timing tests to inject a constant-returning resolver, since they pin the sentinel/cancellation mechanics rather than re-resolution, and added two regression tests (interval lowered and raised mid-stream) that fail against the prior static-resolve signature and pass with the fix. * fix(proxy): keep re-resolving keepalive even when a stream starts disabled Greptile P1: a stream that starts on a deployment with keepalive off (or unset) skipped _iter_with_keepalive entirely at the call site, so a mid-stream fallback to a deployment that enables it never got a chance to activate heartbeats for the rest of that stream, risking the exact load-balancer idle-timeout this feature exists to prevent. _iter_with_keepalive now has an internal fast path for keepalive_seconds <= 0 that still re-resolves after every chunk (no asyncio.create_task/wait overhead while inactive, same cost as a bare async for), so activation from a disabled start works the same way deactivation and interval changes already do. The caller now only skips wrapping entirely when there's no router to ever fall back through in the first place (llm_router is None), rather than whenever the first chunk's deployment happens to start with keepalive off. Added a regression test that starts keepalive_seconds=0, has the resolver enable a short interval on a later chunk, and asserts sentinels appear afterward; it fails against the prior call-site-gated code and passes with the fix. * perf(proxy): memoize keepalive resolution per chunk's model_id _resolve_keepalive_seconds ran a full llm_router.get_deployment() Pydantic rebuild after every streamed chunk, even when keepalive was unconfigured anywhere in the deployment list, since async_data_generator wraps every stream once a router exists. Caching the result by model_id keeps mid-stream fallback re-resolution correct while paying the router lookup once per deployment instead of once per token. * fix(proxy): expire cached keepalive resolution after a bounded TTL veria-ai flagged that caching by model_id alone lets an already-in-flight stream keep evading a live config reload (deployment removed, keepalive disabled, or client override revoked) for the rest of the stream. Expiring the memo after _KEEPALIVE_CACHE_TTL_SECONDS bounds that window instead of freezing the resolved value for the stream's full lifetime, while still avoiding a full deployment rebuild on every chunk in the steady state. Also fixes add_litellm_data_for_backend_llm_call's now-required request_data kwarg in the header-merge test, picked up by rebasing onto litellm_internal_staging. --------- Co-authored-by: Deepanshu Co-authored-by: Claude Sonnet 4.6 --- litellm/proxy/_types.py | 1 + litellm/proxy/litellm_pre_call_utils.py | 17 + litellm/proxy/proxy_server.py | 220 +++++- litellm/types/router.py | 8 + litellm/types/utils.py | 2 + .../proxy_server/test_streaming_helpers.py | 707 +++++++++++++++++- .../proxy/test_litellm_pre_call_utils.py | 49 ++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 14 + 8 files changed, 1004 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index c8ca9dcf57e..dc4f17c7b31 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4031,6 +4031,7 @@ class LitellmDataForBackendLLMCall(TypedDict, total=False): # deliberately tiny value isn't treated as a deployment health signal (see # cooldown_handlers._trigger_cooldown_for_failed_deployment). client_side_timeout: bool + keepalive_seconds: float | None class LitellmMetadataFromRequestHeaders(TypedDict, total=False): diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 66bfc8b81d8..10142a894a1 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -882,6 +882,19 @@ class LiteLLMProxyRequestSetup: return float(stream_timeout_header) return None + @staticmethod + def _get_keepalive_seconds_from_request(headers: Mapping[str, str]) -> float | None: + """ + Get `keepalive_seconds` from the request headers, for clients (e.g. the + Vercel AI SDK) that can set custom headers more easily than extra body + fields. Subject to the same deployment-level allow_client_keepalive_override + gate as the request body field: see _resolve_keepalive_seconds. + """ + keepalive_seconds_header: Final = headers.get("x-litellm-keepalive-seconds", None) + if keepalive_seconds_header is not None: + return float(keepalive_seconds_header) + return None + @staticmethod def _get_num_retries_from_request(headers: dict) -> int | None: """ @@ -1114,6 +1127,10 @@ class LiteLLMProxyRequestSetup: if num_retries is not None: data["num_retries"] = num_retries + keepalive_seconds: Final = LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request(headers) + if keepalive_seconds is not None: + data["keepalive_seconds"] = keepalive_seconds + return data @staticmethod diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index bc980934f9f..baa1579d2c0 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -15,7 +15,7 @@ import threading import time import traceback import warnings -from collections.abc import AsyncGenerator, Callable, Mapping +from collections.abc import AsyncGenerator, AsyncIterator, Callable, Mapping from datetime import datetime, timedelta, timezone from types import MappingProxyType, UnionType from typing import ( @@ -23,6 +23,7 @@ from typing import ( Any, Final, Literal, + NamedTuple, Optional, TypedDict, Union, @@ -7643,6 +7644,200 @@ def _pop_complete_sse_frame(buffer: str) -> tuple[str | None, str]: return buffer[:frame_end], buffer[frame_end:] +_STREAM_KEEPALIVE: Final = object() + +_KEEPALIVE_MIN_SECONDS: Final = 1.0 +_KEEPALIVE_MAX_SECONDS: Final = 300.0 +_EMPTY_MAPPING: Final[Mapping[str, Any]] = MappingProxyType({}) + + +async def _iter_with_keepalive( + aiter: AsyncIterator[Any], + resolve_keepalive_seconds: Callable[[object], float], + keepalive_seconds: float, +) -> AsyncGenerator[Any, None]: + """Wrap `aiter` with idle-gap heartbeats, re-resolving the interval after each + real chunk via `resolve_keepalive_seconds`. A mid-stream router fallback can + swap in a deployment with a different keepalive policy, including one that + newly enables or newly disables heartbeats, partway through the same stream; + re-resolving against each chunk's own identity (rather than trusting the + interval picked before iteration started, or picked the last time it went + inactive) keeps the heartbeat behavior in sync with whichever deployment + actually produced it, in both directions. While the interval is <= 0, no + task is created and no timeout is awaited: a chunk is forwarded the moment + it arrives, at the same cost as a bare `async for`.""" + pending: asyncio.Task[Any] | None = None # rebind-ok: rebound each loop iteration + current_keepalive_seconds = keepalive_seconds # rebind-ok: re-resolved after each chunk + try: + while True: + if current_keepalive_seconds <= 0: + try: + item = await aiter.__anext__() + except StopAsyncIteration: + break + yield item + current_keepalive_seconds = resolve_keepalive_seconds(item) + continue + + if pending is None: + pending = asyncio.create_task(aiter.__anext__()) + done, _ = await asyncio.wait((pending,), timeout=current_keepalive_seconds) + if not done: + yield _STREAM_KEEPALIVE + continue + try: + item = pending.result() + except StopAsyncIteration: + break + finally: + pending = None + yield item + current_keepalive_seconds = resolve_keepalive_seconds(item) + finally: + if pending is not None and not pending.done(): + pending.cancel() + try: + await pending + except asyncio.CancelledError: + pass + + +class _DeploymentKeepaliveConfig(NamedTuple): + keepalive_seconds: Any + allow_client_override: bool + + +def _keepalive_from_deployment_config( + request_data: Mapping[str, Any], response: object +) -> _DeploymentKeepaliveConfig | None: + if llm_router is None: + return None + + hidden: Final = get_hidden_params_dict(response) + model_id: Final = hidden.get("model_id") + if isinstance(model_id, str) and model_id: + deployment: Final = llm_router.get_deployment(model_id=model_id) + # A populated model_id names the specific deployment that served this + # stream. If it no longer resolves (e.g. removed by a config reload + # mid-stream), that's a stale identity, not an absent one: don't fall + # through to guessing via model_name below, since a currently-live + # sibling deployment's config was never what actually served this + # stream. + if deployment is None: + return None + return _DeploymentKeepaliveConfig( + keepalive_seconds=getattr(deployment.litellm_params, "keepalive_seconds", None), + allow_client_override=bool(getattr(deployment.litellm_params, "allow_client_keepalive_override", False)), + ) + + # No model_id at all to pin down which deployment actually served this + # stream: only trust the fallback when every deployment under this + # model_name agrees on both keepalive_seconds and + # allow_client_keepalive_override (including deployments that leave either + # field unset), so a stream never inherits a sibling deployment's policy. + configs: Final = frozenset( + ( + (deployment_dict.get("litellm_params") or _EMPTY_MAPPING).get("keepalive_seconds"), + bool( + (deployment_dict.get("litellm_params") or _EMPTY_MAPPING).get("allow_client_keepalive_override", False) + ), + ) + for deployment_dict in llm_router.get_model_list(model_name=request_data.get("model")) or () + ) + if len(configs) == 1: + keepalive_seconds, allow_client_override = next(iter(configs)) + return _DeploymentKeepaliveConfig( + keepalive_seconds=keepalive_seconds, allow_client_override=allow_client_override + ) + return None + + +def _is_explicit_keepalive_disable(raw: object) -> bool: + if not isinstance(raw, (int, float, str)): + return False + try: + return float(raw) <= 0 + except ValueError: + return False + + +def _resolve_keepalive_seconds(request_data: Mapping[str, Any], response: object = None) -> float: + deployment_config: Final = _keepalive_from_deployment_config(request_data, response) + deployment_raw: Final = deployment_config.keepalive_seconds if deployment_config is not None else None + allow_client_override: Final = deployment_config.allow_client_override if deployment_config is not None else False + + # An operator setting keepalive_seconds: 0 on a deployment is an explicit hard + # disable: an authenticated client must not be able to re-enable heartbeats + # (and the idle-timeout evasion that comes with them) for a deployment the + # operator opted out of, regardless of what the request body asks for. + if _is_explicit_keepalive_disable(deployment_raw): + return 0.0 + + # keepalive_seconds is operator-only unless the deployment explicitly opts in: + # a client can't unilaterally enable heartbeats (and the LB-idle-timeout + # evasion that comes with them) for a deployment that never configured this. + client_supplied: Final = request_data.get("keepalive_seconds") if allow_client_override else None + raw: Final = client_supplied if client_supplied is not None else deployment_raw + try: + value: Final = float(raw) if isinstance(raw, (int, float, str)) else 0.0 + except ValueError: + return 0.0 + if value <= 0: + return 0.0 + clamped: Final = max(_KEEPALIVE_MIN_SECONDS, min(value, _KEEPALIVE_MAX_SECONDS)) + if clamped != value: + verbose_proxy_logger.info( + "keepalive_seconds=%s clamped to %s [min=%s, max=%s]", + value, + clamped, + _KEEPALIVE_MIN_SECONDS, + _KEEPALIVE_MAX_SECONDS, + ) + return clamped + + +_KEEPALIVE_CACHE_TTL_SECONDS: Final = 5.0 + + +def _make_keepalive_resolver(request_data: Mapping[str, Any]) -> Callable[[object], float]: + """Wrap `_resolve_keepalive_seconds` with a memo keyed on the serving + deployment's model_id. The steady-state case (no mid-stream fallback, the + overwhelming majority of streams) sees the same model_id on every chunk, so + this turns the per-chunk cost from a full `llm_router.get_deployment()` + Pydantic rebuild into a cheap hidden-params read once per + `_KEEPALIVE_CACHE_TTL_SECONDS` for that model_id. The cache expires on its + own rather than living for the life of the stream, so an operator's live + config change (disabling keepalive, revoking client override, or removing + the deployment) is observed within a bounded window instead of being able + to be evaded by an already-in-flight stream indefinitely. A missing/empty + model_id can't be trusted as a cache key (see + `_keepalive_from_deployment_config`'s model_name fallback, which reflects + current router state rather than one deployment's fixed identity), so + those chunks always resolve fresh, matching prior behavior exactly. + """ + last_model_id: str | None = None # rebind-ok: memoized identity of the last-resolved chunk + last_value: float = 0.0 # rebind-ok: cached resolution for last_model_id + last_resolved_at: float = float("-inf") # rebind-ok: monotonic timestamp of the last real resolution + + def _resolve(item: object) -> float: + nonlocal last_model_id, last_value, last_resolved_at + model_id = get_hidden_params_dict(item).get("model_id") + now: Final = time.monotonic() + if ( + isinstance(model_id, str) + and model_id + and model_id == last_model_id + and now - last_resolved_at < _KEEPALIVE_CACHE_TTL_SECONDS + ): + return last_value + value: Final = _resolve_keepalive_seconds(request_data, item) + if isinstance(model_id, str) and model_id: + last_model_id, last_value, last_resolved_at = model_id, value, now + return value + + return _resolve + + async def async_data_generator( response, user_api_key_dict: UserAPIKeyAuth, @@ -7691,7 +7886,28 @@ async def async_data_generator( else: stream_iterator = response - async for chunk in stream_iterator: + # A stream can start on a deployment with keepalive off and fall back + # mid-stream to one that enables it: only skip wrapping altogether when + # there's no router to ever fall back through in the first place (in + # which case _resolve_keepalive_seconds can never return non-zero for + # any chunk of this stream), not merely because the first chunk's + # deployment happens to start with it off. + resolve_keepalive_seconds: Final = _make_keepalive_resolver(request_data) + stream_source: Final = ( + _iter_with_keepalive( + stream_iterator.__aiter__(), + resolve_keepalive_seconds, + resolve_keepalive_seconds(response), + ) + if llm_router is not None + else stream_iterator + ) + + async for item in stream_source: + if item is _STREAM_KEEPALIVE: + yield ": ping\n\n" + continue + chunk = cast(Any, item) # cast-ok: sentinel already handled above, item is a real chunk here if needs_per_chunk_hook: ### CALL HOOKS ### - modify outgoing data chunk, _str_so_far = await _apply_streaming_chunk_hooks( diff --git a/litellm/types/router.py b/litellm/types/router.py index b03796fb14f..4f8c133c20b 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -281,6 +281,12 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams): # Deployment budgets max_budget: float | None = None budget_duration: str | None = None + keepalive_seconds: float | None = None + # keepalive_seconds is operator-only by default: a client's request-level + # value is ignored unless the deployment opts in here. Prevents a client + # from unilaterally enabling heartbeats (and the LB-idle-timeout evasion + # that comes with them) for a deployment that never configured them. + allow_client_keepalive_override: bool | None = False use_in_pass_through: bool | None = False use_litellm_proxy: bool | None = False use_chat_completions_api: bool | None = None @@ -457,6 +463,8 @@ class LiteLLMParamsTypedDict(TypedDict, total=False): # deployment budgets max_budget: float | None budget_duration: str | None + keepalive_seconds: float | None + allow_client_keepalive_override: bool | None # per-deployment cooldown override cooldown_time: float | None diff --git a/litellm/types/utils.py b/litellm/types/utils.py index abf9382845b..8c7664257df 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3399,6 +3399,8 @@ all_litellm_params = ( + [ "metadata", "litellm_metadata", + "keepalive_seconds", + "allow_client_keepalive_override", "litellm_trace_id", "litellm_request_debug", "guardrails", diff --git a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py index f7e2d276a2e..15758c595c0 100644 --- a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py +++ b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py @@ -196,9 +196,7 @@ async def test_async_assistants_data_generator_hook_failure_yields_error_chunk( async def _noop_failure(*args, **kwargs): return None - monkeypatch.setattr( - ps.proxy_logging_obj, "async_post_call_streaming_hook", _boom_hook - ) + monkeypatch.setattr(ps.proxy_logging_obj, "async_post_call_streaming_hook", _boom_hook) monkeypatch.setattr(ps.proxy_logging_obj, "post_call_failure_hook", _noop_failure) stream = _FakeAssistantsStream([_simple_chunk()]) @@ -385,9 +383,7 @@ def test_get_streaming_fallback_metadata_no_additional_headers(): def test_get_streaming_fallback_metadata_zero_fallback_count(): stream = _FakeStream( [], - hidden_params={ - "additional_headers": {"x-litellm-attempted-fallbacks": 0} - }, + hidden_params={"additional_headers": {"x-litellm-attempted-fallbacks": 0}}, ) assert _get_streaming_fallback_metadata(stream) == (False, None, []) @@ -558,9 +554,7 @@ async def test_apply_streaming_chunk_hooks_appends_to_str_so_far(monkeypatch): async def _passthrough(*, user_api_key_dict, response, data, str_so_far=None): return response - monkeypatch.setattr( - ps.proxy_logging_obj, "async_post_call_streaming_hook", _passthrough - ) + monkeypatch.setattr(ps.proxy_logging_obj, "async_post_call_streaming_hook", _passthrough) new_chunk, new_str = await _apply_streaming_chunk_hooks( chunk=chunk, @@ -870,9 +864,7 @@ async def test_async_data_generator_mid_stream_exception_yields_error_payload( out.append(line) # First entry is the successful "partial" chunk (bytes), last is the error. - assert any( - isinstance(item, str) and item.startswith('data: {"error":') for item in out - ) + assert any(isinstance(item, str) and item.startswith('data: {"error":') for item in out) # --------------------------------------------------------------------------- @@ -914,3 +906,694 @@ def test_select_data_generator_missing_required_kwarg_raises_type_error(): streaming starts.""" with pytest.raises(TypeError): select_data_generator(response=_async_iter([]), user_api_key_dict=_user_auth()) # type: ignore[call-arg] + + +# --------------------------------------------------------------------------- +# SSE keepalive helpers +# --------------------------------------------------------------------------- + + +from litellm.proxy.proxy_server import ( # noqa: E402 + _iter_with_keepalive, + _keepalive_from_deployment_config, + _make_keepalive_resolver, + _resolve_keepalive_seconds, +) +from litellm.proxy.proxy_server import _KEEPALIVE_MAX_SECONDS, _KEEPALIVE_MIN_SECONDS # noqa: E402 + + +@pytest.mark.asyncio +async def test_iter_with_keepalive_hot_path_no_task_wrapping(): + """When keepalive_seconds <= 0, the generator is a transparent pass-through.""" + chunks = [_simple_chunk(content="a"), _simple_chunk(content="b")] + out = [] + async for item in _iter_with_keepalive(_async_iter(chunks), lambda _: 0, keepalive_seconds=0): + out.append(item) + + assert out == chunks + assert ps._STREAM_KEEPALIVE not in out + + +@pytest.mark.asyncio +async def test_iter_with_keepalive_emits_sentinel_when_stream_stalls(): + """With a short keepalive interval and a stalled upstream, _STREAM_KEEPALIVE + sentinels appear before the delayed chunk arrives. The resolver returns a + constant interval, since this test pins the timing mechanics, not + re-resolution.""" + import asyncio + + async def _slow_stream(): + yield _simple_chunk(content="first") + await asyncio.sleep(0.3) + yield _simple_chunk(content="second") + + items = [] + async for item in _iter_with_keepalive(_slow_stream(), lambda _: 0.05, keepalive_seconds=0.05): + items.append(item) + + sentinels = [i for i in items if i is ps._STREAM_KEEPALIVE] + real_chunks = [i for i in items if i is not ps._STREAM_KEEPALIVE] + + assert len(sentinels) >= 2, f"expected >= 2 sentinels during 0.3s stall; got {len(sentinels)}" + assert len(real_chunks) == 2 + assert real_chunks[0].choices[0].delta.content == "first" + assert real_chunks[1].choices[0].delta.content == "second" + + +@pytest.mark.asyncio +async def test_iter_with_keepalive_cancel_on_early_close(): + """Closing the generator early cancels the in-flight task without raising.""" + import asyncio + + async def _infinite_stream(): + while True: + await asyncio.sleep(10) + yield _simple_chunk() + + gen = _iter_with_keepalive(_infinite_stream(), lambda _: 0.05, keepalive_seconds=0.05) + # Advance once to get the sentinel; then close before the real chunk. + first = await gen.__anext__() + assert first is ps._STREAM_KEEPALIVE + # aclose must not raise, and must drain the cancelled task cleanly. + await gen.aclose() + + +@pytest.mark.asyncio +async def test_iter_with_keepalive_disables_after_fallback_lowers_interval(): + """Greptile P1: a mid-stream router fallback can hand off to a deployment + with a different (or disabled) keepalive policy partway through the same + stream. The interval must be re-resolved against each chunk's own identity, + not the one picked before iteration started, or heartbeats keep using the + pre-fallback deployment's policy for the rest of the stream.""" + import asyncio + + async def _slow_stream(): + yield _simple_chunk(content="first") + await asyncio.sleep(0.3) + yield _simple_chunk(content="second") + + def _resolver(item): + # First chunk resolves under the enabled interval used to start the + # wrapper; every chunk after that resolves as if a fallback disabled it. + return 0.0 if item.choices[0].delta.content == "first" else 999.0 + + items = [] + async for item in _iter_with_keepalive(_slow_stream(), _resolver, keepalive_seconds=0.05): + items.append(item) + + sentinels = [i for i in items if i is ps._STREAM_KEEPALIVE] + real_chunks = [i for i in items if i is not ps._STREAM_KEEPALIVE] + + assert sentinels == [], f"expected no sentinels once the resolver disables keepalive; got {len(sentinels)}" + assert len(real_chunks) == 2 + assert real_chunks[0].choices[0].delta.content == "first" + assert real_chunks[1].choices[0].delta.content == "second" + + +@pytest.mark.asyncio +async def test_iter_with_keepalive_enables_after_fallback_raises_interval(): + """Symmetric case: a mid-stream fallback to a deployment with a *shorter* + keepalive interval must take effect immediately, not stay pinned to the + longer interval the stream started with. The interval used to wait for a + chunk is resolved from the *previous* chunk (the only one seen so far when + that wait begins), so the stall has to follow the fallback chunk rather + than precede it: waiting for "third" is where the shorter interval bites.""" + import asyncio + + async def _slow_stream(): + yield _simple_chunk(content="first") + yield _simple_chunk(content="second") + await asyncio.sleep(0.3) + yield _simple_chunk(content="third") + + def _resolver(item): + # "first" resolves under an interval too long to fire before "second" + # arrives; "second" (the fallback chunk) resolves as if the fallback + # deployment enabled a much shorter interval for everything after it. + return 999.0 if item.choices[0].delta.content == "first" else 0.05 + + items = [] + async for item in _iter_with_keepalive(_slow_stream(), _resolver, keepalive_seconds=999.0): + items.append(item) + + sentinels = [i for i in items if i is ps._STREAM_KEEPALIVE] + real_chunks = [i for i in items if i is not ps._STREAM_KEEPALIVE] + + assert len(sentinels) >= 2, ( + f"expected >= 2 sentinels once the resolver enables a short interval; got {len(sentinels)}" + ) + assert len(real_chunks) == 3 + + +@pytest.mark.asyncio +async def test_iter_with_keepalive_activates_from_a_fully_disabled_start(): + """Greptile P1: a stream can start on a deployment with keepalive off + (keepalive_seconds passed in as 0, not merely a long interval) and fall back + mid-stream to one that enables it. The 0-second start must not be treated as + a one-time decision to skip heartbeats for the rest of the stream: no task + is created while inactive, but every chunk still re-resolves so the fallback + chunk can switch the stream into task-wrapped mode.""" + import asyncio + + async def _slow_stream(): + yield _simple_chunk(content="first") + yield _simple_chunk(content="second") + await asyncio.sleep(0.3) + yield _simple_chunk(content="third") + + def _resolver(item): + # "first" resolves to stay off; "second" (the fallback chunk) resolves + # as if the fallback deployment newly enabled a short interval. + return 0.0 if item.choices[0].delta.content == "first" else 0.05 + + items = [] + async for item in _iter_with_keepalive(_slow_stream(), _resolver, keepalive_seconds=0): + items.append(item) + + sentinels = [i for i in items if i is ps._STREAM_KEEPALIVE] + real_chunks = [i for i in items if i is not ps._STREAM_KEEPALIVE] + + assert len(sentinels) >= 2, ( + f"expected >= 2 sentinels once the resolver activates from a disabled start; got {len(sentinels)}" + ) + assert len(real_chunks) == 3 + + +def test_resolve_keepalive_seconds_client_value_ignored_without_override_permission(monkeypatch): + """keepalive_seconds is operator-only by default: a deployment that hasn't set + allow_client_keepalive_override must not let a client's request-level value + change its behavior at all, since that would let any authenticated client + unilaterally enable heartbeats (and the LB-idle-timeout evasion that comes + with them) for a deployment that never opted in.""" + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = 15.0 + deployment.litellm_params.allow_client_keepalive_override = False + + router = MagicMock() + router.get_deployment.return_value = deployment + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "deploy-locked"} + + result = _resolve_keepalive_seconds({"model": "my-model", "keepalive_seconds": 1}, response=response) + assert result == 15.0 + + +def test_resolve_keepalive_seconds_request_value_wins_when_override_allowed(monkeypatch): + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = None + deployment.litellm_params.allow_client_keepalive_override = True + + router = MagicMock() + router.get_deployment.return_value = deployment + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "deploy-opt-in"} + + result = _resolve_keepalive_seconds({"model": "my-model", "keepalive_seconds": 30}, response=response) + assert result == 30.0 + + +def test_resolve_keepalive_seconds_explicit_zero_disables_when_override_allowed(monkeypatch): + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = 20.0 + deployment.litellm_params.allow_client_keepalive_override = True + + router = MagicMock() + router.get_deployment.return_value = deployment + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "deploy-opt-in"} + + result = _resolve_keepalive_seconds({"model": "my-model", "keepalive_seconds": 0}, response=response) + assert result == 0.0 + + +def test_resolve_keepalive_seconds_clamps_below_minimum(monkeypatch): + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = None + deployment.litellm_params.allow_client_keepalive_override = True + + router = MagicMock() + router.get_deployment.return_value = deployment + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "deploy-opt-in"} + + result = _resolve_keepalive_seconds({"model": "my-model", "keepalive_seconds": 0.001}, response=response) + assert result == _KEEPALIVE_MIN_SECONDS + + +def test_resolve_keepalive_seconds_clamps_above_maximum(monkeypatch): + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = None + deployment.litellm_params.allow_client_keepalive_override = True + + router = MagicMock() + router.get_deployment.return_value = deployment + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "deploy-opt-in"} + + result = _resolve_keepalive_seconds({"model": "my-model", "keepalive_seconds": 9999}, response=response) + assert result == _KEEPALIVE_MAX_SECONDS + + +def test_resolve_keepalive_seconds_non_numeric_returns_zero(monkeypatch): + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = None + deployment.litellm_params.allow_client_keepalive_override = True + + router = MagicMock() + router.get_deployment.return_value = deployment + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "deploy-opt-in"} + + result = _resolve_keepalive_seconds({"model": "my-model", "keepalive_seconds": "not-a-number"}, response=response) + assert result == 0.0 + + +def test_resolve_keepalive_seconds_absent_returns_zero(monkeypatch): + monkeypatch.setattr(ps, "llm_router", None) + result = _resolve_keepalive_seconds({}, response=None) + assert result == 0.0 + + +def test_resolve_keepalive_seconds_deployment_disable_cannot_be_overridden_by_request(monkeypatch): + """A deployment that explicitly sets keepalive_seconds: 0 is a hard operator + disable: an authenticated client must not be able to re-enable heartbeats for + that deployment by passing a positive value in the request body, since that + would let a client evade the deployment's idle-timeout behavior at will. This + holds even if the deployment also grants override permission, since an + explicit disable is a stronger, unconditional signal than an override grant.""" + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = 0 + deployment.litellm_params.allow_client_keepalive_override = True + + router = MagicMock() + router.get_deployment.return_value = deployment + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "deploy-disabled"} + + result = _resolve_keepalive_seconds({"model": "my-model", "keepalive_seconds": 250}, response=response) + assert result == 0.0 + + +def test_keepalive_from_deployment_config_reads_by_model_id(monkeypatch): + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = 45.0 + deployment.litellm_params.allow_client_keepalive_override = True + + router = MagicMock() + router.get_deployment.return_value = deployment + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "deploy-abc"} + + result = _keepalive_from_deployment_config({"model": "my-model"}, response) + assert result == ps._DeploymentKeepaliveConfig(keepalive_seconds=45.0, allow_client_override=True) + router.get_deployment.assert_called_once_with(model_id="deploy-abc") + + +def test_keepalive_from_deployment_config_stale_model_id_does_not_fall_through(monkeypatch): + """A populated model_id names the specific deployment that served the stream. + If that ID no longer resolves (e.g. removed by a config reload mid-stream), + that's a stale identity, not an absent one: it must not fall through to the + model_name fallback, since a currently-live sibling deployment's config was + never what actually served this stream, even if that sibling's config is + unambiguous on its own.""" + from unittest.mock import MagicMock + + router = MagicMock() + router.get_deployment.return_value = None + router.get_model_list.return_value = [ + {"litellm_params": {"keepalive_seconds": 20.0}}, + ] + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "stale-deploy-id"} + + result = _keepalive_from_deployment_config({"model": "slow-model"}, response) + assert result is None + router.get_model_list.assert_not_called() + + +def test_keepalive_from_deployment_config_fallback_by_name(monkeypatch): + from unittest.mock import MagicMock + + router = MagicMock() + router.get_deployment.return_value = None + router.get_model_list.return_value = [ + {"litellm_params": {"keepalive_seconds": 20.0}}, + ] + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {} + + result = _keepalive_from_deployment_config({"model": "slow-model"}, response) + assert result == ps._DeploymentKeepaliveConfig(keepalive_seconds=20.0, allow_client_override=False) + router.get_model_list.assert_called_once_with(model_name="slow-model") + + +def test_keepalive_from_deployment_config_fallback_by_name_agreeing_deployments(monkeypatch): + """Multiple deployments under the same model_name with the same keepalive_seconds + is unambiguous, so the shared value is used even without a model_id.""" + from unittest.mock import MagicMock + + router = MagicMock() + router.get_deployment.return_value = None + router.get_model_list.return_value = [ + {"litellm_params": {"keepalive_seconds": 20.0, "allow_client_keepalive_override": True}}, + {"litellm_params": {"keepalive_seconds": 20.0, "allow_client_keepalive_override": True}}, + ] + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {} + + result = _keepalive_from_deployment_config({"model": "slow-model"}, response) + assert result == ps._DeploymentKeepaliveConfig(keepalive_seconds=20.0, allow_client_override=True) + + +def test_keepalive_from_deployment_config_fallback_by_name_conflicting_deployments(monkeypatch): + """Without a model_id, if deployments under the same model_name disagree on + keepalive_seconds, we can't tell which one served the stream: don't guess and + apply the wrong deployment's interval (or override an explicit disable).""" + from unittest.mock import MagicMock + + router = MagicMock() + router.get_deployment.return_value = None + router.get_model_list.return_value = [ + {"litellm_params": {"keepalive_seconds": 20.0}}, + {"litellm_params": {"keepalive_seconds": 0}}, + ] + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {} + + result = _keepalive_from_deployment_config({"model": "slow-model"}, response) + assert result is None + + +def test_keepalive_from_deployment_config_fallback_by_name_configured_plus_unset(monkeypatch): + """A deployment that leaves keepalive_seconds unset entirely (not explicitly 0) + must not inherit a sibling deployment's configured interval: without a model_id + we can't tell which deployment served the stream, so mixing a configured + deployment with an unconfigured one is just as ambiguous as two conflicting + configured values.""" + from unittest.mock import MagicMock + + router = MagicMock() + router.get_deployment.return_value = None + router.get_model_list.return_value = [ + {"litellm_params": {"keepalive_seconds": 20.0}}, + {"litellm_params": {}}, + ] + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {} + + result = _keepalive_from_deployment_config({"model": "slow-model"}, response) + assert result is None + + +def test_keepalive_from_deployment_config_no_router_returns_none(monkeypatch): + monkeypatch.setattr(ps, "llm_router", None) + result = _keepalive_from_deployment_config({"model": "gpt-4"}, None) + assert result is None + + +def test_make_keepalive_resolver_caches_by_model_id(monkeypatch): + """The steady-state case (no fallback): every chunk shares the same + model_id, so the deployment lookup must happen once, not once per chunk.""" + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = 5.0 + deployment.litellm_params.allow_client_keepalive_override = False + + router = MagicMock() + router.get_deployment.return_value = deployment + monkeypatch.setattr(ps, "llm_router", router) + + resolve = _make_keepalive_resolver({"model": "my-model"}) + + first = _simple_chunk(content="a") + first._hidden_params = {"model_id": "deploy-steady"} + second = _simple_chunk(content="b") + second._hidden_params = {"model_id": "deploy-steady"} + + assert resolve(first) == 5.0 + assert resolve(second) == 5.0 + router.get_deployment.assert_called_once_with(model_id="deploy-steady") + + +def test_make_keepalive_resolver_reresolves_on_model_id_change(monkeypatch): + """A mid-stream fallback changes model_id: the cache must miss and + re-resolve against the new deployment, not keep serving the stale value.""" + from unittest.mock import MagicMock + + before = MagicMock() + before.litellm_params.keepalive_seconds = 5.0 + before.litellm_params.allow_client_keepalive_override = False + + after = MagicMock() + after.litellm_params.keepalive_seconds = 30.0 + after.litellm_params.allow_client_keepalive_override = False + + router = MagicMock() + router.get_deployment.side_effect = lambda model_id: {"deploy-a": before, "deploy-b": after}[model_id] + monkeypatch.setattr(ps, "llm_router", router) + + resolve = _make_keepalive_resolver({"model": "my-model"}) + + chunk_a = _simple_chunk(content="a") + chunk_a._hidden_params = {"model_id": "deploy-a"} + chunk_b = _simple_chunk(content="b") + chunk_b._hidden_params = {"model_id": "deploy-b"} + + assert resolve(chunk_a) == 5.0 + assert resolve(chunk_b) == 30.0 + assert router.get_deployment.call_count == 2 + + +def test_make_keepalive_resolver_missing_model_id_never_cached(monkeypatch): + """Without a model_id there's no reliable cache key (see the model_name + fallback in _keepalive_from_deployment_config), so every chunk must + re-resolve fresh rather than reuse a stale guess.""" + from unittest.mock import MagicMock + + router = MagicMock() + router.get_deployment.return_value = None + router.get_model_list.return_value = [{"litellm_params": {"keepalive_seconds": 12.0}}] + monkeypatch.setattr(ps, "llm_router", router) + + resolve = _make_keepalive_resolver({"model": "slow-model"}) + + chunk_a = _simple_chunk(content="a") + chunk_a._hidden_params = {} + chunk_b = _simple_chunk(content="b") + chunk_b._hidden_params = {} + + assert resolve(chunk_a) == 12.0 + assert resolve(chunk_b) == 12.0 + assert router.get_model_list.call_count == 2 + + +def test_make_keepalive_resolver_expires_cache_after_ttl(monkeypatch): + """An operator's live config change (revoking override, disabling + keepalive, removing the deployment) must be observed within + _KEEPALIVE_CACHE_TTL_SECONDS, not frozen for the rest of an + already-in-flight stream just because the model_id hasn't changed.""" + from unittest.mock import MagicMock + + before = MagicMock() + before.litellm_params.keepalive_seconds = 20.0 + before.litellm_params.allow_client_keepalive_override = False + + after = MagicMock() + after.litellm_params.keepalive_seconds = 0 + after.litellm_params.allow_client_keepalive_override = False + + router = MagicMock() + router.get_deployment.return_value = before + monkeypatch.setattr(ps, "llm_router", router) + + clock = {"t": 0.0} + monkeypatch.setattr(ps.time, "monotonic", lambda: clock["t"]) + + resolve = _make_keepalive_resolver({"model": "my-model"}) + + chunk = _simple_chunk(content="a") + chunk._hidden_params = {"model_id": "deploy-live"} + + assert resolve(chunk) == 20.0 + assert router.get_deployment.call_count == 1 + + # Still within the TTL: same model_id, cached value reused even though + # the router's live config has since changed underneath it. + router.get_deployment.return_value = after + clock["t"] = ps._KEEPALIVE_CACHE_TTL_SECONDS - 0.01 + assert resolve(chunk) == 20.0 + assert router.get_deployment.call_count == 1 + + # Past the TTL: the config-reload disable is now observed. + clock["t"] = ps._KEEPALIVE_CACHE_TTL_SECONDS + 0.01 + assert resolve(chunk) == 0.0 + assert router.get_deployment.call_count == 2 + + +def test_keepalive_seconds_in_all_litellm_params(): + from litellm.types.utils import all_litellm_params + + assert "keepalive_seconds" in all_litellm_params + + +def test_allow_client_keepalive_override_in_all_litellm_params(): + """allow_client_keepalive_override is a deployment-only control flag: if it's + missing from all_litellm_params, it leaks straight through into the actual + provider API call as an unrecognized field and gets rejected (confirmed live + against the real Anthropic API, which returns 'Extra inputs are not + permitted').""" + from litellm.types.utils import all_litellm_params + + assert "allow_client_keepalive_override" in all_litellm_params + + +@pytest.mark.asyncio +async def test_async_data_generator_emits_ping_heartbeat(monkeypatch): + """When keepalive_seconds is set on a deployment that allows client override, + ': ping' frames appear during upstream stalls.""" + import asyncio + from unittest.mock import MagicMock + + _patch_logging_flags(monkeypatch) + monkeypatch.setattr(ps, "_KEEPALIVE_MIN_SECONDS", 0.05) + + router = MagicMock() + router.get_deployment.return_value = None + router.get_model_list.return_value = [{"litellm_params": {"allow_client_keepalive_override": True}}] + monkeypatch.setattr(ps, "llm_router", router) + + async def _slow_response(): + yield _simple_chunk(content="hello") + await asyncio.sleep(0.4) + yield _simple_chunk(content="world") + + out = [] + async for line in async_data_generator( + response=_slow_response(), + user_api_key_dict=_user_auth(), + request_data={"model": "gpt-4", "keepalive_seconds": 0.05}, + ): + out.append(line) + + pings = [item for item in out if item == ": ping\n\n"] + assert len(pings) >= 2, f"expected >= 2 ping frames; got {len(pings)}" + assert out[-1] == "data: [DONE]\n\n" + + +@pytest.mark.asyncio +async def test_async_data_generator_no_keepalive_no_pings(monkeypatch): + """Without keepalive_seconds, no ': ping' frames are emitted.""" + _patch_logging_flags(monkeypatch) + + out = [] + async for line in async_data_generator( + response=_async_iter([_simple_chunk(content="hello")]), + user_api_key_dict=_user_auth(), + request_data={"model": "gpt-4"}, + ): + out.append(line) + + assert ": ping\n\n" not in out + assert out[-1] == "data: [DONE]\n\n" + + +@pytest.mark.asyncio +async def test_async_data_generator_resolves_deployment_once_per_steady_stream(monkeypatch): + """Regression test for the per-chunk resolver cost: a stream where every + real chunk comes from the same deployment (the common, no-fallback case) + must only pay for one `llm_router.get_deployment()` call, not one per + chunk. Before caching, this asserted 1 but got len(chunks) since the + resolver re-ran the full deployment lookup after every single chunk. + + The very first resolve happens on the raw `response` object before any + chunk is yielded; a bare async generator (unlike the real + CustomStreamWrapper this stands in for) can't carry `_hidden_params`, so + that one call goes through the model_name fallback instead of + `get_deployment` — hence it's asserted separately. + """ + from unittest.mock import MagicMock + + _patch_logging_flags(monkeypatch) + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = None + deployment.litellm_params.allow_client_keepalive_override = False + + router = MagicMock() + router.get_deployment.return_value = deployment + router.get_model_list.return_value = [{"litellm_params": {}}] + monkeypatch.setattr(ps, "llm_router", router) + + async def _steady_response(): + for content in ("a", "b", "c", "d", "e"): + chunk = _simple_chunk(content=content) + chunk._hidden_params = {"model_id": "deploy-steady"} + yield chunk + + out = [] + async for line in async_data_generator( + response=_steady_response(), + user_api_key_dict=_user_auth(), + request_data={"model": "gpt-4"}, + ): + out.append(line) + + assert router.get_deployment.call_count == 1 + assert router.get_model_list.call_count == 1 + assert out[-1] == "data: [DONE]\n\n" diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 803094e8d54..f48e1dba601 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -2076,6 +2076,55 @@ def test_get_num_retries_from_request(): assert result == -1 +def test_get_keepalive_seconds_from_request(): + """ + Test LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request method + """ + # Header present with valid float string + headers_with_keepalive = {"x-litellm-keepalive-seconds": "15"} + result = LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request( + headers_with_keepalive + ) + assert result == 15.0 + + # Header not present + result = LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request( + {"Content-Type": "application/json"} + ) + assert result is None + + # Empty headers dictionary + result = LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request({}) + assert result is None + + # Header present with a fractional value + result = LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request( + {"x-litellm-keepalive-seconds": "1.5"} + ) + assert result == 1.5 + + # Header present with invalid value raises ValueError, matching the other + # x-litellm-* numeric header helpers (_get_timeout_from_request, etc.) + with pytest.raises(ValueError): + LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request( + {"x-litellm-keepalive-seconds": "not-a-number"} + ) + + +def test_add_litellm_data_for_backend_llm_call_merges_keepalive_seconds_header(): + """ + The x-litellm-keepalive-seconds header must be merged into the data dict + that add_litellm_data_to_request later data.update()s onto the request body, + the same way x-litellm-timeout/x-litellm-num-retries already are. + """ + result = LiteLLMProxyRequestSetup.add_litellm_data_for_backend_llm_call( + headers={"x-litellm-keepalive-seconds": "20"}, + request_data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + ) + assert result.get("keepalive_seconds") == 20.0 + + def test_add_user_api_key_auth_to_request_metadata(): """ Test that add_user_api_key_auth_to_request_metadata properly adds user API key authentication data to request metadata diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index fa8731d7a16..6323516b126 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -26758,6 +26758,11 @@ export interface components { } | null; /** Adaptive Router Default Model */ adaptive_router_default_model?: string | null; + /** + * Allow Client Keepalive Override + * @default false + */ + allow_client_keepalive_override: boolean | null; /** Annotation Cost Per Page */ annotation_cost_per_page?: number | null; /** Api Base */ @@ -26906,6 +26911,8 @@ export interface components { input_cost_per_video_token?: number | null; /** Itpm */ itpm?: number | null; + /** Keepalive Seconds */ + keepalive_seconds?: number | null; /** Litellm Credential Name */ litellm_credential_name?: string | null; /** Litellm Trace Id */ @@ -35429,6 +35436,11 @@ export interface components { } | null; /** Adaptive Router Default Model */ adaptive_router_default_model?: string | null; + /** + * Allow Client Keepalive Override + * @default false + */ + allow_client_keepalive_override: boolean | null; /** Annotation Cost Per Page */ annotation_cost_per_page?: number | null; /** Api Base */ @@ -35577,6 +35589,8 @@ export interface components { input_cost_per_video_token?: number | null; /** Itpm */ itpm?: number | null; + /** Keepalive Seconds */ + keepalive_seconds?: number | null; /** Litellm Credential Name */ litellm_credential_name?: string | null; /** Litellm Trace Id */