From f5b11b72a6dcc8f4e7a16f00a61bdd124cc171d2 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 8 Jun 2026 22:03:35 +0530 Subject: [PATCH 01/22] feat(proxy): publish /v2/model/info in Swagger OpenAPI spec (#29900) * feat(proxy): publish /v2/model/info in Swagger OpenAPI spec Expose the v2 model info endpoint in /docs by removing include_in_schema=False and documenting query parameters used by the admin UI and proxy CLI consumers. Co-authored-by: Cursor * chore(ui): regenerate schema.d.ts for /v2/model/info OpenAPI docs Co-authored-by: Cursor --------- Co-authored-by: Cursor --- litellm/proxy/proxy_server.py | 46 ++++++++++++++- .../proxy_server/test_routes_model_info.py | 9 +++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 56 ++++++++++++++++--- 3 files changed, 101 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 2c21e19dcec..213f682b8f7 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -11742,10 +11742,8 @@ async def _find_model_by_id( @router.get( "/v2/model/info", - description="v2 - returns models available to the user based on their API key permissions. Shows model info from config.yaml (except api key and api base). Filter to just user-added models with ?user_models_only=true", tags=["model management"], dependencies=[Depends(user_api_key_auth)], - include_in_schema=False, ) async def model_info_v2( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), @@ -11781,7 +11779,49 @@ async def model_info_v2( ), ): """ - BETA ENDPOINT. Might change unexpectedly. Use `/v1/model/info` for now. + Paginated model metadata for proxy deployments (pricing, provider, team access). + + Returns configured router deployments with enriched `model_info` (costs, provider, + context window, etc.). Sensitive fields such as API keys and api_base are omitted. + + Query parameters: + model: Filter to a single public `model_name`. + user_models_only: When true, only return models created by the calling user. + include_team_models: When true, populate `access_via_team_ids` and `direct_access` + on each model and filter to deployments the caller can use. + page / size: Pagination controls (defaults: page=1, size=50). + search: Case-insensitive partial match on model name or team public name. + modelId: Return a single deployment by LiteLLM model id. + teamId: Filter to models with direct access or team membership for this team id. + sortBy / sortOrder: Sort by model_name, created_at, updated_at, costs, or status. + + Example request: + ``` + curl -X GET 'http://localhost:4000/v2/model/info?include_team_models=true&page=1&size=50' \\ + --header 'Authorization: Bearer sk-1234' + ``` + + Example response: + ```json + { + "data": [ + { + "model_name": "gpt-4", + "litellm_params": {"model": "openai/gpt-4.1"}, + "model_info": { + "id": "abc123", + "litellm_provider": "openai", + "access_via_team_ids": ["team-1"], + "direct_access": true + } + } + ], + "total_count": 1, + "current_page": 1, + "total_pages": 1, + "size": 50 + } + ``` """ global llm_model_list, general_settings, user_config_file_path, proxy_config, llm_router diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py b/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py index 98259824378..017f4bd4368 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py @@ -60,6 +60,15 @@ def test_v2_model_info_invalid_page_returns_422(client, auth_as, empty_router): assert "detail" in response.json() +def test_v2_model_info_in_openapi_schema(): + """``GET /v2/model/info`` is published in the proxy OpenAPI/Swagger spec.""" + from litellm.proxy.proxy_server import get_openapi_schema + + schema = get_openapi_schema() + assert "/v2/model/info" in schema["paths"] + assert "get" in schema["paths"]["/v2/model/info"] + + # --------------------------------------------------------------------------- # GET /v1/model/info, GET /model/info # --------------------------------------------------------------------------- diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index c8ec9d3c727..75a14e09852 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -18441,7 +18441,49 @@ export interface paths { }; /** * Model Info V2 - * @description v2 - returns models available to the user based on their API key permissions. Shows model info from config.yaml (except api key and api base). Filter to just user-added models with ?user_models_only=true + * @description Paginated model metadata for proxy deployments (pricing, provider, team access). + * + * Returns configured router deployments with enriched `model_info` (costs, provider, + * context window, etc.). Sensitive fields such as API keys and api_base are omitted. + * + * Query parameters: + * model: Filter to a single public `model_name`. + * user_models_only: When true, only return models created by the calling user. + * include_team_models: When true, populate `access_via_team_ids` and `direct_access` + * on each model and filter to deployments the caller can use. + * page / size: Pagination controls (defaults: page=1, size=50). + * search: Case-insensitive partial match on model name or team public name. + * modelId: Return a single deployment by LiteLLM model id. + * teamId: Filter to models with direct access or team membership for this team id. + * sortBy / sortOrder: Sort by model_name, created_at, updated_at, costs, or status. + * + * Example request: + * ``` + * curl -X GET 'http://localhost:4000/v2/model/info?include_team_models=true&page=1&size=50' \ + * --header 'Authorization: Bearer sk-1234' + * ``` + * + * Example response: + * ```json + * { + * "data": [ + * { + * "model_name": "gpt-4", + * "litellm_params": {"model": "openai/gpt-4.1"}, + * "model_info": { + * "id": "abc123", + * "litellm_provider": "openai", + * "access_via_team_ids": ["team-1"], + * "direct_access": true + * } + * } + * ], + * "total_count": 1, + * "current_page": 1, + * "total_pages": 1, + * "size": 50 + * } + * ``` */ get: operations["model_info_v2_v2_model_info_get"]; put?: never; @@ -43119,13 +43161,13 @@ export interface operations { /** * @description Unified rate-limit error. * - * Every rate-limit condition surfaced by litellm — whether it originated from - * an upstream LLM provider, a vendor batch endpoint, or one of litellm's own - * proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget, - * max-iterations, etc.) — is raised as an instance of this class. + * Every rate-limit condition surfaced by litellm — whether it originated from + * an upstream LLM provider, a vendor batch endpoint, or one of litellm's own + * proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget, + * max-iterations, etc.) — is raised as an instance of this class. * - * The :attr:`category` attribute lets callers distinguish the source. See - * :class:`RateLimitErrorCategory` for the available values. + * The :attr:`category` attribute lets callers distinguish the source. See + * :class:`RateLimitErrorCategory` for the available values. */ 429: { headers: { From ff6cea4833df88690cd282e51654411f0bca0d84 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Mon, 8 Jun 2026 11:25:50 -0700 Subject: [PATCH 02/22] refactor(ui): single source of truth for migrated-page routing (#29949) Consolidate the three hand-synced copies of the migrated-pages map (LEGACY_REDIRECTS in app/page.tsx, MIGRATED_PAGES in the dashboard layout, and MIGRATED_PAGES in leftnav) into one shared module, src/utils/migratedPages.ts, which also owns the migratedHref helper. Delete the unused, incomplete Sidebar2 prototype. No runtime behavior change: the map is still empty and Sidebar2 had no importers, so this is pure deduplication ahead of the per-page App Router migration. Follow-up work will unify the remaining base-URL builders (layout's withBase and page.tsx's redirect) onto migratedHref. --- .../app/(dashboard)/components/Sidebar2.tsx | 472 ------------------ .../src/app/(dashboard)/layout.tsx | 10 +- ui/litellm-dashboard/src/app/page.tsx | 12 +- .../src/components/leftnav.tsx | 24 +- .../src/utils/migratedPages.test.ts | 43 ++ .../src/utils/migratedPages.ts | 25 + 6 files changed, 73 insertions(+), 513 deletions(-) delete mode 100644 ui/litellm-dashboard/src/app/(dashboard)/components/Sidebar2.tsx create mode 100644 ui/litellm-dashboard/src/utils/migratedPages.test.ts create mode 100644 ui/litellm-dashboard/src/utils/migratedPages.ts diff --git a/ui/litellm-dashboard/src/app/(dashboard)/components/Sidebar2.tsx b/ui/litellm-dashboard/src/app/(dashboard)/components/Sidebar2.tsx deleted file mode 100644 index 90f498912a8..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/components/Sidebar2.tsx +++ /dev/null @@ -1,472 +0,0 @@ -"use client"; - -import { Layout, Menu, ConfigProvider } from "antd"; -import { - KeyOutlined, - PlayCircleOutlined, - BlockOutlined, - BarChartOutlined, - TeamOutlined, - BankOutlined, - UserOutlined, - SettingOutlined, - ApiOutlined, - AppstoreOutlined, - DatabaseOutlined, - FileTextOutlined, - LineChartOutlined, - SafetyOutlined, - ExperimentOutlined, - ToolOutlined, - TagsOutlined, - AuditOutlined, -} from "@ant-design/icons"; -// import { -// all_admin_roles, -// rolesWithWriteAccess, -// internalUserRoles, -// isAdminRole, -// } from "../utils/roles"; -// import UsageIndicator from "./usage_indicator"; -import * as React from "react"; -import { useRouter, usePathname } from "next/navigation"; -import { all_admin_roles, internalUserRoles, isAdminRole, rolesWithWriteAccess } from "@/utils/roles"; -import UsageIndicator from "@/components/UsageIndicator"; -import { serverRootPath } from "@/components/networking"; - -const { Sider } = Layout; - -// -------- Types -------- -interface SidebarProps { - accessToken: string | null; - userRole: string; - /** Fallback selection id (legacy), used if path can't be matched */ - defaultSelectedKey: string; - collapsed?: boolean; -} - -interface MenuItemCfg { - key: string; - newTab?: boolean; - page: string; // legacy id; we map this to a path below - label: string; - roles?: string[]; - children?: MenuItemCfg[]; - icon?: React.ReactNode; -} - -/** ---------- Base URL helpers ---------- */ -/** - * Normalizes NEXT_PUBLIC_BASE_URL to either "/" or "/ui/" (always with a trailing slash). - * Supported env values: "" or "ui/". - * Also considers the serverRootPath from the proxy config (e.g., "/my-custom-path"). - */ -const getBasePath = () => { - const raw = process.env.NEXT_PUBLIC_BASE_URL ?? ""; - const trimmed = raw.replace(/^\/+|\/+$/g, ""); // strip leading/trailing slashes - const uiPath = trimmed ? `/${trimmed}/` : "/"; - - // If serverRootPath is set and not "/", prepend it to the UI path - if (serverRootPath && serverRootPath !== "/") { - // Remove trailing slash from serverRootPath and ensure uiPath has no leading slash for proper joining - const cleanServerRoot = serverRootPath.replace(/\/+$/, ""); - const cleanUiPath = uiPath.replace(/^\/+/, ""); - return `${cleanServerRoot}/${cleanUiPath}`; - } - - return uiPath; -}; - -/** Map legacy `page` ids to real app routes (relative, no leading slash). */ -const routeFor = (slug: string): string => { - switch (slug) { - // top level - case "api-keys": - return "virtual-keys"; - case "llm-playground": - return "test-key"; - case "models": - return "models-and-endpoints"; - case "new_usage": - return "usage"; - case "teams": - return "teams"; - case "organizations": - return "organizations"; - case "users": - return "users"; - case "api_ref": - return "api-reference"; - case "model-hub-table": - // If you intend the newer in-dashboard page, use "model-hub". - return "model-hub"; - case "logs": - return "logs"; - case "guardrails": - return "guardrails"; - case "policies": - return "policies"; - case "chat": - return "chat"; - - // tools - case "mcp-servers": - return "tools/mcp-servers"; - case "vector-stores": - return "tools/vector-stores"; - case "byok-demo": - return "tools/byok-demo"; - - // experimental - case "caching": - return "experimental/caching"; - case "prompts": - return "experimental/prompts"; - case "budgets": - return "experimental/budgets"; - case "transform-request": - return "experimental/api-playground"; - case "tag-management": - return "experimental/tag-management"; - case "claude-code-plugins": - return "experimental/claude-code-plugins"; - case "usage": // "Old Usage" - return "experimental/old-usage"; - - // settings - case "general-settings": - return "settings/router-settings"; - case "settings": // "Logging & Alerts" - return "settings/logging-and-alerts"; - case "admin-panel": - return "settings/admin-settings"; - case "ui-theme": - return "settings/ui-theme"; - - default: - // treat as already a relative path - return slug.replace(/^\/+/, ""); - } -}; - -/** Prefix base path ("/" or "/ui/") */ -const toHref = (slugOrPath: string) => { - const base = getBasePath(); // "/" or "/ui/" - const rel = routeFor(slugOrPath).replace(/^\/+|\/+$/g, ""); - return `${base}${rel}`; -}; - -// ----- Menu config (unchanged labels/icons; same appearance) ----- -const menuItems: MenuItemCfg[] = [ - { key: "1", page: "api-keys", label: "Virtual Keys", icon: }, - { - key: "3", - page: "llm-playground", - label: "Test Key", - icon: , - roles: rolesWithWriteAccess, - }, - { - key: "2", - page: "models", - label: "Models + Endpoints", - icon: , - roles: rolesWithWriteAccess, - }, - { - key: "12", - page: "new_usage", - label: "Usage", - icon: , - roles: [...all_admin_roles, ...internalUserRoles], - }, - { key: "6", page: "teams", label: "Teams", icon: }, - { - key: "17", - page: "organizations", - label: "Organizations", - icon: , - roles: all_admin_roles, - }, - { - key: "5", - page: "users", - label: "Internal Users", - icon: , - roles: all_admin_roles, - }, - { key: "14", page: "api-reference", label: "API Reference", icon: }, - { - key: "16", - page: "model-hub-table", - label: "Model Hub", - icon: , - }, - { key: "15", page: "logs", label: "Logs", icon: }, - { - key: "11", - page: "guardrails", - label: "Guardrails", - icon: , - roles: all_admin_roles, - }, - { - key: "28", - page: "policies", - label: "Policies", - icon: , - roles: all_admin_roles, - }, - { - key: "26", - page: "tools", - label: "Tools", - icon: , - children: [ - { key: "18", page: "mcp-servers", label: "MCP Servers", icon: }, - { - key: "21", - page: "vector-stores", - label: "Vector Stores", - icon: , - roles: all_admin_roles, - }, - ], - }, - { - key: "experimental", - page: "experimental", - label: "Experimental", - icon: , - children: [ - { - key: "9", - page: "caching", - label: "Caching", - icon: , - roles: all_admin_roles, - }, - { - key: "25", - page: "prompts", - label: "Prompts", - icon: , - roles: all_admin_roles, - }, - { - key: "10", - page: "budgets", - label: "Budgets", - icon: , - roles: all_admin_roles, - }, - { - key: "20", - page: "transform-request", - label: "API Playground", - icon: , - roles: [...all_admin_roles, ...internalUserRoles], - }, - { - key: "19", - page: "tag-management", - label: "Tag Management", - icon: , - roles: all_admin_roles, - }, - { - key: "27", - page: "claude-code-plugins", - label: "Claude Code Plugins", - icon: , - roles: all_admin_roles, - }, - { key: "4", page: "usage", label: "Old Usage", icon: }, - ], - }, - { - key: "settings", - page: "settings", - label: "Settings", - icon: , - roles: all_admin_roles, - children: [ - { - key: "11", - page: "general-settings", - label: "Router Settings", - icon: , - roles: all_admin_roles, - }, - { - key: "8", - page: "settings", - label: "Logging & Alerts", - icon: , - roles: all_admin_roles, - }, - { - key: "13", - page: "admin-panel", - label: "Admin Settings", - icon: , - roles: all_admin_roles, - }, - { - key: "14", - page: "ui-theme", - label: "UI Theme", - icon: , - roles: all_admin_roles, - }, - ], - }, -]; - -const Sidebar2: React.FC = ({ accessToken, userRole, defaultSelectedKey, collapsed = false }) => { - const router = useRouter(); - const pathname = usePathname() || "/"; - - // ----- Filter by role without mutating originals ----- - const filteredMenuItems = React.useMemo(() => { - return menuItems - .filter((item) => !item.roles || item.roles.includes(userRole)) - .map((item) => ({ - ...item, - children: item.children ? item.children.filter((c) => !c.roles || c.roles.includes(userRole)) : undefined, - })); - }, [userRole]); - - // ----- Compute selected key from current path ----- - const selectedMenuKey = React.useMemo(() => { - const base = getBasePath(); - // strip base prefix and leading slash -> "virtual-keys", "tools/mcp-servers", etc. - const rel = pathname.startsWith(base) ? pathname.slice(base.length) : pathname.replace(/^\/+/, ""); - const relLower = rel.toLowerCase(); - - const matchesPath = (slug: string) => { - const route = routeFor(slug).toLowerCase(); - return relLower === route || relLower.startsWith(`${route}/`); - }; - - // search top-level - for (const item of filteredMenuItems) { - if (!item.children && matchesPath(item.page)) return item.key; - if (item.children) { - for (const child of item.children) { - if (matchesPath(child.page)) return child.key; - } - } - } - - // fallback to legacy defaultSelectedKey mapping - const fallback = filteredMenuItems.find((i) => i.page === defaultSelectedKey)?.key; - if (fallback) return fallback; - - for (const item of filteredMenuItems) { - if (item.children?.some((c) => c.page === defaultSelectedKey)) { - const child = item.children.find((c) => c.page === defaultSelectedKey)!; - return child.key; - } - } - - return "1"; - }, [pathname, filteredMenuItems, defaultSelectedKey]); - - // ----- Navigation ----- - const goTo = (slug: string, newTab?: boolean) => { - const href = toHref(slug); - if (newTab) { - window.open(href, "_blank"); - } else { - router.push(href); - } - }; - - // Wrap label in so every nav item supports right-click → "Open in new tab" - // and Ctrl/Cmd+click to open in a new tab, while preserving SPA navigation for normal clicks. - const renderNavLink = (label: string, page: string, newTab?: boolean): React.ReactNode => { - const href = toHref(page); - return ( - { - if (newTab) { - e.stopPropagation(); - return; - } - if (e.metaKey || e.ctrlKey || e.shiftKey || e.button === 1) { - e.stopPropagation(); - return; - } - e.preventDefault(); - }} - style={{ color: "inherit", textDecoration: "none" }} - > - {label} - - ); - }; - - return ( - - - - ({ - key: item.key, - icon: item.icon, - label: renderNavLink(item.label, item.page, item.newTab), - children: item.children?.map((child) => ({ - key: child.key, - icon: child.icon, - label: renderNavLink(child.label, child.page, child.newTab), - onClick: () => goTo(child.page, child.newTab), - })), - onClick: !item.children ? () => goTo(item.page, item.newTab) : undefined, - }))} - /> - - {isAdminRole(userRole) && !collapsed && } - - - ); -}; - -export default Sidebar2; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx b/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx index a611d619cc1..7dc2b2ad615 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx @@ -7,6 +7,7 @@ import SidebarProvider from "@/app/(dashboard)/components/SidebarProvider"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { useRouter, useSearchParams } from "next/navigation"; import { DebugWarningBanner } from "@/components/DebugWarningBanner"; +import { MIGRATED_PAGES } from "@/utils/migratedPages"; /** ---- BASE URL HELPERS ---- */ function normalizeBasePrefix(raw: string | undefined | null): string { @@ -23,15 +24,6 @@ function withBase(path: string): string { } /** -------------------------------- */ -/** - * Pages that have been migrated to path-based routing under (dashboard)/. - * When the leftnav triggers one of these, navigate to the path route instead - * of the legacy query-param root page. - * - * Key = legacy page id used in leftnav, Value = route segment under (dashboard)/ - */ -const MIGRATED_PAGES: Record = {}; - function LayoutContent({ children }: { children: React.ReactNode }) { const router = useRouter(); const searchParams = useSearchParams(); diff --git a/ui/litellm-dashboard/src/app/page.tsx b/ui/litellm-dashboard/src/app/page.tsx index da6a0d5a76f..f329a00d60d 100644 --- a/ui/litellm-dashboard/src/app/page.tsx +++ b/ui/litellm-dashboard/src/app/page.tsx @@ -54,6 +54,7 @@ import { storeReturnUrl, } from "@/utils/returnUrlUtils"; import { isAdminRole } from "@/utils/roles"; +import { MIGRATED_PAGES } from "@/utils/migratedPages"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { useRouter, useSearchParams } from "next/navigation"; import { Suspense, useEffect, useMemo, useRef, useState } from "react"; @@ -65,13 +66,6 @@ interface ProxySettings { LITELLM_UI_API_DOC_BASE_URL?: string | null; } -/** - * Map of legacy query-param page keys → new path-based route segments. - * When a user visits ?page=, they are redirected to /ui/. - * Add entries here as pages are migrated from the if/else chain to path-based routes. - */ -const LEGACY_REDIRECTS: Record = {}; - function CreateKeyPageContent() { const { authLoading, token, userID, userRole, userEmail, accessToken, premiumUser, setUserRole, setUserEmail } = useAuth(); @@ -202,11 +196,11 @@ function CreateKeyPageContent() { }, [redirectToLogin]); // Redirect legacy query-param pages to their new path-based routes - const isLegacyRedirect = page in LEGACY_REDIRECTS; + const isLegacyRedirect = page in MIGRATED_PAGES; useEffect(() => { if (!authLoading && isLegacyRedirect) { const base = (proxyBaseUrl || "") + "/ui"; - router.replace(`${base}/${LEGACY_REDIRECTS[page]}`); + router.replace(`${base}/${MIGRATED_PAGES[page]}`); } }, [authLoading, isLegacyRedirect, page, router]); diff --git a/ui/litellm-dashboard/src/components/leftnav.tsx b/ui/litellm-dashboard/src/components/leftnav.tsx index 43a28d94046..ff0dac703d3 100644 --- a/ui/litellm-dashboard/src/components/leftnav.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.tsx @@ -43,31 +43,9 @@ import { import NewBadge from "./common_components/NewBadge"; import type { Organization } from "./networking"; import UsageIndicator from "./UsageIndicator"; -import { serverRootPath } from "./networking"; +import { MIGRATED_PAGES, migratedHref } from "@/utils/migratedPages"; const { Sider } = Layout; -/** - * Pages migrated to path-based routing under (dashboard)/. - * Key = legacy page id, Value = route segment. - * Keep in sync with MIGRATED_PAGES in (dashboard)/layout.tsx. - */ -const MIGRATED_PAGES: Record = {}; - -/** Build an absolute href for a migrated page, respecting base URL + serverRootPath. */ -function migratedHref(routeSegment: string): string { - const raw = process.env.NEXT_PUBLIC_BASE_URL ?? ""; - const trimmed = raw.replace(/^\/+|\/+$/g, ""); - let base = trimmed ? `/${trimmed}/` : "/"; - - if (serverRootPath && serverRootPath !== "/") { - const cleanRoot = serverRootPath.replace(/\/+$/, ""); - const cleanBase = base.replace(/^\/+/, ""); - base = `${cleanRoot}/${cleanBase}`; - } - - return `${base}${routeSegment}`; -} - // Define the props type interface SidebarProps { setPage: (page: string) => void; diff --git a/ui/litellm-dashboard/src/utils/migratedPages.test.ts b/ui/litellm-dashboard/src/utils/migratedPages.test.ts new file mode 100644 index 00000000000..2f9ed11b3a0 --- /dev/null +++ b/ui/litellm-dashboard/src/utils/migratedPages.test.ts @@ -0,0 +1,43 @@ +import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; + +describe("migratedHref", () => { + beforeEach(() => { + vi.resetModules(); + vi.unstubAllEnvs(); + }); + + afterEach(() => { + vi.unstubAllEnvs(); + }); + + it("returns an absolute path rooted at / when serverRootPath is /", async () => { + vi.doMock("@/components/networking", () => ({ serverRootPath: "/" })); + const { migratedHref } = await import("./migratedPages"); + + expect(migratedHref("api-reference")).toBe("/api-reference"); + expect(migratedHref("virtual-keys")).toBe("/virtual-keys"); + }); + + it("prefixes a non-root serverRootPath without duplicating slashes", async () => { + vi.doMock("@/components/networking", () => ({ serverRootPath: "/team-x/" })); + const { migratedHref } = await import("./migratedPages"); + + expect(migratedHref("api-reference")).toBe("/team-x/api-reference"); + }); + + it("honors NEXT_PUBLIC_BASE_URL as a base path segment", async () => { + vi.stubEnv("NEXT_PUBLIC_BASE_URL", "ui"); + vi.doMock("@/components/networking", () => ({ serverRootPath: "/" })); + const { migratedHref } = await import("./migratedPages"); + + expect(migratedHref("api-reference")).toBe("/ui/api-reference"); + }); + + it("combines NEXT_PUBLIC_BASE_URL with a non-root serverRootPath", async () => { + vi.stubEnv("NEXT_PUBLIC_BASE_URL", "ui"); + vi.doMock("@/components/networking", () => ({ serverRootPath: "/team-x" })); + const { migratedHref } = await import("./migratedPages"); + + expect(migratedHref("api-reference")).toBe("/team-x/ui/api-reference"); + }); +}); diff --git a/ui/litellm-dashboard/src/utils/migratedPages.ts b/ui/litellm-dashboard/src/utils/migratedPages.ts new file mode 100644 index 00000000000..5421f00bf05 --- /dev/null +++ b/ui/litellm-dashboard/src/utils/migratedPages.ts @@ -0,0 +1,25 @@ +import { serverRootPath } from "@/components/networking"; + +/** + * Single source of truth for pages cut over from the legacy `?page=` switch in + * app/page.tsx to path-based routes under app/(dashboard)/. + * + * Key = legacy page id emitted by the sidebar. Value = route segment under (dashboard)/. + * Add an entry to route the sidebar and deep links to the new path and redirect the + * legacy `?page=` URL; remove it to roll back. Empty until a page is migrated. + */ +export const MIGRATED_PAGES: Record = {}; + +export function migratedHref(routeSegment: string): string { + const raw = process.env.NEXT_PUBLIC_BASE_URL ?? ""; + const trimmed = raw.replace(/^\/+|\/+$/g, ""); + let base = trimmed ? `/${trimmed}/` : "/"; + + if (serverRootPath && serverRootPath !== "/") { + const cleanRoot = serverRootPath.replace(/\/+$/, ""); + const cleanBase = base.replace(/^\/+/, ""); + base = `${cleanRoot}/${cleanBase}`; + } + + return `${base}${routeSegment}`; +} From 26fe26a5c041d0aa03f9e669b255d89b2ab67859 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Mon, 8 Jun 2026 12:12:07 -0700 Subject: [PATCH 03/22] fix(ui/model-hub): render provider icons on the public model hub (#29958) The provider logo base path was relative ("../ui/assets/logos/"). With trailingSlash enabled, the public model hub is served at /ui/model_hub_table/, so the browser resolved the base to /ui/ui/assets/logos/ (a doubled /ui/), which 404s every icon. The authenticated hub renders inside the single-level /ui/ SPA route where the relative path resolves correctly, so only the public hub broke. Make the base root-absolute so it resolves at any route depth. --- .../components/provider_info_helpers.test.tsx | 21 +++++++++++++++++++ .../src/components/provider_info_helpers.tsx | 2 +- 2 files changed, 22 insertions(+), 1 deletion(-) diff --git a/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx b/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx index ab22b0ef49a..a73cf699ac8 100644 --- a/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx +++ b/ui/litellm-dashboard/src/components/provider_info_helpers.test.tsx @@ -97,6 +97,27 @@ describe("provider_info_helpers", () => { }); }); + describe("provider logo asset paths", () => { + // Regression: a relative "../ui/assets/logos/" base resolved to + // "/ui/ui/assets/logos/..." (404) on the public model hub at + // /ui/model_hub_table/, which sits a level below the /ui/ SPA. Root-absolute + // paths resolve correctly at any route depth. + it("should expose every provider logo as a root-absolute /ui path", () => { + const logos = Object.values(providerLogoMap); + expect(logos.length).toBeGreaterThan(0); + logos.forEach((logo) => { + expect(logo.startsWith("/ui/assets/logos/")).toBe(true); + expect(logo).not.toContain("../"); + }); + }); + + it("should resolve a provider logo to a root-absolute path via getProviderLogoAndName", () => { + const { logo } = getProviderLogoAndName("openai"); + expect(logo.startsWith("/ui/assets/logos/")).toBe(true); + expect(logo).not.toContain("../"); + }); + }); + describe("getPlaceholder", () => { it("should return aiml placeholder for AIML provider", () => { expect(getPlaceholder(Providers.AIML)).toBe("aiml/flux-pro/v1.1"); diff --git a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx index 179ce63f457..727eebfe951 100644 --- a/ui/litellm-dashboard/src/components/provider_info_helpers.tsx +++ b/ui/litellm-dashboard/src/components/provider_info_helpers.tsx @@ -216,7 +216,7 @@ export const provider_map: Record = { ZAI: "zai", }; -const asset_logos_folder = "../ui/assets/logos/"; +const asset_logos_folder = "/ui/assets/logos/"; export const providerLogoMap: Record = { [Providers.A2A_Agent]: `${asset_logos_folder}a2a_agent.png`, From 47b383dbbfc9eb3336ee1b7e57a0e7d4d380c46f Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Mon, 8 Jun 2026 12:25:42 -0700 Subject: [PATCH 04/22] fix(ui): keep create guardrail modal open on outside click (#29871) The create guardrail modal used antd's default maskClosable, so clicking outside it dismissed the modal and reset every field the user had entered. Setting maskClosable={false} keeps the modal open; it now closes only via the explicit close button or Cancel, matching the other form modals in the dashboard --- .../guardrails/add_guardrail_form.test.tsx | 43 +++++++++++++++++++ .../guardrails/add_guardrail_form.tsx | 1 + 2 files changed, 44 insertions(+) create mode 100644 ui/litellm-dashboard/src/components/guardrails/add_guardrail_form.test.tsx diff --git a/ui/litellm-dashboard/src/components/guardrails/add_guardrail_form.test.tsx b/ui/litellm-dashboard/src/components/guardrails/add_guardrail_form.test.tsx new file mode 100644 index 00000000000..a13f12e4b42 --- /dev/null +++ b/ui/litellm-dashboard/src/components/guardrails/add_guardrail_form.test.tsx @@ -0,0 +1,43 @@ +import React from "react"; +import { fireEvent, screen } from "@testing-library/react"; +import { renderWithProviders } from "../../../tests/test-utils"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import AddGuardrailForm from "./add_guardrail_form"; + +vi.mock("@/components/networking", () => ({ + createGuardrailCall: vi.fn(), + getGuardrailProviderSpecificParams: vi.fn().mockResolvedValue({}), + getGuardrailUISettings: vi.fn().mockResolvedValue({}), + modelAvailableCall: vi.fn().mockResolvedValue({ data: [] }), +})); + +const renderForm = () => { + const onClose = vi.fn(); + renderWithProviders(); + return { onClose }; +}; + +describe("AddGuardrailForm close behavior", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("does not close when the user clicks outside the modal on the mask", () => { + const { onClose } = renderForm(); + expect(screen.getByText("Create guardrail")).toBeInTheDocument(); + + const wrap = document.querySelector(".ant-modal-wrap") as HTMLElement; + expect(wrap).toBeTruthy(); + fireEvent.mouseDown(wrap); + fireEvent.mouseUp(wrap); + fireEvent.click(wrap); + + expect(onClose).not.toHaveBeenCalled(); + }); + + it("closes when the user clicks the explicit close button", () => { + const { onClose } = renderForm(); + fireEvent.click(screen.getByRole("button", { name: "✕" })); + expect(onClose).toHaveBeenCalledTimes(1); + }); +}); diff --git a/ui/litellm-dashboard/src/components/guardrails/add_guardrail_form.tsx b/ui/litellm-dashboard/src/components/guardrails/add_guardrail_form.tsx index 05a65d1c9e5..41a2dd7f67c 100644 --- a/ui/litellm-dashboard/src/components/guardrails/add_guardrail_form.tsx +++ b/ui/litellm-dashboard/src/components/guardrails/add_guardrail_form.tsx @@ -1144,6 +1144,7 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a title={null} open={visible} onCancel={handleClose} + maskClosable={false} footer={null} width={1000} closable={false} From 728f057c5ec3334604fc65c57c41bac873d4555f Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Mon, 8 Jun 2026 12:25:52 -0700 Subject: [PATCH 05/22] fix(ui): label default key type as "Full Access" on key edit page (#29870) The key edit page showed the default key type (no allowed_routes restriction) as "Default", while the key creation form already labels the same value "Full Access". Align the edit page to the create form so the two surfaces agree on both the label and its description. --- .../src/components/templates/key_edit_view.tsx | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx index d5e410029a7..fe0fa1ab0c2 100644 --- a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx @@ -394,11 +394,11 @@ export function KeyEditView({ } }} > - +
-
Default
+
Full Access
- Can call AI APIs + Management routes + Can call all routes (AI APIs, Management, and read-only)
From 1afc41cb295e78a25a5d67389ac5a831e290ee8d Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Mon, 8 Jun 2026 13:05:12 -0700 Subject: [PATCH 06/22] fix(ui): unify migrated-route URLs and migrate the API Reference page (#29953) * fix(ui): unify migrated-route URLs and cut the API Reference page over to path routing Route all migrated-page navigation through one /ui-prefixed, serverRootPath-aware builder in migratedPages.ts (migratedHref/legacyPageHref/legacyKeyForPathname), replacing the three divergent base-URL constructions that lived in the dashboard layout's withBase, leftnav, and the page.tsx redirect. The previous migratedHref read NEXT_PUBLIC_BASE_URL, which no build sets, so it produced URLs without the /ui prefix the app is served under; every other internal link hardcodes /ui and this now matches that convention. Remove the sidebar's own pushState navigation so the parent (legacy root page or dashboard layout) is the single owner of navigation, fixing the double-navigate that fired when moving between a path route and a legacy ?page= route. Cut API Reference over to its path route: add api_ref -> api-reference to MIGRATED_PAGES and delete its arm from the legacy switch. Visiting /ui/?page=api_ref redirects to /ui/api-reference, the sidebar links to and highlights it, and navigating away returns to the legacy switch. * fix(ui): address review on migrated-page routing Keep the legacy hyphenated ?page=api-reference form working by mapping it to the api-reference route alongside api_ref; the old switch matched both, so a bookmark using the hyphen would otherwise fall through to the Usage default. Add legacyKeyForPathname coverage: a migrated path (with and without trailing slash) resolves to the api_ref sidebar key rather than the alias, a non-migrated path returns null, and a non-root serverRootPath prefix is stripped before matching. * fix(ui): populate serverRootPath from getUiConfig so migrated nav links keep the root path getUiConfig updated proxyBaseUrl but never called updateServerRootPath, so the module-level serverRootPath stayed at its "/" default. Under a custom server_root_path the unified migratedHref/legacyPageHref builders then dropped the prefix and the sidebar produced /ui/api-reference (404) instead of //ui/api-reference. Adds the missing updateServerRootPath call plus a regression test asserting getUiConfig sets serverRootPath and that migratedHref carries the prefix --- .../src/app/(dashboard)/layout.tsx | 36 ++-------- ui/litellm-dashboard/src/app/page.tsx | 19 +++--- .../src/components/leftnav.tsx | 26 ++----- .../src/components/networking.test.ts | 15 ++++ .../src/components/networking.tsx | 1 + .../src/utils/migratedPages.test.ts | 68 +++++++++++++------ .../src/utils/migratedPages.ts | 43 ++++++++---- 7 files changed, 112 insertions(+), 96 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx b/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx index 7dc2b2ad615..5f5c240d025 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/layout.tsx @@ -5,51 +5,29 @@ import Navbar from "@/components/navbar"; import { ThemeProvider } from "@/contexts/ThemeContext"; import SidebarProvider from "@/app/(dashboard)/components/SidebarProvider"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; -import { useRouter, useSearchParams } from "next/navigation"; +import { useRouter, useSearchParams, usePathname } from "next/navigation"; import { DebugWarningBanner } from "@/components/DebugWarningBanner"; -import { MIGRATED_PAGES } from "@/utils/migratedPages"; - -/** ---- BASE URL HELPERS ---- */ -function normalizeBasePrefix(raw: string | undefined | null): string { - const trimmed = (raw ?? "").trim(); - if (!trimmed) return ""; - const core = trimmed.replace(/^\/+/, "").replace(/\/+$/, ""); - return core ? `/${core}/` : "/"; -} -const BASE_PREFIX = normalizeBasePrefix(process.env.NEXT_PUBLIC_BASE_URL); -function withBase(path: string): string { - const body = path.startsWith("/") ? path.slice(1) : path; - const combined = `${BASE_PREFIX}${body}`; - return combined.startsWith("/") ? combined : `/${combined}`; -} -/** -------------------------------- */ +import { MIGRATED_PAGES, migratedHref, legacyPageHref, legacyKeyForPathname } from "@/utils/migratedPages"; function LayoutContent({ children }: { children: React.ReactNode }) { const router = useRouter(); const searchParams = useSearchParams(); + const pathname = usePathname(); const { accessToken } = useAuthorized(); const [sidebarCollapsed, setSidebarCollapsed] = React.useState(false); const [page, setPage] = useState(() => { - return searchParams.get("page") || "api-keys"; + return legacyKeyForPathname(pathname) || searchParams.get("page") || "api-keys"; }); const handleSetPage = (newPage: string) => { - // If the page has been migrated to path routing, navigate there const migratedRoute = MIGRATED_PAGES[newPage]; - if (migratedRoute) { - router.push(withBase(migratedRoute)); - setPage(newPage); - return; - } - - // Otherwise, navigate back to the legacy root page with query params - router.push(withBase(`?page=${newPage}`)); + router.push(migratedRoute ? migratedHref(migratedRoute) : legacyPageHref(newPage)); setPage(newPage); }; useEffect(() => { - setPage(searchParams.get("page") || "api-keys"); - }, [searchParams]); + setPage(legacyKeyForPathname(pathname) || searchParams.get("page") || "api-keys"); + }, [pathname, searchParams]); const toggleSidebar = () => setSidebarCollapsed((v) => !v); diff --git a/ui/litellm-dashboard/src/app/page.tsx b/ui/litellm-dashboard/src/app/page.tsx index f329a00d60d..12dd39a1c21 100644 --- a/ui/litellm-dashboard/src/app/page.tsx +++ b/ui/litellm-dashboard/src/app/page.tsx @@ -1,6 +1,5 @@ "use client"; -import APIReferenceView from "@/app/(dashboard)/api-reference/APIReferenceView"; import SidebarProvider from "@/app/(dashboard)/components/SidebarProvider"; import OldModelDashboard from "@/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView"; import PlaygroundPage from "@/app/(dashboard)/playground/page"; @@ -54,7 +53,7 @@ import { storeReturnUrl, } from "@/utils/returnUrlUtils"; import { isAdminRole } from "@/utils/roles"; -import { MIGRATED_PAGES } from "@/utils/migratedPages"; +import { MIGRATED_PAGES, migratedHref } from "@/utils/migratedPages"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { useRouter, useSearchParams } from "next/navigation"; import { Suspense, useEffect, useMemo, useRef, useState } from "react"; @@ -156,15 +155,16 @@ function CreateKeyPageContent() { return searchParams.get("page") || "api-keys"; }); - // Custom setPage function that updates URL const updatePage = (newPage: string) => { - // Update URL without full page reload + const migratedRoute = MIGRATED_PAGES[newPage]; + if (migratedRoute) { + router.push(migratedHref(migratedRoute)); + setPage(newPage); + return; + } const newSearchParams = new URLSearchParams(searchParams); newSearchParams.set("page", newPage); - - // Use Next.js router to update URL window.history.pushState(null, "", `?${newSearchParams.toString()}`); - setPage(newPage); }; @@ -199,8 +199,7 @@ function CreateKeyPageContent() { const isLegacyRedirect = page in MIGRATED_PAGES; useEffect(() => { if (!authLoading && isLegacyRedirect) { - const base = (proxyBaseUrl || "") + "/ui"; - router.replace(`${base}/${MIGRATED_PAGES[page]}`); + router.replace(migratedHref(MIGRATED_PAGES[page])); } }, [authLoading, isLegacyRedirect, page, router]); @@ -441,8 +440,6 @@ function CreateKeyPageContent() { /> ) : page == "admin-panel" ? ( - ) : page == "api_ref" || page == "api-reference" ? ( - ) : page == "logging-and-alerts" ? ( ) : page == "budgets" ? ( diff --git a/ui/litellm-dashboard/src/components/leftnav.tsx b/ui/litellm-dashboard/src/components/leftnav.tsx index ff0dac703d3..f4946b81d68 100644 --- a/ui/litellm-dashboard/src/components/leftnav.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.tsx @@ -43,7 +43,7 @@ import { import NewBadge from "./common_components/NewBadge"; import type { Organization } from "./networking"; import UsageIndicator from "./UsageIndicator"; -import { MIGRATED_PAGES, migratedHref } from "@/utils/migratedPages"; +import { MIGRATED_PAGES, migratedHref, legacyPageHref } from "@/utils/migratedPages"; const { Sider } = Layout; // Define the props type @@ -418,18 +418,9 @@ const Sidebar: React.FC = ({ // Check if user is a team admin for any team const isTeamAdmin = useMemo(() => isUserTeamAdminForAnyTeam(teams ?? null, userId ?? ""), [teams, userId]); - // Navigate to page helper - const navigateToPage = (page: string) => { - // For migrated pages, just call setPage — the parent layout handles routing - if (MIGRATED_PAGES[page]) { - setPage(page); - return; - } - const newSearchParams = new URLSearchParams(window.location.search); - newSearchParams.set("page", page); - window.history.pushState(null, "", `?${newSearchParams.toString()}`); - setPage(page); - }; + // The parent (legacy root page or dashboard layout) owns navigation for both + // migrated and legacy pages; the sidebar only reports the selected page. + const navigateToPage = (page: string) => setPage(page); // Wrap label in so every nav item supports right-click → "Open in new tab" // and Ctrl/Cmd+click to open in a new tab, while preserving SPA navigation for normal clicks. @@ -447,15 +438,8 @@ const Sidebar: React.FC = ({ ); } - // For migrated pages, generate a path-based href for right-click "Open in new tab" const migratedRoute = MIGRATED_PAGES[page]; - const href = migratedRoute - ? migratedHref(migratedRoute) - : (() => { - const params = new URLSearchParams(window.location.search); - params.set("page", page); - return `?${params.toString()}`; - })(); + const href = migratedRoute ? migratedHref(migratedRoute) : legacyPageHref(page); return ( ({ clearTokenCookies: vi.fn(), @@ -349,6 +350,20 @@ describe("UI config and public endpoints", () => { ); expect(configCall).toBeDefined(); }); + + it("updates serverRootPath so path-based nav links carry the root path", async () => { + const uiConfig = { + server_root_path: "/litellm", + proxy_base_url: "https://example.com", + }; + + setupMockFetch([{ url: "/litellm/.well-known/litellm-ui-config", data: uiConfig }]); + + await Networking.getUiConfig(); + + expect(Networking.serverRootPath).toBe("/litellm"); + expect(migratedHref("api-reference")).toBe("/litellm/ui/api-reference"); + }); }); describe("individualModelHealthCheckCall", () => { diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index d303734dffd..54a7a19362b 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -404,6 +404,7 @@ export const getUiConfig = async () => { * Update the proxy base url and server root path */ console.log("jsonData in getUiConfig:", jsonData); + updateServerRootPath(jsonData.server_root_path); updateProxyBaseUrl(jsonData.server_root_path, jsonData.proxy_base_url); return jsonData; }; diff --git a/ui/litellm-dashboard/src/utils/migratedPages.test.ts b/ui/litellm-dashboard/src/utils/migratedPages.test.ts index 2f9ed11b3a0..e9aceea8148 100644 --- a/ui/litellm-dashboard/src/utils/migratedPages.test.ts +++ b/ui/litellm-dashboard/src/utils/migratedPages.test.ts @@ -1,43 +1,69 @@ -import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; +import { describe, it, expect, vi, beforeEach } from "vitest"; -describe("migratedHref", () => { +describe("migratedHref / legacyPageHref", () => { beforeEach(() => { vi.resetModules(); - vi.unstubAllEnvs(); }); - afterEach(() => { - vi.unstubAllEnvs(); - }); - - it("returns an absolute path rooted at / when serverRootPath is /", async () => { + it("builds a /ui-rooted path when serverRootPath is /", async () => { vi.doMock("@/components/networking", () => ({ serverRootPath: "/" })); - const { migratedHref } = await import("./migratedPages"); + const { migratedHref, legacyPageHref } = await import("./migratedPages"); - expect(migratedHref("api-reference")).toBe("/api-reference"); - expect(migratedHref("virtual-keys")).toBe("/virtual-keys"); + expect(migratedHref("api-reference")).toBe("/ui/api-reference"); + expect(legacyPageHref("models")).toBe("/ui/?page=models"); }); it("prefixes a non-root serverRootPath without duplicating slashes", async () => { vi.doMock("@/components/networking", () => ({ serverRootPath: "/team-x/" })); - const { migratedHref } = await import("./migratedPages"); + const { migratedHref, legacyPageHref } = await import("./migratedPages"); - expect(migratedHref("api-reference")).toBe("/team-x/api-reference"); + expect(migratedHref("api-reference")).toBe("/team-x/ui/api-reference"); + expect(legacyPageHref("models")).toBe("/team-x/ui/?page=models"); }); - it("honors NEXT_PUBLIC_BASE_URL as a base path segment", async () => { - vi.stubEnv("NEXT_PUBLIC_BASE_URL", "ui"); + it("tolerates a leading slash in the route segment", async () => { vi.doMock("@/components/networking", () => ({ serverRootPath: "/" })); const { migratedHref } = await import("./migratedPages"); - expect(migratedHref("api-reference")).toBe("/ui/api-reference"); + expect(migratedHref("/api-reference")).toBe("/ui/api-reference"); }); - it("combines NEXT_PUBLIC_BASE_URL with a non-root serverRootPath", async () => { - vi.stubEnv("NEXT_PUBLIC_BASE_URL", "ui"); - vi.doMock("@/components/networking", () => ({ serverRootPath: "/team-x" })); - const { migratedHref } = await import("./migratedPages"); + it("maps both the api_ref id and the hyphenated alias to the api-reference route", async () => { + vi.doMock("@/components/networking", () => ({ serverRootPath: "/" })); + const { MIGRATED_PAGES } = await import("./migratedPages"); - expect(migratedHref("api-reference")).toBe("/team-x/ui/api-reference"); + expect(MIGRATED_PAGES.api_ref).toBe("api-reference"); + expect(MIGRATED_PAGES["api-reference"]).toBe("api-reference"); + }); +}); + +describe("legacyKeyForPathname", () => { + beforeEach(() => { + vi.resetModules(); + }); + + it("maps a migrated path back to its legacy sidebar key (including trailing slash)", async () => { + vi.doMock("@/components/networking", () => ({ serverRootPath: "/" })); + const { legacyKeyForPathname } = await import("./migratedPages"); + + // Resolves to the sidebar key api_ref, not the hyphenated alias, so highlighting works. + expect(legacyKeyForPathname("/ui/api-reference")).toBe("api_ref"); + expect(legacyKeyForPathname("/ui/api-reference/")).toBe("api_ref"); + }); + + it("returns null for a not-yet-migrated path", async () => { + vi.doMock("@/components/networking", () => ({ serverRootPath: "/" })); + const { legacyKeyForPathname } = await import("./migratedPages"); + + expect(legacyKeyForPathname("/ui/")).toBeNull(); + expect(legacyKeyForPathname("/ui/some-legacy-page")).toBeNull(); + }); + + it("strips a non-root serverRootPath prefix before matching", async () => { + vi.doMock("@/components/networking", () => ({ serverRootPath: "/team-x/" })); + const { legacyKeyForPathname } = await import("./migratedPages"); + + expect(legacyKeyForPathname("/team-x/ui/api-reference")).toBe("api_ref"); + expect(legacyKeyForPathname("/ui/api-reference")).toBeNull(); }); }); diff --git a/ui/litellm-dashboard/src/utils/migratedPages.ts b/ui/litellm-dashboard/src/utils/migratedPages.ts index 5421f00bf05..2c27e4fee64 100644 --- a/ui/litellm-dashboard/src/utils/migratedPages.ts +++ b/ui/litellm-dashboard/src/utils/migratedPages.ts @@ -6,20 +6,35 @@ import { serverRootPath } from "@/components/networking"; * * Key = legacy page id emitted by the sidebar. Value = route segment under (dashboard)/. * Add an entry to route the sidebar and deep links to the new path and redirect the - * legacy `?page=` URL; remove it to roll back. Empty until a page is migrated. + * legacy `?page=` URL; remove it to roll back. */ -export const MIGRATED_PAGES: Record = {}; +export const MIGRATED_PAGES: Record = { + api_ref: "api-reference", + // Legacy alias: older bookmarks used the hyphenated ?page=api-reference form. + "api-reference": "api-reference", +}; -export function migratedHref(routeSegment: string): string { - const raw = process.env.NEXT_PUBLIC_BASE_URL ?? ""; - const trimmed = raw.replace(/^\/+|\/+$/g, ""); - let base = trimmed ? `/${trimmed}/` : "/"; - - if (serverRootPath && serverRootPath !== "/") { - const cleanRoot = serverRootPath.replace(/\/+$/, ""); - const cleanBase = base.replace(/^\/+/, ""); - base = `${cleanRoot}/${cleanBase}`; - } - - return `${base}${routeSegment}`; +function uiBase(): string { + const root = serverRootPath && serverRootPath !== "/" ? `/${serverRootPath.replace(/^\/+|\/+$/g, "")}` : ""; + return `${root}/ui`; +} + +/** Absolute (same-origin) href for a migrated route segment, e.g. "api-reference" -> "/ui/api-reference". */ +export function migratedHref(routeSegment: string): string { + return `${uiBase()}/${routeSegment.replace(/^\/+/, "")}`; +} + +/** Href for a not-yet-migrated page, served by the legacy `?page=` switch at the UI root. */ +export function legacyPageHref(pageKey: string): string { + return `${uiBase()}/?page=${pageKey}`; +} + +/** Reverse-maps a path-routed location back to its legacy page id, e.g. "/ui/api-reference" -> "api_ref". */ +export function legacyKeyForPathname(pathname: string): string | null { + const base = uiBase(); + const rel = (pathname.startsWith(base) ? pathname.slice(base.length) : pathname).replace(/^\/+|\/+$/g, ""); + for (const [key, segment] of Object.entries(MIGRATED_PAGES)) { + if (rel === segment) return key; + } + return null; } From 1528f43d4c5f6c9981d065761834957cf559f7b5 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Mon, 8 Jun 2026 13:35:49 -0700 Subject: [PATCH 07/22] fix(mcp): let non-creator users OAuth into OBO-mode MCP servers from the Tools page (#29867) * fix(ui): let non-creator users OAuth into OBO-mode MCP servers from the Tools page * fix(ui): clear OBO Tools-tab one-shot on navigate-back and gate on credential-status errors --- .../components/mcp_tools/mcp_server_view.tsx | 4 +- .../src/components/mcp_tools/mcp_servers.tsx | 41 +++++- .../components/mcp_tools/mcp_tools.test.tsx | 55 +++++++- .../src/components/mcp_tools/mcp_tools.tsx | 119 +++++++++++++++--- .../src/hooks/mcpOAuthUtils.ts | 9 ++ 5 files changed, 209 insertions(+), 19 deletions(-) diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx index cafecdafd45..fc1fe7cd405 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx @@ -21,6 +21,7 @@ interface MCPServerViewProps { userRole: string | null; userID: string | null; availableAccessGroups: string[]; + initialTabIndex?: number; } export const MCPServerView: React.FC = ({ @@ -32,11 +33,12 @@ export const MCPServerView: React.FC = ({ userRole, userID, availableAccessGroups, + initialTabIndex = 0, }) => { const [editing, setEditing] = useState(isEditing); const [showFullUrl, setShowFullUrl] = useState(false); const [copiedStates, setCopiedStates] = useState>({}); - const [selectedTabIndex, setSelectedTabIndex] = useState(0); + const [selectedTabIndex, setSelectedTabIndex] = useState(initialTabIndex); const handleSuccess = (updated: MCPServer) => { setEditing(false); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx index 4b383f701cb..0a445265e2d 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_servers.tsx @@ -21,6 +21,7 @@ import MCPNetworkSettings from "./MCPNetworkSettings"; import MCPDiscovery from "./mcp_discovery"; import { ByokCredentialModal } from "./ByokCredentialModal"; import { getSecureItem } from "@/utils/secureStorage"; +import { TOOLS_OAUTH_UI_STATE_KEY } from "@/hooks/mcpOAuthUtils"; import UserEnvVarsModal from "./UserEnvVarsModal"; import { listMCPUserEnvVarStatus } from "../networking"; @@ -71,6 +72,23 @@ const compareServers = (a: MCPServer, b: MCPServer, sort: SortKey): number => { const { Text: AntdText, Title: AntdTitle } = Typography; const EDIT_OAUTH_UI_STATE_KEY = "litellm-mcp-oauth-edit-state"; +// Server id stashed by the Tools tab before an OBO OAuth redirect, read once at +// mount so the redirect returns straight to that server's Tools tab. +const readToolsOAuthServerId = (): string | null => { + if (typeof window === "undefined") { + return null; + } + try { + const stored = getSecureItem(TOOLS_OAUTH_UI_STATE_KEY); + if (!stored) { + return null; + } + return JSON.parse(stored)?.serverId ?? null; + } catch { + return null; + } +}; + const { Option } = Select; const MCPServers: React.FC = ({ accessToken, userRole, userID }) => { @@ -103,7 +121,12 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) // state const [serverIdToDelete, setServerToDelete] = useState(null); const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); - const [selectedServerId, setSelectedServerId] = useState(null); + // Server whose Tools tab should be reopened after an OBO OAuth redirect; read + // once from sessionStorage so the restored server selection is correct on the + // first render. Cleared when the user navigates back to the list (handleBack) + // so a later visit to the same server defaults to Overview, not the Tools tab. + const [toolsTabServerId, setToolsTabServerId] = useState(readToolsOAuthServerId); + const [selectedServerId, setSelectedServerId] = useState(toolsTabServerId); const [editServer, setEditServer] = useState(false); const [selectedTeam, setSelectedTeam] = useState("all"); const [selectedMcpAccessGroup, setSelectedMcpAccessGroup] = useState("all"); @@ -178,6 +201,19 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) } }, []); + // The restored server id was consumed by the initializer above; remove the + // one-shot sessionStorage key so a full page reload doesn't reopen the Tools + // tab (removeItem only, no setState). + useEffect(() => { + if (typeof window !== "undefined") { + try { + window.sessionStorage.removeItem(TOOLS_OAUTH_UI_STATE_KEY); + } catch { + // ignore storage errors + } + } + }, []); + // Get unique teams from all servers const uniqueTeams = React.useMemo(() => { if (!serversWithHealth) return []; @@ -338,6 +374,8 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) const handleBack = React.useCallback(() => { setEditServer(false); setSelectedServerId(null); + // Drop the post-redirect one-shot so re-selecting that server opens Overview. + setToolsTabServerId(null); refetch(); }, [refetch]); @@ -483,6 +521,7 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID }) userID={userID} userRole={userRole} availableAccessGroups={uniqueMcpAccessGroups} + initialTabIndex={selectedServerId === toolsTabServerId ? 1 : 0} /> ) : (
diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.test.tsx index 4917f1fcf1d..72ed3984244 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.test.tsx @@ -2,12 +2,13 @@ import { render, screen, waitFor } from "@testing-library/react"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { describe, expect, it, vi, beforeEach } from "vitest"; import MCPToolsViewer from "./mcp_tools"; -import { listMCPTools } from "../networking"; +import { listMCPTools, getMCPOAuthUserCredentialStatus } from "../networking"; import { isTokenValid, getToken } from "@/utils/mcpTokenStore"; vi.mock("../networking", () => ({ listMCPTools: vi.fn(), callMCPTool: vi.fn(), + getMCPOAuthUserCredentialStatus: vi.fn(), })); vi.mock("@/utils/mcpTokenStore", () => ({ @@ -20,6 +21,10 @@ vi.mock("@/hooks/useToolsOAuthFlow", () => ({ useToolsOAuthFlow: () => ({ startOAuthFlow: vi.fn(), status: "idle", error: null }), })); +vi.mock("@/hooks/useUserMcpOAuthFlow", () => ({ + useUserMcpOAuthFlow: () => ({ startOAuthFlow: vi.fn(), status: "idle", error: null }), +})); + const GATE_TEXT = "Authentication required"; // Realistic interactive servers carry a token endpoint; the old heuristic // (`oauth2 && !tokenUrl`) mislabeled exactly these as M2M. Setting it here is @@ -42,6 +47,13 @@ const renderViewer = (props: Record) => , ); +const credStatus = (overrides: Record = {}) => ({ + server_id: "srv-1", + has_credential: true, + is_expired: false, + ...overrides, +}); + describe("MCPToolsViewer auth gate routing", () => { beforeEach(() => { vi.mocked(listMCPTools).mockReset().mockResolvedValue({ tools: [], error: null }); @@ -49,6 +61,8 @@ describe("MCPToolsViewer auth gate routing", () => { vi.mocked(getToken) .mockReset() .mockReturnValue(undefined as any); + // Default: the OBO credential exists and is valid, so OBO servers list tools. + vi.mocked(getMCPOAuthUserCredentialStatus).mockReset().mockResolvedValue(credStatus()); }); it("shows the Authorize gate for a passthrough server with a token endpoint and does not list tools", async () => { @@ -57,6 +71,8 @@ describe("MCPToolsViewer auth gate routing", () => { expect(await screen.findByText(GATE_TEXT)).toBeInTheDocument(); expect(screen.getByRole("button", { name: "Authorize" })).toBeInTheDocument(); expect(vi.mocked(listMCPTools)).not.toHaveBeenCalled(); + // Passthrough must not consult the per-user DB credential. + expect(vi.mocked(getMCPOAuthUserCredentialStatus)).not.toHaveBeenCalled(); }); it("forwards the session token via the x-mcp header for a passthrough server that has one", async () => { @@ -75,7 +91,40 @@ describe("MCPToolsViewer auth gate routing", () => { expect(screen.queryByText(GATE_TEXT)).not.toBeInTheDocument(); }); - it("does not gate an OBO server with a token endpoint; lists with the LiteLLM key and no x-mcp header", async () => { + it("lists tools for an OBO server when the user has a DB credential, with no x-mcp header", async () => { + renderViewer({ oauth2_flow: null, delegate_auth_to_upstream: false }); + + await waitFor(() => expect(vi.mocked(listMCPTools)).toHaveBeenCalledWith("litellm-key", "srv-1", undefined)); + expect(screen.queryByText(GATE_TEXT)).not.toBeInTheDocument(); + }); + + it("shows the Authorize gate for an OBO server when the user has no DB credential and does not list tools", async () => { + vi.mocked(getMCPOAuthUserCredentialStatus).mockResolvedValue(credStatus({ has_credential: false })); + + renderViewer({ oauth2_flow: null, delegate_auth_to_upstream: false }); + + expect(await screen.findByText(GATE_TEXT)).toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Authorize" })).toBeInTheDocument(); + expect(vi.mocked(listMCPTools)).not.toHaveBeenCalled(); + }); + + it("shows the Authorize gate for an OBO server when the credential-status check fails", async () => { + vi.mocked(getMCPOAuthUserCredentialStatus).mockRejectedValue(new Error("network down")); + + renderViewer({ oauth2_flow: null, delegate_auth_to_upstream: false }); + + expect(await screen.findByText(GATE_TEXT)).toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Authorize" })).toBeInTheDocument(); + expect(vi.mocked(listMCPTools)).not.toHaveBeenCalled(); + }); + + it("does not gate an OBO server whose stored token is expired; the list call refreshes it server-side", async () => { + // has_credential=true with is_expired=true must NOT gate: resolve_valid_user_oauth_token + // refreshes from the stored refresh_token on the list call, so the user never reauthorizes. + vi.mocked(getMCPOAuthUserCredentialStatus).mockResolvedValue( + credStatus({ has_credential: true, is_expired: true }), + ); + renderViewer({ oauth2_flow: null, delegate_auth_to_upstream: false }); await waitFor(() => expect(vi.mocked(listMCPTools)).toHaveBeenCalledWith("litellm-key", "srv-1", undefined)); @@ -87,5 +136,7 @@ describe("MCPToolsViewer auth gate routing", () => { await waitFor(() => expect(vi.mocked(listMCPTools)).toHaveBeenCalledWith("litellm-key", "srv-1", undefined)); expect(screen.queryByText(GATE_TEXT)).not.toBeInTheDocument(); + // M2M uses the backend service token, not a per-user DB credential. + expect(vi.mocked(getMCPOAuthUserCredentialStatus)).not.toHaveBeenCalled(); }); }); diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx index 5ec9a3d683a..8957c7b30b8 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx @@ -1,11 +1,14 @@ -import React, { useEffect, useState } from "react"; +import React, { useCallback, useEffect, useState } from "react"; import { useQuery, useMutation } from "@tanstack/react-query"; import { ToolTestPanel } from "./ToolTestPanel"; import { MCPTool, MCPToolsViewerProps, MCPContent, CallMCPToolResponse, getMcpOAuthMode } from "./types"; -import { listMCPTools, callMCPTool } from "../networking"; +import { listMCPTools, callMCPTool, getMCPOAuthUserCredentialStatus } from "../networking"; import { isTokenValid, getToken, removeToken } from "@/utils/mcpTokenStore"; import { sanitizeMcpAliasForHeader } from "@/utils/mcpHeaderUtils"; import { useToolsOAuthFlow } from "@/hooks/useToolsOAuthFlow"; +import { useUserMcpOAuthFlow } from "@/hooks/useUserMcpOAuthFlow"; +import { TOOLS_OAUTH_UI_STATE_KEY } from "@/hooks/mcpOAuthUtils"; +import { setSecureItem } from "@/utils/secureStorage"; import { Card, Title, Text } from "@tremor/react"; import { RobotOutlined, ToolOutlined, SearchOutlined, KeyOutlined, LockOutlined } from "@ant-design/icons"; @@ -31,11 +34,14 @@ const MCPToolsViewer = ({ const [passthroughHeaders, setPassthroughHeaders] = useState>({}); const [showHeaderInput, setShowHeaderInput] = useState(false); - // Only PKCE passthrough uses a browser-held session token (sessionStorage, - // cleared on tab/browser close) and a user-facing auth gate. OBO uses the - // backend-stored per-user token and M2M uses the backend's own service token, - // so neither needs a gate — they list tools with just the LiteLLM key. - const isPassthrough = getMcpOAuthMode({ auth_type, oauth2_flow, delegate_auth_to_upstream }) === "passthrough"; + // PKCE passthrough holds a browser-side session token (sessionStorage) and + // gates tool listing behind it. OBO uses a backend-stored per-user token that + // the user must establish once via an interactive login; we gate on whether + // that DB credential exists. M2M uses the backend's own service token and + // needs no gate. + const oauthMode = getMcpOAuthMode({ auth_type, oauth2_flow, delegate_auth_to_upstream }); + const isPassthrough = oauthMode === "passthrough"; + const isObo = oauthMode === "obo"; const [oauthToken, setOauthToken] = useState(() => isPassthrough && isTokenValid(serverId, userID) ? getToken(serverId, userID)?.access_token ?? null : null, ); @@ -61,6 +67,31 @@ const MCPToolsViewer = ({ onSuccess: setOauthToken, }); + // OBO servers list tools using a per-user token the backend stores in the DB; + // check whether the current user has a valid one so we can prompt them to + // authorize when they don't (otherwise the backend silently returns no tools). + const { + data: oboCredStatus, + isLoading: isLoadingOboCred, + isError: isOboCredError, + refetch: refetchOboCred, + } = useQuery({ + queryKey: ["mcpOauthUserCredStatus", serverId, userID], + queryFn: () => getMCPOAuthUserCredentialStatus(accessToken ?? "", serverId), + enabled: !!accessToken && isObo, + staleTime: 30000, + }); + + // A stored credential is sufficient: the backend proactively refreshes an + // expired or near-expiry token from the stored refresh_token on the next list + // call, so the user only needs to authorize when no credential row exists. If + // the status check itself fails we can't confirm a credential, so surface the + // Authorize gate rather than a silent empty tool list; re-authorizing only + // overwrites the user's own row, so it is safe when a credential did exist. + const hasOboCred = !!oboCredStatus?.has_credential; + const oboNeedsAuth = isObo && !isLoadingOboCred && (isOboCredError || (!!oboCredStatus && !hasOboCred)); + const oboStatusLoading = isObo && isLoadingOboCred; + // Check if this server has extra headers configured const hasExtraHeaders = extraHeaders && extraHeaders.length > 0; @@ -135,8 +166,9 @@ const MCPToolsViewer = ({ } return result; }, - // For OAuth servers, block the query until a session token is available - enabled: !!accessToken && (!isPassthrough || oauthToken !== null), + // Passthrough blocks until a browser session token exists; OBO blocks until + // the user has a valid DB credential (else the backend returns no tools). + enabled: !!accessToken && (isPassthrough ? oauthToken !== null : isObo ? hasOboCred : true), staleTime: 30000, // Consider data fresh for 30 seconds retry: (failureCount, error: any) => { // Don't retry on 401 — token is invalid, user must re-authenticate @@ -145,6 +177,33 @@ const MCPToolsViewer = ({ }, }); + // OBO authorize: same redirect+exchange flow as the admin "Authorize & Fetch" + // and the chat "Connect" button, but persists the token to the per-user DB. + const onOboAuthSuccess = useCallback(() => { + refetchOboCred(); + refetchTools(); + }, [refetchOboCred, refetchTools]); + + const { + startOAuthFlow: startDbOAuthFlow, + status: dbOAuthStatus, + error: dbOAuthError, + } = useUserMcpOAuthFlow({ + accessToken: accessToken ?? "", + serverId, + serverAlias, + onSuccess: onOboAuthSuccess, + }); + + // Stash which server started the redirect so the MCP Servers page can reopen + // this Tools tab on return and let the flow resume to persist the credential. + const startOboAuthorize = useCallback(() => { + try { + setSecureItem(TOOLS_OAUTH_UI_STATE_KEY, JSON.stringify({ serverId })); + } catch (_) {} + startDbOAuthFlow(); + }, [serverId, startDbOAuthFlow]); + // If the tools query fails with 401, the cached OAuth token is invalid — // clear it so the auth gate is shown again and the user can re-authenticate. useEffect(() => { @@ -187,6 +246,13 @@ const MCPToolsViewer = ({ const toolsData = mcpToolsResponse?.tools || []; + // An auth gate replaces the tool list when the user must authenticate first: + // passthrough needs a browser token, OBO needs a stored DB credential. + const authGateActive = (isPassthrough && !oauthToken) || oboNeedsAuth; + // Treat OBO credential-status loading as "tools loading" so the empty state + // doesn't flash before we know whether the user needs to authorize. + const toolsAreaLoading = isLoadingTools || oboStatusLoading; + // Filter tools based on search term const filteredTools = toolsData.filter((tool: MCPTool) => { const searchLower = toolSearchTerm.toLowerCase(); @@ -287,7 +353,7 @@ const MCPToolsViewer = ({ )} - {/* OAuth Auth Gate — shown when token is absent for OAuth servers */} + {/* Passthrough auth gate — browser session token absent */} {isPassthrough && !oauthToken && (
@@ -306,8 +372,31 @@ const MCPToolsViewer = ({
)} + {/* OBO auth gate — only when no credential row exists for this user. + An existing-but-expired token is refreshed server-side on the + list call, so the gate never appears for a stored credential. */} + {oboNeedsAuth && ( +
+ +

Authentication required

+

+ Authenticate with the upstream provider to view available tools +

+ + Authorize + + {dbOAuthError &&

{dbOAuthError}

} +
+ )} + {/* Search Bar — only shown when tools are loaded */} - {!isPassthrough || oauthToken ? ( + {!authGateActive ? ( <> {toolsData.length > 0 && (
@@ -324,7 +413,7 @@ const MCPToolsViewer = ({ )} {/* Loading State */} - {isLoadingTools && ( + {toolsAreaLoading && (
@@ -335,7 +424,7 @@ const MCPToolsViewer = ({ )} {/* Error State */} - {(mcpToolsResponse?.error || mcpToolsError) && !isLoadingTools && !toolsData.length && ( + {(mcpToolsResponse?.error || mcpToolsError) && !toolsAreaLoading && !toolsData.length && (

Error: {mcpToolsResponse?.message || (mcpToolsError as Error)?.message} @@ -344,7 +433,7 @@ const MCPToolsViewer = ({ )} {/* No Tools State */} - {!isLoadingTools && + {!toolsAreaLoading && !mcpToolsResponse?.error && !mcpToolsError && (!toolsData || toolsData.length === 0) && ( @@ -370,7 +459,7 @@ const MCPToolsViewer = ({ )} {/* Tools List */} - {!isLoadingTools && !mcpToolsResponse?.error && toolsData.length > 0 && ( + {!toolsAreaLoading && !mcpToolsResponse?.error && toolsData.length > 0 && ( <> {filteredTools.length === 0 ? (

diff --git a/ui/litellm-dashboard/src/hooks/mcpOAuthUtils.ts b/ui/litellm-dashboard/src/hooks/mcpOAuthUtils.ts index 3aff8af6eef..83b2cb21bed 100644 --- a/ui/litellm-dashboard/src/hooks/mcpOAuthUtils.ts +++ b/ui/litellm-dashboard/src/hooks/mcpOAuthUtils.ts @@ -7,6 +7,15 @@ import { getProxyBaseUrl, serverRootPath } from "@/components/networking"; +/** + * sessionStorage key used to restore the MCP server detail view on the Tools + * tab after a full-page OAuth redirect. The OBO authorize flow redirects to the + * IdP and back to the MCP Servers page; without this the user lands on the + * server list and useUserMcpOAuthFlow never re-mounts to persist the credential. + * Mirrors the admin edit flow's EDIT_OAUTH_UI_STATE_KEY. + */ +export const TOOLS_OAUTH_UI_STATE_KEY = "litellm-mcp-oauth-tools-state"; + /** * Build the OAuth callback URL for the current UI deployment. * From 32c88ca74f29143aadc4cfaa708f89ce1086386c Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 9 Jun 2026 02:19:52 +0530 Subject: [PATCH 08/22] Litellm oss staging 080626 (#29932) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat(bedrock_mantle): add SigV4/IAM auth to Responses API route (fixes #29665) (#29788) * feat(responses): add default no-op sign_request to BaseResponsesAPIConfig * feat(responses): call sign_request after body is final, send signed bytes when signed * feat(bedrock_mantle): add SigV4 sign_request via composed BaseAWSLLM (bearer path) * test(bedrock_mantle): cover SigV4 access-key, AssumeRole, body bytes, region/auth consistency * feat(bedrock_mantle): defer auth to sign_request; validate_environment no longer requires bearer * docs(bedrock_mantle): document SigV4 + Bearer auth on Responses route * test(responses): cover fake-stream signing order and mantle bearer arg/env precedence * fix(bedrock_mantle): wrap all botocore credential errors with both-paths guidance * fix(bedrock_mantle): catch specific credential errors, not all BotoCoreError, so STS transport failures are not masked * fix(bedrock_mantle): sign the compact Responses route too, not just create * fix(github-copilot): route per-model on /v1/responses based on model info (#29747) * feat(focus): add GCS destination for FOCUS export (#29751) * test: add failing tests for FocusGCSDestination * feat: add FocusGCSDestination reusing GCSBucketBase auth * feat: register FocusGCSDestination in factory; export from __init__ * fix(focus): preserve GCS_PATH_SERVICE_ACCOUNT when service_account_json not in config * style: apply Black formatting to gcs_destination and tests * style: apply Black formatting to factory.py * fix(bedrock): omit empty additionalModelRequestFields and system from Converse API payload (#29565) Amazon Nova Pro (and other strict Bedrock models) return 400 Malformed input request when additionalModelRequestFields: {} or system: [] are present in the payload. Both fields are optional in CommonRequestObject (total=False) and must be omitted rather than sent as empty structures. Co-authored-by: shin-berri Co-authored-by: yuneng-jiang Co-authored-by: Claude Sonnet 4.6 * fix(proxy): recognize *.cognitiveservices.azure.com as OpenAI-compatible in pass-through cost tracking (#29730) * fix(proxy): recognize *.cognitiveservices.azure.com as OpenAI-compatible Azure OpenAI resources created via the newer "Azure AI Foundry" / Cognitive Services pathway live on `*.cognitiveservices.azure.com` subdomains, not the older `openai.azure.com`. Both are valid Azure OpenAI surfaces in production today. The OpenAI pass-through cost-tracking handler hard-codes only the older hostname in five places (four `is_openai_*_route` methods on OpenAIPassthroughLoggingHandler, plus is_openai_route on PassThroughEndpointLogging). As a result, calls from newer Azure deployments are silently classified as "not an OpenAI route", the dispatch into the cost-tracking handler is skipped, and tokens/cost never get extracted into LiteLLM_SpendLogs — the row gets written with prompt_tokens=0, completion_tokens=0, spend=0, model='unknown'. Reproduced 2026-06-04 against a real Azure OpenAI deployment on `*.cognitiveservices.azure.com` proxied through LiteLLM v1.88.0. Fix: factor the hostname check into a single helper `_is_openai_compatible_host` listing all three recognized surfaces (api.openai.com, openai.azure.com, cognitiveservices.azure.com), and have all five call sites delegate to it. Purely additive — never weakens recognition for the originally-supported hostnames. Adds a test `test_is_openai_route_recognizes_cognitiveservices_azure_com` that exercises all four `is_openai_*_route` static methods against `*.cognitiveservices.azure.com` URLs (positive cases per route + a small cross-route negative to confirm route-specific path matching still works on the new hostname). Out of scope for this PR (separate followup): - `openai_passthrough_handler` calls chat/completions `transform_response` on Responses API payloads (`output:` not `choices:`), which throws inside the dispatch and drops the SpendLogs row entirely. Recognized + tracked separately. * ci: trigger fresh run Empty commit to re-run checks. The previous auth-and-jwt failure was a transient HuggingFace Hub 429 rate-limit hitting tokenizer downloads in tests/proxy_unit_tests/test_custom_tokenizer_bug.py — unrelated to this PR's scope (hostname recognition in pass-through cost tracking). No code change. --------- Co-authored-by: shin-berri Co-authored-by: yuneng-jiang * fix(responses): preserve forced-function tool_choice name in Responses to Chat transform (#29812) The Responses API forces a specific function with a top-level name ({"type": "function", "name": "X"}), but _transform_tool_choice only handled the nested Chat Completions shape and fell through to returning "required" for the flat form, silently dropping the function name and degrading a forced function call to force-any-tool. Map the flat Responses shape to the nested Chat shape, keeping the "required" fallback when no name is present. * Preserve x-anthropic-billing-header system blocks for first-party Anthropic (#29584) * Preserve x-anthropic-billing-header system blocks for first-party Anthropic PR #20951 strips system blocks beginning with "x-anthropic-billing-header:" for every Anthropic target. That block is how the first-party Anthropic API recognizes Claude Code subscription (OAuth) traffic, so dropping it makes requests that carry only that block, such as the auto-mode tool-safety classifier, fail with a misleading 429 rate_limit_error; normal turns still work because they also carry the "You are Claude Code" identity block. Gate the strip behind should_strip_billing_metadata(), defaulting to False on the first-party AnthropicConfig and AnthropicMessagesConfig so the block is kept, and overridden to True on the providers that reach these transforms and reject the block (Bedrock platform, Vertex, Azure for the chat path; Minimax, Azure, DeepSeek for the messages path). Behavior for those providers is unchanged. * Strip billing header on Bedrock invoke and Vertex messages pass-through Two more subclasses reach the gated strip but inherited keep-by-default. AmazonAnthropicClaudeConfig (Bedrock invoke) calls AnthropicConfig.transform_request, which calls translate_system_message, and VertexAIPartnerModelsAnthropicMessagesConfig (Vertex messages pass-through) calls super().transform_anthropic_messages_request. Override should_strip_billing_metadata() to True on both. Add a parametrized test asserting the flag for every first-party base (False) and provider subclass (True), covering all overrides, plus a translate_system_message regression test for the Bedrock invoke path. * fix(cache): log hashed cache keys (#29890) * fix(ui): save routing groups as list (#29889) * Revert "fix(ui): save routing groups as list (#29889)" (#29928) This reverts commit 9b1f78ffa7a309cabe5e9a7ab5f94d1224d192c9. * feat(parasail): add Parasail as a JSON-configured OpenAI-compatible provider (#29842) * feat(parasail): add Parasail as a JSON-configured OpenAI-compatible provider Registers parasail in the openai_like JSON provider loader with both /v1/chat/completions and /v1/responses support. Parasail's Responses API rejects store:true and any request that omits store, so the loader gains a force_store_false special_handling flag; the parasail entry sets it and the generated Responses config overrides store=false on every call. This keeps callers from hitting "State storage not supported" and matches what Parasail's docs require. Adds the PARASAIL enum value, listing under openai_compatible_providers, provider documentation at docs/my-website/docs/providers/parasail.md, and a focused unit test file under tests/test_litellm/llms/parasail/ that covers JSON registration, chat URL construction, Responses URL construction with PARASAIL_API_BASE override, and the force_store_false regression in both the caller-sent-store=true and caller-omitted cases. * fix(parasail): register in provider_endpoints_support, drop in-repo docs Greptile review feedback. The provider doc belongs in the litellm-docs repo, not this one's docs/my-website tree; removing it here. Adds the parasail entry to provider_endpoints_support.json so the check_provider_folders_documented.py CI check passes (chat_completions and responses true; others false). * fix: normalize Anthropic passthrough server tool usage (#29827) * test(anthropic): cover server_tool_use dict cost tracking * fix: normalize Anthropic server tool usage (cherry picked from commit 982f726bed7d3ec05e463c5dd3d090bebae91d19) * fix: keep server tool usage subscriptable (cherry picked from commit 70280b9b272455b2f974d08bc697f67f929755bf) --------- Co-authored-by: Genmin * fix(proxy): fix typo generic_role_mappoings -> generic_role_mappings in ui_sso.py (#29753) Co-authored-by: shin-berri Co-authored-by: yuneng-jiang * feat(proxy): add disable_budget_reservation general setting (#27639) (#29493) * feat(proxy): add disable_budget_reservation general setting (#27639) * feat(proxy): register disable_budget_reservation in ConfigGeneralSettings (#27639) * docs(proxy): document disable_budget_reservation concurrency tradeoff (#27639) * ci: re-trigger flaky docker build (prisma generate ECONNRESET) * fix(proxy): warn and document budget enforcement tradeoff when disable_budget_reservation is set (#27639) * feat(gemini_tts): adding support to Gemini TTS languageCode parameters (#29623) * Adding support to Gemini TTS Language Code parameters * Mapping Gemini TTS languageCode param in Docstring * Use snake_case for language_code input keyMapping Gemini TTS languageCode param in Docstring * Restoring files modified under enterprise/litellm_enterprise due to lint/formatting checks --------- Co-authored-by: João Garrido * feat(guardrails): capture user and model metadata in CrowdStrike AIDR (#29517) * fix(proxy): require OpenAI path segment for shared Azure Cognitive Services domains Address Greptile review: the `*.cognitiveservices.azure.com` / `*.openai.azure.com` domains are shared by every Azure Cognitive Service (Speech, Vision, Language, ...), so a hostname-only substring match misclassified non-OpenAI Azure traffic as OpenAI routes. - Replace the substring host test with suffix matching (rejects look-alike domains like cognitiveservices.azure.com.attacker.example). - Add `_is_openai_compatible_url` that requires an OpenAI-style path marker (`/openai/` or `/v1/`) on the shared Azure domains, and use it in PassThroughEndpointLogging.is_openai_route (previously hostname-only). - Add negative tests for Azure Speech/Vision paths and look-alike domains. Co-Authored-By: Claude Opus 4.8 * fix: support Responses input in Redis semantic cache (#29581) * fix: support responses input in redis semantic cache * test: cover redis semantic prompt extraction * test: handle blank redis semantic text fallbacks * chore: remove async cache dead statement * test: cover redis semantic cache miss paths * fix: filter sensitive cache lookup kwargs * chore: rerun ci after huggingface rate limit * chore(ui): regenerate dashboard API types (npm run gen:api) Sync src/lib/http/schema.d.ts with the proxy OpenAPI spec: adds the disable_budget_reservation general-settings field and picks up the RateLimitError docstring reindent. Fixes the gen:api CI drift check. Co-Authored-By: Claude Opus 4.8 * test(bedrock): assert empty additionalModelRequestFields is omitted The Converse transformer now drops an empty additionalModelRequestFields block instead of sending it as `{}`. Update test_bedrock_top_k_param so models without top_k support (llama3) assert the key is absent rather than equal to an empty dict. Co-Authored-By: Claude Opus 4.8 --------- Co-authored-by: Kent <72616338+kingdoooo@users.noreply.github.com> Co-authored-by: codgician <15964984+codgician@users.noreply.github.com> Co-authored-by: Praveen Ghuge <95286176+pghuge-cloudwiz@users.noreply.github.com> Co-authored-by: Roi Co-authored-by: shin-berri Co-authored-by: yuneng-jiang Co-authored-by: Claude Sonnet 4.6 Co-authored-by: Liam Scott Co-authored-by: abhay23-AI Co-authored-by: Ceder Dens Co-authored-by: 冯基魁 <56265583+fengjikui@users.noreply.github.com> Co-authored-by: Kai Huang Co-authored-by: rinto <54238243+ririnto@users.noreply.github.com> Co-authored-by: Genmin Co-authored-by: Arnav Bhilwariya Co-authored-by: Armaan Sandhu <74664101+Ar-maan05@users.noreply.github.com> Co-authored-by: João Garrido <48538534+johngarrido@users.noreply.github.com> Co-authored-by: João Garrido Co-authored-by: Kenan Yildirim Co-authored-by: Dávid Balatoni --- litellm/caching/caching.py | 47 +- litellm/caching/redis_semantic_cache.py | 105 +++- litellm/constants.py | 1 + .../focus/destinations/__init__.py | 2 + .../focus/destinations/factory.py | 15 + .../focus/destinations/gcs_destination.py | 74 +++ litellm/llms/anthropic/chat/transformation.py | 22 +- .../messages/transformation.py | 13 +- .../anthropic/messages_transformation.py | 3 + .../llms/azure_ai/anthropic/transformation.py | 3 + .../llms/base_llm/responses/transformation.py | 20 + .../bedrock/chat/converse_transformation.py | 6 +- .../anthropic_claude3_transformation.py | 3 + .../bedrock/claude_platform/transformation.py | 3 + .../responses/transformation.py | 120 ++++- litellm/llms/custom_httpx/llm_http_handler.py | 105 +++- .../llms/deepseek/messages/transformation.py | 3 + .../responses/transformation.py | 45 +- .../llms/minimax/messages/transformation.py | 3 + litellm/llms/openai_like/dynamic_config.py | 19 + litellm/llms/openai_like/providers.json | 9 + .../vertex_and_google_ai_studio_gemini.py | 7 +- .../transformation.py | 3 + .../anthropic/transformation.py | 3 + litellm/proxy/_types.py | 18 + litellm/proxy/auth/user_api_key_auth.py | 12 + .../crowdstrike_aidr/crowdstrike_aidr.py | 18 +- litellm/proxy/management_endpoints/ui_sso.py | 4 +- .../openai_passthrough_logging_handler.py | 98 +++- .../pass_through_endpoints/success_handler.py | 15 +- .../transformation.py | 4 +- litellm/types/llms/vertex_ai.py | 1 + litellm/types/utils.py | 11 +- litellm/utils.py | 8 +- provider_endpoints_support.json | 17 + .../test_bedrock_completion.py | 6 +- tests/test_litellm/caching/test_caching.py | 48 ++ .../caching/test_redis_semantic_cache.py | 465 ++++++++++++++++++ .../focus/test_focus_gcs_destination.py | 180 +++++++ .../test_tool_call_cost_tracking.py | 21 +- .../test_anthropic_chat_transformation.py | 169 +++++++ .../chat/test_converse_transformation.py | 7 +- .../test_bedrock_files_transformation.py | 64 +++ ...bedrock_mantle_responses_transformation.py | 398 ++++++++++++++- .../custom_httpx/test_llm_http_handler.py | 238 +++++++++ .../llms/gemini/test_gemini_tts.py | 92 ++++ ...github_copilot_responses_transformation.py | 216 +++++++- .../llms/parasail/test_parasail.py | 172 +++++++ .../proxy/auth/test_user_api_key_auth.py | 60 +++ .../guardrail_hooks/test_crowdstrike_aidr.py | 121 ++++- ...t_anthropic_passthrough_logging_handler.py | 38 ++ ...test_openai_passthrough_logging_handler.py | 88 ++++ .../test_litellm_completion_responses.py | 22 + tests/test_litellm/types/test_types_utils.py | 46 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 5 + 55 files changed, 3154 insertions(+), 142 deletions(-) create mode 100644 litellm/integrations/focus/destinations/gcs_destination.py create mode 100644 tests/test_litellm/caching/test_caching.py create mode 100644 tests/test_litellm/integrations/focus/test_focus_gcs_destination.py create mode 100644 tests/test_litellm/llms/parasail/test_parasail.py diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index 11733ce4cee..c1afde16250 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -309,9 +309,13 @@ class Cache: param_value = kwargs[param] cache_key += f"{str(param)}: {str(param_value)}" - verbose_logger.debug("\nCreated cache key: %s", cache_key) hashed_cache_key = Cache._get_hashed_cache_key(cache_key) hashed_cache_key = self._add_namespace_to_cache_key(hashed_cache_key, **kwargs) + verbose_logger.debug( + "\nCreated cache key: %s (source material length: %d)", + hashed_cache_key, + len(cache_key), + ) # Remove preset_cache_key from kwargs to avoid "got multiple values" TypeError # when kwargs already contains preset_cache_key from upstream callers kwargs_for_preset = {k: v for k, v in kwargs.items() if k != "preset_cache_key"} @@ -497,6 +501,34 @@ class Cache: return cached_response return cached_result + @staticmethod + def _get_safe_cache_lookup_kwargs(kwargs: Dict[str, Any]) -> Dict[str, Any]: + cache_lookup_kwargs: Dict[str, Any] = {} + for prompt_kwarg in ("messages", "input"): + if prompt_kwarg in kwargs: + cache_lookup_kwargs[prompt_kwarg] = kwargs[prompt_kwarg] + + if isinstance(kwargs.get("metadata"), dict): + cache_lookup_kwargs["metadata"] = {} + + return cache_lookup_kwargs + + @staticmethod + def _update_metadata_from_cache_lookup_kwargs( + original_kwargs: Dict[str, Any], cache_lookup_kwargs: Dict[str, Any] + ) -> None: + original_metadata = original_kwargs.get("metadata") + cache_lookup_metadata = cache_lookup_kwargs.get("metadata") + if not isinstance(original_metadata, dict) or not isinstance( + cache_lookup_metadata, dict + ): + return + + if "semantic-similarity" in cache_lookup_metadata: + original_metadata["semantic-similarity"] = cache_lookup_metadata[ + "semantic-similarity" + ] + def get_cache(self, dynamic_cache_object: Optional[BaseCache] = None, **kwargs): """ Retrieves the cached result for the given arguments. @@ -511,7 +543,6 @@ class Cache: try: # never block execution if self.should_use_cache(**kwargs) is not True: return - messages = kwargs.get("messages", []) if "cache_key" in kwargs: cache_key = kwargs["cache_key"] else: @@ -523,12 +554,19 @@ class Cache: or cache_control_args.get("s-max-age") or float("inf") ) + cache_lookup_kwargs = self._get_safe_cache_lookup_kwargs(kwargs) if dynamic_cache_object is not None: cached_result = dynamic_cache_object.get_cache( - cache_key, messages=messages + cache_key, **cache_lookup_kwargs ) else: - cached_result = self.cache.get_cache(cache_key, messages=messages) + cached_result = self.cache.get_cache( + cache_key, **cache_lookup_kwargs + ) + self._update_metadata_from_cache_lookup_kwargs( + original_kwargs=kwargs, + cache_lookup_kwargs=cache_lookup_kwargs, + ) return self._get_cache_logic( cached_result=cached_result, max_age=max_age ) @@ -549,7 +587,6 @@ class Cache: if self.should_use_cache(**kwargs) is not True: return - kwargs.get("messages", []) if "cache_key" in kwargs: cache_key = kwargs["cache_key"] else: diff --git a/litellm/caching/redis_semantic_cache.py b/litellm/caching/redis_semantic_cache.py index da9e7b1e587..cce4b75795f 100644 --- a/litellm/caching/redis_semantic_cache.py +++ b/litellm/caching/redis_semantic_cache.py @@ -213,6 +213,78 @@ class RedisSemanticCache(BaseCache): ttl = int(ttl) return ttl + @classmethod + def _get_prompt_from_kwargs(cls, **kwargs) -> Optional[str]: + """ + Extract a semantic-cache prompt from chat or Responses API request kwargs. + """ + messages = kwargs.get("messages") + if messages: + return get_str_from_messages(messages) + + if "input" not in kwargs: + return None + + prompt_parts: List[str] = [] + cls._collect_responses_input_text(kwargs.get("input"), prompt_parts) + prompt = "\n".join(prompt_parts).strip() + return prompt or None + + @classmethod + def _collect_responses_input_text(cls, value: Any, prompt_parts: List[str]) -> None: + value = cls._coerce_response_input_value(value) + if value is None: + return + + if isinstance(value, str): + stripped_value = value.strip() + if stripped_value: + prompt_parts.append(stripped_value) + return + + if isinstance(value, (list, tuple)): + for item in value: + cls._collect_responses_input_text(item, prompt_parts) + return + + if isinstance(value, dict): + content = value.get("content") + if content is not None: + cls._collect_responses_input_text(content, prompt_parts) + return + + for text_key in ("text", "output", "input_text", "output_text"): + text_value = value.get(text_key) + if isinstance(text_value, str): + stripped_text = text_value.strip() + if stripped_text: + prompt_parts.append(stripped_text) + return + return + + content = getattr(value, "content", None) + if content is not None: + cls._collect_responses_input_text(content, prompt_parts) + return + + for text_key in ("text", "output", "input_text", "output_text"): + text_value = getattr(value, text_key, None) + if isinstance(text_value, str): + stripped_text = text_value.strip() + if stripped_text: + prompt_parts.append(stripped_text) + return + + @staticmethod + def _coerce_response_input_value(value: Any) -> Any: + model_dump = getattr(value, "model_dump", None) + if callable(model_dump): + return model_dump() + dict_method = getattr(value, "dict", None) + if callable(dict_method): + return dict_method() + return value + def _get_embedding(self, prompt: str) -> List[float]: """ Generate an embedding vector for the given prompt using the configured embedding model. @@ -278,13 +350,11 @@ class RedisSemanticCache(BaseCache): value_str: Optional[str] = None try: - # Extract the prompt from messages - messages = kwargs.get("messages", []) - if not messages: - print_verbose("No messages provided for semantic caching") + prompt = self._get_prompt_from_kwargs(**kwargs) + if prompt is None: + print_verbose("No prompt provided for semantic caching") return - prompt = get_str_from_messages(messages) value_str = str(value) store_kwargs: Dict[str, Any] = { @@ -315,14 +385,12 @@ class RedisSemanticCache(BaseCache): print_verbose(f"Redis semantic-cache get_cache, kwargs: {kwargs}") try: - # Extract the prompt from messages - messages = kwargs.get("messages", []) - if not messages: - print_verbose("No messages provided for semantic cache lookup") + prompt = self._get_prompt_from_kwargs(**kwargs) + if prompt is None: + print_verbose("No prompt provided for semantic cache lookup") kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 return None - prompt = get_str_from_messages(messages) # Check the cache for semantically similar prompts in this exact # LiteLLM cache-key scope. check_kwargs: Dict[str, Any] = { @@ -428,13 +496,11 @@ class RedisSemanticCache(BaseCache): print_verbose(f"Async Redis semantic-cache set_cache, kwargs: {kwargs}") try: - # Extract the prompt from messages - messages = kwargs.get("messages", []) - if not messages: - print_verbose("No messages provided for semantic caching") + prompt = self._get_prompt_from_kwargs(**kwargs) + if prompt is None: + print_verbose("No prompt provided for semantic caching") return - prompt = get_str_from_messages(messages) value_str = str(value) # Generate embedding for the value (response) to cache @@ -471,15 +537,12 @@ class RedisSemanticCache(BaseCache): print_verbose(f"Async Redis semantic-cache get_cache, kwargs: {kwargs}") try: - # Extract the prompt from messages - messages = kwargs.get("messages", []) - if not messages: - print_verbose("No messages provided for semantic cache lookup") + prompt = self._get_prompt_from_kwargs(**kwargs) + if prompt is None: + print_verbose("No prompt provided for semantic cache lookup") kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0 return None - prompt = get_str_from_messages(messages) - # Generate embedding for the prompt prompt_embedding = await self._get_async_embedding(prompt, **kwargs) diff --git a/litellm/constants.py b/litellm/constants.py index 36e578bd323..f10cec034f0 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -831,6 +831,7 @@ openai_compatible_providers: List = [ "nano-gpt", # Nano-GPT - JSON-configured provider "poe", # Poe - JSON-configured provider "chutes", # Chutes - JSON-configured provider + "parasail", # Parasail - JSON-configured provider "featherless_ai", "nscale", "nebius", diff --git a/litellm/integrations/focus/destinations/__init__.py b/litellm/integrations/focus/destinations/__init__.py index 775d3a259d2..e0cd90c1d61 100644 --- a/litellm/integrations/focus/destinations/__init__.py +++ b/litellm/integrations/focus/destinations/__init__.py @@ -2,12 +2,14 @@ from .base import FocusDestination, FocusTimeWindow from .factory import FocusDestinationFactory +from .gcs_destination import FocusGCSDestination from .s3_destination import FocusS3Destination from .vantage_destination import FocusVantageDestination __all__ = [ "FocusDestination", "FocusDestinationFactory", + "FocusGCSDestination", "FocusTimeWindow", "FocusS3Destination", "FocusVantageDestination", diff --git a/litellm/integrations/focus/destinations/factory.py b/litellm/integrations/focus/destinations/factory.py index 706e10624ce..7ce21d4040a 100644 --- a/litellm/integrations/focus/destinations/factory.py +++ b/litellm/integrations/focus/destinations/factory.py @@ -6,6 +6,7 @@ import os from typing import Any, Dict, Optional from .base import FocusDestination +from .gcs_destination import FocusGCSDestination from .s3_destination import FocusS3Destination from .vantage_destination import FocusVantageDestination @@ -29,6 +30,8 @@ class FocusDestinationFactory: return FocusS3Destination(prefix=prefix, config=normalized_config) if provider_lower == "vantage": return FocusVantageDestination(prefix=prefix, config=normalized_config) + if provider_lower == "gcs": + return FocusGCSDestination(prefix=prefix, config=normalized_config) raise NotImplementedError( f"Provider '{provider}' not supported for Focus export" ) @@ -72,6 +75,18 @@ class FocusDestinationFactory: "VANTAGE_INTEGRATION_TOKEN must be provided for Vantage exports" ) return {k: v for k, v in resolved.items() if v is not None} + if provider == "gcs": + resolved = { + "bucket_name": overrides.get("bucket_name") + or os.getenv("FOCUS_GCS_BUCKET_NAME"), + "service_account_json": overrides.get("service_account_json") + or os.getenv("FOCUS_GCS_PATH_SERVICE_ACCOUNT"), + } + if not resolved.get("bucket_name"): + raise ValueError( + "FOCUS_GCS_BUCKET_NAME must be provided for GCS exports" + ) + return {k: v for k, v in resolved.items() if v is not None} raise NotImplementedError( f"Provider '{provider}' not supported for Focus export configuration" ) diff --git a/litellm/integrations/focus/destinations/gcs_destination.py b/litellm/integrations/focus/destinations/gcs_destination.py new file mode 100644 index 00000000000..b04c16c9d32 --- /dev/null +++ b/litellm/integrations/focus/destinations/gcs_destination.py @@ -0,0 +1,74 @@ +"""GCS destination for Focus export — reuses GCSBucketBase auth and httpx client.""" + +from __future__ import annotations + +from datetime import timezone +from typing import Any, Optional + +from litellm._logging import verbose_logger +from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase +from litellm.litellm_core_utils.cloud_storage_security import ( + encode_gcs_object_name_for_url, +) + +from .base import FocusDestination, FocusTimeWindow + + +class FocusGCSDestination(GCSBucketBase, FocusDestination): + """Upload serialized Focus exports to GCS using the GCS JSON API.""" + + def __init__( + self, + *, + prefix: str, + config: Optional[dict[str, Any]] = None, + ) -> None: + config = config or {} + bucket_name = config.get("bucket_name") + if not bucket_name: + raise ValueError("bucket_name must be provided for GCS destination") + super().__init__(bucket_name=bucket_name) + service_account_json = config.get("service_account_json") + if service_account_json is not None: + self.path_service_account_json = service_account_json + self.prefix = prefix.rstrip("/") + + async def deliver( + self, + *, + content: bytes, + time_window: FocusTimeWindow, + filename: str, + ) -> None: + object_name = self._build_object_key(time_window=time_window, filename=filename) + headers = await self.construct_request_headers( + service_account_json=self.path_service_account_json + ) + headers["Content-Type"] = "application/octet-stream" + encoded_name = encode_gcs_object_name_for_url(object_name) + url = ( + f"https://storage.googleapis.com/upload/storage/v1/b/" + f"{self.BUCKET_NAME}/o?uploadType=media&name={encoded_name}" + ) + response = await self.async_httpx_client.post( + url=url, headers=headers, data=content + ) + if response.status_code != 200: + raise RuntimeError( + f"GCS upload failed: status={response.status_code} body={response.text}" + ) + verbose_logger.debug( + "Focus GCS: uploaded %d bytes to gs://%s/%s", + len(content), + self.BUCKET_NAME, + object_name, + ) + + def _build_object_key(self, *, time_window: FocusTimeWindow, filename: str) -> str: + start_utc = time_window.start_time.astimezone(timezone.utc) + date_component = f"date={start_utc.strftime('%Y-%m-%d')}" + parts = [self.prefix, date_component] + if time_window.frequency == "hourly": + parts.append(f"hour={start_utc.strftime('%H')}") + key_prefix = "/".join(filter(None, parts)) + return f"{key_prefix}/{filename}" if key_prefix else filename diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 7949e150c23..3f30d5d6807 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -1607,6 +1607,15 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) return _tool + def should_strip_billing_metadata(self) -> bool: + """ + Whether to drop x-anthropic-billing-header system blocks before sending upstream. + + The first-party Anthropic API uses these blocks for Claude Code attribution, so the + base config keeps them. Providers that reject them (e.g. Bedrock) override this to True. + """ + return False + def translate_system_message( self, messages: List[AllMessageValues] ) -> List[AnthropicSystemMessageContent]: @@ -1614,7 +1623,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): Translate system message to anthropic format. Removes system message from the original list and returns a new list of anthropic system message content. - Filters out system messages containing x-anthropic-billing-header metadata. + When should_strip_billing_metadata() is True, x-anthropic-billing-header system blocks are dropped. """ system_prompt_indices = [] anthropic_system_message_list: List[AnthropicSystemMessageContent] = [] @@ -1626,10 +1635,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): # Skip empty text blocks - Anthropic API raises errors for empty text if not system_message_block["content"]: continue - # Skip system messages containing x-anthropic-billing-header metadata - if system_message_block["content"].startswith( - "x-anthropic-billing-header:" - ): + if self.should_strip_billing_metadata() and system_message_block[ + "content" + ].startswith("x-anthropic-billing-header:"): continue anthropic_system_message_content = AnthropicSystemMessageContent( type="text", @@ -1648,9 +1656,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): text_value = _content.get("text") if _content.get("type") == "text" and not text_value: continue - # Skip system messages containing x-anthropic-billing-header metadata if ( - _content.get("type") == "text" + self.should_strip_billing_metadata() + and _content.get("type") == "text" and text_value and text_value.startswith("x-anthropic-billing-header:") ): diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py index 3a2c09f2183..07e8270b496 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py @@ -84,6 +84,15 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): if isinstance(content, list): _process_content_list(content) + def should_strip_billing_metadata(self) -> bool: + """ + Whether to drop x-anthropic-billing-header system blocks before sending upstream. + + The first-party Anthropic API uses these blocks for Claude Code attribution, so the + base config keeps them. Providers that reject them override this to True. + """ + return False + @staticmethod def _filter_billing_headers_from_system(system_param): """ @@ -286,14 +295,12 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): optional_params=anthropic_messages_optional_request_params, ) - # Filter out x-anthropic-billing-header from system messages system_param = anthropic_messages_optional_request_params.get("system") - if system_param is not None: + if self.should_strip_billing_metadata() and system_param is not None: filtered_system = self._filter_billing_headers_from_system(system_param) if filtered_system is not None and len(filtered_system) > 0: anthropic_messages_optional_request_params["system"] = filtered_system else: - # Remove system parameter if all content was filtered out anthropic_messages_optional_request_params.pop("system", None) # Transform context_management from OpenAI format to Anthropic format if needed diff --git a/litellm/llms/azure_ai/anthropic/messages_transformation.py b/litellm/llms/azure_ai/anthropic/messages_transformation.py index a81218ab76a..59b6ee2b424 100644 --- a/litellm/llms/azure_ai/anthropic/messages_transformation.py +++ b/litellm/llms/azure_ai/anthropic/messages_transformation.py @@ -21,6 +21,9 @@ class AzureAnthropicMessagesConfig(AnthropicMessagesConfig): and Azure endpoint format. """ + def should_strip_billing_metadata(self) -> bool: + return True + def validate_anthropic_messages_environment( self, headers: dict, diff --git a/litellm/llms/azure_ai/anthropic/transformation.py b/litellm/llms/azure_ai/anthropic/transformation.py index e176a4d860e..367ca75c196 100644 --- a/litellm/llms/azure_ai/anthropic/transformation.py +++ b/litellm/llms/azure_ai/anthropic/transformation.py @@ -40,6 +40,9 @@ class AzureAnthropicConfig(AnthropicConfig): def custom_llm_provider(self) -> Optional[str]: return "azure_ai" + def should_strip_billing_metadata(self) -> bool: + return True + def validate_environment( self, headers: dict, diff --git a/litellm/llms/base_llm/responses/transformation.py b/litellm/llms/base_llm/responses/transformation.py index 853eb282758..407d5ad8146 100644 --- a/litellm/llms/base_llm/responses/transformation.py +++ b/litellm/llms/base_llm/responses/transformation.py @@ -62,6 +62,26 @@ class BaseResponsesAPIConfig(ABC): """ return False + def sign_request( + self, + headers: dict, + optional_params: dict, + request_data: dict, + api_base: str, + api_key: Optional[str] = None, + model: Optional[str] = None, + stream: Optional[bool] = None, + fake_stream: Optional[bool] = None, + ) -> Tuple[dict, Optional[bytes]]: + """Sign the request after the body is finalized. + + Default is a no-op (returns headers unchanged, no signed body). Providers + whose endpoint requires request signing (e.g. Bedrock Mantle SigV4) + override this and return the signed body bytes so the handler sends those + exact bytes. + """ + return headers, None + @abstractmethod def get_supported_openai_params(self, model: str) -> list: pass diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 90dfa13e938..ea0326dffd1 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -1649,12 +1649,14 @@ class AmazonConverseConfig(BaseConfig): bedrock_tool_config["toolChoice"] = tool_choice_values data: CommonRequestObject = { - "additionalModelRequestFields": additional_request_params, - "system": system_content_blocks, "inferenceConfig": self._transform_inference_params( inference_params=inference_params ), } + if additional_request_params: + data["additionalModelRequestFields"] = additional_request_params + if system_content_blocks: + data["system"] = system_content_blocks # Handle all config blocks for config_name, config_class in self.get_config_blocks().items(): diff --git a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py index a13336b6c88..4887cbd23be 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py @@ -60,6 +60,9 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig): def custom_llm_provider(self) -> Optional[str]: return "bedrock" + def should_strip_billing_metadata(self) -> bool: + return True + def get_supported_openai_params(self, model: str) -> List[str]: return AnthropicConfig.get_supported_openai_params(self, model) diff --git a/litellm/llms/bedrock/claude_platform/transformation.py b/litellm/llms/bedrock/claude_platform/transformation.py index 0167c457c96..c20dc63444f 100644 --- a/litellm/llms/bedrock/claude_platform/transformation.py +++ b/litellm/llms/bedrock/claude_platform/transformation.py @@ -17,6 +17,9 @@ class BedrockClaudePlatformConfig(BedrockClaudePlatformMixin, AnthropicConfig): def custom_llm_provider(self) -> Optional[str]: return "bedrock" + def should_strip_billing_metadata(self) -> bool: + return True + def validate_environment( self, headers: dict, diff --git a/litellm/llms/bedrock_mantle/responses/transformation.py b/litellm/llms/bedrock_mantle/responses/transformation.py index b63fd0ecdb1..df219091074 100644 --- a/litellm/llms/bedrock_mantle/responses/transformation.py +++ b/litellm/llms/bedrock_mantle/responses/transformation.py @@ -4,14 +4,26 @@ Amazon Bedrock Mantle - Responses API backend. gpt-5.5 / gpt-5.4 on Mantle are exposed ONLY on the `/openai/v1/responses` path (not the standard `/v1/responses`). Payloads and SSE follow the OpenAI Responses spec, so this config inherits OpenAIResponsesAPIConfig and overrides -only the endpoint URL and Bearer authentication. +only the endpoint URL and authentication. -Auth: AWS Bedrock API key as Bearer token (BEDROCK_MANTLE_API_KEY or the -standard AWS_BEARER_TOKEN_BEDROCK), NOT SigV4. +Auth: Bearer token (BEDROCK_MANTLE_API_KEY or the standard +AWS_BEARER_TOKEN_BEDROCK, or litellm_params.api_key) when present; otherwise +AWS SigV4 (service name "bedrock") using the standard credential chain (IAM +role / access key / profile / web identity), signed via the shared +BaseAWSLLM._sign_request after the request body is finalized. """ -from typing import Optional +import re +from typing import Optional, Tuple +from botocore.exceptions import ( + CredentialRetrievalError, + NoCredentialsError, + PartialCredentialsError, + ProfileNotFound, +) + +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.secret_managers.main import get_secret_str from litellm.types.router import GenericLiteLLMParams @@ -29,22 +41,44 @@ _BASE_SUFFIXES_TO_STRIP = ( "/v1", ) +# Standard Mantle host: https://bedrock-mantle..api.aws (group 1 = region). +_MANTLE_HOST_RE = re.compile( + r"^https?://bedrock-mantle\.([^/.]+)\.api\.aws", re.IGNORECASE +) + class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig): + def __init__(self, aws_signer: Optional[BaseAWSLLM] = None): + super().__init__() + self._aws_signer = aws_signer or BaseAWSLLM() + @property def custom_llm_provider(self) -> LlmProviders: return LlmProviders.BEDROCK_MANTLE + @staticmethod + def _resolve_region(params: dict) -> str: + region = params.get("aws_region_name") + if region: + return region + base = params.get("api_base") or get_secret_str("BEDROCK_MANTLE_API_BASE") + if base: + match = _MANTLE_HOST_RE.match(base.rstrip("/")) + if match: + return match.group(1) + return ( + get_secret_str("BEDROCK_MANTLE_REGION") + or get_secret_str("AWS_REGION_NAME") + or get_secret_str("AWS_REGION") + or BEDROCK_MANTLE_DEFAULT_REGION + ) + def get_complete_url( self, api_base: Optional[str], litellm_params: dict, ) -> str: - region = ( - get_secret_str("BEDROCK_MANTLE_REGION") - or get_secret_str("AWS_REGION") - or BEDROCK_MANTLE_DEFAULT_REGION - ) + region = self._resolve_region({**litellm_params, "api_base": api_base}) base = ( api_base or get_secret_str("BEDROCK_MANTLE_API_BASE") @@ -55,6 +89,11 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig): if base.endswith(suffix): base = base[: -len(suffix)] break + # For the standard Mantle host (including the default-region base that + # responses/main.py auto-injects into litellm_params.api_base), pin to the + # single resolved region so aws_region_name wins; preserve custom proxy hosts. + if _MANTLE_HOST_RE.match(base): + base = f"https://bedrock-mantle.{region}.api.aws" return f"{base}/openai/v1/responses" def validate_environment( @@ -66,12 +105,8 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig): or get_secret_str("BEDROCK_MANTLE_API_KEY") or get_secret_str("AWS_BEARER_TOKEN_BEDROCK") ) - if not api_key: - raise ValueError( - "Bedrock Mantle API key is required. Set BEDROCK_MANTLE_API_KEY " - "(or AWS_BEARER_TOKEN_BEDROCK) or pass api_key." - ) - headers["Authorization"] = f"Bearer {api_key}" + if api_key: + headers["Authorization"] = f"Bearer {api_key}" return headers def supports_native_file_search(self) -> bool: @@ -79,3 +114,58 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig): def supports_native_websocket(self) -> bool: return False + + def sign_request( + self, + headers: dict, + optional_params: dict, + request_data: dict, + api_base: str, + api_key: Optional[str] = None, + model: Optional[str] = None, + stream: Optional[bool] = None, + fake_stream: Optional[bool] = None, + ) -> Tuple[dict, Optional[bytes]]: + bearer = ( + api_key + or get_secret_str("BEDROCK_MANTLE_API_KEY") + or get_secret_str("AWS_BEARER_TOKEN_BEDROCK") + ) + if not bearer: + # SigV4 path. Pin the credential-scope region to the region of the actual + # signing URL (api_base, already region-resolved by get_complete_url) so the + # SigV4 scope and the URL host can never disagree. Resolve from api_base first, + # then fall back to the regular precedence. Also drop any caller Authorization + # so _sign_request's restore-original-Authorization step cannot override the + # SigV4 header. + optional_params = { + **optional_params, + "aws_region_name": self._resolve_region( + {**optional_params, "api_base": api_base} + ), + } + headers = {k: v for k, v in headers.items() if k.lower() != "authorization"} + try: + return self._aws_signer._sign_request( + service_name="bedrock", + headers=headers, + optional_params=optional_params, + request_data=request_data, + api_base=api_base, + api_key=bearer, + model=model, + stream=stream, + fake_stream=fake_stream, + ) + except ( + NoCredentialsError, + PartialCredentialsError, + ProfileNotFound, + CredentialRetrievalError, + ) as e: + raise ValueError( + "Bedrock Mantle auth failed: no Bearer token and no usable AWS " + "credentials. Set BEDROCK_MANTLE_API_KEY (or AWS_BEARER_TOKEN_BEDROCK) " + "or pass api_key for Bearer auth, or provide AWS credentials " + "(IAM role / access key / profile / web identity) for SigV4." + ) from e diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 31c772510ba..25424feaeb4 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2318,6 +2318,31 @@ class BaseLLMHTTPHandler: # but never included in the outbound provider payload. request_context["litellm_params"] = dict(litellm_params) + is_stream_request = bool(stream) + if is_stream_request and fake_stream is True: + stream, data = self._prepare_fake_stream_request( + stream=stream, + data=data, + fake_stream=fake_stream, + ) + + # Sign after the body is final (post-transform/normalize/extra_body and post + # fake-stream prep) so signed bytes match what we send. No-op for providers + # that inherit the default sign_request. + headers, signed_body = responses_api_provider_config.sign_request( + headers=headers, + optional_params=dict(litellm_params), + request_data=data, + api_base=api_base, + api_key=litellm_params.api_key, + model=model, + stream=stream, + fake_stream=fake_stream, + ) + body_kwargs: Dict[str, Any] = ( + {"data": signed_body} if signed_body is not None else {"json": data} + ) + ## LOGGING logging_obj.pre_call( input=input, @@ -2330,22 +2355,14 @@ class BaseLLMHTTPHandler: ) try: - if stream: - # For streaming, use stream=True in the request - if fake_stream is True: - stream, data = self._prepare_fake_stream_request( - stream=stream, - data=data, - fake_stream=fake_stream, - ) - + if is_stream_request: response = sync_httpx_client.post( url=api_base, headers=headers, - json=data, timeout=timeout or float(response_api_optional_request_params.get("timeout", 0)), stream=stream, + **body_kwargs, ) if fake_stream is True: return MockResponsesAPIStreamingIterator( @@ -2370,13 +2387,12 @@ class BaseLLMHTTPHandler: call_type=CallTypes.responses.value, ) else: - # For non-streaming requests response = sync_httpx_client.post( url=api_base, headers=headers, - json=data, timeout=timeout or float(response_api_optional_request_params.get("timeout", 0)), + **body_kwargs, ) except Exception as e: raise self._handle_error( @@ -2464,6 +2480,28 @@ class BaseLLMHTTPHandler: # but never included in the outbound provider payload. request_context["litellm_params"] = dict(litellm_params) + is_stream_request = bool(stream) + if is_stream_request and fake_stream is True: + stream, data = self._prepare_fake_stream_request( + stream=stream, + data=data, + fake_stream=fake_stream, + ) + + headers, signed_body = responses_api_provider_config.sign_request( + headers=headers, + optional_params=dict(litellm_params), + request_data=data, + api_base=api_base, + api_key=litellm_params.api_key, + model=model, + stream=stream, + fake_stream=fake_stream, + ) + body_kwargs: Dict[str, Any] = ( + {"data": signed_body} if signed_body is not None else {"json": data} + ) + ## LOGGING logging_obj.pre_call( input=input, @@ -2476,22 +2514,14 @@ class BaseLLMHTTPHandler: ) try: - if stream: - # For streaming, we need to use stream=True in the request - if fake_stream is True: - stream, data = self._prepare_fake_stream_request( - stream=stream, - data=data, - fake_stream=fake_stream, - ) - + if is_stream_request: response = await async_httpx_client.post( url=api_base, headers=headers, - json=data, timeout=timeout or float(response_api_optional_request_params.get("timeout", 0)), stream=stream, + **body_kwargs, ) if fake_stream is True: @@ -2518,13 +2548,12 @@ class BaseLLMHTTPHandler: call_type=CallTypes.responses.value, ) else: - # For non-streaming, proceed as before response = await async_httpx_client.post( url=api_base, headers=headers, - json=data, timeout=timeout or float(response_api_optional_request_params.get("timeout", 0)), + **body_kwargs, ) except Exception as e: @@ -4005,6 +4034,18 @@ class BaseLLMHTTPHandler: ) data = BaseResponsesAPIConfig.normalize_responses_api_request_dict(data) + headers, signed_body = responses_api_provider_config.sign_request( + headers=headers, + optional_params=dict(litellm_params), + request_data=data, + api_base=url, + api_key=litellm_params.api_key, + model=model, + ) + body_kwargs: Dict[str, Any] = ( + {"data": signed_body} if signed_body is not None else {"json": data} + ) + ## LOGGING logging_obj.pre_call( input=input, @@ -4018,7 +4059,7 @@ class BaseLLMHTTPHandler: try: response = sync_httpx_client.post( - url=url, headers=headers, json=data, timeout=timeout + url=url, headers=headers, timeout=timeout, **body_kwargs ) except Exception as e: @@ -4088,6 +4129,18 @@ class BaseLLMHTTPHandler: ) data = BaseResponsesAPIConfig.normalize_responses_api_request_dict(data) + headers, signed_body = responses_api_provider_config.sign_request( + headers=headers, + optional_params=dict(litellm_params), + request_data=data, + api_base=url, + api_key=litellm_params.api_key, + model=model, + ) + body_kwargs: Dict[str, Any] = ( + {"data": signed_body} if signed_body is not None else {"json": data} + ) + ## LOGGING logging_obj.pre_call( input=input, @@ -4101,7 +4154,7 @@ class BaseLLMHTTPHandler: try: response = await async_httpx_client.post( - url=url, headers=headers, json=data, timeout=timeout + url=url, headers=headers, timeout=timeout, **body_kwargs ) except Exception as e: diff --git a/litellm/llms/deepseek/messages/transformation.py b/litellm/llms/deepseek/messages/transformation.py index ad60478960e..63b736ffd1d 100644 --- a/litellm/llms/deepseek/messages/transformation.py +++ b/litellm/llms/deepseek/messages/transformation.py @@ -26,6 +26,9 @@ class DeepSeekAnthropicMessagesConfig(AnthropicMessagesConfig): def custom_llm_provider(self) -> Optional[str]: return "deepseek" + def should_strip_billing_metadata(self) -> bool: + return True + @staticmethod def get_api_key(api_key: Optional[str] = None) -> Optional[str]: return api_key or get_secret_str("DEEPSEEK_API_KEY") or litellm.api_key diff --git a/litellm/llms/github_copilot/responses/transformation.py b/litellm/llms/github_copilot/responses/transformation.py index 0929f95cf43..3406538c774 100644 --- a/litellm/llms/github_copilot/responses/transformation.py +++ b/litellm/llms/github_copilot/responses/transformation.py @@ -2,7 +2,7 @@ GitHub Copilot Responses API Configuration. This module provides the configuration for GitHub Copilot's Responses API, -which is required for models like gpt-5.1-codex that only support the /responses endpoint. +which is required for models like gpt-5.3-codex that only support the /responses endpoint. Implementation based on analysis of the copilot-api project by caozhiyuan: https://github.com/caozhiyuan/copilot-api @@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Any, Dict, Optional, Union import os +import litellm from litellm._logging import verbose_logger from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH from litellm.exceptions import AuthenticationError @@ -22,6 +23,7 @@ from litellm.types.llms.openai import ( ) from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import LlmProviders +from litellm.utils import _cached_get_model_info_helper from ..authenticator import Authenticator from ..common_utils import ( @@ -38,6 +40,47 @@ else: LiteLLMLoggingObj = Any +def github_copilot_supports_responses_api(model: str) -> bool: + """ + Gate native /v1/responses dispatch per github_copilot model. + + Resolution (first match wins): mode "responses" -> True; mode "chat" -> + False (opt-out wins for dual-endpoint models); "/v1/responses" in + supported_endpoints -> True; else False. Unknown model -> False (the bridge + always works since every Copilot model supports /chat/completions). + + Reads merged model info (per-deployment model_info applied via the router's + register_model, which also clears the cache used here). + """ + try: + info = _cached_get_model_info_helper( + model=model, custom_llm_provider="github_copilot" + ) + except Exception as e: + verbose_logger.debug( + "github_copilot_supports_responses_api: get_model_info failed " + "for %s: %s", + model, + e, + ) + return False + + mode = info.get("mode") + if mode == "responses": + return True + if mode == "chat": + return False + + # supported_endpoints is dropped by ModelInfoBase; read it from the raw + # model_cost entry via the resolved key. + key = info.get("key") + raw_info = litellm.model_cost.get(key) if isinstance(key, str) else None + endpoints = ( + raw_info.get("supported_endpoints") if isinstance(raw_info, dict) else None + ) + return isinstance(endpoints, list) and "/v1/responses" in endpoints + + class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): """ Configuration for GitHub Copilot's Responses API. diff --git a/litellm/llms/minimax/messages/transformation.py b/litellm/llms/minimax/messages/transformation.py index 3190a5f5412..57cfcbf0621 100644 --- a/litellm/llms/minimax/messages/transformation.py +++ b/litellm/llms/minimax/messages/transformation.py @@ -28,6 +28,9 @@ class MinimaxMessagesConfig(AnthropicMessagesConfig): def custom_llm_provider(self) -> Optional[str]: return "minimax" + def should_strip_billing_metadata(self) -> bool: + return True + @staticmethod def get_api_key(api_key: Optional[str] = None) -> Optional[str]: """ diff --git a/litellm/llms/openai_like/dynamic_config.py b/litellm/llms/openai_like/dynamic_config.py index fac453447fa..9ed9734edae 100644 --- a/litellm/llms/openai_like/dynamic_config.py +++ b/litellm/llms/openai_like/dynamic_config.py @@ -187,6 +187,7 @@ def create_responses_config_class(provider: SimpleProviderConfig): from litellm.llms.openai_like.responses.transformation import ( OpenAILikeResponsesConfig, ) + from litellm.types.llms.openai import ResponseInputParam from litellm.types.router import GenericLiteLLMParams class JSONProviderResponsesConfig(OpenAILikeResponsesConfig): @@ -223,5 +224,23 @@ def create_responses_config_class(provider: SimpleProviderConfig): api_base = api_base.rstrip("/") return f"{api_base}/responses" + def transform_responses_api_request( + self, + model: str, + input: Union[str, ResponseInputParam], + response_api_optional_request_params: dict, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> dict: + if provider.special_handling.get("force_store_false"): + response_api_optional_request_params["store"] = False + return super().transform_responses_api_request( + model=model, + input=input, + response_api_optional_request_params=response_api_optional_request_params, + litellm_params=litellm_params, + headers=headers, + ) + _responses_config_cache[provider.slug] = JSONProviderResponsesConfig return JSONProviderResponsesConfig diff --git a/litellm/llms/openai_like/providers.json b/litellm/llms/openai_like/providers.json index 49b3801c82f..13d22488838 100644 --- a/litellm/llms/openai_like/providers.json +++ b/litellm/llms/openai_like/providers.json @@ -132,5 +132,14 @@ "param_mappings": { "max_completion_tokens": "max_tokens" } + }, + "parasail": { + "base_url": "https://api.parasail.io/v1", + "api_key_env": "PARASAIL_API_KEY", + "api_base_env": "PARASAIL_API_BASE", + "supported_endpoints": ["/v1/chat/completions", "/v1/responses"], + "special_handling": { + "force_store_false": true + } } } 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 5cd02293f14..7d355a2e908 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 @@ -1111,6 +1111,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): { "voice": "alloy", "format": "mp3", + "language_code": "en-US", } Expected output: @@ -1119,7 +1120,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): prebuiltVoiceConfig: { voiceName: "alloy", } - } + }, + languageCode: "en-US", } """ from litellm.types.llms.vertex_ai import ( @@ -1145,6 +1147,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): voice_config: VoiceConfig = {"prebuiltVoiceConfig": prebuilt_voice_config} speech_config["voiceConfig"] = voice_config + if "language_code" in value: + speech_config["languageCode"] = value["language_code"] + return cast(dict, speech_config) @staticmethod diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py index 1e92754857b..8a92e7ec4a5 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py @@ -17,6 +17,9 @@ from ..output_params_utils import sanitize_vertex_anthropic_output_params class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, VertexBase): + def should_strip_billing_metadata(self) -> bool: + return True + def validate_anthropic_messages_environment( self, headers: dict, diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/transformation.py index c852909d475..ae8bdc55443 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/transformation.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/transformation.py @@ -52,6 +52,9 @@ class VertexAIAnthropicConfig(AnthropicConfig): def custom_llm_provider(self) -> Optional[str]: return "vertex_ai" + def should_strip_billing_metadata(self) -> bool: + return True + def _add_context_management_beta_headers( self, beta_set: set, context_management: dict ) -> None: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 57a7d860baa..fba968dd95d 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2312,6 +2312,24 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="List of MCP server fields that must be filled in for a submission to pass standards checks (e.g. ['description', 'source_url', 'alias']).", ) + disable_budget_reservation: Optional[bool] = Field( + None, + description=( + "If True, disables the optimistic per-request budget reservation " + "introduced in v1.84.0. " + "WARNING: This weakens hard budget enforcement. Without the reservation, " + "a burst of concurrent requests from a single key can each pass the " + "read-time spend check before any of them is charged, allowing a " + "configured budget to be exceeded under high concurrency. " + "Budgets are still evaluated on every request at read time, so " + "an already-exhausted budget is still rejected. " + "Enable only if your deployment is experiencing phantom " + "BudgetExceededError responses caused by leaked reservations " + "(see GitHub issue #27639). " + "A proxy-level WARNING is logged on every request while this flag " + "is active as a reminder that hard enforcement is relaxed." + ), + ) class ConfigYAML(LiteLLMPydanticObjectBase): diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index a970e0ddee8..666c01562b5 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -2425,6 +2425,7 @@ async def _run_centralized_common_checks( # noqa: PLR0915 user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, skip_budget_checks=skip_budget_checks, + general_settings=general_settings, ) @@ -2445,12 +2446,23 @@ async def _reserve_budget_after_common_checks( user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, skip_budget_checks: bool, + general_settings: dict, end_user_id: Optional[str] = None, end_user_object: Optional[LiteLLM_EndUserTable] = None, ) -> None: user_api_key_auth_obj.budget_reservation = None if skip_budget_checks: return + if general_settings.get("disable_budget_reservation") is True: + verbose_proxy_logger.warning( + "disable_budget_reservation is enabled: skipping optimistic budget " + "reservation. Budget enforcement is read-time only — concurrent " + "requests can each pass the spend check before their cost is recorded, " + "so a configured budget may be briefly exceeded under high concurrency. " + "Set disable_budget_reservation to False or remove it to restore " + "hard per-request budget enforcement." + ) + return from litellm.proxy.spend_tracking.budget_reservation import ( reserve_budget_for_request, diff --git a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py index 14d950ecdf4..d1ef165b46e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py +++ b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py @@ -312,11 +312,27 @@ class CrowdStrikeAIDRHandler(CustomGuardrail): event_type = "output" hook_name = "apply_guardrail (response)" - ai_guard_payload = { + ai_guard_payload: dict[str, Any] = { "guard_input": guard_input.model_dump(mode="json"), "event_type": event_type, } + model = inputs.get("model") + if model: + ai_guard_payload["model"] = model + + metadata = request_data.get("litellm_metadata", request_data.get("metadata")) + if isinstance(metadata, Mapping): + user_id = metadata.get("user_api_key_user_id") + if user_id: + ai_guard_payload["user_id"] = user_id + + extra_info: dict[str, str] = {} + user_email = metadata.get("user_api_key_user_email") + if user_email: + extra_info["user_name"] = user_email + ai_guard_payload["extra_info"] = extra_info + ai_guard_response = await self._call_crowdstrike_aidr_guard( ai_guard_payload, hook_name ) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index a9570616850..4812bed2f21 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -1220,7 +1220,7 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]: generic_role_mappings_group_claim = os.getenv( "GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", None ) - generic_role_mappoings_default_role = os.getenv( + generic_role_mappings_default_role = os.getenv( "GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", None ) if generic_role_mappings is not None: @@ -1239,7 +1239,7 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]: role_mappings_data = { "provider": "generic", "group_claim": generic_role_mappings_group_claim, - "default_role": generic_role_mappoings_default_role, + "default_role": generic_role_mappings_default_role, "roles": generic_user_role_mappings_data, } diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py index 29bbb37501f..6dd1f8548eb 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py @@ -32,6 +32,71 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import ( from litellm.types.utils import ImageResponse, LlmProviders, PassthroughCallTypes from litellm.utils import ModelResponse, TextCompletionResponse +# Hostnames that route to OpenAI-compatible APIs. +# +# `api.openai.com` is OpenAI proper. The two Azure domains below are *shared by +# every Azure Cognitive Service* (Speech, Vision, Language, ...), not just Azure +# OpenAI: `openai.azure.com` is the classic Azure OpenAI domain, while +# `cognitiveservices.azure.com` is used by newer "Azure AI Foundry" / +# Cognitive Services-hosted Azure OpenAI deployments. Because the hostname alone +# cannot tell Azure OpenAI apart from the other Cognitive Services on those +# domains, requests there must additionally carry an OpenAI-style path segment. +_OPENAI_HOSTNAMES = ("api.openai.com",) +_AZURE_OPENAI_HOSTNAMES = ("openai.azure.com", "cognitiveservices.azure.com") +# Path markers that identify an Azure request as Azure OpenAI rather than Speech +# / Vision / Language / ... `/openai/` is the native Azure OpenAI path prefix; +# `/v1/` is the OpenAI-v1 surface used by LiteLLM's pass-through routing. Other +# Cognitive Services use service-named prefixes and versions like `/v3.1/`, +# `/v1.0/`, so they do not collide with these markers. +_AZURE_OPENAI_PATH_MARKERS = ("/openai/", "/v1/") + + +def _hostname_matches(hostname: str, suffixes: tuple) -> bool: + """True if hostname equals one of `suffixes` or is a subdomain of it. + + Uses suffix matching (not a bare substring test) so look-alikes such as + `cognitiveservices.azure.com.attacker.example` are not accepted. + """ + return any( + hostname == suffix or hostname.endswith("." + suffix) for suffix in suffixes + ) + + +def _is_openai_compatible_host(hostname: Optional[str]) -> bool: + """True if the hostname is OpenAI proper or one of the Azure OpenAI domains. + + Hostname-only check, kept for the route-level helpers that additionally + require a specific OpenAI path (e.g. `/v1/chat/completions`). When only the + hostname would otherwise gate dispatch, use `_is_openai_compatible_url` so + non-OpenAI Azure Cognitive Services on the shared domains are excluded. + """ + if not hostname: + return False + return _hostname_matches(hostname, _OPENAI_HOSTNAMES) or _hostname_matches( + hostname, _AZURE_OPENAI_HOSTNAMES + ) + + +def _is_openai_compatible_url(url_route: Optional[str]) -> bool: + """True if the URL targets an OpenAI-compatible API surface. + + For the shared Azure Cognitive Services domains we additionally require an + OpenAI-style path segment (`/openai/` or `/v1/`) so non-OpenAI Azure services + (Speech, Vision, Language, ...) on the same domain are not misclassified as + OpenAI routes. + """ + if not url_route: + return False + parsed_url = urlparse(url_route) + hostname = parsed_url.hostname + if not hostname: + return False + if _hostname_matches(hostname, _OPENAI_HOSTNAMES): + return True + if _hostname_matches(hostname, _AZURE_OPENAI_HOSTNAMES): + return any(marker in parsed_url.path for marker in _AZURE_OPENAI_PATH_MARKERS) + return False + class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): """ @@ -52,12 +117,8 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): if not url_route: return False parsed_url = urlparse(url_route) - return bool( - parsed_url.hostname - and ( - "api.openai.com" in parsed_url.hostname - or "openai.azure.com" in parsed_url.hostname - ) + return ( + _is_openai_compatible_host(parsed_url.hostname) and "/v1/chat/completions" in parsed_url.path ) @@ -67,12 +128,8 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): if not url_route: return False parsed_url = urlparse(url_route) - return bool( - parsed_url.hostname - and ( - "api.openai.com" in parsed_url.hostname - or "openai.azure.com" in parsed_url.hostname - ) + return ( + _is_openai_compatible_host(parsed_url.hostname) and "/v1/images/generations" in parsed_url.path ) @@ -82,12 +139,8 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): if not url_route: return False parsed_url = urlparse(url_route) - return bool( - parsed_url.hostname - and ( - "api.openai.com" in parsed_url.hostname - or "openai.azure.com" in parsed_url.hostname - ) + return ( + _is_openai_compatible_host(parsed_url.hostname) and "/v1/images/edits" in parsed_url.path ) @@ -97,13 +150,8 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): if not url_route: return False parsed_url = urlparse(url_route) - return bool( - parsed_url.hostname - and ( - "api.openai.com" in parsed_url.hostname - or "openai.azure.com" in parsed_url.hostname - ) - and ("/v1/responses" in parsed_url.path or "/responses" in parsed_url.path) + return _is_openai_compatible_host(parsed_url.hostname) and ( + "/v1/responses" in parsed_url.path or "/responses" in parsed_url.path ) def _get_user_from_metadata( diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py index 292871bae67..af1d39da020 100644 --- a/litellm/proxy/pass_through_endpoints/success_handler.py +++ b/litellm/proxy/pass_through_endpoints/success_handler.py @@ -434,15 +434,20 @@ class PassThroughEndpointLogging: return False def is_openai_route(self, url_route: str): - """Check if the URL route is an OpenAI API route.""" + """Check if the URL route is an OpenAI API route. + + Uses the URL-aware helper so that non-OpenAI Azure Cognitive Services + (Speech, Vision, Language, ...) sharing the `*.cognitiveservices.azure.com` + / `*.openai.azure.com` domains are not misclassified as OpenAI routes. + """ if not url_route: return False - parsed_url = urlparse(url_route) - return parsed_url.hostname and ( - "api.openai.com" in parsed_url.hostname - or "openai.azure.com" in parsed_url.hostname + from .llm_provider_handlers.openai_passthrough_logging_handler import ( + _is_openai_compatible_url, ) + return _is_openai_compatible_url(url_route) + def is_gemini_route( self, url_route: str, custom_llm_provider: Optional[str] = None ): diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index e2ba8353591..d3d30642216 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -148,7 +148,9 @@ class LiteLLMCompletionResponsesConfig: # which is equivalent to "required" in OpenAI format return "required" elif tool_choice_type == "function": - # function type without name - fall back to required + function_name = tool_choice.get("name") + if function_name: + return {"type": "function", "function": {"name": function_name}} return "required" # Return as-is for unknown formats diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index 51429d0769e..00db7b199b6 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -232,6 +232,7 @@ class VoiceConfig(TypedDict): class SpeechConfig(TypedDict, total=False): voiceConfig: VoiceConfig + languageCode: str class GenerationConfig(TypedDict, total=False): diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 9633cecf96c..c3ea605dd9e 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -1540,6 +1540,11 @@ class ServerToolUse(BaseModel): web_search_requests: Optional[int] = None tool_search_requests: Optional[int] = None + def __getitem__(self, key: str) -> Optional[int]: + if key not in self.__class__.model_fields: + raise KeyError(key) + return getattr(self, key) + class Usage(SafeAttributeModel, CompletionUsage): _cache_creation_input_tokens: int = PrivateAttr( @@ -1570,7 +1575,7 @@ class Usage(SafeAttributeModel, CompletionUsage): completion_tokens_details: Optional[ Union[CompletionTokensDetailsWrapper, dict] ] = None, - server_tool_use: Optional[ServerToolUse] = None, + server_tool_use: Optional[Union[ServerToolUse, dict]] = None, cost: Optional[float] = None, **params, ): @@ -1671,6 +1676,9 @@ class Usage(SafeAttributeModel, CompletionUsage): prompt_tokens_details=_prompt_tokens_details or None, ) + if isinstance(server_tool_use, dict): + server_tool_use = ServerToolUse(**server_tool_use) + if server_tool_use is not None: self.server_tool_use = server_tool_use else: # maintain openai compatibility in usage object if possible @@ -3392,6 +3400,7 @@ class LlmProviders(str, Enum): POE = "poe" CHUTES = "chutes" NEOSANTARA = "neosantara" + PARASAIL = "parasail" XIAOMI_MIMO = "xiaomi_mimo" TENSORMESH = "tensormesh" LITELLM_AGENT = "litellm_agent" diff --git a/litellm/utils.py b/litellm/utils.py index 7312e71bbd1..8d9d0a409c6 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8895,7 +8895,13 @@ class ProviderConfigManager: elif litellm.LlmProviders.XAI == provider: return litellm.XAIResponsesAPIConfig() elif litellm.LlmProviders.GITHUB_COPILOT == provider: - return litellm.GithubCopilotResponsesAPIConfig() + from litellm.llms.github_copilot.responses.transformation import ( + github_copilot_supports_responses_api, + ) + + if model is None or github_copilot_supports_responses_api(model=model): + return litellm.GithubCopilotResponsesAPIConfig() + return None elif litellm.LlmProviders.CHATGPT == provider: return litellm.ChatGPTResponsesAPIConfig() elif litellm.LlmProviders.LITELLM_PROXY == provider: diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index a1ad20fffd1..6caab585ac9 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -1834,6 +1834,23 @@ "search": true } }, + "parasail": { + "display_name": "Parasail (`parasail`)", + "url": "https://docs.litellm.ai/docs/providers/parasail", + "endpoints": { + "chat_completions": true, + "messages": false, + "responses": true, + "embeddings": false, + "image_generations": false, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false, + "a2a": false + } + }, "perplexity": { "display_name": "Perplexity AI (`perplexity`)", "url": "https://docs.litellm.ai/docs/providers/perplexity", diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index 9cf253c379d..fa22ff6b392 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -2712,6 +2712,10 @@ def test_bedrock_top_k_param(model, expected_params): data = json.loads(mock_post.call_args.kwargs["data"]) if "mistral" in model: assert data["top_k"] == 2 + elif expected_params == {}: + # Models that don't support top_k produce no additionalModelRequestFields; + # the empty block is now omitted entirely rather than sent as `{}`. + assert "additionalModelRequestFields" not in data else: assert data["additionalModelRequestFields"] == expected_params @@ -3059,8 +3063,6 @@ async def test_bedrock_max_completion_tokens(model: str): assert request_body == { "messages": [{"role": "user", "content": [{"text": "Hello!"}]}], - "additionalModelRequestFields": {}, - "system": [], "inferenceConfig": {"maxTokens": 10}, } diff --git a/tests/test_litellm/caching/test_caching.py b/tests/test_litellm/caching/test_caching.py new file mode 100644 index 00000000000..02d62a19152 --- /dev/null +++ b/tests/test_litellm/caching/test_caching.py @@ -0,0 +1,48 @@ +import logging +import re + +from litellm.caching.caching import Cache +from litellm.types.caching import LiteLLMCacheType + + +def test_cache_key_debug_log_does_not_include_prompt_material(caplog): + cache = Cache(type=LiteLLMCacheType.LOCAL) + prompt_marker = "secret prompt material " + + with caplog.at_level(logging.DEBUG, logger="LiteLLM"): + cache_key = cache.get_cache_key( + model="gpt-4.1-mini", + messages=[ + {"role": "system", "content": prompt_marker * 100}, + {"role": "user", "content": "hello"}, + ], + tools=[ + { + "type": "function", + "function": { + "name": "lookup", + "parameters": { + "type": "object", + "properties": {"query": {"type": "string"}}, + }, + }, + } + ], + response_format={ + "type": "json_schema", + "json_schema": { + "name": "lookup_response", + "schema": {"type": "object"}, + }, + }, + stream=True, + ) + + assert re.fullmatch(r"[0-9a-f]{64}", cache_key) + + created_cache_key_logs = [ + record.getMessage() for record in caplog.records if "Created cache key:" in record.getMessage() + ] + assert created_cache_key_logs + assert all(prompt_marker not in message for message in created_cache_key_logs) + assert any(cache_key in message for message in created_cache_key_logs) diff --git a/tests/test_litellm/caching/test_redis_semantic_cache.py b/tests/test_litellm/caching/test_redis_semantic_cache.py index b50a35ef50e..13f9d00136d 100644 --- a/tests/test_litellm/caching/test_redis_semantic_cache.py +++ b/tests/test_litellm/caching/test_redis_semantic_cache.py @@ -523,3 +523,468 @@ async def test_redis_semantic_cache_async_set_cache_stores_cache_key_filter( filters={RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key"}, ttl=60, ) + + +def test_redis_semantic_cache_set_cache_uses_responses_string_input(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache) + redis_semantic_cache.llmcache = MagicMock() + redis_semantic_cache._get_cache_filters = MagicMock( + return_value={RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key"} + ) + redis_semantic_cache._get_ttl = MagicMock(return_value=None) + + redis_semantic_cache.set_cache( + key="test_key", + value={"content": "Paris"}, + input="What is the capital of France?", + ) + + redis_semantic_cache.llmcache.store.assert_called_once_with( + "What is the capital of France?", + "{'content': 'Paris'}", + filters={RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key"}, + ) + + +def test_redis_semantic_cache_get_cache_uses_responses_string_input(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache) + redis_semantic_cache.similarity_threshold = 0.8 + redis_semantic_cache.llmcache = MagicMock() + redis_semantic_cache.llmcache.check = MagicMock( + return_value=[ + { + "prompt": "What is the capital of France?", + "response": '{"content": "Paris"}', + "vector_distance": 0.1, + RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key", + } + ] + ) + + with patch.object( + redis_semantic_cache, + "_get_cache_key_filter_expression", + return_value="cache-key-filter", + ): + metadata = {} + result = redis_semantic_cache.get_cache( + key="test_key", + input="What is the capital of France?", + metadata=metadata, + ) + + assert result == {"content": "Paris"} + assert metadata["semantic-similarity"] == pytest.approx(0.9) + redis_semantic_cache.llmcache.check.assert_called_once_with( + prompt="What is the capital of France?", + filter_expression="cache-key-filter", + ) + + +def test_redis_semantic_cache_set_cache_flattens_structured_responses_input(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache) + redis_semantic_cache.llmcache = MagicMock() + redis_semantic_cache._get_cache_filters = MagicMock( + return_value={RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key"} + ) + redis_semantic_cache._get_ttl = MagicMock(return_value=None) + + redis_semantic_cache.set_cache( + key="test_key", + value={"content": "Paris"}, + input=[ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "What is the capital of France?"}, + {"type": "input_text", "text": "Answer briefly."}, + { + "type": "input_image", + "image_url": "https://example.com/paris.png", + }, + ], + } + ], + ) + + redis_semantic_cache.llmcache.store.assert_called_once_with( + "What is the capital of France?\nAnswer briefly.", + "{'content': 'Paris'}", + filters={RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key"}, + ) + + +def test_redis_semantic_cache_prompt_extraction_prefers_messages(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + prompt = RedisSemanticCache._get_prompt_from_kwargs( + messages=[{"content": "message prompt"}], + input="responses prompt", + ) + + assert prompt == "message prompt" + + +def test_redis_semantic_cache_prompt_extraction_handles_model_objects(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + class ModelDumpInput: + def model_dump(self): + return {"content": [{"text": "model dump prompt"}]} + + class DictInput: + def dict(self): + return {"content": [{"output_text": "dict prompt"}]} + + prompt = RedisSemanticCache._get_prompt_from_kwargs( + input=[ + ModelDumpInput(), + DictInput(), + {"content": [{"input_text": "inline prompt"}]}, + {"content": [{"type": "input_image", "image_url": "https://example.com"}]}, + ] + ) + + assert prompt == "model dump prompt\ndict prompt\ninline prompt" + + +def test_redis_semantic_cache_prompt_extraction_returns_none_without_text(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + assert RedisSemanticCache._get_prompt_from_kwargs() is None + assert RedisSemanticCache._get_prompt_from_kwargs(input=None) is None + assert RedisSemanticCache._get_prompt_from_kwargs(input=" ") is None + assert ( + RedisSemanticCache._get_prompt_from_kwargs( + input=[{"type": "input_image", "image_url": "https://example.com"}] + ) + is None + ) + + +def test_redis_semantic_cache_prompt_extraction_skips_blank_dict_text_keys(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + prompt = RedisSemanticCache._get_prompt_from_kwargs( + input={"text": " ", "input_text": "fallback prompt"} + ) + + assert prompt == "fallback prompt" + + +def test_redis_semantic_cache_prompt_extraction_skips_blank_object_text_keys(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + class ResponseInput: + text = " " + input_text = "fallback prompt" + + prompt = RedisSemanticCache._get_prompt_from_kwargs(input=ResponseInput()) + + assert prompt == "fallback prompt" + + +def test_redis_semantic_cache_prompt_extraction_handles_object_content(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + class ResponseInput: + content = [{"text": "object content prompt"}] + + prompt = RedisSemanticCache._get_prompt_from_kwargs(input=ResponseInput()) + + assert prompt == "object content prompt" + + +def test_redis_semantic_cache_set_cache_skips_blank_responses_input(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache) + redis_semantic_cache.llmcache = MagicMock() + + redis_semantic_cache.set_cache( + key="test_key", + value={"content": "Paris"}, + input=" ", + ) + + redis_semantic_cache.llmcache.store.assert_not_called() + + +def test_redis_semantic_cache_get_cache_sets_similarity_on_blank_responses_input(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache) + redis_semantic_cache.llmcache = MagicMock() + metadata = {} + + result = redis_semantic_cache.get_cache( + key="test_key", + input=" ", + metadata=metadata, + ) + + assert result is None + assert metadata["semantic-similarity"] == 0.0 + redis_semantic_cache.llmcache.check.assert_not_called() + + +def test_redis_semantic_cache_get_cache_sets_similarity_when_no_results(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache) + redis_semantic_cache.llmcache = MagicMock() + redis_semantic_cache.llmcache.check = MagicMock(return_value=[]) + + with patch.object( + redis_semantic_cache, + "_get_cache_key_filter_expression", + return_value="cache-key-filter", + ): + metadata = {} + result = redis_semantic_cache.get_cache( + key="test_key", + input="What is the capital of France?", + metadata=metadata, + ) + + assert result is None + assert metadata["semantic-similarity"] == 0.0 + redis_semantic_cache.llmcache.check.assert_called_once_with( + prompt="What is the capital of France?", + filter_expression="cache-key-filter", + ) + + +@pytest.mark.asyncio +async def test_redis_semantic_cache_async_paths_use_responses_string_input(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache) + redis_semantic_cache.similarity_threshold = 0.8 + redis_semantic_cache.llmcache = MagicMock() + redis_semantic_cache.llmcache.astore = AsyncMock() + redis_semantic_cache.llmcache.acheck = AsyncMock( + return_value=[ + { + "prompt": "What is the capital of France?", + "response": '{"content": "Paris"}', + "vector_distance": 0.1, + RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key", + } + ] + ) + redis_semantic_cache._get_cache_filters = MagicMock( + return_value={RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key"} + ) + redis_semantic_cache._get_ttl = MagicMock(return_value=None) + redis_semantic_cache._get_async_embedding = AsyncMock(return_value=[0.1, 0.2, 0.3]) + + await redis_semantic_cache.async_set_cache( + key="test_key", + value={"content": "Paris"}, + input="What is the capital of France?", + ) + + with patch.object( + redis_semantic_cache, + "_get_cache_key_filter_expression", + return_value="cache-key-filter", + ): + metadata = {} + result = await redis_semantic_cache.async_get_cache( + key="test_key", + input="What is the capital of France?", + metadata=metadata, + ) + + redis_semantic_cache.llmcache.astore.assert_called_once_with( + "What is the capital of France?", + "{'content': 'Paris'}", + vector=[0.1, 0.2, 0.3], + filters={RedisSemanticCache.CACHE_KEY_FIELD_NAME: "test_key"}, + ) + assert result == {"content": "Paris"} + assert metadata["semantic-similarity"] == pytest.approx(0.9) + redis_semantic_cache.llmcache.acheck.assert_called_once_with( + prompt="What is the capital of France?", + vector=[0.1, 0.2, 0.3], + filter_expression="cache-key-filter", + ) + + +@pytest.mark.asyncio +async def test_redis_semantic_cache_async_paths_set_similarity_on_misses(): + from litellm.caching.redis_semantic_cache import RedisSemanticCache + + redis_semantic_cache = RedisSemanticCache.__new__(RedisSemanticCache) + redis_semantic_cache.llmcache = MagicMock() + redis_semantic_cache.llmcache.astore = AsyncMock() + redis_semantic_cache.llmcache.acheck = AsyncMock(return_value=[]) + redis_semantic_cache._get_async_embedding = AsyncMock(return_value=[0.1, 0.2, 0.3]) + + await redis_semantic_cache.async_set_cache( + key="test_key", + value={"content": "Paris"}, + input=" ", + ) + + redis_semantic_cache.llmcache.astore.assert_not_called() + redis_semantic_cache._get_async_embedding.assert_not_called() + + blank_metadata = {} + blank_result = await redis_semantic_cache.async_get_cache( + key="test_key", + input=" ", + metadata=blank_metadata, + ) + + assert blank_result is None + assert blank_metadata["semantic-similarity"] == 0.0 + redis_semantic_cache.llmcache.acheck.assert_not_called() + redis_semantic_cache._get_async_embedding.assert_not_called() + + with patch.object( + redis_semantic_cache, + "_get_cache_key_filter_expression", + return_value="cache-key-filter", + ): + miss_metadata = {} + miss_result = await redis_semantic_cache.async_get_cache( + key="test_key", + input="What is the capital of France?", + metadata=miss_metadata, + ) + + assert miss_result is None + assert miss_metadata["semantic-similarity"] == 0.0 + redis_semantic_cache.llmcache.acheck.assert_called_once_with( + prompt="What is the capital of France?", + vector=[0.1, 0.2, 0.3], + filter_expression="cache-key-filter", + ) + + +def test_cache_get_cache_passes_responses_input_to_backend_cache(): + from litellm.caching.caching import Cache + + cache = Cache.__new__(Cache) + cache.cache = MagicMock() + cache.cache.get_cache = MagicMock(return_value=None) + cache.should_use_cache = MagicMock(return_value=True) + cache.get_cache_key = MagicMock(return_value="test_key") + + metadata = {} + cache.get_cache( + input="What is the capital of France?", + metadata=metadata, + cache={}, + ) + + cache.cache.get_cache.assert_called_once_with( + "test_key", + input="What is the capital of France?", + metadata=metadata, + ) + + +def test_cache_get_cache_filters_sensitive_kwargs_from_backend_cache(): + from litellm.caching.caching import Cache + + cache = Cache.__new__(Cache) + cache.cache = MagicMock() + cache.should_use_cache = MagicMock(return_value=True) + cache.get_cache_key = MagicMock(return_value="test_key") + cache._get_cache_logic = MagicMock(return_value={"content": "Paris"}) + + def _cache_hit(_cache_key, **cache_kwargs): + cache_kwargs["metadata"]["semantic-similarity"] = 0.7 + return {"content": "Paris"} + + cache.cache.get_cache = MagicMock(side_effect=_cache_hit) + + metadata = {"user_api_key": "sk-secret", "trace_id": "trace-id"} + result = cache.get_cache( + input="What is the capital of France?", + metadata=metadata, + cache={"s-maxage": 10}, + api_key="sk-secret", + headers={"authorization": "Bearer sk-secret"}, + ) + + assert result == {"content": "Paris"} + assert metadata == { + "user_api_key": "sk-secret", + "trace_id": "trace-id", + "semantic-similarity": 0.7, + } + + forwarded_kwargs = cache.cache.get_cache.call_args.kwargs + assert forwarded_kwargs == { + "input": "What is the capital of France?", + "metadata": {"semantic-similarity": 0.7}, + } + assert forwarded_kwargs["metadata"] is not metadata + cache._get_cache_logic.assert_called_once_with( + cached_result={"content": "Paris"}, + max_age=10, + ) + + +def test_cache_get_cache_filters_sensitive_kwargs_without_metadata(): + from litellm.caching.caching import Cache + + cache = Cache.__new__(Cache) + cache.cache = MagicMock() + cache.cache.get_cache = MagicMock(return_value={"content": "Paris"}) + cache.should_use_cache = MagicMock(return_value=True) + cache.get_cache_key = MagicMock(return_value="test_key") + cache._get_cache_logic = MagicMock(return_value={"content": "Paris"}) + + result = cache.get_cache( + input="What is the capital of France?", + cache={"s-maxage": 10}, + api_key="sk-secret", + headers={"authorization": "Bearer sk-secret"}, + ) + + assert result == {"content": "Paris"} + cache.cache.get_cache.assert_called_once_with( + "test_key", + input="What is the capital of France?", + ) + + +def test_cache_get_cache_passes_responses_input_to_dynamic_cache(): + from litellm.caching.caching import Cache + + cache = Cache.__new__(Cache) + cache.should_use_cache = MagicMock(return_value=True) + cache.get_cache_key = MagicMock(return_value="test_key") + cache._get_cache_logic = MagicMock(return_value={"content": "Paris"}) + dynamic_cache_object = MagicMock() + dynamic_cache_object.get_cache = MagicMock(return_value={"content": "Paris"}) + + metadata = {} + result = cache.get_cache( + dynamic_cache_object=dynamic_cache_object, + input="What is the capital of France?", + metadata=metadata, + cache={}, + ) + + assert result == {"content": "Paris"} + dynamic_cache_object.get_cache.assert_called_once_with( + "test_key", + input="What is the capital of France?", + metadata=metadata, + ) + cache._get_cache_logic.assert_called_once_with( + cached_result={"content": "Paris"}, + max_age=float("inf"), + ) diff --git a/tests/test_litellm/integrations/focus/test_focus_gcs_destination.py b/tests/test_litellm/integrations/focus/test_focus_gcs_destination.py new file mode 100644 index 00000000000..35cdb18326a --- /dev/null +++ b/tests/test_litellm/integrations/focus/test_focus_gcs_destination.py @@ -0,0 +1,180 @@ +"""Tests for FocusGCSDestination.""" + +from __future__ import annotations + +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from litellm.integrations.focus.destinations.base import FocusTimeWindow + + +def _make_window(frequency: str = "hourly") -> FocusTimeWindow: + return FocusTimeWindow( + start_time=datetime(2026, 1, 1, 10, 0, 0, tzinfo=timezone.utc), + end_time=datetime(2026, 1, 1, 11, 0, 0, tzinfo=timezone.utc), + frequency=frequency, + ) + + +@pytest.mark.asyncio +async def test_deliver_posts_to_gcs_upload_endpoint(): + """deliver() must POST raw bytes to the GCS upload endpoint.""" + from litellm.integrations.focus.destinations.gcs_destination import ( + FocusGCSDestination, + ) + + dest = FocusGCSDestination( + prefix="focus_exports", + config={"bucket_name": "my-bucket", "service_account_json": None}, + ) + + mock_response = MagicMock() + mock_response.status_code = 200 + + mock_client = MagicMock() + mock_client.post = AsyncMock(return_value=mock_response) + dest.async_httpx_client = mock_client + + with patch.object( + dest, + "construct_request_headers", + new=AsyncMock(return_value={"Authorization": "Bearer tok-123"}), + ): + await dest.deliver( + content=b"col1,col2\nval1,val2\n", + time_window=_make_window(), + filename="usage_20260101T100000Z_20260101T110000Z.csv", + ) + + mock_client.post.assert_called_once() + call_kwargs = mock_client.post.call_args + url = call_kwargs.kwargs.get("url") or call_kwargs.args[0] + assert "my-bucket" in url + assert "uploadType=media" in url + headers = call_kwargs.kwargs["headers"] + assert headers["Authorization"] == "Bearer tok-123" + + +@pytest.mark.asyncio +async def test_deliver_raises_on_gcs_error(): + """deliver() must raise RuntimeError when GCS returns non-200.""" + from litellm.integrations.focus.destinations.gcs_destination import ( + FocusGCSDestination, + ) + + dest = FocusGCSDestination( + prefix="focus_exports", + config={"bucket_name": "my-bucket"}, + ) + + mock_response = MagicMock() + mock_response.status_code = 403 + mock_response.text = "Permission denied" + + mock_client = MagicMock() + mock_client.post = AsyncMock(return_value=mock_response) + dest.async_httpx_client = mock_client + + with patch.object( + dest, + "construct_request_headers", + new=AsyncMock(return_value={"Authorization": "Bearer tok-bad"}), + ): + with pytest.raises(RuntimeError, match="GCS upload failed"): + await dest.deliver( + content=b"data", + time_window=_make_window(), + filename="usage.csv", + ) + + +def test_build_object_key_hourly(): + """Hourly key must include date= and hour= components.""" + from litellm.integrations.focus.destinations.gcs_destination import ( + FocusGCSDestination, + ) + + dest = FocusGCSDestination(prefix="focus_exports", config={"bucket_name": "b"}) + key = dest._build_object_key( + time_window=_make_window("hourly"), filename="usage.parquet" + ) + + assert key == "focus_exports/date=2026-01-01/hour=10/usage.parquet" + + +def test_build_object_key_daily(): + """Daily key must include date= but not hour=.""" + from litellm.integrations.focus.destinations.gcs_destination import ( + FocusGCSDestination, + ) + + dest = FocusGCSDestination(prefix="focus_exports", config={"bucket_name": "b"}) + window = FocusTimeWindow( + start_time=datetime(2026, 1, 1, 0, 0, 0, tzinfo=timezone.utc), + end_time=datetime(2026, 1, 2, 0, 0, 0, tzinfo=timezone.utc), + frequency="daily", + ) + key = dest._build_object_key(time_window=window, filename="usage.parquet") + + assert key == "focus_exports/date=2026-01-01/usage.parquet" + + +def test_missing_bucket_name_raises(): + """Constructing without bucket_name must raise ValueError.""" + from litellm.integrations.focus.destinations.gcs_destination import ( + FocusGCSDestination, + ) + + with pytest.raises(ValueError, match="bucket_name"): + FocusGCSDestination(prefix="focus_exports", config={}) + + +def test_global_gcs_service_account_not_overwritten_when_absent(monkeypatch): + """service_account_json absent from config must not overwrite GCS_PATH_SERVICE_ACCOUNT. + + GCSBucketBase sets self.path_service_account_json from GCS_PATH_SERVICE_ACCOUNT. + If config has no service_account_json key, we must leave the parent value intact + so deployments using the global credential don't silently fall back to ADC. + """ + monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/global/sa.json") + + from litellm.integrations.focus.destinations.gcs_destination import ( + FocusGCSDestination, + ) + + dest = FocusGCSDestination(prefix="focus_exports", config={"bucket_name": "b"}) + + assert dest.path_service_account_json == "/global/sa.json" + + +def test_explicit_service_account_overrides_global(monkeypatch): + """Explicit service_account_json in config must take precedence over GCS_PATH_SERVICE_ACCOUNT.""" + monkeypatch.setenv("GCS_PATH_SERVICE_ACCOUNT", "/global/sa.json") + + from litellm.integrations.focus.destinations.gcs_destination import ( + FocusGCSDestination, + ) + + dest = FocusGCSDestination( + prefix="focus_exports", + config={"bucket_name": "b", "service_account_json": "/focus/sa.json"}, + ) + + assert dest.path_service_account_json == "/focus/sa.json" + + +def test_factory_creates_gcs_destination(monkeypatch): + """FocusDestinationFactory.create(provider='gcs') must return FocusGCSDestination.""" + monkeypatch.setenv("FOCUS_GCS_BUCKET_NAME", "env-bucket") + + from litellm.integrations.focus.destinations.factory import FocusDestinationFactory + from litellm.integrations.focus.destinations.gcs_destination import ( + FocusGCSDestination, + ) + + dest = FocusDestinationFactory.create(provider="gcs", prefix="focus_exports") + + assert isinstance(dest, FocusGCSDestination) + assert dest.BUCKET_NAME == "env-bucket" diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py index a04f6407e4b..c43291566b6 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_tool_call_cost_tracking.py @@ -1,17 +1,14 @@ -import json import os import sys -from unittest.mock import MagicMock import pytest -from fastapi.testclient import TestClient import litellm from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import ( StandardBuiltInToolCostTracking, ) from litellm.types.llms.openai import FileSearchTool, WebSearchOptions -from litellm.types.utils import ModelInfo, ModelResponse, StandardBuiltInToolsParams +from litellm.types.utils import ModelResponse, StandardBuiltInToolsParams sys.path.insert( 0, os.path.abspath("../../..") @@ -139,6 +136,22 @@ def test_get_cost_for_anthropic_web_search(): assert cost > 0.0 +def test_get_cost_for_anthropic_web_search_with_server_tool_use_dict(): + """ + Anthropic-compatible passthrough responses can construct Usage from a raw + usage payload. Ensure dict server_tool_use values are normalized before + built-in tool cost tracking reads server_tool_use.web_search_requests. + """ + from litellm.types.utils import ServerToolUse, Usage + + usage = Usage(server_tool_use={"web_search_requests": 1}) + + assert isinstance(usage.server_tool_use, ServerToolUse) + assert StandardBuiltInToolCostTracking.response_object_includes_web_search_call( + response_object=None, usage=usage + ) + + @pytest.mark.parametrize( "model", ["gemini/gemini-2.0-flash-001", "gemini-2.0-flash-001"] ) diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py index 4c330312930..75038574c63 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py @@ -5092,6 +5092,175 @@ def test_map_tool_helper_collision_prefers_definitions_over_components_schemas() assert transformed["input_schema"]["properties"]["from_components"] == expected +BILLING_HEADER_BLOCK = { + "type": "text", + "text": "x-anthropic-billing-header: cc_version=1.0.abc; cc_entrypoint=cli; cch=00000;", +} + + +def _system_with_billing_header(real_text: str) -> list: + return [ + { + "role": "system", + "content": [BILLING_HEADER_BLOCK, {"type": "text", "text": real_text}], + } + ] + + +def test_translate_system_message_keeps_billing_header_for_first_party_anthropic(): + config = AnthropicConfig() + assert config.should_strip_billing_metadata() is False + + result = config.translate_system_message( + messages=_system_with_billing_header( + "You are Claude Code, Anthropic's official CLI for Claude." + ) + ) + + texts = [block["text"] for block in result] + assert any(t.startswith("x-anthropic-billing-header:") for t in texts) + assert "You are Claude Code, Anthropic's official CLI for Claude." in texts + + +def test_translate_system_message_strips_billing_header_for_bedrock(): + from litellm.llms.bedrock.claude_platform.transformation import ( + BedrockClaudePlatformConfig, + ) + + config = BedrockClaudePlatformConfig() + assert config.should_strip_billing_metadata() is True + + result = config.translate_system_message( + messages=_system_with_billing_header("real system prompt") + ) + + texts = [block["text"] for block in result] + assert all(not t.startswith("x-anthropic-billing-header:") for t in texts) + assert "real system prompt" in texts + + +def test_anthropic_messages_request_keeps_billing_header_for_first_party(): + from litellm.types.router import GenericLiteLLMParams + + config = AnthropicMessagesConfig() + assert config.should_strip_billing_metadata() is False + + optional_params = { + "max_tokens": 16, + "system": [ + BILLING_HEADER_BLOCK, + {"type": "text", "text": "real system prompt"}, + ], + } + result = config.transform_anthropic_messages_request( + model="claude-3-5-sonnet-latest", + messages=[{"role": "user", "content": "hi"}], + anthropic_messages_optional_request_params=optional_params, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + texts = [block["text"] for block in result["system"]] + assert any(t.startswith("x-anthropic-billing-header:") for t in texts) + + +def test_anthropic_messages_request_strips_billing_header_for_minimax(): + from litellm.llms.minimax.messages.transformation import MinimaxMessagesConfig + from litellm.types.router import GenericLiteLLMParams + + config = MinimaxMessagesConfig() + assert config.should_strip_billing_metadata() is True + + optional_params = { + "max_tokens": 16, + "system": [ + BILLING_HEADER_BLOCK, + {"type": "text", "text": "real system prompt"}, + ], + } + result = config.transform_anthropic_messages_request( + model="MiniMax-M2", + messages=[{"role": "user", "content": "hi"}], + anthropic_messages_optional_request_params=optional_params, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + texts = [block["text"] for block in result.get("system", [])] + assert all(not t.startswith("x-anthropic-billing-header:") for t in texts) + + +def test_translate_system_message_strips_billing_header_for_bedrock_invoke(): + from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import ( + AmazonAnthropicClaudeConfig, + ) + + config = AmazonAnthropicClaudeConfig() + assert config.should_strip_billing_metadata() is True + + result = config.translate_system_message( + messages=_system_with_billing_header("real system prompt") + ) + + texts = [block["text"] for block in result] + assert all(not t.startswith("x-anthropic-billing-header:") for t in texts) + assert "real system prompt" in texts + + +@pytest.mark.parametrize( + "module_path, class_name, expected_strip", + [ + ("litellm.llms.anthropic.chat.transformation", "AnthropicConfig", False), + ( + "litellm.llms.anthropic.experimental_pass_through.messages.transformation", + "AnthropicMessagesConfig", + False, + ), + ( + "litellm.llms.bedrock.claude_platform.transformation", + "BedrockClaudePlatformConfig", + True, + ), + ( + "litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation", + "AmazonAnthropicClaudeConfig", + True, + ), + ( + "litellm.llms.vertex_ai.vertex_ai_partner_models.anthropic.transformation", + "VertexAIAnthropicConfig", + True, + ), + ( + "litellm.llms.azure_ai.anthropic.transformation", + "AzureAnthropicConfig", + True, + ), + ("litellm.llms.minimax.messages.transformation", "MinimaxMessagesConfig", True), + ( + "litellm.llms.azure_ai.anthropic.messages_transformation", + "AzureAnthropicMessagesConfig", + True, + ), + ( + "litellm.llms.deepseek.messages.transformation", + "DeepSeekAnthropicMessagesConfig", + True, + ), + ( + "litellm.llms.vertex_ai.vertex_ai_partner_models.anthropic.experimental_pass_through.transformation", + "VertexAIPartnerModelsAnthropicMessagesConfig", + True, + ), + ], +) +def test_should_strip_billing_metadata_by_provider( + module_path, class_name, expected_strip +): + import importlib + + config_cls = getattr(importlib.import_module(module_path), class_name) + assert config_cls().should_strip_billing_metadata() is expected_strip def test_namespace_tool_flat_nested_tools_are_extracted(): """Codex sends nested tools in flat format {type, name, description, parameters} with no 'function' wrapper. These must be normalized and mapped without raising KeyError: 'function'.""" diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index a6aa35ee6d1..ed978113b8b 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -1467,11 +1467,10 @@ def test_transform_request_with_function_tool(): ) # Verify the structure - assert "additionalModelRequestFields" in request_data - additional_fields = request_data["additionalModelRequestFields"] + # Function tools are not computer use tools, so they don't get anthropic_beta — + # additionalModelRequestFields should be absent (not serialized as empty {}) + assert "additionalModelRequestFields" not in request_data - # Function tools are not computer use tools, so they don't get anthropic_beta - # They are processed through the regular tool config assert "toolConfig" in request_data assert "tools" in request_data["toolConfig"] assert len(request_data["toolConfig"]["tools"]) == 1 diff --git a/tests/test_litellm/llms/bedrock/files/test_bedrock_files_transformation.py b/tests/test_litellm/llms/bedrock/files/test_bedrock_files_transformation.py index ba41fc47e8b..4731be13e78 100644 --- a/tests/test_litellm/llms/bedrock/files/test_bedrock_files_transformation.py +++ b/tests/test_litellm/llms/bedrock/files/test_bedrock_files_transformation.py @@ -128,6 +128,70 @@ class TestBedrockFilesTransformation: # Must have messages assert "messages" in model_input + # Nova Pro rejects empty additionalModelRequestFields / system — they must be absent + assert ( + "additionalModelRequestFields" not in model_input + ), "Nova: empty additionalModelRequestFields must be omitted, not serialized as {}" + assert ( + "system" not in model_input + ), "Nova: empty system must be omitted, not serialized as []" + + def test_nova_batch_jsonl_omits_empty_converse_fields(self): + """ + Regression test: Amazon Nova Pro returns 400 Malformed input request when + additionalModelRequestFields or system are present but empty in the Converse + API payload. The proxy must strip these keys when they carry no data. + """ + from litellm.llms.bedrock.files.transformation import BedrockFilesConfig + + config = BedrockFilesConfig() + + openai_jsonl_content = [ + { + "custom_id": "req-0", + "method": "POST", + "url": "/v1/chat/completions", + "body": { + "model": "us.amazon.nova-pro-v1:0", + "messages": [ + { + "role": "user", + "content": "What is 1 + 1? Answer with just the number.", + } + ], + "max_tokens": 16, + }, + } + ] + + result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content( + openai_jsonl_content + ) + + assert len(result) == 1 + model_input = result[0]["modelInput"] + + assert ( + "additionalModelRequestFields" not in model_input + or model_input["additionalModelRequestFields"] + ), "additionalModelRequestFields must be absent or non-empty — Nova rejects {}" + assert ( + "system" not in model_input or model_input["system"] + ), "system must be absent or non-empty — Nova rejects []" + + # Validate the exact shape AWS accepts + assert model_input == { + "messages": [ + { + "role": "user", + "content": [ + {"text": "What is 1 + 1? Answer with just the number."} + ], + } + ], + "inferenceConfig": {"maxTokens": 16}, + } + def test_nova_image_content_uses_converse_image_blocks(self): """ Test that image_url content blocks are converted to Bedrock Converse diff --git a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py index e2133d56f89..92b5ca7b10b 100644 --- a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py +++ b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_responses_transformation.py @@ -12,6 +12,11 @@ import sys sys.path.insert(0, os.path.abspath("../../../../..")) import pytest +from botocore.exceptions import ( + ConnectTimeoutError, + PartialCredentialsError, + ProfileNotFound, +) import litellm from litellm.llms.bedrock_mantle.responses.transformation import ( @@ -114,16 +119,15 @@ class TestBedrockMantleResponsesAuth: ) assert headers["Authorization"] == "Bearer bearer-key" - def test_missing_key_raises(self, monkeypatch): + def test_missing_bearer_does_not_raise_in_validate_environment(self, monkeypatch): + # SigV4 may still apply, so validate_environment must defer instead of raising. monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) cfg = BedrockMantleResponsesAPIConfig() - with pytest.raises(ValueError, match="Bedrock Mantle API key"): - cfg.validate_environment( - headers={}, - model="openai.gpt-5.5", - litellm_params=GenericLiteLLMParams(), - ) + headers = cfg.validate_environment( + headers={}, model="openai.gpt-5.5", litellm_params=GenericLiteLLMParams() + ) + assert "Authorization" not in headers def test_custom_llm_provider(self): cfg = BedrockMantleResponsesAPIConfig() @@ -261,6 +265,386 @@ def local_cost_map(monkeypatch): litellm.get_model_info.cache_clear() +class TestBedrockMantleResponsesSigV4: + def test_bearer_short_circuits_without_credentials(self, monkeypatch): + from unittest.mock import MagicMock + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + + signer = BaseAWSLLM() + signer.get_credentials = MagicMock( + side_effect=AssertionError("get_credentials must not run for bearer auth") + ) + cfg = BedrockMantleResponsesAPIConfig(aws_signer=signer) + + headers, signed_body = cfg.sign_request( + headers={}, + optional_params={}, + request_data={"input": "hi"}, + api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses", + api_key="bearer-from-config", + ) + assert headers["Authorization"] == "Bearer bearer-from-config" + assert signed_body == b'{"input": "hi"}' + signer.get_credentials.assert_not_called() + + def test_bearer_resolved_from_mantle_env_key(self, monkeypatch): + from unittest.mock import MagicMock + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + monkeypatch.setenv("BEDROCK_MANTLE_API_KEY", "env-bearer") + + signer = BaseAWSLLM() + signer.get_credentials = MagicMock( + side_effect=AssertionError("get_credentials must not run for bearer auth") + ) + cfg = BedrockMantleResponsesAPIConfig(aws_signer=signer) + + headers, _ = cfg.sign_request( + headers={}, + optional_params={}, + request_data={"input": "hi"}, + api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses", + api_key=None, + ) + assert headers["Authorization"] == "Bearer env-bearer" + + def test_bearer_arg_takes_priority_over_mantle_env_key(self, monkeypatch): + # The passed api_key (e.g. litellm_params.api_key) must win over the env + # bearer; a reordered precedence chain would silently use the wrong token. + from unittest.mock import MagicMock + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + monkeypatch.setenv("BEDROCK_MANTLE_API_KEY", "env-bearer") + + signer = BaseAWSLLM() + signer.get_credentials = MagicMock( + side_effect=AssertionError("get_credentials must not run for bearer auth") + ) + cfg = BedrockMantleResponsesAPIConfig(aws_signer=signer) + + headers, _ = cfg.sign_request( + headers={}, + optional_params={}, + request_data={"input": "hi"}, + api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses", + api_key="arg-bearer", + ) + assert headers["Authorization"] == "Bearer arg-bearer" + signer.get_credentials.assert_not_called() + + def test_access_key_produces_sigv4_headers(self, monkeypatch): + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + + cfg = BedrockMantleResponsesAPIConfig(aws_signer=BaseAWSLLM()) + headers, signed_body = cfg.sign_request( + headers={}, + optional_params={ + "aws_access_key_id": "AKIAEXAMPLE", + "aws_secret_access_key": "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0", + "aws_session_token": "session-token-test", + "aws_region_name": "us-east-2", + }, + request_data={"input": "hi"}, + api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses", + api_key=None, + ) + assert headers["Authorization"].startswith("AWS4-HMAC-SHA256") + assert "Credential=AKIAEXAMPLE/" in headers["Authorization"] + assert "/us-east-2/bedrock/aws4_request" in headers["Authorization"] + assert "X-Amz-Date" in headers + assert headers["X-Amz-Security-Token"] == "session-token-test" + assert signed_body == b'{"input": "hi"}' + + def test_assume_role_path_produces_sigv4_headers(self, monkeypatch): + from unittest.mock import MagicMock + from botocore.credentials import Credentials + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + + signer = BaseAWSLLM() + signer.get_credentials = MagicMock( + return_value=Credentials( + access_key="ASIAEXAMPLE", + secret_key="YXNzdW1lZC1yb2xlLXNlY3JldC1hc3N1bWVk", + token="assumed-session-token", + ) + ) + cfg = BedrockMantleResponsesAPIConfig(aws_signer=signer) + + headers, _ = cfg.sign_request( + headers={}, + optional_params={ + "aws_role_name": "arn:aws:iam::000000000000:role/test-role", + "aws_session_name": "litellm-test", + "aws_region_name": "us-east-2", + }, + request_data={"input": "hi"}, + api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses", + api_key=None, + ) + signer.get_credentials.assert_called_once() + call = signer.get_credentials.call_args.kwargs + assert call["aws_role_name"] == "arn:aws:iam::000000000000:role/test-role" + assert call["aws_session_name"] == "litellm-test" + assert headers["Authorization"].startswith("AWS4-HMAC-SHA256") + assert "/us-east-2/bedrock/aws4_request" in headers["Authorization"] + + def test_signed_body_matches_final_data_after_normalize(self, monkeypatch): + """Core regression: the signed bytes must equal the bytes actually sent. + + Sign the *final* data dict and assert the returned signed_body decodes to + exactly that dict, so a later change to the data would break the SigV4 hash. + """ + import json + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + + final_data = {"model": "openai.gpt-5.5", "input": "hi", "max_output_tokens": 16} + cfg = BedrockMantleResponsesAPIConfig(aws_signer=BaseAWSLLM()) + _, signed_body = cfg.sign_request( + headers={}, + optional_params={ + "aws_access_key_id": "AKIAEXAMPLE", + "aws_secret_access_key": "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0", + "aws_region_name": "us-east-2", + }, + request_data=final_data, + api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses", + api_key=None, + ) + assert signed_body is not None + assert json.loads(signed_body) == final_data + + def test_region_comes_from_optional_params(self, monkeypatch): + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + monkeypatch.delenv("AWS_REGION", raising=False) + monkeypatch.delenv("AWS_REGION_NAME", raising=False) + + cfg = BedrockMantleResponsesAPIConfig(aws_signer=BaseAWSLLM()) + headers, _ = cfg.sign_request( + headers={}, + optional_params={ + "aws_access_key_id": "AKIAEXAMPLE", + "aws_secret_access_key": "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0", + "aws_region_name": "eu-west-1", + }, + request_data={"input": "hi"}, + api_base="https://bedrock-mantle.eu-west-1.api.aws/openai/v1/responses", + api_key=None, + ) + assert "/eu-west-1/bedrock/aws4_request" in headers["Authorization"] + + def test_url_region_and_sigv4_region_agree_from_litellm_params(self, monkeypatch): + """Adversarial-review regression: a caller-supplied aws_region_name (no region + env set) must shape BOTH the URL host and the SigV4 credential scope, or the + request is signed for one region and sent to another -> 401. + """ + monkeypatch.delenv("BEDROCK_MANTLE_REGION", raising=False) + monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) + monkeypatch.delenv("AWS_REGION", raising=False) + monkeypatch.delenv("AWS_REGION_NAME", raising=False) + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + params = { + "aws_region_name": "ap-southeast-2", + "aws_access_key_id": "AKIAEXAMPLE", + "aws_secret_access_key": "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0", + } + cfg = BedrockMantleResponsesAPIConfig(aws_signer=BaseAWSLLM()) + url = cfg.get_complete_url(api_base=None, litellm_params=params) + assert ( + url == "https://bedrock-mantle.ap-southeast-2.api.aws/openai/v1/responses" + ) + + headers, _ = cfg.sign_request( + headers={}, + optional_params=params, + request_data={"input": "hi"}, + api_base=url, + api_key=None, + ) + assert "/ap-southeast-2/bedrock/aws4_request" in headers["Authorization"] + + def test_injected_default_region_base_does_not_override_aws_region_name( + self, monkeypatch + ): + """2nd-round adversarial regression: responses/main.py auto-injects + litellm_params.api_base = https://bedrock-mantle..api.aws/v1 (default + region, ignoring aws_region_name). The config must still pin BOTH the URL host + and the SigV4 scope to aws_region_name, or the IAM deployment 401s. A naive + 'resolve region only when api_base is None' fix would fail this test. + """ + monkeypatch.delenv("BEDROCK_MANTLE_REGION", raising=False) + monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) + monkeypatch.delenv("AWS_REGION", raising=False) + monkeypatch.delenv("AWS_REGION_NAME", raising=False) + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + injected_base = "https://bedrock-mantle.us-east-1.api.aws/v1" # default region + params = { + "aws_region_name": "us-east-2", # what the caller actually wants + "api_base": injected_base, + "aws_access_key_id": "AKIAEXAMPLE", + "aws_secret_access_key": "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0", + } + cfg = BedrockMantleResponsesAPIConfig(aws_signer=BaseAWSLLM()) + url = cfg.get_complete_url(api_base=injected_base, litellm_params=params) + assert url == "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses" + + headers, _ = cfg.sign_request( + headers={}, + optional_params=params, + request_data={"input": "hi"}, + api_base=url, + api_key=None, + ) + assert "/us-east-2/bedrock/aws4_request" in headers["Authorization"] + assert "us-east-1" not in headers["Authorization"] + + def test_custom_proxy_host_is_preserved(self, monkeypatch): + """A genuinely custom (non-Mantle) api_base host must be preserved, not rewritten + to a bedrock-mantle host. Only standard Mantle hosts are region-pinned. + """ + monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) + cfg = BedrockMantleResponsesAPIConfig() + url = cfg.get_complete_url( + api_base="https://mantle-proxy.internal.example/openai/v1", + litellm_params={"aws_region_name": "us-east-2"}, + ) + assert url == "https://mantle-proxy.internal.example/openai/v1/responses" + + def test_caller_authorization_does_not_override_sigv4(self, monkeypatch): + """Adversarial-review regression: a caller-supplied Authorization header (e.g. + from extra_headers, surviving the relaxed validate_environment) must not clobber + the SigV4 Authorization that _sign_request would otherwise restore. + """ + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + cfg = BedrockMantleResponsesAPIConfig(aws_signer=BaseAWSLLM()) + headers, _ = cfg.sign_request( + headers={"Authorization": "Bearer stale-caller-token"}, + optional_params={ + "aws_access_key_id": "AKIAEXAMPLE", + "aws_secret_access_key": "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0", + "aws_region_name": "us-east-2", + }, + request_data={"input": "hi"}, + api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses", + api_key=None, + ) + assert headers["Authorization"].startswith("AWS4-HMAC-SHA256") + assert "Bearer stale-caller-token" not in headers["Authorization"] + + def test_no_bearer_and_no_credentials_raises_both_paths(self, monkeypatch): + from unittest.mock import MagicMock + from botocore.exceptions import NoCredentialsError + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + + signer = BaseAWSLLM() + signer.get_credentials = MagicMock(side_effect=NoCredentialsError()) + cfg = BedrockMantleResponsesAPIConfig(aws_signer=signer) + + with pytest.raises(ValueError) as exc: + cfg.sign_request( + headers={}, + optional_params={"aws_region_name": "us-east-2"}, + request_data={"input": "hi"}, + api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses", + api_key=None, + ) + msg = str(exc.value) + assert "Bearer" in msg + assert "SigV4" in msg or "IAM" in msg + + @pytest.mark.parametrize( + "cred_error", + [ + PartialCredentialsError(provider="env", cred_var="aws_secret_access_key"), + ProfileNotFound(profile="missing-profile"), + ], + ) + def test_partial_credentials_raises_both_paths(self, monkeypatch, cred_error): + from unittest.mock import MagicMock + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + + signer = BaseAWSLLM() + signer.get_credentials = MagicMock(side_effect=cred_error) + cfg = BedrockMantleResponsesAPIConfig(aws_signer=signer) + + with pytest.raises(ValueError) as exc: + cfg.sign_request( + headers={}, + optional_params={"aws_region_name": "us-east-2"}, + request_data={"input": "hi"}, + api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses", + api_key=None, + ) + msg = str(exc.value) + assert "Bearer" in msg + assert "SigV4" in msg or "IAM" in msg + + def test_sts_transport_error_is_not_masked_as_credentials(self, monkeypatch): + # An AssumeRole / web-identity flow hits STS over the network, so a transient + # connection error must surface as itself, not be rewritten into the + # "no usable AWS credentials" message that would send the user to fix the + # wrong thing. + from unittest.mock import MagicMock + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + + signer = BaseAWSLLM() + signer.get_credentials = MagicMock( + side_effect=ConnectTimeoutError( + endpoint_url="https://sts.us-east-2.amazonaws.com" + ) + ) + cfg = BedrockMantleResponsesAPIConfig(aws_signer=signer) + + with pytest.raises(ConnectTimeoutError): + cfg.sign_request( + headers={}, + optional_params={ + "aws_role_name": "arn:aws:iam::000000000000:role/test-role", + "aws_region_name": "us-east-2", + }, + request_data={"input": "hi"}, + api_base="https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses", + api_key=None, + ) + + class TestBedrockMantleResponsesPricing: def test_gpt_5_5_pricing_and_mode(self, local_cost_map): info = litellm.get_model_info("bedrock_mantle/openai.gpt-5.5") diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index 279e9730e69..7321abcee46 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -742,3 +742,241 @@ async def test_anthropic_post_retry_reserializes_mutated_body(): assert first_sent == prebuilt # attempt 0 used prebuilt assert second_sent == _json.dumps(request_body) # attempt 1 re-serialized assert "MUTATED" in second_sent # ... the mutated body + + +def test_base_responses_config_sign_request_is_noop_by_default(): + """Default responses sign_request must be a no-op: unchanged headers, no signed body. + + Guards the 15 existing responses providers from accidental signing when the + handler starts calling sign_request. + """ + from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig + + cfg = OpenAIResponsesAPIConfig() + headers = {"Authorization": "Bearer sk-existing"} + out_headers, signed_body = cfg.sign_request( + headers=headers, + optional_params={}, + request_data={"input": "hi"}, + api_base="https://api.openai.com/v1/responses", + ) + assert out_headers == {"Authorization": "Bearer sk-existing"} + assert signed_body is None + + +def _make_responses_handler_call(signed_body): + """Drive BaseLLMHTTPHandler.response_api_handler with a fully mocked provider + config + sync client, returning the kwargs the client.post was called with. + + signed_body=None simulates a no-op (non-signing) provider; bytes simulates a + signing provider (e.g. Bedrock Mantle). + """ + from unittest.mock import MagicMock + from litellm.llms.custom_httpx.http_handler import HTTPHandler + from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler + from litellm.types.router import GenericLiteLLMParams + + provider_config = MagicMock() + provider_config.validate_environment.return_value = {} + provider_config.get_complete_url.return_value = ( + "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses" + ) + provider_config.transform_responses_api_request.return_value = {"input": "hi"} + provider_config.should_fake_stream.return_value = False + provider_config.sign_request.return_value = ({"X-Signed": "1"}, signed_body) + + mock_client = MagicMock(spec=HTTPHandler) + mock_client.post.return_value = MagicMock() + + handler = BaseLLMHTTPHandler() + handler.response_api_handler( + model="openai.gpt-5.5", + input="hi", + responses_api_provider_config=provider_config, + response_api_optional_request_params={}, + custom_llm_provider="bedrock_mantle", + litellm_params=GenericLiteLLMParams(aws_region_name="us-east-2"), + logging_obj=MagicMock(), + client=mock_client, + _is_async=False, + ) + return mock_client.post.call_args.kwargs + + +def test_responses_handler_sends_json_when_not_signed(): + """No-op provider (signed_body is None) -> handler posts json=data, no data= bytes.""" + kwargs = _make_responses_handler_call(signed_body=None) + assert kwargs.get("json") == {"input": "hi"} + assert "data" not in kwargs + + +def test_responses_handler_sends_signed_bytes_when_signed(): + """Signing provider -> handler posts the exact signed bytes via data=, not json=.""" + kwargs = _make_responses_handler_call(signed_body=b'{"input": "hi"}') + assert kwargs.get("data") == b'{"input": "hi"}' + assert "json" not in kwargs + assert kwargs["headers"] == {"X-Signed": "1"} + + +def test_responses_handler_signs_after_fake_stream_prep_strips_stream(): + """Fake-stream signing-order invariant: the bytes SIGNED must equal the bytes SENT. + + In the streaming + fake-stream path the handler first runs + _prepare_fake_stream_request, which pops "stream" out of the body, and only + then calls sign_request. If signing ran before that pop, the signed body + would still carry "stream" while the body sent over the wire would not, + producing a SigV4 payload-hash mismatch (401) for a real Mantle deployment. + We snapshot request_data at sign time and assert "stream" is already gone. + """ + from unittest.mock import MagicMock + from litellm.llms.custom_httpx.http_handler import HTTPHandler + from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler + from litellm.types.llms.openai import ResponsesAPIResponse + from litellm.types.router import GenericLiteLLMParams + + provider_config = MagicMock() + provider_config.validate_environment.return_value = {} + provider_config.get_complete_url.return_value = ( + "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses" + ) + provider_config.transform_responses_api_request.return_value = { + "input": "hi", + "stream": True, + } + provider_config.should_fake_stream.return_value = True + provider_config.transform_response_api_response.return_value = ResponsesAPIResponse( + id="resp_1", + created_at=0, + output=[], + status="completed", + model="openai.gpt-5.5", + ) + + captured = {} + + def _capture_sign(**kwargs): + captured["request_data"] = dict(kwargs["request_data"]) + return ({"X-Signed": "1"}, b'{"input": "hi"}') + + provider_config.sign_request.side_effect = _capture_sign + + mock_client = MagicMock(spec=HTTPHandler) + mock_client.post.return_value = MagicMock() + + handler = BaseLLMHTTPHandler() + handler.response_api_handler( + model="openai.gpt-5.5", + input="hi", + responses_api_provider_config=provider_config, + response_api_optional_request_params={"stream": True}, + custom_llm_provider="bedrock_mantle", + litellm_params=GenericLiteLLMParams(aws_region_name="us-east-2"), + logging_obj=MagicMock(), + client=mock_client, + _is_async=False, + fake_stream=True, + ) + + assert "stream" not in captured["request_data"] + assert "input" in captured["request_data"] + + post_kwargs = mock_client.post.call_args.kwargs + assert post_kwargs.get("data") == b'{"input": "hi"}' + assert "json" not in post_kwargs + assert "stream" in post_kwargs + + +def _make_compact_handler_call(signed_body, is_async): + """Drive (async_)compact_response_api_handler with a fully mocked provider config + + client, returning the kwargs the client.post was called with. + + signed_body=None simulates a no-op (non-signing) provider; bytes simulates a + signing provider (e.g. Bedrock Mantle SigV4 / bearer). + """ + from unittest.mock import MagicMock + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler + from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler + from litellm.types.router import GenericLiteLLMParams + + compact_url = "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses/compact" + provider_config = MagicMock() + provider_config.validate_environment.return_value = {} + provider_config.get_complete_url.return_value = ( + "https://bedrock-mantle.us-east-2.api.aws/openai/v1/responses" + ) + provider_config.transform_compact_response_api_request.return_value = ( + compact_url, + {"model": "openai.gpt-5.5", "input": "hi"}, + ) + provider_config.sign_request.return_value = ({"X-Signed": "1"}, signed_body) + provider_config.transform_compact_response_api_response.return_value = "ok" + + spec = AsyncHTTPHandler if is_async else HTTPHandler + mock_client = MagicMock(spec=spec) + if is_async: + mock_client.post = AsyncMock(return_value=MagicMock()) + else: + mock_client.post.return_value = MagicMock() + + handler = BaseLLMHTTPHandler() + result = handler.compact_response_api_handler( + model="openai.gpt-5.5", + input="hi", + responses_api_provider_config=provider_config, + response_api_optional_request_params={}, + custom_llm_provider="bedrock_mantle", + litellm_params=GenericLiteLLMParams(aws_region_name="us-east-2"), + logging_obj=MagicMock(), + client=mock_client, + _is_async=is_async, + ) + if is_async: + asyncio.run(result) + return provider_config, mock_client.post.call_args.kwargs + + +def test_compact_handler_sends_json_when_not_signed(): + """No-op provider on compact (signed_body is None) -> posts json=data, no data= bytes.""" + provider_config, kwargs = _make_compact_handler_call( + signed_body=None, is_async=False + ) + provider_config.sign_request.assert_called_once() + assert kwargs.get("json") == {"model": "openai.gpt-5.5", "input": "hi"} + assert "data" not in kwargs + + +def test_compact_handler_sends_signed_bytes_when_signed(): + """Signing provider on compact -> posts the signed bytes via data=, not json=. + + Regression for the adversarial-review finding that /responses/compact bypassed + the SigV4 signing hook, so IAM-only Mantle callers sent unsigned bodies. + """ + provider_config, kwargs = _make_compact_handler_call( + signed_body=b'{"model": "openai.gpt-5.5", "input": "hi"}', is_async=False + ) + assert kwargs.get("data") == b'{"model": "openai.gpt-5.5", "input": "hi"}' + assert "json" not in kwargs + assert kwargs["headers"] == {"X-Signed": "1"} + # signing must use the compact endpoint as api_base, not the create URL + assert provider_config.sign_request.call_args.kwargs["api_base"].endswith( + "/openai/v1/responses/compact" + ) + + +def test_async_compact_handler_sends_signed_bytes_when_signed(): + """Async compact must sign identically to sync (same omission in the async twin).""" + provider_config, kwargs = _make_compact_handler_call( + signed_body=b'{"model": "openai.gpt-5.5", "input": "hi"}', is_async=True + ) + assert kwargs.get("data") == b'{"model": "openai.gpt-5.5", "input": "hi"}' + assert "json" not in kwargs + assert kwargs["headers"] == {"X-Signed": "1"} + + +def test_async_compact_handler_sends_json_when_not_signed(): + """Async no-op provider on compact -> posts json=data, no data= bytes.""" + _provider_config, kwargs = _make_compact_handler_call( + signed_body=None, is_async=True + ) + assert kwargs.get("json") == {"model": "openai.gpt-5.5", "input": "hi"} + assert "data" not in kwargs diff --git a/tests/test_litellm/llms/gemini/test_gemini_tts.py b/tests/test_litellm/llms/gemini/test_gemini_tts.py index 65eefca5af1..98f3ac0f4e5 100644 --- a/tests/test_litellm/llms/gemini/test_gemini_tts.py +++ b/tests/test_litellm/llms/gemini/test_gemini_tts.py @@ -80,6 +80,46 @@ class TestGeminiTTSTransformation: assert "responseModalities" in result assert "AUDIO" in result["responseModalities"] + def test_gemini_tts_audio_parameter_mapping_with_language_code(self): + config = GoogleAIStudioGeminiConfig() + + non_default_params = { + "audio": {"voice": "Kore", "format": "pcm16", "language_code": "en-US"} + } + optional_params = {} + + result = config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model="gemini-2.5-flash-preview-tts", + drop_params=False, + ) + + assert "speechConfig" in result + assert result["speechConfig"]["languageCode"] == "en-US" + assert ( + result["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] + == "Kore" + ) + + def test_map_audio_params_language_code(self): + config = GoogleAIStudioGeminiConfig() + + result = config._map_audio_params( + {"voice": "Kore", "format": "pcm16", "language_code": "de-DE"} + ) + + assert result["languageCode"] == "de-DE" + assert result["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Kore" + + def test_map_audio_params_no_language_code(self): + config = GoogleAIStudioGeminiConfig() + + result = config._map_audio_params({"voice": "Kore", "format": "pcm16"}) + + assert "languageCode" not in result + assert result["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Kore" + def test_gemini_tts_audio_parameter_with_existing_modalities(self): """Test audio parameter mapping when modalities already exist""" config = GoogleAIStudioGeminiConfig() @@ -328,5 +368,57 @@ class TestGeminiTTSSpeechConfigInRequestBody: assert "AUDIO" in generation_config["responseModalities"] + @pytest.mark.parametrize( + "model,custom_llm_provider", + [ + ("gemini-2.5-flash-tts", "vertex_ai"), + ("gemini-2.5-flash-tts", "gemini"), + ("gemini-2.5-flash-preview-tts", "vertex_ai"), + ], + ) + def test_language_code_end_to_end_mapping(self, model, custom_llm_provider): + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, + ) + from litellm.llms.vertex_ai.gemini.transformation import ( + _transform_request_body, + ) + + config = VertexGeminiConfig() + + non_default_params = { + "audio": {"voice": "Puck", "format": "pcm16", "language_code": "pt-BR"} + } + optional_params = {} + + mapped_params = config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=model, + drop_params=False, + ) + + assert mapped_params["speechConfig"]["languageCode"] == "pt-BR" + + request_body = _transform_request_body( + messages=[{"role": "user", "content": "Hello world"}], + model=model, + optional_params=mapped_params, + custom_llm_provider=custom_llm_provider, + litellm_params={}, + cached_content=None, + ) + + generation_config = request_body["generationConfig"] + assert generation_config["speechConfig"]["languageCode"] == "pt-BR" + assert ( + generation_config["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"][ + "voiceName" + ] + == "Puck" + ) + assert "AUDIO" in generation_config["responseModalities"] + + if __name__ == "__main__": pytest.main([__file__]) diff --git a/tests/test_litellm/llms/github_copilot/responses/test_github_copilot_responses_transformation.py b/tests/test_litellm/llms/github_copilot/responses/test_github_copilot_responses_transformation.py index 54e7170bb20..17373f24a97 100644 --- a/tests/test_litellm/llms/github_copilot/responses/test_github_copilot_responses_transformation.py +++ b/tests/test_litellm/llms/github_copilot/responses/test_github_copilot_responses_transformation.py @@ -14,6 +14,8 @@ from unittest.mock import patch, MagicMock sys.path.insert(0, os.path.abspath("../../../../..")) import pytest +import litellm +from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map from litellm.types.utils import LlmProviders from litellm.utils import ProviderConfigManager from litellm.llms.github_copilot.responses.transformation import ( @@ -22,13 +24,26 @@ from litellm.llms.github_copilot.responses.transformation import ( from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams +@pytest.fixture(autouse=True) +def use_local_model_cost_map(monkeypatch: pytest.MonkeyPatch): + """Pin litellm.model_cost to the bundled local backup so tests don't depend + on remote catalog fetches (and don't change behavior across remote refreshes).""" + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr( + litellm, "model_cost", get_model_cost_map(url=litellm.model_cost_map_url) + ) + litellm.add_known_models(model_cost_map=litellm.model_cost) + + class TestGithubCopilotResponsesAPITransformation: """Test GitHub Copilot Responses API configuration and transformations""" def test_github_copilot_provider_config_registration(self): - """Test that GitHub Copilot provider returns GithubCopilotResponsesAPIConfig""" + """Test that GitHub Copilot provider returns the native Responses API + config for a Responses-capable catalog model. Exercises the full stack: + catalog lookup -> github_copilot_supports_responses_api -> native config.""" config = ProviderConfigManager.get_provider_responses_api_config( - model="github_copilot/gpt-5.1-codex", + model="github_copilot/gpt-5.3-codex", provider=LlmProviders.GITHUB_COPILOT, ) @@ -373,3 +388,200 @@ class TestGithubCopilotResponsesAPITransformation: # Non-reasoning items should pass through unchanged assert result == message_item + + +class TestGithubCopilotResponsesAPIRouting: + """``ProviderConfigManager.get_provider_responses_api_config`` for github_copilot + returns the native Responses config only when the model has ``mode=responses`` + in the (already-merged) model info; otherwise returns None so the dispatcher + routes through the chat-completions translation bridge.""" + + @patch( + "litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper" + ) + def test_returns_config_when_mode_is_responses(self, mock_get_info): + """``mode=responses`` returns native config.""" + mock_get_info.return_value = {"mode": "responses"} + config = ProviderConfigManager.get_provider_responses_api_config( + model="github_copilot/some-responses-model", + provider=LlmProviders.GITHUB_COPILOT, + ) + assert isinstance(config, GithubCopilotResponsesAPIConfig) + + @patch( + "litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper" + ) + def test_returns_none_when_mode_is_chat(self, mock_get_info): + """``mode=chat`` returns None so dispatcher uses bridge.""" + mock_get_info.return_value = {"mode": "chat"} + config = ProviderConfigManager.get_provider_responses_api_config( + model="github_copilot/some-chat-only-model", + provider=LlmProviders.GITHUB_COPILOT, + ) + assert config is None + + @patch( + "litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper" + ) + def test_returns_none_when_mode_is_unset_and_no_endpoints(self, mock_get_info): + """Entry without ``mode`` and without ``supported_endpoints`` returns None + (conservative default).""" + mock_get_info.return_value = {} + config = ProviderConfigManager.get_provider_responses_api_config( + model="github_copilot/some-model", + provider=LlmProviders.GITHUB_COPILOT, + ) + assert config is None + + def test_returns_config_when_mode_unset_but_endpoints_have_responses(self): + """``mode`` unset but ``supported_endpoints`` declaring /v1/responses + returns native config (endpoint-list fallback for stale-but-correct + catalog entries that lack ``mode``). + + Exercises the real ``_cached_get_model_info_helper`` plumbing via + ``register_model`` (no mock). ``supported_endpoints`` is not carried on + the normalized ``ModelInfoBase`` the helper returns, so the gate must + read it from the raw ``litellm.model_cost`` entry; a mock-based test + would mask that. + """ + litellm.register_model( + { + "github_copilot/test-endpoints-only-model": { + "litellm_provider": "github_copilot", + "max_tokens": 1, + "input_cost_per_token": 0, + "output_cost_per_token": 0, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses", + ], + } + } + ) + config = ProviderConfigManager.get_provider_responses_api_config( + model="github_copilot/test-endpoints-only-model", + provider=LlmProviders.GITHUB_COPILOT, + ) + assert isinstance(config, GithubCopilotResponsesAPIConfig) + + def test_mode_chat_overrides_endpoints_with_responses(self): + """``mode=chat`` is a hard opt-out: forces bridge even when + ``supported_endpoints`` includes /v1/responses. Lets users force the + bridge for dual-endpoint models without clearing endpoint metadata. + + Exercises the real ``_cached_get_model_info_helper`` plumbing via + ``register_model`` (no mock) so the ``mode``-over-endpoints precedence + is verified against the actual model-info resolution. + """ + litellm.register_model( + { + "github_copilot/test-chat-override-model": { + "litellm_provider": "github_copilot", + "max_tokens": 1, + "input_cost_per_token": 0, + "output_cost_per_token": 0, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses", + ], + } + } + ) + config = ProviderConfigManager.get_provider_responses_api_config( + model="github_copilot/test-chat-override-model", + provider=LlmProviders.GITHUB_COPILOT, + ) + assert config is None + + def test_returns_config_when_model_is_none(self): + """Follow-up GET/DELETE operations pass model=None and keep the native + config path (no per-model lookup is possible).""" + config = ProviderConfigManager.get_provider_responses_api_config( + model=None, + provider=LlmProviders.GITHUB_COPILOT, + ) + assert isinstance(config, GithubCopilotResponsesAPIConfig) + + @patch( + "litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper" + ) + def test_returns_none_when_get_model_info_raises(self, mock_get_info): + """Catalog lookup failure (model not registered) returns None + (conservative default; bridge handles unknown models safely).""" + mock_get_info.side_effect = Exception("model not in catalog") + config = ProviderConfigManager.get_provider_responses_api_config( + model="github_copilot/never-seen-model", + provider=LlmProviders.GITHUB_COPILOT, + ) + assert config is None + + @patch( + "litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper" + ) + def test_user_override_via_register_model(self, mock_get_info): + """User-supplied per-deployment ``model_info`` flows through + ``litellm.register_model`` (called by the router) into the merged + catalog read by ``_cached_get_model_info_helper``. Setting ``mode=responses`` + for a model whose catalog entry says ``mode=chat`` therefore opts in + to native dispatch without any per-call argument plumbing.""" + mock_get_info.return_value = {"mode": "responses"} + config = ProviderConfigManager.get_provider_responses_api_config( + model="github_copilot/some-chat-only-model", + provider=LlmProviders.GITHUB_COPILOT, + ) + assert isinstance(config, GithubCopilotResponsesAPIConfig) + + @patch( + "litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper" + ) + def test_realistic_chat_only_entry_returns_none(self, mock_get_info): + """Realistic ``model_prices_and_context_window.json`` shape for a + chat-only Copilot model (e.g. github_copilot/gemini-3.1-pro-preview) + returns None so /v1/responses calls fall back to the bridge.""" + mock_get_info.return_value = { + "litellm_provider": "github_copilot", + "max_input_tokens": 136000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "supported_endpoints": ["/v1/chat/completions"], + "supports_function_calling": True, + "supports_tool_choice": True, + "supports_parallel_function_calling": True, + "supports_vision": True, + "supports_reasoning": True, + } + config = ProviderConfigManager.get_provider_responses_api_config( + model="github_copilot/some-chat-only-model", + provider=LlmProviders.GITHUB_COPILOT, + ) + assert config is None + + @patch( + "litellm.llms.github_copilot.responses.transformation._cached_get_model_info_helper" + ) + def test_realistic_responses_only_entry_returns_config(self, mock_get_info): + """Realistic catalog entry for a Responses-only Copilot model + (e.g. github_copilot/gpt-5.5) returns the native config.""" + mock_get_info.return_value = { + "litellm_provider": "github_copilot", + "max_input_tokens": 272000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "responses", + "supported_endpoints": ["/v1/responses"], + "supports_function_calling": True, + "supports_tool_choice": True, + "supports_parallel_function_calling": True, + "supports_response_schema": True, + "supports_vision": True, + "supports_reasoning": True, + "supports_none_reasoning_effort": True, + "supports_xhigh_reasoning_effort": True, + } + config = ProviderConfigManager.get_provider_responses_api_config( + model="github_copilot/some-responses-only-model", + provider=LlmProviders.GITHUB_COPILOT, + ) + assert isinstance(config, GithubCopilotResponsesAPIConfig) diff --git a/tests/test_litellm/llms/parasail/test_parasail.py b/tests/test_litellm/llms/parasail/test_parasail.py new file mode 100644 index 00000000000..8fb9b22b5f6 --- /dev/null +++ b/tests/test_litellm/llms/parasail/test_parasail.py @@ -0,0 +1,172 @@ +import os +from unittest.mock import patch + +PARASAIL_API_BASE = "https://api.parasail.io/v1" +PARASAIL_RESPONSES_GATEWAY = "https://api-webflux.saas.parasail.io/v1" + + +def test_parasail_json_registry(): + import litellm + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + assert litellm.LlmProviders.PARASAIL.value == "parasail" + assert litellm.LlmProviders("parasail") == litellm.LlmProviders.PARASAIL + assert JSONProviderRegistry.exists("parasail") + config = JSONProviderRegistry.get("parasail") + assert config is not None + assert config.base_url == PARASAIL_API_BASE + assert config.api_key_env == "PARASAIL_API_KEY" + assert config.api_base_env == "PARASAIL_API_BASE" + assert "/v1/chat/completions" in config.supported_endpoints + assert "/v1/responses" in config.supported_endpoints + assert config.special_handling.get("force_store_false") is True + + +def test_parasail_listed_in_openai_compatible_providers(): + from litellm.constants import openai_compatible_providers + + assert "parasail" in openai_compatible_providers + + +def test_parasail_dynamic_config_env_vars(): + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + config = create_config_class(JSONProviderRegistry.get("parasail"))() + + with patch.dict( + os.environ, + { + "PARASAIL_API_KEY": "test-key", + "PARASAIL_API_BASE": PARASAIL_RESPONSES_GATEWAY, + }, + ): + api_base, api_key = config._get_openai_compatible_provider_info(None, None) + + assert api_base == PARASAIL_RESPONSES_GATEWAY + assert api_key == "test-key" + + +def test_parasail_provider_detection_by_prefix(): + from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider + + model, provider, _, api_base = get_llm_provider( + "parasail/parasail-llama-33-70b-fp8" + ) + + assert model == "parasail-llama-33-70b-fp8" + assert provider == "parasail" + assert api_base == PARASAIL_API_BASE + + +def test_parasail_chat_complete_url(): + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + config = create_config_class(JSONProviderRegistry.get("parasail"))() + + assert ( + config.get_complete_url( + api_base=None, + api_key=None, + model="parasail-llama-33-70b-fp8", + optional_params={}, + litellm_params={}, + ) + == f"{PARASAIL_API_BASE}/chat/completions" + ) + + +def test_parasail_responses_api_config(): + from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig + from litellm.utils import ProviderConfigManager + + config = ProviderConfigManager.get_provider_responses_api_config( + provider="parasail", + model="parasail-kimi-k25-elicit", + ) + + assert isinstance(config, OpenAIResponsesAPIConfig) + assert config.custom_llm_provider == "parasail" + assert ( + config.get_complete_url(api_base=None, litellm_params={}) + == f"{PARASAIL_API_BASE}/responses" + ) + + +def test_parasail_responses_api_honors_api_base_override(): + from litellm.utils import ProviderConfigManager + + config = ProviderConfigManager.get_provider_responses_api_config( + provider="parasail", + model="parasail-kimi-k25-elicit", + ) + + with patch.dict( + os.environ, + {"PARASAIL_API_BASE": PARASAIL_RESPONSES_GATEWAY}, + ): + url = config.get_complete_url(api_base=None, litellm_params={}) + + assert url == f"{PARASAIL_RESPONSES_GATEWAY}/responses" + + +def test_parasail_responses_api_forces_store_false_when_caller_sets_true(): + from litellm.types.router import GenericLiteLLMParams + from litellm.utils import ProviderConfigManager + + config = ProviderConfigManager.get_provider_responses_api_config( + provider="parasail", + model="parasail-kimi-k25-elicit", + ) + + request_params: dict = {"store": True, "temperature": 0.2} + transformed = config.transform_responses_api_request( + model="parasail-kimi-k25-elicit", + input="hello", + response_api_optional_request_params=request_params, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert transformed["store"] is False + assert transformed["temperature"] == 0.2 + + +def test_parasail_responses_api_forces_store_false_when_caller_omits_store(): + from litellm.types.router import GenericLiteLLMParams + from litellm.utils import ProviderConfigManager + + config = ProviderConfigManager.get_provider_responses_api_config( + provider="parasail", + model="parasail-kimi-k25-elicit", + ) + + transformed = config.transform_responses_api_request( + model="parasail-kimi-k25-elicit", + input="hello", + response_api_optional_request_params={}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert transformed["store"] is False + + +def test_parasail_responses_api_validate_environment_sets_bearer_token(): + from litellm.types.router import GenericLiteLLMParams + from litellm.utils import ProviderConfigManager + + config = ProviderConfigManager.get_provider_responses_api_config( + provider="parasail", + model="parasail-kimi-k25-elicit", + ) + + with patch.dict(os.environ, {"PARASAIL_API_KEY": "secret-from-env"}): + headers = config.validate_environment( + headers={}, + model="parasail-kimi-k25-elicit", + litellm_params=GenericLiteLLMParams(), + ) + + assert headers["Authorization"] == "Bearer secret-from-env" diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index fa6cc8bed1b..0236646c796 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -112,11 +112,71 @@ async def test_should_clear_stale_budget_reservation_when_budget_checks_skip(): user_api_key_cache=MagicMock(), proxy_logging_obj=MagicMock(), skip_budget_checks=True, + general_settings={}, ) assert user_api_key_auth_obj.budget_reservation is None +@pytest.mark.asyncio +async def test_disable_budget_reservation_skips_reservation(): + """#27639: general_settings.disable_budget_reservation turns off the optimistic Redis + reservation so operators hit by phantom BudgetExceededError can opt out of it.""" + user_api_key_auth_obj = UserAPIKeyAuth(token="test_token") + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request", + new=AsyncMock(return_value={"reserved_cost": 0.5, "entries": []}), + ) as mock_reserve: + await _reserve_budget_after_common_checks( + user_api_key_auth_obj=user_api_key_auth_obj, + request_data={"model": "gpt-4o"}, + route="/v1/chat/completions", + llm_router=None, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + skip_budget_checks=False, + general_settings={"disable_budget_reservation": True}, + ) + + mock_reserve.assert_not_called() + assert user_api_key_auth_obj.budget_reservation is None + + +@pytest.mark.asyncio +async def test_budget_reservation_runs_when_not_disabled(): + """Control for #27639: with the flag absent, the reservation still runs and is stored.""" + user_api_key_auth_obj = UserAPIKeyAuth(token="test_token") + reservation = { + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:test_token"}], + } + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.reserve_budget_for_request", + new=AsyncMock(return_value=reservation), + ) as mock_reserve: + await _reserve_budget_after_common_checks( + user_api_key_auth_obj=user_api_key_auth_obj, + request_data={"model": "gpt-4o"}, + route="/v1/chat/completions", + llm_router=None, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + skip_budget_checks=False, + general_settings={}, + ) + + mock_reserve.assert_awaited_once() + assert user_api_key_auth_obj.budget_reservation == reservation + + @pytest.mark.asyncio async def test_should_not_reuse_cached_key_object_for_request_state(): key_cache = DualCache() diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py index c58c94cbbc7..e7f72ff7a3f 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py @@ -41,7 +41,8 @@ def test_crowdstrike_aidr_guardrail_config() -> None: ) -def test_crowdstrike_aidr_guardrail_config_no_api_key() -> None: +def test_crowdstrike_aidr_guardrail_config_no_api_key(monkeypatch) -> None: + monkeypatch.delenv("CS_AIDR_TOKEN", raising=False) with pytest.raises(CrowdStrikeAIDRGuardrailMissingSecrets): init_guardrails_v2( all_guardrails=[ @@ -59,7 +60,8 @@ def test_crowdstrike_aidr_guardrail_config_no_api_key() -> None: ) -def test_crowdstrike_aidr_guardrail_config_no_api_base() -> None: +def test_crowdstrike_aidr_guardrail_config_no_api_base(monkeypatch) -> None: + monkeypatch.delenv("CS_AIDR_BASE_URL", raising=False) with pytest.raises(CrowdStrikeAIDRGuardrailMissingSecrets): init_guardrails_v2( all_guardrails=[ @@ -412,6 +414,121 @@ async def test_apply_guardrail_response_ok( assert result["texts"] == inputs["texts"] +@pytest.mark.asyncio +async def test_apply_guardrail_sends_user_id_model_and_extra_info( + crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler, +) -> None: + inputs: GenericGuardrailAPIInputs = { + "texts": ["Hello"], + "structured_messages": [{"role": "user", "content": "Hello"}], + "model": "gpt-4o", + } + request_data = { + "messages": inputs["structured_messages"], + "model": "gpt-4o", + "litellm_metadata": { + "user_api_key_user_id": "uid-abc", + "user_api_key_user_email": "alice@example.com", + }, + } + guardrail_endpoint = ( + f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions" + ) + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + return_value=httpx.Response( + status_code=200, + json={"result": {"blocked": False, "transformed": False}}, + request=httpx.Request(method="POST", url=guardrail_endpoint), + ), + ) as mock_method: + await crowdstrike_aidr_guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + ) + + payload = mock_method.call_args.kwargs["json"] + assert payload["user_id"] == "uid-abc" + assert payload["model"] == "gpt-4o" + assert payload["extra_info"] == {"user_name": "alice@example.com"} + + +@pytest.mark.asyncio +async def test_apply_guardrail_empty_extra_info_when_no_email( + crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler, +) -> None: + inputs: GenericGuardrailAPIInputs = { + "texts": ["Hello"], + "structured_messages": [{"role": "user", "content": "Hello"}], + "model": "gemini-flash", + } + request_data = { + "messages": inputs["structured_messages"], + "model": "gemini-flash", + "litellm_metadata": { + "user_api_key_user_id": "uid-no-email", + "user_api_key_user_email": None, + }, + } + guardrail_endpoint = ( + f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions" + ) + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + return_value=httpx.Response( + status_code=200, + json={"result": {"blocked": False, "transformed": False}}, + request=httpx.Request(method="POST", url=guardrail_endpoint), + ), + ) as mock_method: + await crowdstrike_aidr_guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + ) + + payload = mock_method.call_args.kwargs["json"] + assert payload["user_id"] == "uid-no-email" + assert payload["model"] == "gemini-flash" + assert payload["extra_info"] == {} + + +@pytest.mark.asyncio +async def test_apply_guardrail_no_metadata_skips_user_fields( + crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler, +) -> None: + inputs: GenericGuardrailAPIInputs = { + "texts": ["Hello"], + "structured_messages": [{"role": "user", "content": "Hello"}], + } + request_data = {"messages": inputs["structured_messages"]} + guardrail_endpoint = ( + f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions" + ) + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + return_value=httpx.Response( + status_code=200, + json={"result": {"blocked": False, "transformed": False}}, + request=httpx.Request(method="POST", url=guardrail_endpoint), + ), + ) as mock_method: + await crowdstrike_aidr_guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + ) + + payload = mock_method.call_args.kwargs["json"] + assert "user_id" not in payload + assert "model" not in payload + assert "extra_info" not in payload + + @pytest.mark.asyncio async def test_apply_guardrail_request_skipped_messages_stay_aligned( crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler, diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py index f8b6fbde3dc..1114b3df0c2 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_anthropic_passthrough_logging_handler.py @@ -321,6 +321,44 @@ class TestAzureAnthropicCostCalculation: assert call_kwargs["model"] == "azure_ai/claude-sonnet-4-5_gb_20250929" assert call_kwargs["custom_llm_provider"] == "azure_ai" + def test_passthrough_logging_sets_response_cost_with_server_tool_use_dict(self): + from litellm.types.utils import Choices, Message, ModelResponse + + logging_obj = self._create_mock_logging_obj(model="claude-3-7-sonnet-20250219") + logging_obj.get_router_model_id.return_value = None + logging_obj.litellm_params = {} + + response = ModelResponse( + id="test-id", + choices=[ + Choices( + finish_reason="stop", + index=0, + message=Message(content="test", role="assistant"), + ) + ], + created=1234567890, + model="claude-3-7-sonnet-20250219", + usage={ + "prompt_tokens": 10, + "completion_tokens": 5, + "total_tokens": 15, + "server_tool_use": {"web_search_requests": 1}, + }, + ) + + kwargs = AnthropicPassthroughLoggingHandler._create_anthropic_response_logging_payload( + litellm_model_response=response, + model="claude-3-7-sonnet-20250219", + kwargs={}, + start_time=datetime.now(), + end_time=datetime.now(), + logging_obj=logging_obj, + ) + + assert "response_cost" in kwargs + assert kwargs["response_cost"] > 0 + class TestAnthropicBatchPassthroughCostTracking: """Test cases for Anthropic batch passthrough cost tracking functionality""" diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py index bfcaaafd335..3c6af3e528a 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py @@ -257,6 +257,64 @@ class TestOpenAIPassthroughLoggingHandler: ) assert OpenAIPassthroughLoggingHandler.is_openai_responses_route("") == False + def test_is_openai_route_recognizes_cognitiveservices_azure_com(self): + """Azure OpenAI resources created via the newer "Azure AI Foundry" / + Cognitive Services pathway live on `*.cognitiveservices.azure.com` + subdomains rather than the older `openai.azure.com`. All four + is_openai_*_route methods must recognize both Azure subdomains so + cost tracking applies regardless of which Azure naming the user's + resource happens to be on. + """ + cognitive_chat = ( + "https://my-resource.cognitiveservices.azure.com/v1/chat/completions" + ) + cognitive_images_gen = ( + "https://my-resource.cognitiveservices.azure.com/v1/images/generations" + ) + cognitive_images_edit = ( + "https://my-resource.cognitiveservices.azure.com/v1/images/edits" + ) + cognitive_responses = ( + "https://my-resource.cognitiveservices.azure.com/v1/responses" + ) + + assert ( + OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route( + cognitive_chat + ) + is True + ) + assert ( + OpenAIPassthroughLoggingHandler.is_openai_image_generation_route( + cognitive_images_gen + ) + is True + ) + assert ( + OpenAIPassthroughLoggingHandler.is_openai_image_editing_route( + cognitive_images_edit + ) + is True + ) + assert ( + OpenAIPassthroughLoggingHandler.is_openai_responses_route( + cognitive_responses + ) + is True + ) + + # Cross-route negatives still hold for cognitiveservices hosts. + assert ( + OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route( + cognitive_responses + ) + is False + ) + assert ( + OpenAIPassthroughLoggingHandler.is_openai_responses_route(cognitive_chat) + is False + ) + @patch("litellm.completion_cost") @patch( "litellm.litellm_core_utils.litellm_logging.get_standard_logging_object_payload" @@ -766,6 +824,14 @@ class TestOpenAIPassthroughIntegration: == True ) assert self.handler.is_openai_route("https://api.openai.com/v1/models") == True + # Azure OpenAI on the shared Cognitive Services domain, identified by an + # OpenAI-style path segment. + assert ( + self.handler.is_openai_route( + "https://my-resource.cognitiveservices.azure.com/v1/chat/completions" + ) + == True + ) # Negative cases assert ( @@ -782,6 +848,28 @@ class TestOpenAIPassthroughIntegration: self.handler.is_openai_route("https://api.assemblyai.com/v2/transcript") == False ) + # Non-OpenAI Azure Cognitive Services share the `cognitiveservices.azure.com` + # domain but must NOT be classified as OpenAI routes (no OpenAI path segment). + assert ( + self.handler.is_openai_route( + "https://my-resource.cognitiveservices.azure.com/speechtotext/v3.1/recognize" + ) + == False + ) + assert ( + self.handler.is_openai_route( + "https://my-resource.cognitiveservices.azure.com/vision/v3.2/analyze" + ) + == False + ) + # A look-alike domain that merely contains an OpenAI host as a substring + # must be rejected by the suffix-based hostname match. + assert ( + self.handler.is_openai_route( + "https://cognitiveservices.azure.com.attacker.example/v1/chat/completions" + ) + == False + ) assert self.handler.is_openai_route("") == False @patch( diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py index 503a610e016..960fca205ce 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py @@ -949,6 +949,28 @@ class TestToolChoiceTransformation: result = LiteLLMCompletionResponsesConfig._transform_tool_choice(tool_choice) assert result == tool_choice + def test_transform_tool_choice_responses_flat_function_name(self): + """Responses-API forced-function with a top-level name maps to the nested Chat + Completions shape instead of degrading to required and dropping the name""" + result = LiteLLMCompletionResponsesConfig._transform_tool_choice( + {"type": "function", "name": "get_weather"} + ) + assert result == {"type": "function", "function": {"name": "get_weather"}} + + def test_transform_tool_choice_function_without_name_falls_back_to_required(self): + """A function-type dict with no name still falls back to required""" + result = LiteLLMCompletionResponsesConfig._transform_tool_choice( + {"type": "function"} + ) + assert result == "required" + + def test_transform_tool_choice_function_empty_name_falls_back_to_required(self): + """An empty top-level name is falsy and must not produce an empty function name""" + result = LiteLLMCompletionResponsesConfig._transform_tool_choice( + {"type": "function", "name": ""} + ) + assert result == "required" + class TestContentTypeTransformation: """Test content type transformation from Responses API to Chat Completion format""" diff --git a/tests/test_litellm/types/test_types_utils.py b/tests/test_litellm/types/test_types_utils.py index c146847f391..a4074ccdaaa 100644 --- a/tests/test_litellm/types/test_types_utils.py +++ b/tests/test_litellm/types/test_types_utils.py @@ -1,13 +1,9 @@ -import asyncio import os import sys -from typing import Optional -from unittest.mock import AsyncMock, patch import pytest sys.path.insert(0, os.path.abspath("../..")) -import json from litellm.types.utils import HiddenParams @@ -75,6 +71,48 @@ def test_usage_dump(): assert new_usage.prompt_tokens_details.web_search_requests == 1 +def test_usage_server_tool_use_dict_is_coerced_and_round_trips(): + from litellm.types.utils import ServerToolUse, Usage + + current_usage = Usage( + completion_tokens=1, + prompt_tokens=1, + total_tokens=2, + server_tool_use={"web_search_requests": 1}, + ) + + assert isinstance(current_usage.server_tool_use, ServerToolUse) + assert current_usage.server_tool_use.web_search_requests == 1 + + new_usage = Usage(**current_usage.model_dump()) + assert isinstance(new_usage.server_tool_use, ServerToolUse) + assert new_usage.server_tool_use.web_search_requests == 1 + + +def test_usage_converts_server_tool_use_dict(): + from litellm.types.utils import ServerToolUse, Usage + + usage = Usage( + completion_tokens=2, + prompt_tokens=1, + total_tokens=3, + server_tool_use={"web_search_requests": 4, "tool_search_requests": 1}, + ) + + assert isinstance(usage.server_tool_use, ServerToolUse) + assert usage.server_tool_use.web_search_requests == 4 + assert usage.server_tool_use["web_search_requests"] == 4 + assert usage.server_tool_use.tool_search_requests == 1 + with pytest.raises(KeyError): + usage.server_tool_use["unknown_metric"] + + round_trip = Usage(**usage.model_dump()) + assert isinstance(round_trip.server_tool_use, ServerToolUse) + assert round_trip.server_tool_use.web_search_requests == 4 + assert round_trip.server_tool_use["web_search_requests"] == 4 + assert round_trip.server_tool_use.tool_search_requests == 1 + + def test_usage_completion_tokens_details_text_tokens(): from litellm.types.utils import Usage diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 75a14e09852..d9992680c8a 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -22034,6 +22034,11 @@ export interface components { * @description connect to a postgres db - needed for generating temporary keys + tracking spend / key */ database_url?: string | null; + /** + * Disable Budget Reservation + * @description If True, disables the optimistic per-request budget reservation introduced in v1.84.0. WARNING: This weakens hard budget enforcement. Without the reservation, a burst of concurrent requests from a single key can each pass the read-time spend check before any of them is charged, allowing a configured budget to be exceeded under high concurrency. Budgets are still evaluated on every request at read time, so an already-exhausted budget is still rejected. Enable only if your deployment is experiencing phantom BudgetExceededError responses caused by leaked reservations (see GitHub issue #27639). A proxy-level WARNING is logged on every request while this flag is active as a reminder that hard enforcement is relaxed. + */ + disable_budget_reservation?: boolean | null; /** * Enable Public Model Hub * @description Public model hub for users to see what models they have access to, supported openai params, etc. From dfb68a23de435afc4c2503c1caf552f1a116f154 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 9 Jun 2026 02:27:03 +0530 Subject: [PATCH 09/22] feat(galileo): add health check support for UI callback test (#29908) * feat(galileo): add health check support for UI callback test Register galileo in /health/services so the proxy UI callback connection test works. Co-authored-by: Cursor * feat(galileo): verify API key via /current_user health check Call Galileo's current_user endpoint so the UI callback test validates credentials against the provider. Co-authored-by: Cursor * chore(ui): regenerate schema.d.ts for galileo health service Co-authored-by: Cursor * fix(galileo): return IntegrationHealthCheckStatus from async_health_check Fixes mypy assignment error in health_services_endpoint where response was narrowed to IntegrationHealthCheckStatus from earlier branches. Co-authored-by: Cursor * Fix Galileo logging to match Langfuse across all endpoint types. Stop skipping ingest when output is empty and log embeddings with a placeholder so embedding, speech, and other non-text responses are recorded like Langfuse. Co-authored-by: Cursor * fix(galileo): remove unreachable health-check guard and None output sentinel The use_v2_api flag is derived from bool(api_key), so the inner GALILEO_API_KEY check inside the v2 branch could never run; collapse the credential validation into the username/password path with a combined message. _serialize_galileo_output now returns an empty string for None, so _get_galileo_input_output_content always yields a str and the post-call None coalescing guard is no longer needed. * test(galileo): cover async_health_check failure paths and empty model response Add regression tests for the Galileo health check unhealthy branches (missing project id, missing base url, missing credentials, auth failure, and request exception) and for logging a model response with no choices, which now queues an empty output instead of being skipped. --------- Co-authored-by: Cursor Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/integrations/galileo.py | 71 ++++-- .../health_endpoints/_health_endpoints.py | 15 ++ .../test_litellm/integrations/test_galileo.py | 206 +++++++++++++++++- .../health_endpoints/test_health_endpoints.py | 27 +++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 +- 5 files changed, 303 insertions(+), 18 deletions(-) diff --git a/litellm/integrations/galileo.py b/litellm/integrations/galileo.py index 8fef90c24e0..f9ff7e8c7a1 100644 --- a/litellm/integrations/galileo.py +++ b/litellm/integrations/galileo.py @@ -26,6 +26,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) +from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus GALILEO_CLOUD_API_BASE_URL = "https://api.galileo.ai" # Cap the in-memory buffer so persistent flush failures (e.g. Galileo @@ -89,6 +90,52 @@ class GalileoObserve(CustomLogger): return bool(self.api_key) return bool(self.username and self.password) + async def async_health_check(self) -> IntegrationHealthCheckStatus: + try: + if not self.project_id: + return IntegrationHealthCheckStatus( + status="unhealthy", + error_message="GALILEO_PROJECT_ID environment variable not set", + ) + + if not self.base_url: + return IntegrationHealthCheckStatus( + status="unhealthy", + error_message="GALILEO_BASE_URL environment variable not set", + ) + + if not self.use_v2_api and (not self.username or not self.password): + return IntegrationHealthCheckStatus( + status="unhealthy", + error_message=( + "GALILEO_API_KEY or GALILEO_USERNAME and GALILEO_PASSWORD " + "environment variables must be set" + ), + ) + + if not await self._ensure_headers(): + return IntegrationHealthCheckStatus( + status="unhealthy", + error_message="Galileo authentication failed", + ) + + response = await self.async_httpx_handler.get( + url=f"{self.base_url}/current_user", + headers=self.headers, + ) + if response.status_code >= 400: + return IntegrationHealthCheckStatus( + status="unhealthy", + error_message=(f"Galileo API returned HTTP {response.status_code}"), + ) + + return IntegrationHealthCheckStatus(status="healthy", error_message=None) + except Exception as e: + return IntegrationHealthCheckStatus( + status="unhealthy", + error_message=f"Galileo health check failed: {str(e)}", + ) + async def async_set_galileo_headers(self) -> None: galileo_login_response = await self.async_httpx_handler.post( url=f"{self.base_url}/login", @@ -399,9 +446,9 @@ class GalileoObserve(CustomLogger): return prompt @staticmethod - def _serialize_galileo_output(value: Any) -> Optional[str]: + def _serialize_galileo_output(value: Any) -> str: if value is None: - return None + return "" if isinstance(value, str): return value @@ -460,11 +507,11 @@ class GalileoObserve(CustomLogger): response_obj: Any, level: str = "DEFAULT", status_message: Optional[str] = None, - ) -> Tuple[str, Optional[str], Any]: + ) -> Tuple[str, str, Any]: """ Mirror Langfuse _get_langfuse_input_output_content for Galileo ingest. - Returns (input_text, output_text, messages_for_span). output_text None skips ingest. + Returns (input_text, output_text, messages_for_span). """ call_type = kwargs.get("call_type") prompt = self._build_prompt(kwargs) @@ -477,10 +524,11 @@ class GalileoObserve(CustomLogger): return self._prompt_to_input_text(prompt), status_message, prompt if response_obj is not None and ( - call_type == "embedding" + call_type in ("embedding", "aembedding") or isinstance(response_obj, litellm.EmbeddingResponse) ): - return self._prompt_to_input_text(prompt), None, prompt + # Match Langfuse OTEL: log embeddings without serializing vectors. + return self._prompt_to_input_text(prompt), "embedding-output", prompt if response_obj is not None and isinstance(response_obj, litellm.ModelResponse): output = self._get_chat_content_for_galileo(response_obj) @@ -549,7 +597,7 @@ class GalileoObserve(CustomLogger): ): input_val = kwargs.get("input") return ( - self._serialize_galileo_output(input_val) or "", + self._serialize_galileo_output(input_val), self._serialize_galileo_output(response_obj), input_val, ) @@ -574,11 +622,11 @@ class GalileoObserve(CustomLogger): kwargs.get("messages") or [], ) - return self._prompt_to_input_text(prompt), None, kwargs.get("messages") or [] + return self._prompt_to_input_text(prompt), "", kwargs.get("messages") or [] def get_output_str_from_response( self, response_obj: Any, kwargs: Dict[str, Any] - ) -> Optional[str]: + ) -> str: _, output_text, _ = self._get_galileo_input_output_content( kwargs=kwargs, response_obj=response_obj ) @@ -659,11 +707,6 @@ class GalileoObserve(CustomLogger): input_text, output_text, messages = self._get_galileo_input_output_content( kwargs=kwargs, response_obj=response_obj ) - if output_text is None: - verbose_logger.debug( - "Galileo Logger: skipping %s — no text output to log", _call_type - ) - return raw_start = slo.get("startTime") raw_end = slo.get("endTime") diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index c109f374993..6ef8bbc4006 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -129,6 +129,7 @@ services = Union[ "datadog_llm_observability", "generic_api", "arize", + "galileo", "sqs", ], str, @@ -206,6 +207,7 @@ async def health_services_endpoint( # noqa: PLR0915 "datadog_llm_observability", "generic_api", "arize", + "galileo", "sqs", ]: raise HTTPException( @@ -295,6 +297,19 @@ async def health_services_endpoint( # noqa: PLR0915 else "Arize is healthy" ), } + elif service == "galileo": + from litellm.integrations.galileo import GalileoObserve + + galileo_logger = GalileoObserve() + response = await galileo_logger.async_health_check() + return { + "status": response["status"], + "message": ( + response["error_message"] + if response["status"] == "unhealthy" + else "Galileo is healthy" + ), + } elif service == "langfuse": from litellm.integrations.langfuse.langfuse import LangFuseLogger diff --git a/tests/test_litellm/integrations/test_galileo.py b/tests/test_litellm/integrations/test_galileo.py index 8ce5eb776f9..0533b7ca7d1 100644 --- a/tests/test_litellm/integrations/test_galileo.py +++ b/tests/test_litellm/integrations/test_galileo.py @@ -357,12 +357,18 @@ def test_galileo_record_to_v2_span_with_tags_and_offset(): def test_galileo_get_output_str_variants(galileo_v2_env): logger = GalileoObserve() - assert logger.get_output_str_from_response(None, {}) is None + assert logger.get_output_str_from_response(None, {}) == "" assert ( logger.get_output_str_from_response( EmbeddingResponse(), {"call_type": "embedding"} ) - is None + == "embedding-output" + ) + assert ( + logger.get_output_str_from_response( + EmbeddingResponse(), {"call_type": "aembedding"} + ) + == "embedding-output" ) text_resp = TextCompletionResponse() @@ -414,7 +420,7 @@ def test_galileo_get_output_str_variants(galileo_v2_env): {"call_type": "acompletion", "messages": [{"role": "user", "content": "hi"}]}, ) - assert logger.get_output_str_from_response("not-a-supported-type", {}) is None + assert logger.get_output_str_from_response("not-a-supported-type", {}) == "" def test_galileo_get_input_output_error_status_message(galileo_v2_env): @@ -445,6 +451,48 @@ def test_galileo_get_output_str_rerank_response(galileo_v2_env): assert '"relevance_score": 0.98' in output +@pytest.mark.asyncio +async def test_galileo_async_log_success_embedding(galileo_v2_env): + import datetime + + logger = GalileoObserve() + embedding_response = EmbeddingResponse( + data=[{"object": "embedding", "embedding": [0.1, 0.2, 0.3], "index": 0}] + ) + + mock_response = MagicMock() + mock_response.is_success = True + mock_response.status_code = 201 + + with patch.object(logger.async_httpx_handler, "post", return_value=mock_response): + await logger.async_log_success_event( + kwargs={ + "call_type": "aembedding", + "model": "text-embedding-3-small", + "input": "hello world", + "standard_logging_object": { + "call_type": "aembedding", + "model": "text-embedding-3-small", + "prompt_tokens": 2, + "completion_tokens": 0, + "total_tokens": 2, + "response_cost": 0.0, + "startTime": datetime.datetime( + 2026, 5, 25, 12, 0, 0, tzinfo=datetime.timezone.utc + ).timestamp(), + "endTime": datetime.datetime( + 2026, 5, 25, 12, 0, 1, tzinfo=datetime.timezone.utc + ).timestamp(), + }, + }, + response_obj=embedding_response, + start_time=datetime.datetime(2026, 5, 25, 12, 0, 0), + end_time=datetime.datetime(2026, 5, 25, 12, 0, 1), + ) + + assert logger.in_memory_records == [] + + @pytest.mark.asyncio async def test_galileo_async_log_success_rerank(galileo_v2_env): import datetime @@ -524,6 +572,158 @@ def test_galileo_get_ingest_request_legacy(monkeypatch): assert payload["traces"][0]["input"] == "hi" +@pytest.mark.asyncio +async def test_galileo_async_health_check_success(galileo_v2_env): + logger = GalileoObserve() + current_user_resp = MagicMock() + current_user_resp.status_code = 200 + + with patch.object( + logger.async_httpx_handler, "get", new_callable=AsyncMock + ) as mock_get: + mock_get.return_value = current_user_resp + result = await logger.async_health_check() + + assert result["status"] == "healthy" + mock_get.assert_awaited_once_with( + url="https://api.galileo.ai/current_user", + headers={ + "accept": "application/json", + "Content-Type": "application/json", + "Galileo-API-Key": "test-api-key", + }, + ) + + +@pytest.mark.asyncio +async def test_galileo_async_health_check_api_error(galileo_v2_env): + logger = GalileoObserve() + current_user_resp = MagicMock() + current_user_resp.status_code = 401 + + with patch.object( + logger.async_httpx_handler, "get", new_callable=AsyncMock + ) as mock_get: + mock_get.return_value = current_user_resp + result = await logger.async_health_check() + + assert result["status"] == "unhealthy" + assert "HTTP 401" in result["error_message"] + + +@pytest.mark.asyncio +async def test_galileo_async_health_check_missing_project_id(monkeypatch): + monkeypatch.setenv("GALILEO_API_KEY", "test-api-key") + monkeypatch.setenv("GALILEO_BASE_URL", "https://api.galileo.ai") + monkeypatch.delenv("GALILEO_PROJECT_ID", raising=False) + logger = GalileoObserve() + + result = await logger.async_health_check() + + assert result["status"] == "unhealthy" + assert "GALILEO_PROJECT_ID" in result["error_message"] + + +@pytest.mark.asyncio +async def test_galileo_async_health_check_missing_base_url(monkeypatch): + monkeypatch.delenv("GALILEO_API_KEY", raising=False) + monkeypatch.delenv("GALILEO_BASE_URL", raising=False) + monkeypatch.setenv("GALILEO_PROJECT_ID", "p") + monkeypatch.setenv("GALILEO_USERNAME", "u") + monkeypatch.setenv("GALILEO_PASSWORD", "pw") + logger = GalileoObserve() + + result = await logger.async_health_check() + + assert result["status"] == "unhealthy" + assert "GALILEO_BASE_URL" in result["error_message"] + + +@pytest.mark.asyncio +async def test_galileo_async_health_check_missing_credentials(monkeypatch): + monkeypatch.delenv("GALILEO_API_KEY", raising=False) + monkeypatch.delenv("GALILEO_USERNAME", raising=False) + monkeypatch.delenv("GALILEO_PASSWORD", raising=False) + monkeypatch.setenv("GALILEO_PROJECT_ID", "p") + monkeypatch.setenv("GALILEO_BASE_URL", "https://galileo.example") + logger = GalileoObserve() + + result = await logger.async_health_check() + + assert result["status"] == "unhealthy" + assert "GALILEO_USERNAME" in result["error_message"] + + +@pytest.mark.asyncio +async def test_galileo_async_health_check_auth_failed(monkeypatch): + monkeypatch.delenv("GALILEO_API_KEY", raising=False) + monkeypatch.setenv("GALILEO_PROJECT_ID", "p") + monkeypatch.setenv("GALILEO_BASE_URL", "https://galileo.example") + monkeypatch.setenv("GALILEO_USERNAME", "u") + monkeypatch.setenv("GALILEO_PASSWORD", "pw") + logger = GalileoObserve() + + with patch.object( + logger.async_httpx_handler, "post", new_callable=AsyncMock + ) as mock_post: + mock_post.side_effect = Exception("login failed") + result = await logger.async_health_check() + + assert result["status"] == "unhealthy" + assert result["error_message"] == "Galileo authentication failed" + + +@pytest.mark.asyncio +async def test_galileo_async_health_check_request_exception(galileo_v2_env): + logger = GalileoObserve() + + with patch.object( + logger.async_httpx_handler, "get", new_callable=AsyncMock + ) as mock_get: + mock_get.side_effect = Exception("connection refused") + result = await logger.async_health_check() + + assert result["status"] == "unhealthy" + assert "connection refused" in result["error_message"] + + +@pytest.mark.asyncio +async def test_galileo_async_log_success_empty_model_response(galileo_v2_env): + import datetime + + logger = GalileoObserve() + logger.batch_size = 2 + empty_response = ModelResponse(choices=[]) + + await logger.async_log_success_event( + kwargs={ + "call_type": "acompletion", + "model": "gpt-5.2", + "messages": [{"role": "user", "content": "hi"}], + "standard_logging_object": { + "call_type": "acompletion", + "model": "gpt-5.2", + "prompt_tokens": 1, + "completion_tokens": 0, + "total_tokens": 1, + "response_cost": 0.0, + "startTime": datetime.datetime( + 2026, 5, 25, 12, 0, 0, tzinfo=datetime.timezone.utc + ).timestamp(), + "endTime": datetime.datetime( + 2026, 5, 25, 12, 0, 1, tzinfo=datetime.timezone.utc + ).timestamp(), + }, + }, + response_obj=empty_response, + start_time=datetime.datetime(2026, 5, 25, 12, 0, 0), + end_time=datetime.datetime(2026, 5, 25, 12, 0, 1), + ) + + assert len(logger.in_memory_records) == 1 + assert logger.in_memory_records[0]["output_text"] == "" + + @pytest.mark.asyncio async def test_galileo_ensure_headers_v2_missing_key(monkeypatch): monkeypatch.delenv("GALILEO_API_KEY", raising=False) diff --git a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py index 80a4804956c..64c57ab90e3 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -696,6 +696,33 @@ async def test_test_model_connection_falls_back_to_deployments_zero_without_id() assert model_params.get("api_key") == "fake-key-A" +@pytest.mark.asyncio +@pytest.mark.parametrize( + "status,error_message", + [ + ("healthy", ""), + ("unhealthy", "Galileo authentication failed"), + ], +) +async def test_health_services_endpoint_galileo(status, error_message): + with patch("litellm.integrations.galileo.GalileoObserve") as MockGalileoObserve: + mock_instance = MagicMock() + mock_instance.async_health_check = AsyncMock( + return_value={"status": status, "error_message": error_message} + ) + MockGalileoObserve.return_value = mock_instance + + result = await health_services_endpoint(service="galileo") + + if status == "healthy": + assert result["status"] == "healthy" + assert result["message"] == "Galileo is healthy" + else: + assert result["status"] == "unhealthy" + assert result["message"] == error_message + mock_instance.async_health_check.assert_awaited_once() + + @pytest.mark.asyncio async def test_health_services_endpoint_datadog_llm_observability(): """ diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index d9992680c8a..8379d0536a6 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -40422,7 +40422,7 @@ export interface operations { parameters: { query: { /** @description Specify the service being hit. */ - service: ("slack_budget_alerts" | "langfuse" | "langfuse_otel" | "slack" | "openmeter" | "webhook" | "email" | "braintrust" | "datadog" | "datadog_llm_observability" | "generic_api" | "arize" | "sqs") | string; + service: ("slack_budget_alerts" | "langfuse" | "langfuse_otel" | "slack" | "openmeter" | "webhook" | "email" | "braintrust" | "datadog" | "datadog_llm_observability" | "generic_api" | "arize" | "galileo" | "sqs") | string; }; header?: never; path?: never; From 69a7bdb24728eb13280bba5e4d7c3427e8ac0768 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Mon, 8 Jun 2026 14:28:39 -0700 Subject: [PATCH 10/22] fix(model-management): allow deleting a BYOK model after its team is deleted (#29875) * fix(model-management): allow deleting a BYOK model after its team is deleted A team BYOK model (model_info.team_id set) became undeletable once its team was deleted: POST /model/delete ran can_user_make_model_call, which looked the team up and raised 400 "Team id=... does not exist in db" before the delete could run, so the model lingered on the Models + Endpoints page with no way to remove it. Drop the team-existence prerequisite from the delete path. When the model's team still exists the normal auth check runs unchanged; when it is gone a proxy admin may delete the orphan and any other caller gets a 403. The check is fail-closed, so a missing or errored team lookup can only block the delete or require an admin, never grant a non-admin access. Add/update/health keep their team-existence validation. * refactor(model-management): drop redundant team lookup on model delete Move the orphaned-team handling into can_user_make_model_call behind an allow_missing_team flag instead of pre-checking team existence in delete_model. The endpoint no longer issues its own litellm_teamtable lookup, so deleting a model whose team still exists hits the team table once instead of twice. The auth behavior is unchanged: a proxy admin can delete a model whose team was deleted, any other caller gets a 403, and add/update/health keep the strict "team must exist" validation. --- .../model_management_endpoints.py | 14 ++ .../test_model_management_endpoints.py | 184 ++++++++++++++++++ 2 files changed, 198 insertions(+) diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 566ef845333..0cbccfc18ad 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -860,6 +860,7 @@ class ModelManagementAuthChecks: user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, premium_user: bool, + allow_missing_team: bool = False, ) -> Literal[True]: ## Check team model auth if ( @@ -870,6 +871,18 @@ class ModelManagementAuthChecks: where={"team_id": model_params.model_info.team_id} ) if team_obj_row is None: + # The team was deleted. Callers that opt in (e.g. model deletion) may + # act on the orphaned model, but only as a proxy admin -- without the + # team there is no team-admin membership left to verify. + if allow_missing_team: + if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN: + return True + raise HTTPException( + status_code=403, + detail={ + "error": "Only a proxy admin can delete a model whose team has been deleted." + }, + ) raise HTTPException( status_code=400, detail={ @@ -955,6 +968,7 @@ async def delete_model( user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, premium_user=premium_user, + allow_missing_team=True, ) # update DB diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index a9bb2b09a15..2ba604e5da6 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -1915,6 +1915,190 @@ class TestDeleteTeamBYOKModelGhost: mock_refresh.assert_not_awaited() +class TestDeleteModelTeamAuth: + """Team auth on the /model/delete path. + + A model added via /model/new with model_info.team_id is orphaned once its + team is deleted: can_user_make_model_call looked the team up and raised + 'Team id=... does not exist in db' before the delete could run, so the model + was undeletable from the Models + Endpoints page. Without the team, team-admin + membership can't be verified, so a proxy admin (and only a proxy admin) may + delete the orphan; a missing team must never let a non-admin through. The team + is also looked up exactly once -- the auth check must not add a second query. + """ + + def _orphaned_model_mocks(self, team_id, model_id): + db_row = LiteLLM_ProxyModelTable( + model_id=model_id, + model_name=f"model_name_{team_id}_abc-uuid", + litellm_params={"model": "openai/gpt-4.1-nano"}, + model_info={ + "id": model_id, + "team_id": team_id, + "team_public_model_name": "orphaned-gpt", + }, + created_by="admin", + updated_by="admin", + ) + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.litellm_proxymodeltable = AsyncMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( + return_value=db_row + ) + mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + # The team is gone -> every team lookup returns None. + mock_prisma.db.litellm_teamtable = AsyncMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) + mock_prisma.db.litellm_teamtable.update = AsyncMock() + mock_prisma.db.litellm_modeltable = AsyncMock() + mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=[]) + return mock_prisma + + @pytest.mark.asyncio + async def test_proxy_admin_can_delete_model_when_team_deleted(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + ModelInfoDelete, + delete_model as delete_model_endpoint, + ) + + team_id = "deleted-team-xyz" + model_id = "orphaned-byok-1" + mock_prisma = self._orphaned_model_mocks(team_id, model_id) + + admin_user = UserAPIKeyAuth( + user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN + ) + + _PS = "litellm.proxy.proxy_server" + _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" + with ( + patch(f"{_PS}.prisma_client", mock_prisma), + patch(f"{_PS}.store_model_in_db", True), + patch(f"{_PS}.premium_user", True), + patch(f"{_PS}.llm_router", MagicMock()), + patch(f"{_PS}.proxy_logging_obj", MagicMock()), + patch(f"{_PS}.user_api_key_cache", MagicMock()), + patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), + ): + result = await delete_model_endpoint( + model_info=ModelInfoDelete(id=model_id), + user_api_key_dict=admin_user, + ) + + assert "deleted successfully" in result["message"] + mock_prisma.db.litellm_proxymodeltable.delete.assert_awaited_once() + # Team is gone -> no team.models cleanup to do. + mock_prisma.db.litellm_teamtable.update.assert_not_awaited() + + @pytest.mark.asyncio + async def test_non_admin_cannot_delete_model_when_team_deleted(self): + """A missing team must never let a non-admin delete the orphan (no fail-open).""" + from litellm.proxy.management_endpoints.model_management_endpoints import ( + ModelInfoDelete, + delete_model as delete_model_endpoint, + ) + from litellm.proxy.proxy_server import ProxyException + + team_id = "deleted-team-abc" + model_id = "orphaned-byok-2" + mock_prisma = self._orphaned_model_mocks(team_id, model_id) + + non_admin = UserAPIKeyAuth( + user_id="someone", user_role=LitellmUserRoles.INTERNAL_USER + ) + + _PS = "litellm.proxy.proxy_server" + _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" + with ( + patch(f"{_PS}.prisma_client", mock_prisma), + patch(f"{_PS}.store_model_in_db", True), + patch(f"{_PS}.premium_user", True), + patch(f"{_PS}.llm_router", MagicMock()), + patch(f"{_PS}.proxy_logging_obj", MagicMock()), + patch(f"{_PS}.user_api_key_cache", MagicMock()), + patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), + ): + with pytest.raises(ProxyException) as exc_info: + await delete_model_endpoint( + model_info=ModelInfoDelete(id=model_id), + user_api_key_dict=non_admin, + ) + + assert str(exc_info.value.code) == "403" + mock_prisma.db.litellm_proxymodeltable.delete.assert_not_awaited() + + @pytest.mark.asyncio + async def test_live_team_delete_looks_up_team_once(self): + """The auth check must not add a redundant team query on the live-team path.""" + from litellm.proxy.management_endpoints.model_management_endpoints import ( + ModelInfoDelete, + delete_model as delete_model_endpoint, + ) + from litellm.proxy.proxy_server import ProxyException + + team_id = "live-team-1" + model_id = "live-byok-1" + db_row = LiteLLM_ProxyModelTable( + model_id=model_id, + model_name=f"model_name_{team_id}_abc-uuid", + litellm_params={"model": "openai/gpt-4.1-nano"}, + model_info={ + "id": model_id, + "team_id": team_id, + "team_public_model_name": "live-gpt", + }, + created_by="admin", + updated_by="admin", + ) + team_row = LiteLLM_TeamTable( + team_id=team_id, + team_alias="live-team", + members_with_roles=[Member(user_id="admin", role="admin")], + models=["live-gpt"], + ) + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.litellm_proxymodeltable = AsyncMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( + return_value=db_row + ) + mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_teamtable = AsyncMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + mock_prisma.db.litellm_modeltable = AsyncMock() + mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=[]) + + # A team member who is not the team admin: rejected before the delete runs, + # so the only team lookup is the single one inside the auth check. + non_admin = UserAPIKeyAuth( + user_id="someone", user_role=LitellmUserRoles.INTERNAL_USER + ) + + _PS = "litellm.proxy.proxy_server" + _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" + with ( + patch(f"{_PS}.prisma_client", mock_prisma), + patch(f"{_PS}.store_model_in_db", True), + patch(f"{_PS}.premium_user", True), + patch(f"{_PS}.llm_router", MagicMock()), + patch(f"{_PS}.proxy_logging_obj", MagicMock()), + patch(f"{_PS}.user_api_key_cache", MagicMock()), + patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), + ): + with pytest.raises(ProxyException) as exc_info: + await delete_model_endpoint( + model_info=ModelInfoDelete(id=model_id), + user_api_key_dict=non_admin, + ) + + assert str(exc_info.value.code) == "403" + assert mock_prisma.db.litellm_teamtable.find_unique.await_count == 1 + mock_prisma.db.litellm_proxymodeltable.delete.assert_not_awaited() + + class TestGetTeamDeployments: """Tests for _get_team_deployments which filters by model_name prefix + Python-side team_id check.""" From a7ecf6b5b1452410cb83f58054d939561ef322a6 Mon Sep 17 00:00:00 2001 From: milan-berri Date: Tue, 9 Jun 2026 01:09:03 +0300 Subject: [PATCH 11/22] feat(jwt-auth): opt-in fallback to DB team on unresolved JWT claim (#28913) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(jwt-auth): defer to single-team DB fallback on claim mismatch Extends the single-team DB fallback introduced in #26418 to two more cases where it previously could not run: * `find_and_validate_specific_team_id`: when `team_id_jwt_field` is configured and a claim value is present in the token but the team does not exist in the LiteLLM DB (HTTPException 404 from `get_team_object`), return `(None, None)` instead of raising — the auth_builder fallback then attributes the request to the user's single DB team. Only HTTPException is caught; other errors (e.g. "No DB Connected") still propagate. * `find_team_with_model_access`: when none of the `team_ids_jwt_field` groups resolve to a real LiteLLM team, return `(None, None)` instead of raising 403 so the same fallback path runs. If at least one group DID resolve to a team but none granted the requested model, the original 403 is preserved (legitimate access denial — not a claim mismatch). Tracked via the new `any_claim_team_resolved` flag. The strict `is_required_team_id` raise and `enforce_team_based_model_access` raise remain unchanged. Unit tests cover both new soft-fail paths and guard each preserved path (strict required, enforce_team_based, the preserved 403, and the non-HTTPException propagation). Co-authored-by: Cursor * fix(jwt-auth): narrow HTTPException catch to 404 (greptile review) Address Greptile review comments on #28913: * `find_and_validate_specific_team_id`: re-raise HTTPException when `status_code != 404`, pinning the catch to the "team doesn't exist in db" path documented for `get_team_object`. A future change that introduces a different status code (e.g. 403 for a blocked team) will now propagate instead of silently falling through to the single-team DB fallback. * Add `test_find_and_validate_specific_team_id_non_404_http_exception_propagates` parametrised over 400 / 403 / 500 to lock in the contract. Co-authored-by: Cursor * fix(jwt-auth): gate claim-mismatch fallback behind opt-in flag The unresolved-team-claim fallback added in the previous commit weakened the strict claim-based authorization contract by default — an authenticated user whose JWT carries a stale or invalid team claim could still consume their single DB team's models/quota via the fallback. Gate both soft-fail paths in `find_and_validate_specific_team_id` and `find_team_with_model_access` behind a new opt-in flag `team_claim_fallback` on `LiteLLM_JWTAuth` (default False). Default-off preserves the pre-existing strict behavior. Operators who intentionally treat IdP team claims as advisory (e.g. machine tokens whose group claims live in a separate namespace from LiteLLM team_ids) opt in via config. Adds two regression tests guarding the default-off behavior. Co-authored-by: Cursor --------- Co-authored-by: Cursor --- litellm/proxy/_types.py | 10 + litellm/proxy/auth/handle_jwt.py | 43 ++- .../proxy/auth/test_handle_jwt.py | 304 ++++++++++++++++++ 3 files changed, 347 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index fba968dd95d..f1443edf455 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4195,6 +4195,16 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): default=None, description="Optional claim-based routing overrides for JWT-shaped tokens. Matching rules route requests to oauth2 before default JWT flow.", ) + team_claim_fallback: bool = Field( + default=False, + description=( + "If True, when a configured team_id_jwt_field / team_ids_jwt_field " + "claim is present but does not resolve to any known team, defer to " + "the single-team DB fallback (caller's only team membership) " + "instead of raising. Default False preserves strict claim-based " + "authorization." + ), + ) issuers: Optional[List[JWTIssuerConfig]] = Field( default=None, description="Optional issuer-bound JWT validation rules. When a token's `iss` matches a configured issuer, validation uses that issuer's JWKS, audience, and claim mappings. Tokens with an unlisted `iss` fall back to the global JWT_AUDIENCE/JWT_ISSUER validation path — this is additive routing, not an allow-list.", diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 536d2867855..fd6ff2ada7f 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -1299,15 +1299,29 @@ class JWTAuthManager: # First try to get team by team_id if individual_team_id: - team_object = await get_team_object( - team_id=individual_team_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert, - ) - return individual_team_id, team_object + try: + team_object = await get_team_object( + team_id=individual_team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + team_id_upsert=jwt_handler.litellm_jwtauth.team_id_upsert, + ) + return individual_team_id, team_object + except HTTPException as e: + if ( + e.status_code != 404 + or not jwt_handler.litellm_jwtauth.team_claim_fallback + ): + raise + # Claim doesn't map to a known team — defer to fallback. + verbose_proxy_logger.debug( + "JWT team_id claim '%s' did not resolve to a team: %s", + individual_team_id, + e.detail, + ) + return None, None # If no team_id found, try to resolve via team_alias_jwt_field team_alias = jwt_handler.get_team_alias( @@ -1431,6 +1445,7 @@ class JWTAuthManager: ) return None, None + any_claim_team_resolved = False for team_id in team_ids: try: team_object = await get_team_object( @@ -1441,6 +1456,9 @@ class JWTAuthManager: proxy_logging_obj=proxy_logging_obj, ) + if team_object is not None: + any_claim_team_resolved = True + if team_object and team_object.models is not None: team_models = team_object.models if isinstance(team_models, list) and ( @@ -1478,12 +1496,17 @@ class JWTAuthManager: if denied_auth_enforced_pass_through_route: JWTAuthManager._raise_team_passthrough_route_denial(route=route) - if requested_model: + if requested_model and ( + any_claim_team_resolved + or not jwt_handler.litellm_jwtauth.team_claim_fallback + ): + # Claim resolved but no model access, or fallback disabled — deny. raise HTTPException( status_code=403, detail=f"No team has access to the requested model: {requested_model}. Checked teams={team_ids}. Check `/models` to see all available models.", ) + # No claim team resolved and fallback enabled — defer to fallback. return None, None @staticmethod diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index 7dadca24504..63510086f95 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -3182,6 +3182,310 @@ def test_build_decode_kwargs_no_warning_when_scoped( assert matching == [] +# --------------------------------------------------------------------------- +# Defer to single-team DB fallback (PR #26418) when JWT claims are present +# but do not resolve to a LiteLLM team. +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_find_and_validate_specific_team_id_unresolved_claim_returns_none(): + """With `team_claim_fallback=True`: team_id claim is present in the JWT + but the team is missing in the DB — return (None, None) so the + auth_builder single-team fallback can run, instead of raising and + failing auth.""" + from fastapi import HTTPException + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + team_id_jwt_field="team_id", + team_claim_fallback=True, + ) + token = {"sub": "user-1", "team_id": "claim-team-not-in-db"} + + with patch( + "litellm.proxy.auth.handle_jwt.get_team_object", + new_callable=AsyncMock, + ) as mock_get_team: + mock_get_team.side_effect = HTTPException(status_code=404, detail="missing") + + team_id, team_object = await JWTAuthManager.find_and_validate_specific_team_id( + jwt_handler=jwt_handler, + jwt_valid_token=token, + prisma_client=None, + user_api_key_cache=None, + parent_otel_span=None, + proxy_logging_obj=None, + ) + + assert team_id is None + assert team_object is None + + +@pytest.mark.asyncio +async def test_find_team_with_model_access_unresolved_group_claim_returns_none( + monkeypatch, +): + """With `team_claim_fallback=True`: group claim resolves to team_ids that + don't exist in the DB — return (None, None) instead of raising 403, so + the single-team fallback can run.""" + import sys + import types + + from fastapi import HTTPException + + from litellm.router import Router + + router = Router( + model_list=[ + {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}} + ] + ) + proxy_server_module = types.ModuleType("proxy_server") + proxy_server_module.llm_router = router + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) + + async def raise_404(*_args, **_kwargs): + raise HTTPException(status_code=404, detail="missing") + + monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", raise_404) + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_claim_fallback=True) + + team_id, team_object = await JWTAuthManager.find_team_with_model_access( + team_ids={"idp-group-a", "idp-group-b"}, + requested_model="gpt-4o-mini", + route="/chat/completions", + jwt_handler=jwt_handler, + prisma_client=None, + user_api_key_cache=None, + parent_otel_span=None, + proxy_logging_obj=None, + ) + + assert team_id is None + assert team_object is None + + +@pytest.mark.asyncio +async def test_find_and_validate_specific_team_id_non_http_exception_still_propagates(): + """Regression guard: only the 404 HTTPException raised by + `get_team_object` ("team doesn't exist in db") is softened. Other + errors — e.g. "No DB Connected" — must still propagate so operator-side + problems are loud.""" + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="team_id") + token = {"sub": "user-1", "team_id": "some-claim-team"} + + with patch( + "litellm.proxy.auth.handle_jwt.get_team_object", + new_callable=AsyncMock, + ) as mock_get_team: + mock_get_team.side_effect = RuntimeError("simulated infrastructure error") + + with pytest.raises(RuntimeError, match="simulated infrastructure error"): + await JWTAuthManager.find_and_validate_specific_team_id( + jwt_handler=jwt_handler, + jwt_valid_token=token, + prisma_client=None, + user_api_key_cache=None, + parent_otel_span=None, + proxy_logging_obj=None, + ) + + +@pytest.mark.asyncio +async def test_find_and_validate_specific_team_id_non_404_http_exception_propagates(): + """Regression guard: only 404 HTTPException is softened. If + `get_team_object` is ever updated to raise a different HTTP status code + (e.g. 403 for a blocked team), that error must still propagate rather + than silently fall through to the single-team DB fallback.""" + from fastapi import HTTPException + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="team_id") + token = {"sub": "user-1", "team_id": "some-claim-team"} + + for status_code in (400, 403, 500): + with patch( + "litellm.proxy.auth.handle_jwt.get_team_object", + new_callable=AsyncMock, + ) as mock_get_team: + mock_get_team.side_effect = HTTPException( + status_code=status_code, detail="non-404 failure" + ) + + with pytest.raises(HTTPException) as exc_info: + await JWTAuthManager.find_and_validate_specific_team_id( + jwt_handler=jwt_handler, + jwt_valid_token=token, + prisma_client=None, + user_api_key_cache=None, + parent_otel_span=None, + proxy_logging_obj=None, + ) + assert exc_info.value.status_code == status_code + + +@pytest.mark.asyncio +async def test_find_team_with_model_access_enforce_team_based_access_still_raises(): + """Regression guard: when no group claims are present and + `enforce_team_based_model_access` is on, the original 403 still fires — + the new soft-fail only applies to the unresolved-claim path inside the + loop, not to the no-team-claims-at-all path at the top.""" + from fastapi import HTTPException + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(enforce_team_based_model_access=True) + + with pytest.raises(HTTPException) as exc_info: + await JWTAuthManager.find_team_with_model_access( + team_ids=set(), + requested_model="gpt-4o-mini", + route="/chat/completions", + jwt_handler=jwt_handler, + prisma_client=None, + user_api_key_cache=None, + parent_otel_span=None, + proxy_logging_obj=None, + ) + + assert exc_info.value.status_code == 403 + assert "enforce_team_based_model_access" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_find_team_with_model_access_resolved_team_without_model_still_raises_403( + monkeypatch, +): + """Regression guard: when the JWT group claim DOES resolve to a real + LiteLLM team but that team does not grant the requested model, keep the + original 403. Only the unresolved-claim case is softened.""" + import sys + import types + + from fastapi import HTTPException + + from litellm.router import Router + + router = Router( + model_list=[ + {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}}, + { + "model_name": "gpt-3.5-turbo", + "litellm_params": {"model": "gpt-3.5-turbo"}, + }, + ] + ) + proxy_server_module = types.ModuleType("proxy_server") + proxy_server_module.llm_router = router + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) + + team = LiteLLM_TeamTable(team_id="real-team", models=["gpt-3.5-turbo"]) + + async def mock_get_team_object(*_args, **_kwargs): + return team + + monkeypatch.setattr( + "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object + ) + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() + + with pytest.raises(HTTPException) as exc_info: + await JWTAuthManager.find_team_with_model_access( + team_ids={"real-team"}, + requested_model="gpt-4o-mini", + route="/chat/completions", + jwt_handler=jwt_handler, + prisma_client=None, + user_api_key_cache=None, + parent_otel_span=None, + proxy_logging_obj=None, + ) + + assert exc_info.value.status_code == 403 + assert "No team has access to the requested model" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_find_and_validate_specific_team_id_unresolved_claim_default_raises(): + """Default `team_claim_fallback=False`: unresolved team_id claim must + still raise — preserves the strict claim-based authorization boundary + when the operator has not opted in to the fallback.""" + from fastapi import HTTPException + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="team_id") + token = {"sub": "user-1", "team_id": "claim-team-not-in-db"} + + with patch( + "litellm.proxy.auth.handle_jwt.get_team_object", + new_callable=AsyncMock, + ) as mock_get_team: + mock_get_team.side_effect = HTTPException(status_code=404, detail="missing") + + with pytest.raises(HTTPException) as exc_info: + await JWTAuthManager.find_and_validate_specific_team_id( + jwt_handler=jwt_handler, + jwt_valid_token=token, + prisma_client=None, + user_api_key_cache=None, + parent_otel_span=None, + proxy_logging_obj=None, + ) + + assert exc_info.value.status_code == 404 + + +@pytest.mark.asyncio +async def test_find_team_with_model_access_unresolved_group_claim_default_raises( + monkeypatch, +): + """Default `team_claim_fallback=False`: group claims that don't resolve + to any LiteLLM team must still raise 403 — preserves the strict + claim-based authorization boundary.""" + import sys + import types + + from fastapi import HTTPException + + from litellm.router import Router + + router = Router( + model_list=[ + {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}} + ] + ) + proxy_server_module = types.ModuleType("proxy_server") + proxy_server_module.llm_router = router + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) + + async def raise_404(*_args, **_kwargs): + raise HTTPException(status_code=404, detail="missing") + + monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", raise_404) + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() + + with pytest.raises(HTTPException) as exc_info: + await JWTAuthManager.find_team_with_model_access( + team_ids={"idp-group-a", "idp-group-b"}, + requested_model="gpt-4o-mini", + route="/chat/completions", + jwt_handler=jwt_handler, + prisma_client=None, + user_api_key_cache=None, + parent_otel_span=None, + proxy_logging_obj=None, + ) + + assert exc_info.value.status_code == 403 + + # GH #26789: JWT claim user_id must rebind to legacy DB row after fuzzy match. From 9ccda119194415313b3a40acd0dfedf0136b2826 Mon Sep 17 00:00:00 2001 From: milan-berri Date: Tue, 9 Jun 2026 01:14:24 +0300 Subject: [PATCH 12/22] fix(team_endpoints): don't block /team/update on unchanged team budget (#29525) On /team/update for a standalone (no-org) team, _check_user_team_limits() compared the request max_budget against the caller's personal max_budget whenever max_budget was present in the payload. A team admin whose personal budget is lower than the team's budget could not edit any field (tpm_limit, team name, etc.) because the UI re-sends the unchanged max_budget on every update, tripping the personal-budget check. Pass the team's current max_budget into _check_user_team_limits() and skip the personal-budget comparison when the incoming value is unchanged or lower than the team's current budget. Only genuine increases above the team's current budget are still validated against the caller's personal limit, so no over-relaxation. Proxy admins and the org-scoped path are unaffected. Adds two regression tests for the standalone update path (unchanged budget + tpm_limit change, and lowering the budget), both for a caller whose personal budget is below the team budget. Co-authored-by: Cursor --- .../management_endpoints/team_endpoints.py | 48 +++-- .../test_team_endpoints.py | 193 ++++++++++++++++++ 2 files changed, 226 insertions(+), 15 deletions(-) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 7a784ee4622..f2eafcbf839 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -779,6 +779,7 @@ async def _check_user_team_limits( user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, user_api_key_cache: Any, + existing_team_max_budget: Optional[float] = None, ) -> None: """ Check user team limits for standalone teams (not org-scoped). @@ -789,28 +790,44 @@ async def _check_user_team_limits( Should only be called for standalone teams (when organization_id is None). For org-scoped teams, use _check_org_team_limits() instead. + + `existing_team_max_budget` is the team's current `max_budget` on the + /team/update path. When the incoming `max_budget` is unchanged or lower + than the team's current budget, the personal-budget comparison is skipped + so a team admin can edit other fields (e.g. tpm_limit, team name) without + being blocked by a budget the team already has. The UI sends the full team + object on every update, so the unchanged `max_budget` would otherwise fail. """ # Validate team budget against user's max_budget if data.max_budget is not None and user_api_key_dict.user_id is not None: - user_obj = await get_user_object( - user_id=user_api_key_dict.user_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - user_id_upsert=False, + # On /team/update, allow unchanged or lower budgets without checking + # the caller's personal max_budget. Only increases above the team's + # current budget are validated against the user's personal limit. + budget_unchanged_or_lower = ( + existing_team_max_budget is not None + and data.max_budget <= existing_team_max_budget ) - if ( - user_obj is not None - and user_obj.max_budget is not None - and data.max_budget > user_obj.max_budget - ): - raise HTTPException( - status_code=400, - detail={ - "error": f"max budget higher than user max. User max budget={user_obj.max_budget}. User role={user_api_key_dict.user_role}" - }, + if not budget_unchanged_or_lower: + user_obj = await get_user_object( + user_id=user_api_key_dict.user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_id_upsert=False, ) + if ( + user_obj is not None + and user_obj.max_budget is not None + and data.max_budget > user_obj.max_budget + ): + raise HTTPException( + status_code=400, + detail={ + "error": f"max budget higher than user max. User max budget={user_obj.max_budget}. User role={user_api_key_dict.user_role}" + }, + ) + # Validate team models against user's allowed models if data.models is not None and len(user_api_key_dict.models) > 0: for m in data.models: @@ -1824,6 +1841,7 @@ async def update_team( # noqa: PLR0915 user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, + existing_team_max_budget=existing_team_row.max_budget, ) updated_kv = data.json(exclude_unset=True) diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index d580f1f7703..400ed287802 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -4531,6 +4531,199 @@ async def test_update_team_standalone_budget_exceeds_user_limit(): assert "budget" in str(exc_info.value.message).lower() +@pytest.mark.asyncio +async def test_update_team_standalone_unchanged_budget_allowed(): + """ + Test that /team/update for a standalone team does NOT compare against the + caller's personal max_budget when the budget is unchanged. + + This is the LiteLLM UI scenario: the UI sends the full team object on every + update (including the unchanged max_budget). A team admin only changing + tpm_limit should not be blocked by a budget the team already has. + + Scenario: + - User (team admin) has personal max_budget=$100 + - Standalone team exists with current budget=$500 + - User updates tpm_limit and re-sends the unchanged max_budget=$500 + - Expected: Should succeed (budget unchanged, not an increase) + """ + from fastapi import Request + + from litellm.proxy._types import ( + LiteLLM_UserTable, + UpdateTeamRequest, + UserAPIKeyAuth, + ) + from litellm.proxy.management_endpoints.team_endpoints import update_team + + team_admin_user = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="standalone-unchanged-budget-admin", + models=[], + ) + + # UI re-sends the unchanged max_budget alongside the tpm_limit change. + update_request = UpdateTeamRequest( + team_id="standalone-unchanged-budget-123", + max_budget=500.0, # Unchanged from the team's current budget + tpm_limit=50000, + ) + + dummy_request = MagicMock(spec=Request) + + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + ): + # Mock existing standalone team (no organization_id) with budget=$500 + mock_existing_team = MagicMock() + mock_existing_team.team_id = "standalone-unchanged-budget-123" + mock_existing_team.organization_id = None + mock_existing_team.max_budget = 500.0 + mock_existing_team.model_id = None + mock_existing_team.model_dump.return_value = { + "team_id": "standalone-unchanged-budget-123", + "organization_id": None, + "max_budget": 500.0, + "members_with_roles": [ + {"user_id": "standalone-unchanged-budget-admin", "role": "admin"} + ], + } + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( + return_value=mock_existing_team + ) + mock_prisma.jsonify_team_object = lambda db_data: db_data + + # User has a restrictive personal budget that is lower than the team's. + mock_user_obj = LiteLLM_UserTable( + user_id="standalone-unchanged-budget-admin", + max_budget=100.0, + ) + mock_cache.async_get_cache = AsyncMock(return_value=mock_user_obj) + mock_cache.async_set_cache = AsyncMock() + + mock_updated_team = MagicMock() + mock_updated_team.team_id = "standalone-unchanged-budget-123" + mock_updated_team.organization_id = None + mock_updated_team.max_budget = 500.0 + mock_updated_team.litellm_model_table = None + mock_updated_team.model_dump.return_value = { + "team_id": "standalone-unchanged-budget-123", + "organization_id": None, + "max_budget": 500.0, + "tpm_limit": 50000, + } + mock_prisma.db.litellm_teamtable.update = AsyncMock( + return_value=mock_updated_team + ) + + # Should NOT raise - unchanged budget skips the personal-budget check. + result = await update_team( + data=update_request, + http_request=dummy_request, + user_api_key_dict=team_admin_user, + ) + + assert result is not None + assert result["data"].max_budget == 500.0 + + +@pytest.mark.asyncio +async def test_update_team_standalone_lower_budget_allowed(): + """ + Test that /team/update for a standalone team allows lowering the budget + below the team's current value even when the new value still exceeds the + caller's personal max_budget. + + Scenario: + - User (team admin) has personal max_budget=$100 + - Standalone team exists with current budget=$500 + - User lowers team budget to $300 (a decrease, still above user's $100) + - Expected: Should succeed (decrease is not an increase above team budget) + """ + from fastapi import Request + + from litellm.proxy._types import ( + LiteLLM_UserTable, + UpdateTeamRequest, + UserAPIKeyAuth, + ) + from litellm.proxy.management_endpoints.team_endpoints import update_team + + team_admin_user = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id="standalone-lower-budget-admin", + models=[], + ) + + update_request = UpdateTeamRequest( + team_id="standalone-lower-budget-123", + max_budget=300.0, # Lower than current $500, still above user's $100 + ) + + dummy_request = MagicMock(spec=Request) + + with ( + patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, + patch("litellm.proxy.proxy_server.user_api_key_cache") as mock_cache, + patch("litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"), + patch( + "litellm.proxy.proxy_server.create_audit_log_for_update", new=AsyncMock() + ) as mock_audit, + ): + mock_existing_team = MagicMock() + mock_existing_team.team_id = "standalone-lower-budget-123" + mock_existing_team.organization_id = None + mock_existing_team.max_budget = 500.0 + mock_existing_team.model_id = None + mock_existing_team.model_dump.return_value = { + "team_id": "standalone-lower-budget-123", + "organization_id": None, + "max_budget": 500.0, + "members_with_roles": [ + {"user_id": "standalone-lower-budget-admin", "role": "admin"} + ], + } + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( + return_value=mock_existing_team + ) + mock_prisma.jsonify_team_object = lambda db_data: db_data + + mock_user_obj = LiteLLM_UserTable( + user_id="standalone-lower-budget-admin", + max_budget=100.0, + ) + mock_cache.async_get_cache = AsyncMock(return_value=mock_user_obj) + mock_cache.async_set_cache = AsyncMock() + + mock_updated_team = MagicMock() + mock_updated_team.team_id = "standalone-lower-budget-123" + mock_updated_team.organization_id = None + mock_updated_team.max_budget = 300.0 + mock_updated_team.litellm_model_table = None + mock_updated_team.model_dump.return_value = { + "team_id": "standalone-lower-budget-123", + "organization_id": None, + "max_budget": 300.0, + } + mock_prisma.db.litellm_teamtable.update = AsyncMock( + return_value=mock_updated_team + ) + + result = await update_team( + data=update_request, + http_request=dummy_request, + user_api_key_dict=team_admin_user, + ) + + assert result is not None + assert result["data"].max_budget == 300.0 + + @pytest.mark.asyncio async def test_update_team_org_scoped_budget_exceeds_org_limit(): """ From 1c881eee5da8467ca2d07ea5caedf8bdf5b69069 Mon Sep 17 00:00:00 2001 From: milan-berri Date: Tue, 9 Jun 2026 01:54:19 +0300 Subject: [PATCH 13/22] fix(fireworks): enable tool calling for glm-5p1 in model cost map (#29697) glm-5p1 supports native tools on Fireworks; explicit false flags caused drop_params to strip tools and tool_choice before the provider request. Co-authored-by: Cursor --- litellm/model_prices_and_context_window_backup.json | 12 ++++++------ model_prices_and_context_window.json | 12 ++++++------ .../chat/test_fireworks_ai_chat_transformation.py | 10 ++++++---- 3 files changed, 18 insertions(+), 16 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 397f96fdb1e..e765512175b 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -14286,10 +14286,10 @@ "mode": "chat", "output_cost_per_token": 4.4e-06, "source": "https://fireworks.ai/models/fireworks/glm-5p1", - "supports_function_calling": false, + "supports_function_calling": true, "supports_reasoning": true, - "supports_response_schema": false, - "supports_tool_choice": false + "supports_response_schema": true, + "supports_tool_choice": true }, "fireworks_ai/accounts/fireworks/models/gpt-oss-120b": { "input_cost_per_token": 1.5e-07, @@ -14567,10 +14567,10 @@ "mode": "chat", "output_cost_per_token": 4.4e-06, "source": "https://fireworks.ai/models/fireworks/glm-5p1", - "supports_function_calling": false, + "supports_function_calling": true, "supports_reasoning": true, - "supports_response_schema": false, - "supports_tool_choice": false + "supports_response_schema": true, + "supports_tool_choice": true }, "fireworks_ai/kimi-k2p5": { "cache_read_input_token_cost": 1e-07, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index b2836a096b7..c1c05b982f6 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -14286,10 +14286,10 @@ "mode": "chat", "output_cost_per_token": 4.4e-06, "source": "https://fireworks.ai/models/fireworks/glm-5p1", - "supports_function_calling": false, + "supports_function_calling": true, "supports_reasoning": true, - "supports_response_schema": false, - "supports_tool_choice": false + "supports_response_schema": true, + "supports_tool_choice": true }, "fireworks_ai/accounts/fireworks/models/gpt-oss-120b": { "input_cost_per_token": 1.5e-07, @@ -14567,10 +14567,10 @@ "mode": "chat", "output_cost_per_token": 4.4e-06, "source": "https://fireworks.ai/models/fireworks/glm-5p1", - "supports_function_calling": false, + "supports_function_calling": true, "supports_reasoning": true, - "supports_response_schema": false, - "supports_tool_choice": false + "supports_response_schema": true, + "supports_tool_choice": true }, "fireworks_ai/kimi-k2p5": { "cache_read_input_token_cost": 1e-07, diff --git a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py index ca340b5f275..0221db1b23d 100644 --- a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py +++ b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py @@ -127,12 +127,14 @@ def test_get_supported_openai_params_parallel_tool_calls(): config = FireworksAIConfig() supported_params = config.get_supported_openai_params( - "fireworks_ai/accounts/fireworks/models/glm-4p6" + "fireworks_ai/accounts/fireworks/models/glm-5p1" ) assert "parallel_tool_calls" in supported_params + assert "tools" in supported_params + assert "tool_choice" in supported_params unsupported_params = config.get_supported_openai_params( - "fireworks_ai/accounts/fireworks/models/glm-5p1" + "fireworks_ai/accounts/fireworks/models/llama-v3p1-8b-instruct" ) assert "parallel_tool_calls" not in unsupported_params @@ -163,9 +165,9 @@ def test_get_model_info_respects_explicit_fireworks_capabilities(): """Test that get_model_info preserves explicit capability flags from the model map.""" model_info = get_model_info("fireworks_ai/accounts/fireworks/models/glm-5p1") - assert model_info["supports_function_calling"] is False + assert model_info["supports_function_calling"] is True assert model_info["supports_reasoning"] is True - assert model_info["supports_tool_choice"] is False + assert model_info["supports_tool_choice"] is True def test_get_provider_info_omits_false_supports_reasoning(monkeypatch): From dfd6cbc514f7af84179ab126d15cf68833ca410e Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 9 Jun 2026 04:44:30 +0530 Subject: [PATCH 14/22] fix(vertex): propagate Vertex AI metadata in streaming success callbacks (#29899) * fix(vertex): propagate Vertex AI metadata in streaming success callbacks Streaming calls assembled via stream_chunk_builder were missing vertex_ai_grounding_metadata and vertex_ai_url_context_metadata in standard_logging_object.response. Merge metadata from chunks into the assembled response and mirror non-streaming hidden_params on Gemini chunks. Co-authored-by: Cursor * refactor(vertex): move streaming metadata merge into provider config hook Address review feedback by delegating assembled-stream metadata propagation to VertexGeminiConfig via BaseConfig.apply_assembled_streaming_response_metadata, and only write chunk hidden_params when metadata is non-empty. Co-authored-by: Cursor * fix(redaction): scrub Vertex provider metadata when message logging is off Clear vertex_ai_grounding_metadata and related fields from standard logging responses and assembled streaming ModelResponse objects so turn_off_message_logging cannot leak prompt-derived web search queries. Co-authored-by: Cursor * Use assembled model for streaming metadata hook * Fix Vertex metadata redaction bypass in logging callbacks. Scrub Vertex provider fields from litellm_params.metadata.hidden_params during perform_redaction so streaming success_handler merges do not leak prompt-derived metadata when message logging is disabled. Co-authored-by: Cursor * Fix Vertex streaming metadata from hidden params * fix(vertex): mirror vertex_ai_safety_results on assembled streaming responses The non-streaming transform_response stores safety data under vertex_ai_safety_results, but the streaming path only wrote vertex_ai_safety_ratings. Assembled streaming responses therefore never carried vertex_ai_safety_results, so any consumer reading that field saw a silent difference between streaming and non-streaming calls. Set vertex_ai_safety_results alongside vertex_ai_safety_ratings in the shared stream metadata setter and add it to the assembled metadata field list so it propagates through stream_chunk_builder. * fix(streaming): log provider streaming metadata hook failures instead of swallowing them * refactor(vertex): share single Vertex metadata field tuple across redaction and streaming * refactor(vertex): move Vertex metadata redaction helpers into llms/vertex_ai --------- Co-authored-by: Cursor Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/litellm_core_utils/redact_messages.py | 10 ++ .../streaming_chunk_builder_utils.py | 49 ++++++ litellm/llms/base_llm/chat/transformation.py | 8 + litellm/llms/vertex_ai/common_utils.py | 47 +++++- .../vertex_and_google_ai_studio_gemini.py | 77 ++++++++- litellm/main.py | 6 + litellm/types/llms/vertex_ai.py | 9 ++ .../test_litellm_logging.py | 35 ++++ .../test_redact_messages.py | 93 +++++++++++ .../test_streaming_chunk_builder_utils.py | 150 ++++++++++++++++++ ...test_vertex_and_google_ai_studio_gemini.py | 20 +++ 11 files changed, 499 insertions(+), 5 deletions(-) diff --git a/litellm/litellm_core_utils/redact_messages.py b/litellm/litellm_core_utils/redact_messages.py index dbc9cabdc7a..763596336a0 100644 --- a/litellm/litellm_core_utils/redact_messages.py +++ b/litellm/litellm_core_utils/redact_messages.py @@ -17,6 +17,10 @@ from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.core_helpers import ( get_metadata_variable_name_from_kwargs, ) +from litellm.llms.vertex_ai.common_utils import ( + redact_vertex_ai_metadata_from_litellm_params, + redact_vertex_ai_metadata_from_logged_object, +) from litellm.secret_managers.main import str_to_bool from litellm.types.utils import StandardCallbackDynamicParams @@ -119,10 +123,12 @@ def _redact_standard_logging_object(model_call_details: dict): # ResponsesAPIResponse format - redact content in output items if isinstance(response.get("output"), list): _redact_responses_api_output_dict(response["output"], redacted_str) + redact_vertex_ai_metadata_from_logged_object(response) elif isinstance(response, dict) and "choices" in response: # ModelResponse dict format - redact content in choices if isinstance(response.get("choices"), list): _redact_model_response_dict_choices(response["choices"], redacted_str) + redact_vertex_ai_metadata_from_logged_object(response) elif isinstance(response, str): standard_logging_object["response"] = redacted_str else: @@ -164,6 +170,7 @@ def perform_redaction(model_call_details: dict, result): model_call_details["prompt"] = "" model_call_details["input"] = "" _redact_standard_logging_object(model_call_details) + redact_vertex_ai_metadata_from_litellm_params(model_call_details) # Redact streaming response if ( @@ -174,6 +181,7 @@ def perform_redaction(model_call_details: dict, result): if hasattr(_streaming_response, "choices"): for choice in _streaming_response.choices: _redact_choice_content(choice) + redact_vertex_ai_metadata_from_logged_object(_streaming_response) elif hasattr(_streaming_response, "output"): _redact_responses_api_output(_streaming_response.output) # Redact reasoning field in ResponsesAPIResponse @@ -200,12 +208,14 @@ def perform_redaction(model_call_details: dict, result): if hasattr(_result, "choices") and _result.choices is not None: for choice in _result.choices: _redact_choice_content(choice) + redact_vertex_ai_metadata_from_logged_object(_result) elif isinstance(_result, dict) and "choices" in _result: # Handle dict representation of ModelResponse (e.g., from model_dump()) if _result.get("choices") is not None: _redact_model_response_dict_choices( _result["choices"], "redacted-by-litellm" ) + redact_vertex_ai_metadata_from_logged_object(_result) elif isinstance(_result, dict) and "output" in _result: if isinstance(_result.get("output"), list): _redact_responses_api_output_dict( diff --git a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py index fe7c62c3842..6257cce9aec 100644 --- a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py +++ b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py @@ -20,6 +20,7 @@ from litellm.types.utils import ( ServerToolUse, Usage, ) +from litellm._logging import verbose_logger from litellm.utils import print_verbose, token_counter if TYPE_CHECKING: @@ -79,6 +80,54 @@ class ChunkProcessor: model_response._hidden_params = chunk.get("_hidden_params", {}) return model_response + @staticmethod + def apply_provider_assembled_streaming_metadata( + response: ModelResponse, + chunks: List[Any], + logging_obj: Optional[Any] = None, + ) -> None: + if not chunks: + return + + model = getattr(response, "model", None) + if not model: + return + + custom_llm_provider = None + if logging_obj is not None: + custom_llm_provider = logging_obj.model_call_details.get( + "custom_llm_provider" + ) + + try: + from litellm.litellm_core_utils.get_llm_provider_logic import ( + get_llm_provider, + ) + from litellm.types.utils import LlmProviders + from litellm.utils import ProviderConfigManager + + if custom_llm_provider: + provider = LlmProviders(custom_llm_provider) + else: + _, provider_str, _, _ = get_llm_provider(model) + provider = LlmProviders(provider_str) + + provider_config = ProviderConfigManager.get_provider_chat_config( + model=model, + provider=provider, + ) + if provider_config is not None: + provider_config.apply_assembled_streaming_response_metadata( + response=response, + chunks=chunks, + ) + except Exception as e: + verbose_logger.debug( + "apply_provider_assembled_streaming_metadata failed for model=%s: %s", + model, + e, + ) + @staticmethod def _get_chunk_id(chunks: List[Dict[str, Any]]) -> str: """ diff --git a/litellm/llms/base_llm/chat/transformation.py b/litellm/llms/base_llm/chat/transformation.py index 5f35a58ce1f..8f9d5cad7c4 100644 --- a/litellm/llms/base_llm/chat/transformation.py +++ b/litellm/llms/base_llm/chat/transformation.py @@ -442,6 +442,14 @@ class BaseConfig(ABC): """Hook for providers to post-process streaming responses. Default: pass-through.""" return stream + def apply_assembled_streaming_response_metadata( + self, + response: "ModelResponse", + chunks: List[Any], + ) -> None: + """Hook for providers to merge chunk metadata into assembled streaming responses.""" + return None + def calculate_additional_costs( self, model: str, prompt_tokens: int, completion_tokens: int ) -> Optional[dict]: diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index e6e39651109..85c23d8603c 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -12,7 +12,11 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import unpack_defs from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.types.llms.openai import AllMessageValues -from litellm.types.llms.vertex_ai import PartType, Schema +from litellm.types.llms.vertex_ai import ( + VERTEX_AI_PROVIDER_METADATA_FIELDS, + PartType, + Schema, +) from litellm.types.utils import TokenCountResponse from litellm.utils import supports_response_schema, supports_system_messages @@ -27,6 +31,47 @@ class VertexAIError(BaseLLMException): super().__init__(message=message, status_code=status_code, headers=headers) +def redact_vertex_ai_metadata_from_logged_object(obj: Any) -> None: + if isinstance(obj, dict): + for field in VERTEX_AI_PROVIDER_METADATA_FIELDS: + if field in obj: + obj[field] = [] + hidden_params = obj.get("_hidden_params") + if isinstance(hidden_params, dict): + for field in VERTEX_AI_PROVIDER_METADATA_FIELDS: + hidden_params.pop(field, None) + return + + for field in VERTEX_AI_PROVIDER_METADATA_FIELDS: + if hasattr(obj, field): + setattr(obj, field, []) + hidden_params = getattr(obj, "_hidden_params", None) + if isinstance(hidden_params, dict): + for field in VERTEX_AI_PROVIDER_METADATA_FIELDS: + hidden_params.pop(field, None) + + +def redact_vertex_ai_metadata_from_litellm_params(model_call_details: dict) -> None: + """ + success_handler() merges response._hidden_params into + litellm_params.metadata['hidden_params'] before redaction runs, so the Vertex + metadata must be scrubbed from that copy too. + """ + litellm_params = model_call_details.get("litellm_params") + if not isinstance(litellm_params, dict): + return + + for metadata_key in ("metadata", "litellm_metadata"): + metadata = litellm_params.get(metadata_key) + if not isinstance(metadata, dict): + continue + hidden_params = metadata.get("hidden_params") + if not isinstance(hidden_params, dict): + continue + for field in VERTEX_AI_PROVIDER_METADATA_FIELDS: + hidden_params.pop(field, None) + + def vertex_request_labels_from_litellm_params( litellm_params: Optional[dict], ) -> Optional[Dict[str, str]]: 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 7d355a2e908..430a789d2a0 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 @@ -63,6 +63,7 @@ from litellm.types.llms.openai import ( OpenAIChatCompletionFinishReason, ) from litellm.types.llms.vertex_ai import ( + VERTEX_AI_PROVIDER_METADATA_FIELDS, VERTEX_CREDENTIALS_TYPES, Candidates, ContentType, @@ -2258,6 +2259,71 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): citation_metadata, ) + @staticmethod + def _get_stream_chunk_attr(chunk: Any, field_name: str) -> Any: + if isinstance(chunk, dict): + value = chunk.get(field_name) + if value is not None: + return value + model_extra = chunk.get("model_extra") + if isinstance(model_extra, dict): + value = model_extra.get(field_name) + if value is not None: + return value + hidden_params = chunk.get("_hidden_params") + if isinstance(hidden_params, dict): + return hidden_params.get(field_name) + return None + return getattr(chunk, field_name, None) + + @staticmethod + def _set_stream_metadata_on_response( + model_response: Any, + grounding_metadata: List[dict], + url_context_metadata: List[dict], + safety_ratings: List[dict], + citation_metadata: List[dict], + ) -> None: + setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata) # type: ignore + if grounding_metadata: + model_response._hidden_params["vertex_ai_grounding_metadata"] = ( + grounding_metadata + ) + setattr(model_response, "vertex_ai_url_context_metadata", url_context_metadata) # type: ignore + if url_context_metadata: + model_response._hidden_params["vertex_ai_url_context_metadata"] = ( + url_context_metadata + ) + setattr(model_response, "vertex_ai_safety_ratings", safety_ratings) # type: ignore + setattr(model_response, "vertex_ai_safety_results", safety_ratings) # type: ignore + if safety_ratings: + model_response._hidden_params["vertex_ai_safety_ratings"] = safety_ratings + model_response._hidden_params["vertex_ai_safety_results"] = safety_ratings + setattr(model_response, "vertex_ai_citation_metadata", citation_metadata) # type: ignore + if citation_metadata: + model_response._hidden_params["vertex_ai_citation_metadata"] = ( + citation_metadata + ) + + def apply_assembled_streaming_response_metadata( + self, + response: ModelResponse, + chunks: List[Any], + ) -> None: + for field_name in VERTEX_AI_PROVIDER_METADATA_FIELDS: + merged: List[Any] = [] + for chunk in chunks: + value = VertexGeminiConfig._get_stream_chunk_attr(chunk, field_name) + if not value: + continue + if isinstance(value, list): + merged.extend(value) + else: + merged.append(value) + if merged: + setattr(response, field_name, merged) + response._hidden_params[field_name] = merged + @staticmethod def _convert_grounding_metadata_to_annotations( grounding_metadata: List[dict], @@ -3390,10 +3456,13 @@ class ModelResponseIterator: if choice.finish_reason == "stop": choice.finish_reason = "tool_calls" - setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata) # type: ignore - setattr(model_response, "vertex_ai_url_context_metadata", url_context_metadata) # type: ignore - setattr(model_response, "vertex_ai_safety_ratings", safety_ratings) # type: ignore - setattr(model_response, "vertex_ai_citation_metadata", citation_metadata) # type: ignore + VertexGeminiConfig._set_stream_metadata_on_response( + model_response, + grounding_metadata, + url_context_metadata, + safety_ratings, + citation_metadata, + ) return ( grounding_metadata, diff --git a/litellm/main.py b/litellm/main.py index 64891e2def9..1a0d0312d73 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -7761,6 +7761,9 @@ def stream_chunk_builder( # noqa: PLR0915 "cost", logging_obj._response_cost_calculator(result=response), ) + processor.apply_provider_assembled_streaming_metadata( + response, chunks, logging_obj + ) return response tool_call_chunks = [ @@ -7940,6 +7943,9 @@ def stream_chunk_builder( # noqa: PLR0915 usage, "cost", logging_obj._response_cost_calculator(result=response) ) + processor.apply_provider_assembled_streaming_metadata( + response, chunks, logging_obj + ) return response except Exception as e: verbose_logger.exception( diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index 00db7b199b6..b28fee51284 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -758,3 +758,12 @@ class VertexPartnerProvider(str, Enum): llama = "llama" ai21 = "ai21" claude = "claude" + + +VERTEX_AI_PROVIDER_METADATA_FIELDS = ( + "vertex_ai_grounding_metadata", + "vertex_ai_url_context_metadata", + "vertex_ai_safety_ratings", + "vertex_ai_safety_results", + "vertex_ai_citation_metadata", +) diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index d57d8dafdbd..34edd6eccf3 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -2165,6 +2165,41 @@ def test_get_assembled_streaming_response_returns_result_for_streaming(): assert assembled is result +def test_streaming_success_handler_includes_vertex_ai_metadata_in_standard_logging(): + """Assembled streaming responses should include Vertex AI metadata in logging payload.""" + import datetime + + from litellm.types.utils import Choices, Message + + logging_obj = _make_logging_obj(stream=True) + grounding_metadata = [{"webSearchQueries": ["weather in SF"]}] + url_context_metadata = [{"urlMetadata": [{"retrievedUrl": "https://example.com"}]}] + result = ModelResponse( + id="resp-1", + choices=[ + Choices( + index=0, + message=Message(role="assistant", content="hello"), + finish_reason="stop", + ) + ], + model="gemini-2.5-flash", + ) + setattr(result, "vertex_ai_grounding_metadata", grounding_metadata) + setattr(result, "vertex_ai_url_context_metadata", url_context_metadata) + result._hidden_params["vertex_ai_grounding_metadata"] = grounding_metadata + result._hidden_params["vertex_ai_url_context_metadata"] = url_context_metadata + + start = datetime.datetime.now() + end = datetime.datetime.now() + logging_obj.success_handler(result=result, start_time=start, end_time=end) + + payload = logging_obj.model_call_details.get("standard_logging_object") + assert payload is not None + assert payload["response"]["vertex_ai_grounding_metadata"] == grounding_metadata + assert payload["response"]["vertex_ai_url_context_metadata"] == url_context_metadata + + def test_get_assembled_streaming_response_returns_none_for_non_streaming_text_completion(): """Non-streaming TextCompletionResponse should also return None.""" import datetime diff --git a/tests/test_litellm/litellm_core_utils/test_redact_messages.py b/tests/test_litellm/litellm_core_utils/test_redact_messages.py index 60cfff6e4a0..36f220f9a2c 100644 --- a/tests/test_litellm/litellm_core_utils/test_redact_messages.py +++ b/tests/test_litellm/litellm_core_utils/test_redact_messages.py @@ -349,3 +349,96 @@ class TestPerformRedaction: assert redacted.output[0].content[0].text == "redacted-by-litellm" assert response.output[0].content[0].text == "sensitive output" + + def test_redacts_vertex_provider_metadata_in_standard_logging_response(self): + details = { + "standard_logging_object": { + "messages": [{"role": "user", "content": "sensitive prompt"}], + "response": { + "choices": [ + { + "message": { + "content": "sensitive answer", + "role": "assistant", + } + } + ], + "vertex_ai_grounding_metadata": [ + {"webSearchQueries": ["sensitive search term"]} + ], + "vertex_ai_url_context_metadata": [ + {"urlMetadata": [{"retrievedUrl": "https://example.com"}]} + ], + }, + } + } + + perform_redaction(details, None) + + response = details["standard_logging_object"]["response"] + assert response["choices"][0]["message"]["content"] == "redacted-by-litellm" + assert response["vertex_ai_grounding_metadata"] == [] + assert response["vertex_ai_url_context_metadata"] == [] + + def test_redacts_vertex_provider_metadata_on_streaming_model_response(self): + response = litellm.ModelResponse( + id="resp-1", + choices=[ + litellm.Choices( + message=litellm.Message( + content="sensitive answer", + role="assistant", + ) + ) + ], + model="gemini-2.5-flash", + ) + setattr( + response, + "vertex_ai_grounding_metadata", + [{"webSearchQueries": ["sensitive search term"]}], + ) + response._hidden_params["vertex_ai_grounding_metadata"] = [ + {"webSearchQueries": ["sensitive search term"]} + ] + + details = { + "stream": True, + "complete_streaming_response": response, + } + + perform_redaction(details, response) + + assert response.choices[0].message.content == "redacted-by-litellm" + assert getattr(response, "vertex_ai_grounding_metadata") == [] + assert "vertex_ai_grounding_metadata" not in response._hidden_params + + def test_redacts_vertex_provider_metadata_from_metadata_hidden_params(self): + """Streaming success_handler copies _hidden_params into metadata before redaction.""" + details = { + "stream": True, + "litellm_params": { + "metadata": { + "hidden_params": { + "response_cost": 0.01, + "vertex_ai_grounding_metadata": [ + {"webSearchQueries": ["sensitive search term"]} + ], + "vertex_ai_url_context_metadata": [ + {"urlMetadata": [{"retrievedUrl": "https://example.com"}]} + ], + "vertex_ai_safety_ratings": [{"category": "HARM"}], + "vertex_ai_citation_metadata": [{"citations": ["source"]}], + } + } + }, + } + + perform_redaction(details, None) + + hidden_params = details["litellm_params"]["metadata"]["hidden_params"] + assert hidden_params["response_cost"] == 0.01 + assert "vertex_ai_grounding_metadata" not in hidden_params + assert "vertex_ai_url_context_metadata" not in hidden_params + assert "vertex_ai_safety_ratings" not in hidden_params + assert "vertex_ai_citation_metadata" not in hidden_params diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py index e40a0817fd9..77765340c61 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py @@ -613,3 +613,153 @@ def test_stream_chunk_builder_dict_snapshot_preserves_hidden_provider_fields(): assert ( response._hidden_params["provider_specific_fields"]["traffic_type"] == "default" ) + + +def test_stream_chunk_builder_propagates_vertex_ai_metadata_from_chunks(): + """Vertex AI metadata on streaming chunks must appear on assembled response.""" + grounding_metadata = [{"webSearchQueries": ["weather in SF"]}] + url_context_metadata = [{"urlMetadata": [{"retrievedUrl": "https://example.com"}]}] + + chunk1 = ModelResponseStream( + id="chatcmpl-vertex-1", + created=1, + model="gemini-2.5-flash", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta(content="The weather", role="assistant"), + ) + ], + ) + setattr(chunk1, "vertex_ai_grounding_metadata", grounding_metadata) + chunk1._hidden_params["vertex_ai_grounding_metadata"] = grounding_metadata + + chunk2 = ModelResponseStream( + id="chatcmpl-vertex-1", + created=1, + model="gemini-2.5-flash", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason="stop", + index=0, + delta=Delta(content=" is sunny.", role="assistant"), + ) + ], + ) + setattr(chunk2, "vertex_ai_url_context_metadata", url_context_metadata) + chunk2._hidden_params["vertex_ai_url_context_metadata"] = url_context_metadata + + response = stream_chunk_builder(chunks=[chunk1, chunk2]) + assert response is not None + assert getattr(response, "vertex_ai_grounding_metadata") == grounding_metadata + assert getattr(response, "vertex_ai_url_context_metadata") == url_context_metadata + assert response._hidden_params["vertex_ai_grounding_metadata"] == grounding_metadata + assert ( + response._hidden_params["vertex_ai_url_context_metadata"] + == url_context_metadata + ) + + dumped = response.model_dump() + assert dumped["vertex_ai_grounding_metadata"] == grounding_metadata + assert dumped["vertex_ai_url_context_metadata"] == url_context_metadata + + +def test_stream_chunk_builder_uses_assembled_model_for_provider_metadata(): + grounding_metadata = [{"webSearchQueries": ["weather in SF"]}] + + chunk1 = ModelResponseStream( + id="chatcmpl-vertex-router", + created=1, + model="gpt-4o", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta(content="The weather", role="assistant"), + ) + ], + ) + chunk2 = ModelResponseStream( + id="chatcmpl-vertex-router", + created=1, + model="gemini-2.5-flash", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason="stop", + index=0, + delta=Delta(content=" is sunny.", role=None), + ) + ], + ) + setattr(chunk2, "vertex_ai_grounding_metadata", grounding_metadata) + chunk2._hidden_params["vertex_ai_grounding_metadata"] = grounding_metadata + + response = stream_chunk_builder(chunks=[chunk1, chunk2]) + assert response is not None + assert response.model == "gemini-2.5-flash" + assert getattr(response, "vertex_ai_grounding_metadata") == grounding_metadata + + +def test_stream_chunk_builder_propagates_vertex_ai_safety_results(): + """Assembled response must expose safety data under the non-streaming field name.""" + safety_ratings = [ + [{"category": "HARM_CATEGORY_HATE_SPEECH", "probability": "NEGLIGIBLE"}] + ] + + chunk = ModelResponseStream( + id="chatcmpl-vertex-safety", + created=1, + model="gemini-2.5-flash", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason="stop", + index=0, + delta=Delta(content="hello", role="assistant"), + ) + ], + ) + setattr(chunk, "vertex_ai_safety_ratings", safety_ratings) + setattr(chunk, "vertex_ai_safety_results", safety_ratings) + chunk._hidden_params["vertex_ai_safety_ratings"] = safety_ratings + chunk._hidden_params["vertex_ai_safety_results"] = safety_ratings + + response = stream_chunk_builder(chunks=[chunk]) + assert response is not None + assert getattr(response, "vertex_ai_safety_results") == safety_ratings + assert response._hidden_params["vertex_ai_safety_results"] == safety_ratings + assert response.model_dump()["vertex_ai_safety_results"] == safety_ratings + + +def test_stream_chunk_builder_propagates_vertex_ai_metadata_from_dict_chunks(): + """Dict snapshot chunks (model_dump) should also propagate Vertex AI metadata.""" + chunk_dict = ModelResponseStream( + id="chatcmpl-vertex-2", + created=1, + model="gemini-2.5-flash", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason="stop", + index=0, + delta=Delta(content="hello", role="assistant"), + ) + ], + ).model_dump() + chunk_dict["_hidden_params"] = { + "vertex_ai_grounding_metadata": [{"webSearchQueries": ["test query"]}] + } + + response = stream_chunk_builder(chunks=[chunk_dict]) + assert response is not None + assert getattr(response, "vertex_ai_grounding_metadata") == [ + {"webSearchQueries": ["test query"]} + ] + assert response.model_dump()["vertex_ai_grounding_metadata"] == [ + {"webSearchQueries": ["test query"]} + ] 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 0d02521433a..671d7355e8f 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 @@ -1459,6 +1459,26 @@ def test_vertex_ai_process_candidates_with_grounding_metadata(): assert len(result[0]) == 1 +def test_set_stream_metadata_mirrors_non_streaming_safety_field_names(): + safety_ratings = [ + [{"category": "HARM_CATEGORY_HATE_SPEECH", "probability": "NEGLIGIBLE"}] + ] + + model_response = ModelResponse() + VertexGeminiConfig._set_stream_metadata_on_response( + model_response=model_response, + grounding_metadata=[], + url_context_metadata=[], + safety_ratings=safety_ratings, + citation_metadata=[], + ) + + assert getattr(model_response, "vertex_ai_safety_ratings") == safety_ratings + assert getattr(model_response, "vertex_ai_safety_results") == safety_ratings + assert model_response._hidden_params["vertex_ai_safety_ratings"] == safety_ratings + assert model_response._hidden_params["vertex_ai_safety_results"] == safety_ratings + + def test_vertex_ai_tool_call_id_format(): """ Test that tool call IDs have the correct format and length. From f59e4ebc9e98e26b434aa3ee338f41b3f3cf3fa6 Mon Sep 17 00:00:00 2001 From: milan-berri Date: Tue, 9 Jun 2026 02:27:35 +0300 Subject: [PATCH 15/22] fix(ui): show team projects to internal users (#28855) Allow internal users to fetch their backend-scoped project list so the key creation project dropdown can populate for selected teams. --- .../app/(dashboard)/hooks/projects/useProjects.test.ts | 10 +++++++++- .../src/app/(dashboard)/hooks/projects/useProjects.ts | 6 ++++-- 2 files changed, 13 insertions(+), 3 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useProjects.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useProjects.test.ts index 13b9107bdc1..39d1b28303d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useProjects.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useProjects.test.ts @@ -115,8 +115,16 @@ describe("useProjects", () => { expect(global.fetch).not.toHaveBeenCalled(); }); - it("should not fetch when userRole is not an admin role", () => { + it("should fetch when userRole is an internal user role", async () => { mockUseAuthorized.mockReturnValue({ accessToken: "test-token", userRole: "Internal User" }); + (global.fetch as any).mockResolvedValue({ ok: true, json: async () => mockProjects }); + const { result } = renderHook(() => useProjects(), { wrapper: makeWrapper(queryClient) }); + await waitFor(() => expect(result.current.isSuccess).toBe(true)); + expect(global.fetch).toHaveBeenCalled(); + }); + + it("should not fetch when userRole cannot read projects", () => { + mockUseAuthorized.mockReturnValue({ accessToken: "test-token", userRole: "regular_user" }); const { result } = renderHook(() => useProjects(), { wrapper: makeWrapper(queryClient) }); expect(result.current.isFetched).toBe(false); expect(global.fetch).not.toHaveBeenCalled(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useProjects.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useProjects.ts index 7bdc8a4fe6d..c240dbb0170 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useProjects.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useProjects.ts @@ -2,7 +2,7 @@ import { useQuery } from "@tanstack/react-query"; import { createQueryKeys } from "../common/queryKeysFactory"; import { getProxyBaseUrl, getGlobalLitellmHeaderName, deriveErrorMessage, handleError } from "@/components/networking"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; -import { all_admin_roles } from "@/utils/roles"; +import { all_admin_roles, internalUserRoles } from "@/utils/roles"; // ── Types ──────────────────────────────────────────────────────────────────── @@ -42,6 +42,8 @@ export interface ProjectResponse { export const projectKeys = createQueryKeys("projects"); +const projectReaderRoles = [...all_admin_roles, ...internalUserRoles]; + // ── Fetch function ─────────────────────────────────────────────────────────── const fetchProjects = async (accessToken: string): Promise => { @@ -74,6 +76,6 @@ export const useProjects = () => { return useQuery({ queryKey: projectKeys.list({}), queryFn: async () => fetchProjects(accessToken!), - enabled: Boolean(accessToken) && all_admin_roles.includes(userRole!), + enabled: Boolean(accessToken) && projectReaderRoles.includes(userRole!), }); }; From bac2590b39b3a32a25ae4d6a932a342349def497 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Mon, 8 Jun 2026 16:39:21 -0700 Subject: [PATCH 16/22] build(deps): bump pyjwt to 2.13.0 and ws override to 8.20.1 (#29982) Raise the PyJWT floor in pyproject (>=2.13.0,<3.0) and re-resolve uv.lock so the proxy installs 2.13.0 instead of 2.12.0. Bump the ws transitive-version override in the dashboard from 8.19.0 to 8.20.1 and regenerate package-lock; jsdom and openai both dedupe onto the single 8.20.1 copy. Both are routine dependency maintenance bumps to keep pinned versions current. --- pyproject.toml | 2 +- ui/litellm-dashboard/package-lock.json | 6 +++--- ui/litellm-dashboard/package.json | 2 +- uv.lock | 13 ++++++++----- 4 files changed, 13 insertions(+), 10 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 577e800d79b..28e6f48dc4c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -53,7 +53,7 @@ proxy = [ "orjson>=3.11.6,<4.0", "apscheduler>=3.11.2,<4.0", "fastapi-sso>=0.19.0,<1.0", - "PyJWT>=2.12.0,<3.0", + "PyJWT>=2.13.0,<3.0", "python-multipart>=0.0.27,<1.0", "cryptography>=46.0.7,<47.0", "pynacl>=1.6.2,<2.0", diff --git a/ui/litellm-dashboard/package-lock.json b/ui/litellm-dashboard/package-lock.json index 7000efd63b4..b7dd2a6f59b 100644 --- a/ui/litellm-dashboard/package-lock.json +++ b/ui/litellm-dashboard/package-lock.json @@ -13770,9 +13770,9 @@ } }, "node_modules/ws": { - "version": "8.19.0", - "resolved": "https://registry.npmjs.org/ws/-/ws-8.19.0.tgz", - "integrity": "sha512-blAT2mjOEIi0ZzruJfIhb3nps74PRWTCz1IjglWEEpQl5XS/UNama6u2/rjFkDDouqr4L67ry+1aGIALViWjDg==", + "version": "8.20.1", + "resolved": "https://registry.npmjs.org/ws/-/ws-8.20.1.tgz", + "integrity": "sha512-It4dO0K5v//JtTXuPkfEOaI3uUN87iYPnqo/ZzqCoG3g8uhA66QUMs/SrM0YK7/NAu+r4LMh/9dq2A7k+rHs+w==", "devOptional": true, "license": "MIT", "engines": { diff --git a/ui/litellm-dashboard/package.json b/ui/litellm-dashboard/package.json index ca753e59dc5..623389e76ed 100644 --- a/ui/litellm-dashboard/package.json +++ b/ui/litellm-dashboard/package.json @@ -90,7 +90,7 @@ "glob": "13.0.0", "minimatch": "10.2.4", "lodash": "4.18.1", - "ws": "8.19.0", + "ws": "8.20.1", "braces": "3.0.3", "axios": "1.13.6", "postcss": "8.5.13" diff --git a/uv.lock b/uv.lock index 4386fce1b4d..2403a7fbf03 100644 --- a/uv.lock +++ b/uv.lock @@ -9,7 +9,7 @@ resolution-markers = [ ] [options] -exclude-newer = "2026-06-03T21:40:52.018333Z" +exclude-newer = "2026-06-05T23:18:37.734017Z" exclude-newer-span = "P3D" [manifest] @@ -3522,7 +3522,7 @@ requires-dist = [ { name = "prometheus-client", marker = "extra == 'proxy-runtime'", specifier = ">=0.20.0,<1.0" }, { name = "pydantic", specifier = ">=2.10.0,<3.0.0" }, { name = "pydantic-settings", marker = "extra == 'proxy'", specifier = ">=2.14.1,<3.0" }, - { name = "pyjwt", marker = "extra == 'proxy'", specifier = ">=2.12.0,<3.0" }, + { name = "pyjwt", marker = "extra == 'proxy'", specifier = ">=2.13.0,<3.0" }, { name = "pynacl", marker = "extra == 'proxy'", specifier = ">=1.6.2,<2.0" }, { name = "pypdf", marker = "python_full_version < '3.14' and extra == 'proxy-runtime'", specifier = ">=6.10.2,<7.0" }, { name = "pyroscope-io", marker = "sys_platform != 'win32' and extra == 'proxy'", specifier = ">=0.8.16,<1.0" }, @@ -5983,11 +5983,14 @@ wheels = [ [[package]] name = "pyjwt" -version = "2.12.0" +version = "2.13.0" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/a8/10/e8192be5f38f3e8e7e046716de4cae33d56fd5ae08927a823bb916be36c1/pyjwt-2.12.0.tar.gz", hash = "sha256:2f62390b667cd8257de560b850bb5a883102a388829274147f1d724453f8fb02", size = 102511, upload-time = "2026-03-12T17:15:30.831Z" } +dependencies = [ + { name = "typing-extensions", marker = "python_full_version < '3.11'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/3b/81/58d0ac84e1ef3a3843791d6954d94c0b33d526c75eeb1efbce9d0a4c4077/pyjwt-2.13.0.tar.gz", hash = "sha256:41571c89ca91598c79e8ef18a2d07367d4810fbbd6f637794879baf1b7703423", size = 107515, upload-time = "2026-05-21T19:54:36.618Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/15/70/70f895f404d363d291dcf62c12c85fdd47619ad9674ac0f53364d035925a/pyjwt-2.12.0-py3-none-any.whl", hash = "sha256:9bb459d1bdd0387967d287f5656bf7ec2b9a26645d1961628cda1764e087fd6e", size = 29700, upload-time = "2026-03-12T17:15:29.257Z" }, + { url = "https://files.pythonhosted.org/packages/a3/5e/ecf12fdb62546d64385c158514e9b2b671f7832108ef2ecd2020ce0af2d1/pyjwt-2.13.0-py3-none-any.whl", hash = "sha256:66adcc2aff09b3f1bbd95fc1e1577df8ac8723c978552fd43304c8a290ac5728", size = 31274, upload-time = "2026-05-21T19:54:35.362Z" }, ] [package.optional-dependencies] From c24a3603d971c8af8a29d4f6ef2fa16cfd097100 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Mon, 8 Jun 2026 16:55:35 -0700 Subject: [PATCH 17/22] fix(team-management): delete a team's BYOK models when the team is deleted (#29977) A team's BYOK models (rows in LiteLLM_ProxyModelTable with model_info.team_id set) were left orphaned when the team was deleted; they lingered in the database and kept showing on the Models + Endpoints page. delete_team now removes them via a new delete_team_models helper that deletes the rows in one transaction and syncs the in-memory router only after that transaction commits, run before the team rows are deleted so a mid-flight failure never leaves the team gone with its models orphaned --- .../model_management_endpoints.py | 44 +++++- .../management_endpoints/team_endpoints.py | 14 ++ .../test_model_management_endpoints.py | 143 ++++++++++++++++++ .../test_team_endpoints.py | 10 ++ 4 files changed, 209 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 0cbccfc18ad..0be476469a6 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -561,7 +561,7 @@ async def _setup_new_team_model_assignment( async def _get_team_deployments( - team_id: str, prisma_client: PrismaClient + team_id: str, prisma_client: PrismaClient, table: Optional[Any] = None ) -> List[LiteLLM_ProxyModelTable]: """ Fetch all deployments for a given team_id from the database. @@ -572,9 +572,13 @@ async def _get_team_deployments( Note: prisma-client-py 0.11.0 does not support JSON path filtering, so we filter by the model_name prefix (team models use "model_name_{team_id}_*") and confirm team_id in model_info with Python-side filtering. + + Pass ``table`` (a transaction's proxy-model table) to run the read inside an + existing transaction. """ prefix = f"model_name_{team_id}_" - response = await ModelRepository(prisma_client).table.find_many( + table = table or ModelRepository(prisma_client).table + response = await table.find_many( where={ "model_name": {"startswith": prefix}, } @@ -596,6 +600,42 @@ async def _get_team_deployments( return result +async def delete_team_models( + team_ids: List[str], + prisma_client: PrismaClient, + llm_router: Optional[Any], +) -> List[str]: + """ + Delete every BYOK model owned by the given teams, from the DB and the router. + + The DB rows are removed inside a single transaction, so deletion is atomic + across all team_ids. Each team's rows are deleted by the exact model_ids read + in the same transaction, which keeps the deleted set identical to the set + handed to the router. The router is synced only after the transaction commits, + so a rollback can never leave a deployment live in the router without its row. + + Returns the model_ids that were deleted. + """ + deleted_model_ids: List[str] = [] + async with prisma_client.db.tx() as tx: + for team_id in team_ids: + rows = await _get_team_deployments( + team_id, prisma_client, table=tx.litellm_proxymodeltable + ) + model_ids = [row.model_id for row in rows] + if model_ids: + await tx.litellm_proxymodeltable.delete_many( + where={"model_id": {"in": model_ids}} + ) + deleted_model_ids.extend(model_ids) + + if llm_router is not None: + for model_id in deleted_model_ids: + llm_router.delete_deployment(id=model_id) + + return deleted_model_ids + + async def _get_team_public_model_names( team_id: str, prisma_client: PrismaClient, diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index f2eafcbf839..0e69de87ce2 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -3289,6 +3289,20 @@ async def delete_team( await prisma_client.delete_data(team_id_list=data.team_ids, table_name="key") + ## DELETE ASSOCIATED BYOK MODELS + # Runs before the team rows are deleted so a mid-flight failure never leaves + # the team gone with its models orphaned. + from litellm.proxy.management_endpoints.model_management_endpoints import ( + delete_team_models, + ) + from litellm.proxy.proxy_server import llm_router + + await delete_team_models( + team_ids=data.team_ids, + prisma_client=prisma_client, + llm_router=llm_router, + ) + # ## DELETE TEAM MEMBERSHIPS for team_row in team_rows: ### get all team members diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 2ba604e5da6..6a81b1b613b 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -24,6 +24,7 @@ from litellm.proxy.management_endpoints.model_management_endpoints import ( ModelManagementAuthChecks, _get_team_deployments, clear_cache, + delete_team_models, ) from litellm.proxy.utils import PrismaClient from litellm.types.router import Deployment, LiteLLM_Params, updateDeployment @@ -2184,6 +2185,148 @@ class TestGetTeamDeployments: assert result[0] is dep1 +def _model_row(model_id: str, team_id: str): + row = MagicMock() + row.model_id = model_id + row.model_name = f"model_name_{team_id}_{model_id}" + row.model_info = {"team_id": team_id} + return row + + +class _TxProxyModelTable: + """Transactional proxy-model table that records the order of DB writes.""" + + def __init__(self, rows, events): + self._rows = list(rows) + self.events = events + + async def find_many(self, where): + prefix = where["model_name"]["startswith"] + return [r for r in self._rows if r.model_name.startswith(prefix)] + + async def delete_many(self, where): + ids = list(where["model_id"]["in"]) + self.events.append(("delete_many", tuple(ids))) + self._rows = [r for r in self._rows if r.model_id not in ids] + return len(ids) + + +class _TxPrismaClient: + """Minimal prisma stub whose ``db.tx()`` yields a transaction and records commit.""" + + def __init__(self, rows): + self.events: list = [] + self._table = _TxProxyModelTable(rows, self.events) + tx = MagicMock() + tx.litellm_proxymodeltable = self._table + outer = self + + class _TxCM: + async def __aenter__(self): + return tx + + async def __aexit__(self, *exc): + outer.events.append(("commit",)) + return False + + self.db = MagicMock() + self.db.tx = MagicMock(return_value=_TxCM()) + + +class _RecordingRouter: + def __init__(self, events): + self.events = events + self.deleted: list = [] + + def delete_deployment(self, id): # noqa: A002 - matches router signature + self.events.append(("router", id)) + self.deleted.append(id) + + +class TestDeleteTeamModels: + """delete_team_models must remove every team's BYOK models in one transaction + and sync the in-memory router only after that transaction commits.""" + + @pytest.mark.asyncio + async def test_deletes_all_teams_models_and_syncs_router(self): + rows = [_model_row("a1", "team_a"), _model_row("b1", "team_b")] + prisma = _TxPrismaClient(rows) + router = _RecordingRouter(prisma.events) + + deleted = await delete_team_models( + team_ids=["team_a", "team_b"], + prisma_client=prisma, + llm_router=router, + ) + + assert sorted(deleted) == ["a1", "b1"] + assert sorted(router.deleted) == ["a1", "b1"] + + @pytest.mark.asyncio + async def test_router_sync_happens_after_commit(self): + """Race-safety: the router is touched only once the DB transaction has + committed, so a rollback can never leave a deployment without its row.""" + rows = [_model_row("a1", "team_a"), _model_row("b1", "team_b")] + prisma = _TxPrismaClient(rows) + router = _RecordingRouter(prisma.events) + + await delete_team_models( + team_ids=["team_a", "team_b"], prisma_client=prisma, llm_router=router + ) + + commit_idx = prisma.events.index(("commit",)) + router_indices = [i for i, e in enumerate(prisma.events) if e[0] == "router"] + delete_indices = [ + i for i, e in enumerate(prisma.events) if e[0] == "delete_many" + ] + assert router_indices, "router was never synced" + assert all(i > commit_idx for i in router_indices) + assert all(i < commit_idx for i in delete_indices) + + @pytest.mark.asyncio + async def test_only_owning_team_models_deleted(self): + """A row sharing the prefix but a different model_info.team_id is left alone.""" + mine = _model_row("a1", "team_a") + intruder = MagicMock() + intruder.model_id = "x9" + intruder.model_name = "model_name_team_a_x9" + intruder.model_info = {"team_id": "someone_else"} + prisma = _TxPrismaClient([mine, intruder]) + router = _RecordingRouter(prisma.events) + + deleted = await delete_team_models( + team_ids=["team_a"], prisma_client=prisma, llm_router=router + ) + + assert deleted == ["a1"] + assert router.deleted == ["a1"] + + @pytest.mark.asyncio + async def test_no_models_no_writes(self): + prisma = _TxPrismaClient([]) + router = _RecordingRouter(prisma.events) + + deleted = await delete_team_models( + team_ids=["team_a"], prisma_client=prisma, llm_router=router + ) + + assert deleted == [] + assert router.deleted == [] + assert not any(e[0] == "delete_many" for e in prisma.events) + + @pytest.mark.asyncio + async def test_missing_router_is_safe(self): + rows = [_model_row("a1", "team_a")] + prisma = _TxPrismaClient(rows) + + deleted = await delete_team_models( + team_ids=["team_a"], prisma_client=prisma, llm_router=None + ) + + assert deleted == ["a1"] + assert any(e[0] == "delete_many" for e in prisma.events) + + def _build_db_model_for_blocked_test(): from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 400ed287802..06adcf80707 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -6350,6 +6350,14 @@ async def test_delete_team_persists_deleted_teams(monkeypatch): mock_find_many_keys = AsyncMock(return_value=[]) mock_prisma_client.db.litellm_verificationtoken.find_many = mock_find_many_keys + # delete_team now deletes team BYOK models inside a transaction; this team has none. + mock_tx = AsyncMock() + mock_tx.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_tx_cm = MagicMock() + mock_tx_cm.__aenter__ = AsyncMock(return_value=mock_tx) + mock_tx_cm.__aexit__ = AsyncMock(return_value=False) + mock_prisma_client.db.tx = MagicMock(return_value=mock_tx_cm) + monkeypatch.setattr( "litellm.proxy.proxy_server.prisma_client", mock_prisma_client, @@ -8499,6 +8507,8 @@ async def test_new_team_encrypts_callback_vars( assert cv["langfuse_secret_key"] != "sk-real" recovered = decrypt_callback_vars(metadata)["logging"][0]["callback_vars"] assert recovered["langfuse_secret_key"] == "sk-real" + + def _non_admin_auth(): return UserAPIKeyAuth( user_id="u-team-admin", user_role=LitellmUserRoles.INTERNAL_USER From 411bd3da5ba6abb9b2aae173d3274317e5805364 Mon Sep 17 00:00:00 2001 From: milan-berri Date: Tue, 9 Jun 2026 02:59:21 +0300 Subject: [PATCH 18/22] feat(vantage): include organization metadata in FOCUS Tags export (#28184) * feat(vantage): include organization metadata in FOCUS Tags export Join LiteLLM_OrganizationTable when building Vantage/FOCUS export rows so organization_id and organization_alias appear in Tags for org-level filtering. Co-authored-by: Cursor * test(focus): include api_requests in organization Tags tests FocusTransformer now requires api_requests after staging merge; add the column to test fixtures so integrations CI can run the Tags assertions. Co-authored-by: Cursor --------- Co-authored-by: Cursor --- litellm/integrations/focus/database.py | 6 +- litellm/integrations/focus/transformer.py | 2 + .../integrations/focus/test_focus_database.py | 15 +++++ .../integrations/focus/test_transformer.py | 62 +++++++++++++++++++ 4 files changed, 84 insertions(+), 1 deletion(-) create mode 100644 tests/test_litellm/integrations/focus/test_transformer.py diff --git a/litellm/integrations/focus/database.py b/litellm/integrations/focus/database.py index 298254670eb..3ae3f6b53ac 100644 --- a/litellm/integrations/focus/database.py +++ b/litellm/integrations/focus/database.py @@ -80,11 +80,15 @@ class FocusLiteLLMDatabase: vt.team_id, vt.key_alias as api_key_alias, tt.team_alias, - ut.user_email as user_email + ut.user_email as user_email, + COALESCE(vt.organization_id, tt.organization_id) as organization_id, + ot.organization_alias as organization_alias FROM "LiteLLM_DailyUserSpend" dus LEFT JOIN "LiteLLM_VerificationToken" vt ON dus.api_key = vt.token LEFT JOIN "LiteLLM_TeamTable" tt ON vt.team_id = tt.team_id LEFT JOIN "LiteLLM_UserTable" ut ON dus.user_id = ut.user_id + LEFT JOIN "LiteLLM_OrganizationTable" ot + ON ot.organization_id = COALESCE(vt.organization_id, tt.organization_id) {where_clause} ORDER BY dus.date DESC, dus.created_at DESC {limit_clause} diff --git a/litellm/integrations/focus/transformer.py b/litellm/integrations/focus/transformer.py index 8496b7ec159..a17df29b912 100644 --- a/litellm/integrations/focus/transformer.py +++ b/litellm/integrations/focus/transformer.py @@ -12,6 +12,8 @@ from .schema import FOCUS_NORMALIZED_SCHEMA _TAG_KEYS = ( "team_id", "team_alias", + "organization_id", + "organization_alias", "user_id", "user_email", "api_key_alias", diff --git a/tests/test_litellm/integrations/focus/test_focus_database.py b/tests/test_litellm/integrations/focus/test_focus_database.py index 5ee98cc9dd0..d77af2dd170 100644 --- a/tests/test_litellm/integrations/focus/test_focus_database.py +++ b/tests/test_litellm/integrations/focus/test_focus_database.py @@ -72,3 +72,18 @@ async def test_should_reject_invalid_limit(monkeypatch: pytest.MonkeyPatch): await db.get_usage_data(limit="invalid") assert query_mock.await_count == 0 + + +@pytest.mark.asyncio +async def test_should_join_organization_table(monkeypatch: pytest.MonkeyPatch): + db, query_mock = _setup_db(monkeypatch, []) + + await db.get_usage_data() + + query_text, *_ = query_mock.await_args.args + assert ( + "COALESCE(vt.organization_id, tt.organization_id) as organization_id" + in query_text + ) + assert "ot.organization_alias as organization_alias" in query_text + assert 'LEFT JOIN "LiteLLM_OrganizationTable" ot' in query_text diff --git a/tests/test_litellm/integrations/focus/test_transformer.py b/tests/test_litellm/integrations/focus/test_transformer.py new file mode 100644 index 00000000000..4461d19efde --- /dev/null +++ b/tests/test_litellm/integrations/focus/test_transformer.py @@ -0,0 +1,62 @@ +"""Tests for FocusTransformer organization metadata in Tags.""" + +from __future__ import annotations + +import json +from datetime import date + +import polars as pl + +from litellm.integrations.focus.transformer import FocusTransformer + + +def test_should_include_organization_fields_in_tags(): + frame = pl.DataFrame( + { + "date": [date(2024, 1, 2)], + "spend": [1.25], + "api_requests": [1], + "api_key": ["hashed-key"], + "api_key_alias": ["prod-key"], + "model": ["gpt-4o"], + "model_group": ["gpt-4o"], + "custom_llm_provider": ["openai"], + "team_id": ["team-1"], + "team_alias": ["Platform"], + "organization_id": ["org-123"], + "organization_alias": ["Acme Corp"], + "user_id": ["user-1"], + "user_email": ["user@example.com"], + } + ) + + normalized = FocusTransformer().transform(frame) + + tags = json.loads(normalized["Tags"][0]) + assert tags["organization_id"] == "org-123" + assert tags["organization_alias"] == "Acme Corp" + assert tags["team_id"] == "team-1" + + +def test_should_omit_missing_organization_fields_from_tags(): + frame = pl.DataFrame( + { + "date": [date(2024, 1, 2)], + "spend": [0.5], + "api_requests": [1], + "api_key": ["hashed-key"], + "api_key_alias": ["prod-key"], + "model": ["gpt-4o-mini"], + "model_group": ["gpt-4o-mini"], + "custom_llm_provider": ["openai"], + "team_id": ["team-1"], + "team_alias": ["Platform"], + } + ) + + normalized = FocusTransformer().transform(frame) + + tags = json.loads(normalized["Tags"][0]) + assert "organization_id" not in tags + assert "organization_alias" not in tags + assert tags["team_id"] == "team-1" From 1bbaf1c39dda367f5a2b4b6b9ab4cac46d71ab14 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Mon, 8 Jun 2026 17:46:28 -0700 Subject: [PATCH 19/22] fix(guardrails): read CrowdStrike AIDR identity from both metadata bags (#29991) Capture user_id and extra_info from metadata or litellm_metadata. The single-bag read dropped identity whenever a request carried a present litellm_metadata field (null or a user-supplied dict), since /chat/completions routes the authenticated identity into metadata while the guardrail read litellm_metadata first --- .../crowdstrike_aidr/crowdstrike_aidr.py | 14 +++++- .../guardrail_hooks/test_crowdstrike_aidr.py | 50 +++++++++++++++++++ 2 files changed, 62 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py index d1ef165b46e..248202b644c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py +++ b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py @@ -105,6 +105,16 @@ def _extract_text_from_content(content: object) -> str: return "" +def _merge_metadata_bags(request_data: Mapping[str, Any]) -> Optional[dict[str, Any]]: + merged: dict[str, Any] = {} + present = False + for bag in (request_data.get("metadata"), request_data.get("litellm_metadata")): + if isinstance(bag, Mapping): + present = True + merged.update(bag) + return merged if present else None + + class CrowdStrikeAIDRHandler(CustomGuardrail): """ CrowdStrike AIDR AI Guardrail handler to interact with the CrowdStrike AIDR @@ -321,8 +331,8 @@ class CrowdStrikeAIDRHandler(CustomGuardrail): if model: ai_guard_payload["model"] = model - metadata = request_data.get("litellm_metadata", request_data.get("metadata")) - if isinstance(metadata, Mapping): + metadata = _merge_metadata_bags(request_data) + if metadata is not None: user_id = metadata.get("user_api_key_user_id") if user_id: ai_guard_payload["user_id"] = user_id diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py index e7f72ff7a3f..f8fd9a0a185 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_crowdstrike_aidr.py @@ -529,6 +529,56 @@ async def test_apply_guardrail_no_metadata_skips_user_fields( assert "extra_info" not in payload +@pytest.mark.asyncio +@pytest.mark.parametrize( + "litellm_metadata, metadata", + [ + (None, {"user_api_key_user_id": "uid-abc", "user_api_key_user_email": "alice@example.com"}), + ({"trace_id": "t1"}, {"user_api_key_user_id": "uid-abc", "user_api_key_user_email": "alice@example.com"}), + (["unexpected"], {"user_api_key_user_id": "uid-abc", "user_api_key_user_email": "alice@example.com"}), + ({"user_api_key_user_id": "uid-abc", "user_api_key_user_email": "alice@example.com"}, {"trace_id": "t1"}), + ], + ids=["identity_in_metadata_llm_none", "identity_in_metadata_llm_user_dict", "identity_in_metadata_llm_non_mapping", "identity_in_litellm_metadata"], +) +async def test_apply_guardrail_reads_identity_from_either_metadata_bag( + crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler, + litellm_metadata, + metadata, +) -> None: + inputs: GenericGuardrailAPIInputs = { + "texts": ["Hello"], + "structured_messages": [{"role": "user", "content": "Hello"}], + "model": "gpt-4o", + } + request_data = { + "messages": inputs["structured_messages"], + "model": "gpt-4o", + "litellm_metadata": litellm_metadata, + "metadata": metadata, + } + guardrail_endpoint = ( + f"{crowdstrike_aidr_guardrail.api_base}/v1/guard_chat_completions" + ) + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + return_value=httpx.Response( + status_code=200, + json={"result": {"blocked": False, "transformed": False}}, + request=httpx.Request(method="POST", url=guardrail_endpoint), + ), + ) as mock_method: + await crowdstrike_aidr_guardrail.apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type="request", + ) + + payload = mock_method.call_args.kwargs["json"] + assert payload["user_id"] == "uid-abc" + assert payload["extra_info"] == {"user_name": "alice@example.com"} + + @pytest.mark.asyncio async def test_apply_guardrail_request_skipped_messages_stay_aligned( crowdstrike_aidr_guardrail: CrowdStrikeAIDRHandler, From 92817cb65bc351b71543e6f980850953f7c1511e Mon Sep 17 00:00:00 2001 From: tin-berri Date: Mon, 8 Jun 2026 18:13:06 -0700 Subject: [PATCH 20/22] changing expires_in default to use actual slack return details (#29951) --- .../mcp_server/discoverable_endpoints.py | 7 +- .../mcp_server/test_discoverable_endpoints.py | 72 +++++++++++++++++++ 2 files changed, 76 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index ed374635fea..3beddd2c435 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -512,12 +512,13 @@ async def exchange_token_with_server( result = { "access_token": access_token, "token_type": token_response.get("token_type", "Bearer"), - "expires_in": token_response.get("expires_in", 3600), } - if "refresh_token" in token_response and token_response["refresh_token"]: + if token_response.get("expires_in") is not None: + result["expires_in"] = token_response["expires_in"] + if token_response.get("refresh_token"): result["refresh_token"] = token_response["refresh_token"] - if "scope" in token_response and token_response["scope"]: + if token_response.get("scope"): result["scope"] = token_response["scope"] # RFC 6749 §5.1: token responses must not be cached. diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index da66d60aed8..6fd935e3364 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -1,5 +1,6 @@ """Tests for MCP OAuth discoverable endpoints""" +import json from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -2661,3 +2662,74 @@ async def test_token_endpoint_sets_no_store_cache_control(): assert response.headers["cache-control"] == "no-store" assert response.headers["pragma"] == "no-cache" + + +async def _exchange_with_upstream_token_response(upstream_body): + from fastapi import Request + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + exchange_token_with_server, + ) + from litellm.proxy._types import MCPTransport + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="t", + name="t", + server_name="t", + alias="t", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="cid", + client_secret="cs", + authorization_url="https://provider.com/oauth/authorize", + token_url="https://provider.com/oauth/token", + ) + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://litellm.example.com/" + mock_request.headers = {} + + fake_http_response = MagicMock() + fake_http_response.json.return_value = upstream_body + fake_http_response.raise_for_status = MagicMock() + fake_http_client = MagicMock() + fake_http_client.post = AsyncMock(return_value=fake_http_response) + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=fake_http_client, + ): + response = await exchange_token_with_server( + request=mock_request, + mcp_server=server, + grant_type="authorization_code", + code="c", + redirect_uri="http://127.0.0.1:3000/cb", + client_id="cid", + client_secret=None, + code_verifier=None, + ) + return json.loads(response.body) + + +@pytest.mark.asyncio +async def test_token_exchange_omits_expires_in_when_upstream_omits_it(): + """A provider that issues a non-expiring token (e.g. Slack without token + rotation) returns no ``expires_in``. The exchange must mirror that and omit + ``expires_in`` rather than fabricate a 1-hour TTL, so the stored credential + is treated as non-expiring instead of dying after an hour.""" + body = await _exchange_with_upstream_token_response( + {"access_token": "tok", "token_type": "Bearer"} + ) + assert "expires_in" not in body + + +@pytest.mark.asyncio +async def test_token_exchange_passes_through_upstream_expires_in(): + """When the provider does send ``expires_in`` (e.g. Slack with token + rotation), the exchange forwards the real value unchanged.""" + body = await _exchange_with_upstream_token_response( + {"access_token": "tok", "token_type": "Bearer", "expires_in": 43200} + ) + assert body["expires_in"] == 43200 From 424db6a9807d6e830492825598332f9bb7f234c6 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 9 Jun 2026 06:57:04 +0530 Subject: [PATCH 21/22] feat(azure_ai): add MAI-Image-2.5 image generation support (#29688) * feat(azure_ai): add MAI-Image-2.5 image generation support Route azure_ai MAI models to /mai/v1/images/generations and map OpenAI size to width/height for the serverless API. Co-authored-by: Cursor * fix(azure_ai): address MAI image generation review feedback Validate unsupported size values, default width/height independently, add MAI-Image-2.5 pricing, and expand test coverage. @greptileai Co-authored-by: Cursor * feat(azure_ai): add MAI image edit and expand model cost map Add MAI image edit support with usage normalization for Azure response format, and register MAI-Image-2.5-Flash and MAI-Image-2e pricing in the model map. Co-authored-by: Cursor * fix(azure_ai): validate MAI edit size by consuming map iterator Greptile: lazy map() never evaluated int() so values like 1024xabc passed through. Co-authored-by: Cursor * fix(azure_ai): normalize MAI usage in generation response handler Apply normalize_mai_image_usage before building ImageResponse so token-based cost calculation works when Azure returns num_output_tokens fields. Co-authored-by: Cursor * fix(azure_ai): narrow MAI edit size param type for mypy Co-authored-by: Cursor * Fix Azure MAI image response handling * Fix MAI image generation base model routing * fix(azure_ai): preserve zero num_output_tokens in MAI usage normalization * fix(azure_ai): wrap MAI generation response JSON parsing in error handling * fix(azure_ai): build MAI image edit URL correctly for /mai/ root bases * fix(azure_ai): build MAI image generation URL correctly for /mai/ root bases --------- Co-authored-by: Cursor Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/llms/azure/azure.py | 39 +- .../llms/azure/image_generation/__init__.py | 3 + litellm/llms/azure_ai/image_edit/__init__.py | 14 +- .../azure_ai/image_edit/mai_transformation.py | 199 +++++++++ .../azure_ai/image_generation/__init__.py | 4 + .../image_generation/cost_calculator.py | 26 +- .../image_generation/mai_transformation.py | 236 +++++++++++ ...odel_prices_and_context_window_backup.json | 37 ++ model_prices_and_context_window.json | 37 ++ .../test_azure_image_generation_init.py | 16 + .../test_mai_image_edit_transformation.py | 171 ++++++++ .../test_mai_image_generation.py | 380 ++++++++++++++++++ 12 files changed, 1149 insertions(+), 13 deletions(-) create mode 100644 litellm/llms/azure_ai/image_edit/mai_transformation.py create mode 100644 litellm/llms/azure_ai/image_generation/mai_transformation.py create mode 100644 tests/test_litellm/llms/azure_ai/image_edit/test_mai_image_edit_transformation.py create mode 100644 tests/test_litellm/llms/azure_ai/image_generation/test_mai_image_generation.py diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 734b8ecef16..56cf035d0f7 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -43,7 +43,10 @@ from .common_utils import ( process_azure_headers, select_azure_base_url_or_endpoint, ) -from .image_generation import get_azure_image_generation_config +from .image_generation import ( + AzureFoundryMAIImageGenerationConfig, + get_azure_image_generation_config, +) from .image_generation.http_utils import azure_deployment_image_generation_json_body @@ -1097,10 +1100,14 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): ) def create_azure_base_url( - self, azure_client_params: dict, model: Optional[str] + self, + azure_client_params: dict, + model: Optional[str], + base_model: Optional[str] = None, ) -> str: from litellm.llms.azure_ai.image_generation import ( AzureFoundryFluxImageGenerationConfig, + AzureFoundryMAIImageGenerationConfig, ) api_base: str = azure_client_params.get( @@ -1112,6 +1119,12 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): if model is None: model = "" + if AzureFoundryMAIImageGenerationConfig.is_mai_model(base_model or model): + return AzureFoundryMAIImageGenerationConfig.get_mai_image_generation_url( + api_base=api_base, + api_version=api_version, + ) + # Handle FLUX 2 models on Azure AI which use a different URL pattern # e.g., /providers/blackforestlabs/v1/flux-2-pro instead of /openai/deployments/{model}/images/generations if AzureFoundryFluxImageGenerationConfig.is_flux2_model(model): @@ -1153,10 +1166,10 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): if api_base.endswith("/"): api_base = api_base.rstrip("/") api_version: str = azure_client_params.get("api_version", "") - # Use the deployment name (model) for URL construction, not the base_model from data img_gen_api_base = self.create_azure_base_url( azure_client_params=azure_client_params, model=model or data.get("model", ""), + base_model=data.get("model", ""), ) ## LOGGING @@ -1285,9 +1298,10 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): if aimg_generation is True: return self.aimage_generation(data=data, input=input, logging_obj=logging_obj, model_response=model_response, api_key=api_key, client=client, azure_client_params=azure_client_params, timeout=timeout, headers=headers, model=model) # type: ignore - # Use the deployment name (model) for URL construction, not the base_model from data img_gen_api_base = self.create_azure_base_url( - azure_client_params=azure_client_params, model=model + azure_client_params=azure_client_params, + model=model, + base_model=base_model, ) ## LOGGING @@ -1309,6 +1323,21 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): data=data, headers=headers, ) + provider_config = get_azure_image_generation_config( + data.get("model", "dall-e-2") + ) + if isinstance(provider_config, AzureFoundryMAIImageGenerationConfig): + return provider_config.transform_image_generation_response( + model=data.get("model", "dall-e-2"), + raw_response=httpx_response, + model_response=model_response or ImageResponse(), + logging_obj=logging_obj, + request_data=data, + optional_params=data, + litellm_params=data, + encoding=litellm.encoding, + ) + response = httpx_response.json() ## LOGGING diff --git a/litellm/llms/azure/image_generation/__init__.py b/litellm/llms/azure/image_generation/__init__.py index f60e446f0c4..64636bc689d 100644 --- a/litellm/llms/azure/image_generation/__init__.py +++ b/litellm/llms/azure/image_generation/__init__.py @@ -1,4 +1,5 @@ from litellm._logging import verbose_logger +from litellm.llms.azure_ai.image_generation import AzureFoundryMAIImageGenerationConfig from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, ) @@ -24,6 +25,8 @@ def get_azure_image_generation_config(model: str) -> BaseImageGenerationConfig: return AzureDallE2ImageGenerationConfig() elif "dalle3" in model: return AzureDallE3ImageGenerationConfig() + elif AzureFoundryMAIImageGenerationConfig.is_mai_model(model): + return AzureFoundryMAIImageGenerationConfig() else: verbose_logger.debug( f"Using AzureGPTImageGenerationConfig for model: {model}. This follows the gpt-image model format." diff --git a/litellm/llms/azure_ai/image_edit/__init__.py b/litellm/llms/azure_ai/image_edit/__init__.py index e3acd610446..42ece6d19ec 100644 --- a/litellm/llms/azure_ai/image_edit/__init__.py +++ b/litellm/llms/azure_ai/image_edit/__init__.py @@ -1,21 +1,33 @@ from litellm.llms.azure_ai.image_generation.flux_transformation import ( AzureFoundryFluxImageGenerationConfig, ) +from litellm.llms.azure_ai.image_generation.mai_transformation import ( + AzureFoundryMAIImageGenerationConfig, +) from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig from .flux2_transformation import AzureFoundryFlux2ImageEditConfig +from .mai_transformation import AzureFoundryMAIImageEditConfig from .transformation import AzureFoundryFluxImageEditConfig -__all__ = ["AzureFoundryFluxImageEditConfig", "AzureFoundryFlux2ImageEditConfig"] +__all__ = [ + "AzureFoundryFluxImageEditConfig", + "AzureFoundryFlux2ImageEditConfig", + "AzureFoundryMAIImageEditConfig", +] def get_azure_ai_image_edit_config(model: str) -> BaseImageEditConfig: """ Get the appropriate image edit config for an Azure AI model. + - MAI models use /mai/v1/images/edits with multipart form data and size - FLUX 2 models use JSON with base64 image - FLUX 1 models use multipart/form-data """ + if AzureFoundryMAIImageGenerationConfig.is_mai_model(model): + return AzureFoundryMAIImageEditConfig() + # Check if it's a FLUX 2 model if AzureFoundryFluxImageGenerationConfig.is_flux2_model(model): return AzureFoundryFlux2ImageEditConfig() diff --git a/litellm/llms/azure_ai/image_edit/mai_transformation.py b/litellm/llms/azure_ai/image_edit/mai_transformation.py new file mode 100644 index 00000000000..75bfc913a8f --- /dev/null +++ b/litellm/llms/azure_ai/image_edit/mai_transformation.py @@ -0,0 +1,199 @@ +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast + +import httpx +from httpx._types import RequestFiles + +from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo +from litellm.llms.azure_ai.image_generation.mai_transformation import ( + AzureFoundryMAIImageGenerationConfig, +) +from litellm.llms.openai.common_utils import OpenAIError +from litellm.llms.openai.image_edit.transformation import OpenAIImageEditConfig +from litellm.secret_managers.main import get_secret_str +from litellm.types.images.main import ImageEditOptionalRequestParams +from litellm.types.llms.openai import FileTypes +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import ImageResponse +from litellm.utils import convert_to_model_response_object + +if TYPE_CHECKING: + from litellm.litellm_core_utils.logging import Logging as LiteLLMLoggingObj + + +class AzureFoundryMAIImageEditConfig(OpenAIImageEditConfig): + """Azure AI Foundry MAI image editing (e.g. MAI-Image-2.5).""" + + DEFAULT_SIZE = "1024x1024" + + def get_supported_openai_params(self, model: str) -> list: + return ["prompt", "image", "model", "n", "size"] + + def map_openai_params( + self, + image_edit_optional_params: ImageEditOptionalRequestParams, + model: str, + drop_params: bool, + ) -> Dict: + optional_params: Dict[str, Any] = {} + supported_params = self.get_supported_openai_params(model) + + for key, value in dict(image_edit_optional_params).items(): + if value is None or key in optional_params: + continue + + if key in supported_params: + if key == "size" and value: + size_param = cast(str, value) + self._validate_size_param(size_param) + optional_params[key] = size_param + else: + optional_params[key] = value + elif not drop_params: + raise ValueError( + f"Parameter {key} is not supported for model {model}. " + f"Supported parameters are {supported_params}. " + f"Set drop_params=True to drop unsupported parameters." + ) + + if "size" not in optional_params: + optional_params["size"] = self.DEFAULT_SIZE + + return optional_params + + def _validate_size_param(self, size: str) -> None: + known_sizes = { + "1024x1024", + "1792x1024", + "1024x1792", + "512x512", + "256x256", + } + + if size in known_sizes: + return + + if "x" in size: + try: + tuple(map(int, size.lower().split("x", 1))) + return + except ValueError: + raise ValueError( + f"Invalid size format: '{size}'. Expected format 'WIDTHxHEIGHT' (e.g., '1024x1024')." + ) + + raise ValueError( + f"Unsupported size value: '{size}'. " + f"Use a known size (e.g., '1024x1024') or a custom 'WIDTHxHEIGHT' string." + ) + + def validate_environment( + self, + headers: dict, + model: str, + api_key: Optional[str] = None, + litellm_params: Optional[dict] = None, + api_base: Optional[str] = None, + ) -> dict: + api_key = AzureFoundryModelInfo.get_api_key(api_key) + + if not api_key: + raise ValueError( + f"Azure AI API key is required for model {model}. " + "Set AZURE_AI_API_KEY environment variable or pass api_key parameter." + ) + + headers.update({"api-key": api_key}) + return headers + + def get_complete_url( + self, + model: str, + api_base: Optional[str], + litellm_params: dict, + ) -> str: + api_base = AzureFoundryModelInfo.get_api_base(api_base) + + if api_base is None: + raise ValueError( + "Azure AI API base is required. Set AZURE_AI_API_BASE environment variable or pass api_base parameter." + ) + + api_version = ( + litellm_params.get("api_version") + or get_secret_str("AZURE_AI_API_VERSION") + or "preview" + ) + + return AzureFoundryMAIImageGenerationConfig.get_mai_image_edit_url( + api_base=api_base, + api_version=api_version, + ) + + def transform_image_edit_request( + self, + model: str, + prompt: Optional[str], + image: Optional[FileTypes], + image_edit_optional_request_params: Dict, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[Dict, RequestFiles]: + request_params = { + "model": model, + **image_edit_optional_request_params, + } + if prompt is not None: + request_params["prompt"] = prompt + + data_without_files = { + key: value + for key, value in request_params.items() + if key not in ["image", "mask"] + } + files_list: List[Tuple[str, Any]] = [] + + if image is not None: + image_list = [image] if not isinstance(image, list) else image + for _image in image_list: + if _image is not None: + self._add_image_to_files( + files_list=files_list, + image=_image, + field_name="image", + ) + break + + return data_without_files, files_list + + def transform_image_edit_response( + self, + model: str, + raw_response: httpx.Response, + logging_obj: "LiteLLMLoggingObj", + ) -> ImageResponse: + try: + response = raw_response.json() + except Exception: + raise OpenAIError( + message=raw_response.text, status_code=raw_response.status_code + ) + + if "usage" in response: + response["usage"] = ( + AzureFoundryMAIImageGenerationConfig.normalize_mai_image_usage( + response.get("usage") + ) + ) + + logging_obj.post_call( + input="", + api_key="", + additional_args={"complete_input_dict": {}}, + original_response=response, + ) + + return convert_to_model_response_object( + response_object=response, + model_response_object=ImageResponse(), + response_type="image_generation", + ) diff --git a/litellm/llms/azure_ai/image_generation/__init__.py b/litellm/llms/azure_ai/image_generation/__init__.py index cebab3de16e..70821d5d764 100644 --- a/litellm/llms/azure_ai/image_generation/__init__.py +++ b/litellm/llms/azure_ai/image_generation/__init__.py @@ -7,12 +7,14 @@ from .dall_e_2_transformation import AzureFoundryDallE2ImageGenerationConfig from .dall_e_3_transformation import AzureFoundryDallE3ImageGenerationConfig from .flux_transformation import AzureFoundryFluxImageGenerationConfig from .gpt_transformation import AzureFoundryGPTImageGenerationConfig +from .mai_transformation import AzureFoundryMAIImageGenerationConfig __all__ = [ "AzureFoundryFluxImageGenerationConfig", "AzureFoundryGPTImageGenerationConfig", "AzureFoundryDallE2ImageGenerationConfig", "AzureFoundryDallE3ImageGenerationConfig", + "AzureFoundryMAIImageGenerationConfig", ] @@ -24,6 +26,8 @@ def get_azure_ai_image_generation_config(model: str) -> BaseImageGenerationConfi return AzureFoundryDallE2ImageGenerationConfig() elif "dalle3" in model: return AzureFoundryDallE3ImageGenerationConfig() + elif AzureFoundryMAIImageGenerationConfig.is_mai_model(model): + return AzureFoundryMAIImageGenerationConfig() elif "flux" in model: return AzureFoundryFluxImageGenerationConfig() else: diff --git a/litellm/llms/azure_ai/image_generation/cost_calculator.py b/litellm/llms/azure_ai/image_generation/cost_calculator.py index b67de9cb70d..f8c876bb5be 100644 --- a/litellm/llms/azure_ai/image_generation/cost_calculator.py +++ b/litellm/llms/azure_ai/image_generation/cost_calculator.py @@ -1,6 +1,9 @@ from typing import Any import litellm +from litellm.litellm_core_utils.llm_cost_calc.utils import ( + calculate_image_response_cost_from_usage, +) from litellm.types.utils import ImageResponse @@ -9,19 +12,28 @@ def cost_calculator( image_response: Any, ) -> float: """ - Recraft image generation cost calculator + Azure AI image generation cost calculator """ _model_info = litellm.get_model_info( model=model, custom_llm_provider=litellm.LlmProviders.AZURE_AI.value, ) - output_cost_per_image: float = _model_info.get("output_cost_per_image") or 0.0 - num_images: int = 0 + if isinstance(image_response, ImageResponse): + token_based_cost = calculate_image_response_cost_from_usage( + model=model, + image_response=image_response, + custom_llm_provider=litellm.LlmProviders.AZURE_AI.value, + ) + if token_based_cost is not None: + return token_based_cost + + output_cost_per_image: float = _model_info.get("output_cost_per_image") or 0.0 + num_images: int = 0 if image_response.data: num_images = len(image_response.data) return output_cost_per_image * num_images - else: - raise ValueError( - f"image_response must be of type ImageResponse got type={type(image_response)}" - ) + + raise ValueError( + f"image_response must be of type ImageResponse got type={type(image_response)}" + ) diff --git a/litellm/llms/azure_ai/image_generation/mai_transformation.py b/litellm/llms/azure_ai/image_generation/mai_transformation.py new file mode 100644 index 00000000000..071ca9d9895 --- /dev/null +++ b/litellm/llms/azure_ai/image_generation/mai_transformation.py @@ -0,0 +1,236 @@ +from typing import TYPE_CHECKING, Any, Dict, List, Optional + +import httpx + +from litellm.llms.base_llm.image_generation.transformation import ( + BaseImageGenerationConfig, +) +from litellm.llms.openai.common_utils import OpenAIError +from litellm.types.llms.openai import OpenAIImageGenerationOptionalParams +from litellm.types.utils import ImageResponse +from litellm.utils import convert_to_model_response_object + +if TYPE_CHECKING: + from litellm.litellm_core_utils.logging import Logging as LiteLLMLoggingObj + + +class AzureFoundryMAIImageGenerationConfig(BaseImageGenerationConfig): + """Azure AI Foundry MAI image generation (e.g. MAI-Image-2.5).""" + + DEFAULT_WIDTH = 1024 + DEFAULT_HEIGHT = 1024 + + @staticmethod + def get_mai_image_generation_url( + api_base: Optional[str], + api_version: Optional[str], + ) -> str: + if api_base is None: + raise ValueError("api_base is required for Azure AI MAI image generation") + + api_version = api_version or "preview" + path, separator, query = api_base.partition("?") + path = path.rstrip("/") + + if "/mai/" in path: + prefix, _, _ = path.partition("/images/") + path = f"{prefix}/images/generations" + else: + path = f"{path}/mai/v1/images/generations" + + if separator: + return f"{path}?{query}" + return f"{path}?api-version={api_version}" + + @staticmethod + def get_mai_image_edit_url( + api_base: Optional[str], + api_version: Optional[str], + ) -> str: + if api_base is None: + raise ValueError("api_base is required for Azure AI MAI image editing") + + api_version = api_version or "preview" + path, separator, query = api_base.partition("?") + path = path.rstrip("/") + + if "/mai/" in path: + prefix, _, _ = path.partition("/images/") + path = f"{prefix}/images/edits" + else: + path = f"{path}/mai/v1/images/edits" + + if separator: + return f"{path}?{query}" + return f"{path}?api-version={api_version}" + + @staticmethod + def is_mai_model(model: str) -> bool: + model_normalized = model.lower().replace("-", "").replace("_", "") + return "maiimage" in model_normalized + + @staticmethod + def normalize_mai_image_usage(usage: Optional[Dict[str, Any]]) -> Dict[str, Any]: + """Map Azure MAI usage fields to OpenAI ImageUsage schema.""" + if usage is None: + return { + "input_tokens": 0, + "input_tokens_details": {"image_tokens": 0, "text_tokens": 0}, + "output_tokens": 0, + "total_tokens": 0, + } + + normalized_usage = dict(usage) + input_tokens_details = normalized_usage.get("input_tokens_details") + if not isinstance(input_tokens_details, dict): + input_tokens_details = {} + + text_tokens = normalized_usage.get("num_input_text_tokens") + if text_tokens is None: + text_tokens = input_tokens_details.get("text_tokens") + if text_tokens is None: + text_tokens = normalized_usage.get("input_tokens", 0) or 0 + + image_tokens = normalized_usage.get("num_input_image_tokens") + if image_tokens is None: + image_tokens = input_tokens_details.get("image_tokens") + if image_tokens is None: + image_tokens = 0 + + output_tokens = normalized_usage.get("output_tokens") + if output_tokens is None: + output_tokens = normalized_usage.get("num_output_tokens") + if output_tokens is None: + output_tokens = normalized_usage.get("output_image_tokens") + if output_tokens is None: + output_tokens = 0 + + input_tokens = normalized_usage.get("input_tokens") + if input_tokens is None: + input_tokens = text_tokens + image_tokens + + total_tokens = normalized_usage.get("total_tokens") + if total_tokens is None: + total_tokens = input_tokens + output_tokens + + normalized_usage.update( + { + "input_tokens": input_tokens, + "input_tokens_details": { + "image_tokens": image_tokens, + "text_tokens": text_tokens, + }, + "output_tokens": output_tokens, + "total_tokens": total_tokens, + } + ) + return normalized_usage + + def get_supported_openai_params( + self, model: str + ) -> List[OpenAIImageGenerationOptionalParams]: + return ["n", "size"] + + def map_openai_params( + self, + non_default_params: dict, + optional_params: dict, + model: str, + drop_params: bool, + ) -> dict: + supported_params = self.get_supported_openai_params(model) + + for k, v in non_default_params.items(): + if k in optional_params: + continue + + if k in supported_params: + if k == "size" and v: + self._map_size_param(v, optional_params) + else: + optional_params[k] = v + elif k in ("width", "height"): + optional_params[k] = v + elif not drop_params: + raise ValueError( + f"Parameter {k} is not supported for model {model}. " + f"Supported parameters are {supported_params} and width/height. " + f"Set drop_params=True to drop unsupported parameters." + ) + + if "width" not in optional_params: + optional_params["width"] = self.DEFAULT_WIDTH + if "height" not in optional_params: + optional_params["height"] = self.DEFAULT_HEIGHT + + optional_params.pop("size", None) + return optional_params + + def _map_size_param(self, size: str, optional_params: dict) -> None: + size_mapping = { + "1024x1024": (1024, 1024), + "1792x1024": (1792, 1024), + "1024x1792": (1024, 1792), + "512x512": (512, 512), + "256x256": (256, 256), + } + + if size in size_mapping: + width, height = size_mapping[size] + optional_params["width"] = width + optional_params["height"] = height + elif "x" in size: + try: + width, height = map(int, size.lower().split("x")) + optional_params["width"] = width + optional_params["height"] = height + except ValueError: + raise ValueError( + f"Invalid size format: '{size}'. Expected format 'WIDTHxHEIGHT' (e.g., '1024x1024')." + ) + else: + raise ValueError( + f"Unsupported size value: '{size}'. " + f"Use a known size (e.g., '1024x1024') or a custom 'WIDTHxHEIGHT' string." + ) + + def transform_image_generation_response( + self, + model: str, + raw_response: httpx.Response, + model_response: ImageResponse, + logging_obj: "LiteLLMLoggingObj", + request_data: dict, + optional_params: dict, + litellm_params: dict, + encoding: Any, + api_key: Optional[str] = None, + json_mode: Optional[bool] = None, + ) -> ImageResponse: + try: + response = raw_response.json() + except Exception: + raise OpenAIError( + message=raw_response.text, status_code=raw_response.status_code + ) + + if "usage" in response: + response["usage"] = self.normalize_mai_image_usage(response.get("usage")) + + logging_obj.post_call( + input=request_data.get("prompt", ""), + api_key=api_key, + additional_args={"complete_input_dict": request_data}, + original_response=response, + ) + + image_response: ImageResponse = convert_to_model_response_object( + response_object=response, + model_response_object=model_response, + response_type="image_generation", + ) + + width = optional_params.get("width", self.DEFAULT_WIDTH) + height = optional_params.get("height", self.DEFAULT_HEIGHT) + image_response.size = f"{width}x{height}" # type: ignore[assignment] + return image_response diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index e765512175b..282a292ab17 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -6889,6 +6889,43 @@ "/v1/images/generations" ] }, + "azure_ai/MAI-Image-2.5": { + "input_cost_per_image_token": 8e-06, + "input_cost_per_token": 5e-06, + "litellm_provider": "azure_ai", + "mode": "image_generation", + "output_cost_per_image": 0.05, + "output_cost_per_image_token": 4.7e-05, + "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/new-mai-models-in-microsoft-foundry-across-text-image-voice-and-speech/4524632", + "supported_endpoints": [ + "/v1/images/generations", + "/v1/images/edits" + ] + }, + "azure_ai/MAI-Image-2.5-Flash": { + "input_cost_per_image_token": 1.75e-06, + "input_cost_per_token": 1.75e-06, + "litellm_provider": "azure_ai", + "mode": "image_generation", + "output_cost_per_image": 0.0338, + "output_cost_per_image_token": 3.3e-05, + "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/new-mai-models-in-microsoft-foundry-across-text-image-voice-and-speech/4524632", + "supported_endpoints": [ + "/v1/images/generations", + "/v1/images/edits" + ] + }, + "azure_ai/MAI-Image-2e": { + "input_cost_per_token": 5e-06, + "litellm_provider": "azure_ai", + "mode": "image_generation", + "output_cost_per_image": 0.02, + "output_cost_per_image_token": 1.95e-05, + "source": "https://aka.ms/mai-image-2e-foundryblog", + "supported_endpoints": [ + "/v1/images/generations" + ] + }, "azure_ai/Llama-3.2-11B-Vision-Instruct": { "input_cost_per_token": 3.7e-07, "litellm_provider": "azure_ai", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index c1c05b982f6..b0ffc66d03b 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -6889,6 +6889,43 @@ "/v1/images/generations" ] }, + "azure_ai/MAI-Image-2.5": { + "input_cost_per_image_token": 8e-06, + "input_cost_per_token": 5e-06, + "litellm_provider": "azure_ai", + "mode": "image_generation", + "output_cost_per_image": 0.05, + "output_cost_per_image_token": 4.7e-05, + "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/new-mai-models-in-microsoft-foundry-across-text-image-voice-and-speech/4524632", + "supported_endpoints": [ + "/v1/images/generations", + "/v1/images/edits" + ] + }, + "azure_ai/MAI-Image-2.5-Flash": { + "input_cost_per_image_token": 1.75e-06, + "input_cost_per_token": 1.75e-06, + "litellm_provider": "azure_ai", + "mode": "image_generation", + "output_cost_per_image": 0.0338, + "output_cost_per_image_token": 3.3e-05, + "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/new-mai-models-in-microsoft-foundry-across-text-image-voice-and-speech/4524632", + "supported_endpoints": [ + "/v1/images/generations", + "/v1/images/edits" + ] + }, + "azure_ai/MAI-Image-2e": { + "input_cost_per_token": 5e-06, + "litellm_provider": "azure_ai", + "mode": "image_generation", + "output_cost_per_image": 0.02, + "output_cost_per_image_token": 1.95e-05, + "source": "https://aka.ms/mai-image-2e-foundryblog", + "supported_endpoints": [ + "/v1/images/generations" + ] + }, "azure_ai/Llama-3.2-11B-Vision-Instruct": { "input_cost_per_token": 3.7e-07, "litellm_provider": "azure_ai", diff --git a/tests/test_litellm/llms/azure/image_generation/test_azure_image_generation_init.py b/tests/test_litellm/llms/azure/image_generation/test_azure_image_generation_init.py index f49b3f09d6d..a211a69b9c7 100644 --- a/tests/test_litellm/llms/azure/image_generation/test_azure_image_generation_init.py +++ b/tests/test_litellm/llms/azure/image_generation/test_azure_image_generation_init.py @@ -57,6 +57,22 @@ def test_azure_providers_image_generation_json_body_keeps_model(): assert out == data +def test_azure_image_generation_mai_base_model_uses_mai_url(): + azure_chat = AzureChatCompletion() + url = azure_chat.create_azure_base_url( + azure_client_params={ + "azure_endpoint": "https://my-resource.services.ai.azure.com", + "api_version": "preview", + }, + model="image-deployment-alias", + base_model="MAI-Image-2.5", + ) + assert ( + url + == "https://my-resource.services.ai.azure.com/mai/v1/images/generations?api-version=preview" + ) + + def test_azure_image_generation_flattens_extra_body(): """ Test that Azure image generation correctly flattens extra_body parameters. diff --git a/tests/test_litellm/llms/azure_ai/image_edit/test_mai_image_edit_transformation.py b/tests/test_litellm/llms/azure_ai/image_edit/test_mai_image_edit_transformation.py new file mode 100644 index 00000000000..d5256be02d7 --- /dev/null +++ b/tests/test_litellm/llms/azure_ai/image_edit/test_mai_image_edit_transformation.py @@ -0,0 +1,171 @@ +import io +import os +import sys +from unittest.mock import MagicMock + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../../..")) + +from litellm.llms.azure_ai.image_edit import ( + AzureFoundryMAIImageEditConfig, + get_azure_ai_image_edit_config, +) +from litellm.llms.azure_ai.image_generation.mai_transformation import ( + AzureFoundryMAIImageGenerationConfig, +) + + +class TestAzureMAIImageEdit: + def test_get_mai_image_edit_url(self): + url = AzureFoundryMAIImageGenerationConfig.get_mai_image_edit_url( + api_base="https://my-resource.services.ai.azure.com", + api_version="preview", + ) + assert ( + url + == "https://my-resource.services.ai.azure.com/mai/v1/images/edits?api-version=preview" + ) + + def test_get_mai_image_edit_url_rewrites_generation_url(self): + url = AzureFoundryMAIImageGenerationConfig.get_mai_image_edit_url( + api_base=( + "https://my-resource.services.ai.azure.com/mai/v1/images/generations" + "?api-version=preview" + ), + api_version="preview", + ) + assert ( + url + == "https://my-resource.services.ai.azure.com/mai/v1/images/edits?api-version=preview" + ) + + def test_get_mai_image_edit_url_appends_edits_to_mai_root(self): + url = AzureFoundryMAIImageGenerationConfig.get_mai_image_edit_url( + api_base="https://my-resource.services.ai.azure.com/mai/v1", + api_version="preview", + ) + assert ( + url + == "https://my-resource.services.ai.azure.com/mai/v1/images/edits?api-version=preview" + ) + + def test_get_azure_ai_image_edit_config_returns_mai(self): + config = get_azure_ai_image_edit_config("MAI-Image-2.5") + assert isinstance(config, AzureFoundryMAIImageEditConfig) + + def test_validate_environment_uses_api_key_header(self): + config = AzureFoundryMAIImageEditConfig() + headers: dict = {} + config.validate_environment(headers, "MAI-Image-2.5", api_key="test-key") + assert headers["api-key"] == "test-key" + assert "Api-Key" not in headers + + def test_get_complete_url(self): + config = AzureFoundryMAIImageEditConfig() + url = config.get_complete_url( + model="MAI-Image-2.5", + api_base="https://my-resource.services.ai.azure.com", + litellm_params={"api_version": "preview"}, + ) + assert "/mai/v1/images/edits" in url + assert "api-version=preview" in url + + def test_map_openai_params_keeps_size(self): + config = AzureFoundryMAIImageEditConfig() + optional_params = config.map_openai_params( + image_edit_optional_params={"size": "1792x1024", "n": 1}, + model="MAI-Image-2.5", + drop_params=True, + ) + assert optional_params["size"] == "1792x1024" + assert optional_params["n"] == 1 + assert "width" not in optional_params + assert "height" not in optional_params + + def test_map_openai_params_defaults_size(self): + config = AzureFoundryMAIImageEditConfig() + optional_params = config.map_openai_params( + image_edit_optional_params={}, + model="MAI-Image-2.5", + drop_params=True, + ) + assert optional_params["size"] == "1024x1024" + + def test_map_openai_params_unsupported_size_raises(self): + config = AzureFoundryMAIImageEditConfig() + with pytest.raises(ValueError, match="Unsupported size value: 'auto'"): + config.map_openai_params( + image_edit_optional_params={"size": "auto"}, + model="MAI-Image-2.5", + drop_params=True, + ) + + def test_map_openai_params_invalid_size_format_raises(self): + config = AzureFoundryMAIImageEditConfig() + with pytest.raises(ValueError, match="Invalid size format: '1024xabc'"): + config.map_openai_params( + image_edit_optional_params={"size": "1024xabc"}, + model="MAI-Image-2.5", + drop_params=True, + ) + + def test_transform_image_edit_request_uses_image_field(self): + config = AzureFoundryMAIImageEditConfig() + image_bytes = io.BytesIO(b"fake-image-bytes") + + data, files = config.transform_image_edit_request( + model="MAI-Image-2.5", + prompt="Turn this into a studio product shot", + image=image_bytes, + image_edit_optional_request_params={"size": "1024x1024", "n": 1}, + litellm_params={}, + headers={}, + ) + + assert data["model"] == "MAI-Image-2.5" + assert data["prompt"] == "Turn this into a studio product shot" + assert data["size"] == "1024x1024" + assert data["n"] == 1 + assert len(files) == 1 + assert files[0][0] == "image" + assert files[0][0] != "image[]" + + def test_normalize_mai_image_usage_maps_edit_response_fields(self): + usage = AzureFoundryMAIImageGenerationConfig.normalize_mai_image_usage( + { + "num_output_tokens": 1024, + "output_image_tokens": 1024, + } + ) + assert usage["output_tokens"] == 1024 + assert usage["input_tokens"] == 0 + assert usage["total_tokens"] == 1024 + assert usage["input_tokens_details"]["text_tokens"] == 0 + assert usage["input_tokens_details"]["image_tokens"] == 0 + + def test_transform_image_edit_response_parses_mai_usage(self): + config = AzureFoundryMAIImageEditConfig() + raw_response = MagicMock(spec=httpx.Response) + raw_response.status_code = 200 + raw_response.text = "" + raw_response.json.return_value = { + "created": 1780897477, + "data": [{"b64_json": "abc123"}], + "usage": { + "num_output_tokens": 1024, + "output_image_tokens": 1024, + }, + } + + logging_obj = MagicMock() + image_response = config.transform_image_edit_response( + model="MAI-Image-2.5", + raw_response=raw_response, + logging_obj=logging_obj, + ) + + assert image_response.data[0].b64_json == "abc123" + assert image_response.usage.output_tokens == 1024 + assert image_response.usage.total_tokens == 1024 diff --git a/tests/test_litellm/llms/azure_ai/image_generation/test_mai_image_generation.py b/tests/test_litellm/llms/azure_ai/image_generation/test_mai_image_generation.py new file mode 100644 index 00000000000..f7ad333293c --- /dev/null +++ b/tests/test_litellm/llms/azure_ai/image_generation/test_mai_image_generation.py @@ -0,0 +1,380 @@ +import os +import sys +from unittest.mock import MagicMock + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../../..")) + +import litellm +from litellm.llms.azure.azure import AzureChatCompletion +from litellm.llms.azure.image_generation import get_azure_image_generation_config +from litellm.llms.azure.image_generation.http_utils import ( + azure_deployment_image_generation_json_body, +) +from litellm.llms.azure_ai.image_generation import ( + AzureFoundryMAIImageGenerationConfig, + get_azure_ai_image_generation_config, +) +from litellm.llms.azure_ai.image_generation.cost_calculator import ( + cost_calculator as azure_ai_image_cost_calculator, +) +from litellm.types.utils import ( + ImageObject, + ImageResponse, + ImageUsage, + ImageUsageInputTokensDetails, +) +from litellm.utils import get_optional_params_image_gen + + +class TestAzureMAIImageGeneration: + def test_is_mai_model(self): + assert AzureFoundryMAIImageGenerationConfig.is_mai_model("MAI-Image-2.5") + assert AzureFoundryMAIImageGenerationConfig.is_mai_model( + "azure_ai/MAI-Image-2.5" + ) + assert AzureFoundryMAIImageGenerationConfig.is_mai_model("MAI-Image-2.5-Flash") + assert AzureFoundryMAIImageGenerationConfig.is_mai_model("MAI-Image-2e") + assert not AzureFoundryMAIImageGenerationConfig.is_mai_model("flux.2-pro") + assert not AzureFoundryMAIImageGenerationConfig.is_mai_model("MAI-DS-R1") + + def test_mai_flash_and_2e_model_pricing_in_cost_map(self): + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + flash_info = litellm.get_model_info( + model="azure_ai/MAI-Image-2.5-Flash", + custom_llm_provider="azure_ai", + ) + assert flash_info["input_cost_per_token"] == 1.75e-06 + assert flash_info["input_cost_per_image_token"] == 1.75e-06 + assert flash_info["output_cost_per_image_token"] == 3.3e-05 + + image_2e_info = litellm.get_model_info( + model="azure_ai/MAI-Image-2e", + custom_llm_provider="azure_ai", + ) + assert image_2e_info["input_cost_per_token"] == 5e-06 + assert image_2e_info["output_cost_per_image_token"] == 1.95e-05 + + def test_get_mai_image_generation_url(self): + url = AzureFoundryMAIImageGenerationConfig.get_mai_image_generation_url( + api_base="https://my-resource.services.ai.azure.com", + api_version="preview", + ) + assert ( + url + == "https://my-resource.services.ai.azure.com/mai/v1/images/generations?api-version=preview" + ) + + def test_get_mai_image_generation_url_preserves_full_path(self): + api = ( + "https://my-resource.services.ai.azure.com/mai/v1/images/generations" + "?api-version=preview" + ) + url = AzureFoundryMAIImageGenerationConfig.get_mai_image_generation_url( + api_base=api, + api_version="preview", + ) + assert url == api + + def test_get_mai_image_generation_url_appends_generations_to_mai_root(self): + url = AzureFoundryMAIImageGenerationConfig.get_mai_image_generation_url( + api_base="https://my-resource.services.ai.azure.com/mai/v1", + api_version="preview", + ) + assert ( + url + == "https://my-resource.services.ai.azure.com/mai/v1/images/generations?api-version=preview" + ) + + def test_get_azure_ai_image_generation_config_returns_mai(self): + config = get_azure_ai_image_generation_config("MAI-Image-2.5") + assert isinstance(config, AzureFoundryMAIImageGenerationConfig) + + def test_azure_image_generation_config_returns_mai(self): + config = get_azure_image_generation_config("MAI-Image-2.5") + assert isinstance(config, AzureFoundryMAIImageGenerationConfig) + + def test_map_openai_params_size_to_width_height(self): + config = AzureFoundryMAIImageGenerationConfig() + optional_params = config.map_openai_params( + non_default_params={"size": "1024x1024", "n": 1}, + optional_params={}, + model="MAI-Image-2.5", + drop_params=True, + ) + assert optional_params["width"] == 1024 + assert optional_params["height"] == 1024 + assert optional_params["n"] == 1 + assert "size" not in optional_params + + def test_map_openai_params_defaults(self): + config = AzureFoundryMAIImageGenerationConfig() + optional_params = config.map_openai_params( + non_default_params={}, + optional_params={}, + model="MAI-Image-2.5", + drop_params=True, + ) + assert optional_params["width"] == 1024 + assert optional_params["height"] == 1024 + + def test_get_optional_params_image_gen_mai(self): + config = AzureFoundryMAIImageGenerationConfig() + optional_params = get_optional_params_image_gen( + model="MAI-Image-2.5", + size="1792x1024", + n=1, + custom_llm_provider="azure_ai", + provider_config=config, + drop_params=True, + ) + assert optional_params["width"] == 1792 + assert optional_params["height"] == 1024 + assert "size" not in optional_params + + def test_azure_create_azure_base_url_mai(self): + azure_chat = AzureChatCompletion() + url = azure_chat.create_azure_base_url( + azure_client_params={ + "azure_endpoint": "https://my-resource.services.ai.azure.com", + "api_version": "preview", + }, + model="MAI-Image-2.5", + ) + assert "/mai/v1/images/generations" in url + assert "api-version=preview" in url + + def test_mai_json_body_keeps_model(self): + api = ( + "https://my-resource.services.ai.azure.com/mai/v1/images/generations" + "?api-version=preview" + ) + data = { + "model": "MAI-Image-2.5", + "prompt": "A photograph of a red fox", + "width": 1024, + "height": 1024, + "n": 1, + } + out = azure_deployment_image_generation_json_body(api, data) + assert out == data + + def test_map_openai_params_custom_size(self): + config = AzureFoundryMAIImageGenerationConfig() + optional_params = config.map_openai_params( + non_default_params={"size": "768x768"}, + optional_params={}, + model="MAI-Image-2.5", + drop_params=True, + ) + assert optional_params["width"] == 768 + assert optional_params["height"] == 768 + + def test_map_openai_params_width_only_gets_height_default(self): + config = AzureFoundryMAIImageGenerationConfig() + optional_params = config.map_openai_params( + non_default_params={"width": 1792}, + optional_params={}, + model="MAI-Image-2.5", + drop_params=True, + ) + assert optional_params["width"] == 1792 + assert optional_params["height"] == config.DEFAULT_HEIGHT + + def test_map_openai_params_height_only_gets_width_default(self): + config = AzureFoundryMAIImageGenerationConfig() + optional_params = config.map_openai_params( + non_default_params={"height": 1792}, + optional_params={}, + model="MAI-Image-2.5", + drop_params=True, + ) + assert optional_params["width"] == config.DEFAULT_WIDTH + assert optional_params["height"] == 1792 + + def test_map_openai_params_unsupported_size_raises(self): + config = AzureFoundryMAIImageGenerationConfig() + with pytest.raises(ValueError, match="Unsupported size value: 'auto'"): + config.map_openai_params( + non_default_params={"size": "auto"}, + optional_params={}, + model="MAI-Image-2.5", + drop_params=True, + ) + + def test_map_openai_params_invalid_custom_size_raises(self): + config = AzureFoundryMAIImageGenerationConfig() + with pytest.raises(ValueError, match="Invalid size format: '1024xabc'"): + config.map_openai_params( + non_default_params={"size": "1024xabc"}, + optional_params={}, + model="MAI-Image-2.5", + drop_params=True, + ) + + def test_map_openai_params_unsupported_param_raises(self): + config = AzureFoundryMAIImageGenerationConfig() + with pytest.raises(ValueError, match="Parameter quality is not supported"): + config.map_openai_params( + non_default_params={"quality": "hd"}, + optional_params={}, + model="MAI-Image-2.5", + drop_params=False, + ) + + def test_transform_image_generation_response_normalizes_mai_usage(self): + config = AzureFoundryMAIImageGenerationConfig() + raw_response = MagicMock(spec=httpx.Response) + raw_response.json.return_value = { + "created": 1780897477, + "data": [{"b64_json": "abc123"}], + "usage": { + "num_output_tokens": 1024, + "num_input_text_tokens": 22, + "output_image_tokens": 1024, + }, + } + + logging_obj = MagicMock() + image_response = config.transform_image_generation_response( + model="MAI-Image-2.5", + raw_response=raw_response, + model_response=ImageResponse(), + logging_obj=logging_obj, + request_data={"prompt": "A red fox"}, + optional_params={"width": 1024, "height": 1024}, + litellm_params={}, + encoding=None, + ) + + assert image_response.data[0].b64_json == "abc123" + assert image_response.usage.output_tokens == 1024 + assert image_response.usage.input_tokens == 22 + assert image_response.usage.total_tokens == 1046 + + def test_transform_image_generation_response_non_json_raises_openai_error(self): + from litellm.llms.openai.common_utils import OpenAIError + + config = AzureFoundryMAIImageGenerationConfig() + raw_response = MagicMock(spec=httpx.Response) + raw_response.json.side_effect = ValueError("not json") + raw_response.text = "upstream gateway error" + raw_response.status_code = 502 + + with pytest.raises(OpenAIError) as exc_info: + config.transform_image_generation_response( + model="MAI-Image-2.5", + raw_response=raw_response, + model_response=ImageResponse(), + logging_obj=MagicMock(), + request_data={"prompt": "A red fox"}, + optional_params={"width": 1024, "height": 1024}, + litellm_params={}, + encoding=None, + ) + + assert exc_info.value.status_code == 502 + assert exc_info.value.message == "upstream gateway error" + + def test_normalize_mai_usage_preserves_zero_output_tokens(self): + config = AzureFoundryMAIImageGenerationConfig() + normalized = config.normalize_mai_image_usage( + { + "num_output_tokens": 0, + "output_image_tokens": 1024, + "num_input_text_tokens": 22, + } + ) + assert normalized["output_tokens"] == 0 + assert normalized["input_tokens"] == 22 + assert normalized["total_tokens"] == 22 + + def test_azure_sync_image_generation_uses_mai_response_transform(self): + raw_response = MagicMock(spec=httpx.Response) + raw_response.json.return_value = { + "created": 1780897477, + "data": [{"b64_json": "abc123"}], + "usage": { + "num_output_tokens": 1024, + "num_input_text_tokens": 22, + }, + } + + class MAIImageGenerationAzureChatCompletion(AzureChatCompletion): + def make_sync_azure_httpx_request(self, **kwargs): + return raw_response + + logging_obj = MagicMock() + image_response = MAIImageGenerationAzureChatCompletion().image_generation( + prompt="A red fox", + timeout=60.0, + optional_params={"width": 1792, "height": 1024}, + logging_obj=logging_obj, + headers={}, + model="MAI-Image-2.5", + api_key="test-key", + api_base="https://my-resource.services.ai.azure.com", + api_version="preview", + litellm_params={}, + ) + + assert image_response.data[0].b64_json == "abc123" + assert image_response.usage.output_tokens == 1024 + assert image_response.usage.input_tokens == 22 + assert image_response.usage.total_tokens == 1046 + assert image_response.size == "1792x1024" + + def test_mai_image_cost_calculator_token_based(self): + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + model = "azure_ai/MAI-Image-2.5" + model_info = litellm.get_model_info(model=model, custom_llm_provider="azure_ai") + input_text_tokens = 100 + output_image_tokens = 1024 + + image_response = ImageResponse( + data=[ImageObject(b64_json="img1")], + usage=ImageUsage( + input_tokens=input_text_tokens, + input_tokens_details=ImageUsageInputTokensDetails( + text_tokens=input_text_tokens, + image_tokens=0, + ), + output_tokens=output_image_tokens, + total_tokens=input_text_tokens + output_image_tokens, + ), + ) + + cost = azure_ai_image_cost_calculator( + model=model, + image_response=image_response, + ) + + expected_cost = ( + input_text_tokens * model_info["input_cost_per_token"] + + output_image_tokens * model_info["output_cost_per_image_token"] + ) + assert round(cost, 10) == round(expected_cost, 10) + + def test_mai_image_cost_calculator_falls_back_to_flat_image_pricing(self): + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + model = "azure_ai/MAI-Image-2.5" + model_info = litellm.get_model_info(model=model, custom_llm_provider="azure_ai") + image_response = ImageResponse( + data=[ImageObject(b64_json="img1"), ImageObject(b64_json="img2")] + ) + + cost = azure_ai_image_cost_calculator( + model=model, + image_response=image_response, + ) + + assert ( + cost == len(image_response.data or []) * model_info["output_cost_per_image"] + ) + assert cost > 0 From 51ba6e39cd23576b9c2110361f1045782762f3e4 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Mon, 8 Jun 2026 19:58:51 -0700 Subject: [PATCH 22/22] fix(mcp): load MCP tool configuration tools via the OBO/passthrough-aware GET path (#29960) * fix(ui): load MCP tool configuration tools via the OBO/passthrough-aware GET path * fix(mcp): admin-only include_disabled_tools so the settings UI shows toggled-off tools * fix(ui): repopulate MCP server edit form when server data loads after mount (OAuth return) * fix(ui): persist MCP OAuth token on save and return to the Settings tab after authorize * fix(ui): scope MCP OAuth callback to the initiating form so create and edit flows don't cross-talk * fix(ui): derive OAuth-return Settings tab via lazy state init instead of setState-in-effect * Fix MCP OAuth edit token handling --------- Co-authored-by: Cursor Agent --- .../mcp_server/rest_endpoints.py | 31 +- .../mcp_server/test_rest_endpoints.py | 125 +++++++ ui/litellm-dashboard/eslint-suppressions.json | 10 - .../mcp_tools/create_mcp_server.tsx | 2 +- .../mcp_tools/mcp_server_edit.test.tsx | 318 +++++++++++++++++- .../components/mcp_tools/mcp_server_edit.tsx | 123 ++++++- .../components/mcp_tools/mcp_server_view.tsx | 29 +- .../mcp_tools/mcp_tool_configuration.tsx | 23 +- .../src/components/mcp_tools/mcp_tools.tsx | 13 +- .../src/components/networking.tsx | 15 +- .../src/hooks/useMcpOAuthFlow.tsx | 18 + .../src/utils/mcpHeaderUtils.ts | 15 + 12 files changed, 658 insertions(+), 64 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 725f7a335bc..2149f079a3d 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -386,8 +386,15 @@ if MCP_AVAILABLE: raw_headers: Optional[Dict[str, str]] = None, user_api_key_auth: Optional[UserAPIKeyAuth] = None, extra_headers: Optional[Dict[str, str]] = None, + apply_tool_filters: bool = True, ): - """Helper function to get tools for a single server.""" + """Helper function to get tools for a single server. + + When ``apply_tool_filters`` is False the raw server catalog is returned + without the allowed_tools/disallowed_tools gate or the per-key tool + permissions. This is the admin-only configuration view; every runtime + path keeps the default True so callable tools stay filtered. + """ tools = await global_mcp_server_manager._get_tools_from_server( server=server, mcp_auth_header=server_auth_header, @@ -397,6 +404,9 @@ if MCP_AVAILABLE: user_api_key_auth=user_api_key_auth, ) + if not apply_tool_filters: + return _create_tool_response_objects(tools, server.mcp_info) + # Always apply allowed_tools/disallowed_tools so the blacklist is # enforced even when no allowlist is set (matches the SSE/HTTP path). tools = filter_tools_by_allowed_tools(tools, server) @@ -463,6 +473,7 @@ if MCP_AVAILABLE: mcp_auth_header: Optional[str], raw_headers_from_request: dict, user_api_key_dict: UserAPIKeyAuth, + apply_tool_filters: bool = True, ) -> dict: """Handle tool listing for a single server_id request.""" # Resolve a server name to its UUID if needed @@ -527,6 +538,7 @@ if MCP_AVAILABLE: raw_headers_from_request, user_api_key_dict, extra_headers=user_oauth_extra_headers, + apply_tool_filters=apply_tool_filters, ) except MCPUpstreamAuthError: # Surface the upstream 401/403 to the caller so it can emit the @@ -552,6 +564,14 @@ if MCP_AVAILABLE: server_id: Optional[str] = Query( None, description="The server id to list tools for" ), + include_disabled_tools: bool = Query( + False, + description=( + "Admin only. Return the full server tool catalog without the " + "allowed_tools filter or per-key tool permissions, so the MCP " + "settings UI can configure the allowlist. Ignored for non-admins." + ), + ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ) -> dict: """ @@ -579,6 +599,13 @@ if MCP_AVAILABLE: ) try: + # The full catalog (allowlist filter skipped) is admin-only so the + # REST endpoint can't be used to enumerate deliberately-disabled tools. + apply_tool_filters = not ( + include_disabled_tools + and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN + ) + # Extract auth headers from request headers = request.headers raw_headers_from_request = dict(headers) @@ -620,6 +647,7 @@ if MCP_AVAILABLE: mcp_auth_header=mcp_auth_header, raw_headers_from_request=raw_headers_from_request, user_api_key_dict=user_api_key_dict, + apply_tool_filters=apply_tool_filters, ) else: if not allowed_server_ids: @@ -677,6 +705,7 @@ if MCP_AVAILABLE: raw_headers_from_request, user_api_key_dict, extra_headers=user_oauth_extra_headers, + apply_tool_filters=apply_tool_filters, ) list_tools_result.extend(tools_result) except Exception as e: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index caff9ea2d28..47c9396f121 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -501,6 +501,7 @@ class TestListToolsRestAPI: raw_headers=None, user_api_key_auth=None, extra_headers=None, + apply_tool_filters=True, ): captured["called"] = True captured["server"] = server @@ -545,6 +546,78 @@ class TestListToolsRestAPI: assert result["error"] is None assert result["message"] == "Successfully retrieved tools" + async def test_include_disabled_tools_is_admin_only(self, monkeypatch): + """include_disabled_tools skips the allowlist filter only for PROXY_ADMIN; + a non-admin passing it stays filtered so the REST endpoint can't be used + to enumerate deliberately-disabled tools.""" + from litellm.proxy._types import LitellmUserRoles + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["server-1"] + + class StubServer: + alias = "server-1" + server_name = "server-1" + name = "stub" + allowed_tools = ["tool1"] + mcp_info = {"server_name": "stub"} + available_on_public_internet = True + + stub_server = StubServer() + captured = {} + + async def fake_get_tools( + server, server_auth_header, *args, apply_tool_filters=True, **kwargs + ): + captured["apply_tool_filters"] = apply_tool_filters + return ["tool-1"] + + monkeypatch.setattr( + rest_endpoints, + "build_effective_auth_contexts", + fake_contexts, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: stub_server if server_id == "server-1" else None, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints, + "_get_tools_for_single_server", + fake_get_tools, + raising=False, + ) + + request = _build_request(path="/mcp-rest/tools/list", method="GET") + + await rest_endpoints.list_tool_rest_api( + request, + server_id="server-1", + include_disabled_tools=True, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + ) + assert captured["apply_tool_filters"] is False + + await rest_endpoints.list_tool_rest_api( + request, + server_id="server-1", + include_disabled_tools=True, + user_api_key_dict=UserAPIKeyAuth(), + ) + assert captured["apply_tool_filters"] is True + @pytest.mark.parametrize("upstream_status", [401, 403]) async def test_upstream_auth_failure_surfaces_status_and_challenge( self, monkeypatch, upstream_status @@ -649,6 +722,7 @@ class TestListToolsRestAPI: raw_headers=None, user_api_key_auth=None, extra_headers=None, + apply_tool_filters=True, ): captured["called"] = True captured["server_arg"] = server @@ -792,6 +866,7 @@ class TestListToolsRestAPI: raw_headers=None, user_api_key_auth=None, extra_headers=None, + apply_tool_filters=True, ): captured["server"] = server captured["auth_header"] = server_auth_header @@ -1284,6 +1359,56 @@ class TestGetToolsForSingleServer: assert "tool1" not in tool_names assert "tool4" not in tool_names + async def test_apply_tool_filters_false_returns_full_catalog(self, monkeypatch): + """apply_tool_filters=False returns the raw catalog without the server + allowed_tools gate, so the config UI can render disabled tools as off.""" + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.types.mcp import MCPTransport + + class MockTool: + def __init__(self, name): + self.name = name + self.description = name + self.inputSchema = {} + + mock_tools = [MockTool("tool1"), MockTool("tool2"), MockTool("tool3")] + + async def fake_get_tools_from_server(**kwargs): + return mock_tools + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_get_tools_from_server", + fake_get_tools_from_server, + raising=False, + ) + + # Server enforces an allowlist of just tool1. + server = MCPServer( + server_id="test-server-id", + name="test-server", + transport=MCPTransport.sse, + allowed_tools=["tool1"], + ) + user_api_key_dict = UserAPIKeyAuth(api_key="test-key", object_permission=None) + + # Runtime default: only the allowed tool comes back. + filtered = await rest_endpoints._get_tools_for_single_server( + server=server, + server_auth_header=None, + user_api_key_auth=user_api_key_dict, + ) + assert [t.name for t in filtered] == ["tool1"] + + # Config view: full catalog, including the disabled tools. + full = await rest_endpoints._get_tools_for_single_server( + server=server, + server_auth_header=None, + user_api_key_auth=user_api_key_dict, + apply_tool_filters=False, + ) + assert {t.name for t in full} == {"tool1", "tool2", "tool3"} + class TestStdioCommandAllowlist: """Tests for MCP stdio command allowlist validation.""" diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index bbab73b07f1..233741652a9 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -1351,11 +1351,6 @@ "count": 1 } }, - "src/components/mcp_tools/mcp_server_edit.test.tsx": { - "unused-imports/no-unused-imports": { - "count": 1 - } - }, "src/components/mcp_tools/mcp_server_edit.tsx": { "no-restricted-imports": { "count": 1 @@ -1517,11 +1512,6 @@ "count": 1 } }, - "src/components/organisms/RegenerateKeyModal.tsx": { - "react-hooks/set-state-in-effect": { - "count": 1 - } - }, "src/components/organisms/create_key_button.test.tsx": { "@typescript-eslint/no-require-imports": { "count": 2 diff --git a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx index 6a8dc353f21..28c38c459aa 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/create_mcp_server.tsx @@ -188,6 +188,7 @@ const CreateMCPServer: React.FC = ({ } }, onBeforeRedirect: persistCreateUiState, + flowSource: "create", }); React.useEffect(() => { @@ -1088,7 +1089,6 @@ const CreateMCPServer: React.FC = ({
({ updateMCPServer: vi.fn(), listMCPTools: vi.fn().mockResolvedValue({ tools: [], error: null }), + storeMCPOAuthUserCredential: vi.fn().mockResolvedValue({}), })); vi.mock("../molecules/notifications_manager", () => ({ @@ -17,12 +18,13 @@ vi.mock("../molecules/notifications_manager", () => ({ }, })); +const mockOauth: { tokenResponse: any } = { tokenResponse: null }; vi.mock("@/hooks/useMcpOAuthFlow", () => ({ useMcpOAuthFlow: () => ({ startOAuthFlow: vi.fn(), status: "idle", error: null, - tokenResponse: null, + tokenResponse: mockOauth.tokenResponse, }), })); @@ -37,12 +39,19 @@ vi.mock("./MCPPermissionManagement", () => ({ vi.mock("./mcp_tool_configuration", () => ({ default: ({ existingAllowedTools, + externalTools, + externalError, onAllowedToolsChange, onToolAllowlistInteraction, onToolNameToDisplayNameChange, onToolNameToDescriptionChange, }: any) => ( -
+