diff --git a/package.json b/package.json index 8662a0cd..93d16da6 100644 --- a/package.json +++ b/package.json @@ -147,6 +147,7 @@ "mysql2": "^3.11.0", "openpgp": "^6.3.1", "otplib": "^12.0.1", + "parse5": "^7.3.0", "qrcode": "^1.5.4", "rehype-katex": "7.0.1", "rehype-raw": "^7.0.0", diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml index 6a3fca1c..77b77caa 100644 --- a/pnpm-lock.yaml +++ b/pnpm-lock.yaml @@ -134,6 +134,9 @@ importers: otplib: specifier: ^12.0.1 version: 12.0.1 + parse5: + specifier: ^7.3.0 + version: 7.3.0 qrcode: specifier: ^1.5.4 version: 1.5.4 diff --git a/scripts/fetch-frontend.mjs b/scripts/fetch-frontend.mjs index ec3574bd..f3096115 100644 --- a/scripts/fetch-frontend.mjs +++ b/scripts/fetch-frontend.mjs @@ -8,9 +8,10 @@ * 1. FRONTEND_DIST 环境变量:已构建好的 dist 目录路径(最快,CI 缓存场景) * 2. FRONTEND_REPO 环境变量:本地官方前端仓库路径(自动 install + build) * 3. 同级目录 ../OpenList-Frontend(monorepo 布局,自动探测,自动 install + build) - * 4. 默认:下载 npm 上【已发布】的 dist(版本取 registry 的 latest, - * 可用 FRONTEND_VERSION 固定) - * 5. FRONTEND_BUILD_FROM_SOURCE=1:从 Git 克隆前端 main 分支并现构建 + * 4. 默认:下载 npm 上【已发布】且兼容 Worker 初始化协议的 dist(版本取 + * registry 的 latest,可用 FRONTEND_VERSION 固定) + * 5. 发布版尚未包含 Worker 初始化协议,或 FRONTEND_BUILD_FROM_SOURCE=1: + * 从 Git 克隆前端 main 分支并现构建 * * 为什么默认取「已发布 dist」而不是「克隆 main 现构建」: * 前端产物是内容哈希文件名(/assets/index-XXXX.js),而 CDN(jsdelivr / @@ -22,6 +23,10 @@ * 可用(等价 Go Release 版行为)。同时 stampFrontendVersion 会把该版本号写进 * index.html,使 ASSET_URLS 的 $version 正好解析到这份 dist 对应的版本。 * + * 已发布包可能落后于 Worker 后端。只有同时包含 /public/init_status 与 /@init + * 的产物才可使用;否则全新部署会只请求会被 503 拦截的 /public/settings,既不 + * 显示初始化向导也无法创建管理员。遇到这种版本会自动改为构建前端 main。 + * * 用法: * FRONTEND_DIST=/path/to/dist node scripts/fetch-frontend.mjs * FRONTEND_REPO=../OpenList-Frontend node scripts/fetch-frontend.mjs @@ -95,6 +100,32 @@ function requireDist(src) { } } +/** + * 已构建前端是否包含 Worker 首次初始化协议。 + * + * Vite 会保留路由和 API 路径字符串,因此无需执行或反编译 bundle。两项必须 + * 同时存在:init_status 负责识别空数据库,/@init 才能渲染创建管理员的向导。 + */ +function supportsWorkerSetup(src) { + let hasStatus = false + let hasRoute = false + const pending = [src] + while (pending.length > 0 && (!hasStatus || !hasRoute)) { + const current = pending.pop() + for (const entry of fs.readdirSync(current, { withFileTypes: true })) { + const full = path.join(current, entry.name) + if (entry.isDirectory()) { + pending.push(full) + } else if (entry.isFile() && entry.name.endsWith(".js")) { + const code = fs.readFileSync(full, "utf-8") + hasStatus ||= code.includes("/public/init_status") + hasRoute ||= code.includes("/@init") + } + } + } + return hasStatus && hasRoute +} + function replaceDist(src) { console.log(` Copying frontend dist: ${src} -> ${DEST}`) fs.rmSync(DEST, { recursive: true, force: true }) @@ -113,7 +144,10 @@ function replaceDist(src) { function stampFrontendVersion(src) { try { const pkg = JSON.parse( - fs.readFileSync(path.join(path.resolve(src, ".."), "package.json"), "utf-8"), + fs.readFileSync( + path.join(path.resolve(src, ".."), "package.json"), + "utf-8", + ), ) // 只信任官方前端包的版本号:FRONTEND_DIST 可能指向任意目录, // 误读(例如 worker 自身 package.json 的 4.2.3)会戳出错误的 CDN 版本。 @@ -145,7 +179,9 @@ function stampFrontendVersion(src) { function fetchI18n(repo) { const langDir = path.join(repo, "src", "lang") if (!fs.existsSync(langDir)) { - console.warn(" [fetch-frontend] repo missing src/lang, skipping i18n fetch") + console.warn( + " [fetch-frontend] repo missing src/lang, skipping i18n fetch", + ) return } const tmpTar = path.join(os.tmpdir(), `openlist-i18n-${process.pid}.tar.gz`) @@ -172,8 +208,7 @@ function buildLocalRepo(repo) { } const pm = detectPackageManager(abs) const cmd = resolvePmCommand(abs, pm) - const install = (extra = "") => - run(`${cmd} install${extra}`, { cwd: abs }) + const install = (extra = "") => run(`${cmd} install${extra}`, { cwd: abs }) try { install() } catch { @@ -181,7 +216,9 @@ function buildLocalRepo(repo) { // minimumReleaseAge / trustPolicy 供应链复核,registry manifest 缺少 // 平台子包时会误报(如 @crowdin/cli-*-arm64)。lockfile 来自刚克隆的 // 官方前端仓库(HTTPS + 官方分支),属于可信来源,跳过复核安全。 - console.warn(" [fetch-frontend] pnpm install failed (lockfile supply-chain recheck or network issue), retrying once with --trust-lockfile...") + console.warn( + " [fetch-frontend] pnpm install failed (lockfile supply-chain recheck or network issue), retrying once with --trust-lockfile...", + ) install(" --trust-lockfile") } fetchI18n(abs) @@ -239,7 +276,14 @@ async function fetchPublishedDist() { run(`tar -xzf pkg.tgz package/dist package/package.json`, { cwd: tmp }) const src = path.join(tmp, "package", "dist") requireDist(src) + if (!supportsWorkerSetup(src)) { + console.warn( + ` Published frontend ${version} lacks the Worker setup protocol; building ${OFFICIAL_REPO_REF} instead`, + ) + return false + } replaceDist(src) + return true } finally { fs.rmSync(tmp, { recursive: true, force: true }) } @@ -272,17 +316,18 @@ async function main() { return } - // 4. 下载 npm 上已发布的 dist(默认) + // 4. 下载 npm 上已发布且兼容 Worker 初始化协议的 dist(默认) // 从 main 现构建的产物哈希与 CDN 不一致,会让路径 A 失效、npmmirror 之类的 // 镜像完全不可用(详见文件头),故默认改为取已发布产物。 if (process.env.FRONTEND_BUILD_FROM_SOURCE !== "1") { - await fetchPublishedDist() - return + if (await fetchPublishedDist()) return } - // 5. 从 Git 克隆 main 并构建(FRONTEND_BUILD_FROM_SOURCE=1 时使用) + // 5. 从 Git 克隆 main 并构建(显式要求源码构建,或发布版尚未兼容 Worker) const tmp = fs.mkdtempSync(path.join(os.tmpdir(), "openlist-frontend-")) - console.log(` Cloning official frontend: ${OFFICIAL_REPO_URL}#${OFFICIAL_REPO_REF}`) + console.log( + ` Cloning official frontend: ${OFFICIAL_REPO_URL}#${OFFICIAL_REPO_REF}`, + ) try { run( // -c core.autocrlf=false:禁用克隆端的换行符转换。Windows 上 autocrlf diff --git a/src/backend/drivers/autoindex/driver.test.ts b/src/backend/drivers/autoindex/driver.test.ts index ed2596c5..40a4cd4b 100644 --- a/src/backend/drivers/autoindex/driver.test.ts +++ b/src/backend/drivers/autoindex/driver.test.ts @@ -54,6 +54,31 @@ test("parseAutoIndexHTML extracts files, dirs, size and modified", () => { assert.equal(nodes[1].isDir, true) }) +test("parseAutoIndexHTML tolerates malformed Apache directory HTML", () => { + // archive.apache.org/dist/tomcat/ 实际返回过这种重复 html/body/pre、只关闭一层 + // pre 的页面。浏览器与 Go htmlquery 都能解析,AutoIndex 也不能按严格 XML 拒绝。 + const html = ` +Index +Index
+
Parent Directory
+tomcat-10/ 2026-09-15 10:54 -
+
` + + const nodes = parseAutoIndexHTML( + html, + "//pre/a", + "@href", + "string(following-sibling::text()[1])", + "string(following-sibling::text()[2])", + ["Parent Directory"], + ) + + const tomcat = nodes.find((node) => node.name === "tomcat-10") + assert.ok(tomcat) + assert.equal(tomcat.url, "tomcat-10/") + assert.equal(tomcat.isDir, true) +}) + test("normalizeAutoIndexAddition fills defaults", () => { const a = normalizeAutoIndexAddition({ url: "example.com/files" }) assert.equal(a.url, "https://example.com/files/") diff --git a/src/backend/drivers/autoindex/driver.ts b/src/backend/drivers/autoindex/driver.ts index 4516adeb..c05185d5 100644 --- a/src/backend/drivers/autoindex/driver.ts +++ b/src/backend/drivers/autoindex/driver.ts @@ -13,6 +13,7 @@ const DefaultItemXPath = "//pre/a" const DefaultNameXPath = "@href" const DefaultSizeXPath = "string(following-sibling::text()[1])" const DefaultModifiedXPath = "string(following-sibling::text()[2])" +const FETCH_TIMEOUT_MS = 30_000 export function normalizeAutoIndexAddition(a: any): AutoIndexAddition { const norm = { ...(a || {}) } as any @@ -78,7 +79,12 @@ export class AutoIndexDriver implements StorageDriver { async list(virtualPath: string, physicalPath: string): Promise { const baseURL = this.buildDirURL(physicalPath) - const res = await fetch(baseURL) + // 边缘运行时通常会等到平台级超时(EdgeOne 为 120 秒)才中止不可达的 + // 上游请求。显式限制单次目录读取,避免一个失联的 AutoIndex 挂载长期占用 + // 实例并拖累同一服务的其他请求。 + const res = await fetch(baseURL, { + signal: AbortSignal.timeout(FETCH_TIMEOUT_MS), + }) if (!res.ok) { throw new Error(`Failed to fetch ${baseURL}: HTTP ${res.status}`) } diff --git a/src/backend/drivers/autoindex/util.ts b/src/backend/drivers/autoindex/util.ts index 85a80eff..a3d3e474 100644 --- a/src/backend/drivers/autoindex/util.ts +++ b/src/backend/drivers/autoindex/util.ts @@ -1,6 +1,7 @@ // AutoIndex utility functions import * as xpath from "xpath" import { DOMParser } from "@xmldom/xmldom" +import { parse, serialize } from "parse5" import { AutoIndexNode } from "./types" /** @@ -20,19 +21,30 @@ function extractXPathValue(raw: unknown): string | undefined { } // 将宽松 HTML 清洗为可被 XML 解析器(xmldom)解析的良构 XML: -// 1) 移除 DOCTYPE 与注释;2) 把 void 标签(
//
等)转为自闭合。 -// 这样用 text/xml 解析后节点无命名空间,用户配置的无前缀 XPath(如 //pre/a) -// 才能正常匹配——与 Go 侧 htmlquery(golang.org/x/net/html)的行为保持一致。 -// 注意:若直接用 text/html 解析,xmldom 会注入 XHTML 命名空间,导致 //pre/a 匹配不到。 +// 1) 先用 parse5 按 HTML5 容错规则修复未闭合/错位/重复标签; +// 2) 移除 XML 不需要的 DOCTYPE、注释和 script/style; +// 3) 把 void 标签(
//
等)转为自闭合。 +// +// 不能直接把远端 HTML 交给 xmldom 的 XML 模式:真实目录页常包含浏览器可以 +// 正常处理的非严格 HTML。例如 archive.apache.org/dist/tomcat/ 同时有重复的 +// html/body/pre 标签,旧实现会抛 "Opening and ending tag mismatch"。parse5 与 +// Go 版 htmlquery 底层的 HTML parser 一样会容错,再转 XML 后仍可继续使用用户 +// 配置的无前缀 XPath(如 //pre/a)。 +// +// 注意:若直接用 xmldom 的 text/html 模式,它会注入 XHTML 命名空间,导致 +// //pre/a 匹配不到。 const HTML_VOID_TAGS = "area|base|br|col|embed|hr|img|input|link|meta|param|source|track|wbr" function htmlToXml(html: string): string { - let s = html + let s = serialize(parse(html)) s = s.replace(/]*>/gi, "") s = s.replace(//g, "") + // script/style 的 raw text 可以包含 XML 非法的裸 <、&;AutoIndex XPath 不会 + // 依赖这些内容,移除可避免第二阶段 XML 解析被无关脚本或样式破坏。 + s = s.replace(/<(script|style)\b[^>]*>[\s\S]*?<\/\1\s*>/gi, "") s = s.replace( - new RegExp(`<(${HTML_VOID_TAGS})([^>]*?)/?>`, "gi"), + new RegExp(`<(${HTML_VOID_TAGS})([^>]*?)/?\\s*>`, "gi"), "<$1$2/>", ) return s @@ -44,7 +56,7 @@ export function parseAutoIndexHTML( nameXPath: string, sizeXPath: string, modifiedXPath: string, - ignoreNames: string[] + ignoreNames: string[], ): AutoIndexNode[] { // @xmldom/xmldom >= 0.9 已废弃 errorHandler,改用 onError 回调 const doc = new DOMParser({ @@ -146,10 +158,7 @@ export function parseSize(sizeStr: string): number { return Math.round(num * mul) } -export function parseTime( - timeStr: string, - format: string -): string { +export function parseTime(timeStr: string, format: string): string { if (!timeStr) return new Date().toISOString() try { @@ -161,7 +170,9 @@ export function parseTime( const isoMatch = timeStr.match(/(\d{4})-(\d{2})-(\d{2})\s+(\d{2}):(\d{2})/) if (isoMatch) { const [, year, month, day, hour, minute] = isoMatch - return new Date(`${year}-${month}-${day}T${hour}:${minute}:00`).toISOString() + return new Date( + `${year}-${month}-${day}T${hour}:${minute}:00`, + ).toISOString() } // Try common formats diff --git a/src/backend/internal/model/store/driver/kv.ts b/src/backend/internal/model/store/driver/kv.ts index 6962fd03..502a66ee 100644 --- a/src/backend/internal/model/store/driver/kv.ts +++ b/src/backend/internal/model/store/driver/kv.ts @@ -1,10 +1,10 @@ /** * KV 驱动(自动适配 Cloudflare / EdgeOne) - * + * * 支持两种模式: * 1. Binding 模式:直接访问 KV binding(Cloudflare Workers / EdgeOne Edge Functions) * 2. HTTP 代理模式:通过 Edge Function 代理访问(EdgeOne Node Functions) - * + * * 自动检测环境并选择合适的模式。 */ import { sanitizeProxyOrigin } from "../proxy" @@ -71,6 +71,33 @@ function getProxySecret(env?: EnvContext): string | null { } } +/** + * 当前环境是否有充分依据启用 EdgeOne KV HTTP 代理探测。 + * + * `__requestOrigin` 不能作为平台依据:index.ts 会给所有平台和自托管请求注入它。 + * 旧实现只要同时看到 origin 与 JWT_SECRET,就会在 DB_DRIVER=auto 时请求当前站点 + * 的 /kv-list;没有该代理的站点若恰好把未知路径回退成 HTML 200,便会被误判为 + * KV 可用,随后把 HTML 当 JSON 读取,并让 /env_check 错报 ready=true。 + * + * 只有以下情况才允许无 binding 的代理模式: + * 1. 用户显式指定 DB_DRIVER=kv; + * 2. 用户显式提供 EO_KV_URLS; + * 3. EdgeOne Node 云函数的可靠 SCF 运行时标记存在。 + */ +function shouldProbeProxy(env?: any): boolean { + const configured = String(env?.DB_DRIVER || "") + .trim() + .toLowerCase() + const procEnv: any = + typeof process !== "undefined" ? (process as any).env || {} : {} + return Boolean( + configured === "kv" || + env?.EO_KV_URLS || + env?.TENCENTCLOUD_SCF_FUNCTIONNAME || + procEnv.TENCENTCLOUD_SCF_FUNCTIONNAME, + ) +} + /** * Detect missing configuration for KV proxy mode. * @@ -258,9 +285,13 @@ export const kvDriver: Driver = { } // 模式2: HTTP 代理(EdgeOne Node Functions 拿不到 binding) + // __requestOrigin 本身不能证明代理存在;auto 模式下必须有显式配置或 + // EdgeOne Node 的平台标记,否则未知路由的 HTML 200 会造成假阳性。 + if (!shouldProbeProxy(env)) return false const probe = await probeProxy(env) if (!probeToAvailability(probe)) { - if (probe.error) console.error("[DB] KV proxy unavailable: " + probe.error) + if (probe.error) + console.error("[DB] KV proxy unavailable: " + probe.error) return false } return true @@ -272,7 +303,7 @@ export const kvDriver: Driver = { async get(key: string, env?: any): Promise { const kv = getKvBinding(env) - + // 模式1: Binding 模式 if (kv) { // Cloudflare KV 用 get(key, "text"),EdgeOne KV 用 get(key, {type:"text"}), @@ -312,7 +343,7 @@ export const kvDriver: Driver = { // 绑定误返回对象时统一序列化,保持 Driver.get 的 string 契约 return JSON.stringify(value) } - + // 模式2: HTTP 代理模式(无原生 binding → 经 Edge Function 代理) const baseUrl = requireProxyBaseUrl(env) const url = `${baseUrl}/kv-get?key=${encodeURIComponent(key)}` @@ -330,7 +361,7 @@ export const kvDriver: Driver = { throw new Error(`KV proxy get failed: ${response.status}`) } - const data = await response.json() as { value?: string | null } + const data = (await response.json()) as { value?: string | null } // 归一化为 string | null:Edge Function 在错误分支只返回 { error }, // 此时 data.value 为 undefined,不能直接透传(调用方按 === null 判断会漏掉)。 if (data?.value === undefined || data?.value === null) return null @@ -343,13 +374,13 @@ export const kvDriver: Driver = { async put(key: string, value: string, env?: any): Promise { const kv = getKvBinding(env) - + // 模式1: Binding 模式 if (kv) { await kv.put(key, value) return } - + // 模式2: HTTP 代理模式 const baseUrl = requireProxyBaseUrl(env) const url = `${baseUrl}/kv-put` @@ -372,13 +403,13 @@ export const kvDriver: Driver = { async delete(key: string, env?: any): Promise { const kv = getKvBinding(env) - + // 模式1: Binding 模式 if (kv) { await kv.delete(key) return } - + // 模式2: HTTP 代理模式 const baseUrl = requireProxyBaseUrl(env) const url = `${baseUrl}/kv-delete?key=${encodeURIComponent(key)}` @@ -400,7 +431,7 @@ export const kvDriver: Driver = { async list(prefix: string, env?: any): Promise { const kv = getKvBinding(env) - + // 模式1: Binding 模式 if (kv) { // EdgeOne KV list() 语义(依据官方 functions-kv 示例): @@ -430,7 +461,7 @@ export const kvDriver: Driver = { return keys } - + // 模式2: HTTP 代理模式 const baseUrl = requireProxyBaseUrl(env) const url = `${baseUrl}/kv-list?prefix=${encodeURIComponent(prefix)}` @@ -445,7 +476,7 @@ export const kvDriver: Driver = { throw new Error(`KV proxy list failed: ${response.status}`) } - const data = await response.json() as { keys: string[] } + const data = (await response.json()) as { keys: string[] } return data.keys || [] } catch (err) { console.error(`[KV] list(${prefix}) failed:`, err) @@ -455,7 +486,7 @@ export const kvDriver: Driver = { async health(env?: any): Promise { const kv = getKvBinding(env) - + // 模式1: Binding 模式 if (kv) { try { @@ -475,7 +506,7 @@ export const kvDriver: Driver = { } } } - + // 模式2: HTTP 代理模式。 // // 判定比 isAvailable 更严格(见 probeToHealth / probeToAvailability 的注释): diff --git a/src/backend/server/auth.ts b/src/backend/server/auth.ts index e903f475..a14447d9 100644 --- a/src/backend/server/auth.ts +++ b/src/backend/server/auth.ts @@ -42,7 +42,10 @@ const LOGIN_MAX_FAILURES = 5 const LOGIN_MAX_FAILURES_GLOBAL = 20 const LOGIN_LOCK_MS = 15 * 60 * 1000 const LOGIN_MAX_LOCK_MS = 24 * 60 * 60 * 1000 // 最长锁定 24 小时 -const loginFailures = new Map() +const loginFailures = new Map< + string, + { count: number; lockedUntil: number; attempts: number } +>() function clientIpOf(c: Context): string { return ( @@ -78,11 +81,11 @@ function calculateLockDuration(attempts: number): number { async function bumpLoginFailure( key: string, maxFailures: number, - env: any + env: any, ): Promise { const now = Date.now() let rec = loginFailures.get(key) || { count: 0, lockedUntil: 0, attempts: 0 } - + // 尝试从 KV 读取(多实例共享) try { const { getKvBinding } = await import("../internal/model/db") @@ -109,23 +112,23 @@ async function bumpLoginFailure( } catch { // KV 不可用,回退到内存模式 } - + if (rec.lockedUntil > now) return // already locked - + rec.count += 1 rec.attempts += 1 - + if (rec.count >= maxFailures) { const lockDuration = calculateLockDuration(rec.attempts) rec.lockedUntil = now + lockDuration rec.count = 0 console.warn( - `[Auth] Login attempts exceeded for ${key}. Locked for ${Math.round(lockDuration / 60000)} minutes (attempt #${rec.attempts}).` + `[Auth] Login attempts exceeded for ${key}. Locked for ${Math.round(lockDuration / 60000)} minutes (attempt #${rec.attempts}).`, ) } - + loginFailures.set(key, rec) - + // 持久化到 KV try { const { getKvBinding } = await import("../internal/model/db") @@ -144,7 +147,11 @@ async function bumpLoginFailure( } } -async function isLoginLocked(c: Context, username: string, env: any): Promise { +async function isLoginLocked( + c: Context, + username: string, + env: any, +): Promise { // 懒清理:Map 过大时清掉已过锁定期/无锁定的条目,防止无限增长 if (loginFailures.size > 10000) { const now = Date.now() @@ -152,15 +159,15 @@ async function isLoginLocked(c: Context, username: string, env: any): Promise now) return true if (grec && grec.lockedUntil > now) return true return false } -async function recordLoginFailure(c: Context, username: string, env: any): Promise { +async function recordLoginFailure( + c: Context, + username: string, + env: any, +): Promise { await bumpLoginFailure(loginKey(c, username), LOGIN_MAX_FAILURES, env) - await bumpLoginFailure(globalLoginKey(username), LOGIN_MAX_FAILURES_GLOBAL, env) + await bumpLoginFailure( + globalLoginKey(username), + LOGIN_MAX_FAILURES_GLOBAL, + env, + ) } -async function clearLoginFailures(c: Context, username: string, env: any): Promise { +async function clearLoginFailures( + c: Context, + username: string, + env: any, +): Promise { const ipKey = loginKey(c, username) const globalKey = globalLoginKey(username) - + loginFailures.delete(ipKey) loginFailures.delete(globalKey) - + // 从 KV 中删除 try { const { getKvBinding } = await import("../internal/model/db") @@ -230,7 +249,9 @@ async function clearLoginFailures(c: Context, username: string, env: any): Promi // OpenList/AList StaticHash —— 前端 /login/hash 提交的哈希算法。 // 与新密码模块 pkg/password.staticHash 一致,保留旧名以兼容既有调用。 -export async function hashPasswordSHA256(plainPassword: string): Promise { +export async function hashPasswordSHA256( + plainPassword: string, +): Promise { return staticHash(plainPassword) } @@ -245,8 +266,12 @@ export async function verifyUserStaticHash( user: any, inputStatic: string, ): Promise { - const stored = String(user?.password || "").trim().toLowerCase() - const input = String(inputStatic || "").trim().toLowerCase() + const stored = String(user?.password || "") + .trim() + .toLowerCase() + const input = String(inputStatic || "") + .trim() + .toLowerCase() if (!stored || !isHex64(input) || !isHex64(stored)) return false if (user?.salt) { const expect = (await saltedHash(input, String(user.salt))).toLowerCase() @@ -483,7 +508,7 @@ authRouter.post("/login", async (c) => { ) } await clearLoginFailures(c, username, c.env) - + // 生成 JWT Token const payload = { id: matchedUser.id, @@ -494,14 +519,14 @@ authRouter.post("/login", async (c) => { } const secret = await getJwtSecret(c) const token = await sign(payload, secret) - + // 生成 CSRF Token const csrfToken = setCSRFToken(c) - + // 记录审计日志 const auditLogger = getAuditLogger() await auditLogger.logLoginSuccess(c, username) - + return c.json({ code: 200, message: "success", @@ -513,7 +538,7 @@ authRouter.post("/login", async (c) => { // 登录失败 - 记录审计日志 const auditLogger = getAuditLogger() await auditLogger.logLoginFailure(c, username, "Invalid credentials") - + await recordLoginFailure(c, username, c.env) return c.json({ code: 401, message: "Invalid credentials", data: null }, 401) }) @@ -566,7 +591,7 @@ authRouter.post("/login/hash", async (c) => { ) } await clearLoginFailures(c, username, c.env) - + // 生成 JWT Token const payload = { id: matchedUser.id, @@ -577,14 +602,14 @@ authRouter.post("/login/hash", async (c) => { } const secret = await getJwtSecret(c) const token = await sign(payload, secret) - + // 生成 CSRF Token const csrfToken = setCSRFToken(c) - + // 记录审计日志 const auditLogger = getAuditLogger() await auditLogger.logLoginSuccess(c, username) - + return c.json({ code: 200, message: "success", @@ -596,7 +621,7 @@ authRouter.post("/login/hash", async (c) => { // 登录失败 - 记录审计日志 const auditLogger = getAuditLogger() await auditLogger.logLoginFailure(c, username, "Invalid credentials") - + await recordLoginFailure(c, username, c.env) return c.json({ code: 401, message: "Invalid credentials", data: null }, 401) }) @@ -683,14 +708,14 @@ export const logoutHandler = async (c: any) => { // token 无效则无需注销 } } - + // 清除 CSRF Token clearCSRFToken(c) - + // 记录审计日志 const auditLogger = getAuditLogger() await auditLogger.logLogout(c) - + return c.json({ code: 200, message: "success", @@ -733,10 +758,11 @@ authRouter.post("/2fa/generate", async (c) => { } const secret = generateTotpSecret() const otpauth = buildOtpauthUrl(secret, user.username) + const qr = await buildQrImageUrl(otpauth) return c.json({ code: 200, message: "success", - data: { qr: buildQrImageUrl(otpauth), secret }, + data: { qr, secret }, }) }) diff --git a/src/backend/server/auth_2fa.test.ts b/src/backend/server/auth_2fa.test.ts new file mode 100644 index 00000000..54ada784 --- /dev/null +++ b/src/backend/server/auth_2fa.test.ts @@ -0,0 +1,59 @@ +import assert from "node:assert/strict" +import test from "node:test" +import { Hono } from "hono" +import { sign } from "hono/jwt" +import { saveDb } from "../internal/model/db" +import { authRouter } from "./auth" + +test("2fa/generate returns a resolved QR data URL", async () => { + const env: any = { + JWT_SECRET: "test-only-jwt-secret", + } + await saveDb( + { + settings: [], + users: [ + { + id: 1, + username: "admin", + password: "unused", + role: 2, + permission: 0, + base_path: "/", + disabled: false, + }, + ], + storages: [], + shares: [], + }, + env, + { force: true }, + ) + + const token = await sign( + { + id: 1, + username: "admin", + role: 2, + exp: Math.floor(Date.now() / 1000) + 60, + jti: "auth-2fa-generate-test", + }, + env.JWT_SECRET, + ) + const app = new Hono() + app.route("/api/auth", authRouter) + + const res = await app.request( + "/api/auth/2fa/generate", + { + method: "POST", + headers: { Authorization: `Bearer ${token}` }, + }, + env, + ) + assert.equal(res.status, 200) + const json: any = await res.json() + assert.equal(json.code, 200) + assert.match(json.data.qr, /^data:image\/png;base64,/) + assert.match(json.data.secret, /^[A-Z2-7]+$/) +})