From 727a471ae90162a19253e388fd2e8faf3ba1cc93 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Tue, 19 May 2026 10:58:48 -0700 Subject: [PATCH 1/7] [Refactor] UI - Spend Logs: consolidate filter state and extract components (#25847) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * [Refactor] UI - Spend Logs: consolidate filter state, extract components, remove dead code - Lift filter state into index.tsx and pass to hook (removes selectedX vars + sync useEffect) - Move main useQuery into useLogFilterLogic hook (removes isMainQueryEnabled toggle) - Delete dead RequestViewer component (300 lines, replaced by LogDetailsDrawer) - Extract LogsTableToolbar component (search, date range, pagination, live tail) - Extract filter options config to filter_options.ts - Remove dead code: handleRefresh, handleSelectLog, handleCloseDrawer, formatTimeUnit, showFilters/showColumnDropdown state, dropdownRef/filtersRef * Fix PR feedback: use antd Switch instead of Tremor in new file, fix typo * Collapse dual-path filtering into single React Query All 10 filter keys now go through the useQuery — the imperative performSearch / debouncedSearch / backendFilteredLogs path is deleted. Filter values are debounced via useDebouncedValue(300ms) before hitting the query key so text inputs don't fire per-keystroke. Removed: performSearch, debouncedSearch, backendFilteredLogs, lastSearchTimestamp, hasBackendFilters, clientDerivedFilteredLogs, the sort/page/time refetch useEffect, and the filteredLogs chooser memo. * Clean up remaining smells: remove isFetchingDeferred, internalize selectedTimeInterval, fix circular import - Remove useDeferredValue/isButtonLoading — pass logsQuery.isFetching directly - Move selectedTimeInterval into LogsTableToolbar as internal state - Move PaginatedResponse type from index.tsx to log_filter_logic.tsx * Fix quick-select dropdown overlapping sidebar * Fix stale quick-select label after Reset Filters Move selectedTimeInterval back to parent so handleFilterReset can reset it to the 24-hour default. The toolbar receives it as a prop. * refactor useLogFilterLogic tests for controlled-hook + backend-query shape The hook no longer owns filter state or does client-side filtering — it receives filters/setFilters as props and drives filteredLogs from a useQuery over uiSpendLogsCall. Reshape the tests around that contract: introduce a controlled harness that owns filter state, collapse the 10 per-filter assertions into a single it.each over filterKey → API param, and drop the client-side passthrough tests (the .min test file and the "return all logs when no filters" / "empty when logs null" cases) that no longer correspond to any hook behavior. * cover new useLogFilterLogic invariants: activeTab gate, filterByCurrentUser fallback, debounce negative, partial merge Follow-up to the test refactor. Adds coverage for invariants the refactored hook contract introduced but that the first pass didn't assert: - query enablement: expand the single accessToken-null case into an it.each over all four credential props (accessToken, token, userRole, userID), plus a separate test for activeTab !== "request logs" - filterByCurrentUser: when true with a blank User ID filter, the outbound request carries user_id = userID - debounce: also assert the negative case — no call in the first 100ms after a filter change (first waiting out the initial mount fire) - handleFilterChange: partial updates merge without clobbering other filter keys (protects the spread + default-fill semantics) - handleFilterReset: calls setCurrentPage(1) alongside restoring filters * fix typo dropping the live-tail banner border Tailwind silently ignores unknown classes, so border-greem-200 was leaving the auto-refresh banner with only its bg-green-50 fill and no outline. * memoize columns and derived table data in SpendLogsTable The table's columns array, four-pass data pipeline, and sort-change handler were all being rebuilt on every parent render. That made every filter click re-instance all 23 TanStack-Table columns, re-run filter/reduce/map over all rows, and recreate per-row click closures — all before the intentional 300ms debounce timer even got a chance to fire. Local measurement (40 rows, dev mode): filter click → query fires: 1957ms → 1217ms (−38%) Wrap createColumns in useMemo keyed on sortBy/sortOrder, hoist onSortChange into a useCallback, and move the searchedLogs / sessionComposition / sessionRepresentativeMap / filteredData derivations into a single useMemo keyed on filteredLogs.data + searchTerm. These were pre-existing issues on main — not regressions from the hook refactor — but the refactor made them user-visible because the new query debounce put render cost on the critical path. * apply dropdown filters instantly, debounce only text inputs Dropdown selects now bypass the 300ms debounce so a click updates the table immediately. Text inputs (Key Hash, Error Message, Request ID, User ID) still debounce. handleFilterReset also clears the pending debounced value so a half-typed text filter can't re-fire after reset. * fix(ui/spend-logs): restore lost loading/debounce behavior + cover dropped tests Regressions from the spend-logs-view refactor: - debounce the 'Public model / search tool' text filter (was firing a backend query per keystroke) via TEXT_FILTER_KEYS - restore Fetch-button smoothing through table repaint using useDeferredValue on the rendered data (explicit staleness) - show AntDLoadingSpinner during the auth-resolve phase instead of a blank screen on first load - only live-tail-poll while the tab is visible (refetchIntervalInBackground: false) - extract getLiveTailRefetchInterval helper for the poll decision Tests: - LogDetailContent: retries display (>0 / 0 / absent), overhead-absent - log_filter_logic: regression guard that the public-model filter debounces; getLiveTailRefetchInterval unit tests - logs_utils: getTimeRangeDisplay quick-select window labels * test(ui/spend-logs): cover the cold-load auth-not-ready spinner guard Asserts SpendLogsTable shows a loading spinner (not a blank screen) while credentials are unresolved, and renders the table once present. --- .../LogDetailContent.test.tsx | 38 +- .../components/view_logs/LogsTableToolbar.tsx | 243 +++ .../components/view_logs/filter_options.ts | 77 + .../src/components/view_logs/index.test.tsx | 203 +-- .../src/components/view_logs/index.tsx | 1001 ++----------- .../view_logs/log_filter_logic.test.tsx | 1329 ++++++++--------- .../components/view_logs/log_filter_logic.tsx | 397 ++--- .../components/view_logs/logs_utils.test.tsx | 45 + .../view_logs/useLogFilterLogic.min.test.tsx | 62 - 9 files changed, 1303 insertions(+), 2092 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/view_logs/LogsTableToolbar.tsx create mode 100644 ui/litellm-dashboard/src/components/view_logs/filter_options.ts create mode 100644 ui/litellm-dashboard/src/components/view_logs/logs_utils.test.tsx delete mode 100644 ui/litellm-dashboard/tests/view_logs/useLogFilterLogic.min.test.tsx diff --git a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.test.tsx b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.test.tsx index a2da5136755..85a38e26977 100644 --- a/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/LogDetailsDrawer/LogDetailContent.test.tsx @@ -172,12 +172,7 @@ describe("LogDetailContent", () => { }); it("should display loading state when isLoadingDetails is true", () => { - render( - , - ); + render(); expect(screen.getByText("Loading request & response data...")).toBeInTheDocument(); }); @@ -298,6 +293,37 @@ describe("LogDetailContent", () => { expect(screen.getByText("42.50 ms")).toBeInTheDocument(); }); + it("should not display LiteLLM Overhead when litellm_overhead_time_ms is absent from metadata", () => { + render(); + + expect(screen.queryByText("LiteLLM Overhead")).not.toBeInTheDocument(); + }); + + const retriesItem = () => screen.getByText("Retries").closest(".ant-descriptions-item") as HTMLElement; + + it("should display attempted_retries / max_retries for Retries when attempted_retries > 0", () => { + render( + , + ); + + expect(within(retriesItem()).getByText("2 / 3")).toBeInTheDocument(); + }); + + it("should display a green 'None' tag for Retries when attempted_retries is 0", () => { + render(); + + const noneTag = within(retriesItem()).getByText("None"); + expect(noneTag.closest(".ant-tag")).toHaveClass("ant-tag-green"); + }); + + it("should display '-' for Retries when attempted_retries is absent from metadata", () => { + render(); + + expect(within(retriesItem()).getByText("-")).toBeInTheDocument(); + }); + it("should display start and end time in ISO format", () => { render( void; + startTime: string; + onStartTimeChange: (value: string) => void; + endTime: string; + onEndTimeChange: (value: string) => void; + isCustomDate: boolean; + onIsCustomDateChange: (value: boolean) => void; + selectedTimeInterval: { value: number; unit: string }; + onSelectedTimeIntervalChange: (value: { value: number; unit: string }) => void; + isLiveTail: boolean; + onIsLiveTailChange: (value: boolean) => void; + currentPage: number; + onCurrentPageChange: (updater: number | ((prev: number) => number)) => void; + pageSize: number; + isLoading: boolean; + isButtonLoading: boolean; + onRefetch: () => void; + filteredLogs: PaginatedResponse; +} + +export function LogsTableToolbar({ + searchTerm, + onSearchChange, + startTime, + onStartTimeChange, + endTime, + onEndTimeChange, + isCustomDate, + onIsCustomDateChange, + selectedTimeInterval, + onSelectedTimeIntervalChange, + isLiveTail, + onIsLiveTailChange, + currentPage, + onCurrentPageChange, + pageSize, + isLoading, + isButtonLoading, + onRefetch, + filteredLogs, +}: LogsTableToolbarProps) { + const [quickSelectOpen, setQuickSelectOpen] = useState(false); + const quickSelectRef = useRef(null); + + useEffect(() => { + function handleClickOutside(event: MouseEvent) { + if (quickSelectRef.current && !quickSelectRef.current.contains(event.target as Node)) { + setQuickSelectOpen(false); + } + } + document.addEventListener("mousedown", handleClickOutside); + return () => document.removeEventListener("mousedown", handleClickOutside); + }, []); + + const selectedOption = QUICK_SELECT_OPTIONS.find( + (option) => option.value === selectedTimeInterval.value && option.unit === selectedTimeInterval.unit, + ); + const displayLabel = isCustomDate ? getTimeRangeDisplay(isCustomDate, startTime, endTime) : selectedOption?.label; + + return ( + <> +
+
+
+
+ onSearchChange(e.target.value)} + /> + + + +
+ +
+
+ + + {quickSelectOpen && ( +
+
+ {QUICK_SELECT_OPTIONS.map((option) => ( + + ))} +
+ +
+
+ )} +
+ +
+ Live Tail + +
+ + +
+ + {isCustomDate && ( +
+
+ { + onStartTimeChange(e.target.value); + onCurrentPageChange(1); + }} + className="px-3 py-2 border rounded-md text-sm focus:outline-none focus:ring-2 focus:ring-blue-500 focus:border-blue-500" + /> +
+ to +
+ { + onEndTimeChange(e.target.value); + onCurrentPageChange(1); + }} + className="px-3 py-2 border rounded-md text-sm focus:outline-none focus:ring-2 focus:ring-blue-500 focus:border-blue-500" + /> +
+
+ )} +
+ +
+ + Showing {isLoading ? "..." : filteredLogs ? (currentPage - 1) * pageSize + 1 : 0} -{" "} + {isLoading + ? "..." + : filteredLogs + ? Math.min(currentPage * pageSize, filteredLogs.total) + : 0}{" "} + of {isLoading ? "..." : filteredLogs ? filteredLogs.total : 0} results + +
+ + Page {isLoading ? "..." : currentPage} of{" "} + {isLoading ? "..." : filteredLogs ? filteredLogs.total_pages : 1} + + + +
+
+
+
+ {isLiveTail && currentPage === 1 && ( +
+
+ Auto-refreshing every 15 seconds +
+ +
+ )} + + ); +} diff --git a/ui/litellm-dashboard/src/components/view_logs/filter_options.ts b/ui/litellm-dashboard/src/components/view_logs/filter_options.ts new file mode 100644 index 00000000000..59ac58b6745 --- /dev/null +++ b/ui/litellm-dashboard/src/components/view_logs/filter_options.ts @@ -0,0 +1,77 @@ +import FilterTeamDropdown from "../common_components/FilterTeamDropdown"; +import { PaginatedKeyAliasSelect } from "../KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect"; +import { PaginatedModelSelect } from "../ModelSelect/PaginatedModelSelect/PaginatedModelSelect"; +import { FilterOption } from "../molecules/filter"; +import { allEndUsersCall } from "../networking"; +import { ERROR_CODE_OPTIONS } from "./constants"; +import { FILTER_KEYS } from "./log_filter_logic"; + +export function getLogFilterOptions(accessToken: string): FilterOption[] { + return [ + { + name: "Team ID", + label: "Team ID", + customComponent: FilterTeamDropdown, + }, + { + name: "Status", + label: "Status", + isSearchable: false, + options: [ + { label: "Success", value: "success" }, + { label: "Failure", value: "failure" }, + ], + }, + { + name: "Model", + label: "Model", + customComponent: PaginatedModelSelect, + }, + { + name: FILTER_KEYS.PUBLIC_MODEL_OR_SEARCH_TOOL, + label: "Public model / search tool", + isSearchable: false, + }, + { + name: "Key Alias", + label: "Key Alias", + customComponent: PaginatedKeyAliasSelect, + }, + { + name: "End User", + label: "End User", + isSearchable: true, + searchFn: async (searchText: string) => { + const data = await allEndUsersCall(accessToken); + const users = data?.map((u: any) => u.user_id) || []; + const filtered = users.filter((u: string) => u.toLowerCase().includes(searchText.toLowerCase())); + return filtered.map((u: string) => ({ label: u, value: u })); + }, + }, + { + name: "Error Code", + label: "Error Code", + isSearchable: true, + searchFn: async (searchText: string) => { + if (!searchText) return ERROR_CODE_OPTIONS; + const lower = searchText.toLowerCase(); + const filtered = ERROR_CODE_OPTIONS.filter((opt) => opt.label.toLowerCase().includes(lower)); + const isExactValue = ERROR_CODE_OPTIONS.some((opt) => opt.value === searchText.trim()); + if (!isExactValue && searchText.trim()) { + filtered.push({ label: `Use custom code: ${searchText.trim()}`, value: searchText.trim() }); + } + return filtered; + }, + }, + { + name: "Key Hash", + label: "Key Hash", + isSearchable: false, + }, + { + name: "Error Message", + label: "Error Message", + isSearchable: false, + }, + ]; +} diff --git a/ui/litellm-dashboard/src/components/view_logs/index.test.tsx b/ui/litellm-dashboard/src/components/view_logs/index.test.tsx index 7a9a541d3e0..aed194a2972 100644 --- a/ui/litellm-dashboard/src/components/view_logs/index.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/index.test.tsx @@ -1,12 +1,8 @@ -import { render, screen, waitFor } from "@testing-library/react"; +import { screen, waitFor } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; -import moment from "moment"; import { beforeEach, describe, expect, it, vi } from "vitest"; -import SpendLogsTable, { RequestViewer } from "./index"; -import type { LogEntry } from "./columns"; -import type { Row } from "@tanstack/react-table"; +import SpendLogsTable from "./index"; import { renderWithProviders } from "../../../tests/test-utils"; -import { uiSpendLogsCall } from "../networking"; const mockHandleFilterResetFromHook = vi.fn(); vi.mock("./log_filter_logic", async (importOriginal) => { @@ -14,14 +10,8 @@ vi.mock("./log_filter_logic", async (importOriginal) => { return { ...actual, useLogFilterLogic: vi.fn(() => ({ - filters: {}, - filteredLogs: { - data: [], - total: 0, - page: 1, - page_size: 50, - total_pages: 1, - }, + logsQuery: { isLoading: false, isFetching: false, refetch: vi.fn() }, + filteredLogs: { data: [], total: 0, page: 1, page_size: 50, total_pages: 1 }, allTeams: [], handleFilterChange: vi.fn(), handleFilterReset: mockHandleFilterResetFromHook, @@ -50,139 +40,6 @@ vi.mock("../key_team_helpers/filter_helpers", () => ({ fetchAllTeams: vi.fn().mockResolvedValue([]), })); -const baseLogEntry: LogEntry = { - request_id: "chatcmpl-test-id", - api_key: "api-key", - team_id: "team-id", - model: "gpt-4", - model_id: "gpt-4", - call_type: "chat", - spend: 0, - total_tokens: 0, - prompt_tokens: 0, - completion_tokens: 0, - startTime: "2025-11-14T00:00:00Z", - endTime: "2025-11-14T00:00:00Z", - cache_hit: "miss", - request_duration_ms: 1000, - messages: [{ role: "user", content: "hello" }], - response: { status: "ok" }, - metadata: { - status: "success", - additional_usage_values: { - cache_read_input_tokens: 0, - cache_creation_input_tokens: 0, - }, - }, - request_tags: {}, - custom_llm_provider: "openai", - api_base: "https://api.example.com", -}; - -const createRow = (overrides: Partial = {}): Row => - ({ - original: { - ...baseLogEntry, - ...overrides, - }, - }) as unknown as Row; - -describe("Request Viewer", () => { - it("renders the request details heading", () => { - render(); - expect(screen.getByText("Request Details")).toBeInTheDocument(); - }); - - it("should truncate the request id if it is longer than 64 characters", () => { - const LONG_REQUEST_ID = "a".repeat(128); - const TRUNCATED_REQUEST_ID = `${"a".repeat(64)}...`; - render( - , - ); - - expect(screen.getByText(TRUNCATED_REQUEST_ID)).toBeInTheDocument(); - }); - - it("should display LiteLLM Overhead when litellm_overhead_time_ms is present in metadata", () => { - render( - , - ); - - expect(screen.getByText("LiteLLM Overhead:")).toBeInTheDocument(); - expect(screen.getByText("150 ms")).toBeInTheDocument(); - }); - - it("should not display LiteLLM Overhead when litellm_overhead_time_ms is not present in metadata", () => { - render(); - - expect(screen.queryByText("LiteLLM Overhead:")).not.toBeInTheDocument(); - }); - - it("should display retry count when attempted_retries > 0 in metadata", () => { - render( - , - ); - - expect(screen.getByText("Retries:")).toBeInTheDocument(); - expect(screen.getByText("2 / 3")).toBeInTheDocument(); - }); - - it("should display green 'None' tag when attempted_retries is 0", () => { - render( - , - ); - - expect(screen.getByText("Retries:")).toBeInTheDocument(); - expect(screen.getByText("None")).toBeInTheDocument(); - }); - - it("should display '-' for Retries when attempted_retries is not present in metadata", () => { - render(); - - expect(screen.getByText("Retries:")).toBeInTheDocument(); - expect(screen.getByText("-")).toBeInTheDocument(); - }); -}); - describe("SpendLogsTable", () => { const defaultProps = { accessToken: "test-token", @@ -215,7 +72,9 @@ describe("SpendLogsTable", () => { renderWithProviders(); // Open the time range quick select dropdown (button shows current range like "Last 24 Hours") - const quickSelectButton = screen.getByRole("button", { name: /Last 24 Hours|Last 15 Minutes|Last Hour|Last 4 Hours|Last 7 Days/i }); + const quickSelectButton = screen.getByRole("button", { + name: /Last 24 Hours|Last 15 Minutes|Last Hour|Last 4 Hours|Last 7 Days/i, + }); await user.click(quickSelectButton); // Click "Custom Range" to enable custom date selection @@ -241,51 +100,19 @@ describe("SpendLogsTable", () => { }); }); - describe("Quick Select time range", () => { - const waitForWindowSeconds = async (minMinutes: number) => { - let diff = -1; - await waitFor(() => { - const lastCall = vi.mocked(uiSpendLogsCall).mock.calls.at(-1)?.[0]; - if (!lastCall) throw new Error("uiSpendLogsCall was not called"); - diff = moment - .utc(lastCall.end_date, "YYYY-MM-DD HH:mm:ss") - .diff(moment.utc(lastCall.start_date, "YYYY-MM-DD HH:mm:ss"), "seconds"); - // start_date is rounded down to the minute boundary; end_date is current time - expect(diff).toBeGreaterThanOrEqual(minMinutes * 60); - expect(diff).toBeLessThan((minMinutes + 1) * 60); - }); - return diff; - }; + describe("auth-not-ready guard", () => { + it("shows a loading spinner when credentials are not yet resolved", () => { + renderWithProviders(); - it("should pass a ~1-minute window to uiSpendLogsCall when 'Last Minute' is selected", async () => { - const user = userEvent.setup(); - renderWithProviders(); - - await user.click(screen.getByRole("button", { name: /Last 24 Hours/i })); - await user.click(await screen.findByRole("button", { name: "Last Minute" })); - - await waitForWindowSeconds(1); + expect(document.querySelector(".ant-spin")).toBeInTheDocument(); + expect(screen.queryByRole("button", { name: "Reset Filters" })).not.toBeInTheDocument(); }); - it("should pass a ~15-minute window to uiSpendLogsCall when 'Last 15 Minutes' is selected", async () => { - const user = userEvent.setup(); + it("renders the table (no spinner) once all credentials are present", () => { renderWithProviders(); - await user.click(screen.getByRole("button", { name: /Last 24 Hours/i })); - await user.click(await screen.findByRole("button", { name: "Last 15 Minutes" })); - - await waitForWindowSeconds(15); - }); - - it("should update the time-range button label to 'Last Minute' after selecting it", async () => { - const user = userEvent.setup(); - renderWithProviders(); - - await user.click(screen.getByRole("button", { name: /Last 24 Hours/i })); - await user.click(await screen.findByRole("button", { name: "Last Minute" })); - - expect(screen.getByRole("button", { name: "Last Minute" })).toBeInTheDocument(); - expect(screen.queryByRole("button", { name: /Last 24 Hours/i })).not.toBeInTheDocument(); + expect(document.querySelector(".ant-spin")).not.toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Reset Filters" })).toBeInTheDocument(); }); }); }); diff --git a/ui/litellm-dashboard/src/components/view_logs/index.tsx b/ui/litellm-dashboard/src/components/view_logs/index.tsx index 2f9e8fe8780..03d917cd923 100644 --- a/ui/litellm-dashboard/src/components/view_logs/index.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/index.tsx @@ -1,35 +1,22 @@ -import { keepPreviousData, useQuery, useQueryClient } from "@tanstack/react-query"; import moment from "moment"; -import { useCallback, useDeferredValue, useEffect, useRef, useState } from "react"; -import GuardrailViewer from "@/components/view_logs/GuardrailViewer/GuardrailViewer"; -import { formatNumberWithCommas } from "@/utils/dataUtils"; -import { truncateString } from "@/utils/textUtils"; -import { SyncOutlined } from "@ant-design/icons"; -import { Row } from "@tanstack/react-table"; -import { Switch, Tab, TabGroup, TabList, TabPanel, TabPanels } from "@tremor/react"; -import { Button, Tag, Tooltip } from "antd"; +import { useCallback, useDeferredValue, useEffect, useMemo, useState } from "react"; +import { Tab, TabGroup, TabList, TabPanel, TabPanels } from "@tremor/react"; import { internalUserRoles } from "../../utils/roles"; import DeletedKeysPage from "../DeletedKeysPage/DeletedKeysPage"; import DeletedTeamsPage from "../DeletedTeamsPage/DeletedTeamsPage"; -import FilterTeamDropdown from "../common_components/FilterTeamDropdown"; import { KeyResponse } from "../key_team_helpers/key_list"; -import { PaginatedKeyAliasSelect } from "../KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect"; -import { PaginatedModelSelect } from "../ModelSelect/PaginatedModelSelect/PaginatedModelSelect"; -import FilterComponent, { FilterOption } from "../molecules/filter"; -import { allEndUsersCall, keyInfoV1Call, uiSpendLogsCall } from "../networking"; +import FilterComponent from "../molecules/filter"; +import { keyInfoV1Call } from "../networking"; import KeyInfoView from "../templates/key_info_view"; import AuditLogs from "./audit_logs"; import { createColumns, LogEntry, type LogsSortField } from "./columns"; -import { ConfigInfoMessage } from "./ConfigInfoMessage"; -import { AGENT_CALL_TYPES, ERROR_CODE_OPTIONS, MCP_CALL_TYPES, QUICK_SELECT_OPTIONS } from "./constants"; -import { CostBreakdownViewer } from "./CostBreakdownViewer"; -import { ErrorViewer } from "./ErrorViewer"; -import { FILTER_KEYS, useLogFilterLogic } from "./log_filter_logic"; +import { AGENT_CALL_TYPES, MCP_CALL_TYPES } from "./constants"; +import { getLogFilterOptions } from "./filter_options"; +import { useLogFilterLogic, defaultFilters, type LogFilterState } from "./log_filter_logic"; import { LogDetailsDrawer } from "./LogDetailsDrawer"; -import { getTimeRangeDisplay } from "./logs_utils"; -import { RequestResponsePanel } from "./RequestResponsePanel"; +import { LogsTableToolbar } from "./LogsTableToolbar"; import { DataTable } from "./table"; -import { VectorStoreViewer } from "./VectorStoreViewer"; +import { AntDLoadingSpinner } from "../ui/AntDLoadingSpinner"; interface SpendLogsTableProps { accessToken: string | null; @@ -39,45 +26,19 @@ interface SpendLogsTableProps { premiumUser: boolean; } -export interface PaginatedResponse { - data: LogEntry[]; - total: number; - page: number; - page_size: number; - total_pages: number; -} - -export default function SpendLogsTable({ - accessToken, - token, - userRole, - userID, - premiumUser, -}: SpendLogsTableProps) { +export default function SpendLogsTable({ accessToken, token, userRole, userID, premiumUser }: SpendLogsTableProps) { const [searchTerm, setSearchTerm] = useState(""); - const [showFilters, setShowFilters] = useState(false); - const [showColumnDropdown, setShowColumnDropdown] = useState(false); const [currentPage, setCurrentPage] = useState(1); const [pageSize] = useState(50); - const dropdownRef = useRef(null); - const filtersRef = useRef(null); - const quickSelectRef = useRef(null); // New state variables for Start and End Time const [startTime, setStartTime] = useState(moment().subtract(24, "hours").format("YYYY-MM-DDTHH:mm")); const [endTime, setEndTime] = useState(moment().format("YYYY-MM-DDTHH:mm")); const [isCustomDate, setIsCustomDate] = useState(false); - const [quickSelectOpen, setQuickSelectOpen] = useState(false); - const [tempTeamId, setTempTeamId] = useState(""); - const [tempKeyHash, setTempKeyHash] = useState(""); - const [selectedTeamId, setSelectedTeamId] = useState(""); - const [selectedKeyHash, setSelectedKeyHash] = useState(""); - const [selectedModelId, setSelectedModelId] = useState(""); + const [filters, setFilters] = useState(defaultFilters); const [selectedKeyInfo, setSelectedKeyInfo] = useState(null); const [selectedKeyIdInfoView, setSelectedKeyIdInfoView] = useState(null); - const [selectedStatus, setSelectedStatus] = useState(""); - const [selectedEndUser, setSelectedEndUser] = useState(""); const [filterByCurrentUser, setFilterByCurrentUser] = useState(userRole && internalUserRoles.includes(userRole)); const [activeTab, setActiveTab] = useState("request logs"); @@ -88,12 +49,10 @@ export default function SpendLogsTable({ const [sortBy, setSortBy] = useState("startTime"); const [sortOrder, setSortOrder] = useState<"asc" | "desc">("desc"); - // Tracks whether any filter that uses performSearch (backend) is active. - // Used to disable the main query so it doesn't fire redundant unfiltered requests - // when time range / sort / page changes while a backend filter is in effect. - const [isMainQueryEnabled, setIsMainQueryEnabled] = useState(true); - - const queryClient = useQueryClient(); + const [selectedTimeInterval, setSelectedTimeInterval] = useState<{ value: number; unit: string }>({ + value: 24, + unit: "hours", + }); const [isLiveTail, setIsLiveTail] = useState(() => { const storedValue = sessionStorage.getItem("isLiveTail"); @@ -105,11 +64,6 @@ export default function SpendLogsTable({ sessionStorage.setItem("isLiveTail", JSON.stringify(isLiveTail)); }, [isLiveTail]); - const [selectedTimeInterval, setSelectedTimeInterval] = useState<{ value: number; unit: string }>({ - value: 24, - unit: "hours", - }); - useEffect(() => { const fetchKeyInfo = async () => { if (selectedKeyIdInfoView && accessToken) { @@ -126,132 +80,33 @@ export default function SpendLogsTable({ fetchKeyInfo(); }, [selectedKeyIdInfoView, accessToken]); - // Close dropdown when clicking outside - useEffect(() => { - function handleClickOutside(event: MouseEvent) { - if (dropdownRef.current && !dropdownRef.current.contains(event.target as Node)) { - setShowColumnDropdown(false); - } - if (filtersRef.current && !filtersRef.current.contains(event.target as Node)) { - setShowFilters(false); - } - if (quickSelectRef.current && !quickSelectRef.current.contains(event.target as Node)) { - setQuickSelectOpen(false); - } - } - - document.addEventListener("mousedown", handleClickOutside); - return () => document.removeEventListener("mousedown", handleClickOutside); - }, []); - useEffect(() => { if (userRole && internalUserRoles.includes(userRole)) { setFilterByCurrentUser(true); } }, [userRole]); - const LiveTailControls = () => { - return ( -
- Live Tail - -
- ); - }; - - const logs = useQuery({ - queryKey: [ - "logs", - "table", - currentPage, - pageSize, - startTime, - endTime, - selectedTeamId, - selectedKeyHash, - filterByCurrentUser ? userID : null, - selectedStatus, - selectedModelId, - sortBy, - sortOrder, - ], - queryFn: async () => { - if (!accessToken || !token || !userRole || !userID) { - return { - data: [], - total: 0, - page: 1, - page_size: pageSize, - total_pages: 0, - }; - } - - const formattedStartTime = moment(startTime).utc().format("YYYY-MM-DD HH:mm:ss"); - const formattedEndTime = isCustomDate - ? moment(endTime).utc().format("YYYY-MM-DD HH:mm:ss") - : moment().utc().format("YYYY-MM-DD HH:mm:ss"); - - // Get base response from API - // NOTE: We only fetch the list of logs here (lightweight). - // Log details (messages/response) are fetched on-demand when user clicks a row. - const response = await uiSpendLogsCall({ - accessToken, - start_date: formattedStartTime, - end_date: formattedEndTime, - page: currentPage, - page_size: pageSize, - params: { - api_key: selectedKeyHash || undefined, - team_id: selectedTeamId || undefined, - user_id: filterByCurrentUser ? userID ?? undefined : undefined, - end_user: selectedEndUser || undefined, - status_filter: selectedStatus || undefined, - model_id: selectedModelId || undefined, - sort_by: sortBy, - sort_order: sortOrder, - }, - }); - - return response; - }, - enabled: !!accessToken && !!token && !!userRole && !!userID && activeTab === "request logs" && isMainQueryEnabled, - refetchInterval: isLiveTail && currentPage === 1 ? 15000 : false, - placeholderData: keepPreviousData, - refetchIntervalInBackground: true, - }); - - // Defer the transition from "Fetching" to "Fetch" so the button stays loading until - // the table has rendered with the new data (avoids the visual gap where the button - // exits loading state before the table updates) - const isFetchingDeferred = useDeferredValue(logs.isFetching); - const isButtonLoading = logs.isFetching || isFetchingDeferred; - - const logsData = logs.data || { - data: [], - total: 0, - page: 1, - page_size: pageSize || 10, - total_pages: 1, - }; - const { - filters, + logsQuery, filteredLogs, - hasBackendFilters, allTeams, handleFilterChange, handleFilterReset: handleFilterResetFromHook, - refetchWithFilters, } = useLogFilterLogic({ - logs: logsData, accessToken, + token, + userRole, + userID, + filters, + setFilters, + filterByCurrentUser: !!filterByCurrentUser, + activeTab, + isLiveTail, startTime, endTime, pageSize, isCustomDate, setCurrentPage, - userID, - userRole, sortBy, sortOrder, currentPage, @@ -259,7 +114,6 @@ export default function SpendLogsTable({ const handleFilterReset = useCallback(() => { handleFilterResetFromHook(); - // Reset custom time range to default (last 24 hours) setStartTime(moment().subtract(24, "hours").format("YYYY-MM-DDTHH:mm")); setEndTime(moment().format("YYYY-MM-DDTHH:mm")); setIsCustomDate(false); @@ -267,111 +121,102 @@ export default function SpendLogsTable({ setCurrentPage(1); }, [handleFilterResetFromHook]); - // Disable the main query whenever backend filters are active so it doesn't fire - // redundant unfiltered requests when time range / sort / page changes. - useEffect(() => { - setIsMainQueryEnabled(!hasBackendFilters); - }, [hasBackendFilters]); + const handleSortChange = useCallback((newSortBy: LogsSortField, newSortOrder: "asc" | "desc") => { + setSortBy(newSortBy); + setSortOrder(newSortOrder); + setCurrentPage(1); + }, []); - // Sync filter state into the individual selectedX state variables used by the main query - useEffect(() => { - if (!accessToken) return; + const columns = useMemo( + () => createColumns({ sortBy, sortOrder, onSortChange: handleSortChange }), + [sortBy, sortOrder, handleSortChange], + ); - if (filters["Team ID"]) { - setSelectedTeamId(filters["Team ID"]); - } else { - setSelectedTeamId(""); + const filteredData = useMemo(() => { + const searchedLogs = filteredLogs.data.filter((log) => { + const matchesSearch = + !searchTerm || + log.request_id.includes(searchTerm) || + log.model.includes(searchTerm) || + (log.user && log.user.includes(searchTerm)); + + // No need for additional filtering since we're now handling this in the API call + return matchesSearch; + }); + + const sessionCompositionById = searchedLogs.reduce>( + (acc, log) => { + if (!log.session_id) return acc; + if (!acc[log.session_id]) { + acc[log.session_id] = { llm: 0, agent: 0, mcp: 0 }; + } + if (MCP_CALL_TYPES.includes(log.call_type)) { + acc[log.session_id].mcp += 1; + } else if (AGENT_CALL_TYPES.includes(log.call_type)) { + acc[log.session_id].agent += 1; + } else { + acc[log.session_id].llm += 1; + } + return acc; + }, + {}, + ); + + // Build a single-pass map of session_id → representative request_id. + // Prefers an LLM row over an MCP row as the representative. + const sessionRepresentativeMap = new Map(); + for (const log of searchedLogs) { + if (!log.session_id || (log.session_total_count || 1) <= 1) continue; + const isMcp = MCP_CALL_TYPES.includes(log.call_type); + const existing = sessionRepresentativeMap.get(log.session_id); + if (!existing || (existing.isMcp && !isMcp)) { + sessionRepresentativeMap.set(log.session_id, { requestId: log.request_id, isMcp }); + } } - setSelectedStatus(filters["Status"] || ""); - setSelectedModelId(filters["Model"] || ""); - setSelectedEndUser(filters["End User"] || ""); - // Key Alias filtering is handled server-side by performSearch via the key_alias param. - // We intentionally do not translate the alias to a hash here to avoid firing a - // redundant main-query request (api_key=hash) alongside performSearch's key_alias request. - setSelectedKeyHash(filters["Key Hash"] || ""); - }, [filters, accessToken]); + return ( + searchedLogs + .map((log) => { + const sessionComposition = log.session_id ? sessionCompositionById[log.session_id] : undefined; + return { + ...log, + request_duration_ms: log.request_duration_ms, + session_llm_count: sessionComposition?.llm ?? undefined, + session_mcp_count: sessionComposition?.mcp ?? undefined, + session_agent_count: sessionComposition?.agent ?? undefined, + onKeyHashClick: (keyHash: string) => setSelectedKeyIdInfoView(keyHash), + onSessionClick: (sessionId: string) => { + if (sessionId) { + setSelectedSessionId(sessionId); + setSelectedLog(log); + setIsDrawerOpen(true); + } + }, + }; + }) + // Deduplicate multi-call sessions using the pre-built map (O(1) per row). + .filter((log) => { + if (!log.session_id || (log.session_total_count || 1) <= 1) return true; + return sessionRepresentativeMap.get(log.session_id)?.requestId === log.request_id; + }) + ); + }, [filteredLogs.data, searchTerm]); + + // Keep the Fetch button busy until the table has actually committed the new + // rows. `keepPreviousData` leaves logsQuery.isLoading false on refetch, so + // without this the button clears while stale rows are still on screen. + const deferredData = useDeferredValue(filteredData); + const isStale = deferredData !== filteredData; + const isButtonLoading = logsQuery.isFetching || isStale; if (!accessToken || !token || !userRole || !userID) { - return null; + return ( +
+ +
+ ); } - const searchedLogs = filteredLogs.data.filter((log) => { - const matchesSearch = - !searchTerm || - log.request_id.includes(searchTerm) || - log.model.includes(searchTerm) || - (log.user && log.user.includes(searchTerm)); - - // No need for additional filtering since we're now handling this in the API call - return matchesSearch; - }); - - const sessionCompositionById = searchedLogs.reduce>((acc, log) => { - if (!log.session_id) return acc; - if (!acc[log.session_id]) { - acc[log.session_id] = { llm: 0, agent: 0, mcp: 0 }; - } - if (MCP_CALL_TYPES.includes(log.call_type)) { - acc[log.session_id].mcp += 1; - } else if (AGENT_CALL_TYPES.includes(log.call_type)) { - acc[log.session_id].agent += 1; - } else { - acc[log.session_id].llm += 1; - } - return acc; - }, {}); - - // Build a single-pass map of session_id → representative request_id. - // Prefers an LLM row over an MCP row as the representative. - const sessionRepresentativeMap = new Map(); - for (const log of searchedLogs) { - if (!log.session_id || (log.session_total_count || 1) <= 1) continue; - const isMcp = MCP_CALL_TYPES.includes(log.call_type); - const existing = sessionRepresentativeMap.get(log.session_id); - if (!existing || (existing.isMcp && !isMcp)) { - sessionRepresentativeMap.set(log.session_id, { requestId: log.request_id, isMcp }); - } - } - - const filteredData = - searchedLogs - .map((log) => { - const sessionComposition = log.session_id ? sessionCompositionById[log.session_id] : undefined; - return { - ...log, - request_duration_ms: log.request_duration_ms, - session_llm_count: sessionComposition?.llm ?? undefined, - session_mcp_count: sessionComposition?.mcp ?? undefined, - session_agent_count: sessionComposition?.agent ?? undefined, - onKeyHashClick: (keyHash: string) => setSelectedKeyIdInfoView(keyHash), - onSessionClick: (sessionId: string) => { - if (sessionId) { - setSelectedSessionId(sessionId); - setSelectedLog(log); - setIsDrawerOpen(true); - } - }, - }; - }) - // Deduplicate multi-call sessions using the pre-built map (O(1) per row). - .filter((log) => { - if (!log.session_id || (log.session_total_count || 1) <= 1) return true; - return sessionRepresentativeMap.get(log.session_id)?.requestId === log.request_id; - }) || []; - - // Add this function to handle manual refresh - const handleRefresh = () => { - if (hasBackendFilters) { - // When backend filters (e.g. Key Alias) are active the main TanStack Query - // is disabled and its params do not include filter values like key_alias. - // Route through the filter-aware refetch so all active filters are preserved. - refetchWithFilters(); - } else { - logs.refetch(); - } - }; - const handleRowClick = (log: LogEntry) => { // Multi-call session row: open in the same right-side drawer (session mode) if (log.session_id && (log.session_total_count || 1) > 1) { @@ -386,100 +231,6 @@ export default function SpendLogsTable({ setIsDrawerOpen(true); }; - const handleCloseDrawer = () => { - setIsDrawerOpen(false); - setSelectedSessionId(null); - }; - - const handleSelectLog = (log: LogEntry) => { - setSelectedLog(log); - }; - - const logFilterOptions: FilterOption[] = [ - { - name: "Team ID", - label: "Team ID", - customComponent: FilterTeamDropdown, - }, - { - name: "Status", - label: "Status", - isSearchable: false, - options: [ - { label: "Success", value: "success" }, - { label: "Failure", value: "failure" }, - ], - }, - { - name: "Model", - label: "Model", - customComponent: PaginatedModelSelect, - }, - { - name: FILTER_KEYS.PUBLIC_MODEL_OR_SEARCH_TOOL, - label: "Public model / search tool", - isSearchable: false, - }, - { - name: "Key Alias", - label: "Key Alias", - customComponent: PaginatedKeyAliasSelect, - }, - { - name: "End User", - label: "End User", - isSearchable: true, - searchFn: async (searchText: string) => { - if (!accessToken) return []; - const data = await allEndUsersCall(accessToken); - // data if set, is a list of objects, with key = user_id - const users = data?.map((u: any) => u.user_id) || []; - const filtered = users.filter((u: string) => u.toLowerCase().includes(searchText.toLowerCase())); - return filtered.map((u: string) => ({ label: u, value: u })); - }, - }, - { - name: "Error Code", - label: "Error Code", - isSearchable: true, - searchFn: async (searchText: string) => { - if (!searchText) return ERROR_CODE_OPTIONS; - const lower = searchText.toLowerCase(); - const filtered = ERROR_CODE_OPTIONS.filter((opt) => opt.label.toLowerCase().includes(lower)); - const isExactValue = ERROR_CODE_OPTIONS.some((opt) => opt.value === searchText.trim()); - if (!isExactValue && searchText.trim()) { - filtered.push({ label: `Use custom code: ${searchText.trim()}`, value: searchText.trim() }); - } - return filtered; - }, - }, - { - name: "Key Hash", - label: "Key Hash", - isSearchable: false, - }, - { - name: "Error Message", - label: "Error Message", - isSearchable: false, - }, - ]; - - const formatTimeUnit = (value: number, unit: string) => { - if (value === 1) { - if (unit === "minutes") return "minute"; - if (unit === "hours") return "hour"; - if (unit === "days") return "day"; - } - return unit; - }; - - const selectedOption = QUICK_SELECT_OPTIONS.find( - (option) => option.value === selectedTimeInterval.value && option.unit === selectedTimeInterval.unit, - ); - - const displayLabel = isCustomDate ? getTimeRangeDisplay(isCustomDate, startTime, endTime) : selectedOption?.label; - return (
setActiveTab(index === 0 ? "request logs" : "audit logs")}> @@ -505,192 +256,37 @@ export default function SpendLogsTable({ ) : ( <>
-
-
-
-
- setSearchTerm(e.target.value)} - /> - - - -
- -
-
- - - {quickSelectOpen && ( -
-
- {QUICK_SELECT_OPTIONS.map((option) => ( - - ))} -
- -
-
- )} -
- - - - -
- - {isCustomDate && ( -
-
- { - setStartTime(e.target.value); - setCurrentPage(1); - }} - className="px-3 py-2 border rounded-md text-sm focus:outline-none focus:ring-2 focus:ring-blue-500 focus:border-blue-500" - /> -
- to -
- { - setEndTime(e.target.value); - setCurrentPage(1); - }} - className="px-3 py-2 border rounded-md text-sm focus:outline-none focus:ring-2 focus:ring-blue-500 focus:border-blue-500" - /> -
-
- )} -
- -
- - Showing {logs.isLoading ? "..." : filteredLogs ? (currentPage - 1) * pageSize + 1 : 0} -{" "} - {logs.isLoading - ? "..." - : filteredLogs - ? Math.min(currentPage * pageSize, filteredLogs.total) - : 0}{" "} - of {logs.isLoading ? "..." : filteredLogs ? filteredLogs.total : 0} results - -
- - Page {logs.isLoading ? "..." : currentPage} of{" "} - {logs.isLoading ? "..." : filteredLogs ? filteredLogs.total_pages : 1} - - - -
-
-
-
- {isLiveTail && currentPage === 1 && isMainQueryEnabled && ( -
-
- Auto-refreshing every 15 seconds -
- -
- )} + logsQuery.refetch()} + filteredLogs={filteredLogs} + /> { - setSortBy(newSortBy); - setSortOrder(newSortOrder); - setCurrentPage(1); - }, - })} - data={filteredData} + columns={columns} + data={deferredData} onRowClick={handleRowClick} - isLoading={logs.isLoading} + isLoading={logsQuery.isLoading} />
@@ -706,330 +302,29 @@ export default function SpendLogsTable({ premiumUser={premiumUser} /> - - + + + + + + {/* Log Details Drawer */} { + setIsDrawerOpen(false); + setSelectedSessionId(null); + }} logEntry={selectedLog} sessionId={selectedSessionId} accessToken={accessToken} allLogs={filteredData} - onSelectLog={handleSelectLog} + onSelectLog={setSelectedLog} startTime={moment(startTime).utc().format("YYYY-MM-DD HH:mm:ss")} />
); } - -export function RequestViewer({ row }: { row: Row }) { - // Helper function to clean metadata by removing specific fields - const formatData = (input: any) => { - if (typeof input === "string") { - try { - return JSON.parse(input); - } catch { - return input; - } - } - return input; - }; - - // New helper function to get raw request - const getRawRequest = () => { - // First check if proxy_server_request exists in metadata - if (row.original?.proxy_server_request) { - return formatData(row.original.proxy_server_request); - } - // Fall back to messages if proxy_server_request is empty - return formatData(row.original.messages); - }; - - // Extract error information from metadata if available - const metadata = row.original.metadata || {}; - const hasError = metadata.status === "failure"; - const errorInfo = hasError ? metadata.error_information : null; - - // Check if request/response data is missing - const hasMessages = - row.original.messages && - (Array.isArray(row.original.messages) - ? row.original.messages.length > 0 - : Object.keys(row.original.messages).length > 0); - const hasResponse = row.original.response && Object.keys(formatData(row.original.response)).length > 0; - const missingData = !hasMessages && !hasResponse && !hasError; - - // Format the response with error details if present - const formattedResponse = () => { - if (hasError && errorInfo) { - return { - error: { - message: errorInfo.error_message || "An error occurred", - type: errorInfo.error_class || "error", - code: errorInfo.error_code || "unknown", - param: null, - }, - }; - } - return formatData(row.original.response); - }; - - // Extract vector store request metadata if available - const hasVectorStoreData = - metadata.vector_store_request_metadata && - Array.isArray(metadata.vector_store_request_metadata) && - metadata.vector_store_request_metadata.length > 0; - - // Extract guardrail information from metadata if available - const guardrailInfo = row.original.metadata?.guardrail_information; - const guardrailEntries = Array.isArray(guardrailInfo) ? guardrailInfo : guardrailInfo ? [guardrailInfo] : []; - const hasGuardrailData = guardrailEntries.length > 0; - - // Calculate total masked entities if guardrail data exists - const totalMaskedEntities = guardrailEntries.reduce((sum, entry) => { - const maskedCounts = entry?.masked_entity_count; - if (!maskedCounts) { - return sum; - } - return ( - sum + - Object.values(maskedCounts).reduce((acc, count) => (typeof count === "number" ? acc + count : acc), 0) - ); - }, 0); - - const primaryGuardrailLabel = - guardrailEntries.length === 1 - ? guardrailEntries[0]?.guardrail_name ?? "-" - : guardrailEntries.length > 1 - ? `${guardrailEntries.length} guardrails` - : "-"; - - const truncatedRequestId = truncateString(row.original.request_id, 64); - - return ( -
- {/* Combined Info Card */} -
-
-

Request Details

-
-
-
-
- Request ID: - {row.original.request_id.length > 64 ? ( - - {truncatedRequestId} - - ) : ( - {row.original.request_id} - )} -
-
- Model: - {row.original.model} -
-
- Model ID: - {row.original.model_id} -
-
- Call Type: - {row.original.call_type} -
-
- Provider: - {row.original.custom_llm_provider || "-"} -
-
- API Base: - - {row.original.api_base || "-"} - -
- {row?.original?.requester_ip_address && ( -
- IP Address: - {row?.original?.requester_ip_address} -
- )} - {hasGuardrailData && ( -
- Guardrail: -
- {primaryGuardrailLabel} - {totalMaskedEntities > 0 && ( - - {totalMaskedEntities} masked - - )} -
-
- )} -
-
-
- Tokens: - - {row.original.total_tokens} ({row.original.prompt_tokens} prompt tokens +{" "} - {row.original.completion_tokens} completion tokens) - -
-
- Cache Read Tokens: - - {formatNumberWithCommas(row.original.metadata?.additional_usage_values?.cache_read_input_tokens || 0)} - -
-
- Cache Creation Tokens: - - {formatNumberWithCommas(row.original.metadata?.additional_usage_values.cache_creation_input_tokens)} - -
-
- Cost: - ${formatNumberWithCommas(row.original.spend || 0, 6)} -
-
- Cache Hit: - {row.original.cache_hit} -
- -
- Status: - - {(row.original.metadata?.status || "Success").toLowerCase() !== "failure" ? "Success" : "Failure"} - -
-
- Start Time: - {row.original.startTime} -
-
- End Time: - {row.original.endTime} -
-
- Duration: - {row.original.request_duration_ms != null ? (row.original.request_duration_ms / 1000).toFixed(3) : "-"} s. -
- {row.original.metadata?.litellm_overhead_time_ms !== undefined && ( -
- LiteLLM Overhead: - {row.original.metadata.litellm_overhead_time_ms} ms -
- )} -
- Retries: - - {row.original.metadata?.attempted_retries !== undefined && row.original.metadata?.attempted_retries !== null - ? row.original.metadata.attempted_retries > 0 - ? `${row.original.metadata.attempted_retries}${row.original.metadata.max_retries !== undefined && row.original.metadata.max_retries !== null ? ` / ${row.original.metadata.max_retries}` : ''}` - : None - : '-'} - -
-
-
-
- - {/* Cost Breakdown - Show if cost breakdown data is available */} - - - {/* Configuration Info Message - Show when data is missing */} - - - {/* Request/Response Panel */} -
- -
- - {/* Guardrail Data - Show only if present */} - {hasGuardrailData && } - - {/* Vector Store Request Data - Show only if present */} - {hasVectorStoreData && } - - {/* Error Card - Only show for failures */} - {hasError && errorInfo && } - - {/* Tags Card - Only show if there are tags */} - {row.original.request_tags && Object.keys(row.original.request_tags).length > 0 && ( -
-
-

Request Tags

-
-
-
- {Object.entries(row.original.request_tags).map(([key, value]) => ( - - {key}: {String(value)} - - ))} -
-
-
- )} - - {/* Metadata Card - Only show if there's metadata */} - {row.original.metadata && Object.keys(row.original.metadata).length > 0 && ( -
-
-

Metadata

- -
-
-
-              {JSON.stringify(row.original.metadata, null, 2)}
-            
-
-
- )} -
- ); -} diff --git a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx index 17c50771526..cbe37e0b70f 100644 --- a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx @@ -1,10 +1,16 @@ import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { act, renderHook, waitFor } from "@testing-library/react"; -import React, { ReactNode } from "react"; +import React, { ReactNode, useState } from "react"; import { beforeEach, describe, expect, it, vi } from "vitest"; -import type { PaginatedResponse } from "."; -import type { LogEntry, LogsSortField } from "./columns"; -import { useLogFilterLogic } from "./log_filter_logic"; +import type { LogsSortField } from "./columns"; +import { + defaultFilters, + getLiveTailRefetchInterval, + LIVE_TAIL_INTERVAL_MS, + useLogFilterLogic, + type LogFilterState, + type PaginatedResponse, +} from "./log_filter_logic"; vi.mock("../networking", () => ({ uiSpendLogsCall: vi.fn(), @@ -16,767 +22,642 @@ vi.mock("@/components/key_team_helpers/filter_helpers", () => ({ import { uiSpendLogsCall } from "../networking"; -const createLogEntry = (overrides: Partial = {}): LogEntry => -({ - request_id: "req-1", - api_key: "key-1", - team_id: "team-1", - model: "gpt-4", - model_id: "gpt-4", - call_type: "chat", - spend: 0, - total_tokens: 0, - prompt_tokens: 0, - completion_tokens: 0, - startTime: "2025-01-01T00:00:00Z", - endTime: "2025-01-01T00:01:00Z", - cache_hit: "miss", - messages: [], - response: {}, - metadata: {}, - request_tags: {}, - ...overrides, -} as LogEntry); - -const createPaginatedResponse = (data: LogEntry[]): PaginatedResponse => ({ - data, - total: data.length, +const emptyResponse: PaginatedResponse = { + data: [], + total: 0, page: 1, page_size: 50, - total_pages: 1, -}); + total_pages: 0, +}; const defaultProps = { - logs: createPaginatedResponse([]), - accessToken: "test-token", + accessToken: "test-token" as string | null, + token: "test-token" as string | null, + userRole: "Admin" as string | null, + userID: "user-1" as string | null, + filterByCurrentUser: false, + activeTab: "request logs", + isLiveTail: false, startTime: "2025-01-01T00:00:00", endTime: "2025-01-01T23:59:59", isCustomDate: true, - setCurrentPage: vi.fn(), - userID: "user-1", - userRole: "Admin", + sortBy: "startTime" as LogsSortField, + sortOrder: "desc" as "asc" | "desc", + currentPage: 1, }; +type HookOverrides = Partial[0], "filters" | "setFilters">>; + describe("useLogFilterLogic", () => { let queryClient: QueryClient; beforeEach(() => { queryClient = new QueryClient({ - defaultOptions: { - queries: { - retry: false, - }, - }, + defaultOptions: { queries: { retry: false } }, }); vi.clearAllMocks(); - vi.mocked(uiSpendLogsCall).mockResolvedValue({ - data: [], - total: 0, - page: 1, - page_size: 50, - total_pages: 0, - }); + vi.mocked(uiSpendLogsCall).mockResolvedValue(emptyResponse); }); const wrapper = ({ children }: { children: ReactNode }) => React.createElement(QueryClientProvider, { client: queryClient }, children); - it("should return filters, filteredLogs, allTeams, handleFilterChange, and handleFilterReset", () => { - const { result } = renderHook( - () => - useLogFilterLogic({ + function renderFilterHook(overrides: HookOverrides = {}) { + const setCurrentPage = overrides.setCurrentPage ?? vi.fn(); + const rendered = renderHook( + () => { + const [filters, setFilters] = useState(defaultFilters); + const hook = useLogFilterLogic({ ...defaultProps, - logs: createPaginatedResponse([createLogEntry()]), - }), + ...overrides, + filters, + setFilters, + setCurrentPage, + }); + return { ...hook, filters, setFilters }; + }, { wrapper }, ); + return { ...rendered, setCurrentPage }; + } - expect(result.current.filters).toBeDefined(); - expect(result.current.filteredLogs).toBeDefined(); - expect(result.current).toHaveProperty("allTeams"); - expect(result.current.handleFilterChange).toBeDefined(); - expect(result.current.handleFilterReset).toBeDefined(); - }); + describe("return shape", () => { + it("exposes filteredLogs, allTeams, handleFilterChange, handleFilterReset", () => { + const { result } = renderFilterHook(); - it("should initialize filters with all keys empty", () => { - const { result } = renderHook(() => useLogFilterLogic(defaultProps), { wrapper }); - - const filters = result.current.filters; - expect(filters["Team ID"]).toBe(""); - expect(filters["Key Hash"]).toBe(""); - expect(filters["Request ID"]).toBe(""); - expect(filters["Model"]).toBe(""); - expect(filters["User ID"]).toBe(""); - expect(filters["End User"]).toBe(""); - expect(filters["Status"]).toBe(""); - expect(filters["Key Alias"]).toBe(""); - expect(filters["Error Code"]).toBe(""); - expect(filters["Error Message"]).toBe(""); - expect(filters["Public model / search tool"]).toBe(""); - }); - - it("should return all logs when no filters are applied", () => { - const logs = createPaginatedResponse([ - createLogEntry({ request_id: "req-1" }), - createLogEntry({ request_id: "req-2" }), - ]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - expect(result.current.filteredLogs.data).toHaveLength(2); - expect(result.current.filteredLogs.data).toEqual(logs.data); - }); - - it("should filter logs by team_id when Team ID filter is set", () => { - const logs = createPaginatedResponse([ - createLogEntry({ request_id: "req-1", team_id: "team-a" }), - createLogEntry({ request_id: "req-2", team_id: "team-b" }), - createLogEntry({ request_id: "req-3", team_id: "team-a" }), - ]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - act(() => { - result.current.handleFilterChange({ "Team ID": "team-a" }); - }); - - expect(result.current.filteredLogs.data).toHaveLength(2); - expect(result.current.filteredLogs.data.every((log) => log.team_id === "team-a")).toBe(true); - }); - - it("should filter logs by status when Status filter is set to success", () => { - const logs = createPaginatedResponse([ - createLogEntry({ request_id: "req-1", status: "success" }), - createLogEntry({ request_id: "req-2" }), - createLogEntry({ request_id: "req-3", status: "error" }), - ]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - act(() => { - result.current.handleFilterChange({ Status: "success" }); - }); - - expect(result.current.filteredLogs.data).toHaveLength(2); - expect(result.current.filteredLogs.data.every((log) => !log.status || log.status === "success")).toBe(true); - }); - - it("should filter logs by status when Status filter is set to error", () => { - const logs = createPaginatedResponse([ - createLogEntry({ request_id: "req-1", status: "success" }), - createLogEntry({ request_id: "req-2", status: "error" }), - ]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - act(() => { - result.current.handleFilterChange({ Status: "error" }); - }); - - expect(result.current.filteredLogs.data).toHaveLength(1); - expect(result.current.filteredLogs.data[0].status).toBe("error"); - }); - - it("should filter logs by model_id when Model filter is set", async () => { - const filteredLogs = [ - createLogEntry({ request_id: "req-1", model_id: "gpt-4" }), - createLogEntry({ request_id: "req-3", model_id: "gpt-4" }), - ]; - vi.mocked(uiSpendLogsCall).mockResolvedValue( - createPaginatedResponse(filteredLogs), - ); - const logs = createPaginatedResponse([ - createLogEntry({ request_id: "req-1", model_id: "gpt-4" }), - createLogEntry({ request_id: "req-2", model_id: "gpt-3.5" }), - createLogEntry({ request_id: "req-3", model_id: "gpt-4" }), - ]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - act(() => { - result.current.handleFilterChange({ Model: "gpt-4" }); - }); - - await waitFor( - () => { - expect(result.current.filteredLogs.data).toHaveLength(2); - expect(result.current.filteredLogs.data.every((log) => log.model_id === "gpt-4")).toBe(true); - }, - { timeout: 500 }, - ); - }); - - it("should pass model param and filter search-tool rows by spend log model column", async () => { - const searchRows = [ - createLogEntry({ - request_id: "s1", - call_type: "asearch", - model: "tavily-marketing", - model_id: "", - team_id: "team-x", - }), - ]; - vi.mocked(uiSpendLogsCall).mockResolvedValue(createPaginatedResponse(searchRows)); - const logs = createPaginatedResponse([ - ...searchRows, - createLogEntry({ - request_id: "c1", - call_type: "chat", - model: "gpt-4o", - model_id: "mid-1", - team_id: "team-x", - }), - ]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - act(() => { - result.current.handleFilterChange({ "Public model / search tool": "tavily-marketing" }); - }); - - await waitFor( - () => { - expect(result.current.filteredLogs.data).toHaveLength(1); - expect(result.current.filteredLogs.data[0].model).toBe("tavily-marketing"); - expect(result.current.filteredLogs.data[0].call_type).toBe("asearch"); - }, - { timeout: 500 }, - ); - - expect(vi.mocked(uiSpendLogsCall)).toHaveBeenCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ - model: "tavily-marketing", - }), - }), - ); - }); - - it("should filter logs by api_key when Key Hash filter is set", async () => { - const filteredLog = createLogEntry({ request_id: "req-1", api_key: "key-x" }); - vi.mocked(uiSpendLogsCall).mockResolvedValue( - createPaginatedResponse([filteredLog]), - ); - const logs = createPaginatedResponse([ - createLogEntry({ request_id: "req-1", api_key: "key-x" }), - createLogEntry({ request_id: "req-2", api_key: "key-y" }), - ]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - act(() => { - result.current.handleFilterChange({ "Key Hash": "key-x" }); - }); - - await waitFor( - () => { - expect(result.current.filteredLogs.data).toHaveLength(1); - expect(result.current.filteredLogs.data[0].api_key).toBe("key-x"); - }, - { timeout: 500 }, - ); - }); - - it("should filter logs by end_user when End User filter is set", async () => { - const filteredLog = createLogEntry({ request_id: "req-1", end_user: "user-a" }); - vi.mocked(uiSpendLogsCall).mockResolvedValue( - createPaginatedResponse([filteredLog]), - ); - const logs = createPaginatedResponse([ - createLogEntry({ request_id: "req-1", end_user: "user-a" }), - createLogEntry({ request_id: "req-2", end_user: "user-b" }), - ]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - act(() => { - result.current.handleFilterChange({ "End User": "user-a" }); - }); - - await waitFor( - () => { - expect(result.current.filteredLogs.data).toHaveLength(1); - expect(result.current.filteredLogs.data[0].end_user).toBe("user-a"); - }, - { timeout: 500 }, - ); - }); - - it("should filter logs by error_code when Error Code filter is set", async () => { - const filteredLog = createLogEntry({ - request_id: "req-1", - metadata: { error_information: { error_code: "429" } }, - }); - vi.mocked(uiSpendLogsCall).mockResolvedValue( - createPaginatedResponse([filteredLog]), - ); - const logs = createPaginatedResponse([ - createLogEntry({ - request_id: "req-1", - metadata: { error_information: { error_code: "429" } }, - }), - createLogEntry({ - request_id: "req-2", - metadata: { error_information: { error_code: "500" } }, - }), - ]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - act(() => { - result.current.handleFilterChange({ "Error Code": "429" }); - }); - - await waitFor( - () => { - expect(result.current.filteredLogs.data).toHaveLength(1); - expect(result.current.filteredLogs.data[0].metadata?.error_information?.error_code).toBe("429"); - }, - { timeout: 500 }, - ); - }); - - it("should return empty data when logs is null or has no data", () => { - const { result } = renderHook( - () => - useLogFilterLogic({ - ...defaultProps, - logs: { data: [], total: 0, page: 1, page_size: 50, total_pages: 0 }, - }), - { wrapper }, - ); - - expect(result.current.filteredLogs.data).toEqual([]); - expect(result.current.filteredLogs.total).toBe(0); - }); - - it("should reset filters when handleFilterReset is called", () => { - const logs = createPaginatedResponse([createLogEntry()]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - act(() => { - result.current.handleFilterChange({ "Team ID": "team-1", Status: "success" }); - }); - - expect(result.current.filters["Team ID"]).toBe("team-1"); - expect(result.current.filters["Status"]).toBe("success"); - - act(() => { - result.current.handleFilterReset(); - }); - - expect(result.current.filters["Team ID"]).toBe(""); - expect(result.current.filters["Status"]).toBe(""); - }); - - it("should call setCurrentPage with 1 when handleFilterChange is invoked", () => { - const setCurrentPage = vi.fn(); - const logs = createPaginatedResponse([createLogEntry()]); - const { result } = renderHook( - () => useLogFilterLogic({ ...defaultProps, logs, setCurrentPage }), - { wrapper }, - ); - - act(() => { - result.current.handleFilterChange({ "Team ID": "team-1" }); - }); - - expect(setCurrentPage).toHaveBeenCalledWith(1); - }); - - it("should call uiSpendLogsCall when backend filter is set and debounce elapses", async () => { - const logs = createPaginatedResponse([createLogEntry()]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await waitFor( - () => { - expect(uiSpendLogsCall).toHaveBeenCalled(); - }, - { timeout: 500 }, - ); - }); - - it("should not call uiSpendLogsCall when accessToken is null", async () => { - const logs = createPaginatedResponse([createLogEntry()]); - const { result } = renderHook( - () => useLogFilterLogic({ ...defaultProps, logs, accessToken: null }), - { wrapper }, - ); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await new Promise((resolve) => setTimeout(resolve, 350)); - - expect(uiSpendLogsCall).not.toHaveBeenCalled(); - }); - - it("should use backend filtered logs when backend filters are active and API returns data", async () => { - const backendLog = createLogEntry({ request_id: "backend-req" }); - vi.mocked(uiSpendLogsCall).mockResolvedValue( - createPaginatedResponse([backendLog]), - ); - const logs = createPaginatedResponse([createLogEntry({ request_id: "client-req" })]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await waitFor( - () => { - expect(result.current.filteredLogs.data).toHaveLength(1); - expect(result.current.filteredLogs.data[0].request_id).toBe("backend-req"); - }, - { timeout: 500 }, - ); - }); - - it("should call uiSpendLogsCall with request_id when Request ID filter is set", async () => { - vi.mocked(uiSpendLogsCall).mockResolvedValue( - createPaginatedResponse([createLogEntry({ request_id: "req-xyz" })]), - ); - const logs = createPaginatedResponse([createLogEntry()]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - act(() => { - result.current.handleFilterChange({ "Request ID": "req-xyz" }); - }); - - await waitFor( - () => { - expect(uiSpendLogsCall).toHaveBeenCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ request_id: "req-xyz" }), - }), - ); - }, - { timeout: 500 }, - ); - }); - - it("should call uiSpendLogsCall with user_id when User ID filter is set", async () => { - vi.mocked(uiSpendLogsCall).mockResolvedValue( - createPaginatedResponse([createLogEntry()]), - ); - const logs = createPaginatedResponse([createLogEntry()]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - act(() => { - result.current.handleFilterChange({ "User ID": "user-123" }); - }); - - await waitFor( - () => { - expect(uiSpendLogsCall).toHaveBeenCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ user_id: "user-123" }), - }), - ); - }, - { timeout: 500 }, - ); - }); - - it("should call uiSpendLogsCall with error_message when Error Message filter is set", async () => { - vi.mocked(uiSpendLogsCall).mockResolvedValue( - createPaginatedResponse([createLogEntry()]), - ); - const logs = createPaginatedResponse([createLogEntry()]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - act(() => { - result.current.handleFilterChange({ "Error Message": "rate limit exceeded" }); - }); - - await waitFor( - () => { - expect(uiSpendLogsCall).toHaveBeenCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ error_message: "rate limit exceeded" }), - }), - ); - }, - { timeout: 500 }, - ); - }); - - it("should return empty results when backend filters are active but API returns empty", async () => { - vi.mocked(uiSpendLogsCall).mockResolvedValue({ - data: [], - total: 0, - page: 1, - page_size: 50, - total_pages: 0, - }); - const clientLog = createLogEntry({ request_id: "client-req" }); - const logs = createPaginatedResponse([clientLog]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await waitFor( - () => { - expect(uiSpendLogsCall).toHaveBeenCalled(); - }, - { timeout: 500 }, - ); - - expect(result.current.filteredLogs.data).toHaveLength(0); - }); - - it("should refetch when sortBy changes and backend filters are active", async () => { - vi.mocked(uiSpendLogsCall).mockResolvedValue( - createPaginatedResponse([createLogEntry()]), - ); - const logs = createPaginatedResponse([createLogEntry()]); - const { result, rerender } = renderHook( - (props: { sortBy?: LogsSortField }) => - useLogFilterLogic({ ...defaultProps, logs, ...props }), - { wrapper, initialProps: { sortBy: "startTime" as LogsSortField } }, - ); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(1), { - timeout: 500, - }); - - rerender({ sortBy: "spend" as LogsSortField }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(2), { - timeout: 500, - }); - expect(uiSpendLogsCall).toHaveBeenLastCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ sort_by: "spend" }), - }), - ); - }); - - it("should refetch when sortOrder changes and backend filters are active", async () => { - vi.mocked(uiSpendLogsCall).mockResolvedValue( - createPaginatedResponse([createLogEntry()]), - ); - const logs = createPaginatedResponse([createLogEntry()]); - const { result, rerender } = renderHook( - (props: { sortOrder?: "asc" | "desc" }) => - useLogFilterLogic({ ...defaultProps, logs, ...props }), - { wrapper, initialProps: { sortOrder: "desc" } }, - ); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(1), { - timeout: 500, - }); - - rerender({ sortOrder: "asc" }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(2), { - timeout: 500, - }); - expect(uiSpendLogsCall).toHaveBeenLastCalledWith( - expect.objectContaining({ - params: expect.objectContaining({ sort_order: "asc" }), - }), - ); - }); - - it("should refetch when currentPage changes and backend filters are active", async () => { - vi.mocked(uiSpendLogsCall).mockResolvedValue( - createPaginatedResponse([createLogEntry()]), - ); - const logs = createPaginatedResponse([createLogEntry()]); - const { result, rerender } = renderHook( - (props) => useLogFilterLogic({ ...defaultProps, logs, ...props }), - { wrapper, initialProps: { currentPage: 1 } }, - ); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(1), { - timeout: 500, - }); - - rerender({ currentPage: 2 }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(2), { - timeout: 500, - }); - expect(uiSpendLogsCall).toHaveBeenLastCalledWith( - expect.objectContaining({ page: 2 }), - ); - }); - - it("should refetch when startTime changes and backend filters are active", async () => { - vi.mocked(uiSpendLogsCall).mockResolvedValue( - createPaginatedResponse([createLogEntry()]), - ); - const logs = createPaginatedResponse([createLogEntry()]); - const { result, rerender } = renderHook( - (props: { startTime?: string }) => - useLogFilterLogic({ ...defaultProps, logs, ...props }), - { wrapper, initialProps: { startTime: "2025-01-01T00:00:00Z" } }, - ); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(1), { - timeout: 500, - }); - - rerender({ startTime: "2025-01-02T00:00:00Z" }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(2), { - timeout: 500, - }); - expect(uiSpendLogsCall).toHaveBeenLastCalledWith( - expect.objectContaining({ - start_date: "2025-01-02 00:00:00", - }), - ); - }); - - it("should refetch when isCustomDate changes and backend filters are active", async () => { - vi.mocked(uiSpendLogsCall).mockResolvedValue( - createPaginatedResponse([createLogEntry()]), - ); - const logs = createPaginatedResponse([createLogEntry()]); - const { result, rerender } = renderHook( - (props: { isCustomDate?: boolean }) => - useLogFilterLogic({ ...defaultProps, logs, ...props }), - { wrapper, initialProps: { isCustomDate: false } }, - ); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(1), { - timeout: 500, - }); - - rerender({ isCustomDate: true }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(2), { - timeout: 500, + expect(result.current.filteredLogs).toBeDefined(); + expect(result.current).toHaveProperty("allTeams"); + expect(result.current.handleFilterChange).toBeInstanceOf(Function); + expect(result.current.handleFilterReset).toBeInstanceOf(Function); }); }); - it("should not call setCurrentPage when handleFilterChange receives identical filters", async () => { - const setCurrentPage = vi.fn(); - const logs = createPaginatedResponse([createLogEntry()]); - const { result } = renderHook( - () => useLogFilterLogic({ ...defaultProps, logs, setCurrentPage }), - { wrapper }, - ); + describe("handleFilterReset", () => { + it("restores filters to defaults after changes", () => { + const { result } = renderFilterHook(); - act(() => { - result.current.handleFilterChange({ "Team ID": "team-1" }); + act(() => { + result.current.handleFilterChange({ "Team ID": "team-1", Status: "success" }); + }); + + expect(result.current.filters["Team ID"]).toBe("team-1"); + expect(result.current.filters["Status"]).toBe("success"); + + act(() => { + result.current.handleFilterReset(); + }); + + expect(result.current.filters["Team ID"]).toBe(""); + expect(result.current.filters["Status"]).toBe(""); }); - await waitFor(() => expect(setCurrentPage).toHaveBeenCalledTimes(1), { - timeout: 500, + it("calls setCurrentPage(1)", () => { + const setCurrentPage = vi.fn(); + const { result } = renderFilterHook({ setCurrentPage }); + + act(() => { + result.current.handleFilterReset(); + }); + + expect(setCurrentPage).toHaveBeenCalledWith(1); }); - setCurrentPage.mockClear(); + it("triggers a fetch with all filter params undefined", async () => { + vi.mocked(uiSpendLogsCall).mockResolvedValue(emptyResponse); + const { result } = renderFilterHook(); - await act(async () => { - result.current.handleFilterChange({ "Team ID": "team-1" }); - await new Promise((resolve) => setTimeout(resolve, 350)); - }); + act(() => { + result.current.handleFilterChange({ "Key Alias": "alias-1" }); + }); - expect(setCurrentPage).not.toHaveBeenCalled(); - }); + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled(), { timeout: 500 }); - it("should not crash when uiSpendLogsCall throws", async () => { - vi.mocked(uiSpendLogsCall).mockRejectedValue(new Error("Network error")); - const logs = createPaginatedResponse([createLogEntry()]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); + act(() => { + result.current.handleFilterReset(); + }); - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled(), { - timeout: 500, - }); - - expect(result.current.filteredLogs).toBeDefined(); - expect(result.current.filters).toBeDefined(); - }); - - it("should clear backendFilteredLogs when handleFilterReset is called", async () => { - const backendLog = createLogEntry({ request_id: "backend-req" }); - vi.mocked(uiSpendLogsCall).mockResolvedValue( - createPaginatedResponse([backendLog]), - ); - const logs = createPaginatedResponse([createLogEntry({ request_id: "client-req" })]); - const { result } = renderHook(() => useLogFilterLogic({ ...defaultProps, logs }), { wrapper }); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await waitFor( - () => { - expect(result.current.filteredLogs.data[0].request_id).toBe("backend-req"); - }, - { timeout: 500 }, - ); - - act(() => { - result.current.handleFilterReset(); - }); - - expect(result.current.filteredLogs.data).toEqual(logs.data); - expect(result.current.filteredLogs.data[0].request_id).toBe("client-req"); - }); - - it("should pass correct start_date, end_date, sort_by, and sort_order to uiSpendLogsCall", async () => { - vi.mocked(uiSpendLogsCall).mockResolvedValue( - createPaginatedResponse([createLogEntry()]), - ); - const logs = createPaginatedResponse([createLogEntry()]); - const { result } = renderHook( - () => - useLogFilterLogic({ - ...defaultProps, - logs, - startTime: "2025-01-15T00:00:00Z", - endTime: "2025-01-15T23:59:59Z", - isCustomDate: true, - sortBy: "spend", - sortOrder: "asc", - }), - { wrapper }, - ); - - act(() => { - result.current.handleFilterChange({ "Key Alias": "alias-1" }); - }); - - await waitFor( - () => { - expect(uiSpendLogsCall).toHaveBeenCalledWith( - expect.objectContaining({ - start_date: "2025-01-15 00:00:00", - end_date: "2025-01-15 23:59:59", - params: expect.objectContaining({ - sort_by: "spend", - sort_order: "asc", + await waitFor( + () => { + expect(uiSpendLogsCall).toHaveBeenLastCalledWith( + expect.objectContaining({ + params: expect.objectContaining({ + team_id: undefined, + api_key: undefined, + request_id: undefined, + user_id: undefined, + end_user: undefined, + status_filter: undefined, + model_id: undefined, + key_alias: undefined, + error_code: undefined, + error_message: undefined, + }), }), - }), + ); + }, + { timeout: 500 }, + ); + }); + }); + + describe("handleFilterChange", () => { + it("calls setCurrentPage(1) when filters change", () => { + const setCurrentPage = vi.fn(); + const { result } = renderFilterHook({ setCurrentPage }); + + act(() => { + result.current.handleFilterChange({ "Team ID": "team-1" }); + }); + + expect(setCurrentPage).toHaveBeenCalledWith(1); + }); + + it("merges partial updates without clobbering other filter keys", () => { + const { result } = renderFilterHook(); + + act(() => { + result.current.handleFilterChange({ "Team ID": "team-a" }); + }); + expect(result.current.filters["Team ID"]).toBe("team-a"); + + act(() => { + result.current.handleFilterChange({ Model: "gpt-4" }); + }); + + expect(result.current.filters["Team ID"]).toBe("team-a"); + expect(result.current.filters["Model"]).toBe("gpt-4"); + }); + + it("does not call setCurrentPage when filters are identical", async () => { + const setCurrentPage = vi.fn(); + const { result } = renderFilterHook({ setCurrentPage }); + + act(() => { + result.current.handleFilterChange({ "Team ID": "team-1" }); + }); + + await waitFor(() => expect(setCurrentPage).toHaveBeenCalledTimes(1), { timeout: 500 }); + + setCurrentPage.mockClear(); + + await act(async () => { + result.current.handleFilterChange({ "Team ID": "team-1" }); + await new Promise((resolve) => setTimeout(resolve, 350)); + }); + + expect(setCurrentPage).not.toHaveBeenCalled(); + }); + }); + + describe("query params — filter keys", () => { + const filterCases: Array<{ + filterKey: keyof LogFilterState; + paramName: string; + value: string; + }> = [ + { filterKey: "Team ID", paramName: "team_id", value: "team-a" }, + { filterKey: "Key Hash", paramName: "api_key", value: "key-x" }, + { filterKey: "Request ID", paramName: "request_id", value: "req-xyz" }, + { filterKey: "User ID", paramName: "user_id", value: "user-123" }, + { filterKey: "End User", paramName: "end_user", value: "user-a" }, + { filterKey: "Status", paramName: "status_filter", value: "error" }, + { filterKey: "Model", paramName: "model_id", value: "gpt-4" }, + { filterKey: "Public model / search tool", paramName: "model", value: "tavily-marketing" }, + { filterKey: "Error Code", paramName: "error_code", value: "429" }, + { filterKey: "Error Message", paramName: "error_message", value: "rate limit exceeded" }, + ]; + + it.each(filterCases)( + "forwards $filterKey as params.$paramName to uiSpendLogsCall", + async ({ filterKey, paramName, value }) => { + const { result } = renderFilterHook(); + + act(() => { + result.current.handleFilterChange({ [filterKey]: value } as Partial); + }); + + await waitFor( + () => { + expect(uiSpendLogsCall).toHaveBeenCalledWith( + expect.objectContaining({ + params: expect.objectContaining({ [paramName]: value }), + }), + ); + }, + { timeout: 500 }, ); }, - { timeout: 500 }, ); }); + + describe("query params — date & sort", () => { + it("passes start_date, end_date, sort_by, and sort_order to uiSpendLogsCall", async () => { + const { result } = renderFilterHook({ + startTime: "2025-01-15T00:00:00Z", + endTime: "2025-01-15T23:59:59Z", + isCustomDate: true, + sortBy: "spend" as LogsSortField, + sortOrder: "asc", + }); + + act(() => { + result.current.handleFilterChange({ "Key Alias": "alias-1" }); + }); + + await waitFor( + () => { + expect(uiSpendLogsCall).toHaveBeenCalledWith( + expect.objectContaining({ + start_date: "2025-01-15 00:00:00", + end_date: "2025-01-15 23:59:59", + params: expect.objectContaining({ + sort_by: "spend", + sort_order: "asc", + }), + }), + ); + }, + { timeout: 500 }, + ); + }); + }); + + describe("debounce", () => { + it("calls uiSpendLogsCall after the debounce elapses for text filters", async () => { + const { result } = renderFilterHook(); + + act(() => { + result.current.handleFilterChange({ "Key Hash": "hash-1" }); + }); + + await waitFor( + () => + expect(uiSpendLogsCall).toHaveBeenCalledWith( + expect.objectContaining({ + params: expect.objectContaining({ api_key: "hash-1" }), + }), + ), + { timeout: 500 }, + ); + }); + + it("does not call uiSpendLogsCall with a text filter before the debounce elapses", async () => { + const { result } = renderFilterHook(); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled(), { timeout: 500 }); + vi.mocked(uiSpendLogsCall).mockClear(); + + act(() => { + result.current.handleFilterChange({ "Key Hash": "hash-1" }); + }); + + await new Promise((resolve) => setTimeout(resolve, 100)); + expect(uiSpendLogsCall).not.toHaveBeenCalledWith( + expect.objectContaining({ + params: expect.objectContaining({ api_key: "hash-1" }), + }), + ); + + await waitFor( + () => + expect(uiSpendLogsCall).toHaveBeenCalledWith( + expect.objectContaining({ + params: expect.objectContaining({ api_key: "hash-1" }), + }), + ), + { timeout: 500 }, + ); + }); + + it("applies dropdown filter changes without waiting for the debounce", async () => { + const { result } = renderFilterHook(); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled(), { timeout: 500 }); + vi.mocked(uiSpendLogsCall).mockClear(); + + act(() => { + result.current.handleFilterChange({ "Team ID": "team-instant" }); + }); + + await waitFor( + () => + expect(uiSpendLogsCall).toHaveBeenCalledWith( + expect.objectContaining({ + params: expect.objectContaining({ team_id: "team-instant" }), + }), + ), + { timeout: 100 }, + ); + }); + + // Guards the TEXT_FILTER_KEYS fix: this free-text filter must debounce, not fire per keystroke. + it("debounces the 'Public model / search tool' text filter", async () => { + const { result } = renderFilterHook(); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled(), { timeout: 500 }); + vi.mocked(uiSpendLogsCall).mockClear(); + + act(() => { + result.current.handleFilterChange({ "Public model / search tool": "tavily-marketing" }); + }); + + await new Promise((resolve) => setTimeout(resolve, 100)); + expect(uiSpendLogsCall).not.toHaveBeenCalledWith( + expect.objectContaining({ + params: expect.objectContaining({ model: "tavily-marketing" }), + }), + ); + + await waitFor( + () => + expect(uiSpendLogsCall).toHaveBeenCalledWith( + expect.objectContaining({ + params: expect.objectContaining({ model: "tavily-marketing" }), + }), + ), + { timeout: 500 }, + ); + }); + }); + + describe("handleFilterReset", () => { + it("flushes the text-filter debounce so a pending typed value is not sent", async () => { + const { result } = renderFilterHook(); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled(), { timeout: 500 }); + vi.mocked(uiSpendLogsCall).mockClear(); + + act(() => { + result.current.handleFilterChange({ "Key Hash": "pending-hash" }); + }); + + act(() => { + result.current.handleFilterReset(); + }); + + await new Promise((resolve) => setTimeout(resolve, 400)); + + for (const call of vi.mocked(uiSpendLogsCall).mock.calls) { + expect(call[0].params?.api_key).toBeUndefined(); + } + }); + }); + + describe("backend filtered logs", () => { + it("returns the query payload as filteredLogs when backend filters are active", async () => { + const backendLog = { request_id: "backend-req" }; + vi.mocked(uiSpendLogsCall).mockResolvedValue({ + data: [backendLog], + total: 1, + page: 1, + page_size: 50, + total_pages: 1, + } as PaginatedResponse); + + const { result } = renderFilterHook(); + + act(() => { + result.current.handleFilterChange({ "Key Alias": "alias-1" }); + }); + + await waitFor( + () => { + expect(result.current.filteredLogs.data).toHaveLength(1); + expect(result.current.filteredLogs.data[0].request_id).toBe("backend-req"); + }, + { timeout: 500 }, + ); + }); + + it("returns empty data when the API returns an empty payload", async () => { + vi.mocked(uiSpendLogsCall).mockResolvedValue(emptyResponse); + const { result } = renderFilterHook(); + + act(() => { + result.current.handleFilterChange({ "Key Alias": "alias-1" }); + }); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled(), { timeout: 500 }); + + expect(result.current.filteredLogs.data).toHaveLength(0); + }); + }); + + describe("refetch triggers", () => { + it("refetches when sortBy changes", async () => { + const { rerender } = renderHook( + (props: { sortBy: LogsSortField }) => { + const [filters, setFilters] = useState(defaultFilters); + return useLogFilterLogic({ + ...defaultProps, + filters, + setFilters, + setCurrentPage: vi.fn(), + sortBy: props.sortBy, + }); + }, + { wrapper, initialProps: { sortBy: "startTime" } }, + ); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(1), { timeout: 500 }); + + rerender({ sortBy: "spend" }); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(2), { timeout: 500 }); + expect(uiSpendLogsCall).toHaveBeenLastCalledWith( + expect.objectContaining({ + params: expect.objectContaining({ sort_by: "spend" }), + }), + ); + }); + + it("refetches when sortOrder changes", async () => { + const { rerender } = renderHook( + (props: { sortOrder: "asc" | "desc" }) => { + const [filters, setFilters] = useState(defaultFilters); + return useLogFilterLogic({ + ...defaultProps, + filters, + setFilters, + setCurrentPage: vi.fn(), + sortOrder: props.sortOrder, + }); + }, + { wrapper, initialProps: { sortOrder: "desc" } }, + ); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(1), { timeout: 500 }); + + rerender({ sortOrder: "asc" }); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(2), { timeout: 500 }); + expect(uiSpendLogsCall).toHaveBeenLastCalledWith( + expect.objectContaining({ + params: expect.objectContaining({ sort_order: "asc" }), + }), + ); + }); + + it("refetches when currentPage changes", async () => { + const { rerender } = renderHook( + (props: { currentPage: number }) => { + const [filters, setFilters] = useState(defaultFilters); + return useLogFilterLogic({ + ...defaultProps, + filters, + setFilters, + setCurrentPage: vi.fn(), + currentPage: props.currentPage, + }); + }, + { wrapper, initialProps: { currentPage: 1 } }, + ); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(1), { timeout: 500 }); + + rerender({ currentPage: 2 }); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(2), { timeout: 500 }); + expect(uiSpendLogsCall).toHaveBeenLastCalledWith(expect.objectContaining({ page: 2 })); + }); + + it("refetches when startTime changes", async () => { + const { rerender } = renderHook( + (props: { startTime: string }) => { + const [filters, setFilters] = useState(defaultFilters); + return useLogFilterLogic({ + ...defaultProps, + filters, + setFilters, + setCurrentPage: vi.fn(), + startTime: props.startTime, + }); + }, + { wrapper, initialProps: { startTime: "2025-01-01T00:00:00Z" } }, + ); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(1), { timeout: 500 }); + + rerender({ startTime: "2025-01-02T00:00:00Z" }); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(2), { timeout: 500 }); + expect(uiSpendLogsCall).toHaveBeenLastCalledWith(expect.objectContaining({ start_date: "2025-01-02 00:00:00" })); + }); + + it("refetches with a different end_date when isCustomDate toggles", async () => { + const customEndTime = "2025-01-15T23:59:59Z"; + const customEndFormatted = "2025-01-15 23:59:59"; + + const { rerender } = renderHook( + (props: { isCustomDate: boolean }) => { + const [filters, setFilters] = useState(defaultFilters); + return useLogFilterLogic({ + ...defaultProps, + endTime: customEndTime, + filters, + setFilters, + setCurrentPage: vi.fn(), + isCustomDate: props.isCustomDate, + }); + }, + { wrapper, initialProps: { isCustomDate: false } }, + ); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(1), { timeout: 500 }); + const firstEndDate = vi.mocked(uiSpendLogsCall).mock.calls[0][0].end_date; + expect(firstEndDate).not.toBe(customEndFormatted); + + rerender({ isCustomDate: true }); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalledTimes(2), { timeout: 500 }); + expect(vi.mocked(uiSpendLogsCall).mock.calls[1][0].end_date).toBe(customEndFormatted); + }); + }); + + describe("query enablement", () => { + const nullCredentialCases: Array<{ name: string; override: HookOverrides }> = [ + { name: "accessToken", override: { accessToken: null } }, + { name: "token", override: { token: null } }, + { name: "userRole", override: { userRole: null } }, + { name: "userID", override: { userID: null } }, + ]; + + it.each(nullCredentialCases)("does not call uiSpendLogsCall when $name is null", async ({ override }) => { + const { result } = renderFilterHook(override); + + act(() => { + result.current.handleFilterChange({ "Key Alias": "alias-1" }); + }); + + await new Promise((resolve) => setTimeout(resolve, 350)); + + expect(uiSpendLogsCall).not.toHaveBeenCalled(); + }); + + it("does not call uiSpendLogsCall when activeTab is not 'request logs'", async () => { + const { result } = renderFilterHook({ activeTab: "audit logs" }); + + act(() => { + result.current.handleFilterChange({ "Key Alias": "alias-1" }); + }); + + await new Promise((resolve) => setTimeout(resolve, 350)); + + expect(uiSpendLogsCall).not.toHaveBeenCalled(); + }); + }); + + describe("filterByCurrentUser", () => { + it("sends user_id: userID when the User ID filter is blank", async () => { + const { result } = renderFilterHook({ + filterByCurrentUser: true, + userID: "me-123", + }); + + act(() => { + result.current.handleFilterChange({ "Key Alias": "alias-1" }); + }); + + await waitFor( + () => { + expect(uiSpendLogsCall).toHaveBeenCalledWith( + expect.objectContaining({ + params: expect.objectContaining({ user_id: "me-123" }), + }), + ); + }, + { timeout: 500 }, + ); + }); + }); + + describe("error handling", () => { + it("does not crash when uiSpendLogsCall throws", async () => { + vi.mocked(uiSpendLogsCall).mockRejectedValue(new Error("Network error")); + const { result } = renderFilterHook(); + + act(() => { + result.current.handleFilterChange({ "Key Alias": "alias-1" }); + }); + + await waitFor(() => expect(uiSpendLogsCall).toHaveBeenCalled(), { timeout: 500 }); + + expect(result.current.filteredLogs).toBeDefined(); + expect(result.current.filteredLogs.data).toEqual([]); + }); + }); +}); + +describe("getLiveTailRefetchInterval", () => { + it("polls every 15s when live tail is on and on page 1", () => { + expect(getLiveTailRefetchInterval(true, 1)).toBe(LIVE_TAIL_INTERVAL_MS); + }); + + it("does not poll when live tail is off", () => { + expect(getLiveTailRefetchInterval(false, 1)).toBe(false); + }); + + it("does not poll when not on page 1, even with live tail on", () => { + expect(getLiveTailRefetchInterval(true, 2)).toBe(false); + }); }); diff --git a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx index 8f916999c16..d9830a83f6e 100644 --- a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx @@ -1,13 +1,28 @@ import moment from "moment"; -import { useCallback, useEffect, useState, useRef, useMemo } from "react"; +import { useEffect, useMemo, useState } from "react"; import { uiSpendLogsCall } from "../networking"; import { Team } from "../key_team_helpers/key_list"; -import { useQuery } from "@tanstack/react-query"; +import { keepPreviousData, useQuery } from "@tanstack/react-query"; import { fetchAllTeams } from "../../components/key_team_helpers/filter_helpers"; -import { debounce } from "lodash"; import { defaultPageSize } from "../constants"; -import { PaginatedResponse } from "."; -import type { LogsSortField } from "./columns"; +import type { LogEntry, LogsSortField } from "./columns"; + +export interface PaginatedResponse { + data: LogEntry[]; + total: number; + page: number; + page_size: number; + total_pages: number; +} + +function useDebouncedValue(value: T, delayMs: number): [T, React.Dispatch>] { + const [debounced, setDebounced] = useState(value); + useEffect(() => { + const timer = setTimeout(() => setDebounced(value), delayMs); + return () => clearTimeout(timer); + }, [value, delayMs]); + return [debounced, setDebounced]; +} /** Spend log `model` column (LLM public model name or `search_tool_name` for /search). */ export const FILTER_KEYS = { @@ -28,324 +43,188 @@ export const FILTER_KEYS = { export type FilterKey = keyof typeof FILTER_KEYS; export type LogFilterState = Record<(typeof FILTER_KEYS)[FilterKey], string>; +// Keys whose UI is a free-form text input; only these need debouncing. +const TEXT_FILTER_KEYS: readonly (keyof LogFilterState)[] = [ + FILTER_KEYS.KEY_HASH, + FILTER_KEYS.ERROR_MESSAGE, + FILTER_KEYS.REQUEST_ID, + FILTER_KEYS.USER_ID, + FILTER_KEYS.PUBLIC_MODEL_OR_SEARCH_TOOL, +]; + +// Live-tail polls every 15s, but only on page 1 (newest) while live tail is on. +export const LIVE_TAIL_INTERVAL_MS = 15000; +export const getLiveTailRefetchInterval = (isLiveTail: boolean, currentPage: number): number | false => + isLiveTail && currentPage === 1 ? LIVE_TAIL_INTERVAL_MS : false; + +export const defaultFilters: LogFilterState = { + [FILTER_KEYS.TEAM_ID]: "", + [FILTER_KEYS.KEY_HASH]: "", + [FILTER_KEYS.REQUEST_ID]: "", + [FILTER_KEYS.MODEL]: "", + [FILTER_KEYS.PUBLIC_MODEL_OR_SEARCH_TOOL]: "", + [FILTER_KEYS.USER_ID]: "", + [FILTER_KEYS.END_USER]: "", + [FILTER_KEYS.STATUS]: "", + [FILTER_KEYS.KEY_ALIAS]: "", + [FILTER_KEYS.ERROR_CODE]: "", + [FILTER_KEYS.ERROR_MESSAGE]: "", +}; + export function useLogFilterLogic({ - logs, accessToken, - startTime, // Receive from SpendLogsTable - endTime, // Receive from SpendLogsTable + token, + userRole, + userID, + filters, + setFilters, + filterByCurrentUser, + activeTab, + isLiveTail, + startTime, + endTime, pageSize = defaultPageSize, isCustomDate, setCurrentPage, - userID, - userRole, sortBy = "startTime", sortOrder = "desc", currentPage = 1, }: { - logs: PaginatedResponse; accessToken: string | null; + token: string | null; + userRole: string | null; + userID: string | null; + filters: LogFilterState; + setFilters: React.Dispatch>; + filterByCurrentUser: boolean | null; + activeTab: string; + isLiveTail: boolean; startTime: string; endTime: string; pageSize?: number; isCustomDate: boolean; setCurrentPage: (page: number) => void; - userID: string | null; - userRole: string | null; sortBy?: LogsSortField; sortOrder?: "asc" | "desc"; currentPage?: number; }) { - const defaultFilters = useMemo( - () => ({ - [FILTER_KEYS.TEAM_ID]: "", - [FILTER_KEYS.KEY_HASH]: "", - [FILTER_KEYS.REQUEST_ID]: "", - [FILTER_KEYS.MODEL]: "", - [FILTER_KEYS.PUBLIC_MODEL_OR_SEARCH_TOOL]: "", - [FILTER_KEYS.USER_ID]: "", - [FILTER_KEYS.END_USER]: "", - [FILTER_KEYS.STATUS]: "", - [FILTER_KEYS.KEY_ALIAS]: "", - [FILTER_KEYS.ERROR_CODE]: "", - [FILTER_KEYS.ERROR_MESSAGE]: "", - }), - [], - ); + const [debouncedFilters, setDebouncedFilters] = useDebouncedValue(filters, 300); - const [filters, setFilters] = useState(defaultFilters); - const [backendFilteredLogs, setBackendFilteredLogs] = useState(null); - const lastSearchTimestamp = useRef(0); + // Live values for dropdown keys, debounced for text keys. + const effectiveFilters = useMemo(() => { + const merged = { ...filters }; + for (const k of TEXT_FILTER_KEYS) { + merged[k] = debouncedFilters[k]; + } + return merged; + }, [filters, debouncedFilters]); - // Refs that always hold the latest filters and hasBackendFilters values. - // The sort/page/time effect below intentionally omits these from its dep array - // to avoid double-fetches when a filter changes; reading from refs instead of - // the closure prevents stale-closure bugs (e.g. the effect using a snapshot of - // filters taken before the user selected Key Alias). - const filtersRef = useRef(filters); - const hasBackendFiltersRef = useRef(false); - const performSearch = useCallback( - async (filters: LogFilterState, page = 1) => { - if (!accessToken) return; - - console.log("Filters being sent to API:", filters); - const currentTimestamp = Date.now(); - lastSearchTimestamp.current = currentTimestamp; + const logsQuery = useQuery({ + queryKey: [ + "logs", + "table", + currentPage, + pageSize, + startTime, + endTime, + isCustomDate, + effectiveFilters, + filterByCurrentUser ? userID : null, + sortBy, + sortOrder, + ], + queryFn: async () => { + if (!accessToken || !token || !userRole || !userID) { + return { + data: [], + total: 0, + page: 1, + page_size: pageSize, + total_pages: 0, + }; + } const formattedStartTime = moment(startTime).utc().format("YYYY-MM-DD HH:mm:ss"); const formattedEndTime = isCustomDate ? moment(endTime).utc().format("YYYY-MM-DD HH:mm:ss") : moment().utc().format("YYYY-MM-DD HH:mm:ss"); - try { - const response = await uiSpendLogsCall({ - accessToken, - start_date: formattedStartTime, - end_date: formattedEndTime, - page, - page_size: pageSize, - params: { - api_key: filters[FILTER_KEYS.KEY_HASH] || undefined, - team_id: filters[FILTER_KEYS.TEAM_ID] || undefined, - request_id: filters[FILTER_KEYS.REQUEST_ID] || undefined, - user_id: filters[FILTER_KEYS.USER_ID] || undefined, - end_user: filters[FILTER_KEYS.END_USER] || undefined, - status_filter: filters[FILTER_KEYS.STATUS] || undefined, - model_id: filters[FILTER_KEYS.MODEL] || undefined, - model: filters[FILTER_KEYS.PUBLIC_MODEL_OR_SEARCH_TOOL] || undefined, - key_alias: filters[FILTER_KEYS.KEY_ALIAS] || undefined, - error_code: filters[FILTER_KEYS.ERROR_CODE] || undefined, - error_message: filters[FILTER_KEYS.ERROR_MESSAGE] || undefined, - sort_by: sortBy, - sort_order: sortOrder, - }, - }); + const response = await uiSpendLogsCall({ + accessToken, + start_date: formattedStartTime, + end_date: formattedEndTime, + page: currentPage, + page_size: pageSize, + params: { + api_key: effectiveFilters[FILTER_KEYS.KEY_HASH] || undefined, + team_id: effectiveFilters[FILTER_KEYS.TEAM_ID] || undefined, + request_id: effectiveFilters[FILTER_KEYS.REQUEST_ID] || undefined, + user_id: effectiveFilters[FILTER_KEYS.USER_ID] || (filterByCurrentUser ? userID ?? undefined : undefined), + end_user: effectiveFilters[FILTER_KEYS.END_USER] || undefined, + status_filter: effectiveFilters[FILTER_KEYS.STATUS] || undefined, + model_id: effectiveFilters[FILTER_KEYS.MODEL] || undefined, + model: effectiveFilters[FILTER_KEYS.PUBLIC_MODEL_OR_SEARCH_TOOL] || undefined, + key_alias: effectiveFilters[FILTER_KEYS.KEY_ALIAS] || undefined, + error_code: effectiveFilters[FILTER_KEYS.ERROR_CODE] || undefined, + error_message: effectiveFilters[FILTER_KEYS.ERROR_MESSAGE] || undefined, + sort_by: sortBy, + sort_order: sortOrder, + }, + }); - if (currentTimestamp === lastSearchTimestamp.current) { - setBackendFilteredLogs({ - ...response, - data: response.data ?? [], - }); - } - } catch (error) { - console.error("Error searching users:", error); - setBackendFilteredLogs({ - data: [], - total: 0, - page: 1, - page_size: pageSize, - total_pages: 0, - }); - } + return response; }, - [accessToken, startTime, endTime, isCustomDate, pageSize, sortBy, sortOrder], - ); + enabled: !!accessToken && !!token && !!userRole && !!userID && activeTab === "request logs", + refetchInterval: getLiveTailRefetchInterval(isLiveTail, currentPage), + placeholderData: keepPreviousData, + // Only live-tail-poll while the tab is visible. + refetchIntervalInBackground: false, + }); - const debouncedSearch = useMemo( - () => debounce((filters: LogFilterState, page: number) => performSearch(filters, page), 300), - [performSearch], - ); + const filteredLogs: PaginatedResponse = logsQuery.data ?? { + data: [], + total: 0, + page: 1, + page_size: pageSize, + total_pages: 0, + }; - useEffect(() => { - return () => debouncedSearch.cancel(); - }, [debouncedSearch]); - - // Determine when backend filters are active (server-side filtering) - const hasBackendFilters = useMemo( - () => - !!( - filters[FILTER_KEYS.KEY_ALIAS] || - filters[FILTER_KEYS.KEY_HASH] || - filters[FILTER_KEYS.REQUEST_ID] || - filters[FILTER_KEYS.USER_ID] || - filters[FILTER_KEYS.END_USER] || - filters[FILTER_KEYS.ERROR_CODE] || - filters[FILTER_KEYS.ERROR_MESSAGE] || - filters[FILTER_KEYS.MODEL] || - filters[FILTER_KEYS.PUBLIC_MODEL_OR_SEARCH_TOOL] - ), - [filters], - ); - - // Keep refs in sync on every render so the sort/page/time effect always reads - // the latest values without those values being in its dep array. - useEffect(() => { - filtersRef.current = filters; - hasBackendFiltersRef.current = hasBackendFilters; - }, [filters, hasBackendFilters]); - - // Refetch when sort, page, or time range changes (backend filters use their own fetch, not the main query) - useEffect(() => { - if (hasBackendFiltersRef.current && accessToken) { - // Cancel any pending debounced search to prevent it from overwriting this page's results - debouncedSearch.cancel(); - performSearch(filtersRef.current, currentPage); - } - // filters / hasBackendFilters are read via refs — avoids stale-closure bugs - // when sort/page/time changes after a filter (e.g. Key Alias) was set. - // debouncedSearch / performSearch: filter changes go through handleFilterChange - // → debouncedSearch; adding them here would cause double-fetches on filter apply. - // accessToken: stable across sort/page/time changes. - // eslint-disable-next-line react-hooks/exhaustive-deps - }, [sortBy, sortOrder, currentPage, startTime, endTime, isCustomDate]); - - // Compute client-side filtered logs directly from incoming logs and filters - const clientDerivedFilteredLogs: PaginatedResponse = useMemo(() => { - if (!logs || !logs.data) { - return { - data: [], - total: 0, - page: 1, - page_size: pageSize, - total_pages: 0, - }; - } - - // If backend filters are on, don't perform client-side filtering here - if (hasBackendFilters) { - return logs; - } - - let filteredData = [...logs.data]; - - if (filters[FILTER_KEYS.TEAM_ID]) { - filteredData = filteredData.filter((log) => log.team_id === filters[FILTER_KEYS.TEAM_ID]); - } - - if (filters[FILTER_KEYS.STATUS]) { - filteredData = filteredData.filter((log) => { - if (filters[FILTER_KEYS.STATUS] === "success") { - return !log.status || log.status === "success"; - } - return log.status === filters[FILTER_KEYS.STATUS]; - }); - } - - if (filters[FILTER_KEYS.MODEL]) { - filteredData = filteredData.filter((log) => log.model_id === filters[FILTER_KEYS.MODEL]); - } - - if (filters[FILTER_KEYS.PUBLIC_MODEL_OR_SEARCH_TOOL]) { - const m = filters[FILTER_KEYS.PUBLIC_MODEL_OR_SEARCH_TOOL]; - filteredData = filteredData.filter((log) => log.model === m); - } - - if (filters[FILTER_KEYS.KEY_HASH]) { - filteredData = filteredData.filter((log) => log.api_key === filters[FILTER_KEYS.KEY_HASH]); - } - - if (filters[FILTER_KEYS.END_USER]) { - filteredData = filteredData.filter((log) => log.end_user === filters[FILTER_KEYS.END_USER]); - } - - if (filters[FILTER_KEYS.ERROR_CODE]) { - filteredData = filteredData.filter((log) => { - const metadata = log.metadata || {}; - const errorInfo = metadata.error_information; - return errorInfo && errorInfo.error_code === filters[FILTER_KEYS.ERROR_CODE]; - }); - } - - return { - data: filteredData, - total: logs.total, - page: logs.page, - page_size: logs.page_size, - total_pages: logs.total_pages, - }; - }, [logs, filters, hasBackendFilters]); - - // Choose which filtered logs to expose: backend result when active, otherwise client-derived - const filteredLogs: PaginatedResponse = useMemo(() => { - if (hasBackendFilters) { - // When backend filters are active, only show backend results. - // If search hasn't completed yet (null), show empty state rather than - // falling back to unfiltered logs — that caused filtered views to - // display mismatched data when the filter matched zero rows. - if (backendFilteredLogs !== null) { - return backendFilteredLogs; - } - return { - data: [], - total: 0, - page: 1, - page_size: pageSize, - total_pages: 0, - }; - } - return clientDerivedFilteredLogs; - }, [hasBackendFilters, backendFilteredLogs, clientDerivedFilteredLogs]); - - // Fetch all teams and users for potential filter dropdowns (optional, can be adapted) const { data: allTeams } = useQuery({ queryKey: ["allTeamsForLogFilters", accessToken], queryFn: async () => { if (!accessToken) return []; - // Use fetchAllTeams helper function for consistency and abstraction - // Assuming fetchAllTeams returns Team[] directly const teamsData = await fetchAllTeams(accessToken); - return teamsData || []; // Ensure it returns an array + return teamsData || []; }, enabled: !!accessToken, }); - // Update filters state const handleFilterChange = (newFilters: Partial) => { setFilters((prev) => { const updatedFilters = { ...prev, ...newFilters }; - - // Ensure all keys in LogFilterState are present, defaulting to '' if not in newFilters for (const key of Object.keys(defaultFilters) as Array) { if (!(key in updatedFilters)) { updatedFilters[key] = defaultFilters[key]; } } - - // Only call debouncedSearch if filters have actually changed if (JSON.stringify(updatedFilters) !== JSON.stringify(prev)) { setCurrentPage(1); - setBackendFilteredLogs(null); - debouncedSearch(updatedFilters, 1); } - return updatedFilters as LogFilterState; }); }; const handleFilterReset = () => { - // Reset filters state setFilters(defaultFilters); - - // Clear backend filtered logs to ensure fresh render - setBackendFilteredLogs(null); - - // Cancel any in-flight debounced search - debouncedSearch.cancel(); - - // Reset to first page so the unfiltered view starts at page 1 + setDebouncedFilters(defaultFilters); setCurrentPage(1); }; - // Expose a filter-aware refetch so callers (e.g. the manual Fetch button) can - // refresh results while keeping all active backend filters intact. The plain - // `logs.refetch()` in the parent only re-runs the main TanStack Query, which - // does not carry key_alias or other backend-only filter params. - const refetchWithFilters = useCallback( - (page = currentPage) => { - if (hasBackendFilters && accessToken) { - debouncedSearch.cancel(); - performSearch(filters, page); - } - }, - [hasBackendFilters, accessToken, filters, currentPage, performSearch, debouncedSearch], - ); - return { - filters, + logsQuery, filteredLogs, - hasBackendFilters, allTeams, handleFilterChange, handleFilterReset, - refetchWithFilters, }; } diff --git a/ui/litellm-dashboard/src/components/view_logs/logs_utils.test.tsx b/ui/litellm-dashboard/src/components/view_logs/logs_utils.test.tsx new file mode 100644 index 00000000000..b0b74df4c24 --- /dev/null +++ b/ui/litellm-dashboard/src/components/view_logs/logs_utils.test.tsx @@ -0,0 +1,45 @@ +import moment from "moment"; +import { describe, expect, it } from "vitest"; +import { getTimeRangeDisplay } from "./logs_utils"; + +// startTime built relative to "now"; getTimeRangeDisplay computes now() internally. +const ago = (amount: number, unit: moment.unitOfTime.DurationConstructor) => + moment().subtract(amount, unit).toISOString(); + +describe("getTimeRangeDisplay", () => { + it("labels a ~1-minute window as 'Last 1 Minute'", () => { + expect(getTimeRangeDisplay(false, ago(1, "minutes"), "")).toBe("Last 1 Minute"); + }); + + it("labels a ~10-minute window as 'Last 15 Minutes'", () => { + expect(getTimeRangeDisplay(false, ago(10, "minutes"), "")).toBe("Last 15 Minutes"); + }); + + it("labels a ~30-minute window as 'Last Hour'", () => { + expect(getTimeRangeDisplay(false, ago(30, "minutes"), "")).toBe("Last Hour"); + }); + + it("labels a ~2-hour window as 'Last 4 Hours'", () => { + expect(getTimeRangeDisplay(false, ago(2, "hours"), "")).toBe("Last 4 Hours"); + }); + + it("labels a ~10-hour window as 'Last 24 Hours'", () => { + expect(getTimeRangeDisplay(false, ago(10, "hours"), "")).toBe("Last 24 Hours"); + }); + + it("labels a ~3-day window as 'Last 7 Days'", () => { + expect(getTimeRangeDisplay(false, ago(3, "days"), "")).toBe("Last 7 Days"); + }); + + it("falls back to a 'MMM D - MMM D' range beyond 7 days", () => { + const label = getTimeRangeDisplay(false, ago(30, "days"), ""); + expect(label).toMatch(/^[A-Z][a-z]{2} \d{1,2} - [A-Z][a-z]{2} \d{1,2}$/); + }); + + it("renders an explicit start - end range when isCustomDate is true", () => { + const start = "2025-01-02T03:04:00Z"; + const end = "2025-01-05T06:07:00Z"; + const expected = `${moment(start).format("MMM D, h:mm A")} - ${moment(end).format("MMM D, h:mm A")}`; + expect(getTimeRangeDisplay(true, start, end)).toBe(expected); + }); +}); diff --git a/ui/litellm-dashboard/tests/view_logs/useLogFilterLogic.min.test.tsx b/ui/litellm-dashboard/tests/view_logs/useLogFilterLogic.min.test.tsx deleted file mode 100644 index faede84f4c2..00000000000 --- a/ui/litellm-dashboard/tests/view_logs/useLogFilterLogic.min.test.tsx +++ /dev/null @@ -1,62 +0,0 @@ -import React from "react"; -import { describe, it, expect } from "vitest"; -import { render, screen } from "@testing-library/react"; -import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; -import { useLogFilterLogic } from "../../src/components/view_logs/log_filter_logic"; - -// Minimal mocks to avoid real network during hook init -vi.mock("../../src/components/key_team_helpers/filter_helpers", () => ({ - fetchAllKeyAliases: vi.fn().mockResolvedValue([]), - fetchAllTeams: vi.fn().mockResolvedValue([]), -})); - -const createQueryClient = () => - new QueryClient({ - defaultOptions: { queries: { retry: false, gcTime: 0 } }, - }); - -function Harness({ logs }: { logs: any }) { - const { filteredLogs } = useLogFilterLogic({ - logs, - accessToken: "token", - startTime: "2025-01-01 00:00:00", - endTime: "2025-01-02 00:00:00", - pageSize: 50, - isCustomDate: true, - setCurrentPage: () => {}, - userID: "user-1", - userRole: "admin", - }); - - return
{filteredLogs.data.length}
; -} - -describe("useLogFilterLogic (minimal)", () => { - it("useLogFilterLogic minimal: updates filteredLogs when logs change", async () => { - const qc = createQueryClient(); - const logsA = { data: [{ request_id: "a" }], total: 1, page: 1, page_size: 50, total_pages: 1 }; - const logsB = { - data: [{ request_id: "a" }, { request_id: "b" }], - total: 2, - page: 1, - page_size: 50, - total_pages: 1, - }; - - const { rerender } = render( - - - , - ); - - expect(await screen.findByTestId("count")).toHaveTextContent("1"); - - rerender( - - - , - ); - - expect(await screen.findByTestId("count")).toHaveTextContent("2"); - }); -}); From 92de7423efca5756a2cb1bcf3228812628f91960 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 19 May 2026 14:48:30 -0700 Subject: [PATCH 2/7] fix(tests): replace shut-down gpt-4o-audio-preview with gpt-audio-1.5 (#28281) * fix(tests): replace shut-down gpt-4o-audio-preview with gpt-audio-1.5 OpenAI shut down gpt-4o-audio-preview on 2026-05-07, so the live audio calls in test_stream_chunk_builder_openai_audio_output_usage and test_standard_logging_payload_audio now hard-fail with a model-not-found error on every PR. The error was not "openai-internal", so the except block swallowed it and execution fell through to an unbound completion/response (UnboundLocalError). Switch both tests to gpt-audio-1.5, OpenAI's recommended successor (GA, not deprecated, already present in the litellm cost map so the response_cost assertion still resolves). Also broaden the except to skip with the real error in the reason instead of crashing, so a transient upstream blip can't reintroduce the UnboundLocalError. * fix(tests): narrow audio-test skip to model-not-found, re-raise the rest Address review feedback: an unconditional skip on any exception would silently mask a litellm-internal regression in the audio path (broken param transformation, serialization, bad header) instead of failing CI. Skip only on the upstream-unavailable class (model_not_found / "does not exist" / openai-internal) and re-raise everything else, so genuine regressions still fail loudly. The UnboundLocalError is still fixed because the handler either skips or raises - it never falls through. * fix(tests): add budget_exceeded to expected Interaction status enum Staging added budget_exceeded to the Interaction OpenAPI status enum; the staging merge into this branch picked up the spec change but not the matching test update, so test_status_enum_values failed in CI. Align the test's expected list (exact-match by design) with the live spec. * fix(tests): mock HTTP fetch in test_img_url_token_counter The test parameterized a live third-party image URL (blog.purpureus.net) which now 404s, causing get_image_dimensions to fall through to its base64 decode path and crash with 'not enough values to unpack' on every PR run. Mock safe_get with a tiny 1x1 PNG so the URL branch is still exercised without any network dependency. * fix(tests): swap gpt-4o-audio-preview to gpt-audio-1.5 in test_gpt4o_audio OpenAI shut down gpt-4o-audio-preview on 2026-05-07, so both live tests in test_gpt4o_audio.py (test_audio_output_from_model and test_audio_input_to_model) hard-fail model_not_found on every PR. Swap the hardcoded model to OpenAI's successor gpt-audio-1.5 (same chat-completions audio surface; already in the litellm cost map). Mirror the narrowed-skip pattern from the prior audio fixes: skip on model_not_found / does-not-exist / openai-internal, re-raise everything else so genuine litellm regressions still fail CI loudly. --- tests/llm_translation/test_gpt4o_audio.py | 25 +++++++++++----- .../test_custom_callback_input.py | 12 ++++++-- .../test_stream_chunk_builder.py | 12 ++++++-- .../interactions/test_openapi_compliance.py | 6 +++- .../litellm_core_utils/test_token_counter.py | 29 +++++++++++++++++-- 5 files changed, 68 insertions(+), 16 deletions(-) diff --git a/tests/llm_translation/test_gpt4o_audio.py b/tests/llm_translation/test_gpt4o_audio.py index 169fe855163..a50d07406d4 100644 --- a/tests/llm_translation/test_gpt4o_audio.py +++ b/tests/llm_translation/test_gpt4o_audio.py @@ -59,7 +59,7 @@ async def test_audio_output_from_model(stream): litellm.set_verbose = False try: completion = await litellm.acompletion( - model="gpt-4o-audio-preview", + model="gpt-audio-1.5", modalities=["text", "audio"], audio={"voice": "alloy", "format": "pcm16"}, messages=[{"role": "user", "content": "response in 1 word - yes or no"}], @@ -69,8 +69,14 @@ async def test_audio_output_from_model(stream): print(e) pytest.skip("Skipping test due to timeout") except Exception as e: - if "openai-internal" in str(e): - pytest.skip("Skipping test due to openai-internal error") + err = str(e).lower() + if ( + "model_not_found" in err + or "does not exist" in err + or "openai-internal" in err + ): + pytest.skip(f"Skipping - upstream gpt-audio-1.5 unavailable: {e}") + raise if stream is True: await check_streaming_response(completion) @@ -85,7 +91,7 @@ async def test_audio_output_from_model(stream): @pytest.mark.asyncio @pytest.mark.parametrize("stream", [True, False]) -@pytest.mark.parametrize("model", ["gpt-4o-audio-preview"]) # "gpt-4o-audio-preview", +@pytest.mark.parametrize("model", ["gpt-audio-1.5"]) async def test_audio_input_to_model(stream, model): # Fetch the audio file and convert it to a base64 encoded string audio_format = "pcm16" @@ -121,9 +127,14 @@ async def test_audio_input_to_model(stream, model): print(e) pytest.skip("Skipping test due to timeout") except Exception as e: - if "openai-internal" in str(e): - pytest.skip("Skipping test due to openai-internal error") - raise e + err = str(e).lower() + if ( + "model_not_found" in err + or "does not exist" in err + or "openai-internal" in err + ): + pytest.skip(f"Skipping - upstream gpt-audio-1.5 unavailable: {e}") + raise if stream is True: await check_streaming_response(completion) else: diff --git a/tests/local_testing/test_custom_callback_input.py b/tests/local_testing/test_custom_callback_input.py index 545039e60ba..6a4ec9206f7 100644 --- a/tests/local_testing/test_custom_callback_input.py +++ b/tests/local_testing/test_custom_callback_input.py @@ -1125,7 +1125,7 @@ def test_standard_logging_payload_audio(turn_off_message_logging, stream): ) as mock_client: try: response = litellm.completion( - model="gpt-4o-audio-preview", + model="gpt-audio-1.5", modalities=["text", "audio"], audio={"voice": "alloy", "format": "pcm16"}, messages=[ @@ -1134,8 +1134,14 @@ def test_standard_logging_payload_audio(turn_off_message_logging, stream): stream=stream, ) except Exception as e: - if "openai-internal" in str(e): - pytest.skip("Skipping test due to openai-internal error") + err = str(e).lower() + if ( + "model_not_found" in err + or "does not exist" in err + or "openai-internal" in err + ): + pytest.skip(f"Skipping - upstream gpt-audio-1.5 unavailable: {e}") + raise if stream: for chunk in response: diff --git a/tests/local_testing/test_stream_chunk_builder.py b/tests/local_testing/test_stream_chunk_builder.py index 24fdf49c16c..38e04b93f18 100644 --- a/tests/local_testing/test_stream_chunk_builder.py +++ b/tests/local_testing/test_stream_chunk_builder.py @@ -649,7 +649,7 @@ def test_stream_chunk_builder_openai_audio_output_usage(): try: completion = client.chat.completions.create( - model="gpt-4o-audio-preview", + model="gpt-audio-1.5", modalities=["text", "audio"], audio={"voice": "alloy", "format": "pcm16"}, messages=[{"role": "user", "content": "response in 1 word - yes or no"}], @@ -657,8 +657,14 @@ def test_stream_chunk_builder_openai_audio_output_usage(): stream_options={"include_usage": True}, ) except Exception as e: - if "openai-internal" in str(e): - pytest.skip("Skipping test due to openai-internal error") + err = str(e).lower() + if ( + "model_not_found" in err + or "does not exist" in err + or "openai-internal" in err + ): + pytest.skip(f"Skipping - upstream gpt-audio-1.5 unavailable: {e}") + raise chunks = [] for chunk in completion: diff --git a/tests/test_litellm/interactions/test_openapi_compliance.py b/tests/test_litellm/interactions/test_openapi_compliance.py index 1d3b6b8ae1e..aededaaca77 100644 --- a/tests/test_litellm/interactions/test_openapi_compliance.py +++ b/tests/test_litellm/interactions/test_openapi_compliance.py @@ -179,7 +179,10 @@ class TestResponseCompliance: # `status` is an output-only field; validate against the response schema. schema = spec_dict["components"]["schemas"]["Interaction"] status_prop = schema["properties"]["status"] - # Google Interactions API uses lowercase status values (updated Feb 2026) + # Google Interactions API uses lowercase status values (updated Feb 2026). + # Keep this an exact match: this test intentionally breaks CI when + # Google changes the live spec — that breakage is how we get notified + # to review the change. expected_statuses = [ "in_progress", "requires_action", @@ -187,6 +190,7 @@ class TestResponseCompliance: "failed", "cancelled", "incomplete", + "budget_exceeded", ] assert status_prop["enum"] == expected_statuses print(f"✓ Status enum values: {expected_statuses}") diff --git a/tests/test_litellm/litellm_core_utils/test_token_counter.py b/tests/test_litellm/litellm_core_utils/test_token_counter.py index 3aa5f012467..324bace0e96 100644 --- a/tests/test_litellm/litellm_core_utils/test_token_counter.py +++ b/tests/test_litellm/litellm_core_utils/test_token_counter.py @@ -437,13 +437,38 @@ def test_gpt_4o_token_counter(): @pytest.mark.parametrize( "img_url", [ - "https://blog.purpureus.net/assets/blog/personal_key_rotation/simplified-asset-graph.jpg", + "https://example.com/test-image.png", "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAL0AAAC9CAMAAADRCYwCAAAAh1BMVEX///8AAAD8/Pz5+fkEBAT39/cJCQn09PRNTU3y8vIMDAwzMzPe3t7v7+8QEBCOjo7FxcXR0dHn5+elpaWGhoYYGBivr686OjocHBy0tLQtLS1TU1PY2Ni6urpaWlpERER3d3ecnJxoaGiUlJRiYmIlJSU4ODhBQUFycnKAgIDBwcFnZ2chISE7EjuwAAAI/UlEQVR4nO1caXfiOgz1bhJIyAJhX1JoSzv8/9/3LNlpYd4rhX6o4/N8Z2lKM2cURZau5JsQEhERERERERERERERERERERHx/wBjhDPC3OGN8+Cc5JeMuheaETSdO8vZFyCScHtmz2CsktoeMn7rLM1u3h0PMAEhyYX7v/Q9wQvoGdB0hlbzm45lEq/wd6y6G9aezvBk9AXwp1r3LHJIRsh6s2maxaJpmvqgvkC7WFS3loUnaFJtKRVUCEoV/RpCnHRvAsesVQ1hw+vd7Mpo+424tLs72NplkvQgcdrsvXkW/zJWqH/fA0FT84M/xnQJt4to3+ZLuanbM6X5lfXKHosO9COgREqpCR5i86pf2zPS7j9tTj+9nO7bQz3+xGEyGW9zqgQ1tyQ/VsxEDvce/4dcUPNb5OD9yXvR4Z2QisuP0xiGWPnemgugU5q/troHhGEjIF5sTOyW648aC0TssuaaCEsYEIkGzjWXOp3A0vVsf6kgRyqaDk+T7DIVWrb58b2tT5xpUucKwodOD/5LbrZC1ws6YSaBZJ/8xlh+XZSYXaMJ2ezNqjB3IPXuehPcx2U6b4t1dS/xNdFzguUt8ie7arnPeyCZroxLHzGgGdqVcspwafizPWEXBee+9G1OaufGdvNng/9C+gwgZ3PH3r87G6zXTZ5D5De2G2DeFoANXfbACkT+fxBQ22YFsTTJF9hjFVO6VbqxZXko4WJ8s52P4PnuxO5KRzu0/hlix1ySt8iXjgaQ+4IHPA9nVzNkdduM9LFT/Aacj4FtKrHA7iAw602Vnht6R8Vq1IOS+wNMKLYqayAYfRuufQPGeGb7sZogQQoLZrGPgZ6KoYn70Iw30O92BNEDpvwouCFn6wH2uS+EhRb3WF/HObZk3HuxfRQM3Y/Of/VH0n4MKNHZDiZvO9+m/ABALfkOcuar/7nOo7B95ACGVAFaz4jMiJwJhdaHBkySmzlGTu82gr6FSTik2kJvLnY9nOd/D90qcH268m3I/cgI1xg1maE5CuZYaWLH+UHANCIck0yt7Mx5zBm5vVHXHwChsZ35kKqUpmo5Svq5/fzfAI5g2vDtFPYo1HiEA85QrDeGm9g//LG7K0scO3sdpj2CBDgCa+0OFs0bkvVgnnM/QBDwllOMm+cN7vMSHlB7Uu4haHKaTwgGkv8tlK+hP8fzmFuK/RQTpaLPWvbd58yWIo66HHM0OsPoPhVqmtaEVL7N+wYcTLTbb0DLdgp23Eyy2VYJ2N7bkLFAAibtoLPe5sLt6Oa2bvU+zyeMa8wrixO0gRTn9tO9NCSThTLGqcqtsDvphlfmx/cPBZVvw24jg1LE2lPuEo35Mhi58U0I/Ga8n5w+NS8i34MAQLos5B1u0xL1ZvCVYVRw/Fs2q53KLaXJMWwOZZ/4MPYV19bAHmgGDKB6f01xoeJKFbl63q9J34KdaVNPJWztQyRkzA3KNs1AdAEDowMxh10emXTCx75CkurtbY/ZpdNDGdsn2UcHKHsQ8Ai3WZi48IfkvtjOhsLpuIRSKZTX9FA4o+0d6o/zOWqQzVJMynL9NsxhSJOaourq6nBVQBueMSyubsX2xHrmuABZN2Ns9jr5nwLFlLF/2R6atjW/67Yd11YQ1Z+kA9Zk9dPTM/o6dVo6HHVgC0JR8oUfmI93T9u3gvTG94bAH02Y5xeqRcjuwnKCK6Q2+ajl8KXJ3GSh22P3Zfx6S+n008ROhJn+JRIUVu6o7OXl8w1SeyhuqNDwNI7SjbK08QrqPxS95jy4G7nCXVq6G3HNu0LtK5J0e226CfC005WKK9sVvfxI0eUbcnzutfhWe3rpZHM0nZ/ny/N8tanKYlQ6VEW5Xuym8yV1zZX58vwGhZp/5tFfhybZabdbrQYOs8F+xEhmPsb0/nki6kIyVvzZzUASiOrTfF+Sj9bXC7DoJxeiV8tjQL6loSd0yCx7YyB6rPdLx31U2qCG3F/oXIuDuqd6LFO+4DNIJuxFZqSsU0ea88avovFnWKRYFYRQDfCfcGaBCLn4M4A1ntJ5E57vicwqq2enaZEF5nokCYu9TbKqCC5yCDfL+GhLxT4w4xEJs+anqgou8DOY2q8FMryjb2MehC1dRJ9s4g9NXeTwPkWON4RH+FhIe0AWR/S9ekvQ+t70XHeimGF78LzuU7d7PwrswdIG2VpgF8C53qVQsTDtBJc4CdnkQPbnZY9mbPdDFra3PCXBBQ5QBn2aQqtyhvlyYM4Hb2/mdhsxCUen04GZVvIJZw5PAamMOmjzq8Q+dzAKLXDQ3RUZItWsg4t7W2DP+JDrJDymoMH7E5zQtuEpG03GTIjGCW3LQqOYEsXgFc78x76NeRwY6SNM+IfQoh6myJKRBIcLYxZcwscJ/gI2isTBty2Po9IkYzP0/SS4hGlxRjFAG5z1Jt1LckiB57yWvo35EaolbvA+6fBa24xodL2YjsPpTnj3JgJOqhcgOeLVsYYwoK0wjY+m1D3rGc40CukkaHnkEjarlXrF1B9M6ECQ6Ow0V7R7N4G3LfOHAXtymoyXOb4QhaYHJ/gNBJUkxclpSs7DNcgWWDDmM7Ke5MJpGuioe7w5EOvfTunUKRzOh7G2ylL+6ynHrD54oQO3//cN3yVO+5qMVsPZq0CZIOx4TlcJ8+Vz7V5waL+7WekzUpRFMTnnTlSCq3X5usi8qmIleW/rit1+oQZn1WGSU/sKBYEqMNh1mBOc6PhK8yCfKHdUNQk8o/G19ZPTs5MYfai+DLs5vmee37zEyyH48WW3XA6Xw6+Az8lMhci7N/KleToo7PtTKm+RA887Kqc6E9dyqL/QPTugzMHLbLZtJKqKLFfzVWRNJ63c+95uWT/F7R0U5dDVvuS409AJXhJvD0EwWaWdW8UN11u/7+umaYjT8mJtzZwP/MD4r57fihiHlC5fylHfaqnJdro+Dr7DajvO+vi2EwyD70s8nCH71nzIO1l5Zl+v1DMCb5ebvCMkGHvobXy/hPumGLyX0218/3RyD1GRLOuf9u/OGQyDmto32yMiIiIiIiIiIiIiIiIiIiIiIiIiIiIiIiIiIv7GP8YjWPR/czH2AAAAAElFTkSuQmCC", ], ) -def test_img_url_token_counter(img_url): +def test_img_url_token_counter(img_url, monkeypatch): + """ + Verify get_image_dimensions returns valid (width, height) for both an + HTTPS URL and a base64 data URI. The HTTPS branch is exercised with a + mocked HTTP fetch so the test is hermetic - it can't break when a + third-party image URL goes away. + """ + import base64 from litellm.litellm_core_utils.token_counter import get_image_dimensions + # Minimal valid 1x1 PNG, served by the mocked safe_get for the URL case. + _tiny_png = base64.b64decode( + "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNkYAAAAAYAAjCB0C8AAAAASUVORK5CYII=" + ) + + if img_url.startswith(("http://", "https://")): + + class _FakeResponse: + headers = {"Content-Length": str(len(_tiny_png))} + + def read(self): + return _tiny_png + + monkeypatch.setattr( + "litellm.litellm_core_utils.token_counter.safe_get", + lambda client, url, **kw: _FakeResponse(), + ) + width, height = get_image_dimensions(data=img_url) print(width, height) From 1480ec698bc7eb5fa34aea8dda781d6bbc00815d Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Tue, 19 May 2026 15:10:37 -0700 Subject: [PATCH 3/7] chore(ci): bump versions (#28287) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * bump: version 0.4.72 → 0.4.73 * bump: version 1.86.0 → 1.87.0 * uv lock --- litellm-proxy-extras/pyproject.toml | 4 ++-- pyproject.toml | 6 +++--- uv.lock | 4 ++-- 3 files changed, 7 insertions(+), 7 deletions(-) diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml index ee27cc3585e..0654f17ec68 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.72" +version = "0.4.73" 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.72" +version = "0.4.73" version_files = [ "pyproject.toml:^version", "../pyproject.toml:litellm-proxy-extras==", diff --git a/pyproject.toml b/pyproject.toml index f63770105dc..70681c4ed6c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm" -version = "1.86.0" +version = "1.87.0" description = "Library to easily interface with LLM API providers" readme = "README.md" requires-python = ">=3.10, <3.14" @@ -56,7 +56,7 @@ proxy = [ "azure-identity==1.25.2", "azure-storage-blob==12.28.0", "mcp==1.26.0", - "litellm-proxy-extras==0.4.72", + "litellm-proxy-extras==0.4.73", "litellm-enterprise==0.1.41", "RestrictedPython==8.1", "rich==13.9.4", @@ -251,7 +251,7 @@ source-exclude = [ profile = "black" [tool.commitizen] -version = "1.86.0" +version = "1.87.0" version_files = [ "pyproject.toml:^version", ] diff --git a/uv.lock b/uv.lock index f3eaf6ca88c..cafb6664958 100644 --- a/uv.lock +++ b/uv.lock @@ -3189,7 +3189,7 @@ wheels = [ [[package]] name = "litellm" -version = "1.86.0" +version = "1.87.0" source = { editable = "." } dependencies = [ { name = "aiohttp" }, @@ -3539,7 +3539,7 @@ source = { editable = "enterprise" } [[package]] name = "litellm-proxy-extras" -version = "0.4.72" +version = "0.4.73" source = { editable = "litellm-proxy-extras" } [[package]] From cde4b1a94d1272e91b5cf72c717b677b1bbf86a1 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Tue, 19 May 2026 15:31:25 -0700 Subject: [PATCH 4/7] feat: propagate team_id and team_alias to all child OTEL spans (#28273) - Add `_set_team_attributes_on_span` helper to stamp team_id/team_alias onto any span, ensuring these attributes are not limited to the root litellm_request span - Add `_set_team_attributes_from_kwargs` helper to extract team metadata from the standard_logging_object in kwargs and apply them to a span - Apply team attributes to raw request spans via `_maybe_log_raw_request` so downstream consumers can filter traces by team without needing the root span - Apply team attributes to guardrail spans so guardrail activity can be correlated to teams in tracing backends - Apply team attributes to exception logging spans to preserve team context during failure paths - Add comprehensive unit tests covering all new helpers, including edge cases where metadata or standard_logging_object is absent Co-authored-by: Yassin Kortam --- litellm/integrations/opentelemetry.py | 82 +++++ .../integrations/test_opentelemetry.py | 143 +++++++++ .../test_otel_team_attributes_matrix.py | 285 ++++++++++++++++++ 3 files changed, 510 insertions(+) create mode 100644 tests/test_litellm/integrations/test_otel_team_attributes_matrix.py diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index ced15a01660..a70574952b8 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -673,6 +673,15 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): if parent_otel_span is not None: parent_otel_span.set_status(Status(StatusCode.ERROR)) + # Stamp team attributes onto the SERVER (root) span too, so the + # trace root is team-filterable on the failure path like the + # child exception span below. + self._set_team_attributes_on_span( + span=parent_otel_span, + team_id=user_api_key_dict.team_id, + team_alias=user_api_key_dict.team_alias, + ) + # Stamp structured error attrs on the SERVER span itself; the # failure path otherwise only sets its status (_handle_failure # records on the litellm_request child span). Inline import: @@ -709,6 +718,11 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): key="exception", value=str(original_exception), ) + self._set_team_attributes_on_span( + span=exception_logging_span, + team_id=user_api_key_dict.team_id, + team_alias=user_api_key_dict.team_alias, + ) exception_logging_span.set_status(Status(StatusCode.ERROR)) exception_logging_span.end(end_time=self._to_ns(datetime.now())) @@ -1012,6 +1026,10 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): ): parent_span.end(end_time=self._to_ns(end_time)) + # Stamp team attributes onto the SERVER (root) span before it is + # closed, so the trace root carries them like every child span. + self._set_team_attributes_on_proxy_span_from_kwargs(kwargs) + # close the proxy span explicitly from kwargs metadata # after all child spans (litellm_request, guardrail, raw_request) # have been fully recorded and exported. @@ -1070,8 +1088,70 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): ) raw_span.set_status(Status(StatusCode.OK)) self.set_raw_request_attributes(raw_span, kwargs, response_obj) + self._set_team_attributes_from_kwargs(raw_span, kwargs) raw_span.end(end_time=self._to_ns(end_time)) + def _set_team_attributes_on_span( + self, + span: Span, + team_id: Optional[str], + team_alias: Optional[str], + ) -> None: + """Stamp team_id / team_alias onto a span so every child span of a + litellm_request trace carries them, not just the root span. + + Empty strings are treated as absent: a request made with the master + key or a team-less virtual key carries ``user_api_key_team_id=""`` + in ``standard_logging_object.metadata``; propagating that to every + span only adds noise that makes traces look mis-instrumented. + """ + if team_id: + self.safe_set_attribute( + span=span, + key="metadata.user_api_key_team_id", + value=team_id, + ) + if team_alias: + self.safe_set_attribute( + span=span, + key="metadata.user_api_key_team_alias", + value=team_alias, + ) + + def _set_team_attributes_from_kwargs(self, span: Span, kwargs: dict) -> None: + """Pull team_id / team_alias from the standard logging metadata in kwargs and stamp them onto span.""" + std_log = kwargs.get("standard_logging_object") + md: dict = {} + if isinstance(std_log, dict): + md = std_log.get("metadata") or {} + elif std_log is not None: + md = getattr(std_log, "metadata", None) or {} + self._set_team_attributes_on_span( + span=span, + team_id=md.get("user_api_key_team_id"), + team_alias=md.get("user_api_key_team_alias"), + ) + + def _set_team_attributes_on_proxy_span_from_kwargs(self, kwargs: dict) -> None: + """Stamp team attributes onto the proxy SERVER (root) span so the + trace root is filterable by team, not just its children. The root + span is created in auth before the team is resolved and is + otherwise only closed (never re-attributed) on the success path. + + Guarded to the LiteLLM-created proxy span (by name + recording) so + externally provided parent spans are never mutated. + """ + litellm_params = kwargs.get("litellm_params") or {} + metadata = litellm_params.get("metadata") or {} + proxy_span = metadata.get("litellm_parent_otel_span") + if ( + proxy_span is not None + and getattr(proxy_span, "name", None) == LITELLM_PROXY_REQUEST_SPAN_NAME + and hasattr(proxy_span, "is_recording") + and proxy_span.is_recording() + ): + self._set_team_attributes_from_kwargs(proxy_span, kwargs) + def _record_metrics(self, kwargs, response_obj, start_time, end_time): duration_s = (end_time - start_time).total_seconds() params = kwargs.get("litellm_params") or {} @@ -1537,6 +1617,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): value=guardrail_information.get("guardrail_response"), ) + self._set_team_attributes_from_kwargs(guardrail_span, kwargs) + guardrail_span.end(end_time=self._to_ns(end_time_datetime)) def _handle_failure(self, kwargs, response_obj, start_time, end_time): diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py index 87298b5a7a3..6de855262bd 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/test_litellm/integrations/test_opentelemetry.py @@ -88,6 +88,149 @@ class TestOpenTelemetryGuardrails(unittest.TestCase): otel.tracer.start_span.assert_not_called() +class TestOpenTelemetryTeamAttributesOnChildSpans(unittest.TestCase): + """team_id / team_alias must land on every child span of a + litellm_request trace, not only the root litellm_request span.""" + + def _slo_metadata(self): + return { + "user_api_key_team_id": "team-123", + "user_api_key_team_alias": "my-team", + } + + @patch("litellm.integrations.opentelemetry.datetime") + def test_guardrail_span_has_team_attributes(self, mock_datetime): + otel = OpenTelemetry() + otel.tracer = MagicMock() + mock_span = MagicMock() + otel.tracer.start_span.return_value = mock_span + + guardrail_info = { + "guardrail_name": "test_guardrail", + "guardrail_mode": "input", + "guardrail_response": "filtered_content", + "start_time": 1609459200.0, + "end_time": 1609459201.0, + } + kwargs = { + "standard_logging_object": { + "guardrail_information": [guardrail_info], + "metadata": self._slo_metadata(), + } + } + + otel._create_guardrail_span(kwargs=kwargs, context=None) + + mock_span.set_attribute.assert_any_call( + "metadata.user_api_key_team_id", "team-123" + ) + mock_span.set_attribute.assert_any_call( + "metadata.user_api_key_team_alias", "my-team" + ) + + @patch.dict(os.environ, {"OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT": ""}) + @patch("litellm.turn_off_message_logging", False) + def test_raw_request_span_has_team_attributes(self): + otel = OpenTelemetry() + otel.message_logging = True + + mock_tracer = MagicMock() + mock_span = MagicMock() + mock_tracer.start_span.return_value = mock_span + otel.get_tracer_to_use_for_request = MagicMock(return_value=mock_tracer) + otel.set_raw_request_attributes = MagicMock() + otel._to_ns = MagicMock(return_value=1234567890) + + kwargs = { + "litellm_params": {"metadata": {}}, + "standard_logging_object": {"metadata": self._slo_metadata()}, + } + otel._maybe_log_raw_request( + kwargs, {}, datetime.now(), datetime.now(), MagicMock() + ) + + mock_span.set_attribute.assert_any_call( + "metadata.user_api_key_team_id", "team-123" + ) + mock_span.set_attribute.assert_any_call( + "metadata.user_api_key_team_alias", "my-team" + ) + + def test_helper_skips_when_team_values_missing(self): + otel = OpenTelemetry() + mock_span = MagicMock() + + otel._set_team_attributes_on_span(span=mock_span, team_id=None, team_alias=None) + + mock_span.set_attribute.assert_not_called() + + def test_helper_skips_when_team_values_are_empty_strings(self): + """A master-key / team-less request carries user_api_key_team_id='' + in metadata. Propagating '' to every span is noise that makes + traces look mis-instrumented; treat empty as absent.""" + otel = OpenTelemetry() + mock_span = MagicMock() + + otel._set_team_attributes_on_span(span=mock_span, team_id="", team_alias="") + + mock_span.set_attribute.assert_not_called() + + def test_helper_reads_metadata_from_kwargs(self): + otel = OpenTelemetry() + mock_span = MagicMock() + + otel._set_team_attributes_from_kwargs( + mock_span, + {"standard_logging_object": {"metadata": self._slo_metadata()}}, + ) + + mock_span.set_attribute.assert_any_call( + "metadata.user_api_key_team_id", "team-123" + ) + mock_span.set_attribute.assert_any_call( + "metadata.user_api_key_team_alias", "my-team" + ) + + def test_helper_handles_missing_standard_logging_object(self): + otel = OpenTelemetry() + mock_span = MagicMock() + + otel._set_team_attributes_from_kwargs(mock_span, {}) + + mock_span.set_attribute.assert_not_called() + + def test_failure_hook_exception_span_has_team_attributes(self): + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + + otel = OpenTelemetry() + otel.tracer = tracer + server_span = tracer.start_span("Received Proxy Server Request") + + user_api_key_dict = MagicMock() + user_api_key_dict.parent_otel_span = server_span + user_api_key_dict.team_id = "team-123" + user_api_key_dict.team_alias = "my-team" + + asyncio.run( + otel.async_post_call_failure_hook( + request_data={}, + original_exception=ValueError("boom"), + user_api_key_dict=user_api_key_dict, + traceback_str="trace", + ) + ) + + finished = {s.name: s for s in exporter.get_finished_spans()} + exception_span = finished["Failed Proxy Server Request"] + assert exception_span.attributes["metadata.user_api_key_team_id"] == "team-123" + assert ( + exception_span.attributes["metadata.user_api_key_team_alias"] == "my-team" + ) + + class TestOpenTelemetryCostBreakdown(unittest.TestCase): def test_cost_breakdown_emitted_to_otel_span(self): """ diff --git a/tests/test_litellm/integrations/test_otel_team_attributes_matrix.py b/tests/test_litellm/integrations/test_otel_team_attributes_matrix.py new file mode 100644 index 00000000000..1ce55fa7a58 --- /dev/null +++ b/tests/test_litellm/integrations/test_otel_team_attributes_matrix.py @@ -0,0 +1,285 @@ +""" +Matrix test: team_id / team_alias must land on EVERY span of a proxy +request trace, for a representative set of endpoints x HTTP outcomes. + +Endpoints + - /v1/chat/completions (OpenAI-format LLM path) + - /v1/messages (Anthropic-format LLM path) + - /team/info (management/admin path) + +Outcomes + - 2xx success + - 3xx redirect (LLM endpoints never 3xx -> N/A; admin too) + - 4xx client error (auth / validation failure) + - 5xx server error (upstream / DB failure) + +Strategy + These assertions exercise the real OpenTelemetry callback the proxy + invokes for each path, with a SERVER parent span (as + ``user_api_key_auth`` creates) and an in-memory exporter. Each cell + drives the path, then asserts team attributes on every span that path + actually emits. + + - success path -> ``log_success_event`` -> litellm_request + + raw_gen_ai_request + guardrail child spans. + - failure path -> ``async_post_call_failure_hook`` -> Failed Proxy + Server Request exception child span. + + Admin endpoints do not run the LLM success callback, so their only + trace surface is the SERVER span (success) or the exception child span + (failure) -- the cells below assert exactly that. +""" + +import asyncio +import os +import sys +import unittest +from datetime import datetime +from unittest.mock import MagicMock + +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( + InMemorySpanExporter, +) + +sys.path.insert(0, os.path.abspath("../..")) + +from litellm.integrations.opentelemetry import ( + LITELLM_PROXY_REQUEST_SPAN_NAME, + OpenTelemetry, +) + +TEAM_ID = "team-123" +TEAM_ALIAS = "my-team" +TEAM_ID_ATTR = "metadata.user_api_key_team_id" +TEAM_ALIAS_ATTR = "metadata.user_api_key_team_alias" + + +def _make_otel(): + """OTel callback whose every span lands in an in-memory exporter.""" + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + + otel = OpenTelemetry() + otel.tracer = provider.get_tracer(__name__) + # raw_gen_ai_request sub-span is gated on message logging. + otel.message_logging = True + return otel, exporter + + +def _server_span(otel): + """Mirror the SERVER span user_api_key_auth opens per request.""" + return otel.create_litellm_proxy_request_started_span( + start_time=datetime.now(), headers={} + ) + + +def _slo(call_type, with_guardrail=False): + """standard_logging_object the proxy attaches, carrying team metadata.""" + md = { + "user_api_key_team_id": TEAM_ID, + "user_api_key_team_alias": TEAM_ALIAS, + } + slo = {"metadata": md, "call_type": call_type} + if with_guardrail: + slo["guardrail_information"] = [ + { + "guardrail_name": "test_guardrail", + "guardrail_mode": "input", + "guardrail_response": "ok", + "start_time": 1609459200.0, + "end_time": 1609459201.0, + } + ] + return slo + + +def _success_kwargs(call_type, server_span, with_guardrail=True): + """kwargs the success callback receives for an LLM proxy request.""" + return { + "model": "gpt-4.1-mini", + "litellm_call_id": "call-abc", + "call_type": call_type, + "litellm_params": { + "metadata": { + "litellm_parent_otel_span": server_span, + "user_api_key_team_id": TEAM_ID, + "user_api_key_team_alias": TEAM_ALIAS, + } + }, + "standard_logging_object": _slo(call_type, with_guardrail=with_guardrail), + "messages": [{"role": "user", "content": "hi"}], + } + + +def _team_user_api_key_dict(server_span): + d = MagicMock() + d.parent_otel_span = server_span + d.team_id = TEAM_ID + d.team_alias = TEAM_ALIAS + return d + + +def _spans_by_name(exporter): + return {s.name: s for s in exporter.get_finished_spans()} + + +def _assert_team_attrs(span, where): + assert span.attributes.get(TEAM_ID_ATTR) == TEAM_ID, ( + f"{where}: missing/blank {TEAM_ID_ATTR} " + f"(got {span.attributes.get(TEAM_ID_ATTR)!r})" + ) + assert span.attributes.get(TEAM_ALIAS_ATTR) == TEAM_ALIAS, ( + f"{where}: missing/blank {TEAM_ALIAS_ATTR} " + f"(got {span.attributes.get(TEAM_ALIAS_ATTR)!r})" + ) + + +class _Boom(Exception): + """Upstream/DB style 5xx.""" + + status_code = 500 + + +class _ClientErr(Exception): + """Auth/validation style 4xx.""" + + status_code = 401 + + +# --------------------------------------------------------------------------- +# LLM success cells: litellm_request + raw_gen_ai_request + guardrail spans +# --------------------------------------------------------------------------- +class TestLLMSuccessCells(unittest.TestCase): + def _run_success(self, call_type): + otel, exporter = _make_otel() + server_span = _server_span(otel) + kwargs = _success_kwargs(call_type, server_span) + now = datetime.now() + otel.log_success_event(kwargs, {"id": "resp-1"}, now, now) + return _spans_by_name(exporter) + + def test_chat_completions_2xx(self): + spans = self._run_success("completion") + for name in ( + LITELLM_PROXY_REQUEST_SPAN_NAME, + "litellm_request", + "raw_gen_ai_request", + "guardrail", + ): + assert name in spans, f"chat/completions 2xx: missing span {name}" + _assert_team_attrs(spans[name], f"chat/completions 2xx [{name}]") + + def test_v1_messages_2xx(self): + spans = self._run_success("anthropic_messages") + for name in ( + LITELLM_PROXY_REQUEST_SPAN_NAME, + "litellm_request", + "raw_gen_ai_request", + "guardrail", + ): + assert name in spans, f"v1/messages 2xx: missing span {name}" + _assert_team_attrs(spans[name], f"v1/messages 2xx [{name}]") + + +# --------------------------------------------------------------------------- +# LLM failure cells: Failed Proxy Server Request exception child span +# --------------------------------------------------------------------------- +class TestLLMFailureCells(unittest.TestCase): + def _run_failure(self, exc): + """Drive the failure hook, then close the SERVER span (the proxy + closes it after the hook in real flow) so both the exception child + span and the SERVER root span are asserted.""" + otel, exporter = _make_otel() + server_span = _server_span(otel) + asyncio.run( + otel.async_post_call_failure_hook( + request_data={}, + original_exception=exc, + user_api_key_dict=_team_user_api_key_dict(server_span), + traceback_str="tb", + ) + ) + server_span.end() + return _spans_by_name(exporter) + + def _assert_all(self, spans, where): + for name in ("Failed Proxy Server Request", LITELLM_PROXY_REQUEST_SPAN_NAME): + assert name in spans, f"{where}: missing span {name}" + _assert_team_attrs(spans[name], f"{where} [{name}]") + + def test_chat_completions_4xx(self): + self._assert_all( + self._run_failure(_ClientErr("bad key")), "chat/completions 4xx" + ) + + def test_chat_completions_5xx(self): + self._assert_all( + self._run_failure(_Boom("upstream blew up")), "chat/completions 5xx" + ) + + def test_v1_messages_4xx(self): + self._assert_all( + self._run_failure(_ClientErr("bad anthropic key")), "v1/messages 4xx" + ) + + def test_v1_messages_5xx(self): + self._assert_all( + self._run_failure(_Boom("anthropic upstream timeout")), "v1/messages 5xx" + ) + + +# --------------------------------------------------------------------------- +# Admin /team/info cells. +# 2xx: admin path never runs the LLM success callback -> its only trace +# surface is the SERVER span; no child spans are emitted. +# 3xx: management endpoints do not redirect -> N/A (documented, no run). +# 4xx/5xx: proxy_logging post_call_failure_hook -> exception child span. +# --------------------------------------------------------------------------- +class TestAdminTeamInfoCells(unittest.TestCase): + def _run_admin_failure(self, exc): + otel, exporter = _make_otel() + server_span = _server_span(otel) + asyncio.run( + otel.async_post_call_failure_hook( + request_data={}, + original_exception=exc, + user_api_key_dict=_team_user_api_key_dict(server_span), + traceback_str="tb", + ) + ) + server_span.end() + return _spans_by_name(exporter) + + def test_team_info_4xx(self): + spans = self._run_admin_failure(_ClientErr("team not found")) + for name in ("Failed Proxy Server Request", LITELLM_PROXY_REQUEST_SPAN_NAME): + _assert_team_attrs(spans[name], f"/team/info 4xx [{name}]") + + def test_team_info_5xx(self): + spans = self._run_admin_failure(_Boom("db connection lost")) + for name in ("Failed Proxy Server Request", LITELLM_PROXY_REQUEST_SPAN_NAME): + _assert_team_attrs(spans[name], f"/team/info 5xx [{name}]") + + def test_team_info_2xx_only_server_span_no_orphan_children(self): + """Admin success path emits no LLM child spans; nothing to stamp + beyond the SERVER span. This pins that contract so a future + regression that starts emitting child spans here without team + attrs is caught.""" + otel, exporter = _make_otel() + server_span = _server_span(otel) + server_span.end() + spans = _spans_by_name(exporter) + assert set(spans) == { + LITELLM_PROXY_REQUEST_SPAN_NAME + }, f"/team/info 2xx: unexpected child spans {set(spans)}" + + def test_team_info_3xx_not_applicable(self): + """Management endpoints return JSON, never a 3xx redirect.""" + self.skipTest("/team/info has no 3xx redirect path (N/A)") + + +if __name__ == "__main__": + unittest.main() From 3c3d131f017af22e09a1b32d033238a3933085dc Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 20 May 2026 04:20:54 +0530 Subject: [PATCH 5/7] Day 0 support : Gemini 3.5 Flash (#28268) * Add day 0 support for gemini 3.5 flash * Fix pricing * Fix greptile review * Fix failing test * Fix tests * Fix: revert tool removing logic * fix greptile and test --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .../prompt_templates/factory.py | 40 +++- .../llms/vertex_ai/gemini/transformation.py | 4 +- .../vertex_and_google_ai_studio_gemini.py | 128 ++++++++++++- ...odel_prices_and_context_window_backup.json | 178 ++++++++++++++++++ litellm/types/llms/vertex_ai.py | 21 ++- model_prices_and_context_window.json | 178 ++++++++++++++++++ ...test_vertex_and_google_ai_studio_gemini.py | 32 ++++ .../llms/vertex_ai/test_vertex.py | 40 ++-- 8 files changed, 576 insertions(+), 45 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 79d527d1eb8..d6c011990c1 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1233,6 +1233,7 @@ def infer_protocol_value( def _gemini_tool_call_invoke_helper( function_call_params: ChatCompletionToolCallFunctionChunk, + tool_call_id: Optional[str] = None, ) -> Optional[VertexFunctionCall]: name = function_call_params.get("name", "") or "" arguments = function_call_params.get("arguments", "") @@ -1248,6 +1249,10 @@ def _gemini_tool_call_invoke_helper( name=name, args=arguments_dict, ) + if tool_call_id: + clean_id = tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0] + if clean_id: + function_call["id"] = clean_id return function_call @@ -1384,12 +1389,23 @@ def convert_to_gemini_tool_call_invoke( tool_calls = message.get("tool_calls", None) function_call = message.get("function_call", None) + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, + ) + + forward_tool_call_id = bool( + model and VertexGeminiConfig._is_gemini_3_or_newer(model) + ) + if tool_calls is not None: for idx, tool in enumerate(tool_calls): if "function" in tool: gemini_function_call: Optional[VertexFunctionCall] = ( _gemini_tool_call_invoke_helper( - function_call_params=tool["function"] + function_call_params=tool["function"], + tool_call_id=( + tool.get("id") if forward_tool_call_id else None + ), ) ) if gemini_function_call is not None: @@ -1429,10 +1445,6 @@ def convert_to_gemini_tool_call_invoke( thought_signature = provider_fields.get("thought_signature") # If no signature found and model is gemini-3, use dummy signature - from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( - VertexGeminiConfig, - ) - if ( not thought_signature and model @@ -1462,6 +1474,7 @@ def convert_to_gemini_tool_call_invoke( def convert_to_gemini_tool_call_result( # noqa: PLR0915 message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage], last_message_with_tool_calls: Optional[dict], + model: Optional[str] = None, ) -> Union[VertexPartType, List[VertexPartType]]: """ OpenAI message with a tool result looks like: @@ -1602,6 +1615,21 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915 ): name = tool.get("function", {}).get("name", "") + # Echo the OpenAI tool_call_id on functionResponse (strip thought-signature suffix). + # Only Gemini 3+ accepts (and returns) an `id` on function_response parts; + # older Gemini models reject the field with a 400. + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, + ) + + gemini_call_id: Optional[str] = None + if model and VertexGeminiConfig._is_gemini_3_or_newer(model): + raw_tool_call_id = message.get("tool_call_id") + if raw_tool_call_id and isinstance(raw_tool_call_id, str): + stripped_id = raw_tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0] + if stripped_id: + gemini_call_id = stripped_id + if not name: raise Exception( "Missing corresponding tool call for tool response message. Received - message={}, last_message_with_tool_calls={}".format( @@ -1632,6 +1660,8 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915 name=name, response=response_data, # type: ignore ) + if gemini_call_id: + _function_response["id"] = gemini_call_id # Create part with function_response, and optionally inline_data for images (Computer Use) _part: VertexPartType = {"function_response": _function_response} diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index f56992a2502..ea57339e35f 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -1042,7 +1042,9 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915 and messages[msg_i]["role"] in tool_call_message_roles ): _part = convert_to_gemini_tool_call_result( - messages[msg_i], last_message_with_tool_calls # type: ignore + messages[msg_i], # type: ignore + last_message_with_tool_calls, # type: ignore + model=model, ) msg_i += 1 # Handle both single part and list of parts (for Computer Use with images) diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 49c1c335467..e99f69fcd3e 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -280,6 +280,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): - gemini-3-pro-preview - gemini-3-flash - gemini-3-flash-preview (Gemini 3 Flash) + - gemini-3.1-pro-preview, gemini-3.1-flash, gemini-3.1-flash-lite-preview + - gemini-3.5-flash - Any future Gemini 3.x models """ # Check for Gemini 3 models @@ -300,6 +302,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): supported_params = [ "temperature", "top_p", + "top_k", "max_tokens", "max_completion_tokens", "stream", @@ -363,6 +366,66 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): """ return Tools(googleSearch={}) + @staticmethod + def _search_tool_keys() -> set: + return { + VertexToolName.GOOGLE_SEARCH.value, + VertexToolName.GOOGLE_SEARCH_RETRIEVAL.value, + VertexToolName.ENTERPRISE_WEB_SEARCH.value, + VertexToolName.URL_CONTEXT.value, + "google_search", + "google_search_retrieval", + "enterprise_web_search", + "urlContext", + } + + @classmethod + def _drop_search_tools_mixed_with_functions(cls, optional_params: dict) -> None: + """ + Drop search tools from optional_params when mixed with function declarations + and include_server_side_tool_invocations is not enabled. + + Runs after map_openai_params merges tools and web_search_options so both + code paths (single _map_function call vs split tools + web_search_options) + get the same conflict resolution. + """ + if optional_params.get("include_server_side_tool_invocations"): + return + + tools = optional_params.get("tools") + if not isinstance(tools, list) or not tools: + return + + search_tool_keys = cls._search_tool_keys() + has_function_declarations = any( + isinstance(tool, dict) and tool.get("function_declarations") + for tool in tools + ) + if not has_function_declarations: + return + + has_search_tools = any( + isinstance(tool, dict) and any(key in tool for key in search_tool_keys) + for tool in tools + ) + if not has_search_tools: + return + + verbose_logger.warning( + "Vertex AI does not support mixing function declarations with " + "search tools (googleSearch, enterpriseWebSearch, urlContext, " + "googleSearchRetrieval) in the same request. Dropping search " + "tools and keeping function declarations. To use search tools, " + "send a request without function calling tools." + ) + optional_params["tools"] = [ + tool + for tool in tools + if not ( + isinstance(tool, dict) and any(key in tool for key in search_tool_keys) + ) + ] + def _map_service_tier_param(self, value: str, optional_params: dict) -> None: """ Map OpenAI service_tier (string) to Gemini serviceTier. @@ -884,9 +947,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): GeminiThinkingConfig with thinkingLevel and includeThoughts """ # Check if this is gemini-3-flash which supports MINIMAL thinking level - # Covers gemini-3-flash, gemini-3-flash-preview, gemini-3.1-flash, gemini-3.1-flash-lite-preview, etc. + # Covers gemini-3-flash, gemini-3-flash-preview, gemini-3.1-flash, gemini-3.1-flash-lite-preview, + # gemini-3.5-flash, and any future 3.x-flash variants. is_gemini3flash = model and ( - "gemini-3-flash" in model.lower() or "gemini-3.1-flash" in model.lower() + "flash" in model.lower() and "gemini-3" in model.lower() ) is_gemini31pro = model and ("gemini-3.1-pro-preview" in model.lower()) if reasoning_effort == "minimal": @@ -982,8 +1046,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): # Follow provider defaults unless explicitly opted into legacy behavior. if litellm.enable_gemini_default_thinking_level_low is True: is_gemini3flash = ( - "gemini-3-flash-preview" in model.lower() - or "gemini-3-flash" in model.lower() + "gemini-3" in model.lower() and "flash" in model.lower() ) params["thinkingLevel"] = ( "minimal" if is_gemini3flash else "low" @@ -1077,6 +1140,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): model: str, drop_params: bool, ) -> Dict: + gemini_sampling_params_warned: bool = False for param, value in non_default_params.items(): if param == "temperature": if VertexGeminiConfig._is_gemini_3_or_newer(model): @@ -1086,9 +1150,41 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): "can cause infinite loops, degraded reasoning performance, and failure on complex tasks. " "Strongly recommended to use temperature = 1.0 (default)." ) + if not gemini_sampling_params_warned: + verbose_logger.warning( + "DeprecationWarning: `temperature`, `top_p`, and `top_k` continue to " + f"function for Gemini 3+ ({model}) but are planned for removal in a " + "future release. Move sampling guidance into the `system` " + "instructions instead." + ) + gemini_sampling_params_warned = True optional_params["temperature"] = value elif param == "top_p": + if ( + VertexGeminiConfig._is_gemini_3_or_newer(model) + and not gemini_sampling_params_warned + ): + verbose_logger.warning( + "DeprecationWarning: `temperature`, `top_p`, and `top_k` continue to " + f"function for Gemini 3+ ({model}) but are planned for removal in a " + "future release. Move sampling guidance into the `system` " + "instructions instead." + ) + gemini_sampling_params_warned = True optional_params["top_p"] = value + elif param == "top_k": + if ( + VertexGeminiConfig._is_gemini_3_or_newer(model) + and not gemini_sampling_params_warned + ): + verbose_logger.warning( + "DeprecationWarning: `temperature`, `top_p`, and `top_k` continue to " + f"function for Gemini 3+ ({model}) but are planned for removal in a " + "future release. Move sampling guidance into the `system` " + "instructions instead." + ) + gemini_sampling_params_warned = True + optional_params["top_k"] = value elif ( param == "stream" and value is True ): # sending stream = False, can cause it to get passed unchecked and raise issues @@ -1139,11 +1235,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): if _tool_choice_value is not None: optional_params["tool_choice"] = _tool_choice_value elif param == "parallel_tool_calls": - if value is False and not ( - drop_params or litellm.drop_params - ): # if drop params is True, then we should just ignore this - self.validate_parallel_tool_calls(value, non_default_params) - else: + tools_list = non_default_params.get( + "tools", non_default_params.get("functions") + ) + num_tools = len(tools_list) if isinstance(tools_list, list) else 0 + # Gemini does not support parallel_tool_calls=False with multiple + # tools. Drop the param instead of failing — Responses API clients + # often send parallel_tool_calls=false by default. + if not (value is False and num_tools > 1): optional_params["parallel_tool_calls"] = value elif param == "seed": optional_params["seed"] = value @@ -1216,6 +1315,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): if "temperature" not in optional_params: optional_params["temperature"] = 1.0 + self._drop_search_tools_mixed_with_functions(optional_params) + return optional_params def get_mapped_special_auth_params(self) -> dict: @@ -1588,6 +1689,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): } # Extract thought signature if present thought_signature = part.get("thoughtSignature") + # Gemini 3.5+ returns a stable `id` per function call to enable + # strict response matching. Preserve it as the OpenAI + # tool_call_id so it can be echoed back unchanged. + gemini_call_id = part["functionCall"].get("id") if is_function_call is True: function_dict: Dict[str, Any] = dict(_function_chunk) @@ -1605,6 +1710,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): "function": _function_chunk, "index": cumulative_tool_call_idx, } + # Gemini 3.5+ returns a stable native `id`; prefer it over + # the synthetic call_ so the same value can be echoed + # back on the matching `functionResponse`. + if gemini_call_id: + _tool_response_chunk["id"] = gemini_call_id # Embed thought signature in ID for OpenAI client compatibility if thought_signature: _tool_response_chunk["provider_specific_fields"] = { # type: ignore diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index fb5bfa6cf4e..6d7c3eeb0b8 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -15611,6 +15611,64 @@ }, "web_search_billing_unit": "per_query" }, + "vertex_ai/gemini-3.5-flash": { + "cache_read_input_token_cost": 1.5e-07, + "input_cost_per_token": 1.5e-06, + "input_cost_per_audio_token": 1e-06, + "litellm_provider": "vertex_ai", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65535, + "max_pdf_size_mb": 30, + "max_tokens": 65535, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 9e-06, + "output_cost_per_token": 9e-06, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "input_cost_per_token_priority": 2.7e-06, + "input_cost_per_audio_token_priority": 1.8e-06, + "output_cost_per_token_priority": 1.62e-05, + "cache_read_input_token_cost_priority": 2.7e-07, + "supports_service_tier": true, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "vertex_ai/gemini-3.1-pro-preview": { "cache_read_input_token_cost": 2e-07, "cache_read_input_token_cost_above_200k_tokens": 4e-07, @@ -16988,6 +17046,67 @@ }, "web_search_billing_unit": "per_query" }, + "gemini/gemini-3.5-flash": { + "cache_read_input_token_cost": 1.5e-07, + "input_cost_per_audio_token": 1e-06, + "input_cost_per_token": 1.5e-06, + "litellm_provider": "gemini", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65535, + "max_pdf_size_mb": 30, + "max_tokens": 65535, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 9e-06, + "output_cost_per_token": 9e-06, + "rpm": 2000, + "source": "https://ai.google.dev/pricing/gemini-3", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_output": false, + "supports_audio_input": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "tpm": 800000, + "input_cost_per_token_priority": 2.7e-06, + "input_cost_per_audio_token_priority": 1.8e-06, + "output_cost_per_token_priority": 1.62e-05, + "cache_read_input_token_cost_priority": 2.7e-07, + "supports_service_tier": true, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "gemini/gemini-3.1-pro-preview": { "cache_read_input_token_cost": 2e-07, "cache_read_input_token_cost_above_200k_tokens": 4e-07, @@ -17173,6 +17292,65 @@ }, "web_search_billing_unit": "per_query" }, + "gemini-3.5-flash": { + "cache_read_input_token_cost": 1.5e-07, + "input_cost_per_audio_token": 1e-06, + "input_cost_per_token": 1.5e-06, + "litellm_provider": "vertex_ai-language-models", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65535, + "max_pdf_size_mb": 30, + "max_tokens": 65535, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 9e-06, + "output_cost_per_token": 9e-06, + "source": "https://ai.google.dev/pricing/gemini-3", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_output": false, + "supports_audio_input": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "input_cost_per_token_priority": 2.7e-06, + "input_cost_per_audio_token_priority": 1.8e-06, + "output_cost_per_token_priority": 1.62e-05, + "cache_read_input_token_cost_priority": 2.7e-07, + "supports_service_tier": true, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "gemini/gemini-2.5-pro-preview-tts": { "cache_read_input_token_cost": 1.25e-07, "cache_read_input_token_cost_above_200k_tokens": 2.5e-07, diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index 87bf11a9026..b357b64156d 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -14,13 +14,19 @@ from litellm.types.llms.openai import EmbeddingInput GeminiEmbeddingInput = Union[EmbeddingInput, List[List[str]]] -class FunctionResponse(TypedDict): - name: str +class FunctionResponse(TypedDict, total=False): + # `id` correlates this response with the originating `functionCall` part. + # Required by Gemini 3.5+ for strict function-calling response matching. + id: str + name: Required[str] response: Optional[dict] -class FunctionCall(TypedDict): - name: str +class FunctionCall(TypedDict, total=False): + # `id` is returned by Gemini 3.5+ to correlate the corresponding + # `functionResponse`. Older Gemini models omit this field. + id: str + name: Required[str] args: Optional[dict] @@ -45,8 +51,11 @@ class PartType(TypedDict, total=False): media_resolution: Literal["low", "medium", "high"] -class HttpxFunctionCall(TypedDict): - name: str +class HttpxFunctionCall(TypedDict, total=False): + # `id` is returned by Gemini 3.5+ to correlate the corresponding + # `functionResponse`. Older Gemini models omit this field. + id: str + name: Required[str] args: dict diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 94f0f1e78d3..e7a03bb0984 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -15645,6 +15645,64 @@ }, "web_search_billing_unit": "per_query" }, + "vertex_ai/gemini-3.5-flash": { + "cache_read_input_token_cost": 1.5e-07, + "input_cost_per_token": 1.5e-06, + "input_cost_per_audio_token": 1e-06, + "litellm_provider": "vertex_ai", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65535, + "max_pdf_size_mb": 30, + "max_tokens": 65535, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 9e-06, + "output_cost_per_token": 9e-06, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "input_cost_per_token_priority": 2.7e-06, + "input_cost_per_audio_token_priority": 1.8e-06, + "output_cost_per_token_priority": 1.62e-05, + "cache_read_input_token_cost_priority": 2.7e-07, + "supports_service_tier": true, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "vertex_ai/gemini-3.1-pro-preview": { "cache_read_input_token_cost": 2e-07, "cache_read_input_token_cost_above_200k_tokens": 4e-07, @@ -17022,6 +17080,67 @@ }, "web_search_billing_unit": "per_query" }, + "gemini/gemini-3.5-flash": { + "cache_read_input_token_cost": 1.5e-07, + "input_cost_per_audio_token": 1e-06, + "input_cost_per_token": 1.5e-06, + "litellm_provider": "gemini", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65535, + "max_pdf_size_mb": 30, + "max_tokens": 65535, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 9e-06, + "output_cost_per_token": 9e-06, + "rpm": 2000, + "source": "https://ai.google.dev/pricing/gemini-3", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_output": false, + "supports_audio_input": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "tpm": 800000, + "input_cost_per_token_priority": 2.7e-06, + "input_cost_per_audio_token_priority": 1.8e-06, + "output_cost_per_token_priority": 1.62e-05, + "cache_read_input_token_cost_priority": 2.7e-07, + "supports_service_tier": true, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "gemini/gemini-3.1-pro-preview": { "cache_read_input_token_cost": 2e-07, "cache_read_input_token_cost_above_200k_tokens": 4e-07, @@ -17207,6 +17326,65 @@ }, "web_search_billing_unit": "per_query" }, + "gemini-3.5-flash": { + "cache_read_input_token_cost": 1.5e-07, + "input_cost_per_audio_token": 1e-06, + "input_cost_per_token": 1.5e-06, + "litellm_provider": "vertex_ai-language-models", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65535, + "max_pdf_size_mb": 30, + "max_tokens": 65535, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 9e-06, + "output_cost_per_token": 9e-06, + "source": "https://ai.google.dev/pricing/gemini-3", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_output": false, + "supports_audio_input": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "input_cost_per_token_priority": 2.7e-06, + "input_cost_per_audio_token_priority": 1.8e-06, + "output_cost_per_token_priority": 1.62e-05, + "cache_read_input_token_cost_priority": 2.7e-07, + "supports_service_tier": true, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "gemini/gemini-2.5-pro-preview-tts": { "cache_read_input_token_cost": 1.25e-07, "cache_read_input_token_cost_above_200k_tokens": 2.5e-07, diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index 1e0ad04c3c2..db993245426 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -2959,6 +2959,38 @@ def test_vertex_ai_gemini3_tool_combination_no_drop(): assert len(tools) == 3 +def test_vertex_ai_mixed_tools_and_web_search_options_drops_search(): + """ + When function tools and web_search_options are sent separately (Codex-style), + search tools are dropped unless include_server_side_tool_invocations is set. + """ + v = VertexGeminiConfig() + optional_params: dict = {} + non_default_params = { + "tools": [ + { + "type": "function", + "function": {"name": "exec_command", "description": "Run a command"}, + } + ], + "web_search_options": {}, + } + + result = v.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model="gemini-3.5-flash", + drop_params=True, + ) + + assert not result.get("include_server_side_tool_invocations") + tool_keys = set() + for tool in result.get("tools", []): + tool_keys.update(tool.keys()) + assert "function_declarations" in tool_keys + assert "googleSearch" not in tool_keys + + def test_vertex_ai_openai_web_search_tool_transformation(): """ Test that OpenAI-style web_search and web_search_preview tools are transformed to googleSearch. diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex.py b/tests/test_litellm/llms/vertex_ai/test_vertex.py index be0e59e8b7d..ec73e5e42be 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex.py @@ -1517,39 +1517,31 @@ def test_vertex_parallel_tool_calls_true(): assert "tools" in optional_params -def test_vertex_parallel_tool_calls_false_multiple_tools_error(): +def test_vertex_parallel_tool_calls_false_multiple_tools_dropped(): """ - Test that parallel_tool_calls = False with multiple tools raises UnsupportedParamsError - when drop_params is False. + parallel_tool_calls=False with multiple tools is dropped for Gemini + (unsupported upstream). Request should succeed without the param. """ tools = [ {"type": "function", "function": {"name": "get_weather"}}, {"type": "function", "function": {"name": "get_time"}}, ] - with pytest.raises(litellm.utils.UnsupportedParamsError) as excinfo: - get_optional_params( - model="gemini-1.5-pro", - custom_llm_provider="vertex_ai", - tools=tools, - parallel_tool_calls=False, - ) - assert ( - "`parallel_tool_calls=False` is not supported by Gemini when multiple tools are" - in str(excinfo.value) + optional_params = get_optional_params( + model="gemini-1.5-pro", + custom_llm_provider="vertex_ai", + tools=tools, + parallel_tool_calls=False, ) + assert "parallel_tool_calls" not in optional_params + assert "tools" in optional_params - # works when specified as "functions" - with pytest.raises(litellm.utils.UnsupportedParamsError) as excinfo: - get_optional_params( - model="gemini-1.5-pro", - custom_llm_provider="vertex_ai", - functions=tools, - parallel_tool_calls=False, - ) - assert ( - "`parallel_tool_calls=False` is not supported by Gemini when multiple tools are" - in str(excinfo.value) + optional_params = get_optional_params( + model="gemini-1.5-pro", + custom_llm_provider="vertex_ai", + functions=tools, + parallel_tool_calls=False, ) + assert "parallel_tool_calls" not in optional_params def test_vertex_parallel_tool_calls_false_single_tool(): From e59e34bed3670a6894d43129c2af16af28057d03 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 20 May 2026 04:32:03 +0530 Subject: [PATCH 6/7] Gemini managed agents support (#28270) * Add support for environment variable in interactions api * Add sdk support for gemini create agent * Add agents endpoint support via proxy * Add outputs of each api * Add routing for model and agents param * Remove redundant condition in get_provider_agents_api_config LlmProviders.GEMINI.value is literally the string "gemini", so the second clause of the or was checking the exact same thing as the first. Co-authored-by: Sameer Kankute * fix: forward query-param credentials to list/get/delete/versions Gemini agent endpoints The list_gemini_agents, get_gemini_agent, delete_gemini_agent, and list_gemini_agent_versions endpoints previously constructed a hardcoded data dict with no mechanism to pass provider credentials. Unlike create_gemini_agent (POST, reads litellm_params_template from body), these GET/DELETE endpoints gave no way for multi-tenant callers to supply a per-request api_key or other LiteLLM params. Fix: - Add _merge_query_params_into_data() helper that reads query parameters from the request and merges them into the data dict without overwriting already-set keys (e.g. path params like 'name'). - Support a JSON-encoded litellm_params_template query parameter (matching the POST body pattern) as well as flat key=value pairs (e.g. api_key=AIza...). - Apply the helper in all four affected endpoints. - Add 13 unit tests covering the helper and each endpoint. Co-authored-by: Sameer Kankute * fix: pass model=None for managed agent proxy endpoints to prevent agent name polluting data["model"] Endpoints acreate_agent, aget_agent, adelete_agent, and alist_agent_versions were passing model= to base_process_llm_request. This caused common_processing_pre_call_logic to write the agent name into self.data["model"], which then triggered spurious model-alias mapping, rate-limiting lookups, and logging tied to a non-existent model deployment. The agent name is already carried in data["name"] and is passed correctly to the SDK functions (litellm.interactions.agents.*). There is no reason to also set model=; the correct value is model=None for all five managed-agent management routes. Adds tests/test_litellm/proxy/google_endpoints/test_managed_agents_model_param.py to verify all five managed-agent endpoints pass model=None. Co-authored-by: Sameer Kankute * fix: address greptile P1/P2 review comments P1 (router.py): Restore fallback/retry support for acreate_interaction and create_interaction. Both were silently moved to _init_interactions_api_endpoints (direct call, no fallbacks). Moved them back to _ageneric_api_call_with_fallbacks so users with configured fallback models keep retry behaviour. P1 security (agents_endpoints.py): Remove flat query-param credential path (e.g. ?api_key=AIza...) from _merge_query_params_into_data. Credentials in URL query strings appear verbatim in server access logs, CDN edge logs, and browser history. Only the JSON-encoded litellm_params_template query param (matching the POST body pattern) is retained. P2 (interactions/http_handler.py): Extract _BaseHTTPHandler with shared _handle_error, _sync_client, and _async_client helpers. InteractionsHTTPHandler now extends _BaseHTTPHandler. The _async_client reads the provider from litellm_params instead of hardcoding GEMINI. P2 (interactions/agents/http_handler.py): AgentsHTTPHandler now extends InteractionsHTTPHandler (which inherits _BaseHTTPHandler) so all shared HTTP infrastructure is reused rather than duplicated. Removes the hardcoded LlmProviders.GEMINI from the async client path. Co-authored-by: Cursor * fix: address CI failures from greptile review fixes - black: format interactions/agents/main.py and utils.py - tests: update test_gemini_agents_endpoints.py to match new _merge_query_params_into_data behaviour (flat credential params are rejected; only JSON-encoded litellm_params_template is accepted) - ci: add test_gemini_agents_endpoints.py to endpoints-and-responses shard in test-unit-proxy-db.yml so assert-shard-coverage passes - tests: add _initialize_managed_agents_endpoints and _init_managed_agents_api_endpoints test coverage so router_code_coverage passes; also fix TestRouterCreateInteractionRouting to reflect that acreate_interaction now correctly routes through _ageneric_api_call_with_fallbacks (restoring fallback support) Co-authored-by: Cursor * fix: remove InteractionsHTTPHandler._handle_error override to fix type errors AgentsHTTPHandler extends InteractionsHTTPHandler and calls self._handle_error(provider_config=agents_api_config) where agents_api_config is BaseAgentsAPIConfig. Python MRO resolved _handle_error to InteractionsHTTPHandler._handle_error which expected BaseInteractionsAPIConfig, causing 10 mypy arg-type errors in interactions/agents/http_handler.py. Removing the redundant override lets both classes inherit _BaseHTTPHandler._handle_error (provider_config: Any) which is structurally correct for both config types. Co-authored-by: Cursor * fix: agent-only interactions and managed agents provider routing Resolve None custom_llm_provider in agents HTTP client lookup and set custom_llm_provider on GenericLiteLLMParams for all agent CRUD paths. Stop mapping agent names to proxy model routing; route interactions through _init_interactions_api_endpoints with fallbacks only when model is set. Consolidate duplicate router elif branches for interaction APIs. Co-authored-by: Cursor * Fix greptile review * test(agents): add unit tests for managed agents SDK and HTTP handler Adds coverage for the new `litellm.interactions.agents` surface area: - main.py: sync/async entry points (create/list/get/delete/list_versions), provider config lookup, logging-obj helper, async error wrapping - http_handler.py: every CRUD method (sync + async paths), `_is_async` dispatch branches, and provider error mapping through GeminiAgentsConfig - utils.py: get_provider_agents_api_config for supported / unsupported providers Brings patch coverage on these files from <25% to ~100% so codecov/patch is satisfied. Co-authored-by: Mateo Wang * docs(gemini-agents): fix misleading credential-passing examples in GET/DELETE docstrings (#28293) The four GET/DELETE endpoint docstrings (list_gemini_agents, get_gemini_agent, delete_gemini_agent, list_gemini_agent_versions) documented passing per-request credentials as flat query parameters (e.g. ?api_key=AIza...). However, _merge_query_params_into_data only reads the JSON-encoded litellm_params_template query parameter and intentionally ignores flat params (URL query strings appear verbatim in access logs, browser history, and Referer headers). Callers following the documented curl examples would have their credentials silently dropped and hit auth failures against Gemini. Update the examples to use the supported JSON-encoded litellm_params_template query parameter, matching _merge_query_params_into_data's own docstring. Co-authored-by: Cursor Agent Co-authored-by: Mateo Wang * refactor(agents): rename provider-agnostic agent response types Move GeminiAgent{ListResponse,DeleteResult,VersionsResponse} to provider-neutral names (AgentListResponse, AgentDeleteResult, AgentVersionsResponse) so the BaseAgentsAPIConfig interface no longer references Gemini-specific type names. * fix(gemini-agents): close veria-flagged credential-escalation gaps Two high-severity findings from the veria-ai PR review are addressed: 1. **api_base override could leak the shared Gemini key** GeminiAgentsConfig.validate_environment falls back to GOOGLE_API_KEY / GEMINI_API_KEY when no api_key is supplied. Combined with caller-controlled api_base on the proxy CRUD endpoints, an authenticated user could redirect the outbound request to an attacker-controlled host and capture the operator's shared Gemini key from the x-goog-api-key header. The config now refuses env-fallback whenever api_base is explicitly overridden. 2. **Managed-agent CRUD exposed to ordinary LLM keys** The new /v1beta/agents routes live in google_routes (i.e. llm_api_routes), so any non-admin LLM key can reach them. Unlike /v1beta/models/...: generateContent these endpoints are NOT model-routed and have no model_list-supplied credentials, so env-fallback would let any LLM key list / create / delete agents inside the operator's Gemini project. Each endpoint now calls _enforce_caller_supplied_provider_key, which requires non-admin callers to supply their own Gemini api_key via litellm_params_template. Proxy admins keep the env-fallback convenience. Tests cover non-admin rejection, admin allow-through, the api_base override guard, and SDK env-fallback when api_base is not overridden. Co-authored-by: Mateo Wang * test(router): restore strict assert_called_once_with on interactions default-provider test --------- Co-authored-by: Cursor Agent Co-authored-by: Sameer Kankute Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Co-authored-by: Mateo Wang --- .github/workflows/test-unit-proxy-db.yml | 1 + litellm/__init__.py | 12 + litellm/interactions/__init__.py | 27 +- litellm/interactions/agents/__init__.py | 39 ++ litellm/interactions/agents/http_handler.py | 478 ++++++++++++++ litellm/interactions/agents/main.py | 523 ++++++++++++++++ litellm/interactions/agents/utils.py | 23 + litellm/interactions/http_handler.py | 62 +- litellm/interactions/main.py | 23 +- litellm/interactions/utils.py | 1 + litellm/llms/base_llm/agents/__init__.py | 0 .../llms/base_llm/agents/transformation.py | 165 +++++ litellm/llms/gemini/agents/__init__.py | 0 litellm/llms/gemini/agents/transformation.py | 299 +++++++++ .../gemini/interactions/transformation.py | 2 + litellm/proxy/_lazy_features.py | 5 + litellm/proxy/_types.py | 4 + litellm/proxy/agent_endpoints/utils.py | 6 + litellm/proxy/common_request_processing.py | 10 + .../google_endpoints/agents_endpoints.py | 445 +++++++++++++ litellm/proxy/google_endpoints/endpoints.py | 2 +- litellm/proxy/route_llm_request.py | 20 + litellm/router.py | 114 +++- litellm/types/agents.py | 60 ++ litellm/types/interactions/__init__.py | 2 + litellm/types/interactions/generated.py | 3 + .../test_gemini_agents_endpoints.py | 519 ++++++++++++++++ .../interactions/test_agents_http_handler.py | 587 ++++++++++++++++++ .../test_agents_main_and_utils.py | 354 +++++++++++ ...test_gemini_interactions_transformation.py | 50 ++ .../test_interactions_agent_param.py | 136 ++-- .../test_managed_agents_model_param.py | 199 ++++++ .../test_router_interactions_endpoints.py | 156 +++++ 33 files changed, 4246 insertions(+), 81 deletions(-) create mode 100644 litellm/interactions/agents/__init__.py create mode 100644 litellm/interactions/agents/http_handler.py create mode 100644 litellm/interactions/agents/main.py create mode 100644 litellm/interactions/agents/utils.py create mode 100644 litellm/llms/base_llm/agents/__init__.py create mode 100644 litellm/llms/base_llm/agents/transformation.py create mode 100644 litellm/llms/gemini/agents/__init__.py create mode 100644 litellm/llms/gemini/agents/transformation.py create mode 100644 litellm/proxy/google_endpoints/agents_endpoints.py create mode 100644 tests/proxy_unit_tests/test_gemini_agents_endpoints.py create mode 100644 tests/test_litellm/interactions/test_agents_http_handler.py create mode 100644 tests/test_litellm/interactions/test_agents_main_and_utils.py create mode 100644 tests/test_litellm/proxy/google_endpoints/test_managed_agents_model_param.py diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml index 8ee3a1ed0cd..49a36aa23f0 100644 --- a/.github/workflows/test-unit-proxy-db.yml +++ b/.github/workflows/test-unit-proxy-db.yml @@ -215,6 +215,7 @@ jobs: tests/proxy_unit_tests/test_models_fallback_endpoint.py tests/proxy_unit_tests/test_google_endpoint_routing.py tests/proxy_unit_tests/test_google_gemini_proxy_request.py + tests/proxy_unit_tests/test_gemini_agents_endpoints.py tests/proxy_unit_tests/test_get_favicon.py tests/proxy_unit_tests/test_get_image.py tests/proxy_unit_tests/test_ui_path_detection.py diff --git a/litellm/__init__.py b/litellm/__init__.py index c868ae55b4f..1e8f8613fba 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1288,6 +1288,18 @@ from .responses.main import * # Interactions API is available as litellm.interactions module # Usage: litellm.interactions.create(), litellm.interactions.get(), etc. from . import interactions +from .interactions.agents.main import ( + acreate as acreate_agent, + create as create_agent, + alist as alist_agents, + list as list_agents, + aget as aget_agent, + get as get_agent, + adelete as adelete_agent, + delete as delete_agent, + alist_versions as alist_agent_versions, + list_versions as list_agent_versions, +) from .skills.main import ( create_skill, acreate_skill, diff --git a/litellm/interactions/__init__.py b/litellm/interactions/__init__.py index e1125b649a6..ed01462cba6 100644 --- a/litellm/interactions/__init__.py +++ b/litellm/interactions/__init__.py @@ -5,31 +5,40 @@ This module provides SDK methods for Google's Interactions API. Usage: import litellm - + # Create an interaction with a model response = litellm.interactions.create( model="gemini-2.5-flash", input="Hello, how are you?" ) - + # Create an interaction with an agent response = litellm.interactions.create( agent="deep-research-pro-preview-12-2025", input="Research the current state of cancer research" ) - + # Async version response = await litellm.interactions.acreate(...) - + # Get an interaction response = litellm.interactions.get(interaction_id="...") - + # Delete an interaction result = litellm.interactions.delete(interaction_id="...") - + # Cancel an interaction result = litellm.interactions.cancel(interaction_id="...") + # Create a managed agent on the provider side + result = litellm.interactions.agents.create( + name="waverunner", + custom_llm_provider="gemini", + api_key="...", + base_agent="gemini-2.5-flash", + instructions="You are a helpful assistant.", + ) + Methods: - create(): Sync create interaction - acreate(): Async create interaction @@ -39,8 +48,12 @@ Methods: - adelete(): Async delete interaction - cancel(): Sync cancel interaction - acancel(): Async cancel interaction + +Sub-modules: +- agents: Provider-side agent creation (litellm.interactions.agents.create) """ +from litellm.interactions import agents from litellm.interactions.main import ( acancel, acreate, @@ -65,4 +78,6 @@ __all__ = [ # Cancel "cancel", "acancel", + # Sub-modules + "agents", ] diff --git a/litellm/interactions/agents/__init__.py b/litellm/interactions/agents/__init__.py new file mode 100644 index 00000000000..711a54fdcbb --- /dev/null +++ b/litellm/interactions/agents/__init__.py @@ -0,0 +1,39 @@ +""" +litellm.interactions.agents + +Full CRUD SDK for provider-side managed agents (e.g. Gemini v1beta/agents). + + litellm.interactions.agents.create(name=..., ...) + litellm.interactions.agents.list(api_key=...) + litellm.interactions.agents.get(name=..., ...) + litellm.interactions.agents.delete(name=..., ...) + litellm.interactions.agents.list_versions(name=..., ...) + +Async counterparts: acreate, alist, aget, adelete, alist_versions +""" + +from litellm.interactions.agents.main import ( + acreate, + adelete, + aget, + alist, + alist_versions, + create, + delete, + get, + list, + list_versions, +) + +__all__ = [ + "create", + "acreate", + "list", + "alist", + "get", + "aget", + "delete", + "adelete", + "list_versions", + "alist_versions", +] diff --git a/litellm/interactions/agents/http_handler.py b/litellm/interactions/agents/http_handler.py new file mode 100644 index 00000000000..d45ca6f4346 --- /dev/null +++ b/litellm/interactions/agents/http_handler.py @@ -0,0 +1,478 @@ +""" +HTTP handler for the Agents API. + +Extends InteractionsHTTPHandler so that the shared HTTP infrastructure +(_handle_error, _sync_client, _async_client) is reused rather than +duplicated. BaseAgentsAPIConfig stays as pure transform code. +""" + +from typing import Any, Coroutine, Dict, Optional, Union + +import httpx + +from litellm.constants import request_timeout +from litellm.interactions.http_handler import InteractionsHTTPHandler +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.types.agents import ( + AgentCreateResponse, + AgentDeleteResult, + AgentListResponse, + AgentVersionsResponse, +) +from litellm.types.router import GenericLiteLLMParams + + +class AgentsHTTPHandler(InteractionsHTTPHandler): + """HTTP handler for Agents API CRUD requests.""" + + # ------------------------------------------------------------------ # + # CREATE # + # ------------------------------------------------------------------ # + + def create_agent( + self, + agents_api_config: BaseAgentsAPIConfig, + name: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + extra_body: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[HTTPHandler] = None, + _is_async: bool = False, + ) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]: + if _is_async: + return self.async_create_agent( + agents_api_config=agents_api_config, + name=name, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + extra_body=extra_body, + timeout=timeout, + ) + + sync_httpx_client = self._sync_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url = agents_api_config.get_complete_url( + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + data = agents_api_config.transform_create_request( + name=name, litellm_params=dict(litellm_params) + ) + if extra_body: + data.update(extra_body) + + logging_obj.pre_call( + input=name, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": url, + "headers": headers, + }, + ) + try: + response = sync_httpx_client.post( + url=url, headers=headers, json=data, timeout=timeout or request_timeout + ) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call( + original_response=response.text, + additional_args={"complete_input_dict": data}, + ) + return agents_api_config.transform_create_response( + raw_response=response, name=name + ) + + async def async_create_agent( + self, + agents_api_config: BaseAgentsAPIConfig, + name: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + extra_body: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[AsyncHTTPHandler] = None, + ) -> AgentCreateResponse: + async_httpx_client = self._async_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url = agents_api_config.get_complete_url( + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + data = agents_api_config.transform_create_request( + name=name, litellm_params=dict(litellm_params) + ) + if extra_body: + data.update(extra_body) + + logging_obj.pre_call( + input=name, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": url, + "headers": headers, + }, + ) + try: + response = await async_httpx_client.post( + url=url, headers=headers, json=data, timeout=timeout or request_timeout + ) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call( + original_response=response.text, + additional_args={"complete_input_dict": data}, + ) + return agents_api_config.transform_create_response( + raw_response=response, name=name + ) + + # ------------------------------------------------------------------ # + # LIST # + # ------------------------------------------------------------------ # + + def list_agents( + self, + agents_api_config: BaseAgentsAPIConfig, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[HTTPHandler] = None, + _is_async: bool = False, + ) -> Union[AgentListResponse, Coroutine[Any, Any, AgentListResponse]]: + if _is_async: + return self.async_list_agents( + agents_api_config=agents_api_config, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + timeout=timeout, + ) + + sync_httpx_client = self._sync_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url, params = agents_api_config.transform_list_request( + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + logging_obj.pre_call( + input="list_agents", + api_key="", + additional_args={"api_base": url, "headers": headers}, + ) + try: + response = sync_httpx_client.get(url=url, headers=headers, params=params) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call(original_response=response.text, additional_args={}) + return agents_api_config.transform_list_response(raw_response=response) + + async def async_list_agents( + self, + agents_api_config: BaseAgentsAPIConfig, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[AsyncHTTPHandler] = None, + ) -> AgentListResponse: + async_httpx_client = self._async_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url, params = agents_api_config.transform_list_request( + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + logging_obj.pre_call( + input="list_agents", + api_key="", + additional_args={"api_base": url, "headers": headers}, + ) + try: + response = await async_httpx_client.get( + url=url, headers=headers, params=params + ) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call(original_response=response.text, additional_args={}) + return agents_api_config.transform_list_response(raw_response=response) + + # ------------------------------------------------------------------ # + # GET # + # ------------------------------------------------------------------ # + + def get_agent( + self, + agents_api_config: BaseAgentsAPIConfig, + name: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[HTTPHandler] = None, + _is_async: bool = False, + ) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]: + if _is_async: + return self.async_get_agent( + agents_api_config=agents_api_config, + name=name, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + timeout=timeout, + ) + + sync_httpx_client = self._sync_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url, params = agents_api_config.transform_get_request( + name=name, + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + logging_obj.pre_call( + input=name, + api_key="", + additional_args={"api_base": url, "headers": headers}, + ) + try: + response = sync_httpx_client.get(url=url, headers=headers, params=params) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call(original_response=response.text, additional_args={}) + return agents_api_config.transform_get_response( + raw_response=response, name=name + ) + + async def async_get_agent( + self, + agents_api_config: BaseAgentsAPIConfig, + name: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[AsyncHTTPHandler] = None, + ) -> AgentCreateResponse: + async_httpx_client = self._async_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url, params = agents_api_config.transform_get_request( + name=name, + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + logging_obj.pre_call( + input=name, + api_key="", + additional_args={"api_base": url, "headers": headers}, + ) + try: + response = await async_httpx_client.get( + url=url, headers=headers, params=params + ) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call(original_response=response.text, additional_args={}) + return agents_api_config.transform_get_response( + raw_response=response, name=name + ) + + # ------------------------------------------------------------------ # + # DELETE # + # ------------------------------------------------------------------ # + + def delete_agent( + self, + agents_api_config: BaseAgentsAPIConfig, + name: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[HTTPHandler] = None, + _is_async: bool = False, + ) -> Union[AgentDeleteResult, Coroutine[Any, Any, AgentDeleteResult]]: + if _is_async: + return self.async_delete_agent( + agents_api_config=agents_api_config, + name=name, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + timeout=timeout, + ) + + sync_httpx_client = self._sync_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url = agents_api_config.transform_delete_request( + name=name, + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + logging_obj.pre_call( + input=name, + api_key="", + additional_args={"api_base": url, "headers": headers}, + ) + try: + response = sync_httpx_client.delete( + url=url, headers=headers, timeout=timeout or request_timeout + ) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call(original_response=response.text, additional_args={}) + return agents_api_config.transform_delete_response( + raw_response=response, name=name + ) + + async def async_delete_agent( + self, + agents_api_config: BaseAgentsAPIConfig, + name: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[AsyncHTTPHandler] = None, + ) -> AgentDeleteResult: + async_httpx_client = self._async_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url = agents_api_config.transform_delete_request( + name=name, + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + logging_obj.pre_call( + input=name, + api_key="", + additional_args={"api_base": url, "headers": headers}, + ) + try: + response = await async_httpx_client.delete( + url=url, headers=headers, timeout=timeout or request_timeout + ) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call(original_response=response.text, additional_args={}) + return agents_api_config.transform_delete_response( + raw_response=response, name=name + ) + + # ------------------------------------------------------------------ # + # LIST VERSIONS # + # ------------------------------------------------------------------ # + + def list_agent_versions( + self, + agents_api_config: BaseAgentsAPIConfig, + name: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[HTTPHandler] = None, + _is_async: bool = False, + ) -> Union[AgentVersionsResponse, Coroutine[Any, Any, AgentVersionsResponse]]: + if _is_async: + return self.async_list_agent_versions( + agents_api_config=agents_api_config, + name=name, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + timeout=timeout, + ) + + sync_httpx_client = self._sync_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url, params = agents_api_config.transform_list_versions_request( + name=name, + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + logging_obj.pre_call( + input=name, + api_key="", + additional_args={"api_base": url, "headers": headers}, + ) + try: + response = sync_httpx_client.get(url=url, headers=headers, params=params) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call(original_response=response.text, additional_args={}) + return agents_api_config.transform_list_versions_response( + raw_response=response, name=name + ) + + async def async_list_agent_versions( + self, + agents_api_config: BaseAgentsAPIConfig, + name: str, + litellm_params: GenericLiteLLMParams, + logging_obj: LiteLLMLoggingObj, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + client: Optional[AsyncHTTPHandler] = None, + ) -> AgentVersionsResponse: + async_httpx_client = self._async_client(litellm_params, client) + headers = agents_api_config.validate_environment( + headers=extra_headers or {}, litellm_params=dict(litellm_params) + ) + url, params = agents_api_config.transform_list_versions_request( + name=name, + api_base=litellm_params.get("api_base"), + litellm_params=dict(litellm_params), + ) + logging_obj.pre_call( + input=name, + api_key="", + additional_args={"api_base": url, "headers": headers}, + ) + try: + response = await async_httpx_client.get( + url=url, headers=headers, params=params + ) + except Exception as e: + raise self._handle_error(e=e, provider_config=agents_api_config) + + logging_obj.post_call(original_response=response.text, additional_args={}) + return agents_api_config.transform_list_versions_response( + raw_response=response, name=name + ) + + +agents_http_handler = AgentsHTTPHandler() diff --git a/litellm/interactions/agents/main.py b/litellm/interactions/agents/main.py new file mode 100644 index 00000000000..7375fd6273f --- /dev/null +++ b/litellm/interactions/agents/main.py @@ -0,0 +1,523 @@ +""" +LiteLLM Agents API - Main Module + +Usage: + import litellm + + # Create + response = litellm.interactions.agents.create( + name="waverunner", + custom_llm_provider="gemini", + api_key="...", + base_agent="gemini-2.5-flash", + instructions="You are a helpful assistant.", + ) + + # List + response = litellm.interactions.agents.list(api_key="...", custom_llm_provider="gemini") + + # Get + response = litellm.interactions.agents.get(name="waverunner", api_key="...") + + # Delete + result = litellm.interactions.agents.delete(name="waverunner", api_key="...") + + # List versions + result = litellm.interactions.agents.list_versions(name="waverunner", api_key="...") + + # Async versions: acreate, alist, aget, adelete, alist_versions +""" + +import asyncio +import contextvars +from functools import partial +from typing import Any, Coroutine, Dict, Optional, Union + +import httpx + +import litellm +from litellm.interactions.agents.http_handler import agents_http_handler +from litellm.interactions.agents.utils import get_provider_agents_api_config +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.types.agents import ( + AgentCreateResponse, + AgentDeleteResult, + AgentListResponse, + AgentVersionsResponse, +) +from litellm.types.interactions import InteractionEnvironment +from litellm.types.router import GenericLiteLLMParams +from litellm.utils import client + + +# ------------------------------------------------------------------ # +# Shared helpers # +# ------------------------------------------------------------------ # + + +def _get_agents_api_config(custom_llm_provider: str): + config = get_provider_agents_api_config(custom_llm_provider) + if config is None: + raise litellm.BadRequestError( + message=( + f"Provider '{custom_llm_provider}' does not have a native " + "agents API. Use the proxy POST /v1/agents endpoint to store " + "agents locally." + ), + model="", + llm_provider=custom_llm_provider, + ) + return config + + +def _make_logging_obj( + kwargs: Dict[str, Any], + model: str, + custom_llm_provider: str, + call_type: str, + optional_params: Dict[str, Any], +) -> LiteLLMLoggingObj: + litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore + litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None) + litellm_logging_obj.update_from_kwargs( + kwargs=kwargs, + model=model, + optional_params=optional_params, + litellm_params={"litellm_call_id": litellm_call_id}, + custom_llm_provider=custom_llm_provider, + ) + return litellm_logging_obj + + +# ================================================================== # +# CREATE # +# ================================================================== # + + +@client +async def acreate( + name: str, + base_agent: Optional[str] = None, + instructions: Optional[str] = None, + base_environment: Optional[InteractionEnvironment] = None, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + extra_body: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> AgentCreateResponse: + """Async: Create a managed agent on the provider side.""" + local_vars = locals() + try: + loop = asyncio.get_event_loop() + kwargs["acreate_agent"] = True + func = partial( + create, + name=name, + base_agent=base_agent, + instructions=instructions, + base_environment=base_environment, + custom_llm_provider=custom_llm_provider or "gemini", + extra_headers=extra_headers, + extra_body=extra_body, + timeout=timeout, + **kwargs, + ) + ctx = contextvars.copy_context() + init_response = await loop.run_in_executor(None, partial(ctx.run, func)) + if asyncio.iscoroutine(init_response): + return await init_response + return init_response + except Exception as e: + raise litellm.exception_type( + model=name, + custom_llm_provider=custom_llm_provider or "gemini", + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +@client +def create( + name: str, + base_agent: Optional[str] = None, + instructions: Optional[str] = None, + base_environment: Optional[InteractionEnvironment] = None, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + extra_body: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]: + """ + Sync: Create a managed agent on the provider side. + + Args: + name: Name for the agent (required). + base_agent: Base agent to derive from (e.g. "waverunner"). + instructions: System instructions for the agent. + base_environment: Environment to fork from — an env_id string or a + dict like ``{"type": "remote", "sources": [...]}``. + custom_llm_provider: Provider to use, e.g. "gemini". + extra_headers: Additional HTTP headers. + extra_body: Additional request body fields. + timeout: Request timeout. + **kwargs: Forwarded to GenericLiteLLMParams (api_key, api_base, etc.). + """ + local_vars = locals() + custom_llm_provider = ( + custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini" + ) + try: + _is_async = kwargs.pop("acreate_agent", False) is True + if base_agent is not None: + kwargs["base_agent"] = base_agent + if instructions is not None: + kwargs["instructions"] = instructions + if base_environment is not None: + kwargs["base_environment"] = base_environment + kwargs.setdefault("custom_llm_provider", custom_llm_provider) + litellm_params = GenericLiteLLMParams(**kwargs) + logging_obj = _make_logging_obj( + kwargs, name, custom_llm_provider, "create_agent", {} + ) + config = _get_agents_api_config(custom_llm_provider) + return agents_http_handler.create_agent( + agents_api_config=config, + name=name, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + extra_body=extra_body, + timeout=timeout, + _is_async=_is_async, + ) + except Exception as e: + raise litellm.exception_type( + model=name, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +# ================================================================== # +# LIST # +# ================================================================== # + + +@client +async def alist( + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> AgentListResponse: + """Async: List all agents on the provider side.""" + local_vars = locals() + try: + loop = asyncio.get_event_loop() + kwargs["alist_agents"] = True + func = partial( + list, + custom_llm_provider=custom_llm_provider or "gemini", + extra_headers=extra_headers, + timeout=timeout, + **kwargs, + ) + ctx = contextvars.copy_context() + init_response = await loop.run_in_executor(None, partial(ctx.run, func)) + if asyncio.iscoroutine(init_response): + return await init_response + return init_response + except Exception as e: + raise litellm.exception_type( + model="", + custom_llm_provider=custom_llm_provider or "gemini", + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +@client +def list( + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> Union[AgentListResponse, Coroutine[Any, Any, AgentListResponse]]: + """Sync: List all agents on the provider side.""" + local_vars = locals() + custom_llm_provider = ( + custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini" + ) + try: + _is_async = kwargs.pop("alist_agents", False) is True + kwargs.setdefault("custom_llm_provider", custom_llm_provider) + litellm_params = GenericLiteLLMParams(**kwargs) + logging_obj = _make_logging_obj( + kwargs, "", custom_llm_provider, "list_agents", {} + ) + config = _get_agents_api_config(custom_llm_provider) + return agents_http_handler.list_agents( + agents_api_config=config, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + timeout=timeout, + _is_async=_is_async, + ) + except Exception as e: + raise litellm.exception_type( + model="", + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +# ================================================================== # +# GET # +# ================================================================== # + + +@client +async def aget( + name: str, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> AgentCreateResponse: + """Async: Get a specific agent by name.""" + local_vars = locals() + try: + loop = asyncio.get_event_loop() + kwargs["aget_agent"] = True + func = partial( + get, + name=name, + custom_llm_provider=custom_llm_provider or "gemini", + extra_headers=extra_headers, + timeout=timeout, + **kwargs, + ) + ctx = contextvars.copy_context() + init_response = await loop.run_in_executor(None, partial(ctx.run, func)) + if asyncio.iscoroutine(init_response): + return await init_response + return init_response + except Exception as e: + raise litellm.exception_type( + model=name, + custom_llm_provider=custom_llm_provider or "gemini", + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +@client +def get( + name: str, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]: + """Sync: Get a specific agent by name.""" + local_vars = locals() + custom_llm_provider = ( + custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini" + ) + try: + _is_async = kwargs.pop("aget_agent", False) is True + kwargs.setdefault("custom_llm_provider", custom_llm_provider) + litellm_params = GenericLiteLLMParams(**kwargs) + logging_obj = _make_logging_obj( + kwargs, name, custom_llm_provider, "get_agent", {"name": name} + ) + config = _get_agents_api_config(custom_llm_provider) + return agents_http_handler.get_agent( + agents_api_config=config, + name=name, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + timeout=timeout, + _is_async=_is_async, + ) + except Exception as e: + raise litellm.exception_type( + model=name, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +# ================================================================== # +# DELETE # +# ================================================================== # + + +@client +async def adelete( + name: str, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> AgentDeleteResult: + """Async: Delete a specific agent by name.""" + local_vars = locals() + try: + loop = asyncio.get_event_loop() + kwargs["adelete_agent"] = True + func = partial( + delete, + name=name, + custom_llm_provider=custom_llm_provider or "gemini", + extra_headers=extra_headers, + timeout=timeout, + **kwargs, + ) + ctx = contextvars.copy_context() + init_response = await loop.run_in_executor(None, partial(ctx.run, func)) + if asyncio.iscoroutine(init_response): + return await init_response + return init_response + except Exception as e: + raise litellm.exception_type( + model=name, + custom_llm_provider=custom_llm_provider or "gemini", + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +@client +def delete( + name: str, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> Union[AgentDeleteResult, Coroutine[Any, Any, AgentDeleteResult]]: + """Sync: Delete a specific agent by name.""" + local_vars = locals() + custom_llm_provider = ( + custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini" + ) + try: + _is_async = kwargs.pop("adelete_agent", False) is True + kwargs.setdefault("custom_llm_provider", custom_llm_provider) + litellm_params = GenericLiteLLMParams(**kwargs) + logging_obj = _make_logging_obj( + kwargs, name, custom_llm_provider, "delete_agent", {"name": name} + ) + config = _get_agents_api_config(custom_llm_provider) + return agents_http_handler.delete_agent( + agents_api_config=config, + name=name, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + timeout=timeout, + _is_async=_is_async, + ) + except Exception as e: + raise litellm.exception_type( + model=name, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +# ================================================================== # +# LIST VERSIONS # +# ================================================================== # + + +@client +async def alist_versions( + name: str, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> AgentVersionsResponse: + """Async: List versions of a specific agent.""" + local_vars = locals() + try: + loop = asyncio.get_event_loop() + kwargs["alist_agent_versions"] = True + func = partial( + list_versions, + name=name, + custom_llm_provider=custom_llm_provider or "gemini", + extra_headers=extra_headers, + timeout=timeout, + **kwargs, + ) + ctx = contextvars.copy_context() + init_response = await loop.run_in_executor(None, partial(ctx.run, func)) + if asyncio.iscoroutine(init_response): + return await init_response + return init_response + except Exception as e: + raise litellm.exception_type( + model=name, + custom_llm_provider=custom_llm_provider or "gemini", + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) + + +@client +def list_versions( + name: str, + custom_llm_provider: Optional[str] = None, + extra_headers: Optional[Dict[str, Any]] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + **kwargs, +) -> Union[AgentVersionsResponse, Coroutine[Any, Any, AgentVersionsResponse]]: + """Sync: List versions of a specific agent.""" + local_vars = locals() + custom_llm_provider = ( + custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini" + ) + try: + _is_async = kwargs.pop("alist_agent_versions", False) is True + kwargs.setdefault("custom_llm_provider", custom_llm_provider) + litellm_params = GenericLiteLLMParams(**kwargs) + logging_obj = _make_logging_obj( + kwargs, name, custom_llm_provider, "list_agent_versions", {"name": name} + ) + config = _get_agents_api_config(custom_llm_provider) + return agents_http_handler.list_agent_versions( + agents_api_config=config, + name=name, + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers=extra_headers, + timeout=timeout, + _is_async=_is_async, + ) + except Exception as e: + raise litellm.exception_type( + model=name, + custom_llm_provider=custom_llm_provider, + original_exception=e, + completion_kwargs=local_vars, + extra_kwargs=kwargs, + ) diff --git a/litellm/interactions/agents/utils.py b/litellm/interactions/agents/utils.py new file mode 100644 index 00000000000..d16a9597f53 --- /dev/null +++ b/litellm/interactions/agents/utils.py @@ -0,0 +1,23 @@ +""" +Utility functions for the Agents API SDK. +""" + +from typing import Optional + +from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig + + +def get_provider_agents_api_config( + custom_llm_provider: Optional[str], +) -> Optional[BaseAgentsAPIConfig]: + """ + Return a provider-specific BaseAgentsAPIConfig if the provider has a + native agent-creation API, or None otherwise. + """ + from litellm.types.utils import LlmProviders + + if custom_llm_provider == LlmProviders.GEMINI.value: + from litellm.llms.gemini.agents.transformation import GeminiAgentsConfig + + return GeminiAgentsConfig() + return None diff --git a/litellm/interactions/http_handler.py b/litellm/interactions/http_handler.py index 7fead07043f..695da2be89a 100644 --- a/litellm/interactions/http_handler.py +++ b/litellm/interactions/http_handler.py @@ -41,27 +41,55 @@ from litellm.types.interactions import ( from litellm.types.router import GenericLiteLLMParams -class InteractionsHTTPHandler: +class _BaseHTTPHandler: + """ + Shared HTTP infrastructure for LiteLLM handler classes. + + Provides common client resolution and error-mapping helpers so that + handler subclasses (InteractionsHTTPHandler, AgentsHTTPHandler, …) do + not duplicate this boilerplate. + """ + + def _handle_error(self, e: Exception, provider_config: Any) -> Exception: + if isinstance(e, httpx.HTTPStatusError): + return provider_config.get_error_class( + error_message=e.response.text, + status_code=e.response.status_code, + headers=dict(e.response.headers), + ) + return e + + def _sync_client( + self, + litellm_params: GenericLiteLLMParams, + client: Optional[HTTPHandler], + ) -> HTTPHandler: + return client or _get_httpx_client( + params={"ssl_verify": litellm_params.get("ssl_verify", None)} + ) + + def _async_client( + self, + litellm_params: GenericLiteLLMParams, + client: Optional[AsyncHTTPHandler], + ) -> AsyncHTTPHandler: + # GenericLiteLLMParams.get uses getattr; an unset field is None, not the default. + custom_llm_provider = litellm_params.get("custom_llm_provider") or "gemini" + return client or get_async_httpx_client( + llm_provider=litellm.LlmProviders(custom_llm_provider), + params={"ssl_verify": litellm_params.get("ssl_verify", None)}, + ) + + +class InteractionsHTTPHandler(_BaseHTTPHandler): """ HTTP handler for Interactions API requests. """ - def _handle_error( - self, - e: Exception, - provider_config: BaseInteractionsAPIConfig, - ) -> Exception: - """Handle errors from HTTP requests.""" - if isinstance(e, httpx.HTTPStatusError): - error_message = e.response.text - status_code = e.response.status_code - headers = dict(e.response.headers) - return provider_config.get_error_class( - error_message=error_message, - status_code=status_code, - headers=headers, - ) - return e + # _handle_error is inherited from _BaseHTTPHandler (accepts Any provider_config). + # AgentsHTTPHandler also extends this class and passes BaseAgentsAPIConfig, which + # is structurally compatible but a different type — keeping the override here with + # BaseInteractionsAPIConfig would cause type errors in the subclass. # ========================================================= # CREATE INTERACTION diff --git a/litellm/interactions/main.py b/litellm/interactions/main.py index ab429ef6db5..c6eca410fa7 100644 --- a/litellm/interactions/main.py +++ b/litellm/interactions/main.py @@ -48,6 +48,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.types.interactions import ( CancelInteractionResult, DeleteInteractionResult, + InteractionEnvironment, InteractionInput, InteractionsAPIResponse, InteractionsAPIStreamingResponse, @@ -80,6 +81,8 @@ async def acreate( store: Optional[bool] = None, # Background execution background: Optional[bool] = None, + # Agent execution environment ("remote", env id, or remote config object) + environment: Optional[InteractionEnvironment] = None, # Response format response_modalities: Optional[List[str]] = None, response_format: Optional[Dict[str, Any]] = None, @@ -109,6 +112,10 @@ async def acreate( stream: Whether to stream the response store: Whether to store the response for later retrieval background: Whether to run in background + environment: Agent execution environment — ``"remote"``, an existing env id + string, or a config object such as + ``{"type": "remote", "sources": [...]}`` / + ``{"type": "remote", "network": {...}}`` response_modalities: Requested response modalities (TEXT, IMAGE, AUDIO) response_format: JSON schema for response format response_mime_type: MIME type of the response @@ -144,6 +151,7 @@ async def acreate( stream=stream, store=store, background=background, + environment=environment, response_modalities=response_modalities, response_format=response_format, response_mime_type=response_mime_type, @@ -194,6 +202,8 @@ def create( store: Optional[bool] = None, # Background execution background: Optional[bool] = None, + # Agent execution environment ("remote", env id, or remote config object) + environment: Optional[InteractionEnvironment] = None, # Response format response_modalities: Optional[List[str]] = None, response_format: Optional[Dict[str, Any]] = None, @@ -231,6 +241,10 @@ def create( stream: Whether to stream the response store: Whether to store the response for later retrieval background: Whether to run in background + environment: Agent execution environment — ``"remote"``, an existing env id + string, or a config object such as + ``{"type": "remote", "sources": [...]}`` / + ``{"type": "remote", "network": {...}}`` response_modalities: Requested response modalities (TEXT, IMAGE, AUDIO) response_format: JSON schema for response format response_mime_type: MIME type of the response @@ -252,7 +266,14 @@ def create( litellm_params = GenericLiteLLMParams(**kwargs) - if model: + # Routing logic: + # - agent provided (no model, or model accidentally set to agent name) → gemini + # - model provided → resolve provider via get_llm_provider (normal routing) + if agent and model == agent: + model = None + if agent and not model: + custom_llm_provider = custom_llm_provider or "gemini" + elif model: model, custom_llm_provider, _, _ = litellm.get_llm_provider( model=model, custom_llm_provider=custom_llm_provider, diff --git a/litellm/interactions/utils.py b/litellm/interactions/utils.py index 3a18ddf52fe..84437f4d3d8 100644 --- a/litellm/interactions/utils.py +++ b/litellm/interactions/utils.py @@ -15,6 +15,7 @@ INTERACTIONS_API_OPTIONAL_PARAMS = { "stream", "store", "background", + "environment", "response_modalities", "response_format", "response_mime_type", diff --git a/litellm/llms/base_llm/agents/__init__.py b/litellm/llms/base_llm/agents/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/base_llm/agents/transformation.py b/litellm/llms/base_llm/agents/transformation.py new file mode 100644 index 00000000000..508e54cb7ab --- /dev/null +++ b/litellm/llms/base_llm/agents/transformation.py @@ -0,0 +1,165 @@ +""" +Base transformation class for provider-side Agents API. + +Providers that have a native agents CRUD API (e.g. Gemini v1beta/agents) +subclass BaseAgentsAPIConfig and implement the abstract methods. + +The HTTP calls are handled by AgentsHTTPHandler — this class is pure +transform logic (same separation as BaseInteractionsAPIConfig / +InteractionsHTTPHandler). +""" + +from abc import ABC, abstractmethod +from typing import Any, Dict, Optional, Tuple, Union + +import httpx + +from litellm.types.agents import ( + AgentCreateResponse, + AgentDeleteResult, + AgentListResponse, + AgentVersionsResponse, +) + + +class BaseAgentsAPIConfig(ABC): + """ + Minimal interface for providers that expose a native agents CRUD API. + """ + + # ------------------------------------------------------------------ # + # CREATE # + # ------------------------------------------------------------------ # + + @abstractmethod + def get_complete_url( + self, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> str: + """Return the full URL for POST /agents (create).""" + + @abstractmethod + def validate_environment( + self, + headers: Dict[str, str], + litellm_params: Dict[str, Any], + ) -> Dict[str, str]: + """Validate credentials and return auth headers.""" + + @abstractmethod + def transform_create_request( + self, + name: str, + litellm_params: Dict[str, Any], + ) -> Dict[str, Any]: + """Map name + litellm_params to the provider's create-agent body.""" + + @abstractmethod + def transform_create_response( + self, + raw_response: httpx.Response, + name: str, + ) -> AgentCreateResponse: + """Parse create response. Raise on non-2xx.""" + + # ------------------------------------------------------------------ # + # LIST # + # ------------------------------------------------------------------ # + + @abstractmethod + def transform_list_request( + self, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> Tuple[str, Dict[str, Any]]: + """Return (url, query_params) for GET /agents.""" + + @abstractmethod + def transform_list_response( + self, + raw_response: httpx.Response, + ) -> AgentListResponse: + """Parse list-agents response. Raise on non-2xx.""" + + # ------------------------------------------------------------------ # + # GET # + # ------------------------------------------------------------------ # + + @abstractmethod + def transform_get_request( + self, + name: str, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> Tuple[str, Dict[str, Any]]: + """Return (url, query_params) for GET /agents/{name}.""" + + @abstractmethod + def transform_get_response( + self, + raw_response: httpx.Response, + name: str, + ) -> AgentCreateResponse: + """Parse get-agent response. Raise on non-2xx.""" + + # ------------------------------------------------------------------ # + # DELETE # + # ------------------------------------------------------------------ # + + @abstractmethod + def transform_delete_request( + self, + name: str, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> str: + """Return the URL for DELETE /agents/{name}.""" + + @abstractmethod + def transform_delete_response( + self, + raw_response: httpx.Response, + name: str, + ) -> AgentDeleteResult: + """Parse delete-agent response. Raise on non-2xx.""" + + # ------------------------------------------------------------------ # + # LIST VERSIONS # + # ------------------------------------------------------------------ # + + @abstractmethod + def transform_list_versions_request( + self, + name: str, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> Tuple[str, Dict[str, Any]]: + """Return (url, query_params) for GET /agents/{name}/versions.""" + + @abstractmethod + def transform_list_versions_response( + self, + raw_response: httpx.Response, + name: str, + ) -> AgentVersionsResponse: + """Parse list-versions response. Raise on non-2xx.""" + + # ------------------------------------------------------------------ # + # ERROR HANDLING # + # ------------------------------------------------------------------ # + + def get_error_class( + self, + error_message: str, + status_code: int, + headers: Union[dict, httpx.Headers], + ) -> Exception: + """Map HTTP error status codes to provider-specific exceptions.""" + from litellm.llms.base_llm.chat.transformation import BaseLLMException + + return BaseLLMException( + status_code=status_code, + message=error_message, + headers=headers, + ) diff --git a/litellm/llms/gemini/agents/__init__.py b/litellm/llms/gemini/agents/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/gemini/agents/transformation.py b/litellm/llms/gemini/agents/transformation.py new file mode 100644 index 00000000000..150918c4737 --- /dev/null +++ b/litellm/llms/gemini/agents/transformation.py @@ -0,0 +1,299 @@ +""" +Google AI Studio Agents API configuration. + +Proxies the Gemini v1beta Agents API: + POST /v1beta/agents create + GET /v1beta/agents list + GET /v1beta/agents/{name} get + DELETE /v1beta/agents/{name} delete + GET /v1beta/agents/{name}/versions list versions +""" + +from typing import Any, Dict, Optional, Tuple, Union + +import httpx + +from litellm._logging import verbose_logger +from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig +from litellm.llms.gemini.common_utils import GeminiError, GeminiModelInfo +from litellm.types.agents import ( + AgentCreateResponse, + AgentDeleteResult, + AgentListResponse, + AgentVersionsResponse, +) + + +# Keys inside litellm_params that should be forwarded to the Gemini +# create-agent body verbatim. +_GEMINI_AGENT_BODY_KEYS = ("base_agent", "instructions", "base_environment") + +# LiteLLM-internal keys that must never be forwarded to Gemini. +_LITELLM_INTERNAL_KEYS = frozenset( + { + "custom_llm_provider", + "api_key", + "api_base", + "make_public", + "cost_per_query", + "input_cost_per_token", + "output_cost_per_token", + "require_trace_id_on_calls_to_agent", + "require_trace_id_on_calls_by_agent", + "max_iterations", + "max_budget_per_session", + "guardrails", + "is_public", + "agent_name", + "agent_id", + "agent_card_params", + "provider_agent_response", + } +) + + +class GeminiAgentsConfig(BaseAgentsAPIConfig): + """ + Configuration for the Google AI Studio (Gemini) native Agents API. + + Authentication uses x-goog-api-key, resolved from (in order): + 1. litellm_params["api_key"] + 2. GOOGLE_API_KEY env var + 3. GEMINI_API_KEY env var + """ + + @property + def api_version(self) -> str: + return "v1beta" + + def _base_url(self, api_base: Optional[str]) -> str: + return f"{GeminiModelInfo.get_api_base(api_base)}/{self.api_version}" + + # ------------------------------------------------------------------ # + # Shared helpers # + # ------------------------------------------------------------------ # + + def get_error_class( + self, + error_message: str, + status_code: int, + headers: Union[dict, httpx.Headers], + ) -> Exception: + return GeminiError( + message=error_message, + status_code=status_code, + headers=dict(headers), + ) + + def get_complete_url( + self, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> str: + return f"{self._base_url(api_base)}/agents" + + def validate_environment( + self, + headers: Dict[str, str], + litellm_params: Dict[str, Any], + ) -> Dict[str, str]: + headers = dict(headers) + headers["Content-Type"] = "application/json" + explicit_api_key = litellm_params.get("api_key") + # SECURITY: when the caller overrides ``api_base``, refuse to fall back + # to the process-wide GOOGLE_API_KEY / GEMINI_API_KEY env vars. Otherwise + # an authenticated proxy user could set ``api_base`` to an attacker- + # controlled host and have the proxy ship its shared Gemini key in the + # ``x-goog-api-key`` header. + if litellm_params.get("api_base") and not explicit_api_key: + raise ValueError( + "When overriding api_base for Gemini agents, you must also " + "supply an explicit api_key. Falling back to GOOGLE_API_KEY / " + "GEMINI_API_KEY env vars with a custom api_base is refused " + "to prevent leaking the shared provider key to arbitrary hosts." + ) + api_key = GeminiModelInfo.get_api_key(explicit_api_key) + if not api_key: + raise ValueError( + "Google API key is required. " + "Set GOOGLE_API_KEY or GEMINI_API_KEY, or pass api_key." + ) + headers["x-goog-api-key"] = api_key + return headers + + def _raise_for_status(self, raw_response: httpx.Response) -> None: + if not (200 <= raw_response.status_code < 300): + raise GeminiError( + message=raw_response.text, + status_code=raw_response.status_code, + headers=dict(raw_response.headers), + ) + + # ------------------------------------------------------------------ # + # CREATE # + # ------------------------------------------------------------------ # + + def transform_create_request( + self, + name: str, + litellm_params: Dict[str, Any], + ) -> Dict[str, Any]: + body: Dict[str, Any] = {"name": name} + for key in _GEMINI_AGENT_BODY_KEYS: + value = litellm_params.get(key) + if value is not None: + body[key] = value + verbose_logger.debug("GeminiAgentsConfig create body: %s", body) + return body + + def transform_create_response( + self, + raw_response: httpx.Response, + name: str, + ) -> AgentCreateResponse: + """ + Gemini returns: + {"id": "my-agent", "base_agent": "waverunner", + "system_instruction": "...", "base_environment": {...}} + """ + self._raise_for_status(raw_response) + try: + data: Dict[str, Any] = raw_response.json() + except Exception: + verbose_logger.warning( + "GeminiAgentsConfig: non-JSON create response (status=%d).", + raw_response.status_code, + ) + data = {"id": name} + # Gemini uses "id" as the identifier; normalise to both fields. + data.setdefault("id", name) + data.setdefault("name", data["id"]) + verbose_logger.debug("GeminiAgentsConfig create response: %s", data) + return AgentCreateResponse(**data) + + # ------------------------------------------------------------------ # + # LIST # + # ------------------------------------------------------------------ # + + def transform_list_request( + self, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> Tuple[str, Dict[str, Any]]: + url = f"{self._base_url(api_base)}/agents" + params: Dict[str, Any] = {} + if litellm_params.get("page_size"): + params["pageSize"] = litellm_params["page_size"] + if litellm_params.get("page_token"): + params["pageToken"] = litellm_params["page_token"] + return url, params + + def transform_list_response( + self, + raw_response: httpx.Response, + ) -> AgentListResponse: + self._raise_for_status(raw_response) + try: + data = raw_response.json() + except Exception: + data = {} + verbose_logger.debug("GeminiAgentsConfig list response: %s", data) + return AgentListResponse( + agents=data.get("agents", []), + next_page_token=data.get("nextPageToken"), + ) + + # ------------------------------------------------------------------ # + # GET # + # ------------------------------------------------------------------ # + + def transform_get_request( + self, + name: str, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> Tuple[str, Dict[str, Any]]: + url = f"{self._base_url(api_base)}/agents/{name}" + return url, {} + + def transform_get_response( + self, + raw_response: httpx.Response, + name: str, + ) -> AgentCreateResponse: + """Same shape as create response — Gemini returns "id" as identifier.""" + self._raise_for_status(raw_response) + try: + data = raw_response.json() + except Exception: + data = {"id": name} + data.setdefault("id", name) + data.setdefault("name", data["id"]) + verbose_logger.debug("GeminiAgentsConfig get response: %s", data) + return AgentCreateResponse(**data) + + # ------------------------------------------------------------------ # + # DELETE # + # ------------------------------------------------------------------ # + + def transform_delete_request( + self, + name: str, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> str: + return f"{self._base_url(api_base)}/agents/{name}" + + def transform_delete_response( + self, + raw_response: httpx.Response, + name: str, + ) -> AgentDeleteResult: + """Gemini returns an empty body ``{}`` with HTTP 200 on success.""" + self._raise_for_status(raw_response) + verbose_logger.debug( + "GeminiAgentsConfig delete (status=%d) agent '%s'", + raw_response.status_code, + name, + ) + return AgentDeleteResult(name=name, deleted=True) + + # ------------------------------------------------------------------ # + # LIST VERSIONS # + # ------------------------------------------------------------------ # + + def transform_list_versions_request( + self, + name: str, + api_base: Optional[str], + litellm_params: Dict[str, Any], + ) -> Tuple[str, Dict[str, Any]]: + url = f"{self._base_url(api_base)}/agents/{name}/versions" + params: Dict[str, Any] = {} + if litellm_params.get("page_size"): + params["pageSize"] = litellm_params["page_size"] + if litellm_params.get("page_token"): + params["pageToken"] = litellm_params["page_token"] + return url, params + + def transform_list_versions_response( + self, + raw_response: httpx.Response, + name: str, + ) -> AgentVersionsResponse: + """ + Gemini returns: + {"agentVersions": [{"agent": "waverunner", "name": "agents/.../versions/uuid", ...}]} + """ + self._raise_for_status(raw_response) + try: + data = raw_response.json() + except Exception: + data = {} + verbose_logger.debug( + "GeminiAgentsConfig list_versions response for '%s': %s", name, data + ) + return AgentVersionsResponse( + agent_versions=data.get("agentVersions", []), + next_page_token=data.get("nextPageToken"), + ) diff --git a/litellm/llms/gemini/interactions/transformation.py b/litellm/llms/gemini/interactions/transformation.py index 593cbf7c2cf..73435c8db6a 100644 --- a/litellm/llms/gemini/interactions/transformation.py +++ b/litellm/llms/gemini/interactions/transformation.py @@ -64,6 +64,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): "stream", "store", "background", + "environment", "response_modalities", "response_format", "response_mime_type", @@ -142,6 +143,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig): "stream", "store", "background", + "environment", "response_modalities", "response_format", "response_mime_type", diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 9f034575222..a70c5b3a920 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -84,6 +84,11 @@ LAZY_FEATURES: Tuple[LazyFeature, ...] = ( module_path="litellm.proxy.agent_endpoints.endpoints", path_prefixes=("/v1/agents", "/agents", "/agent/"), ), + LazyFeature( + name="gemini_agents", + module_path="litellm.proxy.google_endpoints.agents_endpoints", + path_prefixes=("/v1beta/agents",), + ), LazyFeature( name="a2a", module_path="litellm.proxy.agent_endpoints.a2a_endpoints", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index c989f5dff13..9337aa7c8ea 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -481,6 +481,10 @@ class LiteLLMRoutes(enum.Enum): "/v1beta/interactions/{interaction_id}", "/interactions/{interaction_id}/cancel", "/v1beta/interactions/{interaction_id}/cancel", + # Google Managed Agents API + "/v1beta/agents", + "/v1beta/agents/{name}", + "/v1beta/agents/{name}/versions", ] apply_guardrail_routes = [ diff --git a/litellm/proxy/agent_endpoints/utils.py b/litellm/proxy/agent_endpoints/utils.py index 2b968de54be..393f5934fd9 100644 --- a/litellm/proxy/agent_endpoints/utils.py +++ b/litellm/proxy/agent_endpoints/utils.py @@ -2,6 +2,12 @@ from typing import Dict, Mapping, Optional +# Re-export from the canonical SDK location so the proxy and SDK always +# share the same provider-config lookup logic. +from litellm.interactions.agents.utils import ( # noqa: F401 + get_provider_agents_api_config, +) + def merge_agent_headers( *, diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 038d2d81277..7d2954fd2dd 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -807,6 +807,11 @@ class ProxyBaseLLMRequestProcessing: "aget_interaction", "adelete_interaction", "acancel_interaction", + "acreate_agent", + "alist_agents", + "aget_agent", + "adelete_agent", + "alist_agent_versions", "asend_message", "call_mcp_tool", "acreate_eval", @@ -1074,6 +1079,11 @@ class ProxyBaseLLMRequestProcessing: "aget_interaction", "adelete_interaction", "acancel_interaction", + "acreate_agent", + "alist_agents", + "aget_agent", + "adelete_agent", + "alist_agent_versions", "asend_message", "call_mcp_tool", "acreate_eval", diff --git a/litellm/proxy/google_endpoints/agents_endpoints.py b/litellm/proxy/google_endpoints/agents_endpoints.py new file mode 100644 index 00000000000..779284023a0 --- /dev/null +++ b/litellm/proxy/google_endpoints/agents_endpoints.py @@ -0,0 +1,445 @@ +""" +Google AI Studio Managed Agents API Proxy Endpoints. + +Exposes Gemini's /v1beta/agents surface through the LiteLLM proxy so that +user curl commands transfer 1-to-1 by swapping the host + auth header. + +Routes: + POST /v1beta/agents -> acreate_agent + GET /v1beta/agents -> alist_agents + GET /v1beta/agents/{name} -> aget_agent + DELETE /v1beta/agents/{name} -> adelete_agent + GET /v1beta/agents/{name}/versions -> alist_agent_versions + +These are distinct from the A2A agent registry at /v1/agents. +""" + +import json + +from fastapi import APIRouter, Depends, HTTPException, Request, Response, status +from fastapi.responses import ORJSONResponse + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.proxy.common_utils.http_parsing_utils import ( + _read_request_body, + _safe_get_request_query_params, +) + +router = APIRouter(tags=["gemini managed agents"]) + + +def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool: + return ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN + or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value + ) + + +def _enforce_caller_supplied_provider_key( + data: dict, + user_api_key_dict: UserAPIKeyAuth, +) -> None: + """ + SECURITY: refuse to use the proxy's shared GOOGLE_API_KEY / GEMINI_API_KEY + env fallback for non-admin callers on Gemini managed-agent CRUD endpoints. + + These endpoints are part of ``llm_api_routes`` so any authenticated LLM key + can reach them, but unlike ``/v1beta/models/...:generateContent`` they are + *not* routed through ``model_list`` — the only credential source is either + the per-request ``litellm_params_template`` or the env var fallback. Without + this guard, any ordinary proxy user could list, create, or delete managed + agents inside the operator's Gemini project using the operator's key. + + Proxy admins (master key) keep the env-fallback convenience for ops use. + """ + if _is_proxy_admin(user_api_key_dict): + return + if data.get("api_key"): + return + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=( + "Gemini managed-agent endpoints require a caller-supplied " + "Gemini api_key (via 'litellm_params_template'). Falling back to " + "the proxy's GOOGLE_API_KEY / GEMINI_API_KEY env vars is only " + "permitted for proxy admins." + ), + ) + + +def _merge_query_params_into_data(data: dict, request: Request) -> dict: + """ + For GET/DELETE endpoints that cannot carry a JSON body, read a + JSON-encoded ``litellm_params_template`` query parameter and merge its + contents into *data*, without overwriting keys that are already present + (e.g. path params like ``name`` or the fixed ``custom_llm_provider``). + + This mirrors the ``litellm_params_template`` handling in + ``create_gemini_agent`` and is the supported way for multi-tenant + callers to supply per-request credentials on non-POST endpoints: + + .. code-block:: bash + + curl "http://localhost:4000/v1beta/agents?litellm_params_template=%7B%22api_key%22%3A%22AIza...%22%7D" \\ + -H "Authorization: Bearer sk-..." + + Credentials MUST NOT be passed as plain flat query parameters (e.g. + ``?api_key=AIza...``) because URL query strings appear verbatim in + web-server access logs, CDN edge logs, browser history, and Referer + headers. Use the ``litellm_params_template`` JSON body field on POST + requests, or the JSON-encoded query parameter above for GET/DELETE. + """ + query_params = _safe_get_request_query_params(request) + if not query_params: + return data + + raw_template = query_params.get("litellm_params_template") + if raw_template: + try: + template = ( + json.loads(raw_template) + if isinstance(raw_template, str) + else raw_template + ) + except (json.JSONDecodeError, ValueError): + template = {} + if isinstance(template, dict): + for key, value in template.items(): + data.setdefault(key, value) + + return data + + +def _proxy_server_imports(): + from litellm.proxy.proxy_server import ( # noqa: PLC0415 + general_settings, + llm_router, + proxy_config, + proxy_logging_obj, + select_data_generator, + user_api_base, + user_max_tokens, + user_model, + user_request_timeout, + user_temperature, + version, + ) + + return dict( + general_settings=general_settings, + llm_router=llm_router, + proxy_config=proxy_config, + proxy_logging_obj=proxy_logging_obj, + select_data_generator=select_data_generator, + user_api_base=user_api_base, + user_max_tokens=user_max_tokens, + user_model=user_model, + user_request_timeout=user_request_timeout, + user_temperature=user_temperature, + version=version, + ) + + +@router.post( + "/v1beta/agents", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, +) +async def create_gemini_agent( + request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Create a named custom agent on the Gemini side. + + Example: + ```bash + curl -X POST "http://localhost:4000/v1beta/agents" \\ + -H "Authorization: Bearer sk-..." \\ + -H "Content-Type: application/json" \\ + -d '{ + "name": "my-custom-slides-agent", + "base_agent": "waverunner", + "instructions": "You are a helpful assistant that creates slides.", + "base_environment": { + "type": "remote", + "sources": [ + {"type": "gcs", "source": "gs://eap-templates/slides-skill", + "target": "/.agents/skills/slides-skill"} + ] + } + }' + ``` + """ + srv = _proxy_server_imports() + data = await _read_request_body(request=request) + # Merge litellm_params_template (e.g. custom_llm_provider, api_key) into the request + litellm_params_template = data.pop("litellm_params_template", None) or {} + if isinstance(litellm_params_template, dict): + for key, value in litellm_params_template.items(): + if key not in data: + data[key] = value + data.setdefault("custom_llm_provider", "gemini") + _enforce_caller_supplied_provider_key(data, user_api_key_dict) + + processor = ProxyBaseLLMRequestProcessing(data=data) + try: + return await processor.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="acreate_agent", + proxy_logging_obj=srv["proxy_logging_obj"], + llm_router=srv["llm_router"], + general_settings=srv["general_settings"], + proxy_config=srv["proxy_config"], + select_data_generator=srv["select_data_generator"], + model=None, + user_model=srv["user_model"], + user_temperature=srv["user_temperature"], + user_request_timeout=srv["user_request_timeout"], + user_max_tokens=srv["user_max_tokens"], + user_api_base=srv["user_api_base"], + version=srv["version"], + ) + except Exception as e: + raise await processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=srv["proxy_logging_obj"], + version=srv["version"], + ) + + +@router.get( + "/v1beta/agents", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, +) +async def list_gemini_agents( + request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + List all custom agents on the Gemini side. + + Pass per-request Gemini credentials via the JSON-encoded + ``litellm_params_template`` query parameter. Flat query parameters + (e.g. ``?api_key=AIza...``) are intentionally ignored — see + ``_merge_query_params_into_data`` for the rationale. + + ```bash + curl "http://localhost:4000/v1beta/agents?litellm_params_template=%7B%22api_key%22%3A%22AIza...%22%7D" \\ + -H "Authorization: Bearer sk-..." + ``` + """ + srv = _proxy_server_imports() + data: dict = {"custom_llm_provider": "gemini"} + _merge_query_params_into_data(data, request) + _enforce_caller_supplied_provider_key(data, user_api_key_dict) + + processor = ProxyBaseLLMRequestProcessing(data=data) + try: + return await processor.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="alist_agents", + proxy_logging_obj=srv["proxy_logging_obj"], + llm_router=srv["llm_router"], + general_settings=srv["general_settings"], + proxy_config=srv["proxy_config"], + select_data_generator=srv["select_data_generator"], + model=None, + user_model=srv["user_model"], + user_temperature=srv["user_temperature"], + user_request_timeout=srv["user_request_timeout"], + user_max_tokens=srv["user_max_tokens"], + user_api_base=srv["user_api_base"], + version=srv["version"], + ) + except Exception as e: + raise await processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=srv["proxy_logging_obj"], + version=srv["version"], + ) + + +@router.get( + "/v1beta/agents/{name}", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, +) +async def get_gemini_agent( + request: Request, + name: str, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Get a specific custom agent by name. + + Pass per-request Gemini credentials via the JSON-encoded + ``litellm_params_template`` query parameter. Flat query parameters + (e.g. ``?api_key=AIza...``) are intentionally ignored — see + ``_merge_query_params_into_data`` for the rationale. + + ```bash + curl "http://localhost:4000/v1beta/agents/my-custom-slides-agent?litellm_params_template=%7B%22api_key%22%3A%22AIza...%22%7D" \\ + -H "Authorization: Bearer sk-..." + ``` + """ + srv = _proxy_server_imports() + data = {"name": name, "custom_llm_provider": "gemini"} + _merge_query_params_into_data(data, request) + _enforce_caller_supplied_provider_key(data, user_api_key_dict) + + processor = ProxyBaseLLMRequestProcessing(data=data) + try: + return await processor.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="aget_agent", + proxy_logging_obj=srv["proxy_logging_obj"], + llm_router=srv["llm_router"], + general_settings=srv["general_settings"], + proxy_config=srv["proxy_config"], + select_data_generator=srv["select_data_generator"], + model=None, + user_model=srv["user_model"], + user_temperature=srv["user_temperature"], + user_request_timeout=srv["user_request_timeout"], + user_max_tokens=srv["user_max_tokens"], + user_api_base=srv["user_api_base"], + version=srv["version"], + ) + except Exception as e: + raise await processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=srv["proxy_logging_obj"], + version=srv["version"], + ) + + +@router.delete( + "/v1beta/agents/{name}", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, +) +async def delete_gemini_agent( + request: Request, + name: str, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Delete a custom agent by name. + + Pass per-request Gemini credentials via the JSON-encoded + ``litellm_params_template`` query parameter. Flat query parameters + (e.g. ``?api_key=AIza...``) are intentionally ignored — see + ``_merge_query_params_into_data`` for the rationale. + + ```bash + curl -X DELETE "http://localhost:4000/v1beta/agents/my-custom-slides-agent?litellm_params_template=%7B%22api_key%22%3A%22AIza...%22%7D" \\ + -H "Authorization: Bearer sk-..." + ``` + """ + srv = _proxy_server_imports() + data = {"name": name, "custom_llm_provider": "gemini"} + _merge_query_params_into_data(data, request) + _enforce_caller_supplied_provider_key(data, user_api_key_dict) + + processor = ProxyBaseLLMRequestProcessing(data=data) + try: + return await processor.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="adelete_agent", + proxy_logging_obj=srv["proxy_logging_obj"], + llm_router=srv["llm_router"], + general_settings=srv["general_settings"], + proxy_config=srv["proxy_config"], + select_data_generator=srv["select_data_generator"], + model=None, + user_model=srv["user_model"], + user_temperature=srv["user_temperature"], + user_request_timeout=srv["user_request_timeout"], + user_max_tokens=srv["user_max_tokens"], + user_api_base=srv["user_api_base"], + version=srv["version"], + ) + except Exception as e: + raise await processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=srv["proxy_logging_obj"], + version=srv["version"], + ) + + +@router.get( + "/v1beta/agents/{name}/versions", + dependencies=[Depends(user_api_key_auth)], + response_class=ORJSONResponse, +) +async def list_gemini_agent_versions( + request: Request, + name: str, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + List versions of a custom agent. + + Pass per-request Gemini credentials via the JSON-encoded + ``litellm_params_template`` query parameter. Flat query parameters + (e.g. ``?api_key=AIza...``) are intentionally ignored — see + ``_merge_query_params_into_data`` for the rationale. + + ```bash + curl "http://localhost:4000/v1beta/agents/my-custom-slides-agent/versions?litellm_params_template=%7B%22api_key%22%3A%22AIza...%22%7D" \\ + -H "Authorization: Bearer sk-..." + ``` + """ + srv = _proxy_server_imports() + data = {"name": name, "custom_llm_provider": "gemini"} + _merge_query_params_into_data(data, request) + _enforce_caller_supplied_provider_key(data, user_api_key_dict) + + processor = ProxyBaseLLMRequestProcessing(data=data) + try: + return await processor.base_process_llm_request( + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + route_type="alist_agent_versions", + proxy_logging_obj=srv["proxy_logging_obj"], + llm_router=srv["llm_router"], + general_settings=srv["general_settings"], + proxy_config=srv["proxy_config"], + select_data_generator=srv["select_data_generator"], + model=None, + user_model=srv["user_model"], + user_temperature=srv["user_temperature"], + user_request_timeout=srv["user_request_timeout"], + user_max_tokens=srv["user_max_tokens"], + user_api_base=srv["user_api_base"], + version=srv["version"], + ) + except Exception as e: + raise await processor._handle_llm_api_exception( + e=e, + user_api_key_dict=user_api_key_dict, + proxy_logging_obj=srv["proxy_logging_obj"], + version=srv["version"], + ) diff --git a/litellm/proxy/google_endpoints/endpoints.py b/litellm/proxy/google_endpoints/endpoints.py index 967ac9f0ac4..1f503247bf4 100644 --- a/litellm/proxy/google_endpoints/endpoints.py +++ b/litellm/proxy/google_endpoints/endpoints.py @@ -285,7 +285,7 @@ async def create_interaction( general_settings=general_settings, proxy_config=proxy_config, select_data_generator=select_data_generator, - model=data.get("model") or data.get("agent"), + model=data.get("model"), user_model=user_model, user_temperature=user_temperature, user_request_timeout=user_request_timeout, diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 06fd35448f1..8f6f7084a0c 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -94,6 +94,12 @@ ROUTE_ENDPOINT_MAPPING = { "aget_interaction": "/interactions/{interaction_id}", "adelete_interaction": "/interactions/{interaction_id}", "acancel_interaction": "/interactions/{interaction_id}/cancel", + # Google Managed Agents API routes + "acreate_agent": "/v1beta/agents", + "alist_agents": "/v1beta/agents", + "aget_agent": "/v1beta/agents/{name}", + "adelete_agent": "/v1beta/agents/{name}", + "alist_agent_versions": "/v1beta/agents/{name}/versions", # OpenAI Evals API routes "acreate_eval": "/evals", "alist_evals": "/evals", @@ -311,6 +317,11 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin "aget_interaction", "adelete_interaction", "acancel_interaction", + "acreate_agent", + "alist_agents", + "aget_agent", + "adelete_agent", + "alist_agent_versions", "asend_message", "call_mcp_tool", "acancel_batch", @@ -468,6 +479,15 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin "acancel_interaction", ]: return getattr(llm_router, f"{route_type}")(**data) + # Managed Agents API: these don't need model routing + if route_type in [ + "acreate_agent", + "alist_agents", + "aget_agent", + "adelete_agent", + "alist_agent_versions", + ]: + return getattr(llm_router, f"{route_type}")(**data) if route_type in [ "avideo_list", "avideo_status", diff --git a/litellm/router.py b/litellm/router.py index 420c9b8a816..019f565e5c4 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -1579,6 +1579,44 @@ class Router: cancel_interaction, call_type="cancel_interaction" ) + def _initialize_managed_agents_endpoints(self): + """Initialize Google Managed Agents API endpoints (v1beta/agents).""" + from litellm.interactions.agents import acreate as acreate_agent + from litellm.interactions.agents import adelete as adelete_agent + from litellm.interactions.agents import aget as aget_agent + from litellm.interactions.agents import alist as alist_agents + from litellm.interactions.agents import alist_versions as alist_agent_versions + from litellm.interactions.agents import create as create_agent + from litellm.interactions.agents import delete as delete_agent + from litellm.interactions.agents import get as get_agent + from litellm.interactions.agents import list as list_agents + from litellm.interactions.agents import list_versions as list_agent_versions + + self.acreate_agent = self.factory_function( + acreate_agent, call_type="acreate_agent" + ) + self.create_agent = self.factory_function( + create_agent, call_type="create_agent" + ) + self.alist_agents = self.factory_function( + alist_agents, call_type="alist_agents" + ) + self.list_agents = self.factory_function(list_agents, call_type="list_agents") + self.aget_agent = self.factory_function(aget_agent, call_type="aget_agent") + self.get_agent = self.factory_function(get_agent, call_type="get_agent") + self.adelete_agent = self.factory_function( + adelete_agent, call_type="adelete_agent" + ) + self.delete_agent = self.factory_function( + delete_agent, call_type="delete_agent" + ) + self.alist_agent_versions = self.factory_function( + alist_agent_versions, call_type="alist_agent_versions" + ) + self.list_agent_versions = self.factory_function( + list_agent_versions, call_type="list_agent_versions" + ) + def _initialize_specialized_endpoints(self): """Helper to initialize specialized router endpoints (vector store, OCR, search, video, container, skills, interactions).""" self._initialize_vector_store_endpoints() @@ -1591,6 +1629,7 @@ class Router: self._initialize_container_endpoints() self._initialize_skills_endpoints() self._initialize_interactions_endpoints() + self._initialize_managed_agents_endpoints() def initialize_router_endpoints(self): self._initialize_core_endpoints() @@ -5322,6 +5361,16 @@ class Router: "delete_interaction", "acancel_interaction", "cancel_interaction", + "acreate_agent", + "create_agent", + "alist_agents", + "list_agents", + "aget_agent", + "get_agent", + "adelete_agent", + "delete_agent", + "alist_agent_versions", + "list_agent_versions", ] = "assistants", ): """ @@ -5406,6 +5455,27 @@ class Router: return vector_store_file_sync_wrapper + if call_type in ( + "create_agent", + "list_agents", + "get_agent", + "delete_agent", + "list_agent_versions", + ): + + def managed_agents_sync_wrapper( + custom_llm_provider: Optional[str] = None, + client: Optional[Any] = None, + **kwargs, + ): + if custom_llm_provider and "custom_llm_provider" not in kwargs: + kwargs["custom_llm_provider"] = custom_llm_provider + if "custom_llm_provider" not in kwargs: + kwargs["custom_llm_provider"] = "gemini" + return original_function(**kwargs) + + return managed_agents_sync_wrapper + # Handle asynchronous call types async def async_wrapper( custom_llm_provider: Optional[str] = None, @@ -5469,8 +5539,6 @@ class Router: "alist_skills", "aget_skill", "adelete_skill", - "acreate_interaction", - "create_interaction", ): return await self._ageneric_api_call_with_fallbacks( original_function=original_function, @@ -5530,6 +5598,8 @@ class Router: **kwargs, ) elif call_type in ( + "acreate_interaction", + "create_interaction", "aget_interaction", "adelete_interaction", "acancel_interaction", @@ -5539,6 +5609,18 @@ class Router: custom_llm_provider=custom_llm_provider, **kwargs, ) + elif call_type in ( + "acreate_agent", + "alist_agents", + "aget_agent", + "adelete_agent", + "alist_agent_versions", + ): + return await self._init_managed_agents_api_endpoints( + original_function=original_function, + custom_llm_provider=custom_llm_provider, + **kwargs, + ) return async_wrapper @@ -5643,6 +5725,34 @@ class Router: if custom_llm_provider and "custom_llm_provider" not in kwargs: kwargs["custom_llm_provider"] = custom_llm_provider # Default to gemini for interactions API + if "custom_llm_provider" not in kwargs: + kwargs["custom_llm_provider"] = "gemini" + # If the proxy accidentally passed agent name as model, clear it + if kwargs.get("agent") and kwargs.get("model") == kwargs.get("agent"): + kwargs["model"] = None + # Model-based interactions use deployment routing + fallbacks; agent-only calls + # must not enter model-group lookup (agent name is not a LiteLLM deployment). + if kwargs.get("model"): + return await self._ageneric_api_call_with_fallbacks( + original_function=original_function, + **kwargs, + ) + return await original_function(**kwargs) + + async def _init_managed_agents_api_endpoints( + self, + original_function: Callable, + custom_llm_provider: Optional[str] = None, + **kwargs, + ): + """ + Initialize the Managed Agents API endpoints on the router (v1beta/agents). + + CRUD operations for Gemini managed agents don't need model-based routing, + so we call the original function directly with the custom_llm_provider. + """ + if custom_llm_provider and "custom_llm_provider" not in kwargs: + kwargs["custom_llm_provider"] = custom_llm_provider if "custom_llm_provider" not in kwargs: kwargs["custom_llm_provider"] = "gemini" return await original_function(**kwargs) diff --git a/litellm/types/agents.py b/litellm/types/agents.py index efb2e73bfb5..8556b6bac93 100644 --- a/litellm/types/agents.py +++ b/litellm/types/agents.py @@ -228,6 +228,66 @@ class ListAgentsResponse(BaseModel): agents: List[AgentResponse] +class AgentCreateResponse(LiteLLMPydanticObjectBase): + """ + Response from a provider-side agent creation or get call (e.g. Gemini v1beta/agents). + + Gemini returns ``"id"`` as the agent identifier; we surface both ``id`` + (Gemini's value) and ``name`` (the user-supplied name, equal to ``id`` for + Gemini) so callers can use either. All extra fields returned by the + provider (e.g. ``base_agent``, ``system_instruction``, ``base_environment``) + are preserved via extra="allow". + """ + + id: Optional[str] = None + name: Optional[str] = None + model_config = {"extra": "allow"} + + _hidden_params: dict = PrivateAttr(default_factory=dict) + + +class AgentDeleteResult(LiteLLMPydanticObjectBase): + """Result of a provider-side agent deletion (e.g. Gemini DELETE /v1beta/agents/{name}). + + Gemini returns an empty body ``{}`` on success; we synthesise ``name`` and + ``deleted`` so callers always get a consistent response object. + """ + + name: str + deleted: bool = True + model_config = {"extra": "allow"} + + _hidden_params: dict = PrivateAttr(default_factory=dict) + + +class AgentListResponse(LiteLLMPydanticObjectBase): + """Response from listing agents on the provider side (e.g. Gemini GET /v1beta/agents). + + Gemini returns ``{"agents": [{"id": "..."}, ...]}``; each item is kept as + a plain dict so no fields are silently dropped. + """ + + agents: List[Dict[str, Any]] = [] + next_page_token: Optional[str] = None + model_config = {"extra": "allow"} + + _hidden_params: dict = PrivateAttr(default_factory=dict) + + +class AgentVersionsResponse(LiteLLMPydanticObjectBase): + """Response from listing versions of an agent (e.g. Gemini GET /v1beta/agents/{name}/versions). + + Gemini returns ``{"agentVersions": [...]}``; each version has a ``name`` + field of the form ``agents/{agent_id}/versions/{uuid}``. + """ + + agent_versions: List[Dict[str, Any]] = [] + next_page_token: Optional[str] = None + model_config = {"extra": "allow"} + + _hidden_params: dict = PrivateAttr(default_factory=dict) + + class AgentMakePublicResponse(BaseModel): message: str public_agent_groups: List[str] diff --git a/litellm/types/interactions/__init__.py b/litellm/types/interactions/__init__.py index a3acdc4cb1f..0f934fa0152 100644 --- a/litellm/types/interactions/__init__.py +++ b/litellm/types/interactions/__init__.py @@ -37,6 +37,7 @@ from litellm.types.interactions.generated import ( ImageContent, Interaction, InteractionEvent, + InteractionEnvironment, InteractionInput, InteractionsAPIOptionalRequestParams, InteractionsAPIResponse, @@ -115,6 +116,7 @@ __all__ = [ "ResponseModality", "Annotation", # LiteLLM types + "InteractionEnvironment", "InteractionInput", "InteractionsAPIResponse", "InteractionsAPIStreamingResponse", diff --git a/litellm/types/interactions/generated.py b/litellm/types/interactions/generated.py index ed626b0b7c8..2ce6331b448 100644 --- a/litellm/types/interactions/generated.py +++ b/litellm/types/interactions/generated.py @@ -1257,3 +1257,6 @@ class CancelInteractionResult(BaseLiteLLMOpenAIResponseObject): InteractionTool = Tool InteractionToolChoiceConfig = ToolChoiceConfig InteractionsAPIOptionalRequestParams = Dict[str, Any] + +# Agent interaction execution environment +InteractionEnvironment = Union[str, Dict[str, Any]] diff --git a/tests/proxy_unit_tests/test_gemini_agents_endpoints.py b/tests/proxy_unit_tests/test_gemini_agents_endpoints.py new file mode 100644 index 00000000000..bdac9348f71 --- /dev/null +++ b/tests/proxy_unit_tests/test_gemini_agents_endpoints.py @@ -0,0 +1,519 @@ +""" +Unit tests for litellm/proxy/google_endpoints/agents_endpoints.py + +Focus: verify that list_gemini_agents, get_gemini_agent, delete_gemini_agent, +and list_gemini_agent_versions correctly forward per-request credentials +(api_key, api_base, …) supplied via the JSON-encoded litellm_params_template +query parameter. Flat credential query params (e.g. ?api_key=…) are no +longer accepted — they would appear in server logs. +""" + +import json +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import Request +from fastapi.datastructures import Headers, QueryParams + +sys.path.insert(0, os.path.abspath("../..")) + +from litellm.proxy.google_endpoints.agents_endpoints import ( + _merge_query_params_into_data, +) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_request(query_string: str = "") -> MagicMock: + """Build a minimal mock Request whose query_params match *query_string*.""" + req = MagicMock(spec=Request) + req.query_params = QueryParams(query_string) + req.headers = Headers({}) + return req + + +# --------------------------------------------------------------------------- +# _merge_query_params_into_data – unit tests for the helper +# --------------------------------------------------------------------------- + + +class TestMergeQueryParamsIntoData: + def test_no_query_params_leaves_data_unchanged(self): + data = {"custom_llm_provider": "gemini"} + request = _make_request("") + result = _merge_query_params_into_data(data, request) + assert result == {"custom_llm_provider": "gemini"} + + def test_flat_api_key_is_ignored(self): + """Flat credential params must NOT be merged (they leak into server logs).""" + data = {"custom_llm_provider": "gemini"} + request = _make_request("api_key=AIzaSyTest123") + _merge_query_params_into_data(data, request) + assert "api_key" not in data + assert data["custom_llm_provider"] == "gemini" + + def test_flat_params_are_silently_dropped(self): + """Flat params (including name injection attempts) are ignored entirely.""" + data = {"name": "my-agent", "custom_llm_provider": "gemini"} + request = _make_request("name=INJECTED&api_key=AIzaSyTest") + _merge_query_params_into_data(data, request) + assert data["name"] == "my-agent" + assert "api_key" not in data + + def test_litellm_params_template_json_is_expanded(self): + template = json.dumps( + {"api_key": "AIzaFromTemplate", "api_base": "https://example.com"} + ) + from urllib.parse import quote + + request = _make_request(f"litellm_params_template={quote(template)}") + data = {"custom_llm_provider": "gemini"} + _merge_query_params_into_data(data, request) + assert data["api_key"] == "AIzaFromTemplate" + assert data["api_base"] == "https://example.com" + # The raw template key itself must NOT appear in data + assert "litellm_params_template" not in data + + def test_litellm_params_template_does_not_overwrite_existing(self): + template = json.dumps( + {"api_key": "FromTemplate", "custom_llm_provider": "openai"} + ) + from urllib.parse import quote + + request = _make_request(f"litellm_params_template={quote(template)}") + data = {"custom_llm_provider": "gemini"} + _merge_query_params_into_data(data, request) + # custom_llm_provider was already set; template must not override it + assert data["custom_llm_provider"] == "gemini" + assert data["api_key"] == "FromTemplate" + + def test_invalid_litellm_params_template_json_is_ignored(self): + request = _make_request("litellm_params_template=NOT_VALID_JSON") + data = {"custom_llm_provider": "gemini"} + _merge_query_params_into_data(data, request) + # Bad JSON is silently skipped; other data stays intact + assert data == {"custom_llm_provider": "gemini"} + + def test_template_only_no_flat_params_merged(self): + """Only litellm_params_template is expanded; unknown flat params are dropped.""" + template = json.dumps({"api_key": "FromTemplate"}) + from urllib.parse import quote + + qs = f"litellm_params_template={quote(template)}&vertex_project=my-project" + request = _make_request(qs) + data = {"custom_llm_provider": "gemini"} + _merge_query_params_into_data(data, request) + assert data["api_key"] == "FromTemplate" + # flat vertex_project is ignored since it wasn't in litellm_params_template + assert "vertex_project" not in data + assert "litellm_params_template" not in data + + +# --------------------------------------------------------------------------- +# Endpoint-level smoke tests: data dict is populated before the processor call +# --------------------------------------------------------------------------- + + +@pytest.fixture +def mock_srv(): + """Patch _proxy_server_imports to return lightweight fakes.""" + srv = { + "general_settings": {}, + "llm_router": MagicMock(), + "proxy_config": MagicMock(), + "proxy_logging_obj": MagicMock(), + "select_data_generator": MagicMock(), + "user_api_base": None, + "user_max_tokens": None, + "user_model": None, + "user_request_timeout": None, + "user_temperature": None, + "version": "0.0.0", + } + with patch( + "litellm.proxy.google_endpoints.agents_endpoints._proxy_server_imports", + return_value=srv, + ): + yield srv + + +@pytest.fixture +def user_api_key_dict(): + from litellm.proxy._types import UserAPIKeyAuth + + return UserAPIKeyAuth(api_key="test-key") + + +def _make_endpoint_request(query_string: str = "") -> MagicMock: + req = MagicMock(spec=Request) + req.query_params = QueryParams(query_string) + req.headers = Headers({}) + req.scope = {} + + async def _body(): + return b"" + + req.body = _body + return req + + +@pytest.mark.asyncio +async def test_list_gemini_agents_passes_api_key_to_processor( + mock_srv, user_api_key_dict +): + from urllib.parse import quote + + from litellm.proxy.google_endpoints.agents_endpoints import list_gemini_agents + + template = json.dumps({"api_key": "AIzaListTest"}) + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request(f"litellm_params_template={quote(template)}") + await list_gemini_agents( + request=request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + init_data = MockProcessor.call_args[1]["data"] + assert init_data.get("api_key") == "AIzaListTest" + assert init_data.get("custom_llm_provider") == "gemini" + + +@pytest.mark.asyncio +async def test_get_gemini_agent_passes_api_key_to_processor( + mock_srv, user_api_key_dict +): + from urllib.parse import quote + + from litellm.proxy.google_endpoints.agents_endpoints import get_gemini_agent + + template = json.dumps({"api_key": "AIzaGetTest"}) + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request(f"litellm_params_template={quote(template)}") + await get_gemini_agent( + request=request, + name="my-agent", + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + init_data = MockProcessor.call_args[1]["data"] + assert init_data.get("api_key") == "AIzaGetTest" + assert init_data.get("name") == "my-agent" + assert init_data.get("custom_llm_provider") == "gemini" + + +@pytest.mark.asyncio +async def test_delete_gemini_agent_passes_api_key_to_processor( + mock_srv, user_api_key_dict +): + from urllib.parse import quote + + from litellm.proxy.google_endpoints.agents_endpoints import delete_gemini_agent + + template = json.dumps({"api_key": "AIzaDeleteTest"}) + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request(f"litellm_params_template={quote(template)}") + await delete_gemini_agent( + request=request, + name="my-agent", + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + init_data = MockProcessor.call_args[1]["data"] + assert init_data.get("api_key") == "AIzaDeleteTest" + assert init_data.get("name") == "my-agent" + assert init_data.get("custom_llm_provider") == "gemini" + + +@pytest.mark.asyncio +async def test_list_gemini_agent_versions_passes_api_key_to_processor( + mock_srv, user_api_key_dict +): + from urllib.parse import quote + + from litellm.proxy.google_endpoints.agents_endpoints import ( + list_gemini_agent_versions, + ) + + template = json.dumps({"api_key": "AIzaVersionsTest"}) + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request(f"litellm_params_template={quote(template)}") + await list_gemini_agent_versions( + request=request, + name="my-agent", + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + init_data = MockProcessor.call_args[1]["data"] + assert init_data.get("api_key") == "AIzaVersionsTest" + assert init_data.get("name") == "my-agent" + assert init_data.get("custom_llm_provider") == "gemini" + + +@pytest.mark.asyncio +async def test_get_gemini_agent_name_not_overwritten_by_query_param( + mock_srv, user_api_key_dict +): + """Path-param ``name`` must not be replaced by an attacker-controlled query param.""" + from urllib.parse import quote + + from litellm.proxy.google_endpoints.agents_endpoints import get_gemini_agent + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + # Even if a caller tries to inject "name" via flat query param, it is + # ignored (flat params are not merged). The path-param name wins. + # ``api_key`` is supplied via the JSON template (required for non-admin + # callers — see test_*_non_admin_without_api_key_is_rejected below). + template = json.dumps({"api_key": "AIzaTest"}) + request = _make_endpoint_request( + f"name=INJECTED&litellm_params_template={quote(template)}" + ) + await get_gemini_agent( + request=request, + name="real-agent", + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + init_data = MockProcessor.call_args[1]["data"] + assert init_data["name"] == "real-agent" + + +@pytest.mark.asyncio +async def test_list_agents_template_via_query_param(mock_srv, user_api_key_dict): + """litellm_params_template in query string is expanded.""" + from litellm.proxy.google_endpoints.agents_endpoints import list_gemini_agents + from urllib.parse import quote + + template = json.dumps({"api_key": "TemplateKey", "vertex_project": "proj-x"}) + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request(f"litellm_params_template={quote(template)}") + await list_gemini_agents( + request=request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + + init_data = MockProcessor.call_args[1]["data"] + assert init_data["api_key"] == "TemplateKey" + assert init_data["vertex_project"] == "proj-x" + assert "litellm_params_template" not in init_data + + +# --------------------------------------------------------------------------- +# Security guards (veria-flagged findings) +# --------------------------------------------------------------------------- + + +@pytest.fixture +def proxy_admin_user_api_key_dict(): + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + + return UserAPIKeyAuth( + api_key="sk-admin", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + +@pytest.mark.asyncio +async def test_list_agents_non_admin_without_api_key_is_rejected( + mock_srv, user_api_key_dict +): + """Non-admin callers must supply an explicit api_key — the proxy must not + silently fall back to the operator's shared GOOGLE_API_KEY/GEMINI_API_KEY. + """ + from fastapi import HTTPException + + from litellm.proxy.google_endpoints.agents_endpoints import list_gemini_agents + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request("") + with pytest.raises(HTTPException) as excinfo: + await list_gemini_agents( + request=request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + assert excinfo.value.status_code == 401 + # Processor must never be invoked + instance.base_process_llm_request.assert_not_called() + + +@pytest.mark.asyncio +async def test_delete_agent_non_admin_without_api_key_is_rejected( + mock_srv, user_api_key_dict +): + from fastapi import HTTPException + + from litellm.proxy.google_endpoints.agents_endpoints import delete_gemini_agent + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request("") + with pytest.raises(HTTPException) as excinfo: + await delete_gemini_agent( + request=request, + name="my-agent", + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + assert excinfo.value.status_code == 401 + instance.base_process_llm_request.assert_not_called() + + +@pytest.mark.asyncio +async def test_create_agent_non_admin_without_api_key_is_rejected( + mock_srv, user_api_key_dict +): + from fastapi import HTTPException + + from litellm.proxy.google_endpoints.agents_endpoints import create_gemini_agent + + with ( + patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor, + patch( + "litellm.proxy.google_endpoints.agents_endpoints._read_request_body", + new=AsyncMock(return_value={"name": "agent-1", "base_agent": "waverunner"}), + ), + ): + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request("") + with pytest.raises(HTTPException) as excinfo: + await create_gemini_agent( + request=request, + fastapi_response=MagicMock(), + user_api_key_dict=user_api_key_dict, + ) + assert excinfo.value.status_code == 401 + instance.base_process_llm_request.assert_not_called() + + +@pytest.mark.asyncio +async def test_list_agents_proxy_admin_may_use_env_fallback( + mock_srv, proxy_admin_user_api_key_dict +): + """Proxy admins (master key) keep the env-fallback convenience.""" + from litellm.proxy.google_endpoints.agents_endpoints import list_gemini_agents + + with patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing" + ) as MockProcessor: + instance = MockProcessor.return_value + instance.base_process_llm_request = AsyncMock(return_value=MagicMock()) + + request = _make_endpoint_request("") + await list_gemini_agents( + request=request, + fastapi_response=MagicMock(), + user_api_key_dict=proxy_admin_user_api_key_dict, + ) + + init_data = MockProcessor.call_args[1]["data"] + assert "api_key" not in init_data + instance.base_process_llm_request.assert_awaited_once() + + +def test_validate_environment_rejects_api_base_override_without_explicit_key( + monkeypatch, +): + """SECURITY: caller-supplied api_base must be paired with an explicit + api_key — otherwise the proxy's shared GOOGLE_API_KEY leaks to the + attacker-controlled host via the x-goog-api-key header. + """ + from litellm.llms.gemini.agents.transformation import GeminiAgentsConfig + + # Even if env-fallback is available, api_base override must require api_key. + monkeypatch.setenv("GOOGLE_API_KEY", "AIzaSharedSecret") + + cfg = GeminiAgentsConfig() + with pytest.raises(ValueError, match="api_base"): + cfg.validate_environment( + headers={}, + litellm_params={"api_base": "https://attacker.example"}, + ) + + +def test_validate_environment_allows_api_base_with_explicit_key(monkeypatch): + """api_base override is OK when paired with an explicit api_key.""" + from litellm.llms.gemini.agents.transformation import GeminiAgentsConfig + + monkeypatch.delenv("GOOGLE_API_KEY", raising=False) + monkeypatch.delenv("GEMINI_API_KEY", raising=False) + + cfg = GeminiAgentsConfig() + headers = cfg.validate_environment( + headers={}, + litellm_params={ + "api_base": "https://my-gemini-proxy.example", + "api_key": "AIzaCallerOwned", + }, + ) + assert headers["x-goog-api-key"] == "AIzaCallerOwned" + + +def test_validate_environment_env_fallback_when_no_api_base_override(monkeypatch): + """Without api_base override, env fallback continues to work for SDK use.""" + from litellm.llms.gemini.agents.transformation import GeminiAgentsConfig + + monkeypatch.setenv("GOOGLE_API_KEY", "AIzaFromEnv") + monkeypatch.delenv("GEMINI_API_KEY", raising=False) + + cfg = GeminiAgentsConfig() + headers = cfg.validate_environment(headers={}, litellm_params={}) + assert headers["x-goog-api-key"] == "AIzaFromEnv" diff --git a/tests/test_litellm/interactions/test_agents_http_handler.py b/tests/test_litellm/interactions/test_agents_http_handler.py new file mode 100644 index 00000000000..6947503e0bb --- /dev/null +++ b/tests/test_litellm/interactions/test_agents_http_handler.py @@ -0,0 +1,587 @@ +""" +Unit tests for litellm/interactions/agents/http_handler.py + +These tests exercise both the sync and async branches of every CRUD method +on AgentsHTTPHandler using stub httpx clients, plus the _is_async dispatch +branches, error mapping, and pre/post logging hooks. + +No real HTTP traffic is made. +""" + +import os +import sys +from unittest.mock import AsyncMock, MagicMock + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../..")) + +from litellm.interactions.agents.http_handler import ( + AgentsHTTPHandler, + agents_http_handler, +) +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.llms.gemini.agents.transformation import GeminiAgentsConfig +from litellm.llms.gemini.common_utils import GeminiError +from litellm.types.agents import ( + AgentCreateResponse, + AgentDeleteResult, + AgentListResponse, + AgentVersionsResponse, +) +from litellm.types.router import GenericLiteLLMParams + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_response(status_code: int = 200, json_data=None, text: str = "") -> MagicMock: + """Build a stub httpx-like response.""" + response = MagicMock() + response.status_code = status_code + response.text = text or (str(json_data) if json_data is not None else "") + response.headers = {} + if json_data is not None: + response.json.return_value = json_data + else: + response.json.return_value = {} + return response + + +def _make_sync_client() -> MagicMock: + client = MagicMock(spec=HTTPHandler) + return client + + +def _make_async_client() -> MagicMock: + client = MagicMock(spec=AsyncHTTPHandler) + client.post = AsyncMock() + client.get = AsyncMock() + client.delete = AsyncMock() + return client + + +def _make_logging_obj() -> MagicMock: + return MagicMock() + + +@pytest.fixture +def handler() -> AgentsHTTPHandler: + return AgentsHTTPHandler() + + +@pytest.fixture +def config() -> GeminiAgentsConfig: + return GeminiAgentsConfig() + + +@pytest.fixture +def litellm_params() -> GenericLiteLLMParams: + return GenericLiteLLMParams(api_key="AIza-test") + + +# --------------------------------------------------------------------------- +# Module-level singleton sanity check +# --------------------------------------------------------------------------- + + +def test_module_singleton_is_agents_http_handler_instance(): + assert isinstance(agents_http_handler, AgentsHTTPHandler) + + +# --------------------------------------------------------------------------- +# CREATE +# --------------------------------------------------------------------------- + + +class TestCreateAgent: + def test_sync_returns_parsed_create_response(self, handler, config, litellm_params): + client = _make_sync_client() + client.post.return_value = _make_response( + 200, json_data={"id": "agent-x", "base_agent": "gemini-2.5-flash"} + ) + logging_obj = _make_logging_obj() + + result = handler.create_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=logging_obj, + extra_headers={"X-Test": "1"}, + extra_body={"foo": "bar"}, + client=client, + ) + + assert isinstance(result, AgentCreateResponse) + assert result.id == "agent-x" + client.post.assert_called_once() + kwargs = client.post.call_args.kwargs + assert kwargs["url"].endswith("/v1beta/agents") + assert kwargs["json"]["name"] == "agent-x" + assert kwargs["json"]["foo"] == "bar" + assert kwargs["headers"]["X-Test"] == "1" + logging_obj.pre_call.assert_called_once() + logging_obj.post_call.assert_called_once() + + def test_sync_dispatches_to_async_when_is_async( + self, handler, config, litellm_params + ): + client = _make_sync_client() + + result = handler.create_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + _is_async=True, + ) + + import asyncio + + assert asyncio.iscoroutine(result) + result.close() + + def test_sync_maps_http_error_via_config(self, handler, config, litellm_params): + client = _make_sync_client() + bad = _make_response(404, text="not found") + client.post.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + handler.create_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + @pytest.mark.asyncio + async def test_async_returns_parsed_create_response( + self, handler, config, litellm_params + ): + client = _make_async_client() + client.post.return_value = _make_response( + 200, json_data={"id": "agent-y", "base_agent": "gemini-2.5-flash"} + ) + + result = await handler.async_create_agent( + agents_api_config=config, + name="agent-y", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + extra_body={"baz": "qux"}, + client=client, + ) + + assert isinstance(result, AgentCreateResponse) + assert result.id == "agent-y" + client.post.assert_awaited_once() + + @pytest.mark.asyncio + async def test_async_maps_http_error_via_config( + self, handler, config, litellm_params + ): + client = _make_async_client() + bad = _make_response(500, text="server error") + client.post.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + await handler.async_create_agent( + agents_api_config=config, + name="agent-y", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + +# --------------------------------------------------------------------------- +# LIST +# --------------------------------------------------------------------------- + + +class TestListAgents: + def test_sync_returns_list_response(self, handler, config, litellm_params): + client = _make_sync_client() + client.get.return_value = _make_response( + 200, + json_data={ + "agents": [{"id": "a-1"}, {"id": "a-2"}], + "nextPageToken": "tok", + }, + ) + + result = handler.list_agents( + agents_api_config=config, + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + assert isinstance(result, AgentListResponse) + assert len(result.agents) == 2 + assert result.next_page_token == "tok" + client.get.assert_called_once() + + def test_sync_dispatches_to_async_when_is_async( + self, handler, config, litellm_params + ): + client = _make_sync_client() + + result = handler.list_agents( + agents_api_config=config, + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + _is_async=True, + ) + + import asyncio + + assert asyncio.iscoroutine(result) + result.close() + + def test_sync_maps_http_error_via_config(self, handler, config, litellm_params): + client = _make_sync_client() + bad = _make_response(403, text="forbidden") + client.get.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + handler.list_agents( + agents_api_config=config, + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + @pytest.mark.asyncio + async def test_async_returns_list_response(self, handler, config, litellm_params): + client = _make_async_client() + client.get.return_value = _make_response( + 200, json_data={"agents": [{"id": "a-1"}]} + ) + + result = await handler.async_list_agents( + agents_api_config=config, + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + assert isinstance(result, AgentListResponse) + assert len(result.agents) == 1 + client.get.assert_awaited_once() + + @pytest.mark.asyncio + async def test_async_maps_http_error_via_config( + self, handler, config, litellm_params + ): + client = _make_async_client() + bad = _make_response(429, text="rate limited") + client.get.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + await handler.async_list_agents( + agents_api_config=config, + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + +# --------------------------------------------------------------------------- +# GET +# --------------------------------------------------------------------------- + + +class TestGetAgent: + def test_sync_returns_get_response(self, handler, config, litellm_params): + client = _make_sync_client() + client.get.return_value = _make_response(200, json_data={"id": "agent-x"}) + + result = handler.get_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + assert isinstance(result, AgentCreateResponse) + assert result.id == "agent-x" + kwargs = client.get.call_args.kwargs + assert kwargs["url"].endswith("/v1beta/agents/agent-x") + + def test_sync_dispatches_to_async_when_is_async( + self, handler, config, litellm_params + ): + result = handler.get_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=_make_sync_client(), + _is_async=True, + ) + import asyncio + + assert asyncio.iscoroutine(result) + result.close() + + def test_sync_maps_http_error_via_config(self, handler, config, litellm_params): + client = _make_sync_client() + bad = _make_response(404, text="not found") + client.get.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + handler.get_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + @pytest.mark.asyncio + async def test_async_returns_get_response(self, handler, config, litellm_params): + client = _make_async_client() + client.get.return_value = _make_response(200, json_data={"id": "agent-y"}) + + result = await handler.async_get_agent( + agents_api_config=config, + name="agent-y", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + assert isinstance(result, AgentCreateResponse) + assert result.id == "agent-y" + + @pytest.mark.asyncio + async def test_async_maps_http_error_via_config( + self, handler, config, litellm_params + ): + client = _make_async_client() + bad = _make_response(404, text="not found") + client.get.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + await handler.async_get_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + +# --------------------------------------------------------------------------- +# DELETE +# --------------------------------------------------------------------------- + + +class TestDeleteAgent: + def test_sync_returns_delete_result(self, handler, config, litellm_params): + client = _make_sync_client() + client.delete.return_value = _make_response(200, json_data={}) + + result = handler.delete_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + assert isinstance(result, AgentDeleteResult) + assert result.name == "agent-x" + assert result.deleted is True + kwargs = client.delete.call_args.kwargs + assert kwargs["url"].endswith("/v1beta/agents/agent-x") + + def test_sync_dispatches_to_async_when_is_async( + self, handler, config, litellm_params + ): + result = handler.delete_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=_make_sync_client(), + _is_async=True, + ) + import asyncio + + assert asyncio.iscoroutine(result) + result.close() + + def test_sync_maps_http_error_via_config(self, handler, config, litellm_params): + client = _make_sync_client() + bad = _make_response(403, text="forbidden") + client.delete.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + handler.delete_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + @pytest.mark.asyncio + async def test_async_returns_delete_result(self, handler, config, litellm_params): + client = _make_async_client() + client.delete.return_value = _make_response(200, json_data={}) + + result = await handler.async_delete_agent( + agents_api_config=config, + name="agent-y", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + assert isinstance(result, AgentDeleteResult) + assert result.name == "agent-y" + assert result.deleted is True + + @pytest.mark.asyncio + async def test_async_maps_http_error_via_config( + self, handler, config, litellm_params + ): + client = _make_async_client() + bad = _make_response(500, text="server error") + client.delete.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + await handler.async_delete_agent( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + +# --------------------------------------------------------------------------- +# LIST VERSIONS +# --------------------------------------------------------------------------- + + +class TestListAgentVersions: + def test_sync_returns_versions_response(self, handler, config, litellm_params): + client = _make_sync_client() + client.get.return_value = _make_response( + 200, + json_data={ + "agentVersions": [ + {"agent": "agent-x", "name": "agents/agent-x/versions/v1"} + ], + "nextPageToken": "tok", + }, + ) + + result = handler.list_agent_versions( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + assert isinstance(result, AgentVersionsResponse) + assert len(result.agent_versions) == 1 + assert result.next_page_token == "tok" + kwargs = client.get.call_args.kwargs + assert kwargs["url"].endswith("/v1beta/agents/agent-x/versions") + + def test_sync_dispatches_to_async_when_is_async( + self, handler, config, litellm_params + ): + result = handler.list_agent_versions( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=_make_sync_client(), + _is_async=True, + ) + import asyncio + + assert asyncio.iscoroutine(result) + result.close() + + def test_sync_maps_http_error_via_config(self, handler, config, litellm_params): + client = _make_sync_client() + bad = _make_response(404, text="not found") + client.get.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + handler.list_agent_versions( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + @pytest.mark.asyncio + async def test_async_returns_versions_response( + self, handler, config, litellm_params + ): + client = _make_async_client() + client.get.return_value = _make_response(200, json_data={"agentVersions": []}) + + result = await handler.async_list_agent_versions( + agents_api_config=config, + name="agent-y", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) + + assert isinstance(result, AgentVersionsResponse) + assert result.agent_versions == [] + + @pytest.mark.asyncio + async def test_async_maps_http_error_via_config( + self, handler, config, litellm_params + ): + client = _make_async_client() + bad = _make_response(500, text="server error") + client.get.side_effect = httpx.HTTPStatusError( + "boom", request=MagicMock(), response=bad + ) + + with pytest.raises(GeminiError): + await handler.async_list_agent_versions( + agents_api_config=config, + name="agent-x", + litellm_params=litellm_params, + logging_obj=_make_logging_obj(), + client=client, + ) diff --git a/tests/test_litellm/interactions/test_agents_main_and_utils.py b/tests/test_litellm/interactions/test_agents_main_and_utils.py new file mode 100644 index 00000000000..7c0183d20c6 --- /dev/null +++ b/tests/test_litellm/interactions/test_agents_main_and_utils.py @@ -0,0 +1,354 @@ +""" +Unit tests for litellm/interactions/agents/utils.py and main.py +focused on the managed agents SDK surface added in the +"Gemini managed agents support" PR. + +The tests mock the underlying HTTP handler so they cover the public +sync + async create/list/get/delete/list_versions entry points and the +small helper utilities without touching the network. +""" + +import asyncio +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../..")) + +import litellm +from litellm.interactions.agents import ( + acreate, + adelete, + aget, + alist, + alist_versions, + create, + delete, + get, + list as list_agents, + list_versions, +) +from litellm.interactions.agents.main import ( + _get_agents_api_config, + _make_logging_obj, +) +from litellm.interactions.agents.utils import get_provider_agents_api_config +from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig +from litellm.llms.gemini.agents.transformation import GeminiAgentsConfig + + +_HANDLER_PATH = "litellm.interactions.agents.main.agents_http_handler" + + +# --------------------------------------------------------------------------- +# utils.get_provider_agents_api_config +# --------------------------------------------------------------------------- + + +class TestGetProviderAgentsApiConfig: + def test_returns_gemini_config_for_gemini(self): + cfg = get_provider_agents_api_config("gemini") + assert isinstance(cfg, GeminiAgentsConfig) + assert isinstance(cfg, BaseAgentsAPIConfig) + + @pytest.mark.parametrize( + "provider", ["openai", "anthropic", "bedrock", "vertex_ai", "unknown"] + ) + def test_returns_none_for_non_gemini(self, provider): + assert get_provider_agents_api_config(provider) is None + + def test_returns_none_for_none(self): + assert get_provider_agents_api_config(None) is None + + +# --------------------------------------------------------------------------- +# main._get_agents_api_config +# --------------------------------------------------------------------------- + + +class TestGetAgentsApiConfig: + def test_returns_config_for_gemini(self): + cfg = _get_agents_api_config("gemini") + assert isinstance(cfg, GeminiAgentsConfig) + + def test_raises_bad_request_for_unsupported_provider(self): + with pytest.raises(litellm.BadRequestError) as excinfo: + _get_agents_api_config("openai") + assert "does not have a native" in str(excinfo.value) + + +# --------------------------------------------------------------------------- +# main._make_logging_obj +# --------------------------------------------------------------------------- + + +class TestMakeLoggingObj: + def test_calls_update_from_kwargs_and_returns_same_obj(self): + logging_obj = MagicMock() + kwargs = {"litellm_logging_obj": logging_obj, "litellm_call_id": "abc-123"} + + returned = _make_logging_obj( + kwargs=kwargs, + model="my-agent", + custom_llm_provider="gemini", + call_type="create_agent", + optional_params={"foo": "bar"}, + ) + + assert returned is logging_obj + logging_obj.update_from_kwargs.assert_called_once() + kwargs_call = logging_obj.update_from_kwargs.call_args.kwargs + assert kwargs_call["model"] == "my-agent" + assert kwargs_call["optional_params"] == {"foo": "bar"} + assert kwargs_call["custom_llm_provider"] == "gemini" + assert kwargs_call["litellm_params"]["litellm_call_id"] == "abc-123" + + +# --------------------------------------------------------------------------- +# Sync entry points: create / list / get / delete / list_versions +# --------------------------------------------------------------------------- + + +def _stub_handler(return_value): + """Build a stub AgentsHTTPHandler whose CRUD methods return *return_value*.""" + handler = MagicMock() + handler.create_agent.return_value = return_value + handler.list_agents.return_value = return_value + handler.get_agent.return_value = return_value + handler.delete_agent.return_value = return_value + handler.list_agent_versions.return_value = return_value + return handler + + +class TestSyncEntryPoints: + def test_create_passes_args_to_handler(self): + sentinel = MagicMock(name="create_response") + with patch(_HANDLER_PATH, _stub_handler(sentinel)) as handler: + response = create( + name="waverunner", + base_agent="gemini-2.5-flash", + instructions="be helpful", + base_environment={"type": "remote"}, + custom_llm_provider="gemini", + api_key="AIza-test", + extra_headers={"X-Test": "1"}, + extra_body={"foo": "bar"}, + ) + + assert response is sentinel + handler.create_agent.assert_called_once() + kw = handler.create_agent.call_args.kwargs + assert kw["name"] == "waverunner" + assert kw["_is_async"] is False + assert kw["extra_headers"] == {"X-Test": "1"} + assert kw["extra_body"] == {"foo": "bar"} + assert isinstance(kw["agents_api_config"], GeminiAgentsConfig) + + def test_create_defaults_custom_llm_provider_to_gemini(self): + sentinel = MagicMock(name="create_response") + with patch(_HANDLER_PATH, _stub_handler(sentinel)) as handler: + create(name="agent-x", api_key="AIza") + assert handler.create_agent.call_args.kwargs["_is_async"] is False + cfg = handler.create_agent.call_args.kwargs["agents_api_config"] + assert isinstance(cfg, GeminiAgentsConfig) + + def test_create_raises_for_unsupported_provider(self): + with pytest.raises(litellm.exceptions.BadRequestError): + create(name="agent-x", custom_llm_provider="openai", api_key="sk-x") + + def test_list_passes_args_to_handler(self): + sentinel = MagicMock(name="list_response") + with patch(_HANDLER_PATH, _stub_handler(sentinel)) as handler: + response = list_agents(custom_llm_provider="gemini", api_key="AIza") + assert response is sentinel + handler.list_agents.assert_called_once() + assert handler.list_agents.call_args.kwargs["_is_async"] is False + + def test_get_passes_args_to_handler(self): + sentinel = MagicMock(name="get_response") + with patch(_HANDLER_PATH, _stub_handler(sentinel)) as handler: + response = get(name="waverunner", api_key="AIza") + assert response is sentinel + kw = handler.get_agent.call_args.kwargs + assert kw["name"] == "waverunner" + assert kw["_is_async"] is False + + def test_delete_passes_args_to_handler(self): + sentinel = MagicMock(name="delete_response") + with patch(_HANDLER_PATH, _stub_handler(sentinel)) as handler: + response = delete(name="waverunner", api_key="AIza") + assert response is sentinel + kw = handler.delete_agent.call_args.kwargs + assert kw["name"] == "waverunner" + assert kw["_is_async"] is False + + def test_list_versions_passes_args_to_handler(self): + sentinel = MagicMock(name="versions_response") + with patch(_HANDLER_PATH, _stub_handler(sentinel)) as handler: + response = list_versions(name="waverunner", api_key="AIza") + assert response is sentinel + kw = handler.list_agent_versions.call_args.kwargs + assert kw["name"] == "waverunner" + assert kw["_is_async"] is False + + +# --------------------------------------------------------------------------- +# Async entry points +# --------------------------------------------------------------------------- + + +class TestAsyncEntryPoints: + """Async entry points delegate to their sync counterparts via run_in_executor.""" + + @pytest.mark.asyncio + async def test_acreate_dispatches_with_async_flag(self): + sentinel = MagicMock(name="acreate_response") + + def fake_create_agent(**kwargs): + assert kwargs["_is_async"] is True + assert kwargs["name"] == "waverunner" + return sentinel + + handler = MagicMock() + handler.create_agent.side_effect = fake_create_agent + + with patch(_HANDLER_PATH, handler): + response = await acreate( + name="waverunner", + base_agent="gemini-2.5-flash", + api_key="AIza", + ) + assert response is sentinel + + @pytest.mark.asyncio + async def test_acreate_awaits_coroutine_result(self): + async def _coro(): + return "async-value" + + handler = MagicMock() + handler.create_agent.return_value = _coro() + + with patch(_HANDLER_PATH, handler): + response = await acreate(name="waverunner", api_key="AIza") + + assert response == "async-value" + + @pytest.mark.asyncio + async def test_alist_dispatches_with_async_flag(self): + sentinel = MagicMock(name="alist_response") + + def fake_list_agents(**kwargs): + assert kwargs["_is_async"] is True + return sentinel + + handler = MagicMock() + handler.list_agents.side_effect = fake_list_agents + + with patch(_HANDLER_PATH, handler): + response = await alist(api_key="AIza") + assert response is sentinel + + @pytest.mark.asyncio + async def test_aget_dispatches_with_async_flag(self): + sentinel = MagicMock(name="aget_response") + + def fake_get_agent(**kwargs): + assert kwargs["_is_async"] is True + assert kwargs["name"] == "waverunner" + return sentinel + + handler = MagicMock() + handler.get_agent.side_effect = fake_get_agent + + with patch(_HANDLER_PATH, handler): + response = await aget(name="waverunner", api_key="AIza") + assert response is sentinel + + @pytest.mark.asyncio + async def test_adelete_dispatches_with_async_flag(self): + sentinel = MagicMock(name="adelete_response") + + def fake_delete_agent(**kwargs): + assert kwargs["_is_async"] is True + assert kwargs["name"] == "waverunner" + return sentinel + + handler = MagicMock() + handler.delete_agent.side_effect = fake_delete_agent + + with patch(_HANDLER_PATH, handler): + response = await adelete(name="waverunner", api_key="AIza") + assert response is sentinel + + @pytest.mark.asyncio + async def test_alist_versions_dispatches_with_async_flag(self): + sentinel = MagicMock(name="alist_versions_response") + + def fake_versions(**kwargs): + assert kwargs["_is_async"] is True + assert kwargs["name"] == "waverunner" + return sentinel + + handler = MagicMock() + handler.list_agent_versions.side_effect = fake_versions + + with patch(_HANDLER_PATH, handler): + response = await alist_versions(name="waverunner", api_key="AIza") + assert response is sentinel + + +# --------------------------------------------------------------------------- +# Async error wrapping: exception_type must be invoked +# --------------------------------------------------------------------------- + + +class TestAsyncErrorWrapping: + """If the underlying handler raises, async entry points re-raise via + litellm.exception_type so users get a normalised provider error.""" + + @pytest.mark.asyncio + async def test_acreate_wraps_exception(self): + handler = MagicMock() + handler.create_agent.side_effect = RuntimeError("kaboom") + + with patch(_HANDLER_PATH, handler): + with pytest.raises(Exception): + await acreate(name="waverunner", api_key="AIza") + + @pytest.mark.asyncio + async def test_aget_wraps_exception(self): + handler = MagicMock() + handler.get_agent.side_effect = RuntimeError("kaboom") + + with patch(_HANDLER_PATH, handler): + with pytest.raises(Exception): + await aget(name="waverunner", api_key="AIza") + + @pytest.mark.asyncio + async def test_alist_wraps_exception(self): + handler = MagicMock() + handler.list_agents.side_effect = RuntimeError("kaboom") + + with patch(_HANDLER_PATH, handler): + with pytest.raises(Exception): + await alist(api_key="AIza") + + @pytest.mark.asyncio + async def test_adelete_wraps_exception(self): + handler = MagicMock() + handler.delete_agent.side_effect = RuntimeError("kaboom") + + with patch(_HANDLER_PATH, handler): + with pytest.raises(Exception): + await adelete(name="waverunner", api_key="AIza") + + @pytest.mark.asyncio + async def test_alist_versions_wraps_exception(self): + handler = MagicMock() + handler.list_agent_versions.side_effect = RuntimeError("kaboom") + + with patch(_HANDLER_PATH, handler): + with pytest.raises(Exception): + await alist_versions(name="waverunner", api_key="AIza") diff --git a/tests/test_litellm/interactions/test_gemini_interactions_transformation.py b/tests/test_litellm/interactions/test_gemini_interactions_transformation.py index 0ef97e25e44..2e596a72158 100644 --- a/tests/test_litellm/interactions/test_gemini_interactions_transformation.py +++ b/tests/test_litellm/interactions/test_gemini_interactions_transformation.py @@ -122,6 +122,56 @@ class TestGetCompleteUrl: ) +class TestTransformRequest: + def test_passes_environment_to_request_body(self, config): + request_body = config.transform_request( + model=None, + agent="my-custom-slides-agent", + input=[{"type": "text", "text": "Create a 5-slide presentation about AI trends."}], + optional_params={ + "environment": "remote", + "stream": False, + }, + litellm_params=GenericLiteLLMParams(api_key="test-api-key"), + headers={}, + ) + + assert request_body["agent"] == "my-custom-slides-agent" + assert request_body["environment"] == "remote" + assert request_body["stream"] is False + assert request_body["input"] == [ + {"type": "text", "text": "Create a 5-slide presentation about AI trends."} + ] + + def test_passes_environment_object_to_request_body(self, config): + environment_config = { + "type": "remote", + "sources": [{"type": "gcs", "uri": "gs://bucket/skills.zip"}], + "network": {"egress": "allow_all"}, + } + request_body = config.transform_request( + model=None, + agent="waverunner", + input="What is 2 + 2?", + optional_params={"environment": environment_config}, + litellm_params=GenericLiteLLMParams(api_key="test-api-key"), + headers={}, + ) + + assert request_body["environment"] == environment_config + + def test_passes_existing_environment_id_to_request_body(self, config): + env_id = "env-abc123" + request_body = config.transform_request( + model=None, + agent="my-custom-slides-agent", + input="Continue the presentation.", + optional_params={"environment": env_id}, + litellm_params=GenericLiteLLMParams(api_key="test-api-key"), + headers={}, + ) + + assert request_body["environment"] == env_id class TestStreamingIterator: def _make_iterator(self) -> LiteLLMResponsesInteractionsStreamingIterator: return LiteLLMResponsesInteractionsStreamingIterator( diff --git a/tests/test_litellm/proxy/google_endpoints/test_interactions_agent_param.py b/tests/test_litellm/proxy/google_endpoints/test_interactions_agent_param.py index 1063f59afb6..f3cec320532 100644 --- a/tests/test_litellm/proxy/google_endpoints/test_interactions_agent_param.py +++ b/tests/test_litellm/proxy/google_endpoints/test_interactions_agent_param.py @@ -1,75 +1,103 @@ """ -Test for interactions endpoint agent parameter handling. +Tests for managed-agent interaction routing. -Tests that the /v1beta/interactions endpoint correctly extracts -the `agent` parameter as a fallback when `model` is not provided. +Custom Gemini agents are identified by ``agent`` (name/id), not ``model``. +The proxy must not pass the agent name as ``model`` or LiteLLM may route to +openai/* wildcards instead of Gemini interactions. """ +from unittest.mock import MagicMock, patch + import pytest class TestInteractionsAgentParameter: - """Test agent parameter handling in interactions endpoint.""" + """Proxy endpoint must keep agent and model separate.""" - def test_agent_parameter_fallback_logic(self): - """ - Test the core logic: model or agent extraction. - - This tests the fix in endpoints.py line ~267: - model=data.get("model") or data.get("agent") - """ - # Case 1: Only agent provided (Deep Research use case) + def test_create_interaction_uses_model_only_from_body(self): + """POST /v1beta/interactions: model kwarg is only the request's model field.""" data = { - "agent": "deep-research-pro-preview-12-2025", - "input": "Research quantum computing", - "background": True, + "agent": "mqy-custom-slides-agent", + "input": "hello", } - model = data.get("model") or data.get("agent") - assert model == "deep-research-pro-preview-12-2025" + # Fixed behavior: do NOT fall back agent → model + model_for_routing = data.get("model") + assert model_for_routing is None + assert data.get("agent") == "mqy-custom-slides-agent" - # Case 2: Only model provided (normal use case) + def test_model_field_still_used_when_present(self): data = { "model": "gemini-2.5-flash", - "input": "Hello world", + "input": "hello", } - model = data.get("model") or data.get("agent") - assert model == "gemini-2.5-flash" + model_for_routing = data.get("model") + assert model_for_routing == "gemini-2.5-flash" - # Case 3: Both provided (model takes precedence) - data = { - "model": "gemini-2.5-flash", - "agent": "deep-research-pro-preview-12-2025", - "input": "Test", - } - model = data.get("model") or data.get("agent") - assert model == "gemini-2.5-flash" - # Case 4: Neither provided - data = { - "input": "Test", - } - model = data.get("model") or data.get("agent") - assert model is None +class TestInteractionsAgentOnlyProviderRouting: + """SDK: agent-only create must not call get_llm_provider on the agent name.""" - def test_route_type_in_skip_model_routing_list(self): - """ - Test that acreate_interaction is in the list of routes - that skip model-based routing. + @patch("litellm.interactions.main.interactions_http_handler") + @patch("litellm.interactions.main.get_provider_interactions_api_config") + @patch("litellm.get_llm_provider") + def test_agent_only_skips_get_llm_provider( + self, + mock_get_llm_provider, + mock_get_config, + mock_handler, + ): + from litellm.interactions.main import create + from litellm.types.interactions import InteractionsAPIResponse - This tests the fix in route_llm_request.py. - """ - # The list of routes that skip model routing for interactions - skip_model_routing_routes = [ - "acreate_interaction", - "aget_interaction", - "adelete_interaction", - "acancel_interaction", - ] + mock_get_config.return_value = MagicMock() + mock_handler.create_interaction.return_value = InteractionsAPIResponse( + id="int-1", + status="completed", + object="interaction", + ) - # acreate_interaction should be in the list (this is the fix) - assert "acreate_interaction" in skip_model_routing_routes + logging_obj = MagicMock() + create( + agent="mqy-custom-slides-agent", + input="test", + custom_llm_provider="gemini", + litellm_logging_obj=logging_obj, + ) - # All interaction routes should be covered - assert "aget_interaction" in skip_model_routing_routes - assert "adelete_interaction" in skip_model_routing_routes - assert "acancel_interaction" in skip_model_routing_routes + mock_get_llm_provider.assert_not_called() + call_kwargs = mock_handler.create_interaction.call_args.kwargs + assert call_kwargs["agent"] == "mqy-custom-slides-agent" + assert call_kwargs["model"] is None + assert call_kwargs["custom_llm_provider"] == "gemini" + + @patch("litellm.interactions.main.interactions_http_handler") + @patch("litellm.interactions.main.get_provider_interactions_api_config") + @patch("litellm.get_llm_provider") + def test_proxy_mistake_model_equals_agent_is_corrected( + self, + mock_get_llm_provider, + mock_get_config, + mock_handler, + ): + """If model was wrongly set to the agent name, clear it before the HTTP call.""" + from litellm.interactions.main import create + from litellm.types.interactions import InteractionsAPIResponse + + mock_get_config.return_value = MagicMock() + mock_handler.create_interaction.return_value = InteractionsAPIResponse( + id="int-1", + status="completed", + object="interaction", + ) + + logging_obj = MagicMock() + create( + model="mqy-custom-slides-agent", + agent="mqy-custom-slides-agent", + input="test", + custom_llm_provider="gemini", + litellm_logging_obj=logging_obj, + ) + + mock_get_llm_provider.assert_not_called() + assert mock_handler.create_interaction.call_args.kwargs["model"] is None diff --git a/tests/test_litellm/proxy/google_endpoints/test_managed_agents_model_param.py b/tests/test_litellm/proxy/google_endpoints/test_managed_agents_model_param.py new file mode 100644 index 00000000000..5485d0f2929 --- /dev/null +++ b/tests/test_litellm/proxy/google_endpoints/test_managed_agents_model_param.py @@ -0,0 +1,199 @@ +""" +Tests verifying that managed-agent proxy endpoints never pass the agent name +as the ``model`` parameter to ``base_process_llm_request``. + +Passing ``model=`` would cause ``common_processing_pre_call_logic`` +to write the agent name into ``self.data["model"]``, which triggers spurious +model-alias mapping, rate-limiting lookups, and logging tied to a +non-existent model deployment. The agent name is already carried in +``data["name"]`` and must not pollute the ``model`` slot. +""" + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + + +def _build_agents_client(): + """Build a TestClient whose auth dependency is overridden to a PROXY_ADMIN + user. Using ``dependency_overrides`` is the only reliable way to bypass the + real ``user_api_key_auth`` for FastAPI route tests — patching the module- + level name does not affect the function reference captured by ``Depends``. + The PROXY_ADMIN role also bypasses the caller-supplied-api_key guard so + these tests can focus on the ``model=None`` invariant. + """ + from fastapi import FastAPI + from fastapi.testclient import TestClient + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.google_endpoints.agents_endpoints import router as agents_router + + app = FastAPI() + app.include_router(agents_router) + + async def _fake_user_api_key_auth(): + return UserAPIKeyAuth( + api_key="sk-test", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + + app.dependency_overrides[user_api_key_auth] = _fake_user_api_key_auth + return TestClient(app) + + +def _patch_proxy_server_imports(client=None): + """Return a context-manager that stubs _proxy_server_imports so tests + don't need a running proxy.""" + mock_srv = { + "general_settings": {}, + "llm_router": MagicMock(), + "proxy_config": MagicMock(), + "proxy_logging_obj": MagicMock(), + "select_data_generator": None, + "user_api_base": None, + "user_max_tokens": None, + "user_model": None, + "user_request_timeout": None, + "user_temperature": None, + "version": "0.0.0", + } + return patch( + "litellm.proxy.google_endpoints.agents_endpoints._proxy_server_imports", + return_value=mock_srv, + ) + + +def _patch_base_process(return_value=None): + if return_value is None: + return_value = {"name": "agents/my-agent", "displayName": "My Agent"} + return patch( + "litellm.proxy.google_endpoints.agents_endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request", + new_callable=AsyncMock, + return_value=return_value, + ) + + +def _patch_auth(): + """Deprecated no-op kept for call-site compatibility. + + ``_build_agents_client`` now installs a FastAPI ``dependency_overrides`` + entry that injects a PROXY_ADMIN ``UserAPIKeyAuth``, so individual tests + no longer need to patch the module-level ``user_api_key_auth`` name. + """ + return patch("os.getpid") + + +class TestManagedAgentsModelParam: + """Endpoints must pass model=None, not the agent name, to base_process_llm_request.""" + + def test_create_agent_passes_model_none(self): + """POST /v1beta/agents: model kwarg must be None, not the name field.""" + try: + client = _build_agents_client() + except ImportError as exc: + pytest.skip(f"Skipping: missing dependency {exc}") + + with ( + _patch_proxy_server_imports(), + _patch_base_process() as mock_process, + _patch_auth(), + ): + client.post( + "/v1beta/agents", + json={ + "name": "my-custom-slides-agent", + "base_agent": "waverunner", + "instructions": "Be helpful.", + }, + ) + + mock_process.assert_called_once() + kwargs = mock_process.call_args.kwargs + assert kwargs["model"] is None, ( + f"create_gemini_agent must not pass model={kwargs['model']!r}; " + "the agent name must stay in data['name'], not pollute data['model']" + ) + assert kwargs["route_type"] == "acreate_agent" + + def test_get_agent_passes_model_none(self): + """GET /v1beta/agents/{name}: model kwarg must be None.""" + try: + client = _build_agents_client() + except ImportError as exc: + pytest.skip(f"Skipping: missing dependency {exc}") + + with ( + _patch_proxy_server_imports(), + _patch_base_process() as mock_process, + _patch_auth(), + ): + client.get("/v1beta/agents/my-custom-slides-agent") + + mock_process.assert_called_once() + kwargs = mock_process.call_args.kwargs + assert ( + kwargs["model"] is None + ), f"get_gemini_agent must not pass model={kwargs['model']!r}" + assert kwargs["route_type"] == "aget_agent" + + def test_delete_agent_passes_model_none(self): + """DELETE /v1beta/agents/{name}: model kwarg must be None.""" + try: + client = _build_agents_client() + except ImportError as exc: + pytest.skip(f"Skipping: missing dependency {exc}") + + with ( + _patch_proxy_server_imports(), + _patch_base_process() as mock_process, + _patch_auth(), + ): + client.delete("/v1beta/agents/my-custom-slides-agent") + + mock_process.assert_called_once() + kwargs = mock_process.call_args.kwargs + assert ( + kwargs["model"] is None + ), f"delete_gemini_agent must not pass model={kwargs['model']!r}" + assert kwargs["route_type"] == "adelete_agent" + + def test_list_agent_versions_passes_model_none(self): + """GET /v1beta/agents/{name}/versions: model kwarg must be None.""" + try: + client = _build_agents_client() + except ImportError as exc: + pytest.skip(f"Skipping: missing dependency {exc}") + + with ( + _patch_proxy_server_imports(), + _patch_base_process() as mock_process, + _patch_auth(), + ): + client.get("/v1beta/agents/my-custom-slides-agent/versions") + + mock_process.assert_called_once() + kwargs = mock_process.call_args.kwargs + assert ( + kwargs["model"] is None + ), f"list_gemini_agent_versions must not pass model={kwargs['model']!r}" + assert kwargs["route_type"] == "alist_agent_versions" + + def test_list_agents_already_passes_model_none(self): + """GET /v1beta/agents: existing list endpoint already passes model=None — keep it so.""" + try: + client = _build_agents_client() + except ImportError as exc: + pytest.skip(f"Skipping: missing dependency {exc}") + + with ( + _patch_proxy_server_imports(), + _patch_base_process(return_value={"agents": []}) as mock_process, + _patch_auth(), + ): + client.get("/v1beta/agents") + + mock_process.assert_called_once() + kwargs = mock_process.call_args.kwargs + assert kwargs["model"] is None + assert kwargs["route_type"] == "alist_agents" diff --git a/tests/test_litellm/router_utils/test_router_interactions_endpoints.py b/tests/test_litellm/router_utils/test_router_interactions_endpoints.py index c5468d73810..91bea170458 100644 --- a/tests/test_litellm/router_utils/test_router_interactions_endpoints.py +++ b/tests/test_litellm/router_utils/test_router_interactions_endpoints.py @@ -140,3 +140,159 @@ class TestInitInteractionsApiEndpoints: custom_llm_provider="vertex_ai", ) assert result == {"result": "success"} + + @pytest.mark.asyncio + async def test_init_interactions_api_endpoints_clears_model_when_equals_agent( + self, + ): + """Managed agent interactions must not pass agent name as model to the SDK.""" + router = Router(model_list=[]) + + mock_function = AsyncMock(return_value={"result": "success"}) + + await router._init_interactions_api_endpoints( + original_function=mock_function, + agent="mqy-custom-slides-agent", + model="mqy-custom-slides-agent", + input="hello", + ) + + mock_function.assert_called_once_with( + custom_llm_provider="gemini", + agent="mqy-custom-slides-agent", + model=None, + input="hello", + ) + + +class TestRouterCreateInteractionRouting: + """acreate_interaction routing: agent-only vs model + fallbacks.""" + + @pytest.mark.asyncio + async def test_acreate_interaction_agent_only_uses_init_interactions(self): + """Agent-only create must not use model-group fallback lookup.""" + router = Router( + model_list=[ + { + "model_name": "openai/*", + "litellm_params": {"model": "gpt-4"}, + } + ] + ) + + with ( + patch.object( + router, + "_init_interactions_api_endpoints", + new_callable=AsyncMock, + return_value={"id": "int-1"}, + ) as mock_init, + patch.object( + router, + "_ageneric_api_call_with_fallbacks", + new_callable=AsyncMock, + ) as mock_generic, + ): + result = await router.acreate_interaction( + agent="mqy-custom-slides-agent", + input="hello", + custom_llm_provider="gemini", + ) + + mock_init.assert_called_once() + mock_generic.assert_not_called() + assert result == {"id": "int-1"} + + @pytest.mark.asyncio + async def test_init_interactions_model_uses_generic_fallbacks(self): + """Model-based create uses _ageneric_api_call_with_fallbacks inside _init_interactions.""" + router = Router(model_list=[]) + + with patch.object( + router, + "_ageneric_api_call_with_fallbacks", + new_callable=AsyncMock, + return_value={"id": "int-1"}, + ) as mock_generic: + result = await router._init_interactions_api_endpoints( + original_function=AsyncMock(), + model="gemini-2.5-flash", + input="hello", + custom_llm_provider="gemini", + ) + + mock_generic.assert_called_once() + assert result == {"id": "int-1"} + + +class TestInitializeManagedAgentsEndpoints: + """Tests for _initialize_managed_agents_endpoints.""" + + def test_initialize_managed_agents_endpoints_creates_methods(self): + router = Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4"}, + } + ] + ) + + for method_name in ( + "acreate_agent", + "alist_agents", + "aget_agent", + "adelete_agent", + "alist_agent_versions", + ): + assert hasattr(router, method_name), f"missing {method_name}" + assert callable(getattr(router, method_name)), f"{method_name} not callable" + + def test_initialize_managed_agents_endpoints_can_be_called_directly(self): + router = Router(model_list=[]) + router._initialize_managed_agents_endpoints() + assert callable(router.acreate_agent) + assert callable(router.alist_agents) + + +class TestInitManagedAgentsApiEndpoints: + """Tests for _init_managed_agents_api_endpoints.""" + + @pytest.mark.asyncio + async def test_init_managed_agents_api_endpoints_defaults_to_gemini(self): + router = Router(model_list=[]) + mock_fn = AsyncMock(return_value={"agents": []}) + + await router._init_managed_agents_api_endpoints( + original_function=mock_fn, + ) + + call_kwargs = mock_fn.call_args.kwargs + assert call_kwargs["custom_llm_provider"] == "gemini" + + @pytest.mark.asyncio + async def test_init_managed_agents_api_endpoints_passes_custom_provider(self): + router = Router(model_list=[]) + mock_fn = AsyncMock(return_value={"agents": []}) + + await router._init_managed_agents_api_endpoints( + original_function=mock_fn, + custom_llm_provider="vertex_ai", + ) + + call_kwargs = mock_fn.call_args.kwargs + assert call_kwargs["custom_llm_provider"] == "vertex_ai" + + @pytest.mark.asyncio + async def test_init_managed_agents_api_endpoints_does_not_override_existing_provider( + self, + ): + router = Router(model_list=[]) + mock_fn = AsyncMock(return_value={"agents": []}) + + await router._init_managed_agents_api_endpoints( + original_function=mock_fn, + custom_llm_provider="vertex_ai", + ) + + mock_fn.assert_called_once_with(custom_llm_provider="vertex_ai") From 99a63d5180f7094250ed1d578d0a91c6dca92cdd Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 20 May 2026 22:33:14 +0530 Subject: [PATCH 7/7] feat(gemini): add gemini-3.1-flash-lite model cost map (#28320) * feat(gemini): add gemini-3.1-flash-lite model cost map entries Co-authored-by: Cursor * Update model_prices_and_context_window.json * Update source URL for model pricing information * Sync source URL for gemini-3.1-flash-lite in backup JSON * fix(model_cost_map): add mistral/ministral-8b-2512 entry Mistral rotated the 'mistral/mistral-tiny' alias to return 'ministral-8b-2512' as the response model, which is not in the cost map. This caused test_completion_mistral_api and test_completion_mistral_api_modified_input to fail in completion_cost lookup. Add the entry mirroring the existing openrouter/mistralai/ministral-8b-2512 pricing. * test(cost_calculator): assert output_cost_per_reasoning_token for gemini-3.1-flash-lite * fix(tests): backfill local backup entries into runtime model_cost litellm.model_cost is loaded from LITELLM_MODEL_COST_MAP_URL (pinned to main) at import time, so any pricing entries added to the in-tree backup on this branch aren't visible at test runtime until they also land on main. The Mistral cassette currently returns model=ministral-8b-2512 and the cost-calculator lookup in test_completion_mistral_api / test_completion_mistral_api_modified_input fails despite the entry existing in the local backup. Backfill missing backup entries into litellm.model_cost in the local_testing conftest so these lookups succeed against the cassette state the branch is being tested with. * fix(tests): guard conftest backfill against empty local cost map --------- Co-authored-by: Cursor Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 227 +++++++++++++++++- model_prices_and_context_window.json | 191 +++++++++++++++ tests/local_testing/conftest.py | 14 ++ tests/test_litellm/test_cost_calculator.py | 19 ++ 4 files changed, 450 insertions(+), 1 deletion(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 6d7c3eeb0b8..9ba337da0a5 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -1448,6 +1448,35 @@ "supports_native_structured_output": true, "supports_minimal_reasoning_effort": true }, + "jp.anthropic.claude-sonnet-4-6": { + "cache_creation_input_token_cost": 4.125e-06, + "cache_read_input_token_cost": 3.3e-07, + "input_cost_per_token": 3.3e-06, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1000000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 1.65e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_max_reasoning_effort": true, + "supports_tool_choice": true, + "supports_vision": true, + "tool_use_system_prompt_tokens": 346, + "supports_native_structured_output": true, + "supports_minimal_reasoning_effort": true + }, "anthropic.claude-sonnet-4-20250514-v1:0": { "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, @@ -9602,6 +9631,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "supports_adaptive_thinking": true, "supports_assistant_prefill": true, "supports_computer_use": true, "supports_function_calling": true, @@ -9795,6 +9825,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -9828,6 +9859,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -9861,6 +9893,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -9895,6 +9928,7 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 }, + "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, "supports_function_calling": true, @@ -14883,7 +14917,65 @@ "mode": "chat", "output_cost_per_reasoning_token": 1.5e-06, "output_cost_per_token": 1.5e-06, - "source": "https://ai.google.dev/gemini-api/docs/models", + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_code_execution": true, + "supports_file_search": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query", + "supports_service_tier": true + }, + "gemini-3.1-flash-lite": { + "cache_read_input_token_cost": 4.5e-08, + "cache_read_input_token_cost_per_audio_token": 9e-08, + "input_cost_per_audio_token": 9e-07, + "input_cost_per_token": 4.5e-07, + "litellm_provider": "vertex_ai-language-models", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_pdf_size_mb": 30, + "max_tokens": 65536, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 2.7e-06, + "output_cost_per_token": 2.7e-06, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", "supported_endpoints": [ "/v1/chat/completions", "/v1/completions", @@ -16987,6 +17079,66 @@ "web_search_billing_unit": "per_query", "supports_service_tier": true }, + "gemini/gemini-3.1-flash-lite": { + "cache_read_input_token_cost": 4.5e-08, + "cache_read_input_token_cost_per_audio_token": 9e-08, + "input_cost_per_audio_token": 9e-07, + "input_cost_per_token": 4.5e-07, + "litellm_provider": "gemini", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_pdf_size_mb": 30, + "max_tokens": 65536, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 2.7e-06, + "output_cost_per_token": 2.7e-06, + "rpm": 15, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_code_execution": true, + "supports_file_search": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "tpm": 250000, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query", + "supports_service_tier": true + }, "gemini/gemini-3-flash-preview": { "cache_read_input_token_cost": 5e-08, "input_cost_per_audio_token": 1e-06, @@ -24285,6 +24437,21 @@ "supports_tool_choice": true, "supports_vision": true }, + "mistral/ministral-8b-2512": { + "input_cost_per_token": 1.5e-07, + "litellm_provider": "mistral", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 1.5e-07, + "source": "https://mistral.ai/pricing", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "mistral/mistral-tiny": { "input_cost_per_token": 2.5e-07, "litellm_provider": "mistral", @@ -33605,6 +33772,64 @@ }, "web_search_billing_unit": "per_query" }, + "vertex_ai/gemini-3.1-flash-lite": { + "cache_read_input_token_cost": 4.5e-08, + "cache_read_input_token_cost_per_audio_token": 9e-08, + "input_cost_per_audio_token": 9e-07, + "input_cost_per_token": 4.5e-07, + "litellm_provider": "vertex_ai-language-models", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_pdf_size_mb": 30, + "max_tokens": 65536, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 2.7e-06, + "output_cost_per_token": 2.7e-06, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_code_execution": true, + "supports_file_search": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query", + "supports_service_tier": true + }, "vertex_ai/deep-research-pro-preview-12-2025": { "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index e7a03bb0984..27d6a59740f 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -14957,6 +14957,64 @@ "web_search_billing_unit": "per_query", "supports_service_tier": true }, + "gemini-3.1-flash-lite": { + "cache_read_input_token_cost": 4.5e-08, + "cache_read_input_token_cost_per_audio_token": 9e-08, + "input_cost_per_audio_token": 9e-07, + "input_cost_per_token": 4.5e-07, + "litellm_provider": "vertex_ai-language-models", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_pdf_size_mb": 30, + "max_tokens": 65536, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 2.7e-06, + "output_cost_per_token": 2.7e-06, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_code_execution": true, + "supports_file_search": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query", + "supports_service_tier": true + }, "deep-research-pro-preview-12-2025": { "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, @@ -17021,6 +17079,66 @@ "web_search_billing_unit": "per_query", "supports_service_tier": true }, + "gemini/gemini-3.1-flash-lite": { + "cache_read_input_token_cost": 4.5e-08, + "cache_read_input_token_cost_per_audio_token": 9e-08, + "input_cost_per_audio_token": 9e-07, + "input_cost_per_token": 4.5e-07, + "litellm_provider": "gemini", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_pdf_size_mb": 30, + "max_tokens": 65536, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 2.7e-06, + "output_cost_per_token": 2.7e-06, + "rpm": 15, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_code_execution": true, + "supports_file_search": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "tpm": 250000, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query", + "supports_service_tier": true + }, "gemini/gemini-3-flash-preview": { "cache_read_input_token_cost": 5e-08, "input_cost_per_audio_token": 1e-06, @@ -24319,6 +24437,21 @@ "supports_tool_choice": true, "supports_vision": true }, + "mistral/ministral-8b-2512": { + "input_cost_per_token": 1.5e-07, + "litellm_provider": "mistral", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 1.5e-07, + "source": "https://mistral.ai/pricing", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "mistral/mistral-tiny": { "input_cost_per_token": 2.5e-07, "litellm_provider": "mistral", @@ -33639,6 +33772,64 @@ }, "web_search_billing_unit": "per_query" }, + "vertex_ai/gemini-3.1-flash-lite": { + "cache_read_input_token_cost": 4.5e-08, + "cache_read_input_token_cost_per_audio_token": 9e-08, + "input_cost_per_audio_token": 9e-07, + "input_cost_per_token": 4.5e-07, + "litellm_provider": "vertex_ai-language-models", + "max_audio_length_hours": 8.4, + "max_audio_per_prompt": 1, + "max_images_per_prompt": 3000, + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_pdf_size_mb": 30, + "max_tokens": 65536, + "max_video_length": 1, + "max_videos_per_prompt": 10, + "mode": "chat", + "output_cost_per_reasoning_token": 2.7e-06, + "output_cost_per_token": 2.7e-06, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_code_execution": true, + "supports_file_search": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query", + "supports_service_tier": true + }, "vertex_ai/deep-research-pro-preview-12-2025": { "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, diff --git a/tests/local_testing/conftest.py b/tests/local_testing/conftest.py index 6a746041f15..06637b844b1 100644 --- a/tests/local_testing/conftest.py +++ b/tests/local_testing/conftest.py @@ -22,6 +22,20 @@ sys.path.insert( ) # Adds the parent directory to the system path import litellm +# ``litellm.model_cost`` is loaded at import time from the URL pinned to +# ``main`` (``LITELLM_MODEL_COST_MAP_URL``). The in-tree backup ships with +# this branch and can include pricing entries that main has not yet picked +# up (e.g. an upstream provider rotates a model id and the test cassette +# records the new name). Backfill any entries that are missing from the +# remote-fetched map so cost-calculator lookups in tests succeed against +# the cassette state the branch is being tested with. +from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap + +_local_cost_map = GetModelCostMap.load_local_model_cost_map() +for _k, _v in _local_cost_map.items(): + litellm.model_cost.setdefault(_k, _v) +del _local_cost_map + from tests._vcr_conftest_common import ( # noqa: E402,F401 VerboseReporterState, _pin_multipart_boundary, diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 1be4abbec6e..18ab8a2a07a 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -2059,6 +2059,25 @@ def test_openrouter_gemini_3_1_flash_lite_preview_pricing(): assert model_info["max_output_tokens"] == 65536 +def test_gemini_3_1_flash_lite_pricing(): + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + for model_name in ( + "gemini-3.1-flash-lite", + "gemini/gemini-3.1-flash-lite", + "vertex_ai/gemini-3.1-flash-lite", + ): + model_info = litellm.model_cost.get(model_name) + assert model_info is not None, f"Missing model pricing entry: {model_name}" + assert model_info["input_cost_per_token"] == 4.5e-07 + assert model_info["input_cost_per_audio_token"] == 9e-07 + assert model_info["output_cost_per_token"] == 2.7e-06 + assert model_info["output_cost_per_reasoning_token"] == 2.7e-06 + assert model_info["cache_read_input_token_cost"] == 4.5e-08 + assert model_info["max_input_tokens"] == 1048576 + + def test_custom_pricing_applies_cache_read_input_cost(): """ Bug 1 reproduction: custom_cost_per_token with cache_read_input_token_cost