From 71d1bfb70a7f9ad2cac9e4e2090f029791fbd2e7 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Thu, 10 Sep 2026 15:55:03 -0700 Subject: [PATCH 01/16] feat(ui): link the User, Team and Created By cells on the Virtual Keys page The key detail page already walks out to the user, team and org behind a key, but the Virtual Keys table rendered those same values as dead text, so getting to a team meant copying its alias and searching the Teams page. User, Team and Created By now render through the shared IdentityCell with an href, the same hover-highlight-and-chevron affordance the Key column already uses. Sentinel ids do not get a link, since they have no detail page to open. Rather than repeat that check at every call site, teamDetailHref and userDetailHref now return undefined for "litellm-dashboard" and "default_user_id", the way modelGroupHref already does for model grants, and EntityLink falls back to plain text when it has no href, the way BadgeLink already does. Both sentinels move into src/utils/sentinels.ts instead of staying as string literals scattered across components. Claude-Session: https://claude.ai/code/session_01NfwfQhamRNnSqgXMUjf3h4 --- .../_components/AccessGroupsDetailsPage.tsx | 2 +- .../VirtualKeysPage/VirtualKeysTable.test.tsx | 71 +++++++++++++++++++ .../VirtualKeysPage/keyTableColumns.tsx | 45 ++++++------ .../DefaultProxyAdminTag.tsx | 5 +- .../common_components/LabeledField.tsx | 3 +- .../src/components/shared/EntityLink.test.tsx | 6 ++ .../src/components/shared/EntityLink.tsx | 14 +++- .../src/utils/entityLinks.test.ts | 24 ++++++- ui/litellm-dashboard/src/utils/entityLinks.ts | 7 +- ui/litellm-dashboard/src/utils/sentinels.ts | 3 + 10 files changed, 147 insertions(+), 33 deletions(-) create mode 100644 ui/litellm-dashboard/src/utils/sentinels.ts diff --git a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsDetailsPage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsDetailsPage.tsx index 1eeebe4ebba..118a8655d5a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsDetailsPage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsDetailsPage.tsx @@ -54,7 +54,7 @@ function ResourceBadge({ fallback, }: { resource: AccessGroupResource; - href: string; + href?: string; fallback: (id: string) => string; }) { const badge = ( diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx index 881b0b93ff9..57786d71c66 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx @@ -473,6 +473,77 @@ it("should display 'Default Proxy Admin' for user_id when value is 'default_user }); }); +describe("entity links out of the key rows", () => { + const keyRow = async () => (await screen.findByText("Test Key Alias")).closest("tr") as HTMLElement; + + const enableCreatedByColumn = async (user: ReturnType) => { + await user.click(screen.getByRole("button", { name: "Columns" })); + await user.click(await screen.findByText("Created By")); + await user.keyboard("{Escape}"); + }; + + it("points the User and Team cells at their detail pages", async () => { + renderWithProviders(); + + const row = await keyRow(); + expect(within(row).getByRole("link", { name: "user@example.com" })).toHaveAttribute( + "href", + "/ui/users?user=user-1", + ); + expect(within(row).getByRole("link", { name: "Test Team" })).toHaveAttribute("href", "/ui/teams?team=team-1"); + }); + + it("points the Created By cell at the creator's detail page", async () => { + mockUseKeys.mockReturnValue( + keysResult([ + { + ...mockKey, + created_by: "creator-1", + created_by_user: { user_id: "creator-1", user_email: "creator@example.com", user_alias: "The Creator" }, + }, + ]), + ); + const user = userEvent.setup(); + renderWithProviders(); + await enableCreatedByColumn(user); + + const row = await keyRow(); + expect(within(row).getByRole("link", { name: "The Creator" })).toHaveAttribute("href", "/ui/users?user=creator-1"); + }); + + it("leaves the default_user_id placeholder unlinked even once it resolves to a named user", async () => { + const placeholder = { user_id: "default_user_id", user_email: "admin@example.com", user_alias: "Proxy Admin" }; + mockUseKeys.mockReturnValue( + keysResult([ + { + ...mockKey, + user_id: placeholder.user_id, + user_email: placeholder.user_email, + user: placeholder, + created_by: placeholder.user_id, + created_by_user: placeholder, + }, + ]), + ); + const user = userEvent.setup(); + renderWithProviders(); + await enableCreatedByColumn(user); + + const row = await keyRow(); + expect(within(row).getAllByText("Proxy Admin")).toHaveLength(2); + expect(within(row).queryByRole("link", { name: "Proxy Admin" })).not.toBeInTheDocument(); + }); + + it("leaves the litellm-dashboard session team unlinked", async () => { + mockUseKeys.mockReturnValue(keysResult([{ ...mockKey, team_id: "litellm-dashboard" }])); + renderWithProviders(); + + const row = await keyRow(); + expect(within(row).getByText("litellm-dashboard")).toBeInTheDocument(); + expect(within(row).queryByRole("link", { name: "litellm-dashboard" })).not.toBeInTheDocument(); + }); +}); + it("should render table without crashing when models is null", async () => { mockUseKeys.mockReturnValue(keysResult([{ ...mockKey, models: null as unknown as string[] }])); diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx index 0865fe76519..c87bc90fb47 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx @@ -16,6 +16,8 @@ import { StatusBadge, type StatusTone, } from "@/components/shared/table_cells"; +import { teamDetailHref, userDetailHref } from "@/utils/entityLinks"; +import { DEFAULT_PROXY_ADMIN_USER_ID } from "@/utils/sentinels"; import DefaultProxyAdminTag from "../common_components/DefaultProxyAdminTag"; import { KeyResponse, Team } from "../key_team_helpers/key_list"; @@ -27,6 +29,8 @@ interface KeyStatus { tooltip?: string; } +const ENTITY_CELL_TITLE_CLASSES = "font-mono text-xs font-normal"; + const SPEND_BUDGET_SORT_FIELDS: DataTableSortField[] = [ { id: "spend", label: "Spend" }, { id: "max_budget", label: "Budget" }, @@ -74,7 +78,7 @@ const UserPopoverCell = ({ width: number; }) => { const displayValue = userAlias || userEmail || userId; - const isDefaultAdmin = userId === "default_user_id"; + const isDefaultAdmin = userId === DEFAULT_PROXY_ADMIN_USER_ID; const popoverContent = (
@@ -95,28 +99,21 @@ const UserPopoverCell = ({
); - if (isDefaultAdmin && !userAlias && !userEmail) { - return ( - - }> - - - {popoverContent} - + const trigger = + isDefaultAdmin && !userAlias && !userEmail ? ( + + ) : ( + ); - } return ( - - } - > - {displayValue || "-"} + }> + {trigger} {popoverContent} @@ -201,12 +198,12 @@ export const getKeyTableColumns = ({ const teamId = info.getValue() as string | null; if (!teamId) return "-"; const team = allTeams.find((t) => t.team_id === teamId); - const displayValue = team?.team_alias || teamId; - const width = info.cell.column.getSize(); return ( - - {displayValue} - + ); }, }, diff --git a/ui/litellm-dashboard/src/components/common_components/DefaultProxyAdminTag.tsx b/ui/litellm-dashboard/src/components/common_components/DefaultProxyAdminTag.tsx index 9ec24bb929b..308f722a8f2 100644 --- a/ui/litellm-dashboard/src/components/common_components/DefaultProxyAdminTag.tsx +++ b/ui/litellm-dashboard/src/components/common_components/DefaultProxyAdminTag.tsx @@ -1,13 +1,12 @@ import { Badge } from "@/components/ui/badge"; - -const DEFAULT_USER_ID = "default_user_id"; +import { DEFAULT_PROXY_ADMIN_USER_ID } from "@/utils/sentinels"; interface DefaultProxyAdminTagProps { userId: string | null | undefined; } export default function DefaultProxyAdminTag({ userId }: DefaultProxyAdminTagProps) { - if (userId === DEFAULT_USER_ID) { + if (userId === DEFAULT_PROXY_ADMIN_USER_ID) { return Default Proxy Admin; } diff --git a/ui/litellm-dashboard/src/components/common_components/LabeledField.tsx b/ui/litellm-dashboard/src/components/common_components/LabeledField.tsx index 9f45b05f306..99046cfa84d 100644 --- a/ui/litellm-dashboard/src/components/common_components/LabeledField.tsx +++ b/ui/litellm-dashboard/src/components/common_components/LabeledField.tsx @@ -2,6 +2,7 @@ import React from "react"; import CopyButton from "@/components/shared/CopyButton"; import { EntityLink } from "@/components/shared/EntityLink"; import { cx } from "@/lib/cva.config"; +import { DEFAULT_PROXY_ADMIN_USER_ID } from "@/utils/sentinels"; import DefaultProxyAdminTag from "./DefaultProxyAdminTag"; interface LabeledFieldProps { @@ -24,7 +25,7 @@ export default function LabeledField({ defaultUserIdCheck = false, }: LabeledFieldProps) { const isEmpty = !value; - const isDefaultUser = defaultUserIdCheck && value === "default_user_id"; + const isDefaultUser = defaultUserIdCheck && value === DEFAULT_PROXY_ADMIN_USER_ID; const displayValue = isEmpty ? "-" : value; const isCopyable = copyable && !isEmpty && !isDefaultUser; const isLink = href != null && !isEmpty && !isDefaultUser; diff --git a/ui/litellm-dashboard/src/components/shared/EntityLink.test.tsx b/ui/litellm-dashboard/src/components/shared/EntityLink.test.tsx index a02f1699c30..3d6fc9a435b 100644 --- a/ui/litellm-dashboard/src/components/shared/EntityLink.test.tsx +++ b/ui/litellm-dashboard/src/components/shared/EntityLink.test.tsx @@ -25,6 +25,12 @@ describe("EntityLink", () => { expect(push).toHaveBeenCalledWith("/ui/users?user=u1"); }); + it("renders the label as plain text when there is no href to point at", () => { + render(default_user_id); + expect(screen.queryByRole("link")).not.toBeInTheDocument(); + expect(screen.getByText("default_user_id")).toBeInTheDocument(); + }); + it("leaves modified clicks to the browser so new-tab shortcuts keep working", async () => { const user = userEvent.setup(); render(alice); diff --git a/ui/litellm-dashboard/src/components/shared/EntityLink.tsx b/ui/litellm-dashboard/src/components/shared/EntityLink.tsx index 4054b929943..4885835a42c 100644 --- a/ui/litellm-dashboard/src/components/shared/EntityLink.tsx +++ b/ui/litellm-dashboard/src/components/shared/EntityLink.tsx @@ -19,12 +19,24 @@ export function useEntityLinkClick(href: string): (e: React.MouseEvent) => void } interface EntityLinkProps { - href: string; + href?: string; className?: string; children: React.ReactNode; } export function EntityLink({ href, className, children }: EntityLinkProps) { + if (!href) { + return {children}; + } + + return ( + + {children} + + ); +} + +function LinkedEntity({ href, className, children }: EntityLinkProps & { href: string }) { const handleClick = useEntityLinkClick(href); return ( diff --git a/ui/litellm-dashboard/src/utils/entityLinks.test.ts b/ui/litellm-dashboard/src/utils/entityLinks.test.ts index 47161a903ed..231413b213a 100644 --- a/ui/litellm-dashboard/src/utils/entityLinks.test.ts +++ b/ui/litellm-dashboard/src/utils/entityLinks.test.ts @@ -2,7 +2,29 @@ import { describe, expect, it, vi } from "vitest"; vi.mock("@/components/networking", () => ({ serverRootPath: "" })); -import { modelGroupHref } from "./entityLinks"; +import { modelGroupHref, teamDetailHref, userDetailHref } from "./entityLinks"; + +describe("userDetailHref", () => { + it("targets the users page filtered to the encoded user id", () => { + expect(userDetailHref("user-1")).toMatch(/\/users\?user=user-1$/); + expect(userDetailHref("a b/c")).toMatch(/\?user=a%20b%2Fc$/); + }); + + it("returns no href for the proxy admin placeholder, which has no user page", () => { + expect(userDetailHref("default_user_id")).toBeUndefined(); + }); +}); + +describe("teamDetailHref", () => { + it("targets the teams page filtered to the encoded team id", () => { + expect(teamDetailHref("team-1")).toMatch(/\/teams\?team=team-1$/); + expect(teamDetailHref("a b/c")).toMatch(/\?team=a%20b%2Fc$/); + }); + + it("returns no href for the Admin UI session team, which has no team page", () => { + expect(teamDetailHref("litellm-dashboard")).toBeUndefined(); + }); +}); describe("modelGroupHref", () => { it("targets the models page filtered to the encoded model group", () => { diff --git a/ui/litellm-dashboard/src/utils/entityLinks.ts b/ui/litellm-dashboard/src/utils/entityLinks.ts index b8c70bdda48..33f2aa34976 100644 --- a/ui/litellm-dashboard/src/utils/entityLinks.ts +++ b/ui/litellm-dashboard/src/utils/entityLinks.ts @@ -1,3 +1,4 @@ +import { DEFAULT_PROXY_ADMIN_USER_ID, UI_TEAM_ID } from "@/utils/sentinels"; import { uiHref } from "@/utils/uiHref"; const MODEL_GRANT_SENTINELS: ReadonlySet = new Set([ @@ -6,7 +7,8 @@ const MODEL_GRANT_SENTINELS: ReadonlySet = new Set([ "no-default-models", ]); -export function teamDetailHref(teamId: string): string { +export function teamDetailHref(teamId: string): string | undefined { + if (teamId === UI_TEAM_ID) return undefined; return `${uiHref("teams")}?team=${encodeURIComponent(teamId)}`; } @@ -14,7 +16,8 @@ export function keyDetailHref(keyToken: string): string { return `${uiHref("api-keys")}?key=${encodeURIComponent(keyToken)}`; } -export function userDetailHref(userId: string): string { +export function userDetailHref(userId: string): string | undefined { + if (userId === DEFAULT_PROXY_ADMIN_USER_ID) return undefined; return `${uiHref("users")}?user=${encodeURIComponent(userId)}`; } diff --git a/ui/litellm-dashboard/src/utils/sentinels.ts b/ui/litellm-dashboard/src/utils/sentinels.ts new file mode 100644 index 00000000000..da6d09dfdd7 --- /dev/null +++ b/ui/litellm-dashboard/src/utils/sentinels.ts @@ -0,0 +1,3 @@ +export const DEFAULT_PROXY_ADMIN_USER_ID = "default_user_id"; + +export const UI_TEAM_ID = "litellm-dashboard"; From 311441ced19b1963f8ad22b758a8ae6171105a74 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Thu, 10 Sep 2026 16:08:57 -0700 Subject: [PATCH 02/16] bump: litellm-proxy-extras 0.4.95 -> 0.4.96 --- litellm-proxy-extras/pyproject.toml | 4 ++-- pyproject.toml | 2 +- uv.lock | 4 ++-- 3 files changed, 5 insertions(+), 5 deletions(-) diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml index 91b4e4a7ba1..7d4c78088f1 100644 --- a/litellm-proxy-extras/pyproject.toml +++ b/litellm-proxy-extras/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm-proxy-extras" -version = "0.4.95" +version = "0.4.96" description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package." readme = "README.md" requires-python = ">=3.9" @@ -26,7 +26,7 @@ required-version = ">=0.10.9" module-root = "" [tool.commitizen] -version = "0.4.95" +version = "0.4.96" version_files = [ "pyproject.toml:^version", "../pyproject.toml:litellm-proxy-extras==", diff --git a/pyproject.toml b/pyproject.toml index 04f2f3fd1dd..d33d693f794 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -67,7 +67,7 @@ proxy = [ "azure-identity>=1.25.2,<2.0", "azure-storage-blob>=12.28.0,<13.0", "mcp>=1.28.1,<2.0", - "litellm-proxy-extras==0.4.95", + "litellm-proxy-extras==0.4.96", "litellm-enterprise==0.1.66", "RestrictedPython>=8.5,<9.0", "rich>=13.9.4,<14.0", diff --git a/uv.lock b/uv.lock index 0fe787645a2..cedc505acdf 100644 --- a/uv.lock +++ b/uv.lock @@ -10,7 +10,7 @@ resolution-markers = [ ] [options] -exclude-newer = "2026-09-06T00:40:30.433549Z" +exclude-newer = "2026-09-07T23:09:03.362777Z" exclude-newer-span = "P3D" [manifest] @@ -4772,7 +4772,7 @@ source = { editable = "enterprise" } [[package]] name = "litellm-proxy-extras" -version = "0.4.95" +version = "0.4.96" source = { editable = "litellm-proxy-extras" } [[package]] From 760043b5337a83f0356cd717e5b555229980deec Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Thu, 10 Sep 2026 16:57:53 -0700 Subject: [PATCH 03/16] fix(reset_budget_job): reset end users by budget link, not by user id The cascade zeroed end-user spend with a single update_many whose where clause enumerated every dependent user id. Prisma compiles that IN-list into one prepared statement carrying one bind variable per customer, and PostgreSQL caps a statement at 32,767 of them. Once a shared budget had more dependents than that the statement could not be parsed at all, so the atomic cascade rolled back, budget_reset_at never advanced, and the tier stayed due on every later tick forever. Customers sitting at their cap were blocked indefinitely with only a recurring log line to show for it. End users now match on budget_id like every other gated table, plus a NULL-budget_id branch for the implicitly created rows that carry no link and ride the default tier. The statement's bind count now tracks the number of expiring tiers rather than the customer population, so a reset costs the same whether a budget has ten dependents or a million. Fixes #40564 Claude-Session: https://claude.ai/code/session_01Hn5E8Jz1LjGLFyiYxBRcBW --- .../proxy/common_utils/reset_budget_job.py | 46 ++++++------ .../test_proxy_budget_reset.py | 6 +- .../common_utils/test_reset_budget_job.py | 75 ++++++++++++++++--- 3 files changed, 88 insertions(+), 39 deletions(-) diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index f2648c8466e..35e74418628 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -224,33 +224,31 @@ def _queue_budget_linked_resets( def _queue_enduser_resets(writes: LinkedSpendResetWrites, cascade: "_BudgetCascade") -> None: - """End users are matched by id rather than budget link: rows with no - budget_id ride the default budget tier (litellm.max_end_user_budget_id). - Zero-before-decrement ordering matters here too (see - _queue_budget_linked_resets).""" - if not cascade.rollover_caps: - if cascade.endusers: - writes.queue_spend_zero( - where={"user_id": {"in": [row.user_id for row in cascade.endusers]}} - ) # mutable-ok: prisma where filter must be a dict + """End users reset on the budget link like every other gated table, plus a + NULL-budget_id branch: rows created implicitly persist no link and ride the + default tier (litellm.max_end_user_budget_id). + + Matching on the link rather than enumerating user ids keeps a statement's + bind count proportional to the expiring tiers instead of the customer + population, which past ~32,700 dependents exceeds PostgreSQL's per-statement + bind ceiling and wedges the cascade permanently (#40564). + """ + _queue_budget_linked_resets(writes, cascade, extra=_SPENT_ROWS_WHERE) + default_budget_id: Final = litellm.max_end_user_budget_id + if default_budget_id is None or default_budget_id not in cascade.budget_ids: return - tiered: Final = tuple((row.budget_id or litellm.max_end_user_budget_id, row.user_id) for row in cascade.endusers) - for budget_id, cap in cascade.rollover_caps.items(): - if not ( - user_ids := [uid for bid, uid in tiered if bid == budget_id] - ): # mutable-ok: prisma "in" filter takes a list - continue + cap: Final = cascade.rollover_caps.get(default_budget_id) + if cap is None: writes.queue_spend_zero( - where={"user_id": {"in": user_ids}, "spend": {"lte": cap}} + where={"budget_id": None, **_SPENT_ROWS_WHERE} ) # mutable-ok: prisma where filter must be a dict - writes.queue_spend_decrement( - where={"user_id": {"in": user_ids}, "spend": {"gt": cap}}, amount=cap - ) # mutable-ok: prisma where filter must be a dict - plain: Final = [ - uid for bid, uid in tiered if bid is None or bid not in cascade.rollover_caps - ] # mutable-ok: prisma "in" filter takes a list - if plain: - writes.queue_spend_zero(where={"user_id": {"in": plain}}) # mutable-ok: prisma where filter must be a dict + return + writes.queue_spend_zero( + where={"budget_id": None, "spend": {"gt": 0, "lte": cap}} + ) # mutable-ok: prisma where filter must be a dict + writes.queue_spend_decrement( + where={"budget_id": None, "spend": {"gt": cap}}, amount=cap + ) # mutable-ok: prisma where filter must be a dict @dataclass(frozen=True, slots=True) diff --git a/tests/litellm_utils_tests/test_proxy_budget_reset.py b/tests/litellm_utils_tests/test_proxy_budget_reset.py index 9f6e1f4c3f7..4103536950d 100644 --- a/tests/litellm_utils_tests/test_proxy_budget_reset.py +++ b/tests/litellm_utils_tests/test_proxy_budget_reset.py @@ -401,7 +401,7 @@ async def test_reset_budget_endusers_are_zeroed_with_the_budget_window_advance() enduser_writes = [c for c in batch_calls if c["table"] == "enduser"] assert len(enduser_writes) == 1 - assert enduser_writes[0]["where"]["user_id"]["in"] == [f"user{i}" for i in range(1, 7)] + assert enduser_writes[0]["where"] == {"budget_id": {"in": ["budget1"]}, "spend": {"gt": 0}} assert enduser_writes[0]["data"] == {"spend": 0} budget_writes = [c for c in batch_calls if c["table"] == "budget"] @@ -602,7 +602,7 @@ async def test_reset_budget_continues_other_categories_on_failure(): assert len([c for c in batch_calls if c["table"] == "team_membership"]) == 1 enduser_writes = [c for c in batch_calls if c["table"] == "enduser"] assert len(enduser_writes) == 1 - assert enduser_writes[0]["where"] == {"user_id": {"in": ["user1"]}} + assert enduser_writes[0]["where"] == {"budget_id": {"in": ["budget1"]}, "spend": {"gt": 0}} assert enduser_writes[0]["data"] == {"spend": 0} # Check the new batch write path: 2 keys + 1 user (user1 failed) + 2 teams. @@ -1031,7 +1031,7 @@ async def test_service_logger_endusers_success(): enduser_writes = [c for c in batch_calls if c["table"] == "enduser"] assert len(enduser_writes) == 1 - assert enduser_writes[0]["where"] == {"user_id": {"in": ["user1", "user2"]}} + assert enduser_writes[0]["where"] == {"budget_id": {"in": ["budget1"]}, "spend": {"gt": 0}} proxy_logging_obj.service_logging_obj.async_service_success_hook.assert_called_once() ( diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index 56c0efb41d2..560953f0b51 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -403,7 +403,7 @@ def test_reset_budget_for_enduser(reset_budget_job, mock_prisma_client): { "table": "enduser", "op": "update_many", - "where": {"user_id": {"in": ["test-enduser-1"]}}, + "where": {"budget_id": {"in": ["test-budget-1"]}, "spend": {"gt": 0}}, "data": {"spend": 0}, } ] @@ -504,7 +504,7 @@ def test_reset_budget_all(reset_budget_job, mock_prisma_client): { "table": "enduser", "op": "update_many", - "where": {"user_id": {"in": ["test-enduser-1"]}}, + "where": {"budget_id": {"in": ["test-budget-1"]}, "spend": {"gt": 0}}, "data": {"spend": 0}, } ] @@ -524,6 +524,7 @@ _LINKED_TABLE_CASES = [ ("org", {"budget_id": {"in": ["7d-budget-tier"]}, "spend": {"gt": 0}}), ("tag", {"budget_id": {"in": ["7d-budget-tier"]}, "spend": {"gt": 0}}), ("model_access_group", {"budget_id": {"in": ["7d-budget-tier"]}, "spend": {"gt": 0}}), + ("enduser", {"budget_id": {"in": ["7d-budget-tier"]}, "spend": {"gt": 0}}), ] @@ -553,6 +554,48 @@ def test_budget_table_reset_zeroes_spend_on_every_linked_table( assert writes[0]["data"] == {"spend": 0} +_POSTGRES_MAX_BIND_VARIABLES: Final = 32767 + + +def _bind_count(where: Dict[str, Any]) -> int: + """Bind variables one prisma where-clause compiles to: each scalar is one + placeholder and an ``in`` list contributes one per element.""" + return sum(len(value["in"]) if isinstance(value, dict) and "in" in value else 1 for value in where.values()) + + +@pytest.mark.parametrize("population", [3, 40_000], ids=["small", "over-pg-bind-ceiling"]) +def test_enduser_reset_bind_count_does_not_scale_with_population(reset_budget_job, mock_prisma_client, population): + """Regression for #40564. + + Enumerating every dependent user id put one bind variable per customer into + a single prepared statement. Past PostgreSQL's ceiling the statement could + not be parsed at all, so the whole atomic cascade rolled back, + budget_reset_at never advanced, and the tier stayed due on every later tick + forever. Matching on the budget link keeps the statement the same size no + matter how many customers share a tier. + """ + budget = _budget_row(budget_id="shared-tier", budget_duration="1d") + mock_prisma_client.data["budget"] = [budget] + mock_prisma_client.data["enduser"] = [ + types.SimpleNamespace( + spend=1.0, + litellm_budget_table=budget, + user_id=f"cust-{index:08d}", + budget_id="shared-tier", + ) + for index in range(population) + ] + + asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) + + writes = _batch_writes(mock_prisma_client, "enduser") + assert [_bind_count(write["where"]) for write in writes] == [2], ( + f"the cascade must not enumerate {population} user ids: past " + f"{_POSTGRES_MAX_BIND_VARIABLES} binds PostgreSQL refuses the statement, got {writes[:1]}" + ) + assert _batch_writes(mock_prisma_client, "budget")[0]["data"]["budget_reset_at"] is not None + + def test_budget_table_reset_writes_nothing_when_no_budget_is_due(reset_budget_job, mock_prisma_client): """Nothing due means no transaction is opened at all.""" asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) @@ -720,14 +763,22 @@ def test_reset_budget_resets_endusers_with_null_budget_id(reset_budget_job, mock asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) - # Both end users are zeroed by the same committed statement. - enduser_writes = _batch_writes(mock_prisma_client, "enduser") - assert len(enduser_writes) == 1, f"Expected a single enduser write, got {enduser_writes}" - assert set(enduser_writes[0]["where"]["user_id"]["in"]) == { - "enduser-explicit", - "enduser-implicit", - } - assert enduser_writes[0]["data"] == {"spend": 0} + # Both end users are zeroed: the linked rows on the tier's budget_id, the + # implicit ones on the NULL branch that stands in for the default tier. + assert _batch_writes(mock_prisma_client, "enduser") == [ + { + "table": "enduser", + "op": "update_many", + "where": {"budget_id": {"in": [default_budget_id]}, "spend": {"gt": 0}}, + "data": {"spend": 0}, + }, + { + "table": "enduser", + "op": "update_many", + "where": {"budget_id": None, "spend": {"gt": 0}}, + "data": {"spend": 0}, + }, + ] # Verify find_many was called to fetch NULL-budget-id end users find_many_calls = mock_prisma_client.db.litellm_endusertable.find_many_calls @@ -3043,13 +3094,13 @@ def test_budget_cascade_carries_enduser_overage_when_rollover_enabled( assert { "table": "enduser", "op": "update_many", - "where": {"user_id": {"in": ["enduser-roll"]}, "spend": {"gt": 10.0}}, + "where": {"budget_id": "budget-roll", "spend": {"gt": 10.0}}, "data": {"spend": {"decrement": 10.0}}, } in enduser_writes assert { "table": "enduser", "op": "update_many", - "where": {"user_id": {"in": ["enduser-roll"]}, "spend": {"lte": 10.0}}, + "where": {"budget_id": "budget-roll", "spend": {"gt": 0, "lte": 10.0}}, "data": {"spend": 0}, } in enduser_writes From ae018825357ddafcf0da15d903d10df0f2581c59 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 10 Sep 2026 17:14:13 -0700 Subject: [PATCH 04/16] feat(proxy): offload spend tracking to a pod-local collector sidecar (#40545) * feat(proxy): offload spend tracking to a pod-local spend worker sidecar py-spy on the gateway showed the post-response _PROXY_track_cost_callback, spend-log and DBSpendUpdateWriter work running on the inference workers' event loop, so a DB or Redis stall backed up the request path. When LITELLM_SPEND_WORKER_ENABLED=true, _ProxyDBLogger serializes one compact typed SpendEvent per success and hands it to a SpendEventProducer that ships it over a unix socket (default) or loopback-only TCP to a sidecar started as `python -m gateway.spend_worker`. The sidecar runs the unchanged _ProxyDBLogger pipeline against the pod's PgBouncer (pooled_database_url). When the sidecar is unreachable, the buffer is full, or the gateway shuts down with events still queued or in flight, the producer applies LITELLM_SPEND_WORKER_ON_UNAVAILABLE (fallback in-process, or drop). The sidecar half-closes producers on SIGTERM and drains, the producer treats EOF as unavailable, and the gateway flushes buffered spend counters on shutdown. The sidecar honors LITELLM_LOG so its writes are visible in its own process log. Helm: both charts gain an opt-in spend-worker sidecar container sharing an emptyDir socket dir, and the componentized chart's HPA uses a ContainerResource CPU metric scoped to the gateway container so sidecar CPU does not drive inference scaling. * feat(terraform): opt-in spend-worker sidecar for the AWS and GCP gateway stacks Adds spend_worker_* inputs to both modules. On ECS Fargate the sidecar is a second, non-essential container in the gateway task; on Cloud Run it is a second container in the gateway service. Both listen on loopback TCP, share the gateway's DB/Redis/secret env, and set LITELLM_JOB_ROLE=spend_worker. Disabled by default. Plan-only tests cover both, and the terraform CI workflow now runs the gcp module too Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): retrieve a completed batch in the in-process spend path test The base now defers cost tracking for batches that are still in flight, so an in_progress batch never reaches update_database Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(proxy): rename the spend worker sidecar to collector Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): run the collector from the installed litellm package and finish in-flight fallbacks on shutdown The sidecar command becomes python -m litellm.proxy.collector so the classic image, whose runtime stage copies only the installed package, can run it. The module now assembles DATABASE_URL and the pod-local pgbouncer URL itself, replacing gateway/collector.py The componentized collector sidecar inherits gateway.volumeMounts so custom CA mounts reach it. SpendEventProducer shields an in-progress fallback from the writer task cancellation so close() no longer loses an event already handed to the in-process pipeline Helpers used across modules (address_argument, should_store_prompts_and_responses_in_spend_logs, flush_spend_counters_on_shutdown) become public so the change adds no reportPrivateUsage errors Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * ci(terraform): drop the gcp job duplicated by the aws/gcp matrix Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(collector): keep metrics env off the classic sidecar and reject shared loopback ports The classic chart no longer hands PROMETHEUS_METRICS_PORT and the billing metrics env to the collector container, and gives it the same /.npm scratch mount as the proxy on a read-only root. AWS and GCP now refuse a plan where the spend collector and the metrics sidecar bind the same loopback port. A regression test drives a sidecar crash mid-stream on asyncio and uvloop and checks no event is billed by both the sidecar and the in-process fallback; the producer docstring spells out why a failed drain() cannot double count Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(proxy): format pooled_database_url after the pgbouncer rebase Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): keep the cache-hit preset key and survive dead producers on collector drain Cache hits updated the logging object after the early return, so the offloaded spend event carried preset_cache_key=None and the collector re-hashed reconstructed kwargs. Also guard write_eof() against producer transports uvloop already closed so one dead connection cannot abort the drain Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(terraform): keep the gcp collector port off the metrics sidecar health port Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): collector connects to Postgres directly under IAM or Entra token auth Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): mark the collector's DATABASE_URL as pooled when it uses the pod's pgbouncer Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .github/workflows/test-terraform-modules.yml | 40 +- helm/litellm-helm/templates/_helpers.tpl | 160 +++++++ helm/litellm-helm/templates/deployment.yaml | 180 +++----- helm/litellm-helm/templates/hpa.yaml | 10 + helm/litellm-helm/tests/collector_tests.yaml | 272 ++++++++++++ helm/litellm-helm/values.yaml | 42 ++ helm/litellm/templates/_helpers.tpl | 31 ++ .../litellm/templates/gateway/deployment.yaml | 56 ++- helm/litellm/templates/gateway/hpa.yaml | 10 + helm/litellm/tests/collector_tests.yaml | 204 +++++++++ helm/litellm/values.yaml | 36 ++ litellm/caching/caching_handler.py | 4 +- litellm/proxy/collector.py | 220 +++++++++ litellm/proxy/db/pgbouncer.py | 16 + .../proxy/hooks/proxy_track_cost_callback.py | 54 ++- litellm/proxy/proxy_server.py | 39 +- litellm/proxy/spend_tracking/spend_event.py | 418 ++++++++++++++++++ .../spend_tracking/spend_event_producer.py | 338 ++++++++++++++ .../spend_tracking/spend_tracking_utils.py | 24 +- litellm/utils.py | 3 + terraform/litellm/README.md | 1 + terraform/litellm/aws/README.md | 33 ++ terraform/litellm/aws/ecs.tf | 82 +++- .../litellm/aws/tests/collector.tftest.hcl | 144 ++++++ terraform/litellm/aws/variables.tf | 70 +++ terraform/litellm/gcp/README.md | 37 ++ terraform/litellm/gcp/cloudrun.tf | 85 +++- .../litellm/gcp/tests/collector.tftest.hcl | 165 +++++++ terraform/litellm/gcp/variables.tf | 70 +++ tests/test_gateway/test_launch.py | 6 +- .../caching/test_caching_handler.py | 35 ++ .../hooks/test_proxy_track_cost_callback.py | 255 ++++++++++- .../proxy/proxy_server/test_lifecycle.py | 32 ++ .../proxy/spend_tracking/test_spend_event.py | 213 +++++++++ .../test_spend_event_producer.py | 359 +++++++++++++++ .../test_spend_tracking_utils.py | 122 +++-- tests/test_litellm/proxy/test_collector.py | 229 ++++++++++ 37 files changed, 3865 insertions(+), 230 deletions(-) create mode 100644 helm/litellm-helm/tests/collector_tests.yaml create mode 100644 helm/litellm/tests/collector_tests.yaml create mode 100644 litellm/proxy/collector.py create mode 100644 litellm/proxy/spend_tracking/spend_event.py create mode 100644 litellm/proxy/spend_tracking/spend_event_producer.py create mode 100644 terraform/litellm/aws/tests/collector.tftest.hcl create mode 100644 terraform/litellm/gcp/tests/collector.tftest.hcl create mode 100644 tests/test_litellm/proxy/spend_tracking/test_spend_event.py create mode 100644 tests/test_litellm/proxy/spend_tracking/test_spend_event_producer.py create mode 100644 tests/test_litellm/proxy/test_collector.py diff --git a/.github/workflows/test-terraform-modules.yml b/.github/workflows/test-terraform-modules.yml index 52006d9b578..e6896604b7f 100644 --- a/.github/workflows/test-terraform-modules.yml +++ b/.github/workflows/test-terraform-modules.yml @@ -25,13 +25,17 @@ concurrency: cancel-in-progress: true jobs: - aws-module: - name: fmt, validate, test (aws) + module: + name: fmt, validate, test (${{ matrix.module }}) runs-on: ubuntu-latest timeout-minutes: 15 + strategy: + fail-fast: false + matrix: + module: [aws, gcp] defaults: run: - working-directory: terraform/litellm/aws + working-directory: terraform/litellm/${{ matrix.module }} steps: - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 with: @@ -51,35 +55,7 @@ jobs: - name: validate run: terraform validate - # Plan-only, mock_provider-backed: no AWS credentials, no API calls. + # Plan-only, mock_provider-backed: no cloud credentials, no API calls. - name: test run: terraform test - gcp-module: - name: fmt, validate, test (gcp) - runs-on: ubuntu-latest - timeout-minutes: 15 - defaults: - run: - working-directory: terraform/litellm/gcp - steps: - - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 - with: - persist-credentials: false - - - uses: hashicorp/setup-terraform@b9cd54a3c349d3f38e8881555d616ced269862dd # v3.1.2 - with: - terraform_version: 1.13.3 - terraform_wrapper: false - - - name: fmt - run: terraform fmt -recursive -check -diff - - - name: init - run: terraform init -backend=false -input=false - - - name: validate - run: terraform validate - - - name: test - run: terraform test diff --git a/helm/litellm-helm/templates/_helpers.tpl b/helm/litellm-helm/templates/_helpers.tpl index 8f2acb20fce..9630633912e 100644 --- a/helm/litellm-helm/templates/_helpers.tpl +++ b/helm/litellm-helm/templates/_helpers.tpl @@ -161,3 +161,163 @@ taken before the change, which by that point no longer exists. {{- fail (printf "postgresql.image.tag must be pinned to an explicit version when db.deployStandalone is true (got %q). An unpinned tag can start a different PostgreSQL major against the existing data directory, which makes the database unreadable and is not recoverable in place. Crossing a major version requires a dump and restore." $tag) -}} {{- end -}} {{- end -}} + +{{/* +Environment shared by the proxy container and the opt-in collector sidecar: +database, pgbouncer, master key, redis, user envVars. Both containers must see +the same DATABASE_URL and REDIS_* so the sidecar reaches the pod's pgbouncer +and the same spend transaction buffer. +*/}} +{{- define "litellm.proxyEnv" -}} +- name: HOST + value: "{{ .Values.listen | default "0.0.0.0" }}" +- name: PORT + value: {{ .Values.service.port | quote}} +{{- if .Values.db.deployStandalone }} +- name: DATABASE_USERNAME + valueFrom: + secretKeyRef: + name: {{ include "litellm.fullname" . }}-dbcredentials + key: username +- name: DATABASE_PASSWORD + valueFrom: + secretKeyRef: + name: {{ include "litellm.fullname" . }}-dbcredentials + key: password +- name: DATABASE_HOST + value: {{ .Release.Name }}-postgresql +- name: DATABASE_NAME + value: litellm +{{- else if .Values.db.useExisting }} +- name: DATABASE_USERNAME + valueFrom: + secretKeyRef: + name: {{ .Values.db.secret.name }} + key: {{ .Values.db.secret.usernameKey }} +- name: DATABASE_PASSWORD + valueFrom: + secretKeyRef: + name: {{ .Values.db.secret.name }} + key: {{ .Values.db.secret.passwordKey }} +- name: DATABASE_HOST + {{- if .Values.db.secret.endpointKey }} + valueFrom: + secretKeyRef: + name: {{ .Values.db.secret.name }} + key: {{ .Values.db.secret.endpointKey }} + {{- else }} + value: {{ .Values.db.endpoint }} + {{- end }} +- name: DATABASE_NAME + value: {{ .Values.db.database }} +- name: DATABASE_URL + value: {{ .Values.db.url | quote }} +{{- end }} +{{- if and .Values.db.useExisting .Values.db.readReplicaUrl .Values.db.secret.readReplicaEndpointKey (not .Values.db.secret.readReplicaUrlKey) }} +- name: DATABASE_READER_HOST + valueFrom: + secretKeyRef: + name: {{ .Values.db.secret.name }} + key: {{ .Values.db.secret.readReplicaEndpointKey }} +{{- end }} +{{- if and .Values.db.useExisting .Values.db.secret.readReplicaUrlKey }} +- name: DATABASE_URL_READ_REPLICA + valueFrom: + secretKeyRef: + name: {{ .Values.db.secret.name }} + key: {{ .Values.db.secret.readReplicaUrlKey }} +{{- else if .Values.db.readReplicaUrl }} +- name: DATABASE_URL_READ_REPLICA + value: {{ .Values.db.readReplicaUrl | quote }} +{{- end }} +{{- if .Values.db.connectionPool.enabled }} +- name: LITELLM_PGBOUNCER_ENABLED + value: "true" +- name: LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS + value: {{ .Values.db.connectionPool.maxDbConnections | quote }} +- name: LITELLM_PGBOUNCER_MAX_CLIENT_CONN + value: {{ .Values.db.connectionPool.maxClientConn | quote }} +{{- end }} +- name: PROXY_MASTER_KEY + valueFrom: + secretKeyRef: + name: {{ .Values.masterkeySecretName | default (printf "%s-masterkey" (include "litellm.fullname" .)) }} + key: {{ .Values.masterkeySecretKey | default "masterkey" }} +{{- if .Values.redis.enabled }} +- name: REDIS_HOST + value: {{ include "litellm.redis.serviceName" . }} +- name: REDIS_PORT + value: {{ include "litellm.redis.port" . | quote }} +- name: REDIS_PASSWORD + valueFrom: + secretKeyRef: + name: {{ include "redis.secretName" .Subcharts.redis }} + key: {{include "redis.secretPasswordKey" .Subcharts.redis }} +{{- end }} +{{- /* + Inject LITELLM_LOG only when envVars does not already define it. +*/}} +{{- if and .Values.logLevel (not (hasKey (default dict .Values.envVars) "LITELLM_LOG")) }} +- name: LITELLM_LOG + value: {{ .Values.logLevel | quote }} +{{- end }} +{{- if .Values.envVars }} +{{- range $key, $val := .Values.envVars }} +- name: {{ $key }} + value: {{ $val | quote }} +{{- end }} +{{- end }} +{{- with .Values.extraEnvVars }} +{{ toYaml . }} +{{- end }} +{{- if .Values.migrationJob.enabled }} +# Schema updates are owned by the dedicated migrations Job; skip +# the proxy's startup `prisma db push` so N replicas don't race +# one DB on every rollout. Placed last (after envVars and +# extraEnvVars) so this override can't be silently shadowed by a +# user-supplied DISABLE_SCHEMA_UPDATE under last-wins duplicate-env +# semantics — same pattern the migrations Job uses. +- name: DISABLE_SCHEMA_UPDATE + value: "true" +{{- end }} +{{- end -}} + +{{/* +Proxy-only metering and metrics env. The collector sidecar serves no HTTP +traffic, so it gets neither. +*/}} +{{- define "litellm.proxyMetricsEnv" -}} +{{- if .Values.billingMetrics.enabled }} +{{ include "litellm.billingMetricsEnv" . }} +{{- end }} +{{- if .Values.metricsServer.enabled }} +{{- if eq (int .Values.metricsServer.port) (int .Values.service.port) }} +{{- fail "metricsServer.port must differ from service.port" }} +{{- end }} +- name: PROMETHEUS_METRICS_PORT + value: {{ .Values.metricsServer.port | quote }} +{{- end }} +{{- end -}} + +{{/* +Directory of the collector's unix socket, shared between the two containers +through an emptyDir. Empty when the sidecar is off or uses 127.0.0.1 TCP. +*/}} +{{- define "litellm.collector.socketDir" -}} +{{- if and .Values.collector.enabled (hasPrefix "unix://" .Values.collector.address) -}} +{{- dir (trimPrefix "unix://" .Values.collector.address) -}} +{{- end -}} +{{- end -}} + +{{- define "litellm.collectorEnv" -}} +- name: LITELLM_COLLECTOR_ENABLED + value: "true" +- name: LITELLM_COLLECTOR_ADDRESS + value: {{ .Values.collector.address | quote }} +- name: LITELLM_COLLECTOR_BUFFER_SIZE + value: {{ .Values.collector.bufferSize | quote }} +- name: LITELLM_COLLECTOR_ON_UNAVAILABLE + value: {{ .Values.collector.onUnavailable | quote }} +- name: LITELLM_COLLECTOR_DRAIN_TIMEOUT_SECONDS + value: {{ .Values.collector.drainTimeoutSeconds | quote }} +{{- end -}} diff --git a/helm/litellm-helm/templates/deployment.yaml b/helm/litellm-helm/templates/deployment.yaml index 834071eb9b2..cf7b3f8a38d 100644 --- a/helm/litellm-helm/templates/deployment.yaml +++ b/helm/litellm-helm/templates/deployment.yaml @@ -56,126 +56,10 @@ spec: image: "{{ .Values.image.repository }}:{{ .Values.image.tag | default .Chart.AppVersion }}" imagePullPolicy: {{ .Values.image.pullPolicy }} env: - - name: HOST - value: "{{ .Values.listen | default "0.0.0.0" }}" - - name: PORT - value: {{ .Values.service.port | quote}} - {{- if .Values.db.deployStandalone }} - - name: DATABASE_USERNAME - valueFrom: - secretKeyRef: - name: {{ include "litellm.fullname" . }}-dbcredentials - key: username - - name: DATABASE_PASSWORD - valueFrom: - secretKeyRef: - name: {{ include "litellm.fullname" . }}-dbcredentials - key: password - - name: DATABASE_HOST - value: {{ .Release.Name }}-postgresql - - name: DATABASE_NAME - value: litellm - {{- else if .Values.db.useExisting }} - - name: DATABASE_USERNAME - valueFrom: - secretKeyRef: - name: {{ .Values.db.secret.name }} - key: {{ .Values.db.secret.usernameKey }} - - name: DATABASE_PASSWORD - valueFrom: - secretKeyRef: - name: {{ .Values.db.secret.name }} - key: {{ .Values.db.secret.passwordKey }} - - name: DATABASE_HOST - {{- if .Values.db.secret.endpointKey }} - valueFrom: - secretKeyRef: - name: {{ .Values.db.secret.name }} - key: {{ .Values.db.secret.endpointKey }} - {{- else }} - value: {{ .Values.db.endpoint }} - {{- end }} - - name: DATABASE_NAME - value: {{ .Values.db.database }} - - name: DATABASE_URL - value: {{ .Values.db.url | quote }} - {{- end }} - {{- if and .Values.db.useExisting .Values.db.readReplicaUrl .Values.db.secret.readReplicaEndpointKey (not .Values.db.secret.readReplicaUrlKey) }} - - name: DATABASE_READER_HOST - valueFrom: - secretKeyRef: - name: {{ .Values.db.secret.name }} - key: {{ .Values.db.secret.readReplicaEndpointKey }} - {{- end }} - {{- if and .Values.db.useExisting .Values.db.secret.readReplicaUrlKey }} - - name: DATABASE_URL_READ_REPLICA - valueFrom: - secretKeyRef: - name: {{ .Values.db.secret.name }} - key: {{ .Values.db.secret.readReplicaUrlKey }} - {{- else if .Values.db.readReplicaUrl }} - - name: DATABASE_URL_READ_REPLICA - value: {{ .Values.db.readReplicaUrl | quote }} - {{- end }} - {{- if .Values.db.connectionPool.enabled }} - - name: LITELLM_PGBOUNCER_ENABLED - value: "true" - - name: LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS - value: {{ .Values.db.connectionPool.maxDbConnections | quote }} - - name: LITELLM_PGBOUNCER_MAX_CLIENT_CONN - value: {{ .Values.db.connectionPool.maxClientConn | quote }} - {{- end }} - - name: PROXY_MASTER_KEY - valueFrom: - secretKeyRef: - name: {{ .Values.masterkeySecretName | default (printf "%s-masterkey" (include "litellm.fullname" .)) }} - key: {{ .Values.masterkeySecretKey | default "masterkey" }} - {{- if .Values.redis.enabled }} - - name: REDIS_HOST - value: {{ include "litellm.redis.serviceName" . }} - - name: REDIS_PORT - value: {{ include "litellm.redis.port" . | quote }} - - name: REDIS_PASSWORD - valueFrom: - secretKeyRef: - name: {{ include "redis.secretName" .Subcharts.redis }} - key: {{include "redis.secretPasswordKey" .Subcharts.redis }} - {{- end }} - {{- /* - Inject LITELLM_LOG only when envVars does not already define it. - */}} - {{- if and .Values.logLevel (not (hasKey (default dict .Values.envVars) "LITELLM_LOG")) }} - - name: LITELLM_LOG - value: {{ .Values.logLevel | quote }} - {{- end }} - {{- if .Values.envVars }} - {{- range $key, $val := .Values.envVars }} - - name: {{ $key }} - value: {{ $val | quote }} - {{- end }} - {{- end }} - {{- with .Values.extraEnvVars }} - {{- toYaml . | nindent 12 }} - {{- end }} - {{- if .Values.billingMetrics.enabled }} - {{- include "litellm.billingMetricsEnv" . | nindent 12 }} - {{- end }} - {{- if .Values.metricsServer.enabled }} - {{- if eq (int .Values.metricsServer.port) (int .Values.service.port) }} - {{- fail "metricsServer.port must differ from service.port" }} - {{- end }} - - name: PROMETHEUS_METRICS_PORT - value: {{ .Values.metricsServer.port | quote }} - {{- end }} - {{- if .Values.migrationJob.enabled }} - # Schema updates are owned by the dedicated migrations Job; skip - # the proxy's startup `prisma db push` so N replicas don't race - # one DB on every rollout. Placed last (after envVars and - # extraEnvVars) so this override can't be silently shadowed by a - # user-supplied DISABLE_SCHEMA_UPDATE under last-wins duplicate-env - # semantics — same pattern the migrations Job uses. - - name: DISABLE_SCHEMA_UPDATE - value: "true" + {{- include "litellm.proxyEnv" . | nindent 12 }} + {{- include "litellm.proxyMetricsEnv" . | nindent 12 }} + {{- if .Values.collector.enabled }} + {{- include "litellm.collectorEnv" . | nindent 12 }} {{- end }} envFrom: {{- range .Values.environmentSecrets }} @@ -253,6 +137,10 @@ spec: {{- if .Values.billingMetrics.enabled }} {{- include "litellm.billingMetricsVolumeMounts" . | nindent 12 }} {{- end }} + {{- if include "litellm.collector.socketDir" . }} + - name: collector-socket + mountPath: {{ include "litellm.collector.socketDir" . }} + {{- end }} {{- with .Values.volumeMounts }} {{- toYaml . | nindent 12 }} {{- end }} @@ -260,6 +148,53 @@ spec: lifecycle: {{- toYaml . | nindent 12 }} {{- end }} + {{- if .Values.collector.enabled }} + - name: {{ include "litellm.name" . }}-collector + securityContext: + {{- toYaml .Values.securityContext | nindent 12 }} + image: "{{ .Values.image.repository }}:{{ .Values.image.tag | default .Chart.AppVersion }}" + imagePullPolicy: {{ .Values.image.pullPolicy }} + command: {{ toYaml .Values.collector.command | nindent 12 }} + env: + {{- include "litellm.proxyEnv" . | nindent 12 }} + {{- include "litellm.collectorEnv" . | nindent 12 }} + - name: LITELLM_JOB_ROLE + value: collector + {{- if not (hasKey (default dict .Values.envVars) "CONFIG_FILE_PATH") }} + - name: CONFIG_FILE_PATH + value: /etc/litellm/config.yaml + {{- end }} + envFrom: + {{- range .Values.environmentSecrets }} + - secretRef: + name: {{ . }} + {{- end }} + {{- range .Values.environmentConfigMaps }} + - configMapRef: + name: {{ . }} + {{- end }} + resources: + {{- toYaml .Values.collector.resources | nindent 12 }} + volumeMounts: + - name: litellm-config + mountPath: /etc/litellm/config.yaml + subPath: config.yaml + {{- if include "litellm.collector.socketDir" . }} + - name: collector-socket + mountPath: {{ include "litellm.collector.socketDir" . }} + {{- end }} + {{ if .Values.securityContext.readOnlyRootFilesystem }} + - name: tmp + mountPath: /tmp + - name: cache + mountPath: /.cache + - name: npm + mountPath: /.npm + {{- end }} + {{- with .Values.volumeMounts }} + {{- toYaml . | nindent 12 }} + {{- end }} + {{- end }} {{- with .Values.extraContainers }} {{- tpl (toYaml .) $ | nindent 8 }} {{- end }} @@ -288,6 +223,11 @@ spec: {{- if .Values.billingMetrics.enabled }} {{- include "litellm.billingMetricsVolumes" . | nindent 8 }} {{- end }} + {{- if include "litellm.collector.socketDir" . }} + - name: collector-socket + emptyDir: + sizeLimit: 1Mi + {{- end }} {{- with .Values.volumes }} {{- toYaml . | nindent 8 }} {{- end }} diff --git a/helm/litellm-helm/templates/hpa.yaml b/helm/litellm-helm/templates/hpa.yaml index 010c4095ab3..a651f916d21 100644 --- a/helm/litellm-helm/templates/hpa.yaml +++ b/helm/litellm-helm/templates/hpa.yaml @@ -18,6 +18,15 @@ spec: {{- end }} metrics: {{- if .Values.autoscaling.targetCPUUtilizationPercentage }} + {{- if and .Values.collector.enabled .Values.collector.scaleOnProxyContainerCpu }} + - type: ContainerResource + containerResource: + name: cpu + container: {{ include "litellm.name" . }} + target: + type: Utilization + averageUtilization: {{ .Values.autoscaling.targetCPUUtilizationPercentage }} + {{- else }} - type: Resource resource: name: cpu @@ -25,6 +34,7 @@ spec: type: Utilization averageUtilization: {{ .Values.autoscaling.targetCPUUtilizationPercentage }} {{- end }} + {{- end }} {{- if .Values.autoscaling.targetMemoryUtilizationPercentage }} - type: Resource resource: diff --git a/helm/litellm-helm/tests/collector_tests.yaml b/helm/litellm-helm/tests/collector_tests.yaml new file mode 100644 index 00000000000..0340b1161b7 --- /dev/null +++ b/helm/litellm-helm/tests/collector_tests.yaml @@ -0,0 +1,272 @@ +suite: test collector sidecar +templates: + - deployment.yaml + - hpa.yaml + - configmap-litellm.yaml +tests: + - it: should run the proxy alone with no collector env by default + template: deployment.yaml + asserts: + - lengthEqual: + path: spec.template.spec.containers + count: 1 + - notContains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_COLLECTOR_ENABLED + value: "true" + - notContains: + path: spec.template.spec.volumes + content: + name: collector-socket + any: true + + - it: should add the sidecar on the same image and point both containers at the unix socket + template: deployment.yaml + set: + image.tag: test + db.connectionPool.enabled: true + collector.enabled: true + collector.resources: + requests: + cpu: 500m + memory: 1Gi + limits: + cpu: "1" + memory: 2Gi + asserts: + - lengthEqual: + path: spec.template.spec.containers + count: 2 + - equal: + path: spec.template.spec.containers[1].name + value: litellm-collector + - equal: + path: spec.template.spec.containers[1].image + value: ghcr.io/berriai/litellm:test + - equal: + path: spec.template.spec.containers[1].command + value: [python, -m, litellm.proxy.collector] + - equal: + path: spec.template.spec.containers[1].resources.requests.cpu + value: 500m + - equal: + path: spec.template.spec.containers[1].resources.limits.memory + value: 2Gi + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_COLLECTOR_ENABLED + value: "true" + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_COLLECTOR_ADDRESS + value: unix:///var/run/litellm/collector.sock + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_COLLECTOR_BUFFER_SIZE + value: "1000" + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_COLLECTOR_ON_UNAVAILABLE + value: fallback + - notContains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_JOB_ROLE + value: collector + - contains: + path: spec.template.spec.containers[1].env + content: + name: LITELLM_JOB_ROLE + value: collector + - contains: + path: spec.template.spec.containers[1].env + content: + name: LITELLM_COLLECTOR_ADDRESS + value: unix:///var/run/litellm/collector.sock + - contains: + path: spec.template.spec.containers[1].env + content: + name: CONFIG_FILE_PATH + value: /etc/litellm/config.yaml + - contains: + path: spec.template.spec.containers[1].env + content: + name: DATABASE_HOST + value: RELEASE-NAME-postgresql + - contains: + path: spec.template.spec.containers[1].env + content: + name: DATABASE_PASSWORD + valueFrom: + secretKeyRef: + name: RELEASE-NAME-litellm-dbcredentials + key: password + - contains: + path: spec.template.spec.containers[1].env + content: + name: LITELLM_PGBOUNCER_ENABLED + value: "true" + - contains: + path: spec.template.spec.containers[1].env + content: + name: LITELLM_PGBOUNCER_MAX_DB_CONNECTIONS + value: "20" + - contains: + path: spec.template.spec.containers[0].volumeMounts + content: + name: collector-socket + mountPath: /var/run/litellm + - contains: + path: spec.template.spec.containers[1].volumeMounts + content: + name: collector-socket + mountPath: /var/run/litellm + - contains: + path: spec.template.spec.containers[1].volumeMounts + content: + name: litellm-config + mountPath: /etc/litellm/config.yaml + subPath: config.yaml + - contains: + path: spec.template.spec.volumes + content: + name: collector-socket + emptyDir: + sizeLimit: 1Mi + + - it: should skip the socket volume and pass the policy through on tcp transport + template: deployment.yaml + set: + collector.enabled: true + collector.address: tcp://127.0.0.1:4100 + collector.onUnavailable: drop + collector.bufferSize: 50 + envVars: + CONFIG_FILE_PATH: /custom/config.yaml + asserts: + - lengthEqual: + path: spec.template.spec.containers + count: 2 + - notContains: + path: spec.template.spec.containers[1].env + content: + name: CONFIG_FILE_PATH + value: /etc/litellm/config.yaml + - contains: + path: spec.template.spec.containers[1].env + content: + name: CONFIG_FILE_PATH + value: /custom/config.yaml + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_COLLECTOR_ADDRESS + value: tcp://127.0.0.1:4100 + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_COLLECTOR_ON_UNAVAILABLE + value: drop + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_COLLECTOR_BUFFER_SIZE + value: "50" + - notContains: + path: spec.template.spec.volumes + content: + name: collector-socket + any: true + + - it: should keep metrics and billing env on the proxy container only + template: deployment.yaml + set: + collector.enabled: true + metricsServer.enabled: true + metricsServer.port: 9090 + billingMetrics.enabled: true + billingMetrics.endpoint: https://metering.example.com + billingMetrics.secretName: billing-mtls + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: PROMETHEUS_METRICS_PORT + value: "9090" + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_BILLING_METRICS_ENDPOINT + value: https://metering.example.com + - notContains: + path: spec.template.spec.containers[1].env + content: + name: PROMETHEUS_METRICS_PORT + any: true + - notContains: + path: spec.template.spec.containers[1].env + content: + name: LITELLM_BILLING_METRICS_ENDPOINT + any: true + - notContains: + path: spec.template.spec.containers[1].volumeMounts + content: + name: billing-metrics-mtls + any: true + + - it: should give the sidecar the same scratch mounts as the proxy on a read-only root + template: deployment.yaml + set: + collector.enabled: true + securityContext.readOnlyRootFilesystem: true + asserts: + - contains: + path: spec.template.spec.containers[1].volumeMounts + content: + name: npm + mountPath: /.npm + - contains: + path: spec.template.spec.containers[1].volumeMounts + content: + name: cache + mountPath: /.cache + - contains: + path: spec.template.spec.containers[1].volumeMounts + content: + name: tmp + mountPath: /tmp + + - it: should keep the pod-wide cpu metric unless asked to scale on the proxy container + template: hpa.yaml + set: + autoscaling.enabled: true + collector.enabled: true + asserts: + - equal: { path: "spec.metrics[0].type", value: Resource } + - equal: { path: "spec.metrics[0].resource.name", value: cpu } + + - it: should scale on the proxy container's cpu only when opted in + template: hpa.yaml + set: + autoscaling.enabled: true + collector.enabled: true + collector.scaleOnProxyContainerCpu: true + asserts: + - equal: { path: "spec.metrics[0].type", value: ContainerResource } + - equal: { path: "spec.metrics[0].containerResource.name", value: cpu } + - equal: { path: "spec.metrics[0].containerResource.container", value: litellm } + - equal: { path: "spec.metrics[0].containerResource.target.averageUtilization", value: 60 } + - isNull: { path: "spec.metrics[0].resource" } + + - it: should not switch to the container metric while the sidecar is off + template: hpa.yaml + set: + autoscaling.enabled: true + collector.scaleOnProxyContainerCpu: true + asserts: + - equal: { path: "spec.metrics[0].type", value: Resource } diff --git a/helm/litellm-helm/values.yaml b/helm/litellm-helm/values.yaml index db596e2d68e..fcee331a5aa 100644 --- a/helm/litellm-helm/values.yaml +++ b/helm/litellm-helm/values.yaml @@ -190,6 +190,48 @@ metricsServer: enabled: false port: 4001 +# Opt-in sidecar that runs the post-response spend pipeline (cost calculation, +# spend logs, spend counters, budget reservation reconciliation) so the proxy's +# uvicorn workers only serialise a compact typed event and go back to serving +# inference. Same image and tag as the proxy, second container in the same pod, +# fed over loopback (a unix socket on a shared emptyDir, or 127.0.0.1 TCP). It +# reuses the pod's in-container pgbouncer (db.connectionPool) and the same Redis +# spend transaction buffer, so the per-pod DB connection budget is unchanged. +# Delivery is at-most-once inside the pod: events already handed to the sidecar +# are lost if it crashes before writing them; events the workers could not hand +# over follow onUnavailable. Both containers drain on SIGTERM within +# terminationGracePeriodSeconds +collector: + enabled: false + # unix:////.sock (the becomes a shared emptyDir) or tcp://127.0.0.1: + address: unix:///var/run/litellm/collector.sock + # Events each uvicorn worker holds in memory while the sidecar is slow or restarting + bufferSize: 1000 + # fallback: run the pipeline in the worker when the sidecar is unreachable or the + # buffer is full (spend stays exact, that request costs proxy CPU again) + # drop: count and discard the event instead (spend under-reports) + onUnavailable: fallback + # How long the workers keep pushing buffered events on shutdown, and how long the + # sidecar keeps serving its open connections after SIGTERM + drainTimeoutSeconds: 10 + command: + - python + - -m + - litellm.proxy.collector + # Sized independently of the proxy container; the pipeline is CPU bound + resources: {} + # requests: + # cpu: 500m + # memory: 1Gi + # limits: + # cpu: "1" + # memory: 2Gi + # When autoscaling.enabled, swap the pod-wide cpu Resource metric for an + # autoscaling/v2 ContainerResource metric on the proxy container only, so the + # sidecar's CPU never scales inference replicas. Needs Kubernetes 1.30+ (or the + # HPAContainerMetrics feature gate on 1.27 to 1.29) + scaleOnProxyContainerCpu: false + resources: {} # Unset by default so the chart installs on small clusters such as Minikube, and so an diff --git a/helm/litellm/templates/_helpers.tpl b/helm/litellm/templates/_helpers.tpl index dc47bc0b0a1..692a799e783 100644 --- a/helm/litellm/templates/_helpers.tpl +++ b/helm/litellm/templates/_helpers.tpl @@ -457,3 +457,34 @@ ImplementationSpecific {{- end -}} {{- define "litellm.gateway.prometheusMultiprocDir" -}}/tmp/litellm_prometheus_multiproc{{- end -}} + +{{/* +Directory of the collector's unix socket, shared by the gateway and +collector containers through an emptyDir. Empty when the sidecar is off +or gateway.collector.address is a tcp://127.0.0.1: address. +*/}} +{{- define "litellm.gateway.collectorSocketDir" -}} +{{- if and .Values.gateway.collector.enabled (hasPrefix "unix://" .Values.gateway.collector.address) -}} +{{- dir (trimPrefix "unix://" .Values.gateway.collector.address) -}} +{{- end -}} +{{- end -}} + +{{/* +LITELLM_COLLECTOR_* env shared by the producer (gateway container) and the +consumer (collector container), so both agree on the transport and the +shutdown drain window. +*/}} +{{- define "litellm.gateway.collectorEnv" -}} +{{- with .Values.gateway.collector }} +- name: LITELLM_COLLECTOR_ENABLED + value: "true" +- name: LITELLM_COLLECTOR_ADDRESS + value: {{ .address | quote }} +- name: LITELLM_COLLECTOR_BUFFER_SIZE + value: {{ .bufferSize | quote }} +- name: LITELLM_COLLECTOR_ON_UNAVAILABLE + value: {{ .onUnavailable | quote }} +- name: LITELLM_COLLECTOR_DRAIN_TIMEOUT_SECONDS + value: {{ .drainTimeoutSeconds | quote }} +{{- end }} +{{- end -}} diff --git a/helm/litellm/templates/gateway/deployment.yaml b/helm/litellm/templates/gateway/deployment.yaml index 1e9041f0a33..5c1a089b50a 100644 --- a/helm/litellm/templates/gateway/deployment.yaml +++ b/helm/litellm/templates/gateway/deployment.yaml @@ -74,8 +74,11 @@ spec: - name: PROMETHEUS_MULTIPROC_DIR value: {{ include "litellm.gateway.prometheusMultiprocDir" . }} {{- end }} + {{- if .Values.gateway.collector.enabled }} + {{- include "litellm.gateway.collectorEnv" . | nindent 12 }} + {{- end }} {{- include "litellm.envFrom" .Values.gateway | nindent 10 }} - {{- if or .Values.gateway.config.create .Values.gateway.volumeMounts .Values.billingMetrics.enabled .Values.gateway.metricsServer.enabled }} + {{- if or .Values.gateway.config.create .Values.gateway.volumeMounts .Values.billingMetrics.enabled .Values.gateway.metricsServer.enabled (include "litellm.gateway.collectorSocketDir" .) }} volumeMounts: {{- if .Values.gateway.config.create }} - name: gateway-config @@ -86,6 +89,10 @@ spec: - name: prometheus-multiproc mountPath: {{ include "litellm.gateway.prometheusMultiprocDir" . }} {{- end }} + {{- if include "litellm.gateway.collectorSocketDir" . }} + - name: collector-socket + mountPath: {{ include "litellm.gateway.collectorSocketDir" . }} + {{- end }} {{- if .Values.billingMetrics.enabled }} {{- include "litellm.billingMetricsVolumeMounts" . | nindent 12 }} {{- end }} @@ -145,10 +152,50 @@ spec: resources: {{- toYaml .Values.gateway.metricsServer.resources | nindent 12 }} {{- end }} + {{- if .Values.gateway.collector.enabled }} + - name: collector + image: "{{ .Values.gateway.image.repository }}:{{ .Values.gateway.image.tag | default .Chart.AppVersion }}" + imagePullPolicy: {{ .Values.gateway.image.pullPolicy }} + {{- with .Values.gateway.securityContext }} + securityContext: + {{- toYaml . | nindent 12 }} + {{- end }} + command: + - python + - -m + - litellm.proxy.collector + env: + {{- include "litellm.serverEnv" (dict "root" $ "component" .Values.gateway) | nindent 12 }} + {{- if .Values.gateway.config.create }} + - name: CONFIG_FILE_PATH + value: /app/config/config.yaml + {{- end }} + {{- include "litellm.gateway.collectorEnv" . | nindent 12 }} + - name: LITELLM_JOB_ROLE + value: collector + {{- include "litellm.envFrom" .Values.gateway | nindent 10 }} + {{- if or .Values.gateway.config.create .Values.gateway.volumeMounts (include "litellm.gateway.collectorSocketDir" .) }} + volumeMounts: + {{- if .Values.gateway.config.create }} + - name: gateway-config + mountPath: /app/config/config.yaml + subPath: config.yaml + {{- end }} + {{- if include "litellm.gateway.collectorSocketDir" . }} + - name: collector-socket + mountPath: {{ include "litellm.gateway.collectorSocketDir" . }} + {{- end }} + {{- with .Values.gateway.volumeMounts }} + {{- toYaml . | nindent 12 }} + {{- end }} + {{- end }} + resources: + {{- toYaml .Values.gateway.collector.resources | nindent 12 }} + {{- end }} {{- with .Values.gateway.extraContainers }} {{- tpl (toYaml .) $ | nindent 8 }} {{- end }} - {{- if or .Values.gateway.config.create .Values.gateway.volumes .Values.billingMetrics.enabled .Values.gateway.metricsServer.enabled }} + {{- if or .Values.gateway.config.create .Values.gateway.volumes .Values.billingMetrics.enabled .Values.gateway.metricsServer.enabled (include "litellm.gateway.collectorSocketDir" .) }} volumes: {{- if .Values.gateway.config.create }} - name: gateway-config @@ -159,6 +206,11 @@ spec: - name: prometheus-multiproc emptyDir: {} {{- end }} + {{- if include "litellm.gateway.collectorSocketDir" . }} + - name: collector-socket + emptyDir: + sizeLimit: 1Mi + {{- end }} {{- if .Values.billingMetrics.enabled }} {{- include "litellm.billingMetricsVolumes" . | nindent 8 }} {{- end }} diff --git a/helm/litellm/templates/gateway/hpa.yaml b/helm/litellm/templates/gateway/hpa.yaml index 4183079cd31..e7094e96106 100644 --- a/helm/litellm/templates/gateway/hpa.yaml +++ b/helm/litellm/templates/gateway/hpa.yaml @@ -15,6 +15,15 @@ spec: maxReplicas: {{ .Values.gateway.hpa.maxReplicas }} metrics: {{- if .Values.gateway.hpa.targetCPUUtilizationPercentage }} + {{- if and .Values.gateway.collector.enabled .Values.gateway.collector.scaleOnGatewayContainerCpu }} + - type: ContainerResource + containerResource: + name: cpu + container: gateway + target: + type: Utilization + averageUtilization: {{ .Values.gateway.hpa.targetCPUUtilizationPercentage }} + {{- else }} - type: Resource resource: name: cpu @@ -22,6 +31,7 @@ spec: type: Utilization averageUtilization: {{ .Values.gateway.hpa.targetCPUUtilizationPercentage }} {{- end }} + {{- end }} {{- if .Values.gateway.hpa.targetMemoryUtilizationPercentage }} - type: Resource resource: diff --git a/helm/litellm/tests/collector_tests.yaml b/helm/litellm/tests/collector_tests.yaml new file mode 100644 index 00000000000..b1199a16793 --- /dev/null +++ b/helm/litellm/tests/collector_tests.yaml @@ -0,0 +1,204 @@ +suite: test gateway collector sidecar +templates: + - gateway/configmap.yaml + - gateway/deployment.yaml + - gateway/hpa.yaml +values: + - ./values/required.yaml +tests: + - it: adds no sidecar, env, volume or container metric when the collector is off + asserts: + - lengthEqual: + path: spec.template.spec.containers + count: 1 + template: gateway/deployment.yaml + - notContains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_COLLECTOR_ENABLED + value: "true" + template: gateway/deployment.yaml + - notContains: + path: spec.template.spec.volumes + content: + name: collector-socket + any: true + template: gateway/deployment.yaml + - equal: + path: spec.metrics[0].type + value: Resource + template: gateway/hpa.yaml + + - it: runs the collector as a sidecar sharing env, config and a unix socket emptyDir, and scales on the gateway container only + set: + gateway.collector.enabled: true + gateway.collector.bufferSize: 250 + gateway.collector.onUnavailable: drop + gateway.image.tag: v1.102.0 + gateway.numWorkers: 4 + gateway.extraEnv: + - name: LITELLM_PGBOUNCER_ENABLED + value: "true" + gateway.envSecrets: + - litellm-license + gateway.volumes: + - name: redis-ca + secret: + secretName: redis-ca + gateway.volumeMounts: + - name: redis-ca + mountPath: /etc/litellm/redis-ca + readOnly: true + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_COLLECTOR_ADDRESS + value: unix:///var/run/litellm/collector.sock + template: gateway/deployment.yaml + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_COLLECTOR_BUFFER_SIZE + value: "250" + template: gateway/deployment.yaml + - contains: + path: spec.template.spec.containers[0].env + content: + name: LITELLM_COLLECTOR_ON_UNAVAILABLE + value: drop + template: gateway/deployment.yaml + - contains: + path: spec.template.spec.containers[0].volumeMounts + content: + name: collector-socket + mountPath: /var/run/litellm + template: gateway/deployment.yaml + - equal: + path: spec.template.spec.containers[1].name + value: collector + template: gateway/deployment.yaml + - equal: + path: spec.template.spec.containers[1].image + value: ghcr.io/berriai/litellm-gateway:v1.102.0 + template: gateway/deployment.yaml + - equal: + path: spec.template.spec.containers[1].command + value: + - python + - -m + - litellm.proxy.collector + template: gateway/deployment.yaml + - contains: + path: spec.template.spec.containers[1].env + content: + name: LITELLM_JOB_ROLE + value: collector + template: gateway/deployment.yaml + - contains: + path: spec.template.spec.containers[1].env + content: + name: CONFIG_FILE_PATH + value: /app/config/config.yaml + template: gateway/deployment.yaml + - contains: + path: spec.template.spec.containers[1].env + content: + name: LITELLM_PGBOUNCER_ENABLED + value: "true" + template: gateway/deployment.yaml + - contains: + path: spec.template.spec.containers[1].env + content: + name: DATABASE_HOST + value: postgres.example.com + template: gateway/deployment.yaml + - contains: + path: spec.template.spec.containers[1].env + content: + name: LITELLM_COLLECTOR_ADDRESS + value: unix:///var/run/litellm/collector.sock + template: gateway/deployment.yaml + - notContains: + path: spec.template.spec.containers[1].env + content: + name: NUM_WORKERS + any: true + template: gateway/deployment.yaml + - equal: + path: spec.template.spec.containers[1].envFrom + value: + - secretRef: + name: litellm-license + template: gateway/deployment.yaml + - contains: + path: spec.template.spec.containers[1].volumeMounts + content: + name: gateway-config + mountPath: /app/config/config.yaml + subPath: config.yaml + template: gateway/deployment.yaml + - contains: + path: spec.template.spec.containers[1].volumeMounts + content: + name: collector-socket + mountPath: /var/run/litellm + template: gateway/deployment.yaml + - contains: + path: spec.template.spec.containers[1].volumeMounts + content: + name: redis-ca + mountPath: /etc/litellm/redis-ca + readOnly: true + template: gateway/deployment.yaml + - equal: + path: spec.template.spec.containers[1].resources.limits.cpu + value: "1" + template: gateway/deployment.yaml + - contains: + path: spec.template.spec.volumes + content: + name: collector-socket + emptyDir: + sizeLimit: 1Mi + template: gateway/deployment.yaml + - equal: + path: spec.metrics[0] + value: + type: ContainerResource + containerResource: + name: cpu + container: gateway + target: + type: Utilization + averageUtilization: 70 + template: gateway/hpa.yaml + + - it: uses loopback tcp without a socket volume and keeps the pod-wide cpu metric when asked + set: + gateway.collector.enabled: true + gateway.collector.address: tcp://127.0.0.1:4010 + gateway.collector.scaleOnGatewayContainerCpu: false + asserts: + - contains: + path: spec.template.spec.containers[1].env + content: + name: LITELLM_COLLECTOR_ADDRESS + value: tcp://127.0.0.1:4010 + template: gateway/deployment.yaml + - notContains: + path: spec.template.spec.volumes + content: + name: collector-socket + any: true + template: gateway/deployment.yaml + - notContains: + path: spec.template.spec.containers[1].volumeMounts + content: + name: collector-socket + any: true + template: gateway/deployment.yaml + - equal: + path: spec.metrics[0].type + value: Resource + template: gateway/hpa.yaml diff --git a/helm/litellm/values.yaml b/helm/litellm/values.yaml index 677e21a5a51..1f67b984513 100644 --- a/helm/litellm/values.yaml +++ b/helm/litellm/values.yaml @@ -314,6 +314,42 @@ gateway: labels: {} interval: 15s scrapeTimeout: 10s + # Opt-in `collector` sidecar (same image, `python -m litellm.proxy.collector`) + # that runs the post-response spend pipeline (cost calculation, spend logs, + # spend counters, budget reservation reconciliation) so the uvicorn workers + # only serialise a compact event over loopback and go back to serving + # requests. It shares the pod's env, proxy config, in-container pgbouncer and + # Redis spend buffer, so the per-pod DB connection budget is unchanged. + # Delivery is at-most-once inside the pod: events already handed over are + # lost if the sidecar dies before writing them; events the workers cannot + # hand over follow `onUnavailable`. + collector: + enabled: false + # unix:////.sock (the becomes a shared emptyDir) or + # tcp://127.0.0.1: + address: unix:///var/run/litellm/collector.sock + # Events each uvicorn worker holds in memory while the sidecar is slow or + # restarting. + bufferSize: 1000 + # fallback: run the pipeline in the worker when the sidecar is unreachable + # or the buffer is full (spend stays exact, that request costs gateway CPU + # again). drop: count and discard the event instead (spend under-reports). + onUnavailable: fallback + # How long the workers keep pushing buffered events on shutdown, and how + # long the sidecar keeps serving open connections after SIGTERM. + drainTimeoutSeconds: 10 + # Sized independently of the gateway container; the pipeline is CPU bound. + resources: + requests: + cpu: 500m + memory: 1Gi + limits: + cpu: "1" + memory: 2Gi + # With hpa.targetCPUUtilizationPercentage set, scale on an autoscaling/v2 + # ContainerResource metric of the `gateway` container only, so the + # sidecar's CPU never drives inference replicas. Needs Kubernetes 1.30+. + scaleOnGatewayContainerCpu: true image: repository: ghcr.io/berriai/litellm-gateway tag: "" # defaults to .Chart.AppVersion diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 0de88eacaa5..139dcf058d2 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -1217,7 +1217,9 @@ class LLMCachingHandler: } if litellm.cache is not None: - litellm_params["preset_cache_key"] = litellm.cache._get_preset_cache_key_from_kwargs(**kwargs) + litellm_params["preset_cache_key"] = ( + self.preset_cache_key or litellm.cache._get_preset_cache_key_from_kwargs(**kwargs) + ) else: litellm_params["preset_cache_key"] = None diff --git a/litellm/proxy/collector.py b/litellm/proxy/collector.py new file mode 100644 index 00000000000..3ff2a83860b --- /dev/null +++ b/litellm/proxy/collector.py @@ -0,0 +1,220 @@ +"""Collector sidecar: consume spend events from the pod's inference workers and run the cost pipeline. + +Runs the proxy startup lifespan (config, Prisma, Redis transaction buffer, scheduled spend flushes) +without serving HTTP, then listens on ``LITELLM_COLLECTOR_ADDRESS`` for newline-delimited spend +events. Each event goes through the unchanged ``_ProxyDBLogger._PROXY_track_cost_callback``, so +spend logs, spend counters, budget reservation reconciliation and cache updates happen exactly as +they would in-process, just in this container. Events are handled in order per producer connection +(one per uvicorn worker); a slow pipeline fills the socket buffer and the producer's bounded queue, +which is the backpressure that triggers its fallback or drop policy. ``SIGTERM`` stops accepting +connections, half-closes every producer connection so the producers switch to their unavailable +policy, finishes the events already sent, then runs the proxy shutdown (which flushes the buffered +spend transactions). + +``DATABASE_URL`` is assembled from the same ``DATABASE_*`` inputs as the proxy container, and when +``LITELLM_PGBOUNCER_ENABLED`` is set it points at the PgBouncer that container already runs on the +pod's loopback, so the sidecar must see the same env as the proxy. Under ``IAM_TOKEN_DB_AUTH`` or +``AZURE_POSTGRESQL_AUTH`` that PgBouncer only accepts the token the proxy container minted, so the +sidecar goes to Postgres directly and mints its own. Works from any image that has ``litellm`` +installed: + + python -m litellm.proxy.collector [--address unix:///path.sock] +""" + +import asyncio +import logging +import os +import signal +import sys +from collections.abc import Awaitable, Callable, Mapping, Sequence +from pathlib import Path +from typing import Final + +from litellm._logging import verbose_logger, verbose_proxy_logger, verbose_router_logger +from litellm.proxy.db.db_url_settings import DatabaseURLSettings +from litellm.proxy.db.pgbouncer import ( + PgBouncerError, + PgBouncerSettings, + export_pooled_database_url, + pooled_database_url, +) +from litellm.proxy.spend_tracking.spend_event_producer import ( + COLLECTOR_JOB_ROLE, + AddressError, + CollectorAddress, + CollectorSettings, + TcpAddress, + UnixAddress, + parse_collector_address, +) + +MAX_EVENT_BYTES: Final = 64 * 1024 * 1024 + + +class SpendEventConsumer: + """Accepts producer connections and runs ``handler`` on every line each one sends, in order.""" + + def __init__(self, handler: Callable[[bytes], Awaitable[None]]) -> None: + self._handler = handler + self._open_connections: set[asyncio.StreamWriter] = set() # mutable-ok: live producer connections + self._idle = asyncio.Event() + self._idle.set() + self._received = 0 + self._handled = 0 + self._failed = 0 + + @property + def received(self) -> int: + return self._received + + @property + def handled(self) -> int: + return self._handled + + @property + def failed(self) -> int: + return self._failed + + async def serve(self, address: CollectorAddress) -> asyncio.Server: + match address: + case UnixAddress(path=path): + socket_path: Final = Path(path) + socket_path.parent.mkdir(parents=True, exist_ok=True) + socket_path.unlink(missing_ok=True) + return await asyncio.start_unix_server(self._on_connection, path=path, limit=MAX_EVENT_BYTES) + case TcpAddress(host=host, port=port): + return await asyncio.start_server(self._on_connection, host=host, port=port, limit=MAX_EVENT_BYTES) + + async def _on_connection(self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: + self._open_connections.add(writer) + self._idle.clear() + try: + while line := await reader.readline(): + if not line.endswith(b"\n"): + verbose_proxy_logger.error("collector: discarding truncated spend event (%d bytes)", len(line)) + break + self._received += 1 + await self._handle(line) + except (ConnectionError, asyncio.IncompleteReadError, asyncio.LimitOverrunError) as error: + verbose_proxy_logger.warning("collector: producer connection ended abnormally: %s", error) + finally: + writer.close() + self._open_connections.discard(writer) + if not self._open_connections: + self._idle.set() + + async def _handle(self, line: bytes) -> None: + try: + await self._handler(line) + self._handled += 1 + except Exception: # noqa: BLE001 # the cost pipeline raises anything; one bad event must not stop the sidecar + self._failed += 1 + verbose_proxy_logger.exception("collector: spend event failed") + + async def drain(self, timeout: float) -> int: + """Half-close every producer connection, then keep reading until each producer hangs up or ``timeout``. + + Returns how many producer connections were still open when the timeout hit. + """ + for writer in tuple(self._open_connections): + if writer.is_closing() or not writer.can_write_eof(): + continue + try: + writer.write_eof() + except (OSError, RuntimeError) as error: + verbose_proxy_logger.debug("collector: producer already gone before half-close: %s", error) + try: + await asyncio.wait_for(self._idle.wait(), timeout) + except TimeoutError: + pass + return len(self._open_connections) + + +def _install_stop_signals(loop: asyncio.AbstractEventLoop, stop: asyncio.Event) -> None: + for signum in (signal.SIGTERM, signal.SIGINT): + loop.add_signal_handler(signum, stop.set) + + +async def run_collector(address: CollectorAddress, drain_timeout: float) -> None: + from fastapi import FastAPI + + from litellm.proxy.hooks.proxy_track_cost_callback import run_spend_event + from litellm.proxy.proxy_server import proxy_startup_event + + stop: Final = asyncio.Event() + _install_stop_signals(asyncio.get_running_loop(), stop) + consumer: Final = SpendEventConsumer(handler=run_spend_event) + async with proxy_startup_event(FastAPI()): + server: Final = await consumer.serve(address) + verbose_proxy_logger.info("collector: listening on %s", address) + await stop.wait() + server.close() + still_open: Final = await consumer.drain(drain_timeout) + verbose_proxy_logger.info( + "collector: stopping. received=%d handled=%d failed=%d connections_cut=%d", + consumer.received, + consumer.handled, + consumer.failed, + still_open, + ) + + +def address_argument(argv: Sequence[str], default: str) -> str | AddressError: + match tuple(argv): + case (): + return default + case ("--address", value): + return value + case _: + return AddressError(f"usage: python -m litellm.proxy.collector [--address ADDRESS], got {tuple(argv)}") + + +def apply_log_level(litellm_log: str | None) -> None: + """Mirror the proxy's ``LITELLM_LOG`` handling: the sidecar has no CLI flags to turn logging on.""" + level: Final = logging.getLevelNamesMapping().get((litellm_log or "").upper()) + if level is None: + return + for logger in (verbose_logger, verbose_router_logger, verbose_proxy_logger): + logger.setLevel(level) + + +def pod_pgbouncer_database_url( + pgbouncer: PgBouncerSettings, environ: Mapping[str, str], *, token_auth: bool +) -> str | PgBouncerError | None: + """The proxy container's PgBouncer URL for ``environ["DATABASE_URL"]``, or None to connect to Postgres directly. + + Direct is the answer when PgBouncer is off, and also under token auth: that PgBouncer's auth file + only holds the token its own container minted, which this container cannot present. + """ + if not pgbouncer.enabled or token_auth: + return None + upstream_url: Final = environ.get("DATABASE_URL") + if upstream_url is None: + return PgBouncerError("LITELLM_PGBOUNCER_ENABLED is set but no DATABASE_URL could be assembled") + return pooled_database_url(upstream_url, pgbouncer) + + +def main(argv: Sequence[str]) -> None: + os.environ.setdefault("LITELLM_JOB_ROLE", COLLECTOR_JOB_ROLE) + apply_log_level(os.environ.get("LITELLM_LOG")) + database: Final = DatabaseURLSettings.from_env() + database.apply_to_env() + pooled: Final = pod_pgbouncer_database_url( + PgBouncerSettings(), + os.environ, + token_auth=database.iam_token_db_auth or database.azure_postgresql_auth, + ) + if isinstance(pooled, PgBouncerError): + sys.exit(f"LiteLLM collector: cannot use the pod's pgbouncer: {pooled.reason}") + if pooled is not None: + export_pooled_database_url(pooled) + settings: Final = CollectorSettings() + raw_address: Final = address_argument(argv, default=settings.address) + address: Final = raw_address if isinstance(raw_address, AddressError) else parse_collector_address(raw_address) + if isinstance(address, AddressError): + sys.exit(f"LiteLLM collector: {address.reason}") + asyncio.run(run_collector(address, drain_timeout=settings.drain_timeout_seconds)) + + +if __name__ == "__main__": + main(sys.argv[1:]) diff --git a/litellm/proxy/db/pgbouncer.py b/litellm/proxy/db/pgbouncer.py index eac80f7a712..c9fbabd3585 100644 --- a/litellm/proxy/db/pgbouncer.py +++ b/litellm/proxy/db/pgbouncer.py @@ -280,6 +280,22 @@ def plan_pgbouncer( ) +def pooled_database_url(upstream_url: str, settings: PgBouncerSettings) -> str | PgBouncerError: + """The loopback URL of a PgBouncer another container in the pod already runs for ``upstream_url``. + + Only the container that started PgBouncer knows the pool user's password, so + this logs in as the upstream user, whom the auth file lists as well. + """ + plan: Final = plan_pgbouncer(upstream_url, settings, runtime_dir=Path("/nonexistent"), run_as_user=None) + if isinstance(plan, PgBouncerError): + return plan + password: Final = urllib.parse.urlsplit(upstream_url).password or "" + credentials: Final = f"{urllib.parse.quote(plan.upstream_user, safe='')}:{password}" + return urllib.parse.urlunsplit( + urllib.parse.urlsplit(plan.pooled_url)._replace(netloc=f"{credentials}@{PGBOUNCER_LISTEN_ADDR}:{settings.port}") + ) + + def _write_private(path: Path, content: str, run_as_user: str | None) -> None: with open(os.open(path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600), "w", encoding="utf-8") as handle: handle.write(content) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index c4fba8ecf9e..00406ad436e 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -27,6 +27,16 @@ from litellm.proxy.db.db_spend_update_writer import ( get_llm_router, ) from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup +from litellm.proxy.spend_tracking.spend_event import ( + ObjectMapping, + SpendEventBuildError, + SpendEventDecodeError, + build_spend_event, + decode_spend_event, + is_offloadable_success, + spend_event_callback_args, +) +from litellm.proxy.spend_tracking.spend_event_producer import SpendEventProducer from litellm.proxy.spend_tracking.spend_log_error_logger import ( should_suppress_spend_log_tracebacks, spend_log_error, @@ -34,6 +44,7 @@ from litellm.proxy.spend_tracking.spend_log_error_logger import ( from litellm.proxy.spend_tracking.spend_tracking_utils import ( _sanitize_error_information_for_spend_logs, get_request_model_access_groups, + should_store_prompts_and_responses_in_spend_logs, ) from litellm.proxy.utils import ProxyUpdateSpend from litellm.types.utils import ( @@ -71,8 +82,43 @@ _CAPTURED_IDENTITY_CALL_TYPES: Final[frozenset[str]] = frozenset( class _ProxyDBLogger(CustomLogger): - async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): - await self._PROXY_track_cost_callback(kwargs, response_obj, start_time, end_time) + def __init__( + self, + spend_event_producer: SpendEventProducer | None = None, + *, + turn_off_message_logging: bool = False, + message_logging: bool = True, + ) -> None: + super().__init__(turn_off_message_logging=turn_off_message_logging, message_logging=message_logging) + self.spend_event_producer = spend_event_producer + + async def async_log_success_event( + self, kwargs: ObjectMapping, response_obj: object, start_time: datetime, end_time: datetime + ) -> None: + if self.spend_event_producer is None or not is_offloadable_success(response_obj): + await self._PROXY_track_cost_callback(kwargs, response_obj, start_time, end_time) + return + event: Final = build_spend_event( + kwargs, + response_obj, + start_time, + end_time, + store_bodies=should_store_prompts_and_responses_in_spend_logs(), + ) + if isinstance(event, SpendEventBuildError): + verbose_proxy_logger.warning("collector: tracking cost in-process, event not buildable: %s", event.reason) + await self._PROXY_track_cost_callback(kwargs, response_obj, start_time, end_time) + return + await self.spend_event_producer.publish(event) + + async def run_spend_event(self, line: bytes) -> None: + """Run the unchanged cost pipeline on a serialized spend event (sidecar consumer and in-process fallback).""" + event: Final = decode_spend_event(line) + if isinstance(event, SpendEventDecodeError): + verbose_proxy_logger.error("collector: discarding undecodable spend event: %s", event.reason) + return + args: Final = spend_event_callback_args(event) + await self._PROXY_track_cost_callback(args.kwargs, args.response_obj, args.start_time, args.end_time) async def async_post_call_failure_hook( self, @@ -503,6 +549,10 @@ def _write_spend_metadata_to_kwargs(kwargs: dict, metadata: dict) -> None: bucket[key] = value +async def run_spend_event(line: bytes) -> None: + await _ProxyDBLogger().run_spend_event(line) + + def _is_unbilled_interaction_response(completion_response: object) -> bool: from litellm.interactions.background_cost_polling import missing_usage_is_expected from litellm.types.interactions import InteractionsAPIResponse diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 7b64f65c79b..0f24acb8bb4 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -469,7 +469,7 @@ from litellm.proxy.hooks.model_max_budget_limiter import ( from litellm.proxy.hooks.prompt_injection_detection import ( _OPTIONAL_PromptInjectionDetection, ) -from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger +from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger, run_spend_event from litellm.proxy.image_endpoints.endpoints import router as image_router from litellm.proxy.list_api.common import ( PROBLEM_TYPE_BASE, @@ -592,6 +592,11 @@ from litellm.proxy.plugin_routes import ( from litellm.proxy.plugin_routes import ( router as plugin_router, ) +from litellm.proxy.spend_tracking.spend_event_producer import ( + CollectorSettings, + SpendEventProducer, + build_spend_event_producer, +) from litellm.types.proxy.management_endpoints.management_v1 import ProblemDetail try: @@ -943,6 +948,17 @@ def cleanup_router_config_variables(): heuristic_v1_tuning_baselines = None +async def flush_spend_counters_on_shutdown() -> None: + if prisma_client is None: + return + try: + await proxy_logging_obj.db_spend_update_writer.db_update_spend_transaction_handler( + prisma_client=prisma_client, n_retry_times=3, proxy_logging_obj=proxy_logging_obj + ) + except Exception as e: # noqa: BLE001 # shutdown must continue even if the commit fails + verbose_proxy_logger.exception("Error flushing spend counters on shutdown: %s", e) + + async def _flush_spend_logs_queue_on_shutdown() -> None: if prisma_client is None: return @@ -1378,6 +1394,10 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: except Exception as e: verbose_proxy_logger.error("Error stopping DB health watchdog task: %s", e) + await _drain_spend_event_producer_on_shutdown() + + await flush_spend_counters_on_shutdown() + await _flush_spend_logs_queue_on_shutdown() await proxy_config.stop_config_sync_subscriber() @@ -2457,16 +2477,27 @@ def load_from_azure_key_vault(use_azure_key_vault: bool = False): ) +spend_event_producer: SpendEventProducer | None = None + + def cost_tracking(): - global prisma_client + global prisma_client, spend_event_producer if prisma_client is not None: from litellm.integrations.shadow_eval_logger import ShadowEvalLogger - litellm.logging_callback_manager.add_litellm_callback(_ProxyDBLogger()) - litellm.logging_callback_manager.add_litellm_async_success_callback(_ProxyDBLogger()) + spend_event_producer = build_spend_event_producer(CollectorSettings(), fallback=run_spend_event) + litellm.logging_callback_manager.add_litellm_callback(_ProxyDBLogger(spend_event_producer)) + litellm.logging_callback_manager.add_litellm_async_success_callback(_ProxyDBLogger(spend_event_producer)) litellm.logging_callback_manager.add_litellm_callback(ShadowEvalLogger()) +async def _drain_spend_event_producer_on_shutdown() -> None: + if spend_event_producer is None: + return + await spend_event_producer.close(drain_timeout=CollectorSettings().drain_timeout_seconds) + verbose_proxy_logger.info("collector: producer drained on shutdown. stats=%s", spend_event_producer.stats()) + + # Bounds authoritative DB re-reads when enforcing a budget against a # stale-low spend counter: at most one DB read per counter per window. SPEND_DB_FLOOR_CACHE_TTL_SECONDS: Final = 5 diff --git a/litellm/proxy/spend_tracking/spend_event.py b/litellm/proxy/spend_tracking/spend_event.py new file mode 100644 index 00000000000..53f26346f85 --- /dev/null +++ b/litellm/proxy/spend_tracking/spend_event.py @@ -0,0 +1,418 @@ +"""Compact, typed success event handed from an inference worker to the collector. + +``build_spend_event`` runs on the inference worker right after ``Logging.async_success_handler`` +has built the ``standard_logging_object`` (so the cost is already known). It validates the success +callback's ``kwargs`` into the projection ``_PROXY_track_cost_callback`` and +``DBSpendUpdateWriter.update_database`` actually read: identities and metadata, timings, usage, the +standard logging payload without its prompt/response bodies, and the tool names. The request +messages, the raw ``proxy_server_request`` body and the full response travel only when spend logs +are configured to store prompts and responses. The cache key is the preset key the caching layer +already computed, never a fresh hash over the request body. + +``spend_event_callback_args`` rebuilds the ``(kwargs, response_obj, start_time, end_time)`` tuple +the existing cost pipeline consumes, so the sidecar runs the unchanged pipeline against the event. +""" + +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from datetime import datetime +from types import MappingProxyType +from typing import Final, Literal, TypeAlias + +from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError +from typing_extensions import NotRequired, ReadOnly, TypedDict + +import litellm +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.db.spend_log_tool_index import response_tool_call_names +from litellm.types.interactions import InteractionsAPIResponse +from litellm.types.utils import LiteLLMBatch, Usage + +SPEND_EVENT_VERSION: Final = 1 +CACHE_OFF_KEY: Final = "Cache OFF" + +ObjectMapping: TypeAlias = Mapping[str, object] + +_UNSERIALIZABLE_METADATA_KEYS: Final = frozenset({"user_api_key_auth", "litellm_parent_otel_span"}) +_STANDARD_LOGGING_BODY_KEYS: Final = frozenset({"messages", "response"}) +_STANDARD_LOGGING_DROPPED_KEYS: Final = frozenset({"model_parameters"}) +_NOT_OFFLOADED_RESPONSE_TYPES: Final = (LiteLLMBatch, InteractionsAPIResponse) + + +class _LitellmParams(TypedDict, total=False): + api_base: ReadOnly[str | None] + custom_llm_provider: ReadOnly[str | None] + litellm_call_id: ReadOnly[str | None] + user_api_key_end_user_id: ReadOnly[str | None] + metadata: ReadOnly[ObjectMapping | None] + litellm_metadata: ReadOnly[ObjectMapping | None] + proxy_server_request: ReadOnly[ObjectMapping | None] + preset_cache_key: ReadOnly[str | None] + + +class _DynamicParams(TypedDict, total=False): + turn_off_message_logging: ReadOnly[bool | None] + + +class _RequestBody(TypedDict, total=False): + tools: ReadOnly[Sequence[ObjectMapping] | None] + + +class _PassthroughPayload(TypedDict, total=False): + request_body: ReadOnly[_RequestBody | None] + + +class _ToolCallFunction(TypedDict): + name: ReadOnly[str] + arguments: ReadOnly[str] + + +class _ToolCall(TypedDict): + id: ReadOnly[str | None] + type: ReadOnly[Literal["function"]] + function: ReadOnly[_ToolCallFunction] + + +class _ToolCallMessage(TypedDict): + role: ReadOnly[Literal["assistant"]] + content: ReadOnly[None] + tool_calls: ReadOnly[Sequence[_ToolCall]] + + +class _ToolCallChoice(TypedDict): + index: ReadOnly[int] + finish_reason: ReadOnly[Literal["tool_calls"]] + message: ReadOnly[_ToolCallMessage] + + +class CompactResponse(TypedDict, total=False): + """What the spend pipeline reads off a response: its id, usage and which tools it called.""" + + id: ReadOnly[object] + model: ReadOnly[object] + usage: ReadOnly[object] + usage_info: ReadOnly[object] + status: ReadOnly[object] + background: ReadOnly[object] + choices: ReadOnly[Sequence[_ToolCallChoice]] + + +class _SuccessKwargs(TypedDict, total=False): + """The success callback's ``kwargs`` (``Logging.model_call_details``), validated and projected.""" + + litellm_call_id: ReadOnly[str | None] + call_type: ReadOnly[str | None] + model: ReadOnly[str | None] + custom_llm_provider: ReadOnly[str | None] + stream: ReadOnly[bool | None] + complete_streaming_response: ReadOnly[object] + cache_hit: ReadOnly[bool | None] + response_cost: ReadOnly[float | None] + completion_start_time: ReadOnly[datetime | None] + agent_id: ReadOnly[str | None] + litellm_trace_id: ReadOnly[str | None] + litellm_params: ReadOnly[_LitellmParams] + standard_logging_object: ReadOnly[ObjectMapping | None] + standard_callback_dynamic_params: ReadOnly[_DynamicParams | None] + combined_usage_object: ReadOnly[Usage | None] + realtime_tools: ReadOnly[Sequence[object] | None] + realtime_tool_calls: ReadOnly[Sequence[object] | None] + tools: ReadOnly[Sequence[ObjectMapping] | None] + passthrough_logging_payload: ReadOnly[_PassthroughPayload | None] + + +class _FunctionToolFunction(TypedDict): + name: ReadOnly[str] + + +class _FunctionTool(TypedDict): + type: ReadOnly[Literal["function"]] + function: ReadOnly[_FunctionToolFunction] + + +class SpendCallbackKwargs(TypedDict): + """The ``kwargs`` handed to ``_PROXY_track_cost_callback`` on the sidecar.""" + + litellm_call_id: ReadOnly[str | None] + call_type: ReadOnly[str | None] + model: ReadOnly[str | None] + custom_llm_provider: ReadOnly[str | None] + stream: ReadOnly[bool | None] + cache_hit: ReadOnly[bool | None] + response_cost: ReadOnly[float | None] + completion_start_time: ReadOnly[datetime | None] + agent_id: ReadOnly[str | None] + litellm_trace_id: ReadOnly[str | None] + litellm_params: ReadOnly[_LitellmParams] + standard_logging_object: ReadOnly[ObjectMapping | None] + standard_callback_dynamic_params: ReadOnly[_DynamicParams | None] + combined_usage_object: ReadOnly[Usage | None] + realtime_tools: ReadOnly[Sequence[object] | None] + realtime_tool_calls: ReadOnly[Sequence[object] | None] + tools: ReadOnly[Sequence[_FunctionTool] | None] + complete_streaming_response: NotRequired[ReadOnly[CompactResponse | None]] + + +_NO_LITELLM_PARAMS: Final[_LitellmParams] = {} +_SUCCESS_KWARGS: Final = TypeAdapter(_SuccessKwargs) +_OBJECT_MAPPING: Final = TypeAdapter(ObjectMapping) +_COMPACT_RESPONSE: Final = TypeAdapter(CompactResponse) + + +class SpendEvent(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + + version: Literal[1] + litellm_call_id: str | None + call_type: str | None + model: str | None + custom_llm_provider: str | None + stream: bool | None + complete_streaming_response: bool + cache_hit: bool | None + response_cost: float | None + start_time: datetime + end_time: datetime + completion_start_time: datetime | None + agent_id: str | None + litellm_trace_id: str | None + litellm_params: _LitellmParams + standard_logging_object: ObjectMapping | None + standard_callback_dynamic_params: _DynamicParams | None + response: CompactResponse | None + combined_usage: ObjectMapping | None + realtime_tools: Sequence[object] | None + realtime_tool_calls: Sequence[object] | None + request_tool_names: tuple[str, ...] + + +@dataclass(frozen=True, slots=True) +class SpendEventCallbackArgs: + kwargs: SpendCallbackKwargs + response_obj: CompactResponse | None + start_time: datetime + end_time: datetime + + +@dataclass(frozen=True, slots=True) +class SpendEventBuildError: + reason: str + + +@dataclass(frozen=True, slots=True) +class SpendEventDecodeError: + reason: str + + +def is_offloadable_success(response_obj: object) -> bool: + """Batch retrieves and interaction polls branch on the concrete response class, so they stay in-process.""" + return not isinstance(response_obj, _NOT_OFFLOADED_RESPONSE_TYPES) + + +def _json_fallback(value: object) -> str: + return str(value) + + +def _mapping_or_none(value: object) -> ObjectMapping | None: + try: + return _OBJECT_MAPPING.validate_python(value) + except ValidationError: + return None + + +def _drop_keys(mapping: ObjectMapping, keys: frozenset[str]) -> ObjectMapping: + return MappingProxyType({key: value for key, value in mapping.items() if key not in keys}) + + +def _budget_reservation(metadata: ObjectMapping) -> ObjectMapping | None: + """The admission-time reservation, wherever the request setup left it, so the sidecar can reconcile it.""" + direct: Final = _mapping_or_none(metadata.get("user_api_key_budget_reservation")) + if direct is not None: + return direct + auth: Final = metadata.get("user_api_key_auth") + if isinstance(auth, UserAPIKeyAuth): + return auth.budget_reservation + auth_mapping: Final = _mapping_or_none(auth) + return _mapping_or_none(auth_mapping.get("budget_reservation")) if auth_mapping is not None else None + + +def _metadata_for_event( + metadata: ObjectMapping | None, budget_reservation: ObjectMapping | None +) -> ObjectMapping | None: + if metadata is None: + return None + kept: Final = _drop_keys(metadata, _UNSERIALIZABLE_METADATA_KEYS) + if budget_reservation is None: + return kept + return MappingProxyType({**kept, "user_api_key_budget_reservation": budget_reservation}) + + +def _litellm_params_for_event( + litellm_params: _LitellmParams, cache_key: str | None, store_bodies: bool +) -> _LitellmParams: + metadata: Final = litellm_params.get("metadata") + litellm_metadata: Final = litellm_params.get("litellm_metadata") + budget_reservation: Final = next( + ( + reservation + for source in (litellm_metadata, metadata) + if source is not None and (reservation := _budget_reservation(source)) is not None + ), + None, + ) + projected: Final[_LitellmParams] = { + "api_base": litellm_params.get("api_base"), + "custom_llm_provider": litellm_params.get("custom_llm_provider"), + "litellm_call_id": litellm_params.get("litellm_call_id"), + "user_api_key_end_user_id": litellm_params.get("user_api_key_end_user_id"), + "metadata": _metadata_for_event(metadata, budget_reservation), + "litellm_metadata": _metadata_for_event(litellm_metadata, budget_reservation), + "proxy_server_request": litellm_params.get("proxy_server_request") if store_bodies else None, + "preset_cache_key": cache_key, + } + return projected + + +def _standard_logging_for_event(sl_object: ObjectMapping | None, store_bodies: bool) -> ObjectMapping | None: + if sl_object is None: + return None + dropped: Final = ( + _STANDARD_LOGGING_DROPPED_KEYS if store_bodies else _STANDARD_LOGGING_DROPPED_KEYS | _STANDARD_LOGGING_BODY_KEYS + ) + return _drop_keys(sl_object, dropped) + + +def _tool_call(name: str) -> _ToolCall: + tool_call: Final[_ToolCall] = {"id": None, "type": "function", "function": {"name": name, "arguments": "{}"}} + return tool_call + + +def _tool_call_choice(names: Sequence[str]) -> _ToolCallChoice: + choice: Final[_ToolCallChoice] = { + "index": 0, + "finish_reason": "tool_calls", + "message": {"role": "assistant", "content": None, "tool_calls": tuple(_tool_call(name) for name in names)}, + } + return choice + + +def _compact_response(response_obj: object) -> CompactResponse | None: + """Usage, identity and tool calls of the response, in chat-completions shape, without the content.""" + dumped: Final = response_obj.model_dump() if isinstance(response_obj, BaseModel) else _mapping_or_none(response_obj) + if dumped is None: + return None + scalars: Final = _COMPACT_RESPONSE.validate_python(_drop_keys(dumped, frozenset({"choices"}))) + tool_call_names: Final = response_tool_call_names(response_obj) + if not tool_call_names: + return scalars + with_tool_calls: Final[CompactResponse] = {**scalars, "choices": (_tool_call_choice(tool_call_names),)} + return with_tool_calls + + +def _tool_name(tool: ObjectMapping) -> str | None: + """Chat tools nest the name under ``function``; Anthropic and Responses API tools keep it at the top.""" + function: Final = _mapping_or_none(tool.get("function")) + name: Final = function.get("name") if function is not None else tool.get("name") + return name.strip() if isinstance(name, str) and name.strip() else None + + +def _request_tool_names(kwargs: _SuccessKwargs) -> tuple[str, ...]: + passthrough: Final = kwargs.get("passthrough_logging_payload") + request_body: Final = passthrough.get("request_body") if passthrough is not None else None + passthrough_tools: Final = request_body.get("tools") if request_body is not None else None + return tuple( + name + for source in (kwargs.get("tools"), passthrough_tools) + if source is not None + for tool in source + if (name := _tool_name(tool)) is not None + ) + + +def preset_spend_log_cache_key(litellm_params: _LitellmParams) -> str | None: + """The key the caching layer already stored in ``litellm_params``, or ``Cache OFF``; never hashes the body.""" + if litellm.cache is None: + return CACHE_OFF_KEY + return litellm_params.get("preset_cache_key") + + +def _function_tool(name: str) -> _FunctionTool: + tool: Final[_FunctionTool] = {"type": "function", "function": {"name": name}} + return tool + + +def build_spend_event( + raw_kwargs: ObjectMapping, response_obj: object, start_time: datetime, end_time: datetime, store_bodies: bool +) -> bytes | SpendEventBuildError: + """Validate the success callback's kwargs and serialize the event once, as a single JSON line.""" + try: + kwargs: Final = _SUCCESS_KWARGS.validate_python(raw_kwargs) + except ValidationError as error: + return SpendEventBuildError(reason=str(error)) + litellm_params: Final = kwargs.get("litellm_params", _NO_LITELLM_PARAMS) + sl_object: Final = kwargs.get("standard_logging_object") + cache_key: Final = preset_spend_log_cache_key(litellm_params) + response_cost: Final = sl_object.get("response_cost") if sl_object is not None else kwargs.get("response_cost") + combined_usage: Final = kwargs.get("combined_usage_object") + event: Final = SpendEvent( + version=SPEND_EVENT_VERSION, + litellm_call_id=kwargs.get("litellm_call_id"), + call_type=kwargs.get("call_type"), + model=kwargs.get("model"), + custom_llm_provider=kwargs.get("custom_llm_provider"), + stream=kwargs.get("stream"), + complete_streaming_response="complete_streaming_response" in kwargs, + cache_hit=kwargs.get("cache_hit"), + response_cost=response_cost if isinstance(response_cost, (int, float)) else None, + start_time=start_time, + end_time=end_time, + completion_start_time=kwargs.get("completion_start_time"), + agent_id=kwargs.get("agent_id"), + litellm_trace_id=kwargs.get("litellm_trace_id"), + litellm_params=_litellm_params_for_event(litellm_params, cache_key, store_bodies), + standard_logging_object=_standard_logging_for_event(sl_object, store_bodies), + standard_callback_dynamic_params=kwargs.get("standard_callback_dynamic_params"), + response=_compact_response(response_obj), + combined_usage=combined_usage.model_dump() if combined_usage is not None else None, + realtime_tools=kwargs.get("realtime_tools"), + realtime_tool_calls=kwargs.get("realtime_tool_calls"), + request_tool_names=_request_tool_names(kwargs), + ) + return event.model_dump_json(fallback=_json_fallback).encode() + b"\n" + + +def decode_spend_event(line: bytes) -> SpendEvent | SpendEventDecodeError: + try: + return SpendEvent.model_validate_json(line) + except ValidationError as error: + return SpendEventDecodeError(reason=str(error)) + + +def spend_event_callback_args(event: SpendEvent) -> SpendEventCallbackArgs: + """The ``(kwargs, response_obj, start_time, end_time)`` the in-process cost callback receives.""" + tools: Final = tuple(_function_tool(name) for name in event.request_tool_names) + kwargs: Final[SpendCallbackKwargs] = { + "litellm_call_id": event.litellm_call_id, + "call_type": event.call_type, + "model": event.model, + "custom_llm_provider": event.custom_llm_provider, + "stream": event.stream, + "cache_hit": event.cache_hit, + "response_cost": event.response_cost, + "completion_start_time": event.completion_start_time, + "agent_id": event.agent_id, + "litellm_trace_id": event.litellm_trace_id, + "litellm_params": event.litellm_params, + "standard_logging_object": event.standard_logging_object, + "standard_callback_dynamic_params": event.standard_callback_dynamic_params, + "combined_usage_object": Usage.model_validate(event.combined_usage) + if event.combined_usage is not None + else None, + "realtime_tools": event.realtime_tools, + "realtime_tool_calls": event.realtime_tool_calls, + "tools": tools or None, + } + if not event.complete_streaming_response: + return SpendEventCallbackArgs(kwargs, event.response, event.start_time, event.end_time) + streaming_kwargs: Final[SpendCallbackKwargs] = {**kwargs, "complete_streaming_response": event.response} + return SpendEventCallbackArgs(streaming_kwargs, event.response, event.start_time, event.end_time) diff --git a/litellm/proxy/spend_tracking/spend_event_producer.py b/litellm/proxy/spend_tracking/spend_event_producer.py new file mode 100644 index 00000000000..20c7f177af0 --- /dev/null +++ b/litellm/proxy/spend_tracking/spend_event_producer.py @@ -0,0 +1,338 @@ +"""Fire-and-forget push of serialized spend events from an inference worker to the pod-local sidecar. + +``LITELLM_COLLECTOR_ENABLED=true`` turns the push on in the gateway; the sidecar process sets +``LITELLM_JOB_ROLE=collector`` and always runs the pipeline in-process. Events queue in a bounded +in-memory buffer that a single writer task flushes over a unix socket or loopback TCP connection. +When the sidecar is unreachable, the buffer is full, or the connection breaks mid-write, each affected +event follows ``LITELLM_COLLECTOR_ON_UNAVAILABLE``: ``fallback`` runs the existing cost pipeline in +the worker, ``drop`` counts it and moves on. Transitions are logged with the counters, so a sidecar +outage is visible without scraping anything. + +Delivery is at-most-once: a sidecar crash loses the events the kernel already took from its socket. +A sidecar that stops gracefully half-closes each connection first (EOF towards the producer) and +keeps reading until the producer hangs up, so the producer switches to the unavailable policy without +losing the events in flight. A write that fails part-way follows the unavailable policy without double +counting: ``drain()`` only fails while part of the line is still buffered in this process, so the +sidecar can at most have read a truncated line, which it discards. When the gateway itself stops with +the writer stuck mid-send, only an event whose bytes are still in the producer's write buffer follows +the unavailable policy; the connection is aborted first so the sidecar discards the truncated line +instead of also counting it. Events from one uvicorn worker are handled in the order it produced them; +events from different workers interleave, exactly like the in-process callbacks do today. +""" + +import asyncio +import ipaddress +import time +from collections.abc import Awaitable, Callable +from dataclasses import dataclass +from typing import Final, Literal, TypeAlias +from urllib.parse import urlsplit + +from pydantic import AliasChoices, Field +from pydantic_settings import BaseSettings, SettingsConfigDict + +from litellm._logging import verbose_proxy_logger + +COLLECTOR_ENV_PREFIX: Final = "LITELLM_COLLECTOR_" +COLLECTOR_JOB_ROLE: Final = "collector" +DEFAULT_COLLECTOR_ADDRESS: Final = "unix:///var/run/litellm/collector.sock" +RECONNECT_BACKOFF_SECONDS: Final = 1.0 +DROP_LOG_EVERY: Final = 1000 + +UnavailablePolicy: TypeAlias = Literal["fallback", "drop"] +PublishOutcome: TypeAlias = Literal["queued", "fallback", "dropped"] + + +class CollectorSettings(BaseSettings): + """``LITELLM_COLLECTOR_*`` env vars, shared by the gateway producer and the sidecar consumer.""" + + model_config = SettingsConfigDict( + env_prefix=COLLECTOR_ENV_PREFIX, case_sensitive=False, extra="ignore", frozen=True, populate_by_name=True + ) + + enabled: bool = False + address: str = DEFAULT_COLLECTOR_ADDRESS + buffer_size: int = Field(default=1000, ge=1) + on_unavailable: UnavailablePolicy = "fallback" + drain_timeout_seconds: float = Field(default=10.0, gt=0) + connect_timeout_seconds: float = Field(default=1.0, gt=0) + job_role: str | None = Field(default=None, validation_alias=AliasChoices("LITELLM_JOB_ROLE")) + + @property + def produces(self) -> bool: + return self.enabled and self.job_role != COLLECTOR_JOB_ROLE + + +@dataclass(frozen=True, slots=True) +class UnixAddress: + path: str + + +@dataclass(frozen=True, slots=True) +class TcpAddress: + host: str + port: int + + +@dataclass(frozen=True, slots=True) +class AddressError: + reason: str + + +CollectorAddress: TypeAlias = UnixAddress | TcpAddress + + +def _is_loopback(host: str) -> bool: + try: + return ipaddress.ip_address(host).is_loopback + except ValueError: + return host == "localhost" + + +def parse_collector_address(address: str) -> CollectorAddress | AddressError: + """``unix:///path/to.sock`` or ``tcp://127.0.0.1:port``; the socket carries unauthenticated spend events.""" + parsed: Final = urlsplit(address) + if parsed.scheme == "unix" and parsed.path: + return UnixAddress(path=parsed.path) + if parsed.scheme == "tcp" and parsed.hostname and parsed.port is not None: + if not _is_loopback(parsed.hostname): + return AddressError(reason=f"tcp collector address must be a loopback host, got {address!r}") + return TcpAddress(host=parsed.hostname, port=parsed.port) + return AddressError(reason=f"expected unix:///path or tcp://127.0.0.1:port, got {address!r}") + + +async def open_collector_connection( + address: CollectorAddress, timeout: float +) -> tuple[asyncio.StreamReader, asyncio.StreamWriter]: + match address: + case UnixAddress(path=path): + return await asyncio.wait_for(asyncio.open_unix_connection(path), timeout) + case TcpAddress(host=host, port=port): + return await asyncio.wait_for(asyncio.open_connection(host, port), timeout) + + +def build_spend_event_producer( + settings: CollectorSettings, fallback: Callable[[bytes], Awaitable[None]] +) -> "SpendEventProducer | None": + """The gateway producer for these settings, or ``None`` when the pipeline stays in-process.""" + if not settings.produces: + return None + address: Final = parse_collector_address(settings.address) + if isinstance(address, AddressError): + verbose_proxy_logger.error("collector: %s; running the spend pipeline in-process", address.reason) + return None + verbose_proxy_logger.info( + "collector: offloading spend tracking to %s (buffer=%d, on_unavailable=%s)", + settings.address, + settings.buffer_size, + settings.on_unavailable, + ) + return SpendEventProducer( + address=address, + on_unavailable=settings.on_unavailable, + buffer_size=settings.buffer_size, + connect_timeout=settings.connect_timeout_seconds, + fallback=fallback, + ) + + +@dataclass(frozen=True, slots=True) +class _Connection: + reader: asyncio.StreamReader + writer: asyncio.StreamWriter + + @property + def alive(self) -> bool: + return not self.writer.is_closing() and not self.reader.at_eof() + + +@dataclass(frozen=True, slots=True) +class SpendEventProducerStats: + queued: int + sent: int + fallback: int + dropped: int + connected: bool + + +class SpendEventProducer: + """Bounded buffer plus one writer task per process; see the module docstring for the contract.""" + + def __init__( + self, + address: CollectorAddress, + on_unavailable: UnavailablePolicy, + buffer_size: int, + connect_timeout: float, + fallback: Callable[[bytes], Awaitable[None]], + clock: Callable[[], float] = time.monotonic, + open_connection: Callable[ + [CollectorAddress, float], Awaitable[tuple[asyncio.StreamReader, asyncio.StreamWriter]] + ] = open_collector_connection, + ) -> None: + self._address = address + self._on_unavailable = on_unavailable + self._buffer_size = buffer_size + self._connect_timeout = connect_timeout + self._fallback = fallback + self._clock = clock + self._open_connection = open_connection + self._queue: asyncio.Queue[bytes] | None = None + self._writer_task: asyncio.Task[None] | None = None + self._connection: _Connection | None = None + self._in_flight: bytes | None = None + self._closing = False + self._next_connect_at = 0.0 + self._queued = 0 + self._sent = 0 + self._fallback_count = 0 + self._dropped = 0 + + def stats(self) -> SpendEventProducerStats: + return SpendEventProducerStats( + queued=self._queued, + sent=self._sent, + fallback=self._fallback_count, + dropped=self._dropped, + connected=self._connection is not None, + ) + + async def publish(self, line: bytes) -> PublishOutcome: + """Hand one serialized event to the writer task, or apply the unavailable policy right away.""" + if self._closing or self._clock() < self._next_connect_at: + return await self._unavailable(line, "sidecar unreachable") + queue: Final = self._ensure_writer() + try: + queue.put_nowait(line) + except asyncio.QueueFull: + return await self._unavailable(line, "buffer full") + self._queued += 1 + return "queued" + + async def close(self, drain_timeout: float) -> None: + """Flush the buffer for up to ``drain_timeout`` seconds, then apply the unavailable policy to the rest.""" + self._closing = True + queue: Final = self._queue + task: Final = self._writer_task + if queue is None or task is None: + return + try: + await asyncio.wait_for(queue.join(), drain_timeout) + except asyncio.TimeoutError: + verbose_proxy_logger.warning( + "collector: %s events still buffered after %.1fs drain timeout", queue.qsize(), drain_timeout + ) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + unsent: Final = self._take_unsent() + await self._disconnect() + if unsent is not None: + await self._unavailable(unsent, "shutdown") + while not queue.empty(): + await self._unavailable(queue.get_nowait(), "shutdown") + + def _take_unsent(self) -> bytes | None: + """The in-flight event if any of its bytes never left this process, aborting the half-written connection.""" + in_flight: Final = self._in_flight + self._in_flight = None + connection: Final = self._connection + if in_flight is None: + return None + if connection is None: + return in_flight + if connection.writer.transport.get_write_buffer_size() == 0: + return None + connection.writer.transport.abort() + return in_flight + + def _ensure_writer(self) -> asyncio.Queue[bytes]: + if self._queue is None: + self._queue = asyncio.Queue(maxsize=self._buffer_size) + if self._writer_task is None or self._writer_task.done(): + self._writer_task = asyncio.get_running_loop().create_task(self._run_writer(self._queue)) + return self._queue + + async def _run_writer(self, queue: asyncio.Queue[bytes]) -> None: + while True: + line = await queue.get() + try: + await self._send(line) + finally: + queue.task_done() + + async def _send(self, line: bytes) -> None: + self._in_flight = line + connection: Final = await self._connect() + if connection is None: + self._in_flight = None + await self._unavailable(line, "sidecar unreachable") + return + try: + connection.writer.write(line) + await connection.writer.drain() + except (ConnectionError, OSError, RuntimeError) as error: # uvloop: RuntimeError on a closed transport + self._in_flight = None + await self._disconnect() + self._next_connect_at = self._clock() + RECONNECT_BACKOFF_SECONDS + await self._unavailable(line, f"write failed: {error}") + return + self._in_flight = None + self._sent += 1 + + async def _connect(self) -> _Connection | None: + if self._connection is not None and self._connection.alive: + return self._connection + await self._disconnect() + if self._clock() < self._next_connect_at: + return None + try: + reader, writer = await self._open_connection(self._address, self._connect_timeout) + except (ConnectionError, OSError, asyncio.TimeoutError) as error: + self._next_connect_at = self._clock() + RECONNECT_BACKOFF_SECONDS + verbose_proxy_logger.warning( + "collector: cannot reach %s (%s); applying %s policy for %.0fs. stats=%s", + self._address, + error, + self._on_unavailable, + RECONNECT_BACKOFF_SECONDS, + self.stats(), + ) + return None + self._connection = _Connection(reader=reader, writer=writer) + verbose_proxy_logger.info("collector: connected to %s. stats=%s", self._address, self.stats()) + return self._connection + + async def _disconnect(self) -> None: + connection: Final = self._connection + self._connection = None + if connection is None: + return + connection.writer.close() + try: + await connection.writer.wait_closed() + except (ConnectionError, OSError): + pass + + async def _unavailable(self, line: bytes, reason: str) -> PublishOutcome: + if self._on_unavailable == "fallback": + self._fallback_count += 1 + fallback: Final = asyncio.ensure_future(self._run_fallback(line, reason)) + try: + await asyncio.shield(fallback) + except asyncio.CancelledError: + await fallback + raise + return "fallback" + self._dropped += 1 + if self._dropped % DROP_LOG_EVERY == 1: + verbose_proxy_logger.warning("collector: dropping spend event (%s). stats=%s", reason, self.stats()) + return "dropped" + + async def _run_fallback(self, line: bytes, reason: str) -> None: + try: + await self._fallback(line) + except Exception: # noqa: BLE001 # one failing event must not kill the writer task + verbose_proxy_logger.exception("collector: in-process fallback failed (%s)", reason) diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index dad9a2a8777..f0f38358cf0 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -554,10 +554,12 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs additional_usage_values["cache_creation_input_tokens"] = cache_write_tokens clean_metadata["additional_usage_values"] = additional_usage_values - if litellm.cache is not None: - cache_key = litellm.cache.get_cache_key(**kwargs) - else: + if litellm.cache is None: cache_key = "Cache OFF" + elif litellm_params.get("preset_cache_key") is not None: + cache_key = litellm_params["preset_cache_key"] + else: + cache_key = litellm.cache.get_cache_key(**kwargs) if cache_hit is True: import time @@ -864,7 +866,7 @@ def _get_messages_for_spend_logs_payload( standard_logging_payload: StandardLoggingPayload | None, metadata: dict | None = None, ) -> str: - if _should_store_prompts_and_responses_in_spend_logs(): + if should_store_prompts_and_responses_in_spend_logs(): if standard_logging_payload is not None: call_type: Final = standard_logging_payload.get("call_type", "") if call_type == "_arealtime": @@ -1114,7 +1116,7 @@ def _sanitize_guardrail_information_for_spend_logs( here to match OTEL's defensive read pattern; otherwise iteration would yield the dict's keys and crash the whole spend-log write. """ - if guardrail_information is None or _should_store_prompts_and_responses_in_spend_logs(): + if guardrail_information is None or should_store_prompts_and_responses_in_spend_logs(): return guardrail_information entries: Final = [guardrail_information] if isinstance(guardrail_information, dict) else guardrail_information return [_redact_prompt_fields_in_guardrail_entry(entry) for entry in entries if isinstance(entry, dict)] @@ -1186,7 +1188,7 @@ def _sanitize_error_information_for_spend_logs( sanitized = cast(dict, {**error_information}) - if not _should_store_prompts_and_responses_in_spend_logs(): + if not should_store_prompts_and_responses_in_spend_logs(): for field in ("error_message", "traceback"): value = sanitized.get(field) if isinstance(value, str): @@ -1263,11 +1265,11 @@ def _get_proxy_server_request_for_spend_logs_payload( kwargs: dict | None = None, ) -> str: """ - Only store if _should_store_prompts_and_responses_in_spend_logs() is True + Only store if should_store_prompts_and_responses_in_spend_logs() is True If turn_off_message_logging is enabled, redact messages in the request body. """ - if _should_store_prompts_and_responses_in_spend_logs(): + if should_store_prompts_and_responses_in_spend_logs(): _proxy_server_request: Final = cast(dict | None, litellm_params.get("proxy_server_request", EMPTY_MAPPING)) if _proxy_server_request is not None: _request_body = _proxy_server_request.get("body", EMPTY_MAPPING) or EMPTY_MAPPING @@ -1317,7 +1319,7 @@ def _get_vector_store_request_for_spend_logs_payload( """ If user does not want to store prompts and responses, then remove the content from the vector store request metadata """ - if _should_store_prompts_and_responses_in_spend_logs(): + if should_store_prompts_and_responses_in_spend_logs(): return vector_store_request_metadata # if user does not want to store prompts and responses, then remove the content from the vector store request metadata @@ -1341,7 +1343,7 @@ def _get_response_for_spend_logs_payload( ) -> str: if payload is None: return "{}" - if _should_store_prompts_and_responses_in_spend_logs(): + if should_store_prompts_and_responses_in_spend_logs(): response_obj: object = payload.get("response") if response_obj is None: return "{}" @@ -1389,7 +1391,7 @@ def _get_response_for_spend_logs_payload( return "{}" -def _should_store_prompts_and_responses_in_spend_logs() -> bool: +def should_store_prompts_and_responses_in_spend_logs() -> bool: from litellm.proxy.proxy_server import general_settings from litellm.secret_managers.main import get_secret_bool diff --git a/litellm/utils.py b/litellm/utils.py index 69ee730ff30..aced4c9b312 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1843,6 +1843,9 @@ def client(original_function): elif _caching_handler_response.embedding_all_elements_cache_hit is True: return _caching_handler_response.final_embedding_cached_response + if _llm_caching_handler.preset_cache_key is not None: + logging_obj.litellm_params["preset_cache_key"] = _llm_caching_handler.preset_cache_key + # CHECK MAX TOKENS if ( kwargs.get("max_tokens", None) is not None diff --git a/terraform/litellm/README.md b/terraform/litellm/README.md index d4b40741052..7cc6e4c08ba 100644 --- a/terraform/litellm/README.md +++ b/terraform/litellm/README.md @@ -183,6 +183,7 @@ only where the underlying cloud forces it. | Extra secret-backed env | `gateway_extra_secrets`, `backend_extra_secrets` (ARNs) | `gateway_extra_secrets`, `backend_extra_secrets` (resource IDs) | | Uvicorn `--workers` on gateway | `gateway_num_workers` | `gateway_num_workers` | | OpenTelemetry v2 (opt-in) | `otel_endpoint`, `otel_exporter`, `otel_environment_name`, `otel_capture_message_content`, `otel_headers_secret_arn` | `otel_endpoint`, `otel_exporter`, `otel_environment_name`, `otel_capture_message_content`, `otel_headers_secret` | +| Collector sidecar (opt-in) | `collector_enabled`, `collector_port`, `collector_cpu`, `collector_memory`, `collector_buffer_size`, `collector_on_unavailable`, `collector_drain_timeout_seconds` | same names; `collector_cpu` / `collector_memory` take Cloud Run strings | Each module stamps its own stack-identity tag (`litellm:stack` on AWS, `litellm-stack` on GCP — GCP label keys forbid colons) plus diff --git a/terraform/litellm/aws/README.md b/terraform/litellm/aws/README.md index 5ca45944c54..67f1270d0c2 100644 --- a/terraform/litellm/aws/README.md +++ b/terraform/litellm/aws/README.md @@ -345,6 +345,39 @@ ten tasks handle 4,200,000,000 tokens in a minute, `tokens / 60` is `ceil(10 * 7000000 / 6000000) = 12`. Container Insights must be enabled on the cluster for `RunningTaskCount` to exist +### Collector sidecar + +`collector_enabled = true` adds a second container to the gateway task +that runs `python -m litellm.proxy.collector` from the gateway image, and sets +`LITELLM_COLLECTOR_ENABLED=true` on the gateway so its uvicorn workers +ship spend events (SpendLogs writes, key/team/user spend updates, budget +alerts) to the sidecar instead of running that pipeline in the request +path. This is the Terraform counterpart of helm's `gateway.collector`. +The default (`false`) leaves the task definition exactly as before. + +Fargate tasks share one network namespace, so the sidecar listens on +loopback TCP (`tcp://127.0.0.1:${collector_port}`, default 4010) instead +of the Unix socket helm uses; the proxy rejects any non-loopback address. +The sidecar gets the same database, Redis, master-key, license, proxy +config, and `gateway_extra_env` / `gateway_extra_secrets` values as the +gateway container, runs with `LITELLM_JOB_ROLE=collector`, and is +non-essential with an ECS restart policy, so a sidecar crash restarts it in +place while the gateway falls back to in-process spend tracking. + +```hcl +collector_enabled = true +# collector_cpu = 512 # carved out of gateway_cpu +# collector_memory = 2048 # MiB, carved out of gateway_memory +# collector_buffer_size = 1000 +# collector_on_unavailable = "fallback" # or "drop" +# collector_drain_timeout_seconds = 10 +``` + +Both sidecar reservations must leave room for the gateway container inside +`gateway_cpu` / `gateway_memory` (the plan fails otherwise). Service +autoscaling keeps tracking the whole task's CPU and memory, sidecar +included + ## Tenant deployment Every resource the stack creates is named `${tenant}-litellm-${env}` (or diff --git a/terraform/litellm/aws/ecs.tf b/terraform/litellm/aws/ecs.tf index 7176921e8fa..2b235c2bad5 100644 --- a/terraform/litellm/aws/ecs.tf +++ b/terraform/litellm/aws/ecs.tf @@ -278,6 +278,62 @@ locals { "${local.proxy_config_fetch_cmd} && ${local.backend_launch_cmd}" ] } : {} + + collector_address = "tcp://127.0.0.1:${var.collector_port}" + collector_env = var.collector_enabled ? [ + { name = "LITELLM_COLLECTOR_ENABLED", value = "true" }, + { name = "LITELLM_COLLECTOR_ADDRESS", value = local.collector_address }, + { name = "LITELLM_COLLECTOR_BUFFER_SIZE", value = tostring(var.collector_buffer_size) }, + { name = "LITELLM_COLLECTOR_ON_UNAVAILABLE", value = var.collector_on_unavailable }, + { name = "LITELLM_COLLECTOR_DRAIN_TIMEOUT_SECONDS", value = tostring(var.collector_drain_timeout_seconds) }, + ] : [] + + gateway_environment = concat( + local.shared_env, + local.gateway_otel_env, + local.billing_metrics_env, + local.gateway_extra_env_list, + local.proxy_config_env, + local.metrics_env, + local.gateway_pool_env, + local.collector_env, + ) + + collector_launch_cmd = "exec python -m litellm.proxy.collector" + collector_command = [ + local.proxy_config_enabled ? "${local.proxy_config_fetch_cmd} && ${local.collector_launch_cmd}" : local.collector_launch_cmd + ] + + collector_container = var.collector_enabled ? [{ + name = "collector" + image = var.gateway_image + essential = false + cpu = var.collector_cpu + memory = var.collector_memory + + restartPolicy = { enabled = true } + + entryPoint = ["sh", "-c"] + command = local.collector_command + environment = concat( + local.shared_env, + local.gateway_extra_env_list, + local.proxy_config_env, + local.gateway_pool_env, + local.collector_env, + [{ name = "LITELLM_JOB_ROLE", value = "collector" }], + ) + secrets = concat(local.shared_secrets, local.gateway_extra_secrets_list) + + logConfiguration = { + logDriver = "awslogs" + options = { + awslogs-group = aws_cloudwatch_log_group.gateway.name + awslogs-region = var.region + awslogs-stream-prefix = "collector" + } + } + }] : [] } # ---------- Gateway ---------- @@ -309,6 +365,16 @@ resource "aws_ecs_task_definition" "gateway" { condition = !var.gateway_connection_pool_enabled || local.database_enabled error_message = "gateway_connection_pool_enabled needs a database: set create_database = true or pass database_url." } + + precondition { + condition = !var.collector_enabled || (var.collector_cpu < var.gateway_cpu && var.collector_memory < var.gateway_memory) + error_message = "collector_cpu and collector_memory are carved out of gateway_cpu / gateway_memory and must leave room for the gateway container." + } + + precondition { + condition = !var.collector_enabled || var.gateway_metrics_port == null || var.collector_port != var.gateway_metrics_port + error_message = "collector_port and gateway_metrics_port must differ: both sidecars bind loopback in the same task." + } } family = "${local.name}-gateway" @@ -327,17 +393,9 @@ resource "aws_ecs_task_definition" "gateway" { essential = true portMappings = [{ containerPort = 4000, protocol = "tcp" }] - environment = concat( - local.shared_env, - local.gateway_otel_env, - local.billing_metrics_env, - local.gateway_extra_env_list, - local.proxy_config_env, - local.metrics_env, - local.gateway_pool_env, - ) - secrets = concat(local.shared_secrets, local.gateway_extra_secrets_list) - mountPoints = local.metrics_mount_points + environment = local.gateway_environment + secrets = concat(local.shared_secrets, local.gateway_extra_secrets_list) + mountPoints = local.metrics_mount_points # Container-level healthCheck intentionally omitted — the wolfi # runtime image doesn't ship curl/wget. The ALB target group polls @@ -354,7 +412,7 @@ resource "aws_ecs_task_definition" "gateway" { }, local.gateway_proxy_overrides, ) - ], local.gateway_metrics_container)) + ], local.gateway_metrics_container, local.collector_container)) dynamic "volume" { for_each = local.metrics_enabled ? [1] : [] diff --git a/terraform/litellm/aws/tests/collector.tftest.hcl b/terraform/litellm/aws/tests/collector.tftest.hcl new file mode 100644 index 00000000000..1465130232f --- /dev/null +++ b/terraform/litellm/aws/tests/collector.tftest.hcl @@ -0,0 +1,144 @@ +# Plan-only coverage for the opt-in collector sidecar in the gateway task. +# The rendered container_definitions JSON is unknown at plan time (it embeds +# Aurora/ElastiCache endpoints and secret ARNs), so the assertions target the +# locals it is built from. Run from terraform/litellm/aws with `terraform test`. + +mock_provider "aws" { + mock_data "aws_iam_policy_document" { + defaults = { + json = "{\"Version\":\"2012-10-17\",\"Statement\":[]}" + } + } +} +mock_provider "random" {} + +variables { + region = "us-east-1" + tenant = "acme" + env = "test" + allow_plaintext_alb = true + azs = ["us-east-1a", "us-east-1b"] +} + +run "disabled_by_default_leaves_the_task_untouched" { + command = plan + + assert { + condition = length(local.collector_container) == 0 + error_message = "The gateway task must stay single-container unless collector_enabled is set." + } + + assert { + condition = !anytrue([for e in local.gateway_environment : startswith(e.name, "LITELLM_COLLECTOR_")]) + error_message = "No LITELLM_COLLECTOR_* env may reach the gateway while the sidecar is disabled." + } +} + +run "enabled_adds_a_sidecar_that_shares_the_gateway_transport" { + command = plan + + variables { + collector_enabled = true + collector_port = 4321 + collector_buffer_size = 250 + collector_on_unavailable = "drop" + gateway_extra_env = { OPENAI_API_BASE = "https://example.invalid" } + gateway_extra_secrets = { OPENAI_API_KEY = "arn:aws:secretsmanager:us-east-1:111122223333:secret:openai-AbCdEf" } + } + + assert { + condition = length(local.collector_container) == 1 && local.collector_container[0].name == "collector" + error_message = "Enabling the sidecar must add exactly one collector container." + } + + assert { + condition = alltrue([ + for env in [local.gateway_environment, local.collector_container[0].environment] : ( + { for e in env : e.name => e.value }["LITELLM_COLLECTOR_ENABLED"] == "true" && + { for e in env : e.name => e.value }["LITELLM_COLLECTOR_ADDRESS"] == "tcp://127.0.0.1:4321" && + { for e in env : e.name => e.value }["LITELLM_COLLECTOR_BUFFER_SIZE"] == "250" && + { for e in env : e.name => e.value }["LITELLM_COLLECTOR_ON_UNAVAILABLE"] == "drop" && + { for e in env : e.name => e.value }["LITELLM_COLLECTOR_DRAIN_TIMEOUT_SECONDS"] == "10" + ) + ]) + error_message = "Gateway and sidecar must agree on the loopback address and the collector knobs." + } + + assert { + condition = ( + local.collector_container[0].image == var.gateway_image && + local.collector_container[0].entryPoint == ["sh", "-c"] && + local.collector_container[0].command == ["exec python -m litellm.proxy.collector"] && + local.collector_container[0].essential == false && + local.collector_container[0].restartPolicy.enabled == true && + { for e in local.collector_container[0].environment : e.name => e.value }["LITELLM_JOB_ROLE"] == "collector" + ) + error_message = "The sidecar must run litellm.proxy.collector from the gateway image as a restartable, non-essential collector." + } + + assert { + condition = ( + { for e in local.collector_container[0].environment : e.name => e.value }["OPENAI_API_BASE"] == "https://example.invalid" && + contains([for e in local.collector_container[0].environment : e.name], "DATABASE_HOST") && + contains([for e in local.collector_container[0].environment : e.name], "REDIS_HOST") && + contains([for s in local.collector_container[0].secrets : s.name], "LITELLM_MASTER_KEY") && + contains([for s in local.collector_container[0].secrets : s.name], "OPENAI_API_KEY") + ) + error_message = "The sidecar must receive the gateway's database, Redis, and shared secrets plus gateway_extra_env / gateway_extra_secrets." + } + + assert { + condition = !contains(keys(local.collector_container[0]), "portMappings") + error_message = "The sidecar must not expose a port to the task's load balancer." + } + + assert { + condition = local.collector_container[0].cpu == 512 && local.collector_container[0].memory == 2048 + error_message = "The sidecar defaults must mirror helm's collector resources (500m / 2Gi)." + } +} + +run "proxy_config_is_fetched_by_the_sidecar_too" { + command = plan + + variables { + collector_enabled = true + proxy_config = { model_list = [] } + } + + assert { + condition = ( + startswith(local.collector_container[0].command[0], local.proxy_config_fetch_cmd) && + endswith(local.collector_container[0].command[0], "exec python -m litellm.proxy.collector") && + contains([for e in local.collector_container[0].environment : e.name], "CONFIG_FILE_PATH") + ) + error_message = "The sidecar must pull the proxy config from S3 before starting, like the gateway does." + } +} + +run "sidecar_must_leave_room_for_the_gateway" { + command = plan + + variables { + collector_enabled = true + collector_cpu = 1024 + } + + expect_failures = [ + aws_ecs_task_definition.gateway, + ] +} + +run "sidecars_must_not_share_a_loopback_port" { + command = plan + + variables { + collector_enabled = true + collector_port = 4001 + gateway_metrics_port = 4001 + } + + expect_failures = [ + aws_ecs_task_definition.gateway, + ] +} diff --git a/terraform/litellm/aws/variables.tf b/terraform/litellm/aws/variables.tf index f57bb1ab7f6..580a0cc657a 100644 --- a/terraform/litellm/aws/variables.tf +++ b/terraform/litellm/aws/variables.tf @@ -808,3 +808,73 @@ variable "billing_metrics_ca_cert_pem" { default = "" sensitive = true } + +# ---------- Collector sidecar ---------- +# +# Opt-in offload of spend tracking from the gateway's uvicorn workers to a +# `python -m litellm.proxy.collector` sidecar in the same Fargate task (helm's +# `gateway.collector`). Fargate awsvpc tasks share one network namespace, +# so the sidecar listens on loopback TCP. Disabled (the default) adds nothing +# to the task definition. + +variable "collector_enabled" { + description = "Run the collector sidecar next to the gateway container and have the gateway ship spend events to it (sets LITELLM_COLLECTOR_ENABLED=true on both). Autoscaling still targets the whole task's CPU/memory, sidecar included." + type = bool + default = false +} + +variable "collector_port" { + description = "Loopback TCP port the sidecar listens on (LITELLM_COLLECTOR_ADDRESS=tcp://127.0.0.1:)." + type = number + default = 4010 + + validation { + condition = var.collector_port >= 1024 && var.collector_port <= 65535 && var.collector_port != 4000 + error_message = "collector_port must be in 1024-65535 and not 4000." + } +} + +variable "collector_cpu" { + description = "CPU units reserved for the sidecar container, carved out of gateway_cpu. Matches helm's collector.resources.requests.cpu (500m)." + type = number + default = 512 +} + +variable "collector_memory" { + description = "Hard memory limit (MiB) for the sidecar container, carved out of gateway_memory. Matches helm's collector.resources.limits.memory (2Gi)." + type = number + default = 2048 +} + +variable "collector_buffer_size" { + description = "Per-worker in-memory queue of spend events waiting to be shipped to the sidecar (LITELLM_COLLECTOR_BUFFER_SIZE)." + type = number + default = 1000 + + validation { + condition = var.collector_buffer_size >= 1 + error_message = "collector_buffer_size must be >= 1." + } +} + +variable "collector_on_unavailable" { + description = "What the gateway does with spend events when the sidecar is unreachable or the buffer is full (LITELLM_COLLECTOR_ON_UNAVAILABLE): `fallback` runs the pipeline in-process, `drop` discards them." + type = string + default = "fallback" + + validation { + condition = contains(["fallback", "drop"], var.collector_on_unavailable) + error_message = "collector_on_unavailable must be one of: fallback, drop." + } +} + +variable "collector_drain_timeout_seconds" { + description = "Seconds a gateway worker waits on shutdown for its buffered spend events to reach the sidecar (LITELLM_COLLECTOR_DRAIN_TIMEOUT_SECONDS)." + type = number + default = 10 + + validation { + condition = var.collector_drain_timeout_seconds > 0 + error_message = "collector_drain_timeout_seconds must be > 0." + } +} diff --git a/terraform/litellm/gcp/README.md b/terraform/litellm/gcp/README.md index 9c71d3e15b7..23cfe9afc24 100644 --- a/terraform/litellm/gcp/README.md +++ b/terraform/litellm/gcp/README.md @@ -321,6 +321,43 @@ launcher reads these variables, starts the pooler once per instance before uvicorn forks the workers and hands them its loopback `DATABASE_URL`. It also honours `KEEPALIVE_TIMEOUT` from `gateway_extra_env` the way the image does +### Collector sidecar + +`collector_enabled = true` adds a `spend-collector` container to the gateway +Cloud Run service that runs `python -m litellm.proxy.collector` from the gateway +image, and sets `LITELLM_COLLECTOR_ENABLED=true` on the gateway so its +uvicorn workers ship spend events (SpendLogs writes, key/team/user spend +updates, budget alerts) to the sidecar instead of running that pipeline in +the request path. This is the Terraform counterpart of helm's +`gateway.collector`. The default (`false`) leaves the service exactly as +before. It is independent of the metrics sidecars above, whose GMP scraper +already owns the `collector` container name. + +Containers in one Cloud Run instance share localhost, so the sidecar listens +on loopback TCP (`tcp://127.0.0.1:${collector_port}`, default 4010) +instead of the Unix socket helm uses; the proxy rejects any non-loopback +address. The sidecar runs the same Redis CA + `DATABASE_URL` bootstrap as +the gateway container, gets the same database, Redis, master-key, license, +proxy config, and `gateway_extra_env` / `gateway_extra_secrets` values, and +runs with `LITELLM_JOB_ROLE=collector`. When it is unreachable the +gateway falls back to in-process spend tracking. + +```hcl +collector_enabled = true +# collector_cpu = "1000m" # added on top of gateway_cpu +# collector_memory = "2Gi" # added on top of gateway_memory +# collector_buffer_size = 1000 +# collector_on_unavailable = "fallback" # or "drop" +# collector_drain_timeout_seconds = 10 +``` + +Cloud Run allocates CPU per instance while requests are in flight, and the +sidecar shares that allocation. Spend events are shipped right after each +response, so this works with request-based billing, but keep +`gateway_min_instances >= 1` if spend must keep draining while an instance +is otherwise idle. Variable names match the AWS stack; only the resource +units differ (Cloud Run strings vs Fargate units) + ## Tenant deployment Every resource the stack creates is named `${tenant}-litellm-${env}` (or diff --git a/terraform/litellm/gcp/cloudrun.tf b/terraform/litellm/gcp/cloudrun.tf index 78d4ffa2152..d0b32a367d6 100644 --- a/terraform/litellm/gcp/cloudrun.tf +++ b/terraform/litellm/gcp/cloudrun.tf @@ -177,6 +177,34 @@ locals { [local.backend_launch_cmd], )) + collector_address = "tcp://127.0.0.1:${var.collector_port}" + collector_env_kv = var.collector_enabled ? [ + { name = "LITELLM_COLLECTOR_ENABLED", value = "true" }, + { name = "LITELLM_COLLECTOR_ADDRESS", value = local.collector_address }, + { name = "LITELLM_COLLECTOR_BUFFER_SIZE", value = tostring(var.collector_buffer_size) }, + { name = "LITELLM_COLLECTOR_ON_UNAVAILABLE", value = var.collector_on_unavailable }, + { name = "LITELLM_COLLECTOR_DRAIN_TIMEOUT_SECONDS", value = tostring(var.collector_drain_timeout_seconds) }, + ] : [] + + gateway_env_kv = concat(local.shared_env_kv, local.gateway_otel_env_kv, local.billing_metrics_env_kv, local.gateway_extra_env_kv, local.proxy_config_env, local.metrics_env_kv, local.gateway_pool_env, local.collector_env_kv) + gateway_env_secrets = concat(local.shared_env_secrets, local.otel_env_secrets, local.billing_metrics_env_secrets, local.gateway_extra_secret_kv) + + collector_env_kv_all = concat( + local.shared_env_kv, + local.gateway_extra_env_kv, + local.proxy_config_env, + local.gateway_pool_env, + local.collector_env_kv, + [{ name = "LITELLM_JOB_ROLE", value = "collector" }], + ) + collector_env_secrets = concat(local.shared_env_secrets, local.gateway_extra_secret_kv) + + collector_args = join(" && ", concat( + local.redis_ca_fragment, + local.database_url_fragment, + ["exec python -m litellm.proxy.collector"], + )) + # Env shipped to the migrations Job. The migrations image runs run.py # which assembles DATABASE_URL from these discrete vars itself, so we # only need writer-side DB env (no read replica, no proxy_config, no @@ -203,6 +231,13 @@ resource "google_cloud_run_v2_service" "gateway" { labels = local.labels deletion_protection = false + lifecycle { + precondition { + condition = !var.collector_enabled || var.gateway_metrics_port == null || var.collector_port != var.gateway_metrics_port + error_message = "collector_port and gateway_metrics_port must differ: both sidecars bind loopback in the same instance." + } + } + template { service_account = google_service_account.runtime.email max_instance_request_concurrency = var.gateway_max_instance_request_concurrency @@ -235,7 +270,7 @@ resource "google_cloud_run_v2_service" "gateway" { } dynamic "env" { - for_each = concat(local.shared_env_kv, local.gateway_otel_env_kv, local.billing_metrics_env_kv, local.gateway_extra_env_kv, local.proxy_config_env, local.metrics_env_kv, local.gateway_pool_env) + for_each = local.gateway_env_kv content { name = env.value.name value = env.value.value @@ -243,7 +278,7 @@ resource "google_cloud_run_v2_service" "gateway" { } dynamic "env" { - for_each = concat(local.shared_env_secrets, local.otel_env_secrets, local.billing_metrics_env_secrets, local.gateway_extra_secret_kv) + for_each = local.gateway_env_secrets content { name = env.value.name value_source { @@ -357,6 +392,52 @@ resource "google_cloud_run_v2_service" "gateway" { } } + dynamic "containers" { + for_each = var.collector_enabled ? [1] : [] + content { + name = "spend-collector" + image = local.gateway_image + command = ["sh", "-c"] + args = [local.collector_args] + + resources { + limits = { + cpu = var.collector_cpu + memory = var.collector_memory + } + } + + dynamic "env" { + for_each = local.collector_env_kv_all + content { + name = env.value.name + value = env.value.value + } + } + + dynamic "env" { + for_each = local.collector_env_secrets + content { + name = env.value.name + value_source { + secret_key_ref { + secret = env.value.secret + version = env.value.version + } + } + } + } + + dynamic "volume_mounts" { + for_each = local.proxy_config_enabled ? [1] : [] + content { + name = local.proxy_config_volume + mount_path = local.proxy_config_mount_path + } + } + } + } + dynamic "volumes" { for_each = local.proxy_config_enabled ? [1] : [] content { diff --git a/terraform/litellm/gcp/tests/collector.tftest.hcl b/terraform/litellm/gcp/tests/collector.tftest.hcl new file mode 100644 index 00000000000..7a5ea781acb --- /dev/null +++ b/terraform/litellm/gcp/tests/collector.tftest.hcl @@ -0,0 +1,165 @@ +# Plan-only coverage for the opt-in collector sidecar on the gateway Cloud +# Run service. `mock_provider` keeps this offline: no GCP credentials, no API +# calls. Run from terraform/litellm/gcp with `terraform test`. + +mock_provider "google" {} +mock_provider "google-beta" {} +mock_provider "random" {} + +variables { + project_id = "acme-test" + region = "us-central1" + tenant = "acme" + env = "test" + allow_plaintext_lb = true +} + +run "disabled_by_default_leaves_the_service_untouched" { + command = plan + + assert { + condition = [for c in google_cloud_run_v2_service.gateway[0].template[0].containers : c.name] == ["gateway"] + error_message = "The gateway service must stay single-container unless collector_enabled is set." + } + + assert { + condition = !anytrue([ + for e in google_cloud_run_v2_service.gateway[0].template[0].containers[0].env : startswith(e.name, "LITELLM_COLLECTOR_") + ]) + error_message = "No LITELLM_COLLECTOR_* env may reach the gateway while the sidecar is disabled." + } +} + +run "enabled_adds_a_sidecar_that_shares_the_gateway_transport" { + command = plan + + variables { + collector_enabled = true + collector_port = 4321 + collector_buffer_size = 250 + collector_on_unavailable = "drop" + collector_cpu = "500m" + collector_memory = "1Gi" + gateway_extra_env = { OPENAI_API_BASE = "https://example.invalid" } + gateway_extra_secrets = { OPENAI_API_KEY = "projects/acme-test/secrets/openai-api-key" } + } + + assert { + condition = [for c in google_cloud_run_v2_service.gateway[0].template[0].containers : c.name] == ["gateway", "spend-collector"] + error_message = "Enabling the sidecar must append a spend-collector container after the gateway container." + } + + assert { + condition = alltrue([ + for c in google_cloud_run_v2_service.gateway[0].template[0].containers : ( + { for e in c.env : e.name => e.value }["LITELLM_COLLECTOR_ENABLED"] == "true" && + { for e in c.env : e.name => e.value }["LITELLM_COLLECTOR_ADDRESS"] == "tcp://127.0.0.1:4321" && + { for e in c.env : e.name => e.value }["LITELLM_COLLECTOR_BUFFER_SIZE"] == "250" && + { for e in c.env : e.name => e.value }["LITELLM_COLLECTOR_ON_UNAVAILABLE"] == "drop" && + { for e in c.env : e.name => e.value }["LITELLM_COLLECTOR_DRAIN_TIMEOUT_SECONDS"] == "10" + ) + ]) + error_message = "Gateway and sidecar must agree on the loopback address and the collector knobs." + } + + assert { + condition = ( + google_cloud_run_v2_service.gateway[0].template[0].containers[1].image == local.gateway_image && + google_cloud_run_v2_service.gateway[0].template[0].containers[1].command == tolist(["sh", "-c"]) && + endswith(google_cloud_run_v2_service.gateway[0].template[0].containers[1].args[0], " && exec python -m litellm.proxy.collector") && + strcontains(google_cloud_run_v2_service.gateway[0].template[0].containers[1].args[0], "export DATABASE_URL=") && + strcontains(google_cloud_run_v2_service.gateway[0].template[0].containers[1].args[0], "REDIS_SSL_CA_CERTS") && + { for e in google_cloud_run_v2_service.gateway[0].template[0].containers[1].env : e.name => e.value }["LITELLM_JOB_ROLE"] == "collector" + ) + error_message = "The sidecar must run litellm.proxy.collector from the gateway image with the same Redis CA + DATABASE_URL bootstrap as the gateway." + } + + assert { + condition = ( + length(google_cloud_run_v2_service.gateway[0].template[0].containers[1].ports) == 0 && + google_cloud_run_v2_service.gateway[0].template[0].containers[1].resources[0].limits.cpu == "500m" && + google_cloud_run_v2_service.gateway[0].template[0].containers[1].resources[0].limits.memory == "1Gi" + ) + error_message = "The sidecar must not claim the ingress port and must carry its own resource limits." + } + + assert { + condition = ( + { for e in google_cloud_run_v2_service.gateway[0].template[0].containers[1].env : e.name => e.value }["OPENAI_API_BASE"] == "https://example.invalid" && + contains([for e in google_cloud_run_v2_service.gateway[0].template[0].containers[1].env : e.name], "DATABASE_HOST") && + contains([for e in google_cloud_run_v2_service.gateway[0].template[0].containers[1].env : e.name], "REDIS_HOST") && + contains([for e in google_cloud_run_v2_service.gateway[0].template[0].containers[1].env : e.name if length(e.value_source) > 0], "LITELLM_MASTER_KEY") && + contains([for e in google_cloud_run_v2_service.gateway[0].template[0].containers[1].env : e.name if length(e.value_source) > 0], "DATABASE_PASSWORD") && + contains([for e in google_cloud_run_v2_service.gateway[0].template[0].containers[1].env : e.name if length(e.value_source) > 0], "OPENAI_API_KEY") + ) + error_message = "The sidecar must receive the gateway's database, Redis, and Secret Manager env plus gateway_extra_env / gateway_extra_secrets." + } +} + +run "coexists_with_the_metrics_sidecars" { + command = plan + + variables { + collector_enabled = true + gateway_metrics_port = 4001 + } + + assert { + condition = [for c in google_cloud_run_v2_service.gateway[0].template[0].containers : c.name] == ["gateway", "metrics", "collector", "spend-collector"] + error_message = "The spend collector must keep its own container name next to the GMP metrics collector." + } + + assert { + condition = ( + { for e in google_cloud_run_v2_service.gateway[0].template[0].containers[0].env : e.name => e.value }["PROMETHEUS_MULTIPROC_DIR"] == local.metrics_multiproc_dir && + { for e in google_cloud_run_v2_service.gateway[0].template[0].containers[0].env : e.name => e.value }["LITELLM_COLLECTOR_ENABLED"] == "true" + ) + error_message = "The gateway container must keep both the metrics and the collector env when both sidecars are on." + } +} + +run "sidecars_must_not_share_a_loopback_port" { + command = plan + + variables { + collector_enabled = true + collector_port = 4001 + gateway_metrics_port = 4001 + } + + expect_failures = [ + google_cloud_run_v2_service.gateway, + ] +} + +run "collector_cannot_take_the_metrics_sidecar_health_port" { + command = plan + + variables { + collector_enabled = true + collector_port = 13133 + } + + expect_failures = [ + var.collector_port, + ] +} + +run "proxy_config_is_mounted_into_the_sidecar_too" { + command = plan + + variables { + collector_enabled = true + proxy_config = { model_list = [] } + } + + assert { + condition = alltrue([ + for c in google_cloud_run_v2_service.gateway[0].template[0].containers : ( + [for m in c.volume_mounts : m.name] == [local.proxy_config_volume] && + contains([for e in c.env : e.name], "CONFIG_FILE_PATH") + ) + ]) + error_message = "Both containers must mount the proxy-config GCS volume and point CONFIG_FILE_PATH at it." + } +} diff --git a/terraform/litellm/gcp/variables.tf b/terraform/litellm/gcp/variables.tf index f298a0431c0..412f919ab89 100644 --- a/terraform/litellm/gcp/variables.tf +++ b/terraform/litellm/gcp/variables.tf @@ -656,3 +656,73 @@ variable "billing_metrics_ca_cert_pem" { default = "" sensitive = true } + +# ---------- Collector sidecar ---------- +# +# Opt-in offload of spend tracking from the gateway's uvicorn workers to a +# `python -m litellm.proxy.collector` sidecar container in the same Cloud Run +# instance (helm's `gateway.collector`, mirrors the AWS stack). Containers +# in one instance share localhost, so the sidecar listens on loopback TCP. +# Disabled (the default) adds nothing to the service. + +variable "collector_enabled" { + description = "Run the collector sidecar next to the gateway container and have the gateway ship spend events to it (sets LITELLM_COLLECTOR_ENABLED=true on both). The sidecar shares the instance's request-based CPU allocation, so pair it with a non-zero gateway_min_instances if spend must keep flowing between requests." + type = bool + default = false +} + +variable "collector_port" { + description = "Loopback TCP port the sidecar listens on (LITELLM_COLLECTOR_ADDRESS=tcp://127.0.0.1:)." + type = number + default = 4010 + + validation { + condition = var.collector_port >= 1024 && var.collector_port <= 65535 && !contains([4000, 13133], var.collector_port) + error_message = "collector_port must be in 1024-65535 and not 4000 (the gateway port) or 13133 (the metrics sidecar health port)." + } +} + +variable "collector_cpu" { + description = "Cloud Run CPU limit for the sidecar container, on top of gateway_cpu. Matches helm's collector.resources.limits.cpu." + type = string + default = "1000m" +} + +variable "collector_memory" { + description = "Cloud Run memory limit for the sidecar container, on top of gateway_memory. Matches helm's collector.resources.limits.memory." + type = string + default = "2Gi" +} + +variable "collector_buffer_size" { + description = "Per-worker in-memory queue of spend events waiting to be shipped to the sidecar (LITELLM_COLLECTOR_BUFFER_SIZE)." + type = number + default = 1000 + + validation { + condition = var.collector_buffer_size >= 1 + error_message = "collector_buffer_size must be >= 1." + } +} + +variable "collector_on_unavailable" { + description = "What the gateway does with spend events when the sidecar is unreachable or the buffer is full (LITELLM_COLLECTOR_ON_UNAVAILABLE): `fallback` runs the pipeline in-process, `drop` discards them." + type = string + default = "fallback" + + validation { + condition = contains(["fallback", "drop"], var.collector_on_unavailable) + error_message = "collector_on_unavailable must be one of: fallback, drop." + } +} + +variable "collector_drain_timeout_seconds" { + description = "Seconds a gateway worker waits on shutdown for its buffered spend events to reach the sidecar (LITELLM_COLLECTOR_DRAIN_TIMEOUT_SECONDS)." + type = number + default = 10 + + validation { + condition = var.collector_drain_timeout_seconds > 0 + error_message = "collector_drain_timeout_seconds must be > 0." + } +} diff --git a/tests/test_gateway/test_launch.py b/tests/test_gateway/test_launch.py index 6039aaafc2b..a783ce6ac7e 100644 --- a/tests/test_gateway/test_launch.py +++ b/tests/test_gateway/test_launch.py @@ -3,6 +3,7 @@ import socket import sys import textwrap import urllib.parse +from collections.abc import Iterator from pathlib import Path from typing import Final, cast from unittest.mock import MagicMock, patch @@ -66,7 +67,7 @@ def _query(url: str) -> dict[str, str]: @pytest.fixture -def password_env(monkeypatch: pytest.MonkeyPatch) -> dict[str, str]: +def password_env(monkeypatch: pytest.MonkeyPatch) -> Iterator[dict[str, str]]: for var in ( "DATABASE_URL", "IAM_TOKEN_DB_AUTH", @@ -78,7 +79,8 @@ def password_env(monkeypatch: pytest.MonkeyPatch) -> dict[str, str]: monkeypatch.delenv(var) for var, value in DB_ENV.items(): monkeypatch.setenv(var, value) - return dict(DB_ENV) + yield dict(DB_ENV) + os.environ.pop("DATABASE_URL", None) def _minted_iam_token(token: str): diff --git a/tests/test_litellm/caching/test_caching_handler.py b/tests/test_litellm/caching/test_caching_handler.py index 8e0bc200012..071b99850f6 100644 --- a/tests/test_litellm/caching/test_caching_handler.py +++ b/tests/test_litellm/caching/test_caching_handler.py @@ -658,3 +658,38 @@ def test_async_cache_write_completes_when_asyncio_run_closes_the_loop(monkeypatc asyncio.run(_short_lived_script()) assert len(writes) == 1 + + +@pytest.mark.asyncio +async def test_cache_hit_records_the_looked_up_key_as_the_preset_cache_key(monkeypatch): + """The spend log for a cache hit must reuse the key the lookup already computed instead of hashing again.""" + import litellm + from litellm.caching.caching import Cache + from litellm.types.utils import CallTypes + + async def acompletion(**kwargs): + return None + + monkeypatch.setattr(litellm, "cache", Cache(type="local")) + kwargs = {"model": "gpt-5.4", "messages": [{"role": "user", "content": "hello"}], "caching": True} + await litellm.cache.async_add_cache( + litellm.ModelResponse(choices=[{"message": {"role": "assistant", "content": "hi"}}]), **kwargs + ) + handler = LLMCachingHandler(original_function=acompletion, request_kwargs=kwargs, start_time=datetime.now()) + logging_obj = _build_logging_obj(CallTypes.acompletion.value, stream=False) + logging_obj.async_success_handler = AsyncMock() + + hit = await handler._async_get_cache( + model="gpt-5.4", + original_function=acompletion, + logging_obj=logging_obj, + start_time=datetime.now(), + call_type=CallTypes.acompletion.value, + kwargs=kwargs, + args=(), + ) + + assert hit is not None and hit.cached_result is not None + assert handler.preset_cache_key is not None + assert logging_obj.litellm_params["preset_cache_key"] == handler.preset_cache_key + assert hit.cached_result._hidden_params["cache_key"] == handler.preset_cache_key diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index eff892f2d80..c96cfc3ee4a 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -5,14 +5,20 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import SpendLogsPayload, UserAPIKeyAuth +from litellm.proxy.collector import SpendEventConsumer +from litellm.proxy.db.spend_log_tool_index import response_tool_call_names from litellm.proxy.hooks.proxy_track_cost_callback import ( _get_budget_reservation_from_metadata, _ProxyDBLogger, _should_track_cost_callback, _update_database_and_spend_counters, + run_spend_event, ) -from litellm.types.utils import CallTypes, Usage +from litellm.proxy.spend_tracking.spend_event import SpendEventDecodeError, build_spend_event, decode_spend_event +from litellm.proxy.spend_tracking.spend_event_producer import SpendEventProducer, UnixAddress +from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload +from litellm.types.utils import CallTypes, LiteLLMBatch, ModelResponse, Usage @pytest.mark.asyncio @@ -2096,3 +2102,248 @@ async def test_spend_counters_keep_every_granted_group_when_the_deployment_is_un ) assert charged == ("premium", "tier0") + + +def _offload_kwargs() -> dict: + big_prompt = "x" * 10_000 + reservation = {"reserved_cost": 0.5, "entries": [{"counter_key": "key:hash-1", "reserved_cost": 0.5}]} + return { + "litellm_call_id": "call-1", + "call_type": "acompletion", + "model": "gpt-4o", + "custom_llm_provider": "openai", + "stream": False, + "cache_hit": None, + "response_cost": 0.0125, + "completion_start_time": datetime(2026, 1, 1, 0, 0, 1), + "messages": [{"role": "user", "content": big_prompt}], + "tools": [{"type": "function", "function": {"name": "get_weather", "parameters": {"type": "object"}}}], + "litellm_params": { + "api_base": "https://api.openai.com", + "preset_cache_key": None, + "proxy_server_request": {"body": {"messages": [{"role": "user", "content": big_prompt}]}}, + "metadata": { + "user_api_key": "hash-1", + "user_api_key_hash": "hash-1", + "user_api_key_alias": "alias-1", + "user_api_key_user_id": "user-1", + "user_api_key_team_id": "team-1", + "user_api_key_org_id": "org-1", + "user_api_key_end_user_id": "end-user-1", + "user_api_key_auth": UserAPIKeyAuth(api_key="hash-1", budget_reservation=reservation), + "model_group": "gpt-4o", + "model_info": {"id": "deployment-1"}, + "tags": ["tag-a"], + }, + }, + "standard_logging_object": { + "id": "chatcmpl-1", + "trace_id": "trace-1", + "response_cost": 0.0125, + "model": "gpt-4o-2024-08-06", + "model_id": "deployment-1", + "model_group": "gpt-4o", + "api_base": "https://api.openai.com", + "custom_llm_provider": "openai", + "prompt_tokens": 5000, + "completion_tokens": 4000, + "total_tokens": 9000, + "request_tags": ["tag-a"], + "request_model_access_groups": ["premium"], + "messages": [{"role": "user", "content": big_prompt}], + "response": {"choices": [{"message": {"content": "y" * 10_000}}]}, + "model_parameters": {"temperature": 0.1}, + "metadata": { + "user_api_key_hash": "hash-1", + "user_api_key_end_user_id": "end-user-1", + "usage_object": {"prompt_tokens": 5000, "completion_tokens": 4000, "total_tokens": 9000}, + }, + "hidden_params": {"litellm_overhead_time_ms": 3}, + "model_map_information": {}, + "cost_breakdown": {"input_cost": 0.0125, "output_cost": 0.0}, + }, + } + + +def _offload_response() -> ModelResponse: + return ModelResponse( + id="chatcmpl-1", + model="gpt-4o-2024-08-06", + choices=[ + { + "index": 0, + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "call-1", "type": "function", "function": {"name": "get_weather", "arguments": "{}"}} + ], + }, + "finish_reason": "tool_calls", + } + ], + usage=Usage(prompt_tokens=5000, completion_tokens=4000, total_tokens=9000), + ) + + +class _RecordingHandler: + def __init__(self) -> None: + self.lines: list[bytes] = [] # mutable-ok: test double records the events the sidecar received + + async def __call__(self, line: bytes) -> None: + self.lines.append(line) + + +async def _no_fallback(line: bytes) -> None: + raise AssertionError("the sidecar was reachable, nothing should fall back") + + +@pytest.mark.asyncio +async def test_async_log_success_event_hands_the_sidecar_a_compact_event_and_skips_the_pipeline(tmp_path): + handler = _RecordingHandler() + consumer = SpendEventConsumer(handler) + address = UnixAddress(path=str(tmp_path / "spend.sock")) + server = await consumer.serve(address) + producer = SpendEventProducer( + address=address, on_unavailable="fallback", buffer_size=10, connect_timeout=1.0, fallback=_no_fallback + ) + logger = _ProxyDBLogger(producer) + + with ( + patch( # test-quality-ok: the callback imports this from proxy_server inside its body, so there is no injection seam + "litellm.proxy.proxy_server.proxy_logging_obj" + ) as mock_proxy_logging, + patch( # test-quality-ok: same function-body import, no injection seam + "litellm.proxy.proxy_server.increment_spend_counters", new_callable=AsyncMock + ) as counters, + ): + mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock() + await logger.async_log_success_event(_offload_kwargs(), _offload_response(), datetime.now(), datetime.now()) + await producer.close(drain_timeout=5.0) + server.close() + assert await consumer.drain(timeout=5.0) == 0 + + mock_proxy_logging.db_spend_update_writer.update_database.assert_not_awaited() + counters.assert_not_awaited() + assert producer.stats().sent == 1 + assert len(handler.lines) == 1 + assert len(handler.lines[0]) < 4_000 + event = decode_spend_event(handler.lines[0]) + assert not isinstance(event, SpendEventDecodeError) + assert event.litellm_params["metadata"]["user_api_key_team_id"] == "team-1" + assert event.response_cost == 0.0125 + + +@pytest.mark.asyncio +async def test_async_log_success_event_keeps_batch_retrieves_in_process(): + producer = SpendEventProducer( + address=UnixAddress(path="/nonexistent/spend.sock"), + on_unavailable="drop", + buffer_size=10, + connect_timeout=1.0, + fallback=_no_fallback, + ) + logger = _ProxyDBLogger(producer) + kwargs = {**_offload_kwargs(), "call_type": CallTypes.aretrieve_batch.value} + completed_batch = LiteLLMBatch( + id="batch_abc", + completion_window="24h", + created_at=1, + endpoint="/v1/chat/completions", + input_file_id="file-in", + output_file_id="file-out", + object="batch", + status="completed", + ) + + with ( + patch( # test-quality-ok: the callback imports this from proxy_server inside its body, so there is no injection seam + "litellm.proxy.proxy_server.proxy_logging_obj" + ) as mock_proxy_logging, + patch( # test-quality-ok: same function-body import, no injection seam + "litellm.proxy.proxy_server.increment_spend_counters", new_callable=AsyncMock + ), + patch( # test-quality-ok: same function-body import, no injection seam + "litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock + ), + ): + mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock(return_value=True) + mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock() + await logger.async_log_success_event(kwargs, completed_batch, datetime.now(), datetime.now()) + + mock_proxy_logging.db_spend_update_writer.update_database.assert_awaited_once() + assert producer.stats().queued == 0 + + +async def _spend_row_written_by(run) -> tuple[SpendLogsPayload, dict, tuple[str, ...]]: + with ( + patch( # test-quality-ok: the callback imports this from proxy_server inside its body, so there is no injection seam + "litellm.proxy.proxy_server.proxy_logging_obj" + ) as mock_proxy_logging, + patch( # test-quality-ok: same function-body import, no injection seam + "litellm.proxy.proxy_server.increment_spend_counters", new_callable=AsyncMock + ) as counters, + patch( # test-quality-ok: same function-body import, no injection seam + "litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock + ), + ): + mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock(return_value=True) + mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock() + await run() + mock_proxy_logging.db_spend_update_writer.update_database.assert_awaited_once() + written = mock_proxy_logging.db_spend_update_writer.update_database.await_args.kwargs + counters.assert_awaited_once() + counted = dict(counters.await_args.kwargs) + + row = get_logging_payload( + kwargs=written["kwargs"], + response_obj=written["completion_response"], + start_time=written["start_time"], + end_time=written["end_time"], + ) + return row, counted, response_tool_call_names(written["completion_response"]) + + +@pytest.mark.asyncio +async def test_sidecar_writes_the_same_spend_row_and_counters_as_the_in_process_path(): + start_time = datetime(2026, 1, 1, 0, 0, 0) + end_time = datetime(2026, 1, 1, 0, 0, 2) + + async def in_process() -> None: + await _ProxyDBLogger().async_log_success_event(_offload_kwargs(), _offload_response(), start_time, end_time) + + async def via_sidecar() -> None: + line = build_spend_event(_offload_kwargs(), _offload_response(), start_time, end_time, store_bodies=False) + assert isinstance(line, bytes) + await run_spend_event(line) + + in_process_row, in_process_counters, in_process_tools = await _spend_row_written_by(in_process) + sidecar_row, sidecar_counters, sidecar_tools = await _spend_row_written_by(via_sidecar) + + assert sidecar_row == in_process_row + assert in_process_row["spend"] == 0.0125 + assert in_process_row["team_id"] == "team-1" + assert in_process_row["end_user"] == "end-user-1" + assert in_process_row["total_tokens"] == 9000 + assert in_process_row["model_id"] == "deployment-1" + assert in_process_row["request_tags"] == '["tag-a"]' + assert in_process_row["messages"] == "{}" + assert in_process_row["response"] == "{}" + assert sidecar_counters == in_process_counters + assert in_process_counters["token"] == "hash-1" + assert in_process_counters["response_cost"] == 0.0125 + assert in_process_counters["budget_reservation"]["reserved_cost"] == 0.5 + assert in_process_counters["model_access_groups"] == ("premium",) + assert sidecar_tools == in_process_tools == ("get_weather",) + + +@pytest.mark.asyncio +async def test_sidecar_ignores_an_undecodable_event(): # test-quality-ok: a discarded event has no observable output other than the DB writer never being reached + with ( + patch( # test-quality-ok: the callback imports this from proxy_server inside its body, so there is no injection seam + "litellm.proxy.proxy_server.proxy_logging_obj" + ) as mock_proxy_logging + ): + mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock() + await run_spend_event(b"garbage\n") + mock_proxy_logging.db_spend_update_writer.update_database.assert_not_awaited() diff --git a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py index dafe686c4f2..deb7289d2d1 100644 --- a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py +++ b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py @@ -253,6 +253,38 @@ async def test_flush_spend_logs_queue_on_shutdown_swallows_drain_errors(monkeypa await ps._flush_spend_logs_queue_on_shutdown() +@pytest.mark.asyncio +async def test_flush_spend_counters_on_shutdown_commits_buffered_spend(monkeypatch): + fake_prisma = MagicMock() + monkeypatch.setattr(ps, "prisma_client", fake_prisma, raising=False) + commit = AsyncMock() + monkeypatch.setattr(ps.proxy_logging_obj.db_spend_update_writer, "db_update_spend_transaction_handler", commit) + + await ps.flush_spend_counters_on_shutdown() + + observed = { + "commit_calls": commit.await_count, + "commit_prisma": commit.await_args.kwargs["prisma_client"] is fake_prisma, + "commit_proxy_logging": commit.await_args.kwargs["proxy_logging_obj"] is ps.proxy_logging_obj, + } + assert observed == {"commit_calls": 1, "commit_prisma": True, "commit_proxy_logging": True} + + +@pytest.mark.asyncio +async def test_flush_spend_counters_on_shutdown_logs_and_swallows_commit_errors(monkeypatch, caplog): + monkeypatch.setattr(ps, "prisma_client", MagicMock(), raising=False) + monkeypatch.setattr( + ps.proxy_logging_obj.db_spend_update_writer, + "db_update_spend_transaction_handler", + AsyncMock(side_effect=RuntimeError("db gone")), + ) + + with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"): + await ps.flush_spend_counters_on_shutdown() + + assert "Error flushing spend counters on shutdown: db gone" in caplog.text + + # --------------------------------------------------------------------------- # _initialize_shared_aiohttp_session # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_event.py b/tests/test_litellm/proxy/spend_tracking/test_spend_event.py new file mode 100644 index 00000000000..ff449235582 --- /dev/null +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_event.py @@ -0,0 +1,213 @@ +import json +from datetime import datetime +from typing import Final + +import pytest + +import litellm +from litellm.caching.caching import Cache +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.spend_tracking.spend_event import ( + CACHE_OFF_KEY, + SpendEventBuildError, + SpendEventDecodeError, + build_spend_event, + decode_spend_event, + is_offloadable_success, + spend_event_callback_args, +) +from litellm.types.utils import LiteLLMBatch, ModelResponse, Usage + +_BIG_PROMPT: Final = "x" * 20_000 +_RESERVATION: Final = { + "reserved_cost": 0.5, + "entries": [{"counter_key": "key:hash", "reserved_cost": 0.5}], + "finalized": False, + "input_cost": 0.1, + "input_tokens": 5000, +} + + +def _response(tool_name: str | None = None) -> ModelResponse: + tool_calls: Final = ( + [{"id": "call-1", "type": "function", "function": {"name": tool_name, "arguments": "{}"}}] + if tool_name is not None + else None + ) + return ModelResponse( + id="chatcmpl-1", + model="gpt-4o-2024-08-06", + choices=[ + { + "index": 0, + "message": {"role": "assistant", "content": "y" * 20_000, "tool_calls": tool_calls}, + "finish_reason": "tool_calls" if tool_name else "stop", + } + ], + usage=Usage(prompt_tokens=5000, completion_tokens=4000, total_tokens=9000), + ) + + +def _success_kwargs(preset_cache_key: str | None = "preset-key") -> dict: + return { + "litellm_call_id": "call-1", + "call_type": "acompletion", + "model": "gpt-4o", + "custom_llm_provider": "openai", + "stream": False, + "cache_hit": None, + "response_cost": 0.0125, + "completion_start_time": datetime(2026, 1, 1, 0, 0, 1), + "messages": [{"role": "user", "content": _BIG_PROMPT}], + "tools": [{"type": "function", "function": {"name": "get_weather", "parameters": {"type": "object"}}}], + "litellm_params": { + "api_base": "https://api.openai.com", + "preset_cache_key": preset_cache_key, + "proxy_server_request": {"body": {"messages": [{"role": "user", "content": _BIG_PROMPT}]}}, + "metadata": { + "user_api_key": "hash-1", + "user_api_key_user_id": "user-1", + "user_api_key_team_id": "team-1", + "user_api_key_org_id": "org-1", + "user_api_key_end_user_id": "end-user-1", + "user_api_key_auth": UserAPIKeyAuth(api_key="hash-1", budget_reservation=dict(_RESERVATION)), + "model_group": "gpt-4o", + "model_info": {"id": "deployment-1"}, + "tags": ["tag-a"], + "litellm_parent_otel_span": object(), + }, + }, + "standard_logging_object": { + "response_cost": 0.0125, + "model": "gpt-4o-2024-08-06", + "model_id": "deployment-1", + "request_tags": ["tag-a"], + "request_model_access_groups": ["premium"], + "messages": [{"role": "user", "content": _BIG_PROMPT}], + "response": {"choices": [{"message": {"content": "y" * 20_000}}]}, + "model_parameters": {"temperature": 0.1}, + "metadata": {"user_api_key_hash": "hash-1", "usage_object": {"prompt_tokens": 5000}}, + "hidden_params": {"litellm_overhead_time_ms": 3}, + "model_map_information": {}, + }, + } + + +def _build(kwargs: dict, response: object, store_bodies: bool = False) -> bytes: + line: Final = build_spend_event( + kwargs, response, datetime(2026, 1, 1), datetime(2026, 1, 1, 0, 0, 2), store_bodies=store_bodies + ) + assert isinstance(line, bytes) + return line + + +def test_event_is_compact_and_omits_bodies_by_default(): + line: Final = _build(_success_kwargs(), _response(tool_name="get_weather")) + assert line.endswith(b"\n") + assert len(line) < 4_000 + assert _BIG_PROMPT.encode() not in line + assert b"yyyy" not in line + decoded: Final = json.loads(line) + assert "messages" not in decoded["standard_logging_object"] + assert "response" not in decoded["standard_logging_object"] + assert decoded["litellm_params"]["proxy_server_request"] is None + + +def test_event_carries_bodies_when_spend_logs_store_them(): + line: Final = _build(_success_kwargs(), _response(), store_bodies=True) + decoded: Final = json.loads(line) + assert decoded["standard_logging_object"]["messages"][0]["content"] == _BIG_PROMPT + assert decoded["standard_logging_object"]["response"]["choices"][0]["message"]["content"] == "y" * 20_000 + assert decoded["litellm_params"]["proxy_server_request"]["body"]["messages"][0]["content"] == _BIG_PROMPT + + +def test_round_trip_preserves_identity_usage_reservation_and_tools(): + line: Final = _build(_success_kwargs(), _response(tool_name="get_weather")) + event: Final = decode_spend_event(line) + assert not isinstance(event, SpendEventDecodeError) + args: Final = spend_event_callback_args(event) + + metadata: Final = args.kwargs["litellm_params"]["metadata"] + assert metadata is not None + assert (metadata["user_api_key"], metadata["user_api_key_team_id"], metadata["user_api_key_org_id"]) == ( + "hash-1", + "team-1", + "org-1", + ) + assert metadata["user_api_key_budget_reservation"] == _RESERVATION + assert "user_api_key_auth" not in metadata + assert "litellm_parent_otel_span" not in metadata + assert args.kwargs["standard_logging_object"]["request_model_access_groups"] == ["premium"] + assert args.kwargs["standard_logging_object"]["response_cost"] == 0.0125 + assert args.kwargs["tools"] == ({"type": "function", "function": {"name": "get_weather"}},) + assert args.kwargs["completion_start_time"] == datetime(2026, 1, 1, 0, 0, 1) + assert (args.start_time, args.end_time) == (datetime(2026, 1, 1), datetime(2026, 1, 1, 0, 0, 2)) + assert args.response_obj is not None + assert args.response_obj["id"] == "chatcmpl-1" + assert args.response_obj["usage"]["prompt_tokens"] == 5000 + assert args.response_obj["usage"]["completion_tokens"] == 4000 + tool_calls: Final = args.response_obj["choices"][0]["message"]["tool_calls"] + assert [call["function"]["name"] for call in tool_calls] == ["get_weather"] + assert "complete_streaming_response" not in args.kwargs + + +def test_streaming_event_reconstructs_complete_streaming_response(): + kwargs: Final = {**_success_kwargs(), "stream": True, "complete_streaming_response": _response()} + event: Final = decode_spend_event(_build(kwargs, _response())) + assert not isinstance(event, SpendEventDecodeError) + args: Final = spend_event_callback_args(event) + assert args.kwargs["stream"] is True + assert args.kwargs["complete_streaming_response"] == args.response_obj + + +class _HashingCache(Cache): + def __init__(self) -> None: + pass + + def get_cache_key(self, **kwargs) -> str: + raise AssertionError("the fast path must not hash the request body") + + +@pytest.mark.parametrize( + ("cache", "preset", "expected"), + [ + (None, "preset-key", CACHE_OFF_KEY), + (_HashingCache(), "preset-key", "preset-key"), + (_HashingCache(), None, None), + ], +) +def test_event_reuses_preset_cache_key_and_never_hashes(monkeypatch, cache, preset, expected): + monkeypatch.setattr(litellm, "cache", cache) + decoded: Final = json.loads(_build(_success_kwargs(preset_cache_key=preset), _response())) + assert decoded["litellm_params"]["preset_cache_key"] == expected + + +def test_unbuildable_kwargs_fall_back_to_in_process_tracking(): + kwargs: Final = {**_success_kwargs(), "response_cost": "not-a-number"} + assert isinstance( + build_spend_event(kwargs, _response(), datetime.now(), datetime.now(), False), SpendEventBuildError + ) + + +def test_undecodable_line_is_an_error_value(): + assert isinstance(decode_spend_event(b'{"version": 2}\n'), SpendEventDecodeError) + assert isinstance(decode_spend_event(b"not json\n"), SpendEventDecodeError) + + +def test_batch_retrieves_stay_in_process(): + assert is_offloadable_success(_response()) is True + assert is_offloadable_success(None) is True + assert ( + is_offloadable_success( + LiteLLMBatch( + id="batch-1", + completion_window="24h", + created_at=1, + endpoint="/v1/chat/completions", + input_file_id="f", + object="batch", + status="completed", + ) + ) + is False + ) diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_event_producer.py b/tests/test_litellm/proxy/spend_tracking/test_spend_event_producer.py new file mode 100644 index 00000000000..adfb1251d1c --- /dev/null +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_event_producer.py @@ -0,0 +1,359 @@ +import asyncio +from collections.abc import Awaitable, Callable +from pathlib import Path +from typing import Final + +import pytest +import uvloop + +from litellm.proxy.spend_tracking.spend_event_producer import ( + AddressError, + CollectorAddress, + CollectorSettings, + SpendEventProducer, + TcpAddress, + UnixAddress, + build_spend_event_producer, + open_collector_connection, + parse_collector_address, +) + + +class _Sidecar: + """A unix-socket server that records every line it receives, standing in for the collector.""" + + def __init__(self, path: Path, reads: bool = True, limit: int = 2**16) -> None: + self.path = path + self.reads = reads + self.limit = limit + self.lines: list[bytes] = [] # mutable-ok: test double records what the producer sent + self._server: asyncio.Server | None = None + self._stopped = asyncio.Event() + self._connections: list[asyncio.StreamWriter] = [] # mutable-ok: test double tracks peers to hang up on + + async def __aenter__(self) -> "_Sidecar": + self._server = await asyncio.start_unix_server(self._on_connection, path=str(self.path), limit=self.limit) + return self + + async def __aexit__(self, *exc: object) -> None: + self._stopped.set() + await self.hang_up() + + async def hang_up(self) -> None: + """Exit the way a stopped sidecar does: stop listening and close every producer connection.""" + assert self._server is not None + self._server.close() + for connection in self._connections: + connection.close() + await connection.wait_closed() + await self._server.wait_closed() + + async def _on_connection(self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: + self._connections.append(writer) + if not self.reads: + await self._stopped.wait() + return + while line := await reader.readline(): + self.lines.append(line) + writer.close() + + +class _CrashingSidecar(_Sidecar): + """Bills a few lines, then dies mid-stream with the producer's backlog still queued behind them.""" + + def __init__(self, path: Path, lines_before_crash: int) -> None: + super().__init__(path, limit=2**20) + self._lines_before_crash = lines_before_crash + + async def _on_connection(self, reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: + self._connections.append(writer) + for _ in range(self._lines_before_crash): + self.lines.append(await reader.readline()) + writer.transport.abort() + + +class _Fallback: + def __init__(self) -> None: + self.lines: list[bytes] = [] # mutable-ok: test double records what fell back to in-process + + async def __call__(self, line: bytes) -> None: + self.lines.append(line) + + +class _GatedFallback(_Fallback): + """A fallback that blocks, like a slow database write, until the test releases it.""" + + def __init__(self) -> None: + super().__init__() + self.started = asyncio.Event() + self.release = asyncio.Event() + + async def __call__(self, line: bytes) -> None: + self.started.set() + await self.release.wait() + await super().__call__(line) + + +class _StalledDrainWriter(asyncio.StreamWriter): + """Hands bytes to the real transport but never wakes ``drain()``: the loop iteration between a flush + completing and the writer task resuming, frozen in place.""" + + def __init__(self, real: asyncio.StreamWriter, reader: asyncio.StreamReader) -> None: + super().__init__(real.transport, real.transport.get_protocol(), reader, asyncio.get_running_loop()) + self._real_writer_whose_finalizer_would_close_the_transport = real + + async def drain(self) -> None: + await asyncio.Event().wait() + + +async def _open_with_stalled_drain( + address: CollectorAddress, timeout: float +) -> tuple[asyncio.StreamReader, asyncio.StreamWriter]: + reader, writer = await open_collector_connection(address, timeout) + return reader, _StalledDrainWriter(writer, reader) + + +def _producer( + path: Path, + fallback: _Fallback, + on_unavailable="fallback", + buffer_size: int = 100, + open_connection: Callable[ + [CollectorAddress, float], Awaitable[tuple[asyncio.StreamReader, asyncio.StreamWriter]] + ] = open_collector_connection, +) -> SpendEventProducer: + return SpendEventProducer( + address=UnixAddress(path=str(path)), + on_unavailable=on_unavailable, + buffer_size=buffer_size, + connect_timeout=1.0, + fallback=fallback, + open_connection=open_connection, + ) + + +def test_parse_collector_address(): + assert parse_collector_address("unix:///var/run/litellm/collector.sock") == UnixAddress( + path="/var/run/litellm/collector.sock" + ) + assert parse_collector_address("tcp://127.0.0.1:4100") == TcpAddress(host="127.0.0.1", port=4100) + assert parse_collector_address("tcp://localhost:4100") == TcpAddress(host="localhost", port=4100) + assert parse_collector_address("tcp://[::1]:4100") == TcpAddress(host="::1", port=4100) + assert isinstance(parse_collector_address("redis://localhost:6379"), AddressError) + assert isinstance(parse_collector_address("tcp://127.0.0.1"), AddressError) + + +@pytest.mark.parametrize("address", ["tcp://0.0.0.0:4100", "tcp://10.0.0.5:4100", "tcp://collector.svc:4100"]) +def test_tcp_address_outside_loopback_is_refused(address: str): + """The socket has no authentication, so anything reachable from outside the pod would accept forged spend.""" + error: Final = parse_collector_address(address) + assert isinstance(error, AddressError) + assert "loopback" in error.reason + assert build_spend_event_producer(CollectorSettings(enabled=True, address=address), _Fallback()) is None + + +def test_gateway_produces_only_when_enabled_and_not_the_sidecar_itself(): + fallback: Final = _Fallback() + assert build_spend_event_producer(CollectorSettings(enabled=False), fallback) is None + assert build_spend_event_producer(CollectorSettings(enabled=True, job_role="collector"), fallback) is None + assert build_spend_event_producer(CollectorSettings(enabled=True, address="redis://x"), fallback) is None + assert isinstance(build_spend_event_producer(CollectorSettings(enabled=True), fallback), SpendEventProducer) + + +def test_settings_read_the_documented_env(monkeypatch): + monkeypatch.setenv("LITELLM_COLLECTOR_ENABLED", "true") + monkeypatch.setenv("LITELLM_COLLECTOR_ADDRESS", "tcp://127.0.0.1:4100") + monkeypatch.setenv("LITELLM_COLLECTOR_BUFFER_SIZE", "50") + monkeypatch.setenv("LITELLM_COLLECTOR_ON_UNAVAILABLE", "drop") + monkeypatch.setenv("LITELLM_JOB_ROLE", "collector") + settings: Final = CollectorSettings() + assert (settings.enabled, settings.address, settings.buffer_size, settings.on_unavailable) == ( + True, + "tcp://127.0.0.1:4100", + 50, + "drop", + ) + assert settings.produces is False + + +@pytest.mark.asyncio +async def test_events_reach_the_sidecar_once_and_in_order(tmp_path: Path): + fallback: Final = _Fallback() + async with _Sidecar(tmp_path / "spend.sock") as sidecar: + producer: Final = _producer(sidecar.path, fallback) + outcomes: Final = [await producer.publish(f"event-{i}\n".encode()) for i in range(20)] + await producer.close(drain_timeout=5.0) + await asyncio.sleep(0.05) + + assert outcomes == ["queued"] * 20 + assert sidecar.lines == [f"event-{i}\n".encode() for i in range(20)] + assert fallback.lines == [] + stats: Final = producer.stats() + assert (stats.queued, stats.sent, stats.fallback, stats.dropped) == (20, 20, 0, 0) + + +@pytest.mark.asyncio +async def test_unreachable_sidecar_falls_back_in_process_and_backs_off(tmp_path: Path): + fallback: Final = _Fallback() + producer: Final = _producer(tmp_path / "missing.sock", fallback) + first: Final = await producer.publish(b"event-1\n") + await asyncio.sleep(0.05) + second: Final = await producer.publish(b"event-2\n") + await producer.close(drain_timeout=5.0) + + assert first == "queued" + assert second == "fallback" + assert fallback.lines == [b"event-1\n", b"event-2\n"] + stats: Final = producer.stats() + assert (stats.sent, stats.fallback, stats.dropped, stats.connected) == (0, 2, 0, False) + + +@pytest.mark.parametrize("loop_factory", [asyncio.new_event_loop, uvloop.new_event_loop], ids=["asyncio", "uvloop"]) +def test_sidecar_hang_up_falls_back_instead_of_losing_events( + tmp_path: Path, loop_factory: Callable[[], asyncio.AbstractEventLoop] +): + async def scenario() -> tuple[list[bytes], list[bytes], tuple[int, int, int]]: + fallback: Final = _Fallback() + sidecar: Final = _Sidecar(tmp_path / "spend.sock") + async with sidecar: + producer: Final = _producer(sidecar.path, fallback) + await producer.publish(b"event-1\n") + await asyncio.sleep(0.05) + await sidecar.hang_up() + await asyncio.sleep(0.05) + await producer.publish(b"event-2\n") + await producer.close(drain_timeout=5.0) + stats: Final = producer.stats() + return sidecar.lines, fallback.lines, (stats.sent, stats.fallback, stats.dropped) + + with asyncio.Runner(loop_factory=loop_factory) as runner: + sidecar_lines, fallback_lines, counts = runner.run(scenario()) + + assert sidecar_lines == [b"event-1\n"] + assert fallback_lines == [b"event-2\n"] + assert counts == (1, 1, 0) + + +@pytest.mark.parametrize("loop_factory", [asyncio.new_event_loop, uvloop.new_event_loop], ids=["asyncio", "uvloop"]) +def test_mid_stream_crash_never_bills_an_event_on_both_sides( + tmp_path: Path, loop_factory: Callable[[], asyncio.AbstractEventLoop] +): + """Events large enough to straddle the kernel buffer, a sidecar that reads some and then drops the socket: a + failed write may only fall back when the sidecar cannot have read the whole line.""" + events: Final = tuple(f"event-{i:03d}-".encode() + b"x" * 65536 + b"\n" for i in range(64)) + + async def scenario() -> tuple[list[bytes], list[bytes], tuple[int, int, int]]: + fallback: Final = _Fallback() + sidecar: Final = _CrashingSidecar(tmp_path / "spend.sock", lines_before_crash=3) + async with sidecar: + producer: Final = _producer(sidecar.path, fallback) + for event in events: + assert await producer.publish(event) == "queued" + await asyncio.sleep(0.2) + await producer.close(drain_timeout=5.0) + stats: Final = producer.stats() + return sidecar.lines, fallback.lines, (stats.sent, stats.fallback, stats.dropped) + + with asyncio.Runner(loop_factory=loop_factory) as runner: + sidecar_lines, fallback_lines, counts = runner.run(scenario()) + + assert sidecar_lines == list(events[:3]) + assert set(sidecar_lines).isdisjoint(fallback_lines) + assert len(fallback_lines) == len(set(fallback_lines)) + assert fallback_lines[-1] == events[-1] + assert counts[0] + counts[1] == len(events) and counts[2] == 0 + assert counts[0] >= len(sidecar_lines) + + +@pytest.mark.asyncio +async def test_drain_timeout_hands_the_in_flight_event_to_fallback(tmp_path: Path): + """A sidecar that stops reading leaves one event half-written; cancelling the writer must not lose it.""" + fallback: Final = _Fallback() + stuck: Final = b"x" * (4 * 1024 * 1024) + b"\n" + async with _Sidecar(tmp_path / "spend.sock", reads=False) as sidecar: + producer: Final = _producer(sidecar.path, fallback) + assert await producer.publish(stuck) == "queued" + await asyncio.sleep(0.1) + await producer.close(drain_timeout=0.2) + + assert fallback.lines == [stuck] + stats: Final = producer.stats() + assert (stats.sent, stats.fallback, stats.connected) == (0, 1, False) + + +@pytest.mark.asyncio +async def test_shutdown_lets_the_writer_finish_a_fallback_already_in_progress(tmp_path: Path): + """Cancelling the writer while it runs the pipeline in-process must neither lose nor repeat that event.""" + fallback: Final = _GatedFallback() + producer: Final = _producer(tmp_path / "missing.sock", fallback) + assert await producer.publish(b"event-1\n") == "queued" + await asyncio.wait_for(fallback.started.wait(), 5.0) + closing: Final = asyncio.ensure_future(producer.close(drain_timeout=0.05)) + await asyncio.sleep(0.2) + assert fallback.lines == [] + fallback.release.set() + await asyncio.wait_for(closing, 5.0) + + assert fallback.lines == [b"event-1\n"] + assert producer.stats().fallback == 1 + + +@pytest.mark.asyncio +async def test_shutdown_does_not_replay_an_event_the_kernel_already_took(tmp_path: Path): + """Cancelling a drain whose bytes already left the process must not run the event a second time in-process.""" + fallback: Final = _Fallback() + async with _Sidecar(tmp_path / "spend.sock") as sidecar: + producer: Final = _producer(sidecar.path, fallback, open_connection=_open_with_stalled_drain) + assert await producer.publish(b"event-1\n") == "queued" + await asyncio.sleep(0.1) + await producer.close(drain_timeout=0.2) + await asyncio.sleep(0.05) + + assert sidecar.lines == [b"event-1\n"] + assert fallback.lines == [] + stats: Final = producer.stats() + assert (stats.fallback, stats.dropped, stats.connected) == (0, 0, False) + + +@pytest.mark.asyncio +async def test_drop_policy_counts_instead_of_running_in_process(tmp_path: Path): + fallback: Final = _Fallback() + producer: Final = _producer(tmp_path / "missing.sock", fallback, on_unavailable="drop") + await producer.publish(b"event-1\n") + await producer.close(drain_timeout=5.0) + assert await producer.publish(b"event-2\n") == "dropped" + + assert fallback.lines == [] + assert producer.stats().dropped == 2 + + +@pytest.mark.asyncio +async def test_full_buffer_applies_the_unavailable_policy_immediately(tmp_path: Path): + fallback: Final = _Fallback() + async with _Sidecar(tmp_path / "spend.sock") as sidecar: + producer: Final = _producer(sidecar.path, fallback, buffer_size=2) + outcomes: Final = [await producer.publish(f"event-{i}\n".encode()) for i in range(3)] + await producer.close(drain_timeout=5.0) + await asyncio.sleep(0.05) + + assert outcomes == ["queued", "queued", "fallback"] + assert fallback.lines == [b"event-2\n"] + assert sidecar.lines == [b"event-0\n", b"event-1\n"] + + +@pytest.mark.asyncio +async def test_close_flushes_buffered_events_then_refuses_new_ones(tmp_path: Path): + fallback: Final = _Fallback() + async with _Sidecar(tmp_path / "spend.sock") as sidecar: + producer: Final = _producer(sidecar.path, fallback) + for i in range(50): + await producer.publish(f"event-{i}\n".encode()) + assert sidecar.lines == [] + await producer.close(drain_timeout=5.0) + await asyncio.sleep(0.05) + after_close: Final = await producer.publish(b"late\n") + + assert len(sidecar.lines) == 50 + assert after_close == "fallback" + assert fallback.lines == [b"late\n"] + assert producer.stats().connected is False 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 d84d63ce1a4..df113197ec6 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 @@ -38,9 +38,9 @@ from litellm.proxy.spend_tracking.spend_tracking_utils import ( _sanitize_error_information_for_spend_logs, _sanitize_guardrail_information_for_spend_logs, _sanitize_request_body_for_spend_logs_payload, - _should_store_prompts_and_responses_in_spend_logs, get_logging_payload, get_spend_logs_id, + should_store_prompts_and_responses_in_spend_logs, ) from litellm.proxy.utils import hash_token from litellm.types.utils import ( @@ -107,6 +107,48 @@ def test_get_logging_payload_maps_openai_cached_tokens_to_cache_read_input_token assert additional_usage_values["prompt_tokens_details"]["cached_tokens"] == 123 +class _HashingCache(litellm.Cache): + def __init__(self) -> None: + pass + + def get_cache_key(self, **kwargs) -> str: + raise AssertionError("a preset cache key must be reused instead of hashing the request") + + +def _cache_key_in_spend_log(monkeypatch: pytest.MonkeyPatch, cache: litellm.Cache | None, preset: str | None) -> str: + monkeypatch.setattr(litellm, "cache", cache) + payload: Final = get_logging_payload( + kwargs={ + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "x" * 10_000}], + "litellm_params": {"metadata": {"user_api_key": "test-key"}, "preset_cache_key": preset}, + }, + response_obj=litellm.ModelResponse(id="chatcmpl-test", choices=[], usage=litellm.Usage()), + start_time=datetime.datetime.now(timezone.utc), + end_time=datetime.datetime.now(timezone.utc), + ) + return payload["cache_key"] + + +def test_get_logging_payload_reuses_the_preset_cache_key_instead_of_hashing_the_body(monkeypatch): + assert _cache_key_in_spend_log(monkeypatch, _HashingCache(), "preset-key") == "preset-key" + + +def test_get_logging_payload_records_cache_off_without_hashing(monkeypatch): + assert _cache_key_in_spend_log(monkeypatch, None, None) == "Cache OFF" + + +def test_get_logging_payload_still_hashes_when_caching_is_on_and_no_preset_key_exists(monkeypatch): + class _RecordingCache(litellm.Cache): + def __init__(self) -> None: + pass + + def get_cache_key(self, **kwargs) -> str: + return "hashed-from-" + kwargs["model"] + + assert _cache_key_in_spend_log(monkeypatch, _RecordingCache(), None) == "hashed-from-gpt-4o-mini" + + _TRACE_ONLY_STANDARD_LOGGING: Final = cast( StandardLoggingPayload, { @@ -613,11 +655,11 @@ def test_sanitize_request_body_for_spend_logs_payload_circular_reference(): assert sanitized == {"b": {"a": {}}} # Should return empty dict for circular reference -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_get_vector_store_request_for_spend_logs_payload_store_prompts_true( mock_should_store, ): - # When _should_store_prompts_and_responses_in_spend_logs returns True + # When should_store_prompts_and_responses_in_spend_logs returns True mock_should_store.return_value = True # Sample vector store request metadata @@ -631,11 +673,11 @@ def test_get_vector_store_request_for_spend_logs_payload_store_prompts_true( assert result[0]["vector_store_search_response"]["data"][0]["content"][0]["text"] == "sensitive information" -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_get_vector_store_request_for_spend_logs_payload_store_prompts_false( mock_should_store, ): - # When _should_store_prompts_and_responses_in_spend_logs returns False + # When should_store_prompts_and_responses_in_spend_logs returns False mock_should_store.return_value = False # Sample vector store request metadata @@ -651,7 +693,7 @@ def test_get_vector_store_request_for_spend_logs_payload_store_prompts_false( assert result[0]["vector_store_search_response"]["data"][0]["content"][0]["type"] == "text" -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_get_vector_store_request_for_spend_logs_payload_null_input(mock_should_store): # When input is None mock_should_store.return_value = False @@ -659,7 +701,7 @@ def test_get_vector_store_request_for_spend_logs_payload_null_input(mock_should_ assert result is None -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_get_messages_for_spend_logs_realtime_returns_messages(mock_should_store): """ Test that _get_messages_for_spend_logs_payload returns messages @@ -686,7 +728,7 @@ def test_get_messages_for_spend_logs_realtime_returns_messages(mock_should_store assert parsed[1]["content"] == "What is the weather today?" -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_get_messages_for_spend_logs_strips_null_bytes(mock_should_store): """Regression for PostgreSQL 22P05: NUL bytes must be stripped from messages.""" mock_should_store.return_value = True @@ -703,7 +745,7 @@ def test_get_messages_for_spend_logs_strips_null_bytes(mock_should_store): assert parsed[0]["content"] == "helloworld" -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_get_messages_for_spend_logs_realtime_empty_when_disabled(mock_should_store): """ Test that _get_messages_for_spend_logs_payload returns '{}' for realtime calls @@ -721,7 +763,7 @@ def test_get_messages_for_spend_logs_realtime_empty_when_disabled(mock_should_st assert result == "{}" -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_get_messages_for_spend_logs_non_realtime_returns_empty(mock_should_store): """ Test that _get_messages_for_spend_logs_payload returns '{}' for non-realtime @@ -739,7 +781,7 @@ def test_get_messages_for_spend_logs_non_realtime_returns_empty(mock_should_stor assert result == "{}" -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_get_response_for_spend_logs_payload_truncates_large_base64(mock_should_store): from litellm.constants import MAX_STRING_LENGTH_PROMPT_IN_DB @@ -767,7 +809,7 @@ def test_get_response_for_spend_logs_payload_truncates_large_base64(mock_should_ assert parsed["data"][0]["other_field"] == "value" -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_get_response_for_spend_logs_payload_strips_null_bytes(mock_should_store): """Regression for PostgreSQL 22P05: NUL bytes must be stripped from response.""" mock_should_store.return_value = True @@ -780,7 +822,7 @@ def test_get_response_for_spend_logs_payload_strips_null_bytes(mock_should_store assert json.loads(response_json)["content"] == "answerhere" -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_get_response_for_spend_logs_payload_truncates_large_embedding( mock_should_store, ): @@ -835,7 +877,7 @@ def test_truncation_includes_db_safeguard_note(): ) -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_response_truncation_logs_info_message(mock_should_store): """ Test that when response is truncated before DB storage, an info log is emitted @@ -857,7 +899,7 @@ def test_response_truncation_logs_info_message(mock_should_store): assert "response was truncated" in log_msg -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_request_body_truncation_logs_info_message(mock_should_store): """ Test that when request body is truncated before DB storage, an info log is emitted. @@ -1520,7 +1562,7 @@ def test_get_logging_payload_handles_missing_overhead_gracefully(): ) -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_spend_logs_redacts_request_and_response_when_turn_off_message_logging_enabled( mock_should_store, ): @@ -1584,7 +1626,7 @@ def test_should_store_prompts_and_responses_in_spend_logs_case_insensitive_strin mock_get_secret_bool, ): """ - Test that _should_store_prompts_and_responses_in_spend_logs handles + Test that should_store_prompts_and_responses_in_spend_logs handles case-insensitive string values for store_prompts_in_spend_logs in general_settings. """ # Test case-insensitive string "true" variations @@ -1594,7 +1636,7 @@ def test_should_store_prompts_and_responses_in_spend_logs_case_insensitive_strin {"store_prompts_in_spend_logs": true_value}, ): mock_get_secret_bool.return_value = False # Ensure env var is False - result = _should_store_prompts_and_responses_in_spend_logs() + result = should_store_prompts_and_responses_in_spend_logs() assert result is True, f"Expected True for '{true_value}', got {result}" # Test boolean True @@ -1603,7 +1645,7 @@ def test_should_store_prompts_and_responses_in_spend_logs_case_insensitive_strin {"store_prompts_in_spend_logs": True}, ): mock_get_secret_bool.return_value = False - result = _should_store_prompts_and_responses_in_spend_logs() + result = should_store_prompts_and_responses_in_spend_logs() assert result is True, f"Expected True for boolean True, got {result}" # Test that non-true values fall back to environment variable @@ -1614,22 +1656,22 @@ def test_should_store_prompts_and_responses_in_spend_logs_case_insensitive_strin ): # When env var is True, should return True mock_get_secret_bool.return_value = True - result = _should_store_prompts_and_responses_in_spend_logs() + result = should_store_prompts_and_responses_in_spend_logs() assert result is True, f"Expected True (from env var) for '{false_value}', got {result}" # When env var is False, should return False mock_get_secret_bool.return_value = False - result = _should_store_prompts_and_responses_in_spend_logs() + result = should_store_prompts_and_responses_in_spend_logs() assert result is False, f"Expected False (from env var) for '{false_value}', got {result}" # Test when general_settings doesn't have the key at all with patch("litellm.proxy.proxy_server.general_settings", {}): mock_get_secret_bool.return_value = True - result = _should_store_prompts_and_responses_in_spend_logs() + result = should_store_prompts_and_responses_in_spend_logs() assert result is True, "Expected True (from env var) when key missing, got False" mock_get_secret_bool.return_value = False - result = _should_store_prompts_and_responses_in_spend_logs() + result = should_store_prompts_and_responses_in_spend_logs() assert result is False, "Expected False (from env var) when key missing, got True" @@ -1662,7 +1704,7 @@ def test_get_spend_logs_metadata_guardrail_info_fallback_from_metadata(): assert result["guardrail_information"] is None -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_sanitize_guardrail_information_redacts_all_prompt_carrying_fields_when_flag_false( mock_should_store, ): @@ -1698,7 +1740,7 @@ def test_sanitize_guardrail_information_redacts_all_prompt_carrying_fields_when_ assert entry["guardrail_action"] == "NONE" -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_sanitize_guardrail_information_redacts_prompt_fields_when_flag_false( mock_should_store, ): @@ -1762,7 +1804,7 @@ def test_sanitize_guardrail_information_redacts_prompt_fields_when_flag_false( } -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_sanitize_guardrail_information_preserves_guardrail_usage_when_flag_false( mock_should_store, ): @@ -1794,7 +1836,7 @@ def test_sanitize_guardrail_information_preserves_guardrail_usage_when_flag_fals assert entry["guardrail_usage"] == {"topicPolicyUnits": 1, "contentPolicyUnits": 1, "wordPolicyUnits": 0} -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_sanitize_guardrail_information_passthrough_when_flag_true( mock_should_store, ): @@ -1817,13 +1859,13 @@ def test_sanitize_guardrail_information_passthrough_when_flag_true( assert result == guardrail_info -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_sanitize_guardrail_information_none_passthrough(mock_should_store): mock_should_store.return_value = False assert _sanitize_guardrail_information_for_spend_logs(None) is None -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_sanitize_guardrail_information_normalizes_bare_dict_input(mock_should_store): """ Regression: xecguard (xecguard.py:246) assigns a bare dict to @@ -1855,7 +1897,7 @@ def test_sanitize_guardrail_information_normalizes_bare_dict_input(mock_should_s assert entry["start_time"] == 1.0 -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_sanitize_guardrail_information_drops_non_dict_items_in_list(mock_should_store): """ A stray non-dict item in the list (e.g. from a buggy caller that @@ -1874,7 +1916,7 @@ def test_sanitize_guardrail_information_drops_non_dict_items_in_list(mock_should assert result == [{"guardrail_name": "x", "guardrail_response": REDACTED_BY_LITELM_STRING}] -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_sanitize_guardrail_information_preserves_absent_prompt_fields(mock_should_store): """ Entries that never carried guardrail_request or guardrail_response must @@ -2290,7 +2332,7 @@ def test_sanitize_request_body_strips_secret_fields(): assert sanitized["messages"] == [{"role": "user", "content": "hi"}] -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_proxy_server_request_payload_excludes_secret_fields(mock_should_store): """ End-to-end test: when the proxy_server_request body contains @@ -2370,7 +2412,7 @@ def test_redact_prompt_leaks_empty_string(): assert _redact_prompt_leaks_in_error_string("") == "" -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_sanitize_error_information_redacts_when_not_storing_prompts( mock_should_store, ): @@ -2398,7 +2440,7 @@ def test_sanitize_error_information_redacts_when_not_storing_prompts( assert sanitized["llm_provider"] == "openai" -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_sanitize_error_information_skips_redaction_when_storing_prompts( mock_should_store, ): @@ -2420,7 +2462,7 @@ def test_sanitize_error_information_skips_redaction_when_storing_prompts( assert REDACTED_BY_LITELM_STRING not in sanitized["error_message"] -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_sanitize_error_information_caps_size_regardless_of_prompt_flag( mock_should_store, ): @@ -2451,7 +2493,7 @@ def test_sanitize_error_information_none_passthrough(): assert _sanitize_error_information_for_spend_logs(None) is None -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_sanitize_error_information_reproduces_lit_2992(mock_should_store): # Mirrors the reproduced row body from LIT-2992 — a RateLimitError whose # message embeds 178 pydantic validation errors, each carrying a full @@ -2536,7 +2578,7 @@ def test_redact_prompt_leaks_handles_unterminated_value(): assert REDACTED_BY_LITELM_STRING in redacted -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_sanitize_error_information_redacts_traceback_when_not_storing_prompts( mock_should_store, ): @@ -2568,7 +2610,7 @@ def test_sanitize_error_information_redacts_traceback_when_not_storing_prompts( assert "ValueError: invalid request" in sanitized["traceback"] -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_sanitize_error_information_skips_traceback_redaction_when_storing_prompts( mock_should_store, ): @@ -2676,7 +2718,7 @@ def test_redact_prompt_leaks_combined_quoted_key_and_pydantic_assignment(): assert redacted.count(REDACTED_BY_LITELM_STRING) >= 2 -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_sanitize_error_information_redacts_pydantic_assignment_form( mock_should_store, ): @@ -3272,7 +3314,7 @@ def test_get_logging_payload_hashes_bearer_prefixed_api_key(): ) -@patch("litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs") +@patch("litellm.proxy.spend_tracking.spend_tracking_utils.should_store_prompts_and_responses_in_spend_logs") def test_sanitize_guardrail_information_preserves_headroom_compression_token_stats( mock_should_store, ): diff --git a/tests/test_litellm/proxy/test_collector.py b/tests/test_litellm/proxy/test_collector.py new file mode 100644 index 00000000000..ac3e1f5e566 --- /dev/null +++ b/tests/test_litellm/proxy/test_collector.py @@ -0,0 +1,229 @@ +import asyncio +import logging +from collections.abc import Callable, Iterator +from pathlib import Path +from typing import Final + +import pytest +import uvloop + +from litellm._logging import verbose_logger, verbose_proxy_logger, verbose_router_logger +from litellm.proxy.collector import ( + SpendEventConsumer, + address_argument, + apply_log_level, + pod_pgbouncer_database_url, +) +from litellm.proxy.db.pgbouncer import PgBouncerError, PgBouncerSettings +from litellm.proxy.spend_tracking.spend_event_producer import ( + AddressError, + SpendEventProducer, + TcpAddress, + UnixAddress, + open_collector_connection, +) + + +class _Handler: + def __init__(self, fail_on: bytes | None = None) -> None: + self.lines: list[bytes] = [] # mutable-ok: test double records the events the consumer handed over + self._fail_on = fail_on + + async def __call__(self, line: bytes) -> None: + if line == self._fail_on: + raise RuntimeError("pipeline failed") + self.lines.append(line) + + +async def _no_fallback(line: bytes) -> None: + raise AssertionError(f"unexpected fallback for {line!r}") + + +class _Fallback: + def __init__(self) -> None: + self.lines: list[bytes] = [] # mutable-ok: test double records the events run in-process + + async def __call__(self, line: bytes) -> None: + self.lines.append(line) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("transport", ["unix", "tcp"]) +async def test_consumer_handles_each_producer_line_once_in_order(tmp_path: Path, transport: str): + handler: Final = _Handler(fail_on=b"event-3\n") + consumer: Final = SpendEventConsumer(handler) + server: Final = await consumer.serve( + UnixAddress(path=str(tmp_path / "spend.sock")) if transport == "unix" else TcpAddress("127.0.0.1", 0) + ) + address: Final = ( + UnixAddress(path=str(tmp_path / "spend.sock")) + if transport == "unix" + else TcpAddress("127.0.0.1", server.sockets[0].getsockname()[1]) + ) + producer: Final = SpendEventProducer( + address=address, on_unavailable="fallback", buffer_size=100, connect_timeout=1.0, fallback=_no_fallback + ) + for i in range(6): + await producer.publish(f"event-{i}\n".encode()) + await producer.close(drain_timeout=5.0) + + server.close() + assert await consumer.drain(timeout=5.0) == 0 + assert handler.lines == [f"event-{i}\n".encode() for i in range(6) if i != 3] + assert (consumer.received, consumer.handled, consumer.failed) == (6, 5, 1) + + +@pytest.mark.asyncio +async def test_consumer_discards_a_truncated_trailing_event(tmp_path: Path): + handler: Final = _Handler() + consumer: Final = SpendEventConsumer(handler) + address: Final = UnixAddress(path=str(tmp_path / "spend.sock")) + server: Final = await consumer.serve(address) + _, writer = await open_collector_connection(address, timeout=1.0) + writer.write(b"whole\npartial-without-newline") + await writer.drain() + writer.close() + await writer.wait_closed() + await asyncio.sleep(0.05) + + server.close() + assert await consumer.drain(timeout=5.0) == 0 + assert handler.lines == [b"whole\n"] + assert consumer.received == 1 + + +@pytest.mark.asyncio +async def test_drain_reports_producers_still_connected_after_the_timeout(tmp_path: Path): + consumer: Final = SpendEventConsumer(_Handler()) + address: Final = UnixAddress(path=str(tmp_path / "spend.sock")) + server: Final = await consumer.serve(address) + _, writer = await open_collector_connection(address, timeout=1.0) + await asyncio.sleep(0.05) + + server.close() + assert await consumer.drain(timeout=0.1) == 1 + writer.close() + await writer.wait_closed() + assert await consumer.drain(timeout=5.0) == 0 + + +@pytest.mark.asyncio +async def test_graceful_stop_hands_the_producer_over_to_its_fallback_without_losing_events(tmp_path: Path): + handler: Final = _Handler() + fallback: Final = _Fallback() + consumer: Final = SpendEventConsumer(handler) + address: Final = UnixAddress(path=str(tmp_path / "spend.sock")) + server: Final = await consumer.serve(address) + producer: Final = SpendEventProducer( + address=address, on_unavailable="fallback", buffer_size=100, connect_timeout=1.0, fallback=fallback + ) + await producer.publish(b"event-1\n") + await asyncio.sleep(0.05) + + server.close() + draining: Final = asyncio.ensure_future(consumer.drain(timeout=5.0)) + await asyncio.sleep(0.05) + await producer.publish(b"event-2\n") + await producer.close(drain_timeout=5.0) + + assert await draining == 0 + assert handler.lines == [b"event-1\n"] + assert fallback.lines == [b"event-2\n"] + assert (producer.stats().sent, producer.stats().fallback) == (1, 1) + + +@pytest.mark.parametrize("loop_factory", [asyncio.new_event_loop, uvloop.new_event_loop], ids=["asyncio", "uvloop"]) +def test_drain_still_hands_over_live_producers_when_another_connection_already_died( + tmp_path: Path, loop_factory: Callable[[], asyncio.AbstractEventLoop] +): + """A transport the loop force-closed under a busy handler must not abort the half-close of the others.""" + + async def scenario() -> tuple[int, list[bytes]]: + release: Final = asyncio.Event() + + async def slow_handler(line: bytes) -> None: + await release.wait() + + consumer: Final = SpendEventConsumer(slow_handler) + address: Final = UnixAddress(path=str(tmp_path / "spend.sock")) + server: Final = await consumer.serve(address) + _, dead = await open_collector_connection(address, timeout=1.0) + dead.write(b"stuck\n") + await dead.drain() + await asyncio.sleep(0.05) + for connection in consumer._open_connections: # pyright: ignore[reportPrivateUsage] # force-close like uvloop does on a socket error + connection.transport.close() + dead.close() + fallback: Final = _Fallback() + producer: Final = SpendEventProducer( + address=address, on_unavailable="fallback", buffer_size=100, connect_timeout=1.0, fallback=fallback + ) + await producer.publish(b"event-1\n") + await asyncio.sleep(0.05) + + server.close() + draining: Final = asyncio.ensure_future(consumer.drain(timeout=0.5)) + await asyncio.sleep(0.05) + await producer.publish(b"event-2\n") + await producer.close(drain_timeout=5.0) + still_open: Final = await draining + release.set() + await asyncio.sleep(0.05) + return still_open, fallback.lines + + with asyncio.Runner(loop_factory=loop_factory) as runner: + still_open, fallback_lines = runner.run(scenario()) + + assert still_open == 2 + assert fallback_lines == [b"event-2\n"] + + +def test_address_argument(): + assert address_argument((), default="unix:///tmp/x.sock") == "unix:///tmp/x.sock" + assert address_argument(("--address", "tcp://127.0.0.1:4100"), default="unix:///tmp/x.sock") == ( + "tcp://127.0.0.1:4100" + ) + assert isinstance(address_argument(("--listen", "x"), default="unix:///tmp/x.sock"), AddressError) + + +def test_pod_pgbouncer_database_url_points_at_the_proxy_containers_pooler(): + """With pgbouncer on, the sidecar must not open its own upstream connections but share the pod's pooler.""" + upstream: Final = "postgresql://u:p@db.internal:5432/litellm?schema=public" + environ: Final = {"DATABASE_URL": upstream} + assert pod_pgbouncer_database_url(PgBouncerSettings(enabled=False), environ, token_auth=False) is None + assert ( + pod_pgbouncer_database_url(PgBouncerSettings(enabled=True, port=6543), environ, token_auth=False) + == "postgresql://u:p@127.0.0.1:6543/litellm?schema=public&pgbouncer=true" + ) + assert isinstance(pod_pgbouncer_database_url(PgBouncerSettings(enabled=True), {}, token_auth=False), PgBouncerError) + + +def test_pod_pgbouncer_database_url_goes_direct_under_token_auth(): + """The proxy's pgbouncer only knows the token that container minted, so the sidecar must mint its own upstream.""" + iam_upstream: Final = "postgresql://u@db.internal:5432/litellm?schema=public" + assert ( + pod_pgbouncer_database_url(PgBouncerSettings(enabled=True), {"DATABASE_URL": iam_upstream}, token_auth=True) + is None + ) + assert pod_pgbouncer_database_url(PgBouncerSettings(enabled=True), {}, token_auth=True) is None + + +@pytest.fixture +def restore_log_levels() -> Iterator[None]: + loggers: Final = (verbose_logger, verbose_router_logger, verbose_proxy_logger) + levels: Final = tuple(logger.level for logger in loggers) + yield + for logger, level in zip(loggers, levels, strict=True): + logger.setLevel(level) + + +@pytest.mark.usefixtures("restore_log_levels") +@pytest.mark.parametrize( + ("litellm_log", "expected"), + [("DEBUG", logging.DEBUG), ("info", logging.INFO), (None, logging.WARNING), ("loud", logging.WARNING)], +) +def test_apply_log_level_mirrors_the_proxy_env_contract(litellm_log: str | None, expected: int): + verbose_proxy_logger.setLevel(logging.WARNING) + apply_log_level(litellm_log) + assert verbose_proxy_logger.isEnabledFor(expected) + assert not verbose_proxy_logger.isEnabledFor(expected - 10) From 960fc4b114815cd84ac08f2f98d675286e1c9d67 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 10 Sep 2026 17:16:15 -0700 Subject: [PATCH 05/16] feat(newrelic): export team max and remaining budget gauges to the Metric API (#40542) Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../integrations/newrelic/newrelic_metrics.py | 39 +++++++- litellm/types/integrations/newrelic.py | 19 +++- .../newrelic/test_newrelic_metrics.py | 92 ++++++++++++++++++- 3 files changed, 141 insertions(+), 9 deletions(-) diff --git a/litellm/integrations/newrelic/newrelic_metrics.py b/litellm/integrations/newrelic/newrelic_metrics.py index 25dbfc2bdb2..da952b78d3f 100644 --- a/litellm/integrations/newrelic/newrelic_metrics.py +++ b/litellm/integrations/newrelic/newrelic_metrics.py @@ -5,8 +5,9 @@ NR Reference API: https://docs.newrelic.com/docs/data-apis/ingest-apis/metric-ap `async_log_success_event` / `async_log_failure_event` queue one record per request; at flush the queue is aggregated by (team, model group, model, provider, status) -into count/summary metrics. `interval.ms` is the real window between flushes, -computed at flush time. +into count/summary metrics, plus one max/remaining budget gauge pair per team +taken from the team's latest record. `interval.ms` is the real window between +flushes, computed at flush time. Team-scoped by construction: the ingest key is injected explicitly and there is deliberately no environment-variable fallback, so a team's metrics are never sent @@ -47,11 +48,14 @@ from litellm.types.integrations.newrelic import ( NEWRELIC_METRIC_PROMPT_TOKENS, NEWRELIC_METRIC_REQUEST_DURATION_MS, NEWRELIC_METRIC_REQUESTS, + NEWRELIC_METRIC_TEAM_MAX_BUDGET, + NEWRELIC_METRIC_TEAM_REMAINING_BUDGET, NEWRELIC_METRIC_TOTAL_TOKENS, NEWRELIC_METRICS_MAX_BATCH_SIZE, NEWRELIC_METRICS_MAX_DRAIN_PASSES, NEWRELIC_METRICS_MAX_RETRY_QUEUE_SIZE, NewRelicCountMetric, + NewRelicGaugeMetric, NewRelicMetric, NewRelicMetricCommon, NewRelicMetricEnvelope, @@ -98,6 +102,8 @@ def _metric_record_from_payload(standard_logging_object: StandardLoggingPayload) completion_tokens=int(standard_logging_object.get("completion_tokens") or 0), total_tokens=int(standard_logging_object.get("total_tokens") or 0), duration_ms=float(standard_logging_object.get("response_time") or 0.0) * 1000.0, + team_max_budget=metadata.get("user_api_key_team_max_budget") if metadata else None, + team_spend=metadata.get("user_api_key_team_spend") if metadata else None, ) @@ -140,6 +146,33 @@ def _bucket_metrics(bucket_records: tuple[NewRelicMetricRecord, ...]) -> tuple[N return (*count_metrics, summary_metric) +def _team_budget_gauges(record: NewRelicMetricRecord) -> tuple[NewRelicMetric, ...]: + team_max_budget: Final = record.team_max_budget + if team_max_budget is None: + return () + attributes: Final[Mapping[str, str]] = { # mutable-ok: JSON leaf; safe_dumps stringifies MappingProxyType + key: value[:NEWRELIC_METRIC_ATTRIBUTE_MAX_LEN] + for key, value in (("team_id", record.team_id), ("team_alias", record.team_alias)) + if value + } + remaining_budget: Final = team_max_budget - (record.team_spend or 0.0) - record.response_cost + return ( + NewRelicGaugeMetric( + name=NEWRELIC_METRIC_TEAM_MAX_BUDGET, type="gauge", value=team_max_budget, attributes=attributes + ), + NewRelicGaugeMetric( + name=NEWRELIC_METRIC_TEAM_REMAINING_BUDGET, type="gauge", value=remaining_budget, attributes=attributes + ), + ) + + +def _team_budget_metrics(records: tuple[NewRelicMetricRecord, ...]) -> tuple[NewRelicMetric, ...]: + latest_by_team: Final[Mapping[str, NewRelicMetricRecord]] = MappingProxyType( + {record.team_id: record for record in records if record.team_id} + ) + return tuple(gauge for record in latest_by_team.values() for gauge in _team_budget_gauges(record)) + + def build_metric_payload( records: tuple[NewRelicMetricRecord, ...], *, @@ -158,7 +191,7 @@ def build_metric_payload( "timestamp": int(window_start * 1000), "interval.ms": interval_ms, } - return (NewRelicMetricEnvelope(common=common, metrics=metrics),) + return (NewRelicMetricEnvelope(common=common, metrics=(*metrics, *_team_budget_metrics(records))),) class NewRelicMetricsLogger(CustomBatchLogger): diff --git a/litellm/types/integrations/newrelic.py b/litellm/types/integrations/newrelic.py index 36e4d02c2a8..b5905ad0b93 100644 --- a/litellm/types/integrations/newrelic.py +++ b/litellm/types/integrations/newrelic.py @@ -27,9 +27,9 @@ NEWRELIC_METRIC_ENDPOINT_BY_REGION: Final[Mapping[str, str]] = MappingProxyType( NEWRELIC_DEFAULT_REGION: Final = "us" #: Metric API caps a payload at 2000 data points / 1MB compressed; each queued -#: record expands to at most 6 metrics, so cap the per-flush record count well -#: below that. -NEWRELIC_METRICS_MAX_BATCH_SIZE: Final = 250 +#: record expands to at most 6 bucket metrics plus 2 team budget gauges (8), so cap +#: the per-flush record count well below 2000 / 8. +NEWRELIC_METRICS_MAX_BATCH_SIZE: Final = 200 #: Hard cap on records retained across failed flushes (5xx/network requeue). #: Beyond this the oldest records are dropped. @@ -48,6 +48,8 @@ NEWRELIC_METRIC_PROMPT_TOKENS: Final = "litellm.tokens.prompt" NEWRELIC_METRIC_COMPLETION_TOKENS: Final = "litellm.tokens.completion" NEWRELIC_METRIC_TOTAL_TOKENS: Final = "litellm.tokens.total" NEWRELIC_METRIC_REQUEST_DURATION_MS: Final = "litellm.request.duration_ms" +NEWRELIC_METRIC_TEAM_MAX_BUDGET: Final = "litellm.team.max_budget" +NEWRELIC_METRIC_TEAM_REMAINING_BUDGET: Final = "litellm.team.remaining_budget" class NewRelicSummaryValue(TypedDict): @@ -66,6 +68,13 @@ class NewRelicCountMetric(TypedDict): attributes: ReadOnly[Mapping[str, str]] +class NewRelicGaugeMetric(TypedDict): + name: ReadOnly[str] + type: ReadOnly[Literal["gauge"]] + value: ReadOnly[float] + attributes: ReadOnly[Mapping[str, str]] + + class NewRelicSummaryMetric(TypedDict): name: ReadOnly[str] type: ReadOnly[Literal["summary"]] @@ -73,7 +82,7 @@ class NewRelicSummaryMetric(TypedDict): attributes: ReadOnly[Mapping[str, str]] -NewRelicMetric = NewRelicCountMetric | NewRelicSummaryMetric +NewRelicMetric = NewRelicCountMetric | NewRelicGaugeMetric | NewRelicSummaryMetric #: ``interval.ms`` has a dot in it, so the functional TypedDict form is required. @@ -108,6 +117,8 @@ class NewRelicMetricRecord: completion_tokens: int total_tokens: int duration_ms: float + team_max_budget: float | None = None + team_spend: float | None = None @property def bucket_key(self) -> tuple[str, str, str, str, str, str]: diff --git a/tests/test_litellm/integrations/newrelic/test_newrelic_metrics.py b/tests/test_litellm/integrations/newrelic/test_newrelic_metrics.py index 9c75e0b0a47..a9074e89806 100644 --- a/tests/test_litellm/integrations/newrelic/test_newrelic_metrics.py +++ b/tests/test_litellm/integrations/newrelic/test_newrelic_metrics.py @@ -24,6 +24,8 @@ from litellm.types.integrations.newrelic import ( NEWRELIC_METRIC_PROMPT_TOKENS, NEWRELIC_METRIC_REQUEST_DURATION_MS, NEWRELIC_METRIC_REQUESTS, + NEWRELIC_METRIC_TEAM_MAX_BUDGET, + NEWRELIC_METRIC_TEAM_REMAINING_BUDGET, NEWRELIC_METRIC_TOTAL_TOKENS, NewRelicMetricRecord, ) @@ -40,6 +42,8 @@ def _record( completion_tokens=20, total_tokens=30, duration_ms=100.0, + team_max_budget=None, + team_spend=None, ) -> NewRelicMetricRecord: return NewRelicMetricRecord( team_id=team_id, @@ -53,12 +57,28 @@ def _record( completion_tokens=completion_tokens, total_tokens=total_tokens, duration_ms=duration_ms, + team_max_budget=team_max_budget, + team_spend=team_spend, ) -def _standard_logging_object(team_id="team-a", response_cost=0.25) -> dict: +def _standard_logging_object( + team_id="team-a", response_cost=0.25, team_max_budget: float | None = None, team_spend: float | None = None +) -> dict: + budget_metadata = { + key: value + for key, value in ( + ("user_api_key_team_max_budget", team_max_budget), + ("user_api_key_team_spend", team_spend), + ) + if value is not None + } return { - "metadata": {"user_api_key_team_id": team_id, "user_api_key_team_alias": f"{team_id}-alias"}, + "metadata": { + "user_api_key_team_id": team_id, + "user_api_key_team_alias": f"{team_id}-alias", + **budget_metadata, + }, "model_group": "gpt-4o-group", "model": "gpt-4o", "custom_llm_provider": "openai", @@ -204,6 +224,74 @@ class TestBuildMetricPayload: assert "model_group" not in attributes +class TestTeamBudgetGauges: + def test_latest_record_per_team_drives_one_gauge_pair(self): + records = ( + _record(team_id="team-a", model="gpt-4o", response_cost=0.5, team_max_budget=100.0, team_spend=10.0), + _record(team_id="team-a", model="claude-4", response_cost=2.0, team_max_budget=100.0, team_spend=10.5), + _record(team_id="team-b", response_cost=1.0, team_max_budget=None, team_spend=3.0), + _record(team_id="", team_alias="", response_cost=1.0, team_max_budget=50.0, team_spend=1.0), + ) + payload = build_metric_payload(records, window_start=1_000.0, now=1_005.0) + + max_budget_gauges = _metrics_by_name(payload, NEWRELIC_METRIC_TEAM_MAX_BUDGET) + remaining_gauges = _metrics_by_name(payload, NEWRELIC_METRIC_TEAM_REMAINING_BUDGET) + assert [(m["type"], m["value"], m["attributes"]) for m in max_budget_gauges] == [ + ("gauge", 100.0, {"team_id": "team-a", "team_alias": "team-a-alias"}) + ] + assert [(m["type"], m["attributes"]) for m in remaining_gauges] == [ + ("gauge", {"team_id": "team-a", "team_alias": "team-a-alias"}) + ] + assert remaining_gauges[0]["value"] == pytest.approx(100.0 - 10.5 - 2.0) + assert len(_metrics_by_name(payload, NEWRELIC_METRIC_COST_USD)) == 4 + + def test_missing_team_spend_counts_only_this_request(self): + payload = build_metric_payload( + (_record(response_cost=0.25, team_max_budget=10.0, team_spend=None),), window_start=1_000.0, now=1_005.0 + ) + + assert _metrics_by_name(payload, NEWRELIC_METRIC_TEAM_REMAINING_BUDGET)[0]["value"] == pytest.approx(9.75) + + @pytest.mark.asyncio + async def test_budget_gauges_reach_the_metric_api_from_standard_logging_metadata(self): + logger = _make_logger() + logger.async_client.post = AsyncMock(return_value=_response(202)) + slo = _standard_logging_object(response_cost=0.25, team_max_budget=20.0, team_spend=4.5) + + await logger.async_log_success_event( + kwargs={"standard_logging_object": slo}, response_obj={}, start_time=None, end_time=None + ) + await logger.flush_queue() + + body = json.loads(gzip.decompress(logger.async_client.post.await_args.kwargs["data"]).decode("utf-8")) + by_name = {m["name"]: m for m in body[0]["metrics"]} + assert by_name[NEWRELIC_METRIC_TEAM_MAX_BUDGET] == { + "name": NEWRELIC_METRIC_TEAM_MAX_BUDGET, + "type": "gauge", + "value": 20.0, + "attributes": {"team_id": "team-a", "team_alias": "team-a-alias"}, + } + assert by_name[NEWRELIC_METRIC_TEAM_REMAINING_BUDGET]["type"] == "gauge" + assert by_name[NEWRELIC_METRIC_TEAM_REMAINING_BUDGET]["value"] == pytest.approx(15.25) + assert by_name[NEWRELIC_METRIC_COST_USD]["value"] == 0.25 + + @pytest.mark.asyncio + async def test_no_budget_metadata_sends_no_gauges(self): + logger = _make_logger() + logger.async_client.post = AsyncMock(return_value=_response(202)) + + await logger.async_log_success_event( + kwargs={"standard_logging_object": _standard_logging_object()}, + response_obj={}, + start_time=None, + end_time=None, + ) + await logger.flush_queue() + + body = json.loads(gzip.decompress(logger.async_client.post.await_args.kwargs["data"]).decode("utf-8")) + assert {m["type"] for m in body[0]["metrics"]} == {"count", "summary"} + + class TestQueueAndFlush: @pytest.mark.asyncio async def test_log_event_queues_record_from_standard_logging_object(self): From 1f7c4d6784411efa2c9622aa6ea7dcf3e0320e89 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Thu, 10 Sep 2026 17:18:56 -0700 Subject: [PATCH 06/16] test: fix stale completion response fixtures --- tests/llm_translation/test_azure_o_series.py | 8 ++++++++ tests/llm_translation/test_azure_openai.py | 5 +++++ tests/llm_translation/test_openai.py | 6 ++++++ tests/local_testing/test_completion.py | 9 +++++++-- 4 files changed, 26 insertions(+), 2 deletions(-) diff --git a/tests/llm_translation/test_azure_o_series.py b/tests/llm_translation/test_azure_o_series.py index 1a2d672af71..ce7e614cbe2 100644 --- a/tests/llm_translation/test_azure_o_series.py +++ b/tests/llm_translation/test_azure_o_series.py @@ -159,15 +159,23 @@ def test_azure_o_series_routing(): def test_openai_o_series_max_retries_0(mock_get_openai_client): import litellm + mock_get_openai_client.return_value.chat.completions.with_raw_response.create.return_value.headers = {} + mock_get_openai_client.return_value.chat.completions.with_raw_response.create.return_value.parse.return_value = ( + ModelResponse(choices=[{"message": {"role": "assistant", "content": "Hello"}}]) + ) litellm.set_verbose = True response = litellm.completion( model="azure/o1-preview", messages=[{"role": "user", "content": "hi"}], max_retries=0, + api_key="fake-key", + api_base="https://fake-azure.openai.azure.com", + api_version="2024-10-21", ) mock_get_openai_client.assert_called_once() assert mock_get_openai_client.call_args.kwargs["max_retries"] == 0 + assert response.choices[0].message.content == "Hello" @pytest.mark.asyncio diff --git a/tests/llm_translation/test_azure_openai.py b/tests/llm_translation/test_azure_openai.py index 0fa72b45ed8..e6528e77749 100644 --- a/tests/llm_translation/test_azure_openai.py +++ b/tests/llm_translation/test_azure_openai.py @@ -335,6 +335,10 @@ def test_azure_gpt_4o_with_tool_call_and_response_format(api_version): ] with patch.object(client.chat.completions.with_raw_response, "create") as mock_post: + mock_post.return_value.headers = {} + mock_post.return_value.parse.return_value = litellm.ModelResponse( + choices=[{"message": {"role": "assistant", "content": InvestigationOutput().model_dump_json()}}] + ) response = litellm.completion( model="azure/gpt-4.1-mini", messages=[ @@ -362,6 +366,7 @@ def test_azure_gpt_4o_with_tool_call_and_response_format(api_version): assert "response_format" in mock_post.call_args.kwargs else: assert "response_format" not in mock_post.call_args.kwargs + assert response.choices[0].message.content == InvestigationOutput().model_dump_json() def test_map_openai_params(): diff --git a/tests/llm_translation/test_openai.py b/tests/llm_translation/test_openai.py index 2b9abdec5d0..af4ba85d58e 100644 --- a/tests/llm_translation/test_openai.py +++ b/tests/llm_translation/test_openai.py @@ -292,15 +292,21 @@ class TestOpenAIChatCompletion(BaseLLMChatTest): def test_openai_max_retries_0(mock_get_openai_client): import litellm + mock_get_openai_client.return_value.chat.completions.with_raw_response.create.return_value.headers = {} + mock_get_openai_client.return_value.chat.completions.with_raw_response.create.return_value.parse.return_value = ( + ModelResponse(choices=[{"message": {"role": "assistant", "content": "Hello"}}]) + ) litellm.set_verbose = True response = litellm.completion( model="gpt-4o-mini", messages=[{"role": "user", "content": "hi"}], max_retries=0, + api_key="fake-key", ) mock_get_openai_client.assert_called_once() assert mock_get_openai_client.call_args.kwargs["max_retries"] == 0 + assert response.choices[0].message.content == "Hello" @patch("litellm.main.openai_chat_completions._get_openai_client") diff --git a/tests/local_testing/test_completion.py b/tests/local_testing/test_completion.py index f8f23ea015a..43ed57f63af 100644 --- a/tests/local_testing/test_completion.py +++ b/tests/local_testing/test_completion.py @@ -3999,10 +3999,14 @@ def test_completion_novita_ai(): openai_client = OpenAI(api_key="fake-key") with patch.object( - openai_client.chat.completions, "create", new=MagicMock() + openai_client.chat.completions.with_raw_response, "create" ) as mock_call: + mock_call.return_value.headers = {} + mock_call.return_value.parse.return_value = litellm.ModelResponse( + choices=[{"message": {"role": "assistant", "content": "Hello"}}] + ) try: - completion( + response = completion( model="novita/meta-llama/llama-3.3-70b-instruct", messages=messages, client=openai_client, @@ -4010,6 +4014,7 @@ def test_completion_novita_ai(): ) mock_call.assert_called_once() + assert response.choices[0].message.content == "Hello" # Verify model is passed correctly assert ( From c7a41c35d5cdf2fc2d270c6fee6c793729eeee06 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 11 Sep 2026 00:48:25 +0000 Subject: [PATCH 07/16] perf(mock): emit admission-time usage chunk on streaming mock_response (#40637) * perf(mock): emit admission-time usage chunk on streaming mock_response Streaming mock_response chunks carried no usage, so the chunk builder re-tokenized the whole prompt in Python after the stream ended even when budget reservation had already counted it at admission. The mock streaming generators now yield a final usage-only chunk carrying the admission prompt count (same completion count as the non-streaming path). Without an admission count the old tokenizer fallback stays. * fix(mock): type the mock stream generators and keep the usage chunk on the content stream id Review follow-up: the usage-only chunk was built with a fresh id, so CustomStreamWrapper switched response_id for the finish-reason and usage chunks. It now copies the content stream id. The generators also get full parameter and return annotations. --------- Co-authored-by: yassin --- litellm/main.py | 4 +- litellm/utils.py | 35 +++- .../test_streaming_chunk_builder_utils.py | 49 ++++++ tests/test_litellm/test_main.py | 158 ++++++++++++++++++ tests/test_litellm/test_utils.py | 90 ++++++++++ 5 files changed, 329 insertions(+), 7 deletions(-) diff --git a/litellm/main.py b/litellm/main.py index 1a4beb787bc..10fb32828f1 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -939,7 +939,7 @@ def mock_completion( if kwargs.get("acompletion", False) is True: return CustomStreamWrapper( completion_stream=async_mock_completion_streaming_obj( - model_response, mock_response=mock_response, model=model, n=n + model_response, mock_response=mock_response, model=model, n=n, prompt_tokens=prompt_tokens ), model=model, custom_llm_provider="openai", @@ -947,7 +947,7 @@ def mock_completion( ) return CustomStreamWrapper( completion_stream=mock_completion_streaming_obj( - model_response, mock_response=mock_response, model=model, n=n + model_response, mock_response=mock_response, model=model, n=n, prompt_tokens=prompt_tokens ), model=model, custom_llm_provider="openai", diff --git a/litellm/utils.py b/litellm/utils.py index aced4c9b312..b93938a480e 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -66,6 +66,7 @@ from litellm.constants import ( DEFAULT_EMBEDDING_PARAM_VALUES, DEFAULT_MAX_LRU_CACHE_SIZE, DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT, + DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT, DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET, DEFAULT_TRIM_RATIO, FUNCTION_DEFINITION_TOKEN_COUNT, @@ -278,7 +279,7 @@ except (ImportError, AttributeError, TypeError): # Convert to str (if necessary) claude_json_str = json.dumps(json_data) import importlib.metadata -from collections.abc import Callable, Iterable, Mapping, Sequence +from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast, get_args from litellm import utils as litellm_utils @@ -7001,7 +7002,26 @@ class TextCompletionStreamWrapper: raise StopAsyncIteration -def mock_completion_streaming_obj(model_response, mock_response, model, n: int | None = None): +def mock_stream_usage_chunk(model_response: ModelResponseStream, model: str, prompt_tokens: int) -> ModelResponseStream: + return ModelResponseStream( + id=model_response.id, + choices=[], # mutable-ok: ModelResponseStream only treats a list as explicit choices, a tuple gets a default choice + model=model, + usage=Usage( + prompt_tokens=prompt_tokens, + completion_tokens=DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT, + total_tokens=prompt_tokens + DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT, + ), + ) + + +def mock_completion_streaming_obj( + model_response: ModelResponseStream, + mock_response: str | MockException | ModelResponseStream, + model: str, + n: int | None = None, + prompt_tokens: int | None = None, +) -> Iterator[ModelResponseStream]: if isinstance(mock_response, litellm.MockException): raise mock_response if isinstance(mock_response, ModelResponseStream): @@ -7021,14 +7041,17 @@ def mock_completion_streaming_obj(model_response, mock_response, model, n: int | _all_choices.append(_streaming_choice) model_response.choices = _all_choices yield model_response + if prompt_tokens is not None: + yield mock_stream_usage_chunk(model_response, model=model, prompt_tokens=prompt_tokens) async def async_mock_completion_streaming_obj( - model_response, + model_response: ModelResponseStream, mock_response: str | MockException | ModelResponseStream, - model, + model: str, n: int | None = None, -): + prompt_tokens: int | None = None, +) -> AsyncIterator[ModelResponseStream]: if isinstance(mock_response, litellm.MockException): raise mock_response if isinstance(mock_response, ModelResponseStream): @@ -7048,6 +7071,8 @@ async def async_mock_completion_streaming_obj( _all_choices.append(_streaming_choice) model_response.choices = _all_choices yield model_response + if prompt_tokens is not None: + yield mock_stream_usage_chunk(model_response, model=model, prompt_tokens=prompt_tokens) ########## Reading Config File ############################ diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py index 626b8a63b20..efe4209c1c9 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py @@ -1554,3 +1554,52 @@ def test_stream_chunk_builder_reads_role_from_first_frame_with_choices() -> None assert response is not None assert response.choices[0].message.role == "user" assert response.choices[0].message.content == "Hi" + + +def _fail_prompt_token_count() -> int: + raise AssertionError("prompt tokens must come from the usage chunk, not the tokenizer") + + +def test_calculate_usage_reads_prompt_tokens_from_mock_stream_usage_chunk_without_tokenizer_fallback() -> None: + from litellm.utils import mock_completion_streaming_obj + + chunks: Final = list( + mock_completion_streaming_obj( + ModelResponseStream(model="gpt-5.4-mini"), + mock_response="ok", + model="gpt-5.4-mini", + prompt_tokens=51234, + ) + ) + assert chunks[-1].choices == [] + + usage: Final = ChunkProcessor(chunks=chunks).calculate_usage( + chunks=chunks, + model="gpt-5.4-mini", + completion_output="ok", + count_prompt_tokens=_fail_prompt_token_count, + ) + + assert usage.prompt_tokens == 51234 + assert usage.completion_tokens == chunks[-1].usage.completion_tokens + assert usage.total_tokens == 51234 + usage.completion_tokens + + +def test_calculate_usage_falls_back_to_prompt_counter_when_mock_stream_has_no_admission_count() -> None: + from litellm.utils import mock_completion_streaming_obj + + chunks: Final = list( + mock_completion_streaming_obj( + ModelResponseStream(model="gpt-5.4-mini"), mock_response="ok", model="gpt-5.4-mini" + ) + ) + assert all(chunk.choices for chunk in chunks) + + usage: Final = ChunkProcessor(chunks=chunks).calculate_usage( + chunks=chunks, + model="gpt-5.4-mini", + completion_output="ok", + count_prompt_tokens=lambda: 77, + ) + + assert usage.prompt_tokens == 77 diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index f8841af9750..a36ca229981 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -2421,6 +2421,164 @@ def test_mock_completion_usage_falls_back_to_default_without_admission_count(): assert response.usage.prompt_tokens == litellm_main.DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT +_ADMISSION_INPUT_TOKENS: Final = 51234 +_ADMISSION_METADATA: Final = { + "user_api_key_budget_reservation": {"reserved_cost": 1.0, "input_tokens": _ADMISSION_INPUT_TOKENS} +} +_MOCK_STREAM_MESSAGES: Final = [{"role": "user", "content": "hello " * 200}] +_STREAM_CHUNK_BUILDER_TOKEN_COUNTER: Final = "litellm.litellm_core_utils.streaming_chunk_builder_utils.token_counter" + + +def _prompt_token_counter_calls(token_counter: MagicMock) -> list[object]: + return [call for call in token_counter.call_args_list if call.kwargs.get("messages") is not None] + + +def _client_usage_chunks(chunks: list[ModelResponseStream]) -> list[Usage]: + return [chunk.usage for chunk in chunks if getattr(chunk, "usage", None) is not None] + + +@pytest.mark.parametrize("n", (None, 2)) +def test_mock_completion_stream_usage_reports_admission_input_tokens_without_tokenizer_fallback(n: int | None): + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + chunks = list( + litellm.completion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + n=n, + stream_options={"include_usage": True}, + metadata=_ADMISSION_METADATA, + ) + ) + + usage_chunks = _client_usage_chunks(chunks) + assert len(usage_chunks) == 1 + assert usage_chunks[0].prompt_tokens == _ADMISSION_INPUT_TOKENS + assert usage_chunks[0].completion_tokens == litellm_main.DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT + assert usage_chunks[0].total_tokens == _ADMISSION_INPUT_TOKENS + usage_chunks[0].completion_tokens + assert _prompt_token_counter_calls(token_counter) == [] + assert all(chunk.choices for chunk in chunks[:-1]) + assert {chunk.id for chunk in chunks} == {chunks[0].id} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("n", (None, 2)) +async def test_mock_acompletion_stream_usage_reports_admission_input_tokens_without_tokenizer_fallback( + n: int | None, +): + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + response = await litellm.acompletion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + n=n, + stream_options={"include_usage": True}, + litellm_metadata=_ADMISSION_METADATA, + ) + chunks = [chunk async for chunk in response] + + usage_chunks = _client_usage_chunks(chunks) + assert len(usage_chunks) == 1 + assert usage_chunks[0].prompt_tokens == _ADMISSION_INPUT_TOKENS + assert usage_chunks[0].total_tokens == _ADMISSION_INPUT_TOKENS + usage_chunks[0].completion_tokens + assert _prompt_token_counter_calls(token_counter) == [] + assert all(chunk.choices for chunk in chunks[:-1]) + assert {chunk.id for chunk in chunks} == {chunks[0].id} + + +def test_mock_completion_stream_without_include_usage_hides_usage_chunk_but_logs_admission_count(): + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + chunks = list( + litellm.completion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + metadata=_ADMISSION_METADATA, + ) + ) + + assert _client_usage_chunks(chunks) == [] + assert all(len(chunk.choices) == 1 for chunk in chunks) + assert chunks[-1]._hidden_params["usage"].prompt_tokens == _ADMISSION_INPUT_TOKENS + assert _prompt_token_counter_calls(token_counter) == [] + + +def test_mock_completion_stream_without_admission_count_falls_back_to_tokenizer(): + expected_prompt_tokens: Final = litellm.token_counter(model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES) + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + chunks = list( + litellm.completion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + stream_options={"include_usage": True}, + metadata={"user_api_key_budget_reservation": {"reserved_cost": 1.0}}, + ) + ) + + usage_chunks = _client_usage_chunks(chunks) + assert len(usage_chunks) == 1 + assert usage_chunks[0].prompt_tokens == expected_prompt_tokens + assert usage_chunks[0].total_tokens == expected_prompt_tokens + usage_chunks[0].completion_tokens + assert len(_prompt_token_counter_calls(token_counter)) >= 1 + + +@pytest.mark.asyncio +async def test_mock_acompletion_stream_without_admission_count_falls_back_to_tokenizer(): + expected_prompt_tokens: Final = litellm.token_counter(model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES) + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + response = await litellm.acompletion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + stream_options={"include_usage": True}, + ) + chunks = [chunk async for chunk in response] + + usage_chunks = _client_usage_chunks(chunks) + assert len(usage_chunks) == 1 + assert usage_chunks[0].prompt_tokens == expected_prompt_tokens + assert len(_prompt_token_counter_calls(token_counter)) >= 1 + + +def test_mock_completion_stream_and_non_stream_report_the_same_admission_usage(): + non_stream = litellm.completion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + metadata=_ADMISSION_METADATA, + ) + chunks = list( + litellm.completion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + stream_options={"include_usage": True}, + metadata=_ADMISSION_METADATA, + ) + ) + + stream_usage: Final = _client_usage_chunks(chunks)[0] + assert (non_stream.usage.prompt_tokens, non_stream.usage.completion_tokens, non_stream.usage.total_tokens) == ( + stream_usage.prompt_tokens, + stream_usage.completion_tokens, + stream_usage.total_tokens, + ) + + def test_mock_completion_stream_with_model_response(): """Test that mock_completion correctly handles stream=True with a ModelResponse as mock_response.""" from litellm import completion diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 8b186be43e5..2bdde43ff2d 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -15,6 +15,7 @@ from jsonschema import validate import litellm from litellm._internal_context import is_internal_call +from litellm.constants import DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT from litellm._logging import ( CorrelationContextFilter, JsonFormatter, @@ -6208,3 +6209,92 @@ def test_load_credentials_from_list_fills_kwargs_from_the_loaded_credential_with "api_key": "sk-from-db", } assert _credential_warnings(caplog) == [] + + +_MOCK_STREAM_ID: Final = "chatcmpl-mock-stream" +_ChunkSnapshot = tuple[str, tuple[str | None, ...], Usage | None] + + +def _snapshot(chunk: ModelResponseStream) -> _ChunkSnapshot: + return chunk.id, tuple(choice.delta.content for choice in chunk.choices), getattr(chunk, "usage", None) + + +def _mock_stream_snapshots(mock_response: object, prompt_tokens: int | None) -> list[_ChunkSnapshot]: + from litellm.utils import mock_completion_streaming_obj + + return [ + _snapshot(chunk) + for chunk in mock_completion_streaming_obj( + ModelResponseStream(id=_MOCK_STREAM_ID, model="gpt-5.4-mini"), + mock_response=mock_response, + model="gpt-5.4-mini", + prompt_tokens=prompt_tokens, + ) + ] + + +async def _async_mock_stream_snapshots(mock_response: object, prompt_tokens: int | None) -> list[_ChunkSnapshot]: + from litellm.utils import async_mock_completion_streaming_obj + + return [ + _snapshot(chunk) + async for chunk in async_mock_completion_streaming_obj( + ModelResponseStream(id=_MOCK_STREAM_ID, model="gpt-5.4-mini"), + mock_response=mock_response, + model="gpt-5.4-mini", + prompt_tokens=prompt_tokens, + ) + ] + + +_CONTENT_SNAPSHOTS: Final = [(_MOCK_STREAM_ID, (content,), None) for content in ("hel", "lo ", "wor", "ld")] + + +def _assert_trailing_usage_chunk(snapshots: list[_ChunkSnapshot], prompt_tokens: int) -> None: + assert snapshots[:-1] == _CONTENT_SNAPSHOTS + chunk_id, choices, usage = snapshots[-1] + assert chunk_id == _MOCK_STREAM_ID + assert choices == () + assert usage is not None + assert usage.prompt_tokens == prompt_tokens + assert usage.completion_tokens == DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT + assert usage.total_tokens == prompt_tokens + usage.completion_tokens + + +@pytest.mark.parametrize("prompt_tokens", (51234, 0)) +def test_mock_completion_streaming_obj_emits_usage_chunk_with_admission_prompt_tokens(prompt_tokens: int) -> None: + _assert_trailing_usage_chunk(_mock_stream_snapshots("hello world", prompt_tokens), prompt_tokens) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("prompt_tokens", (51234, 0)) +async def test_async_mock_completion_streaming_obj_emits_usage_chunk_with_admission_prompt_tokens( + prompt_tokens: int, +) -> None: + _assert_trailing_usage_chunk(await _async_mock_stream_snapshots("hello world", prompt_tokens), prompt_tokens) + + +def test_mock_completion_streaming_obj_emits_no_usage_chunk_without_admission_prompt_tokens() -> None: + assert _mock_stream_snapshots("hello world", None) == _CONTENT_SNAPSHOTS + + +@pytest.mark.asyncio +async def test_async_mock_completion_streaming_obj_emits_no_usage_chunk_without_admission_prompt_tokens() -> None: + assert await _async_mock_stream_snapshots("hello world", None) == _CONTENT_SNAPSHOTS + + +def test_mock_completion_streaming_obj_passes_prebuilt_stream_chunk_through_without_usage_chunk() -> None: + prebuilt: Final = ModelResponseStream( + model="gpt-5.4-mini", choices=[StreamingChoices(index=0, delta=Delta(role="assistant", content="prebuilt"))] + ) + + assert _mock_stream_snapshots(prebuilt, 51234) == [(prebuilt.id, ("prebuilt",), None)] + + +@pytest.mark.asyncio +async def test_async_mock_completion_streaming_obj_raises_mock_exception_before_usage_chunk() -> None: + mock_exception: Final = litellm.MockException( + status_code=500, message="boom", llm_provider="openai", model="gpt-5.4-mini" + ) + with pytest.raises(litellm.MockException): + await _async_mock_stream_snapshots(mock_exception, 51234) From 985ac6b6a514d8abec939b897720326a8404eb69 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 10 Sep 2026 17:48:46 -0700 Subject: [PATCH 08/16] fix(rate_limiter): skip non-Latin-1 x-litellm-priority header on /v1/messages (#40636) A team or key priority that is not Latin-1 encodable (for example CJK text) was attached as a response header by the dynamic rate limiter v3 post-call hook, and Starlette then raised UnicodeEncodeError while writing headers, turning a successful /v1/messages call into HTTP 500. The header is now omitted for such values while x-litellm-rate-limiter-version and the v3 rate limit headers are still attached. Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/hooks/dynamic_rate_limiter_v3.py | 12 ++++++++- .../hooks/test_dynamic_rate_limiter_v3.py | 27 +++++++++++++++++++ 2 files changed, 38 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py index a074f02f4e8..0339cf4dfea 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py @@ -59,6 +59,10 @@ def _get_priority_settings() -> "PriorityReservationSettings": return settings +def _is_latin1_encodable(value: object) -> bool: + return all(ord(char) < 256 for char in str(value)) + + class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): """ Saturation-aware priority-based rate limiter using v3 infrastructure. @@ -666,7 +670,13 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): if response_has_hidden_params(response): priority: Final = self._get_priority_from_user_api_key_dict(user_api_key_dict=user_api_key_dict) additional_headers: Final = ensure_response_additional_headers(response) - additional_headers["x-litellm-priority"] = priority or "default" + priority_header: Final = priority or "default" + if _is_latin1_encodable(priority_header): + additional_headers["x-litellm-priority"] = priority_header + else: + verbose_proxy_logger.debug( + "Skipping x-litellm-priority header: priority %r is not Latin-1 encodable", priority + ) additional_headers["x-litellm-rate-limiter-version"] = "v3" return response diff --git a/tests/test_litellm/proxy/hooks/test_dynamic_rate_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_dynamic_rate_limiter_v3.py index 0cd6b4ede9c..527449bbc48 100644 --- a/tests/test_litellm/proxy/hooks/test_dynamic_rate_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_dynamic_rate_limiter_v3.py @@ -1918,3 +1918,30 @@ async def test_post_call_success_hook_leaves_raw_provider_dict_untouched(): ) assert response == {"id": "msg_123", "type": "message", "role": "assistant", "content": []} + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("team_metadata", "expected_priority_header"), + [ + ({"priority": "优先"}, None), + ({"priority": "high"}, "high"), + ({}, "default"), + ], +) +async def test_post_call_success_hook_priority_header_is_always_http_encodable(team_metadata, expected_priority_header): + from starlette.responses import Response + + handler = DynamicRateLimitHandler(internal_usage_cache=DualCache()) + response = {"id": "msg_123", "type": "message", "role": "assistant", "content": [], "_hidden_params": {}} + + await handler.async_post_call_success_hook( + data={"model": "anthropic-haiku"}, + user_api_key_dict=UserAPIKeyAuth(team_id="team-1", team_metadata=team_metadata), + response=response, + ) + + additional_headers = response["_hidden_params"]["additional_headers"] + http_response = Response(headers={key: str(value) for key, value in additional_headers.items()}) + assert http_response.headers.get("x-litellm-priority") == expected_priority_header + assert http_response.headers["x-litellm-rate-limiter-version"] == "v3" From 5ea2f9698295a0e53735a9f8a001262de5a7dee6 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Thu, 10 Sep 2026 17:59:22 -0700 Subject: [PATCH 09/16] feat(ui): link the Organization cell on the Virtual Keys page too Same treatment as User, Team and Created By in the previous commit: the Organization column rendered the alias as dead text, so it now goes through IdentityCell with an orgDetailHref. Claude-Session: https://claude.ai/code/session_01NfwfQhamRNnSqgXMUjf3h4 --- .../VirtualKeysPage/VirtualKeysTable.test.tsx | 19 +++++++++++++++++-- .../VirtualKeysPage/keyTableColumns.tsx | 12 ++++++------ 2 files changed, 23 insertions(+), 8 deletions(-) diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx index 57786d71c66..c303dba697b 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/VirtualKeysTable.test.tsx @@ -476,12 +476,14 @@ it("should display 'Default Proxy Admin' for user_id when value is 'default_user describe("entity links out of the key rows", () => { const keyRow = async () => (await screen.findByText("Test Key Alias")).closest("tr") as HTMLElement; - const enableCreatedByColumn = async (user: ReturnType) => { + const enableColumn = async (user: ReturnType, title: string) => { await user.click(screen.getByRole("button", { name: "Columns" })); - await user.click(await screen.findByText("Created By")); + await user.click(await screen.findByText(title)); await user.keyboard("{Escape}"); }; + const enableCreatedByColumn = (user: ReturnType) => enableColumn(user, "Created By"); + it("points the User and Team cells at their detail pages", async () => { renderWithProviders(); @@ -493,6 +495,19 @@ describe("entity links out of the key rows", () => { expect(within(row).getByRole("link", { name: "Test Team" })).toHaveAttribute("href", "/ui/teams?team=team-1"); }); + it("points the Organization cell at the org's detail page", async () => { + mockUseKeys.mockReturnValue(keysResult([{ ...mockKey, org_id: "org-1" }])); + const user = userEvent.setup(); + renderWithProviders(); + await enableColumn(user, "Organization"); + + const row = await keyRow(); + expect(within(row).getByRole("link", { name: "Test Organization" })).toHaveAttribute( + "href", + "/ui/organizations?org=org-1", + ); + }); + it("points the Created By cell at the creator's detail page", async () => { mockUseKeys.mockReturnValue( keysResult([ diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx index c87bc90fb47..9dd4b1c2d59 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx @@ -16,7 +16,7 @@ import { StatusBadge, type StatusTone, } from "@/components/shared/table_cells"; -import { teamDetailHref, userDetailHref } from "@/utils/entityLinks"; +import { orgDetailHref, teamDetailHref, userDetailHref } from "@/utils/entityLinks"; import { DEFAULT_PROXY_ADMIN_USER_ID } from "@/utils/sentinels"; import DefaultProxyAdminTag from "../common_components/DefaultProxyAdminTag"; @@ -218,12 +218,12 @@ export const getKeyTableColumns = ({ const orgId = info.getValue() as string | null; if (!orgId) return "-"; const org = organizations.find((o) => o.organization_id === orgId); - const displayValue = org?.organization_alias || orgId; - const width = info.cell.column.getSize(); return ( - - {displayValue} - + ); }, }, From d1a1cda144391eefecd21ecc6d8d096af04f8782 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Thu, 10 Sep 2026 18:07:12 -0700 Subject: [PATCH 10/16] test: respect optional logging payload fields --- tests/local_testing/test_custom_callback_input.py | 4 ++-- tests/logging_callback_tests/test_datadog.py | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/tests/local_testing/test_custom_callback_input.py b/tests/local_testing/test_custom_callback_input.py index f0f24a6e6b2..834570091bd 100644 --- a/tests/local_testing/test_custom_callback_input.py +++ b/tests/local_testing/test_custom_callback_input.py @@ -1076,7 +1076,7 @@ def test_standard_logging_payload(model, turn_off_message_logging): ) ) - keys_list = list(StandardLoggingPayload.__annotations__.keys()) + keys_list = list(StandardLoggingPayload.__required_keys__) for k in keys_list: assert ( @@ -1190,7 +1190,7 @@ def test_standard_logging_payload_audio(turn_off_message_logging, stream): ) ) - keys_list = list(StandardLoggingPayload.__annotations__.keys()) + keys_list = list(StandardLoggingPayload.__required_keys__) for k in keys_list: assert ( diff --git a/tests/logging_callback_tests/test_datadog.py b/tests/logging_callback_tests/test_datadog.py index 83a652e8884..7ac9ac0b5ad 100644 --- a/tests/logging_callback_tests/test_datadog.py +++ b/tests/logging_callback_tests/test_datadog.py @@ -270,7 +270,7 @@ async def test_datadog_logging_http_request(): message = json.loads(body[0]["message"]) print("logged message", json.dumps(message, indent=4)) - expected_message_fields = StandardLoggingPayload.__annotations__.keys() + expected_message_fields = StandardLoggingPayload.__required_keys__ for field in expected_message_fields: assert field in message, f"Field '{field}' is missing from the message" From 4fbe2276a1ece2f58d047c60660cfe630fdbbede Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 10 Sep 2026 18:15:24 -0700 Subject: [PATCH 11/16] fix(logging): finish response metadata before the sync logging thread reads it (#39869) * fix(logging): finish response metadata before the sync logging thread reads it The async and sync client wrappers handed the response to the threaded success handler before computing its cost, call id, and api_base, so that thread inserted into the same metadata dict the request coroutine was still iterating and a finished chat completion turned into a 500 (dictionary changed size during iteration). Metadata is now finalized first, and the merge and header copies snapshot their dicts before iterating. * fix(logging): snapshot metadata with a dict copy and drop redundant comment Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(logging): copy metadata via dict.copy and dedupe Final import Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/litellm_core_utils/litellm_logging.py | 39 +++--- litellm/utils.py | 112 +++++++++++------- .../test_litellm_logging.py | 59 +++++++++ tests/test_litellm/test_utils.py | 46 +++++++ 4 files changed, 193 insertions(+), 63 deletions(-) diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 893bfeefff8..6ee68ab21c5 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -295,6 +295,9 @@ def _get_provider_request_id(original_exception: Exception) -> str | None: # Cache custom pricing keys as frozenset for O(1) lookups instead of looping through 49 keys _CUSTOM_PRICING_KEYS: Final[frozenset[str]] = frozenset(CustomPricingLiteLLMParams.model_fields.keys()) _MODEL_INFO_CUSTOM_PRICING_KEYS: Final[frozenset[str]] = _CUSTOM_PRICING_KEYS | DEPLOYMENT_SCOPED_PRICING_FIELDS +_UNSERIALIZABLE_METADATA_KEYS: Final[frozenset[str]] = frozenset( + ("user_api_key_auth", "user_api_key_budget_reservation") +) sentry_sdk_instance = None capture_exception = None @@ -5386,23 +5389,23 @@ class StandardLoggingPayloadSetup: Returns: dict: Merged metadata with user API key fields taking precedence """ - merged_metadata: Final[dict] = {} - - # Start with metadata (user API key fields) - but skip non-serializable objects - if litellm_params.get("metadata") and isinstance(litellm_params.get("metadata"), dict): - for key, value in litellm_params["metadata"].items(): - # Skip non-serializable objects like UserAPIKeyAuth - if key in {"user_api_key_auth", "user_api_key_budget_reservation"}: - continue - merged_metadata[key] = value - - # Then merge litellm_metadata (model-related fields) - this will NOT overwrite existing keys - if litellm_params.get("litellm_metadata") and isinstance(litellm_params.get("litellm_metadata"), dict): - for key, value in litellm_params["litellm_metadata"].items(): - if key not in merged_metadata: # Don't overwrite existing keys from metadata - merged_metadata[key] = value - - return merged_metadata + metadata: Final = litellm_params.get("metadata") + litellm_metadata: Final = litellm_params.get("litellm_metadata") + user_metadata: Final = MappingProxyType( + { + key: value + for key, value in (metadata.copy().items() if isinstance(metadata, dict) else ()) + if key not in _UNSERIALIZABLE_METADATA_KEYS + } + ) + model_metadata: Final = MappingProxyType( + { + key: value + for key, value in (litellm_metadata.copy().items() if isinstance(litellm_metadata, dict) else ()) + if key not in user_metadata + } + ) + return {**user_metadata, **model_metadata} # mutable-ok: function contract returns a plain dict @staticmethod def get_standard_logging_metadata( @@ -5660,7 +5663,7 @@ class StandardLoggingPayloadSetup: additional_logging_headers[key] = additiona_headers[_key] # Preserve all remaining headers verbatim (e.g. llm_provider-x-request-id) - for k, v in additiona_headers.items(): + for k, v in additiona_headers.copy().items(): if k.lower() not in typed_keys: additional_logging_headers[k] = v diff --git a/litellm/utils.py b/litellm/utils.py index b93938a480e..a765e1b1246 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1196,6 +1196,47 @@ def function_setup( raise e +def _dispatch_success_logging( + logging_obj: LiteLLMLoggingObject, + result: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + is_completion_with_fallbacks: bool, + is_litellm_internal_call: bool, +) -> None: + if not is_litellm_internal_call: + if getattr(logging_obj, "_defer_async_logging", False): + + def _enqueue_deferred_logging() -> None: + asyncio.create_task( + _client_async_logging_helper( + logging_obj=logging_obj, + result=result, + start_time=start_time, + end_time=end_time, + is_completion_with_fallbacks=is_completion_with_fallbacks, + ) + ) + + logging_obj._enqueue_deferred_logging = _enqueue_deferred_logging + else: + asyncio.create_task( + _client_async_logging_helper( + logging_obj=logging_obj, + result=result, + start_time=start_time, + end_time=end_time, + is_completion_with_fallbacks=is_completion_with_fallbacks, + ) + ) + + logging_obj.handle_sync_success_callbacks_for_async_calls( + result=result, + start_time=start_time, + end_time=end_time, + ) + + async def _client_async_logging_helper( logging_obj: LiteLLMLoggingObject, result, @@ -1663,6 +1704,16 @@ def client(original_function): kwargs=kwargs, ) + _update_response_metadata: Final = getattr(sys.modules[__name__], "update_response_metadata") + _update_response_metadata( + result=result, + logging_obj=logging_obj, + model=model, + kwargs=kwargs, + start_time=start_time, + end_time=end_time, + ) + # LOG SUCCESS - handle streaming success logging in the _next_ object, remove `handle_success` once it's deprecated verbose_logger.info("Wrapper: Completed Call, calling success_handler") # Copy the current context to propagate it to the background thread @@ -1677,15 +1728,6 @@ def client(original_function): end_time, ) # RETURN RESULT - update_response_metadata = getattr(sys.modules[__name__], "update_response_metadata") - update_response_metadata( - result=result, - logging_obj=logging_obj, - model=model, - kwargs=kwargs, - start_time=start_time, - end_time=end_time, - ) return result except Exception as e: call_type = original_function.__name__ @@ -1945,48 +1987,20 @@ def client(original_function): args=args, ) - # LOG SUCCESS - handle streaming success logging in the _next_ object - # Internal sub-calls (e.g. emulated file-search steps) share the - # parent's logging obj; skip async logging here so only the outer call bills once. - # NOTE: streaming requests return early (before this point) via - # CustomStreamWrapper, so this block is non-streaming only. - if not _is_litellm_internal_call: - if getattr(logging_obj, "_defer_async_logging", False): - - def _enqueue_deferred_logging() -> None: - asyncio.create_task( - _client_async_logging_helper( - logging_obj=logging_obj, - result=result, - start_time=start_time, - end_time=end_time, - is_completion_with_fallbacks=is_completion_with_fallbacks, - ) - ) - - logging_obj._enqueue_deferred_logging = _enqueue_deferred_logging - else: - asyncio.create_task( - _client_async_logging_helper( - logging_obj=logging_obj, - result=result, - start_time=start_time, - end_time=end_time, - is_completion_with_fallbacks=is_completion_with_fallbacks, - ) - ) - - logging_obj.handle_sync_success_callbacks_for_async_calls( - result=result, - start_time=start_time, - end_time=end_time, - ) # REBUILD EMBEDDING CACHING if ( isinstance(result, EmbeddingResponse) and _caching_handler_response is not None and _caching_handler_response.final_embedding_cached_response is not None ): + _dispatch_success_logging( + logging_obj=logging_obj, + result=result, + start_time=start_time, + end_time=end_time, + is_completion_with_fallbacks=is_completion_with_fallbacks, + is_litellm_internal_call=_is_litellm_internal_call, + ) return _llm_caching_handler._combine_cached_embedding_response_with_api_result( _caching_handler_response=_caching_handler_response, embedding_response=result, @@ -2002,6 +2016,14 @@ def client(original_function): start_time=start_time, end_time=end_time, ) + _dispatch_success_logging( + logging_obj=logging_obj, + result=result, + start_time=start_time, + end_time=end_time, + is_completion_with_fallbacks=is_completion_with_fallbacks, + is_litellm_internal_call=_is_litellm_internal_call, + ) return result except Exception as e: diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index aeaa8f18c26..2f6339dcdbb 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -3,6 +3,7 @@ import contextlib import datetime import os import sys +from collections.abc import Callable from typing import Final, Literal from unittest.mock import AsyncMock, MagicMock, patch @@ -6945,3 +6946,61 @@ def test_classifier_audit_is_not_added_to_other_calls(logging_obj, call_type, or logging_obj.model_call_details["litellm_params"] = {"metadata": {"internal_call_origin": origin}} logging_obj.pre_call(input=[], api_key=None, additional_args={"complete_input_dict": {"input": "embedding"}}) assert logging_obj.classifier_input is None + + +def _run_while_a_thread_grows(target: dict, read: Callable[[], None], reads: int) -> None: + import itertools + import threading + + stop: Final = threading.Event() + + def grow() -> None: + for counter in itertools.count(): + if stop.is_set(): + return + key: Final = f"late_{counter % 64}" + if key in target: + del target[key] + else: + target[key] = counter + + writer: Final = threading.Thread(target=grow, daemon=True) + previous_interval: Final = sys.getswitchinterval() + sys.setswitchinterval(1e-6) + writer.start() + try: + for _ in range(reads): + read() + finally: + stop.set() + writer.join(timeout=5) + sys.setswitchinterval(previous_interval) + + +def test_merge_litellm_metadata_survives_a_thread_growing_metadata_mid_merge(): + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + + metadata: Final = {f"key_{i}": i for i in range(2000)} + litellm_params: Final = {"metadata": metadata, "litellm_metadata": {"model_group": "gpt"}} + + def read() -> None: + merged: Final = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params) + assert merged["key_1999"] == 1999 + assert merged["model_group"] == "gpt" + + _run_while_a_thread_grows(metadata, read, reads=300) + + +def test_get_additional_headers_survives_a_thread_growing_headers_mid_copy(): + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + + headers: Final = {f"llm_provider-x-custom-{i}": str(i) for i in range(2000)} + headers["x-ratelimit-remaining-requests"] = "7" + + def read() -> None: + copied: Final = StandardLoggingPayloadSetup.get_additional_headers(headers) + assert copied is not None + assert copied["x_ratelimit_remaining_requests"] == 7 + assert copied["llm_provider-x-custom-1999"] == "1999" + + _run_while_a_thread_grows(headers, read, reads=300) diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 2bdde43ff2d..7753cb7d770 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -1,9 +1,12 @@ import asyncio +import contextlib import json import logging import os +import queue import threading from datetime import datetime, timedelta, timezone +from collections.abc import Iterator from typing import Final from unittest.mock import AsyncMock, MagicMock, patch @@ -6298,3 +6301,46 @@ async def test_async_mock_completion_streaming_obj_raises_mock_exception_before_ ) with pytest.raises(litellm.MockException): await _async_mock_stream_snapshots(mock_exception, 51234) + + + +@contextlib.contextmanager +def _recording_hidden_params_at_submit(submit_target: str) -> "Iterator[queue.SimpleQueue[dict[str, object]]]": + seen: Final = queue.SimpleQueue() + + def record_submit(_fn, *args, **_kwargs): + response: Final = next(arg for arg in args if isinstance(arg, litellm.ModelResponse)) + seen.put(dict(response._hidden_params)) + return MagicMock() + + with patch(submit_target, side_effect=record_submit): + yield seen + + +@pytest.mark.asyncio +async def test_acompletion_finishes_response_metadata_before_handing_the_response_to_the_logging_thread(monkeypatch): + monkeypatch.setattr(litellm, "success_callback", [lambda kwargs, response, start_time, end_time: None]) + with _recording_hidden_params_at_submit("litellm.litellm_core_utils.litellm_logging.executor.submit") as seen: + await litellm.acompletion( + model="gpt-5.5", + messages=[{"role": "user", "content": "hi"}], + mock_response="Hello there!", + num_retries=0, + ) + snapshot: Final = seen.get_nowait() + assert snapshot["litellm_call_id"] + assert snapshot["response_cost"] is not None + assert snapshot["api_base"] + + +def test_completion_finishes_response_metadata_before_handing_the_response_to_the_logging_thread(): + with _recording_hidden_params_at_submit("litellm.utils.executor.submit") as seen: + litellm.completion( + model="gpt-5.5", + messages=[{"role": "user", "content": "hi"}], + mock_response="Hello there!", + ) + snapshot: Final = seen.get_nowait() + assert snapshot["litellm_call_id"] + assert snapshot["response_cost"] is not None + assert snapshot["api_base"] From 880ccc76a5bcc89b07a7ca337662a902bf56083e Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 10 Sep 2026 18:37:39 -0700 Subject: [PATCH 12/16] fix(streaming): keep admitted mock streams alive with empty stream_options and honor zero prompt counts (#40650) * fix(streaming): keep usage-only chunks from crashing streams with empty stream_options The usage-only chunk branch in CustomStreamWrapper.chunk_creator indexed stream_options["include_usage"] directly, so a caller passing stream_options={} hit a KeyError that surfaced as MidStreamFallbackError. Streaming mock_response with an admission input_tokens count (#40637) now always emits such a chunk, which made the crash reachable. Reuse the send_stream_usage policy computed at init instead. Also annotate the #40637 test bindings with Final. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(streaming): report admitted zero prompt tokens instead of recounting in mock streams Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../litellm_core_utils/streaming_handler.py | 7 +- litellm/main.py | 9 ++ tests/test_litellm/test_main.py | 143 ++++++++++++++---- 3 files changed, 131 insertions(+), 28 deletions(-) diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index db23929e0c3..6a5a8832cc6 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -236,9 +236,11 @@ class CustomStreamWrapper: stream_options=None, make_call: Callable | None = None, _response_headers: dict | httpx.Headers | None = None, + count_prompt_tokens: Callable[[], int] | None = None, ): self.model = model self.make_call = make_call + self.count_prompt_tokens = count_prompt_tokens self.custom_llm_provider = custom_llm_provider self.logging_obj: LiteLLMLoggingObject = logging_obj self.completion_stream = completion_stream @@ -1641,7 +1643,7 @@ class CustomStreamWrapper: except Exception: model_response.choices[0].delta = Delta() else: - if self.stream_options is not None and self.stream_options["include_usage"] is True: + if self.send_stream_usage is True: model_response.choices = [] return model_response self._record_usage_only_chunk(model_response=model_response) @@ -1996,6 +1998,7 @@ class CustomStreamWrapper: chunks=self.chunks, messages=self.messages, logging_obj=self.logging_obj, + count_prompt_tokens=self.count_prompt_tokens, ) except Exception as e: # stream_chunk_builder can re-raise (as APIError) on large agentic @@ -2248,6 +2251,7 @@ class CustomStreamWrapper: chunks=self.chunks, messages=self.messages, logging_obj=self.logging_obj, + count_prompt_tokens=self.count_prompt_tokens, ) except Exception as e: # see sync __next__: a raise from stream_chunk_builder inside this @@ -2371,6 +2375,7 @@ class CustomStreamWrapper: chunks=self.chunks, messages=self.messages if isinstance(self.messages, list) else None, logging_obj=self.logging_obj, + count_prompt_tokens=self.count_prompt_tokens, ) if partial_response is None: return diff --git a/litellm/main.py b/litellm/main.py index 10fb32828f1..17edafcdfca 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -850,6 +850,12 @@ def admission_input_tokens(kwargs: Mapping[str, object]) -> int | None: ) +def admitted_prompt_token_counter(prompt_tokens: int | None) -> Callable[[], int] | None: + if prompt_tokens is None: + return None + return lambda: prompt_tokens + + def mock_completion( model: str, messages: list, @@ -935,6 +941,7 @@ def mock_completion( if stream is True: model_response = ModelResponseStream() + count_prompt_tokens: Final = admitted_prompt_token_counter(prompt_tokens) # don't try to access stream object, if kwargs.get("acompletion", False) is True: return CustomStreamWrapper( @@ -944,6 +951,7 @@ def mock_completion( model=model, custom_llm_provider="openai", logging_obj=logging, + count_prompt_tokens=count_prompt_tokens, ) return CustomStreamWrapper( completion_stream=mock_completion_streaming_obj( @@ -952,6 +960,7 @@ def mock_completion( model=model, custom_llm_provider="openai", logging_obj=logging, + count_prompt_tokens=count_prompt_tokens, ) if isinstance(mock_response, litellm.MockException): raise mock_response diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index a36ca229981..f71225c6fc5 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -2422,9 +2422,13 @@ def test_mock_completion_usage_falls_back_to_default_without_admission_count(): _ADMISSION_INPUT_TOKENS: Final = 51234 -_ADMISSION_METADATA: Final = { - "user_api_key_budget_reservation": {"reserved_cost": 1.0, "input_tokens": _ADMISSION_INPUT_TOKENS} -} + + +def _admission_metadata(input_tokens: int) -> dict[str, object]: # mutable-ok: logging writes into metadata + return {"user_api_key_budget_reservation": {"reserved_cost": 1.0, "input_tokens": input_tokens}} + + +_ADMISSION_METADATA: Final = _admission_metadata(_ADMISSION_INPUT_TOKENS) _MOCK_STREAM_MESSAGES: Final = [{"role": "user", "content": "hello " * 200}] _STREAM_CHUNK_BUILDER_TOKEN_COUNTER: Final = "litellm.litellm_core_utils.streaming_chunk_builder_utils.token_counter" @@ -2440,7 +2444,7 @@ def _client_usage_chunks(chunks: list[ModelResponseStream]) -> list[Usage]: @pytest.mark.parametrize("n", (None, 2)) def test_mock_completion_stream_usage_reports_admission_input_tokens_without_tokenizer_fallback(n: int | None): with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - chunks = list( + chunks: Final = list( litellm.completion( model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES, @@ -2453,7 +2457,7 @@ def test_mock_completion_stream_usage_reports_admission_input_tokens_without_tok ) ) - usage_chunks = _client_usage_chunks(chunks) + usage_chunks: Final = _client_usage_chunks(chunks) assert len(usage_chunks) == 1 assert usage_chunks[0].prompt_tokens == _ADMISSION_INPUT_TOKENS assert usage_chunks[0].completion_tokens == litellm_main.DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT @@ -2469,7 +2473,7 @@ async def test_mock_acompletion_stream_usage_reports_admission_input_tokens_with n: int | None, ): with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - response = await litellm.acompletion( + response: Final = await litellm.acompletion( model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES, mock_response="ok", @@ -2479,9 +2483,9 @@ async def test_mock_acompletion_stream_usage_reports_admission_input_tokens_with stream_options={"include_usage": True}, litellm_metadata=_ADMISSION_METADATA, ) - chunks = [chunk async for chunk in response] + chunks: Final = [chunk async for chunk in response] - usage_chunks = _client_usage_chunks(chunks) + usage_chunks: Final = _client_usage_chunks(chunks) assert len(usage_chunks) == 1 assert usage_chunks[0].prompt_tokens == _ADMISSION_INPUT_TOKENS assert usage_chunks[0].total_tokens == _ADMISSION_INPUT_TOKENS + usage_chunks[0].completion_tokens @@ -2492,7 +2496,7 @@ async def test_mock_acompletion_stream_usage_reports_admission_input_tokens_with def test_mock_completion_stream_without_include_usage_hides_usage_chunk_but_logs_admission_count(): with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - chunks = list( + chunks: Final = list( litellm.completion( model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES, @@ -2509,10 +2513,48 @@ def test_mock_completion_stream_without_include_usage_hides_usage_chunk_but_logs assert _prompt_token_counter_calls(token_counter) == [] +def test_mock_completion_stream_with_empty_stream_options_completes_and_logs_admission_count(): + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + chunks: Final = list( + litellm.completion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + stream_options={}, + metadata=_ADMISSION_METADATA, + ) + ) + + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "ok" + assert _client_usage_chunks(chunks) == [] + assert _prompt_token_counter_calls(token_counter) == [] + + +@pytest.mark.asyncio +async def test_mock_acompletion_stream_with_empty_stream_options_completes_and_logs_admission_count(): + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + response: Final = await litellm.acompletion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + stream_options={}, + litellm_metadata=_ADMISSION_METADATA, + ) + chunks: Final = [chunk async for chunk in response] + + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "ok" + assert _client_usage_chunks(chunks) == [] + assert _prompt_token_counter_calls(token_counter) == [] + + def test_mock_completion_stream_without_admission_count_falls_back_to_tokenizer(): expected_prompt_tokens: Final = litellm.token_counter(model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES) with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - chunks = list( + chunks: Final = list( litellm.completion( model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES, @@ -2524,7 +2566,7 @@ def test_mock_completion_stream_without_admission_count_falls_back_to_tokenizer( ) ) - usage_chunks = _client_usage_chunks(chunks) + usage_chunks: Final = _client_usage_chunks(chunks) assert len(usage_chunks) == 1 assert usage_chunks[0].prompt_tokens == expected_prompt_tokens assert usage_chunks[0].total_tokens == expected_prompt_tokens + usage_chunks[0].completion_tokens @@ -2535,7 +2577,7 @@ def test_mock_completion_stream_without_admission_count_falls_back_to_tokenizer( async def test_mock_acompletion_stream_without_admission_count_falls_back_to_tokenizer(): expected_prompt_tokens: Final = litellm.token_counter(model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES) with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: - response = await litellm.acompletion( + response: Final = await litellm.acompletion( model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES, mock_response="ok", @@ -2543,40 +2585,87 @@ async def test_mock_acompletion_stream_without_admission_count_falls_back_to_tok stream=True, stream_options={"include_usage": True}, ) - chunks = [chunk async for chunk in response] + chunks: Final = [chunk async for chunk in response] - usage_chunks = _client_usage_chunks(chunks) + usage_chunks: Final = _client_usage_chunks(chunks) assert len(usage_chunks) == 1 assert usage_chunks[0].prompt_tokens == expected_prompt_tokens assert len(_prompt_token_counter_calls(token_counter)) >= 1 -def test_mock_completion_stream_and_non_stream_report_the_same_admission_usage(): - non_stream = litellm.completion( +def _usage_triple(usage: Usage) -> tuple[int, int, int]: + return (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) + + +@pytest.mark.parametrize("input_tokens", (_ADMISSION_INPUT_TOKENS, 0)) +def test_mock_completion_stream_and_non_stream_report_the_same_admission_usage(input_tokens: int): + metadata: Final = _admission_metadata(input_tokens) + non_stream: Final = litellm.completion( model="openai/gpt-5.4-mini", messages=_MOCK_STREAM_MESSAGES, mock_response="ok", api_key="mock", - metadata=_ADMISSION_METADATA, + metadata=metadata, ) - chunks = list( - litellm.completion( + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + chunks: Final = list( + litellm.completion( + model="openai/gpt-5.4-mini", + messages=_MOCK_STREAM_MESSAGES, + mock_response="ok", + api_key="mock", + stream=True, + stream_options={"include_usage": True}, + metadata=metadata, + ) + ) + + assert _usage_triple(non_stream.usage) == _usage_triple(_client_usage_chunks(chunks)[0]) + assert non_stream.usage.prompt_tokens == input_tokens + assert _prompt_token_counter_calls(token_counter) == [] + + +@pytest.mark.asyncio +async def test_mock_acompletion_stream_reports_zero_admission_input_tokens_without_tokenizer_fallback(): + with patch(_STREAM_CHUNK_BUILDER_TOKEN_COUNTER, wraps=litellm.token_counter) as token_counter: + response: Final = await litellm.acompletion( model="openai/gpt-5.4-mini", - messages=_MOCK_STREAM_MESSAGES, + messages=[{"role": "user", "content": ""}], mock_response="ok", api_key="mock", stream=True, stream_options={"include_usage": True}, - metadata=_ADMISSION_METADATA, + litellm_metadata=_admission_metadata(0), + ) + chunks: Final = [chunk async for chunk in response] + + usage_chunks: Final = _client_usage_chunks(chunks) + assert len(usage_chunks) == 1 + assert _usage_triple(usage_chunks[0]) == (0, usage_chunks[0].completion_tokens, usage_chunks[0].completion_tokens) + assert _prompt_token_counter_calls(token_counter) == [] + + +def test_mock_text_completion_stream_and_non_stream_report_the_same_zero_admission_usage(): + metadata: Final = _admission_metadata(0) + non_stream: Final = litellm.text_completion( + model="openai/gpt-5.4-mini", prompt="", mock_response="ok", api_key="mock", metadata=metadata + ) + chunks: Final = list( + litellm.text_completion( + model="openai/gpt-5.4-mini", + prompt="", + mock_response="ok", + api_key="mock", + stream=True, + stream_options={"include_usage": True}, + metadata=metadata, ) ) - stream_usage: Final = _client_usage_chunks(chunks)[0] - assert (non_stream.usage.prompt_tokens, non_stream.usage.completion_tokens, non_stream.usage.total_tokens) == ( - stream_usage.prompt_tokens, - stream_usage.completion_tokens, - stream_usage.total_tokens, - ) + stream_usages: Final = tuple(chunk.usage for chunk in chunks if getattr(chunk, "usage", None) is not None) + assert len(stream_usages) == 1 + assert _usage_triple(non_stream.usage) == _usage_triple(stream_usages[0]) + assert non_stream.usage.prompt_tokens == 0 def test_mock_completion_stream_with_model_response(): From 56e2d8846da375ae157e27d1a789030882066f2d Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Thu, 10 Sep 2026 18:08:13 -0700 Subject: [PATCH 13/16] feat(ui): link the entity cells on the team detail page's keys table The team detail page's Virtual Keys table showed Organization ID, User Email, User ID and Created By as dead text, so getting from a key to the org or user behind it meant copying an id and searching for it. Those four cells now render as links, reusing the sentinel-aware href helpers, so default_user_id and the litellm-dashboard team stay plain text instead of pointing at pages that do not exist. The Created By cell was a verbatim copy of the Virtual Keys page's user popover, so that moved into the shared table_cells kit and both tables now use the one implementation. Claude-Session: https://claude.ai/code/session_01NfwfQhamRNnSqgXMUjf3h4 --- .../VirtualKeysPage/keyTableColumns.tsx | 62 +----------- .../shared/table_cells/UserPopoverCell.tsx | 62 ++++++++++++ .../components/shared/table_cells/index.ts | 1 + .../team/TeamVirtualKeysTable.test.tsx | 62 ++++++++++++ .../components/team/TeamVirtualKeysTable.tsx | 94 ++++++++----------- 5 files changed, 168 insertions(+), 113 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/shared/table_cells/UserPopoverCell.tsx diff --git a/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx b/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx index 9dd4b1c2d59..6eea77ae827 100644 --- a/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/VirtualKeysPage/keyTableColumns.tsx @@ -9,17 +9,17 @@ import { HoverCard, HoverCardContent, HoverCardTrigger } from "@/components/ui/h import { Skeleton } from "@/components/ui/skeleton"; import { DateCell, + ENTITY_CELL_TITLE_CLASSES, IdCell, IdentityCell, ModelsCell, SpendBudgetCell, StatusBadge, + UserPopoverCell, type StatusTone, } from "@/components/shared/table_cells"; -import { orgDetailHref, teamDetailHref, userDetailHref } from "@/utils/entityLinks"; -import { DEFAULT_PROXY_ADMIN_USER_ID } from "@/utils/sentinels"; +import { orgDetailHref, teamDetailHref } from "@/utils/entityLinks"; -import DefaultProxyAdminTag from "../common_components/DefaultProxyAdminTag"; import { KeyResponse, Team } from "../key_team_helpers/key_list"; import { Organization } from "../networking"; @@ -29,8 +29,6 @@ interface KeyStatus { tooltip?: string; } -const ENTITY_CELL_TITLE_CLASSES = "font-mono text-xs font-normal"; - const SPEND_BUDGET_SORT_FIELDS: DataTableSortField[] = [ { id: "spend", label: "Spend" }, { id: "max_budget", label: "Budget" }, @@ -66,60 +64,6 @@ const getKeyStatus = (key: KeyResponse): KeyStatus => { }; }; -const UserPopoverCell = ({ - userAlias, - userEmail, - userId, - width, -}: { - userAlias: string | null; - userEmail: string | null; - userId: string | null; - width: number; -}) => { - const displayValue = userAlias || userEmail || userId; - const isDefaultAdmin = userId === DEFAULT_PROXY_ADMIN_USER_ID; - - const popoverContent = ( -
- {[ - { label: "User Alias", value: userAlias }, - { label: "User Email", value: userEmail }, - { label: "User ID", value: userId }, - ].map(({ label, value }) => ( -
- {label} - {value ? ( - - ) : ( - - - )} -
- ))} -
- ); - - const trigger = - isDefaultAdmin && !userAlias && !userEmail ? ( - - ) : ( - - ); - - return ( - - }> - {trigger} - - {popoverContent} - - ); -}; - const InfoHeader = ({ label, tooltip }: { label: string; tooltip: string }) => ( {label} diff --git a/ui/litellm-dashboard/src/components/shared/table_cells/UserPopoverCell.tsx b/ui/litellm-dashboard/src/components/shared/table_cells/UserPopoverCell.tsx new file mode 100644 index 00000000000..223d66d7421 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/table_cells/UserPopoverCell.tsx @@ -0,0 +1,62 @@ +"use client"; + +import DefaultProxyAdminTag from "@/components/common_components/DefaultProxyAdminTag"; +import { HoverCard, HoverCardContent, HoverCardTrigger } from "@/components/ui/hover-card"; +import { userDetailHref } from "@/utils/entityLinks"; +import { DEFAULT_PROXY_ADMIN_USER_ID } from "@/utils/sentinels"; + +import { IdCell } from "./id_cell"; +import { IdentityCell } from "./identity_cell"; + +export const ENTITY_CELL_TITLE_CLASSES = "font-mono text-xs font-normal"; + +interface UserPopoverCellProps { + userAlias: string | null; + userEmail: string | null; + userId: string | null; + width: number; +} + +export function UserPopoverCell({ userAlias, userEmail, userId, width }: UserPopoverCellProps) { + const displayValue = userAlias || userEmail || userId; + const isDefaultAdmin = userId === DEFAULT_PROXY_ADMIN_USER_ID; + + const popoverContent = ( +
+ {[ + { label: "User Alias", value: userAlias }, + { label: "User Email", value: userEmail }, + { label: "User ID", value: userId }, + ].map(({ label, value }) => ( +
+ {label} + {value ? ( + + ) : ( + - + )} +
+ ))} +
+ ); + + const trigger = + isDefaultAdmin && !userAlias && !userEmail ? ( + + ) : ( + + ); + + return ( + + }> + {trigger} + + {popoverContent} + + ); +} diff --git a/ui/litellm-dashboard/src/components/shared/table_cells/index.ts b/ui/litellm-dashboard/src/components/shared/table_cells/index.ts index 99b54f9aad0..34e3701eed4 100644 --- a/ui/litellm-dashboard/src/components/shared/table_cells/index.ts +++ b/ui/litellm-dashboard/src/components/shared/table_cells/index.ts @@ -13,3 +13,4 @@ export { ModelsCell } from "./models_cell"; export { MoneyCell } from "./money_cell"; export { SpendBudgetCell } from "./spend_budget_cell"; export { StatusBadge, type StatusTone } from "./status_badge"; +export { UserPopoverCell, ENTITY_CELL_TITLE_CLASSES } from "./UserPopoverCell"; diff --git a/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.test.tsx b/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.test.tsx index 0d9d0988aa1..76df0d12539 100644 --- a/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.test.tsx @@ -10,6 +10,8 @@ vi.mock("@/app/(dashboard)/hooks/keys/useKeys", () => ({ useKeys: vi.fn(), })); +vi.mock("next/navigation", () => ({ useRouter: () => ({ push: vi.fn() }) })); + vi.mock("../key_team_helpers/fetch_available_models_team_key", () => ({ getModelDisplayName: vi.fn((model: string) => model), })); @@ -384,4 +386,64 @@ describe("TeamVirtualKeysTable", () => { expect(screen.getByText("Key Info View")).toBeInTheDocument(); }); }); + + describe("entity links out of the key rows", () => { + const renderRow = async (key: KeyResponse, organization: Organization | null = null) => { + mockUseKeys.mockReturnValue({ + data: { keys: [key], total_count: 1, current_page: 1, total_pages: 1 } as KeysResponse, + isPending: false, + isFetching: false, + refetch: vi.fn(), + } as any); + renderWithProviders(); + return (await screen.findByText(key.key_alias as string)).closest("tr") as HTMLElement; + }; + + it("points the Organization ID cell at the org's detail page", async () => { + const row = await renderRow(createMockKey({ organization_id: null }), mockOrganization); + expect(within(row).getByRole("link", { name: "org-123" })).toHaveAttribute( + "href", + "/ui/organizations?org=org-123", + ); + }); + + it("points the User Email and User ID cells at the owning user's detail page", async () => { + const row = await renderRow( + createMockKey({ user_id: "user-1", user: { user_id: "user-1", user_email: "alice@example.com" } }), + ); + expect(within(row).getByRole("link", { name: "alice@example.com" })).toHaveAttribute( + "href", + "/ui/users?user=user-1", + ); + expect(within(row).getByRole("link", { name: "user-1" })).toHaveAttribute("href", "/ui/users?user=user-1"); + }); + + it("points the Created By cell at the creator's detail page", async () => { + const row = await renderRow( + createMockKey({ + created_by: "creator-1", + created_by_user: { user_id: "creator-1", user_email: "creator@example.com", user_alias: "The Creator" }, + }), + ); + expect(within(row).getByRole("link", { name: "The Creator" })).toHaveAttribute( + "href", + "/ui/users?user=creator-1", + ); + }); + + it("leaves the default_user_id placeholder unlinked in the User ID and Created By cells", async () => { + const placeholder = { user_id: "default_user_id", user_email: "admin@example.com", user_alias: "Proxy Admin" }; + const row = await renderRow( + createMockKey({ + user_id: placeholder.user_id, + user: placeholder, + created_by: placeholder.user_id, + created_by_user: placeholder, + }), + ); + expect(within(row).getByText("Default Proxy Admin")).toBeInTheDocument(); + expect(within(row).getByText("Proxy Admin")).toBeInTheDocument(); + expect(within(row).queryByRole("link", { name: "Proxy Admin" })).not.toBeInTheDocument(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.tsx b/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.tsx index b0380255b95..5b1b71e060e 100644 --- a/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamVirtualKeysTable.tsx @@ -1,8 +1,14 @@ "use client"; import { useKeys } from "@/app/(dashboard)/hooks/keys/useKeys"; import { SimpleTooltip } from "@/components/ui/tooltip"; -import CopyButton from "@/components/shared/CopyButton"; -import { DateCell, IdCell, MoneyCell } from "@/components/shared/table_cells"; +import { + DateCell, + ENTITY_CELL_TITLE_CLASSES, + IdCell, + IdentityCell, + MoneyCell, + UserPopoverCell, +} from "@/components/shared/table_cells"; import { DataTable, DataTableFilterDrawer, @@ -11,8 +17,9 @@ import { DataTableToolbar, } from "@/components/shared/DataTable"; import { Badge } from "@/components/ui/badge"; -import { HoverCard, HoverCardContent, HoverCardTrigger } from "@/components/ui/hover-card"; import { Input } from "@/components/ui/input"; +import { orgDetailHref, userDetailHref } from "@/utils/entityLinks"; +import { DEFAULT_PROXY_ADMIN_USER_ID } from "@/utils/sentinels"; import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; import { useDebouncedValue } from "@tanstack/react-pacer/debouncer"; import { ColumnDef, ColumnFiltersState, OnChangeFn, PaginationState, SortingState } from "@tanstack/react-table"; @@ -168,7 +175,15 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi header: "Organization ID", size: 140, enableSorting: false, - cell: (info) => (info.getValue() ? info.renderValue() : "-"), + cell: (info) => { + const orgId = info.getValue() as string | null; + if (!orgId) return "-"; + return ( + + + + ); + }, }, { id: "user_email", @@ -179,9 +194,14 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi cell: (info) => { const user = info.getValue() as { user_email?: string } | undefined; const value = user?.user_email; + const userId = info.row.original.user_id; return ( - {value ?? "-"} + ); }, @@ -194,10 +214,16 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi enableSorting: false, cell: (info) => { const userId = info.getValue() as string | null; - const displayValue = userId === "default_user_id" ? "Default Proxy Admin" : userId; + if (userId === DEFAULT_PROXY_ADMIN_USER_ID) { + return ; + } return ( - - {displayValue ?? "-"} + + ); }, @@ -221,53 +247,13 @@ export function TeamVirtualKeysTable({ teamId, teamAlias, organization }: TeamVi const userId = info.getValue() as string | null; if (!userId) return "-"; const { created_by_user } = info.row.original; - const userAlias = created_by_user?.user_alias ?? null; - const userEmail = created_by_user?.user_email ?? null; - const isDefaultAdmin = userId === "default_user_id"; - const displayValue = userAlias || userEmail || userId; - - const popoverContent = ( -
- {[ - { label: "User Alias", value: userAlias }, - { label: "User Email", value: userEmail }, - { label: "User ID", value: userId }, - ].map(({ label, value }) => ( -
- {label} - {value ? ( - - {value} - - - ) : ( - - - )} -
- ))} -
- ); - - if (isDefaultAdmin && !userAlias && !userEmail) { - return ( - - }> - - - {popoverContent} - - ); - } - return ( - - } - > - {displayValue} - - {popoverContent} - + ); }, }, From 06b259e0920792600c35c9951297fb83cd236c35 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Thu, 10 Sep 2026 18:31:19 -0700 Subject: [PATCH 14/16] fix(ui): name the popover copy buttons after the field they copy The shared user popover copied alias, email and ID through three copy buttons that all announced themselves as "Copy ID", so a screen reader could not tell them apart. IdCell now takes the label, defaulting to the old text everywhere else. Also drops the closest("tr") the new link tests used, which put the testing-library/no-node-access budget over its ceiling, and asserts the sentinel row leaves User Email and the admin badge unlinked too. Claude-Session: https://claude.ai/code/session_01NfwfQhamRNnSqgXMUjf3h4 --- .../shared/table_cells/UserPopoverCell.tsx | 2 +- .../shared/table_cells/id_cell.test.tsx | 8 ++++++++ .../components/shared/table_cells/id_cell.tsx | 4 +++- .../team/TeamVirtualKeysTable.test.tsx | 20 ++++++++++--------- 4 files changed, 23 insertions(+), 11 deletions(-) diff --git a/ui/litellm-dashboard/src/components/shared/table_cells/UserPopoverCell.tsx b/ui/litellm-dashboard/src/components/shared/table_cells/UserPopoverCell.tsx index 223d66d7421..eb786e45541 100644 --- a/ui/litellm-dashboard/src/components/shared/table_cells/UserPopoverCell.tsx +++ b/ui/litellm-dashboard/src/components/shared/table_cells/UserPopoverCell.tsx @@ -31,7 +31,7 @@ export function UserPopoverCell({ userAlias, userEmail, userId, width }: UserPop
{label} {value ? ( - + ) : ( - )} diff --git a/ui/litellm-dashboard/src/components/shared/table_cells/id_cell.test.tsx b/ui/litellm-dashboard/src/components/shared/table_cells/id_cell.test.tsx index 395c1815dd0..c715210fdde 100644 --- a/ui/litellm-dashboard/src/components/shared/table_cells/id_cell.test.tsx +++ b/ui/litellm-dashboard/src/components/shared/table_cells/id_cell.test.tsx @@ -71,6 +71,14 @@ describe("IdCell", () => { expect(rowClick).not.toHaveBeenCalled(); }); + it("names the copy button after the field it copies", async () => { + const user = userEvent.setup(); + render(); + expect(screen.queryByRole("button", { name: "Copy ID" })).not.toBeInTheDocument(); + await user.click(screen.getByRole("button", { name: "Copy User Email" })); + expect(copyToClipboardMock).toHaveBeenCalledWith("alice@example.com"); + }); + it("passes dataTestId through to the id element", () => { render(); expect(screen.getByTestId("key-id-cell")).toHaveTextContent("k-1"); diff --git a/ui/litellm-dashboard/src/components/shared/table_cells/id_cell.tsx b/ui/litellm-dashboard/src/components/shared/table_cells/id_cell.tsx index 1b109a86106..33c7f835e64 100644 --- a/ui/litellm-dashboard/src/components/shared/table_cells/id_cell.tsx +++ b/ui/litellm-dashboard/src/components/shared/table_cells/id_cell.tsx @@ -15,6 +15,7 @@ interface IdCellProps { variant?: IdCellVariant; onClick?: (value: string) => void; copyable?: boolean; + copyLabel?: string; truncate?: boolean; fallback?: string; tooltip?: React.ReactNode; @@ -39,6 +40,7 @@ export function IdCell({ variant = "pill", onClick, copyable = false, + copyLabel = "Copy ID", truncate = true, fallback = "-", tooltip, @@ -80,7 +82,7 @@ export function IdCell({ {withTooltip}