Merge pull request #17443 from BerriAI/litellm_v2_login

[Feature] New Login Page
This commit is contained in:
yuneng-jiang 2025-12-03 16:23:47 -08:00 • committed by GitHub
commit 73824c278a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 451 additions and 68 deletions

View file

@ -8318,6 +8318,53 @@ async def login(request: Request): # noqa: PLR0915
return redirect_response
@router.post(
"/v2/login", include_in_schema=False
) # hidden helper for UI logins via API
async def login_v2(request: Request): # noqa: PLR0915
global premium_user, general_settings, master_key
from litellm.proxy.auth.login_utils import authenticate_user, create_ui_token_object
from litellm.proxy.utils import get_custom_url
body = await request.json()
username = str(body.get("username"))
password = str(body.get("password"))
login_result = await authenticate_user(
username=username,
password=password,
master_key=master_key,
prisma_client=prisma_client,
)
returned_ui_token_object = create_ui_token_object(
login_result=login_result,
general_settings=general_settings,
premium_user=premium_user,
)
import jwt
jwt_token = jwt.encode( # type: ignore
cast(dict, returned_ui_token_object),
master_key,
algorithm="HS256",
)
litellm_dashboard_ui = get_custom_url(str(request.base_url))
if litellm_dashboard_ui.endswith("/"):
litellm_dashboard_ui += "ui/"
else:
litellm_dashboard_ui += "/ui/"
litellm_dashboard_ui += "?login=success"
json_response = JSONResponse(
content={"redirect_url": litellm_dashboard_ui},
status_code=status.HTTP_200_OK,
)
json_response.set_cookie(key="token", value=jwt_token)
return json_response
@app.get("/onboarding/get_token", include_in_schema=False)
async def onboarding(invite_link: str, request: Request):
"""

View file

@ -71,6 +71,58 @@ def client_no_auth():
return TestClient(app)
def test_login_v2_returns_redirect_url_and_sets_cookie(monkeypatch):
mock_login_result = {"user_id": "test-user"}
mock_prisma_client = MagicMock()
mock_authenticate_user = AsyncMock(return_value=mock_login_result)
mock_create_ui_token_object = MagicMock(return_value={"user_id": "test-user"})
mock_jwt_encode = MagicMock(return_value="signed-token")
monkeypatch.setattr(
"litellm.proxy.auth.login_utils.authenticate_user",
mock_authenticate_user,
)
monkeypatch.setattr(
"litellm.proxy.auth.login_utils.create_ui_token_object",
mock_create_ui_token_object,
)
monkeypatch.setattr("jwt.encode", mock_jwt_encode)
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
client = TestClient(app)
response = client.post(
"/v2/login",
json={"username": "alice", "password": "secret"},
)
assert response.status_code == 200
assert (
response.json()
== {"redirect_url": "http://testserver/ui/?login=success"}
)
assert response.cookies.get("token") == "signed-token"
mock_authenticate_user.assert_awaited_once_with(
username="alice",
password="secret",
master_key="test-master-key",
prisma_client=mock_prisma_client,
)
mock_create_ui_token_object.assert_called_once_with(
login_result=mock_login_result,
general_settings={},
premium_user=False,
)
mock_jwt_encode.assert_called_once_with(
{"user_id": "test-user"},
"test-master-key",
algorithm="HS256",
)
@pytest.mark.asyncio
async def test_initialize_scheduled_jobs_credentials(monkeypatch):
"""

View file

@ -0,0 +1,11 @@
import { useMutation } from "@tanstack/react-query";
import { loginCall, LoginRequest } from "@/components/networking";
export const useLogin = () => {
return useMutation({
mutationFn: async ({ username, password }: LoginRequest) => {
const result = await loginCall(username, password);
return result;
},
});
};

View file

@ -0,0 +1,14 @@
import { getUiConfig, LiteLLMWellKnownUiConfig } from "@/components/networking";
import { useQuery } from "@tanstack/react-query";
import { createQueryKeys } from "../common/queryKeysFactory";
const uiConfigKeys = createQueryKeys("uiConfig");
export const useUIConfig = () => {
return useQuery<LiteLLMWellKnownUiConfig>({
queryKey: uiConfigKeys.list({}),
queryFn: async () => await getUiConfig(),
staleTime: 24 * 60 * 60 * 1000, // 24 hours - data rarely changes
gcTime: 24 * 60 * 60 * 1000, // 24 hours - keep in cache for 24 hours
});
};

View file

