From 28c1e431968d61ea7caf1d82e351aca765ec83a1 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Fri, 14 Aug 2026 00:01:16 -0700 Subject: [PATCH 01/25] feat(ui): standardize the Teams page header --- .../_components/AccessGroupsPage.tsx | 4 +- .../budgets/_components/budget_panel.tsx | 4 +- .../projects/_components/ProjectsPage.tsx | 4 +- .../src/components/Teams.test.tsx | 37 ++++++--- ui/litellm-dashboard/src/components/Teams.tsx | 47 +++++------ .../VirtualKeysPage/VirtualKeysTable.tsx | 4 +- .../shared/LegacyPageHeader.test.tsx | 33 ++++++++ .../components/shared/LegacyPageHeader.tsx | 25 ++++++ .../src/components/shared/PageHeader.test.tsx | 80 +++++++++++++++---- .../src/components/shared/PageHeader.tsx | 63 +++++++++++---- 10 files changed, 224 insertions(+), 77 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/shared/LegacyPageHeader.test.tsx create mode 100644 ui/litellm-dashboard/src/components/shared/LegacyPageHeader.tsx diff --git a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsPage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsPage.tsx index f37acb3d85a..aeff249fdd3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsPage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsPage.tsx @@ -3,7 +3,7 @@ import { useDeleteAccessGroup } from "@/app/(dashboard)/hooks/accessGroups/useDe import { Plus, SearchIcon, X } from "lucide-react"; import { useMemo, useState } from "react"; import DeleteResourceModal from "@/components/common_components/DeleteResourceModal"; -import { PageHeader } from "@/components/shared/PageHeader"; +import { LegacyPageHeader } from "@/components/shared/LegacyPageHeader"; import { Button } from "@/components/ui/button"; import { InputGroup, InputGroupAddon, InputGroupButton, InputGroupInput } from "@/components/ui/input-group"; import { AccessGroupDetail } from "./AccessGroupsDetailsPage"; @@ -61,7 +61,7 @@ export function AccessGroupsPage() { return (
- = ({ accessToken }) => { return (
- } title="Budgets" subtitle="Spend, TPM and RPM limits you can assign to customers." diff --git a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsPage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsPage.tsx index 91ad9f847f5..dc18a05edca 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsPage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/projects/_components/ProjectsPage.tsx @@ -3,7 +3,7 @@ import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; import { Plus, SearchIcon, X } from "lucide-react"; import { parseAsString, useQueryState } from "nuqs"; import { useMemo, useState } from "react"; -import { PageHeader } from "@/components/shared/PageHeader"; +import { LegacyPageHeader } from "@/components/shared/LegacyPageHeader"; import { Button } from "@/components/ui/button"; import { InputGroup, InputGroupAddon, InputGroupButton, InputGroupInput } from "@/components/ui/input-group"; import { CreateProjectModal } from "./ProjectModals/CreateProjectModal"; @@ -56,7 +56,7 @@ export function ProjectsPage() { return (
- { expect(onUrlUpdate.mock.calls.at(-1)![0].searchParams.has("team")).toBe(false); await waitFor(() => expect(screen.queryByTestId("team-info-view")).not.toBeInTheDocument()); }); + + it("should preserve the legacy inset for the team detail view", async () => { + renderWithQueryClient(, { + searchParams: "?team=team-from-url", + }); + + await waitFor(() => expect(mockTeamInfoView).toHaveBeenCalled()); + expect(screen.getByRole("main")).toHaveClass("px-12", "py-6"); + }); }); describe("Teams - Create Team CTA is grouped with the tabs on the left", () => { @@ -521,22 +530,28 @@ describe("Teams - Create Team CTA is grouped with the tabs on the left", () => { mockUseOrganizations.mockReturnValue({ data: [] }); }); - it("renders the Create Team button inside the tab bar, ahead of the tabs", () => { - const { container } = renderWithQueryClient(); + it("should render the Create Team button inside the tab bar, ahead of the tabs", () => { + renderWithQueryClient(); - const createButton = screen.getByTestId("create-team-button"); - const tabNav = container.querySelector(".ant-tabs-nav"); + const tabNav = screen.getByRole("tablist"); + const createButton = within(tabNav).getByTestId("create-team-button"); + const firstTab = within(tabNav).getByRole("tab", { name: "Your Teams" }); + const tabs = tabNav.closest(".ant-tabs"); - // The CTA lives in the tab bar's left slot, not the standalone page header. - expect(tabNav).not.toBeNull(); - expect(tabNav!.contains(createButton)).toBe(true); - - // It reads as the left end of the cluster: it precedes the first tab in DOM order. - const firstTab = screen.getByRole("tab", { name: "Your Teams" }); + expect(screen.getByRole("main")).toHaveClass("p-8"); + expect(within(tabNav).getByRole("separator")).toBeInTheDocument(); expect(createButton.compareDocumentPosition(firstTab) & Node.DOCUMENT_POSITION_FOLLOWING).toBeTruthy(); + expect(tabs).toHaveClass( + "[&>.ant-tabs-nav]:!mb-6", + "[&>.ant-tabs-nav]:before:!border-b-0", + "[&_.ant-tabs-ink-bar]:!h-0.5", + "[&_.ant-tabs-tab]:!py-[7px]", + "[&_.ant-tabs-tab+_.ant-tabs-tab]:!ml-[22px]", + "[&_.ant-tabs-tab-active]:font-semibold", + ); }); - it("omits the Create Team CTA for a role that cannot manage teams", () => { + it("should omit the Create Team CTA for a role that cannot manage teams", () => { renderWithQueryClient(); expect(screen.queryByTestId("create-team-button")).not.toBeInTheDocument(); }); diff --git a/ui/litellm-dashboard/src/components/Teams.tsx b/ui/litellm-dashboard/src/components/Teams.tsx index becbe0e2b48..5b79b067412 100644 --- a/ui/litellm-dashboard/src/components/Teams.tsx +++ b/ui/litellm-dashboard/src/components/Teams.tsx @@ -6,7 +6,7 @@ import TeamSSOSettings from "@/components/TeamSSOSettings"; import { isProxyAdminRole } from "@/utils/roles"; import { InfoCircleOutlined } from "@ant-design/icons"; import { Accordion, AccordionBody, AccordionHeader, TextInput } from "@tremor/react"; -import { Button, Form, Input, Layout, Modal, Select, Switch, Tabs, theme, Tooltip, Typography } from "antd"; +import { Button, Form, Input, Layout, Modal, Select, Switch, Tabs, Tooltip, Typography } from "antd"; import { Plus, Users } from "lucide-react"; import React, { useEffect, useState } from "react"; import { useQuery, useQueryClient } from "@tanstack/react-query"; @@ -403,7 +403,6 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser return false; }; - const { token } = theme.useToken(); const { Text } = Typography; const { Content } = Layout; @@ -474,7 +473,7 @@ const Teams: React.FC = ({ accessToken, userID, userRole, premiumUser ]; return ( - + {selectedTeamId ? ( = ({ accessToken, userID, userRole, premiumUser premiumUser={premiumUser} /> ) : ( - <> -
- } - title="Teams" - subtitle="Manage teams, members, and their access to models and budgets" + } + title="Teams" + subtitle="Manage teams, members, and their access to models and budgets" + primaryAction={ + canCreateOrManageTeams(userRole, userID, organizations) ? ( + setIsTeamModalVisible(true)} data-testid="create-team-button"> + + Create Team + + ) : undefined + } + tabs={({ leadingControls }) => ( + -
- - - setIsTeamModalVisible(true)} data-testid="create-team-button"> - - Create Team - -
-
- ) : undefined, - }} - /> - + )} + /> )} {canCreateOrManageTeams(userRole, userID, organizations) && ( diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx index fa0360c0dda..b278b4f675d 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.tsx @@ -12,7 +12,7 @@ import { DataTableToolbar, } from "@/components/shared/DataTable"; import { SearchSelect } from "@/components/shared/SearchSelect"; -import { PageHeader } from "@/components/shared/PageHeader"; +import { LegacyPageHeader } from "@/components/shared/LegacyPageHeader"; import { Input } from "@/components/ui/input"; import { useDebouncedValue } from "@tanstack/react-pacer/debouncer"; import { ColumnFiltersState, OnChangeFn, PaginationState, SortingState } from "@tanstack/react-table"; @@ -172,7 +172,7 @@ export function VirtualKeysTable({ headerActions }: VirtualKeysTableProps) { return (
- } title="Virtual Keys" subtitle="Every key that authenticates requests to the gateway." diff --git a/ui/litellm-dashboard/src/components/shared/LegacyPageHeader.test.tsx b/ui/litellm-dashboard/src/components/shared/LegacyPageHeader.test.tsx new file mode 100644 index 00000000000..a0081c1c7f1 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/LegacyPageHeader.test.tsx @@ -0,0 +1,33 @@ +import { renderWithProviders, screen } from "@/../tests/test-utils"; +import { describe, expect, it } from "vitest"; + +import { LegacyPageHeader } from "./LegacyPageHeader"; + +describe("LegacyPageHeader", () => { + it("should render the title as a heading", () => { + renderWithProviders(); + + expect(screen.getByRole("heading", { name: "Virtual Keys" })).toBeInTheDocument(); + }); + + it("should render the optional identity and actions", () => { + renderWithProviders( + Key icon} + actions={} + />, + ); + + expect(screen.getByText("Every key that authenticates requests")).toBeInTheDocument(); + expect(screen.getByText("Key icon")).toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Create New Key" })).toBeInTheDocument(); + }); + + it("should omit optional actions when none are provided", () => { + renderWithProviders(); + + expect(screen.queryByRole("button")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/shared/LegacyPageHeader.tsx b/ui/litellm-dashboard/src/components/shared/LegacyPageHeader.tsx new file mode 100644 index 00000000000..43979ad00b4 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/LegacyPageHeader.tsx @@ -0,0 +1,25 @@ +"use client"; + +import * as React from "react"; + +interface LegacyPageHeaderProps { + title: React.ReactNode; + subtitle?: React.ReactNode; + icon?: React.ReactNode; + actions?: React.ReactNode; +} + +export function LegacyPageHeader({ title, subtitle, icon, actions }: LegacyPageHeaderProps) { + return ( +
+
+ {icon != null && {icon}} +
+

{title}

+ {subtitle != null &&

{subtitle}

} +
+
+ {actions != null &&
{actions}
} +
+ ); +} diff --git a/ui/litellm-dashboard/src/components/shared/PageHeader.test.tsx b/ui/litellm-dashboard/src/components/shared/PageHeader.test.tsx index f7a313271da..3741542abad 100644 --- a/ui/litellm-dashboard/src/components/shared/PageHeader.test.tsx +++ b/ui/litellm-dashboard/src/components/shared/PageHeader.test.tsx @@ -1,31 +1,77 @@ -import { render, screen } from "@testing-library/react"; +import { renderWithProviders, screen, within } from "@/../tests/test-utils"; import { describe, expect, it } from "vitest"; import { PageHeader } from "./PageHeader"; +const identity = { + icon: Teams icon, + title: "Teams", + subtitle: "Manage teams, members, and their access to models and budgets", +}; + describe("PageHeader", () => { - it("renders the title as a heading", () => { - render(); - expect(screen.getByRole("heading", { name: "Virtual Keys" })).toBeInTheDocument(); + it("should render the page identity", () => { + renderWithProviders(); + + expect(screen.getByRole("heading", { name: "Teams" })).toBeInTheDocument(); + expect(screen.getByText("Teams icon").parentElement).toHaveAttribute("aria-hidden", "true"); + expect(screen.getByText(identity.subtitle)).toBeInTheDocument(); }); - it("renders the subtitle, icon, and actions when provided", () => { - render( + it("should apply the standard title and subtext typography", () => { + renderWithProviders(); + + const icon = screen.getByText("Teams icon").parentElement; + expect(screen.getByRole("heading", { name: "Teams" })).toHaveClass("text-2xl", "font-semibold", "tracking-tight"); + expect(screen.getByText(identity.subtitle)).toHaveClass("mt-1.5", "text-sm", "text-muted-foreground"); + expect(icon).toHaveClass("size-5", "[&_svg]:size-5", "[&_svg]:stroke-[1.75]"); + expect(icon?.parentElement).toHaveClass("gap-2.5"); + }); + + it("should render the primary action, divider, tabs, and utilities in the standard control row", () => { + renderWithProviders( } - actions={} + {...identity} + primaryAction={} + tabs={ +
+ +
+ } + utilities={} />, ); - expect(screen.getByText("Every key that authenticates requests")).toBeInTheDocument(); - expect(screen.getByTestId("icon")).toBeInTheDocument(); - expect(screen.getByRole("button", { name: "Create New Key" })).toBeInTheDocument(); + + const controls = screen.getByRole("group", { name: "Page controls" }); + expect(controls).toHaveClass("mt-5", "h-9"); + expect(within(controls).getByRole("separator")).toHaveClass("mx-4", "h-6"); + expect(controls).toHaveTextContent("Create TeamYour TeamsRefresh"); }); - it("omits the optional slots when not provided", () => { - render(); - expect(screen.queryByRole("button")).not.toBeInTheDocument(); - expect(document.querySelector("p")).toBeNull(); + it("should omit the divider when tabs are absent", () => { + renderWithProviders(Create Team} />); + + expect(screen.queryByRole("separator")).not.toBeInTheDocument(); + }); + + it("should provide standard controls to an embedded tab shell", () => { + renderWithProviders( + Create Team} + tabs={({ leadingControls, utilities }) => ( +
+ {leadingControls} + + {utilities} +
+ )} + utilities={} + />, + ); + + const tabs = screen.getByRole("tablist"); + expect(within(tabs).getByRole("separator")).toBeInTheDocument(); + expect(tabs).toHaveTextContent("Create TeamYour TeamsRefresh"); }); }); diff --git a/ui/litellm-dashboard/src/components/shared/PageHeader.tsx b/ui/litellm-dashboard/src/components/shared/PageHeader.tsx index e314e8e8bc2..81092821efc 100644 --- a/ui/litellm-dashboard/src/components/shared/PageHeader.tsx +++ b/ui/litellm-dashboard/src/components/shared/PageHeader.tsx @@ -2,24 +2,57 @@ import * as React from "react"; -interface PageHeaderProps { - title: React.ReactNode; - subtitle?: React.ReactNode; - icon?: React.ReactNode; - actions?: React.ReactNode; +import { ToolbarSeparator } from "./ToolbarSeparator"; + +interface EmbeddedTabsSlots { + leadingControls: React.ReactNode; + utilities: React.ReactNode; } -export function PageHeader({ title, subtitle, icon, actions }: PageHeaderProps) { - return ( -
-
- {icon != null && {icon}} -
-

{title}

- {subtitle != null &&

{subtitle}

} -
+interface PageHeaderProps { + title: React.ReactNode; + subtitle: React.ReactNode; + icon: React.ReactNode; + primaryAction?: React.ReactNode; + tabs?: React.ReactNode | ((slots: EmbeddedTabsSlots) => React.ReactNode); + utilities?: React.ReactNode; +} + +export function PageHeader({ title, subtitle, icon, primaryAction, tabs, utilities }: PageHeaderProps) { + const leadingControls = + primaryAction == null ? null : ( +
+ {primaryAction} + {tabs != null && }
- {actions != null &&
{actions}
} + ); + const utilityControls = utilities == null ? null :
{utilities}
; + const hasControlRow = primaryAction != null || tabs != null || utilities != null; + + return ( +
+
+ +

{title}

+
+

{subtitle}

+ + {typeof tabs === "function" ? ( +
{tabs({ leadingControls, utilities: utilityControls })}
+ ) : ( + hasControlRow && ( +
+ {leadingControls} + {tabs} + {utilityControls != null &&
{utilityControls}
} +
+ ) + )}
); } From 20a3a16c2f6ba7a7d5755cf56740d5180973ede0 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 19 Aug 2026 14:25:43 -0700 Subject: [PATCH 02/25] fix(proxy): populate deployment fields on failed-request spend logs from the standard logging payload --- .../spend_tracking/spend_tracking_utils.py | 30 ++++++-- .../test_spend_tracking_utils.py | 77 +++++++++++++++++++ 2 files changed, 102 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 3146d8bccfb..0b56f0d8246 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -216,6 +216,15 @@ def _extract_usage_for_ocr_call(response_obj: Any, response_obj_dict: dict) -> d return {} +def _sl_attribution_fallback( + standard_logging_payload: StandardLoggingPayload | None, + field: Literal["model_id", "model_group", "api_base", "custom_llm_provider"], +) -> str: + if standard_logging_payload is None: + return "" + return standard_logging_payload.get(field) or "" + + def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogsPayload: if kwargs is None: kwargs = {} @@ -288,8 +297,15 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs ): # use 'tags' from standard logging payload instead request_tags = safe_dumps(standard_logging_payload["request_tags"]) - _model_id: Final = metadata.get("model_info", {}).get("id", "") - _model_group: Final = metadata.get("model_group", "") + _model_id: Final = metadata.get("model_info", {}).get("id", "") or _sl_attribution_fallback( + standard_logging_payload, "model_id" + ) + _model_group: Final = metadata.get("model_group", "") or _sl_attribution_fallback( + standard_logging_payload, "model_group" + ) + _api_base: Final = litellm_params.get("api_base", "") or _sl_attribution_fallback( + standard_logging_payload, "api_base" + ) # Extract overhead from hidden_params if available litellm_overhead_time_ms = None @@ -389,7 +405,11 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs # Extract agent_id for A2A requests (set directly on model_call_details) agent_id: Final[str | None] = kwargs.get("agent_id") or metadata.get("agent_id") - custom_llm_provider: Final = kwargs.get("custom_llm_provider") + custom_llm_provider: Final = ( + kwargs.get("custom_llm_provider") + or _sl_attribution_fallback(standard_logging_payload, "custom_llm_provider") + or None + ) raw_model: Final = cast(str, kwargs.get("model") or "") model_name: Final = reconstruct_model_name(raw_model, custom_llm_provider, metadata or {}) @@ -414,13 +434,13 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs completion_tokens=usage.get("completion_tokens", standard_logging_completion_tokens), request_tags=request_tags, end_user=end_user_id or "", - api_base=litellm_params.get("api_base", ""), + api_base=_api_base, model_group=_model_group, model_id=_model_id, mcp_namespaced_tool_name=mcp_namespaced_tool_name, agent_id=agent_id, requester_ip_address=clean_metadata.get("requester_ip_address", None), - custom_llm_provider=kwargs.get("custom_llm_provider", ""), + custom_llm_provider=custom_llm_provider or "", messages=_get_messages_for_spend_logs_payload( standard_logging_payload=standard_logging_payload, metadata=metadata ), diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index e5add059260..9710dc44e99 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -3164,3 +3164,80 @@ def test_batch_cost_row_id_is_stable_across_repeated_accounting(): ] assert ids[0] == ids[1] == "batch_same_batch_cost" + + +def _make_failed_request_standard_logging_payload() -> StandardLoggingPayload: + base: Final = _make_standard_logging_payload_with_usage_object(usage_object={}) + return cast( + StandardLoggingPayload, + { + **base, + "status": "failure", + "call_type": "aresponses", + "model_id": "mid-123", + "model_group": "group-x", + "api_base": "https://api.openai.com/v1/responses", + "custom_llm_provider": "openai", + }, + ) + + +def test_get_logging_payload_failed_request_falls_back_to_standard_logging_payload(): + """Failed-request kwargs from the proxy failure hook carry no deployment info + (LIT-5795), so the attribution columns must come from the failure-time + standard_logging_object.""" + payload = get_logging_payload( + kwargs={ + "model": "group-x", + "litellm_params": {"metadata": {"user_api_key": "test-key", "status": "failure"}}, + "standard_logging_object": _make_failed_request_standard_logging_payload(), + }, + response_obj={}, + start_time=datetime.datetime.now(timezone.utc), + end_time=datetime.datetime.now(timezone.utc), + ) + assert payload["model_id"] == "mid-123" + assert payload["model_group"] == "group-x" + assert payload["api_base"] == "https://api.openai.com/v1/responses" + assert payload["custom_llm_provider"] == "openai" + + +def test_get_logging_payload_request_kwargs_win_over_standard_logging_payload(): + payload = get_logging_payload( + kwargs={ + "model": "group-y", + "custom_llm_provider": "anthropic", + "litellm_params": { + "api_base": "https://kwargs.example.com", + "metadata": { + "user_api_key": "test-key", + "model_group": "kwargs-group", + "model_info": {"id": "kwargs-mid"}, + }, + }, + "standard_logging_object": _make_failed_request_standard_logging_payload(), + }, + response_obj={}, + start_time=datetime.datetime.now(timezone.utc), + end_time=datetime.datetime.now(timezone.utc), + ) + assert payload["model_id"] == "kwargs-mid" + assert payload["model_group"] == "kwargs-group" + assert payload["api_base"] == "https://kwargs.example.com" + assert payload["custom_llm_provider"] == "anthropic" + + +def test_get_logging_payload_failed_request_without_standard_logging_payload_leaves_fields_empty(): + payload = get_logging_payload( + kwargs={ + "model": "group-x", + "litellm_params": {"metadata": {"user_api_key": "test-key", "status": "failure"}}, + }, + response_obj={}, + start_time=datetime.datetime.now(timezone.utc), + end_time=datetime.datetime.now(timezone.utc), + ) + assert payload["model_id"] == "" + assert payload["model_group"] == "" + assert payload["api_base"] == "" + assert payload["custom_llm_provider"] == "" From 7a6a677b72cd3822e3fe0ae29329fa8b0e17d577 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 19 Aug 2026 15:09:05 -0700 Subject: [PATCH 03/25] feat(proxy): enqueued-token rate limiting for batches with refund on completion and cancellation --- litellm/constants.py | 6 + litellm/proxy/hooks/batch_enqueued_tokens.py | 357 ++++++++++++++++++ litellm/proxy/hooks/batch_rate_limiter.py | 80 +++- .../hooks/parallel_request_limiter_v3.py | 42 +++ tests/e2e/batches/test_batches_e2e.py | 146 ++++++- .../coverage_registry/quota_management.yaml | 3 + tests/e2e/models.py | 1 + .../proxy/hooks/test_batch_enqueued_tokens.py | 190 ++++++++++ .../proxy/hooks/test_batch_file_validation.py | 192 ++++++++++ .../hooks/test_parallel_request_limiter_v3.py | 120 ++++++ 10 files changed, 1133 insertions(+), 4 deletions(-) create mode 100644 litellm/proxy/hooks/batch_enqueued_tokens.py create mode 100644 tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py diff --git a/litellm/constants.py b/litellm/constants.py index 39a49e55f0d..d7ddacf5fac 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1766,6 +1766,12 @@ PTU_LAPSED_ALERT_LIMIT: Final[int] = 10 # one is seconds old, so a few minutes separates them. PTU_PRUNE_SKEW_GRACE_SECONDS: Final[int] = 300 +# How long enqueued-token reservations for batches live without a refund. Providers +# complete or expire batches within their completion window (24h for OpenAI), so a +# reservation still unrefunded after 8 days belongs to a batch whose terminal state +# was never observed (e.g. proxy restart); expiry returns the tokens to the caller. +BATCH_ENQUEUED_TOKEN_TTL_SECONDS: Final[int] = 8 * 24 * 60 * 60 + # Shared read-only empty mapping, for defaulting optional Mapping parameters without # constructing a fresh mutable dict at each call site. EMPTY_MAPPING: Final = MappingProxyType({}) diff --git a/litellm/proxy/hooks/batch_enqueued_tokens.py b/litellm/proxy/hooks/batch_enqueued_tokens.py new file mode 100644 index 00000000000..040c8b687a7 --- /dev/null +++ b/litellm/proxy/hooks/batch_enqueued_tokens.py @@ -0,0 +1,357 @@ +""" +Enqueued-token accounting for batch submissions. + +Opt-in via ``batch_enqueued_token_limit`` in key or team metadata: batch +submissions reserve their estimated token count against a long-lived +enqueued-token allowance instead of the per-minute rate-limit windows, and +the reservation is refunded when the batch reaches a terminal state +(completed, failed, expired, or cancelled). +""" + +import asyncio +from collections.abc import Awaitable, Mapping, Sequence +from dataclasses import dataclass +from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeAlias + +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError + +from litellm._logging import verbose_proxy_logger +from litellm.constants import BATCH_ENQUEUED_TOKEN_TTL_SECONDS +from litellm.proxy._types import UserAPIKeyAuth + +if TYPE_CHECKING: + from opentelemetry.trace import Span as _Span + + from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache + + Span = _Span + InternalUsageCache = _InternalUsageCache + +BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY: Final = "batch_enqueued_token_limit" + +BATCH_ENQUEUED_REFUND_STATUSES: Final[frozenset[str]] = frozenset( + {"completed", "complete", "failed", "expired", "cancelled", "cancelling"} +) + +ScopeKey: TypeAlias = Literal["api_key", "team"] + +RESERVE_ENQUEUED_TOKENS_SCRIPT: Final = """ +local amount = tonumber(ARGV[1]) +local ttl = tonumber(ARGV[2]) +for i = 1, #KEYS do + local limit = tonumber(ARGV[2 + i]) + local current = tonumber(redis.call('GET', KEYS[i]) or '0') + if current + amount > limit then + return {0, i - 1, current} + end +end +for i = 1, #KEYS do + redis.call('INCRBY', KEYS[i], amount) + redis.call('EXPIRE', KEYS[i], ttl) +end +return {1, -1, 0} +""" + +REFUND_ENQUEUED_TOKENS_SCRIPT: Final = """ +local amount = tonumber(ARGV[1]) +for i = 1, #KEYS do + local updated = redis.call('DECRBY', KEYS[i], amount) + if updated <= 0 then + redis.call('DEL', KEYS[i]) + end +end +return 1 +""" + +SAVE_RESERVATION_SCRIPT: Final = """ +redis.call('SET', KEYS[1], ARGV[1], 'EX', tonumber(ARGV[2])) +return 1 +""" + +POP_RESERVATION_SCRIPT: Final = """ +local value = redis.call('GET', KEYS[1]) +if value then + redis.call('DEL', KEYS[1]) +end +return value +""" + + +@dataclass(frozen=True, slots=True) +class BatchEnqueuedTokenScope: + key: ScopeKey + value: str + limit: int + + +@dataclass(frozen=True, slots=True) +class BatchEnqueuedTokenReservation: + tokens: int + scopes: tuple[BatchEnqueuedTokenScope, ...] + + +@dataclass(frozen=True, slots=True) +class BatchEnqueuedTokenOverLimit: + scope: BatchEnqueuedTokenScope + enqueued: int + + +BatchEnqueuedTokenOutcome: TypeAlias = BatchEnqueuedTokenReservation | BatchEnqueuedTokenOverLimit + +_LIMIT_ADAPTER: Final = TypeAdapter(Annotated[int, Field(gt=0)]) +_RESERVE_RESULT_ADAPTER: Final = TypeAdapter(tuple[int, int, int]) +_POPPED_VALUE_ADAPTER: Final = TypeAdapter(str | bytes | None) +_STORED_COUNTER_ADAPTER: Final = TypeAdapter(int | None) +_RESERVATION_ADAPTER: Final = TypeAdapter(BatchEnqueuedTokenReservation) + + +class _ScriptRunner(Protocol): + def __call__(self, keys: Sequence[str], args: Sequence[str | bytes | int | float]) -> Awaitable[object]: ... + + +def _read_metadata_limit(metadata: Mapping[str, object] | None) -> int | None: + if not metadata: + return None + raw: Final = metadata.get(BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY) + if raw is None: + return None + try: + return _LIMIT_ADAPTER.validate_python(raw) + except ValidationError: + verbose_proxy_logger.warning( + "Ignoring invalid %s value %r; expected a positive integer", + BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY, + raw, + ) + return None + + +def resolve_batch_enqueued_token_scopes( + user_api_key_dict: UserAPIKeyAuth, +) -> tuple[BatchEnqueuedTokenScope, ...]: + key_limit: Final = _read_metadata_limit(user_api_key_dict.metadata) + team_limit: Final = _read_metadata_limit(user_api_key_dict.team_metadata) + candidates: Final = ( + BatchEnqueuedTokenScope(key="api_key", value=user_api_key_dict.api_key, limit=key_limit) + if key_limit is not None and user_api_key_dict.api_key + else None, + BatchEnqueuedTokenScope(key="team", value=user_api_key_dict.team_id, limit=team_limit) + if team_limit is not None and user_api_key_dict.team_id + else None, + ) + return tuple(scope for scope in candidates if scope is not None) + + +def canonical_provider_batch_id(batch_id: str) -> str: + from litellm.proxy.openai_files_endpoints.common_utils import ( + _is_base64_encoded_unified_file_id, # pyright: ignore[reportPrivateUsage] # canonical unified-id decoder has no public wrapper + get_batch_id_from_unified_batch_id, + get_original_file_id, + ) + + decoded: Final = _is_base64_encoded_unified_file_id(batch_id) + if isinstance(decoded, str): + if "llm_batch_id" in decoded or "generic_response_id" in decoded: + return get_batch_id_from_unified_batch_id(decoded) + return decoded + return get_original_file_id(batch_id) + + +class _BatchResponseView(BaseModel): + model_config = ConfigDict(extra="ignore") + + id: str + status: str + object: Literal["batch"] + + +def batch_response_view(response: object) -> _BatchResponseView | None: + try: + return _BatchResponseView.model_validate(response, from_attributes=True) + except ValidationError: + return None + + +class BatchEnqueuedTokenStore: + """Tracks enqueued batch tokens per scope, plus per-batch reservation records for refunds. + + Counters and records live in Redis (via atomic Lua scripts) when Redis is + configured; otherwise a single-process in-memory fallback guarded by one + asyncio lock is used. Everything expires after + ``BATCH_ENQUEUED_TOKEN_TTL_SECONDS`` so a crash between submission and the + terminal-state refund can never leak tokens forever. + """ + + def __init__(self, internal_usage_cache: "InternalUsageCache") -> None: + self.internal_usage_cache = internal_usage_cache + self._lock = asyncio.Lock() + redis_cache = internal_usage_cache.dual_cache.redis_cache + self._reserve_script: _ScriptRunner | None = ( + redis_cache.async_register_script(RESERVE_ENQUEUED_TOKENS_SCRIPT) if redis_cache is not None else None + ) + self._refund_script: _ScriptRunner | None = ( + redis_cache.async_register_script(REFUND_ENQUEUED_TOKENS_SCRIPT) if redis_cache is not None else None + ) + self._save_script: _ScriptRunner | None = ( + redis_cache.async_register_script(SAVE_RESERVATION_SCRIPT) if redis_cache is not None else None + ) + self._pop_script: _ScriptRunner | None = ( + redis_cache.async_register_script(POP_RESERVATION_SCRIPT) if redis_cache is not None else None + ) + + @staticmethod + def _counter_key(scope: BatchEnqueuedTokenScope) -> str: + return f"batch_enqueued_tokens:{scope.key}:{scope.value}" + + @staticmethod + def _record_key(batch_id: str) -> str: + return f"batch_enqueued_token_reservation:{batch_id}" + + async def reserve( + self, + tokens: int, + scopes: tuple[BatchEnqueuedTokenScope, ...], + litellm_parent_otel_span: "Span | None" = None, + ) -> BatchEnqueuedTokenOutcome: + if tokens <= 0 or not scopes: + return BatchEnqueuedTokenReservation(tokens=max(tokens, 0), scopes=scopes) + if self._reserve_script is not None: + try: + raw_result = await self._reserve_script( + tuple(self._counter_key(scope) for scope in scopes), + (tokens, BATCH_ENQUEUED_TOKEN_TTL_SECONDS, *(scope.limit for scope in scopes)), + ) + result = _RESERVE_RESULT_ADAPTER.validate_python(raw_result) + if result[0] == 1: + return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes) + return BatchEnqueuedTokenOverLimit(scope=scopes[result[1]], enqueued=result[2]) + except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory counters + verbose_proxy_logger.warning( + "Redis enqueued-token reserve failed, falling back to in-memory: %s", str(e) + ) + return await self._reserve_in_memory(tokens=tokens, scopes=scopes, span=litellm_parent_otel_span) + + async def _reserve_in_memory( + self, + tokens: int, + scopes: tuple[BatchEnqueuedTokenScope, ...], + span: "Span | None", + ) -> BatchEnqueuedTokenOutcome: + async with self._lock: + currents: Final = tuple([await self._get_local_counter(scope, span) for scope in scopes]) + for scope, current in zip(scopes, currents): + if current + tokens > scope.limit: + return BatchEnqueuedTokenOverLimit(scope=scope, enqueued=current) + for scope, current in zip(scopes, currents): + await self._set_local_counter(scope, current + tokens, span) + return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes) + + async def refund( + self, + reservation: BatchEnqueuedTokenReservation, + litellm_parent_otel_span: "Span | None" = None, + ) -> None: + if reservation.tokens <= 0 or not reservation.scopes: + return + if self._refund_script is not None: + try: + await self._refund_script( + tuple(self._counter_key(scope) for scope in reservation.scopes), + (reservation.tokens,), + ) + except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory counters + verbose_proxy_logger.warning( + "Redis enqueued-token refund failed, falling back to in-memory: %s", str(e) + ) + else: + return + async with self._lock: + for scope in reservation.scopes: + current = await self._get_local_counter(scope, litellm_parent_otel_span) + remaining = current - reservation.tokens + if remaining <= 0: + self.internal_usage_cache.dual_cache.in_memory_cache.delete_cache(key=self._counter_key(scope)) + else: + await self._set_local_counter(scope, remaining, litellm_parent_otel_span) + + async def save_reservation( + self, + batch_id: str, + reservation: BatchEnqueuedTokenReservation, + litellm_parent_otel_span: "Span | None" = None, + ) -> None: + serialized: Final = _RESERVATION_ADAPTER.dump_json(reservation).decode("utf-8") + if self._save_script is not None: + try: + await self._save_script( + (self._record_key(batch_id),), + (serialized, BATCH_ENQUEUED_TOKEN_TTL_SECONDS), + ) + except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory record + verbose_proxy_logger.warning( + "Redis enqueued-token reservation save failed, falling back to in-memory: %s", str(e) + ) + else: + return + await self.internal_usage_cache.async_set_cache( + key=self._record_key(batch_id), + value=serialized, + ttl=BATCH_ENQUEUED_TOKEN_TTL_SECONDS, + litellm_parent_otel_span=litellm_parent_otel_span, + local_only=True, + ) + + async def pop_reservation( + self, + batch_id: str, + litellm_parent_otel_span: "Span | None" = None, + ) -> BatchEnqueuedTokenReservation | None: + raw: object = None + if self._pop_script is not None: + try: + raw = _POPPED_VALUE_ADAPTER.validate_python(await self._pop_script((self._record_key(batch_id),), ())) + except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory record + verbose_proxy_logger.warning( + "Redis enqueued-token reservation pop failed, falling back to in-memory: %s", str(e) + ) + raw = await self._pop_local_record(batch_id, litellm_parent_otel_span) + else: + raw = await self._pop_local_record(batch_id, litellm_parent_otel_span) + if raw is None: + return None + try: + if isinstance(raw, (str, bytes)): + return _RESERVATION_ADAPTER.validate_json(raw) + return _RESERVATION_ADAPTER.validate_python(raw) + except ValidationError: + verbose_proxy_logger.warning("Discarding malformed enqueued-token reservation record for %s", batch_id) + return None + + async def _pop_local_record(self, batch_id: str, span: "Span | None") -> object: + async with self._lock: + stored = await self.internal_usage_cache.async_get_cache( + key=self._record_key(batch_id), + litellm_parent_otel_span=span, + local_only=True, + ) + if stored is None: + return None + self.internal_usage_cache.dual_cache.in_memory_cache.delete_cache(key=self._record_key(batch_id)) + return stored + + async def _get_local_counter(self, scope: BatchEnqueuedTokenScope, span: "Span | None") -> int: + stored = await self.internal_usage_cache.async_get_cache( + key=self._counter_key(scope), + litellm_parent_otel_span=span, + local_only=True, + ) + return _STORED_COUNTER_ADAPTER.validate_python(stored) or 0 + + async def _set_local_counter(self, scope: BatchEnqueuedTokenScope, value: int, span: "Span | None") -> None: + await self.internal_usage_cache.async_set_cache( + key=self._counter_key(scope), + value=value, + ttl=BATCH_ENQUEUED_TOKEN_TTL_SECONDS, + litellm_parent_otel_span=span, + local_only=True, + ) diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index d6229fb80a6..5b814ad28fd 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -46,9 +46,16 @@ from litellm.proxy.common_utils.proxy_rate_limit_error import ( ProxyRateLimitError, map_v3_rate_limit_type, ) +from litellm.proxy.hooks.batch_enqueued_tokens import ( + BatchEnqueuedTokenOverLimit, + BatchEnqueuedTokenReservation, + BatchEnqueuedTokenScope, + resolve_batch_enqueued_token_scopes, +) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY, + get_or_create_request_stash, ) from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit @@ -291,6 +298,7 @@ class _PROXY_BatchRateLimiter(CustomLogger): self, data: dict, user_api_key_dict: UserAPIKeyAuth, + has_enqueued_scopes: bool = False, ) -> tuple[bool, list["RateLimitDescriptor"] | None]: """ Skip downloading batch input files when the operator disabled batch @@ -343,8 +351,10 @@ class _PROXY_BatchRateLimiter(CustomLogger): user_api_key_dict=user_api_key_dict, data=data, ) - if not self._has_applicable_batch_rate_limits(descriptors) and not self._project_has_any_io_token_limits( - user_api_key_dict + if ( + not has_enqueued_scopes + and not self._has_applicable_batch_rate_limits(descriptors) + and not self._project_has_any_io_token_limits(user_api_key_dict) ): verbose_proxy_logger.debug("Skipping batch input file processing: no rate limits configured") return True, None @@ -511,6 +521,59 @@ class _PROXY_BatchRateLimiter(CustomLogger): return file_id, fetch_kwargs + async def _reserve_batch_enqueued_tokens( + self, + user_api_key_dict: UserAPIKeyAuth, + data: Mapping[str, object], + batch_usage: BatchFileUsage, + scopes: tuple[BatchEnqueuedTokenScope, ...], + ) -> None: + """Reserve the batch's estimated tokens against the caller's enqueued-token allowance. + + Runs instead of the per-minute counter charge when the key or team + opted in via ``batch_enqueued_token_limit`` metadata. The reservation + is stashed on the request so the v3 limiter's post-call hooks can + persist it (keyed by the provider batch id) and refund it when the + batch reaches a terminal state. + """ + outcome: Final = await self.parallel_request_limiter.batch_enqueued_token_store.reserve( + tokens=batch_usage.total_tokens, + scopes=scopes, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + ) + match outcome: + case BatchEnqueuedTokenOverLimit(): + self._raise_enqueued_limit_error(over_limit=outcome, data=data, batch_usage=batch_usage) + case BatchEnqueuedTokenReservation(): + get_or_create_request_stash().batch_enqueued_reservation = outcome + + def _raise_enqueued_limit_error( + self, + over_limit: BatchEnqueuedTokenOverLimit, + data: Mapping[str, object], + batch_usage: BatchFileUsage, + ) -> NoReturn: + scope: Final = over_limit.scope + remaining: Final = max(0, scope.limit - over_limit.enqueued) + detail: Final = ( + f"Batch enqueued token limit exceeded for {scope.key}: {scope.value}. " + f"Batch requires {batch_usage.total_tokens} tokens but only {remaining} enqueued tokens remaining " + f"out of {scope.limit} enqueued token limit. " + f"Tokens free up as running batches complete or are cancelled." + ) + raw_model: Final = data.get("model") + resolved_model, llm_provider = resolve_llm_provider_for_rate_limit( + raw_model if isinstance(raw_model, str) else None + ) + raise ProxyRateLimitError( + detail=detail, + headers=MappingProxyType({"rate_limit_type": "tokens"}), + category=RateLimitErrorCategory.LITELLM_BATCH_RATE_LIMIT, + rate_limit_type=map_v3_rate_limit_type("tokens"), + model=resolved_model, + llm_provider=llm_provider, + ) + def _raise_rate_limit_error( self, status: "RateLimitStatus", @@ -1039,8 +1102,9 @@ class _PROXY_BatchRateLimiter(CustomLogger): verbose_proxy_logger.debug("No input_file_id in batch request, skipping rate limiting") return data + enqueued_scopes: Final = resolve_batch_enqueued_token_scopes(user_api_key_dict) should_skip, batch_rate_limit_descriptors = self._should_skip_batch_input_file_processing( - data=data, user_api_key_dict=user_api_key_dict + data=data, user_api_key_dict=user_api_key_dict, has_enqueued_scopes=bool(enqueued_scopes) ) if should_skip: return data @@ -1066,6 +1130,16 @@ class _PROXY_BatchRateLimiter(CustomLogger): data["_batch_token_count"] = batch_usage.total_tokens data["_batch_request_count"] = batch_usage.request_count + if enqueued_scopes: + await self._reserve_batch_enqueued_tokens( + user_api_key_dict=user_api_key_dict, + data=data, + batch_usage=batch_usage, + scopes=enqueued_scopes, + ) + verbose_proxy_logger.debug("Batch enqueued-token reservation succeeded") + return data + # Directly increment counters by batch amounts (check happens atomically) # This will raise HTTPException if limits are exceeded await self._check_and_increment_batch_counters( diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 2d799ded752..b295da69255 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -44,6 +44,13 @@ from litellm.proxy.common_utils.proxy_rate_limit_error import ( ProxyRateLimitError, map_v3_rate_limit_type, ) +from litellm.proxy.hooks.batch_enqueued_tokens import ( + BATCH_ENQUEUED_REFUND_STATUSES, + BatchEnqueuedTokenReservation, + BatchEnqueuedTokenStore, + batch_response_view, + canonical_provider_batch_id, +) from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject, ResponseAPIUsage @@ -515,6 +522,7 @@ class RequestRateLimiterStash: otpm_reserved_window_identities: frozenset[tuple[str, str, Literal["redis", "local"]]] = field( default_factory=frozenset ) + batch_enqueued_reservation: BatchEnqueuedTokenReservation | None = None reservation_released: bool = False @@ -619,6 +627,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): # Batch rate limiter (lazy loaded) self._batch_rate_limiter: CallTypeRateLimiter | None = None + self.batch_enqueued_token_store = BatchEnqueuedTokenStore(internal_usage_cache=internal_usage_cache) # Serializes multi-phase check+increment sequences (batch + dynamic # limiters) within this process to close the TOCTOU window between @@ -4673,6 +4682,32 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): except Exception as e: verbose_proxy_logger.exception("Error in rate limit post-call hook: %s", e) + try: + await self._handle_batch_enqueued_post_call(user_api_key_dict=user_api_key_dict, response=response) + except Exception as e: # noqa: BLE001 # post-call batch accounting must never fail the response + verbose_proxy_logger.exception("Error in batch enqueued-token post-call hook: %s", e) + + async def _handle_batch_enqueued_post_call(self, user_api_key_dict: UserAPIKeyAuth, response: object) -> None: + view: Final = batch_response_view(response) + if view is None: + return + span: Final = user_api_key_dict.parent_otel_span + stash: Final = get_request_stash() + if stash is not None and stash.batch_enqueued_reservation is not None: + await self.batch_enqueued_token_store.save_reservation( + batch_id=canonical_provider_batch_id(view.id), + reservation=stash.batch_enqueued_reservation, + litellm_parent_otel_span=span, + ) + stash.batch_enqueued_reservation = None + if view.status in BATCH_ENQUEUED_REFUND_STATUSES: + popped: Final = await self.batch_enqueued_token_store.pop_reservation( + batch_id=canonical_provider_batch_id(view.id), + litellm_parent_otel_span=span, + ) + if popped is not None: + await self.batch_enqueued_token_store.refund(reservation=popped, litellm_parent_otel_span=span) + async def async_post_call_failure_hook( self, request_data: dict, @@ -4706,6 +4741,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) stash.parallel_slot = None + if stash.batch_enqueued_reservation is not None: + await self.batch_enqueued_token_store.refund( + reservation=stash.batch_enqueued_reservation, + litellm_parent_otel_span=user_api_key_dict.parent_otel_span, + ) + stash.batch_enqueued_reservation = None + if stash.reservation_released: return reserved_tokens: Final = stash.reserved_tokens diff --git a/tests/e2e/batches/test_batches_e2e.py b/tests/e2e/batches/test_batches_e2e.py index 1376bdbed38..53bf9739983 100644 --- a/tests/e2e/batches/test_batches_e2e.py +++ b/tests/e2e/batches/test_batches_e2e.py @@ -17,6 +17,7 @@ from __future__ import annotations import json import os +import re import time from datetime import datetime, timedelta, timezone from typing import Callable @@ -57,7 +58,7 @@ from e2e_http import ( unwrap, ) from lifecycle import ResourceManager -from models import KeyGenerateBody, LiteLLMParamsBody, SpendLogRow +from models import KeyGenerateBody, KeyMetadata, LiteLLMParamsBody, SpendLogRow pytestmark = pytest.mark.e2e @@ -685,6 +686,149 @@ class TestBatchRateLimitErrorMapping: ) +BATCH_ENQUEUED_HEADROOM_TOKENS = 100_000 +_BATCH_REQUIRES_TOKENS = re.compile(r"Batch requires (\d+) tokens") + + +class TestBatchEnqueuedTokenLimit: + """Opt-in enqueued-token allowance governs batch submission instead of RPM/TPM. + + A key whose metadata carries batch_enqueued_token_limit reserves the batch's + token estimate against that allowance at create time: per-minute limits no + longer gate batch submission, exhausting the allowance rejects the create + before it reaches the provider, and cancelling a running batch refunds its + reservation so blocked submissions go through again (LIT-5273). + """ + + def _upload_batch_file( + self, client: BatchClient, resources: ResourceManager, key: str + ) -> FileObject: + file = unwrap( + client.upload_file( + content=_multi_request_jsonl("gpt-4o-mini", BATCH_RL_REQUEST_LINES), + form=FileUploadForm(purpose="batch"), + model=OPENAI_BATCH_MODEL, + key=key, + ) + ) + resources.defer(quietly(lambda: client.delete_file(file.id, key=key))) + return file + + def _generate_enqueued_key( + self, + client: BatchClient, + resources: ResourceManager, + *, + limit: int, + marker: str, + rpm_limit: int | None = None, + ) -> str: + key = client.proxy.generate_key( + KeyGenerateBody( + models=[], + rpm_limit=rpm_limit, + user_id=f"e2e-batch-enq-{marker}-{unique_marker()}", + metadata=KeyMetadata(batch_enqueued_token_limit=limit), + ) + ) + resources.defer(lambda: client.proxy.delete_key(key)) + return key + + @pytest.mark.covers( + "quota_management.ratelimit.batch_enqueued_tokens.accepts_over_rpm", + exercised_on=["batches"], + ) + def test_enqueued_allowance_accepts_batch_over_key_rpm( + self, client: BatchClient, resources: ResourceManager, batch_deployments: None + ) -> None: + key = self._generate_enqueued_key( + client, + resources, + limit=BATCH_ENQUEUED_HEADROOM_TOKENS, + marker="rpm", + rpm_limit=BATCH_RL_RPM_LIMIT, + ) + file = self._upload_batch_file(client, resources, key) + + created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key) + + assert created.status_code != 429, ( + f"enqueued-token allowance must govern batch submission instead of the " + f"key RPM ({BATCH_RL_RPM_LIMIT} < {BATCH_RL_REQUEST_LINES} rows); " + f"got 429: {created.body[:400]}" + ) + require_successful_call(created) + batch = BatchObject.model_validate_json(created.body) + resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key))) + + @pytest.mark.covers( + "quota_management.ratelimit.batch_enqueued_tokens.blocks_when_exhausted", + exercised_on=["batches"], + ) + @pytest.mark.covers( + "quota_management.ratelimit.batch_enqueued_tokens.refunds_on_cancel", + exercised_on=["batches"], + ) + def test_exhausted_allowance_blocks_until_cancel_refunds( + self, client: BatchClient, resources: ResourceManager, batch_deployments: None + ) -> None: + sizing_key = self._generate_enqueued_key( + client, resources, limit=1, marker="size" + ) + sizing_file = self._upload_batch_file(client, resources, sizing_key) + sized = client.create_batch( + body=BatchCreateBody(input_file_id=sizing_file.id), key=sizing_key + ) + assert sized.status_code == 429, ( + f"a 1-token allowance must reject any batch before it reaches the " + f"provider, got {sized.status_code}: {sized.body[:400]}" + ) + assert "batch enqueued token limit exceeded" in sized.body.lower(), ( + f"429 body must name the enqueued token limit, got: {sized.body[:400]}" + ) + requires = _BATCH_REQUIRES_TOKENS.search(sized.body) + assert requires is not None, ( + f"429 body must report the batch token requirement so callers can size " + f"allowances, got: {sized.body[:400]}" + ) + batch_tokens = int(requires.group(1)) + assert batch_tokens > 1 + + key = self._generate_enqueued_key( + client, resources, limit=batch_tokens + batch_tokens // 2, marker="refund" + ) + file = self._upload_batch_file(client, resources, key) + + first = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key) + require_successful_call(first) + first_batch = BatchObject.model_validate_json(first.body) + resources.defer(quietly(lambda: client.cancel_batch(first_batch.id, key=key))) + + blocked = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key) + assert blocked.status_code == 429, ( + f"second batch must not fit the remaining allowance while the first is " + f"enqueued, got {blocked.status_code}: {blocked.body[:400]}" + ) + assert "batch enqueued token limit exceeded" in blocked.body.lower(), ( + f"429 body must name the enqueued token limit, got: {blocked.body[:400]}" + ) + + cancelled = cancel_batch(client, first_batch.id, key=key, provider=None) + assert cancelled.status in {"cancelling", "cancelled"}, ( + f"cancel must reach a cancel state for the refund to fire, " + f"got {cancelled.status}" + ) + + retried = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key) + assert retried.status_code != 429, ( + f"cancelling the first batch must refund its reservation so the retry " + f"fits the allowance, got 429: {retried.body[:400]}" + ) + require_successful_call(retried) + retry_batch = BatchObject.model_validate_json(retried.body) + resources.defer(quietly(lambda: client.cancel_batch(retry_batch.id, key=key))) + + ASSUME_ROLE_RAW_MODEL = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" diff --git a/tests/e2e/coverage_registry/quota_management.yaml b/tests/e2e/coverage_registry/quota_management.yaml index 2dfa7adddea..4b8aa1da002 100644 --- a/tests/e2e/coverage_registry/quota_management.yaml +++ b/tests/e2e/coverage_registry/quota_management.yaml @@ -2,6 +2,9 @@ # litellm/proxy/hooks/ + litellm/proxy/auth/auth_checks.py + litellm/proxy/spend_tracking/. - {id: quota_management.ratelimit.rpm.blocks_over_limit, module: quota_management, tier: P0, behavior: ratelimit, variant: rpm, assertions: [blocks_over_limit], exercised_on: [chat_completions, messages], source: "parallel_request_limiter_v3.py", rationale: "v3 limiter enforces RPM per key/team/model; 429 on breach"} - {id: quota_management.ratelimit.batch_rpm.blocks_over_limit, module: quota_management, tier: P0, behavior: ratelimit, variant: batch_rpm, assertions: [blocks_over_limit], exercised_on: [batches], source: "batch_rate_limiter.py", rationale: "Batch create that exceeds key RPM returns mapped 429 with retry-after"} +- {id: quota_management.ratelimit.batch_enqueued_tokens.accepts_over_rpm, module: quota_management, tier: P0, behavior: ratelimit, variant: batch_enqueued_tokens, assertions: [accepts_over_rpm], exercised_on: [batches], source: "batch_rate_limiter.py + batch_enqueued_tokens.py", rationale: "Key with an enqueued-token allowance submits a batch whose row count exceeds its RPM and the create is accepted"} +- {id: quota_management.ratelimit.batch_enqueued_tokens.blocks_when_exhausted, module: quota_management, tier: P0, behavior: ratelimit, variant: batch_enqueued_tokens, assertions: [blocks_when_exhausted], exercised_on: [batches], source: "batch_rate_limiter.py + batch_enqueued_tokens.py", rationale: "Batch create is rejected with a 429 naming the enqueued token limit once the allowance cannot fit the file"} +- {id: quota_management.ratelimit.batch_enqueued_tokens.refunds_on_cancel, module: quota_management, tier: P0, behavior: ratelimit, variant: batch_enqueued_tokens, assertions: [refunds_on_cancel], exercised_on: [batches], source: "batch_rate_limiter.py + batch_enqueued_tokens.py", rationale: "Cancelling a running batch returns its reserved tokens so a previously blocked submission succeeds"} - {id: quota_management.ratelimit.tpm.blocks_over_limit, module: quota_management, tier: P0, behavior: ratelimit, variant: tpm, assertions: [blocks_over_limit], exercised_on: [chat_completions, messages], source: "parallel_request_limiter_v3.py", rationale: "v3 limiter enforces TPM per key/team/model; 429 on breach"} - {id: quota_management.ratelimit.tpm.excludes_cached_tokens, module: quota_management, tier: P0, behavior: ratelimit, variant: tpm, assertions: [excludes_cached_tokens], exercised_on: [chat_completions], source: "parallel_request_limiter_v3.py:_get_total_tokens_from_usage", rationale: "Cached prompt tokens must not count toward TPM (LIT-1930)"} - {id: quota_management.ratelimit.redis_backed.blocks_over_limit, module: quota_management, tier: P0, behavior: ratelimit, variant: redis_backed, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "parallel_request_limiter_v3.py", rationale: "With Redis configured, RPM still enforces 429 across the shared limiter path customers run multi-replica"} diff --git a/tests/e2e/models.py b/tests/e2e/models.py index df5cb841fad..150be8966ee 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -46,6 +46,7 @@ class KeyLoggingCallback(BaseModel): class KeyMetadata(BaseModel): logging: list[KeyLoggingCallback] | None = None priority: str | None = None + batch_enqueued_token_limit: int | None = None class ObjectPermission(BaseModel): diff --git a/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py new file mode 100644 index 00000000000..a917206f33a --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py @@ -0,0 +1,190 @@ +""" +LIT-5273: enqueued-token accounting for batch submissions. + +Covers the ``BatchEnqueuedTokenStore`` (reserve / refund / reservation +records), the metadata-driven scope resolution, and the batch-id and +response-shape helpers the v3 limiter's post-call hooks rely on. +""" + +import base64 +import socket +import uuid +from types import SimpleNamespace + +import pytest + +from litellm.caching.caching import DualCache +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.hooks.batch_enqueued_tokens import ( + BatchEnqueuedTokenOverLimit, + BatchEnqueuedTokenReservation, + BatchEnqueuedTokenScope, + BatchEnqueuedTokenStore, + batch_response_view, + canonical_provider_batch_id, + resolve_batch_enqueued_token_scopes, +) +from litellm.proxy.utils import InternalUsageCache + + +def _in_memory_store() -> BatchEnqueuedTokenStore: + return BatchEnqueuedTokenStore(internal_usage_cache=InternalUsageCache(DualCache(default_in_memory_ttl=60))) + + +def _scope(limit: int, key: str = "api_key") -> BatchEnqueuedTokenScope: + return BatchEnqueuedTokenScope(key=key, value=f"{key}-{uuid.uuid4().hex}", limit=limit) + + +def test_scope_resolution_reads_key_and_team_metadata(): + user = UserAPIKeyAuth( + api_key="hashed-key", + metadata={"batch_enqueued_token_limit": 100}, + team_id="team-1", + team_metadata={"batch_enqueued_token_limit": "150"}, + ) + scopes = resolve_batch_enqueued_token_scopes(user) + assert scopes == ( + BatchEnqueuedTokenScope(key="api_key", value="hashed-key", limit=100), + BatchEnqueuedTokenScope(key="team", value="team-1", limit=150), + ) + + +def test_scope_resolution_returns_empty_without_opt_in(): + assert resolve_batch_enqueued_token_scopes(UserAPIKeyAuth(api_key="k")) == () + assert resolve_batch_enqueued_token_scopes(UserAPIKeyAuth(api_key="k", metadata={}, team_metadata=None)) == () + + +@pytest.mark.parametrize("bad_value", ["not-a-number", 0, -5, None, [1000]]) +def test_scope_resolution_ignores_invalid_limits(bad_value): + user = UserAPIKeyAuth(api_key="k", metadata={"batch_enqueued_token_limit": bad_value}) + assert resolve_batch_enqueued_token_scopes(user) == () + + +def test_scope_resolution_skips_team_scope_without_team_id(): + user = UserAPIKeyAuth(api_key="k", team_metadata={"batch_enqueued_token_limit": 100}) + assert resolve_batch_enqueued_token_scopes(user) == () + + +@pytest.mark.asyncio +async def test_reserve_rejects_once_allowance_is_exhausted(): + store = _in_memory_store() + scope = _scope(limit=100) + first = await store.reserve(tokens=80, scopes=(scope,)) + assert isinstance(first, BatchEnqueuedTokenReservation) + second = await store.reserve(tokens=30, scopes=(scope,)) + assert second == BatchEnqueuedTokenOverLimit(scope=scope, enqueued=80) + third = await store.reserve(tokens=20, scopes=(scope,)) + assert isinstance(third, BatchEnqueuedTokenReservation) + + +@pytest.mark.asyncio +async def test_reserve_is_all_or_nothing_across_scopes(): + store = _in_memory_store() + key_scope = _scope(limit=100, key="api_key") + team_scope = _scope(limit=50, key="team") + over = await store.reserve(tokens=60, scopes=(key_scope, team_scope)) + assert over == BatchEnqueuedTokenOverLimit(scope=team_scope, enqueued=0) + exact_fit = await store.reserve(tokens=50, scopes=(key_scope, team_scope)) + assert isinstance(exact_fit, BatchEnqueuedTokenReservation) + + +@pytest.mark.asyncio +async def test_refund_restores_allowance_and_never_goes_negative(): + store = _in_memory_store() + scope = _scope(limit=100) + reservation = await store.reserve(tokens=30, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + await store.refund(reservation) + await store.refund(reservation) + refill = await store.reserve(tokens=100, scopes=(scope,)) + assert isinstance(refill, BatchEnqueuedTokenReservation) + assert isinstance(await store.reserve(tokens=1, scopes=(scope,)), BatchEnqueuedTokenOverLimit) + + +@pytest.mark.asyncio +async def test_reservation_record_roundtrip_pops_exactly_once(): + store = _in_memory_store() + scope = _scope(limit=100) + reservation = await store.reserve(tokens=40, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + await store.save_reservation("batch_abc", reservation) + assert await store.pop_reservation("batch_abc") == reservation + assert await store.pop_reservation("batch_abc") is None + assert await store.pop_reservation("batch_never_saved") is None + + +@pytest.mark.asyncio +async def test_zero_token_reserve_charges_nothing(): + store = _in_memory_store() + scope = _scope(limit=100) + empty = await store.reserve(tokens=0, scopes=(scope,)) + assert empty == BatchEnqueuedTokenReservation(tokens=0, scopes=(scope,)) + full = await store.reserve(tokens=100, scopes=(scope,)) + assert isinstance(full, BatchEnqueuedTokenReservation) + + +def test_canonical_provider_batch_id_passes_raw_ids_through(): + assert canonical_provider_batch_id("batch_abc123") == "batch_abc123" + + +def test_canonical_provider_batch_id_decodes_unified_batch_ids(): + unified = "litellm_proxy;model_id:m-1;llm_batch_id:batch_prov_9;llm_output_file_id:file-9" + encoded = base64.urlsafe_b64encode(unified.encode()).decode().rstrip("=") + assert canonical_provider_batch_id(encoded) == "batch_prov_9" + + +def test_canonical_provider_batch_id_decodes_model_embedded_ids(): + from litellm.proxy.openai_files_endpoints.common_utils import encode_file_id_with_model + + encoded = encode_file_id_with_model(file_id="batch_prov_7", model="my-alias", id_type="batch") + assert canonical_provider_batch_id(encoded) == "batch_prov_7" + + +def test_batch_response_view_accepts_batch_objects_only(): + batch = SimpleNamespace(id="batch_1", status="completed", object="batch") + view = batch_response_view(batch) + assert view is not None and view.id == "batch_1" and view.status == "completed" + assert batch_response_view({"id": "chatcmpl-1", "object": "chat.completion"}) is None + assert batch_response_view(None) is None + assert batch_response_view("batch_1") is None + + +def _local_redis_port() -> int | None: + for port in (6379,): + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock: + sock.settimeout(0.2) + if sock.connect_ex(("127.0.0.1", port)) == 0: + return port + return None + + +@pytest.mark.asyncio +@pytest.mark.skipif(_local_redis_port() is None, reason="requires a local Redis on 6379 for the Lua script path") +async def test_redis_lua_path_full_lifecycle(): + from litellm.caching.redis_cache import RedisCache + + port = _local_redis_port() + redis_cache = RedisCache(host="127.0.0.1", port=port) + store = BatchEnqueuedTokenStore( + internal_usage_cache=InternalUsageCache(DualCache(redis_cache=redis_cache, default_in_memory_ttl=60)) + ) + key_scope = _scope(limit=100, key="api_key") + team_scope = _scope(limit=50, key="team") + + over = await store.reserve(tokens=60, scopes=(key_scope, team_scope)) + assert over == BatchEnqueuedTokenOverLimit(scope=team_scope, enqueued=0) + + reservation = await store.reserve(tokens=50, scopes=(key_scope, team_scope)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + assert isinstance(await store.reserve(tokens=1, scopes=(key_scope, team_scope)), BatchEnqueuedTokenOverLimit) + + batch_id = f"batch_{uuid.uuid4().hex}" + await store.save_reservation(batch_id, reservation) + popped = await store.pop_reservation(batch_id) + assert popped == reservation + assert await store.pop_reservation(batch_id) is None + + await store.refund(popped) + refill = await store.reserve(tokens=50, scopes=(key_scope, team_scope)) + assert isinstance(refill, BatchEnqueuedTokenReservation) + await store.refund(refill) diff --git a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py index 1ce1a2f3e51..8698ee8ba12 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py +++ b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py @@ -2094,3 +2094,195 @@ def test_estimate_entry_output_tokens_multiplies_candidate_count(body_extra, exp } assert rate_limiter._estimate_entry_output_tokens(entry, None) == expected + + +# --------------------------------------------------------------------------- +# LIT-5273: enqueued-token limits govern batch submission when opted in +# --------------------------------------------------------------------------- + + +def _enqueued_rate_limiter(): + from litellm import DualCache + from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, + ) + from litellm.proxy.utils import InternalUsageCache + + local_cache = DualCache(default_in_memory_ttl=60) + internal_usage_cache = InternalUsageCache(local_cache) + parallel_request_limiter = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=internal_usage_cache) + rate_limiter = _PROXY_BatchRateLimiter( + internal_usage_cache=internal_usage_cache, + parallel_request_limiter=parallel_request_limiter, + ) + return rate_limiter, local_cache + + +_ENQUEUED_BATCH_FILE_CONTENT = ( + b'{"body": {"model": "gpt-4o-mini", "max_tokens": 40, "messages": [{"role": "user", "content": "hi"}]}}\n' + b'{"body": {"model": "gpt-4o-mini", "max_tokens": 40, "messages": [{"role": "user", "content": "hi"}]}}\n' + b'{"body": {"model": "gpt-4o-mini", "max_tokens": 40, "messages": [{"role": "user", "content": "hi"}]}}\n' +) + + +def _enqueued_batch_patches(): + mock_content = MagicMock() + mock_content.content = _ENQUEUED_BATCH_FILE_CONTENT + afile_content_mock = AsyncMock(return_value=mock_content) + return afile_content_mock, ( + patch("litellm.proxy.proxy_server.general_settings", {}), + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch( + "litellm.proxy.openai_files_endpoints.common_utils.get_credentials_for_model", + return_value={"custom_llm_provider": "openai"}, + ), + ) + + +@pytest.mark.asyncio +async def test_enqueued_limit_accepts_batch_over_per_minute_limits(): + """The headline LIT-5273 behavior: a key that opted into an enqueued-token + allowance submits a batch whose row count and token count both exceed its + per-minute RPM/TPM limits, and the batch is accepted (repeatedly) because + only the enqueued allowance governs. Without the opt-in the same key is + rejected on RPM before the batch reaches the provider.""" + from litellm.proxy.hooks.parallel_request_limiter_v3 import get_request_stash + + rate_limiter, local_cache = _enqueued_rate_limiter() + afile_content_mock, patches = _enqueued_batch_patches() + + legacy_user = UserAPIKeyAuth(api_key="sk-legacy-rpm", models=["*"], rpm_limit=1, tpm_limit=10) + opted_in_user = UserAPIKeyAuth( + api_key="sk-enqueued-rpm", + models=["*"], + rpm_limit=1, + tpm_limit=10, + metadata={"batch_enqueued_token_limit": 100000}, + ) + + with patches[0], patches[1], patches[2], patch("litellm.afile_content", new=afile_content_mock): + with pytest.raises(HTTPException) as legacy_exc: + await rate_limiter.async_pre_call_hook( + user_api_key_dict=legacy_user, + cache=local_cache, + data={"input_file_id": "file-abc123", "model": "gpt-4o-mini"}, + call_type="acreate_batch", + ) + assert legacy_exc.value.status_code == 429 + + first_data = {"input_file_id": "file-abc123", "model": "gpt-4o-mini"} + result = await rate_limiter.async_pre_call_hook( + user_api_key_dict=opted_in_user, + cache=local_cache, + data=first_data, + call_type="acreate_batch", + ) + assert result is first_data + stash = get_request_stash() + assert stash is not None and stash.batch_enqueued_reservation is not None + assert stash.batch_enqueued_reservation.tokens == first_data["_batch_token_count"] > 0 + + second = await rate_limiter.async_pre_call_hook( + user_api_key_dict=opted_in_user, + cache=local_cache, + data={"input_file_id": "file-abc123", "model": "gpt-4o-mini"}, + call_type="acreate_batch", + ) + assert second is not None + + +@pytest.mark.asyncio +async def test_enqueued_limit_rejects_when_allowance_is_exhausted(): + """Submissions are rejected pre-provider once the enqueued allowance can't + fit the batch, even for a key with no per-minute limits at all (which + previously skipped batch rate limiting entirely).""" + rate_limiter, local_cache = _enqueued_rate_limiter() + afile_content_mock, patches = _enqueued_batch_patches() + + sizing_user = UserAPIKeyAuth( + api_key="sk-enqueued-sizing", models=["*"], metadata={"batch_enqueued_token_limit": 1000000} + ) + with patches[0], patches[1], patches[2], patch("litellm.afile_content", new=afile_content_mock): + sizing_data = {"input_file_id": "file-abc123", "model": "gpt-4o-mini"} + await rate_limiter.async_pre_call_hook( + user_api_key_dict=sizing_user, + cache=local_cache, + data=sizing_data, + call_type="acreate_batch", + ) + batch_tokens = sizing_data["_batch_token_count"] + assert batch_tokens > 0 + + capped_user = UserAPIKeyAuth( + api_key="sk-enqueued-capped", + models=["*"], + metadata={"batch_enqueued_token_limit": batch_tokens + batch_tokens // 2}, + ) + await rate_limiter.async_pre_call_hook( + user_api_key_dict=capped_user, + cache=local_cache, + data={"input_file_id": "file-abc123", "model": "gpt-4o-mini"}, + call_type="acreate_batch", + ) + with pytest.raises(HTTPException) as exc: + await rate_limiter.async_pre_call_hook( + user_api_key_dict=capped_user, + cache=local_cache, + data={"input_file_id": "file-abc123", "model": "gpt-4o-mini"}, + call_type="acreate_batch", + ) + + assert exc.value.status_code == 429 + assert "Batch enqueued token limit exceeded for api_key" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_enqueued_team_limit_applies_to_batch_submission(): + rate_limiter, local_cache = _enqueued_rate_limiter() + afile_content_mock, patches = _enqueued_batch_patches() + + team_user = UserAPIKeyAuth( + api_key="sk-enqueued-team-key", + models=["*"], + team_id="team-enqueued-batch", + team_metadata={"batch_enqueued_token_limit": 10}, + ) + with patches[0], patches[1], patches[2], patch("litellm.afile_content", new=afile_content_mock): + with pytest.raises(HTTPException) as exc: + await rate_limiter.async_pre_call_hook( + user_api_key_dict=team_user, + cache=local_cache, + data={"input_file_id": "file-abc123", "model": "gpt-4o-mini"}, + call_type="acreate_batch", + ) + + assert exc.value.status_code == 429 + assert "Batch enqueued token limit exceeded for team: team-enqueued-batch" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_disable_flag_still_skips_batch_processing_with_enqueued_limits(): + rate_limiter, local_cache = _enqueued_rate_limiter() + afile_content_mock, _ = _enqueued_batch_patches() + + opted_in_user = UserAPIKeyAuth( + api_key="sk-enqueued-disabled", + models=["*"], + metadata={"batch_enqueued_token_limit": 10}, + ) + with ( + patch("litellm.proxy.proxy_server.general_settings", {"disable_batch_input_file_rate_limiting": True}), + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch("litellm.afile_content", new=afile_content_mock), + ): + data = {"input_file_id": "file-abc123", "model": "gpt-4o-mini"} + result = await rate_limiter.async_pre_call_hook( + user_api_key_dict=opted_in_user, + cache=local_cache, + data=data, + call_type="acreate_batch", + ) + + assert result is data + afile_content_mock.assert_not_awaited() diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 54366226dfb..68c219fd72d 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -5858,3 +5858,123 @@ async def test_conflicting_token_limits_cannot_bypass_tpm_reservation(): ) assert exc_info.value.status_code == 429 + + +# --------------------------------------------------------------------------- +# LIT-5273: batch enqueued-token reservations in the post-call hooks +# --------------------------------------------------------------------------- + + +def _enqueued_test_handler() -> _PROXY_MaxParallelRequestsHandler: + return _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache(default_in_memory_ttl=60))) + + +def _batch_response(batch_id: str, status: str): + from types import SimpleNamespace + + return SimpleNamespace(id=batch_id, status=status, object="batch") + + +@pytest.mark.asyncio +async def test_success_hook_persists_batch_enqueued_reservation_and_refunds_on_completion(): + from litellm.proxy.hooks.batch_enqueued_tokens import ( + BatchEnqueuedTokenOverLimit, + BatchEnqueuedTokenReservation, + BatchEnqueuedTokenScope, + ) + + handler = _enqueued_test_handler() + store = handler.batch_enqueued_token_store + scope = BatchEnqueuedTokenScope(key="api_key", value="hashed-enqueued-key", limit=100) + user = UserAPIKeyAuth(api_key="hashed-enqueued-key") + + reservation = await store.reserve(tokens=60, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + get_or_create_request_stash().batch_enqueued_reservation = reservation + + await handler.async_post_call_success_hook( + data={}, user_api_key_dict=user, response=_batch_response("batch_enq_1", "validating") + ) + assert get_request_stash().batch_enqueued_reservation is None + assert isinstance(await store.reserve(tokens=50, scopes=(scope,)), BatchEnqueuedTokenOverLimit) + + await handler.async_post_call_success_hook( + data={}, user_api_key_dict=user, response=_batch_response("batch_enq_1", "completed") + ) + refill = await store.reserve(tokens=40, scopes=(scope,)) + assert isinstance(refill, BatchEnqueuedTokenReservation) + + await handler.async_post_call_success_hook( + data={}, user_api_key_dict=user, response=_batch_response("batch_enq_1", "completed") + ) + assert isinstance(await store.reserve(tokens=70, scopes=(scope,)), BatchEnqueuedTokenOverLimit) + + +@pytest.mark.asyncio +async def test_success_hook_refunds_batch_enqueued_reservation_on_cancellation(): + from litellm.proxy.hooks.batch_enqueued_tokens import ( + BatchEnqueuedTokenReservation, + BatchEnqueuedTokenScope, + ) + + handler = _enqueued_test_handler() + store = handler.batch_enqueued_token_store + scope = BatchEnqueuedTokenScope(key="team", value="team-enqueued", limit=100) + user = UserAPIKeyAuth(api_key="hashed-enqueued-key", team_id="team-enqueued") + + reservation = await store.reserve(tokens=90, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + get_or_create_request_stash().batch_enqueued_reservation = reservation + await handler.async_post_call_success_hook( + data={}, user_api_key_dict=user, response=_batch_response("batch_enq_2", "validating") + ) + + await handler.async_post_call_success_hook( + data={}, user_api_key_dict=user, response=_batch_response("batch_enq_2", "cancelling") + ) + assert isinstance(await store.reserve(tokens=100, scopes=(scope,)), BatchEnqueuedTokenReservation) + + +@pytest.mark.asyncio +async def test_failure_hook_refunds_stashed_batch_enqueued_reservation(): + from litellm.proxy.hooks.batch_enqueued_tokens import ( + BatchEnqueuedTokenReservation, + BatchEnqueuedTokenScope, + ) + + handler = _enqueued_test_handler() + store = handler.batch_enqueued_token_store + scope = BatchEnqueuedTokenScope(key="api_key", value="hashed-failing-key", limit=100) + user = UserAPIKeyAuth(api_key="hashed-failing-key") + + reservation = await store.reserve(tokens=80, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + get_or_create_request_stash().batch_enqueued_reservation = reservation + + await handler.async_post_call_failure_hook( + request_data={}, original_exception=Exception("guardrail rejected"), user_api_key_dict=user + ) + assert get_request_stash().batch_enqueued_reservation is None + assert isinstance(await store.reserve(tokens=100, scopes=(scope,)), BatchEnqueuedTokenReservation) + + +@pytest.mark.asyncio +async def test_success_hook_leaves_stash_untouched_for_non_batch_responses(): + from litellm.proxy.hooks.batch_enqueued_tokens import ( + BatchEnqueuedTokenReservation, + BatchEnqueuedTokenScope, + ) + + handler = _enqueued_test_handler() + store = handler.batch_enqueued_token_store + scope = BatchEnqueuedTokenScope(key="api_key", value="hashed-chat-key", limit=100) + user = UserAPIKeyAuth(api_key="hashed-chat-key") + + reservation = await store.reserve(tokens=10, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + get_or_create_request_stash().batch_enqueued_reservation = reservation + + await handler.async_post_call_success_hook( + data={}, user_api_key_dict=user, response=ModelResponse(usage=Usage(total_tokens=5)) + ) + assert get_request_stash().batch_enqueued_reservation == reservation From b477d0967a4bf1f4a869c70ddf5a740f1762ae86 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 19 Aug 2026 15:34:38 -0700 Subject: [PATCH 04/25] fix(proxy): lift standard_logging_object onto request_data before the logging object is popped --- litellm/proxy/utils.py | 52 ++++++++++--------- tests/test_litellm/proxy/test_proxy_utils.py | 53 ++++++++++++++++++++ 2 files changed, 82 insertions(+), 23 deletions(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index a743526e975..58ac7882cb8 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -517,6 +517,34 @@ def _failure_usage_to_lift( return estimated_usage, 0.0 +_EMPTY_LIFT: Final = MappingProxyType({}) + + +def _failure_fields_to_lift(request_data: Mapping[str, object]) -> Mapping[str, object]: + """Failure-path callbacks run after ``litellm_logging_obj`` is popped from + request_data (it is not serialisable), so the caller merges these fields + onto request_data first: the first-handoff instant for preprocessing + latency, recovered or estimated usage for token counts, and the standard + logging object for deployment attribution on failed-request spend logs.""" + _logging_obj: Final = request_data.get("litellm_logging_obj") + if _logging_obj is None: + return _EMPTY_LIFT + _model_call_details: Final = getattr(_logging_obj, "model_call_details", {}) + _first_handoff: Final = _model_call_details.get("first_api_call_start_time") + _usage_to_lift: Final = _failure_usage_to_lift( + model_call_details=_model_call_details, + request_body=request_data, + dispatched=_first_handoff is not None, + ) + _entries: Final = ( + ("first_api_call_start_time", _first_handoff), + ("combined_usage_object", None if _usage_to_lift is None else _usage_to_lift[0]), + ("response_cost", None if _usage_to_lift is None else _usage_to_lift[1]), + ("standard_logging_object", _model_call_details.get("standard_logging_object")), + ) + return MappingProxyType({key: value for key, value in _entries if value is not None}) + + @dataclass(frozen=True) class _CallbackCapabilities: """Cached per-hook capability flags derived from ``litellm.callbacks``. @@ -2294,29 +2322,7 @@ class ProxyLogging: original_exception=original_exception, ) - # Lift the first-handoff instant onto request_data (top-level - # internal key, not metadata) so failure-path callbacks can still - # compute preprocessing latency after the logging object is popped. - _logging_obj: Final = request_data.get("litellm_logging_obj") - if _logging_obj is not None: - _model_call_details: Final = getattr(_logging_obj, "model_call_details", {}) - _first_handoff: Final = _model_call_details.get("first_api_call_start_time") - if _first_handoff is not None: - request_data["first_api_call_start_time"] = _first_handoff - - # Lift recovered partial-stream usage, or an estimated input-side - # usage for a dispatched failure, onto request_data so the - # failure-path spend callbacks (which run after the logging object - # is popped) record real token counts instead of zero. - _usage_to_lift: Final = _failure_usage_to_lift( - model_call_details=_model_call_details, - request_body=request_data, - dispatched=_first_handoff is not None, - ) - if _usage_to_lift is not None: - _lifted_usage, _lifted_cost = _usage_to_lift - request_data["combined_usage_object"] = _lifted_usage - request_data["response_cost"] = _lifted_cost + request_data.update(_failure_fields_to_lift(request_data)) # Remove before callbacks iterate — not serialisable request_data.pop("litellm_logging_obj", None) diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index d6cf0e30139..c80130b44da 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -478,6 +478,59 @@ class TestPostCallFailureHookLiftsRecoveredPartialSpend: assert "response_cost" not in request_data +class TestPostCallFailureHookLiftsStandardLoggingObject: + """Failure callbacks read standard_logging_object from request_data, but + post_call_failure_hook pops litellm_logging_obj before they run. The hook + must lift the logging obj's standard_logging_object onto request_data so + failed-request spend logs keep deployment attribution (LIT-5795). + """ + + async def _run(self, request_data): + from unittest.mock import AsyncMock, patch + + from litellm.proxy._types import UserAPIKeyAuth + + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + proxy_logging_obj.alert_types = [] + with patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()): + await proxy_logging_obj.post_call_failure_hook( + request_data=request_data, + original_exception=Exception("boom"), + user_api_key_dict=UserAPIKeyAuth(), + ) + + @pytest.mark.asyncio + async def test_lifts_standard_logging_object(self): + sl_object = {"model_id": "mid-123", "model_group": "group-x"} + logging_obj = MagicMock() + logging_obj.model_call_details = {"standard_logging_object": sl_object} + request_data = {"litellm_logging_obj": logging_obj, "metadata": {}} + await self._run(request_data) + assert request_data["standard_logging_object"] is sl_object + assert "litellm_logging_obj" not in request_data + + @pytest.mark.asyncio + async def test_logging_obj_value_overwrites_preexisting_key(self): + authoritative = {"model_id": "from-logging-obj"} + logging_obj = MagicMock() + logging_obj.model_call_details = {"standard_logging_object": authoritative} + request_data = { + "litellm_logging_obj": logging_obj, + "standard_logging_object": {"model_id": "client-injected"}, + "metadata": {}, + } + await self._run(request_data) + assert request_data["standard_logging_object"] is authoritative + + @pytest.mark.asyncio + async def test_no_standard_logging_object_is_noop(self): + logging_obj = MagicMock() + logging_obj.model_call_details = {} + request_data = {"litellm_logging_obj": logging_obj, "metadata": {}} + await self._run(request_data) + assert "standard_logging_object" not in request_data + + class TestPostCallFailureHookEstimatesDispatchedInputTokens: """A non-stream request that failed after dispatch (timeout, provider error) consumed provider-billed input tokens but recovered no usage. From 160d3dac42ddbcde36a538303bfc9612f6e02487 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 19 Aug 2026 15:40:15 -0700 Subject: [PATCH 05/25] fix(proxy): issue enqueued-token Lua calls one key at a time for Redis Cluster compatibility --- litellm/proxy/hooks/batch_enqueued_tokens.py | 85 +++++++++++-------- .../proxy/hooks/test_batch_enqueued_tokens.py | 60 ++++++++++++- 2 files changed, 109 insertions(+), 36 deletions(-) diff --git a/litellm/proxy/hooks/batch_enqueued_tokens.py b/litellm/proxy/hooks/batch_enqueued_tokens.py index 040c8b687a7..e484b6a59fe 100644 --- a/litellm/proxy/hooks/batch_enqueued_tokens.py +++ b/litellm/proxy/hooks/batch_enqueued_tokens.py @@ -38,27 +38,20 @@ ScopeKey: TypeAlias = Literal["api_key", "team"] RESERVE_ENQUEUED_TOKENS_SCRIPT: Final = """ local amount = tonumber(ARGV[1]) local ttl = tonumber(ARGV[2]) -for i = 1, #KEYS do - local limit = tonumber(ARGV[2 + i]) - local current = tonumber(redis.call('GET', KEYS[i]) or '0') - if current + amount > limit then - return {0, i - 1, current} - end +local limit = tonumber(ARGV[3]) +local current = tonumber(redis.call('GET', KEYS[1]) or '0') +if current + amount > limit then + return {0, current} end -for i = 1, #KEYS do - redis.call('INCRBY', KEYS[i], amount) - redis.call('EXPIRE', KEYS[i], ttl) -end -return {1, -1, 0} +local updated = redis.call('INCRBY', KEYS[1], amount) +redis.call('EXPIRE', KEYS[1], ttl) +return {1, updated} """ REFUND_ENQUEUED_TOKENS_SCRIPT: Final = """ -local amount = tonumber(ARGV[1]) -for i = 1, #KEYS do - local updated = redis.call('DECRBY', KEYS[i], amount) - if updated <= 0 then - redis.call('DEL', KEYS[i]) - end +local updated = redis.call('DECRBY', KEYS[1], tonumber(ARGV[1])) +if updated <= 0 then + redis.call('DEL', KEYS[1]) end return 1 """ @@ -99,7 +92,7 @@ class BatchEnqueuedTokenOverLimit: BatchEnqueuedTokenOutcome: TypeAlias = BatchEnqueuedTokenReservation | BatchEnqueuedTokenOverLimit _LIMIT_ADAPTER: Final = TypeAdapter(Annotated[int, Field(gt=0)]) -_RESERVE_RESULT_ADAPTER: Final = TypeAdapter(tuple[int, int, int]) +_RESERVE_RESULT_ADAPTER: Final = TypeAdapter(tuple[int, int]) _POPPED_VALUE_ADAPTER: Final = TypeAdapter(str | bytes | None) _STORED_COUNTER_ADAPTER: Final = TypeAdapter(int | None) _RESERVATION_ADAPTER: Final = TypeAdapter(BatchEnqueuedTokenReservation) @@ -175,9 +168,11 @@ def batch_response_view(response: object) -> _BatchResponseView | None: class BatchEnqueuedTokenStore: """Tracks enqueued batch tokens per scope, plus per-batch reservation records for refunds. - Counters and records live in Redis (via atomic Lua scripts) when Redis is - configured; otherwise a single-process in-memory fallback guarded by one - asyncio lock is used. Everything expires after + Counters and records live in Redis when Redis is configured, through + single-key Lua scripts issued one scope at a time (Redis Cluster safe: no + cross-slot commands), with an over-limit scope rolling back the scopes + reserved before it; otherwise a single-process in-memory fallback guarded + by one asyncio lock is used. Everything expires after ``BATCH_ENQUEUED_TOKEN_TTL_SECONDS`` so a crash between submission and the terminal-state refund can never leak tokens forever. """ @@ -215,22 +210,44 @@ class BatchEnqueuedTokenStore: ) -> BatchEnqueuedTokenOutcome: if tokens <= 0 or not scopes: return BatchEnqueuedTokenReservation(tokens=max(tokens, 0), scopes=scopes) - if self._reserve_script is not None: + reserve_script: Final = self._reserve_script + refund_script: Final = self._refund_script + if reserve_script is not None and refund_script is not None: try: - raw_result = await self._reserve_script( - tuple(self._counter_key(scope) for scope in scopes), - (tokens, BATCH_ENQUEUED_TOKEN_TTL_SECONDS, *(scope.limit for scope in scopes)), - ) - result = _RESERVE_RESULT_ADAPTER.validate_python(raw_result) - if result[0] == 1: - return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes) - return BatchEnqueuedTokenOverLimit(scope=scopes[result[1]], enqueued=result[2]) + return await self._reserve_via_redis(reserve_script, refund_script, tokens=tokens, scopes=scopes) except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory counters verbose_proxy_logger.warning( "Redis enqueued-token reserve failed, falling back to in-memory: %s", str(e) ) return await self._reserve_in_memory(tokens=tokens, scopes=scopes, span=litellm_parent_otel_span) + async def _reserve_via_redis( + self, + reserve_script: _ScriptRunner, + refund_script: _ScriptRunner, + tokens: int, + scopes: tuple[BatchEnqueuedTokenScope, ...], + ) -> BatchEnqueuedTokenOutcome: + for index, scope in enumerate(scopes): + raw_result = await reserve_script( + (self._counter_key(scope),), + (tokens, BATCH_ENQUEUED_TOKEN_TTL_SECONDS, scope.limit), + ) + result = _RESERVE_RESULT_ADAPTER.validate_python(raw_result) + if result[0] != 1: + await self._refund_via_redis(refund_script, tokens=tokens, scopes=scopes[:index]) + return BatchEnqueuedTokenOverLimit(scope=scope, enqueued=result[1]) + return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes) + + async def _refund_via_redis( + self, + refund_script: _ScriptRunner, + tokens: int, + scopes: tuple[BatchEnqueuedTokenScope, ...], + ) -> None: + for scope in scopes: + await refund_script((self._counter_key(scope),), (tokens,)) + async def _reserve_in_memory( self, tokens: int, @@ -253,12 +270,10 @@ class BatchEnqueuedTokenStore: ) -> None: if reservation.tokens <= 0 or not reservation.scopes: return - if self._refund_script is not None: + refund_script: Final = self._refund_script + if refund_script is not None: try: - await self._refund_script( - tuple(self._counter_key(scope) for scope in reservation.scopes), - (reservation.tokens,), - ) + await self._refund_via_redis(refund_script, tokens=reservation.tokens, scopes=reservation.scopes) except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory counters verbose_proxy_logger.warning( "Redis enqueued-token refund failed, falling back to in-memory: %s", str(e) diff --git a/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py index a917206f33a..1bb4798eaa0 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py +++ b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py @@ -9,7 +9,9 @@ response-shape helpers the v3 limiter's post-call hooks rely on. import base64 import socket import uuid -from types import SimpleNamespace +from collections.abc import Mapping, Sequence +from types import MappingProxyType, SimpleNamespace +from typing import Final import pytest @@ -123,6 +125,62 @@ async def test_zero_token_reserve_charges_nothing(): assert isinstance(full, BatchEnqueuedTokenReservation) +class _SingleKeyRedisFake: + """Emulates the Redis script path one single-key call at a time, recording every call.""" + + def __init__(self) -> None: + self.script_calls: tuple[tuple[str, tuple[str, ...]], ...] = () + self.counters: Mapping[str, int] = MappingProxyType({}) + + def async_register_script(self, script: str): + kind: Final = "reserve" if "INCRBY" in script else "refund" if "DECRBY" in script else "record" + + async def run(keys: Sequence[str], args: Sequence[str | bytes | int | float]) -> object: + self.script_calls = (*self.script_calls, (kind, tuple(keys))) + return self._run(kind, tuple(keys), tuple(args)) + + return run + + def _run(self, kind: str, keys: tuple[str, ...], args: tuple[str | bytes | int | float, ...]) -> object: + if kind == "reserve": + amount, limit = int(args[0]), int(args[2]) + current: Final = self.counters.get(keys[0], 0) + if current + amount > limit: + return (0, current) + self.counters = MappingProxyType({**self.counters, keys[0]: current + amount}) + return (1, current + amount) + if kind == "refund": + remaining: Final = self.counters.get(keys[0], 0) - int(args[0]) + self.counters = MappingProxyType( + {key: value for key, value in self.counters.items() if key != keys[0]} + if remaining <= 0 + else {**self.counters, keys[0]: remaining} + ) + return 1 + raise AssertionError(f"unexpected {kind} script call for keys {keys}") + + +@pytest.mark.asyncio +async def test_redis_reserve_issues_single_key_calls_and_rolls_back_on_over_limit(): + fake = _SingleKeyRedisFake() + store = BatchEnqueuedTokenStore( + internal_usage_cache=InternalUsageCache(DualCache(redis_cache=fake, default_in_memory_ttl=60)) + ) + key_scope = _scope(limit=100, key="api_key") + team_scope = _scope(limit=50, key="team") + + over = await store.reserve(tokens=60, scopes=(key_scope, team_scope)) + assert over == BatchEnqueuedTokenOverLimit(scope=team_scope, enqueued=0) + assert tuple(kind for kind, _ in fake.script_calls) == ("reserve", "reserve", "refund") + assert not fake.counters + + fits = await store.reserve(tokens=50, scopes=(key_scope, team_scope)) + assert isinstance(fits, BatchEnqueuedTokenReservation) + await store.refund(fits) + assert not fake.counters + assert all(len(keys) == 1 for _, keys in fake.script_calls) + + def test_canonical_provider_batch_id_passes_raw_ids_through(): assert canonical_provider_batch_id("batch_abc123") == "batch_abc123" From 50896f21b346a08d99047a0f4e72f16e1e7dc54f Mon Sep 17 00:00:00 2001 From: mateo-berri Date: Wed, 19 Aug 2026 16:16:08 -0700 Subject: [PATCH 06/25] fix(batch_enqueued_tokens): roll back partial reserves, route refunds by backend, lowercase terminal statuses Reserve-script failures now roll back the scopes already incremented before re-raising into the in-memory fallback, so a partial redis outage no longer leaks counter increments that shrink the shared allowance. Reservations record which backend granted them, so a refund never debits redis counters an in-memory grant did not charge. Terminal-status matching is now case-insensitive because the Bedrock async-invoke retrieve path returns raw AWS-cased statuses like Completed. --- litellm/proxy/hooks/batch_enqueued_tokens.py | 59 +++++++++++++++---- .../hooks/parallel_request_limiter_v3.py | 2 +- .../proxy/hooks/test_batch_enqueued_tokens.py | 41 ++++++++++++- .../hooks/test_parallel_request_limiter_v3.py | 27 +++++++++ 4 files changed, 117 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/hooks/batch_enqueued_tokens.py b/litellm/proxy/hooks/batch_enqueued_tokens.py index e484b6a59fe..570291aa452 100644 --- a/litellm/proxy/hooks/batch_enqueued_tokens.py +++ b/litellm/proxy/hooks/batch_enqueued_tokens.py @@ -77,10 +77,14 @@ class BatchEnqueuedTokenScope: limit: int +ReservationBackend: TypeAlias = Literal["redis", "memory"] + + @dataclass(frozen=True, slots=True) class BatchEnqueuedTokenReservation: tokens: int scopes: tuple[BatchEnqueuedTokenScope, ...] + backend: ReservationBackend = "redis" @dataclass(frozen=True, slots=True) @@ -170,9 +174,10 @@ class BatchEnqueuedTokenStore: Counters and records live in Redis when Redis is configured, through single-key Lua scripts issued one scope at a time (Redis Cluster safe: no - cross-slot commands), with an over-limit scope rolling back the scopes - reserved before it; otherwise a single-process in-memory fallback guarded - by one asyncio lock is used. Everything expires after + cross-slot commands), with an over-limit or failing scope rolling back the + scopes reserved before it; otherwise a single-process in-memory fallback + guarded by one asyncio lock is used. Reservations remember which backend + granted them so a refund never debits counters the grant did not charge. Everything expires after ``BATCH_ENQUEUED_TOKEN_TTL_SECONDS`` so a crash between submission and the terminal-state refund can never leak tokens forever. """ @@ -229,15 +234,49 @@ class BatchEnqueuedTokenStore: scopes: tuple[BatchEnqueuedTokenScope, ...], ) -> BatchEnqueuedTokenOutcome: for index, scope in enumerate(scopes): - raw_result = await reserve_script( - (self._counter_key(scope),), - (tokens, BATCH_ENQUEUED_TOKEN_TTL_SECONDS, scope.limit), + result = await self._run_reserve_script( + reserve_script, + refund_script, + tokens=tokens, + scope=scope, + already_reserved=scopes[:index], ) - result = _RESERVE_RESULT_ADAPTER.validate_python(raw_result) if result[0] != 1: await self._refund_via_redis(refund_script, tokens=tokens, scopes=scopes[:index]) return BatchEnqueuedTokenOverLimit(scope=scope, enqueued=result[1]) - return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes) + return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes, backend="redis") + + async def _run_reserve_script( + self, + reserve_script: _ScriptRunner, + refund_script: _ScriptRunner, + tokens: int, + scope: BatchEnqueuedTokenScope, + already_reserved: tuple[BatchEnqueuedTokenScope, ...], + ) -> tuple[int, int]: + try: + raw_result: Final = await reserve_script( + (self._counter_key(scope),), + (tokens, BATCH_ENQUEUED_TOKEN_TTL_SECONDS, scope.limit), + ) + return _RESERVE_RESULT_ADAPTER.validate_python(raw_result) + except Exception: + await self._rollback_partial_reserve(refund_script, tokens=tokens, scopes=already_reserved) + raise + + async def _rollback_partial_reserve( + self, + refund_script: _ScriptRunner, + tokens: int, + scopes: tuple[BatchEnqueuedTokenScope, ...], + ) -> None: + try: + await self._refund_via_redis(refund_script, tokens=tokens, scopes=scopes) + except Exception as e: # noqa: BLE001 # best-effort rollback: the leak is TTL-bounded and only tightens the allowance + verbose_proxy_logger.warning( + "Rollback of partially reserved enqueued tokens failed; leaked increments expire with the TTL: %s", + str(e), + ) async def _refund_via_redis( self, @@ -261,7 +300,7 @@ class BatchEnqueuedTokenStore: return BatchEnqueuedTokenOverLimit(scope=scope, enqueued=current) for scope, current in zip(scopes, currents): await self._set_local_counter(scope, current + tokens, span) - return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes) + return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes, backend="memory") async def refund( self, @@ -271,7 +310,7 @@ class BatchEnqueuedTokenStore: if reservation.tokens <= 0 or not reservation.scopes: return refund_script: Final = self._refund_script - if refund_script is not None: + if reservation.backend == "redis" and refund_script is not None: try: await self._refund_via_redis(refund_script, tokens=reservation.tokens, scopes=reservation.scopes) except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory counters diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index b295da69255..1e65da5b867 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -4700,7 +4700,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): litellm_parent_otel_span=span, ) stash.batch_enqueued_reservation = None - if view.status in BATCH_ENQUEUED_REFUND_STATUSES: + if view.status.lower() in BATCH_ENQUEUED_REFUND_STATUSES: popped: Final = await self.batch_enqueued_token_store.pop_reservation( batch_id=canonical_provider_batch_id(view.id), litellm_parent_otel_span=span, diff --git a/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py index 1bb4798eaa0..bc924c32ba3 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py +++ b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py @@ -128,9 +128,10 @@ async def test_zero_token_reserve_charges_nothing(): class _SingleKeyRedisFake: """Emulates the Redis script path one single-key call at a time, recording every call.""" - def __init__(self) -> None: + def __init__(self, fail_reserve_keys: frozenset[str] = frozenset()) -> None: self.script_calls: tuple[tuple[str, tuple[str, ...]], ...] = () self.counters: Mapping[str, int] = MappingProxyType({}) + self.fail_reserve_keys = fail_reserve_keys def async_register_script(self, script: str): kind: Final = "reserve" if "INCRBY" in script else "refund" if "DECRBY" in script else "record" @@ -143,6 +144,8 @@ class _SingleKeyRedisFake: def _run(self, kind: str, keys: tuple[str, ...], args: tuple[str | bytes | int | float, ...]) -> object: if kind == "reserve": + if keys[0] in self.fail_reserve_keys: + raise ConnectionError(f"simulated redis failure for {keys[0]}") amount, limit = int(args[0]), int(args[2]) current: Final = self.counters.get(keys[0], 0) if current + amount > limit: @@ -181,6 +184,42 @@ async def test_redis_reserve_issues_single_key_calls_and_rolls_back_on_over_limi assert all(len(keys) == 1 for _, keys in fake.script_calls) +@pytest.mark.asyncio +async def test_partial_redis_reserve_failure_rolls_back_and_grants_in_memory(): + key_scope = _scope(limit=100, key="api_key") + team_scope = _scope(limit=50, key="team") + fake = _SingleKeyRedisFake(fail_reserve_keys=frozenset({f"batch_enqueued_tokens:team:{team_scope.value}"})) + store = BatchEnqueuedTokenStore( + internal_usage_cache=InternalUsageCache(DualCache(redis_cache=fake, default_in_memory_ttl=60)) + ) + + outcome = await store.reserve(tokens=10, scopes=(key_scope, team_scope)) + assert isinstance(outcome, BatchEnqueuedTokenReservation) + assert outcome.backend == "memory" + assert tuple(kind for kind, _ in fake.script_calls) == ("reserve", "reserve", "refund") + assert not fake.counters + + await store.refund(outcome) + assert tuple(kind for kind, _ in fake.script_calls) == ("reserve", "reserve", "refund") + + refilled = await store.reserve(tokens=50, scopes=(team_scope,)) + assert isinstance(refilled, BatchEnqueuedTokenReservation) + assert refilled.backend == "memory" + + +@pytest.mark.asyncio +async def test_pop_reservation_defaults_legacy_records_to_redis_backend(): + store = _in_memory_store() + legacy = '{"tokens": 5, "scopes": [{"key": "api_key", "value": "k", "limit": 10}]}' + store.internal_usage_cache.dual_cache.in_memory_cache.set_cache( + key="batch_enqueued_token_reservation:batch_legacy", value=legacy + ) + popped = await store.pop_reservation("batch_legacy") + assert popped == BatchEnqueuedTokenReservation( + tokens=5, scopes=(BatchEnqueuedTokenScope(key="api_key", value="k", limit=10),), backend="redis" + ) + + def test_canonical_provider_batch_id_passes_raw_ids_through(): assert canonical_provider_batch_id("batch_abc123") == "batch_abc123" diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 68c219fd72d..fee892ad342 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -5935,6 +5935,33 @@ async def test_success_hook_refunds_batch_enqueued_reservation_on_cancellation() assert isinstance(await store.reserve(tokens=100, scopes=(scope,)), BatchEnqueuedTokenReservation) +@pytest.mark.asyncio +async def test_success_hook_refunds_on_provider_cased_terminal_status(): + from litellm.proxy.hooks.batch_enqueued_tokens import ( + BatchEnqueuedTokenOverLimit, + BatchEnqueuedTokenReservation, + BatchEnqueuedTokenScope, + ) + + handler = _enqueued_test_handler() + store = handler.batch_enqueued_token_store + scope = BatchEnqueuedTokenScope(key="api_key", value="hashed-enqueued-key", limit=100) + user = UserAPIKeyAuth(api_key="hashed-enqueued-key") + + reservation = await store.reserve(tokens=90, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + get_or_create_request_stash().batch_enqueued_reservation = reservation + await handler.async_post_call_success_hook( + data={}, user_api_key_dict=user, response=_batch_response("batch_enq_cased", "InProgress") + ) + assert isinstance(await store.reserve(tokens=100, scopes=(scope,)), BatchEnqueuedTokenOverLimit) + + await handler.async_post_call_success_hook( + data={}, user_api_key_dict=user, response=_batch_response("batch_enq_cased", "Completed") + ) + assert isinstance(await store.reserve(tokens=100, scopes=(scope,)), BatchEnqueuedTokenReservation) + + @pytest.mark.asyncio async def test_failure_hook_refunds_stashed_batch_enqueued_reservation(): from litellm.proxy.hooks.batch_enqueued_tokens import ( From dce207add43a6ece53a5a0a12458dd7edf03ed3e Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 19 Aug 2026 16:19:26 -0700 Subject: [PATCH 07/25] fix(proxy): strip client standard_logging_object and zero-fill unknown recovered cost on the failure path Auth and pass-through failures reach post_call_failure_hook with the raw request body unstripped, so a client-supplied standard_logging_object could feed the new attribution fallback when the logging object carries none. Pop the key before the lift so only the logging object may supply it. Also coalesce a None recovered cost to 0.0 so the lift always overwrites any client-supplied response_cost, matching the merge base's clobber semantics. --- litellm/proxy/utils.py | 5 ++- tests/test_litellm/proxy/test_proxy_utils.py | 34 ++++++++++++++++++++ 2 files changed, 38 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 58ac7882cb8..acafc600b80 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -539,7 +539,7 @@ def _failure_fields_to_lift(request_data: Mapping[str, object]) -> Mapping[str, _entries: Final = ( ("first_api_call_start_time", _first_handoff), ("combined_usage_object", None if _usage_to_lift is None else _usage_to_lift[0]), - ("response_cost", None if _usage_to_lift is None else _usage_to_lift[1]), + ("response_cost", None if _usage_to_lift is None else (_usage_to_lift[1] or 0.0)), ("standard_logging_object", _model_call_details.get("standard_logging_object")), ) return MappingProxyType({key: value for key, value in _entries if value is not None}) @@ -2322,6 +2322,9 @@ class ProxyLogging: original_exception=original_exception, ) + # Auth and pass-through failures reach this hook with the raw request + # body unstripped, so only the logging object may supply this key. + request_data.pop("standard_logging_object", None) request_data.update(_failure_fields_to_lift(request_data)) # Remove before callbacks iterate — not serialisable diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index c80130b44da..71af44d6682 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -468,6 +468,23 @@ class TestPostCallFailureHookLiftsRecoveredPartialSpend: assert request_data["response_cost"] == 3.5e-05 assert "litellm_logging_obj" not in request_data + @pytest.mark.asyncio + async def test_recovered_usage_without_cost_clobbers_client_cost_with_zero(self): + from litellm.types.utils import Usage + + recovered_usage = Usage(prompt_tokens=30, completion_tokens=1, total_tokens=31) + logging_obj = MagicMock() + logging_obj.model_call_details = {"combined_usage_object": recovered_usage} + request_data = { + "litellm_logging_obj": logging_obj, + "response_cost": 999.0, + "metadata": {}, + } + await self._run(request_data) + + assert request_data["combined_usage_object"] is recovered_usage + assert request_data["response_cost"] == 0.0 + @pytest.mark.asyncio async def test_no_recovered_usage_is_noop(self): logging_obj = MagicMock() @@ -522,6 +539,23 @@ class TestPostCallFailureHookLiftsStandardLoggingObject: await self._run(request_data) assert request_data["standard_logging_object"] is authoritative + @pytest.mark.asyncio + async def test_client_supplied_key_is_stripped_when_logging_obj_supplies_none(self): + spoofed = {"model_id": "client-injected"} + request_data = {"standard_logging_object": spoofed, "metadata": {}} + await self._run(request_data) + assert "standard_logging_object" not in request_data + + logging_obj = MagicMock() + logging_obj.model_call_details = {} + request_data_with_obj = { + "litellm_logging_obj": logging_obj, + "standard_logging_object": spoofed, + "metadata": {}, + } + await self._run(request_data_with_obj) + assert "standard_logging_object" not in request_data_with_obj + @pytest.mark.asyncio async def test_no_standard_logging_object_is_noop(self): logging_obj = MagicMock() From 4333d528136571ae11a2f998ac0fafede111d520 Mon Sep 17 00:00:00 2001 From: mateo-berri Date: Wed, 19 Aug 2026 16:36:27 -0700 Subject: [PATCH 08/25] fix(batch_enqueued_tokens): scope in-memory refunds to the granting worker In-memory grants now record an owner token, and a refund only debits local counters when the popping worker is the one that granted them, so a terminal response handled elsewhere can no longer shrink another worker's unrelated fallback reservations. A Redis-granted refund that fails no longer falls back to decrementing local counters either: the leaked Redis increments expire with the TTL and only tighten the allowance. --- litellm/proxy/hooks/batch_enqueued_tokens.py | 40 +++++++++++----- .../proxy/hooks/test_batch_enqueued_tokens.py | 47 ++++++++++++++++++- 2 files changed, 74 insertions(+), 13 deletions(-) diff --git a/litellm/proxy/hooks/batch_enqueued_tokens.py b/litellm/proxy/hooks/batch_enqueued_tokens.py index 570291aa452..bbc007bb00d 100644 --- a/litellm/proxy/hooks/batch_enqueued_tokens.py +++ b/litellm/proxy/hooks/batch_enqueued_tokens.py @@ -9,6 +9,7 @@ the reservation is refunded when the batch reaches a terminal state """ import asyncio +import uuid from collections.abc import Awaitable, Mapping, Sequence from dataclasses import dataclass from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeAlias @@ -85,6 +86,7 @@ class BatchEnqueuedTokenReservation: tokens: int scopes: tuple[BatchEnqueuedTokenScope, ...] backend: ReservationBackend = "redis" + owner: str = "" @dataclass(frozen=True, slots=True) @@ -177,7 +179,8 @@ class BatchEnqueuedTokenStore: cross-slot commands), with an over-limit or failing scope rolling back the scopes reserved before it; otherwise a single-process in-memory fallback guarded by one asyncio lock is used. Reservations remember which backend - granted them so a refund never debits counters the grant did not charge. Everything expires after + granted them, and in-memory grants also remember the granting worker, so a + refund never debits counters the grant did not charge. Everything expires after ``BATCH_ENQUEUED_TOKEN_TTL_SECONDS`` so a crash between submission and the terminal-state refund can never leak tokens forever. """ @@ -185,6 +188,7 @@ class BatchEnqueuedTokenStore: def __init__(self, internal_usage_cache: "InternalUsageCache") -> None: self.internal_usage_cache = internal_usage_cache self._lock = asyncio.Lock() + self._owner_token = uuid.uuid4().hex redis_cache = internal_usage_cache.dual_cache.redis_cache self._reserve_script: _ScriptRunner | None = ( redis_cache.async_register_script(RESERVE_ENQUEUED_TOKENS_SCRIPT) if redis_cache is not None else None @@ -300,7 +304,7 @@ class BatchEnqueuedTokenStore: return BatchEnqueuedTokenOverLimit(scope=scope, enqueued=current) for scope, current in zip(scopes, currents): await self._set_local_counter(scope, current + tokens, span) - return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes, backend="memory") + return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes, backend="memory", owner=self._owner_token) async def refund( self, @@ -309,16 +313,14 @@ class BatchEnqueuedTokenStore: ) -> None: if reservation.tokens <= 0 or not reservation.scopes: return - refund_script: Final = self._refund_script - if reservation.backend == "redis" and refund_script is not None: - try: - await self._refund_via_redis(refund_script, tokens=reservation.tokens, scopes=reservation.scopes) - except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory counters - verbose_proxy_logger.warning( - "Redis enqueued-token refund failed, falling back to in-memory: %s", str(e) - ) - else: - return + if reservation.backend == "redis": + await self._refund_redis_reservation(reservation) + return + if reservation.owner != self._owner_token: + verbose_proxy_logger.warning( + "Skipping enqueued-token refund granted in another worker's memory; its counters expire with the TTL" + ) + return async with self._lock: for scope in reservation.scopes: current = await self._get_local_counter(scope, litellm_parent_otel_span) @@ -328,6 +330,20 @@ class BatchEnqueuedTokenStore: else: await self._set_local_counter(scope, remaining, litellm_parent_otel_span) + async def _refund_redis_reservation(self, reservation: BatchEnqueuedTokenReservation) -> None: + refund_script: Final = self._refund_script + if refund_script is None: + verbose_proxy_logger.warning( + "No Redis client for a Redis-granted enqueued-token refund; leaked increments expire with the TTL" + ) + return + try: + await self._refund_via_redis(refund_script, tokens=reservation.tokens, scopes=reservation.scopes) + except Exception as e: # noqa: BLE001 # best-effort refund: the leak is TTL-bounded and only tightens the allowance + verbose_proxy_logger.warning( + "Redis enqueued-token refund failed; leaked increments expire with the TTL: %s", str(e) + ) + async def save_reservation( self, batch_id: str, diff --git a/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py index bc924c32ba3..d6ae200dfaf 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py +++ b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py @@ -128,10 +128,15 @@ async def test_zero_token_reserve_charges_nothing(): class _SingleKeyRedisFake: """Emulates the Redis script path one single-key call at a time, recording every call.""" - def __init__(self, fail_reserve_keys: frozenset[str] = frozenset()) -> None: + def __init__( + self, + fail_reserve_keys: frozenset[str] = frozenset(), + fail_refund_keys: frozenset[str] = frozenset(), + ) -> None: self.script_calls: tuple[tuple[str, tuple[str, ...]], ...] = () self.counters: Mapping[str, int] = MappingProxyType({}) self.fail_reserve_keys = fail_reserve_keys + self.fail_refund_keys = fail_refund_keys def async_register_script(self, script: str): kind: Final = "reserve" if "INCRBY" in script else "refund" if "DECRBY" in script else "record" @@ -153,6 +158,8 @@ class _SingleKeyRedisFake: self.counters = MappingProxyType({**self.counters, keys[0]: current + amount}) return (1, current + amount) if kind == "refund": + if keys[0] in self.fail_refund_keys: + raise ConnectionError(f"simulated redis failure for {keys[0]}") remaining: Final = self.counters.get(keys[0], 0) - int(args[0]) self.counters = MappingProxyType( {key: value for key, value in self.counters.items() if key != keys[0]} @@ -207,6 +214,44 @@ async def test_partial_redis_reserve_failure_rolls_back_and_grants_in_memory(): assert refilled.backend == "memory" +@pytest.mark.asyncio +async def test_memory_refund_skips_reservations_granted_by_another_worker(): + store = _in_memory_store() + scope = _scope(limit=100) + reservation = await store.reserve(tokens=60, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + assert reservation.backend == "memory" + assert reservation.owner + + foreign: Final = BatchEnqueuedTokenReservation( + tokens=60, scopes=reservation.scopes, backend="memory", owner="another-worker" + ) + await store.refund(foreign) + assert await store.reserve(tokens=50, scopes=(scope,)) == BatchEnqueuedTokenOverLimit(scope=scope, enqueued=60) + + await store.refund(reservation) + assert isinstance(await store.reserve(tokens=100, scopes=(scope,)), BatchEnqueuedTokenReservation) + + +@pytest.mark.asyncio +async def test_failed_redis_refund_leaves_local_counters_untouched(): + scope = _scope(limit=100) + counter_key: Final = f"batch_enqueued_tokens:api_key:{scope.value}" + fake = _SingleKeyRedisFake(fail_refund_keys=frozenset({counter_key})) + store = BatchEnqueuedTokenStore( + internal_usage_cache=InternalUsageCache(DualCache(redis_cache=fake, default_in_memory_ttl=60)) + ) + + reservation = await store.reserve(tokens=60, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + assert reservation.backend == "redis" + store.internal_usage_cache.dual_cache.in_memory_cache.set_cache(key=counter_key, value=45) + + await store.refund(reservation) + assert store.internal_usage_cache.dual_cache.in_memory_cache.get_cache(key=counter_key) == 45 + assert fake.counters == {counter_key: 60} + + @pytest.mark.asyncio async def test_pop_reservation_defaults_legacy_records_to_redis_backend(): store = _in_memory_store() From 504112d5ca0f94dde42ec4a736bb18a6036c76d9 Mon Sep 17 00:00:00 2001 From: mateo-berri Date: Wed, 19 Aug 2026 16:52:46 -0700 Subject: [PATCH 09/25] fix(batch_enqueued_tokens): keep the over-limit verdict on rollback failure, find locally saved records on pop A Redis over-limit verdict now survives a failing rollback DECRBY instead of escaping into the in-memory fallback and granting tokens the counter already rejected; the unrolled increments expire with the TTL. pop_reservation now falls through to the local record when the Redis pop succeeds but finds nothing, so a reservation saved in memory after a transient Redis save failure still refunds on cancel or completion. --- litellm/proxy/hooks/batch_enqueued_tokens.py | 29 +++++----- .../proxy/hooks/test_batch_enqueued_tokens.py | 54 ++++++++++++++++++- 2 files changed, 70 insertions(+), 13 deletions(-) diff --git a/litellm/proxy/hooks/batch_enqueued_tokens.py b/litellm/proxy/hooks/batch_enqueued_tokens.py index bbc007bb00d..9d6a599af4a 100644 --- a/litellm/proxy/hooks/batch_enqueued_tokens.py +++ b/litellm/proxy/hooks/batch_enqueued_tokens.py @@ -246,7 +246,7 @@ class BatchEnqueuedTokenStore: already_reserved=scopes[:index], ) if result[0] != 1: - await self._refund_via_redis(refund_script, tokens=tokens, scopes=scopes[:index]) + await self._rollback_partial_reserve(refund_script, tokens=tokens, scopes=scopes[:index]) return BatchEnqueuedTokenOverLimit(scope=scope, enqueued=result[1]) return BatchEnqueuedTokenReservation(tokens=tokens, scopes=scopes, backend="redis") @@ -376,17 +376,10 @@ class BatchEnqueuedTokenStore: batch_id: str, litellm_parent_otel_span: "Span | None" = None, ) -> BatchEnqueuedTokenReservation | None: - raw: object = None - if self._pop_script is not None: - try: - raw = _POPPED_VALUE_ADAPTER.validate_python(await self._pop_script((self._record_key(batch_id),), ())) - except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory record - verbose_proxy_logger.warning( - "Redis enqueued-token reservation pop failed, falling back to in-memory: %s", str(e) - ) - raw = await self._pop_local_record(batch_id, litellm_parent_otel_span) - else: - raw = await self._pop_local_record(batch_id, litellm_parent_otel_span) + redis_raw: Final = await self._pop_redis_record(batch_id) + raw: Final = ( + redis_raw if redis_raw is not None else await self._pop_local_record(batch_id, litellm_parent_otel_span) + ) if raw is None: return None try: @@ -397,6 +390,18 @@ class BatchEnqueuedTokenStore: verbose_proxy_logger.warning("Discarding malformed enqueued-token reservation record for %s", batch_id) return None + async def _pop_redis_record(self, batch_id: str) -> str | bytes | None: + pop_script: Final = self._pop_script + if pop_script is None: + return None + try: + return _POPPED_VALUE_ADAPTER.validate_python(await pop_script((self._record_key(batch_id),), ())) + except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory record + verbose_proxy_logger.warning( + "Redis enqueued-token reservation pop failed, falling back to in-memory: %s", str(e) + ) + return None + async def _pop_local_record(self, batch_id: str, span: "Span | None") -> object: async with self._lock: stored = await self.internal_usage_cache.async_get_cache( diff --git a/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py index d6ae200dfaf..5167856359d 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py +++ b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py @@ -132,14 +132,21 @@ class _SingleKeyRedisFake: self, fail_reserve_keys: frozenset[str] = frozenset(), fail_refund_keys: frozenset[str] = frozenset(), + fail_save_keys: frozenset[str] = frozenset(), ) -> None: self.script_calls: tuple[tuple[str, tuple[str, ...]], ...] = () self.counters: Mapping[str, int] = MappingProxyType({}) + self.records: Mapping[str, str] = MappingProxyType({}) self.fail_reserve_keys = fail_reserve_keys self.fail_refund_keys = fail_refund_keys + self.fail_save_keys = fail_save_keys def async_register_script(self, script: str): - kind: Final = "reserve" if "INCRBY" in script else "refund" if "DECRBY" in script else "record" + kind: Final = ( + "reserve" + if "INCRBY" in script + else "refund" if "DECRBY" in script else "save" if "SET" in script else "pop" + ) async def run(keys: Sequence[str], args: Sequence[str | bytes | int | float]) -> object: self.script_calls = (*self.script_calls, (kind, tuple(keys))) @@ -167,6 +174,15 @@ class _SingleKeyRedisFake: else {**self.counters, keys[0]: remaining} ) return 1 + if kind == "save": + if keys[0] in self.fail_save_keys: + raise ConnectionError(f"simulated redis failure for {keys[0]}") + self.records = MappingProxyType({**self.records, keys[0]: str(args[0])}) + return 1 + if kind == "pop": + popped: Final = self.records.get(keys[0]) + self.records = MappingProxyType({key: value for key, value in self.records.items() if key != keys[0]}) + return popped raise AssertionError(f"unexpected {kind} script call for keys {keys}") @@ -214,6 +230,42 @@ async def test_partial_redis_reserve_failure_rolls_back_and_grants_in_memory(): assert refilled.backend == "memory" +@pytest.mark.asyncio +async def test_over_limit_verdict_survives_a_failing_rollback(): + key_scope = _scope(limit=100, key="api_key") + team_scope = _scope(limit=5, key="team") + key_counter: Final = f"batch_enqueued_tokens:api_key:{key_scope.value}" + fake = _SingleKeyRedisFake(fail_refund_keys=frozenset({key_counter})) + store = BatchEnqueuedTokenStore( + internal_usage_cache=InternalUsageCache(DualCache(redis_cache=fake, default_in_memory_ttl=60)) + ) + + outcome = await store.reserve(tokens=10, scopes=(key_scope, team_scope)) + assert outcome == BatchEnqueuedTokenOverLimit(scope=team_scope, enqueued=0) + assert fake.counters == {key_counter: 10} + + +@pytest.mark.asyncio +async def test_pop_falls_back_to_local_record_when_redis_pop_finds_nothing(): + scope = _scope(limit=100) + record_key: Final = "batch_enqueued_token_reservation:batch_local_record" + fake = _SingleKeyRedisFake(fail_save_keys=frozenset({record_key})) + store = BatchEnqueuedTokenStore( + internal_usage_cache=InternalUsageCache(DualCache(redis_cache=fake, default_in_memory_ttl=60)) + ) + + reservation = await store.reserve(tokens=60, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + await store.save_reservation("batch_local_record", reservation) + assert not fake.records + + popped = await store.pop_reservation("batch_local_record") + assert popped == reservation + await store.refund(popped) + assert not fake.counters + assert await store.pop_reservation("batch_local_record") is None + + @pytest.mark.asyncio async def test_memory_refund_skips_reservations_granted_by_another_worker(): store = _in_memory_store() From 919bf1a09785c020dd859605cb1074bfa6c42383 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 19 Aug 2026 16:53:02 -0700 Subject: [PATCH 10/25] fix(proxy): strip client standard_logging_object before the failure logging handler --- litellm/proxy/utils.py | 8 +++-- tests/test_litellm/proxy/test_proxy_utils.py | 35 ++++++++++++++++++++ 2 files changed, 40 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index acafc600b80..19dcd64c58d 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -2309,6 +2309,11 @@ class ProxyLogging: ) ) + # Auth and pass-through failure bodies are unstripped client input, and + # the logging handler below flattens body keys into model_call_details, + # so drop the key before it can masquerade as the built payload. + request_data.pop("standard_logging_object", None) + ### LOGGING ### if self._is_proxy_only_llm_api_error( original_exception=original_exception, @@ -2322,9 +2327,6 @@ class ProxyLogging: original_exception=original_exception, ) - # Auth and pass-through failures reach this hook with the raw request - # body unstripped, so only the logging object may supply this key. - request_data.pop("standard_logging_object", None) request_data.update(_failure_fields_to_lift(request_data)) # Remove before callbacks iterate — not serialisable diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 71af44d6682..07877514b69 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -556,6 +556,41 @@ class TestPostCallFailureHookLiftsStandardLoggingObject: await self._run(request_data_with_obj) assert "standard_logging_object" not in request_data_with_obj + @pytest.mark.asyncio + async def test_pass_through_failure_never_relifts_client_supplied_key(self): + from datetime import datetime + from unittest.mock import AsyncMock, patch + + from fastapi import HTTPException + + from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.proxy._types import UserAPIKeyAuth + + logging_obj = Logging( + model="claude-haiku-4-5", + messages=[{"role": "user", "content": "hi"}], + stream=False, + call_type="pass_through_endpoint", + start_time=datetime.now(), + litellm_call_id="test-call-id", + function_id="test-function-id", + ) + request_data = { + "litellm_logging_obj": logging_obj, + "standard_logging_object": {"model_id": "client-injected"}, + "metadata": {}, + } + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + proxy_logging_obj.alert_types = [] + with patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()): + await proxy_logging_obj.post_call_failure_hook( + request_data=request_data, + original_exception=HTTPException(status_code=401, detail="unauthorized"), + user_api_key_dict=UserAPIKeyAuth(request_route="/v1/chat/completions"), + ) + assert "standard_logging_object" not in request_data + assert "standard_logging_object" not in logging_obj.model_call_details + @pytest.mark.asyncio async def test_no_standard_logging_object_is_noop(self): logging_obj = MagicMock() From 6b17b8a8f2489085f8923b5a84cdb79412108c24 Mon Sep 17 00:00:00 2001 From: hiraku-miyoshi Date: Wed, 19 Aug 2026 17:22:15 -0700 Subject: [PATCH 11/25] fix(proxy): restrict batch_enqueued_token_limit metadata writes to proxy admins The field replaces the standard RPM/TPM checks for batch submissions, so a key holder or team admin writing it could pick their own batch quota. Mirrors the output-token-estimate admin gate: change-based, so resending the stored value stays allowed, and enforced on key generate, update, bulk team-key update, regenerate, and team new/update. --- litellm/constants.py | 5 + litellm/proxy/auth/auth_utils.py | 47 +++- litellm/proxy/hooks/batch_enqueued_tokens.py | 6 +- .../key_management_endpoints.py | 27 +++ .../management_endpoints/team_endpoints.py | 17 +- .../test_key_management_endpoints.py | 201 ++++++++++++++++++ .../test_team_endpoints.py | 57 +++++ 7 files changed, 354 insertions(+), 6 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index d7ddacf5fac..facfc6f7c19 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1772,6 +1772,11 @@ PTU_PRUNE_SKEW_GRACE_SECONDS: Final[int] = 300 # was never observed (e.g. proxy restart); expiry returns the tokens to the caller. BATCH_ENQUEUED_TOKEN_TTL_SECONDS: Final[int] = 8 * 24 * 60 * 60 +# Key/team metadata field that opts batches into enqueued-token limiting. Only proxy +# admins may write it: when present it replaces the standard RPM/TPM checks for +# batch submissions. +BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY: Final = "batch_enqueued_token_limit" + # Shared read-only empty mapping, for defaulting optional Mapping parameters without # constructing a fresh mutable dict at each call site. EMPTY_MAPPING: Final = MappingProxyType({}) diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index a105bf19458..883f986f6fd 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -12,7 +12,12 @@ from pydantic import PositiveInt, TypeAdapter, ValidationError import litellm from litellm import Router, provider_list from litellm._logging import verbose_proxy_logger -from litellm.constants import MINIMUM_CUSTOM_KEY_LENGTH, STANDARD_CUSTOMER_ID_HEADERS +from litellm.constants import ( + BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY, + EMPTY_MAPPING, + MINIMUM_CUSTOM_KEY_LENGTH, + STANDARD_CUSTOMER_ID_HEADERS, +) from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.litellm_core_utils.url_utils import ( SSRFError, @@ -1169,6 +1174,46 @@ def enforce_output_token_estimates_are_admin_only( ) +class BatchEnqueuedTokenLimitRequest(Protocol): + """The shape of any management request that can carry a batch enqueued-token limit.""" + + @property + def metadata(self) -> Mapping[str, object] | None: ... + + @property + def model_fields_set(self) -> Collection[str]: ... + + +def enforce_batch_enqueued_token_limit_is_admin_only( + data: BatchEnqueuedTokenLimitRequest, + existing_metadata: Mapping[str, object] | None, + user_api_key_dict: UserAPIKeyAuth, + entity: Literal["key", "team"], +) -> None: + """Only a proxy admin may change a key or team's batch enqueued-token limit. + + When set, ``batch_enqueued_token_limit`` replaces the standard RPM/TPM checks + for batch submissions, so a holder-writable copy would let a caller lift their + own batch quota. Gated on the resulting value rather than on presence, so a + form resending the stored value stays a no-op. + """ + if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: + return + stored: Final[Mapping[str, object]] = existing_metadata or EMPTY_MAPPING + requested: Final[Mapping[str, object]] = ( + (data.metadata or EMPTY_MAPPING) if "metadata" in data.model_fields_set else stored + ) + if requested.get(BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY) == stored.get(BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY): + return + raise HTTPException( + status_code=403, + detail={ # mutable-ok: HTTPException.detail has no immutable form + "error": f"Only proxy admins can set {BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY} on a {entity}. " + "It replaces the standard rate limit checks for batch submissions." + }, + ) + + def get_model_rate_limit_from_metadata( user_api_key_dict: UserAPIKeyAuth, metadata_accessor_key: Literal["team_metadata", "organization_metadata", "project_metadata"], diff --git a/litellm/proxy/hooks/batch_enqueued_tokens.py b/litellm/proxy/hooks/batch_enqueued_tokens.py index 9d6a599af4a..54b0111682d 100644 --- a/litellm/proxy/hooks/batch_enqueued_tokens.py +++ b/litellm/proxy/hooks/batch_enqueued_tokens.py @@ -1,7 +1,7 @@ """ Enqueued-token accounting for batch submissions. -Opt-in via ``batch_enqueued_token_limit`` in key or team metadata: batch +Opt-in via admin-set ``batch_enqueued_token_limit`` in key or team metadata: batch submissions reserve their estimated token count against a long-lived enqueued-token allowance instead of the per-minute rate-limit windows, and the reservation is refunded when the batch reaches a terminal state @@ -17,7 +17,7 @@ from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, TypeAlias from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError from litellm._logging import verbose_proxy_logger -from litellm.constants import BATCH_ENQUEUED_TOKEN_TTL_SECONDS +from litellm.constants import BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY, BATCH_ENQUEUED_TOKEN_TTL_SECONDS from litellm.proxy._types import UserAPIKeyAuth if TYPE_CHECKING: @@ -28,8 +28,6 @@ if TYPE_CHECKING: Span = _Span InternalUsageCache = _InternalUsageCache -BATCH_ENQUEUED_TOKEN_LIMIT_METADATA_KEY: Final = "batch_enqueued_token_limit" - BATCH_ENQUEUED_REFUND_STATUSES: Final[frozenset[str]] = frozenset( {"completed", "complete", "failed", "expired", "cancelled", "cancelling"} ) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 67a836b8c92..ade568ad07f 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -57,6 +57,7 @@ from litellm.proxy.auth.auth_checks import ( ) from litellm.proxy.auth.auth_utils import ( abbreviate_api_key, + enforce_batch_enqueued_token_limit_is_admin_only, enforce_output_token_estimates_are_admin_only, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @@ -901,6 +902,12 @@ async def _common_key_generation_helper( user_api_key_dict=user_api_key_dict, entity="key", ) + enforce_batch_enqueued_token_limit_is_admin_only( + data=data, + existing_metadata=None, + user_api_key_dict=user_api_key_dict, + entity="key", + ) if data.metadata is not None and data.metadata.get("service_account_id") is not None and data.team_id is None: await validate_team_id_used_in_service_account_request( @@ -2296,6 +2303,14 @@ async def _process_single_key_update( prisma_client=prisma_client, ) + _existing_row_metadata: Final = getattr(existing_key_row, "metadata", None) + enforce_batch_enqueued_token_limit_is_admin_only( + data=update_key_request, + existing_metadata=_existing_row_metadata if isinstance(_existing_row_metadata, dict) else None, + user_api_key_dict=user_api_key_dict, + entity="key", + ) + # Check team member permissions if prisma_client is not None: await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint( @@ -2558,6 +2573,12 @@ async def _validate_update_key_data( user_api_key_dict=user_api_key_dict, entity="key", ) + enforce_batch_enqueued_token_limit_is_admin_only( + data=data, + existing_metadata=_existing_metadata if isinstance(_existing_metadata, dict) else None, + user_api_key_dict=user_api_key_dict, + entity="key", + ) # Personal-key bypass: the caller both created the key AND still owns it # (user_id == caller). Checking only created_by would let a demoted admin @@ -4754,6 +4775,12 @@ async def _execute_virtual_key_regeneration( user_api_key_dict=user_api_key_dict, entity="key", ) + enforce_batch_enqueued_token_limit_is_admin_only( + data=data, + existing_metadata=_existing_key_metadata if isinstance(_existing_key_metadata, dict) else None, + user_api_key_dict=user_api_key_dict, + entity="key", + ) new_token: Final = await get_new_token(data=data) new_token_hash: Final = hash_token(new_token) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 95632d7cb35..31006958ba1 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -85,7 +85,10 @@ from litellm.proxy.auth.auth_checks import ( get_team_object, get_user_object, ) -from litellm.proxy.auth.auth_utils import enforce_output_token_estimates_are_admin_only +from litellm.proxy.auth.auth_utils import ( + enforce_batch_enqueued_token_limit_is_admin_only, + enforce_output_token_estimates_are_admin_only, +) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.callback_utils import encrypt_callback_vars from litellm.proxy.common_utils.json_merge_patch import apply_json_merge_patch @@ -1303,6 +1306,12 @@ async def new_team( user_api_key_dict=user_api_key_dict, entity="team", ) + enforce_batch_enqueued_token_limit_is_admin_only( + data=data, + existing_metadata=None, + user_api_key_dict=user_api_key_dict, + entity="team", + ) # Check if license is over limit total_teams: Final = await _team_db(prisma_client).count() @@ -2007,6 +2016,12 @@ async def update_team( user_api_key_dict=user_api_key_dict, entity="team", ) + enforce_batch_enqueued_token_limit_is_admin_only( + data=data, + existing_metadata=_existing_team_metadata if isinstance(_existing_team_metadata, dict) else None, + user_api_key_dict=user_api_key_dict, + entity="team", + ) _check_passthrough_routes_caller_permission(data, user_api_key_dict, entity="team") diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index bdf09a95e4b..93a4788caf8 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -15515,6 +15515,207 @@ async def test_regenerate_key_output_token_estimate_lowered_rejected_for_non_adm assert "Only proxy admins can set" in str(exc.value.detail) +_BATCH_LIMIT = "batch_enqueued_token_limit" + + +@pytest.mark.parametrize( + "label, request_body, existing_metadata, allowed", + [ + ("set on a key with none stored", {"metadata": {_BATCH_LIMIT: 50000}}, None, False), + ("raised above the stored limit", {"metadata": {_BATCH_LIMIT: 200000}}, {_BATCH_LIMIT: 100000}, False), + ("cleared by replacing the blob", {"metadata": {}}, {_BATCH_LIMIT: 100000}, False), + ("resent unchanged", {"metadata": {_BATCH_LIMIT: 100000}}, {_BATCH_LIMIT: 100000}, True), + ("left untouched", {}, {_BATCH_LIMIT: 100000}, True), + ], +) +def test_batch_enqueued_token_limit_admin_gate_matrix(label, request_body, existing_metadata, allowed): + """A non-admin may only leave a key's stored batch enqueued-token limit as it is. + + When set, the limit replaces the standard RPM/TPM checks for batch + submissions, so a key holder writing it would pick their own batch quota. + Resending the stored value is what the edit form produces on every save + and has to stay allowed. + """ + from litellm.proxy.auth.auth_utils import ( + enforce_batch_enqueued_token_limit_is_admin_only, + ) + + def _call(caller): + enforce_batch_enqueued_token_limit_is_admin_only( + data=UpdateKeyRequest(key="sk-1", **request_body), + existing_metadata=existing_metadata, + user_api_key_dict=caller, + entity="key", + ) + + non_admin = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-non-admin", + user_id="alice", + ) + if allowed: + _call(non_admin) + else: + with pytest.raises(HTTPException) as exc: + _call(non_admin) + assert exc.value.status_code == 403 + assert "Only proxy admins can set" in str(exc.value.detail) + + _call( + UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin", + ) + ) + + +@pytest.mark.asyncio +async def test_generate_key_batch_enqueued_token_limit_rejected_for_non_admin(): + """A non-admin self-minting a key with the limit would replace the standard + batch RPM/TPM checks with a cap of their own choosing.""" + with patch("litellm.proxy.proxy_server.prisma_client", AsyncMock()): + with pytest.raises(HTTPException) as exc: + await _common_key_generation_helper( + data=GenerateKeyRequest(metadata={_BATCH_LIMIT: 100000}, rpm_limit=2), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-alice", + user_id="alice", + ), + litellm_changed_by=None, + team_table=None, + ) + assert int(getattr(exc.value, "status_code", 0)) == 403 + assert "Only proxy admins can set" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_update_key_batch_enqueued_token_limit_raised_rejected_for_non_admin(monkeypatch): + """/key/update is reachable by the key's own holder, so the gate has to + fire inside the update path itself rather than only at generation.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + token = "d1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd" + _wire_update_key_fn(monkeypatch, _estimate_key_row(token, {_BATCH_LIMIT: 100000})) + + mock_request = MagicMock() + mock_request.query_params = {} + + with pytest.raises(ProxyException) as exc: + await update_key_fn( + request=mock_request, + data=UpdateKeyRequest(key=token, metadata={_BATCH_LIMIT: 10**12}), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-internal", + user_id="internal_user", + ), + litellm_changed_by=None, + ) + + assert str(exc.value.code) == "403" + assert "Only proxy admins can set" in str(exc.value.message) + + +@pytest.mark.asyncio +async def test_update_key_batch_enqueued_token_limit_unchanged_allows_non_admin_edit(monkeypatch): + """The edit form resends every field it renders, so gating on presence + would 403 a key owner renaming a key that carries an admin-set limit.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + token = "e1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd" + _wire_update_key_fn(monkeypatch, _estimate_key_row(token, {_BATCH_LIMIT: 100000})) + + mock_request = MagicMock() + mock_request.query_params = {} + + result = await update_key_fn( + request=mock_request, + data=UpdateKeyRequest(key=token, key_alias="my-alias", metadata={_BATCH_LIMIT: 100000}), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-internal", + user_id="internal_user", + ), + litellm_changed_by=None, + ) + + assert result is not None + + +@pytest.mark.asyncio +async def test_regenerate_key_batch_enqueued_token_limit_rejected_for_non_admin(): + """/key/regenerate runs the request body through prepare_key_update_data + exactly as an update does, so it is a third write path into the field.""" + from litellm.proxy._types import RegenerateKeyRequest + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _execute_virtual_key_regeneration, + ) + + token = "f1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd" + key_in_db = LiteLLM_VerificationToken( + token=token, + user_id="internal_user", + metadata={_BATCH_LIMIT: 100000}, + ) + + with pytest.raises(HTTPException) as exc: + await _execute_virtual_key_regeneration( + prisma_client=AsyncMock(), + key_in_db=key_in_db, + hashed_api_key=token, + key="sk-original", + data=RegenerateKeyRequest(key="sk-original", metadata={_BATCH_LIMIT: 10**12}), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-internal", + user_id="internal_user", + ), + litellm_changed_by=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + ) + + assert exc.value.status_code == 403 + assert "Only proxy admins can set" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_bulk_key_update_batch_enqueued_token_limit_rejected_for_non_admin(): + """Bulk team-key updates run through _process_single_key_update, not + /key/update's validator, so the gate must also live on that path.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _process_single_key_update, + ) + + token = "a2b2c3d4e5f6789012345678901234567890123456789012345678901234abcd" + existing = _estimate_key_row(token, {_BATCH_LIMIT: 100000}) + + with pytest.raises(HTTPException) as exc: + await _process_single_key_update( + update_key_request=UpdateKeyRequest(key=token, metadata={_BATCH_LIMIT: 10**12}), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-internal", + user_id="internal_user", + ), + litellm_changed_by=None, + prisma_client=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + llm_router=None, + existing_key_row=existing, + ) + + assert exc.value.status_code == 403 + assert "Only proxy admins can set" in str(exc.value.detail) + + @pytest.mark.asyncio async def test_execute_virtual_key_regeneration_stamps_settings_updated_at(): """Regenerate rewrites the key's config, so it must move settings_updated_at.""" diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index db39fcd3799..05072a0a6d7 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -11649,6 +11649,63 @@ async def test_new_team_output_token_estimate_rejected_for_non_admin(): assert "on a team" in str(exc.value.message) +_TEAM_BATCH_LIMIT = "batch_enqueued_token_limit" + + +@pytest.mark.asyncio +async def test_update_team_batch_enqueued_token_limit_raised_rejected_for_team_admin(): + """_verify_team_access admits a team admin, so the gate has to fire inside + update_team itself to keep the team's batch quota admin-owned.""" + import contextlib + from unittest.mock import Mock + + from fastapi import Request + + from litellm.proxy._types import UpdateTeamRequest + from litellm.proxy.management_endpoints.team_endpoints import update_team + + with contextlib.ExitStack() as stack: + _wire_update_team(stack, {_TEAM_BATCH_LIMIT: 100000}) + with pytest.raises(ProxyException) as exc: + await update_team( + data=UpdateTeamRequest(team_id="test_team_id", metadata={_TEAM_BATCH_LIMIT: 10**12}), + http_request=Mock(spec=Request), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-team-admin", + user_id="team-admin", + ), + ) + + assert str(exc.value.code) == "403" + assert "on a team" in str(exc.value.message) + + +@pytest.mark.asyncio +async def test_new_team_batch_enqueued_token_limit_rejected_for_non_admin(): + """/team/new is the other write path into the same stored metadata.""" + from unittest.mock import Mock + + from fastapi import Request + + from litellm.proxy._types import NewTeamRequest + from litellm.proxy.management_endpoints.team_endpoints import new_team + + with pytest.raises(ProxyException) as exc: + await new_team( + data=NewTeamRequest(team_alias="t", metadata={_TEAM_BATCH_LIMIT: 100000}), + http_request=Mock(spec=Request), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-alice", + user_id="alice", + ), + ) + + assert str(exc.value.code) == "403" + assert "on a team" in str(exc.value.message) + + @pytest.mark.asyncio async def test_get_team_daily_activity_aggregated_scopes_and_flags(mock_db_client): """The aggregated endpoint must apply the same non-admin key scoping as the From 5ab20c3678756bb5a04ff2131e1de11a2619583c Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 19 Aug 2026 17:47:06 -0700 Subject: [PATCH 12/25] fix(batch_enqueued_tokens): tombstone popped Redis reservation records so local ghosts cannot double-refund --- litellm/proxy/hooks/batch_enqueued_tokens.py | 14 +++++-- .../proxy/hooks/test_batch_enqueued_tokens.py | 41 ++++++++++++++++++- 2 files changed, 50 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/hooks/batch_enqueued_tokens.py b/litellm/proxy/hooks/batch_enqueued_tokens.py index 54b0111682d..82d4e7c66fc 100644 --- a/litellm/proxy/hooks/batch_enqueued_tokens.py +++ b/litellm/proxy/hooks/batch_enqueued_tokens.py @@ -62,8 +62,8 @@ return 1 POP_RESERVATION_SCRIPT: Final = """ local value = redis.call('GET', KEYS[1]) -if value then - redis.call('DEL', KEYS[1]) +if value and value ~= '' then + redis.call('SET', KEYS[1], '', 'EX', tonumber(ARGV[1])) end return value """ @@ -375,6 +375,12 @@ class BatchEnqueuedTokenStore: litellm_parent_otel_span: "Span | None" = None, ) -> BatchEnqueuedTokenReservation | None: redis_raw: Final = await self._pop_redis_record(batch_id) + if redis_raw is not None and not redis_raw: + # The Redis pop tombstones popped records in place, so a hit on the empty + # tombstone means the batch was already refunded elsewhere; a local copy + # left behind by a save that raised after landing must not refund again. + await self._pop_local_record(batch_id, litellm_parent_otel_span) + return None raw: Final = ( redis_raw if redis_raw is not None else await self._pop_local_record(batch_id, litellm_parent_otel_span) ) @@ -393,7 +399,9 @@ class BatchEnqueuedTokenStore: if pop_script is None: return None try: - return _POPPED_VALUE_ADAPTER.validate_python(await pop_script((self._record_key(batch_id),), ())) + return _POPPED_VALUE_ADAPTER.validate_python( + await pop_script((self._record_key(batch_id),), (BATCH_ENQUEUED_TOKEN_TTL_SECONDS,)) + ) except Exception as e: # noqa: BLE001 # any Redis failure must fall back to the in-memory record verbose_proxy_logger.warning( "Redis enqueued-token reservation pop failed, falling back to in-memory: %s", str(e) diff --git a/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py index 5167856359d..440b3050a39 100644 --- a/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py +++ b/tests/test_litellm/proxy/hooks/test_batch_enqueued_tokens.py @@ -133,6 +133,7 @@ class _SingleKeyRedisFake: fail_reserve_keys: frozenset[str] = frozenset(), fail_refund_keys: frozenset[str] = frozenset(), fail_save_keys: frozenset[str] = frozenset(), + raise_after_landing_save_keys: frozenset[str] = frozenset(), ) -> None: self.script_calls: tuple[tuple[str, tuple[str, ...]], ...] = () self.counters: Mapping[str, int] = MappingProxyType({}) @@ -140,12 +141,13 @@ class _SingleKeyRedisFake: self.fail_reserve_keys = fail_reserve_keys self.fail_refund_keys = fail_refund_keys self.fail_save_keys = fail_save_keys + self.raise_after_landing_save_keys = raise_after_landing_save_keys def async_register_script(self, script: str): kind: Final = ( "reserve" if "INCRBY" in script - else "refund" if "DECRBY" in script else "save" if "SET" in script else "pop" + else "refund" if "DECRBY" in script else "pop" if "GET" in script else "save" ) async def run(keys: Sequence[str], args: Sequence[str | bytes | int | float]) -> object: @@ -178,10 +180,13 @@ class _SingleKeyRedisFake: if keys[0] in self.fail_save_keys: raise ConnectionError(f"simulated redis failure for {keys[0]}") self.records = MappingProxyType({**self.records, keys[0]: str(args[0])}) + if keys[0] in self.raise_after_landing_save_keys: + raise TimeoutError(f"simulated redis timeout after landing for {keys[0]}") return 1 if kind == "pop": popped: Final = self.records.get(keys[0]) - self.records = MappingProxyType({key: value for key, value in self.records.items() if key != keys[0]}) + if popped: + self.records = MappingProxyType({**self.records, keys[0]: ""}) return popped raise AssertionError(f"unexpected {kind} script call for keys {keys}") @@ -266,6 +271,38 @@ async def test_pop_falls_back_to_local_record_when_redis_pop_finds_nothing(): assert await store.pop_reservation("batch_local_record") is None +@pytest.mark.asyncio +async def test_local_ghost_left_by_landed_save_never_refunds_twice(): + scope = _scope(limit=100) + record_key: Final = "batch_enqueued_token_reservation:batch_ghost" + fake = _SingleKeyRedisFake(raise_after_landing_save_keys=frozenset({record_key})) + store = BatchEnqueuedTokenStore( + internal_usage_cache=InternalUsageCache(DualCache(redis_cache=fake, default_in_memory_ttl=60)) + ) + + reservation = await store.reserve(tokens=60, scopes=(scope,)) + assert isinstance(reservation, BatchEnqueuedTokenReservation) + await store.save_reservation("batch_ghost", reservation) + assert fake.records[record_key] + + first = await store.pop_reservation("batch_ghost") + assert first == reservation + await store.refund(first) + assert not fake.counters + assert fake.records[record_key] == "" + + assert await store.pop_reservation("batch_ghost") is None + assert ( + await store.internal_usage_cache.async_get_cache( + key=record_key, litellm_parent_otel_span=None, local_only=True + ) + is None + ) + assert await store.pop_reservation("batch_ghost") is None + assert isinstance(await store.reserve(tokens=100, scopes=(scope,)), BatchEnqueuedTokenReservation) + assert fake.counters[f"batch_enqueued_tokens:{scope.key}:{scope.value}"] == 100 + + @pytest.mark.asyncio async def test_memory_refund_skips_reservations_granted_by_another_worker(): store = _in_memory_store() From 0ab172575703387896d998c108bd32e6ca4bf57f Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Wed, 19 Aug 2026 18:19:09 -0700 Subject: [PATCH 13/25] refactor(ui): migrate the model and router settings pages off antd (#37523) * refactor(ui): migrate the model and router settings pages off antd Converts the add model flow, credential panels, model settings and router settings onto the shadcn primitives, moves the mapping table onto the shared DataTable, and drops the dead uploadProps prop chain that only existed to carry antd's UploadProps type. * fix(ui): split comma-separated custom technical keywords into one term each --- ui/litellm-dashboard/eslint-suppressions.json | 131 ------- .../panels/AddModelPanel.integration.test.tsx | 14 +- .../panels/AddModelPanel.tsx | 4 - .../panels/LlmCredentialsPanel.tsx | 12 +- .../vertexCredentialsUpload.test.ts | 47 --- .../vertexCredentialsUpload.ts | 37 -- .../src/components/ModelInfoEditForm.tsx | 27 +- .../Fallbacks/EditFallbacks.test.tsx | 17 +- .../Fallbacks/FallbackGroupConfig.tsx | 80 ++--- .../Fallbacks/FallbackSelectionForm.test.tsx | 18 +- .../Fallbacks/FallbackSelectionForm.tsx | 82 +++-- .../add_model/AdaptiveRoutingConfig.tsx | 110 +++--- .../add_model/AddModelForm.test.tsx | 40 ++- .../src/components/add_model/AddModelForm.tsx | 5 +- .../add_model/ClassificationMethodConfig.tsx | 294 ++++++++-------- .../add_model/ComplexityRouterConfig.test.tsx | 60 ++-- .../add_model/ComplexityRouterConfig.tsx | 319 +++++++++--------- .../add_model/EscalationKeywords.tsx | 26 +- .../components/add_model/KeywordTierRules.tsx | 138 ++++---- .../add_model/SemanticKeywordMatching.tsx | 50 +-- .../add_model/add_auto_router_tab.test.tsx | 95 +++--- .../add_model/add_auto_router_tab.tsx | 96 +++--- .../add_model/advanced_settings.tsx | 146 ++++---- .../conditional_public_model_name.tsx | 31 +- .../add_model/litellm_model_name.tsx | 36 +- .../provider_specific_fields.test.tsx | 65 +++- .../add_model/provider_specific_fields.tsx | 150 ++++---- .../edit_auto_router_modal.test.tsx | 22 +- .../model_add/CredentialModal.test.tsx | 14 +- .../components/model_add/CredentialModal.tsx | 5 +- .../model_add/CredentialsPanel.test.tsx | 5 +- .../components/model_add/CredentialsPanel.tsx | 10 +- .../ModelSettingsModal.test.tsx | 4 +- .../ModelSettingsModal/ModelSettingsModal.tsx | 30 +- .../src/components/model_info_view.test.tsx | 23 +- .../RouterSettingsForm.test.tsx | 26 +- .../RoutingStrategySelector.test.tsx | 53 ++- .../RoutingStrategySelector.tsx | 34 +- .../components/router_settings/index.test.tsx | 52 +-- .../components/shared/SearchSelect.test.tsx | 7 + .../shared/form/UtcDateTimeInput.tsx | 53 +++ .../src/components/ui/combobox.tsx | 3 +- 42 files changed, 1153 insertions(+), 1318 deletions(-) delete mode 100644 ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/vertexCredentialsUpload.test.ts delete mode 100644 ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/vertexCredentialsUpload.ts create mode 100644 ui/litellm-dashboard/src/components/shared/form/UtcDateTimeInput.tsx diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index cafdb2fa7dc..2d68ea2aa68 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -57,9 +57,6 @@ "local/no-complex-jsx-arrow": { "count": 1 }, - "no-restricted-imports": { - "count": 1 - }, "react-hooks/immutability": { "count": 1 } @@ -199,11 +196,6 @@ "count": 1 } }, - "src/app/(dashboard)/guardrails-monitor/_components/GuardrailConfig.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.tsx": { "no-nested-ternary": { "count": 3 @@ -550,11 +542,6 @@ "count": 1 } }, - "src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.test.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx": { "no-nested-ternary": { "count": 1 @@ -570,11 +557,6 @@ "count": 1 } }, - "src/app/(dashboard)/mcp-servers/_components/PassthroughAuthorizeSection.test.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/app/(dashboard)/mcp-servers/_components/ToolTestPanel.tsx": { "react-hooks/set-state-in-effect": { "count": 1 @@ -782,11 +764,6 @@ "count": 1 } }, - "src/app/(dashboard)/playground/components/compareUI/components/ModelSelector.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx": { "local/no-complex-jsx-arrow": { "count": 2 @@ -987,9 +964,6 @@ "local/filename-pascal-case": { "count": 1 }, - "no-restricted-imports": { - "count": 1 - }, "react-hooks/immutability": { "count": 1 } @@ -1257,11 +1231,6 @@ "count": 1 } }, - "src/app/(dashboard)/vector-stores/_components/CreateVectorStore.tsx": { - "no-restricted-imports": { - "count": 2 - } - }, "src/app/(dashboard)/vector-stores/_components/VectorStoreForm.tsx": { "no-nested-ternary": { "count": 2 @@ -1400,9 +1369,6 @@ "no-nested-ternary": { "count": 1 }, - "no-restricted-imports": { - "count": 1 - }, "react-hooks/set-state-in-effect": { "count": 1 } @@ -1437,18 +1403,7 @@ "count": 1 } }, - "src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx": { - "local/no-complex-jsx-arrow": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, "src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.tsx": { - "no-restricted-imports": { - "count": 1 - }, "react-hooks/set-state-in-effect": { "count": 1 } @@ -1488,11 +1443,6 @@ "count": 1 } }, - "src/components/UsagePage/components/EntityUsage/TopKeyView.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/components/UsagePage/utils/value_formatters.tsx": { "local/filename-pascal-case": { "count": 1 @@ -1511,11 +1461,6 @@ "count": 1 } }, - "src/components/add_model/AddModelForm.test.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/components/add_model/AddModelForm.tsx": { "local/no-complex-jsx-arrow": { "count": 1 @@ -1548,9 +1493,6 @@ "src/components/add_model/advanced_settings.tsx": { "local/filename-pascal-case": { "count": 1 - }, - "no-restricted-imports": { - "count": 1 } }, "src/components/add_model/auto_router_connection_test.tsx": { @@ -1570,9 +1512,6 @@ "local/no-complex-jsx-arrow": { "count": 1 }, - "no-restricted-imports": { - "count": 1 - }, "react-hooks/set-state-in-effect": { "count": 2 } @@ -1593,9 +1532,6 @@ }, "no-nested-ternary": { "count": 1 - }, - "no-restricted-imports": { - "count": 2 } }, "src/components/add_model/model_connection_test.tsx": { @@ -1613,9 +1549,6 @@ "no-nested-ternary": { "count": 3 }, - "no-restricted-imports": { - "count": 1 - }, "react-hooks/immutability": { "count": 3 } @@ -1625,15 +1558,7 @@ "count": 1 } }, - "src/components/agent_management/AgentSelector.test.tsx": { - "react/display-name": { - "count": 1 - } - }, "src/components/agent_management/AgentSelector.tsx": { - "no-restricted-imports": { - "count": 1 - }, "prefer-const": { "count": 1 } @@ -1724,11 +1649,6 @@ "count": 1 } }, - "src/components/common_components/AccessGroupSelector.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/components/common_components/DeleteResourceModal.tsx": { "react-hooks/set-state-in-effect": { "count": 1 @@ -1739,16 +1659,6 @@ "count": 1 } }, - "src/components/common_components/MetadataKeyValueFields.test.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, - "src/components/common_components/MetadataKeyValueFields.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/components/common_components/ModelAliasManager.tsx": { "react-hooks/set-state-in-effect": { "count": 1 @@ -1759,11 +1669,6 @@ "count": 1 } }, - "src/components/common_components/RateLimitTypeFormItem.test.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/components/common_components/budget_duration_dropdown.tsx": { "local/filename-pascal-case": { "count": 1 @@ -1772,9 +1677,6 @@ "src/components/common_components/check_openapi_schema.tsx": { "local/filename-pascal-case": { "count": 1 - }, - "no-restricted-imports": { - "count": 2 } }, "src/components/common_components/fetch_teams.tsx": { @@ -1845,16 +1747,6 @@ "count": 1 } }, - "src/components/key_team_helpers/BudgetFallbacksEditor.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, - "src/components/key_team_helpers/BudgetWindowsEditor.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/components/key_team_helpers/fetch_available_models_team_key.tsx": { "local/filename-pascal-case": { "count": 1 @@ -1953,11 +1845,6 @@ "count": 1 } }, - "src/components/model_dashboard/ModelSettingsModal/ModelSettingsModal.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/components/model_filters.tsx": { "local/filename-pascal-case": { "count": 1 @@ -2111,11 +1998,6 @@ "count": 1 } }, - "src/components/router_settings/RoutingStrategySelector.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/components/router_settings/index.tsx": { "local/filename-pascal-case": { "count": 1 @@ -2264,11 +2146,6 @@ "count": 1 } }, - "src/components/team/LoggingSettings.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/components/team/TeamInfo.tsx": { "max-lines": { "count": 1 @@ -2276,9 +2153,6 @@ "no-nested-ternary": { "count": 3 }, - "no-restricted-imports": { - "count": 1 - }, "react-hooks/set-state-in-effect": { "count": 1 } @@ -2614,11 +2488,6 @@ "count": 2 } }, - "src/contexts/AntdGlobalProvider.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/contexts/AuthContext.tsx": { "react-hooks/set-state-in-effect": { "count": 1 diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.integration.test.tsx index b762f006261..19e1e3aa8bd 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.integration.test.tsx @@ -219,7 +219,7 @@ describe("AddModelPanel submit payload contract", () => { const { user, openAdvanced, fillRequired, submit } = await setup(); await fillRequired(); await openAdvanced(); - await user.click(screen.getByLabelText("Custom Pricing")); + await user.click(screen.getByRole("switch", { name: "Custom Pricing" })); await user.type(await screen.findByLabelText("Input Cost (per 1M tokens)"), "3"); await user.type(screen.getByLabelText("Output Cost (per 1M tokens)"), "9"); await submit(); @@ -241,7 +241,7 @@ describe("AddModelPanel submit payload contract", () => { const { user, openAdvanced, fillRequired, submit } = await setup(); await fillRequired(); await openAdvanced(); - await user.click(screen.getByLabelText("Cache Control Injection Points")); + await user.click(screen.getByRole("switch", { name: "Cache Control Injection Points" })); await screen.findByText("Add Injection Point"); await submit(); @@ -260,7 +260,7 @@ describe("AddModelPanel submit payload contract", () => { const { user, openAdvanced, fillRequired, submit } = await setup(); await fillRequired(); await openAdvanced(); - await user.click(screen.getByLabelText("Cache Control Injection Points")); + await user.click(screen.getByRole("switch", { name: "Cache Control Injection Points" })); await screen.findByText("Add Injection Point"); await user.click(screen.getByText("Select a role")); await user.click(await screen.findByText("System")); @@ -385,7 +385,7 @@ describe("AddModelPanel behaviours the removed Advanced Settings form instance n const { user, openAdvanced, fillRequired, submit } = await setup(); await fillRequired(); await openAdvanced(); - await user.click(screen.getByLabelText("Use in pass through routes")); + await user.click(screen.getByRole("switch", { name: "Use in pass through routes" })); expect(screen.getByLabelText("LiteLLM Params")).toHaveValue(""); await submit(); @@ -401,11 +401,11 @@ describe("AddModelPanel behaviours the removed Advanced Settings form instance n const { user, openAdvanced, fillRequired, submit } = await setup(); await fillRequired(); await openAdvanced(); - await user.click(screen.getByLabelText("Custom Pricing")); + await user.click(screen.getByRole("switch", { name: "Custom Pricing" })); await user.type(await screen.findByLabelText("Input Cost (per 1M tokens)"), "3"); - await user.click(screen.getByLabelText("Custom Pricing")); + await user.click(screen.getByRole("switch", { name: "Custom Pricing" })); await waitFor(() => expect(screen.queryByLabelText("Input Cost (per 1M tokens)")).not.toBeInTheDocument()); - await user.click(screen.getByLabelText("Custom Pricing")); + await user.click(screen.getByRole("switch", { name: "Custom Pricing" })); expect(await screen.findByLabelText("Input Cost (per 1M tokens)")).toHaveValue("3"); await submit(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.tsx index 59d4f95c038..1443c065d9b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/AddModelPanel.tsx @@ -15,7 +15,6 @@ import { useModelCostMap } from "@/app/(dashboard)/hooks/models/useModelCostMap" import { useCredentials } from "@/app/(dashboard)/hooks/credentials/useCredentials"; import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; -import { vertexCredentialsUploadProps } from "@/app/(dashboard)/models-and-endpoints/vertexCredentialsUpload"; const INITIAL_VALUES: MountedFormValues = { litellm_credential_name: null }; @@ -60,9 +59,6 @@ export default function AddModelPanel() { providerModels={providerModels} setProviderModelsFn={(provider) => setProviderModels(getProviderModels(provider, modelCostMapData))} getPlaceholder={getPlaceholder} - uploadProps={vertexCredentialsUploadProps({ - setFieldsValue: (values) => form.setValue("vertex_credentials", values.vertex_credentials), - })} showAdvancedSettings={showAdvancedSettings} setShowAdvancedSettings={setShowAdvancedSettings} teams={teams ?? null} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/LlmCredentialsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/LlmCredentialsPanel.tsx index 7251da7c3c3..c71ed177416 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/LlmCredentialsPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/panels/LlmCredentialsPanel.tsx @@ -1,17 +1,7 @@ "use client"; -import { useForm } from "react-hook-form"; import CredentialsPanel from "@/components/model_add/CredentialsPanel"; -import type { MountedFormValues } from "@/components/common_components/MountedFormField"; -import { vertexCredentialsUploadProps } from "@/app/(dashboard)/models-and-endpoints/vertexCredentialsUpload"; export default function LlmCredentialsPanel() { - const form = useForm(); - return ( - form.setValue("vertex_credentials", values.vertex_credentials), - })} - /> - ); + return ; } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/vertexCredentialsUpload.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/vertexCredentialsUpload.test.ts deleted file mode 100644 index e02e4974353..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/vertexCredentialsUpload.test.ts +++ /dev/null @@ -1,47 +0,0 @@ -import { waitFor } from "@testing-library/react"; -import { beforeEach, describe, expect, it, vi } from "vitest"; - -import { toast } from "@/lib/toast"; - -import { vertexCredentialsUploadProps } from "./vertexCredentialsUpload"; - -const makeForm = () => ({ setFieldsValue: vi.fn() }); - -describe("vertexCredentialsUploadProps", () => { - beforeEach(() => { - vi.clearAllMocks(); - }); - - it("reads a JSON credential file into the vertex_credentials field without uploading it", async () => { - const form = makeForm(); - const props = vertexCredentialsUploadProps(form as never); - const file = new File(['{"project_id":"example"}'], "vertex.json", { type: "application/json" }); - - expect(props.beforeUpload?.(file as never, [file] as never)).toBe(false); - - await waitFor(() => { - expect(form.setFieldsValue).toHaveBeenCalledWith({ vertex_credentials: '{"project_id":"example"}' }); - }); - }); - - it("ignores non-JSON files", async () => { - const form = makeForm(); - const props = vertexCredentialsUploadProps(form as never); - const file = new File(["not json"], "vertex.txt", { type: "text/plain" }); - - expect(props.beforeUpload?.(file as never, [file] as never)).toBe(false); - - await new Promise((resolve) => setTimeout(resolve, 0)); - expect(form.setFieldsValue).not.toHaveBeenCalled(); - }); - - it("reports completed and failed upload states", () => { - const props = vertexCredentialsUploadProps(makeForm() as never); - - props.onChange?.({ file: { name: "vertex.json", status: "done" } } as never); - props.onChange?.({ file: { name: "vertex.json", status: "error" } } as never); - - expect(toast.success).toHaveBeenCalledWith("vertex.json file uploaded successfully"); - expect(toast.fromError).toHaveBeenCalledWith("vertex.json file upload failed."); - }); -}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/vertexCredentialsUpload.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/vertexCredentialsUpload.ts deleted file mode 100644 index f658ec88c1c..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/vertexCredentialsUpload.ts +++ /dev/null @@ -1,37 +0,0 @@ -import type { ComponentProps } from "react"; - -import { toast } from "@/lib/toast"; -import type CredentialsPanel from "@/components/model_add/CredentialsPanel"; - -interface VertexCredentialsForm { - setFieldsValue: (values: { vertex_credentials: string }) => void; -} - -type UploadProps = ComponentProps["uploadProps"]; - -export function vertexCredentialsUploadProps(form: VertexCredentialsForm): UploadProps { - return { - name: "file", - accept: ".json", - pastable: false, - beforeUpload: (file) => { - if (file.type === "application/json") { - const reader = new FileReader(); - reader.onload = (event) => { - if (event.target) { - form.setFieldsValue({ vertex_credentials: event.target.result as string }); - } - }; - reader.readAsText(file); - } - return false; - }, - onChange(info) { - if (info.file.status === "done") { - toast.success(`${info.file.name} file uploaded successfully`); - } else if (info.file.status === "error") { - toast.fromError(`${info.file.name} file upload failed.`); - } - }, - }; -} diff --git a/ui/litellm-dashboard/src/components/ModelInfoEditForm.tsx b/ui/litellm-dashboard/src/components/ModelInfoEditForm.tsx index fff732fe270..d56e65237eb 100644 --- a/ui/litellm-dashboard/src/components/ModelInfoEditForm.tsx +++ b/ui/litellm-dashboard/src/components/ModelInfoEditForm.tsx @@ -1,8 +1,6 @@ "use client"; import { zodResolver } from "@hookform/resolvers/zod"; -// eslint-disable-next-line no-restricted-imports -- the dashboard has no shadcn date-time picker; the PTU window fields need one -import { DatePicker } from "antd"; import { CircleHelp } from "lucide-react"; import type { Dayjs } from "dayjs"; import * as React from "react"; @@ -11,6 +9,7 @@ import { z } from "zod/v4"; import { TagsInput } from "@/app/(dashboard)/guardrails/_components/content_filter/TagsInput"; import { FormField } from "@/components/shared/form/FormField"; +import { UtcDateTimeInput } from "@/components/shared/form/UtcDateTimeInput"; import { Badge } from "@/components/ui/badge"; import { Button } from "@/components/ui/button"; import { Input } from "@/components/ui/input"; @@ -292,9 +291,16 @@ const Display: React.FC<{ children: React.ReactNode }> = ({ children }) => (
{children}
); -const FieldLabel: React.FC<{ children: React.ReactNode }> = ({ children }) => ( -

{children}

-); +const FIELD_LABEL_CLASS = "text-sm font-medium text-foreground"; + +const FieldLabel: React.FC<{ htmlFor?: string; children: React.ReactNode }> = ({ htmlFor, children }) => + htmlFor === undefined ? ( +

{children}

+ ) : ( + + ); const Hint: React.FC<{ text: string }> = ({ text }) => ( @@ -463,13 +469,14 @@ const ModelInfoEditForm: React.FC = ({ {ptuCostAttributionEnabled && PTU_EDIT_FIELDS.map((ptuField) => (
- {ptuField.label} + {ptuField.label} {isEditing ? ( {({ value, onChange, ...control }) => ptuField.input === "number" ? ( = ({ min={ptuField.isCount ? 1 : 0} /> ) : ( - ) diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.test.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.test.tsx index 225c308af91..c0db38eb187 100644 --- a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.test.tsx +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.test.tsx @@ -1,5 +1,5 @@ import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; -import { render, screen, waitFor } from "@testing-library/react"; +import { render, screen, waitFor, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import EditFallbacks, { Fallbacks } from "./EditFallbacks"; @@ -49,10 +49,9 @@ describe("EditFallbacks", () => { it("prefills the existing fallback chain for the primary model", async () => { setup(); - await waitFor(() => { - expect(screen.getByText("gpt-3.5-turbo")).toBeInTheDocument(); - expect(screen.getByText("claude-3-opus")).toBeInTheDocument(); - }); + const chain = await screen.findByRole("list", { name: "Fallback chain" }); + expect(within(chain).getByText("gpt-3.5-turbo")).toBeInTheDocument(); + expect(within(chain).getByText("claude-3-opus")).toBeInTheDocument(); }); it("removes a fallback model and saves only the edited entry", async () => { @@ -61,8 +60,8 @@ describe("EditFallbacks", () => { const onClose = vi.fn(); setup({ onChange, onClose }); - await screen.findByText("gpt-3.5-turbo"); - await user.click(screen.getByTestId("remove-fallback-gpt-3.5-turbo")); + const chain = await screen.findByRole("list", { name: "Fallback chain" }); + await user.click(within(chain).getByRole("button", { name: "Remove gpt-3.5-turbo" })); await user.click(screen.getByRole("button", { name: /save changes/i })); @@ -77,8 +76,8 @@ describe("EditFallbacks", () => { const onChange = vi.fn().mockResolvedValue(undefined); setup({ fallbackEntry: { "gpt-4": ["gpt-3.5-turbo"] }, onChange }); - await screen.findByText("gpt-3.5-turbo"); - await user.click(screen.getByTestId("remove-fallback-gpt-3.5-turbo")); + const chain = await screen.findByRole("list", { name: "Fallback chain" }); + await user.click(within(chain).getByRole("button", { name: "Remove gpt-3.5-turbo" })); const saveButton = screen.getByRole("button", { name: /save changes/i }); expect(saveButton).toBeDisabled(); diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx index e3a36f4fbf2..129f509dcd9 100644 --- a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx @@ -3,10 +3,10 @@ * Handles primary model selection and fallback chain configuration */ -import { SimpleTooltip } from "@/components/ui/tooltip"; -import { Select } from "antd"; +import { MultiSelect } from "@/components/shared/MultiSelect"; +import { SearchSelect } from "@/components/shared/SearchSelect"; import { AlertCircle, ArrowDown, X } from "lucide-react"; -import React from "react"; +import React, { useId } from "react"; export interface FallbackGroup { id: string; @@ -64,25 +64,24 @@ export function FallbackGroupConfig({ }; const canAddMoreFallbacks = group.fallbackModels.length < maxFallbacks; + const primaryModelInputId = useId(); return (
{/* Primary Model Section */}
-
) : ( - group.fallbackModels.map((modelValue, index) => { - return ( -
+ {group.fallbackModels.map((modelValue, index) => ( +
  • @@ -182,15 +154,15 @@ export function FallbackGroupConfig({ -
  • - ); - }) + + ))} + )}
    diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.test.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.test.tsx index ac9b4d98aed..738040149f0 100644 --- a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.test.tsx +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.test.tsx @@ -1,4 +1,4 @@ -import { render, screen } from "@testing-library/react"; +import { render, screen, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { FallbackSelectionForm } from "./FallbackSelectionForm"; @@ -69,7 +69,7 @@ describe("FallbackSelectionForm", () => { , ); - const addTabButton = screen.getByRole("button", { name: /add tab/i }); + const addTabButton = screen.getByRole("button", { name: /add fallback group/i }); await user.click(addTabButton); expect(mockOnGroupsChange).toHaveBeenCalledTimes(1); @@ -98,7 +98,7 @@ describe("FallbackSelectionForm", () => { maxGroups={5} />, ); - expect(screen.queryByRole("button", { name: /add tab/i })).not.toBeInTheDocument(); + expect(screen.queryByRole("button", { name: /add fallback group/i })).not.toBeInTheDocument(); }); it("should show add tab button when below maxGroups with custom maxGroups", () => { @@ -111,7 +111,7 @@ describe("FallbackSelectionForm", () => { maxGroups={3} />, ); - expect(screen.getByRole("button", { name: /add tab/i })).toBeInTheDocument(); + expect(screen.getByRole("button", { name: /add fallback group/i })).toBeInTheDocument(); }); it("should call onGroupsChange when a group is removed", async () => { @@ -124,7 +124,8 @@ describe("FallbackSelectionForm", () => { , ); - const removeButtons = screen.getAllByRole("tab", { name: "remove" }); + const removeButtons = screen.getAllByRole("button", { name: /^remove /i }); + expect(removeButtons).toHaveLength(2); await user.click(removeButtons[0]); expect(mockOnGroupsChange).toHaveBeenCalledTimes(1); @@ -139,7 +140,7 @@ describe("FallbackSelectionForm", () => { render( , ); - expect(screen.getByText("Select primary model")).toBeInTheDocument(); + expect(screen.getByRole("combobox", { name: /primary model/i })).toHaveValue(""); expect(screen.getByText("Primary Model")).toBeInTheDocument(); }); @@ -150,7 +151,8 @@ describe("FallbackSelectionForm", () => { ); expect(screen.getByRole("tab", { name: "gpt-4" })).toBeInTheDocument(); expect(screen.getAllByText("gpt-4").length).toBeGreaterThan(0); - expect(screen.getByText("gpt-3.5-turbo")).toBeInTheDocument(); + const chain = screen.getByRole("list", { name: "Fallback chain" }); + expect(within(chain).getByText("gpt-3.5-turbo")).toBeInTheDocument(); }); it("should not add group when add button clicked at maxGroups", () => { @@ -168,7 +170,7 @@ describe("FallbackSelectionForm", () => { />, ); - expect(screen.queryByRole("button", { name: /add tab/i })).not.toBeInTheDocument(); + expect(screen.queryByRole("button", { name: /add fallback group/i })).not.toBeInTheDocument(); expect(mockOnGroupsChange).not.toHaveBeenCalled(); }); }); diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.tsx index 10bbcd9ba3b..2639171fbe8 100644 --- a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.tsx +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.tsx @@ -5,8 +5,8 @@ */ import { Button } from "@/components/ui/button"; -import { Tabs } from "antd"; -import { Plus } from "lucide-react"; +import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; +import { Plus, X } from "lucide-react"; import React, { useEffect, useState } from "react"; import { toast } from "@/lib/toast"; import { FallbackGroup, FallbackGroupConfig } from "./FallbackGroupConfig"; @@ -76,23 +76,8 @@ export function FallbackSelectionForm({ onGroupsChange(newGroups); }; - // Generate tab items - const items = groups.map((group, index) => { - const label = group.primaryModel ? group.primaryModel : `Group ${index + 1}`; - return { - key: group.id, - label: label, - closable: groups.length > 1, // Only allow closing if there's more than 1 group - children: ( - - ), - }; - }); + const groupLabel = (group: FallbackGroup, index: number) => + group.primaryModel ? group.primaryModel : `Group ${index + 1}`; if (groups.length === 0) { return ( @@ -107,22 +92,47 @@ export function FallbackSelectionForm({ } return ( - { - if (action === "add") handleAddGroup(); - else if (action === "remove" && groups.length > 1) { - handleRemoveGroup(targetKey as string); - } - }} - items={items} - className="fallback-tabs" - tabBarStyle={{ - marginBottom: 0, - }} - hideAdd={groups.length >= maxGroups} - /> + +
    + + {groups.map((group, index) => ( +
    + 1 ? "pr-9" : "pr-4"}`} + > + {groupLabel(group, index)} + + {groups.length > 1 && ( + + )} +
    + ))} +
    + {groups.length < maxGroups && ( + + )} +
    + {groups.map((group) => ( + + + + ))} +
    ); } diff --git a/ui/litellm-dashboard/src/components/add_model/AdaptiveRoutingConfig.tsx b/ui/litellm-dashboard/src/components/add_model/AdaptiveRoutingConfig.tsx index 720b6f88db3..137ef297d56 100644 --- a/ui/litellm-dashboard/src/components/add_model/AdaptiveRoutingConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AdaptiveRoutingConfig.tsx @@ -1,4 +1,9 @@ -import { Card, InputNumber, Radio, Slider, Space, Switch, Typography } from "antd"; +import { Card, CardContent } from "@/components/ui/card"; +import { Input } from "@/components/ui/input"; +import { Label } from "@/components/ui/label"; +import { RadioGroup, RadioGroupItem } from "@/components/ui/radio-group"; +import { Slider } from "@/components/ui/slider"; +import { Switch } from "@/components/ui/switch"; import React from "react"; import { AdaptiveEligible, @@ -7,8 +12,6 @@ import { DEFAULT_TIER_DISTANCE_PENALTY, } from "./ComplexityRouterConfig"; -const { Text } = Typography; - interface AdaptiveRoutingConfigProps { value: ComplexityRouterConfigValue; onChange: (value: ComplexityRouterConfigValue) => void; @@ -45,84 +48,91 @@ const AdaptiveRoutingConfig: React.FC = ({ value, on return ( <> -
    - - Enable adaptive bandit selection -
    - + + When disabled, each request always uses the model assigned to its classified tier. - + - - How Adaptive Routing Works - - - It learns from how each conversation actually goes: does the user have to rephrase or correct the model, does - it get stuck repeating itself, does it run out of tool calls, does the user seem satisfied. Combined with - cost, this live feedback shifts future routing toward the models that are actually working well, and improves - as more conversations come in. Until there's enough feedback, it defaults to the classified tier's - model. - + + How Adaptive Routing Works + + It learns from how each conversation actually goes: does the user have to rephrase or correct the model, + does it get stuck repeating itself, does it run out of tool calls, does the user seem satisfied. Combined + with cost, this live feedback shifts future routing toward the models that are actually working well, and + improves as more conversations come in. Until there's enough feedback, it defaults to the classified + tier's model. + + {value.adaptive && (
    - + Quality vs. Cost ({Math.round(adaptiveWeights.quality * 100)}% quality /{" "} {Math.round(adaptiveWeights.cost * 100)}% cost) - + `${v}% quality / ${100 - (v ?? 0)}% cost` }} + value={[Math.round(adaptiveWeights.quality * 100)]} + onValueChange={(next) => handleQualityWeightChange(Array.isArray(next) ? next[0] : next)} /> - + Higher quality weight favors more capable (pricier) models; higher cost weight favors cheaper models when the bandit has feedback to act on. Recommended: 30% quality / 70% cost split. - +
    - - Eligible Model Pool - - Eligible Model Pool + handleAdaptiveEligibleChange(e.target.value)} + onValueChange={(eligible: unknown) => handleAdaptiveEligibleChange(eligible as AdaptiveEligible)} className="w-full" > - - - All tiers (soft floor){" "} - — router can pick across tiers, depending on the best fit for the prompt - - - Classified tier only{" "} - — router can only pick models within tier - - - +
    + + +
    +
    {adaptiveEligible === "all" && (
    - - Tier Distance Penalty - - Tier Distance Penalty + + handleTierDistancePenaltyChange(event.target.value === "" ? null : event.target.valueAsNumber) + } min={0} step={0.1} - style={{ width: "100%" }} + className="w-full" /> - + Score penalty applied per tier-step away from the classified tier. - +
    )}
    diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.test.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.test.tsx index 26bd9af94a8..340fd811e97 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.test.tsx @@ -1,6 +1,5 @@ import { renderHook, screen, waitFor, renderWithProviders } from "../../../tests/test-utils"; import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event"; -import type { UploadProps } from "antd/es/upload"; import { describe, expect, it, vi } from "vitest"; import type { Team } from "../key_team_helpers/key_list"; import type { CredentialItem } from "../networking"; @@ -157,11 +156,6 @@ const createTestProps = (userRole = "proxy_admin", userId = "user-1", isTeamAdmi }, ]; - const uploadProps: UploadProps = { - beforeUpload: () => false, - showUploadList: false, - }; - return { form, registry, @@ -176,7 +170,6 @@ const createTestProps = (userRole = "proxy_admin", userId = "user-1", isTeamAdmi showAdvancedSettings: false, teams, credentials, - uploadProps, userRole, userId, }; @@ -318,6 +311,35 @@ describe("AddModelForm", () => { expect(await screen.findByRole("button", { name: "Add Model" })).toBeInTheDocument(); }); + describe("the enterprise gate on the Team-BYOK switch", () => { + const renderForm = async (premiumUser: boolean) => { + const mockUseAuthorized = vi.mocked(await import("@/app/(dashboard)/hooks/useAuthorized")); + mockUseAuthorized.default.mockReturnValue(mockAuthorizedUser("proxy_admin", "user-1", premiumUser)); + renderWithProviders(); + return screen.findByRole("switch", { name: "Team-BYOK Model" }); + }; + + it("explains the gate on hover even though the switch it sits on is disabled", async () => { + const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never }); + const teamOnlySwitch = await renderForm(false); + expect(teamOnlySwitch).toHaveAttribute("aria-disabled", "true"); + + await user.hover(teamOnlySwitch); + + expect(await screen.findByText(/enterprise-only feature/)).toBeInTheDocument(); + }); + + it("says nothing on hover once the user is premium", async () => { + const user = userEvent.setup(); + const teamOnlySwitch = await renderForm(true); + expect(teamOnlySwitch).not.toHaveAttribute("aria-disabled", "true"); + + await user.hover(teamOnlySwitch); + + expect(screen.queryByText(/enterprise-only feature/)).not.toBeInTheDocument(); + }); + }); + describe("cache control bindings reach the parent form store", () => { const renderWithForm = async () => { const mockUseAuthorized = vi.mocked(await import("@/app/(dashboard)/hooks/useAuthorized")); @@ -331,11 +353,11 @@ describe("AddModelForm", () => { user, openCacheControl: async () => { await user.click(await screen.findByText("Advanced Settings")); - await user.click(screen.getByLabelText("Cache Control Injection Points")); + await user.click(screen.getByRole("switch", { name: "Cache Control Injection Points" })); await screen.findByText("Add Injection Point"); }, closeCacheControl: async () => { - await user.click(screen.getByLabelText("Cache Control Injection Points")); + await user.click(screen.getByRole("switch", { name: "Cache Control Injection Points" })); await waitFor(() => expect(screen.queryByText("Add Injection Point")).not.toBeInTheDocument()); }, mountedValues: async (): Promise> => props.mountedValues(), diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx index 65efd6aed51..b6dddf43588 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx @@ -9,7 +9,6 @@ import { Select as AntdSelect, Card, Col, Row, Tooltip, Typography } from "antd" import { Info } from "lucide-react"; import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert"; import { Button } from "@/components/ui/button"; -import type { UploadProps } from "antd/es/upload"; import React, { useEffect, useMemo, useState } from "react"; import { FormProvider, useWatch, type UseFormReturn } from "react-hook-form"; import TeamDropdown from "../common_components/team_dropdown"; @@ -44,7 +43,6 @@ interface AddModelFormProps { providerModels: string[]; setProviderModelsFn: (provider: Providers) => void; getPlaceholder: (provider: Providers) => string; - uploadProps: UploadProps; showAdvancedSettings: boolean; setShowAdvancedSettings: (show: boolean) => void; teams: Team[] | null; @@ -71,7 +69,6 @@ const AddModelForm: React.FC = ({ providerModels, setProviderModelsFn, getPlaceholder, - uploadProps, showAdvancedSettings, setShowAdvancedSettings, teams, @@ -311,7 +308,7 @@ const AddModelForm: React.FC = ({ OR
    - + )}
    diff --git a/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx index 019df7d4fa9..3c384aed1b3 100644 --- a/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx @@ -1,6 +1,13 @@ import { Info } from "lucide-react"; import { SimpleTooltip } from "@/components/ui/tooltip"; -import { Select as AntdSelect, Card, InputNumber, Radio, Space, Switch, Typography } from "antd"; +import { MultiSelect } from "@/components/shared/MultiSelect"; +import { SearchSelect } from "@/components/shared/SearchSelect"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { Card, CardContent } from "@/components/ui/card"; +import { Input } from "@/components/ui/input"; +import { Label } from "@/components/ui/label"; +import { RadioGroup, RadioGroupItem } from "@/components/ui/radio-group"; +import { Switch } from "@/components/ui/switch"; import React from "react"; import ClassifierPromptEditor from "./ClassifierPromptEditor"; import HeuristicScoringConfig from "./HeuristicScoringConfig"; @@ -21,8 +28,6 @@ import { effectiveTierLabel, } from "./ComplexityRouterConfig"; -const { Text } = Typography; - const DEFAULT_SCORING_EXPLANATION = "The router scores each request across 7 dimensions: token count, code presence, reasoning markers, technical " + "terms, simple indicators, multi-step patterns, and question complexity. The weighted score determines the tier:"; @@ -87,36 +92,35 @@ const HowClassificationWorks: React.FC<{ value: ComplexityRouterConfigValue }> = return ( - - How Classification Works - - - {scoringExplanation(value)} - - {ranges && ( -
      -
    • - {effectiveTierLabel("SIMPLE", value.tier_labels)}: Score < {ranges.simpleMedium} -
    • -
    • - {effectiveTierLabel("MEDIUM", value.tier_labels)}: Score {ranges.simpleMedium} -{" "} - {ranges.mediumComplex} -
    • -
    • - {effectiveTierLabel("COMPLEX", value.tier_labels)}: Score {ranges.mediumComplex} -{" "} - {ranges.complexReasoning} -
    • -
    • - {effectiveTierLabel("REASONING", value.tier_labels)}: Score > {ranges.complexReasoning}{" "} - (or 2+ reasoning markers with a score of at least {ranges.reasoningOverrideFloor}) -
    • -
    - )} - {!ranges && isError && ( - - The tier score ranges could not be loaded from the proxy. - - )} + + How Classification Works + {scoringExplanation(value)} + {ranges && ( +
      +
    • + {effectiveTierLabel("SIMPLE", value.tier_labels)}: Score < {ranges.simpleMedium} +
    • +
    • + {effectiveTierLabel("MEDIUM", value.tier_labels)}: Score {ranges.simpleMedium} -{" "} + {ranges.mediumComplex} +
    • +
    • + {effectiveTierLabel("COMPLEX", value.tier_labels)}: Score {ranges.mediumComplex} -{" "} + {ranges.complexReasoning} +
    • +
    • + {effectiveTierLabel("REASONING", value.tier_labels)}: Score >{" "} + {ranges.complexReasoning} (or 2+ reasoning markers with a score of at least{" "} + {ranges.reasoningOverrideFloor}) +
    • +
    + )} + {!ranges && isError && ( + + The tier score ranges could not be loaded from the proxy. + + )} +
    ); }; @@ -247,61 +251,64 @@ const ClassificationMethodConfig: React.FC = ({ return ( <> - handleClassifierTypeChange(e.target.value)} + onValueChange={(classifierType: unknown) => handleClassifierTypeChange(classifierType as ClassifierType)} className="w-full" > - - - Heuristic{" "} - (default) — rule-based scoring, no API calls, <1ms latency - - - LLM Classifier{" "} - — use a model to decide the tier (e.g. a small/fast model) - - - +
    + + +
    + {value.classifier_type === "llm" && (
    - - Classifier Model - - Classifier Model + - {classifierModelMissing && ( - - A classifier model is required - - )} + {classifierModelMissing && A classifier model is required}
    - - Timeout (ms) - - Timeout (ms) + + handleClassifierTimeoutChange(event.target.value === "" ? null : event.target.valueAsNumber) + } min={1} - style={{ width: "100%" }} + className="w-full" /> - + How long the classifier call has before it fails and the fallback below takes over. - +
    - Classification Rubric + Classification Rubric @@ -310,28 +317,37 @@ const ClassificationMethodConfig: React.FC = ({ content={usesCustomPrompt ? "Your custom prompt replaces the built-in rubric entirely" : undefined} className="w-full" > - ({ + - + {usesCustomPrompt ? "Not in use: the custom prompt below is the classifier's entire rubric." : CLASSIFICATION_RUBRIC_DESCRIPTIONS[classificationRubric].description} - +
    - - Classifier Prompt - + Classifier Prompt = ({ />
    - - If the classifier fails - - If the classifier fails + handleClassifierFallbackChange(e.target.value)} + onValueChange={(fallback: unknown) => handleClassifierFallbackChange(fallback as ClassifierFallback)} > - - - Score with the heuristic{" "} - — right when the classifier grades complexity too - - +
    + + +
    +
    + Applies when the classifier call errors, times out, or returns an unparseable response. - +
    - - Context Window Size - - Context Window Size + + handleClassifierContextWindowSizeChange(event.target.value === "" ? null : event.target.valueAsNumber) + } min={0} - style={{ width: "100%" }} + className="w-full" /> - + Number of prior user turns (tool output and harness reminders excluded) sent to the classifier as context, so a referring follow-up like "now do the same for the streaming path" is classified against what it refers to. Set to 0 to send only the current message. - +
    - - Context Per-Turn Character Limit - - Context Per-Turn Character Limit + + handleClassifierContextPerTurnCharsChange(event.target.value === "" ? null : event.target.valueAsNumber) + } min={1} - style={{ width: "100%" }} + className="w-full" /> - - Prior turns longer than this are truncated. - + Prior turns longer than this are truncated.
    - Include Assistant Turns + Include Assistant Turns
    - + Let the classifier read the assistant's replies, so difficulty the model stated rather than the user stays visible: a plan the assistant calls complex, approved with "yes", is classified on the work being approved. Context Window Size then counts the last N turns across both roles rather than the last N user turns. - +
    )} @@ -429,25 +449,29 @@ const ClassificationMethodConfig: React.FC = ({ {value.classifier_type === "heuristic" && (
    - Custom Technical Keywords + Custom Technical Keywords
    - + Optional: Add terms to the built-in list to improve classification accuracy on the technical dimension. (e.g., udp, kafka, terraform). - - + ({ label: keyword, value: keyword }))} value={customTechnicalKeywords ?? []} - onChange={(keywords: string[]) => onCustomTechnicalKeywordsChange?.(keywords)} - placeholder="Type a keyword and press Enter, or paste a comma-separated list" - tokenSeparators={[","]} - open={false} - suffixIcon={null} - style={{ width: "100%" }} - allowClear + onValueChange={(keywords: string[]) => + onCustomTechnicalKeywordsChange?.( + Array.from( + new Set(keywords.flatMap((keyword) => keyword.split(",").map((part) => part.trim())).filter(Boolean)), + ), + ) + } + placeholder="Type a keyword and press Enter" + emptyText="Type to add a keyword" + allowCustomValues + className="w-full" />
    )} diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx index d0fda9962e5..5f5ae703b0e 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx @@ -280,11 +280,28 @@ describe("ComplexityRouterConfig", () => { ); fireEvent.click(screen.getByText("Advanced: Classification Method")); const keywordsSection = screen.getByText("Custom Technical Keywords").closest("div")?.parentElement as HTMLElement; - const input = within(keywordsSection).getByRole("combobox"); - fireEvent.change(input, { target: { value: "udp," } }); + await user.type(within(keywordsSection).getByRole("combobox"), "udp"); + await user.click(await screen.findByText('Create "udp"')); expect(onCustomTechnicalKeywordsChange).toHaveBeenCalledWith(["udp"]); }); + it("splits a comma-separated keyword entry into one keyword per token", async () => { + const user = userEvent.setup(); + const onCustomTechnicalKeywordsChange = vi.fn(); + renderWithProviders( + , + ); + fireEvent.click(screen.getByText("Advanced: Classification Method")); + const keywordsSection = screen.getByText("Custom Technical Keywords").closest("div")?.parentElement as HTMLElement; + await user.type(within(keywordsSection).getByRole("combobox"), "udp, kafka ,terraform"); + await user.click(await screen.findByText('Create "udp, kafka ,terraform"')); + expect(onCustomTechnicalKeywordsChange).toHaveBeenCalledWith(["udp", "kafka", "terraform"]); + }); + it("should render an empty state when no keyword tier rules exist", () => { renderWithProviders(); fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching")); @@ -314,9 +331,7 @@ describe("ComplexityRouterConfig", () => { expect(newRules[0]).toMatchObject({ keywords: [], tier: "COMPLEX" }); }); - // The dropdown is closed, so antd has nothing for Enter to select and the word would only land - // on blur. Submitting used to provide that blur; it no longer can while the row reads as empty. - it("commits a typed keyword on Enter, with the dropdown closed", async () => { + it("commits a typed keyword to the rule it was typed into", async () => { const user = userEvent.setup(); const onKeywordTierRulesChange = vi.fn(); renderWithProviders( @@ -329,7 +344,8 @@ describe("ComplexityRouterConfig", () => { fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching")); const field = screen.getByText("Keywords 1").closest("div") as HTMLElement; - await user.type(within(field).getByRole("combobox"), "invoice{enter}"); + await user.type(within(field).getByRole("combobox"), "invoice"); + await user.click(await screen.findByText('Create "invoice"')); expect(onKeywordTierRulesChange).toHaveBeenCalledWith([{ id: "rule-1", keywords: ["invoice"], tier: "COMPLEX" }]); }); @@ -482,7 +498,7 @@ describe("ComplexityRouterConfig classifier fallback", () => { }; renderWithProviders(); fireEvent.click(screen.getByText("Advanced: Classification Method")); - expect(screen.getByRole("radio", { name: /Route to the default model/ })).toBeDisabled(); + expect(screen.getByRole("radio", { name: /Route to the default model/ })).toHaveAttribute("aria-disabled", "true"); }); it("hides the fallback choice for the heuristic classifier, which has nothing to fall back from", () => { @@ -584,8 +600,8 @@ describe("ComplexityRouterConfig classifier rubric", () => { it("records the chat preset the operator picks", async () => { const onChange = openClassificationPanel(llmValue); - fireEvent.mouseDown(screen.getByRole("combobox", { name: "Classification Rubric" })); - await userEvent.click(await screen.findByTitle("Chat")); + await userEvent.click(screen.getByRole("combobox", { name: "Classification Rubric" })); + await userEvent.click(await screen.findByRole("option", { name: "Chat" })); expect(onChange).toHaveBeenCalledWith( expect.objectContaining({ classifier_llm_config: expect.objectContaining({ classification_rubric: "chat" }) }), ); @@ -678,7 +694,7 @@ describe("ComplexityRouterConfig tier labels", () => { />, ); fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching")); - expect(screen.getByTitle("Deep")).toBeInTheDocument(); + expect(screen.getByRole("combobox", { name: "Route keyword rule 1 to tier" })).toHaveTextContent("Deep"); }); }); @@ -716,7 +732,7 @@ describe("ComplexityRouterConfig default model", () => { it("shows what the tiers currently imply, so an untouched router still names its default", () => { renderWithProviders(); - expect(screen.getByText("Derived from tiers: gpt-3.5-turbo")).toBeInTheDocument(); + expect(getDefaultModelSelect()).toHaveAttribute("placeholder", "Derived from tiers: gpt-3.5-turbo"); }); it("asks for a model rather than naming a derived one when no tier holds one", () => { @@ -725,7 +741,7 @@ describe("ComplexityRouterConfig default model", () => { tiers: { SIMPLE: [], MEDIUM: [], COMPLEX: [], REASONING: [] }, }; renderWithProviders(); - expect(screen.getByText("Add a model to the Simple or Medium tier")).toBeInTheDocument(); + expect(getDefaultModelSelect()).toHaveAttribute("placeholder", "Add a model to the Simple or Medium tier"); }); it("records a pinned model", async () => { @@ -734,7 +750,7 @@ describe("ComplexityRouterConfig default model", () => { renderWithProviders(); await user.click(getDefaultModelSelect()); - await user.click((await screen.findAllByTitle("claude-3-opus")).slice(-1)[0]); + await user.click(await screen.findByRole("option", { name: "claude-3-opus" })); expect(onChange).toHaveBeenCalledWith(expect.objectContaining({ default_model: "claude-3-opus" })); }); @@ -745,8 +761,7 @@ describe("ComplexityRouterConfig default model", () => { const pinned: ComplexityRouterConfigValue = { ...defaultValue, default_model: "claude-3-opus" }; renderWithProviders(); - // eslint-disable-next-line local/no-antd-class-selectors -- antd marks the clear affordance aria-hidden, so no accessible query reaches it - await user.click(document.querySelector(".ant-select-clear") as HTMLElement); + await user.click(screen.getByRole("button", { name: "Clear" })); expect(onChange).toHaveBeenCalledWith(expect.objectContaining({ default_model: undefined })); }); @@ -754,10 +769,7 @@ describe("ComplexityRouterConfig default model", () => { it("shows a pinned model as the selection instead of the tier-derived one", () => { const pinned: ComplexityRouterConfigValue = { ...defaultValue, default_model: "claude-3-opus" }; renderWithProviders(); - expect( - // eslint-disable-next-line local/no-antd-class-selectors -- the tier selects show the same model as a tag, so the assertion has to scope to this select's root, which antd exposes only as a class - within(getDefaultModelSelect().closest(".ant-select") as HTMLElement).getByTitle("claude-3-opus"), - ).toBeInTheDocument(); + expect(getDefaultModelSelect()).toHaveValue("claude-3-opus"); }); it("unlocks the default model fallback on a pin alone, with no tier to derive from", () => { @@ -770,7 +782,7 @@ describe("ComplexityRouterConfig default model", () => { }; renderWithProviders(); fireEvent.click(screen.getByText("Advanced: Classification Method")); - expect(screen.getByRole("radio", { name: /Route to the default model/ })).toBeEnabled(); + expect(screen.getByRole("radio", { name: /Route to the default model/ })).not.toHaveAttribute("aria-disabled"); }); it("names the resolved default on the fallback option, so the destination is not a guess", () => { @@ -827,9 +839,9 @@ describe("plan-mode override", () => { />, ); openPanel(); - fireEvent.mouseDown(await screen.findByRole("combobox", { name: "Plan-mode minimum tier" })); - expect(await screen.findByTitle("Medium")).toBeInTheDocument(); - expect(screen.queryByTitle("Reasoning")).not.toBeInTheDocument(); + await userEvent.click(await screen.findByRole("combobox", { name: "Plan-mode minimum tier" })); + expect(await screen.findByRole("option", { name: "Medium" })).toBeInTheDocument(); + expect(screen.queryByRole("option", { name: "Reasoning" })).not.toBeInTheDocument(); }); it("disables the toggle until some tier has models", async () => { @@ -840,6 +852,6 @@ describe("plan-mode override", () => { />, ); openPanel(); - expect(await screen.findByRole("switch", { name: switchName })).toBeDisabled(); + expect(await screen.findByRole("switch", { name: switchName })).toHaveAttribute("aria-disabled", "true"); }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx index ad5260c2c9c..ab6b1d401ce 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx @@ -1,6 +1,13 @@ -import { Info } from "lucide-react"; import { SimpleTooltip } from "@/components/ui/tooltip"; -import { Select as AntdSelect, Card, Collapse, Divider, Input, Space, Switch, Typography } from "antd"; +import { MultiSelect } from "@/components/shared/MultiSelect"; +import { SearchSelect } from "@/components/shared/SearchSelect"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { ChevronRight, Info, X } from "lucide-react"; +import { Switch } from "@/components/ui/switch"; +import { Card, CardContent } from "@/components/ui/card"; +import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible"; +import { InputGroup, InputGroupAddon, InputGroupButton, InputGroupInput } from "@/components/ui/input-group"; +import { Separator } from "@/components/ui/separator"; import React from "react"; import { ModelGroup } from "@/components/llm_calls/fetch_models"; import AdaptiveRoutingConfig from "./AdaptiveRoutingConfig"; @@ -13,8 +20,6 @@ import { type DimensionWeights, type TierBoundaries, type TokenThresholds } from export type { DimensionWeights, TierBoundaries, TokenThresholds }; -const { Text } = Typography; - export const DEFAULT_CLASSIFIER_TIMEOUT_MS = 3000; export const DEFAULT_TIER_DISTANCE_PENALTY = 0.5; export const DEFAULT_CLASSIFIER_CONTEXT_WINDOW_SIZE = 3; @@ -218,6 +223,9 @@ const ComplexityRouterConfig: React.FC = ({ showValidationErrors = false, }) => { const planModeTiers = planModeEligibleTiers(value.tiers); + const planModeTierOptions = tierOptions(value.tier_labels).filter((option) => + (planModeTiers as string[]).includes(option.value), + ); const derivedDefaultModel = resolveComplexityDefaultModel(value.tiers); const defaultModel = resolveComplexityDefaultModel(value.tiers, value.default_model); @@ -251,128 +259,119 @@ const ComplexityRouterConfig: React.FC = ({ return (
    - - - Complexity Tier Configuration - +
    +

    Complexity Tier Configuration

    - +
    - + The complexity router automatically classifies requests by complexity using rule-based scoring (no API calls, <1ms latency). Configure which model(s) handle each tier. - + - + Rename a tier to use your own vocabulary in the dashboard and your spend logs. Renaming doesn't change how requests are classified, and callers never see these names. {value.classifier_type === "llm" && " Your classifier model reads these names, so clearer ones can sharpen its choices."} - + - {TIER_KEYS.map((tier, index) => { - const tierInfo = TIER_DESCRIPTIONS[tier]; - const label = effectiveTierLabel(tier, value.tier_labels); - const tierMissing = showValidationErrors && value.tiers[tier].length === 0; - return ( -
    - {index > 0 && } -
    -
    - - {label} Tier - - - - - - Tier {index + 1} of {TIER_KEYS.length} · {tier} - + + {TIER_KEYS.map((tier, index) => { + const tierInfo = TIER_DESCRIPTIONS[tier]; + const label = effectiveTierLabel(tier, value.tier_labels); + const tierMissing = showValidationErrors && value.tiers[tier].length === 0; + return ( +
    + {index > 0 && } +
    +
    + {label} Tier + + + + + Tier {index + 1} of {TIER_KEYS.length} · {tier} + +
    + Examples: {tierInfo.examples} + + handleTierLabelChange(tier, event.target.value)} + placeholder={`Display name (default: ${tierInfo.label})`} + aria-label={`Display name for the ${tierInfo.label} tier`} + /> + {value.tier_labels?.[tier] && ( + + handleTierLabelChange(tier, "")} + > + + + + )} + + handleTierChange(tier, models)} + placeholder={`Select model(s) for ${label.toLowerCase()} queries`} + emptyText="No models found" + className={tierMissing ? "w-full border-destructive" : "w-full"} + /> + {value.tiers[tier].length > 1 && ( + + Multiple models selected — the router randomly picks among them per request (or Thompson-samples + within the pool when adaptive routing is on). + + )} + {tierMissing && The {label} tier is required}
    - - Examples: {tierInfo.examples} - - handleTierLabelChange(tier, event.target.value)} - placeholder={`Display name (default: ${tierInfo.label})`} - aria-label={`Display name for the ${tierInfo.label} tier`} - style={{ marginBottom: 8 }} - allowClear - /> - handleTierChange(tier, models)} - placeholder={`Select model(s) for ${label.toLowerCase()} queries`} - showSearch - style={{ width: "100%" }} - options={modelOptions} - status={tierMissing ? "error" : undefined} - /> - {value.tiers[tier].length > 1 && ( - - Multiple models selected — the router randomly picks among them per request (or Thompson-samples - within the pool when adaptive routing is on). - - )} - {tierMissing && ( - - The {label} tier is required - - )}
    -
    - ); - })} - + ); + })} + -
    -
    - - Default Model - - - - +
    +
    + Default Model + + + +
    + + + Used when the tier the request lands in has no model, and when the classifier fails with "Route to + the default model" selected. +
    - - - Used when the tier the request lands in has no model, and when the classifier fails with "Route to the - default model" selected. - -
    + - + - + {[ { key: "classifier", - label: ( - - Advanced: Classification Method - - ), + label: Advanced: Classification Method, children: ( = ({ }, { key: "adaptive", - label: ( - - Advanced: Adaptive Routing - - ), + label: Advanced: Adaptive Routing, children: , }, { key: "affinity", - label: ( - - Advanced: Affinity - - ), + label: Advanced: Affinity, children: ( <>
    onChange({ ...value, deployment_affinity: deploymentAffinity })} + onCheckedChange={(deploymentAffinity) => + onChange({ ...value, deployment_affinity: deploymentAffinity }) + } aria-label="Pin a session to one deployment per model group" /> - Pin a session to one deployment per model group + Pin a session to one deployment per model group
    - + Keeps a session on the same deployment within a group, so provider prompt caches stay warm. Turn off to load-balance every turn. - +
    onChange({ ...value, session_affinity: sessionAffinity })} + onCheckedChange={(sessionAffinity) => onChange({ ...value, session_affinity: sessionAffinity })} aria-label="Pin a session to its first model" /> - Pin a session to its first model + Pin a session to its first model
    - + Keeps a session on its first turn's model instead of re-classifying each turn. Also pins the deployment. - + ), }, { key: "plan-mode", - label: ( - - Advanced: Plan-Mode Override - - ), + label: Advanced: Plan-Mode Override, children: ( <>
    + onCheckedChange={(enabled) => onChange({ ...value, plan_mode_min_tier: enabled ? planModeTiers.at(-1) : undefined }) } aria-label="Route plan-mode requests to a minimum tier" /> - Route plan-mode requests to a minimum tier + Route plan-mode requests to a minimum tier
    - + Requests from coding agents in plan mode (Claude Code, GitHub Copilot) route to at least this tier. The classifier still wins when it picks higher, and the override only lasts while plan mode is active. {planModeTiers.length === 0 && " Add models to a tier to enable this."} - + {value.plan_mode_min_tier !== undefined && (
    - - (planModeTiers as string[]).includes(option.value), - )} - onChange={(tier: string) => onChange({ ...value, plan_mode_min_tier: tier })} - /> + onValueChange={(tier: string | null) => tier && onChange({ ...value, plan_mode_min_tier: tier })} + > + + + + + {planModeTierOptions.map((option) => ( + + {option.label} + + ))} + +
    )} @@ -473,23 +469,22 @@ const ComplexityRouterConfig: React.FC = ({ }, { key: "response", - label: ( - - Advanced: Response Format - - ), + label: Advanced: Response Format, children: ( <>
    onChange({ ...value, return_raw_model_name: returnRawModelName })} + onCheckedChange={(returnRawModelName) => + onChange({ ...value, return_raw_model_name: returnRawModelName }) + } + aria-label="Return raw model name" /> - Return raw model name + Return raw model name
    - + Return the resolved underlying model name in responses instead of the autorouter alias. - + ), }, @@ -497,11 +492,7 @@ const ComplexityRouterConfig: React.FC = ({ ? [ { key: "escalation", - label: ( - - Advanced: Escalation Keywords - - ), + label: Advanced: Escalation Keywords, children: , }, ] @@ -510,11 +501,7 @@ const ComplexityRouterConfig: React.FC = ({ ? [ { key: "keyword-semantic", - label: ( - - Advanced: Keyword/Semantic Matching - - ), + label: Advanced: Keyword/Semantic Matching, children: ( <> {onKeywordTierRulesChange && ( @@ -524,9 +511,7 @@ const ComplexityRouterConfig: React.FC = ({ tierLabels={value.tier_labels} /> )} - {onKeywordTierRulesChange && onSemanticMatchingEnabledChange && ( - - )} + {onKeywordTierRulesChange && onSemanticMatchingEnabledChange && } {onSemanticMatchingEnabledChange && ( = ({ }, ] : []), - ]} - /> + ].map(({ key, label, children }) => ( + + + + {label} + + {children} + + ))} +
    ); }; diff --git a/ui/litellm-dashboard/src/components/add_model/EscalationKeywords.tsx b/ui/litellm-dashboard/src/components/add_model/EscalationKeywords.tsx index a7c1b7462ea..b1bb25deb13 100644 --- a/ui/litellm-dashboard/src/components/add_model/EscalationKeywords.tsx +++ b/ui/litellm-dashboard/src/components/add_model/EscalationKeywords.tsx @@ -1,10 +1,8 @@ import { Info } from "lucide-react"; import { SimpleTooltip } from "@/components/ui/tooltip"; -import { Select as AntdSelect, Typography } from "antd"; +import { MultiSelect } from "@/components/shared/MultiSelect"; import React from "react"; -const { Text } = Typography; - export const DEFAULT_ESCALATION_KEYWORDS = ["LITELLM ESCALATE"]; interface EscalationKeywordsProps { @@ -16,28 +14,24 @@ const EscalationKeywords: React.FC = ({ keywords, onCha return (
    - - Escalation Keywords - +

    Escalation Keywords

    - + Optional: when a user message contains one of these phrases, the request is bumped one tier higher than it would otherwise route to. Matching is case-sensitive, so "LITELLM ESCALATE" only fires on the exact, shouted form. Leave empty to disable. - - + ({ label: keyword, value: keyword }))} value={keywords} - onChange={onChange} + onValueChange={onChange} placeholder="e.g., LITELLM ESCALATE" - tokenSeparators={[","]} - open={false} - suffixIcon={null} - style={{ width: "100%" }} - allowClear + emptyText="Type to add a phrase" + allowCustomValues + className="w-full" />
    ); diff --git a/ui/litellm-dashboard/src/components/add_model/KeywordTierRules.tsx b/ui/litellm-dashboard/src/components/add_model/KeywordTierRules.tsx index 08d00606987..f7edd26f0d8 100644 --- a/ui/litellm-dashboard/src/components/add_model/KeywordTierRules.tsx +++ b/ui/litellm-dashboard/src/components/add_model/KeywordTierRules.tsx @@ -1,14 +1,14 @@ -import { Info, Plus, Trash2 } from "lucide-react"; +import { Inbox, Info, Plus, Trash2 } from "lucide-react"; import { SimpleTooltip } from "@/components/ui/tooltip"; -import { Card, Empty, Select as AntdSelect, Typography } from "antd"; +import { MultiSelect } from "@/components/shared/MultiSelect"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { Card, CardContent } from "@/components/ui/card"; import { Button } from "@/components/ui/button"; import React from "react"; import { emptyKeywordTierRuleIndexes } from "./complexity_router_keywords"; import { tierOptions } from "./complexity_router_tiers"; -const { Text } = Typography; - export type ComplexityTier = "SIMPLE" | "MEDIUM" | "COMPLEX" | "REASONING"; export interface KeywordTierRule { @@ -28,29 +28,9 @@ interface KeywordTierRulesProps { // there is no failed attempt left to surface it. const KeywordTierRules: React.FC = ({ rules, onChange, tierLabels }) => { const emptyRuleIndexes = new Set(emptyKeywordTierRuleIndexes(rules)); - const [drafts, setDrafts] = React.useState>({}); - - const setDraft = (id: string, text: string) => setDrafts((current) => ({ ...current, [id]: text })); - - // The dropdown is kept closed, which leaves antd nothing for Enter to select, so a typed keyword - // would only become a tag on blur. Submitting used to supply that blur; the button is disabled - // while the row reads as empty, so Enter has to commit the word itself or the row cannot be filled. - const commitDraft = (rule: KeywordTierRule) => { - const keyword = (drafts[rule.id] ?? "").trim(); - if (!keyword) return; - updateRule(rule.id, { keywords: [...rule.keywords, keyword] }); - setDraft(rule.id, ""); - }; - - const commitDraftOnEnter = (rule: KeywordTierRule) => (event: React.KeyboardEvent) => { - if (event.key !== "Enter") return; - event.preventDefault(); - commitDraft(rule); - }; const replaceKeywords = (rule: KeywordTierRule) => (keywords: string[]) => { updateRule(rule.id, { keywords }); - setDraft(rule.id, ""); }; const addRule = () => { @@ -69,9 +49,7 @@ const KeywordTierRules: React.FC = ({ rules, onChange, ti
    - - Keyword Tier Overrides - +

    Keyword Tier Overrides

    @@ -81,67 +59,71 @@ const KeywordTierRules: React.FC = ({ rules, onChange, ti Add keyword rule
    - + Optional: route requests containing specific keywords directly to a tier, e.g. route "invoice, refund, billing" to the medium tier. - + {rules.length === 0 ? ( - + +
    +
    +
    ) : (
    {rules.map((rule, index) => ( - -
    -
    - - Keywords {index + 1} - - setDraft(rule.id, text)} - onInputKeyDown={commitDraftOnEnter(rule)} - onBlur={() => commitDraft(rule)} - placeholder="e.g., invoice, refund, billing" - tokenSeparators={[","]} - open={false} - suffixIcon={null} - style={{ width: "100%" }} - allowClear - status={emptyRuleIndexes.has(index) ? "error" : undefined} - /> - {emptyRuleIndexes.has(index) && ( - - At least one keyword is required - - )} + + +
    +
    + Keywords {index + 1} + ({ label: keyword, value: keyword }))} + value={rule.keywords} + onValueChange={replaceKeywords(rule)} + placeholder="e.g., invoice, refund, billing" + emptyText="Type to add a keyword" + allowCustomValues + className={emptyRuleIndexes.has(index) ? "w-full border-destructive" : "w-full"} + /> + {emptyRuleIndexes.has(index) && ( + At least one keyword is required + )} +
    +
    + Route to tier + +
    +
    -
    - - Route to tier - - updateRule(rule.id, { tier })} - options={tierOptions(tierLabels)} - style={{ width: "100%" }} - /> -
    - -
    + ))}
    diff --git a/ui/litellm-dashboard/src/components/add_model/SemanticKeywordMatching.tsx b/ui/litellm-dashboard/src/components/add_model/SemanticKeywordMatching.tsx index 7f8a71b1b8b..9da0bf394e5 100644 --- a/ui/litellm-dashboard/src/components/add_model/SemanticKeywordMatching.tsx +++ b/ui/litellm-dashboard/src/components/add_model/SemanticKeywordMatching.tsx @@ -1,11 +1,11 @@ import { Info } from "lucide-react"; import { SimpleTooltip } from "@/components/ui/tooltip"; -import { InputNumber, Select as AntdSelect, Switch, Typography } from "antd"; +import { SearchSelect } from "@/components/shared/SearchSelect"; +import { Input } from "@/components/ui/input"; +import { Switch } from "@/components/ui/switch"; import React from "react"; import { ModelGroup } from "@/components/llm_calls/fetch_models"; -const { Text } = Typography; - const DEFAULT_MATCH_THRESHOLD = 0.5; interface SemanticKeywordMatchingProps { @@ -41,49 +41,49 @@ const SemanticKeywordMatching: React.FC = ({
    - Semantic keyword matching + Semantic keyword matching
    - + Uses same keyword-tier pairs as above and overrides direct keyword matching. Adds latency based on embedding model network request. - +
    - +
    {enabled && (
    - Embedding model - Embedding model + - {embeddingModelMissing && ( - - An embedding model is required - - )} + {embeddingModelMissing && An embedding model is required}
    - Minimum match score - Minimum match score + onMatchThresholdChange(value ?? DEFAULT_MATCH_THRESHOLD)} + onChange={(event) => + onMatchThresholdChange(event.target.value === "" ? DEFAULT_MATCH_THRESHOLD : event.target.valueAsNumber) + } min={0} max={1} step={0.05} - style={{ width: "100%" }} + className="w-full" /> - Match only at or above this similarity score. + Match only at or above this similarity score.
    )} diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx index e3a7482b02d..e45422dee08 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx @@ -27,7 +27,7 @@ const ALL_FAMILY_MODELS: ModelGroup[] = [ const ANTHROPIC_ONLY_MODEL = ANTHROPIC_TIERS.COMPLEX[0]; const openTemplateDropdown = (): void => { - fireEvent.mouseDown(within(screen.getByTestId("template-selector")).getByRole("combobox")); + fireEvent.click(screen.getByTestId("template-selector")); }; // Detailed Configuration is collapsed by default, so any test reaching into it (a tier select, an @@ -36,15 +36,32 @@ const expandDetailedConfiguration = (): void => { fireEvent.click(screen.getByTestId("detailed-configuration-toggle")); }; -const visibleOptions = (): HTMLElement[] => - // eslint-disable-next-line local/no-antd-class-selectors -- antd puts role="option" only on a hidden mirror list of raw values; the visible options carry no role, no aria-disabled, and only a tooltip in title - Array.from(document.querySelectorAll(".ant-select-item-option")); +const visibleOptions = (): HTMLElement[] => screen.queryAllByRole("option"); const optionByLabel = (label: string): HTMLElement | undefined => visibleOptions().find((el) => el.textContent?.startsWith(label)); -// eslint-disable-next-line local/no-antd-class-selectors -- antd signals option disabled state only through this class -const isOptionDisabled = (option: HTMLElement): boolean => option.classList.contains("ant-select-item-option-disabled"); +const isOptionDisabled = (option: HTMLElement): boolean => option.getAttribute("aria-disabled") === "true"; + +const selectTemplate = async (label: string): Promise => { + await userEvent.click(optionByLabel(label)!); +}; + +// Opens the dropdown only when it is closed, since openTemplateDropdown toggles: waiting on a +// second preset in the same test would otherwise close the list out from under the poll. +const waitForPresetEnabled = async (label: string) => { + if (visibleOptions().length === 0) openTemplateDropdown(); + await waitFor(() => { + expect(isOptionDisabled(optionByLabel(label)!)).toBe(false); + }); +}; + +// The keyword field is a combobox that offers whatever is typed as a "Create ..." entry, so a +// keyword only lands on the rule once that entry is picked. +const addKeyword = async (user: ReturnType, field: HTMLElement, keyword: string) => { + await user.type(within(field).getByRole("combobox"), keyword); + await user.click(await screen.findByText(`Create "${keyword}"`)); +}; const { mockFetchAvailableModels, mockFetchAllModelDeployments } = vi.hoisted(() => ({ mockFetchAvailableModels: vi.fn(), @@ -202,10 +219,7 @@ describe("AddAutoRouterTab", () => { await user.click(screen.getByRole("button", { name: /add keyword rule/i })); expect(screen.getByRole("button", { name: /add auto router/i })).toBeDisabled(); - await user.type( - within(screen.getByText("Keywords 1").closest("div") as HTMLElement).getByRole("combobox"), - "invoice{enter}", - ); + await addKeyword(user, screen.getByText("Keywords 1").closest("div") as HTMLElement, "invoice"); expect(screen.getByRole("button", { name: /add auto router/i })).toBeEnabled(); expect(screen.queryByText("At least one keyword is required")).not.toBeInTheDocument(); @@ -221,10 +235,7 @@ describe("AddAutoRouterTab", () => { expandDetailedConfiguration(); await user.click(screen.getByText("Advanced: Keyword/Semantic Matching")); await user.click(screen.getByRole("button", { name: /add keyword rule/i })); - await user.type( - within(screen.getByText("Keywords 1").closest("div") as HTMLElement).getByRole("combobox"), - "invoice{enter}", - ); + await addKeyword(user, screen.getByText("Keywords 1").closest("div") as HTMLElement, "invoice"); await user.click(screen.getByRole("button", { name: /add keyword rule/i })); expect(await screen.findAllByText("At least one keyword is required")).toHaveLength(1); @@ -242,7 +253,7 @@ describe("AddAutoRouterTab", () => { await user.click(screen.getByText("Advanced: Keyword/Semantic Matching")); await user.click(screen.getByRole("button", { name: /add keyword rule/i })); const keywordsField = screen.getByText("Keywords 1").closest("div") as HTMLElement; - await user.type(within(keywordsField).getByRole("combobox"), "invoice{enter}"); + await addKeyword(user, keywordsField, "invoice"); await user.click(screen.getByRole("button", { name: /add auto router/i })); await waitFor(() => expect(handleAddAutoRouterSubmit).toHaveBeenCalled()); @@ -407,7 +418,7 @@ describe("AddAutoRouterTab", () => { await user.click(screen.getByText("Advanced: Keyword/Semantic Matching")); await user.click(screen.getByRole("button", { name: /add keyword rule/i })); const keywordsField = screen.getByText("Keywords 1").closest("div") as HTMLElement; - await user.type(within(keywordsField).getByRole("combobox"), "invoice{enter}"); + await addKeyword(user, keywordsField, "invoice"); await user.click(screen.getByTestId("auto-router-test-routing-btn")); await user.type(await screen.findByTestId("auto-router-routing-test-prompt"), "reconcile this invoice"); @@ -452,17 +463,6 @@ describe("AddAutoRouterTab", () => { }); describe("template presets", () => { - // Opens the dropdown once, then waits out the useQuery load: an open antd Select re-renders its - // already-mounted options in place as state changes, so polling only re-reads the DOM here. - // Re-firing the open/close mousedown on every poll (calling openTemplateDropdown inside the - // waitFor callback) fights the dropdown's own open/close animation and hangs the test. - const waitForPresetEnabled = async (label: string) => { - openTemplateDropdown(); - await waitFor(() => { - expect(isOptionDisabled(optionByLabel(label)!)).toBe(false); - }); - }; - it("disables every preset while the model list is loading", async () => { let resolveModels: (models: ModelGroup[]) => void = () => {}; mockFetchAvailableModels.mockImplementation( @@ -537,7 +537,7 @@ describe("AddAutoRouterTab", () => { renderWithProviders(); await waitForPresetEnabled("Anthropic Family"); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); expect(screen.queryByText("Advanced: Keyword/Semantic Matching")).not.toBeInTheDocument(); expect( @@ -548,11 +548,11 @@ describe("AddAutoRouterTab", () => { ).toBeInTheDocument(); }); - it("expands detailed configuration when Custom Configuration is chosen", () => { + it("expands detailed configuration when Custom Configuration is chosen", async () => { renderWithProviders(); openTemplateDropdown(); - fireEvent.click(optionByLabel("Custom Configuration")!); + await selectTemplate("Custom Configuration"); expect(screen.getByText("Advanced: Keyword/Semantic Matching")).toBeInTheDocument(); }); @@ -561,7 +561,7 @@ describe("AddAutoRouterTab", () => { mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS); renderWithProviders(); await waitForPresetEnabled("Anthropic Family"); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); expect(screen.queryByText("Advanced: Keyword/Semantic Matching")).not.toBeInTheDocument(); fireEvent.click(screen.getByTestId("detailed-configuration-toggle")); @@ -578,7 +578,7 @@ describe("AddAutoRouterTab", () => { renderWithProviders(); await waitForPresetEnabled("Anthropic Family"); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); await user.type(screen.getByPlaceholderText(/smart_router/i), "anthropic-router"); await user.click(screen.getByRole("button", { name: /add auto router/i })); @@ -600,7 +600,7 @@ describe("AddAutoRouterTab", () => { const { container } = renderWithProviders(); await waitForPresetEnabled("Anthropic Family"); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); fireEvent.change(screen.getByPlaceholderText(/smart_router/i), { target: { value: "stale-model-router" } }); expect(screen.getByRole("button", { name: /add auto router/i })).toBeEnabled(); @@ -622,26 +622,16 @@ describe("AddAutoRouterTab", () => { describe("default model pin", () => { const PINNED_MODEL = "pinned-default-model"; - const waitForPresetEnabled = async (label: string) => { - openTemplateDropdown(); - await waitFor(() => { - expect(isOptionDisabled(optionByLabel(label)!)).toBe(false); - }); - }; - const applyPresetAndPin = async (user: ReturnType) => { await waitForPresetEnabled("Anthropic Family"); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); // Applying a preset collapses Detailed Configuration, so the default model row is behind it. expandDetailedConfiguration(); const defaultModel = screen.getByRole("combobox", { name: "Default model" }); await user.click(defaultModel); - // antd virtualizes the option list and jsdom gives every row zero height, so options past - // the first window never render. Typing filters the list down to the pin instead of relying - // on its index, which adding a preset to the bundled JSON shifts. await user.type(defaultModel, PINNED_MODEL); - await user.click((await screen.findAllByTitle(PINNED_MODEL)).slice(-1)[0]); + await user.click(await screen.findByRole("option", { name: PINNED_MODEL })); }; beforeEach(() => { @@ -692,19 +682,12 @@ describe("AddAutoRouterTab", () => { mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS); }); - const waitForPresetEnabled = async (label: string) => { - openTemplateDropdown(); - await waitFor(() => { - expect(isOptionDisabled(optionByLabel(label)!)).toBe(false); - }); - }; - it("omits plan_mode_min_tier from the payload when never touched", async () => { const user = userEvent.setup(); renderWithProviders(); await waitForPresetEnabled("Anthropic Family"); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); await user.type(screen.getByPlaceholderText(/smart_router/i), "no-plan-router"); await user.click(screen.getByRole("button", { name: /add auto router/i })); @@ -719,7 +702,7 @@ describe("AddAutoRouterTab", () => { renderWithProviders(); await waitForPresetEnabled("Anthropic Family"); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); expandDetailedConfiguration(); await user.click(screen.getByText("Advanced: Plan-Mode Override")); await user.click(await screen.findByRole("switch", { name: "Route plan-mode requests to a minimum tier" })); @@ -773,7 +756,7 @@ describe("AddAutoRouterTab", () => { await waitFor(() => { expect(isOptionDisabled(optionByLabel("Anthropic Family")!)).toBe(false); }); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); expect(screen.getByText("Advanced: Keyword/Semantic Matching")).toBeInTheDocument(); @@ -862,7 +845,7 @@ describe("AddAutoRouterTab", () => { await waitFor(() => { expect(isOptionDisabled(optionByLabel("Anthropic Family")!)).toBe(false); }); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); await user.type(screen.getByPlaceholderText(/smart_router/i), "wildcard-router"); await user.click(screen.getByRole("button", { name: /add auto router/i })); diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx index a071b92fcd5..64d1519f915 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx @@ -1,13 +1,14 @@ import React, { useEffect, useState } from "react"; import { useQuery } from "@tanstack/react-query"; import { useWatch } from "react-hook-form"; -import { Card, Select as AntdSelect } from "antd"; +import { Card } from "antd"; import { ChevronDown, ChevronRight, CircleHelp } from "lucide-react"; import { z } from "zod/v4"; import { FieldGroup } from "@/components/shared/form/field"; import { FormField } from "@/components/shared/form/FormField"; import { Button } from "@/components/ui/button"; import { Input } from "@/components/ui/input"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; import { useZodForm } from "@/lib/forms/useZodForm"; @@ -289,6 +290,14 @@ const AddAutoRouterTab: React.FC = ({ [presetAvailability], ); + const templateItems = React.useMemo( + () => [ + ...sortedPresetOptions.map(({ preset }) => ({ value: preset.key, label: preset.label })), + { value: "custom", label: "Custom Configuration" }, + ], + [sortedPresetOptions], + ); + const applyPrefill = (prefill: PresetPrefill) => { setComplexityRouterConfig(prefill.complexityRouterConfig); setCustomTechnicalKeywords(prefill.customTechnicalKeywords); @@ -486,49 +495,52 @@ const AddAutoRouterTab: React.FC = ({
    - handlePresetChange(presetKey ?? undefined)} > - {sortedPresetOptions.map(({ preset, availability: presetState }) => { - const disabledHint = presetDisabledHint(presetState); - const isDisabled = disabledHint !== null; - const hintClass = isPresetHintAlarming(presetState) - ? "text-red-500 dark:text-red-400" - : "text-muted-foreground"; - const matchedHint = - presetState.kind === "available" && presetState.viaDeployments ? "Matches your deployments" : null; + + + + + {sortedPresetOptions.map(({ preset, availability: presetState }) => { + const disabledHint = presetDisabledHint(presetState); + const hintClass = isPresetHintAlarming(presetState) + ? "text-red-500 dark:text-red-400" + : "text-muted-foreground"; + const matchedHint = + presetState.kind === "available" && presetState.viaDeployments + ? "Matches your deployments" + : null; - return ( - -
    -
    {preset.label}
    -
    {preset.description}
    - {disabledHint &&
    {disabledHint}
    } - {matchedHint && ( -
    {matchedHint}
    - )} -
    -
    - ); - })} - -
    -
    Custom Configuration
    -
    Define your auto router from scratch
    -
    -
    -
    + return ( + +
    +
    {preset.label}
    +
    {preset.description}
    + {disabledHint &&
    {disabledHint}
    } + {matchedHint && ( +
    {matchedHint}
    + )} +
    +
    + ); + })} + +
    +
    Custom Configuration
    +
    Define your auto router from scratch
    +
    +
    + + {modelsUnverifiable && (
    Could not load available models.{" "} diff --git a/ui/litellm-dashboard/src/components/add_model/advanced_settings.tsx b/ui/litellm-dashboard/src/components/add_model/advanced_settings.tsx index 2d35a772a63..5bc0dcd143d 100644 --- a/ui/litellm-dashboard/src/components/add_model/advanced_settings.tsx +++ b/ui/litellm-dashboard/src/components/add_model/advanced_settings.tsx @@ -1,14 +1,18 @@ import React from "react"; -import { Switch, Select, Tooltip, DatePicker } from "antd"; +import { MultiSelect } from "@/components/shared/MultiSelect"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { Switch } from "@/components/ui/switch"; +import { SimpleTooltip } from "@/components/ui/tooltip"; +import type { Dayjs } from "dayjs"; import { ChevronDown, Info } from "lucide-react"; import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible"; import { Input } from "@/components/ui/input"; -import { Row, Col, Typography } from "antd"; -import TextArea from "antd/es/input/TextArea"; +import { Textarea } from "@/components/ui/textarea"; import { Team } from "../key_team_helpers/key_list"; import { antdRules } from "../common_components/antdFormRules"; import { labelWithHint } from "@/components/shared/form/LabelWithHint"; import { MountedFormField } from "../common_components/MountedFormField"; +import { UtcDateTimeInput } from "@/components/shared/form/UtcDateTimeInput"; import CacheControlInjectionPoints, { CACHE_CONTROL_LABEL, CACHE_CONTROL_TOOLTIP, @@ -30,7 +34,6 @@ import { PTU_END_FIELD, } from "../../utils/ptuValidation"; import { usePtuCostAttributionEnabled } from "@/app/(dashboard)/hooks/uiSettings/usePtuCostAttributionEnabled"; -const { Link } = Typography; interface AdvancedSettingsProps { showAdvancedSettings: boolean; @@ -51,6 +54,11 @@ const USAGE_COST_FIELDS = [ const REVALIDATED_WHEN_PTU_COUNT_CHANGES = [PTU_RATE_FIELD, PTU_START_FIELD, ...USAGE_COST_FIELDS]; +const PRICING_MODEL_ITEMS = [ + { value: "per_token", label: "Per Million Tokens" }, + { value: "per_second", label: "Per Second" }, +] as const; + const validateNumber = (_: unknown, value: unknown) => { if (!value) { return Promise.resolve(); @@ -79,6 +87,14 @@ const AdvancedSettings: React.FC = ({ const [showCacheControl, setShowCacheControl] = React.useState(false); const ptuCostAttributionEnabled = usePtuCostAttributionEnabled(); + const handlePricingModelChange = + (onChange: (value: string) => void) => + (value: "per_token" | "per_second" | null): void => { + if (value === null) return; + onChange(value); + setPricingModel(value); + }; + return ( <> @@ -93,11 +109,10 @@ const AdvancedSettings: React.FC = ({ { + onCheckedChange={(checked) => { control.onChange(checked); setCustomPricing(checked); }} - className="bg-gray-600" /> )} @@ -107,7 +122,7 @@ const AdvancedSettings: React.FC = ({ label={ Attached Knowledge Bases (RAG){" "} - + = ({ > - + } className="mt-4" @@ -137,7 +152,7 @@ const AdvancedSettings: React.FC = ({ label={ Guardrails{" "} - + = ({ > - + } className="mt-4" help="Select existing guardrails. Go to 'Guardrails' tab to create new guardrails." > {(control) => ( - ({ value: tag.name, label: tag.name, - title: tag.description || tag.name, + description: tag.description || undefined, }))} + allowCustomValues /> )} @@ -249,11 +262,9 @@ const AdvancedSettings: React.FC = ({ className="mb-4" > {(control) => ( - @@ -273,11 +284,9 @@ const AdvancedSettings: React.FC = ({ className="mb-4" > {(control) => ( - @@ -291,19 +300,21 @@ const AdvancedSettings: React.FC = ({ {(control) => ( )} @@ -401,20 +412,20 @@ const AdvancedSettings: React.FC = ({ "Use in pass through routes", Allow using these credentials in pass through routes.{" "} - + Learn more - + , )} className="mb-4 mt-4" > {(control) => ( - + )} @@ -427,11 +438,10 @@ const AdvancedSettings: React.FC = ({ { + onCheckedChange={(checked) => { control.onChange(checked); setShowCacheControl(checked); }} - className="bg-gray-600" /> )} @@ -456,7 +466,7 @@ const AdvancedSettings: React.FC = ({ rules={{ validate: antdRules({ validator: formItemValidateJSON }) }} > {(control) => ( -