@ -0,0 +1,154 @@
"use client";
import { useLogin } from "@/app/(dashboard)/hooks/login/useLogin";
import { useUIConfig } from "@/app/(dashboard)/hooks/uiConfig/useUIConfig";
import LoadingScreen from "@/components/common_components/LoadingScreen";
import { getProxyBaseUrl } from "@/components/networking";
import { getCookie } from "@/utils/cookieUtils";
import { isJwtExpired } from "@/utils/jwtUtils";
import { InfoCircleOutlined } from "@ant-design/icons";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { Alert, Button, Card, Form, Input, Space, Typography } from "antd";
import { useRouter } from "next/navigation";
import { useEffect, useState } from "react";
function LoginPageContent() {
const [username, setUsername] = useState("");
const [password, setPassword] = useState("");
const [isLoading, setIsLoading] = useState(true);
const { isLoading: isConfigLoading } = useUIConfig();
const loginMutation = useLogin();
const router = useRouter();
useEffect(() => {
if (isConfigLoading) {
return;
}
const rawToken = getCookie("token");
if (rawToken && !isJwtExpired(rawToken)) {
router.replace(`${getProxyBaseUrl()}/ui`);
return;
}
setIsLoading(false);
}, [isConfigLoading, router]);
const handleSubmit = () => {
loginMutation.mutate(
{ username, password },
{
onSuccess: (data) => {
router.push(data.redirect_url);
},
},
);
};
const error = loginMutation.error instanceof Error ? loginMutation.error.message : null;
const isLoginLoading = loginMutation.isPending;
const { Title, Text, Paragraph } = Typography;
if (isConfigLoading || isLoading) {
return <LoadingScreen />;
}
return (
<div className="min-h-screen flex items-center justify-center bg-gray-50">
<Card className="w-full max-w-lg shadow-md">
<Space direction="vertical" size="middle" className="w-full">
<div className="text-center">
<Title level={2}>🚅 LiteLLM</Title>
</div>
<div className="text-center">
<Title level={3}>Login</Title>
<Text type="secondary">Access your LiteLLM Admin UI.</Text>
</div>
<Alert
message="Default Credentials"
description={
<>
<Paragraph className="text-sm">
By default, Username is <code className="bg-gray-100 px-1 py-0.5 rounded text-xs">admin</code> and
Password is your set LiteLLM Proxy
<code className="bg-gray-100 px-1 py-0.5 rounded text-xs">MASTER_KEY</code>.
</Paragraph>
<Paragraph className="text-sm">
Need to set UI credentials or SSO?{" "}
<a href="https://docs.litellm.ai/docs/proxy/ui" target="_blank" rel="noopener noreferrer">
Check the documentation
</a>
.
</Paragraph>
</>
}
type="info"
icon={<InfoCircleOutlined />}
showIcon
/>
{error && <Alert message={error} type="error" showIcon />}
<Form onFinish={handleSubmit} layout="vertical" requiredMark={true}>
<Form.Item
label="Username"
name="username"
rules={[{ required: true, message: "Please enter your username" }]}
>
<Input
placeholder="Enter your username"
autoComplete="username"
value={username}
onChange={(e) => setUsername(e.target.value)}
disabled={isLoginLoading}
size="large"
className="rounded-md border-gray-300"
/>
</Form.Item>
<Form.Item
label="Password"
name="password"
rules={[{ required: true, message: "Please enter your password" }]}
>
<Input.Password
placeholder="Enter your password"
autoComplete="current-password"
value={password}
onChange={(e) => setPassword(e.target.value)}
disabled={isLoginLoading}
size="large"
/>
</Form.Item>
<Form.Item>
<Button
type="primary"
htmlType="submit"
loading={isLoginLoading}
disabled={isLoginLoading}
block
size="large"
>
{isLoginLoading ? "Logging in..." : "Login"}
</Button>
</Form.Item>
</Form>
</Space>
</Card>
</div>
);
}
export default function LoginPage() {
const queryClient = new QueryClient();
return (
<QueryClientProvider client={queryClient}>
<LoginPageContent />
</QueryClientProvider>
);
}

View file

@ -1,49 +1,47 @@
"use client";
import React, { Suspense, useEffect, useState } from "react";
import { useSearchParams } from "next/navigation";
import { jwtDecode } from "jwt-decode";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { Team } from "@/components/key_team_helpers/key_list";
import Navbar from "@/components/navbar";
import { ThemeProvider } from "@/contexts/ThemeContext";
import UserDashboard from "@/components/user_dashboard";
import OldModelDashboard from "@/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView";
import ViewUserDashboard from "@/components/view_users";
import Organizations from "@/components/organizations";
import { fetchOrganizations } from "@/components/organizations";
import AdminPanel from "@/components/admins";
import Settings from "@/components/settings";
import GeneralSettings from "@/components/general_settings";
import PassThroughSettings from "@/components/pass_through_settings";
import BudgetPanel from "@/components/budgets/budget_panel";
import SpendLogsTable from "@/components/view_logs";
import ModelHubTable from "@/components/model_hub_table";
import PublicModelHub from "@/components/public_model_hub";
import NewUsagePage from "@/components/new_usage";
import APIReferenceView from "@/app/(dashboard)/api-reference/APIReferenceView";
import PlaygroundPage from "@/app/(dashboard)/playground/page";
import Usage from "@/components/usage";
import CacheDashboard from "@/components/cache_dashboard";
import { getUiConfig, proxyBaseUrl, setGlobalLitellmHeaderName } from "@/components/networking";
import { Organization } from "@/components/networking";
import GuardrailsPanel from "@/components/guardrails";
import AgentsPanel from "@/components/agents";
import PromptsPanel from "@/components/prompts";
import TransformRequestPanel from "@/components/transform_request";
import { fetchUserModels } from "@/components/organisms/create_key_button";
import { fetchTeams } from "@/components/common_components/fetch_teams";
import { MCPServers } from "@/components/mcp_tools";
import TagManagement from "@/components/tag_management";
import VectorStoreManagement from "@/components/vector_store_management";
import UIThemeSettings from "@/components/ui_theme_settings";
import { CostTrackingSettings } from "@/components/CostTrackingSettings";
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
import { cx } from "@/lib/cva.config";
import SidebarProvider from "@/app/(dashboard)/components/SidebarProvider";
import OldModelDashboard from "@/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView";
import PlaygroundPage from "@/app/(dashboard)/playground/page";
import AdminPanel from "@/components/admins";
import AgentsPanel from "@/components/agents";
import BudgetPanel from "@/components/budgets/budget_panel";
import CacheDashboard from "@/components/cache_dashboard";
import { fetchTeams } from "@/components/common_components/fetch_teams";
import LoadingScreen from "@/components/common_components/LoadingScreen";
import { CostTrackingSettings } from "@/components/CostTrackingSettings";
import GeneralSettings from "@/components/general_settings";
import GuardrailsPanel from "@/components/guardrails";
import { Team } from "@/components/key_team_helpers/key_list";
import { MCPServers } from "@/components/mcp_tools";
import ModelHubTable from "@/components/model_hub_table";
import Navbar from "@/components/navbar";
import { getUiConfig, Organization, proxyBaseUrl, setGlobalLitellmHeaderName } from "@/components/networking";
import NewUsagePage from "@/components/new_usage";
import OldTeams from "@/components/OldTeams";
import { fetchUserModels } from "@/components/organisms/create_key_button";
import Organizations, { fetchOrganizations } from "@/components/organizations";
import PassThroughSettings from "@/components/pass_through_settings";
import PromptsPanel from "@/components/prompts";
import PublicModelHub from "@/components/public_model_hub";
import { SearchTools } from "@/components/search_tools";
import Settings from "@/components/settings";
import TagManagement from "@/components/tag_management";
import TransformRequestPanel from "@/components/transform_request";
import UIThemeSettings from "@/components/ui_theme_settings";
import Usage from "@/components/usage";
import UserDashboard from "@/components/user_dashboard";
import VectorStoreManagement from "@/components/vector_store_management";
import SpendLogsTable from "@/components/view_logs";
import ViewUserDashboard from "@/components/view_users";
import { ThemeProvider } from "@/contexts/ThemeContext";
import { isJwtExpired } from "@/utils/jwtUtils";
import { isAdminRole } from "@/utils/roles";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { jwtDecode } from "jwt-decode";
import { useSearchParams } from "next/navigation";
import { Suspense, useEffect, useState } from "react";
function getCookie(name: string) {
// Safer cookie read + decoding; handles '=' inside values
@ -62,19 +60,6 @@ function deleteCookie(name: string, path = "/") {
document.cookie = `${name}=; Max-Age=0; Path=${path}`;
}
function isJwtExpired(token: string): boolean {
try {
const decoded: any = jwtDecode(token);
if (decoded && typeof decoded.exp === "number") {
return decoded.exp * 1000 <= Date.now();
}
return false;
} catch {
// If we can't decode, treat as invalid/expired
return true;
}
}
function formatUserRole(userRole: string) {
if (!userRole) {
return "Undefined Role";
@ -112,19 +97,6 @@ interface ProxySettings {
const queryClient = new QueryClient();
function LoadingScreen() {
return (
<div className={cx("h-screen", "flex items-center justify-center gap-4")}>
<div className="text-lg font-medium py-2 pr-4 border-r border-r-gray-200">🚅 LiteLLM</div>
<div className="flex items-center justify-center gap-2">
<UiLoadingSpinner className="size-4" />
<span className="text-gray-600 text-sm">Loading...</span>
</div>
</div>
);
}
export default function CreateKeyPage() {
const [userRole, setUserRole] = useState("");
const [premiumUser, setPremiumUser] = useState(false);

View file

@ -0,0 +1,10 @@
import { render, screen } from "@testing-library/react";
import { describe, it, expect } from "vitest";
import LoadingScreen from "./LoadingScreen";
describe("LoadingScreen", () => {
it("should render", () => {
render(<LoadingScreen />);
expect(screen.getByText("Loading...")).toBeInTheDocument();
});
});

View file

@ -0,0 +1,15 @@
import { cx } from "@/lib/cva.config";
import { UiLoadingSpinner } from "../ui/ui-loading-spinner";
export default function LoadingScreen() {
return (
<div className={cx("h-screen", "flex items-center justify-center gap-4")}>
<div className="text-lg font-medium py-2 pr-4 border-r border-r-gray-200">🚅 LiteLLM</div>
<div className="flex items-center justify-center gap-2">
<UiLoadingSpinner className="size-4" />
<span className="text-gray-600 text-sm">Loading...</span>
</div>
</div>
);
}

View file

@ -124,7 +124,7 @@ export interface PromptSpec {
prompt_info: PromptInfo;
created_at?: string;
updated_at?: string;
version?: number; // Explicit version number for version history
version?: number; // Explicit version number for version history
}
export interface PromptTemplateBase {
@ -7414,7 +7414,11 @@ interface RegisterMcpOAuthClientPayload {
token_endpoint_auth_method?: string;
}
export const registerMcpOAuthClient = async (accessToken: string, serverId: string, payload: RegisterMcpOAuthClientPayload) => {
export const registerMcpOAuthClient = async (
accessToken: string,
serverId: string,
payload: RegisterMcpOAuthClientPayload,
) => {
const base = getProxyBaseUrl();
const normalizedServerId = encodeURIComponent(serverId.trim());
const url = `${base}/v1/mcp/server/oauth/${normalizedServerId}/register`;
@ -7424,7 +7428,7 @@ export const registerMcpOAuthClient = async (accessToken: string, serverId: stri
headers: {
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
"Accept": "application/json, text/event-stream",
Accept: "application/json, text/event-stream",
},
body: JSON.stringify(payload),
});
@ -7969,3 +7973,40 @@ const deriveErrorMessage = (errorData: any): string => {
JSON.stringify(errorData)
);
};
export interface LoginRequest {
username: string;
password: string;
}
export interface LoginResponse {
redirect_url: string;
}
export const loginCall = async (username: string, password: string): Promise<LoginResponse> => {
const proxyBaseUrl = getProxyBaseUrl();
const loginUrl = proxyBaseUrl ? `${proxyBaseUrl}/v2/login` : "/v2/login";
const body = JSON.stringify({
username,
password,
});
const response = await fetch(loginUrl, {
method: "POST",
body,
credentials: "include",
headers: {
"Content-Type": "application/json",
},
});
if (!response.ok) {
const errorData = await response.json();
const errorMessage = deriveErrorMessage(errorData);
throw new Error(errorMessage);
}
const data = await response.json();
return data;
};

View file

@ -0,0 +1,53 @@
import { describe, expect, it, vi, beforeEach, afterEach } from "vitest";
import { isJwtExpired } from "./jwtUtils";
import { jwtDecode } from "jwt-decode";
vi.mock("jwt-decode");
describe("jwtUtils", () => {
beforeEach(() => {
vi.clearAllMocks();
});
afterEach(() => {
vi.restoreAllMocks();
});
it("should return true if the token is expired", () => {
const mockDateNow = 1716838401000;
vi.spyOn(Date, "now").mockReturnValue(mockDateNow);
vi.mocked(jwtDecode).mockReturnValue({
exp: Math.floor(mockDateNow / 1000) - 1,
user_id: "test",
});
expect(isJwtExpired("any-token")).toBe(true);
});
it("should return false if the token is not expired", () => {
const mockDateNow = 1716838401000;
vi.spyOn(Date, "now").mockReturnValue(mockDateNow);
vi.mocked(jwtDecode).mockReturnValue({
exp: Math.floor(mockDateNow / 1000) + 1000,
user_id: "test",
});
expect(isJwtExpired("any-token")).toBe(false);
});
it("should return false if the token does not have an exp field", () => {
vi.mocked(jwtDecode).mockReturnValue({
user_id: "test",
});
expect(isJwtExpired("any-token")).toBe(false);
});
it("should return true if jwtDecode throws an error", () => {
vi.mocked(jwtDecode).mockImplementation(() => {
throw new Error("Invalid token");
});
expect(isJwtExpired("invalid-token")).toBe(true);
});
});

View file

@ -0,0 +1,14 @@
import { jwtDecode } from "jwt-decode";
export function isJwtExpired(token: string): boolean {
try {
const decoded: any = jwtDecode(token);
if (decoded && typeof decoded.exp === "number") {
return decoded.exp * 1000 <= Date.now();
}
return false;
} catch {
// If we can't decode, treat as invalid/expired
return true;
}
}