个人信息
-
-
ℹ️
-
-
个人资料编辑暂未开放
-
当前仅展示登录用户信息,资料修改接口接入后再开放保存。
-
-
-
+
+
+
微信账号
+
+
+
💬
+
+ {wechatBound ? (
+ <>
+
+ 已绑定微信{user?.wechat_nickname ? `(${user.wechat_nickname})` : ""}
+
+
可使用微信扫码登录本账号
+ >
+ ) : (
+ <>
+
未绑定微信
+
绑定后可使用微信扫码快速登录
+ >
+ )}
+
+
+
+ {wechatBound ? (
+
+ ) : (
+
+ )}
+
+
+
+
+
setWechatBindOpen(false)}
+ onBindSuccess={handleBindSuccess}
+ />
)
}
diff --git a/apps/web/src/router/ProtectedRoute.tsx b/apps/web/src/router/ProtectedRoute.tsx
index 2c49851ff..943511cc3 100644
--- a/apps/web/src/router/ProtectedRoute.tsx
+++ b/apps/web/src/router/ProtectedRoute.tsx
@@ -6,10 +6,16 @@ import { useAuthStore } from "@/store/authStore"
export const ProtectedRoute = ({ children }: { children: React.ReactNode }) => {
const isAuthenticated = useAuthStore((state) => state.isAuthenticated)
const hasAccessToken = Boolean(localStorage.getItem("access_token"))
+ const profileCompleted = useAuthStore((state) => state.user?.profile_completed !== false)
if (!isAuthenticated || !hasAccessToken) {
return
}
+ // 微信新用户未完成昵称引导时,禁止进入主界面
+ if (!profileCompleted) {
+ return
+ }
+
return <>{children}>
}
diff --git a/apps/web/src/router/appRoutes.tsx b/apps/web/src/router/appRoutes.tsx
index 2858217c3..579b1b9a6 100644
--- a/apps/web/src/router/appRoutes.tsx
+++ b/apps/web/src/router/appRoutes.tsx
@@ -1,6 +1,7 @@
import { Navigate, type RouteObject } from "react-router-dom"
import MainLayout from "@/components/layout/MainLayout"
import { ProtectedRoute } from "./ProtectedRoute"
+import { lazyRoute } from "./lazyRoute"
/**
* 受保护的 /app 子路由
@@ -13,202 +14,118 @@ const appChildren: RouteObject[] = [
},
{
path: "dashboard",
- lazy: () =>
- import("@/pages/dashboard/Dashboard").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/dashboard/Dashboard")),
},
{
path: "assets",
- lazy: () =>
- import("@/pages/assets/AssetLibrary").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/assets/AssetLibrary")),
},
{
path: "titles",
- lazy: () =>
- import("@/pages/titles/TitleLibrary").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/titles/TitleLibrary")),
},
{
path: "voices",
- lazy: () =>
- import("@/pages/voices/VoiceLibrary").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/voices/VoiceLibrary")),
},
{
path: "templates",
- lazy: () =>
- import("@/pages/templates/TemplateLibrary").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/templates/TemplateLibrary")),
},
{
path: "generate",
- lazy: () =>
- import("@/pages/generate/GeneratePage").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/generate/GeneratePage")),
},
{
path: "history",
- lazy: () =>
- import("@/pages/history/TaskHistory").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/history/TaskHistory")),
},
{
path: "products",
- lazy: () =>
- import("@/pages/products/ProductLibrary").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/products/ProductLibrary")),
},
{
path: "products/:id",
- lazy: () =>
- import("@/pages/products/ProductDetail").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/products/ProductDetail")),
},
{
path: "tasks",
- lazy: () =>
- import("@/pages/tasks/TaskCenter").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/tasks/TaskCenter")),
},
{
path: "editing-planner",
- lazy: () =>
- import("@/pages/editing-planner/EditingPlanner").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/editing-planner/EditingPlanner")),
},
{
path: "my-templates",
- lazy: () =>
- import("@/pages/my-templates/MyTemplates").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/my-templates/MyTemplates")),
},
{
path: "voice-clone",
- lazy: () =>
- import("@/pages/voice-clone/VoiceClone").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/voice-clone/VoiceClone")),
},
{
path: "voice-materials",
- lazy: () =>
- import("@/pages/voice-materials/VoiceMaterialLibrary").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/voice-materials/VoiceMaterialLibrary")),
},
{
path: "my-voices",
- lazy: () =>
- import("@/pages/my-voices/MyVoices").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/my-voices/MyVoices")),
},
{
path: "accounts",
- lazy: () =>
- import("@/pages/accounts/Accounts").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/accounts/Accounts")),
},
{
path: "duplication",
- lazy: () =>
- import("@/pages/duplication/DuplicationUpload").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/duplication/DuplicationUpload")),
},
{
path: "duplication/results",
- lazy: () =>
- import("@/pages/duplication/DuplicationResults").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/duplication/DuplicationResults")),
},
{
path: "duplication/:id",
- lazy: () =>
- import("@/pages/duplication/DuplicationDetail").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/duplication/DuplicationDetail")),
},
{
path: "subscription",
- lazy: () =>
- import("@/pages/subscription/Plans").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/subscription/Plans")),
},
{
path: "subscription/upgrade",
- lazy: () =>
- import("@/pages/subscription/UpgradeSubscription").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/subscription/UpgradeSubscription")),
},
{
path: "subscription/billing",
- lazy: () =>
- import("@/pages/subscription/Billing").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/subscription/Billing")),
},
{
path: "profile",
- lazy: () =>
- import("@/pages/profile/Settings").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/profile/Settings")),
},
{
path: "admin",
children: [
{
index: true,
- lazy: () =>
- import("@/pages/admin/AdminComingSoon").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/admin/AdminComingSoon")),
},
{
path: "users",
- lazy: () =>
- import("@/pages/admin/AdminComingSoon").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/admin/AdminComingSoon")),
},
{
path: "analytics",
- lazy: () =>
- import("@/pages/admin/AdminComingSoon").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/admin/AdminComingSoon")),
},
{
path: "monitor",
- lazy: () =>
- import("@/pages/admin/AdminComingSoon").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/admin/AdminComingSoon")),
},
{
path: "logs",
- lazy: () =>
- import("@/pages/admin/AdminComingSoon").then((m) => ({
- Component: m.default,
- })),
+ lazy: lazyRoute(() => import("@/pages/admin/AdminComingSoon")),
},
],
},
diff --git a/apps/web/src/router/lazyRoute.ts b/apps/web/src/router/lazyRoute.ts
new file mode 100644
index 000000000..2a1970935
--- /dev/null
+++ b/apps/web/src/router/lazyRoute.ts
@@ -0,0 +1,40 @@
+import type { LazyRouteFunction, RouteObject } from "react-router-dom"
+import { isChunkLoadError } from "@/utils/chunkLoadError"
+
+/**
+ * 给 React Router data router 的路由懒加载包一层自动重试:
+ *
+ * - 网络抖动 / 瞬态失败:自动重试最多 2 次(间隔 300ms / 800ms),用户无感恢复
+ * - 发版后旧 chunk 404(chunk 文件名已不存在):重试也拿不到旧文件名,
+ * 重试耗尽后抛出,由全局 ChunkErrorBoundary 捕获并引导整页刷新
+ * (刷新后 index.html 是 no-cache 的,会拿到新 chunk 引用)
+ */
+const RETRY_DELAYS_MS = [300, 800]
+const RETRY_COUNT = RETRY_DELAYS_MS.length
+
+const sleep = (ms: number) => new Promise((r) => setTimeout(r, ms))
+
+export const lazyRoute = (
+ factory: () => Promise<{ default: React.ComponentType }>,
+): LazyRouteFunction
=> {
+ return async () => {
+ let lastError: unknown
+ for (let attempt = 0; attempt <= RETRY_COUNT; attempt++) {
+ try {
+ const mod = await factory()
+ if (!mod.default) {
+ throw new Error("lazyRoute: 目标模块缺少 default 导出")
+ }
+ return { Component: mod.default }
+ } catch (err) {
+ lastError = err
+ // 非 chunk 加载错误(代码 bug 等)立即抛出,不浪费重试
+ if (!isChunkLoadError(err)) throw err
+ if (attempt < RETRY_COUNT) {
+ await sleep(RETRY_DELAYS_MS[attempt])
+ }
+ }
+ }
+ throw lastError
+ }
+}
diff --git a/apps/web/src/router/publicRoutes.tsx b/apps/web/src/router/publicRoutes.tsx
index e610b941d..48c7245fb 100644
--- a/apps/web/src/router/publicRoutes.tsx
+++ b/apps/web/src/router/publicRoutes.tsx
@@ -5,6 +5,8 @@ import Register from "@/pages/auth/Register"
import ForgotPassword from "@/pages/auth/ForgotPassword"
import ResetPassword from "@/pages/auth/ResetPassword"
import WechatCallback from "@/pages/auth/WechatCallback"
+import WechatOnboarding from "@/pages/auth/WechatOnboarding"
+import WechatBindCallback from "@/pages/auth/WechatBindCallback"
import { useAuthStore } from "@/store/authStore"
/** 首页路由组件:已登录跳 dashboard,未登录显示落地页 */
@@ -45,4 +47,12 @@ export const publicRoutes: RouteObject[] = [
path: "/auth/wechat/callback",
element: ,
},
+ {
+ path: "/auth/wechat/bind/callback",
+ element: ,
+ },
+ {
+ path: "/welcome/wechat",
+ element: ,
+ },
]
diff --git a/apps/web/src/store/authStore.ts b/apps/web/src/store/authStore.ts
index cef495665..3f6a2991c 100644
--- a/apps/web/src/store/authStore.ts
+++ b/apps/web/src/store/authStore.ts
@@ -13,6 +13,12 @@ interface User {
display_name: string
is_email_verified: boolean
email_verified: boolean
+ wechat_bound?: boolean
+ wechat_nickname?: string
+ avatar_url?: string
+ phone?: string
+ phone_verified?: boolean
+ profile_completed?: boolean
}
interface AuthState {
diff --git a/apps/web/src/test/api/assets.test.ts b/apps/web/src/test/api/assets.test.ts
index 71f5788d7..cce914b2e 100644
--- a/apps/web/src/test/api/assets.test.ts
+++ b/apps/web/src/test/api/assets.test.ts
@@ -242,6 +242,20 @@ describe("assets API", () => {
await expect(completeDirectUpload({ name: "test-item" })).resolves.not.toThrow()
})
+ it("请求体携带 file_size(后端同名兜底去重的大小校验依赖它)", async () => {
+ await completeDirectUpload({
+ project_id: "p-1",
+ library_id: "l-1",
+ storage_key: "uploads/k.mp4",
+ file_size: 12345,
+ } as never)
+ const completeCalls = mockPost.mock.calls.filter(
+ ([u]: [string]) => u === "/upload/direct/complete",
+ )
+ expect(completeCalls).toHaveLength(1)
+ expect(completeCalls[0][1]).toMatchObject({ file_size: 12345 })
+ })
+
it("should reject on API error", async () => {
mockGet.mockRejectedValue(new Error("Network error"))
mockPost.mockRejectedValue(new Error("Network error"))
@@ -259,6 +273,114 @@ describe("assets API", () => {
})
})
+ describe("uploadAssetDirect skip_transfer 短路", () => {
+ it("prepare 返回 skip_transfer=true → 直接返回 duplicated,不调 transfer/complete", async () => {
+ mockPost.mockImplementation((url: string) => {
+ if (url === "/upload/direct/prepare") {
+ return Promise.resolve({
+ data: {
+ upload_url: "https://oss/x",
+ method: "POST",
+ storage_key: "uploads/skip/y.mp4",
+ expires_at: "2099",
+ fields: {},
+ max_size_bytes: 1e9,
+ asset_id: "existing-asset",
+ skip_transfer: true,
+ duplicated: true,
+ },
+ })
+ }
+ if (url === "/upload/direct/complete") {
+ throw new Error("complete 不应被调用")
+ }
+ throw new Error("unexpected url " + url)
+ })
+ const putSpy = vi.spyOn(globalThis, "XMLHttpRequest")
+ const file = new File(["x"], "x.mp4", { type: "video/mp4" })
+ const result = await uploadAssetDirect({ file, library_id: "lib-1" })
+ expect(result.duplicated).toBe(true)
+ expect(result.asset_id).toBe("existing-asset")
+ // complete 未被调用(mockPost 只记录 prepare,complete 若调用会抛 "不应被调用")
+ const completeCalls = mockPost.mock.calls.filter(
+ ([u]: [string]) => u === "/upload/direct/complete",
+ )
+ expect(completeCalls).toHaveLength(0)
+ putSpy.mockRestore()
+ })
+
+ it("prepare 返回 skip_transfer=false → 走老流程(complete 被调用)", async () => {
+ mockPost.mockImplementation((url: string) => {
+ if (url === "/upload/direct/prepare") {
+ return Promise.resolve({
+ data: {
+ upload_url: "https://oss/x",
+ method: "POST",
+ storage_key: "uploads/normal/y.mp4",
+ expires_at: "2099",
+ fields: {},
+ max_size_bytes: 1e9,
+ asset_id: "new-asset",
+ },
+ })
+ }
+ if (url === "/upload/direct/complete") {
+ return Promise.resolve({
+ data: {
+ storage_key: "uploads/normal/y.mp4",
+ ingest_job_id: "job-1",
+ url: "https://oss/y.mp4",
+ duplicated: false,
+ asset_id: "new-asset",
+ },
+ })
+ }
+ throw new Error("unexpected url " + url)
+ })
+ // mock XMLHttpRequest:send 之后下一 tick 触发 onload 让 transfer 立即成功
+ const origOpen = XMLHttpRequest.prototype.open
+ const origSend = XMLHttpRequest.prototype.send
+ const origSetReadyState = Object.getOwnPropertyDescriptor(
+ XMLHttpRequest.prototype,
+ "readyState",
+ ) as PropertyDescriptor | undefined
+ const origStatus = Object.getOwnPropertyDescriptor(XMLHttpRequest.prototype, "status")
+ Object.defineProperty(XMLHttpRequest.prototype, "readyState", {
+ configurable: true,
+ writable: true,
+ value: 4,
+ })
+ Object.defineProperty(XMLHttpRequest.prototype, "status", {
+ configurable: true,
+ writable: true,
+ value: 200,
+ })
+ XMLHttpRequest.prototype.open = vi.fn() as unknown as typeof origOpen
+ XMLHttpRequest.prototype.send = vi.fn(function (this: XMLHttpRequest) {
+ // 下一 tick 触发 onload(模拟 XHR 异步完成)
+ setTimeout(() => this.onload?.(new ProgressEvent("load")), 0)
+ }) as unknown as typeof origSend
+ const file = new File(["x"], "x.mp4", { type: "video/mp4" })
+ const result = await uploadAssetDirect({ file, library_id: "lib-1" })
+ expect(result.duplicated).toBeFalsy()
+ expect(result.asset_id).toBe("new-asset")
+ const completeCalls = mockPost.mock.calls.filter(
+ ([u]: [string]) => u === "/upload/direct/complete",
+ )
+ expect(completeCalls).toHaveLength(1)
+ // complete 请求必须带上 file_size,否则后端同名兜底会误杀同名新视频
+ expect(completeCalls[0][1]).toMatchObject({ file_size: file.size })
+ XMLHttpRequest.prototype.open = origOpen
+ XMLHttpRequest.prototype.send = origSend
+ if (origSetReadyState) {
+ Object.defineProperty(XMLHttpRequest.prototype, "readyState", origSetReadyState)
+ }
+ if (origStatus) {
+ Object.defineProperty(XMLHttpRequest.prototype, "status", origStatus)
+ }
+ })
+ })
+
describe("getIngestJob", () => {
it("should resolve successfully", async () => {
await expect(getIngestJob("test-jobId")).resolves.not.toThrow()
diff --git a/apps/web/src/test/api/uploadDedup.test.ts b/apps/web/src/test/api/uploadDedup.test.ts
new file mode 100644
index 000000000..526df30fb
--- /dev/null
+++ b/apps/web/src/test/api/uploadDedup.test.ts
@@ -0,0 +1,130 @@
+/**
+ * 上传去重/幂等工具单测(Issue #1714)
+ */
+import { describe, it, expect, vi } from "vitest"
+import {
+ computeFileHash,
+ findDuplicateInQueue,
+ HASH_FULL_READ_LIMIT,
+ HASH_SAMPLE_CHUNK,
+ makeClientUploadId,
+ makeFileFingerprint,
+} from "@/api/assets/uploadDedup"
+
+const makeFile = (name: string, size = 100, lastModified = 1_700_000_000_000) =>
+ new File([new Uint8Array(size)], name, { type: "video/mp4", lastModified })
+
+describe("makeFileFingerprint", () => {
+ it("同一文件(name+size+lastModified 相同)指纹一致", () => {
+ const a = makeFile("a.mp4", 1000, 12345)
+ const b = makeFile("a.mp4", 1000, 12345)
+ expect(makeFileFingerprint(a)).toBe(makeFileFingerprint(b))
+ })
+
+ it("文件名/大小/修改时间任一不同指纹即不同", () => {
+ const base = makeFile("a.mp4", 1000, 100)
+ expect(makeFileFingerprint(base)).not.toBe(makeFileFingerprint(makeFile("b.mp4", 1000, 100)))
+ expect(makeFileFingerprint(base)).not.toBe(makeFileFingerprint(makeFile("a.mp4", 1001, 100)))
+ expect(makeFileFingerprint(base)).not.toBe(makeFileFingerprint(makeFile("a.mp4", 1000, 101)))
+ })
+})
+
+describe("findDuplicateInQueue", () => {
+ const queue = [
+ { fileKey: "k1", status: "preparing" },
+ { fileKey: "k2", status: "uploading" },
+ { fileKey: "k3", status: "ingesting" },
+ { fileKey: "k4", status: "done" },
+ { fileKey: "k5", status: "error" },
+ ]
+
+ it("在途状态(preparing/uploading/ingesting/done)命中重复", () => {
+ expect(findDuplicateInQueue(queue, "k1")?.status).toBe("preparing")
+ expect(findDuplicateInQueue(queue, "k2")?.status).toBe("uploading")
+ expect(findDuplicateInQueue(queue, "k3")?.status).toBe("ingesting")
+ expect(findDuplicateInQueue(queue, "k4")?.status).toBe("done")
+ })
+
+ it("未命中返回 null", () => {
+ expect(findDuplicateInQueue(queue, "missing")).toBeNull()
+ })
+
+ it("排除 error 状态后,失败项不算重复(允许重新激活)", () => {
+ expect(findDuplicateInQueue(queue, "k5", ["error"])).toBeNull()
+ })
+
+ it("同时排除 done 后,已完成项也不算重复", () => {
+ expect(findDuplicateInQueue(queue, "k4", ["error", "done"])).toBeNull()
+ // 但在途的仍然命中
+ expect(findDuplicateInQueue(queue, "k1", ["error", "done"])).not.toBeNull()
+ })
+})
+
+describe("makeClientUploadId", () => {
+ it("生成带前缀且互不相同的幂等 token", () => {
+ const ids = new Set(Array.from({ length: 20 }, () => makeClientUploadId()))
+ expect(ids.size).toBe(20)
+ for (const id of ids) expect(id.startsWith("up_")).toBe(true)
+ })
+})
+
+describe("computeFileHash", () => {
+ it("相同内容 hash 一致、不同内容 hash 不同", async () => {
+ const f1 = makeFile("a.mp4", 4096)
+ const f2 = makeFile("b.mp4", 4096)
+ // 两个文件都是 0 填充,内容相同 → hash 一致
+ expect(await computeFileHash(f1)).toBe(await computeFileHash(f2))
+
+ const f3 = new File([new Uint8Array(4096).fill(7)], "c.mp4", { type: "video/mp4" })
+ expect(await computeFileHash(f1)).not.toBe(await computeFileHash(f3))
+ })
+
+ it("返回 64 位十六进制(SHA-256,与后端 file_hash 长度一致)", async () => {
+ const hash = await computeFileHash(makeFile("a.mp4", 1024))
+ expect(hash).toMatch(/^[0-9a-f]{64}$/)
+ })
+})
+
+describe("computeFileHash 大文件抽样(>64MB)", () => {
+ it("抽样路径正常返回 64 位 hex,且大小不同则 hash 不同", async () => {
+ // mock 一个「声称」300MB 的 File:slice 返回小 buffer 即可,不真分配 300MB
+ const makeBig = (declaredSize: number, head: number) => {
+ const f = new File([new Uint8Array([head, 2, 3])], "big.mov", { type: "video/quicktime" })
+ Object.defineProperty(f, "size", { value: declaredSize, configurable: true })
+ // slice 仍按真实内容返回小片段(头尾片段内容由底层小 buffer 决定)
+ return f
+ }
+ const h1 = await computeFileHash(makeBig(300 * 1024 * 1024, 1))
+ const h2 = await computeFileHash(makeBig(301 * 1024 * 1024, 1))
+ expect(h1).toMatch(/^[0-9a-f]{64}$/)
+ // 声明大小不同 → 写入的 64 位 size 字段不同 → hash 必须不同(锁定 setBigUint64 路径)
+ expect(h1).not.toBe(h2)
+ })
+
+ it("≤64MB 走全量读取(slice 一次覆盖整个文件)", async () => {
+ const f = new File([new Uint8Array(1024).fill(9)], "full.mp4", { type: "video/mp4" })
+ Object.defineProperty(f, "size", { value: HASH_FULL_READ_LIMIT, configurable: true })
+ const sliceSpy = vi.spyOn(f, "slice")
+ await computeFileHash(f)
+ // 全量路径:唯一一次 slice 为 (0, size)
+ expect(sliceSpy).toHaveBeenCalledTimes(1)
+ expect(sliceSpy).toHaveBeenCalledWith(0, HASH_FULL_READ_LIMIT)
+ sliceSpy.mockRestore()
+ })
+
+ it(">64MB 只读取头尾各 16MB 抽样,绝不整文件读入内存", async () => {
+ const f = new File([new Uint8Array(1024).fill(9)], "big.mp4", { type: "video/mp4" })
+ Object.defineProperty(f, "size", { value: HASH_FULL_READ_LIMIT + 1, configurable: true })
+ const sliceSpy = vi.spyOn(f, "slice")
+ await computeFileHash(f)
+ // 抽样路径:两次 slice —— 头部 (0, 16MB) 与尾部 (size-16MB, size)
+ expect(sliceSpy).toHaveBeenCalledTimes(2)
+ expect(sliceSpy).toHaveBeenNthCalledWith(1, 0, HASH_SAMPLE_CHUNK)
+ expect(sliceSpy).toHaveBeenNthCalledWith(
+ 2,
+ HASH_FULL_READ_LIMIT + 1 - HASH_SAMPLE_CHUNK,
+ HASH_FULL_READ_LIMIT + 1,
+ )
+ sliceSpy.mockRestore()
+ })
+})
diff --git a/apps/web/src/test/api/wxLogin.test.ts b/apps/web/src/test/api/wxLogin.test.ts
new file mode 100644
index 000000000..5cc38730d
--- /dev/null
+++ b/apps/web/src/test/api/wxLogin.test.ts
@@ -0,0 +1,64 @@
+import { describe, expect, it, vi, beforeEach, afterEach } from "vitest"
+
+describe("wxLogin 工具", () => {
+ describe("parseWxAuthUrl", () => {
+ it("从微信授权链接解析出 appid/redirect_uri/state(redirect_uri 解码)", async () => {
+ const { parseWxAuthUrl } = await import("@/api/auth/wxLogin")
+ const authUrl =
+ "https://open.weixin.qq.com/connect/qrconnect?appid=wxb7ae80b48e53980d" +
+ "&redirect_uri=https%3A%2F%2Fstaging.xiaoxiajianji.com%2Fauth%2Fwechat%2Fcallback" +
+ "&response_type=code&scope=snsapi_login&state=abc123#wechat_redirect"
+ const params = parseWxAuthUrl(authUrl)
+ expect(params).not.toBeNull()
+ expect(params?.appid).toBe("wxb7ae80b48e53980d")
+ expect(params?.redirect_uri).toBe("https://staging.xiaoxiajianji.com/auth/wechat/callback")
+ expect(params?.state).toBe("abc123")
+ })
+
+ it("链接里缺 state 时回退使用 stateFallback", async () => {
+ const { parseWxAuthUrl } = await import("@/api/auth/wxLogin")
+ const authUrl =
+ "https://open.weixin.qq.com/connect/qrconnect?appid=wx123" +
+ "&redirect_uri=https%3A%2F%2Fexample.com%2Fcb"
+ const params = parseWxAuthUrl(authUrl, "fallback-state")
+ expect(params?.state).toBe("fallback-state")
+ })
+
+ it("缺 appid 或 redirect_uri 时返回 null(调用方应回退整页跳转)", async () => {
+ const { parseWxAuthUrl } = await import("@/api/auth/wxLogin")
+ expect(parseWxAuthUrl("https://open.weixin.qq.com/connect/qrconnect?appid=wx123")).toBeNull()
+ expect(parseWxAuthUrl("not a url")).toBeNull()
+ })
+ })
+
+ describe("loadWxLoginScript", () => {
+ beforeEach(() => {
+ vi.resetModules()
+ document.head.querySelectorAll("script[src*='wxLogin']").forEach((el) => el.remove())
+ delete (window as unknown as { WxLogin?: unknown }).WxLogin
+ })
+ afterEach(() => {
+ vi.restoreAllMocks()
+ })
+
+ it("window.WxLogin 已存在时直接复用,不重复插入 script", async () => {
+ const fakeCtor = vi.fn()
+ ;(window as unknown as { WxLogin: unknown }).WxLogin = fakeCtor
+ const { loadWxLoginScript } = await import("@/api/auth/wxLogin")
+ const ctor = await loadWxLoginScript()
+ expect(ctor).toBe(fakeCtor)
+ expect(document.head.querySelector("script[src*='wxLogin']")).toBeNull()
+ })
+
+ it("脚本 onerror 时 reject(调用方据此回退整页跳转)", async () => {
+ const { loadWxLoginScript } = await import("@/api/auth/wxLogin")
+ const promise = loadWxLoginScript()
+ const script = document.head.querySelector(
+ "script[src*='wxLogin']",
+ ) as HTMLScriptElement | null
+ expect(script).not.toBeNull()
+ script?.dispatchEvent(new Event("error"))
+ await expect(promise).rejects.toThrow(/加载失败/)
+ })
+ })
+})
diff --git a/apps/web/src/test/components/ChunkErrorBoundary.test.tsx b/apps/web/src/test/components/ChunkErrorBoundary.test.tsx
new file mode 100644
index 000000000..c049dcb95
--- /dev/null
+++ b/apps/web/src/test/components/ChunkErrorBoundary.test.tsx
@@ -0,0 +1,79 @@
+import { describe, it, expect, beforeEach, afterEach, vi } from "vitest"
+import { render, screen, fireEvent } from "@testing-library/react"
+import { Button } from "antd"
+import { useState } from "react"
+import ChunkErrorBoundary from "@/components/common/ChunkErrorBoundary"
+import * as chunkUtils from "@/utils/chunkLoadError"
+
+// reload 函数 mock 掉(jsdom 不支持真实 window.location.reload)
+vi.mock("@/utils/chunkLoadError", async (importOriginal) => {
+ const actual = await importOriginal()
+ return {
+ ...actual,
+ reloadForChunkError: vi.fn(),
+ goHomeRecover: vi.fn(),
+ }
+})
+const { reloadForChunkError, goHomeRecover } = vi.mocked(chunkUtils)
+
+/** 渲染时直接抛错的子组件 */
+const Boom: React.FC<{ error: Error }> = ({ error }) => {
+ throw error
+}
+
+/** 点击按钮后才抛 chunk 错误的子组件 */
+const ChunkBoomButton: React.FC = () => {
+ const [boom, setBoom] = useState(false)
+ if (boom) {
+ throw new TypeError("Failed to fetch dynamically imported module: /assets/x.js")
+ }
+ return
+}
+
+const renderBoundary = (ui: React.ReactNode) =>
+ render({ui})
+
+beforeEach(() => {
+ sessionStorage.clear()
+ vi.clearAllMocks()
+ // error boundary 捕获后 React 会打 error log,静默掉
+ vi.spyOn(console, "error").mockImplementation(() => {})
+})
+
+afterEach(() => {
+ vi.restoreAllMocks()
+ sessionStorage.clear()
+})
+
+describe("ChunkErrorBoundary", () => {
+ it("正常渲染 children", () => {
+ renderBoundary(hello-child
)
+ expect(screen.getByText("hello-child")).toBeInTheDocument()
+ })
+
+ it("首次捕获 chunk 错误 → 自动刷新(reloadForChunkError)并显示自动刷新提示", () => {
+ renderBoundary()
+ fireEvent.click(screen.getByText("boom"))
+ expect(reloadForChunkError).toHaveBeenCalledTimes(1)
+ expect(screen.getByText(/正在自动刷新/)).toBeInTheDocument()
+ })
+
+ it("已刷新过仍失败 → 不再自动刷新,显示手动兜底按钮", () => {
+ // 模拟"本会话已经自动刷新过一次"
+ sessionStorage.setItem("chunk_error_reloaded_at", String(Date.now()))
+ renderBoundary(
+ ,
+ )
+ expect(reloadForChunkError).not.toHaveBeenCalled()
+ expect(screen.getByText("系统已更新")).toBeInTheDocument()
+ // 点击兜底按钮 → goHomeRecover(跳首页,不刷新当前 URL)
+ fireEvent.click(screen.getByText("刷新并返回首页"))
+ expect(goHomeRecover).toHaveBeenCalledTimes(1)
+ })
+
+ it("非 chunk 错误 → 显示通用错误页,不触发 chunk 自动刷新", () => {
+ renderBoundary()
+ expect(reloadForChunkError).not.toHaveBeenCalled()
+ expect(screen.getByText("页面出现异常")).toBeInTheDocument()
+ })
+})
diff --git a/apps/web/src/test/components/WechatQrModal.test.tsx b/apps/web/src/test/components/WechatQrModal.test.tsx
new file mode 100644
index 000000000..c3910e0d5
--- /dev/null
+++ b/apps/web/src/test/components/WechatQrModal.test.tsx
@@ -0,0 +1,165 @@
+import { describe, expect, it, vi, beforeEach, afterEach } from "vitest"
+import { render, screen, waitFor, cleanup, fireEvent } from "@testing-library/react"
+import WechatQrModal from "@/components/auth/WechatQrModal"
+
+const { mockWxLoginCtor, mockGetAuthUrl, mockGetBindUrl, mockGetCurrentUser } = vi.hoisted(() => ({
+ mockWxLoginCtor: vi.fn(),
+ mockGetAuthUrl: vi.fn(),
+ mockGetBindUrl: vi.fn(),
+ mockGetCurrentUser: vi.fn(),
+}))
+
+vi.mock("@/api/auth", () => ({
+ getWechatAuthUrl: (...args: unknown[]) => mockGetAuthUrl(...args),
+ getWechatBindUrl: (...args: unknown[]) => mockGetBindUrl(...args),
+ getCurrentUser: (...args: unknown[]) => mockGetCurrentUser(...args),
+ normalizeUser: (u: unknown) => u,
+}))
+
+vi.mock("@/api/auth/wxLogin", () => ({
+ loadWxLoginScript: vi.fn(async () => mockWxLoginCtor),
+ parseWxAuthUrl: vi.fn(() => ({
+ appid: "wxb7ae80b48e53980d",
+ redirect_uri: "https://staging.xiaoxiajianji.com/auth/wechat/callback",
+ state: "state-from-url",
+ })),
+}))
+
+vi.mock("@/api/auth/tokenRefresh", () => ({
+ scheduleProactiveRefresh: vi.fn(),
+ cancelProactiveRefresh: vi.fn(),
+}))
+
+const { mockSetAuth, mockSetUser } = vi.hoisted(() => ({
+ mockSetAuth: vi.fn(),
+ mockSetUser: vi.fn(),
+}))
+vi.mock("@/store/authStore", () => ({
+ useAuthStore: (selector: (s: unknown) => unknown) =>
+ selector({ setAuth: mockSetAuth, setUser: mockSetUser }),
+}))
+
+const AUTH_URL =
+ "https://open.weixin.qq.com/connect/qrconnect?appid=wxb7ae80b48e53980d" +
+ "&redirect_uri=https%3A%2F%2Fstaging.xiaoxiajianji.com%2Fauth%2Fwechat%2Fcallback&state=st123"
+
+const postMessage = (data: Record) =>
+ window.dispatchEvent(new MessageEvent("message", { data, origin: window.location.origin }))
+
+beforeEach(() => {
+ vi.clearAllMocks()
+ mockGetAuthUrl.mockResolvedValue({ auth_url: AUTH_URL, state: "st123" })
+ mockGetBindUrl.mockResolvedValue({ auth_url: AUTH_URL, state: "st123" })
+ mockGetCurrentUser.mockResolvedValue({ id: 1, display_name: "测试用户" })
+ localStorage.clear()
+})
+
+afterEach(() => cleanup())
+
+describe("WechatQrModal", () => {
+ it("open=false 时不渲染弹窗内容", () => {
+ render()
+ expect(screen.queryByText("微信扫码登录")).toBeNull()
+ })
+
+ it("登录场景:open 后请求授权链接、写入 state、用 WxLogin 渲染二维码", async () => {
+ render()
+ await waitFor(() => expect(mockGetAuthUrl).toHaveBeenCalledTimes(1))
+ expect(localStorage.getItem("wechat_state")).toBe("st123")
+ await waitFor(() => expect(mockWxLoginCtor).toHaveBeenCalledTimes(1))
+ expect(mockWxLoginCtor).toHaveBeenCalledWith(
+ expect.objectContaining({
+ self_redirect: true,
+ appid: "wxb7ae80b48e53980d",
+ scope: "snsapi_login",
+ state: "state-from-url",
+ redirect_uri: "https://staging.xiaoxiajianji.com/auth/wechat/callback",
+ }),
+ )
+ expect(screen.getByText(/请使用微信扫描二维码登录/)).toBeTruthy()
+ })
+
+ it("绑定场景:请求 bind/url 且写入 wechat_bind_state", async () => {
+ render()
+ await waitFor(() => expect(mockGetBindUrl).toHaveBeenCalledTimes(1))
+ expect(mockGetAuthUrl).not.toHaveBeenCalled()
+ expect(localStorage.getItem("wechat_bind_state")).toBe("st123")
+ await waitFor(() => expect(mockWxLoginCtor).toHaveBeenCalledTimes(1))
+ })
+
+ it("获取授权链接失败时弹窗内展示错误并提供刷新", async () => {
+ mockGetAuthUrl.mockRejectedValueOnce({
+ response: { status: 500, data: { detail: "微信服务内部错误" } },
+ })
+ render()
+ expect(await screen.findByText(/微信服务内部错误/)).toBeTruthy()
+ expect(screen.getByText("刷新二维码")).toBeTruthy()
+ // 点刷新后重新请求
+ fireEvent.click(screen.getByText("刷新二维码"))
+ await waitFor(() => expect(mockGetAuthUrl).toHaveBeenCalledTimes(2))
+ })
+
+ it("登录成功消息:同步登录态并回调 onLoginSuccess(needOnboarding)", async () => {
+ const onSuccess = vi.fn()
+ localStorage.setItem("access_token", "tok-123")
+ render()
+ await waitFor(() => expect(mockWxLoginCtor).toHaveBeenCalledTimes(1))
+
+ postMessage({
+ source: "xiaoxia-wechat-qr",
+ scene: "login",
+ success: true,
+ payload: { needOnboarding: true },
+ })
+
+ await waitFor(() => expect(onSuccess).toHaveBeenCalledWith(true))
+ expect(mockGetCurrentUser).toHaveBeenCalled()
+ expect(mockSetAuth).toHaveBeenCalledWith(expect.objectContaining({ id: 1 }), "tok-123", null)
+ })
+
+ it("登录失败消息:弹窗内展示回调页透传的真实原因", async () => {
+ render()
+ await waitFor(() => expect(mockWxLoginCtor).toHaveBeenCalledTimes(1))
+
+ postMessage({
+ source: "xiaoxia-wechat-qr",
+ scene: "login",
+ success: false,
+ detail: "微信登录失败:state 已过期或已被使用",
+ })
+
+ expect(await screen.findByText(/state 已过期或已被使用/)).toBeTruthy()
+ })
+
+ it("绑定成功消息:刷新用户并回调 onBindSuccess", async () => {
+ const onBindSuccess = vi.fn()
+ render()
+ await waitFor(() => expect(mockWxLoginCtor).toHaveBeenCalledTimes(1))
+
+ postMessage({ source: "xiaoxia-wechat-qr", scene: "bind", success: true })
+
+ await waitFor(() => expect(onBindSuccess).toHaveBeenCalledTimes(1))
+ expect(mockSetUser).toHaveBeenCalled()
+ })
+
+ it("忽略跨源消息和其他场景的消息", async () => {
+ const onSuccess = vi.fn()
+ render()
+ await waitFor(() => expect(mockWxLoginCtor).toHaveBeenCalledTimes(1))
+
+ // 跨源
+ window.dispatchEvent(
+ new MessageEvent("message", {
+ data: { source: "xiaoxia-wechat-qr", scene: "login", success: true },
+ origin: "https://evil.example.com",
+ }),
+ )
+ // 场景不符(bind 消息发给 login 弹窗)
+ postMessage({ source: "xiaoxia-wechat-qr", scene: "bind", success: true })
+ // 无协议标识
+ postMessage({ foo: "bar" })
+
+ await new Promise((r) => setTimeout(r, 50))
+ expect(onSuccess).not.toHaveBeenCalled()
+ })
+})
diff --git a/apps/web/src/test/pages/Settings.test.tsx b/apps/web/src/test/pages/Settings.test.tsx
index e0614b0bf..3d0711b64 100644
--- a/apps/web/src/test/pages/Settings.test.tsx
+++ b/apps/web/src/test/pages/Settings.test.tsx
@@ -1,8 +1,8 @@
-import { describe, expect, it, vi } from "vitest"
-import { render, screen } from "@testing-library/react"
+import { describe, expect, it, vi, beforeEach } from "vitest"
+import { render, screen, fireEvent, waitFor } from "@testing-library/react"
import { MemoryRouter } from "react-router-dom"
+import { QueryClient, QueryClientProvider } from "@tanstack/react-query"
-// mock PageHead 简单mock
vi.mock("@/components/layout/PageHead", () => ({
default: ({ title, description }: { title: string; description?: string }) => (
@@ -12,51 +12,142 @@ vi.mock("@/components/layout/PageHead", () => ({
),
}))
+const mockSetUser = vi.fn()
+const mockInvalidate = vi.fn()
+let authState: Record
= {
+ user: {
+ id: "1",
+ user_id: "1",
+ username: "testuser",
+ email: "test@example.com",
+ display_name: "Test User",
+ wechat_bound: false,
+ },
+ isAuthenticated: true,
+ setUser: mockSetUser,
+}
+
vi.mock("@/store/authStore", () => ({
- useAuthStore: (selector: (state: any) => any) =>
- selector({
+ useAuthStore: (selector: (state: unknown) => unknown) => selector(authState),
+}))
+
+const getCurrentUserMock = vi.fn(async () => authState.user as Record)
+const updateProfileMock = vi.fn()
+const getWechatBindUrlMock = vi.fn(async () => ({
+ auth_url: "https://wx.example/auth",
+ state: "s1",
+}))
+const unbindWechatMock = vi.fn(async () => ({ success: true }))
+
+vi.mock("@/api/auth", () => ({
+ getCurrentUser: () => getCurrentUserMock(),
+ updateProfile: (d: unknown) => updateProfileMock(d),
+ getWechatBindUrl: () => getWechatBindUrlMock(),
+ unbindWechat: () => unbindWechatMock(),
+}))
+
+vi.mock("antd", async () => {
+ const actual = await vi.importActual("antd")
+ return { ...actual, message: { success: vi.fn(), error: vi.fn() } }
+})
+
+import Settings from "@/pages/profile/Settings"
+
+const queryClient = new QueryClient({
+ defaultOptions: { queries: { retry: false }, mutations: { retry: false } },
+})
+
+const renderPage = () =>
+ render(
+
+
+
+
+ ,
+ )
+
+describe("Settings Page", () => {
+ beforeEach(() => {
+ vi.clearAllMocks()
+ authState = {
user: {
id: "1",
user_id: "1",
username: "testuser",
email: "test@example.com",
display_name: "Test User",
- is_email_verified: true,
- email_verified: true,
+ wechat_bound: false,
},
isAuthenticated: true,
- }),
-}))
-
-import Settings from "@/pages/profile/Settings"
-
-describe("Settings Page", () => {
- it("should render without crashing", () => {
- render(
-
-
- ,
- )
- expect(screen.getByText("个人设置")).toBeTruthy()
+ setUser: mockSetUser,
+ }
})
- it("should display user info", () => {
- render(
-
-
- ,
- )
+ it("渲染个人设置与用户信息", () => {
+ renderPage()
+ expect(screen.getByText("个人设置")).toBeTruthy()
expect(screen.getByDisplayValue("testuser")).toBeTruthy()
expect(screen.getByDisplayValue("test@example.com")).toBeTruthy()
})
- it("should show save button is disabled", () => {
- render(
-
-
- ,
+ it("未绑定时显示绑定微信按钮,点击跳转微信授权", async () => {
+ renderPage()
+ expect(screen.getByText("未绑定微信")).toBeTruthy()
+ const btn = screen.getByText("绑定微信")
+ fireEvent.click(btn)
+ await waitFor(() => {
+ expect(getWechatBindUrlMock).toHaveBeenCalled()
+ expect(localStorage.getItem("wechat_bind_state")).toBe("s1")
+ })
+ })
+
+ it("已绑定时显示状态与解绑按钮,确认后调解绑接口", async () => {
+ authState.user = {
+ ...(authState.user as object),
+ wechat_bound: true,
+ wechat_nickname: "微信昵称",
+ } as never
+ renderPage()
+ expect(screen.getByText(/已绑定微信/)).toBeTruthy()
+ fireEvent.click(
+ screen.getByText(
+ (_, el) => el?.tagName === "BUTTON" && (el.textContent ?? "").replace(/\s/g, "") === "解绑",
+ ),
)
- const button = screen.getByText("保存暂未开放")
- expect(button).toBeTruthy()
+ // antd Modal.confirm 弹确认框(标题+内容均含"解绑微信",用 role=dialog 内的确认按钮)
+ await waitFor(() => {
+ expect(document.querySelector(".ant-modal-confirm")).toBeTruthy()
+ })
+ fireEvent.click(
+ screen.getByText(
+ (_, el) => el?.tagName === "BUTTON" && (el.textContent ?? "").includes("确定解绑"),
+ ),
+ )
+ await waitFor(() => {
+ expect(unbindWechatMock).toHaveBeenCalled()
+ })
+ })
+
+ it("修改昵称后保存按钮可用,点击调用更新接口", async () => {
+ renderPage()
+ const saveBtn = screen.getByText(
+ (_, el) => el?.tagName === "BUTTON" && (el.textContent ?? "").replace(/\s/g, "") === "保存",
+ )
+ expect(saveBtn.closest("button")?.disabled).toBe(true)
+ fireEvent.change(screen.getByDisplayValue("Test User"), {
+ target: { value: "新昵称" },
+ })
+ await waitFor(() => {
+ expect(saveBtn.closest("button")?.disabled).toBe(false)
+ })
+ updateProfileMock.mockResolvedValueOnce({
+ id: "1",
+ display_name: "新昵称",
+ wechat_bound: false,
+ })
+ fireEvent.click(saveBtn)
+ await waitFor(() => {
+ expect(updateProfileMock).toHaveBeenCalledWith({ display_name: "新昵称" })
+ })
})
})
diff --git a/apps/web/src/test/pages/assets/useAssetUpload.test.tsx b/apps/web/src/test/pages/assets/useAssetUpload.test.tsx
index 3bc612669..2ad8d634c 100644
--- a/apps/web/src/test/pages/assets/useAssetUpload.test.tsx
+++ b/apps/web/src/test/pages/assets/useAssetUpload.test.tsx
@@ -26,22 +26,32 @@ interface FakeHandle {
fields: Record
max_size_bytes: number
asset_id: string
+ duplicated?: boolean
+ skip_transfer?: boolean
}
transfer: ReturnType
complete: ReturnType
/** 手动结束传输(transfer 被调用后挂载);finish(true) 以失败结束 */
finish: (fail?: boolean) => void
+ /** complete 已被调用的次数 */
+ completeCalls: { resolve: () => void; reject: (err: unknown) => void }[]
}
-let activeTransfers = 0
-let maxConcurrent = 0
-
/**
- * 创建一个假 handle:transfer 返回挂起的 promise,
- * finish 槽位在 transfer executor 同步执行时挂载,测试中调用 finish() 控制成败
+ * 创建一个假 handle:
+ * - transfer 返回挂起的 promise,finish()/finish(true) 控制成败
+ * - complete 每次调用返回独立的挂起 promise,由 completeCalls 记录控制,
+ * 成功调 resolve(idx) / 失败调 reject(idx)(模拟超时)
*/
-const makeFakeHandle = (opts: { id: string; duplicated?: boolean; failTransfer?: boolean }) => {
- const h = {
+const makeFakeHandle = (opts: {
+ id: string
+ duplicated?: boolean
+ failTransfer?: boolean
+ completeAuto?: boolean
+ /** prepare 阶段就命中去重:prepare 响应 skip_transfer/duplicated=true */
+ prepareDedup?: boolean
+}) => {
+ const h: FakeHandle = {
prepared: {
upload_url: "https://oss.example.com/u",
method: "POST",
@@ -50,17 +60,43 @@ const makeFakeHandle = (opts: { id: string; duplicated?: boolean; failTransfer?:
fields: {},
max_size_bytes: 2_000_000_000,
asset_id: opts.id,
+ duplicated: opts.prepareDedup ? true : undefined,
+ skip_transfer: opts.prepareDedup ? true : undefined,
},
transfer: vi.fn(),
- complete: vi.fn().mockResolvedValue({
- storage_key: "uploads/x/y.mp4",
- ingest_job_id: opts.duplicated ? "" : "job-1",
- url: "https://oss.example.com/u",
- duplicated: opts.duplicated,
- asset_id: opts.id,
- }),
- finish: (() => {}) as (fail?: boolean) => void,
+ complete: vi.fn(),
+ finish: () => {},
+ completeCalls: [],
}
+
+ h.complete.mockImplementation(
+ () =>
+ new Promise<{
+ storage_key: string
+ ingest_job_id: string
+ url: string
+ duplicated: boolean
+ asset_id: string
+ }>((resolve, reject) => {
+ h.completeCalls.push({
+ resolve: () =>
+ resolve({
+ storage_key: "uploads/x/y.mp4",
+ ingest_job_id: opts.duplicated ? "" : `job-${opts.id}`,
+ url: "https://oss.example.com/u",
+ duplicated: !!opts.duplicated,
+ asset_id: opts.id,
+ }),
+ reject,
+ })
+ // 默认立即成功,保持旧用例简单
+ if (opts.completeAuto !== false) {
+ const idx = h.completeCalls.length - 1
+ Promise.resolve().then(() => h.completeCalls[idx]?.resolve())
+ }
+ }),
+ )
+
h.transfer.mockImplementation(
() =>
new Promise((_resolve, reject) => {
@@ -78,17 +114,24 @@ const makeFakeHandle = (opts: { id: string; duplicated?: boolean; failTransfer?:
type FakeHandleLike = ReturnType
+let activeTransfers = 0
+let maxConcurrent = 0
+
/** prepare mock:调用序号生成稳定 id,立即把 handle(含 finish 槽位)推入数组 */
const installPrepareMock = (
handles: FakeHandleLike[],
- optOverrides?: (id: string) => { duplicated?: boolean; failTransfer?: boolean },
+ optOverrides?: (id: string) => {
+ duplicated?: boolean
+ failTransfer?: boolean
+ completeAuto?: boolean
+ },
) => {
let callNo = 0
;(prepareDirectUploadHandle as unknown as ReturnType).mockImplementation(
async () => {
const id = `asset-${callNo++}`
const overrides = optOverrides?.(id) ?? {}
- const h = makeFakeHandle({ id, ...overrides })
+ const h = makeFakeHandle({ id, completeAuto: true, ...overrides })
handles.push(h)
await new Promise((r) => setTimeout(r, 10))
return h
@@ -187,6 +230,11 @@ describe("useAssetUpload", () => {
})
await waitFor(() => expect(result.current.uploadItems[0].status).toBe("error"))
+ // 失败卡片记录失败阶段与完整错误原因(不再只显示"上传失败")
+ const failed = result.current.uploadItems[0]
+ expect(failed.failedStage).toBe("transfer")
+ expect(failed.error).toContain("OSS boom")
+
// 重试:重新 prepare(handles[1] 成功)
const tempId = result.current.uploadItems[0].tempId
await act(async () => {
@@ -223,4 +271,142 @@ describe("useAssetUpload", () => {
expect(result.current.uploadItems[0].status).toBe("done")
})
})
+ it("同一文件多次选择不重复入队(指纹去重)", async () => {
+ const handles: FakeHandleLike[] = []
+ installPrepareMock(handles)
+
+ const { result } = renderHook(() => useAssetUpload({ effectiveLibId: "lib-1" }), {
+ wrapper: createWrapper(),
+ })
+
+ // 同一文件(name+size+lastModified 完全一致)第一次入队
+ const sameFile = mp4("same.mp4")
+ await act(async () => {
+ result.current.enqueueUploads([sameFile])
+ })
+ await waitFor(() => expect(handles.length).toBe(1))
+ expect(result.current.uploadItems).toHaveLength(1)
+
+ // transfer 挂起期间,再次选择同一文件(模拟用户反复点选/拖拽)
+ await act(async () => {
+ result.current.enqueueUploads([sameFile])
+ })
+ await act(async () => {
+ result.current.enqueueUploads([sameFile])
+ })
+ // 队列表只有 1 项、prepare 只有 1 次
+ expect(result.current.uploadItems).toHaveLength(1)
+ expect(handles.length).toBe(1)
+
+ // 完成后再次重复选择(已 done):仍然不新增
+ await act(async () => {
+ handles[0].finish()
+ })
+ await waitFor(() => expect(result.current.uploadItems[0].status).toBe("done"))
+ await act(async () => {
+ result.current.enqueueUploads([sameFile])
+ })
+ expect(result.current.uploadItems).toHaveLength(1)
+ expect(handles.length).toBe(1)
+
+ // 不同文件正常入队
+ await act(async () => {
+ result.current.enqueueUploads([mp4("other.mp4")])
+ })
+ await waitFor(() => expect(handles.length).toBe(2))
+ expect(result.current.uploadItems).toHaveLength(2)
+ })
+
+ it("complete 失败(超时)后重试:只重发 complete,不重新 prepare/直传", async () => {
+ const handles: FakeHandleLike[] = []
+ installPrepareMock(handles, () => ({ completeAuto: false }))
+
+ const { result } = renderHook(() => useAssetUpload({ effectiveLibId: "lib-1" }), {
+ wrapper: createWrapper(),
+ })
+
+ await act(async () => {
+ result.current.enqueueUploads([mp4("slow.mp4")])
+ })
+ await waitFor(() => expect(handles.length).toBe(1))
+ await waitFor(() => expect(handles[0].transfer).toHaveBeenCalled())
+ await act(async () => {
+ handles[0].finish()
+ })
+ // complete 被调用但挂起
+ await waitFor(() => expect(handles[0].complete).toHaveBeenCalledTimes(1))
+
+ // 模拟 complete 超时(后端记录可能已建成)
+ await act(async () => {
+ handles[0].completeCalls[0]?.reject(new Error("complete timeout (ECONNABORTED)"))
+ })
+ const tempId = result.current.uploadItems[0].tempId
+ await waitFor(() => {
+ const it = result.current.uploadItems.find((x) => x.tempId === tempId)
+ expect(it?.status).toBe("error")
+ expect(it?.failedStage).toBe("complete")
+ // 卡片同时展示真实失败原因与"重试不会重新上传"提示
+ expect(it?.error).toContain("complete timeout")
+ expect(it?.error).toContain("不会重新上传文件")
+ })
+
+ // 点重试:pump 复用 handle,只再调一次 complete(transfer/prepare 不重复)
+ await act(async () => {
+ result.current.retryUpload(tempId)
+ })
+ await waitFor(() => expect(handles[0].complete).toHaveBeenCalledTimes(2))
+ expect(handles.length).toBe(1) // 没有重新 prepare
+ expect(handles[0].transfer).toHaveBeenCalledTimes(1) // 没有重新直传
+
+ // 第二次 complete 成功
+ await act(async () => {
+ handles[0].completeCalls[1]?.resolve()
+ })
+ await waitFor(() => {
+ expect(result.current.uploadItems.find((x) => x.tempId === tempId)?.status).toBe("done")
+ })
+ })
+
+ it("prepare 阶段失败:标记 prepare 阶段并保留后端错误明细", async () => {
+ ;(prepareDirectUploadHandle as unknown as ReturnType).mockRejectedValueOnce({
+ isAxiosError: true,
+ response: { status: 500, data: { detail: "签名服务内部错误" } },
+ message: "Request failed with status code 500",
+ })
+
+ const { result } = renderHook(() => useAssetUpload({ effectiveLibId: "lib-1" }), {
+ wrapper: createWrapper(),
+ })
+
+ await act(async () => {
+ result.current.enqueueUploads([mp4("prep-fail.mp4")])
+ })
+ await waitFor(() => expect(result.current.uploadItems[0]?.status).toBe("error"))
+ const it = result.current.uploadItems[0]
+ expect(it.failedStage).toBe("prepare")
+ expect(it.error).toContain("签名服务内部错误")
+ })
+ it("prepare 返回 skip_transfer=true 时立即跳过 transfer+complete,标记 done+duplicated", async () => {
+ const h = makeFakeHandle({ id: "a-skip", prepareDedup: true })
+ ;(prepareDirectUploadHandle as unknown as ReturnType).mockImplementation(
+ async () => h,
+ )
+
+ const { result } = renderHook(() => useAssetUpload({ effectiveLibId: "lib-1" }), {
+ wrapper: createWrapper(),
+ })
+
+ await act(async () => {
+ result.current.enqueueUploads([mp4("skip-transfer.mp4")])
+ })
+
+ await waitFor(() => {
+ expect(h.transfer).not.toHaveBeenCalled()
+ expect(h.complete).not.toHaveBeenCalled()
+ const it = result.current.uploadItems[0]
+ expect(it?.status).toBe("done")
+ expect(it?.duplicated).toBe(true)
+ expect(it?.assetId).toBe("a-skip")
+ })
+ })
})
diff --git a/apps/web/src/test/pages/auth/WechatBindCallback.test.tsx b/apps/web/src/test/pages/auth/WechatBindCallback.test.tsx
new file mode 100644
index 000000000..c9cadc752
--- /dev/null
+++ b/apps/web/src/test/pages/auth/WechatBindCallback.test.tsx
@@ -0,0 +1,144 @@
+import { describe, expect, it, vi, beforeEach, afterEach } from "vitest"
+import { render, screen, waitFor, cleanup } from "@testing-library/react"
+import { MemoryRouter } from "react-router-dom"
+import WechatBindCallback from "@/pages/auth/WechatBindCallback"
+
+const mockNavigate = vi.fn()
+const mockSetUser = vi.fn()
+const mockParams = new URLSearchParams({ code: "bind_code", state: "bind_state" })
+const mockSearchParams = [mockParams] as const
+
+const localStorageStore: Record = {}
+vi.spyOn(Storage.prototype, "getItem").mockImplementation((key) => localStorageStore[key] || null)
+vi.spyOn(Storage.prototype, "setItem").mockImplementation((key, val) => {
+ localStorageStore[key] = val
+})
+vi.spyOn(Storage.prototype, "removeItem").mockImplementation((key) => {
+ delete localStorageStore[key]
+})
+
+let bindError: unknown = null
+const mockBindResult = { user: { id: "u1", wechat_bound: true } }
+
+vi.mock("react-router-dom", async () => {
+ const actual = await vi.importActual("react-router-dom")
+ return {
+ ...actual,
+ useNavigate: () => mockNavigate,
+ useSearchParams: () => mockSearchParams,
+ }
+})
+
+vi.mock("@/api/auth", () => ({
+ bindWechat: vi.fn(async () => {
+ if (bindError) throw bindError
+ return mockBindResult
+ }),
+ normalizeUser: (u: unknown) => u,
+}))
+
+vi.mock("@/store/authStore", () => ({
+ useAuthStore: (selector: (state: unknown) => unknown) => selector({ setUser: mockSetUser }),
+}))
+
+// iframe 场景:默认非 iframe;用例可 mockReturnValue(true)
+const { mockIsInIframe, mockPostResult } = vi.hoisted(() => ({
+ mockIsInIframe: vi.fn(() => false),
+ mockPostResult: vi.fn(),
+}))
+vi.mock("@/components/auth/WechatQrModal/messages", () => ({
+ isInIframe: () => mockIsInIframe(),
+ postWechatQrResult: (...args: unknown[]) => mockPostResult(...args),
+}))
+
+const renderPage = () =>
+ render(
+
+
+ ,
+ )
+
+describe("WechatBindCallback Page", () => {
+ afterEach(() => {
+ cleanup()
+ })
+
+ beforeEach(() => {
+ vi.clearAllMocks()
+ mockIsInIframe.mockReturnValue(false)
+ bindError = null
+ Array.from(mockParams.keys()).forEach((k) => mockParams.delete(k))
+ mockParams.set("code", "bind_code")
+ mockParams.set("state", "bind_state")
+ localStorageStore.wechat_bind_state = "bind_state"
+ })
+
+ it("绑定成功跳转设置页并携带 success 标记", async () => {
+ renderPage()
+ await waitFor(() => {
+ expect(mockNavigate).toHaveBeenCalledWith("/app/profile?wechat_bind=success", {
+ replace: true,
+ })
+ })
+ expect(mockSetUser).toHaveBeenCalled()
+ })
+
+ it("本地无 wechat_bind_state(微信内/跨浏览器)不再误杀,绑定正常完成", async () => {
+ delete localStorageStore.wechat_bind_state
+ renderPage()
+ await waitFor(() => {
+ expect(mockNavigate).toHaveBeenCalledWith("/app/profile?wechat_bind=success", {
+ replace: true,
+ })
+ })
+ })
+
+ it("后端报错(微信已被其他账号绑定)时页面透传真实原因,不静默跳走", async () => {
+ bindError = {
+ isAxiosError: true,
+ response: { status: 409, data: { detail: "该微信已绑定其他账号" } },
+ message: "Request failed with status code 409",
+ }
+ renderPage()
+ await waitFor(() => {
+ expect(screen.getByText(/该微信已绑定其他账号/)).toBeTruthy()
+ })
+ expect(mockNavigate).not.toHaveBeenCalled()
+ })
+
+ it("缺少 code/state 时提示无效回调", async () => {
+ mockParams.delete("code")
+ renderPage()
+ await waitFor(() => {
+ expect(screen.getByText(/无效的回调参数/)).toBeTruthy()
+ })
+ })
+
+ describe("iframe(弹窗内嵌二维码)场景", () => {
+ it("绑定成功时 postMessage 通知父窗口,不做 navigate", async () => {
+ mockIsInIframe.mockReturnValue(true)
+ renderPage()
+ await waitFor(() => {
+ expect(mockPostResult).toHaveBeenCalledWith("bind", true)
+ })
+ expect(mockSetUser).toHaveBeenCalled()
+ expect(mockNavigate).not.toHaveBeenCalled()
+ })
+
+ it("绑定失败时把真实原因 postMessage 给父窗口", async () => {
+ mockIsInIframe.mockReturnValue(true)
+ bindError = {
+ isAxiosError: true,
+ response: { status: 409, data: { detail: "该微信已绑定其他账号" } },
+ }
+ renderPage()
+ await waitFor(() => {
+ expect(mockPostResult).toHaveBeenCalledWith("bind", false, {
+ detail: expect.stringContaining("该微信已绑定其他账号"),
+ })
+ })
+ expect(screen.queryByText(/返回设置/)).toBeNull()
+ expect(mockNavigate).not.toHaveBeenCalled()
+ })
+ })
+})
diff --git a/apps/web/src/test/pages/auth/WechatCallback.test.tsx b/apps/web/src/test/pages/auth/WechatCallback.test.tsx
index fe558f930..7aa16f3c4 100644
--- a/apps/web/src/test/pages/auth/WechatCallback.test.tsx
+++ b/apps/web/src/test/pages/auth/WechatCallback.test.tsx
@@ -1,79 +1,223 @@
-import { describe, expect, it, vi, beforeEach } from "vitest"
-import { render, screen } from "@testing-library/react"
+import { describe, expect, it, vi, beforeEach, afterEach } from "vitest"
+import { render, screen, waitFor, cleanup } from "@testing-library/react"
import { MemoryRouter } from "react-router-dom"
import WechatCallback from "@/pages/auth/WechatCallback"
+const mockNavigate = vi.fn()
+const mockSetAuth = vi.fn()
+
+// useSearchParams 返回模块级稳定引用(数组元素同一 URLSearchParams 实例),
+// 避免每次 render 返回新数组/新实例导致 useEffect 依赖变化重跑
+const mockParams = new URLSearchParams({ code: "test_code", state: "test_state" })
+const mockSearchParams = [mockParams] as const
+const mockAuthState = { setAuth: mockSetAuth }
+
+// 文件级 localStorage mock(避免每个用例重复 spy 导致链式污染)
+const localStorageStore: Record = {}
+vi.spyOn(Storage.prototype, "getItem").mockImplementation((key) => localStorageStore[key] || null)
+vi.spyOn(Storage.prototype, "setItem").mockImplementation((key, val) => {
+ localStorageStore[key] = val
+})
+vi.spyOn(Storage.prototype, "removeItem").mockImplementation((key) => {
+ delete localStorageStore[key]
+})
+
+let mockCallbackResult: Record = {}
+let mockCurrentUser: Record = {}
+let callbackError: unknown = null
+
vi.mock("react-router-dom", async () => {
const actual = await vi.importActual("react-router-dom")
return {
...actual,
- useNavigate: () => vi.fn(),
- useSearchParams: () => [new URLSearchParams({ code: "test_code", state: "test_state" })],
+ useNavigate: () => mockNavigate,
+ useSearchParams: () => mockSearchParams,
}
})
vi.mock("@/api/auth", () => ({
- wechatCallback: vi.fn(() => new Promise(() => {})), // pending promise,保持loading
- getCurrentUser: vi.fn(),
+ wechatCallback: vi.fn(async () => {
+ if (callbackError) throw callbackError
+ return mockCallbackResult
+ }),
+ getCurrentUser: vi.fn(async () => mockCurrentUser),
normalizeUser: (u: unknown) => u,
}))
+vi.mock("@/api/auth/tokenRefresh", () => ({
+ scheduleProactiveRefresh: vi.fn(),
+ cancelProactiveRefresh: vi.fn(),
+}))
+
vi.mock("@/store/authStore", () => ({
- useAuthStore: () => ({
- setAuth: vi.fn(),
- }),
+ useAuthStore: (selector: (state: unknown) => unknown) => selector({ setAuth: mockSetAuth }),
}))
-vi.mock("@/components/auth/BindContactModal", () => ({
- default: ({ open }: { open: boolean }) => (
-
- BindContactModal
-
- ),
+// iframe 场景:默认非 iframe;用例可 mockReturnValue(true)
+const { mockIsInIframe, mockPostResult } = vi.hoisted(() => ({
+ mockIsInIframe: vi.fn(() => false),
+ mockPostResult: vi.fn(),
+}))
+vi.mock("@/components/auth/WechatQrModal/messages", () => ({
+ isInIframe: () => mockIsInIframe(),
+ postWechatQrResult: (...args: unknown[]) => mockPostResult(...args),
}))
-vi.mock("antd", async () => {
- const actual = await vi.importActual("antd")
- return {
- ...actual,
- message: {
- success: vi.fn(),
- error: vi.fn(),
- },
- }
-})
+const renderPage = () =>
+ render(
+
+
+ ,
+ )
describe("WechatCallback Page", () => {
+ afterEach(() => {
+ cleanup()
+ })
+
beforeEach(() => {
- // mock localStorage,设置wechat_state匹配,让校验通过
- const store: Record = {
- wechat_state: "test_state",
+ vi.clearAllMocks()
+ mockIsInIframe.mockReturnValue(false)
+ callbackError = null
+ // 默认正常回调参数;用例可改写 mockParams 模拟 error 重定向
+ Array.from(mockParams.keys()).forEach((k) => mockParams.delete(k))
+ mockParams.set("code", "test_code")
+ mockParams.set("state", "test_state")
+ localStorageStore.wechat_state = "test_state"
+ mockCallbackResult = {
+ access_token: "at",
+ refresh_token: "rt",
+ is_new_user: false,
}
- vi.spyOn(Storage.prototype, "getItem").mockImplementation((key) => store[key] || null)
- vi.spyOn(Storage.prototype, "setItem").mockImplementation((key, val) => {
- store[key] = val
+ mockCurrentUser = {
+ id: "u1",
+ display_name: "老用户",
+ profile_completed: true,
+ }
+ })
+
+ it("老用户登录成功跳转首页/来源页", async () => {
+ renderPage()
+ await waitFor(() => {
+ expect(mockNavigate).toHaveBeenCalledWith("/", { replace: true })
})
- vi.spyOn(Storage.prototype, "removeItem").mockImplementation((key) => {
- delete store[key]
+ expect(mockSetAuth).toHaveBeenCalled()
+ })
+
+ it("新用户(is_new_user)跳转昵称引导页", async () => {
+ mockCallbackResult = { access_token: "at", refresh_token: "rt", is_new_user: true }
+ mockCurrentUser = { id: "u2", display_name: "微信用户", profile_completed: false }
+ renderPage()
+ await waitFor(() => {
+ expect(mockNavigate).toHaveBeenCalledWith("/welcome/wechat", { replace: true })
})
})
- it("should render without crashing", () => {
- const { container } = render(
-
-
- ,
- )
- expect(container).toBeTruthy()
+ it("is_new_user=false 但 profile_completed=false(上次中断)也跳引导页", async () => {
+ mockCallbackResult = { access_token: "at", refresh_token: "rt", is_new_user: false }
+ mockCurrentUser = { id: "u3", display_name: "微信用户", profile_completed: false }
+ renderPage()
+ await waitFor(() => {
+ expect(mockNavigate).toHaveBeenCalledWith("/welcome/wechat", { replace: true })
+ })
})
- it("should show loading state while processing", () => {
- render(
-
-
- ,
- )
- // wechatCallback 返回 pending promise,所以应该显示 loading
- expect(screen.getByText("正在登录...")).toBeTruthy()
+ it("本地无 wechat_state(微信内打开/跨浏览器场景)不再误杀,正常完成登录", async () => {
+ delete localStorageStore.wechat_state
+ renderPage()
+ await waitFor(() => {
+ expect(mockNavigate).toHaveBeenCalledWith("/", { replace: true })
+ })
+ // state 已被清理
+ expect(localStorageStore.wechat_state).toBeUndefined()
+ })
+
+ it("后端返回 detail 错误时,页面透传真实原因(不再吞成通用提示)", async () => {
+ callbackError = {
+ isAxiosError: true,
+ response: { status: 400, data: { detail: "微信授权码已过期,请重新扫码" } },
+ message: "Request failed with status code 400",
+ }
+ renderPage()
+ await waitFor(() => {
+ expect(screen.getByText(/微信授权码已过期,请重新扫码/)).toBeTruthy()
+ })
+ expect(screen.queryByText(/^微信登录失败,请重试$/)).toBeNull()
+ expect(mockNavigate).not.toHaveBeenCalled()
+ })
+
+ it("微信重定向带 error(用户拒绝授权)时展示授权失败原因", async () => {
+ for (const k of Array.from(mockParams.keys())) mockParams.delete(k)
+ mockParams.set("error", "access_denied")
+ mockParams.set("error_description", "The+user+denied+the+request")
+ renderPage()
+ await waitFor(() => {
+ expect(screen.getByText(/微信授权失败/)).toBeTruthy()
+ expect(screen.getByText(/access_denied/)).toBeTruthy()
+ })
+ expect(mockNavigate).not.toHaveBeenCalled()
+ })
+
+ it("缺少 code/state 参数时提示无效回调", async () => {
+ mockParams.delete("code")
+ renderPage()
+ await waitFor(() => {
+ expect(screen.getByText(/无效的回调参数/)).toBeTruthy()
+ })
+ })
+
+ it("处理中显示 loading", () => {
+ renderPage()
+ expect(screen.getByText("微信登录中...")).toBeTruthy()
+ })
+
+ describe("iframe(弹窗内嵌二维码)场景", () => {
+ it("登录成功时 postMessage 通知父窗口(needOnboarding=false),不做 navigate", async () => {
+ mockIsInIframe.mockReturnValue(true)
+ renderPage()
+ await waitFor(() => {
+ expect(mockPostResult).toHaveBeenCalledWith("login", true, { needOnboarding: false })
+ })
+ expect(mockSetAuth).toHaveBeenCalled()
+ expect(mockNavigate).not.toHaveBeenCalled()
+ })
+
+ it("新用户成功时上报 needOnboarding=true", async () => {
+ mockIsInIframe.mockReturnValue(true)
+ mockCallbackResult = { access_token: "at", refresh_token: "rt", is_new_user: true }
+ renderPage()
+ await waitFor(() => {
+ expect(mockPostResult).toHaveBeenCalledWith("login", true, { needOnboarding: true })
+ })
+ expect(mockNavigate).not.toHaveBeenCalled()
+ })
+
+ it("后端报错时把真实原因 postMessage 给父窗口,页面不渲染错误/按钮", async () => {
+ mockIsInIframe.mockReturnValue(true)
+ callbackError = {
+ isAxiosError: true,
+ response: { status: 400, data: { detail: "state 已过期或已被使用" } },
+ }
+ renderPage()
+ await waitFor(() => {
+ expect(mockPostResult).toHaveBeenCalledWith("login", false, {
+ detail: expect.stringContaining("state 已过期或已被使用"),
+ })
+ })
+ expect(screen.queryByText(/返回登录/)).toBeNull()
+ expect(mockNavigate).not.toHaveBeenCalled()
+ })
+
+ it("微信重定向 error(拒绝授权)在 iframe 内也上报父窗口", async () => {
+ mockIsInIframe.mockReturnValue(true)
+ for (const k of Array.from(mockParams.keys())) mockParams.delete(k)
+ mockParams.set("error", "access_denied")
+ renderPage()
+ await waitFor(() => {
+ expect(mockPostResult).toHaveBeenCalledWith("login", false, {
+ detail: expect.stringContaining("access_denied"),
+ })
+ })
+ })
})
})
diff --git a/apps/web/src/test/pages/auth/WechatOnboarding.test.tsx b/apps/web/src/test/pages/auth/WechatOnboarding.test.tsx
new file mode 100644
index 000000000..548d7f52e
--- /dev/null
+++ b/apps/web/src/test/pages/auth/WechatOnboarding.test.tsx
@@ -0,0 +1,186 @@
+import { describe, expect, it, vi, beforeEach, afterEach } from "vitest"
+import { render, screen, fireEvent, waitFor, cleanup } from "@testing-library/react"
+import { MemoryRouter } from "react-router-dom"
+import { QueryClient, QueryClientProvider } from "@tanstack/react-query"
+import WechatOnboarding from "@/pages/auth/WechatOnboarding"
+
+const queryClient = new QueryClient({
+ defaultOptions: { queries: { retry: false }, mutations: { retry: false } },
+})
+
+const mockNavigate = vi.fn()
+const mockSetUser = vi.fn()
+let updateProfileMock = vi.fn()
+
+vi.mock("react-router-dom", async () => {
+ const actual = await vi.importActual("react-router-dom")
+ return { ...actual, useNavigate: () => mockNavigate }
+})
+
+let authState: Record = {}
+vi.mock("@/store/authStore", () => ({
+ useAuthStore: (selector: (state: unknown) => unknown) => selector(authState),
+}))
+
+vi.mock("@/api/auth", () => ({
+ updateProfile: (data: { display_name: string }) => updateProfileMock(data),
+}))
+
+vi.mock("antd", async () => {
+ const actual = await vi.importActual("antd")
+ return { ...actual, message: { success: vi.fn(), error: vi.fn() } }
+})
+
+const renderPage = () =>
+ render(
+
+
+
+
+ ,
+ )
+
+describe("WechatOnboarding 昵称引导页", () => {
+ afterEach(() => {
+ cleanup()
+ })
+
+ beforeEach(() => {
+ vi.clearAllMocks()
+ authState = {
+ isAuthenticated: true,
+ user: { id: "u1", display_name: "", profile_completed: false },
+ setUser: mockSetUser,
+ }
+ localStorage.setItem("access_token", "at")
+ updateProfileMock = vi.fn(async (data: { display_name: string }) => ({
+ id: "u1",
+ display_name: data.display_name,
+ profile_completed: true,
+ }))
+ })
+
+ it("未登录时跳转登录页", () => {
+ authState = {
+ isAuthenticated: false,
+ user: null,
+ setUser: mockSetUser,
+ }
+ localStorage.removeItem("access_token")
+ renderPage()
+ expect(mockNavigate).not.toHaveBeenCalled()
+ // Navigate 组件渲染即生效;这里断言页面不含昵称表单
+ expect(screen.queryByText("进入小虾智剪")).toBeNull()
+ })
+
+ it("资料已完善的用户跳 dashboard", () => {
+ authState = {
+ isAuthenticated: true,
+ user: { id: "u1", display_name: "已起名", profile_completed: true },
+ setUser: mockSetUser,
+ }
+ renderPage()
+ expect(screen.queryByText("进入小虾智剪")).toBeNull()
+ })
+
+ it("昵称输入框不预填,必须用户自己输入", () => {
+ renderPage()
+ expect(screen.getByText("欢迎使用微信登录,请先设置您的昵称")).toBeTruthy()
+ expect((screen.getByPlaceholderText("请输入您的昵称") as HTMLInputElement).value).toBe("")
+ })
+
+ it("新用户可见昵称表单并能提交", async () => {
+ renderPage()
+ expect(screen.getByText("欢迎使用微信登录,请先设置您的昵称")).toBeTruthy()
+
+ fireEvent.change(screen.getByPlaceholderText("请输入您的昵称"), {
+ target: { value: "小虾用户" },
+ })
+ fireEvent.click(screen.getByText("进入小虾智剪"))
+
+ await waitFor(() => {
+ expect(updateProfileMock).toHaveBeenCalledWith({ display_name: "小虾用户" })
+ })
+ await waitFor(() => {
+ expect(mockSetUser).toHaveBeenCalled()
+ expect(mockNavigate).toHaveBeenCalledWith("/app/dashboard", { replace: true })
+ })
+ })
+
+ it("昵称为空时不允许提交(表单校验)", async () => {
+ renderPage()
+ fireEvent.change(screen.getByPlaceholderText("请输入您的昵称"), {
+ target: { value: " " },
+ })
+ fireEvent.click(screen.getByText("进入小虾智剪"))
+ // 等待表单校验
+ await waitFor(
+ () => {
+ expect(updateProfileMock).not.toHaveBeenCalled()
+ },
+ { timeout: 1000 },
+ )
+ })
+
+ it("连点提交按钮只触发一次请求(防重复提交)", async () => {
+ // mutation 挂起不立即完成,模拟慢网络下连续双击
+ let resolveSubmit: (v: unknown) => void = () => {}
+ updateProfileMock = vi.fn(
+ () =>
+ new Promise((resolve) => {
+ resolveSubmit = resolve
+ }),
+ )
+ renderPage()
+ fireEvent.change(screen.getByPlaceholderText("请输入您的昵称"), {
+ target: { value: "小虾用户" },
+ })
+ const btn = screen.getByText("进入小虾智剪")
+ fireEvent.click(btn)
+ // 第一次点击后立即再点(此时重渲染/loading 可能还没生效)
+ fireEvent.click(btn)
+ fireEvent.click(btn)
+ await waitFor(() => {
+ expect(updateProfileMock).toHaveBeenCalledTimes(1)
+ })
+ // 释放挂起的 Promise,避免泄漏
+ resolveSubmit({ id: "u1", display_name: "小虾用户", profile_completed: true })
+ })
+
+ it("提交失败后守卫复位,允许再次提交", async () => {
+ updateProfileMock = vi
+ .fn()
+ .mockRejectedValueOnce({
+ isAxiosError: true,
+ response: { status: 500, data: { detail: "服务内部错误" } },
+ })
+ .mockResolvedValueOnce({ id: "u1", display_name: "小虾用户", profile_completed: true })
+ renderPage()
+ fireEvent.change(screen.getByPlaceholderText("请输入您的昵称"), {
+ target: { value: "小虾用户" },
+ })
+ fireEvent.click(screen.getByText("进入小虾智剪"))
+ await waitFor(() => {
+ expect(updateProfileMock).toHaveBeenCalledTimes(1)
+ })
+ // 失败后再点一次,应能重新提交
+ fireEvent.click(screen.getByText("进入小虾智剪"))
+ await waitFor(() => {
+ expect(updateProfileMock).toHaveBeenCalledTimes(2)
+ })
+ })
+
+ it("提交失败显示错误且不跳转", async () => {
+ updateProfileMock = vi.fn(async () => {
+ throw new Error("500")
+ })
+ renderPage()
+ fireEvent.change(screen.getByPlaceholderText("请输入您的昵称"), {
+ target: { value: "小虾用户" },
+ })
+ fireEvent.click(screen.getByText("进入小虾智剪"))
+ await waitFor(() => {
+ expect(mockNavigate).not.toHaveBeenCalled()
+ })
+ })
+})
diff --git a/apps/web/src/test/pages/generate/smoke.test.tsx b/apps/web/src/test/pages/generate/smoke.test.tsx
index 475514f67..9e58e0c6b 100755
--- a/apps/web/src/test/pages/generate/smoke.test.tsx
+++ b/apps/web/src/test/pages/generate/smoke.test.tsx
@@ -19,6 +19,10 @@ import "@/pages/generate/GeneratePage"
import "@/pages/generate/components/Step2MaterialSelect"
import "@/pages/generate/components/Step4TitleSettings"
import "@/pages/generate/components/Step5VoiceSelect"
+import "@/pages/generate/components/Step3VoiceWithMode"
+import "@/pages/generate/components/CanvasPreviewGrid"
+import "@/pages/generate/components/BatchGenerationGrid"
+import "@/pages/generate/components/PreviewCountModal"
import "@/pages/generate/components/PreviewVideoPanel"
import "@/pages/generate/components/GenerateResultPanel"
import "@/pages/generate/components/GenerateStepContent"
@@ -45,6 +49,7 @@ describe("GeneratePage module smoke test", () => {
})
})
import "@/pages/generate/hooks/useGenerateVideo"
+import "@/pages/generate/hooks/useBatchCovers"
import "@/pages/generate/hooks/usePreviewAssets"
import "@/pages/generate/hooks/useSegmentScheduler"
import "@/pages/generate/hooks/generate-video/useGenerationPolling"
diff --git a/apps/web/src/test/pages/generate/title-library-autocomplete.test.tsx b/apps/web/src/test/pages/generate/title-library-autocomplete.test.tsx
new file mode 100644
index 000000000..b8a05188d
--- /dev/null
+++ b/apps/web/src/test/pages/generate/title-library-autocomplete.test.tsx
@@ -0,0 +1,126 @@
+/**
+ * TitleLibraryAutoComplete 单测(Issue #1737)
+ *
+ * 覆盖:
+ * - 聚焦空输入框 → 下拉立即展开,展示标题库全部标题(原生 AutoComplete 聚焦不展开,此为本工单核心修复)
+ * - 输入关键词 → 下拉只显示匹配项
+ * - 点击下拉项 → onChange 回填所选标题
+ * - 自由输入自定义标题 → onChange 正常透传,不被下拉干扰
+ * - 标题库为空 → 聚焦不展开(不出"暂无数据"空壳)
+ * - 选中后下拉关闭
+ */
+import { describe, it, expect, vi } from "vitest"
+import { render, screen, waitFor, fireEvent } from "@testing-library/react"
+import userEvent from "@testing-library/user-event"
+import TitleLibraryAutoComplete from "@/pages/generate/components/title/TitleLibraryAutoComplete"
+
+const OPTIONS = [
+ { label: "永康这家面馆绝了", value: "永康这家面馆绝了" },
+ { label: "永康美食探店vlog", value: "永康美食探店vlog" },
+ { label: "萌宠日常第一天", value: "萌宠日常第一天" },
+]
+
+function renderBox(initialValue = "", opts = OPTIONS) {
+ const onChange = vi.fn()
+ const result = render(
+ ,
+ )
+ return { onChange, ...result }
+}
+
+/** 聚焦输入框(combobox role) */
+function focusInput() {
+ const input = screen.getByRole("combobox") as HTMLInputElement
+ fireEvent.focus(input)
+ return input
+}
+
+/** 取下拉中实际可见的选项(rc-virtual-list 渲染为 .ant-select-item-option;role=option 的 listbox 是 a11y 哨兵) */
+function getVisibleOptions(): HTMLElement[] {
+ const dropdown = document.querySelector(".ant-select-dropdown:not(.ant-select-dropdown-hidden)")
+ if (!dropdown) return []
+ return Array.from(dropdown.querySelectorAll(".ant-select-item-option")) as HTMLElement[]
+}
+
+describe("TitleLibraryAutoComplete (#1737)", () => {
+ it("聚焦空输入框时下拉展开并展示标题库全部标题", async () => {
+ renderBox()
+ expect(screen.queryByRole("listbox")).not.toBeInTheDocument()
+
+ focusInput()
+
+ await screen.findByRole("listbox")
+ await waitFor(() => expect(getVisibleOptions()).toHaveLength(3))
+ const options = getVisibleOptions()
+ expect(options[0]).toHaveTextContent("永康这家面馆绝了")
+ expect(options[2]).toHaveTextContent("萌宠日常第一天")
+ })
+
+ it("输入关键词时下拉只显示匹配项", async () => {
+ const user = userEvent.setup()
+ renderBox()
+ const input = screen.getByRole("combobox")
+ await user.click(input)
+ await screen.findByRole("listbox")
+
+ await user.type(input, "永康")
+ await waitFor(() => expect(getVisibleOptions()).toHaveLength(2))
+ const options = getVisibleOptions()
+ expect(options.every((o) => o.textContent?.includes("永康"))).toBe(true)
+ })
+
+ it("点击下拉项后 onChange 回填标题且下拉关闭", async () => {
+ const user = userEvent.setup()
+ const { onChange } = renderBox()
+ const input = screen.getByRole("combobox") as HTMLInputElement
+ await user.click(input)
+ await screen.findByRole("listbox")
+
+ await user.click(screen.getByText("萌宠日常第一天"))
+
+ await waitFor(() => {
+ expect(onChange).toHaveBeenCalledWith("萌宠日常第一天")
+ })
+ await waitFor(() => {
+ expect(screen.queryByRole("listbox")).not.toBeInTheDocument()
+ })
+ })
+
+ it("自由输入自定义标题时 onChange 正常透传(不被下拉干扰)", async () => {
+ const user = userEvent.setup()
+ const { onChange } = renderBox()
+ const input = screen.getByRole("combobox")
+ await user.click(input)
+
+ await user.type(input, "我自己编的标题XYZ")
+ await waitFor(() => {
+ expect(onChange).toHaveBeenCalledWith("我自己编的标题XYZ")
+ })
+ // 输入无匹配关键词,下拉无 option 时不阻塞输入
+ expect(input).toHaveValue("我自己编的标题XYZ")
+ })
+
+ it("标题库为空时聚焦不展开下拉", async () => {
+ renderBox("", [])
+ focusInput()
+ // 等一帧确认没有 listbox
+ await new Promise((r) => setTimeout(r, 50))
+ expect(screen.queryByRole("listbox")).not.toBeInTheDocument()
+ })
+
+ it("渲染下拉箭头图标作为可选择提示", () => {
+ const { container } = renderBox()
+ // antd 后缀图标在 .ant-select-arrow 内
+ expect(container.querySelector(".ant-select-arrow")).toBeInTheDocument()
+ })
+
+ it("有初始值时输入框正常展示", () => {
+ renderBox("已有标题")
+ expect(screen.getByRole("combobox")).toHaveValue("已有标题")
+ })
+})
diff --git a/apps/web/src/test/router/lazyRoute.test.ts b/apps/web/src/test/router/lazyRoute.test.ts
new file mode 100644
index 000000000..8e64fac7a
--- /dev/null
+++ b/apps/web/src/test/router/lazyRoute.test.ts
@@ -0,0 +1,44 @@
+import { describe, it, expect, vi, afterEach } from "vitest"
+import { lazyRoute } from "@/router/lazyRoute"
+
+const chunkErr = () => new TypeError("Failed to fetch dynamically imported module: /assets/x.js")
+
+/** fake 模块 */
+const Comp = function Comp() {}
+const factoryOk = vi.fn(async () => ({ default: Comp }))
+
+afterEach(() => {
+ vi.clearAllMocks()
+})
+
+describe("lazyRoute", () => {
+ it("首次成功直接返回 Component", async () => {
+ const result = await lazyRoute(factoryOk)()
+ expect(result).toEqual({ Component: Comp })
+ expect(factoryOk).toHaveBeenCalledTimes(1)
+ })
+
+ it("chunk 失败重试:前两次失败、第三次成功 → 不抛出", async () => {
+ const f = vi
+ .fn()
+ .mockRejectedValueOnce(chunkErr())
+ .mockRejectedValueOnce(chunkErr())
+ .mockResolvedValueOnce({ default: Comp })
+
+ const result = await lazyRoute(f as never)()
+ expect(result).toEqual({ Component: Comp })
+ expect(f).toHaveBeenCalledTimes(3)
+ })
+
+ it("chunk 失败重试 2 次仍失败 → 抛出", async () => {
+ const f = vi.fn().mockRejectedValue(chunkErr())
+ await expect(lazyRoute(f as never)()).rejects.toThrow(/dynamically imported/)
+ expect(f).toHaveBeenCalledTimes(3)
+ })
+
+ it("非 chunk 错误立即抛出,不重试", async () => {
+ const f = vi.fn().mockRejectedValue(new Error("业务模块内部报错"))
+ await expect(lazyRoute(f as never)()).rejects.toThrow("业务模块内部报错")
+ expect(f).toHaveBeenCalledTimes(1)
+ })
+})
diff --git a/apps/web/src/test/utils/chunkLoadError.test.ts b/apps/web/src/test/utils/chunkLoadError.test.ts
new file mode 100644
index 000000000..b7a5b9997
--- /dev/null
+++ b/apps/web/src/test/utils/chunkLoadError.test.ts
@@ -0,0 +1,80 @@
+import { describe, it, expect, beforeEach, afterEach, vi } from "vitest"
+import {
+ getChunkReloadedAt,
+ goHomeRecover,
+ isChunkLoadError,
+ reloadForChunkError,
+} from "@/utils/chunkLoadError"
+
+describe("isChunkLoadError", () => {
+ it("识别 Vite 动态 import 失败", () => {
+ const err = new TypeError(
+ "Failed to fetch dynamically imported module: https://x/assets/AssetLibrary-abc.js",
+ )
+ expect(isChunkLoadError(err)).toBe(true)
+ })
+
+ it("识别 Webpack 风格 ChunkLoadError", () => {
+ const err = new Error("Loading chunk 12 failed.")
+ err.name = "ChunkLoadError"
+ expect(isChunkLoadError(err)).toBe(true)
+ })
+
+ it("识别字符串形式错误", () => {
+ expect(isChunkLoadError("Error loading dynamically imported module")).toBe(true)
+ })
+
+ it("普通错误不命中", () => {
+ expect(isChunkLoadError(new Error("Cannot read properties of undefined"))).toBe(false)
+ expect(isChunkLoadError(null)).toBe(false)
+ expect(isChunkLoadError(undefined)).toBe(false)
+ expect(isChunkLoadError({ status: 500 })).toBe(false)
+ })
+})
+
+describe("reload 标记", () => {
+ beforeEach(() => {
+ sessionStorage.clear()
+ // jsdom 未实现真实导航,reload 仅打 "not implemented" 警告,静默掉
+ vi.spyOn(console, "error").mockImplementation(() => {})
+ })
+ afterEach(() => {
+ vi.restoreAllMocks()
+ sessionStorage.clear()
+ })
+
+ it("无标记返回 null", () => {
+ expect(getChunkReloadedAt()).toBeNull()
+ })
+
+ it("reloadForChunkError 写入刷新标记", () => {
+ expect(() => reloadForChunkError()).not.toThrow()
+ expect(getChunkReloadedAt()).not.toBeNull()
+ })
+
+ it("标记过期(>10min)返回 null", () => {
+ sessionStorage.setItem("chunk_error_reloaded_at", String(Date.now() - 11 * 60 * 1000))
+ expect(getChunkReloadedAt()).toBeNull()
+ })
+
+ it("goHomeRecover 清掉标记", () => {
+ reloadForChunkError()
+ expect(getChunkReloadedAt()).not.toBeNull()
+ expect(() => goHomeRecover()).not.toThrow()
+ expect(sessionStorage.getItem("chunk_error_reloaded_at")).toBeNull()
+ })
+
+ it("sessionStorage 抛异常(无痕模式)时降级不崩溃", () => {
+ const spy = vi.spyOn(Storage.prototype, "getItem").mockImplementation(() => {
+ throw new Error("Storage disabled")
+ })
+ const setSpy = vi.spyOn(Storage.prototype, "setItem").mockImplementation(() => {
+ throw new Error("Storage disabled")
+ })
+ expect(getChunkReloadedAt()).toBeNull()
+ expect(() => reloadForChunkError()).not.toThrow()
+ expect(() => goHomeRecover()).not.toThrow()
+ spy.mockRestore()
+ setSpy.mockRestore()
+ })
+})
diff --git a/apps/web/src/utils/chunkLoadError.ts b/apps/web/src/utils/chunkLoadError.ts
new file mode 100644
index 000000000..7dce61a66
--- /dev/null
+++ b/apps/web/src/utils/chunkLoadError.ts
@@ -0,0 +1,84 @@
+/**
+ * 发版后旧标签页懒加载 chunk 失效(白屏)的识别与恢复工具。
+ *
+ * 背景:页面 React Router 的 lazy 动态 import,发版后旧 chunk 文件名被删除,
+ * 停留在旧标签页的用户点菜单时 import 404,抛出
+ * "Failed to fetch dynamically imported module"(Vite)/ ChunkLoadError,
+ * 不捕获就是整页白屏。
+ */
+
+/** sessionStorage 标记:最近已经为 chunk 失效自动刷新过一次(带时间戳,10min 有效) */
+const RELOAD_FLAG_KEY = "chunk_error_reloaded_at"
+/** 标记有效期:超过后允许再次自动刷新,避免用户手动正常刷新后标记永久残留 */
+const RELOAD_FLAG_TTL_MS = 10 * 60 * 1000
+
+/**
+ * Storage 在 Safari 无痕模式 / 禁用 Cookie 的浏览器 / 严格 iframe 策略下
+ * 访问可能抛异常;此处统一容错,拿不到存储就降级为"无标记",绝不能让
+ * 错误边界本身因读存储而崩溃。
+ */
+const safeStorage = {
+ getItem: (key: string): string | null => {
+ try {
+ return sessionStorage.getItem(key)
+ } catch {
+ return null
+ }
+ },
+ setItem: (key: string, value: string): void => {
+ try {
+ sessionStorage.setItem(key, value)
+ } catch {
+ /* 存储不可用时静默降级:仅丢失"已刷新"标记,不影响恢复动作 */
+ }
+ },
+ removeItem: (key: string): void => {
+ try {
+ sessionStorage.removeItem(key)
+ } catch {
+ /* ignore */
+ }
+ },
+}
+
+/** 判断错误是否为懒加载 chunk 加载失败(发版 404 / 网络中断 / 动态 import 失败) */
+export const isChunkLoadError = (error: unknown): boolean => {
+ if (!error) return false
+ // Vite: Failed to fetch dynamically imported module: /assets/xxx-yyy.js
+ // Webpack: ChunkLoadError: Loading chunk xxx failed.
+ const needle =
+ error instanceof Error
+ ? `${error.name} ${error.message}`
+ : typeof error === "string"
+ ? error
+ : ""
+ return /failed to fetch dynamically imported module|chunkloaderror|loading chunk \d+ failed|error loading dynamically imported module|importing a module script failed/i.test(
+ needle,
+ )
+}
+
+/** 读取上次自动刷新时间戳;过期或不存在返回 null */
+export const getChunkReloadedAt = (): number | null => {
+ const raw = safeStorage.getItem(RELOAD_FLAG_KEY)
+ if (!raw) return null
+ const ts = Number(raw)
+ if (!Number.isFinite(ts)) return null
+ if (Date.now() - ts > RELOAD_FLAG_TTL_MS) return null
+ return ts
+}
+
+/** 标记"已为 chunk 失效自动刷新过",然后刷新页面 */
+export const reloadForChunkError = (): void => {
+ safeStorage.setItem(RELOAD_FLAG_KEY, String(Date.now()))
+ window.location.reload()
+}
+
+/**
+ * 硬恢复:清掉标记后回到首页(整页导航,不是当前 URL 刷新)。
+ * - chunk 失效兜底:回到首页会拉取最新 index.html,彻底脱离旧 chunk 引用
+ * - 非 chunk 的页面级崩溃:跳首页能绕开当前报错路由,避免"刷新-再崩"死循环
+ */
+export const goHomeRecover = (): void => {
+ safeStorage.removeItem(RELOAD_FLAG_KEY)
+ window.location.href = "/"
+}
diff --git a/apps/worker/video_processing/dedup.py b/apps/worker/video_processing/dedup.py
index b4c6aaba3..7e6ebdf06 100755
--- a/apps/worker/video_processing/dedup.py
+++ b/apps/worker/video_processing/dedup.py
@@ -31,25 +31,65 @@ SCENE_CHANGE_THRESHOLD = 30 # 灰度差异阈值
MIN_KEYFRAME_INTERVAL_SEC = 1.0 # 最小关键帧间隔(秒)
MAX_KEYFRAMES = 30 # 最大关键帧数
MIN_KEYFRAMES = 5 # 最小关键帧数
+FINGERPRINT_SAMPLE_INTERVAL_SEC = 1.0 # 指纹采样间隔(秒):密集均匀采样,保证两视频时序可对齐
+FINGERPRINT_MAX_SAMPLES = 30 # 长视频采样数上限(超过后采样间隔自动放宽)
LONG_VIDEO_SEGMENT_SEC = 30 # 长视频每段秒数
LONG_VIDEO_DURATION_THRESHOLD_SEC = 180 # 3 分钟阈值
MIN_FRAMES_PER_SEGMENT = 2 # 长视频每段最少帧数
-# ── 滑动窗口匹配常量 ────────────────────────────────────────────
-SEGMENT_MATCH_THRESHOLD = 8 # 帧匹配汉明距离阈值
-MIN_CONSECUTIVE_MATCHES = 5 # 最少连续匹配帧数
+# ── 滑动窗口匹配常量(Issue #1702 二次校准) ─────────────────────
+# 阈值经 staging 真实数据两轮回归校准(worker 容器内离线实验):
+# 第一轮(2026-09-05):同源对 <=12 命中 4/11,异源最小距离 24 → 定 12;
+# 第二轮(2026-09-05,证据视频 B->A 仍漏检):扩大样本到该用户全部
+# 15 个真实成片(13 个异源候选)实测:
+# - 同源成片对(A 20s / B、C 各 11.75s,1s 密集采样):
+# B->A 中位数距离 14,<=16 命中 8/11=0.73;C->A 8/11=0.73
+# - 异源成片对(13 个真实视频):每帧全局最近邻最小距离 18,
+# <=16 命中帧数全部为 0(最近邻 18 仅个别帧,中位数 22~28)
+# 12 漏掉同源降重对(降重滤镜/字幕/画面扰动把距离从 ~8 推到 14~16);
+# 16 对同源命中 0.73+ 且与异源分布(最近邻 >=18)仍有 >=2bit 安全裕度,
+# 异源 <=16 命中 0 帧,无误报空间。
+PHASH_THRESHOLD = 16
+SEGMENT_MATCH_THRESHOLD = PHASH_THRESHOLD # 片段匹配阈值与帧匹配统一(#1702:阈值常量统一来源)
+MIN_CONSECUTIVE_MATCHES = 5 # 连续匹配默认门槛;短视频自适应 min(5, max(2, 分片数//2))
MAX_GAP = 2 # 允许的最大间隙帧数
+NEIGHBOR_WINDOW = 1 # 分片时序对齐:允许 ±1 邻接偏移(1s 密集采样下即 ±1s,缓解切点不一致)
# ── 融合判定常量 ────────────────────────────────────────────────
PHASH_WEIGHT = 0.7 # pHash 权重
HISTOGRAM_WEIGHT = 0.3 # 直方图权重
-MATCH_RATIO_THRESHOLD = 0.7 # 至少 70% 帧匹配
+MATCH_RATIO_THRESHOLD = 0.7 # 全片重复(is_duplicate)至少 70% 帧匹配
+PARTIAL_COVERAGE_THRESHOLD = 0.5 # 局部复用覆盖率 >=50% 也判全片重复
DUPLICATE_THRESHOLD = 0.70 # 融合后相似度阈值
+# ── 降重裁剪规避常量(Issue #1702) ─────────────────────────────
+# 成片强制 2-5% random_edge_crop 降重只服务外部平台;自查重指纹取中心 90%
+# 区域,使两次不同裁剪的同源画面 pHash 距离回到同分布。
+FINGERPRINT_CENTER_CROP_RATIO = 0.90
+
# ── 感知哈希 & 颜色直方图工具函数 ────────────────────────────────
+def center_crop_frame(image: np.ndarray, ratio: float = FINGERPRINT_CENTER_CROP_RATIO) -> np.ndarray:
+ """取画面中心 ratio 比例区域(裁除四边边缘)。
+
+ 查重指纹用:random_edge_crop 降重(2-5% 四边随机裁剪)会让同源画面 pHash
+ 位翻转 12-16,污染自查重(Issue #1702)。算 pHash/颜色直方图前先居中裁除
+ 边缘 10%,两次不同裁剪的同源画面中心区域基本重合,指纹不再被降重污染。
+ 降重只服务外部平台,不影响内部查重。
+ """
+ if image is None or image.size == 0:
+ return image
+ h, w = image.shape[:2]
+ ch, cw = int(h * ratio), int(w * ratio)
+ if ch <= 0 or cw <= 0 or (ch >= h and cw >= w):
+ return image
+ y0 = (h - ch) // 2
+ x0 = (w - cw) // 2
+ return image[y0 : y0 + ch, x0 : x0 + cw]
+
+
def compute_phash(image: np.ndarray, hash_size: int = 8) -> str:
"""计算图像的感知哈希(pHash),基于 DCT(离散余弦变换)。
@@ -101,11 +141,17 @@ def hamming_distance(hash1: str, hash2: str) -> int:
def compute_color_histogram(image: np.ndarray, bins: int = 32) -> list[float]:
- """Compute color histogram for an image."""
+ """Compute BGR color histogram for an image.
+
+ Issue #1702: 每个通道独立做 NORM_L1 归一化(通道内 Σ=1,是概率分布),
+ 三通道拼接存储。Bhattacharyya 系数对拼接向量直接 Σ√(a*b) 会得到
+ 3 通道之和(范围 [0,3],实测 ~14.9 是旧 L2 归一化的错误结果),
+ 消费方 _bhattacharyya_coefficient 按通道数平均归一到 [0,1]。
+ """
hist = []
for i in range(3):
h = cv2.calcHist([image], [i], None, [bins], [0, 256])
- h = cv2.normalize(h, h).flatten()
+ h = cv2.normalize(h, h, norm_type=cv2.NORM_L1).flatten()
hist.extend(h)
return hist
@@ -210,6 +256,30 @@ def detect_keyframe_timestamps(
return keyframe_times
+def sample_fingerprint_timestamps(
+ duration: float,
+ *,
+ interval_sec: float = FINGERPRINT_SAMPLE_INTERVAL_SEC,
+ max_samples: int = FINGERPRINT_MAX_SAMPLES,
+) -> list[float]:
+ """指纹采样时间戳:固定间隔密集均匀采样(Issue #1702)。
+
+ 动态场景检测抽帧(#1659)在两个同源视频上会各自取到不同时刻,切点/取帧
+ 错位让对齐帧的 pHash 距离都很大(实测同源对最小距离 12 且配对时序错乱)。
+ 改为固定 1s 间隔均匀采样后,复用片段的帧时刻天然对齐,配合 ±1 邻接窗口
+ 即可检出同源/局部复用。长视频(>max_samples*interval)自动放宽间隔到
+ duration/max_samples,保证分片数有上限。
+ """
+ if duration <= 0:
+ return []
+ step = interval_sec
+ n_uniform = int(duration / step)
+ if n_uniform > max_samples:
+ step = duration / max_samples
+ count = max(1, int(duration / step))
+ return [step * (i + 0.5) for i in range(count)]
+
+
# ── 数据类 ──────────────────────────────────────────────────────
@@ -297,23 +367,33 @@ def find_duplicate_segments(
target_chunks: list,
*,
match_threshold: int = SEGMENT_MATCH_THRESHOLD,
- min_consecutive: int = MIN_CONSECUTIVE_MATCHES,
+ min_consecutive: Optional[int] = None,
max_gap: int = MAX_GAP,
+ neighbor_window: int = NEIGHBOR_WINDOW,
) -> list[DuplicateSegment]:
- """滑动窗口时序匹配:找出两组分片之间的重复片段。
+ """滑动窗口时序匹配:找出两组分片之间的重复片段(Issue #1702 重构)。
算法:
- 1. 对每个 query chunk,找到 target 中汉明距离最小的 chunk
- 2. 距离 <= match_threshold 视为匹配
- 3. 找连续匹配的 run(允许 max_gap 帧间隙)
- 4. 连续匹配数 >= min_consecutive 的 run 报告为重复片段
+ 1. 构建 query×target 全量汉明距离矩阵;每个 query chunk 保留所有
+ 距离 <= match_threshold 的候选 target 分片(与帧匹配判定同一阈值)。
+ 2. 时序一致贪心对齐:沿 query 时序推进,run 内优先选择与上一匹配帧
+ 目标序号连贯(|delta| <= neighbor_window+1,允许 ±1 邻接/时序偏移
+ 对齐——1s 密集采样下相邻帧 pHash 接近,最近邻在目标相邻帧间
+ 正/反向跳变均属正常,缓解场景切割切点、取帧错位、局部倒退)的
+ 候选;同距时偏好小索引(最早对齐位置)。
+ 3. 连贯匹配中允许 <= max_gap 帧间隙桥接;断裂后另起新 run——天然
+ 支持局部片段复用(复用片段可出现在任意时序位置,各成独立片段)。
+ 4. 连续匹配帧数 >= min_consecutive 的 run 报为重复片段。短视频自适应:
+ min_consecutive = min(5, max(2, len(query_chunks)//2));n=1 时
+ 不形成片段,由调用方匹配帧回退兜底。
Args:
query_chunks: 查询视频的分片列表(FingerprintChunk 或 dict)
target_chunks: 目标视频的分片列表
- match_threshold: 汉明距离匹配阈值
- min_consecutive: 最少连续匹配帧数
+ match_threshold: 汉明距离匹配阈值(统一常量 PHASH_THRESHOLD)
+ min_consecutive: 最少连续匹配帧数;None 时按短视频自适应
max_gap: 允许的最大间隙帧数
+ neighbor_window: 时序对齐允许的目标分片序号邻接窗口(正/反向均允许)
Returns:
DuplicateSegment 列表
@@ -321,95 +401,94 @@ def find_duplicate_segments(
if not query_chunks or not target_chunks:
return []
- def _get_phash(chunk) -> str:
+ def _get(chunk, key):
if isinstance(chunk, dict):
- return chunk["phash_binary"]
- return chunk.phash_binary
+ return chunk[key]
+ return getattr(chunk, key)
- def _get_start(chunk) -> int:
- if isinstance(chunk, dict):
- return chunk["start_time_ms"]
- return chunk.start_time_ms
+ n, m = len(query_chunks), len(target_chunks)
+ q_ph = [_get(c, "phash_binary") for c in query_chunks]
+ t_ph = [_get(c, "phash_binary") for c in target_chunks]
- def _get_end(chunk) -> int:
- if isinstance(chunk, dict):
- return chunk["end_time_ms"]
- return chunk.end_time_ms
+ # Step 1: 全量距离矩阵。每个 query chunk 保留所有 <= 阈值的候选 target,
+ # 按距离升序;同距时小索引优先(取最早的对齐位置,贪心连贯推进时最保守,
+ # 不会越过复用片段末端;重复 hash 的连续帧由 Step 2 的连贯性窗口约束)。
+ candidates: list[list[tuple[int, int]]] = [] # 每 query 帧: [(target_idx, dist), ...]
+ for i in range(n):
+ dists = [hamming_distance(q_ph[i], t_ph[j]) for j in range(m)]
+ cand = [(j, d) for j, d in enumerate(dists) if d <= match_threshold]
+ cand.sort(key=lambda x: (x[1], x[0]))
+ candidates.append(cand)
- # Step 1: 逐帧匹配
- frame_matches: list[tuple[bool, int, int]] = [] # (is_match, min_dist, best_target_idx)
- for qc in query_chunks:
- qc_phash = _get_phash(qc)
- best_dist = 64
- best_idx = 0
- for j, tc in enumerate(target_chunks):
- d = hamming_distance(qc_phash, _get_phash(tc))
- if d < best_dist:
- best_dist = d
- best_idx = j
- frame_matches.append((best_dist <= match_threshold, best_dist, best_idx))
+ # 短视频自适应连续匹配门槛(Issue #1702 工单公式):
+ # MIN_CONSECUTIVE_MATCHES = min(5, max(2, 分片数//2))。
+ # n=1 时门槛为 2 不形成片段,由 _evaluate_candidate 的匹配帧回退
+ # (temporal_coverage 按匹配帧占比估计)兜底检出,不回归。
+ if min_consecutive is None:
+ min_consecutive = min(MIN_CONSECUTIVE_MATCHES, max(2, n // 2))
- # Step 2: 找连续匹配的 runs
- runs: list[tuple[int, int]] = [] # list of (start_idx, end_idx)
- run_start = None
+ # Step 2: 时序一致贪心对齐。
+ # run 内偏好与上一匹配帧目标序号连贯(|delta| <= neighbor_window+1,
+ # 支持 ±1 邻接窗口/时序偏移对齐,正反向抖动均允许)的候选;
+ # 无连贯候选时关闭旧 run。
+ # 这天然支持局部片段复用:同一 query 视频中多个复用片段各自形成独立 run。
+ frame_matches: list[tuple[bool, int, int]] = []
+ runs: list[tuple[int, int]] = []
+ run_start: Optional[int] = None
+ run_last_t: Optional[int] = None
gap_count = 0
- for i, (is_match, _dist, _idx) in enumerate(frame_matches):
- if is_match:
+ def _matching_count(a: int, b: int) -> int:
+ return sum(1 for k in range(a, b + 1) if frame_matches[k][0])
+
+ def _close_run(a: int, b: int) -> None:
+ if b >= a and _matching_count(a, b) >= min_consecutive:
+ runs.append((a, b))
+
+ for i in range(n):
+ cand = candidates[i]
+ if run_last_t is None:
+ chosen = cand[0] if cand else None
+ else:
+ chosen = next(
+ (c for c in cand if abs(c[0] - run_last_t) <= neighbor_window + 1),
+ None,
+ )
+
+ if chosen is not None:
+ tidx, dist = chosen
+ frame_matches.append((True, dist, tidx))
if run_start is None:
run_start = i
- gap_count = 0 # 重置间隙
+ gap_count = 0
+ run_last_t = tidx
else:
+ frame_matches.append((False, match_threshold + 1, -1))
if run_start is not None:
gap_count += 1
if gap_count > max_gap:
- # 中断当前 run
- run_end = i - gap_count # 最后一个匹配帧的索引
- # 计算 run 内的实际匹配帧数(总跨度 - 间隙数)
- total_gaps = sum(1 for k in range(run_start, run_end + 1) if not frame_matches[k][0])
- matching_count = (run_end - run_start + 1) - total_gaps
- if matching_count >= min_consecutive:
- runs.append((run_start, run_end))
- run_start = None
- gap_count = 0
+ # 非匹配帧从 i-gap_count+1 开始,run 结束于其前一帧
+ _close_run(run_start, i - gap_count)
+ run_start, run_last_t, gap_count = None, None, 0
- # 处理末尾 run
if run_start is not None:
- last_idx = len(frame_matches) - 1
- # 回退找到最后一个匹配帧的位置(跳过尾部非匹配帧)
+ last_idx = n - 1
while last_idx >= run_start and not frame_matches[last_idx][0]:
last_idx -= 1
- if last_idx >= run_start:
- # 计算 run 内的总间隙数
- total_gaps = sum(1 for k in range(run_start, last_idx + 1) if not frame_matches[k][0])
- matching_count = (last_idx - run_start + 1) - total_gaps
- if matching_count >= min_consecutive:
- runs.append((run_start, last_idx))
+ _close_run(run_start, last_idx)
# Step 3: 构建 DuplicateSegment
segments: list[DuplicateSegment] = []
for start, end in runs:
- query_start = _get_start(query_chunks[start])
- query_end = _get_end(query_chunks[end])
-
- # 取目标范围(按最佳匹配的目标 chunk 时间范围)
target_indices = [frame_matches[k][2] for k in range(start, end + 1) if frame_matches[k][0]]
- if target_indices:
- t_min = min(target_indices)
- t_max = max(target_indices)
- target_start = _get_start(target_chunks[t_min])
- target_end = _get_end(target_chunks[t_max])
- else:
- target_start = _get_start(target_chunks[0])
- target_end = _get_end(target_chunks[-1])
-
- avg_dist = sum(frame_matches[k][1] for k in range(start, end + 1)) / (end - start + 1)
+ t_min, t_max = min(target_indices), max(target_indices)
+ avg_dist = sum(frame_matches[k][1] for k in range(start, end + 1) if frame_matches[k][0]) / len(target_indices)
segments.append(
DuplicateSegment(
- query_start_ms=query_start,
- query_end_ms=query_end,
- target_start_ms=target_start,
- target_end_ms=target_end,
+ query_start_ms=_get(query_chunks[start], "start_time_ms"),
+ query_end_ms=_get(query_chunks[end], "end_time_ms"),
+ target_start_ms=_get(target_chunks[t_min], "start_time_ms"),
+ target_end_ms=_get(target_chunks[t_max], "end_time_ms"),
avg_distance=avg_dist,
)
)
@@ -423,7 +502,9 @@ def find_duplicate_segments(
class VideoDeduplicator:
"""Video deduplication using multiple fingerprint methods."""
- PHASH_THRESHOLD = 8 # Issue #1658: pHash 汉明距离阈值由 10 收紧到 8,降低不同视频误判率
+ # Issue #1702: 阈值统一来源为模块常量 PHASH_THRESHOLD(#1658 曾收紧到 8,
+ # 后经 staging 真实同源/异源指纹分布重新校准,见 test_phash_threshold_calibration_1702)。
+ PHASH_THRESHOLD = PHASH_THRESHOLD
HISTOGRAM_THRESHOLD = 0.85
@staticmethod
@@ -447,27 +528,36 @@ class VideoDeduplicator:
# 单帧不视为坏指纹(短视频或抽帧不足)
if len(phashes) == 1:
return False
- # 多帧但所有 phash 完全相同 → 黑屏/纯色视频
+ # Issue #1702: 旧逻辑"所有 phash 完全相同即判黑屏"会误杀短视频——
+ # 11s 视频只有几个不同镜头时,相邻 1s 采样帧可能 phash 完全一致(内容
+ # 连续但非黑屏)。黑屏的特征是「大量帧全部无内容」,要求至少 8 帧
+ # 且相同帧占比 >=80% 才判坏;短视频(<8 帧)只有真正单值时交给
+ # _bhattacharyya/融合分兜底,不因"帧都一样"直接跳过。
+ if len(phashes) < 8:
+ return False
unique = set(phashes)
- if len(unique) == 1:
+ same_ratio = sum(1 for x in phashes if x == phashes[0]) / len(phashes)
+ if len(unique) == 1 and same_ratio >= 0.8:
return True
- # 多帧但所有 phash 之间的汉明距离都极小(<3)→ 近似黑屏
+ # 多帧但所有唯一 phash 之间的汉明距离都极小(<3)且占比 >=80% → 近似黑屏
phash_list = list(unique)
- if len(phash_list) >= 2:
- all_distances = []
- for i in range(len(phash_list)):
- for j in range(i + 1, len(phash_list)):
- all_distances.append(hamming_distance(phash_list[i], phash_list[j]))
+ if len(phash_list) >= 2 and same_ratio >= 0.8:
+ all_distances = [
+ hamming_distance(phash_list[i], phash_list[j])
+ for i in range(len(phash_list))
+ for j in range(i + 1, len(phash_list))
+ ]
if all_distances and max(all_distances) < 3:
return True
return False
def compute_fingerprint(self, video_path: str) -> VideoFingerprint:
- """Compute video fingerprint using dynamic keyframe detection.
+ """Compute video fingerprint using dense uniform sampling.
- 使用 detect_keyframe_timestamps() 检测内容感知关键帧,
- 在每个关键帧处取帧计算 pHash + color_histogram。
- 同时保留 MD5 计算和分片数据结构。
+ Issue #1702: 使用 sample_fingerprint_timestamps() 固定 1s 间隔密集均匀
+ 采样(替代动态场景检测抽帧),保证两个同源视频复用片段的帧时刻天然
+ 对齐;每帧取中心 90% 区域(center_crop_frame)计算 pHash + color_histogram,
+ 绕开 random_edge_crop 降重裁剪污染;MD5 仍基于原始帧。
"""
cap = cv2.VideoCapture(video_path)
if not cap.isOpened():
@@ -481,8 +571,8 @@ class VideoDeduplicator:
cap.release()
- # 1. 检测关键帧时间戳
- keyframe_times = detect_keyframe_timestamps(video_path)
+ # 1. 固定间隔密集采样(Issue #1702:替代动态场景检测,保证跨视频时序对齐)
+ keyframe_times = sample_fingerprint_timestamps(duration)
if not keyframe_times:
return VideoFingerprint(
@@ -506,12 +596,15 @@ class VideoDeduplicator:
if not ret:
continue
- # MD5 计算
+ # MD5 计算(基于原始帧,指纹文件级去重不受裁剪影响)
_, buffer = cv2.imencode(".jpg", frame)
md5_hash.update(buffer)
- phash = compute_phash(frame)
- hist = compute_color_histogram(frame)
+ # Issue #1702: pHash / 颜色直方图基于中心 90% 区域,绕开 random_edge_crop
+ # 降重裁剪对指纹的污染(降重只服务外部平台,不污染自查重)。
+ fp_frame = center_crop_frame(frame)
+ phash = compute_phash(fp_frame)
+ hist = compute_color_histogram(fp_frame)
# 计算分片时间范围(从前一个关键帧到下一个关键帧的中点)
prev_boundary = keyframe_times[i - 1] * 1000 if i > 0 else 0
@@ -564,12 +657,22 @@ class VideoDeduplicator:
@staticmethod
def _bhattacharyya_coefficient(hist_a: list[float], hist_b: list[float]) -> float:
- """Bhattacharyya 系数:Σ √(a[i] * b[i]),范围 [0, 1],1=完全相同。"""
+ """Bhattacharyya 系数(概率分布版,范围 [0,1],1=完全相同)。
+
+ Issue #1702: compute_color_histogram 输出 3 通道拼接、每通道独立 NORM_L1
+ (单通道 Σ=1,三通道拼接向量 Σ=3)。旧实现直接 Σ√(a*b) 对三通道拼接向量
+ 算出 ~3(旧 L2 归一化更是算出 ~14.9),不是合法的概率系数。
+ 这里按两个直方图各自的总量归一:BC = Σ√(a*b) / √(Σa·Σb)。
+ - 单通道概率分布(Σa=Σb=1):分母 1,与旧测试/教科书定义一致;
+ - 三通道拼接(Σa=Σb=3):分母 3,结果在 [0,1]。
+ """
min_len = min(len(hist_a), len(hist_b))
- a = hist_a[:min_len]
- b = hist_b[:min_len]
- # 纯标准库计算(不依赖 numpy);max(0.0, ...) 防御上游异常负值导致 sqrt domain error
- return float(sum(math.sqrt(max(0.0, ai * bi)) for ai, bi in zip(a, b, strict=False)))
+ a = [max(0.0, float(x)) for x in hist_a[:min_len]]
+ b = [max(0.0, float(x)) for x in hist_b[:min_len]]
+ # max(0.0, ...) 防御上游异常负值导致 sqrt domain error
+ coeff = sum(math.sqrt(ai * bi) for ai, bi in zip(a, b, strict=False))
+ norm = math.sqrt(sum(a) * sum(b))
+ return float(coeff / norm) if norm > 0 else 0.0
@staticmethod
def _compute_histogram_similarity(
@@ -611,6 +714,72 @@ class VideoDeduplicator:
hist_similarity = VideoDeduplicator._compute_histogram_similarity(hist_a, hist_b) if hist_b else 0.5
return PHASH_WEIGHT * phash_similarity + HISTOGRAM_WEIGHT * hist_similarity
+ @staticmethod
+ def _evaluate_candidate(
+ fingerprint: VideoFingerprint,
+ existing_phashes: list[str],
+ existing_histograms: list,
+ existing_chunk_objects: list,
+ *,
+ query_duration_sec: float,
+ ) -> dict:
+ """评估新视频指纹与单个候选视频的相似度(Issue #1702 共享逻辑)。
+
+ 指标:
+ - min_distances / frame_match_rate:每个新分片到候选视频全局最近邻的汉明距离,
+ 分母取两视频分片数的较小值(支持局部片段复用:短视频复用长视频片段时不被长视频分母稀释)。
+ - temporal_coverage:时序一致连续匹配片段总时长 / 新视频时长(局部复用主指标)。
+ - fusion:pHash 中位数距离 + 颜色直方图的加权融合分。
+
+ Returns:
+ {frame_match_rate, temporal_coverage, segments, median_distance,
+ fusion, matching_frames, min_distances}
+ """
+ query_phashes = fingerprint.keyframe_phashes or []
+ if not query_phashes or not existing_phashes:
+ return {
+ "frame_match_rate": 0.0,
+ "temporal_coverage": 0.0,
+ "segments": [],
+ "median_distance": 64,
+ "fusion": 0.0,
+ "matching_frames": 0,
+ "min_distances": [],
+ }
+
+ min_distances = [min(hamming_distance(ph, ep) for ep in existing_phashes) for ph in query_phashes]
+ matching_frames = sum(1 for d in min_distances if d <= PHASH_THRESHOLD)
+ # 分母取 min(两视频分片数):局部复用时(如 B 的 5 片复用 A 9 片中的若干片)
+ # 命中帧占比不因候选视频更长而被稀释。
+ frame_match_rate = matching_frames / min(len(query_phashes), len(existing_phashes))
+
+ segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects)
+ duration_ms = query_duration_sec * 1000 if query_duration_sec else 0
+ if duration_ms > 0 and segments:
+ covered_ms = sum(s.query_end_ms - s.query_start_ms for s in segments)
+ temporal_coverage = min(covered_ms / duration_ms, 1.0)
+ elif matching_frames > 0:
+ # 无连续片段(时序连贯性不足)时,按匹配帧占比估计覆盖:
+ # 密集 1s 采样下每个分片≈1s 等权时间片,匹配帧数≈命中秒数。
+ temporal_coverage = min(frame_match_rate, 1.0)
+ else:
+ temporal_coverage = 0.0
+
+ median_distance = statistics.median(min_distances) if min_distances else 64
+ fusion = VideoDeduplicator._compute_fusion_score(
+ median_distance, fingerprint.color_histograms, existing_histograms
+ )
+
+ return {
+ "frame_match_rate": frame_match_rate,
+ "temporal_coverage": temporal_coverage,
+ "segments": segments,
+ "median_distance": median_distance,
+ "fusion": fusion,
+ "matching_frames": matching_frames,
+ "min_distances": min_distances,
+ }
+
def check_duplicate(
self,
fingerprint: VideoFingerprint,
@@ -620,6 +789,7 @@ class VideoDeduplicator:
scope: str = "project",
user_id: str = "",
duration_sec: float = 0,
+ exclude_video_id: str | None = None,
) -> Optional[dict]:
"""检查视频是否与已有视频重复。
@@ -636,6 +806,9 @@ class VideoDeduplicator:
scope: "project" 项目内查重(默认),"user" 跨项目全局查重
user_id: 用户 ID(scope="user" 时使用)
duration_sec: 视频时长(秒),用于时长预过滤 ±15%
+ exclude_video_id: 排除的视频 ID(查重自身时用)。recompute-dedup
+ 重算时视频记录已存在,不排除会自匹配(距离 0 分最高)导致
+ duplicate_of 指向自己(Issue #1702 连带修复)。
Returns:
重复信息字典(含 duplicate, duplicate_of, reason, similarity, duplicate_segments),
@@ -643,13 +816,22 @@ class VideoDeduplicator:
"""
video_repo = SQLAlchemyGeneratedVideoRepository(session)
if scope == "user" and user_id:
- dur_min = duration_sec * 0.85 if duration_sec > 0 else 0
- dur_max = duration_sec * 1.15 if duration_sec > 0 else 0
- existing_videos = video_repo.list_by_user(user_id, duration_min=dur_min, duration_max=dur_max)
+ # Issue #1702: 不做 ±15% 时长预过滤。旧逻辑按 duration_sec 缩小候选窗口,
+ # 但局部片段复用的两个视频时长必然不同(证据视频 20s vs 11s,差 42%),
+ # ±15% 窗口让同源视频互相不可见 → is_duplicate 恒 False。
+ # 全量遍历同用户视频(与 compute_duplicate_rate 口径一致),异源视频由
+ # fusion/temporal_coverage 阈值天然过滤(校准:异源最小汉明距离 24)。
+ existing_videos = video_repo.list_by_user(user_id)
else:
existing_videos = video_repo.list_by_project(project_id)
+ best_score = 0.0
+ best_result: Optional[dict] = None
+
for existing in existing_videos:
+ # 排除自身(recompute 时当前视频已在候选列表里,否则自匹配距离 0 必最高分)
+ if exclude_video_id and existing.id == exclude_video_id:
+ continue
if not existing.video_fingerprint:
continue
@@ -677,61 +859,70 @@ class VideoDeduplicator:
if not existing_phashes:
continue
- # 计算每个新关键帧到已有关键帧的最小汉明距离
- min_distances = []
- for phash in fingerprint.keyframe_phashes:
- distances = [hamming_distance(phash, ep) for ep in existing_phashes]
- min_distances.append(min(distances))
-
- # 帧匹配比例检查
- matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD)
- match_ratio = matching_frames / len(min_distances) if min_distances else 0
- if match_ratio < MATCH_RATIO_THRESHOLD:
- continue
-
- # 中位数距离
- median_distance = statistics.median(min_distances) if min_distances else 64
- if median_distance >= self.PHASH_THRESHOLD:
- continue
-
- # 直方图融合(chunk 表优先,回退 JSON 字段;JSON NULL 显式回退空列表)
+ # 直方图 / 分片对象(chunk 表优先,回退 JSON 字段;JSON NULL 显式回退空列表)
if chunk_data:
existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")]
+ existing_chunk_objects = chunk_data
else:
existing_histograms = ef.get("color_histograms") or []
+ existing_chunk_objects = [
+ {"phash_binary": pp, "start_time_ms": 0, "end_time_ms": 0} for pp in existing_phashes
+ ]
- combined_score = self._compute_fusion_score(
- median_distance, fingerprint.color_histograms, existing_histograms
+ # Issue #1702: 统一评估每个候选(含局部片段复用),不再用
+ # "frame_match_rate<0.7 整条跳过" 的硬门槛——局部复用(如 B 结尾 2s
+ # ≈ A 中间 2s)帧比例天然低,但 coverage 能检出。
+ ev = self._evaluate_candidate(
+ fingerprint,
+ existing_phashes,
+ existing_histograms,
+ existing_chunk_objects,
+ query_duration_sec=fingerprint.duration,
+ )
+ logger.debug(
+ "check_duplicate candidate=%s min_distances=%s frame_match_rate=%.3f "
+ "temporal_coverage=%.3f median=%.1f fusion=%.3f segments=%d",
+ existing.id,
+ ev["min_distances"],
+ ev["frame_match_rate"],
+ ev["temporal_coverage"],
+ ev["median_distance"],
+ ev["fusion"],
+ len(ev["segments"]),
)
- if combined_score < DUPLICATE_THRESHOLD:
- continue
-
- # 滑动窗口时序匹配:获取具体重复片段
- existing_chunk_objects = (
- chunk_data
- if chunk_data
- else [{"phash_binary": p, "start_time_ms": 0, "end_time_ms": 0} for p in existing_phashes]
+ # 全片重复判定:融合分过阈 且(帧匹配比例 >=70% 或 局部覆盖 >=50%)
+ is_full_duplicate = ev["fusion"] >= DUPLICATE_THRESHOLD and (
+ ev["frame_match_rate"] >= MATCH_RATIO_THRESHOLD or ev["temporal_coverage"] >= PARTIAL_COVERAGE_THRESHOLD
)
- segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects)
-
- return {
- "duplicate": True,
- "duplicate_of": existing.id,
- "reason": "phash_histogram_fusion",
- "similarity": combined_score,
- "duplicate_segments": [
- {
- "query_start_ms": s.query_start_ms,
- "query_end_ms": s.query_end_ms,
- "target_start_ms": s.target_start_ms,
- "target_end_ms": s.target_end_ms,
- "avg_distance": round(s.avg_distance, 2),
- }
- for s in segments
- ],
- }
+ if is_full_duplicate and ev["fusion"] > best_score:
+ best_score = ev["fusion"]
+ best_result = {
+ "duplicate": True,
+ "duplicate_of": existing.id,
+ "reason": "phash_histogram_fusion",
+ "similarity": ev["fusion"],
+ "duplicate_segments": [
+ {
+ "query_start_ms": s.query_start_ms,
+ "query_end_ms": s.query_end_ms,
+ "target_start_ms": s.target_start_ms,
+ "target_end_ms": s.target_end_ms,
+ "avg_distance": round(s.avg_distance, 2),
+ }
+ for s in ev["segments"]
+ ],
+ }
+ if best_result:
+ return best_result
+ logger.info(
+ "check_duplicate no match (project=%s scope=%s): %d candidates evaluated, best_fusion=%.3f",
+ project_id,
+ scope,
+ len(existing_videos),
+ best_score,
+ )
return None
def check_batch_duplicate(
@@ -763,6 +954,9 @@ class VideoDeduplicator:
video_repo = SQLAlchemyGeneratedVideoRepository(session)
batch_videos = video_repo.list_by_batch(batch_id)
+ best_score = 0.0
+ best_result: Optional[dict] = None
+
for existing in batch_videos:
if existing.id == current_video_id:
continue
@@ -796,59 +990,59 @@ class VideoDeduplicator:
if not existing_phashes:
continue
- min_distances = []
- for phash in fingerprint.keyframe_phashes:
- distances = [hamming_distance(phash, ep) for ep in existing_phashes]
- min_distances.append(min(distances))
-
- # 帧匹配比例检查
- matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD)
- match_ratio = matching_frames / len(min_distances) if min_distances else 0
- if match_ratio < MATCH_RATIO_THRESHOLD:
- continue
-
- median_distance = statistics.median(min_distances) if min_distances else 64
- if median_distance >= self.PHASH_THRESHOLD:
- continue
-
- # 直方图融合(chunk 表优先,回退 JSON 字段;JSON NULL 显式回退空列表)
if chunk_data:
existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")]
+ existing_chunk_objects = chunk_data
else:
existing_histograms = ef.get("color_histograms") or []
+ existing_chunk_objects = [
+ {"phash_binary": pp, "start_time_ms": 0, "end_time_ms": 0} for pp in existing_phashes
+ ]
- combined_score = self._compute_fusion_score(
- median_distance, fingerprint.color_histograms, existing_histograms
+ ev = self._evaluate_candidate(
+ fingerprint,
+ existing_phashes,
+ existing_histograms,
+ existing_chunk_objects,
+ query_duration_sec=fingerprint.duration,
+ )
+ logger.debug(
+ "check_batch_duplicate candidate=%s min_distances=%s frame_match_rate=%.3f "
+ "temporal_coverage=%.3f median=%.1f fusion=%.3f segments=%d",
+ existing.id,
+ ev["min_distances"],
+ ev["frame_match_rate"],
+ ev["temporal_coverage"],
+ ev["median_distance"],
+ ev["fusion"],
+ len(ev["segments"]),
)
- if combined_score < DUPLICATE_THRESHOLD:
- continue
-
- # 滑动窗口时序匹配
- existing_chunk_objects = (
- chunk_data
- if chunk_data
- else [{"phash_binary": p, "start_time_ms": 0, "end_time_ms": 0} for p in existing_phashes]
+ is_full_duplicate = ev["fusion"] >= DUPLICATE_THRESHOLD and (
+ ev["frame_match_rate"] >= MATCH_RATIO_THRESHOLD or ev["temporal_coverage"] >= PARTIAL_COVERAGE_THRESHOLD
)
- segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects)
-
- return {
- "duplicate": True,
- "duplicate_of": existing.id,
- "reason": "batch_phash_histogram_fusion",
- "similarity": combined_score,
- "duplicate_segments": [
- {
- "query_start_ms": s.query_start_ms,
- "query_end_ms": s.query_end_ms,
- "target_start_ms": s.target_start_ms,
- "target_end_ms": s.target_end_ms,
- "avg_distance": round(s.avg_distance, 2),
- }
- for s in segments
- ],
- }
+ if is_full_duplicate and ev["fusion"] > best_score:
+ best_score = ev["fusion"]
+ best_result = {
+ "duplicate": True,
+ "duplicate_of": existing.id,
+ "reason": "batch_phash_histogram_fusion",
+ "similarity": ev["fusion"],
+ "duplicate_segments": [
+ {
+ "query_start_ms": s.query_start_ms,
+ "query_end_ms": s.query_end_ms,
+ "target_start_ms": s.target_start_ms,
+ "target_end_ms": s.target_end_ms,
+ "avg_distance": round(s.avg_distance, 2),
+ }
+ for s in ev["segments"]
+ ],
+ }
+ if best_result:
+ return best_result
+ logger.info("check_batch_duplicate no match (batch=%s): best_fusion=%.3f", batch_id, best_score)
return None
def compute_duplicate_rate(
@@ -897,8 +1091,7 @@ class VideoDeduplicator:
max_duplicate_rate = 0.0
max_visual_similarity = 0.0
match_count = 0
-
- total_duration_ms = fingerprint.duration if fingerprint.duration else 0
+ evaluated = 0
for existing in existing_videos:
if current_video_id and existing.id == current_video_id:
@@ -933,57 +1126,63 @@ class VideoDeduplicator:
if not existing_phashes or not fingerprint.keyframe_phashes:
continue
- min_distances = []
- for phash in fingerprint.keyframe_phashes:
- distances = [hamming_distance(phash, ep) for ep in existing_phashes]
- min_distances.append(min(distances))
-
- # frame_match_rate
- total_frames = len(min_distances)
- if total_frames == 0:
- continue
- matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD)
- frame_match_rate = matching_frames / total_frames
-
- # 帧匹配比例太低则跳过
- if frame_match_rate < 0.3:
- continue
-
- # temporal_coverage_rate via find_duplicate_segments
- existing_chunk_objects = (
- chunk_data
- if chunk_data
- else [{"phash_binary": p, "start_time_ms": 0, "end_time_ms": 0} for p in existing_phashes]
- )
- segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects)
-
- if total_duration_ms > 0 and segments:
- covered_ms = sum(s.query_end_ms - s.query_start_ms for s in segments)
- temporal_coverage_rate = min(covered_ms / total_duration_ms, 1.0)
- else:
- temporal_coverage_rate = 0.0
-
- # duplicate_rate = 0.4 * frame_match_rate + 0.6 * temporal_coverage_rate
- dup_rate = (frame_match_rate * 0.4 + temporal_coverage_rate * 0.6) * 100
-
- # visual_similarity (融合相似度,归一化 0~1)
- median_distance = statistics.median(min_distances) if min_distances else 64
+ # 直方图 / 分片对象(chunk 表优先,回退 JSON 字段;JSON NULL 显式回退空列表)
if chunk_data:
existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")]
+ existing_chunk_objects = chunk_data
else:
- # JSON NULL 显式回退空列表
existing_histograms = ef.get("color_histograms") or []
+ existing_chunk_objects = [
+ {"phash_binary": pp, "start_time_ms": 0, "end_time_ms": 0} for pp in existing_phashes
+ ]
- visual_sim = self._compute_fusion_score(median_distance, fingerprint.color_histograms, existing_histograms)
+ # Issue #1702: 统一评估;frame_match_rate 分母为 min(两视频分片数),
+ # temporal_coverage 时长量纲在 _evaluate_candidate 内统一为毫秒。
+ ev = self._evaluate_candidate(
+ fingerprint,
+ existing_phashes,
+ existing_histograms,
+ existing_chunk_objects,
+ query_duration_sec=fingerprint.duration,
+ )
+ evaluated += 1
+ logger.debug(
+ "compute_duplicate_rate candidate=%s min_distances=%s frame_match_rate=%.3f "
+ "temporal_coverage=%.3f median=%.1f fusion=%.3f segments=%d",
+ existing.id,
+ ev["min_distances"],
+ ev["frame_match_rate"],
+ ev["temporal_coverage"],
+ ev["median_distance"],
+ ev["fusion"],
+ len(ev["segments"]),
+ )
- # 判定是否为重复(融合分数超过阈值)
- if visual_sim >= DUPLICATE_THRESHOLD:
+ # Issue #1702: 去掉 "frame_match_rate<0.3 整条跳过" 硬门槛——
+ # 局部片段复用帧比例天然低;coverage 为主指标,0 匹配自然得 0 分。
+ # duplicate_rate = 0.4 * frame_match_rate + 0.6 * temporal_coverage
+ dup_rate = (min(ev["frame_match_rate"], 1.0) * 0.4 + ev["temporal_coverage"] * 0.6) * 100
+
+ # 全片重复计数与 check_duplicate 判定口径一致
+ if ev["fusion"] >= DUPLICATE_THRESHOLD and (
+ ev["frame_match_rate"] >= MATCH_RATIO_THRESHOLD or ev["temporal_coverage"] >= PARTIAL_COVERAGE_THRESHOLD
+ ):
match_count += 1
if dup_rate > max_duplicate_rate:
max_duplicate_rate = dup_rate
- max_visual_similarity = visual_sim
+ max_visual_similarity = ev["fusion"]
+ logger.info(
+ "compute_duplicate_rate done (project=%s scope=%s): evaluated=%d max_rate=%.2f%% "
+ "max_visual_sim=%.3f matches=%d",
+ project_id,
+ scope,
+ evaluated,
+ max_duplicate_rate,
+ max_visual_similarity,
+ match_count,
+ )
return {
"duplicate_rate": round(max(max_duplicate_rate, 0.0), 2),
"visual_similarity": round(max_visual_similarity, 4),
@@ -999,18 +1198,20 @@ def _save_fingerprint_chunks(
session: Session,
) -> None:
"""将指纹分片数据批量写入 video_fingerprint_chunks 表。幂等:已有数据时跳过。"""
- # 幂等检查:已有分片数据则跳过
- existing_count = (
- session.query(VideoFingerprintChunkModel).filter(VideoFingerprintChunkModel.video_id == video_id).count()
- )
- if existing_count > 0:
- logger.debug("Fingerprint chunks already exist for video %s (%d chunks), skipping", video_id, existing_count)
- return
-
if not fingerprint.chunks:
logger.warning("No chunks in fingerprint for video %s, skipping chunk save", video_id)
return
+ # Issue #1702: recompute-dedup 重算时指纹算法已变(中心裁剪 + 新阈值),
+ # 旧分片必须替换而非跳过(旧实现"有数据就跳过"导致重算不刷新分片表)。
+ deleted = (
+ session.query(VideoFingerprintChunkModel)
+ .filter(VideoFingerprintChunkModel.video_id == video_id)
+ .delete(synchronize_session=False)
+ )
+ if deleted:
+ logger.info("Replaced %d stale fingerprint chunks for video %s", deleted, video_id)
+
chunk_models = fingerprint.to_chunk_models(video_id, project_id, user_id)
session.bulk_save_objects(chunk_models)
logger.info("Saved %d fingerprint chunks for video %s", len(chunk_models), video_id)
@@ -1032,9 +1233,17 @@ def check_duplicate_task(self: Task, generated_video_id: str) -> dict:
raise ValueError(f"Generated video {generated_video_id} not found")
local_path = os.path.join(temp_dir, f"{generated_video_id}.mp4")
- storage_service.download_file(
- f"projects/{video.project_id}/generated/{generated_video_id}/{generated_video_id}.mp4", local_path
- )
+ # Issue #1702: recompute 走的是 OSS 重新下载路径(正常生成流程用本地渲染文件,
+ # 不经此任务)。成片真实 OSS key 是生成时的
+ # generated/projects/{pid}/tasks/{task_id}/rendered_*.mp4(见 generation.py
+ # _upload_and_record),旧代码硬编码 projects/{pid}/generated/{vid}/{vid}.mp4
+ # 这个从不存在的 key,导致所有 recompute 任务下载 404、查重数据永远无法重算。
+ # 优先从 file_url 解析真实 key,旧 key 模式仅作回退。
+ download_key = getattr(video, "file_url", "") or ""
+ if not download_key:
+ download_key = f"projects/{video.project_id}/generated/{generated_video_id}/{generated_video_id}.mp4"
+ logger.warning("video %s has no file_url, falling back to legacy key %s", generated_video_id, download_key)
+ storage_service.download_file(download_key, local_path)
fingerprint = deduplicator.compute_fingerprint(local_path)
@@ -1045,7 +1254,10 @@ def check_duplicate_task(self: Task, generated_video_id: str) -> dict:
session,
scope="user",
user_id=video.user_id,
- duration_sec=fingerprint.duration / 1000 if fingerprint.duration else 0,
+ # Issue #1702: fingerprint.duration 单位已经是秒,旧代码 /1000 导致
+ # ±15% 时长预过滤窗口缩到 ~0.013s,scope=user 的跨项目查重永远返回 None。
+ duration_sec=fingerprint.duration if fingerprint.duration else 0,
+ exclude_video_id=generated_video_id,
)
video.video_fingerprint = fingerprint.to_dict()
diff --git a/apps/worker/video_processing/dedup_helpers.py b/apps/worker/video_processing/dedup_helpers.py
index b7e6d4965..d0ee22da8 100755
--- a/apps/worker/video_processing/dedup_helpers.py
+++ b/apps/worker/video_processing/dedup_helpers.py
@@ -92,7 +92,8 @@ def create_video_record_and_dedup(
logger.warning("Failed to save fingerprint chunks for %s: %s", video_id, chunk_err)
# (a) 历史成片查重(跨项目全局 + 时长预过滤)
- duration_sec = fingerprint.duration / 1000 if fingerprint.duration else 0
+ # Issue #1702: fingerprint.duration 单位是秒,旧代码 /1000 让时长预过滤失效
+ duration_sec = fingerprint.duration if fingerprint.duration else 0
duplicate_result = deduplicator.check_duplicate(
fingerprint,
project_id,
@@ -100,6 +101,7 @@ def create_video_record_and_dedup(
scope="user",
user_id=user_id,
duration_sec=duration_sec,
+ exclude_video_id=video_id,
)
# (b) 批次内查重(仅当有 batch_id 时)
diff --git a/apps/worker/worker_app/celery_app.py b/apps/worker/worker_app/celery_app.py
index d953c23f4..34d40ad09 100755
--- a/apps/worker/worker_app/celery_app.py
+++ b/apps/worker/worker_app/celery_app.py
@@ -6,6 +6,24 @@ celery_app = Celery(settings.worker_name)
celery_app.conf.broker_url = settings.broker_url
celery_app.conf.result_backend = settings.result_backend
celery_app.conf.broker_connection_retry_on_startup = True
+
+# #1714 队列隔离:generation(高优,独占 worker)/ transcode(素材转码)/ celery(默认)
+from packages.shared.celery_queues import ( # noqa: E402
+ GENERATION_WORKER_PREFETCH_MULTIPLIER,
+ apply_queue_settings,
+)
+
+apply_queue_settings(celery_app)
+# 长渲染任务预取 1,避免任务被预取占住导致调度不均
+celery_app.conf.worker_prefetch_multiplier = GENERATION_WORKER_PREFETCH_MULTIPLIER
+celery_app.conf.task_acks_late = True # worker 崩溃时未完成任务重回队列,由执行前守卫丢弃作废消息
+# worker 进程被 OOM/容器硬杀时拒绝 ack,消息留在队列由其他 worker 接手
+celery_app.conf.task_reject_on_worker_lost = True
+# Redis broker 消息可见性超时(#1714):acks_late 下,消息被预取后 visibility_timeout
+# 内未 ack 才会重投。长任务(ingest HEVC 转码 20-30 分钟、生成硬超时 11 分钟)
+# 必须远大于最长执行时间,否则正常任务会在执行中被误重投;4 小时覆盖最长转码 + 余量。
+celery_app.conf.broker_transport_options = {"visibility_timeout": 4 * 60 * 60}
+
celery_app.conf.imports = (
"worker_app.tasks.health",
"worker_app.tasks.ingest",
@@ -22,10 +40,25 @@ celery_app.conf.imports = (
)
# Celery Beat 定时任务调度
+# 注:worker 单实例内嵌 beat(entrypoint-worker.sh -B),定时任务不会重复执行
celery_app.conf.beat_schedule = {
+ # pending 任务超时清理:worker 停止消费后,卡 pending 的任务 15 分钟内释放限流名额
"cleanup-stale-pending-tasks": {
"task": "worker.cleanup_stale_pending_tasks",
+ "schedule": 300.0, # 每 5 分钟(秒)
+ "options": {"expires": 240}, # 4 分钟过期,避免堆积
+ },
+ # running 孤儿任务巡检:容器重启/进程被杀后卡 running 的任务,20 分钟无更新则判失败
+ "cleanup-stale-running-tasks": {
+ "task": "worker.cleanup_stale_running_tasks",
+ "schedule": 300.0, # 每 5 分钟(秒)
+ "options": {"expires": 240},
+ },
+ # 上传/转码链路孤儿巡检:worker 重启丢 prefetch 消息后,卡 pending/processing
+ # 的 ingest_job + asset 占位超时标终态(#1714)。转码任务较长,10 分钟一轮
+ "cleanup-stale-ingest-jobs": {
+ "task": "worker.cleanup_stale_ingest_jobs",
"schedule": 600.0, # 每 10 分钟(秒)
- "options": {"expires": 300}, # 5 分钟过期,避免堆积
+ "options": {"expires": 540},
},
}
diff --git a/apps/worker/worker_app/tasks/_startup.py b/apps/worker/worker_app/tasks/_startup.py
index c841f4406..afe329704 100644
--- a/apps/worker/worker_app/tasks/_startup.py
+++ b/apps/worker/worker_app/tasks/_startup.py
@@ -7,11 +7,83 @@ from worker_app.db import SessionLocal
logger = logging.getLogger(__name__)
-# 孤儿任务超时阈值:渲染任务超过此时间未更新则视为卡死
-ORPHAN_TASK_TIMEOUT_MINUTES = 10
-# Pending 任务超时阈值:pending 任务在队列中等待超过此时间则自动清理
-PENDING_TASK_TIMEOUT_MINUTES = 30
+def cleanup_stale_running_with_session(repo, timeout_minutes: int) -> int:
+ """清理超时未更新的 running GenerationTask(可注入 repo 的纯核心,便于单测)。
+
+ Returns:
+ 清理的任务数量
+ """
+ return len(cleanup_stale_running_with_session_ids(repo, timeout_minutes))
+
+
+def cleanup_stale_running_with_session_ids(repo, timeout_minutes: int) -> list[tuple[str, str]]:
+ """同 cleanup_stale_running_with_session,返回 [(task_id, celery_task_id), ...]。"""
+ fn = getattr(repo, "cleanup_stale_running_with_ids", None)
+ if fn is not None:
+ return fn(timeout_minutes)
+ # 旧仓储无 _with_ids 方法:降级为计数,无法撤销消息(执行前状态守卫兜底)
+ count = repo.cleanup_stale_running(timeout_minutes)
+ return [("", "") for _ in range(count)]
+
+
+def cleanup_stale_pending_with_session(repo, timeout_minutes: int) -> int:
+ """清理超时 pending GenerationTask(可注入 repo 的纯核心,便于单测)。
+
+ Returns:
+ 清理的任务数量
+ """
+ return len(cleanup_stale_pending_with_session_ids(repo, timeout_minutes))
+
+
+def cleanup_stale_pending_with_session_ids(repo, timeout_minutes: int) -> list[tuple[str, str]]:
+ """同 cleanup_stale_pending_with_session,返回 [(task_id, celery_task_id), ...]。"""
+ fn = getattr(repo, "cleanup_stale_pending_with_ids", None)
+ if fn is not None:
+ return fn(timeout_minutes)
+ count = repo.cleanup_stale_pending(timeout_minutes)
+ return [("", "") for _ in range(count)]
+
+
+def _revoke_and_purge_stale_messages(items: list[tuple[str, str]]) -> int:
+ """把清理掉的任务对应的 Celery 消息撤销并从 Redis 队列清除(#1714)。
+
+ 防止「DB 已标 failed,但队列消息还在 → 重投执行 → 非法状态转换 → 半成品」。
+ 失败不阻断清理流程(执行前状态守卫是第二道防线)。
+ """
+ biz_ids = [tid for tid, _ in items if tid]
+ celery_ids = [cid for _, cid in items if cid]
+ if not biz_ids and not celery_ids:
+ return 0
+ try:
+ from worker_app.celery_app import celery_app as app
+ from worker_app.core.config import get_settings
+
+ from packages.shared.celery_orphan_guard import revoke_and_purge
+
+ broker_url = get_settings().broker_url
+ return revoke_and_purge(
+ app,
+ broker_url,
+ business_task_ids=biz_ids,
+ celery_task_ids=celery_ids,
+ )
+ except Exception as e: # noqa: BLE001
+ logger.error("撤销作废任务队列消息失败(执行前守卫仍会兜底): %s", e, exc_info=True)
+ return 0
+
+
+# 孤儿任务超时阈值:running 任务超过此时间无进度更新则视为卡死。
+# 依据:worker.generate_video 硬超时 time_limit=11 分钟,正常任务不可能超过;
+# 20 分钟阈值覆盖硬超时 + 重试 + 余量,绝不误杀正常任务。
+ORPHAN_TASK_TIMEOUT_MINUTES = 20
+
+# Pending 任务超时阈值:任务创建后超过此时间仍未开始执行则判死。
+# 注意区分 running 孤儿阈值(20 分钟):pending 是「排队等待」时间,
+# 队列积压(如 20+ 转码任务)时视频生成可能正常排队较久,阈值必须放宽,
+# 避免正常排队任务被误杀。队列隔离(#1714)后 generation 队列独占 worker,
+# 理论上排队极短;保留 45 分钟作为兜底,覆盖 worker 短暂停止消费的场景。
+PENDING_TASK_TIMEOUT_MINUTES = 45
def cleanup_orphan_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> int: # pragma: no cover
@@ -32,11 +104,16 @@ def cleanup_orphan_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) ->
try:
session = SessionLocal()
- repo = SQLAlchemyGenerationTaskRepository(session)
- count = repo.cleanup_stale_running(timeout_minutes)
- session.close()
+ try:
+ repo = SQLAlchemyGenerationTaskRepository(session)
+ items = cleanup_stale_running_with_session_ids(repo, timeout_minutes)
+ finally:
+ session.close()
+ count = len(items)
if count > 0:
logger.warning("清理了 %d 个超时的孤儿 GenerationTask(超过 %d 分钟未更新)", count, timeout_minutes)
+ purged = _revoke_and_purge_stale_messages(items)
+ logger.info("孤儿任务对应队列消息撤销/清除完成: %d 条", purged)
else:
logger.info("无孤儿 GenerationTask 需要清理")
return count
@@ -70,7 +147,9 @@ def cleanup_stale_jobs(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> in
.all()
)
count = 0
+ stale_items: list[tuple[str, str]] = []
for model in stale_jobs:
+ stale_items.append((model.id, getattr(model, "celery_task_id", "") or ""))
model.status = JobStatus.FAILED.value
model.error_message = f"任务执行中断(超过 {timeout_minutes} 分钟未更新)"
count += 1
@@ -80,6 +159,8 @@ def cleanup_stale_jobs(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> in
else:
logger.info("无孤儿 Job 需要清理")
session.close()
+ if count > 0:
+ _revoke_and_purge_generation(stale_items)
return count
except Exception as e:
logger.error("清理孤儿 Job 失败: %s", e, exc_info=True)
@@ -105,9 +186,12 @@ def cleanup_stale_pending_tasks(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINU
session = SessionLocal()
try:
repo = SQLAlchemyGenerationTaskRepository(session)
- count = repo.cleanup_stale_pending(timeout_minutes)
+ items = cleanup_stale_pending_with_session_ids(repo, timeout_minutes)
+ count = len(items)
if count > 0:
logger.warning("清理了 %d 个超时的 pending GenerationTask(超过 %d 分钟未处理)", count, timeout_minutes)
+ purged = _revoke_and_purge_stale_messages(items)
+ logger.info("超时 pending 任务对应队列消息撤销/清除完成: %d 条", purged)
else:
logger.info("无超时 pending GenerationTask 需要清理")
return count
@@ -118,6 +202,31 @@ def cleanup_stale_pending_tasks(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINU
session.close()
+def _revoke_and_purge_generation(items: list[tuple[str, str]]) -> int:
+ """撤销 Job 表孤儿任务(TTS/配音等)的队列消息,队列覆盖全部已知队列。"""
+ biz_ids = [tid for tid, _ in items if tid]
+ celery_ids = [cid for _, cid in items if cid]
+ if not biz_ids and not celery_ids:
+ return 0
+ try:
+ from worker_app.celery_app import celery_app as app
+ from worker_app.core.config import get_settings
+
+ from packages.shared.celery_orphan_guard import revoke_and_purge
+
+ broker_url = get_settings().broker_url
+ return revoke_and_purge(
+ app,
+ broker_url,
+ business_task_ids=biz_ids,
+ celery_task_ids=celery_ids,
+ queue_names=("generation", "transcode", "celery"),
+ )
+ except Exception as e: # noqa: BLE001
+ logger.error("撤销 Job 队列消息失败: %s", e, exc_info=True)
+ return 0
+
+
def cleanup_all_stale_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> dict: # pragma: no cover
"""统一清理所有超时的孤儿任务。
@@ -148,3 +257,32 @@ def _on_worker_ready(sender, **kwargs): # pragma: no cover
result = cleanup_all_stale_tasks()
total = result["generation_tasks"] + result["jobs"]
logger.info("Worker 启动清理完成,共清理 %d 个孤儿任务", total)
+
+
+@worker_ready.connect
+def _recover_stuck_ingest_jobs_on_ready(sender, **kwargs): # pragma: no cover
+ """Worker 启动完成后恢复卡死在 processing 的 ingest_job(#1714)。
+
+ 容器重启/进程 OOM 导致 transcode 队列 unacked 消息未重投时,processing
+ ingest_job 会永久卡死。启动时扫描 processing 超 10 分钟的 job,CAS 重置
+ pending 并重新派单;Redis 锁保证同容器 generation/transcode 双 worker
+ 只有一个执行恢复。旧消息若后来重投,ingest_asset 执行前守卫会丢弃。
+ """
+ try:
+ from packages.application.ingest_orphan_cleanup import (
+ make_redis_recovery_lock,
+ recover_stuck_ingest_jobs_on_startup,
+ )
+
+ session = SessionLocal()
+ try:
+ recovered = recover_stuck_ingest_jobs_on_startup(
+ session,
+ lock_acquire=make_redis_recovery_lock(),
+ stuck_minutes=10,
+ )
+ finally:
+ session.close()
+ logger.info("Worker 启动 ingest 恢复完成,共重新派单 %d 个卡死任务", recovered)
+ except Exception as e: # noqa: BLE001 — 启动恢复失败不能阻断 worker 起服
+ logger.error("启动 ingest 恢复扫描失败(beat 巡检仍会兜底标 failed): %s", e, exc_info=True)
diff --git a/apps/worker/worker_app/tasks/cleanup.py b/apps/worker/worker_app/tasks/cleanup.py
index 28d5e7e22..ff29840dc 100644
--- a/apps/worker/worker_app/tasks/cleanup.py
+++ b/apps/worker/worker_app/tasks/cleanup.py
@@ -1,17 +1,27 @@
"""定期清理任务 — Celery Beat 调度。
包含:
-- cleanup_stale_pending_tasks: 定期清理卡在 pending 超时的 generation_tasks
+- cleanup_stale_pending_tasks: 定期清理卡在 pending 超时的 generation_tasks(worker 停止消费时占位)
+- cleanup_stale_running_tasks: 定期清理卡在 running 超时的 generation_tasks(容器重启/进程被杀后的孤儿)
"""
import logging
from celery import shared_task
from worker_app.tasks._startup import (
+ ORPHAN_TASK_TIMEOUT_MINUTES,
PENDING_TASK_TIMEOUT_MINUTES,
+ cleanup_orphan_tasks,
+ cleanup_stale_jobs,
cleanup_stale_pending_tasks,
)
+from packages.application.ingest_orphan_cleanup import (
+ ASSET_ORPHAN_TIMEOUT_MINUTES,
+ INGEST_PENDING_TIMEOUT_MINUTES,
+ INGEST_PROCESSING_TIMEOUT_MINUTES,
+)
+
logger = logging.getLogger(__name__)
@@ -19,12 +29,14 @@ logger = logging.getLogger(__name__)
def scheduled_cleanup_stale_pending(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINUTES) -> dict:
"""Celery Beat 调度的定期任务:清理超时的 pending 任务。
- 每 10 分钟执行一次(由 celery_app.py 的 beat_schedule 配置),
+ 每 5 分钟执行一次(由 celery_app.py 的 beat_schedule 配置),
查找所有 status='pending' 且 created_at < NOW() - timeout_minutes
- 的 generation_tasks,批量更新为 failed。
+ 的 generation_tasks,批量更新为 failed,释放限流名额;同时 revoke 并清除
+ Redis 队列中对应的 Celery 消息,杜绝作废消息重投执行(#1714)。
Args:
- timeout_minutes: 超时时间(分钟),默认 30 分钟
+ timeout_minutes: 超时时间(分钟),默认 45 分钟(pending 排队阈值放宽,
+ 与 running 孤儿 20 分钟区分,避免正常排队任务被误杀)
Returns:
{"cleaned": int}
@@ -33,3 +45,83 @@ def scheduled_cleanup_stale_pending(timeout_minutes: int = PENDING_TASK_TIMEOUT_
if count > 0:
logger.info("[Beat] 清理了 %d 个超时 pending 任务(超时阈值 %d 分钟)", count, timeout_minutes)
return {"cleaned": count}
+
+
+@shared_task(name="worker.cleanup_stale_running_tasks")
+def scheduled_cleanup_stale_running(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> dict:
+ """Celery Beat 调度的定期任务:清理超时的 running 孤儿任务。
+
+ 每 5 分钟执行一次。worker_ready 信号只在 worker 启动时清一次,
+ 若 worker 没重启但任务卡死(上传挂起、进程 OOM 被内核杀掉等),
+ 任务会永久卡在 running 占位。此任务做持续兜底:
+ 查找 status='running' 且 updated_at < NOW() - timeout_minutes 的任务,
+ 标记为 failed(原因:容器重启/超时中断),同时清理 Job 表孤儿。
+
+ Args:
+ timeout_minutes: 超时时间(分钟),默认 20 分钟
+ (worker.generate_video 硬超时 11 分钟,正常任务不可能超过 20 分钟)
+
+ Returns:
+ {"generation_tasks": int, "jobs": int}
+ """
+ gen_count = cleanup_orphan_tasks(timeout_minutes)
+ job_count = cleanup_stale_jobs(timeout_minutes)
+ total = gen_count + job_count
+ if total > 0:
+ logger.warning(
+ "[Beat] 清理孤儿任务: running GenerationTask=%d, Job=%d(超时阈值 %d 分钟)",
+ gen_count,
+ job_count,
+ timeout_minutes,
+ )
+ return {"generation_tasks": gen_count, "jobs": job_count}
+
+
+@shared_task(name="worker.cleanup_stale_ingest_jobs")
+def scheduled_cleanup_stale_ingest_jobs(
+ processing_timeout_minutes: int = INGEST_PROCESSING_TIMEOUT_MINUTES,
+ pending_timeout_minutes: int = INGEST_PENDING_TIMEOUT_MINUTES,
+ orphan_asset_timeout_minutes: int = ASSET_ORPHAN_TIMEOUT_MINUTES,
+) -> dict:
+ """Celery Beat 调度:清理上传/转码链路(IngestJob + Asset)孤儿记录。
+
+ 每 10 分钟执行一次。worker 容器重启/进程 OOM 时,已 prefetch 的 transcode
+ celery 消息会丢失(队列里也不存在),ingest_job 永久卡 pending/processing、
+ asset 永久卡 processing/uploading,没有兜底永远不会恢复(#1714)。
+
+ - ingest_job processing > processing_timeout_minutes / pending > pending_timeout_minutes
+ → 标 failed;关联 asset 占位(processing/uploading)联动标 error
+ - 无 ingest_job 关联、created_at > orphan_asset_timeout_minutes 的占位 asset
+ → 标 error
+ - 作废 celery 消息 revoke + 物理清除(防重投,执行前守卫是第二道防线)
+ """
+ from worker_app.db import SessionLocal
+
+ from packages.application.ingest_orphan_cleanup import (
+ cleanup_orphan_processing_assets,
+ cleanup_stale_ingest_jobs,
+ revoke_stale_ingest_messages,
+ )
+
+ session = SessionLocal()
+ try:
+ job_items, asset_ids = cleanup_stale_ingest_jobs(
+ session,
+ processing_timeout_minutes=processing_timeout_minutes,
+ pending_timeout_minutes=pending_timeout_minutes,
+ )
+ orphan_asset_ids = cleanup_orphan_processing_assets(session, timeout_minutes=orphan_asset_timeout_minutes)
+ finally:
+ session.close()
+
+ purged = revoke_stale_ingest_messages(job_items) if job_items else 0
+ total_jobs = len(job_items)
+ total_assets = len(set(asset_ids) | set(orphan_asset_ids))
+ if total_jobs or total_assets:
+ logger.warning(
+ "[Beat] 清理 ingest 链路孤儿: stale_jobs=%d, assets→error=%d, 队列清除消息=%d",
+ total_jobs,
+ total_assets,
+ purged,
+ )
+ return {"stale_jobs": total_jobs, "assets_to_error": total_assets, "purged_messages": purged}
diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py
index 35e195bef..9a0a755e7 100644
--- a/apps/worker/worker_app/tasks/generation.py
+++ b/apps/worker/worker_app/tasks/generation.py
@@ -22,6 +22,8 @@ from worker_app.celery_app import celery_app
from worker_app.db import SessionLocal
from worker_app.tasks.generation_plan_builder import build_error_info as _build_error_info
+from packages.shared.celery_orphan_guard import TERMINAL_STATUS_VALUES
+
OUTPUT_WIDTH = 1280
OUTPUT_HEIGHT = 720
OUTPUT_FPS = 25.0
@@ -658,6 +660,34 @@ def generate_video(self, task_id: str) -> dict:
finally:
_session.close()
+ # ── 0. 执行前状态守卫(#1714):任务已被超时清理/孤儿恢复标记为终态时,
+ # 这是作废消息(worker 崩溃前未 ack 的旧消息重投/重复投递),直接丢弃,
+ # 不进入渲染,杜绝 failed→running 非法转换后继续跑产出半成品。
+ if gen_task is not None and gen_task.status.value in TERMINAL_STATUS_VALUES:
+ logger.warning(
+ "[task_id=%s] 任务状态已为 %s,丢弃作废消息,不执行渲染",
+ task_id,
+ gen_task.status.value,
+ )
+ return {
+ "status": "discarded",
+ "task_id": task_id,
+ "reason": f"task already terminal: {gen_task.status.value}",
+ }
+
+ # 标记任务为 running —— 必须成功:状态机非法转换(如 failed→running)说明
+ # 任务已被作废,安全中止,禁止继续执行。
+ if not _update_task_status(task_id, "mark_processing"):
+ logger.error(
+ "[task_id=%s] 标记 running 失败(任务可能已被作废/取消),安全中止,不执行渲染",
+ task_id,
+ )
+ return {
+ "status": "discarded",
+ "task_id": task_id,
+ "reason": "claim failed (invalid state transition)",
+ }
+
# 记录接收任务日志
if gen_task:
gen_task.append_log(
@@ -669,8 +699,6 @@ def generate_video(self, task_id: str) -> dict:
)
_flush_logs(task_id, gen_task)
- # 标记任务为 running
- _update_task_status(task_id, "mark_processing")
_update_task_progress(task_id, 10, "任务启动")
try:
diff --git a/apps/worker/worker_app/tasks/ingest.py b/apps/worker/worker_app/tasks/ingest.py
index 7b48f0f45..4d7094952 100755
--- a/apps/worker/worker_app/tasks/ingest.py
+++ b/apps/worker/worker_app/tasks/ingest.py
@@ -359,6 +359,54 @@ def validate_transcode_output(
return True
+def _original_key_from_storage_key(storage_key: str) -> str:
+ """从可能被 HEVC 转码改写的 storage_key 还原原始 key。
+
+ 转码成功后 key 形如 uploads//IMG_2282_h264.MOV,
+ 占位 asset 以原始 key uploads//IMG_2282.MOV 创建。
+ """
+ if not storage_key:
+ return storage_key
+ _p = Path(storage_key)
+ if _p.stem.endswith("_h264"):
+ return str(_p.parent / (_p.stem[: -len("_h264")] + _p.suffix))
+ return storage_key
+
+
+def _resolve_placeholder_asset(asset_repo, job, original_storage_key):
+ """找到 complete 阶段创建的 PROCESSING 占位 asset(Issue #1714)。
+
+ HEVC 转码成功后 job.storage_key 会被改写为 *_h264 新 key,旧实现用新 key
+ 回查占位必然落空,进而兜底新建一条 READY 记录,导致原占位永久卡 processing。
+
+ 查找优先级:
+ 1. job.asset_id(complete 派单时透传的占位 id,最可靠,不依赖 key);
+ 2. 原始 storage_key(占位记录以原始 key 创建);
+ 3. 当前 job.storage_key(未转码/降级场景与原始 key 相同)。
+
+ 找不到返回 None(旧链路兼容,由调用方兜底新建并告警)。
+ """
+ asset_id = getattr(job, "asset_id", "") or ""
+ if asset_id:
+ try:
+ found = asset_repo.find_by_id(asset_id)
+ if found is not None:
+ return found
+ except Exception as find_err:
+ logger.warning("占位 asset 按 id 查询失败 asset_id=%s: %s", asset_id, find_err)
+ for key in (original_storage_key, getattr(job, "storage_key", "")):
+ if not key:
+ continue
+ try:
+ found = asset_repo.find_by_storage_key(key)
+ except Exception:
+ logger.warning("find_by_storage_key not available, trying fallback lookup")
+ found = None
+ if found is not None:
+ return found
+ return None
+
+
@celery_app.task(name="worker.ingest_asset")
def ingest_asset(job_id: str) -> dict:
"""
@@ -380,6 +428,23 @@ def ingest_asset(job_id: str) -> dict:
if job is None:
return {"status": "failed", "error": "job not found"}
+ # ── 执行前状态守卫(#1714):job 已终态(失败/完成)说明这是作废消息
+ # (超时清理标记 failed 后旧消息重投、或重复投递),直接丢弃不执行,
+ # 避免重复转码、重复回写。processing 是本任务自己第一次置位前的旧消息
+ # 极少见,保守起见也丢弃(processing 的占位由恢复流程处理)。
+ current_status = job.status.value if hasattr(job.status, "value") else str(job.status)
+ if current_status in ("failed", "completed"):
+ logger.warning(
+ "[ingest job_id=%s] 任务状态已为 %s,丢弃作废消息,不执行转码",
+ job_id,
+ current_status,
+ )
+ return {"status": "discarded", "job_id": job_id, "reason": f"job already terminal: {current_status}"}
+
+ # 记录原始 storage_key:HEVC 转码成功后 job.storage_key 会改写为 *_h264,
+ # 而 complete 阶段的占位 asset 始终以原始 key 创建,关联回写必须保留它。
+ original_storage_key = job.storage_key
+
# Update job status to PROCESSING
job.status = IngestJobStatus.PROCESSING
job.updated_at = datetime.now(timezone.utc)
@@ -624,22 +689,41 @@ def ingest_asset(job_id: str) -> dict:
error_reason,
)
- asset = Asset.create(
- project_id=job.project_id,
- library_id=job.library_id,
- name=filename,
- storage_key=job.storage_key,
- mime_type=mime_type,
- metadata={"source": "upload", "ingest_error": error_reason},
- file_size=int(metadata.get("size_bytes", 0)),
- duration=float(metadata.get("duration", 0)),
- width=int(metadata.get("width", 0)),
- height=int(metadata.get("height", 0)),
- codec=metadata.get("codec") or None,
- status=AssetStatus.ERROR,
- file_hash=job.file_hash,
- )
- asset_repo.create(asset)
+ placeholder = _resolve_placeholder_asset(asset_repo, job, original_storage_key)
+ if placeholder is not None:
+ # 回写占位记录:标 ERROR(Issue #1714:禁止新建第二条导致占位孤儿)
+ asset = placeholder
+ asset.mime_type = mime_type
+ asset.metadata = {"source": "upload", "ingest_error": error_reason}
+ asset.file_size = int(metadata.get("size_bytes", 0))
+ asset.duration = float(metadata.get("duration", 0)) or None
+ asset.width = int(metadata.get("width", 0)) or None
+ asset.height = int(metadata.get("height", 0)) or None
+ codec_val = metadata.get("codec")
+ if codec_val:
+ asset.codec = str(codec_val)
+ asset.status = AssetStatus.ERROR
+ asset.updated_at = datetime.now(timezone.utc)
+ asset_repo.update(asset)
+ else:
+ # 旧链路兜底:无占位记录(如历史 job 重跑)才新建
+ logger.warning("无效素材且未找到占位记录,兜底新建 ERROR asset: job_id=%s", job_id)
+ asset = Asset.create(
+ project_id=job.project_id,
+ library_id=job.library_id,
+ name=filename,
+ storage_key=job.storage_key,
+ mime_type=mime_type,
+ metadata={"source": "upload", "ingest_error": error_reason},
+ file_size=int(metadata.get("size_bytes", 0)),
+ duration=float(metadata.get("duration", 0)),
+ width=int(metadata.get("width", 0)),
+ height=int(metadata.get("height", 0)),
+ codec=metadata.get("codec") or None,
+ status=AssetStatus.ERROR,
+ file_hash=job.file_hash,
+ )
+ asset_repo.create(asset)
# Update job status to FAILED
job.status = IngestJobStatus.FAILED
@@ -656,16 +740,19 @@ def ingest_asset(job_id: str) -> dict:
"error": error_reason,
}
- # 查找已存在的 Asset 记录(由 API 端在上传完成时立即创建为 PROCESSING 状态)
- existing_asset = None
- try:
- existing_asset = asset_repo.find_by_storage_key(job.storage_key)
- except Exception:
- logger.warning("find_by_storage_key not available, trying fallback lookup")
+ # 查找 complete 阶段创建的占位 Asset 记录(Issue #1714)。
+ # 必须用原始 storage_key / job.asset_id 关联——HEVC 转码后 job.storage_key
+ # 已改写为 *_h264,用新 key 回查占位必然落空,旧实现因此兜底新建 READY 记录,
+ # 导致原 PROCESSING 占位永久卡住(每个 HEVC 视频产生两条记录)。
+ existing_asset = _resolve_placeholder_asset(asset_repo, job, original_storage_key)
if existing_asset is None:
- # 兜底:如果 API 端没有预先创建 Asset(旧版本兼容),则创建新记录
- logger.info("No pre-created asset found for storage_key=%s, creating new", job.storage_key)
+ # 兜底:仅当确实没有占位记录(旧版本 API / 历史 job 重跑)才新建。
+ logger.warning(
+ "No placeholder asset found for job_id=%s original_key=%s, creating new",
+ job_id,
+ original_storage_key,
+ )
metadata["source"] = "upload"
asset = Asset.create(
project_id=job.project_id,
@@ -685,8 +772,13 @@ def ingest_asset(job_id: str) -> dict:
)
asset_repo.create(asset)
else:
- # 更新已有的 Asset 记录,补充元数据并将状态改为 READY
+ # 更新占位记录:补充元数据、置 READY。转码成功时 storage_key 同步改写为
+ # *_h264(播放/下载走转码产物),原始 key 记入 metadata 可溯源。
asset = existing_asset
+ if job.storage_key != asset.storage_key:
+ metadata["original_storage_key"] = asset.storage_key
+ metadata["hevc_transcoded"] = True
+ asset.storage_key = job.storage_key
asset.mime_type = mime_type
metadata["source"] = "upload"
asset.metadata = metadata
@@ -737,9 +829,28 @@ def ingest_asset(job_id: str) -> dict:
job_repo.update(job)
# 将上传时创建的占位 Asset(PROCESSING/UPLOADING)标记为 ERROR,
- # 避免素材永远卡在中间状态
+ # 避免素材永远卡在中间状态。转码可能已把 job.storage_key 改写为
+ # *_h264,需用 asset_id / 原始 key 多路径关联占位(Issue #1714)。
try:
- existing = asset_repo.find_by_storage_key(job.storage_key)
+ existing = None
+ _asset_id = getattr(job, "asset_id", "") or ""
+ if _asset_id:
+ try:
+ existing = asset_repo.find_by_id(_asset_id)
+ except Exception:
+ existing = None
+ if existing is None:
+ _candidate_keys = [
+ _original_key_from_storage_key(job.storage_key),
+ job.storage_key,
+ ]
+ for _key in _candidate_keys:
+ try:
+ existing = asset_repo.find_by_storage_key(_key)
+ except Exception:
+ existing = None
+ if existing is not None:
+ break
if existing and existing.status in (
AssetStatus.PROCESSING,
AssetStatus.UPLOADING,
diff --git a/deploy/configs/.env.production b/deploy/configs/.env.production
index 55ab74ff3..f14279411 100644
--- a/deploy/configs/.env.production
+++ b/deploy/configs/.env.production
@@ -214,3 +214,10 @@ MEDIAKIT_TIMEOUT=60
# ==================== 监控(可选)====================
# Sentry DSN(取消注释并填入实际值以启用错误追踪)
# SENTRY_DSN=${SENTRY_DSN}
+
+
+# ==================== 微信开放平台 OAuth(网页扫码登录)====================
+# 回调域名:xiaoxiajianji.com(微信开放平台已配置)
+WECHAT_OPEN_APP_ID=${WECHAT_APP_ID}
+WECHAT_OPEN_APP_SECRET=${WECHAT_APP_SECRET}
+WECHAT_OPEN_REDIRECT_URI=https://saas.xiaoxiajianji.com/auth/wechat/callback
diff --git a/deploy/configs/.env.staging b/deploy/configs/.env.staging
index 538417191..9e09b91eb 100644
--- a/deploy/configs/.env.staging
+++ b/deploy/configs/.env.staging
@@ -231,3 +231,10 @@ DASHSCOPE_API_KEY=${DASHSCOPE_API_KEY}
MEDIAKIT_API_KEY=${MEDIAKIT_API_KEY}
MEDIAKIT_BASE_URL=https://mediakit.cn-beijing.volces.com/api/v1
MEDIAKIT_TIMEOUT=60
+
+
+# ==================== 微信开放平台 OAuth(网页扫码登录)====================
+# 回调域名:xiaoxiajianji.com(微信开放平台已配置)
+WECHAT_OPEN_APP_ID=${WECHAT_APP_ID}
+WECHAT_OPEN_APP_SECRET=${WECHAT_APP_SECRET}
+WECHAT_OPEN_REDIRECT_URI=https://staging.xiaoxiajianji.com/auth/wechat/callback
diff --git a/deploy/configs/nginx-production.conf b/deploy/configs/nginx-production.conf
index 70b4b1a02..1944f8f1f 100644
--- a/deploy/configs/nginx-production.conf
+++ b/deploy/configs/nginx-production.conf
@@ -14,6 +14,10 @@ server {
# SPA routing - index.html 禁止缓存,确保每次获取最新版本
location / {
try_files $uri /index.html;
+ # HTML 文档(含 try_files 回退的 SPA 路由,如 /login /app/dashboard)一律 no-cache,
+ # 每次校验 ETag/Last-Modified,保证发版后旧标签页重新加载拿到新 chunk 引用;
+ # 带 hash 的静态资源由下方 ~* \.(js|css...) location 优先匹配,不受影响、保持 immutable
+ add_header Cache-Control "no-cache" always;
}
# API proxy — Production 环境代理到 production API 容器
diff --git a/deploy/configs/nginx-staging.conf b/deploy/configs/nginx-staging.conf
index cc6cc4ab9..9521dbb42 100644
--- a/deploy/configs/nginx-staging.conf
+++ b/deploy/configs/nginx-staging.conf
@@ -21,6 +21,10 @@ server {
# SPA fallback
location / {
try_files $uri /index.html;
+ # HTML 文档(含 try_files 回退的 SPA 路由,如 /login /app/dashboard)一律 no-cache,
+ # 每次校验 ETag/Last-Modified,保证发版后旧标签页重新加载拿到新 chunk 引用;
+ # 带 hash 的静态资源由下方 ~* \.(js|css...) location 优先匹配,不受影响、保持 immutable
+ add_header Cache-Control "no-cache" always;
}
# API proxy — Staging 环境代理到 staging API 容器
diff --git a/infra/docker/compose.yml b/infra/docker/compose.yml
index 91297553b..154e9fad1 100755
--- a/infra/docker/compose.yml
+++ b/infra/docker/compose.yml
@@ -115,6 +115,8 @@ services:
APP_VERSION: ${APP_VERSION:-unknown}
WORKER_CONCURRENCY: ${WORKER_CONCURRENCY:-4}
WORKER_MAX_TASKS_PER_CHILD: ${WORKER_MAX_TASKS_PER_CHILD:-100}
+ # #1714 队列隔离:generation 队列独占 worker(默认并发 2),其余并发给转码
+ GENERATION_CONCURRENCY: ${GENERATION_CONCURRENCY:-2}
GENERATED_FILES_DIR: /app/generated
GENERATED_FILES_URL_PREFIX: /generated-files
PUBLIC_API_BASE_URL: ${PUBLIC_API_BASE_URL:-https://api.xiaoxiajianji.com}
@@ -128,11 +130,11 @@ services:
# 健康检查配置
# 注:celery inspect ping 依赖 broker 连接,在容器内不可靠,改用进程检查
healthcheck:
- test: ["CMD-SHELL", "grep -q celery /proc/1/cmdline || exit 1"]
+ test: ["CMD-SHELL", "pgrep -f 'celery.*worker' | head -n1 >/dev/null 2>&1 || grep -q celery /proc/1/cmdline || exit 1"]
interval: 30s
timeout: 10s
retries: 3
- start_period: 30s
+ start_period: 40s
logging: *default-logging
@@ -140,7 +142,8 @@ services:
# 资源限制建议(生产环境建议启用)
# =========================================
# 注意: Worker 需要处理视频,建议分配更多资源
- # 并发 4 时需要 4C8G 以上,确保视频渲染不 OOM
+ # #1714 队列隔离后容器内运行 generation + transcode 两个 worker 进程,
+ # 总并发 = WORKER_CONCURRENCY(默认 4),4C8G 以上确保视频渲染不 OOM
deploy:
resources:
limits:
diff --git a/infra/docker/deploy-production-registry.sh b/infra/docker/deploy-production-registry.sh
index f06faec2a..664e71f6e 100755
--- a/infra/docker/deploy-production-registry.sh
+++ b/infra/docker/deploy-production-registry.sh
@@ -146,6 +146,7 @@ docker run -d \
-e APP_ENV=production \
-e APP_VERSION="$IMAGE_TAG" \
-e WORKER_CONCURRENCY="${WORKER_CONCURRENCY:-4}" \
+ -e GENERATION_CONCURRENCY="${GENERATION_CONCURRENCY:-2}" \
-e WORKER_MAX_TASKS_PER_CHILD=100 \
-e GENERATED_FILES_DIR=/app/generated \
-e GENERATED_FILES_URL_PREFIX=/generated-files \
@@ -154,7 +155,7 @@ docker run -d \
--restart unless-stopped \
--cpus 2 \
--memory 2g \
- --health-cmd "sh -c \"grep -q celery /proc/1/cmdline || exit 1\"" \
+ --health-cmd "sh -c \"pgrep -f 'celery.*worker' >/dev/null 2>&1 || grep -q celery /proc/1/cmdline || exit 1\"" \
--health-interval 30s \
--health-timeout 10s \
--health-retries 3 \
diff --git a/infra/docker/deploy-staging-registry.sh b/infra/docker/deploy-staging-registry.sh
index 9d3f5c729..b444b9ad9 100755
--- a/infra/docker/deploy-staging-registry.sh
+++ b/infra/docker/deploy-staging-registry.sh
@@ -109,13 +109,14 @@ docker run -d \
-e APP_ENV=staging \
-e APP_VERSION="$IMAGE_TAG" \
-e WORKER_CONCURRENCY="${WORKER_CONCURRENCY:-4}" \
+ -e GENERATION_CONCURRENCY="${GENERATION_CONCURRENCY:-2}" \
-e WORKER_MAX_TASKS_PER_CHILD=100 \
-e GENERATED_FILES_DIR=/app/generated \
-e GENERATED_FILES_URL_PREFIX=/generated-files \
-v "$GENERATED_DIR:/app/generated" \
--restart unless-stopped \
--label com.centurylinklabs.watchtower.enable=true \
- --health-cmd "sh -c \"grep -q celery /proc/1/cmdline || exit 1\"" \
+ --health-cmd "sh -c \"pgrep -f 'celery.*worker' >/dev/null 2>&1 || grep -q celery /proc/1/cmdline || exit 1\"" \
--health-interval 30s \
--health-timeout 10s \
--health-retries 3 \
diff --git a/infra/docker/entrypoint-worker.sh b/infra/docker/entrypoint-worker.sh
index f2e208d96..0a74b474a 100755
--- a/infra/docker/entrypoint-worker.sh
+++ b/infra/docker/entrypoint-worker.sh
@@ -1,18 +1,64 @@
#!/bin/bash
-# Worker 启动脚本 — 支持 WORKER_CONCURRENCY 环境变量
-# 未设置时默认 2(保持向后兼容)
+# Worker 启动脚本 — #1714 队列隔离
+#
+# 部署约束:worker 容器单实例(replicas=1),容器内启动两个 celery 进程:
+# 1. generation-worker:独占消费 generation 队列(用户视频生成,高优先级),
+# 内嵌 celery beat(-B),定时清理任务只在一个进程里跑,避免重复执行;
+# 2. transcode-worker:消费 transcode + celery 默认队列(素材转码/分类/查重/
+# 配音/下载等后台任务)。
+# 转码队列积压时,generation 队列仍有独立 worker 立即领取视频生成任务。
+#
+# 环境变量:
+# WORKER_CONCURRENCY 总并发槽参考(默认 4);生成 worker 并发默认 2,
+# 可用 GENERATION_CONCURRENCY 覆盖
+# GENERATION_CONCURRENCY generation worker 并发(默认 2)
+# TRANSCODE_CONCURRENCY transcode worker 并发(默认 = WORKER_CONCURRENCY - 2,最小 1)
+# WORKER_MAX_TASKS_PER_CHILD 每个子进程最大任务数(默认 100)
set -e
-CONCURRENCY="${WORKER_CONCURRENCY:-2}"
+CONCURRENCY="${WORKER_CONCURRENCY:-4}"
+MAX_TASKS="${WORKER_MAX_TASKS_PER_CHILD:-100}"
-# ⚠️ 部署约束:此 Worker 必须且只能运行单实例(replicas=1)
-# -B 标志嵌入 celery beat,beat 负责定期触发 pending 超时清理等定时任务
-# 多实例部署会导致每个 Worker 独立运行 Beat,造成定时任务重复执行
-# 若需横向扩展 Worker,必须将 Beat 拆分为独立服务(celery beat -A worker_app.celery_app)
-exec celery \
+GEN_CONCURRENCY="${GENERATION_CONCURRENCY:-2}"
+if [ -z "$TRANSCODE_CONCURRENCY" ]; then
+ TRANS_CONCURRENCY=$((CONCURRENCY - GEN_CONCURRENCY))
+ if [ "$TRANS_CONCURRENCY" -lt 1 ]; then
+ TRANS_CONCURRENCY=1
+ fi
+else
+ TRANS_CONCURRENCY="$TRANSCODE_CONCURRENCY"
+fi
+
+echo "Starting generation worker (queue=generation, concurrency=$GEN_CONCURRENCY, beat embedded)"
+celery \
-A worker_app.celery_app \
worker \
--loglevel=info \
"-B" \
- "--concurrency=${CONCURRENCY}"
+ -s /tmp/celerybeat-schedule \
+ -Q generation \
+ "--concurrency=${GEN_CONCURRENCY}" \
+ "--max-tasks-per-child=${MAX_TASKS}" \
+ -n generation@%h &
+GEN_PID=$!
+
+echo "Starting transcode worker (queues=transcode,celery, concurrency=$TRANS_CONCURRENCY)"
+celery \
+ -A worker_app.celery_app \
+ worker \
+ --loglevel=info \
+ -Q transcode,celery \
+ "--concurrency=${TRANS_CONCURRENCY}" \
+ "--max-tasks-per-child=${MAX_TASKS}" \
+ -n transcode@%h &
+TRANS_PID=$!
+
+# 任一进程退出则终止另一个,让容器整体重启(restart: unless-stopped)
+trap 'echo "Shutting down workers..."; kill -TERM $GEN_PID $TRANS_PID 2>/dev/null || true' TERM INT
+
+wait -n $GEN_PID $TRANS_PID
+EXIT_CODE=$?
+echo "One worker exited (code=$EXIT_CODE), stopping the other..."
+kill -TERM $GEN_PID $TRANS_PID 2>/dev/null || true
+exit $EXIT_CODE
diff --git a/infra/docker/nginx-production.conf b/infra/docker/nginx-production.conf
index c80cfa7b5..c5269f003 100755
--- a/infra/docker/nginx-production.conf
+++ b/infra/docker/nginx-production.conf
@@ -16,6 +16,10 @@ server {
# 注意:不能加 $uri/,否则 /assets 等与构建产物目录同名的路由会被当成目录访问,返回 403
location / {
try_files $uri /index.html;
+ # HTML 文档(含 try_files 回退的 SPA 路由,如 /login /app/dashboard)一律 no-cache,
+ # 每次校验 ETag/Last-Modified,保证发版后旧标签页重新加载拿到新 chunk 引用;
+ # 带 hash 的静态资源由下方 ~* \.(js|css...) location 优先匹配,不受影响、保持 immutable
+ add_header Cache-Control "no-cache" always;
}
# API proxy
diff --git a/infra/docker/nginx-staging.conf b/infra/docker/nginx-staging.conf
index d92cdb789..f6ec40cda 100755
--- a/infra/docker/nginx-staging.conf
+++ b/infra/docker/nginx-staging.conf
@@ -23,6 +23,10 @@ server {
location / {
try_files $uri /index.html;
+ # HTML 文档(含 try_files 回退的 SPA 路由,如 /login /app/dashboard)一律 no-cache,
+ # 每次校验 ETag/Last-Modified,保证发版后旧标签页重新加载拿到新 chunk 引用;
+ # 带 hash 的静态资源由下方 ~* \.(js|css...) location 优先匹配,不受影响、保持 immutable
+ add_header Cache-Control "no-cache" always;
}
# API proxy
diff --git a/infra/docker/nginx.conf b/infra/docker/nginx.conf
index a5b581f7b..bd8a1fe68 100755
--- a/infra/docker/nginx.conf
+++ b/infra/docker/nginx.conf
@@ -33,6 +33,10 @@ server {
location / {
try_files $uri /index.html;
+ # HTML 文档(含 try_files 回退的 SPA 路由,如 /login /app/dashboard)一律 no-cache,
+ # 每次校验 ETag/Last-Modified,保证发版后旧标签页重新加载拿到新 chunk 引用;
+ # 带 hash 的静态资源由下方 ~* \.(js|css...) location 优先匹配,不受影响、保持 immutable
+ add_header Cache-Control "no-cache" always;
}
# API proxy
diff --git a/packages/adapters/in_memory/asset_repository.py b/packages/adapters/in_memory/asset_repository.py
index cc5bd76b2..3e6486db8 100755
--- a/packages/adapters/in_memory/asset_repository.py
+++ b/packages/adapters/in_memory/asset_repository.py
@@ -146,3 +146,44 @@ class InMemoryAssetRepository:
if asset.library_id == library_id and asset.file_hash == file_hash:
return asset
return None
+
+ def find_by_library_and_client_upload_id(
+ self,
+ library_id: str,
+ client_upload_id: str,
+ ) -> Asset | None:
+ """按素材库 + 客户端幂等 token 查找已有素材。"""
+ if not client_upload_id:
+ return None
+ for asset in self._assets.values():
+ if asset.library_id == library_id and getattr(asset, "client_upload_id", "") == client_upload_id:
+ return asset
+ return None
+
+ def find_recent_active_by_library_and_name(
+ self,
+ library_id: str,
+ name: str,
+ within_minutes: int = 30,
+ file_size: int = 0,
+ ) -> Asset | None:
+ """兜底去重:同库 + 同文件名(+同大小)且近期活动状态的素材。"""
+ from datetime import datetime, timedelta, timezone
+
+ if not name:
+ return None
+ from packages.domain import AssetStatus
+
+ cutoff = datetime.now(timezone.utc) - timedelta(minutes=within_minutes)
+ candidates = [
+ a
+ for a in self._assets.values()
+ if a.library_id == library_id
+ and a.name == name
+ and a.status in (AssetStatus.UPLOADING, AssetStatus.PROCESSING)
+ and a.created_at >= cutoff
+ and (not file_size or file_size <= 0 or a.file_size == file_size)
+ ]
+ if not candidates:
+ return None
+ return max(candidates, key=lambda a: a.created_at)
diff --git a/packages/adapters/sqlalchemy_impl/asset_repository.py b/packages/adapters/sqlalchemy_impl/asset_repository.py
index 146ab706d..fad219727 100755
--- a/packages/adapters/sqlalchemy_impl/asset_repository.py
+++ b/packages/adapters/sqlalchemy_impl/asset_repository.py
@@ -134,6 +134,7 @@ class SQLAlchemyAssetRepository:
quality_score=asset.quality_score,
uploaded_by_user_id=asset.uploaded_by_user_id or "system",
file_hash=asset.file_hash or None,
+ client_upload_id=asset.client_upload_id or None,
created_at=asset.created_at,
updated_at=now,
)
@@ -163,6 +164,8 @@ class SQLAlchemyAssetRepository:
model.quality_score = asset.quality_score
model.uploaded_by_user_id = asset.uploaded_by_user_id or model.uploaded_by_user_id
model.file_hash = asset.file_hash or model.file_hash
+ if getattr(model, "client_upload_id", None) is None and asset.client_upload_id:
+ model.client_upload_id = asset.client_upload_id
model.updated_at = datetime.now(timezone.utc)
self.session.flush()
self._sync_asset_tags(asset.id, asset.tag_ids)
@@ -388,6 +391,7 @@ class SQLAlchemyAssetRepository:
quality_score=model.quality_score,
uploaded_by_user_id=model.uploaded_by_user_id,
file_hash=model.file_hash or "",
+ client_upload_id=getattr(model, "client_upload_id", None) or "",
metadata=metadata,
tag_ids=tag_ids,
created_at=model.created_at,
@@ -452,3 +456,58 @@ class SQLAlchemyAssetRepository:
if model is None:
return None
return self._to_domain(model)
+
+ def find_by_library_and_client_upload_id(
+ self,
+ library_id: str,
+ client_upload_id: str,
+ ) -> Asset | None:
+ """按素材库 + 客户端幂等 token 查找已有素材(complete 幂等)。"""
+ if not client_upload_id:
+ return None
+ model = (
+ self.session.query(AssetModel)
+ .filter(
+ AssetModel.asset_library_id == library_id,
+ AssetModel.client_upload_id == client_upload_id,
+ )
+ .first()
+ )
+ if model is None:
+ return None
+ return self._to_domain(model)
+
+ def find_recent_active_by_library_and_name(
+ self,
+ library_id: str,
+ name: str,
+ within_minutes: int = 30,
+ file_size: int = 0,
+ ) -> Asset | None:
+ """兜底去重:同库 + 同文件名(+同大小)且近期仍处活动状态(uploading/processing)的素材。
+
+ 用于旧客户端未传 file_hash/client_upload_id 时,防止 complete 超时重试
+ 反复创建 PROCESSING 占位记录。只命中"活动中"的近期记录,READY 历史素材不拦。
+
+ 严格模式(#1714 误杀修复):file_size 必须 > 0 且与记录大小严格一致;
+ file_size=0(大小未知)时直接返回 None——宁可漏判(极端情况下多建一条
+ 占位)也不可仅凭同名 + processing 误杀内容全新的视频。
+ """
+ from datetime import datetime, timedelta, timezone
+
+ if not name:
+ return None
+ if not file_size or file_size <= 0:
+ return None
+ cutoff = datetime.now(timezone.utc) - timedelta(minutes=within_minutes)
+ query = self.session.query(AssetModel).filter(
+ AssetModel.asset_library_id == library_id,
+ AssetModel.name == name,
+ AssetModel.status.in_([AssetStatus.UPLOADING.value, AssetStatus.PROCESSING.value]),
+ AssetModel.created_at >= cutoff,
+ AssetModel.file_size == file_size,
+ )
+ model = query.order_by(AssetModel.created_at.desc()).first()
+ if model is None:
+ return None
+ return self._to_domain(model)
diff --git a/packages/adapters/sqlalchemy_impl/generation_task_repository.py b/packages/adapters/sqlalchemy_impl/generation_task_repository.py
index 38aa44f9d..893e6a582 100755
--- a/packages/adapters/sqlalchemy_impl/generation_task_repository.py
+++ b/packages/adapters/sqlalchemy_impl/generation_task_repository.py
@@ -38,6 +38,7 @@ def _to_domain(model: GenerationTaskModel) -> GenerationTask:
bgm_config=dict(getattr(model, "bgm_config", {}) or {}),
is_preview=bool(getattr(model, "is_preview", False)),
source_task_id=getattr(model, "source_task_id", "") or "",
+ celery_task_id=getattr(model, "celery_task_id", "") or "",
output_width=getattr(model, "output_width", 1280) or 1280,
output_height=getattr(model, "output_height", 720) or 720,
cover_url=getattr(model, "cover_url", "") or "",
@@ -82,6 +83,7 @@ class SQLAlchemyGenerationTaskRepository:
bgm_config=task.bgm_config or {},
is_preview=task.is_preview or False,
source_task_id=task.source_task_id or "",
+ celery_task_id=getattr(task, "celery_task_id", "") or "",
output_width=task.output_width,
output_height=task.output_height,
cover_url=task.cover_url or "",
@@ -138,6 +140,52 @@ class SQLAlchemyGenerationTaskRepository:
.count()
)
+ def count_running_by_user(self, user_id: str) -> int:
+ """统计指定用户处于 running 状态的任务数(用于限流提示展示)。"""
+ return (
+ self.session.query(GenerationTaskModel)
+ .filter(
+ GenerationTaskModel.created_by_user_id == user_id,
+ GenerationTaskModel.status == GenerationTaskStatus.RUNNING.value,
+ )
+ .count()
+ )
+
+ def count_running_total(self) -> int:
+ """统计全局处于 running 状态的任务数(worker 实际在执行的任务数)。"""
+ return (
+ self.session.query(GenerationTaskModel)
+ .filter(GenerationTaskModel.status == GenerationTaskStatus.RUNNING.value)
+ .count()
+ )
+
+ def estimate_avg_duration_seconds(self, limit: int = 20, default_seconds: float = 120.0) -> float:
+ """估算最近完成任务的平均耗时(秒),用于 429 限流提示的等待预估。
+
+ 取最近 N 条 completed 任务的 (completed_at - started_at) 平均值;
+ 无足够历史数据时返回 default_seconds。
+ 用 Python 侧计算差值,避免 SQLite/PostgreSQL 方言差异。
+ """
+ rows = (
+ self.session.query(GenerationTaskModel.started_at, GenerationTaskModel.completed_at)
+ .filter(
+ GenerationTaskModel.status == GenerationTaskStatus.COMPLETED.value,
+ GenerationTaskModel.started_at.isnot(None),
+ GenerationTaskModel.completed_at.isnot(None),
+ )
+ .order_by(GenerationTaskModel.completed_at.desc())
+ .limit(limit)
+ .all()
+ )
+ durations = [
+ (completed - started).total_seconds()
+ for started, completed in rows
+ if completed and started and (completed - started).total_seconds() > 0
+ ]
+ if not durations:
+ return default_seconds
+ return sum(durations) / len(durations)
+
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]:
models = (
self.session.query(GenerationTaskModel)
@@ -269,6 +317,7 @@ class SQLAlchemyGenerationTaskRepository:
if hasattr(model, "is_preview"):
model.is_preview = task.is_preview or False
model.source_task_id = task.source_task_id or ""
+ model.celery_task_id = getattr(task, "celery_task_id", "") or model.celery_task_id or ""
model.output_width = task.output_width
model.output_height = task.output_height
model.cover_url = task.cover_url or ""
@@ -280,12 +329,14 @@ class SQLAlchemyGenerationTaskRepository:
def cleanup_stale_running(self, timeout_minutes: int = 10) -> int:
"""清理超时未更新的 running 任务(孤儿任务)。
- 将 status=running 且 updated_at 超过 timeout_minutes 分钟未更新的任务
- 标记为 failed,error_message 标记为任务执行中断。
-
Returns:
- 清理的任务数量
+ 清理的任务数量(仅计数,保持旧签名兼容)
"""
+ items = self.cleanup_stale_running_with_ids(timeout_minutes)
+ return len(items)
+
+ def cleanup_stale_running_with_ids(self, timeout_minutes: int = 10) -> list[tuple[str, str]]:
+ """同 cleanup_stale_running,但返回 [(task_id, celery_task_id), ...] 供撤销队列消息。"""
from datetime import timedelta
cutoff = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes)
@@ -298,8 +349,10 @@ class SQLAlchemyGenerationTaskRepository:
.all()
)
if not models:
- return 0
+ return []
+ result: list[tuple[str, str]] = []
for model in models:
+ result.append((model.id, getattr(model, "celery_task_id", "") or ""))
model.status = GenerationTaskStatus.FAILED.value
model.error_message = "任务执行中断(worker重启/超时)"
model.error_info = {
@@ -309,43 +362,43 @@ class SQLAlchemyGenerationTaskRepository:
}
model.completed_at = datetime.now(timezone.utc)
self.session.commit()
- return len(models)
+ return result
def cleanup_stale_pending(self, timeout_minutes: int = 30) -> int:
"""清理超时的 pending 任务(未被 Worker 拉取的任务)。
- 全局任务队列有 pending 数量上限,长期卡在 pending 的任务会占满队列,
- 导致新用户无法创建任务。将超时的 pending 任务标记为 failed。
-
- Args:
- timeout_minutes: 超时时间(分钟),默认 30 分钟
-
Returns:
- 清理的任务数量
+ 清理的任务数量(仅计数,保持旧签名兼容)
"""
+ items = self.cleanup_stale_pending_with_ids(timeout_minutes)
+ return len(items)
+
+ def cleanup_stale_pending_with_ids(self, timeout_minutes: int = 30) -> list[tuple[str, str]]:
+ """同 cleanup_stale_pending,但返回 [(task_id, celery_task_id), ...] 供撤销队列消息。"""
from datetime import timedelta
cutoff = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes)
- error_info = {
- "error_type": "PendingTimeout",
- "message": f"任务在 pending 状态停留超过 {timeout_minutes} 分钟,自动清理",
- "failed_at": datetime.now(timezone.utc).isoformat(),
- }
- count = (
+ models = (
self.session.query(GenerationTaskModel)
.filter(
GenerationTaskModel.status == GenerationTaskStatus.PENDING.value,
GenerationTaskModel.created_at < cutoff,
)
- .update(
- {
- GenerationTaskModel.status: GenerationTaskStatus.FAILED.value,
- GenerationTaskModel.error_message: "pending timeout: auto cleanup",
- GenerationTaskModel.error_info: error_info,
- GenerationTaskModel.completed_at: datetime.now(timezone.utc),
- },
- synchronize_session=False,
- )
+ .all()
)
+ if not models:
+ return []
+ error_info = {
+ "error_type": "PendingTimeout",
+ "message": f"任务在 pending 状态停留超过 {timeout_minutes} 分钟,自动清理",
+ "failed_at": datetime.now(timezone.utc).isoformat(),
+ }
+ result: list[tuple[str, str]] = []
+ for model in models:
+ result.append((model.id, getattr(model, "celery_task_id", "") or ""))
+ model.status = GenerationTaskStatus.FAILED.value
+ model.error_message = "pending timeout: auto cleanup"
+ model.error_info = error_info
+ model.completed_at = datetime.now(timezone.utc)
self.session.commit()
- return count
+ return result
diff --git a/packages/adapters/sqlalchemy_impl/ingest_job_repository.py b/packages/adapters/sqlalchemy_impl/ingest_job_repository.py
index f16dc735d..7c42fdeee 100644
--- a/packages/adapters/sqlalchemy_impl/ingest_job_repository.py
+++ b/packages/adapters/sqlalchemy_impl/ingest_job_repository.py
@@ -18,6 +18,8 @@ class SQLAlchemyIngestJobRepository:
error_message=job.error_message,
result_asset_id=job.result_asset_id,
file_hash=job.file_hash,
+ asset_id=job.asset_id or "",
+ celery_task_id=getattr(job, "celery_task_id", "") or "",
created_at=job.created_at,
updated_at=job.updated_at,
)
@@ -38,6 +40,8 @@ class SQLAlchemyIngestJobRepository:
error_message=model.error_message,
result_asset_id=model.result_asset_id,
file_hash=model.file_hash or "",
+ asset_id=getattr(model, "asset_id", "") or "",
+ celery_task_id=getattr(model, "celery_task_id", "") or "",
created_at=model.created_at,
updated_at=model.updated_at,
)
@@ -54,6 +58,12 @@ class SQLAlchemyIngestJobRepository:
model.error_message = job.error_message
model.result_asset_id = job.result_asset_id
model.file_hash = job.file_hash
+ model.storage_key = job.storage_key
+ if job.asset_id:
+ model.asset_id = job.asset_id
+ celery_tid = getattr(job, "celery_task_id", "")
+ if celery_tid:
+ model.celery_task_id = celery_tid
model.updated_at = job.updated_at
self.session.commit()
return job
diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py
index 8e71fcb32..04ff2381d 100755
--- a/packages/adapters/sqlalchemy_impl/models.py
+++ b/packages/adapters/sqlalchemy_impl/models.py
@@ -38,6 +38,7 @@ class UserModel(Base):
phone = Column(String(32), nullable=True, unique=True, index=True)
phone_verified = Column(Boolean, nullable=False, default=False)
binding_completed_at = Column(DateTime, nullable=True)
+ profile_completed = Column(Boolean, nullable=False, default=True, server_default="true")
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
@@ -94,6 +95,7 @@ class AssetModel(Base):
quality_score = Column(Float, nullable=True)
uploaded_by_user_id = Column(String(36), nullable=False)
file_hash = Column(String(64), nullable=True, index=True)
+ client_upload_id = Column(String(64), nullable=True, index=True)
extra_meta = Column("metadata", JSON, nullable=False, default=dict)
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc), index=True)
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
@@ -244,6 +246,8 @@ class IngestJobModel(Base):
error_message = Column(Text, nullable=False, default="")
result_asset_id = Column(String(36), nullable=False, default="")
file_hash = Column(String(64), nullable=True, index=True)
+ asset_id = Column(String(36), nullable=False, default="", index=True)
+ celery_task_id = Column(String(64), nullable=False, default="", server_default="")
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
@@ -296,6 +300,7 @@ class GenerationTaskModel(Base):
resolution = Column(String(20), nullable=False, default="")
is_preview = Column(Boolean, nullable=False, default=False, index=True)
source_task_id = Column(String(32), nullable=False, default="", index=True)
+ celery_task_id = Column(String(64), nullable=False, default="", server_default="")
output_width = Column(Integer, nullable=False, default=1280)
output_height = Column(Integer, nullable=False, default=720)
cover_url = Column(String(1000), nullable=False, default="")
diff --git a/packages/adapters/sqlalchemy_impl/user_repository.py b/packages/adapters/sqlalchemy_impl/user_repository.py
index ac9631aa9..4bafb2b60 100755
--- a/packages/adapters/sqlalchemy_impl/user_repository.py
+++ b/packages/adapters/sqlalchemy_impl/user_repository.py
@@ -38,6 +38,7 @@ class SQLAlchemyUserRepository(UserRepository):
model.phone = user.phone
model.phone_verified = user.phone_verified
model.binding_completed_at = user.binding_completed_at
+ model.profile_completed = user.profile_completed
model.created_at = user.created_at
self.session.commit()
@@ -113,5 +114,6 @@ class SQLAlchemyUserRepository(UserRepository):
phone=model.phone,
phone_verified=model.phone_verified or False,
binding_completed_at=model.binding_completed_at,
+ profile_completed=model.profile_completed if model.profile_completed is not None else True,
created_at=model.created_at,
)
diff --git a/packages/application/auth/wechat_bind_use_case.py b/packages/application/auth/wechat_bind_use_case.py
new file mode 100644
index 000000000..4b9cce0bb
--- /dev/null
+++ b/packages/application/auth/wechat_bind_use_case.py
@@ -0,0 +1,115 @@
+"""
+微信账号绑定/解绑 Use Case(已登录用户场景)
+
+与 wechat_sync_use_case(登录/注册,系统级)不同:
+- bind:把微信 openid/unionid 绑定到【当前登录账号】,不创建新用户;
+ 微信身份若已绑定其他账号则冲突(409)。
+- unbind:解除当前账号的微信绑定;若账号没有其他登录方式(手机/邮箱/密码),
+ 解绑后将无法登录,因此拒绝解绑。
+"""
+
+from __future__ import annotations
+
+from typing import Optional
+
+from packages.domain.entities import User
+
+
+class WechatBindRequest:
+ """微信绑定请求"""
+
+ def __init__(self, user_id: str, openid: str, unionid: str = ""):
+ self.user_id = user_id
+ self.openid = (openid or "").strip()
+ self.unionid = (unionid or "").strip()
+
+
+class WechatBindResult:
+ """微信绑定/解绑结果"""
+
+ def __init__(self, user: User):
+ self.user = user
+
+
+class WechatBindUseCase:
+ """已登录用户绑定微信用例"""
+
+ def __init__(self, user_repository):
+ self.user_repository = user_repository
+
+ def bind(self, request: WechatBindRequest) -> tuple[Optional[WechatBindResult], Optional[str], int]:
+ """
+ 绑定微信到当前登录账号。
+
+ Returns:
+ (结果, 错误信息, http状态码) - 成功时错误信息为 None、状态码为 200;
+ 冲突返回 409,客户端/服务端错误返回 400/404。
+ """
+ if not request.openid:
+ return None, "缺少微信 openid", 400
+
+ user = self.user_repository.find_by_id(request.user_id)
+ if user is None:
+ return None, "当前用户不存在", 404
+
+ # 已绑定同一个微信:幂等成功
+ if user.wechat_openid == request.openid:
+ return WechatBindResult(user=user), None, 200
+
+ # 当前账号已绑定其他微信
+ if user.wechat_openid:
+ return None, "当前账号已绑定微信,请先解绑", 409
+
+ # openid 已被其他账号占用
+ existing = self.user_repository.find_by_wechat_openid(request.openid)
+ if existing is not None and existing.id != user.id:
+ return None, "该微信已绑定其他账号,请先在原账号解绑", 409
+
+ # unionid 冲突:同主体微信已绑其他账号
+ if request.unionid:
+ existing_union = self.user_repository.find_by_wechat_unionid(request.unionid)
+ if existing_union is not None and existing_union.id != user.id:
+ return None, "该微信主体已绑定其他账号,请先在原账号解绑", 409
+
+ user.wechat_openid = request.openid
+ if request.unionid and not user.wechat_unionid:
+ user.wechat_unionid = request.unionid
+ self.user_repository.save(user)
+
+ return WechatBindResult(user=user), None, 200
+
+
+class WechatUnbindUseCase:
+ """已登录用户解绑微信用例"""
+
+ def __init__(self, user_repository):
+ self.user_repository = user_repository
+
+ def unbind(self, user_id: str) -> tuple[Optional[WechatBindResult], Optional[str], int]:
+ """
+ 解除当前账号的微信绑定。
+
+ 解绑前置条件:账号必须还有其他登录方式(密码 / 已验证手机 / 真实邮箱),
+ 否则解绑后将永远无法登录。
+ """
+ user = self.user_repository.find_by_id(user_id)
+ if user is None:
+ return None, "当前用户不存在", 404
+
+ if not user.wechat_openid:
+ return None, "当前账号未绑定微信", 400
+
+ # 守卫:解绑后账号必须仍有可实际使用的登录方式。
+ # 注意:微信注册用户带的是【随机密码】(用户不知道、无法用密码登录,
+ # 且 @wechat.local 占位邮箱收不到重置邮件),故 password_hash 不作为兜底依据,
+ # 口径与 /auth/me 的 binding_complete 一致。
+ has_phone = bool(user.phone and user.phone_verified)
+ has_real_email = bool(user.email and user.email_verified and "@wechat.local" not in user.email)
+ if not (has_phone or has_real_email):
+ return None, "账号需要至少一种其他登录方式(已验证手机或真实邮箱)后才能解绑微信", 400
+
+ user.wechat_openid = None
+ user.wechat_unionid = None
+ self.user_repository.save(user)
+
+ return WechatBindResult(user=user), None, 200
diff --git a/packages/application/auth/wechat_oauth_service.py b/packages/application/auth/wechat_oauth_service.py
index 16cdf5d47..08e4416ce 100755
--- a/packages/application/auth/wechat_oauth_service.py
+++ b/packages/application/auth/wechat_oauth_service.py
@@ -20,6 +20,7 @@ import requests
logger = logging.getLogger(__name__)
STATE_TTL_SECONDS = 600 # state 有效期 10 分钟
+STATE_KEY_PREFIX = "wechat:state:" # Redis key 前缀(独立逻辑命名空间)
class MemoryStateStore:
@@ -53,6 +54,80 @@ class MemoryStateStore:
del self._states[s]
+class RedisStateStore:
+ """Redis state 存储(多实例/容器重启安全)。
+
+ 复用现有 Redis(celery broker 同实例),key 前缀 wechat:state:,
+ TTL 10 分钟,SET NX EX + GETDEL 保证一次性消费。
+ Redis 不可用时降级为内存存储,保证登录流程不中断(单节点场景)。
+ """
+
+ def __init__(
+ self,
+ redis_url: str = "",
+ ttl_seconds: int = STATE_TTL_SECONDS,
+ key_prefix: str = STATE_KEY_PREFIX,
+ client=None,
+ ):
+ self._ttl = ttl_seconds
+ self._prefix = key_prefix
+ self._fallback = MemoryStateStore(ttl_seconds=ttl_seconds)
+ self._redis = None
+ if client is not None:
+ # 测试/显式注入
+ self._redis = client
+ return
+ try:
+ import redis
+
+ self._redis = redis.Redis.from_url(redis_url, decode_responses=True)
+ self._redis.ping()
+ logger.info(
+ "微信 state 存储使用 Redis: %s db=%s",
+ self._redis.connection_pool.connection_kwargs.get("host"),
+ self._redis.connection_pool.connection_kwargs.get("db"),
+ )
+ except Exception as e: # noqa: BLE001 — Redis 不可用降级内存,登录流程不中断
+ logger.warning("微信 state Redis 不可用,降级为内存存储: %s", e)
+ self._redis = None
+
+ def _key(self, state: str) -> str:
+ return f"{self._prefix}{state}"
+
+ def put(self, state: str) -> None:
+ if self._redis is None:
+ self._fallback.put(state)
+ return
+ try:
+ # SET key 1 NX EX ttl:不存在才写入,自带过期
+ self._redis.set(self._key(state), "1", nx=True, ex=self._ttl)
+ except Exception as e: # noqa: BLE001
+ logger.warning("微信 state 写入 Redis 失败,降级内存: %s", e)
+ self._fallback.put(state)
+
+ # Lua:原子读取并删除(单线程执行),兼容所有 Redis 版本(GETDEL 需 6.2+)
+ _CONSUME_LUA = """
+local v = redis.call('GET', KEYS[1])
+if v then redis.call('DEL', KEYS[1]) end
+return v
+"""
+
+ def verify_and_consume(self, state: str) -> bool:
+ if self._redis is None:
+ return self._fallback.verify_and_consume(state)
+ try:
+ try:
+ val = self._redis.eval(self._CONSUME_LUA, 1, self._key(state))
+ except Exception: # noqa: BLE001 — eval 不可用时退化 GET+DELETE
+ val = self._redis.get(self._key(state))
+ if val is not None:
+ self._redis.delete(self._key(state))
+ return val is not None
+ except Exception as e: # noqa: BLE001
+ logger.warning("微信 state 校验 Redis 失败,降级内存: %s", e)
+ return self._fallback.verify_and_consume(state)
+
+
@dataclass
class WechatUserInfo:
"""微信用户信息"""
@@ -158,6 +233,8 @@ class WechatOAuthService:
"grant_type": "authorization_code",
}
token_resp = requests.get(token_url, params=token_params, timeout=10)
+ # 微信响应头不带 charset,requests 默认按 ISO-8859-1 解码会导致中文乱码
+ token_resp.encoding = "utf-8"
token_data = token_resp.json()
if "errcode" in token_data and token_data["errcode"] != 0:
@@ -176,6 +253,8 @@ class WechatOAuthService:
"lang": "zh_CN",
}
user_resp = requests.get(user_url, params=user_params, timeout=10)
+ # 同上:显式 UTF-8 解码,保证中文昵称/unionid 等不乱码
+ user_resp.encoding = "utf-8"
user_data = user_resp.json()
if "errcode" in user_data and user_data["errcode"] != 0:
@@ -200,7 +279,29 @@ class WechatOAuthService:
return None, "微信登录处理失败"
+# 模块级单例:state 存储必须跨请求共享,否则 /wechat/url 生成的 state
+# 与 /wechat/callback 校验时不在同一个 MemoryStateStore,回调必然 400。
+# 多实例部署时应替换为 Redis state store(单容器多 worker 也需如此)。
+_oauth_service_singleton: WechatOAuthService | None = None
+
+
+def _build_default_state_store():
+ """默认 state 存储:优先 Redis(多实例/重启安全),不可用由 store 内部降级内存。"""
+ redis_url = ""
+ try:
+ from app.config import get_settings
+
+ redis_url = get_settings().CELERY_BROKER_URL or get_settings().REDIS_URL
+ except Exception: # noqa: BLE001 — API 配置不可用时退回环境变量
+ redis_url = os.environ.get("CELERY_BROKER_URL", "") or os.environ.get("REDIS_URL", "")
+ if redis_url:
+ return RedisStateStore(redis_url)
+ return MemoryStateStore()
+
+
def get_wechat_oauth_service() -> WechatOAuthService:
- """获取微信 OAuth 服务单例"""
- # TODO: 可替换为 Redis state store
- return WechatOAuthService()
+ """获取微信 OAuth 服务单例(state store 跨请求共享)"""
+ global _oauth_service_singleton
+ if _oauth_service_singleton is None:
+ _oauth_service_singleton = WechatOAuthService(state_store=_build_default_state_store())
+ return _oauth_service_singleton
diff --git a/packages/application/auth/wechat_sync_use_case.py b/packages/application/auth/wechat_sync_use_case.py
index 7f343eb7b..bb96dc1d5 100644
--- a/packages/application/auth/wechat_sync_use_case.py
+++ b/packages/application/auth/wechat_sync_use_case.py
@@ -200,6 +200,8 @@ class WechatSyncUseCase:
email_verified=True, # 微信登录视为已验证
wechat_openid=request.openid,
wechat_unionid=request.unionid or None,
+ # 微信新建用户首次登录需引导设置昵称
+ profile_completed=False,
)
self.user_repository.save(user)
diff --git a/packages/application/ingest_jobs.py b/packages/application/ingest_jobs.py
index 75a708de7..92f062ac2 100644
--- a/packages/application/ingest_jobs.py
+++ b/packages/application/ingest_jobs.py
@@ -12,6 +12,8 @@ class SubmitIngestJobCommand:
library_id: str
storage_key: str
file_hash: str = ""
+ asset_id: str = ""
+ celery_task_id: str = ""
class SubmitIngestJobUseCase:
@@ -24,5 +26,7 @@ class SubmitIngestJobUseCase:
library_id=command.library_id,
storage_key=command.storage_key,
file_hash=command.file_hash,
+ asset_id=command.asset_id,
+ celery_task_id=command.celery_task_id,
)
return self.ingest_job_repository.create(job)
diff --git a/packages/application/ingest_orphan_cleanup.py b/packages/application/ingest_orphan_cleanup.py
new file mode 100644
index 000000000..22fbb09a4
--- /dev/null
+++ b/packages/application/ingest_orphan_cleanup.py
@@ -0,0 +1,310 @@
+"""上传/转码链路(IngestJob + Asset)孤儿清理核心逻辑。
+
+#1714:generation 链路有 cleanup_stale_running/pending 兜底,但上传链路
+(ingest_jobs + assets)没有。worker 容器重启/进程 OOM 时,已 prefetch 的
+celery 消息会丢失(transcode 队列 worker_prefetch_multiplier=1,消息预取后
+宕机即丢失,Redis 队列里也不再存在),导致:
+
+- ingest_jobs.status 永久卡 pending/processing
+- assets.status 永久卡 processing/uploading(complete 阶段预建的占位)
+
+本模块提供纯核心(session 注入,便于单测):超时阈值内无更新的记录
+批量标终态(job→failed、asset→error),并返回 (job_id, celery_task_id)
+列表供调用方 revoke + purge 残留队列消息。
+"""
+
+from __future__ import annotations
+
+import logging
+from datetime import datetime, timedelta, timezone
+from typing import Any, Callable
+
+logger = logging.getLogger(__name__)
+
+# ingest_job PROCESSING 超时阈值:ingest 任务包含下载 + ffprobe + HEVC 转码
+# (1GB 视频约 10-20 分钟)+ 回传 OSS,正常任务可能跑 20-30 分钟;
+# 60 分钟阈值覆盖大文件转码 + 抖动,绝不误杀正常任务。
+INGEST_PROCESSING_TIMEOUT_MINUTES = 60
+
+# ingest_job PENDING 超时阈值:transcode 队列 concurrency=1,队列积压时
+# 正常排队可能较久;90 分钟覆盖 worker 短暂停消费 + 排队。
+INGEST_PENDING_TIMEOUT_MINUTES = 90
+
+# Asset 占位超时阈值:无关联 ingest_job 的孤儿占位(complete 预建后派单失败等),
+# 阈值放宽到 120 分钟,避免与 ingest_job 生命周期错杀。
+ASSET_ORPHAN_TIMEOUT_MINUTES = 120
+
+_TERMINAL_JOB_STATUSES = ("failed", "completed")
+_TERMINAL_ASSET_STATUSES = ("ready", "error", "deleted")
+
+
+def _now() -> datetime:
+ return datetime.now(timezone.utc)
+
+
+def cleanup_stale_ingest_jobs(
+ session: Any,
+ *,
+ processing_timeout_minutes: int = INGEST_PROCESSING_TIMEOUT_MINUTES,
+ pending_timeout_minutes: int = INGEST_PENDING_TIMEOUT_MINUTES,
+ commit: bool = True,
+) -> tuple[list[tuple[str, str]], list[str]]:
+ """清理超时卡 pending/processing 的 ingest_jobs,并联动关联 asset。
+
+ Args:
+ session: SQLAlchemy session(或提供 query/commit 的鸭子类型)
+ processing_timeout_minutes: processing 状态超时阈值
+ pending_timeout_minutes: pending 状态超时阈值
+ commit: 是否提交事务
+
+ Returns:
+ (job_items, asset_ids)
+ - job_items: [(job_id, celery_task_id), ...] 供 revoke/purge
+ - asset_ids: 被联动标记为 error 的 asset id 列表
+ """
+ from packages.adapters.sqlalchemy_impl.models import AssetModel, IngestJobModel
+
+ now = _now()
+ processing_cutoff = now - timedelta(minutes=processing_timeout_minutes)
+ pending_cutoff = now - timedelta(minutes=pending_timeout_minutes)
+
+ stale_jobs = (
+ session.query(IngestJobModel)
+ .filter(
+ IngestJobModel.status.in_(["pending", "processing"]),
+ (
+ (IngestJobModel.status == "processing") & (IngestJobModel.updated_at < processing_cutoff)
+ | (IngestJobModel.status == "pending") & (IngestJobModel.created_at < pending_cutoff)
+ ),
+ )
+ .all()
+ )
+
+ job_items: list[tuple[str, str]] = []
+ asset_ids: list[str] = []
+ stale_asset_models: list[Any] = []
+ for job_model in stale_jobs:
+ ref_time = job_model.updated_at or job_model.created_at
+ if ref_time.tzinfo is None: # SQLite 读回 naive datetime 的防御
+ ref_time = ref_time.replace(tzinfo=timezone.utc)
+ stale_minutes = int((now - ref_time).total_seconds() // 60)
+ job_model.status = "failed"
+ job_model.error_message = (
+ f"转码任务执行中断(超过超时阈值未更新,疑似 worker 重启/进程退出,已卡死 {stale_minutes} 分钟)"
+ )
+ job_model.updated_at = now
+ job_items.append((job_model.id, getattr(job_model, "celery_task_id", "") or ""))
+ if job_model.asset_id:
+ asset_ids.append(job_model.asset_id)
+
+ if asset_ids:
+ stale_asset_models = (
+ session.query(AssetModel)
+ .filter(
+ AssetModel.id.in_(asset_ids),
+ AssetModel.status.in_(["processing", "uploading"]),
+ )
+ .all()
+ )
+ for asset_model in stale_asset_models:
+ asset_model.status = "error"
+ asset_model.updated_at = now
+
+ if commit and (job_items or stale_asset_models):
+ session.commit()
+
+ if job_items:
+ logger.warning(
+ "[ingest-cleanup] 清理 %d 个超时 ingest_job(processing>%dm / pending>%dm),联动 %d 个 asset 标 error",
+ len(job_items),
+ processing_timeout_minutes,
+ pending_timeout_minutes,
+ len(stale_asset_models),
+ )
+ return job_items, [a.id for a in stale_asset_models]
+
+
+def cleanup_orphan_processing_assets(
+ session: Any,
+ *,
+ timeout_minutes: int = ASSET_ORPHAN_TIMEOUT_MINUTES,
+ commit: bool = True,
+) -> list[str]:
+ """清理无 ingest_job 关联、超时卡 processing/uploading 的孤儿 asset 占位。
+
+ complete 阶段预建 asset 后若派单失败(或 direct 上传 complete 后
+ 未触发 ingest),占位会永久卡住。这类 asset 没有对应 ingest_job,
+ 只能按 created_at 超时兜底标 error。
+ """
+ from packages.adapters.sqlalchemy_impl.models import AssetModel, IngestJobModel
+
+ cutoff = _now() - timedelta(minutes=timeout_minutes)
+ orphan_assets = (
+ session.query(AssetModel)
+ .outerjoin(IngestJobModel, IngestJobModel.asset_id == AssetModel.id)
+ .filter(
+ AssetModel.status.in_(["processing", "uploading"]),
+ AssetModel.created_at < cutoff,
+ IngestJobModel.id.is_(None),
+ )
+ .all()
+ )
+ for asset_model in orphan_assets:
+ asset_model.status = "error"
+ asset_model.updated_at = _now()
+ if commit and orphan_assets:
+ session.commit()
+ logger.warning("[ingest-cleanup] 清理 %d 个无 job 关联的超时孤儿 asset 占位", len(orphan_assets))
+ return [a.id for a in orphan_assets]
+
+
+def revoke_stale_ingest_messages(
+ job_items: list[tuple[str, str]],
+ *,
+ celery_app_factory: Callable[[], Any] | None = None,
+ broker_url_factory: Callable[[], str] | None = None,
+) -> int:
+ """revoke + 物理清理 ingest 作废消息(transcode/celery 队列)。
+
+ 消息可能已在 worker 宕机时丢失(队列里查不到),那也无害;
+ 若消息还在(极端重复投递),物理清除防止重投执行。
+ 失败不阻断清理(ingest_asset 的执行前状态守卫是第二道防线)。
+ """
+ biz_ids = [jid for jid, _ in job_items if jid]
+ celery_ids = [cid for _, cid in job_items if cid]
+ if not biz_ids and not celery_ids:
+ return 0
+ try:
+ from packages.shared.celery_orphan_guard import revoke_and_purge
+
+ app = celery_app_factory() if celery_app_factory else None
+ broker_url = broker_url_factory() if broker_url_factory else ""
+ if app is None or not broker_url:
+ from worker_app.celery_app import celery_app as _app
+ from worker_app.core.config import get_settings
+
+ app = _app
+ broker_url = get_settings().broker_url
+ return revoke_and_purge(
+ app,
+ broker_url,
+ business_task_ids=biz_ids,
+ celery_task_ids=celery_ids,
+ queue_names=("transcode", "celery"),
+ )
+ except Exception as e: # noqa: BLE001
+ logger.error("撤销作废 ingest 队列消息失败(执行前守卫仍会兜底): %s", e, exc_info=True)
+ return 0
+
+
+# ── worker 启动恢复(#1714)──────────────────────────────────────────────
+#
+# task_acks_late=True 下,worker 崩溃/容器重启时未 ack 的消息理论上会在
+# visibility_timeout 到期后重新投递;但 prefork 进程异常、部署窗口跨
+# visibility 配置边界等场景仍可能留下卡在 processing 的 ingest_job
+# (staging 实证:03:16 派单、03:45 置 processing 后 worker 重启,
+# unacked 消息未重投,任务永久卡死)。启动时做一次显式恢复扫描兜底。
+#
+# 恢复策略:processing 超过 stuck_minutes(默认 10 分钟,部署中跨进程
+# 交接的正常窗口 < 10 分钟,不会误抢别的 worker 正在执行的任务)的 job,
+# CAS 重置为 pending 并重新 send_task;旧消息若后来重投,ingest_asset
+# 的执行前守卫会把状态不匹配的旧 celery 消息丢弃。
+
+
+def recover_stuck_ingest_jobs_on_startup(
+ session: Any,
+ *,
+ send_task: Callable[..., Any] | None = None,
+ update_celery_task_id: Callable[[str, str], None] | None = None,
+ lock_acquire: Callable[[], bool] | None = None,
+ stuck_minutes: int = 10,
+ commit: bool = True,
+) -> int:
+ """worker 启动时把卡在 processing 超时的 ingest_job 重新派单。
+
+ Args:
+ session: SQLAlchemy session
+ send_task: celery send_task 可调用(注入便于测试);不传则用 worker celery_app
+ update_celery_task_id: 回写新 celery task id 的回调(job_id, new_task_id)
+ lock_acquire: 分布式锁获取回调(多 worker 进程同时启动时只允许一个恢复);
+ 返回 False 表示未抢到锁,本次跳过
+ stuck_minutes: processing 超过该分钟数视为卡死
+
+ Returns:
+ 重新派单的 job 数
+ """
+ if lock_acquire is not None and not lock_acquire():
+ logger.info("[ingest-recover] 未抢到恢复锁,跳过(另一进程正在恢复)")
+ return 0
+
+ from packages.adapters.sqlalchemy_impl.models import IngestJobModel
+
+ cutoff = _now() - timedelta(minutes=stuck_minutes)
+ stuck_jobs = (
+ session.query(IngestJobModel)
+ .filter(IngestJobModel.status == "processing", IngestJobModel.updated_at < cutoff)
+ .order_by(IngestJobModel.updated_at.asc())
+ .all()
+ )
+
+ if not stuck_jobs:
+ logger.info("[ingest-recover] 无卡死 processing ingest_job 需要恢复")
+ return 0
+
+ if send_task is None:
+ from worker_app.celery_app import celery_app as _app
+
+ send_task = _app.send_task
+
+ recovered = 0
+ for job_model in stuck_jobs:
+ # CAS:只有仍是 processing 才重置(并发/旧消息已回写终态时不碰)
+ updated = (
+ session.query(IngestJobModel)
+ .filter(IngestJobModel.id == job_model.id, IngestJobModel.status == "processing")
+ .update({"status": "pending", "error_message": "", "updated_at": _now()})
+ )
+ if not updated:
+ continue
+ try:
+ result = send_task("worker.ingest_asset", args=[job_model.id])
+ new_task_id = getattr(result, "id", "") or ""
+ except Exception as e: # noqa: BLE001
+ logger.error("[ingest-recover] 重新派单失败 job_id=%s: %s", job_model.id, e)
+ continue
+ if new_task_id:
+ job_model.celery_task_id = new_task_id
+ if update_celery_task_id is not None:
+ update_celery_task_id(job_model.id, new_task_id)
+ logger.warning(
+ "[ingest-recover] 卡死 ingest_job %s 已重置 pending 并重新派单 (new celery task=%s)",
+ job_model.id,
+ new_task_id,
+ )
+ recovered += 1
+
+ if commit and recovered:
+ session.commit()
+ logger.warning("[ingest-recover] 启动恢复完成,共重新派单 %d 个卡死 ingest_job", recovered)
+ return recovered
+
+
+def make_redis_recovery_lock(lock_key: str = "ingest:recover:startup", ttl_seconds: int = 300):
+ """构造基于 Redis SET NX 的恢复锁工厂(多 worker 进程互斥)。
+
+ 返回一个无参 callable,调用时尝试抢锁:抢到返回 True,未抢到返回 False。
+ Redis 不可用时不阻断启动恢复(返回 True,恢复逻辑自身有 CAS 幂等保护)。
+ """
+
+ def _acquire() -> bool:
+ try:
+ import redis as redis_lib
+ from worker_app.core.config import get_settings
+
+ client = redis_lib.Redis.from_url(get_settings().broker_url)
+ return bool(client.set(lock_key, "1", nx=True, ex=ttl_seconds))
+ except Exception as e: # noqa: BLE001
+ logger.warning("[ingest-recover] Redis 锁不可用,降级为无锁执行(CAS 兜底): %s", e)
+ return True
+
+ return _acquire
diff --git a/packages/domain/entities.py b/packages/domain/entities.py
index 4df342e4d..b03997a16 100755
--- a/packages/domain/entities.py
+++ b/packages/domain/entities.py
@@ -57,6 +57,8 @@ class User:
phone: str | None = None
phone_verified: bool = False
binding_completed_at: datetime | None = None
+ # 资料是否已完善(微信新用户首次设置昵称后置 True;邮箱注册默认 True)
+ profile_completed: bool = True
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
@@ -174,6 +176,7 @@ class Asset:
quality_score: float | None = None
uploaded_by_user_id: str = ""
file_hash: str = ""
+ client_upload_id: str = ""
metadata: dict[str, Any] = field(default_factory=dict)
tag_ids: list[str] = field(default_factory=list)
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
@@ -208,6 +211,7 @@ class Asset:
quality_score: float | None = None,
uploaded_by_user_id: str = "",
file_hash: str = "",
+ client_upload_id: str = "",
) -> "Asset":
clean_name = name.strip()
if not clean_name:
@@ -235,6 +239,7 @@ class Asset:
quality_score=quality_score,
uploaded_by_user_id=uploaded_by_user_id.strip(),
file_hash=file_hash.strip(),
+ client_upload_id=client_upload_id.strip(),
metadata=metadata or {},
tag_ids=[],
)
@@ -266,6 +271,8 @@ class IngestJob:
error_message: str = ""
result_asset_id: str = ""
file_hash: str = ""
+ asset_id: str = ""
+ celery_task_id: str = ""
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
@@ -276,6 +283,8 @@ class IngestJob:
library_id: str,
storage_key: str,
file_hash: str = "",
+ asset_id: str = "",
+ celery_task_id: str = "",
) -> "IngestJob":
if not project_id.strip():
raise ValueError("project_id 不能为空")
@@ -289,4 +298,6 @@ class IngestJob:
library_id=library_id.strip(),
storage_key=storage_key.strip(),
file_hash=file_hash.strip(),
+ asset_id=asset_id.strip(),
+ celery_task_id=celery_task_id.strip(),
)
diff --git a/packages/domain/generation_task.py b/packages/domain/generation_task.py
index a0e46ac37..ad00c3a70 100755
--- a/packages/domain/generation_task.py
+++ b/packages/domain/generation_task.py
@@ -117,6 +117,7 @@ class GenerationTask:
bgm_config: dict = field(default_factory=dict)
is_preview: bool = False
source_task_id: str = ""
+ celery_task_id: str = ""
output_width: int = 1280
output_height: int = 720
cover_url: str = ""
diff --git a/packages/ports/asset_repository.py b/packages/ports/asset_repository.py
index 9a9c830ad..b92c19fa4 100755
--- a/packages/ports/asset_repository.py
+++ b/packages/ports/asset_repository.py
@@ -125,3 +125,23 @@ class AssetRepository(ABC):
) -> Asset | None:
"""按素材库 + 文件哈希查找已有素材(去重检测)。"""
pass
+
+ @abstractmethod
+ def find_by_library_and_client_upload_id(
+ self,
+ library_id: str,
+ client_upload_id: str,
+ ) -> Asset | None:
+ """按素材库 + 客户端幂等 token 查找已有素材(complete 幂等)。"""
+ pass
+
+ @abstractmethod
+ def find_recent_active_by_library_and_name(
+ self,
+ library_id: str,
+ name: str,
+ within_minutes: int = 30,
+ file_size: int = 0,
+ ) -> Asset | None:
+ """兜底去重:同库 + 同文件名(+同大小)且近期仍在 uploading/processing 的素材。"""
+ pass
diff --git a/packages/ports/generation_task_repository.py b/packages/ports/generation_task_repository.py
index 5c21200e5..cc79f0370 100755
--- a/packages/ports/generation_task_repository.py
+++ b/packages/ports/generation_task_repository.py
@@ -20,6 +20,12 @@ class GenerationTaskRepository(Protocol):
def count_pending_total(self) -> int: ...
+ def count_running_by_user(self, user_id: str) -> int: ...
+
+ def count_running_total(self) -> int: ...
+
+ def estimate_avg_duration_seconds(self, limit: int = 20, default_seconds: float = 120.0) -> float: ...
+
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]: ...
def list_by_source_edit_plan(self, plan_id: str) -> list[GenerationTask]: ...
diff --git a/packages/shared/celery_orphan_guard.py b/packages/shared/celery_orphan_guard.py
new file mode 100644
index 000000000..123404f06
--- /dev/null
+++ b/packages/shared/celery_orphan_guard.py
@@ -0,0 +1,222 @@
+"""孤儿任务消息撤销与执行前状态守卫(API / Worker 共享)。
+
+#1714 / #1710 缺陷修复:超时清理/孤儿恢复把 DB 任务标记为 failed/cancelled
+后,Redis 队列里对应的 Celery 消息仍然存在;worker 重启或重新拉取时该消息
+被再次执行,状态机抛「非法状态转换: failed → running」,旧实现打印 ERROR 后
+继续跑,最终产出半成品。
+
+防御两道:
+1. 清理任务标 failed 时,调用 revoke_and_purge() 撤销(celery revoke 广播,
+ 通知在线 worker 丢弃)并直接扫描 Redis 队列移除消息体(worker 下线期间
+ 队列中的消息 revoke 广播收不到,必须物理移除);
+2. 任务真正开始业务逻辑前,调用 ensure_task_claimable() 校验 DB 状态,
+ 非 pending 的消息直接丢弃(抛 StaleTaskDiscarded,task 捕获后安全返回,
+ 不进入渲染/转码,不产出半成品)。
+"""
+
+from __future__ import annotations
+
+import base64
+import json
+import logging
+from collections.abc import Callable, Iterable
+from typing import Any
+
+logger = logging.getLogger(__name__)
+
+
+class StaleTaskDiscarded(Exception):
+ """任务消息已作废(DB 中任务已是终态),应安全中止、丢弃消息。"""
+
+ def __init__(self, task_id: str, status: str):
+ self.task_id = task_id
+ self.status = status
+ super().__init__(f"任务 {task_id} 已是终态 {status},丢弃重复/作废消息")
+
+
+# 终态状态值集合:处于这些状态的任务消息一律不执行
+TERMINAL_STATUS_VALUES = frozenset({"failed", "cancelled", "completed"})
+
+
+def ensure_task_claimable(
+ task_id: str,
+ get_status: Callable[[str], str | None],
+ *,
+ task_label: str = "任务",
+) -> str:
+ """执行前守卫:任务必须处于可领取状态(pending)。
+
+ Args:
+ task_id: 业务任务 ID
+ get_status: 回调,返回 DB 中任务当前状态字符串;返回 None 表示任务不存在
+ task_label: 日志用任务类型名
+
+ Returns:
+ 当前状态字符串(pending);任务不存在时返回空串(由调用方处理 not found)
+
+ Raises:
+ StaleTaskDiscarded: 任务已是终态(failed/cancelled/completed),消息必须丢弃
+ """
+ status = get_status(task_id)
+ if status is None:
+ return ""
+ if status in TERMINAL_STATUS_VALUES:
+ logger.warning("[%s] task_id=%s 状态已为 %s,消息作废,丢弃不执行", task_label, task_id, status)
+ raise StaleTaskDiscarded(task_id, status)
+ return status
+
+
+def _extract_business_ids(raw: bytes) -> tuple[str | None, str | None]:
+ """从 Redis 中的 Celery 消息提取 (celery 消息 ID, 业务任务 ID)。
+
+ Redis transport 存储格式为 JSON 信封:
+ {"body": base64(json), "headers": {"id": , "task": , ...}, ...}
+ body 解码后 Celery task 协议为 [args, kwargs, embed];
+ generate_video / ingest_asset 均以 args=[业务任务ID] 投递。
+
+ 无法解析时返回 (None, None)(保守保留该消息,绝不误删)。
+ """
+ try:
+ envelope = json.loads(raw)
+ celery_id = None
+ headers = envelope.get("headers") or {}
+ if isinstance(headers, dict):
+ celery_id = headers.get("id")
+ body = envelope.get("body")
+ if not body:
+ return celery_id, None
+ decoded = base64.b64decode(body)
+ payload = json.loads(decoded)
+ # 两种 body 形态:
+ # 1. 标准 Celery task 消息:[args, kwargs, embed] 三元组 → 业务 ID 在 payload[0][0]
+ # 2. 裸 producer 发布:body 即 args 数组 ["biz-id"] → 业务 ID 在 payload[0]
+ args = None
+ if isinstance(payload, dict):
+ args = payload.get("args")
+ elif isinstance(payload, (list, tuple)) and payload:
+ first = payload[0]
+ if isinstance(first, (list, tuple)):
+ args = first # 三元组:[args, kwargs, embed]
+ else:
+ args = payload # body 本身就是 args
+ if isinstance(args, (list, tuple)) and args and args[0] is not None:
+ return celery_id, str(args[0])
+ return celery_id, None
+ except Exception:
+ return None, None
+
+
+def purge_stale_messages_from_queues(
+ broker_url: str,
+ queue_names: Iterable[str],
+ business_task_ids: Iterable[str] = (),
+ celery_task_ids: Iterable[str] = (),
+) -> int:
+ """扫描 Redis 队列,移除作废任务的待消费消息。
+
+ 同时按业务任务 ID(消息 args[0])和 celery 消息 ID(headers.id)匹配,
+ 任一命中即移除。未命中或无法解析的消息原样保留(保持相对顺序)。
+
+ Returns:
+ 实际移除的消息条数
+ """
+ biz_ids = {bid for bid in business_task_ids if bid}
+ msg_ids = {mid for mid in celery_task_ids if mid}
+ if not biz_ids and not msg_ids:
+ return 0
+
+ try:
+ import redis
+ except ImportError:
+ logger.warning("redis-py 不可用,跳过队列消息清理")
+ return 0
+
+ try:
+ client = redis.Redis.from_url(broker_url)
+ client.ping()
+ except Exception as e:
+ logger.warning("连接 Redis 清理作废消息失败: %s", e)
+ return 0
+
+ removed_total = 0
+ try:
+ for queue in queue_names:
+ removed_total += _purge_one_queue(client, queue, biz_ids, msg_ids)
+ finally:
+ try:
+ client.close()
+ except Exception:
+ pass
+ if removed_total:
+ logger.info(
+ "从 Redis 队列移除 %d 条作废消息(biz=%s, celery=%s)",
+ removed_total,
+ sorted(biz_ids),
+ sorted(msg_ids),
+ )
+ return removed_total
+
+
+def _purge_one_queue(client: Any, queue_name: str, biz_ids: set[str], msg_ids: set[str]) -> int:
+ try:
+ raw_messages = client.lrange(queue_name, 0, -1)
+ except Exception as e:
+ logger.warning("读取队列 %s 失败: %s", queue_name, e)
+ return 0
+ if not raw_messages:
+ return 0
+
+ keep: list[bytes] = []
+ removed = 0
+ for raw in raw_messages:
+ celery_id, biz_id = _extract_business_ids(raw)
+ hit = (biz_id is not None and biz_id in biz_ids) or (celery_id is not None and celery_id in msg_ids)
+ if hit:
+ removed += 1
+ continue
+ keep.append(raw)
+
+ if removed:
+ try:
+ pipe = client.pipeline()
+ pipe.delete(queue_name)
+ if keep:
+ pipe.rpush(queue_name, *keep)
+ pipe.execute()
+ except Exception as e:
+ logger.warning("重写队列 %s 失败: %s", queue_name, e)
+ return 0
+ return removed
+
+
+def revoke_and_purge(
+ celery_app: Any,
+ broker_url: str,
+ business_task_ids: Iterable[str] = (),
+ celery_task_ids: Iterable[str] = (),
+ *,
+ queue_names: Iterable[str] = ("generation", "transcode", "celery"),
+) -> int:
+ """撤销作废任务:revoke 广播(在线 worker)+ 物理清理 Redis 队列消息。
+
+ Args:
+ celery_app: Celery app 实例(worker 端 worker_app.celery_app.celery_app)
+ broker_url: Redis broker URL
+ business_task_ids: 业务任务 ID(generation_tasks.id / ingest_jobs.id)
+ celery_task_ids: 入队时记录的 celery 消息 ID
+ queue_names: 需要扫描清理的队列名
+
+ Returns:
+ 从队列中实际移除的消息条数
+ """
+ for tid in celery_task_ids:
+ if not tid:
+ continue
+ try:
+ celery_app.control.revoke(tid)
+ except Exception as e:
+ logger.warning("revoke celery 消息 %s 失败: %s", tid, e)
+
+ return purge_stale_messages_from_queues(
+ broker_url, queue_names, business_task_ids=business_task_ids, celery_task_ids=celery_task_ids
+ )
diff --git a/packages/shared/celery_queues.py b/packages/shared/celery_queues.py
new file mode 100644
index 000000000..c9f4e45e6
--- /dev/null
+++ b/packages/shared/celery_queues.py
@@ -0,0 +1,58 @@
+"""Celery 队列定义与路由配置(API / Worker 共享)。
+
+#1714 队列隔离:用户等待的视频生成任务路由到高优先级 `generation` 队列,
+由专用 worker 进程独占消费;素材入库/转码等后台批量任务路由到 `transcode`
+队列;其余杂项任务走默认 `celery` 队列。转码队列积压时,视频生成任务
+仍能被 generation worker 立即领取执行,不会排队。
+
+队列说明:
+- generation: 用户提交的视频生成/预览渲染(延迟敏感,资源消耗大)
+- transcode: 素材入库(HEVC 转码)、AI 分类、素材查重(批量、可排队)
+- celery(默认): 配音、语音、下载缩略图、定时清理等杂项
+"""
+
+from __future__ import annotations
+
+from kombu import Queue
+
+# ── 队列名常量(生产端与消费端共用,禁止拼写漂移) ──
+QUEUE_GENERATION = "generation"
+QUEUE_TRANSCODE = "transcode"
+QUEUE_DEFAULT = "celery"
+
+# Worker 消费的队列列表(顺序即优先级:高优队列排在前面)
+WORKER_QUEUES = (QUEUE_GENERATION, QUEUE_TRANSCODE, QUEUE_DEFAULT)
+
+# 队列声明:持久化队列,broker 重启不丢消息
+task_queues = (
+ Queue(QUEUE_GENERATION, routing_key=QUEUE_GENERATION, durable=True),
+ Queue(QUEUE_TRANSCODE, routing_key=QUEUE_TRANSCODE, durable=True),
+ Queue(QUEUE_DEFAULT, routing_key=QUEUE_DEFAULT, durable=True),
+)
+
+# ── 任务路由表:task name → 队列 ──
+# 键支持 celery 标准通配符。
+task_routes = {
+ # 高优先级:用户等待的视频生成
+ "worker.generate_video": {"queue": QUEUE_GENERATION},
+ # 后台批量:素材入库/转码 + AI 分类 + 素材查重,积压不影响生成
+ "worker.ingest_asset": {"queue": QUEUE_TRANSCODE},
+ "worker.classify_asset": {"queue": QUEUE_TRANSCODE},
+ "worker.process_duplication_check": {"queue": QUEUE_TRANSCODE},
+ "worker.check_duplicate": {"queue": QUEUE_TRANSCODE},
+}
+
+# 生成任务的预取数:渲染是长任务,预取 1 避免任务被某个 worker 占住不调度
+GENERATION_WORKER_PREFETCH_MULTIPLIER = 1
+
+
+def apply_queue_settings(app) -> None:
+ """把队列隔离配置应用到 Celery app(API 生产端与 Worker 消费端都要调用)。
+
+ 配置 task_queues / task_routes / task_default_queue。生产端靠 task_routes
+ 把消息投递到对应队列;消费端靠 task_queues 声明自己消费哪些队列
+ (实际消费集由启动参数 -Q 控制)。
+ """
+ app.conf.task_queues = task_queues
+ app.conf.task_routes = task_routes
+ app.conf.task_default_queue = QUEUE_DEFAULT
diff --git a/scripts/render_env.sh b/scripts/render_env.sh
index 42930f6e1..a5f128045 100644
--- a/scripts/render_env.sh
+++ b/scripts/render_env.sh
@@ -57,7 +57,7 @@ if [ "$TARGET_ENV" = "staging" ]; then
fi
# 共用 secrets 直接导出(如果存在)
-SHARED_SECRETS="OSS_ACCESS_KEY_ID OSS_ACCESS_KEY_SECRET COSYVOICE_API_KEY DASHSCOPE_API_KEY MEDIAKIT_API_KEY"
+SHARED_SECRETS="OSS_ACCESS_KEY_ID OSS_ACCESS_KEY_SECRET COSYVOICE_API_KEY DASHSCOPE_API_KEY MEDIAKIT_API_KEY WECHAT_APP_ID WECHAT_APP_SECRET"
for var in $SHARED_SECRETS; do
value="${!var:-}"
# 已经在环境中了,无需额外操作
diff --git a/start-worker.ps1 b/start-worker.ps1
index a63e9e659..86f1e7197 100644
--- a/start-worker.ps1
+++ b/start-worker.ps1
@@ -14,4 +14,4 @@ Write-Host "`n启动 Celery Worker..." -ForegroundColor Yellow
Write-Host "监听任务队列: Redis (47.98.113.167:6379)" -ForegroundColor Cyan
Write-Host "`n按 Ctrl+C 停止服务`n" -ForegroundColor Gray
-celery -A celery_app worker --loglevel=info --pool=solo
+celery -A celery_app worker --loglevel=info --pool=solo -Q generation,transcode,celery
diff --git a/tests/unit/test_1677_batch_variants.py b/tests/unit/test_1677_batch_variants.py
new file mode 100644
index 000000000..bf67b5b62
--- /dev/null
+++ b/tests/unit/test_1677_batch_variants.py
@@ -0,0 +1,507 @@
+"""Issue #1677 多视频批量生成 — 变体独立配置与批量预览/批量生成测试。
+
+覆盖:
+- 批量预览:preview_count=N 一次创建 N 个独立任务,返回变体数组
+- 变体克隆链路:N 个预览/正式任务各自关联独立克隆 plan
+- 变体独立配置:titles[]/voice_library_ids[]/cover_urls[] 按变体注入
+- 长度校验:数组长度必须为 1 或 N(共用或独立),非法长度报错
+- N=1 向后兼容:旧字段单值行为不变
+"""
+
+from datetime import datetime, timezone
+from unittest.mock import MagicMock, patch
+
+import pytest
+from app.core.task_enqueue import GlobalQueueFull, UserPendingLimitExceeded
+from app.schemas.generation_task import (
+ BatchPreviewGenerationTaskResponse,
+ CreateGenerationTaskRequest,
+ CreatePreviewGenerationTaskRequest,
+)
+
+from packages.domain import GenerationTask
+from packages.domain.generation_task import GenerationTaskStatus
+
+# ════════════════════════════════════════════════════════════════════════════
+# 辅助构造
+# ════════════════════════════════════════════════════════════════════════════
+
+
+def _make_user(user_id="test_user_001"):
+ mock_user = MagicMock()
+ mock_user.id = user_id
+ auth = MagicMock()
+ auth.user = mock_user
+ return auth
+
+
+def _make_task(task_id=None, status=GenerationTaskStatus.PENDING, source_plan_id=None):
+ task = GenerationTask.create(
+ project_id="",
+ asset_library_id="",
+ template_id="tpl_001",
+ asset_ids=["asset_1"],
+ )
+ if task_id:
+ task.id = task_id
+ task.status = status
+ task.is_preview = True
+ task.source_edit_plan_id = source_plan_id or ""
+ task.voice_library_id = ""
+ task.title_config = {}
+ task.cover_url = ""
+ return task
+
+
+def _make_preview_request(**kwargs):
+ defaults = {
+ "template_id": "tpl_001",
+ "asset_ids": ["asset_1", "asset_2"],
+ }
+ defaults.update(kwargs)
+ return CreatePreviewGenerationTaskRequest(**defaults)
+
+
+def _repo_mock():
+ repo = MagicMock()
+ repo.count_pending_by_user.return_value = 0
+ repo.count_pending_total.return_value = 0
+ repo.get.side_effect = lambda tid: None
+ return repo
+
+
+# ════════════════════════════════════════════════════════════════════════════
+# Schema 校验:变体数组长度
+# ════════════════════════════════════════════════════════════════════════════
+
+
+class TestVariantArrayValidation:
+ """变体数组字段长度校验。"""
+
+ def test_preview_titles_length_matches_count(self):
+ """titles 长度 = preview_count 合法"""
+ req = _make_preview_request(preview_count=3, titles=["标题A", "标题B", "标题C"])
+ assert len(req.titles) == 3
+
+ def test_preview_titles_single_shared(self):
+ """titles 长度 1 = 所有变体共用,合法"""
+ req = _make_preview_request(preview_count=3, titles=["共用标题"])
+ assert req.titles == ["共用标题"]
+
+ def test_preview_titles_wrong_length_raises(self):
+ """titles 长度 2 与 preview_count=3 不匹配 → 报错"""
+ with pytest.raises(ValueError, match="titles"):
+ _make_preview_request(preview_count=3, titles=["A", "B"])
+
+ def test_preview_voice_ids_wrong_length_raises(self):
+ """voice_library_ids 长度非法 → 报错"""
+ from pydantic import ValidationError
+
+ with pytest.raises(ValidationError, match="voice_library_ids"):
+ _make_preview_request(preview_count=4, voice_library_ids=["v1", "v2"])
+
+ def test_preview_empty_arrays_ok(self):
+ """空数组(回退单值字段)合法"""
+ req = _make_preview_request(preview_count=3)
+ assert req.titles == []
+ assert req.voice_library_ids == []
+ assert req.cover_urls == []
+
+ def test_generation_titles_length_matches_count(self):
+ """正式生成 titles 长度 = count 合法"""
+ req = CreateGenerationTaskRequest(
+ template_id="tpl_1",
+ asset_ids=["a1"],
+ count=3,
+ titles=["A", "B", "C"],
+ )
+ assert len(req.titles) == 3
+
+ def test_generation_arrays_wrong_length_raises(self):
+ """正式生成 cover_urls 长度与 count 不匹配 → 报错"""
+ from pydantic import ValidationError
+
+ with pytest.raises(ValidationError, match="cover_urls"):
+ CreateGenerationTaskRequest(
+ template_id="tpl_1",
+ asset_ids=["a1"],
+ count=3,
+ cover_urls=["c1", "c2"],
+ )
+
+ def test_generation_single_count_no_arrays(self):
+ """N=1 且不传数组:完全旧行为"""
+ req = CreateGenerationTaskRequest(template_id="tpl_1", asset_ids=["a1"])
+ assert req.count == 1
+ assert req.titles == []
+ assert req.voice_library_ids == []
+ assert req.cover_urls == []
+
+
+# ════════════════════════════════════════════════════════════════════════════
+# 批量预览路由
+# ════════════════════════════════════════════════════════════════════════════
+
+
+class TestBatchPreviewRoute:
+ """POST /preview 批量变体。"""
+
+ def test_preview_count_1_returns_single_item_array(self):
+ """N=1 返回 items 长度 1 的批量响应(结构统一)"""
+ from app.api.routes.generation_preview import create_preview_generation_task
+
+ task = _make_task(task_id="task_1")
+ repo = _repo_mock()
+ with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
+ MockUC.return_value.execute.return_value = task
+ with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True):
+ resp = create_preview_generation_task(
+ _make_preview_request(preview_count=1),
+ authenticated_user=_make_user(),
+ generation_task_repository=repo,
+ db=MagicMock(),
+ )
+ assert isinstance(resp, BatchPreviewGenerationTaskResponse)
+ assert resp.total == 1
+ assert len(resp.items) == 1
+ assert resp.items[0].task_id == "task_1"
+ assert resp.items[0].variant_index == 0
+
+ def test_preview_count_3_creates_three_independent_tasks(self):
+ """N=3 创建 3 个独立任务,返回 3 个变体,task_id 各不相同"""
+ from app.api.routes.generation_preview import create_preview_generation_task
+
+ tasks = [_make_task(task_id=f"task_{i}") for i in range(3)]
+ repo = _repo_mock()
+ with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
+ MockUC.return_value.execute.side_effect = tasks
+ with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True):
+ resp = create_preview_generation_task(
+ _make_preview_request(preview_count=3),
+ authenticated_user=_make_user(),
+ generation_task_repository=repo,
+ db=MagicMock(),
+ )
+ assert resp.total == 3
+ task_ids = [item.task_id for item in resp.items]
+ assert task_ids == ["task_0", "task_1", "task_2"]
+ assert len(set(task_ids)) == 3
+ for i, item in enumerate(resp.items):
+ assert item.variant_index == i
+
+ def test_preview_count_3_clones_three_variant_plans(self):
+ """有源 plan 时,N=3 克隆 3 个独立变体 plan(预览全部克隆,不用源 plan)"""
+ from app.api.routes.generation_preview import create_preview_generation_task
+
+ tasks = [_make_task(task_id=f"task_{i}", source_plan_id="source_plan") for i in range(3)]
+ repo = _repo_mock()
+ cloned_plan_ids = ["clone_1", "clone_2", "clone_3"]
+ with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
+ MockUC.return_value.execute.side_effect = tasks
+ with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True):
+ with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc:
+ clone_results = [MagicMock(id=pid) for pid in cloned_plan_ids]
+ MockPlanSvc.return_value.clone_plan_for_variant.side_effect = clone_results
+ create_preview_generation_task(
+ _make_preview_request(preview_count=3),
+ authenticated_user=_make_user(),
+ generation_task_repository=repo,
+ db=MagicMock(),
+ )
+ # 克隆被调用 3 次
+ assert MockPlanSvc.return_value.clone_plan_for_variant.call_count == 3
+ # 每个任务关联到不同的克隆 plan
+ for i, task in enumerate(tasks):
+ assert task.source_edit_plan_id == cloned_plan_ids[i]
+
+ def test_preview_variant_titles_injected_per_variant(self):
+ """titles[] 按变体注入 title_config.text"""
+ from app.api.routes.generation_preview import create_preview_generation_task
+
+ tasks = [_make_task(task_id=f"task_{i}") for i in range(3)]
+ repo = _repo_mock()
+ captured_commands = []
+ with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
+
+ def _execute(cmd):
+ captured_commands.append(cmd)
+ return tasks[len(captured_commands) - 1]
+
+ MockUC.return_value.execute.side_effect = _execute
+ with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True):
+ create_preview_generation_task(
+ _make_preview_request(
+ preview_count=3,
+ title_config={"font": "黑体", "position": "bottom"},
+ titles=["标题A", "标题B", "标题C"],
+ ),
+ authenticated_user=_make_user(),
+ generation_task_repository=repo,
+ db=MagicMock(),
+ )
+ assert len(captured_commands) == 3
+ assert captured_commands[0].title_config["text"] == "标题A"
+ assert captured_commands[1].title_config["text"] == "标题B"
+ assert captured_commands[2].title_config["text"] == "标题C"
+ # 样式全局共用
+ assert all(c.title_config["font"] == "黑体" for c in captured_commands)
+
+ def test_preview_shared_title_when_single_length(self):
+ """titles 长度 1 = 所有变体共用同一标题"""
+ from app.api.routes.generation_preview import create_preview_generation_task
+
+ tasks = [_make_task(task_id=f"task_{i}") for i in range(3)]
+ repo = _repo_mock()
+ captured = []
+ with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
+
+ def _execute(cmd):
+ captured.append(cmd)
+ return tasks[len(captured) - 1]
+
+ MockUC.return_value.execute.side_effect = _execute
+ with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True):
+ create_preview_generation_task(
+ _make_preview_request(preview_count=3, titles=["共用标题"]),
+ authenticated_user=_make_user(),
+ generation_task_repository=repo,
+ db=MagicMock(),
+ )
+ assert all(c.title_config["text"] == "共用标题" for c in captured)
+
+ def test_preview_independent_voice_per_variant(self):
+ """voice_library_ids[] 按变体注入独立配音"""
+ from app.api.routes.generation_preview import create_preview_generation_task
+
+ tasks = [_make_task(task_id=f"task_{i}") for i in range(3)]
+ repo = _repo_mock()
+ captured = []
+ with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
+
+ def _execute(cmd):
+ captured.append(cmd)
+ return tasks[len(captured) - 1]
+
+ MockUC.return_value.execute.side_effect = _execute
+ with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True):
+ create_preview_generation_task(
+ _make_preview_request(
+ preview_count=3,
+ voice_library_ids=["voice_a", "voice_b", "voice_c"],
+ ),
+ authenticated_user=_make_user(),
+ generation_task_repository=repo,
+ db=MagicMock(),
+ )
+ assert [c.voice_library_id for c in captured] == ["voice_a", "voice_b", "voice_c"]
+
+ def test_preview_voice_fallback_to_single_field(self):
+ """voice_library_ids 为空时回退 voice_library_id 单值字段(向后兼容)"""
+ from app.api.routes.generation_preview import create_preview_generation_task
+
+ task = _make_task(task_id="task_1")
+ repo = _repo_mock()
+ captured = []
+ with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
+
+ def _execute(cmd):
+ captured.append(cmd)
+ return task
+
+ MockUC.return_value.execute.side_effect = _execute
+ with patch("app.api.routes.generation_preview.safe_enqueue_generation_task", return_value=True):
+ create_preview_generation_task(
+ _make_preview_request(voice_library_id="legacy_voice"),
+ authenticated_user=_make_user(),
+ generation_task_repository=repo,
+ db=MagicMock(),
+ )
+ assert captured[0].voice_library_id == "legacy_voice"
+
+ def test_preview_queue_limit_checks_total_count(self):
+ """限流预检查按变体总数计:用户 pending + N 超限 → 429"""
+ from app.api.routes.generation_preview import create_preview_generation_task
+ from fastapi import HTTPException
+
+ repo = MagicMock()
+ repo.count_pending_by_user.return_value = 3
+ repo.count_pending_total.return_value = 0
+ with pytest.raises(HTTPException) as exc:
+ create_preview_generation_task(
+ _make_preview_request(preview_count=5),
+ authenticated_user=_make_user(),
+ generation_task_repository=repo,
+ db=MagicMock(),
+ )
+ assert exc.value.status_code == 429
+
+ def test_preview_clone_failure_marks_all_failed(self):
+ """克隆变体 plan 失败 → 已创建任务全部标记 failed 并 500"""
+ from app.api.routes.generation_preview import create_preview_generation_task
+ from fastapi import HTTPException
+
+ tasks = [_make_task(task_id=f"task_{i}", source_plan_id="source_plan") for i in range(3)]
+ repo = _repo_mock()
+ with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
+ MockUC.return_value.execute.side_effect = tasks
+ with patch("app.services.edit_plan_service.EditPlanService") as MockPlanSvc:
+ MockPlanSvc.return_value.clone_plan_for_variant.side_effect = RuntimeError("db down")
+ with pytest.raises(HTTPException) as exc:
+ create_preview_generation_task(
+ _make_preview_request(preview_count=3),
+ authenticated_user=_make_user(),
+ generation_task_repository=repo,
+ db=MagicMock(),
+ )
+ assert exc.value.status_code == 500
+ # 所有已创建任务都被标记 failed
+ assert all(t.status == GenerationTaskStatus.FAILED for t in tasks)
+
+
+# ════════════════════════════════════════════════════════════════════════════
+# 批量正式生成:变体配置注入
+# ════════════════════════════════════════════════════════════════════════════
+
+
+class TestBatchGenerationVariantConfig:
+ """POST /tasks count=N 时变体独立配置。"""
+
+ def _call_create_tasks(self, request, repo=None):
+ from app.api.routes.generation_tasks import create_generation_task
+
+ repo = repo or MagicMock()
+ repo.count_pending_by_user.return_value = 0
+ repo.count_pending_total.return_value = 0
+ repo.update.return_value = None
+
+ # 模板模式:asset_repository.find_by_id 返回 None(无 project 关联,
+ # 纯模板模式 project_id/library_id 都为空),避免 MagicMock 属性污染
+ asset_repo = MagicMock()
+ asset_repo.find_by_id.return_value = None
+
+ # db.query().filter()...first() 返回 None:不走兜底关联编辑计划
+ db = MagicMock()
+ db.query.return_value.filter.return_value.order_by.return_value.first.return_value = None
+
+ return create_generation_task(
+ request,
+ authenticated_user=_make_user(),
+ generation_task_repository=repo,
+ project_repository=MagicMock(),
+ asset_library_repository=MagicMock(),
+ asset_repository=asset_repo,
+ db=db,
+ )
+
+ def test_count_3_variant_titles_voices_covers_injected(self):
+ """count=3:titles/voice_library_ids/cover_urls 按变体注入"""
+ from app.api.routes import generation_tasks as routes
+
+ tasks = [_make_task(task_id=f"gen_{i}") for i in range(3)]
+ captured = []
+ with patch.object(routes, "CreateGenerationTaskUseCase") as MockUC:
+
+ def _execute(cmd):
+ captured.append(cmd)
+ t = tasks[len(captured) - 1]
+ t.title_config = cmd.title_config
+ t.voice_library_id = cmd.voice_library_id
+ t.cover_url = cmd.cover_url
+ return t
+
+ MockUC.return_value.execute.side_effect = _execute
+ with patch.object(routes, "safe_enqueue_generation_task", return_value=True):
+ req = CreateGenerationTaskRequest(
+ template_id="tpl_1",
+ asset_ids=["a1"],
+ count=3,
+ title_config={"font": "宋体"},
+ titles=["成片标题1", "成片标题2", "成片标题3"],
+ voice_library_ids=["v1", "v2", "v3"],
+ cover_urls=["http://c1", "http://c2", "http://c3"],
+ )
+ resp = self._call_create_tasks(req)
+ assert resp.total == 3
+ assert [c.title_config["text"] for c in captured] == ["成片标题1", "成片标题2", "成片标题3"]
+ assert [c.voice_library_id for c in captured] == ["v1", "v2", "v3"]
+ assert [c.cover_url for c in captured] == ["http://c1", "http://c2", "http://c3"]
+ # 样式共用
+ assert all(c.title_config["font"] == "宋体" for c in captured)
+
+ def test_count_1_legacy_fields_unchanged(self):
+ """N=1 不传数组:旧字段 voice_library_id/cover_url/title_config 行为不变"""
+ from app.api.routes import generation_tasks as routes
+
+ task = _make_task(task_id="gen_1")
+ task.is_preview = False
+ captured = []
+ with patch.object(routes, "CreateGenerationTaskUseCase") as MockUC:
+
+ def _execute(cmd):
+ captured.append(cmd)
+ return task
+
+ MockUC.return_value.execute.side_effect = _execute
+ with patch.object(routes, "safe_enqueue_generation_task", return_value=True):
+ req = CreateGenerationTaskRequest(
+ template_id="tpl_1",
+ asset_ids=["a1"],
+ count=1,
+ voice_library_id="legacy_voice",
+ cover_url="http://legacy-cover",
+ title_config={"text": "旧标题", "font": "黑体"},
+ )
+ resp = self._call_create_tasks(req)
+ assert resp.total == 1
+ assert captured[0].voice_library_id == "legacy_voice"
+ assert captured[0].cover_url == "http://legacy-cover"
+ assert captured[0].title_config["text"] == "旧标题"
+
+ def test_count_3_shared_single_value_arrays(self):
+ """数组长度 1:3 个变体共用同一配音/封面"""
+ from app.api.routes import generation_tasks as routes
+
+ tasks = [_make_task(task_id=f"gen_{i}") for i in range(3)]
+ captured = []
+ with patch.object(routes, "CreateGenerationTaskUseCase") as MockUC:
+
+ def _execute(cmd):
+ captured.append(cmd)
+ return tasks[len(captured) - 1]
+
+ MockUC.return_value.execute.side_effect = _execute
+ with patch.object(routes, "safe_enqueue_generation_task", return_value=True):
+ req = CreateGenerationTaskRequest(
+ template_id="tpl_1",
+ asset_ids=["a1"],
+ count=3,
+ voice_library_ids=["shared_voice"],
+ cover_urls=["http://shared"],
+ )
+ self._call_create_tasks(req)
+ assert all(c.voice_library_id == "shared_voice" for c in captured)
+ assert all(c.cover_url == "http://shared" for c in captured)
+
+
+class TestVariantValueHelper:
+ """_variant_value 取值逻辑。"""
+
+ def test_empty_returns_fallback(self):
+ from app.api.routes.generation_preview import _variant_value
+
+ assert _variant_value([], 0, fallback="fb") == "fb"
+
+ def test_single_length_shared(self):
+ from app.api.routes.generation_preview import _variant_value
+
+ assert _variant_value(["only"], 5) == "only"
+
+ def test_indexed_access(self):
+ from app.api.routes.generation_preview import _variant_value
+
+ assert _variant_value(["a", "b", "c"], 1) == "b"
+
+ def test_index_out_of_range_fallback(self):
+ from app.api.routes.generation_preview import _variant_value
+
+ assert _variant_value(["a", "b"], 9, fallback="x") == "x"
diff --git a/tests/unit/test_asset_repo_fallback_dedup_1714.py b/tests/unit/test_asset_repo_fallback_dedup_1714.py
new file mode 100644
index 000000000..d68890773
--- /dev/null
+++ b/tests/unit/test_asset_repo_fallback_dedup_1714.py
@@ -0,0 +1,93 @@
+"""#1714 find_recent_active_by_library_and_name 严格模式测试。
+
+file_size=0(未知)时必须返回 None(宁可漏判不可误杀);
+大小严格匹配;只命中近期 UPLOADING/PROCESSING 记录。
+"""
+
+import sys
+from datetime import datetime, timedelta, timezone
+from pathlib import Path
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
+
+from sqlalchemy import create_engine # noqa: E402
+from sqlalchemy.orm import sessionmaker # noqa: E402
+
+from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository # noqa: E402
+from packages.adapters.sqlalchemy_impl.models import Base # noqa: E402
+from packages.domain import Asset, AssetStatus # noqa: E402
+
+
+def _repository():
+ engine = create_engine("sqlite:///:memory:")
+ Base.metadata.create_all(engine)
+ session = sessionmaker(bind=engine)()
+ return SQLAlchemyAssetRepository(session)
+
+
+def _mk_asset(name="IMG_2285.MOV", file_size=5_000_000, status=AssetStatus.PROCESSING, minutes_ago=5):
+ asset = Asset.create(
+ project_id="proj-1",
+ library_id="lib-1",
+ name=name,
+ storage_key=f"uploads/x/{name}",
+ mime_type="video/quicktime",
+ file_size=file_size,
+ )
+ asset.status = status
+ asset.created_at = datetime.now(timezone.utc) - timedelta(minutes=minutes_ago)
+ return asset
+
+
+def test_returns_none_when_file_size_zero():
+ """file_size=0(大小未知)直接返回 None——不许仅凭同名 + processing 判重。"""
+ repo = _repository()
+ repo.create(_mk_asset(file_size=0))
+
+ result = repo.find_recent_active_by_library_and_name(library_id="lib-1", name="IMG_2285.MOV", file_size=0)
+ assert result is None
+
+
+def test_matches_when_name_size_strict_equal():
+ """同名 + 同大小 + processing 近期记录 → 命中。"""
+ repo = _repository()
+ repo.create(_mk_asset(file_size=5_000_000))
+
+ result = repo.find_recent_active_by_library_and_name(library_id="lib-1", name="IMG_2285.MOV", file_size=5_000_000)
+ assert result is not None
+ assert result.name == "IMG_2285.MOV"
+
+
+def test_no_match_when_same_name_but_different_size():
+ """同名但大小不同 → 不命中(内容全新的视频不能误杀)。"""
+ repo = _repository()
+ repo.create(_mk_asset(file_size=5_000_000))
+
+ result = repo.find_recent_active_by_library_and_name(library_id="lib-1", name="IMG_2285.MOV", file_size=9_999_999)
+ assert result is None
+
+
+def test_no_match_ready_history_even_with_same_size():
+ """READY 历史同名素材不命中(允许再次上传同名文件)。"""
+ repo = _repository()
+ repo.create(_mk_asset(file_size=5_000_000, status=AssetStatus.READY))
+
+ result = repo.find_recent_active_by_library_and_name(library_id="lib-1", name="IMG_2285.MOV", file_size=5_000_000)
+ assert result is None
+
+
+def test_no_match_when_window_expired():
+ """超过 30 分钟窗口的活动记录不命中。"""
+ repo = _repository()
+ repo.create(_mk_asset(file_size=5_000_000, minutes_ago=45))
+
+ result = repo.find_recent_active_by_library_and_name(
+ library_id="lib-1", name="IMG_2285.MOV", within_minutes=30, file_size=5_000_000
+ )
+ assert result is None
+
+
+def test_returns_none_when_name_empty():
+ repo = _repository()
+ result = repo.find_recent_active_by_library_and_name(library_id="lib-1", name="", file_size=100)
+ assert result is None
diff --git a/tests/unit/test_bad_fingerprint_filter.py b/tests/unit/test_bad_fingerprint_filter.py
index 05a18cb08..975a1ae8a 100644
--- a/tests/unit/test_bad_fingerprint_filter.py
+++ b/tests/unit/test_bad_fingerprint_filter.py
@@ -85,17 +85,18 @@ class TestIsBadFingerprint:
assert VideoDeduplicator._is_bad_fingerprint(["abcdef0123456789"]) is False
def test_all_identical_phashes_is_bad(self):
- """多帧但所有 phash 完全相同 → 黑屏/纯色视频。"""
- phashes = ["aaaaaaaaaaaaaaaa"] * 5
+ """>=8 帧且所有 phash 完全相同 → 黑屏/纯色视频(#1702:短帧不误杀)。"""
+ phashes = ["aaaaaaaaaaaaaaaa"] * 10
assert VideoDeduplicator._is_bad_fingerprint(phashes) is True
- def test_two_identical_phashes_is_bad(self):
- """两帧完全相同也视为坏指纹。"""
- assert VideoDeduplicator._is_bad_fingerprint(["bbbbbbbbbbbbbbbb", "bbbbbbbbbbbbbbbb"]) is True
+ def test_short_identical_phashes_not_bad(self):
+ """<8 帧完全相同不判坏——短视频内容连续时相邻采样帧 phash 天然相同(#1702)。"""
+ assert VideoDeduplicator._is_bad_fingerprint(["bbbbbbbbbbbbbbbb"] * 5) is False
+ assert VideoDeduplicator._is_bad_fingerprint(["bbbbbbbbbbbbbbbb", "bbbbbbbbbbbbbbbb"]) is False
def test_all_very_similar_phashes_is_bad(self):
- """多帧 phash 之间的汉明距离都 < 3 → 近似黑屏。"""
- phashes = ["0000000000000000", "0000000000000001", "0000000000000002"]
+ """>=8 帧 phash 之间的汉明距离都 < 3 且高占比 → 近似黑屏。"""
+ phashes = ["0000000000000000"] * 8 + ["0000000000000001", "0000000000000002"]
assert VideoDeduplicator._is_bad_fingerprint(phashes) is True
def test_diverse_phashes_is_good(self):
@@ -122,7 +123,9 @@ class TestIsBadFingerprint:
"""已知黑屏视频的 phash 特征(全零或均匀分布)。"""
assert VideoDeduplicator._is_bad_fingerprint(["0000000000000000"] * 10) is True
assert VideoDeduplicator._is_bad_fingerprint(["ffffffffffffffff"] * 8) is True
- assert VideoDeduplicator._is_bad_fingerprint(["9999999999999966"] * 6) is True
+ assert VideoDeduplicator._is_bad_fingerprint(["9999999999999966"] * 8) is True
+ # <8 帧不判坏(#1702 短视频保护)
+ assert VideoDeduplicator._is_bad_fingerprint(["9999999999999966"] * 5) is False
# ── Helper ──────────────────────────────────────────────────────
@@ -151,13 +154,13 @@ class TestCheckDuplicateBadFingerprint:
deduplicator = VideoDeduplicator()
mock_session = MagicMock()
- black_screen = _make_existing_video("vid-black", "md5_black", ["aaaaaaaaaaaaaaaa"] * 5)
+ black_screen = _make_existing_video("vid-black", "md5_black", ["aaaaaaaaaaaaaaaa"] * 10)
mock_repo = MagicMock()
mock_repo.list_by_user.return_value = [black_screen]
fingerprint = VideoFingerprint(
md5="md5_normal",
- keyframe_phashes=["aaaaaaaaaaaaaaaa"] * 5,
+ keyframe_phashes=["aaaaaaaaaaaaaaaa"] * 10,
color_histograms=[],
duration=10.0,
resolution=(1280, 720),
@@ -206,7 +209,7 @@ class TestCheckDuplicateBadFingerprint:
deduplicator = VideoDeduplicator()
mock_session = MagicMock()
- black_screen = _make_existing_video("vid-black", "same_md5", ["aaaaaaaaaaaaaaaa"] * 5)
+ black_screen = _make_existing_video("vid-black", "same_md5", ["aaaaaaaaaaaaaaaa"] * 10)
mock_repo = MagicMock()
mock_repo.list_by_user.return_value = [black_screen]
@@ -283,8 +286,8 @@ class TestComputeDuplicateRateBadFingerprint:
mock_session = MagicMock()
videos = [
- _make_existing_video("vid-b1", "md5_b1", ["aaaaaaaaaaaaaaaa"] * 5),
- _make_existing_video("vid-b2", "md5_b2", ["bbbbbbbbbbbbbbbb"] * 5),
+ _make_existing_video("vid-b1", "md5_b1", ["aaaaaaaaaaaaaaaa"] * 10),
+ _make_existing_video("vid-b2", "md5_b2", ["cccccccccccccccc"] * 5), # hamming(a,c)=32 > PHASH_THRESHOLD
]
mock_repo = MagicMock()
mock_repo.list_by_user.return_value = videos
diff --git a/tests/unit/test_celery_queue_isolation_1714.py b/tests/unit/test_celery_queue_isolation_1714.py
new file mode 100644
index 000000000..1c473eada
--- /dev/null
+++ b/tests/unit/test_celery_queue_isolation_1714.py
@@ -0,0 +1,176 @@
+"""#1714 队列隔离 + 作废消息清除 单元测试。
+
+覆盖:
+1. task_routes:generate_video → generation,ingest_asset/classify/duplication → transcode
+2. purge_stale_messages_from_queues:Redis 队列中作废任务消息被物理移除,未命中保留
+3. revoke_and_purge:revoke 广播 + 队列清理同时生效
+4. ensure_task_claimable:终态任务抛 StaleTaskDiscarded,pending 放行
+"""
+
+from __future__ import annotations
+
+from unittest.mock import MagicMock
+
+import pytest
+from celery import Celery
+
+from packages.shared.celery_orphan_guard import (
+ StaleTaskDiscarded,
+ _extract_business_ids,
+ ensure_task_claimable,
+ purge_stale_messages_from_queues,
+ revoke_and_purge,
+)
+from packages.shared.celery_queues import (
+ QUEUE_GENERATION,
+ QUEUE_TRANSCODE,
+ apply_queue_settings,
+ task_routes,
+)
+
+BROKER_URL = "redis://localhost:6379/15"
+TEST_QUEUES = ("_test_gen_q", "_test_transcode_q")
+
+
+# ── 1. 路由表 ──────────────────────────────────────────────────────────
+
+
+def test_routes_send_generation_to_generation_queue():
+ assert task_routes["worker.generate_video"]["queue"] == QUEUE_GENERATION
+
+
+def test_routes_send_ingest_to_transcode_queue():
+ assert task_routes["worker.ingest_asset"]["queue"] == QUEUE_TRANSCODE
+ assert task_routes["worker.classify_asset"]["queue"] == QUEUE_TRANSCODE
+ assert task_routes["worker.process_duplication_check"]["queue"] == QUEUE_TRANSCODE
+ assert task_routes["worker.check_duplicate"]["queue"] == QUEUE_TRANSCODE
+
+
+def test_apply_queue_settings_configures_celery_app():
+ app = Celery("test-routes")
+ apply_queue_settings(app)
+ queue_names = {q.name for q in app.conf.task_queues}
+ assert queue_names == {"generation", "transcode", "celery"}
+ assert app.conf.task_default_queue == "celery"
+
+
+# ── Redis 队列消息清理(需要本地 redis;不可用时 skip) ─────────────────
+
+
+def _redis_available() -> bool:
+ try:
+ import redis
+
+ return bool(redis.Redis.from_url(BROKER_URL).ping())
+ except Exception:
+ return False
+
+
+@pytest.fixture()
+def redis_client():
+ import redis
+
+ client = redis.Redis.from_url(BROKER_URL)
+ for q in TEST_QUEUES:
+ client.delete(q)
+ yield client
+ for q in TEST_QUEUES:
+ client.delete(q)
+
+
+def _publish(app: Celery, queue: str, celery_id: str, business_id: str) -> None:
+ from kombu import Queue
+ from kombu.pools import producers
+
+ with app.connection_for_write() as conn:
+ with producers[conn].acquire(block=True) as prod:
+ prod.publish(
+ (business_id,),
+ exchange="",
+ routing_key=queue,
+ serializer="json",
+ headers={"id": celery_id, "task": "worker.generate_video"},
+ retry=False,
+ delivery_mode=1,
+ declare=[Queue(queue, routing_key=queue, durable=False)],
+ )
+
+
+@pytest.mark.skipif(not _redis_available(), reason="本地 redis 不可用")
+def test_purge_removes_stale_business_message_and_keeps_others(redis_client):
+ app = Celery("test-purge")
+ app.conf.broker_url = BROKER_URL
+ _publish(app, TEST_QUEUES[0], "celery-1", "task-KEEP-A")
+ _publish(app, TEST_QUEUES[0], "celery-2", "task-STALE-B")
+ _publish(app, TEST_QUEUES[0], "celery-3", "task-KEEP-C")
+ _publish(app, TEST_QUEUES[1], "celery-4", "task-STALE-B") # 同一业务任务在转码队列?不应出现但验证全队列扫描
+
+ removed = purge_stale_messages_from_queues(BROKER_URL, TEST_QUEUES, business_task_ids={"task-STALE-B"})
+ assert removed == 2
+
+ remaining = []
+ for raw in redis_client.lrange(TEST_QUEUES[0], 0, -1):
+ _celery_id, biz_id = _extract_business_ids(raw)
+ remaining.append(biz_id)
+ assert set(remaining) == {"task-KEEP-A", "task-KEEP-C"}
+ assert redis_client.llen(TEST_QUEUES[1]) == 0
+
+
+@pytest.mark.skipif(not _redis_available(), reason="本地 redis 不可用")
+def test_purge_matches_by_celery_message_id(redis_client):
+ app = Celery("test-purge-msg-id")
+ app.conf.broker_url = BROKER_URL
+ _publish(app, TEST_QUEUES[0], "celery-stale-id", "task-X")
+ _publish(app, TEST_QUEUES[0], "celery-good-id", "task-Y")
+
+ removed = purge_stale_messages_from_queues(BROKER_URL, TEST_QUEUES, celery_task_ids={"celery-stale-id"})
+ assert removed == 1
+ assert redis_client.llen(TEST_QUEUES[0]) == 1
+
+
+@pytest.mark.skipif(not _redis_available(), reason="本地 redis 不可用")
+def test_revoke_and_purge_calls_control_revoke(redis_client):
+ app = Celery("test-revoke")
+ app.conf.broker_url = BROKER_URL
+ app.control = MagicMock()
+ _publish(app, TEST_QUEUES[0], "celery-revoke-1", "task-R")
+
+ removed = revoke_and_purge(
+ app,
+ BROKER_URL,
+ business_task_ids={"task-R"},
+ celery_task_ids={"celery-revoke-1"},
+ queue_names=TEST_QUEUES,
+ )
+ assert removed == 1
+ app.control.revoke.assert_called_once_with("celery-revoke-1")
+
+
+def test_purge_empty_ids_is_noop():
+ assert purge_stale_messages_from_queues(BROKER_URL, TEST_QUEUES) == 0
+
+
+# ── 2. 执行前状态守卫 ──────────────────────────────────────────────────
+
+
+def test_guard_allows_pending():
+ status = ensure_task_claimable("t1", lambda _id: "pending", task_label="generation")
+ assert status == "pending"
+
+
+def test_guard_rejects_failed():
+ with pytest.raises(StaleTaskDiscarded) as exc:
+ ensure_task_claimable("t2", lambda _id: "failed", task_label="generation")
+ assert exc.value.task_id == "t2"
+ assert exc.value.status == "failed"
+
+
+def test_guard_rejects_cancelled_and_completed():
+ with pytest.raises(StaleTaskDiscarded):
+ ensure_task_claimable("t3", lambda _id: "cancelled")
+ with pytest.raises(StaleTaskDiscarded):
+ ensure_task_claimable("t4", lambda _id: "completed")
+
+
+def test_guard_missing_task_returns_empty():
+ assert ensure_task_claimable("t5", lambda _id: None) == ""
diff --git a/tests/unit/test_cleanup_ingest_beat_1714.py b/tests/unit/test_cleanup_ingest_beat_1714.py
new file mode 100644
index 000000000..f7c5bf11d
--- /dev/null
+++ b/tests/unit/test_cleanup_ingest_beat_1714.py
@@ -0,0 +1,75 @@
+"""#1714 beat 任务 scheduled_cleanup_stale_ingest_jobs 薄封装测试。
+
+mock SessionLocal 和清理核心,验证 beat 任务正确串联
+cleanup_stale_ingest_jobs → cleanup_orphan_processing_assets → revoke 消息。
+"""
+
+from __future__ import annotations
+
+import os
+import sys
+from pathlib import Path
+from unittest.mock import MagicMock, patch
+
+os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret")
+os.environ.setdefault("DATABASE_URL", "sqlite:///test_beat.db")
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "worker"))
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
+sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
+
+import worker_app.tasks.cleanup as cleanup # noqa: E402
+
+
+def test_beat_cleanup_calls_core_and_revokes():
+ """beat 任务串联三个核心步骤,返回汇总计数。"""
+ fake_session = MagicMock()
+
+ with (
+ patch("worker_app.db.SessionLocal", return_value=fake_session) as m_db,
+ patch(
+ "packages.application.ingest_orphan_cleanup.cleanup_stale_ingest_jobs",
+ return_value=([("job-1", "cel-1"), ("job-2", "")], ["a-1"]),
+ ) as m_jobs,
+ patch(
+ "packages.application.ingest_orphan_cleanup.cleanup_orphan_processing_assets",
+ return_value=["a-2"],
+ ) as m_assets,
+ patch(
+ "packages.shared.celery_orphan_guard.revoke_and_purge",
+ return_value=1,
+ ) as m_revoke,
+ ):
+ result = cleanup.scheduled_cleanup_stale_ingest_jobs()
+
+ m_db.assert_called_once()
+ m_jobs.assert_called_once()
+ assert m_jobs.call_args.kwargs["processing_timeout_minutes"] == 60
+ m_assets.assert_called_once()
+ m_revoke.assert_called_once()
+ # 队列名只传 transcode/celery(不传 generation)
+ assert m_revoke.call_args.kwargs["queue_names"] == ("transcode", "celery")
+ fake_session.close.assert_called_once()
+ assert result == {"stale_jobs": 2, "assets_to_error": 2, "purged_messages": 1}
+
+
+def test_beat_cleanup_no_op_when_nothing_stale():
+ """无孤儿时不调 revoke,返回全 0。"""
+ fake_session = MagicMock()
+
+ with (
+ patch("worker_app.db.SessionLocal", return_value=fake_session),
+ patch(
+ "packages.application.ingest_orphan_cleanup.cleanup_stale_ingest_jobs",
+ return_value=([], []),
+ ),
+ patch(
+ "packages.application.ingest_orphan_cleanup.cleanup_orphan_processing_assets",
+ return_value=[],
+ ),
+ patch("packages.shared.celery_orphan_guard.revoke_and_purge") as m_revoke,
+ ):
+ result = cleanup.scheduled_cleanup_stale_ingest_jobs()
+
+ m_revoke.assert_not_called()
+ assert result == {"stale_jobs": 0, "assets_to_error": 0, "purged_messages": 0}
diff --git a/tests/unit/test_dedup_1702_zero_rate_fix.py b/tests/unit/test_dedup_1702_zero_rate_fix.py
new file mode 100644
index 000000000..49f4e4871
--- /dev/null
+++ b/tests/unit/test_dedup_1702_zero_rate_fix.py
@@ -0,0 +1,548 @@
+"""Issue #1702 — 查重率恒为 0% 修复:单测.
+
+覆盖验收要求:
+1. 同源不同裁剪的两个视频能检出非 0 相似度(指纹中心裁剪绕开降重 + 阈值校准)
+2. 局部片段复用(B 结尾 2s ≈ A 中间 2s)能检出
+3. 异源视频不误报(相似度接近 0)
+4. N=1 现有流程不回归
+5. P1 确定性 bug:时长预过滤单位 /1000、直方图归一化、temporal_coverage 量纲、阈值比较统一
+6. P0:±1 邻接对齐、短视频自适应连续门槛
+7. P2:0 匹配也要落日志
+"""
+
+from __future__ import annotations
+
+import logging
+import sys
+from pathlib import Path
+from unittest.mock import MagicMock, patch
+
+import pytest
+
+sys.modules.setdefault("cv2", MagicMock())
+
+ROOT = Path(__file__).resolve().parents[2]
+sys.path.insert(0, str(ROOT / "apps" / "worker"))
+sys.path.insert(0, str(ROOT / "packages"))
+
+
+from video_processing.dedup import ( # noqa: E402
+ PHASH_THRESHOLD,
+ SEGMENT_MATCH_THRESHOLD,
+ FingerprintChunk,
+ VideoDeduplicator,
+ VideoFingerprint,
+ find_duplicate_segments,
+)
+
+# ── helpers ────────────────────────────────────────────────────
+
+
+def _h(d: int) -> str:
+ """64-bit phash with exactly d bits set vs zero hash."""
+ bits = ["0"] * 64
+ for i in range(d):
+ bits[i] = "1"
+ return f"{int(''.join(bits), 2):016x}"
+
+
+def _chunk(phash: str, t0: float, t1: float):
+
+ return FingerprintChunk(
+ start_time_ms=int(t0 * 1000),
+ end_time_ms=int(t1 * 1000),
+ phash_binary=phash,
+ color_histogram=[],
+ frame_count=1,
+ )
+
+
+def _fingerprint(phashes, duration, chunks=None, md5="fp-md5-x"):
+
+ return VideoFingerprint(
+ md5=md5,
+ keyframe_phashes=list(phashes),
+ color_histograms=[],
+ duration=duration,
+ resolution=(1280, 720),
+ chunks=chunks or [],
+ )
+
+
+def _video(vid, phashes, duration=10.0, project_id="proj1"):
+ from packages.domain import GeneratedVideo
+
+ return GeneratedVideo(
+ id=vid,
+ project_id=project_id,
+ generation_task_id=f"task-{vid}",
+ name=f"video-{vid}.mp4",
+ file_url=f"https://example.com/{vid}.mp4",
+ file_size=1000,
+ duration=duration,
+ width=1280,
+ height=720,
+ fps=25.0,
+ video_fingerprint={"md5": f"md5-{vid}", "keyframe_phashes": list(phashes)},
+ )
+
+
+def _rate(deduplicator, fp, videos, session=None):
+ session_magic = MagicMock()
+ # 分片表无数据 -> 回退 JSON keyframe_phashes
+ session_magic.query.return_value.filter.return_value.order_by.return_value.all.return_value = []
+ with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
+ repo = MockRepo.return_value
+ repo.list_by_project.return_value = videos
+ repo.list_by_user.return_value = videos
+ return deduplicator.compute_duplicate_rate(fp, "proj1", "new-vid", session_magic, scope="project")
+
+
+def _check(deduplicator, fp, videos, scope="project", **kw):
+ session_magic = MagicMock()
+ session_magic.query.return_value.filter.return_value.order_by.return_value.all.return_value = []
+ with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
+ repo = MockRepo.return_value
+ repo.list_by_project.return_value = videos
+ repo.list_by_user.return_value = videos
+ return deduplicator.check_duplicate(fp, "proj1", session_magic, scope=scope, **kw)
+
+
+# ── P0-1/P0-2: 同源不同裁剪(距离 6~10)检出非 0 ──────────────
+
+
+class TestSameSourceDifferentCrop:
+ """同源成片:random_edge_crop 后 pHash 距离 6~10,应检出非 0 相似度。"""
+
+ def test_same_source_high_similarity_detected(self):
+
+ ddp = VideoDeduplicator()
+ # 新视频 5 个分片,每个 phash 与已有视频对应分片距离 6(< 阈值)
+ base = [_h(0) for _ in range(5)]
+ new = [_h(6) for _ in range(5)]
+ existing = _video("v-old", base, duration=11.0)
+ chunks = [_chunk(h, i * 2.2, (i + 1) * 2.2) for i, h in enumerate(new)]
+ fp = _fingerprint(new, 11.0, chunks=chunks)
+
+ result = _rate(ddp, fp, [existing], MagicMock())
+ assert result["duplicate_rate"] > 0
+ assert result["visual_similarity"] > 0
+
+ def test_same_source_distance_at_threshold_still_detected(self):
+ """距离正好等于阈值(<=)也要算匹配——阈值比较统一为 <=。"""
+
+ assert PHASH_THRESHOLD <= 16, "阈值应经真实数据校准保持在能检出同源裁剪/降重对的范围(#1702 二次校准为 16)"
+ ddp = VideoDeduplicator()
+ base = [_h(0) for _ in range(6)]
+ new = [_h(PHASH_THRESHOLD) for _ in range(6)]
+ existing = _video("v-old", base, duration=12.0)
+ chunks = [_chunk(h, i * 2, (i + 1) * 2) for i, h in enumerate(new)]
+ fp = _fingerprint(new, 12.0, chunks=chunks)
+
+ result = _rate(ddp, fp, [existing], MagicMock())
+ assert result["duplicate_rate"] > 0
+
+
+# ── P0-2: 局部片段复用(B 结尾 2s ≈ A 中间 2s) ────────────────
+
+
+class TestPartialReuse:
+ def test_partial_reuse_tail_overlap_detected(self):
+ """新视频 6 片,最后 2 片命中已有视频中间 2 片(距离 4),其余不匹配。
+
+ 旧逻辑 frame_match_rate=2/6≈0.33(<0.3 硬跳过边界)+ MIN_CONSECUTIVE=5
+ 导致完全检不出;新逻辑 coverage 为主指标 + 自适应门槛应检出。
+ """
+
+ ddp = VideoDeduplicator()
+ # 已有 8 片:索引 3、4 是被复用的镜头
+ old = [_h(20 + i) for i in range(8)]
+ # 新视频 6 片:最后 2 片对应 old[3], old[4],距离 4;其余距离 30
+ new = [_h(50 + i) for i in range(4)] + [_h(4)] * 2
+ # 让 new[4] 与 old[3] 距离 4、new[5] 与 old[4] 距离 4(构造近似)
+ new[4] = f"{int('1' * 4 + '0' * 60, 2):016x}"
+ new[5] = f"{int('1' * 4 + '0' * 60, 2):016x}"
+ old[3] = _h(0)
+ old[4] = _h(0)
+
+ existing = _video("v-old", old, duration=16.0)
+ chunks = [_chunk(h, i * 2, (i + 1) * 2) for i, h in enumerate(new)]
+ fp = _fingerprint(new, 12.0, chunks=chunks)
+
+ result = _rate(ddp, fp, [existing], MagicMock())
+ # 局部复用:duplicate_rate 必须非 0
+ assert result["duplicate_rate"] > 0
+
+ def test_short_video_adaptive_consecutive_threshold(self):
+ """11s/5 片短视频:MIN_CONSECUTIVE 自适应 min(5, max(2, 5//2))=2,
+ 2 片连续命中即报片段(旧值 5 让短视频永远无法报片段)。"""
+
+ q = [
+ FingerprintChunk(0, 2000, "f" * 16, []),
+ FingerprintChunk(2000, 4000, "0" * 16, []),
+ FingerprintChunk(4000, 6000, f"{int('11110000', 2):016x}", []),
+ ]
+ t = [
+ FingerprintChunk(0, 2000, "f" * 16, []),
+ FingerprintChunk(2000, 4000, "0" * 16, []),
+ FingerprintChunk(4000, 6000, "e" * 16, []),
+ ]
+ # 3 片视频自适应门槛 = min(5, max(2, 3//2)) = 2
+ segs = find_duplicate_segments(q, t)
+ assert len(segs) >= 1
+
+
+# ── P0-3: ±1 邻接窗口对齐 ─────────────────────────────────────
+
+
+class TestNeighborAlignment:
+ def test_neighbor_window_absorbs_boundary_jitter(self):
+ """切点错位导致目标索引偏移 ±1 时,连续匹配不应被中断。"""
+
+ q = [FingerprintChunk(i * 1000, (i + 1) * 1000, f"{i:016x}", []) for i in range(4)]
+ # 目标:前 3 片与 q 相同,但第 3 片最佳匹配偏移 +1(t[4]),t[3] 是无关内容
+ t_hashes = [f"{i:016x}" for i in range(3)] + ["f" * 16, f"{3:016x}"]
+ t = [FingerprintChunk(i * 1000, (i + 1) * 1000, h, []) for i, h in enumerate(t_hashes)]
+ segs = find_duplicate_segments(q, t)
+ # q[0],q[1] 精确匹配 t[0],t[1];q[2]->t[2];q[3]->t[4](步进 2,窗口 ±1 内)
+ assert len(segs) >= 1
+ assert segs[0].query_end_ms >= 3000
+
+
+# ── P0-5 / 验收:异源不误报 ───────────────────────────────────
+
+
+class TestDifferentSourceNoFalsePositive:
+ def test_unrelated_videos_near_zero(self):
+
+ ddp = VideoDeduplicator()
+ # 异源:所有分片距离 >= 20
+ old = [_h(40 + i * 3 % 20) for i in range(6)]
+ new = [_h(0 + i) for i in range(6)]
+ existing = _video("v-old", old, duration=12.0)
+ chunks = [_chunk(h, i * 2, (i + 1) * 2) for i, h in enumerate(new)]
+ fp = _fingerprint(new, 12.0, chunks=chunks)
+
+ result = _rate(ddp, fp, [existing], MagicMock())
+ assert result["duplicate_rate"] == 0
+ assert result["visual_similarity"] < 0.7
+ assert result["match_count"] == 0
+
+ def test_check_duplicate_returns_none_for_unrelated(self):
+
+ ddp = VideoDeduplicator()
+ old = [_h(40 + i) for i in range(6)]
+ new = [_h(i) for i in range(6)]
+ existing = _video("v-old", old, duration=12.0)
+ fp = _fingerprint(new, 12.0)
+
+ result = _check(ddp, fp, [existing])
+ assert result is None
+
+
+# ── N=1 不回归 ────────────────────────────────────────────────
+
+
+class TestSingleChunkNoRegression:
+ def test_single_chunk_identical_detected(self):
+
+ ddp = VideoDeduplicator()
+ h = _h(2)
+ existing = _video("v-old", [h], duration=3.0)
+ chunks = [_chunk(h, 0, 3000)]
+ fp = _fingerprint([h], 3.0, chunks=chunks)
+ result = _rate(ddp, fp, [existing], MagicMock())
+ assert result["duplicate_rate"] > 0
+
+ def test_single_chunk_md5_exact_match(self):
+
+ ddp = VideoDeduplicator()
+ existing = _video("v-old", [_h(0)], duration=3.0)
+ existing.video_fingerprint["md5"] = "same"
+ fp = _fingerprint([_h(0)], 3.0, md5="same")
+ result = _check(ddp, fp, [existing])
+ assert result is not None
+ assert result["reason"] == "exact_md5_match"
+
+
+# ── P1-6: 时长预过滤单位 bug ──────────────────────────────────
+
+
+class TestDurationPrefilterUnit:
+ def test_user_scope_skips_duration_prefilter(self):
+ """Issue #1702: scope=user 跨项目查重不做 ±15% 时长预过滤。
+
+ 旧逻辑 duration/1000 单位 bug 先修成秒,但 ±15% 窗口与局部片段复用
+ 根本矛盾——复用片段的两个视频时长必然不同(证据视频 20s vs 11s 差 42%),
+ 窗口内找不到对方导致 is_duplicate 恒 False。最终口径:scope=user 全量
+ 遍历同用户视频(与 compute_duplicate_rate 一致),不传 duration_min/max。
+ """
+
+ ddp = VideoDeduplicator()
+ fp = _fingerprint([_h(0)], 13.5)
+ session_magic = MagicMock()
+ session_magic.query.return_value.filter.return_value.order_by.return_value.all.return_value = []
+ with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
+ repo = MockRepo.return_value
+ repo.list_by_user.return_value = []
+ ddp.check_duplicate(fp, "proj1", session_magic, scope="user", user_id="u1", duration_sec=fp.duration)
+ args, kwargs = repo.list_by_user.call_args
+ # 全量查询:不带任何时长过滤参数(局部复用必须跨时长比较)
+ assert "duration_min" not in kwargs
+ assert "duration_max" not in kwargs
+ assert args == ("u1",) or args == ()
+
+
+# ── P1-7: 颜色直方图归一化 ────────────────────────────────────
+
+
+class TestHistogramNormalization:
+ def test_bhattacharyya_coefficient_in_unit_range(self):
+ """Bhattacharyya 系数必须在 [0,1](旧 L2 + 3 通道拼接算出 ~14.9)。"""
+
+ # 3 通道拼接、每通道概率分布(Σ=1)
+ hist_a = [0.5, 0.5] + [0.0] * 94 + [0.5, 0.5] + [0.0] * 94 + [0.5, 0.5] + [0.0] * 94
+ # 长度裁剪到 96(3 通道 × 32 bins)
+ hist_a = ([0.5, 0.5] + [0.0] * 30) * 3
+ hist_b = ([0.5, 0.5] + [0.0] * 30) * 3
+
+ coeff = VideoDeduplicator._bhattacharyya_coefficient(hist_a, hist_b)
+ assert 0.0 <= coeff <= 1.0
+ assert coeff > 0.99 # 完全相同 -> 1.0
+
+ def test_bhattacharyya_disjoint_hist_low(self):
+
+ hist_a = ([1.0] + [0.0] * 31) * 3
+ hist_b = ([0.0] * 31 + [1.0]) * 3
+ coeff = VideoDeduplicator._bhattacharyya_coefficient(hist_a, hist_b)
+ assert coeff < 0.05
+
+
+# ── P1-8: temporal_coverage 量纲 ──────────────────────────────
+
+
+class TestTemporalCoverageUnits:
+ def test_coverage_uses_milliseconds(self):
+ """命中片段 6s / 视频 12s -> coverage=0.5;旧 bug 把 duration(秒)当毫秒,
+ covered_ms(6000)/duration(12) = 500 -> min(1.0)=1.0 误判 100% 覆盖。"""
+
+ ddp = VideoDeduplicator()
+ old = [_h(0) for _ in range(6)]
+ new = [_h(0) for _ in range(3)] + [_h(30) for _ in range(3)]
+ existing = _video("v-old", old, duration=12.0)
+ # 新视频 12s,前 6s(3 片)与 old 相同
+ chunks = [_chunk(h, i * 2, (i + 1) * 2) for i, h in enumerate(new)]
+ fp = _fingerprint(new, 12.0, chunks=chunks)
+ result = _rate(ddp, fp, [existing], MagicMock())
+ # coverage 应约 0.5(3 片 × 2s = 6s / 12s),duplicate_rate ≈ (0.5*0.4 + 0.5*0.6)*100 = 50
+ assert 30 < result["duplicate_rate"] < 70
+
+
+# ── P1-9: 阈值比较统一 ────────────────────────────────────────
+
+
+class TestThresholdConsistency:
+ def test_frame_and_segment_thresholds_same_source(self):
+
+ assert SEGMENT_MATCH_THRESHOLD == PHASH_THRESHOLD
+ assert VideoDeduplicator.PHASH_THRESHOLD == PHASH_THRESHOLD
+
+
+# ── P2: 0 匹配也要有日志痕迹 ──────────────────────────────────
+
+
+class TestZeroMatchLogging:
+ def test_no_match_emits_info_log(self, caplog):
+
+ ddp = VideoDeduplicator()
+ old = [_h(40 + i) for i in range(5)]
+ existing = _video("v-old", old, duration=10.0)
+ fp = _fingerprint([_h(i) for i in range(5)], 10.0)
+
+ with caplog.at_level(logging.INFO, logger="video_processing.dedup"):
+ result = _check(ddp, fp, [existing])
+ assert result is None
+ assert any("no match" in r.message for r in caplog.records)
+
+
+# ── recompute 任务下载路径(#1702 连带修复:旧硬编码 key 404) ─────
+
+
+class TestRecomputeDownloadPath:
+ """recompute-dedup 走 check_duplicate_task,需要从 OSS 重新下载成片。
+
+ 旧代码硬编码 projects/{pid}/generated/{vid}/{vid}.mp4(从不存在),
+ 真实 key 在 file_url:generated/projects/{pid}/tasks/{tid}/rendered_*.mp4。
+ """
+
+ def test_task_downloads_from_file_url(self):
+ import inspect
+
+ import video_processing.dedup as dedup_mod
+
+ source = inspect.getsource(dedup_mod.check_duplicate_task)
+ # 下载 key 必须来自 video.file_url
+ assert 'getattr(video, "file_url"' in source or "video.file_url" in source
+ # 旧的硬编码 key 只能作为回退存在,不能是主路径
+ assert "falling back to legacy key" in source
+ # download_file 接收的是派生 key 而非硬编码 f-string
+ assert "storage_service.download_file(download_key" in source
+ assert '/generated/{generated_video_id}/{generated_video_id}.mp4"' not in source.replace(
+ 'download_key = f"projects/{video.project_id}/generated/{generated_video_id}/{generated_video_id}.mp4"',
+ "",
+ )
+
+
+# ── check_duplicate 排除自身(#1702 连带修复:recompute 自匹配) ─────
+
+
+class TestCheckDuplicateExcludesSelf:
+ def test_exclude_video_id_skips_self_match(self):
+ """recompute 时当前视频已在候选列表:自匹配距离 0 分会让 duplicate_of
+ 指向自己。exclude_video_id 必须跳过自身,返回真实的其他匹配或 None。
+ """
+
+ ddp = VideoDeduplicator()
+ h = _h(0)
+ # 候选列表里同时放「自己」(完全相同)和一个异源视频
+ self_video = _video("v-self", [h], duration=10.0)
+ other_video = _video("v-other", [_h(40 + i) for i in range(3)], duration=10.0)
+ fp = _fingerprint([h], 10.0)
+
+ session = MagicMock()
+ session.query.return_value.filter.return_value.order_by.return_value.all.return_value = []
+
+ # 不传 exclude → 自匹配命中(错误行为复现)
+ with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
+ MockRepo.return_value.list_by_project.return_value = [self_video, other_video]
+ result = ddp.check_duplicate(fp, "proj1", session)
+ assert result is not None and result["duplicate_of"] == "v-self"
+
+ # 传 exclude_video_id → 跳过自己,异源不匹配 → None
+ with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
+ MockRepo.return_value.list_by_project.return_value = [self_video, other_video]
+ result = ddp.check_duplicate(fp, "proj1", session, exclude_video_id="v-self")
+ assert result is None
+
+ # 排除自己后,真实同源其他视频仍能检出
+ real_dup = _video("v-real", [h], duration=10.0)
+ with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
+ MockRepo.return_value.list_by_project.return_value = [self_video, real_dup]
+ result = ddp.check_duplicate(fp, "proj1", session, exclude_video_id="v-self")
+ assert result is not None and result["duplicate_of"] == "v-real"
+
+
+# ── 阈值 16 二次校准 + 时序抖动对齐(#1702 第二轮真实数据校准) ──────
+
+
+class TestThreshold16Calibration:
+ """二次校准:staging 15 个真实成片实测——同源降重对中位数距离 14、
+ <=16 命中 8/11=0.73;异源 13 个候选每帧全局最近邻最小距离 18、<=16
+ 命中全 0。阈值 16 检出同源且异源零误报(>=2bit 安全裕度)。"""
+
+ def test_threshold_calibrated_to_16(self):
+ assert PHASH_THRESHOLD == 16
+
+ @staticmethod
+ def _variant(phash: str, d: int) -> str:
+ """在 phash 基础上翻转恰好 d 个低位 bit → 与原哈希汉明距离恰为 d。"""
+ v = int(phash, 16)
+ for b in range(d):
+ v ^= 1 << b
+ return f"{v:016x}"
+
+ def test_distance_18_unrelated_not_matched(self):
+ """距离 18(异源实测最小最近邻距离)不判匹配,距离 16 判匹配。"""
+ ddp = VideoDeduplicator()
+ # 多样化 base(相邻帧各不相同,避免黑屏过滤器)
+ base = [_h(i + 4) for i in range(8)]
+ near = [self._variant(h, 16) for h in base] # 同源降重:每帧距离恰 16
+ far = [self._variant(h, 18) for h in base] # 异源边界:每帧距离恰 18
+
+ fp_near = _fingerprint(near, 8.0, chunks=[_chunk(h, i, i + 1) for i, h in enumerate(near)])
+ fp_far = _fingerprint(far, 8.0, chunks=[_chunk(h, i, i + 1) for i, h in enumerate(far)])
+
+ r_near = _rate(ddp, fp_near, [_video("v-base", base, duration=8.0)])
+ r_far = _rate(ddp, fp_far, [_video("v-base", base, duration=8.0)])
+
+ assert r_near["duplicate_rate"] > 0, "距离16的同源降重对必须检出"
+ assert r_far["duplicate_rate"] == 0.0, "距离18的异源对不得误报"
+ assert r_far["match_count"] == 0
+
+ def test_deduped_pair_frame_match_rate_over_threshold(self):
+ """真实场景比例:11 帧中 8 帧距离 <=16(0.73 >= 0.7),
+ 其余 3 帧异源距离(>=18)——frame_match_rate 必须过 0.7 门槛。"""
+ ddp = VideoDeduplicator()
+ base = [_h(i + 4) for i in range(11)]
+ near = [self._variant(h, 14) for h in base[:8]] # 中位数 14 的同源降重帧
+ # 异源帧用完全不同前缀(与 base 距离 >=30)
+ far = [_h(52 + i) for i in range(3)]
+ query = near + far
+
+ fp = _fingerprint(query, 11.0, chunks=[_chunk(h, i, i + 1) for i, h in enumerate(query)])
+ r = _rate(ddp, fp, [_video("v-base", base, duration=11.0)])
+ # frame_match_rate=8/11=0.73、时序片段覆盖 ~0.73
+ # → duplicate_rate = 0.4*0.73+0.6*0.73 ≈ 73%(空直方图回退下 fusion=0.6965
+ # 略低于 is_duplicate 的 0.70 判定阈值,故此处断言查重率而非 match_count;
+ # 真实视频带颜色直方图时 fusion≈0.80,staging A-C 实测 is_duplicate=True)
+ assert r["duplicate_rate"] >= 70.0
+
+
+class TestTemporalJitterAlignment:
+ """时序对齐允许目标索引正/反向 ±(neighbor_window+1) 抖动。
+
+ 密集 1s 采样下相邻帧 pHash 接近,全局最近邻会在目标相邻帧间
+ 正负 1 跳变(场景切割/取帧错位/局部倒退);旧逻辑只允许正向
+ delta,把同源连续匹配拆碎,min_consecutive 门槛够不上而漏检。
+ """
+
+ def test_backward_jitter_keeps_run_continuous(self):
+ """匹配目标索引序列 0,1,2,1,2,3(含一次 -1 倒退)应保持同一 run。"""
+ from video_processing.dedup import find_duplicate_segments
+
+ # 构造 target 相邻帧 pHash 相同(距离0),query 帧的最近邻在
+ # target[1]/target[2] 之间抖动;全部 <= 阈值
+ t_hash = _h(0)
+ other = _h(40)
+ # target: 帧0-3 相同场景,帧4+ 异源
+ t_chunks = [_chunk(t_hash, i, i + 1) for i in range(4)] + [_chunk(other, i, i + 1) for i in range(4, 8)]
+ # query 6 帧同场景(最近邻会落到 target 0~3,索引可正可负)
+ q_chunks = [_chunk(t_hash, i, i + 1) for i in range(6)]
+
+ segments = find_duplicate_segments(q_chunks, t_chunks)
+ assert segments, "含 ±1 时序抖动的连续匹配必须形成片段"
+ # 6 帧匹配 >= min_consecutive(min(5,max(2,6//2))=5),报为一个片段
+ assert len(segments) == 1
+ seg = segments[0]
+ assert seg.query_end_ms - seg.query_start_ms >= 5000
+
+ def test_large_backward_jump_breaks_run(self):
+ """目标索引倒退 > neighbor_window+1(如从 5 跳回 0)不属于抖动,
+ 不桥接为同一片段;孤立短匹配 < min_consecutive 不报片段。"""
+ from video_processing.dedup import find_duplicate_segments
+
+ # 异源段:9-bit 不重叠段(相邻段隔 3 bit),跨段距离 18~24 > 阈值 16
+ def _bit_seg(start):
+ bits = ["0"] * 64
+ for b in range(9):
+ bits[start + b] = "1"
+ return f"{int(''.join(bits), 2):016x}"
+
+ t_hash = _bit_seg(0) # 复用场景:bit 0-8
+ t_other = [_bit_seg(22 + 4 * i) for i in range(4)] # target 异源段
+ q_other = [_bit_seg(40 + 4 * i) for i in range(3)] # query 异源段
+ # target: 帧0 同场景;帧1-4 异源;帧5-6 同场景
+ t_chunks = (
+ [_chunk(t_hash, 0, 1)]
+ + [_chunk(t_other[i - 1], i, i + 1) for i in range(1, 5)]
+ + [_chunk(t_hash, i, i + 1) for i in range(5, 7)]
+ )
+ # query: 帧0 匹配 target[0];帧1-3 异源(与 target 任何帧距离 >16);帧4-5 匹配 target[5,6]
+ q_chunks = (
+ [_chunk(t_hash, 0, 1)]
+ + [_chunk(q_other[i - 1], i, i + 1) for i in range(1, 4)]
+ + [_chunk(t_hash, i, i + 1) for i in range(4, 6)]
+ )
+ segments = find_duplicate_segments(q_chunks, t_chunks)
+ # 两段各 1、2 帧 < min_consecutive=5 → 不报片段(大跳跃不桥接)
+ assert segments == []
diff --git a/tests/unit/test_dedup_engine.py b/tests/unit/test_dedup_engine.py
index e9ace015d..35ce57d0e 100644
--- a/tests/unit/test_dedup_engine.py
+++ b/tests/unit/test_dedup_engine.py
@@ -358,11 +358,11 @@ class TestVideoDeduplicatorCheckDuplicate:
finally:
self._restore_repo(mod, orig)
- def test_first_match_returned(self, deduplicator, mock_session):
- """返回第一个通过阈值的匹配(非最优匹配)。"""
- # vid-1: 距离=2 bits(0x03 XOR 0x01 = 0x02 → 1 bit),通过阈值
+ def test_highest_score_match_returned(self, deduplicator, mock_session):
+ """Issue #1702: 遍历所有候选取融合分最高者(旧逻辑首个过阈即返回)。"""
+ # vid-1: 距离=1 bit(0x03 XOR 0x01 = 0x02 → 1 bit),通过阈值
vid1 = self._make_existing_video("vid-1", "md5_1", phashes=["0000000000000003"])
- # vid-2: 距离=0 bits(完全匹配)
+ # vid-2: 距离=0 bits(完全匹配),融合分更高
vid2 = self._make_existing_video("vid-2", "md5_2", phashes=["0000000000000001"])
mock_repo = MagicMock()
@@ -380,8 +380,8 @@ class TestVideoDeduplicatorCheckDuplicate:
try:
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
assert result is not None
- # 返回第一个通过阈值的匹配(vid-1 距离=1 < 10)
- assert result["duplicate_of"] == "vid-1"
+ # 两个候选都过阈,返回融合分最高的 vid-2(距离 0 < 1)
+ assert result["duplicate_of"] == "vid-2"
finally:
self._restore_repo(mod, orig)
diff --git a/tests/unit/test_dedup_pure.py b/tests/unit/test_dedup_pure.py
index ab70cdb15..a80787daa 100755
--- a/tests/unit/test_dedup_pure.py
+++ b/tests/unit/test_dedup_pure.py
@@ -185,11 +185,14 @@ class TestBhattacharyyaCoefficient:
"""_bhattacharyya_coefficient Bhattacharyya 系数测试."""
def test_identical_histograms(self):
- """完全相同的直方图系数为1.0."""
- hist = [0.5, 0.5, 0.0, 0.3]
+ """完全相同的直方图系数为1.0(#1702:按 Σ 归一,概率分布语义)。"""
+ hist = [0.5, 0.5, 0.0, 0.0] # Σ=1 的概率分布
bc = VideoDeduplicator._bhattacharyya_coefficient(hist, hist)
- # Σ √(a[i]*a[i]) = Σ a[i] = 1.0 (normalized)
- assert bc == pytest.approx(sum(h for h in hist))
+ assert bc == pytest.approx(1.0)
+ # 非归一化输入也归一到 1.0(三通道拼接 Σ=3 的等价情形)
+ hist3 = [0.5, 0.5, 0.0, 0.3]
+ bc3 = VideoDeduplicator._bhattacharyya_coefficient(hist3, hist3)
+ assert bc3 == pytest.approx(1.0)
def test_zero_histograms(self):
"""全零直方图系数为0."""
@@ -202,10 +205,10 @@ class TestBhattacharyyaCoefficient:
assert bc == pytest.approx(0.0)
def test_different_lengths(self):
- """不同长度直方图取最小长度对齐."""
+ """不同长度直方图取最小长度对齐,并按各自总量归一(#1702 概率分布语义)。"""
+ # 对齐到前 2 维:coeff = 2,norm = √(Σa·Σb) = √(2·2) = 2 → 1.0
bc = VideoDeduplicator._bhattacharyya_coefficient([1.0, 1.0, 0.0, 0.0], [1.0, 1.0])
- # 对齐到前2维: √(1*1) + √(1*1) = 2.0
- assert bc == pytest.approx(2.0)
+ assert bc == pytest.approx(1.0)
def test_known_value(self):
"""已知值验证."""
diff --git a/tests/unit/test_dedup_v2.py b/tests/unit/test_dedup_v2.py
index 8fe0e8540..d0f7c6efe 100644
--- a/tests/unit/test_dedup_v2.py
+++ b/tests/unit/test_dedup_v2.py
@@ -103,6 +103,7 @@ from video_processing.dedup import ( # noqa: E402
MIN_CONSECUTIVE_MATCHES,
MIN_KEYFRAME_INTERVAL_SEC,
MIN_KEYFRAMES,
+ PHASH_THRESHOLD,
PHASH_WEIGHT,
SCENE_CHANGE_THRESHOLD,
SEGMENT_MATCH_THRESHOLD,
@@ -269,22 +270,17 @@ class TestFindDuplicateSegments:
注意:使用不同的 hash 对,确保后半部分帧距离 > 阈值。
"""
same_hash = "aaaaaaaaaaaaaaaa"
- # 4 帧匹配,后面 6 帧各自不同(在 query 和 target 中使用不同 hash)
+ # 4 帧匹配,后面 6 帧用与匹配哈希距离 32 的不匹配哈希(> PHASH_THRESHOLD=16)
+ nomatch_hash = "cccccccccccccccc" # hamming(aaaa, cccc)=32
chunks_a = [_make_chunk(i * 1000, (i + 1) * 1000, same_hash) for i in range(4)] + [
- _make_chunk(i * 1000, (i + 1) * 1000, "bbbbbbbbbbbbbbbb") for i in range(4, 10)
+ _make_chunk(i * 1000, (i + 1) * 1000, nomatch_hash) for i in range(4, 10)
]
chunks_b = [_make_chunk(i * 1000, (i + 1) * 1000, same_hash) for i in range(4)] + [
- _make_chunk(i * 1000, (i + 1) * 1000, "cccccccccccccccc") for i in range(4, 10)
+ _make_chunk(i * 1000, (i + 1) * 1000, nomatch_hash) for i in range(4, 10)
]
-
- # hamming("bbbb...", "cccc...") should be > 8 (SEGMENT_MATCH_THRESHOLD)
- # b=1011, c=1100 → 4 bits differ per hex digit × 16 digits = 64 bits total? No...
- # Actually: hamming_distance("bbbbbbbbbbbbbbbb", "cccccccccccccccc")
- # b=0xb=1011, c=0xc=1100 → XOR=0111=0x7 → 3 bits per digit × 16 = 48
- # That's > 8 so won't match
-
+ # hamming(aaaa..., cccc...) = 32 > PHASH_THRESHOLD(16),后半段不匹配;
+ # 前 4 帧匹配 < min_consecutive=5,不形成片段
segments = find_duplicate_segments(chunks_a, chunks_b)
- # 只有 4 帧匹配(< min_consecutive=5),所以不报告
assert segments == []
def test_max_gap_behavior(self):
@@ -293,10 +289,12 @@ class TestFindDuplicateSegments:
关键:间隙帧必须在 query 和 target 中使用不同 hash,使其真正不匹配。
"""
match_hash = "aaaaaaaaaaaaaaaa"
- gap_hash_a = "bbbbbbbbbbbbbbbb" # query 端
- gap_hash_b = "cccccccccccccccc" # target 端(与 query 端距离 > 8)
- tail_hash_a = "dddddddddddddddd"
- tail_hash_b = "eeeeeeeeeeeeeeee"
+ # 间隙/尾部哈希与 match_hash 及彼此之间汉明距离均 >64 (> PHASH_THRESHOLD=16),
+ # 确保在 ±(neighbor_window+1) 时序抖动对齐窗口内也不会误匹配
+ gap_hash_a = "ffffffffffffffff" # hamming(a,f)=128
+ gap_hash_b = "9999999999999999" # hamming(a,9)=128, hamming(f,9)=128
+ tail_hash_a = "7777777777777777" # hamming(a,7)=192
+ tail_hash_b = "1111111111111111" # hamming(a,1)=192, hamming(7,1)=128
# 5 帧匹配, 1 帧间隙, 3 帧匹配, 5 帧不匹配
hashes_a = [match_hash] * 5 + [gap_hash_a] + [match_hash] * 3 + [tail_hash_a] * 5
@@ -318,10 +316,10 @@ class TestFindDuplicateSegments:
def test_max_gap_exceeded(self):
"""间隙超过 max_gap → 分成两段."""
match_hash = "aaaaaaaaaaaaaaaa"
- gap_hash_a = "bbbbbbbbbbbbbbbb"
- gap_hash_b = "cccccccccccccccc"
- tail_hash_a = "dddddddddddddddd"
- tail_hash_b = "eeeeeeeeeeeeeeee"
+ gap_hash_a = "ffffffffffffffff" # hamming(a,f)=128
+ gap_hash_b = "9999999999999999" # hamming(a,9)=128
+ tail_hash_a = "7777777777777777" # hamming(a,7)=192
+ tail_hash_b = "1111111111111111" # hamming(a,1)=192
# 5 帧匹配, 3 帧间隙 (> max_gap=2), 5 帧匹配, 5 帧不匹配
hashes_a = [match_hash] * 5 + [gap_hash_a] * 3 + [match_hash] * 5 + [tail_hash_a] * 5
@@ -484,8 +482,10 @@ class TestBackwardCompatibility:
chunks_b = [{"phash_binary": "aaaaaaaaaaaaaaaa", "start_time_ms": 0, "end_time_ms": 5000}]
segments = find_duplicate_segments(chunks_a, chunks_b)
- # 1 帧 < min_consecutive=5,不会报重复
- assert segments == []
+ # Issue #1702: 自适应门槛 min(5, max(2, 1//2))=2,1 帧不成段;
+ # N=1 的检出由 _evaluate_candidate 匹配帧回退兜底(见 test_dedup_1702)。
+ # 这里只要求不崩溃。
+ assert isinstance(segments, list)
# ── TestConstants ───────────────────────────────────────────────
@@ -495,8 +495,11 @@ class TestConstants:
"""常量值验证 — 使用已在模块顶部导入的常量,避免重新 import."""
def test_segment_match_threshold(self):
- # 从已导入的 find_duplicate_segments 默认参数间接验证
- assert SEGMENT_MATCH_THRESHOLD == 8
+ # Issue #1702 二次校准:阈值经 staging 真实数据两轮回归——
+ # 第一轮同源 4/11、异源 min=24 定 12;第二轮扩样本(15 个真实成片)
+ # 同源降重对中位数距离 14、<=16 命中 8/11=0.73,异源 13 个候选
+ # <=16 命中全 0、最近邻最小距离 18 → 校准为 16。
+ assert SEGMENT_MATCH_THRESHOLD == PHASH_THRESHOLD == 16
def test_min_consecutive_matches(self):
assert MIN_CONSECUTIVE_MATCHES == 5
diff --git a/tests/unit/test_duplicate_rate_scope.py b/tests/unit/test_duplicate_rate_scope.py
index 4cf6f02ce..6fd9a0d08 100644
--- a/tests/unit/test_duplicate_rate_scope.py
+++ b/tests/unit/test_duplicate_rate_scope.py
@@ -132,9 +132,14 @@ class TestCheckDuplicateScopeUser:
class TestDurationPrefilter:
- """test_duration_prefilter:时长 ±15% 过滤."""
+ """Issue #1702: scope=user 跨项目查重不做时长预过滤。
- def test_duration_prefilter_passes_correct_range(self):
+ 局部片段复用的两个视频时长必然不同(证据视频 20s vs 11s,差 42%),
+ 旧的 ±15% 窗口会让同源视频互相不可见 → is_duplicate 恒 False。
+ 全量遍历同用户视频,异源视频由 fusion/temporal_coverage 阈值天然过滤。
+ """
+
+ def test_user_scope_no_duration_filter(self):
from video_processing.dedup import VideoDeduplicator
deduplicator = VideoDeduplicator()
@@ -153,12 +158,13 @@ class TestDurationPrefilter:
duration_sec=30.0,
)
- # Should pass duration_min=25.5, duration_max=34.5 (30 ± 15%)
+ # scope=user 全量遍历:位置参数只传 user_id,kwargs 不含时长过滤
call_args = mock_repo.list_by_user.call_args
- assert call_args[1]["duration_min"] == pytest.approx(25.5, abs=0.1)
- assert call_args[1]["duration_max"] == pytest.approx(34.5, abs=0.1)
+ assert call_args[0] == ("user1",)
+ assert "duration_min" not in call_args[1]
+ assert "duration_max" not in call_args[1]
- def test_no_duration_prefilter_when_zero(self):
+ def test_user_scope_no_duration_filter_when_zero(self):
from video_processing.dedup import VideoDeduplicator
deduplicator = VideoDeduplicator()
@@ -178,8 +184,34 @@ class TestDurationPrefilter:
)
call_args = mock_repo.list_by_user.call_args
- assert call_args[1]["duration_min"] == 0
- assert call_args[1]["duration_max"] == 0
+ assert "duration_min" not in call_args[1]
+ assert "duration_max" not in call_args[1]
+
+ def test_project_scope_also_no_duration_filter(self):
+ """scope=project 走 list_by_project,本来就不做时长过滤。"""
+ from video_processing.dedup import VideoDeduplicator
+
+ deduplicator = VideoDeduplicator()
+ fingerprint = _make_fingerprint(duration_ms=30000)
+ session = MagicMock()
+
+ with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
+ mock_repo = MockRepo.return_value
+ mock_repo.list_by_project.return_value = []
+ deduplicator.check_duplicate(
+ fingerprint,
+ "proj1",
+ session,
+ scope="project",
+ user_id="user1",
+ duration_sec=30.0,
+ )
+
+ mock_repo.list_by_project.assert_called_once()
+ call_args = mock_repo.list_by_project.call_args
+ assert call_args[0] == ("proj1",)
+ assert "duration_min" not in call_args[1]
+ assert "duration_max" not in call_args[1]
class TestComputeDuplicateRateFormula:
diff --git a/tests/unit/test_enqueue_persists_celery_id_1714.py b/tests/unit/test_enqueue_persists_celery_id_1714.py
new file mode 100644
index 000000000..09ecd1949
--- /dev/null
+++ b/tests/unit/test_enqueue_persists_celery_id_1714.py
@@ -0,0 +1,57 @@
+"""#1714:入队成功后 celery 消息 ID 必须持久化到任务行(供清理时 revoke)。"""
+
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+from unittest.mock import MagicMock
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
+
+from app.core import task_enqueue # noqa: E402
+
+
+class _FakeTask:
+ def __init__(self):
+ self.id = "task-enqueue-1"
+ self.status = "pending"
+ self.celery_task_id = ""
+
+ def mark_failed(self, msg): # noqa: ARG002
+ self.status = "failed"
+
+
+class _FakeRepo:
+ def __init__(self):
+ self.updated = None
+
+ def count_pending_total(self):
+ return 0
+
+ def count_pending_by_user(self, user_id): # noqa: ARG002
+ return 0
+
+ def update(self, task):
+ self.updated = task
+ return task
+
+
+def test_safe_enqueue_persists_celery_message_id(monkeypatch):
+ fake_result = MagicMock()
+ fake_result.id = "celery-msg-id-enqueue-999"
+ mock_celery = MagicMock()
+ mock_celery.send_task.return_value = fake_result
+ monkeypatch.setattr(task_enqueue, "celery_app", mock_celery)
+
+ task = _FakeTask()
+ repo = _FakeRepo()
+
+ ok = task_enqueue.safe_enqueue_generation_task(task, repo, user_id="u1")
+ assert ok is True
+ # celery_task_id 已持久化
+ assert task.celery_task_id == "celery-msg-id-enqueue-999"
+ assert repo.updated is task
+ mock_celery.send_task.assert_called_once()
+ args, kwargs = mock_celery.send_task.call_args
+ assert args[0] == "worker.generate_video"
+ assert kwargs.get("args") == [task.id]
diff --git a/tests/unit/test_fingerprint_chunks.py b/tests/unit/test_fingerprint_chunks.py
index 12635d72b..ccaf649b5 100644
--- a/tests/unit/test_fingerprint_chunks.py
+++ b/tests/unit/test_fingerprint_chunks.py
@@ -3,7 +3,7 @@
覆盖:
- 分片策略:60秒视频 → 30片,120秒视频 → 24片
- VideoFingerprint.to_chunk_models() 输出正确
-- _save_fingerprint_chunks 幂等性(已有数据跳过)
+- _save_fingerprint_chunks 替换语义(Issue #1702:重算时先删旧分片再写入)
- to_dict() 向后兼容
"""
@@ -169,11 +169,15 @@ class TestVideoFingerprintToChunkModels:
assert models == []
-class TestSaveFingerprintChunksIdempotent:
- """测试 _save_fingerprint_chunks 幂等性。"""
+class TestSaveFingerprintChunksReplace:
+ """测试 _save_fingerprint_chunks 替换语义(Issue #1702)。
- def test_save_skips_existing(self):
- """已有分片数据时跳过写入。"""
+ 重算查重时指纹算法已升级(中心裁剪 + 新采样/阈值),旧分片必须先删除
+ 再写入新分片,否则 recompute-dedup 永远读到旧指纹、修复对存量视频不生效。
+ """
+
+ def test_save_replaces_existing(self):
+ """已有分片数据时:先删除旧分片,再写入新分片。"""
fp = VideoFingerprint(
md5="abc",
keyframe_phashes=["a1b2"],
@@ -186,16 +190,22 @@ class TestSaveFingerprintChunksIdempotent:
)
session = MagicMock()
- # Mock: 已有 1 条分片数据
- session.query.return_value.filter.return_value.count.return_value = 1
+ # Mock: 删除旧分片返回 3(旧算法留下的 3 条分片)
+ session.query.return_value.filter.return_value.delete.return_value = 3
_save_fingerprint_chunks(fp, video_id="v1", project_id="p1", user_id="u1", session=session)
- # bulk_save_objects 不应被调用
- session.bulk_save_objects.assert_not_called()
+ # 必须先执行删除
+ session.query.return_value.filter.return_value.delete.assert_called_once()
+ # 新分片必须写入
+ session.bulk_save_objects.assert_called_once()
+ saved_models = session.bulk_save_objects.call_args[0][0]
+ assert len(saved_models) == 1
+ assert saved_models[0].video_id == "v1"
+ assert saved_models[0].phash_binary == "a1b2"
def test_save_writes_new(self):
- """无分片数据时写入。"""
+ """无旧分片时直接写入。"""
fp = VideoFingerprint(
md5="abc",
keyframe_phashes=["a1b2"],
@@ -208,12 +218,12 @@ class TestSaveFingerprintChunksIdempotent:
)
session = MagicMock()
- # Mock: 无分片数据
- session.query.return_value.filter.return_value.count.return_value = 0
+ # Mock: 无旧分片
+ session.query.return_value.filter.return_value.delete.return_value = 0
_save_fingerprint_chunks(fp, video_id="v1", project_id="p1", user_id="u1", session=session)
- # bulk_save_objects 应被调用一次
+ session.query.return_value.filter.return_value.delete.assert_called_once()
session.bulk_save_objects.assert_called_once()
saved_models = session.bulk_save_objects.call_args[0][0]
assert len(saved_models) == 1
@@ -221,7 +231,7 @@ class TestSaveFingerprintChunksIdempotent:
assert saved_models[0].phash_binary == "a1b2"
def test_save_skips_no_chunks(self):
- """指纹无 chunks 时跳过。"""
+ """指纹无 chunks 时跳过(不删不写)。"""
fp = VideoFingerprint(
md5="abc",
keyframe_phashes=[],
@@ -232,11 +242,11 @@ class TestSaveFingerprintChunksIdempotent:
)
session = MagicMock()
- session.query.return_value.filter.return_value.count.return_value = 0
_save_fingerprint_chunks(fp, video_id="v1", project_id="p1", user_id="u1", session=session)
- # bulk_save_objects 不应被调用
+ # 无 chunks:不查询、不删除、不写入
+ session.query.assert_not_called()
session.bulk_save_objects.assert_not_called()
diff --git a/tests/unit/test_generation_preview.py b/tests/unit/test_generation_preview.py
index 696e889aa..69a3444bc 100644
--- a/tests/unit/test_generation_preview.py
+++ b/tests/unit/test_generation_preview.py
@@ -725,8 +725,12 @@ class TestCreatePreviewRoute:
generation_task_repository=repo,
db=MagicMock(),
)
- assert resp.task_id == "preview_task_001"
- assert resp.status == "pending"
+ # 批量响应:N=1 时 items 长度为 1
+ assert resp.total == 1
+ assert len(resp.items) == 1
+ assert resp.items[0].task_id == "preview_task_001"
+ assert resp.items[0].status == "pending"
+ assert resp.items[0].variant_index == 0
def test_user_pending_limit_exceeded(self):
"""用户待处理任务超限 → 429"""
@@ -807,7 +811,13 @@ class TestCreatePreviewRoute:
repo.count_pending_total.return_value = 0
task = _make_task()
- from fastapi import HTTPException
+
+ # 模拟 mark_failed 真实更新任务状态(_mark_task_failed 内部调用)
+ def _set_failed(error_message="", **_kwargs):
+ task.status = GenerationTaskStatus.FAILED
+ task.error_message = error_message
+
+ task.mark_failed.side_effect = _set_failed
with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
MockUC.return_value.execute.return_value = task
@@ -815,14 +825,15 @@ class TestCreatePreviewRoute:
"app.api.routes.generation_preview.safe_enqueue_generation_task",
return_value=False,
):
- with pytest.raises(HTTPException) as exc_info:
- create_preview_generation_task(
- self._make_request(),
- authenticated_user=_make_user(),
- generation_task_repository=repo,
- db=MagicMock(),
- )
- assert exc_info.value.status_code == 500
+ resp = create_preview_generation_task(
+ self._make_request(),
+ authenticated_user=_make_user(),
+ generation_task_repository=repo,
+ db=MagicMock(),
+ )
+ # 入队失败:任务被标记 failed(mark_failed 设置错误信息),响应正常返回
+ assert resp.total == 1
+ assert resp.items[0].status == "failed"
def test_enqueue_raises_user_limit(self):
"""safe_enqueue 抛出 UserPendingLimitExceeded → 429"""
@@ -833,6 +844,12 @@ class TestCreatePreviewRoute:
task = _make_task()
from fastapi import HTTPException
+ def _set_failed_limit(error_message="", **_kwargs):
+ task.status = GenerationTaskStatus.FAILED
+ task.error_message = error_message or "待处理任务超限"
+
+ task.mark_failed.side_effect = _set_failed_limit
+
with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
MockUC.return_value.execute.return_value = task
with patch(
@@ -846,6 +863,7 @@ class TestCreatePreviewRoute:
generation_task_repository=repo,
db=MagicMock(),
)
+ # 全部变体入队失败且错误消息含"待处理任务" → 429
assert exc_info.value.status_code == 429
def test_enqueue_raises_global_queue_full(self):
@@ -857,6 +875,12 @@ class TestCreatePreviewRoute:
task = _make_task()
from fastapi import HTTPException
+ def _set_failed_queue(error_message="", **_kwargs):
+ task.status = GenerationTaskStatus.FAILED
+ task.error_message = error_message or "系统队列已满"
+
+ task.mark_failed.side_effect = _set_failed_queue
+
with patch("app.api.routes.generation_preview.CreateGenerationTaskUseCase") as MockUC:
MockUC.return_value.execute.return_value = task
with patch(
@@ -870,6 +894,7 @@ class TestCreatePreviewRoute:
generation_task_repository=repo,
db=MagicMock(),
)
+ # 全部变体入队失败且错误消息含"队列" → 503
assert exc_info.value.status_code == 503
diff --git a/tests/unit/test_ingest_hevc_orphan_1714.py b/tests/unit/test_ingest_hevc_orphan_1714.py
new file mode 100644
index 000000000..3701ea381
--- /dev/null
+++ b/tests/unit/test_ingest_hevc_orphan_1714.py
@@ -0,0 +1,317 @@
+"""Issue #1714:HEVC 转码后禁止兜底新建重复 READY 记录,必须回写占位 asset。
+
+覆盖:
+- 转码成功 + 占位 asset 存在(按原始 key 找到)→ 更新占位为 READY、
+ storage_key 改写为 *_h264,绝不 create 新记录(回归 P1 孤儿 PROCESSING bug)
+- job.asset_id 透传时优先按 id 关联占位(即使 key 对不上也能命中)
+- 无占位记录(旧链路)→ 兜底新建(保留兼容)
+- 非 HEVC:占位同样被更新为 READY,不新建
+- 无效媒体:占位标记为 ERROR,不新建 ERROR 记录
+- ingest 异常:占位(按还原后的原始 key)标记 ERROR
+"""
+
+from __future__ import annotations
+
+import sys
+import tempfile
+from pathlib import Path
+from types import SimpleNamespace
+from unittest.mock import MagicMock, patch
+
+# ── 与 test_ingest_hevc_transcode_task.py 相同的 worker 模块加载方式 ──
+_SAVED_MODULES_KEYS = set(sys.modules.keys())
+
+_mock_db_module = MagicMock()
+_mock_db_module.SessionLocal = MagicMock()
+sys.modules["worker_app.db"] = _mock_db_module
+sys.modules["worker_app.core.config"] = MagicMock()
+
+_mock_celery_module = MagicMock()
+
+
+def _passthrough_decorator(*args, **kwargs):
+ if len(args) == 1 and callable(args[0]):
+ return args[0]
+ return lambda f: f
+
+
+_mock_celery_module.celery_app.task = MagicMock(side_effect=_passthrough_decorator)
+sys.modules["worker_app.celery_app"] = _mock_celery_module
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "worker"))
+
+import pytest # noqa: E402
+from worker_app.tasks import ingest as ingest_mod # noqa: E402
+
+from packages.domain import Asset, AssetStatus # noqa: E402
+
+for _key in list(sys.modules.keys()):
+ if _key not in _SAVED_MODULES_KEYS and not _key.startswith("video_processing"):
+ del sys.modules[_key]
+del _SAVED_MODULES_KEYS
+
+
+# ── 假仓储 ──────────────────────────────────────────────────────────────
+class _FakeJobRepo:
+ def __init__(self, job):
+ self.job = job
+ self.updated = None
+
+ def get(self, job_id):
+ return self.job
+
+ def update(self, job):
+ self.updated = job
+ return job
+
+
+class _FakeAssetRepo:
+ """记录 create 调用;find_* 按内部 assets 列表查询。"""
+
+ def __init__(self, assets: list[Asset] | None = None):
+ self.assets = list(assets or [])
+ self.created: list[Asset] = []
+ self.updated: list[Asset] = []
+
+ def create(self, asset: Asset) -> Asset:
+ self.created.append(asset)
+ self.assets.append(asset)
+ return asset
+
+ def update(self, asset: Asset) -> Asset:
+ self.updated.append(asset)
+ return asset
+
+ def find_by_id(self, asset_id: str) -> Asset | None:
+ return next((a for a in self.assets if a.id == asset_id), None)
+
+ def find_by_storage_key(self, storage_key: str) -> Asset | None:
+ return next((a for a in self.assets if a.storage_key == storage_key), None)
+
+
+def _make_job(asset_id: str = "", storage_key: str = "uploads/proj/IMG_2282.MOV"):
+ return SimpleNamespace(
+ id="job-1",
+ project_id="proj-1",
+ library_id="lib-1",
+ storage_key=storage_key,
+ file_hash="hash-1",
+ asset_id=asset_id,
+ status=None,
+ error_message=None,
+ result_asset_id=None,
+ updated_at=None,
+ )
+
+
+def _make_placeholder(storage_key: str = "uploads/proj/IMG_2282.MOV", asset_id: str = "asset-ph"):
+ return Asset(
+ id=asset_id,
+ project_id="proj-1",
+ library_id="lib-1",
+ name="IMG_2282.MOV",
+ storage_key=storage_key,
+ mime_type="video/quicktime",
+ status=AssetStatus.PROCESSING,
+ file_hash="hash-1",
+ )
+
+
+def _video_metadata(codec="hevc"):
+ return {
+ "codec": codec,
+ "width": 1920,
+ "height": 1080,
+ "duration": 10.0,
+ "size_bytes": 5 * 1024 * 1024,
+ }
+
+
+@pytest.fixture
+def transcode_env(tmp_path):
+ """HEVC 转码成功的标准 mock 环境(同 test_ingest_hevc_transcode_task)。"""
+ local_file = tmp_path / "local_hevc.MOV"
+ local_file.write_bytes(b"fake-hevc-source")
+ tc_out = tmp_path / "transcode_out_h264.mp4"
+
+ control = {
+ "validate_ok": True,
+ "tc_out": tc_out,
+ "local_file": local_file,
+ "download_ok": True,
+ "extract_success": True,
+ "codec": "hevc",
+ "raise_in_flow": None,
+ }
+
+ def fake_ntf(*args, **kwargs):
+ mock_file = MagicMock()
+ mock_file.name = str(tc_out) if kwargs.get("suffix") == "_h264.mp4" else str(local_file)
+ mock_file.close = MagicMock()
+ mock_file.__enter__.return_value = mock_file
+ mock_file.__exit__.return_value = False
+ return mock_file
+
+ def fake_subprocess_run(cmd, **kwargs):
+ if cmd and cmd[0] == "ffmpeg" and "libx264" in cmd:
+ Path(cmd[-1]).write_bytes(b"fake-h264-output")
+ return SimpleNamespace(returncode=0, stderr="")
+ return SimpleNamespace(returncode=0, stdout="", stderr="")
+
+ control["patchers"] = {
+ "session": patch.object(ingest_mod, "SessionLocal", return_value=MagicMock()),
+ "download": patch.object(ingest_mod, "download_asset", side_effect=lambda *a, **kw: control["download_ok"]),
+ "upload": patch("video_processing.oss_helpers.upload_to_oss", return_value="https://oss/x"),
+ "metadata": patch.object(
+ ingest_mod,
+ "extract_media_metadata",
+ side_effect=lambda path, mt: (
+ (_video_metadata("h264"), control["extract_success"])
+ if Path(path).name == tc_out.name
+ else (_video_metadata(control["codec"]), control["extract_success"])
+ ),
+ ),
+ "validate": patch.object(
+ ingest_mod, "validate_transcode_output", side_effect=lambda p, portrait: control["validate_ok"]
+ ),
+ "subprocess": patch.object(ingest_mod.subprocess, "run", side_effect=fake_subprocess_run),
+ "ntf": patch.object(tempfile, "NamedTemporaryFile", side_effect=fake_ntf),
+ "thumb": patch(
+ "video_processing.thumbnail_generator.extract_first_frame",
+ side_effect=RuntimeError("skip thumb"),
+ ),
+ }
+ return control
+
+
+def _start(control, job, assets):
+ job_repo = _FakeJobRepo(job)
+ asset_repo = _FakeAssetRepo(assets)
+ patchers = dict(control["patchers"])
+ patchers["job_repo"] = patch.object(ingest_mod, "SQLAlchemyIngestJobRepository", return_value=job_repo)
+ patchers["asset_repo"] = patch.object(ingest_mod, "SQLAlchemyAssetRepository", return_value=asset_repo)
+ started = {name: p.start() for name, p in patchers.items()}
+ return started, job_repo, asset_repo
+
+
+def _stop(control):
+ for p in control["patchers"].values():
+ p.stop()
+
+
+class TestHEVCTranscodePlaceholderRewrite:
+ def test_transcode_success_updates_placeholder_no_duplicate_ready(self, transcode_env):
+ """转码成功 → 占位 asset 原地更新为 READY + storage_key 改写 _h264,禁止新建。"""
+ control = transcode_env
+ placeholder = _make_placeholder()
+ job = _make_job() # 旧 job 无 asset_id,靠原始 key 关联
+ mocks, job_repo, asset_repo = _start(control, job, [placeholder])
+ try:
+ result = ingest_mod.ingest_asset("job-1")
+ finally:
+ _stop(control)
+
+ assert result["status"] == "completed"
+ # 核心断言 1:没有新建任何 READY 记录(旧 bug 会 create 一条 _h264 READY)
+ assert asset_repo.created == [], "转码回写不得新建 asset 记录"
+ # 核心断言 2:占位被更新为 READY,且 storage_key 已是 _h264
+ assert len(asset_repo.updated) == 1
+ updated = asset_repo.updated[0]
+ assert updated.id == placeholder.id
+ assert updated.status == AssetStatus.READY
+ assert updated.storage_key == "uploads/proj/IMG_2282_h264.MOV"
+ assert updated.metadata.get("hevc_transcoded") is True
+ assert updated.metadata.get("original_storage_key") == "uploads/proj/IMG_2282.MOV"
+ # job 关联到同一条 asset
+ assert job_repo.updated.result_asset_id == placeholder.id
+ assert job_repo.updated.storage_key == "uploads/proj/IMG_2282_h264.MOV"
+
+ def test_placeholder_resolved_by_job_asset_id(self, transcode_env):
+ """job.asset_id 透传时优先按 id 关联(即使 storage_key 对不上也命中)。"""
+ control = transcode_env
+ placeholder = _make_placeholder(storage_key="uploads/different/key.MOV", asset_id="asset-by-id")
+ job = _make_job(asset_id="asset-by-id")
+ _, _, asset_repo = _start(control, job, [placeholder])
+ try:
+ result = ingest_mod.ingest_asset("job-1")
+ finally:
+ _stop(control)
+
+ assert result["status"] == "completed"
+ assert asset_repo.created == []
+ assert len(asset_repo.updated) == 1
+ assert asset_repo.updated[0].id == "asset-by-id"
+ assert asset_repo.updated[0].status == AssetStatus.READY
+
+ def test_no_placeholder_fallback_creates_ready(self, transcode_env):
+ """旧链路无占位记录 → 兜底新建 READY(兼容保留,但必须是唯一一条)。"""
+ control = transcode_env
+ job = _make_job(asset_id="")
+ _, _, asset_repo = _start(control, job, [])
+ try:
+ result = ingest_mod.ingest_asset("job-1")
+ finally:
+ _stop(control)
+
+ assert result["status"] == "completed"
+ assert len(asset_repo.created) == 1
+ created = asset_repo.created[0]
+ assert created.status == AssetStatus.READY
+ assert created.storage_key == "uploads/proj/IMG_2282_h264.MOV"
+ assert asset_repo.updated == []
+
+ def test_non_hevc_placeholder_updated_no_create(self, transcode_env):
+ """非 HEVC(h264)不转码:占位按原始 key 找到并更新 READY,不新建。"""
+ control = transcode_env
+ control["codec"] = "h264"
+ placeholder = _make_placeholder()
+ job = _make_job()
+ mocks, _, asset_repo = _start(control, job, [placeholder])
+ try:
+ result = ingest_mod.ingest_asset("job-1")
+ finally:
+ _stop(control)
+
+ assert result["status"] == "completed"
+ assert asset_repo.created == []
+ assert len(asset_repo.updated) == 1
+ updated = asset_repo.updated[0]
+ assert updated.status == AssetStatus.READY
+ assert updated.storage_key == "uploads/proj/IMG_2282.MOV" # 未转码,key 不变
+ mocks["upload"].assert_not_called()
+
+ def test_invalid_media_marks_placeholder_error_no_create(self, transcode_env):
+ """无效媒体:占位标记 ERROR 并 update,禁止再 create 一条 ERROR。"""
+ control = transcode_env
+ control["download_ok"] = False # 下载失败 → extract_success=False → 无效媒体路径
+ placeholder = _make_placeholder()
+ job = _make_job()
+ _, job_repo, asset_repo = _start(control, job, [placeholder])
+ try:
+ result = ingest_mod.ingest_asset("job-1")
+ finally:
+ _stop(control)
+
+ assert result["status"] == "failed"
+ assert asset_repo.created == [], "无效媒体不得新建 ERROR 记录"
+ assert len(asset_repo.updated) == 1
+ assert asset_repo.updated[0].id == placeholder.id
+ assert asset_repo.updated[0].status == AssetStatus.ERROR
+ assert job_repo.updated.result_asset_id == placeholder.id
+
+ def test_exception_path_marks_placeholder_error(self, transcode_env):
+ """ingest 主流程抛异常(如元数据提取炸了)→ 占位按原始 key 找到并标 ERROR。"""
+ control = transcode_env
+ placeholder = _make_placeholder()
+ job = _make_job()
+ started, _, asset_repo = _start(control, job, [placeholder])
+ started["metadata"].side_effect = RuntimeError("boom in flow")
+ try:
+ result = ingest_mod.ingest_asset("job-1")
+ finally:
+ _stop(control)
+
+ assert result["status"] == "failed"
+ # 异常路径把占位标 ERROR(旧实现用被改写的 _h264 key 回查会落空)
+ error_marked = [a for a in asset_repo.assets if a.id == placeholder.id and a.status == AssetStatus.ERROR]
+ assert error_marked, "异常路径必须把占位 asset 标为 ERROR"
diff --git a/tests/unit/test_ingest_orphan_cleanup_1714.py b/tests/unit/test_ingest_orphan_cleanup_1714.py
new file mode 100644
index 000000000..cceccf8d4
--- /dev/null
+++ b/tests/unit/test_ingest_orphan_cleanup_1714.py
@@ -0,0 +1,262 @@
+"""#1714 上传/转码链路(IngestJob + Asset)孤儿清理测试。
+
+场景:worker 容器重启/进程 OOM 时,已 prefetch 的 transcode celery 消息丢失,
+ingest_job 永久卡 pending/processing、asset 永久卡 processing/uploading。
+"""
+
+from __future__ import annotations
+
+import os
+import sys
+from datetime import datetime, timedelta, timezone
+from pathlib import Path
+
+import pytest
+
+os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret")
+os.environ.setdefault("DATABASE_URL", "sqlite:///test_ingest_orphan.db")
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "worker"))
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
+sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
+
+from sqlalchemy import create_engine # noqa: E402
+from sqlalchemy.orm import sessionmaker # noqa: E402
+
+from packages.adapters.sqlalchemy_impl.models import AssetModel, Base, IngestJobModel # noqa: E402
+from packages.application.ingest_orphan_cleanup import ( # noqa: E402
+ cleanup_orphan_processing_assets,
+ cleanup_stale_ingest_jobs,
+)
+
+
+@pytest.fixture()
+def session():
+ engine = create_engine("sqlite:///:memory:")
+ Base.metadata.create_all(bind=engine)
+ Session = sessionmaker(bind=engine)
+ db = Session()
+ yield db
+ db.close()
+
+
+def _mk_job(session, *, status="processing", celery_task_id="cel-1", asset_id="a-1", minutes_ago=90):
+ now = datetime.now(timezone.utc)
+ job = IngestJobModel(
+ id=f"job-{minutes_ago}-{status}-{celery_task_id}",
+ project_id="p-1",
+ library_id="lib-1",
+ storage_key="uploads/x/IMG_2285.MOV",
+ status=status,
+ asset_id=asset_id,
+ celery_task_id=celery_task_id,
+ created_at=now - timedelta(minutes=minutes_ago),
+ updated_at=now - timedelta(minutes=minutes_ago),
+ )
+ session.add(job)
+ session.commit()
+ return job
+
+
+def _mk_asset(session, *, id="a-1", status="processing", minutes_ago=90, file_size=0):
+ now = datetime.now(timezone.utc)
+ asset = AssetModel(
+ id=id,
+ project_id="p-1",
+ asset_library_id="lib-1",
+ name="IMG_2285.MOV",
+ file_type="video",
+ file_size=file_size,
+ file_url="https://example.com/x.mov",
+ storage_key="uploads/x/IMG_2285.MOV",
+ status=status,
+ uploaded_by_user_id="u-1",
+ created_at=now - timedelta(minutes=minutes_ago),
+ updated_at=now - timedelta(minutes=minutes_ago),
+ )
+ session.add(asset)
+ session.commit()
+ return asset
+
+
+class TestCleanupStaleIngestJobs:
+ def test_stale_processing_job_marked_failed_and_asset_to_error(self, session):
+ """processing 超 60 分钟 → job failed,关联 processing asset → error。"""
+ _mk_asset(session, id="a-1", status="processing")
+ _mk_job(session, status="processing", celery_task_id="cel-dead", asset_id="a-1", minutes_ago=90)
+
+ items, asset_ids = cleanup_stale_ingest_jobs(session, processing_timeout_minutes=60, pending_timeout_minutes=90)
+
+ assert len(items) == 1
+ assert items[0] == ("job-90-processing-cel-dead", "cel-dead")
+ assert asset_ids == ["a-1"]
+ db_job = session.query(IngestJobModel).one()
+ assert db_job.status == "failed"
+ assert "中断" in db_job.error_message
+ db_asset = session.query(AssetModel).one()
+ assert db_asset.status == "error"
+
+ def test_stale_pending_job_marked_failed(self, session):
+ """pending 超 90 分钟(从未被消费)→ job failed。"""
+ _mk_asset(session, id="a-2", status="uploading")
+ _mk_job(session, status="pending", celery_task_id="", asset_id="a-2", minutes_ago=120)
+
+ items, asset_ids = cleanup_stale_ingest_jobs(session, processing_timeout_minutes=60, pending_timeout_minutes=90)
+
+ assert len(items) == 1
+ assert items[0][1] == "" # 无 celery task id
+ assert session.query(IngestJobModel).one().status == "failed"
+ assert session.query(AssetModel).one().status == "error"
+
+ def test_recent_processing_job_not_touched(self, session):
+ """processing 仅 10 分钟(正常转码中)→ 不误杀。"""
+ _mk_asset(session, id="a-3", status="processing", minutes_ago=10)
+ _mk_job(session, status="processing", celery_task_id="cel-live", asset_id="a-3", minutes_ago=10)
+
+ items, asset_ids = cleanup_stale_ingest_jobs(session, processing_timeout_minutes=60, pending_timeout_minutes=90)
+
+ assert items == []
+ assert asset_ids == []
+ assert session.query(IngestJobModel).one().status == "processing"
+ assert session.query(AssetModel).one().status == "processing"
+
+ def test_recent_pending_job_not_touched(self, session):
+ """pending 仅 30 分钟(队列积压排队中)→ 不误杀。"""
+ _mk_job(session, status="pending", asset_id="", minutes_ago=30)
+
+ items, _ = cleanup_stale_ingest_jobs(session, processing_timeout_minutes=60, pending_timeout_minutes=90)
+
+ assert items == []
+ assert session.query(IngestJobModel).one().status == "pending"
+
+ def test_terminal_job_not_touched(self, session):
+ """已 completed/failed 的 job 不动。"""
+ _mk_job(session, status="completed", celery_task_id="", asset_id="", minutes_ago=999)
+ _mk_job(session, status="failed", celery_task_id="", asset_id="", minutes_ago=999)
+
+ items, _ = cleanup_stale_ingest_jobs(session)
+
+ assert items == []
+ statuses = sorted(j.status for j in session.query(IngestJobModel).all())
+ assert statuses == ["completed", "failed"]
+
+ def test_ready_asset_not_demoted(self, session):
+ """关联 asset 已是 ready(转码其实成功了,仅 job 回写失败)→ 不降级为 error。"""
+ _mk_asset(session, id="a-4", status="ready")
+ _mk_job(session, status="processing", celery_task_id="cel-x", asset_id="a-4", minutes_ago=90)
+
+ _, asset_ids = cleanup_stale_ingest_jobs(session)
+
+ assert asset_ids == [] # ready 不动
+ assert session.query(AssetModel).one().status == "ready"
+
+
+class TestCleanupOrphanProcessingAssets:
+ def test_orphan_asset_without_job_marked_error(self, session):
+ """无 ingest_job 关联、created 超 120 分钟的 processing 占位 → error。"""
+ _mk_asset(session, id="orphan-1", status="processing", minutes_ago=150)
+
+ ids = cleanup_orphan_processing_assets(session, timeout_minutes=120)
+
+ assert ids == ["orphan-1"]
+ assert session.query(AssetModel).one().status == "error"
+
+ def test_asset_with_active_job_not_touched(self, session):
+ """有 processing job 关联的 asset 不由本函数处理(归 cleanup_stale_ingest_jobs)。"""
+ _mk_asset(session, id="a-5", status="processing", minutes_ago=150)
+ _mk_job(session, status="processing", asset_id="a-5", minutes_ago=150)
+
+ ids = cleanup_orphan_processing_assets(session, timeout_minutes=120)
+
+ assert ids == []
+ assert session.query(AssetModel).one().status == "processing"
+
+ def test_recent_orphan_asset_not_touched(self, session):
+ """无 job 但才创建 30 分钟 → 可能 complete 刚建、job 派单中,不动。"""
+ _mk_asset(session, id="orphan-2", status="processing", minutes_ago=30)
+
+ ids = cleanup_orphan_processing_assets(session, timeout_minutes=120)
+
+ assert ids == []
+ assert session.query(AssetModel).one().status == "processing"
+
+
+class TestRecoverStuckIngestJobsOnStartup:
+ def test_stuck_processing_job_requeued(self, session):
+ """processing 超 10 分钟 → 重置 pending 并重新 send_task,回写新 celery id。"""
+ from types import SimpleNamespace
+
+ from packages.application.ingest_orphan_cleanup import recover_stuck_ingest_jobs_on_startup
+
+ job = _mk_job(session, status="processing", celery_task_id="old-cel-1", asset_id="a-1", minutes_ago=30)
+
+ sent = []
+
+ def fake_send_task(name, args=None, **kw):
+ sent.append((name, args))
+ return SimpleNamespace(id="new-cel-9")
+
+ updated_ids = []
+ recovered = recover_stuck_ingest_jobs_on_startup(
+ session,
+ send_task=fake_send_task,
+ update_celery_task_id=lambda jid, cid: updated_ids.append((jid, cid)),
+ stuck_minutes=10,
+ )
+
+ assert recovered == 1
+ assert sent == [("worker.ingest_asset", [job.id])]
+ refreshed = session.query(IngestJobModel).filter_by(id=job.id).one()
+ assert refreshed.status == "pending"
+ assert refreshed.celery_task_id == "new-cel-9"
+ assert updated_ids == [(job.id, "new-cel-9")]
+
+ def test_recent_processing_job_not_touched(self, session):
+ """processing 仅 5 分钟(正常转码中/部署交接窗口)→ 不抢。"""
+ from packages.application.ingest_orphan_cleanup import recover_stuck_ingest_jobs_on_startup
+
+ _mk_job(session, status="processing", celery_task_id="live", asset_id="", minutes_ago=5)
+
+ sent = []
+ recovered = recover_stuck_ingest_jobs_on_startup(
+ session,
+ send_task=lambda *a, **k: sent.append(a),
+ stuck_minutes=10,
+ )
+
+ assert recovered == 0
+ assert sent == []
+ assert session.query(IngestJobModel).one().status == "processing"
+
+ def test_lock_not_acquired_skips(self, session):
+ """未抢到分布式锁(另一 worker 正在恢复)→ 跳过。"""
+ from packages.application.ingest_orphan_cleanup import recover_stuck_ingest_jobs_on_startup
+
+ _mk_job(session, status="processing", celery_task_id="x", asset_id="", minutes_ago=30)
+
+ recovered = recover_stuck_ingest_jobs_on_startup(
+ session,
+ send_task=lambda *a, **k: None,
+ lock_acquire=lambda: False,
+ stuck_minutes=10,
+ )
+
+ assert recovered == 0
+ assert session.query(IngestJobModel).one().status == "processing"
+
+ def test_pending_and_terminal_not_requeued(self, session):
+ """pending/已终态 job 不在恢复范围。"""
+ from packages.application.ingest_orphan_cleanup import recover_stuck_ingest_jobs_on_startup
+
+ _mk_job(session, status="pending", celery_task_id="", asset_id="", minutes_ago=60)
+ _mk_job(session, status="failed", celery_task_id="", asset_id="", minutes_ago=60)
+
+ recovered = recover_stuck_ingest_jobs_on_startup(
+ session,
+ send_task=lambda *a, **k: None,
+ stuck_minutes=10,
+ )
+
+ assert recovered == 0
+ statuses = sorted(j.status for j in session.query(IngestJobModel).all())
+ assert statuses == ["failed", "pending"]
diff --git a/tests/unit/test_orphan_guard_purge_mocked_1714.py b/tests/unit/test_orphan_guard_purge_mocked_1714.py
new file mode 100644
index 000000000..29122959d
--- /dev/null
+++ b/tests/unit/test_orphan_guard_purge_mocked_1714.py
@@ -0,0 +1,360 @@
+"""#1714:孤儿消息撤销/清理逻辑测试(mock redis,CI 无真实 redis 时也产生覆盖)。
+
+覆盖 packages/shared/celery_orphan_guard.py:
+- _extract_business_ids:三元组 body / 裸 args body / dict args / headers 提取 /
+ 无 body / 坏 JSON / 坏 base64 / 空 args
+- _purge_one_queue:biz id 命中、celery id 命中、未命中保序(重写 rpush)、
+ lrange 异常、重写异常、空队列
+- purge_stale_messages_from_queues:空 ids 早退、redis 未安装、连接失败、
+ 正常清理并 close
+- revoke_and_purge:revoke 逐消息调用、revoke 异常不阻断、空 id 跳过
+- ensure_task_claimable:任务不存在返回空串、终态抛错、pending 放行
+"""
+
+from __future__ import annotations
+
+import base64
+import json
+import sys
+import types
+from pathlib import Path
+from unittest.mock import MagicMock
+
+import pytest
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
+
+from packages.shared import celery_orphan_guard as guard # noqa: E402
+
+
+def _envelope(celery_id: str | None, body_payload) -> bytes:
+ """构造 Redis transport 存储的 celery 消息(JSON 信封)。"""
+ if body_payload is None:
+ body = None
+ else:
+ body = base64.b64encode(json.dumps(body_payload).encode()).decode()
+ envelope = {"body": body, "headers": {"id": celery_id, "task": "worker.generate_video"}}
+ return json.dumps(envelope).encode()
+
+
+# ── _extract_business_ids ───────────────────────────────────────────────
+
+
+def test_extract_ids_standard_tuple_body():
+ raw = _envelope("celery-1", [["biz-task-1"], {}, {"callbacks": None}])
+ assert guard._extract_business_ids(raw) == ("celery-1", "biz-task-1")
+
+
+def test_extract_ids_bare_args_body():
+ raw = _envelope("celery-2", ["biz-task-2"])
+ assert guard._extract_business_ids(raw) == ("celery-2", "biz-task-2")
+
+
+def test_extract_ids_dict_body_with_args():
+ raw = _envelope("celery-3", {"args": ["biz-task-3"], "kwargs": {}})
+ assert guard._extract_business_ids(raw) == ("celery-3", "biz-task-3")
+
+
+def test_extract_ids_non_dict_headers_returns_celery_id_none():
+ raw = json.dumps({"body": base64.b64encode(json.dumps([["biz-4"]]).encode()).decode(), "headers": "x"}).encode()
+ celery_id, biz_id = guard._extract_business_ids(raw)
+ assert celery_id is None
+ assert biz_id == "biz-4"
+
+
+def test_extract_ids_no_body_returns_celery_id_only():
+ raw = json.dumps({"headers": {"id": "celery-5"}}).encode()
+ assert guard._extract_business_ids(raw) == ("celery-5", None)
+
+
+def test_extract_ids_empty_args_returns_no_biz_id():
+ raw = _envelope("celery-6", [[], {}, {}])
+ assert guard._extract_business_ids(raw) == ("celery-6", None)
+
+
+def test_extract_ids_args_first_none_returns_no_biz_id():
+ raw = _envelope("celery-7", [[None], {}, {}])
+ assert guard._extract_business_ids(raw) == ("celery-7", None)
+
+
+def test_extract_ids_bad_json_returns_none_none():
+ assert guard._extract_business_ids(b"not-json{") == (None, None)
+
+
+def test_extract_ids_bad_base64_returns_none_none():
+ raw = json.dumps({"body": "!!!not-base64!!!", "headers": {"id": "c"}}).encode()
+ assert guard._extract_business_ids(raw) == (None, None)
+
+
+def test_extract_ids_int_arg_coerced_to_str():
+ raw = _envelope("celery-9", [[12345], {}, {}])
+ celery_id, biz_id = guard._extract_business_ids(raw)
+ assert celery_id == "celery-9"
+ assert biz_id == "12345"
+
+
+# ── _purge_one_queue ────────────────────────────────────────────────────
+
+
+def _queue_with_messages(*payloads: bytes):
+ """返回 list-backed mock redis client(记录当前队列内容)。"""
+ client = MagicMock()
+ store: dict[str, list[bytes]] = {"q": list(payloads)}
+
+ def lrange(name, start, end): # noqa: ARG001
+ return list(store.get(name, []))
+
+ client.lrange.side_effect = lrange
+
+ pipe = MagicMock()
+ pipe.delete.side_effect = lambda name: store.pop(name, None)
+ pipe.rpush.side_effect = lambda name, *items: store.setdefault(name, []).extend(items)
+ client.pipeline.return_value = pipe
+ return client, store, pipe
+
+
+def test_purge_one_queue_removes_by_biz_id_and_keeps_order():
+ stale = _envelope("c-stale", [["biz-stale"], {}, {}])
+ keep1 = _envelope("c-keep-1", [["biz-keep-1"], {}, {}])
+ keep2 = _envelope("c-keep-2", [["biz-keep-2"], {}, {}])
+ client, store, pipe = _queue_with_messages(keep1, stale, keep2)
+
+ removed = guard._purge_one_queue(client, "q", {"biz-stale"}, set())
+ assert removed == 1
+ # 队列被 delete + rpush 重写,未命中消息保持相对顺序
+ pipe.delete.assert_called_once_with("q")
+ pipe.rpush.assert_called_once()
+ args, _ = pipe.rpush.call_args
+ assert args[0] == "q"
+ assert list(args[1:]) == [keep1, keep2]
+ pipe.execute.assert_called_once()
+
+
+def test_purge_one_queue_removes_by_celery_message_id():
+ stale = _envelope("celery-xyz", [["biz-whatever"], {}, {}])
+ keep = _envelope("celery-aaa", [["biz-keep"], {}, {}])
+ client, store, pipe = _queue_with_messages(stale, keep)
+
+ removed = guard._purge_one_queue(client, "q", set(), {"celery-xyz"})
+ assert removed == 1
+ args, _ = pipe.rpush.call_args
+ assert list(args[1:]) == [keep]
+
+
+def test_purge_one_queue_no_hit_no_rewrite():
+ msg1 = _envelope("c1", [["b1"], {}, {}])
+ msg2 = _envelope("c2", [["b2"], {}, {}])
+ client, store, pipe = _queue_with_messages(msg1, msg2)
+
+ removed = guard._purge_one_queue(client, "q", {"other"}, {"other-c"})
+ assert removed == 0
+ # 没有命中:不重写队列
+ pipe.delete.assert_not_called()
+ pipe.rpush.assert_not_called()
+
+
+def test_purge_one_queue_all_removed_deletes_without_rpush():
+ stale = _envelope("c-stale", [["biz-stale"], {}, {}])
+ client, store, pipe = _queue_with_messages(stale)
+
+ removed = guard._purge_one_queue(client, "q", {"biz-stale"}, set())
+ assert removed == 1
+ pipe.delete.assert_called_once_with("q")
+ pipe.rpush.assert_not_called()
+
+
+def test_purge_one_queue_lrange_exception_returns_zero():
+ client = MagicMock()
+ client.lrange.side_effect = RuntimeError("redis down")
+ assert guard._purge_one_queue(client, "q", {"b"}, set()) == 0
+
+
+def test_purge_one_queue_empty_queue_returns_zero():
+ client = MagicMock()
+ client.lrange.return_value = []
+ assert guard._purge_one_queue(client, "q", {"b"}, set()) == 0
+ client.pipeline.assert_not_called()
+
+
+def test_purge_one_queue_rewrite_exception_returns_zero():
+ stale = _envelope("c-stale", [["biz-stale"], {}, {}])
+ client, store, pipe = _queue_with_messages(stale)
+ pipe.execute.side_effect = RuntimeError("write fail")
+
+ removed = guard._purge_one_queue(client, "q", {"biz-stale"}, set())
+ assert removed == 0
+
+
+def test_purge_one_queue_unparseable_message_conservatively_kept():
+ stale = _envelope("c-stale", [["biz-stale"], {}, {}])
+ garbage = b"garbage-not-a-message"
+ client, store, pipe = _queue_with_messages(garbage, stale)
+
+ removed = guard._purge_one_queue(client, "q", {"biz-stale"}, set())
+ assert removed == 1
+ args, _ = pipe.rpush.call_args
+ # 无法解析的消息保守保留,绝不误删
+ assert list(args[1:]) == [garbage]
+
+
+# ── purge_stale_messages_from_queues ────────────────────────────────────
+
+
+def test_purge_queues_no_ids_returns_zero_without_connecting():
+ assert guard.purge_stale_messages_from_queues("redis://x", ("q",)) == 0
+
+
+def test_purge_queues_blank_ids_filtered_out():
+ assert guard.purge_stale_messages_from_queues("redis://x", ("q",), business_task_ids=["", None]) == 0
+
+
+def test_purge_queues_redis_not_installed(monkeypatch):
+ """redis-py 不可用(ImportError)时安全返回 0。"""
+ import builtins
+
+ real_import = builtins.__import__
+
+ def fake_import(name, *args, **kwargs):
+ if name == "redis":
+ raise ImportError("no redis")
+ return real_import(name, *args, **kwargs)
+
+ monkeypatch.setattr(builtins, "__import__", fake_import)
+ assert guard.purge_stale_messages_from_queues("redis://x", ("q",), business_task_ids=["b1"]) == 0
+
+
+def test_purge_queues_connection_failure_returns_zero():
+ fake_redis = types.ModuleType("redis")
+
+ class _FakeRedis:
+ @classmethod
+ def from_url(cls, url): # noqa: ARG003
+ client = MagicMock()
+ client.ping.side_effect = ConnectionError("connect refused")
+ return client
+
+ fake_redis.Redis = _FakeRedis
+ sys.modules["redis"] = fake_redis
+ try:
+ assert guard.purge_stale_messages_from_queues("redis://x", ("q",), celery_task_ids=["c1"]) == 0
+ finally:
+ sys.modules.pop("redis", None)
+
+
+def test_purge_queues_happy_path_closes_client():
+ stale = _envelope("c-stale", [["biz-stale"], {}, {}])
+ fake_redis = types.ModuleType("redis")
+
+ client = MagicMock()
+ client.lrange.return_value = [stale]
+ pipe = MagicMock()
+ client.pipeline.return_value = pipe
+
+ class _FakeRedis:
+ @classmethod
+ def from_url(cls, url): # noqa: ARG003
+ return client
+
+ fake_redis.Redis = _FakeRedis
+ sys.modules["redis"] = fake_redis
+ try:
+ removed = guard.purge_stale_messages_from_queues(
+ "redis://x", ("generation", "transcode"), business_task_ids=["biz-stale"]
+ )
+ finally:
+ sys.modules.pop("redis", None)
+
+ # mock client 对两个队列都返回同一条作废消息 → 各移除 1 条
+ assert removed == 2
+ client.ping.assert_called_once()
+ client.close.assert_called_once()
+ # 两个队列都扫描
+ assert client.lrange.call_count == 2
+
+
+def test_purge_queues_close_exception_swallowed():
+ fake_redis = types.ModuleType("redis")
+
+ client = MagicMock()
+ client.lrange.return_value = []
+ client.close.side_effect = RuntimeError("close fail")
+
+ class _FakeRedis:
+ @classmethod
+ def from_url(cls, url): # noqa: ARG003
+ return client
+
+ fake_redis.Redis = _FakeRedis
+ sys.modules["redis"] = fake_redis
+ try:
+ removed = guard.purge_stale_messages_from_queues("redis://x", ("q",), celery_task_ids=["c1"])
+ finally:
+ sys.modules.pop("redis", None)
+ assert removed == 0
+
+
+# ── revoke_and_purge ────────────────────────────────────────────────────
+
+
+def test_revoke_and_purge_revokes_each_message(monkeypatch):
+ fake_app = MagicMock()
+ purge_mock = MagicMock(return_value=2)
+ monkeypatch.setattr(guard, "purge_stale_messages_from_queues", purge_mock)
+
+ removed = guard.revoke_and_purge(
+ fake_app,
+ "redis://x",
+ business_task_ids=["b1"],
+ celery_task_ids=["c1", "c2"],
+ queue_names=("generation",),
+ )
+ assert removed == 2
+ assert fake_app.control.revoke.call_count == 2
+ fake_app.control.revoke.assert_any_call("c1")
+ fake_app.control.revoke.assert_any_call("c2")
+ purge_mock.assert_called_once_with(
+ "redis://x", ("generation",), business_task_ids=["b1"], celery_task_ids=["c1", "c2"]
+ )
+
+
+def test_revoke_and_purge_revoke_exception_does_not_block(monkeypatch):
+ fake_app = MagicMock()
+ fake_app.control.revoke.side_effect = RuntimeError("broadcast fail")
+ purge_mock = MagicMock(return_value=0)
+ monkeypatch.setattr(guard, "purge_stale_messages_from_queues", purge_mock)
+
+ removed = guard.revoke_and_purge(fake_app, "redis://x", celery_task_ids=["c1"])
+ assert removed == 0
+ purge_mock.assert_called_once()
+
+
+def test_revoke_and_purge_skips_blank_ids(monkeypatch):
+ fake_app = MagicMock()
+ purge_mock = MagicMock(return_value=0)
+ monkeypatch.setattr(guard, "purge_stale_messages_from_queues", purge_mock)
+
+ guard.revoke_and_purge(fake_app, "redis://x", celery_task_ids=["", None])
+ fake_app.control.revoke.assert_not_called()
+
+
+# ── ensure_task_claimable ───────────────────────────────────────────────
+
+
+def test_ensure_claimable_missing_task_returns_empty():
+ assert guard.ensure_task_claimable("t1", lambda _tid: None) == ""
+
+
+def test_ensure_claimable_terminal_raises():
+ with pytest.raises(guard.StaleTaskDiscarded) as exc_info:
+ guard.ensure_task_claimable("t1", lambda _tid: "failed")
+ assert exc_info.value.task_id == "t1"
+ assert exc_info.value.status == "failed"
+
+
+def test_ensure_claimable_cancelled_raises():
+ with pytest.raises(guard.StaleTaskDiscarded):
+ guard.ensure_task_claimable("t1", lambda _tid: "cancelled")
+
+
+def test_ensure_claimable_pending_passes():
+ assert guard.ensure_task_claimable("t1", lambda _tid: "pending") == "pending"
diff --git a/tests/unit/test_patch_me_profile_1718.py b/tests/unit/test_patch_me_profile_1718.py
new file mode 100644
index 000000000..d2444f77b
--- /dev/null
+++ b/tests/unit/test_patch_me_profile_1718.py
@@ -0,0 +1,222 @@
+"""#1718:PATCH /auth/me 资料更新接口测试。
+
+覆盖:
+- 正常更新昵称并落库
+- strip 生效(前后空白去除)
+- 纯空白/超长 -> 422(pydantic 校验)
+- 首次设置昵称 profile_completed False->True
+- 已完成用户重复提交幂等(仍 True)
+- 未登录由 get_current_user 依赖保证 401(框架行为,这里验证路由声明了该依赖)
+- 响应结构 {user: {...}} 含 wechat_bound/profile_completed 全字段
+- 微信新建用户 profile_completed 默认 False(wechat_sync _create_wechat_user)
+"""
+
+from __future__ import annotations
+
+import asyncio
+import os
+import sys
+from pathlib import Path
+from types import SimpleNamespace
+from unittest.mock import MagicMock
+
+import pytest
+from pydantic import ValidationError
+
+os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
+os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
+
+from app.api.routes import auth as auth_route # noqa: E402
+
+from packages.adapters.in_memory.user_repository import InMemoryUserRepository # noqa: E402
+from packages.domain.entities import User # noqa: E402
+
+
+def _auth_user(user):
+ return SimpleNamespace(user=user, session_id="s-1", token_type="user_auth")
+
+
+def _make_user(**kw):
+ defaults = dict(
+ id="u-1",
+ email="user@example.com",
+ username="user",
+ display_name="微信用户",
+ password_hash="x",
+ email_verified=True,
+ profile_completed=False,
+ )
+ defaults.update(kw)
+ return User(**defaults)
+
+
+# ---------- 请求体校验 ----------
+
+
+def test_display_name_strips_whitespace():
+ req = auth_route.UpdateProfileRequest(display_name=" ying123 ")
+ assert req.display_name == "ying123"
+
+
+def test_display_name_blank_rejected():
+ with pytest.raises(ValidationError) as exc:
+ auth_route.UpdateProfileRequest(display_name=" ")
+ assert "空白" in str(exc.value)
+
+
+def test_display_name_empty_rejected():
+ with pytest.raises(ValidationError):
+ auth_route.UpdateProfileRequest(display_name="")
+
+
+def test_display_name_too_long_rejected():
+ with pytest.raises(ValidationError) as exc:
+ auth_route.UpdateProfileRequest(display_name="甲" * 21)
+ assert "1-20" in str(exc.value)
+
+
+def test_display_name_max_length_accepted():
+ req = auth_route.UpdateProfileRequest(display_name="甲" * 20)
+ assert req.display_name == "甲" * 20
+
+
+# ---------- 路由逻辑 ----------
+
+
+def test_patch_me_updates_display_name_and_persists():
+ user = _make_user()
+ repo = InMemoryUserRepository()
+ repo.save(user)
+
+ resp = asyncio.run(
+ auth_route.update_current_user_profile(
+ auth_route.UpdateProfileRequest(display_name=" ying123 "),
+ current_user=_auth_user(user),
+ user_repository=repo,
+ )
+ )
+ assert resp.user.display_name == "ying123"
+ assert resp.user.profile_completed is True
+ assert resp.user.wechat_bound is False
+ # 落库验证
+ fresh = repo.find_by_id("u-1")
+ assert fresh.display_name == "ying123"
+ assert fresh.profile_completed is True
+
+
+def test_patch_me_first_time_sets_profile_completed_true():
+ user = _make_user(profile_completed=False)
+ repo = InMemoryUserRepository()
+ repo.save(user)
+ assert repo.find_by_id("u-1").profile_completed is False
+
+ asyncio.run(
+ auth_route.update_current_user_profile(
+ auth_route.UpdateProfileRequest(display_name="小虾"),
+ current_user=_auth_user(user),
+ user_repository=repo,
+ )
+ )
+ assert repo.find_by_id("u-1").profile_completed is True
+
+
+def test_patch_me_idempotent_for_completed_user():
+ user = _make_user(display_name="老名字", profile_completed=True)
+ repo = InMemoryUserRepository()
+ repo.save(user)
+
+ resp = asyncio.run(
+ auth_route.update_current_user_profile(
+ auth_route.UpdateProfileRequest(display_name="新名字"),
+ current_user=_auth_user(user),
+ user_repository=repo,
+ )
+ )
+ assert resp.user.profile_completed is True
+ assert resp.user.display_name == "新名字"
+ # 再提交一次同样内容,不报错、状态稳定
+ resp2 = asyncio.run(
+ auth_route.update_current_user_profile(
+ auth_route.UpdateProfileRequest(display_name="新名字"),
+ current_user=_auth_user(repo.find_by_id("u-1")),
+ user_repository=repo,
+ )
+ )
+ assert resp2.user.profile_completed is True
+
+
+def test_patch_me_response_contains_all_me_fields():
+ user = _make_user(wechat_openid="wx-1", phone="13800000000", phone_verified=True)
+ repo = InMemoryUserRepository()
+ repo.save(user)
+
+ resp = asyncio.run(
+ auth_route.update_current_user_profile(
+ auth_route.UpdateProfileRequest(display_name="昵称"),
+ current_user=_auth_user(user),
+ user_repository=repo,
+ )
+ )
+ payload = resp.user.model_dump()
+ for field in (
+ "user_id",
+ "email",
+ "username",
+ "display_name",
+ "email_verified",
+ "phone",
+ "phone_verified",
+ "binding_complete",
+ "wechat_bound",
+ "profile_completed",
+ ):
+ assert field in payload, f"missing field {field}"
+ assert payload["wechat_bound"] is True
+ assert payload["phone"] == "13800000000"
+
+
+def test_get_me_includes_profile_completed_flag():
+ # 未完成
+ u = _make_user(profile_completed=False)
+ resp = asyncio.run(auth_route.get_current_user_info(authenticated_user=_auth_user(u)))
+ assert resp.profile_completed is False
+ assert resp.wechat_bound is False
+
+ # 已完成 + 已绑微信
+ u2 = _make_user(profile_completed=True, wechat_openid="wx-9")
+ resp2 = asyncio.run(auth_route.get_current_user_info(authenticated_user=_auth_user(u2)))
+ assert resp2.profile_completed is True
+ assert resp2.wechat_bound is True
+
+
+def test_patch_me_requires_auth_dependency():
+ # 路由签名必须依赖 get_current_user,未携带 token 时框架返回 401
+ params = (
+ auth_route.update_current_user_profile.__wrapped__
+ if hasattr(auth_route.update_current_user_profile, "__wrapped__")
+ else auth_route.update_current_user_profile
+ )
+ import inspect
+
+ sig = inspect.signature(params)
+ dep = sig.parameters.get("current_user")
+ assert dep is not None
+ assert dep.default is not None and getattr(dep.default, "dependency", None) is auth_route.get_current_user
+
+
+def test_wechat_new_user_created_with_profile_completed_false():
+ # 微信同步建号:新用户 profile_completed=False(需引导设置昵称)
+ from packages.application.auth.wechat_sync_use_case import (
+ WechatSyncRequest,
+ WechatSyncUseCase,
+ )
+
+ repo = InMemoryUserRepository()
+ # session_store 用 mock,不依赖 redis
+ use_case = WechatSyncUseCase(user_repository=repo, session_store=MagicMock(), jwt_secret_key="test-secret")
+ resp, err = use_case.execute(WechatSyncRequest(openid="wx-new-openid", nickname="微信测试", source="web"))
+ assert err is None
+ user = repo.find_by_id(resp.user_id)
+ assert user.profile_completed is False
diff --git a/tests/unit/test_persist_celery_id_routes_1714.py b/tests/unit/test_persist_celery_id_routes_1714.py
new file mode 100644
index 000000000..ddfa89dda
--- /dev/null
+++ b/tests/unit/test_persist_celery_id_routes_1714.py
@@ -0,0 +1,259 @@
+"""#1714:入队后 celery_task_id 持久化路径覆盖(routes / enqueue / celery_app / 仓储)。
+
+CI 无 redis、不走完整 HTTP 流程,这些 try/except 与早退分支此前覆盖率为 0。
+用真实 SQLite 仓储 + monkeypatch celery_app.send_task 直接驱动路由函数:
+- routes/ingest_jobs.submit_ingest_job:正常持久化 + 持久化异常吞掉不影响响应
+- routes/task_center.retry_project_task(ingest 分支):重试后持久化 + 异常吞掉
+- routes/upload._persist_celery_task_id:空 id 早退 + 异常吞掉
+- core/task_enqueue.safe_enqueue_generation_task:持久化失败仅 warning,入队仍 True
+- core/celery_app:apply_queue_settings 抛异常时 API 启动不炸
+- adapters/ingest_job_repository.update:写 celery_task_id 分支落库
+"""
+
+from __future__ import annotations
+
+import importlib
+import importlib.util
+import os
+import sys
+from pathlib import Path
+from types import SimpleNamespace
+from unittest.mock import MagicMock
+
+os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
+os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
+
+API_PATH = str(Path(__file__).resolve().parents[2] / "apps" / "api")
+if API_PATH not in sys.path:
+ sys.path.insert(0, API_PATH)
+
+import pytest # noqa: E402
+from app.api.routes import ingest_jobs as ingest_jobs_route # noqa: E402
+from app.api.routes import task_center as task_center_route # noqa: E402
+from app.api.routes import upload as upload_route # noqa: E402
+from app.schemas.ingest_job import SubmitIngestJobRequest # noqa: E402
+from sqlalchemy import create_engine # noqa: E402
+from sqlalchemy.orm import sessionmaker # noqa: E402
+
+from packages.adapters.sqlalchemy_impl.ingest_job_repository import ( # noqa: E402
+ SQLAlchemyIngestJobRepository,
+)
+from packages.adapters.sqlalchemy_impl.models import Base # noqa: E402
+from packages.domain import IngestJob, IngestJobStatus # noqa: E402
+
+
+def _ingest_repo():
+ engine = create_engine("sqlite:///:memory:")
+ Base.metadata.create_all(engine)
+ session = sessionmaker(bind=engine)()
+ return SQLAlchemyIngestJobRepository(session), session
+
+
+def _fake_celery_result(task_id: str = "celery-route-msg-1"):
+ result = MagicMock()
+ result.id = task_id
+ return result
+
+
+# ── routes/ingest_jobs.submit_ingest_job ────────────────────────────────
+
+
+def test_submit_ingest_job_persists_celery_task_id(monkeypatch):
+ repo, session = _ingest_repo()
+ monkeypatch.setattr(ingest_jobs_route.celery_app, "send_task", MagicMock(return_value=_fake_celery_result()))
+
+ request = SubmitIngestJobRequest(project_id="proj-1", library_id="lib-1", storage_key="uploads/x.mov")
+ response = ingest_jobs_route.submit_ingest_job(request, ingest_job_repository=repo)
+
+ assert response.status == "pending"
+ saved = repo.get(response.id)
+ assert saved.celery_task_id == "celery-route-msg-1"
+
+
+def test_submit_ingest_job_persist_failure_swallowed(monkeypatch):
+ repo, _ = _ingest_repo()
+
+ class _BoomRepo:
+ def __init__(self, inner):
+ self.inner = inner
+
+ def create(self, job):
+ return self.inner.create(job)
+
+ def get(self, job_id):
+ return self.inner.get(job_id)
+
+ def update(self, job): # noqa: ARG002
+ raise RuntimeError("db write fail")
+
+ boom_repo = _BoomRepo(repo)
+ monkeypatch.setattr(ingest_jobs_route.celery_app, "send_task", MagicMock(return_value=_fake_celery_result()))
+
+ request = SubmitIngestJobRequest(project_id="proj-1", library_id="lib-1", storage_key="uploads/y.mov")
+ # 持久化异常被吞掉,主流程(响应)不受影响
+ response = ingest_jobs_route.submit_ingest_job(request, ingest_job_repository=boom_repo)
+ assert response.id
+ assert response.status == "pending"
+
+
+# ── routes/task_center.retry_project_task(ingest 分支) ────────────────
+
+
+def _auth_user():
+ user = SimpleNamespace(id="user-1")
+ return SimpleNamespace(user=user, session_id=None, token_type=None)
+
+
+def test_retry_ingest_job_persists_celery_task_id(monkeypatch):
+ repo, session = _ingest_repo()
+ # 造一条 failed 的 ingest job
+ job = IngestJob.create(project_id="proj-1", library_id="lib-1", storage_key="uploads/z.mov")
+ job.status = IngestJobStatus.FAILED
+ repo.create(job)
+
+ monkeypatch.setattr(
+ task_center_route.celery_app, "send_task", MagicMock(return_value=_fake_celery_result("celery-retry-1"))
+ )
+
+ response = task_center_route.retry_project_task(
+ "ingest", job.id, authenticated_user=_auth_user(), ingest_job_repository=repo
+ )
+ assert response.task_type == "ingest"
+ new_id = response.id.split("ingest:")[1]
+ retried = repo.get(new_id)
+ assert retried is not None
+ assert retried.celery_task_id == "celery-retry-1"
+
+
+def test_retry_ingest_job_persist_failure_swallowed(monkeypatch):
+ repo, _ = _ingest_repo()
+ job = IngestJob.create(project_id="proj-1", library_id="lib-1", storage_key="uploads/w.mov")
+ job.status = IngestJobStatus.FAILED
+ repo.create(job)
+
+ real_update = repo.update
+
+ def _update_that_booms(entity):
+ # 仅在写 celery_task_id 的那次 update 抛错(新建 job 后路由内的持久化)
+ if getattr(entity, "celery_task_id", ""):
+ raise RuntimeError("db write fail")
+ return real_update(entity)
+
+ repo.update = _update_that_booms # type: ignore[method-assign]
+ monkeypatch.setattr(task_center_route.celery_app, "send_task", MagicMock(return_value=_fake_celery_result()))
+
+ # 持久化异常吞掉,重试接口仍正常返回
+ response = task_center_route.retry_project_task(
+ "ingest", job.id, authenticated_user=_auth_user(), ingest_job_repository=repo
+ )
+ assert response.task_type == "ingest"
+
+
+# ── routes/upload._persist_celery_task_id ───────────────────────────────
+
+
+def test_upload_persist_helper_empty_id_early_return():
+ repo = MagicMock()
+ job = MagicMock()
+ upload_route._persist_celery_task_id(repo, job, "")
+ repo.update.assert_not_called()
+ upload_route._persist_celery_task_id(repo, job, None) # type: ignore[arg-type]
+ repo.update.assert_not_called()
+
+
+def test_upload_persist_helper_exception_swallowed():
+ repo = MagicMock()
+ repo.update.side_effect = RuntimeError("db fail")
+ job = MagicMock()
+ # 不抛异常
+ upload_route._persist_celery_task_id(repo, job, "celery-upload-1")
+ repo.update.assert_called_once()
+ assert job.celery_task_id == "celery-upload-1"
+
+
+# ── core/task_enqueue:持久化失败仅 warning ─────────────────────────────
+
+
+def test_safe_enqueue_persist_failure_still_returns_true(monkeypatch):
+ from app.core import task_enqueue
+
+ class _FakeTask:
+ def __init__(self):
+ self.id = "task-enqueue-persist-fail"
+ self.status = "pending"
+ self.celery_task_id = ""
+
+ def mark_failed(self, msg): # noqa: ARG002
+ self.status = "failed"
+
+ class _FakeRepo:
+ def count_pending_total(self):
+ return 0
+
+ def count_pending_by_user(self, user_id): # noqa: ARG002
+ return 0
+
+ def update(self, task): # noqa: ARG002
+ raise RuntimeError("persist celery_task_id failed")
+
+ fake_result = MagicMock()
+ fake_result.id = "celery-enqueue-fail-1"
+ mock_celery = MagicMock()
+ mock_celery.send_task.return_value = fake_result
+ monkeypatch.setattr(task_enqueue, "celery_app", mock_celery)
+
+ task = _FakeTask()
+ repo = _FakeRepo()
+ ok = task_enqueue.safe_enqueue_generation_task(task, repo, user_id="u1")
+ # 持久化失败不影响入队结果
+ assert ok is True
+ mock_celery.send_task.assert_called_once()
+
+
+# ── core/celery_app:队列配置失败不阻断 API 启动 ────────────────────────
+
+
+def test_api_celery_app_survives_queue_settings_failure(monkeypatch):
+ """apply_queue_settings 抛异常时 API 启动不炸(core/celery_app.py 的 try/except 分支)。
+
+ 通过让 `from packages.shared.celery_queues import apply_queue_settings` 本身
+ 抛异常来触发 except 分支;用全新模块名 reload,不替换已被其他模块持有的
+ app.core.celery_app 模块对象,避免污染 task_enqueue 等导入方。
+ """
+ import builtins
+
+ real_import = builtins.__import__
+
+ def _failing_import(name, globals=None, locals=None, fromlist=(), level=0): # noqa: A002
+ if name == "packages.shared.celery_queues" and "apply_queue_settings" in (fromlist or ()):
+ raise RuntimeError("config boom")
+ return real_import(name, globals, locals, fromlist, level)
+
+ monkeypatch.setattr(builtins, "__import__", _failing_import)
+
+ spec = importlib.util.find_spec("app.core.celery_app")
+ fresh_mod = importlib.util.module_from_spec(spec)
+ spec.loader.exec_module(fresh_mod) # 异常在模块内被 try/except 吞掉
+ assert fresh_mod.celery_app is not None
+ assert fresh_mod.celery_app.main == "xiaoxia-saas-api"
+
+ # 已加载的原模块对象不受影响(无 reload 污染)
+ import app.core.celery_app as api_celery_mod
+
+ assert api_celery_mod.celery_app is not None
+
+
+# ── 仓储:update 写 celery_task_id 落库 ─────────────────────────────────
+
+
+def test_ingest_repo_update_persists_celery_task_id():
+ repo, session = _ingest_repo()
+ job = IngestJob.create(project_id="proj-1", library_id="lib-1", storage_key="uploads/repo.mov")
+ repo.create(job)
+
+ job.celery_task_id = "celery-repo-update-1"
+ repo.update(job)
+
+ session.expire_all()
+ saved = repo.get(job.id)
+ assert saved.celery_task_id == "celery-repo-update-1"
diff --git a/tests/unit/test_phash_threshold_calibration_1658.py b/tests/unit/test_phash_threshold_calibration_1658.py
index 6e01b4a6a..cfefdd077 100644
--- a/tests/unit/test_phash_threshold_calibration_1658.py
+++ b/tests/unit/test_phash_threshold_calibration_1658.py
@@ -102,6 +102,7 @@ from video_processing.dedup import ( # noqa: E402
DUPLICATE_THRESHOLD,
HISTOGRAM_WEIGHT,
MATCH_RATIO_THRESHOLD,
+ PHASH_THRESHOLD,
PHASH_WEIGHT,
VideoDeduplicator,
)
@@ -128,11 +129,17 @@ _ZERO_HIST = [0.0] * 96 # 全黑视频的全零直方图(有效数据)
class TestThresholdCalibration:
- """pHash 阈值由 10 收紧到 8(Issue #1658)。"""
+ """pHash 阈值校准(#1658 收紧到 8,#1702 两轮真实数据重校准 12→16)。
- def test_phash_threshold_is_8(self):
- """PHASH_THRESHOLD 必须为 8(旧值 10 会放过 8~9 汉明距离的不同视频)。"""
- assert VideoDeduplicator.PHASH_THRESHOLD == 8
+ #1702 第一轮 staging 离线实验:同帧两次 2-5% 随机裁剪距离 4~10;同源成片
+ (密集 1s 采样)<=12 命中 4/11、异源成片最小距离 24 → 初定 12。
+ #1702 第二轮(证据视频 B->A 仍漏检)扩样本到该用户 15 个真实成片实测:
+ 同源降重对中位数距离 14、<=16 命中 8/11=0.73;异源 13 个候选 <=16 命中
+ 全 0、每帧全局最近邻最小距离 18 → 校准为 16(与异源仍有 >=2bit 裕度)。
+ """
+
+ def test_phash_threshold_is_calibrated(self):
+ assert VideoDeduplicator.PHASH_THRESHOLD == PHASH_THRESHOLD == 16
def test_match_ratio_threshold_constant(self):
assert MATCH_RATIO_THRESHOLD == 0.7
@@ -144,22 +151,22 @@ class TestThresholdCalibration:
assert PHASH_WEIGHT == 0.7
assert HISTOGRAM_WEIGHT == 0.3
- def test_threshold_tightening_excludes_distance_8_and_9(self):
- """距离 8、9 的帧:旧阈值 10 下算匹配,新阈值 8 下不算匹配。
+ def test_threshold_matching_semantics(self):
+ """阈值比较统一为 <=(帧匹配与片段匹配同一口径)。
- 场景:5 个关键帧距离为 [7, 7, 7, 9, 9]。
- - 旧阈值 10:5 帧全部 < 10 → match_ratio = 1.0(误放过)
- - 新阈值 8:仅 3 帧 < 8 → match_ratio = 0.6 < 0.7(正确跳过)
+ 场景:5 个关键帧距离为 [10, 14, 16, 18, 26]。
+ - <=16(#1702 二次校准阈值):3 帧匹配 → 0.6 < 0.7,被帧比例门槛
+ 拦截(异源安全边界:真实数据异源最近邻最小距离 18,<=16 命中 0)
+ - 距离正好 16 的同源降重帧应算匹配(< 与 <= 口径统一)
"""
- distances = [7, 7, 7, 9, 9]
+ distances = [10, 14, 16, 18, 26]
+ matched = sum(1 for d in distances if d <= VideoDeduplicator.PHASH_THRESHOLD)
+ assert matched == 3
+ assert matched / len(distances) == 0.6
+ assert matched / len(distances) < MATCH_RATIO_THRESHOLD
- matched_old = sum(1 for d in distances if d < 10)
- assert matched_old == 5 # 旧行为:全匹配 → 误判风险
-
- matched_new = sum(1 for d in distances if d < VideoDeduplicator.PHASH_THRESHOLD)
- assert matched_new == 3
- assert matched_new / len(distances) == 0.6
- assert matched_new / len(distances) < MATCH_RATIO_THRESHOLD # 被帧比例门槛拦截
+ # 异源安全边界(实测最小距离 18)及以上绝不匹配
+ assert not any(d <= VideoDeduplicator.PHASH_THRESHOLD for d in (18, 24, 26, 30))
# ── TestComputeFusionScore:统一融合得分方法 ────────────────────
diff --git a/tests/unit/test_prepare_dedup_1714.py b/tests/unit/test_prepare_dedup_1714.py
new file mode 100644
index 000000000..cb0ad9ff8
--- /dev/null
+++ b/tests/unit/test_prepare_dedup_1714.py
@@ -0,0 +1,471 @@
+"""#1714 prepare_direct_upload 去重 + 预建 asset 测试。
+
+覆盖 4 类用例:
+- 第一次上传:prepare 返回 duplicated=false + asset_id 非空
+- 第二次同 hash:prepare 返回 duplicated=true, skip_transfer=true
+- 同 client_upload_id 重试:prepare 也直接跳过
+- file_hash 空:走老逻辑,duplicated=false,无 asset_id
+
+以及:
+- pre-create 的 PROCESSING 占位不被"文件名兜底去重"误命中
+- _create_pending_asset find-or-create 复用现有记录
+"""
+
+from __future__ import annotations
+
+import os
+import sys
+from pathlib import Path
+from types import SimpleNamespace
+from unittest.mock import MagicMock
+
+import pytest
+
+os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret")
+os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
+sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
+
+from apps.api.app.api.routes import upload as upload_route # noqa: E402
+from packages.domain.entities import Asset, AssetLibrary, AssetLibraryKind, AssetStatus, Project # noqa: E402
+
+# ---------------------------------------------------------------------------
+# Fake repository
+# ---------------------------------------------------------------------------
+
+
+class _FakeAssetRepo:
+ """内存 asset 仓储:实现 prepare/complete 去重需要的所有方法。"""
+
+ def __init__(self):
+ self.assets = {} # id -> Asset
+ self.saved = 0
+ self.updated = 0
+
+ def create(self, asset):
+ self.assets[asset.id] = asset
+ self.saved += 1
+ return asset
+
+ def update(self, asset):
+ self.assets[asset.id] = asset
+ self.updated += 1
+ return asset
+
+ def find_by_id(self, asset_id):
+ return self.assets.get(asset_id)
+
+ def find_by_library_and_file_hash(self, library_id, file_hash):
+ if not file_hash:
+ return None
+ for a in self.assets.values():
+ if a.library_id == library_id and a.file_hash == file_hash:
+ return a
+ return None
+
+ def find_by_library_and_client_upload_id(self, library_id, client_upload_id):
+ if not client_upload_id:
+ return None
+ for a in self.assets.values():
+ if a.library_id == library_id and a.client_upload_id == client_upload_id:
+ return a
+ return None
+
+ def find_recent_active_by_library_and_name(self, library_id, name, within_minutes=30, file_size=0):
+ return None
+
+
+def _make_asset(**kw):
+ defaults = dict(
+ project_id="p-1",
+ library_id="lib-1",
+ name="existing.mp4",
+ storage_key="uploads/old/existing.mp4",
+ mime_type="video/mp4",
+ status=AssetStatus.READY,
+ file_hash="existinghash",
+ )
+ defaults.update(kw)
+ return Asset(id=defaults.pop("id", "existing-asset"), **defaults)
+
+
+def _make_pending(**kw):
+ defaults = dict(
+ project_id="p-1",
+ library_id="lib-1",
+ name="test.mp4",
+ storage_key="uploads/abc/test.mp4",
+ mime_type="video/mp4",
+ status=AssetStatus.PROCESSING,
+ file_hash="abc123",
+ )
+ defaults.update(kw)
+ return Asset(id=defaults.pop("id", "pending-asset"), **defaults)
+
+
+def _user():
+ return SimpleNamespace(user=SimpleNamespace(id="user-1"), session_id="s", token_type="t")
+
+
+class _StubProjectRepo:
+ def __init__(self, project):
+ self._p = project
+
+ def get(self, pid):
+ return self._p if self._p.id == pid else None
+
+ def find_by_id(self, pid):
+ return self._p if self._p.id == pid else None
+
+
+class _StubLibraryRepo:
+ def __init__(self, lib):
+ self._lib = lib
+
+ def find_by_project(self, pid, kind=None):
+ if self._lib.project_id == pid:
+ return [self._lib]
+ return []
+
+
+_FIXTURE_PROJECT = Project(id="p-1", owner_user_id="user-1", name="proj", description="")
+_FIXTURE_LIBRARY = AssetLibrary(
+ id="lib-1", project_id="p-1", name="videos", kind=AssetLibraryKind.VIDEO, asset_count=0, total_size=0
+)
+
+
+def _storage():
+ s = MagicMock()
+ s.create_direct_upload_post.return_value = {
+ "url": "https://bucket.oss.example.com",
+ "method": "POST",
+ "storage_key": "uploads/abc/test.mp4",
+ "expires_at": "2026-01-01T00:00:00Z",
+ "fields": {"key": "uploads/abc/test.mp4"},
+ }
+ return s
+
+
+# ---------------------------------------------------------------------------
+# 场景 1:第一次上传(无 file_hash)
+# ---------------------------------------------------------------------------
+
+
+@pytest.mark.asyncio
+async def test_prepare_first_upload_no_hash_returns_no_dedup():
+ repo = _FakeAssetRepo()
+ req = SimpleNamespace(
+ project_id="p-1",
+ library_id="lib-1",
+ filename="test.mp4",
+ content_type="video/mp4",
+ file_size=1024,
+ file_hash="",
+ client_upload_id="",
+ )
+ resp = await upload_route.prepare_direct_upload(
+ request=req,
+ authenticated_user=_user(),
+ project_repository=_StubProjectRepo(_FIXTURE_PROJECT),
+ asset_library_repository=_StubLibraryRepo(_FIXTURE_LIBRARY),
+ asset_repository=repo,
+ storage_service=_storage(),
+ )
+ assert resp.duplicated is False
+ assert resp.skip_transfer is False
+ assert resp.asset_id == "" # file_hash 空,不预建
+ assert repo.saved == 0
+
+
+# ---------------------------------------------------------------------------
+# 场景 2:第一次上传带 file_hash → duplicated=false + asset_id 非空
+# ---------------------------------------------------------------------------
+
+
+@pytest.mark.asyncio
+async def test_prepare_first_upload_with_hash_creates_pending():
+ repo = _FakeAssetRepo()
+ req = SimpleNamespace(
+ project_id="p-1",
+ library_id="lib-1",
+ filename="test.mp4",
+ content_type="video/mp4",
+ file_size=1024,
+ file_hash="abc123",
+ client_upload_id="",
+ )
+ resp = await upload_route.prepare_direct_upload(
+ request=req,
+ authenticated_user=_user(),
+ project_repository=_StubProjectRepo(_FIXTURE_PROJECT),
+ asset_library_repository=_StubLibraryRepo(_FIXTURE_LIBRARY),
+ asset_repository=repo,
+ storage_service=_storage(),
+ )
+ assert resp.duplicated is False
+ assert resp.skip_transfer is False
+ assert resp.asset_id != ""
+ # 预建记录确实落库
+ assert repo.saved == 1
+ pending = repo.find_by_id(resp.asset_id)
+ assert pending is not None
+ assert pending.file_hash == "abc123"
+ assert pending.status == AssetStatus.PROCESSING
+
+
+# ---------------------------------------------------------------------------
+# 场景 3:第二次同 hash → duplicated=true, skip_transfer=true
+# ---------------------------------------------------------------------------
+
+
+@pytest.mark.asyncio
+async def test_prepare_second_upload_same_hash_returns_duplicated():
+ repo = _FakeAssetRepo()
+ repo.create(_make_pending(file_hash="abc123", id="existing-asset"))
+ req = SimpleNamespace(
+ project_id="p-1",
+ library_id="lib-1",
+ filename="test.mp4",
+ content_type="video/mp4",
+ file_size=1024,
+ file_hash="abc123",
+ client_upload_id="",
+ )
+ resp = await upload_route.prepare_direct_upload(
+ request=req,
+ authenticated_user=_user(),
+ project_repository=_StubProjectRepo(_FIXTURE_PROJECT),
+ asset_library_repository=_StubLibraryRepo(_FIXTURE_LIBRARY),
+ asset_repository=repo,
+ storage_service=_storage(),
+ )
+ assert resp.duplicated is True
+ assert resp.skip_transfer is True
+ assert resp.asset_id == "existing-asset"
+ assert resp.upload_url == "" # 未签名 OSS
+ # 未新增记录
+ assert repo.saved == 1 # 只有初始那条
+
+
+# ---------------------------------------------------------------------------
+# 场景 4:同 client_upload_id 重试 → 直接跳过
+# ---------------------------------------------------------------------------
+
+
+@pytest.mark.asyncio
+async def test_prepare_retry_same_client_upload_id_skips():
+ repo = _FakeAssetRepo()
+ repo.create(
+ _make_pending(
+ file_hash="abc123",
+ client_upload_id="cuid-xyz",
+ id="existing-asset",
+ )
+ )
+ # 即使 file_hash 不同(理论上不会),client_upload_id 命中也直接跳过
+ req = SimpleNamespace(
+ project_id="p-1",
+ library_id="lib-1",
+ filename="test.mp4",
+ content_type="video/mp4",
+ file_size=1024,
+ file_hash="different-hash",
+ client_upload_id="cuid-xyz",
+ )
+ resp = await upload_route.prepare_direct_upload(
+ request=req,
+ authenticated_user=_user(),
+ project_repository=_StubProjectRepo(_FIXTURE_PROJECT),
+ asset_library_repository=_StubLibraryRepo(_FIXTURE_LIBRARY),
+ asset_repository=repo,
+ storage_service=_storage(),
+ )
+ assert resp.duplicated is True
+ assert resp.skip_transfer is True
+ assert resp.asset_id == "existing-asset"
+
+
+# ---------------------------------------------------------------------------
+# 兜底:文件名兜底去重不误命中 PROCESSING 占位
+# ---------------------------------------------------------------------------
+
+
+def test_filename_fallback_does_not_match_processing_pending():
+ """_find_duplicate_asset 按文件名兜底时,不能命中 pre-create 的 PROCESSING 记录。"""
+ repo = _FakeAssetRepo()
+ repo.create(_make_pending(id="p1"))
+ result = upload_route._find_duplicate_asset(
+ repo,
+ library_id="lib-1",
+ file_hash="", # 无 hash
+ client_upload_id="", # 无 cuid
+ filename="test.mp4", # 同名
+ file_size=1024,
+ )
+ assert result is None # PROCESSING 占位不被兜底命中
+
+
+def test_filename_fallback_matches_stable_ready_record():
+ """READY 状态的已存在记录能被文件名兜底命中。"""
+ repo = _FakeAssetRepo()
+ repo.create(_make_asset(status=AssetStatus.READY, id="ready-asset"))
+ # 伪造 find_recent_active_by_library_and_name 返回 READY 记录
+ repo.find_recent_active_by_library_and_name = lambda **kw: repo.assets["ready-asset"]
+ result = upload_route._find_duplicate_asset(
+ repo,
+ library_id="lib-1",
+ file_hash="",
+ client_upload_id="",
+ filename="existing.mp4",
+ file_size=1024,
+ )
+ assert result is not None
+ assert result.id == "ready-asset"
+
+
+# ---------------------------------------------------------------------------
+# _create_pending_asset find-or-create
+# ---------------------------------------------------------------------------
+
+
+def test_create_pending_asset_reuses_existing_by_hash():
+ """_create_pending_asset:file_hash 命中现有 PROCESSING 记录则复用,不新建。"""
+ repo = _FakeAssetRepo()
+ repo.create(_make_pending(file_hash="abc123", client_upload_id="", id="p1"))
+ # 复用
+ result = upload_route._create_pending_asset(
+ asset_repository=repo,
+ project_id="p-1",
+ library_id="lib-1",
+ storage_key="uploads/new/test.mp4",
+ filename="test.mp4",
+ mime_type="video/mp4",
+ user_id="user-1",
+ file_hash="abc123",
+ client_upload_id="cuid-new",
+ )
+ assert result.id == "p1"
+ assert repo.saved == 1 # 没新增
+ assert repo.updated >= 1 # 字段补齐触发 update
+ assert result.client_upload_id == "cuid-new"
+
+
+def test_create_pending_asset_creates_when_no_match():
+ """无匹配时正常新建。"""
+ repo = _FakeAssetRepo()
+ result = upload_route._create_pending_asset(
+ asset_repository=repo,
+ project_id="p-1",
+ library_id="lib-1",
+ storage_key="uploads/new/test.mp4",
+ filename="test.mp4",
+ mime_type="video/mp4",
+ user_id="user-1",
+ file_hash="newhash",
+ client_upload_id="newcuid",
+ )
+ assert result.id != ""
+ assert result.file_hash == "newhash"
+ assert result.client_upload_id == "newcuid"
+ assert repo.saved == 1
+
+
+# ---------------------------------------------------------------------------
+# 兜底去重:PROCESSING 占位 hash 不同时跳过
+# ---------------------------------------------------------------------------
+
+
+def test_filename_fallback_skips_processing_with_different_hash():
+ """PROCESSING/UPLOADING 占位记录仅当 hash 一致(或占位无 hash)才命中;hash 不同跳过。"""
+ repo = _FakeAssetRepo()
+ repo.create(_make_pending(id="p1", file_hash="oldhash"))
+ repo.find_recent_active_by_library_and_name = lambda **kw: repo.assets["p1"]
+ result = upload_route._find_duplicate_asset(
+ repo,
+ library_id="lib-1",
+ file_hash="differenthash", # 新上传内容不同
+ client_upload_id="",
+ filename="test.mp4",
+ file_size=1024,
+ )
+ assert result is None
+
+
+def test_filename_fallback_matches_processing_with_same_hash():
+ """PROCESSING 占位 hash 与请求一致时命中(重试场景)。"""
+ repo = _FakeAssetRepo()
+ repo.create(_make_pending(id="p1", file_hash="samehash"))
+ repo.find_recent_active_by_library_and_name = lambda **kw: repo.assets["p1"]
+ result = upload_route._find_duplicate_asset(
+ repo,
+ library_id="lib-1",
+ file_hash="samehash",
+ client_upload_id="",
+ filename="test.mp4",
+ file_size=1024,
+ )
+ assert result is not None
+ assert result.id == "p1"
+
+
+# ---------------------------------------------------------------------------
+# prepare 预建失败降级:不阻塞签名
+# ---------------------------------------------------------------------------
+
+
+@pytest.mark.asyncio
+async def test_prepare_pending_asset_create_failure_degrades_gracefully():
+ """预建 asset 抛异常时,prepare 仍正常返回签名(duplicated=False, asset_id 空)。"""
+
+ class _BrokenRepo(_FakeAssetRepo):
+ def create(self, asset):
+ raise RuntimeError("db down")
+
+ repo = _BrokenRepo()
+ req = SimpleNamespace(
+ project_id="p-1",
+ library_id="lib-1",
+ filename="test.mp4",
+ content_type="video/mp4",
+ file_size=1024,
+ file_hash="abc123",
+ client_upload_id="cuid-1",
+ )
+ resp = await upload_route.prepare_direct_upload(
+ request=req,
+ authenticated_user=_user(),
+ project_repository=_StubProjectRepo(_FIXTURE_PROJECT),
+ asset_library_repository=_StubLibraryRepo(_FIXTURE_LIBRARY),
+ asset_repository=repo,
+ storage_service=_storage(),
+ )
+ assert resp.duplicated is False
+ assert resp.skip_transfer is False
+ assert resp.asset_id == "" # 预建失败,降级无 asset_id
+ assert resp.upload_url != "" # 签名仍正常返回
+
+
+def test_create_pending_asset_update_failure_swallowed():
+ """复用占位记录时字段补齐 update 抛异常被吞掉,不阻塞返回。"""
+
+ class _UpdateBrokenRepo(_FakeAssetRepo):
+ def update(self, asset):
+ raise RuntimeError("db down")
+
+ repo = _UpdateBrokenRepo()
+ repo.create(_make_pending(file_hash="abc123", client_upload_id="", id="p1"))
+ result = upload_route._create_pending_asset(
+ asset_repository=repo,
+ project_id="p-1",
+ library_id="lib-1",
+ storage_key="uploads/new/test.mp4",
+ filename="test.mp4",
+ mime_type="video/mp4",
+ user_id="user-1",
+ file_hash="abc123",
+ client_upload_id="cuid-new",
+ file_size=1024,
+ )
+ assert result.id == "p1" # 仍复用,不抛异常
+ assert repo.saved == 1
diff --git a/tests/unit/test_stale_task_revoke_1714.py b/tests/unit/test_stale_task_revoke_1714.py
new file mode 100644
index 000000000..4ac086402
--- /dev/null
+++ b/tests/unit/test_stale_task_revoke_1714.py
@@ -0,0 +1,204 @@
+"""Issue #1714:孤儿/超时清理标记 failed 时必须撤销并清除 Redis 队列消息。
+
+覆盖:
+- cleanup_stale_pending_with_session_ids:超时 pending 标记 failed 并返回
+ (task_id, celery_task_id),worker 清理流程据此 revoke + purge 队列消息
+- 队列中对应业务任务的 celery 消息被物理移除(作废消息不会重投执行)
+- 旧仓储(无 _with_ids 方法)降级为计数模式,不抛异常
+- cleanup_stale_running_with_ids 同样返回 id 列表
+"""
+
+from __future__ import annotations
+
+import sys
+from datetime import datetime, timedelta, timezone
+from pathlib import Path
+from unittest.mock import MagicMock
+
+import pytest
+from sqlalchemy import create_engine, text
+from sqlalchemy.orm import sessionmaker
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
+
+from packages.adapters.sqlalchemy_impl.generation_task_repository import ( # noqa: E402
+ SQLAlchemyGenerationTaskRepository,
+)
+from packages.adapters.sqlalchemy_impl.models import Base # noqa: E402
+from packages.domain import GenerationTask, GenerationTaskStatus # noqa: E402
+
+BROKER_URL = "redis://localhost:6379/15"
+TEST_QUEUE = "_test_revoke_q"
+
+
+def _repository():
+ engine = create_engine("sqlite:///:memory:")
+ Base.metadata.create_all(engine)
+ session = sessionmaker(bind=engine)()
+ return SQLAlchemyGenerationTaskRepository(session), session, engine
+
+
+def _make_task(**kwargs) -> GenerationTask:
+ defaults = dict(project_id="proj-1", asset_library_id="lib-1", created_by_user_id="user-1")
+ defaults.update(kwargs)
+ return GenerationTask.create(**defaults)
+
+
+def _redis_available() -> bool:
+ try:
+ import redis
+
+ return bool(redis.Redis.from_url(BROKER_URL).ping())
+ except Exception:
+ return False
+
+
+# ── 仓储层:返回 ids ────────────────────────────────────────────────────
+
+
+def test_cleanup_stale_pending_returns_ids_with_celery_task_id():
+ repo, _, engine = _repository()
+ task = _make_task()
+ task.celery_task_id = "celery-msg-id-001"
+ repo.create(task)
+ # created_at 改到 60 分钟前
+ with engine.connect() as conn:
+ conn.execute(
+ text("UPDATE generation_tasks SET created_at = :ts WHERE id = :id"),
+ {"ts": datetime.now(timezone.utc) - timedelta(minutes=60), "id": task.id},
+ )
+ conn.commit()
+
+ items = repo.cleanup_stale_pending_with_ids(timeout_minutes=45)
+ assert len(items) == 1
+ biz_id, celery_id = items[0]
+ assert biz_id == task.id
+ assert celery_id == "celery-msg-id-001"
+
+ saved = repo.get(task.id)
+ assert saved.status == GenerationTaskStatus.FAILED
+
+
+def test_cleanup_stale_running_returns_ids():
+ repo, _, engine = _repository()
+ task = _make_task()
+ repo.create(task)
+ task.mark_processing()
+ task.celery_task_id = "celery-msg-id-002"
+ repo.update(task)
+ with engine.connect() as conn:
+ conn.execute(
+ text("UPDATE generation_tasks SET updated_at = :ts WHERE id = :id"),
+ {"ts": datetime.now(timezone.utc) - timedelta(minutes=60), "id": task.id},
+ )
+ conn.commit()
+
+ items = repo.cleanup_stale_running_with_ids(timeout_minutes=20)
+ assert len(items) == 1
+ assert items[0][0] == task.id
+ assert items[0][1] == "celery-msg-id-002"
+ assert repo.get(task.id).status == GenerationTaskStatus.FAILED
+
+
+def test_legacy_repo_without_with_ids_falls_back_to_count():
+ """旧仓储只有 cleanup_stale_pending(返回 int)时降级可用,不抛异常。"""
+ # worker 模块加载(标准 mock 模式)
+ saved = set(sys.modules.keys())
+ mock_db = MagicMock()
+ mock_db.SessionLocal = MagicMock()
+ sys.modules["worker_app.db"] = mock_db
+ sys.modules["worker_app.core.config"] = MagicMock()
+ mock_celery = MagicMock()
+ mock_celery.celery_app.task = MagicMock(
+ side_effect=(lambda *a, **k: (a[0] if a and callable(a[0]) else (lambda f: f)))
+ )
+ sys.modules["worker_app.celery_app"] = mock_celery
+ worker_path = str(Path(__file__).resolve().parents[2] / "apps" / "worker")
+ if worker_path not in sys.path:
+ sys.path.insert(0, worker_path)
+
+ from worker_app.tasks import _startup # noqa: E402
+
+ class LegacyRepo:
+ def cleanup_stale_pending(self, timeout_minutes): # noqa: ARG002
+ return 3
+
+ def cleanup_stale_running(self, timeout_minutes): # noqa: ARG002
+ return 2
+
+ items_p = _startup.cleanup_stale_pending_with_session_ids(LegacyRepo(), 45)
+ items_r = _startup.cleanup_stale_running_with_session_ids(LegacyRepo(), 20)
+ assert len(items_p) == 3
+ assert len(items_r) == 2
+
+ for key in list(sys.modules.keys()):
+ if key not in saved and not key.startswith("video_processing"):
+ del sys.modules[key]
+
+
+# ── 端到端:清理 → 队列消息被移除(作废消息不重投) ────────────────────
+
+
+@pytest.mark.skipif(not _redis_available(), reason="本地 redis 不可用")
+def test_stale_pending_cleanup_purges_redis_message():
+ """任务标 failed 后,其在 Redis 队列里的 celery 消息被清除,不会被重投。"""
+ import redis
+ from celery import Celery
+ from kombu import Queue
+ from kombu.pools import producers
+
+ from packages.shared.celery_orphan_guard import purge_stale_messages_from_queues
+
+ repo, _, engine = _repository()
+ task = _make_task()
+ task.celery_task_id = "celery-stale-xyz"
+ repo.create(task)
+ with engine.connect() as conn:
+ conn.execute(
+ text("UPDATE generation_tasks SET created_at = :ts WHERE id = :id"),
+ {"ts": datetime.now(timezone.utc) - timedelta(minutes=60), "id": task.id},
+ )
+ conn.commit()
+
+ # 模拟该任务的 celery 消息仍在 generation 队列里(worker 下线期间未消费)
+ client = redis.Redis.from_url(BROKER_URL)
+ client.delete(TEST_QUEUE)
+ app = Celery("test-e2e-revoke")
+ app.conf.broker_url = BROKER_URL
+ with app.connection_for_write() as conn:
+ with producers[conn].acquire(block=True) as prod:
+ # 作废任务消息
+ prod.publish(
+ (task.id,),
+ exchange="",
+ routing_key=TEST_QUEUE,
+ serializer="json",
+ headers={"id": "celery-stale-xyz", "task": "worker.generate_video"},
+ retry=False,
+ delivery_mode=1,
+ declare=[Queue(TEST_QUEUE, routing_key=TEST_QUEUE, durable=False)],
+ )
+ # 另一条正常任务消息(必须保留)
+ prod.publish(
+ ("other-task-id",),
+ exchange="",
+ routing_key=TEST_QUEUE,
+ serializer="json",
+ headers={"id": "celery-keep", "task": "worker.generate_video"},
+ retry=False,
+ delivery_mode=1,
+ )
+
+ assert client.llen(TEST_QUEUE) == 2
+
+ # 执行清理(与 worker beat 相同流程:标 failed → 拿 ids → purge)
+ items = repo.cleanup_stale_pending_with_ids(timeout_minutes=45)
+ biz_ids = [bid for bid, _ in items]
+ celery_ids = [cid for _, cid in items if cid]
+ removed = purge_stale_messages_from_queues(
+ BROKER_URL, (TEST_QUEUE,), business_task_ids=biz_ids, celery_task_ids=celery_ids
+ )
+
+ assert removed == 1
+ assert client.llen(TEST_QUEUE) == 1 # 正常任务消息保留
+ client.delete(TEST_QUEUE)
diff --git a/tests/unit/test_task_discard_guard_1714.py b/tests/unit/test_task_discard_guard_1714.py
new file mode 100644
index 000000000..faa1cfd96
--- /dev/null
+++ b/tests/unit/test_task_discard_guard_1714.py
@@ -0,0 +1,243 @@
+"""Issue #1714:任务执行前状态守卫 — 已作废消息必须丢弃,禁止非法转换后继续跑。
+
+覆盖:
+- ingest_asset:job 已 failed/completed 时直接返回 discarded,不下载、不转码、不回写
+- generate_video:GenerationTask 已 failed 时返回 discarded,不进入渲染
+- generate_video:pending → running 标记失败(非法转换)时安全中止
+"""
+
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+from types import SimpleNamespace
+from unittest.mock import MagicMock
+
+# ── worker 模块标准加载方式 ──
+# 显式保存将要覆盖的注入键旧值:全量收集时更早的测试文件(如
+# test_ingest_validation.py)可能已向 sys.modules 注入 worker_app.* mock,
+# 导入完成后必须精确恢复旧值,否则本文件的 bind 感知透传装饰器会残留,
+# 污染后续懒加载短路径 worker_app.celery_app 的 worker 测试。
+_INJECTED_KEYS = ("worker_app.db", "worker_app.core.config", "worker_app.celery_app")
+_SAVED_MODULE_VALUES = {k: sys.modules.get(k) for k in _INJECTED_KEYS}
+_SAVED_MODULES_KEYS = set(sys.modules.keys())
+
+_mock_db_module = MagicMock()
+_mock_db_module.SessionLocal = MagicMock()
+sys.modules["worker_app.db"] = _mock_db_module
+sys.modules["worker_app.core.config"] = MagicMock()
+
+_mock_celery_module = MagicMock()
+
+
+def _passthrough_decorator(*args, **kwargs):
+ if len(args) == 1 and callable(args[0]):
+ return args[0]
+ bind = kwargs.get("bind", False)
+
+ def _wrap(f):
+ if bind:
+ # 模拟 celery bind=True:task(task_id) 调用时注入 self(MagicMock)
+ return lambda *a, **kw: f(MagicMock(), *a, **kw)
+ return f
+
+ return _wrap
+
+
+_mock_celery_module.celery_app.task = MagicMock(side_effect=_passthrough_decorator)
+sys.modules["worker_app.celery_app"] = _mock_celery_module
+
+_WORKER_PATH = str(Path(__file__).resolve().parents[2] / "apps" / "worker")
+sys.path.insert(0, _WORKER_PATH)
+
+import pytest # noqa: E402
+from worker_app.tasks import ingest as ingest_mod # noqa: E402
+
+# video_processing 相关 mock(generation 模块导入链)
+for _mod_name in [
+ "video_processing",
+ "video_processing.ffmpeg_utils",
+ "video_processing.oss_helpers",
+]:
+ sys.modules.setdefault(_mod_name, MagicMock())
+
+from worker_app.tasks import generation as gen_mod # noqa: E402
+
+from packages.domain import IngestJobStatus # noqa: E402
+
+# 模块导入完成后立即清理:删除本次 import 新引入的模块缓存(本模块已通过名字绑定
+# 持有 ingest_mod/gen_mod/IngestJobStatus,删除缓存不影响调用),再把三个注入键
+# 精确恢复为注入前的旧值(旧值不存在则移除),杜绝 mock 残留污染其他 worker 测试。
+for _key in list(sys.modules.keys()):
+ if _key not in _SAVED_MODULES_KEYS and not _key.startswith("video_processing"):
+ del sys.modules[_key]
+for _k, _v in _SAVED_MODULE_VALUES.items():
+ if _v is None:
+ sys.modules.pop(_k, None)
+ else:
+ sys.modules[_k] = _v
+del _SAVED_MODULES_KEYS, _SAVED_MODULE_VALUES
+
+
+# ── ingest 守卫 ────────────────────────────────────────────────────────
+
+
+class _FakeJobRepo:
+ def __init__(self, job):
+ self.job = job
+
+ def get(self, job_id):
+ return self.job
+
+
+def _make_ingest_job(status):
+ job = MagicMock()
+ job.id = "job-stale-1"
+ job.storage_key = "uploads/proj/stale.mov"
+ job.status = status
+ job.file_hash = "h"
+ job.asset_id = ""
+ return job
+
+
+def test_ingest_discards_failed_job_message():
+ """job 已 failed:消息丢弃,不进入下载/转码/回写。"""
+ job = _make_ingest_job(IngestJobStatus.FAILED)
+ fake_session = MagicMock()
+ _mock_db_module.SessionLocal = MagicMock(return_value=fake_session)
+
+ # SQLAlchemy 仓储构造返回 fake
+ fake_job_repo = _FakeJobRepo(job)
+ fake_asset_repo = MagicMock()
+
+ orig_job_repo = ingest_mod.SQLAlchemyIngestJobRepository
+ orig_asset_repo = ingest_mod.SQLAlchemyAssetRepository
+ ingest_mod.SQLAlchemyIngestJobRepository = MagicMock(return_value=fake_job_repo)
+ ingest_mod.SQLAlchemyAssetRepository = MagicMock(return_value=fake_asset_repo)
+ try:
+ result = ingest_mod.ingest_asset("job-stale-1")
+ finally:
+ ingest_mod.SQLAlchemyIngestJobRepository = orig_job_repo
+ ingest_mod.SQLAlchemyAssetRepository = orig_asset_repo
+
+ assert result["status"] == "discarded"
+ # 没有任何 update / commit / 下载动作
+ fake_session.commit.assert_not_called()
+ fake_asset_repo.create.assert_not_called()
+
+
+def test_ingest_discards_completed_job_message():
+ job = _make_ingest_job(IngestJobStatus.COMPLETED)
+ fake_session = MagicMock()
+ _mock_db_module.SessionLocal = MagicMock(return_value=fake_session)
+ fake_job_repo = _FakeJobRepo(job)
+
+ orig = ingest_mod.SQLAlchemyIngestJobRepository
+ ingest_mod.SQLAlchemyIngestJobRepository = MagicMock(return_value=fake_job_repo)
+ ingest_mod.SQLAlchemyAssetRepository = MagicMock(return_value=MagicMock())
+ try:
+ result = ingest_mod.ingest_asset("job-stale-1")
+ finally:
+ ingest_mod.SQLAlchemyIngestJobRepository = orig
+
+ assert result["status"] == "discarded"
+
+
+# ── generation 守卫 ────────────────────────────────────────────────────
+
+
+def _make_gen_task(status_value: str):
+ from packages.domain import GenerationTask
+
+ task = GenerationTask.create(project_id="p", asset_library_id="l", created_by_user_id="u")
+ task.status = type(task.status)(status_value)
+ return task
+
+
+def test_generate_video_discards_failed_task(monkeypatch):
+ """GenerationTask 已 failed:直接 discarded,不加载渲染数据。"""
+ failed_task = _make_gen_task("failed")
+
+ fake_repo = MagicMock()
+ fake_repo.get.return_value = failed_task
+
+ fake_session = MagicMock()
+ _mock_db_module.SessionLocal = MagicMock(return_value=fake_session)
+
+ import packages.adapters.sqlalchemy_impl.generation_task_repository as gen_repo_mod
+
+ orig = gen_repo_mod.SQLAlchemyGenerationTaskRepository
+ gen_repo_mod.SQLAlchemyGenerationTaskRepository = MagicMock(return_value=fake_repo)
+
+ update_status_mock = MagicMock(return_value=False)
+ monkeypatch.setattr(gen_mod, "_update_task_status", update_status_mock)
+ monkeypatch.setattr(
+ gen_mod,
+ "_load_task_info",
+ lambda task_id: {
+ "project_id": "p",
+ "template_id": "",
+ "task_asset_ids": [],
+ "batch_id": "",
+ "user_id": "u",
+ "mode": "one_take",
+ },
+ )
+ monkeypatch.setattr(gen_mod, "_flush_logs", lambda *a, **k: None)
+
+ task_fn = gen_mod.generate_video
+ if hasattr(task_fn, "__wrapped__"):
+ task_fn = task_fn.__wrapped__
+ try:
+ result = task_fn("task-stale-1")
+ finally:
+ gen_repo_mod.SQLAlchemyGenerationTaskRepository = orig
+
+ assert result["status"] == "discarded"
+ # 状态守卫命中终态,根本不应尝试 mark_processing
+ update_status_mock.assert_not_called()
+
+
+def test_generate_video_aborts_when_claim_fails(monkeypatch):
+ """pending 但 mark_processing 返回 False(状态机非法转换)时安全中止。"""
+ pending_task = _make_gen_task("pending")
+
+ fake_repo = MagicMock()
+ fake_repo.get.return_value = pending_task
+ fake_session = MagicMock()
+ _mock_db_module.SessionLocal = MagicMock(return_value=fake_session)
+
+ import packages.adapters.sqlalchemy_impl.generation_task_repository as gen_repo_mod
+
+ orig = gen_repo_mod.SQLAlchemyGenerationTaskRepository
+ gen_repo_mod.SQLAlchemyGenerationTaskRepository = MagicMock(return_value=fake_repo)
+
+ monkeypatch.setattr(
+ gen_mod,
+ "_load_task_info",
+ lambda task_id: {
+ "project_id": "p",
+ "template_id": "",
+ "task_asset_ids": [],
+ "batch_id": "",
+ "user_id": "u",
+ "mode": "one_take",
+ },
+ )
+ monkeypatch.setattr(gen_mod, "_flush_logs", lambda *a, **k: None)
+ # 模拟 mark_processing 失败(failed→running 非法转换被 _update_task_status 吞掉返回 False)
+ update_status_mock = MagicMock(return_value=False)
+ monkeypatch.setattr(gen_mod, "_update_task_status", update_status_mock)
+ render_mock = MagicMock(side_effect=AssertionError("must not render"))
+ monkeypatch.setattr(gen_mod, "_render_from_edit_plan", render_mock)
+
+ task_fn = gen_mod.generate_video
+ if hasattr(task_fn, "__wrapped__"):
+ task_fn = task_fn.__wrapped__
+ try:
+ result = task_fn("task-claim-fail")
+ finally:
+ gen_repo_mod.SQLAlchemyGenerationTaskRepository = orig
+
+ assert result["status"] == "discarded"
+ render_mock.assert_not_called()
diff --git a/tests/unit/test_task_fault_tolerance_1709.py b/tests/unit/test_task_fault_tolerance_1709.py
new file mode 100644
index 000000000..756543d2a
--- /dev/null
+++ b/tests/unit/test_task_fault_tolerance_1709.py
@@ -0,0 +1,315 @@
+"""Issue #1709 任务容错:孤儿任务恢复 + 429 限流结构化提示。
+
+覆盖:
+1. 仓储层:count_running_by_user/count_running_total 计数正确(预览/正式任务都计入)
+2. 仓储层:estimate_avg_duration_seconds 耗时估算(有历史/无历史)
+3. 限流核心:build_rate_limit_detail 返回结构化 code/message/排队数/预计等待
+4. worker 侧:cleanup_stale_running/pending 核心函数——中断任务被重置为 failed
+ 且原因写明(容器重启/超时中断),正常任务不受影响
+"""
+
+import sys
+from datetime import datetime, timedelta, timezone
+from pathlib import Path
+from unittest.mock import MagicMock
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "worker"))
+
+# 预注入 mock worker_app.db,防止真实数据库连接初始化(与其他 worker 测试同模式)
+_mock_db = MagicMock()
+_mock_db.SessionLocal = MagicMock()
+sys.modules.setdefault("worker_app.db", _mock_db)
+
+from app.core import task_enqueue # noqa: E402
+from sqlalchemy import create_engine, text # noqa: E402
+from sqlalchemy.orm import sessionmaker # noqa: E402
+from worker_app.tasks import _startup # noqa: E402
+
+from packages.adapters.sqlalchemy_impl.generation_task_repository import ( # noqa: E402
+ SQLAlchemyGenerationTaskRepository,
+)
+from packages.adapters.sqlalchemy_impl.models import Base # noqa: E402
+from packages.domain import GenerationTask, GenerationTaskStatus # noqa: E402
+
+
+def _repository():
+ engine = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False})
+ Base.metadata.create_all(engine)
+ session = sessionmaker(bind=engine)()
+ return SQLAlchemyGenerationTaskRepository(session), session, engine
+
+
+def _make_task(**kwargs) -> GenerationTask:
+ defaults = dict(
+ project_id="proj-1",
+ asset_library_id="lib-1",
+ created_by_user_id="user-1",
+ )
+ defaults.update(kwargs)
+ return GenerationTask.create(**defaults)
+
+
+def _age_task(engine, task_id, *, updated_minutes=None, created_minutes=None):
+ """用 SQL 直接把 updated_at/created_at 改到过去(模拟孤儿任务)。"""
+ sets, params = [], {"id": task_id}
+ if updated_minutes is not None:
+ sets.append("updated_at = :uts")
+ params["uts"] = datetime.now(timezone.utc) - timedelta(minutes=updated_minutes)
+ if created_minutes is not None:
+ sets.append("created_at = :cts")
+ params["cts"] = datetime.now(timezone.utc) - timedelta(minutes=created_minutes)
+ with engine.connect() as conn:
+ conn.execute(text(f"UPDATE generation_tasks SET {', '.join(sets)} WHERE id = :id"), params)
+ conn.commit()
+
+
+# ---------------------------------------------------------------------------
+# 1. running 计数(限流"渲染中"数量)
+# ---------------------------------------------------------------------------
+
+
+def test_count_running_by_user_mix_statuses():
+ """count_running_by_user 只统计该用户 running,不含 pending/completed/failed。"""
+ repo, _, _ = _repository()
+ t1 = _make_task(project_id="p1")
+ repo.create(t1) # pending
+ t2 = _make_task(project_id="p2")
+ repo.create(t2)
+ t2.mark_processing()
+ repo.update(t2)
+ t3 = _make_task(project_id="p3")
+ repo.create(t3)
+ t3.mark_processing()
+ repo.update(t3)
+ t4 = _make_task(project_id="p4")
+ repo.create(t4)
+ t4.mark_processing()
+ repo.update(t4)
+ t4.mark_completed()
+ repo.update(t4)
+ t5 = _make_task(project_id="p5", created_by_user_id="user-2")
+ repo.create(t5)
+ t5.mark_processing()
+ repo.update(t5)
+
+ assert repo.count_running_by_user("user-1") == 2
+ assert repo.count_running_by_user("user-2") == 1
+ assert repo.count_running_total() == 3
+
+
+def test_count_running_total_empty():
+ repo, _, _ = _repository()
+ assert repo.count_running_total() == 0
+ assert repo.count_running_by_user("nobody") == 0
+
+
+def test_preview_tasks_counted_in_running():
+ """预览任务(is_preview=True,工单实测卡 80% 的那种)同样计入 running。"""
+ repo, _, _ = _repository()
+ t = _make_task(is_preview=True)
+ repo.create(t)
+ t.mark_processing()
+ repo.update(t)
+ assert repo.count_running_by_user("user-1") == 1
+ assert repo.count_running_total() == 1
+
+
+# ---------------------------------------------------------------------------
+# 2. 平均耗时估算(429 等待预估依据)
+# ---------------------------------------------------------------------------
+
+
+def _complete_task(repo, engine, task, duration_seconds: float):
+ repo.create(task)
+ task.mark_processing()
+ repo.update(task)
+ task.mark_completed()
+ repo.update(task)
+ now = datetime.now(timezone.utc)
+ with engine.connect() as conn:
+ conn.execute(
+ text("UPDATE generation_tasks SET started_at = :s, completed_at = :c WHERE id = :id"),
+ {"s": now - timedelta(seconds=duration_seconds), "c": now, "id": task.id},
+ )
+ conn.commit()
+
+
+def test_estimate_avg_duration_with_history():
+ """有历史完成任务时返回平均耗时(秒)。"""
+ repo, _, engine = _repository()
+ _complete_task(repo, engine, _make_task(project_id="p1"), 60.0)
+ _complete_task(repo, engine, _make_task(project_id="p2"), 180.0)
+
+ avg = repo.estimate_avg_duration_seconds(default_seconds=120.0)
+ assert 119.0 < avg < 121.0 # (60+180)/2 = 120
+
+
+def test_estimate_avg_duration_no_history_returns_default():
+ """无历史数据时返回默认值。"""
+ repo, _, _ = _repository()
+ assert repo.estimate_avg_duration_seconds(default_seconds=90.0) == 90.0
+
+
+# ---------------------------------------------------------------------------
+# 3. build_rate_limit_detail 结构化提示(前端区分"排队"与"创建失败")
+# ---------------------------------------------------------------------------
+
+
+def test_user_rate_limit_detail_structure():
+ """429 用户限流:返回 USER_QUEUE_FULL + 排队/渲染数 + 预计等待。"""
+ repo, _, _ = _repository()
+ for i in range(2): # 2 个渲染中
+ t = _make_task(project_id=f"rp{i}")
+ repo.create(t)
+ t.mark_processing()
+ repo.update(t)
+
+ exc = task_enqueue.UserPendingLimitExceeded(user_id="user-1", pending_count=3, limit=3)
+ detail = task_enqueue.build_rate_limit_detail(exc, repo, scope="user")
+
+ assert detail["code"] == task_enqueue.ERROR_CODE_USER_QUEUE_FULL
+ assert detail["queued_count"] == 3
+ assert detail["running_count"] == 2
+ assert detail["limit"] == 3
+ assert detail["estimated_wait_seconds"] > 0
+ assert "排队" in detail["message"]
+ assert "user-1" not in detail["message"] # 不泄露内部 ID
+
+
+def test_global_rate_limit_detail_structure():
+ """503 全局繁忙:返回 SYSTEM_QUEUE_FULL。"""
+ repo, _, _ = _repository()
+ exc = task_enqueue.GlobalQueueFull(pending_count=20, limit=20)
+ detail = task_enqueue.build_rate_limit_detail(exc, repo, scope="global")
+
+ assert detail["code"] == task_enqueue.ERROR_CODE_SYSTEM_QUEUE_FULL
+ assert detail["queued_count"] == 20
+ assert detail["limit"] == 20
+ assert detail["estimated_wait_seconds"] > 0
+ assert "系统繁忙" in detail["message"]
+
+
+def test_wait_estimate_uses_concurrency():
+ """等待预估:排队 8 个 / 并发 4 = 2 批 × 平均耗时。"""
+
+ class FakeRepo:
+ def estimate_avg_duration_seconds(self, limit=20, default_seconds=120.0):
+ return 100.0
+
+ wait = task_enqueue._estimate_wait_seconds(8, FakeRepo())
+ assert wait == 200 # ceil(8/4)=2 批 × 100 秒
+
+
+def test_wait_estimate_repo_without_methods_uses_default():
+ """仓储没有新方法(旧 mock/鸭子类型)时用默认 120 秒兜底,不抛错。"""
+
+ class LegacyRepo:
+ """只实现旧接口的仓储(模拟未升级的调用方)。"""
+
+ def count_pending_total(self):
+ return 0
+
+ wait = task_enqueue._estimate_wait_seconds(4, LegacyRepo())
+ assert wait == 120 # ceil(4/4)=1 批 × 120 默认
+
+
+def test_rate_limit_detail_running_count_falls_back_to_zero():
+ """仓储不支持 running 计数时,running_count 优雅降级为 0。"""
+
+ class LegacyRepo:
+ def count_pending_total(self):
+ return 0
+
+ exc = task_enqueue.GlobalQueueFull(pending_count=20, limit=20)
+ detail = task_enqueue.build_rate_limit_detail(exc, LegacyRepo(), scope="global")
+ assert detail["running_count"] == 0
+ assert detail["code"] == task_enqueue.ERROR_CODE_SYSTEM_QUEUE_FULL
+
+
+# ---------------------------------------------------------------------------
+# 4. worker 清理核心:中断任务被重置(worker 重启/超时恢复)
+# ---------------------------------------------------------------------------
+
+
+def test_worker_cleanup_resets_interrupted_running_task():
+ """模拟 worker 重启:running 超 20 分钟无更新的任务被重置为 failed,原因写明。"""
+ repo, _, engine = _repository()
+
+ t = _make_task(is_preview=True) # 预览任务
+ repo.create(t)
+ t.mark_processing() # running
+ repo.update(t)
+ _age_task(engine, t.id, updated_minutes=25) # 25 分钟无进度更新
+
+ cleaned = _startup.cleanup_stale_running_with_session(repo, 20)
+ assert cleaned == 1
+
+ saved = repo.get(t.id)
+ assert saved.status == GenerationTaskStatus.FAILED
+ assert "中断" in saved.error_message
+ assert saved.error_info.get("error_type") == "WorkerInterrupted"
+ assert saved.completed_at is not None
+
+
+def test_worker_cleanup_keeps_healthy_running_task():
+ """正常运行中(5 分钟前有更新)的任务不被误杀。"""
+ repo, _, engine = _repository()
+
+ t = _make_task()
+ repo.create(t)
+ t.mark_processing()
+ repo.update(t)
+ _age_task(engine, t.id, updated_minutes=5)
+
+ assert _startup.cleanup_stale_running_with_session(repo, 20) == 0
+ assert repo.get(t.id).status == GenerationTaskStatus.RUNNING
+
+
+def test_worker_cleanup_resets_stale_pending_task():
+ """卡 pending 超 15 分钟(worker 停止消费)的任务被重置,释放限流名额。"""
+ repo, _, engine = _repository()
+
+ t = _make_task(is_preview=True)
+ repo.create(t) # 一直 pending
+ _age_task(engine, t.id, created_minutes=20)
+
+ cleaned = _startup.cleanup_stale_pending_with_session(repo, 15)
+ assert cleaned == 1
+
+ saved = repo.get(t.id)
+ assert saved.status == GenerationTaskStatus.FAILED
+ assert saved.error_info.get("error_type") == "PendingTimeout"
+ # 释放名额后 pending 计数归零,新请求不再被 429 误伤
+ assert repo.count_pending_total() == 0
+
+
+def test_worker_cleanup_pending_keeps_recent():
+ """刚创建 3 分钟的 pending 任务不清理。"""
+ repo, _, engine = _repository()
+
+ t = _make_task()
+ repo.create(t)
+ _age_task(engine, t.id, created_minutes=3)
+
+ assert _startup.cleanup_stale_pending_with_session(repo, 15) == 0
+ assert repo.get(t.id).status == GenerationTaskStatus.PENDING
+
+
+def test_worker_cleanup_multiple_orphans_all_reset():
+ """3 个卡死 running 任务(工单实测:3 个预览卡 80% 超 10 小时)全部恢复。"""
+ repo, _, engine = _repository()
+
+ ids = []
+ for i in range(3):
+ t = _make_task(project_id=f"p{i}", is_preview=True)
+ repo.create(t)
+ t.mark_processing()
+ repo.update(t)
+ _age_task(engine, t.id, updated_minutes=600) # 10 小时
+ ids.append(t.id)
+
+ cleaned = _startup.cleanup_stale_running_with_session(repo, 20)
+ assert cleaned == 3
+ for tid in ids:
+ assert repo.get(tid).status == GenerationTaskStatus.FAILED
diff --git a/tests/unit/test_task_queue_limit.py b/tests/unit/test_task_queue_limit.py
index dc4e56b41..b78baaf96 100644
--- a/tests/unit/test_task_queue_limit.py
+++ b/tests/unit/test_task_queue_limit.py
@@ -164,7 +164,9 @@ class TestSafeEnqueueWithLimits:
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
assert result is True
mock_celery.assert_called_once_with("worker.generate_video", args=["task-1"])
- assert len(repo.updated_tasks) == 0 # 成功不需要更新状态
+ # 成功入队后持久化 celery 消息 ID(#1714:孤儿清理据此 revoke/清队列)
+ assert len(repo.updated_tasks) == 1
+ assert task.celery_task_id
def test_user_limit_rejected_with_failed_status(self, mock_celery):
"""用户超限:任务标记为 failed,抛 UserPendingLimitExceeded。"""
@@ -335,7 +337,9 @@ class TestPostEnqueueFinalCheck:
assert result is True
mock_celery.assert_called_once()
assert task.status == "pending" # 状态没变
- assert len(repo.updated_tasks) == 0 # 没更新 DB
+ # 入队成功后持久化 celery_task_id(#1714),业务状态不变
+ assert len(repo.updated_tasks) == 1
+ assert task.celery_task_id
def test_post_enqueue_no_user_id_skips_user_check(self, mock_celery):
"""不传 user_id 时,入队后校验也跳过用户级,只查全局。"""
diff --git a/tests/unit/test_upload_complete_idempotency_1714.py b/tests/unit/test_upload_complete_idempotency_1714.py
new file mode 100644
index 000000000..c0436b065
--- /dev/null
+++ b/tests/unit/test_upload_complete_idempotency_1714.py
@@ -0,0 +1,422 @@
+"""Issue #1714:POST /upload/direct/complete 幂等 + multipart 幂等。
+
+覆盖:
+- 同 client_upload_id 重复 complete → 只建一条 asset、不重复派 ingest job
+- 同 file_hash 重复 complete → 返回已存在记录
+- 旧客户端不传 hash/token:近期同库同名 processing 占位 → 兜底幂等返回
+- 旧客户端不传 hash/token:READY 历史同名 → 不兜底(正常新建)
+- 兜底窗口外(>30 分钟)→ 不兜底
+- 旧仓储(无新方法)鸭子类型降级 → 不报错、正常新建
+- 重复 complete 时即使 OSS 已无文件(file_exists=False)也返回已存在记录
+ (模拟 complete 超时后 OSS 侧对象已过期/清理,重试仍不重复建库)
+- multipart 上传:同 client_upload_id 重复提交 → 第二次直接 duplicated,不再传 OSS
+"""
+
+from __future__ import annotations
+
+import os
+import sys
+from datetime import datetime, timedelta, timezone
+from pathlib import Path
+from unittest.mock import MagicMock
+
+os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
+os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
+
+from fastapi import FastAPI # noqa: E402
+from fastapi.testclient import TestClient # noqa: E402
+
+from packages.domain import Asset, AssetLibrary, AssetLibraryKind, AssetStatus, IngestJob, Project # noqa: E402
+
+
+class StubProjectRepository:
+ def __init__(self, projects: dict | None = None):
+ self._projects = projects or {}
+
+ def get(self, project_id: str):
+ return self._projects.get(project_id)
+
+ def find_by_id(self, project_id: str):
+ return self._projects.get(project_id)
+
+
+class StubAssetLibraryRepository:
+ def __init__(self, libraries: dict | None = None):
+ self._libraries = libraries or {}
+
+ def find_by_project(self, project_id: str, kind=None) -> list:
+ return list(self._libraries.values())
+
+
+class StubAssetRepository:
+ """支持三种幂等查询的内存仓储,并统计 create 次数。"""
+
+ def __init__(self, assets: list[Asset] | None = None):
+ self._assets = list(assets or [])
+ self.created: list[Asset] = []
+
+ def find_by_library_and_file_hash(self, library_id: str, file_hash: str) -> Asset | None:
+ if not file_hash:
+ return None
+ return next((a for a in self._assets if a.library_id == library_id and a.file_hash == file_hash), None)
+
+ def find_by_library_and_client_upload_id(self, library_id: str, client_upload_id: str) -> Asset | None:
+ if not client_upload_id:
+ return None
+ return next(
+ (a for a in self._assets if a.library_id == library_id and a.client_upload_id == client_upload_id),
+ None,
+ )
+
+ def find_recent_active_by_library_and_name(
+ self, library_id: str, name: str, within_minutes: int = 30, file_size: int = 0
+ ) -> Asset | None:
+ # 严格模式(#1714):大小未知(0)直接不命中,宁可漏判不可误杀
+ if not file_size or file_size <= 0:
+ return None
+ cutoff = datetime.now(timezone.utc) - timedelta(minutes=within_minutes)
+ candidates = [
+ a
+ for a in self._assets
+ if a.library_id == library_id
+ and a.name == name
+ and a.status in (AssetStatus.UPLOADING, AssetStatus.PROCESSING)
+ and a.created_at >= cutoff
+ and a.file_size == file_size
+ ]
+ return max(candidates, key=lambda a: a.created_at) if candidates else None
+
+ def create(self, asset: Asset) -> Asset:
+ self._assets.append(asset)
+ self.created.append(asset)
+ return asset
+
+ def update(self, asset: Asset) -> Asset:
+ return asset
+
+
+class LegacyStubAssetRepository:
+ """旧仓储:只有 file_hash 去重,没有新方法(鸭子类型降级验证)。"""
+
+ def __init__(self, assets: list[Asset] | None = None):
+ self._assets = list(assets or [])
+ self.created: list[Asset] = []
+
+ def find_by_library_and_file_hash(self, library_id: str, file_hash: str) -> Asset | None:
+ if not file_hash:
+ return None
+ return next((a for a in self._assets if a.library_id == library_id and a.file_hash == file_hash), None)
+
+ def create(self, asset: Asset) -> Asset:
+ self._assets.append(asset)
+ self.created.append(asset)
+ return asset
+
+
+class StubIngestJobRepository:
+ def __init__(self):
+ self._jobs: dict[str, IngestJob] = {}
+ self.created_count = 0
+
+ def create(self, job: IngestJob) -> IngestJob:
+ self._jobs[job.id] = job
+ self.created_count += 1
+ return job
+
+ def get(self, job_id: str) -> IngestJob | None:
+ return self._jobs.get(job_id)
+
+ def update(self, job: IngestJob) -> IngestJob:
+ self._jobs[job.id] = job
+ return job
+
+
+def _make_project() -> Project:
+ return Project(id="proj-1", name="Test Project", owner_user_id="user-1")
+
+
+def _make_library() -> AssetLibrary:
+ return AssetLibrary(id="lib-1", name="Test Library", project_id="proj-1", kind=AssetLibraryKind.VIDEO)
+
+
+def _build_app(asset_repo=None, ingest_repo=None, storage=None):
+ from app.api.routes.upload import router
+ from app.auth import AuthenticatedUser, get_current_user
+ from app.core.storage import get_storage_service
+ from app.dependencies import (
+ get_asset_library_repository,
+ get_asset_repository,
+ get_ingest_job_repository,
+ get_project_repository,
+ )
+
+ app = FastAPI()
+ app.include_router(router, prefix="/api/v1")
+
+ project_repo = StubProjectRepository({"proj-1": _make_project()})
+ library_repo = StubAssetLibraryRepository({"lib-1": _make_library()})
+ asset_repo = asset_repo or StubAssetRepository()
+ ingest_repo = ingest_repo or StubIngestJobRepository()
+
+ storage = storage or MagicMock()
+ storage.is_configured = True
+ storage._normalize_storage_key = lambda key: key
+ storage.file_exists = MagicMock(return_value=True)
+ storage.upload_file = MagicMock(return_value="https://oss.example.com/file.mp4")
+ storage.get_url = MagicMock(return_value="https://oss.example.com/file.mp4")
+
+ mock_user = MagicMock(spec=AuthenticatedUser)
+ mock_user.id = "user-1"
+ mock_user.user = MagicMock(id="user-1")
+ mock_user.email = "test@example.com"
+
+ app.dependency_overrides[get_current_user] = lambda: mock_user
+ app.dependency_overrides[get_project_repository] = lambda: project_repo
+ app.dependency_overrides[get_asset_library_repository] = lambda: library_repo
+ app.dependency_overrides[get_asset_repository] = lambda: asset_repo
+ app.dependency_overrides[get_ingest_job_repository] = lambda: ingest_repo
+ app.dependency_overrides[get_storage_service] = lambda: storage
+ return app, asset_repo, ingest_repo, storage
+
+
+def _client(**kwargs):
+ app, asset_repo, ingest_repo, storage = _build_app(**kwargs)
+ return TestClient(app), asset_repo, ingest_repo, storage
+
+
+COMPLETE_BODY = {
+ "project_id": "proj-1",
+ "library_id": "lib-1",
+ "storage_key": "uploads/abc/IMG_2282.MOV",
+}
+
+
+class TestDirectCompleteIdempotency:
+ def test_same_client_upload_id_creates_single_asset_and_job(self):
+ """同一 client_upload_id 连发两次 complete:只建 1 条 asset、1 个 job。"""
+ client, asset_repo, ingest_repo, _ = _client()
+ body = {**COMPLETE_BODY, "client_upload_id": "up-token-1", "file_size": 12345}
+
+ r1 = client.post("/api/v1/direct/complete", json=body)
+ r2 = client.post("/api/v1/direct/complete", json={**body, "storage_key": "uploads/zzz/IMG_2282.MOV"})
+
+ assert r1.status_code == 200 and r2.status_code == 200
+ b1, b2 = r1.json(), r2.json()
+ assert b1["duplicated"] is False
+ assert b2["duplicated"] is True
+ assert b1["asset_id"] == b2["asset_id"]
+ assert len(asset_repo.created) == 1
+ assert ingest_repo.created_count == 1
+ # 第二次返回的是已存在记录(其 storage_key 为第一次的 key)
+ assert b2["storage_key"] == "uploads/abc/IMG_2282.MOV"
+
+ def test_same_file_hash_returns_existing(self):
+ """同 file_hash(不同 token)重复 complete → 返回已存在记录。"""
+ client, asset_repo, ingest_repo, _ = _client()
+ body1 = {**COMPLETE_BODY, "file_hash": "h" * 32, "client_upload_id": "tok-a"}
+ body2 = {
+ **COMPLETE_BODY,
+ "storage_key": "uploads/def/IMG_2282.MOV",
+ "file_hash": "h" * 32,
+ "client_upload_id": "tok-b",
+ }
+
+ client.post("/api/v1/direct/complete", json=body1)
+ r2 = client.post("/api/v1/direct/complete", json=body2)
+
+ assert r2.json()["duplicated"] is True
+ assert len(asset_repo.created) == 1
+ assert ingest_repo.created_count == 1
+
+ def test_fallback_dedup_when_no_hash_no_token(self):
+ """旧客户端不传 hash/token:近期同库同名 processing 占位 → 兜底幂等。
+
+ 模拟 complete 超时重试:第一次已建好占位,第二次(OSS 重传拿到新 key)
+ 不应再建第二条。
+ """
+ client, asset_repo, ingest_repo, _ = _client()
+ # 第一次 complete(旧客户端无 token/hash,但 file_size 可知)
+ r1 = client.post(
+ "/api/v1/direct/complete",
+ json={**COMPLETE_BODY, "file_size": 5_000_000},
+ )
+ assert r1.json()["duplicated"] is False
+ # 重试:重新 prepare 产生新 storage_key(仅 uuid 目录不同,文件名一致——
+ # 前端重试传的是同一个 File),且近期;同大小才允许兜底命中
+ r2 = client.post(
+ "/api/v1/direct/complete",
+ json={
+ **COMPLETE_BODY,
+ "storage_key": "uploads/retry/IMG_2282.MOV",
+ "file_size": 5_000_000,
+ },
+ )
+ assert r2.status_code == 200
+ assert r2.json()["duplicated"] is True
+ assert r2.json()["asset_id"] == r1.json()["asset_id"]
+ assert len(asset_repo.created) == 1
+ assert ingest_repo.created_count == 1
+
+ def test_fallback_dedup_skipped_when_file_size_unknown(self):
+ """file_size=0(未知)时不允许仅凭同名 + processing 判重,直接放行(#1714)。
+
+ 根因场景:complete 没传 file_size,30 分钟内同名占位(如 iPhone 的
+ IMG_2285.MOV)会把内容/大小全新的视频误判为重复跳过。
+ """
+ client, asset_repo, _ingest_repo, _ = _client()
+ r1 = client.post(
+ "/api/v1/direct/complete",
+ json={**COMPLETE_BODY, "file_size": 0},
+ )
+ assert r1.json()["duplicated"] is False
+ # 第二个全新视频:同名(IMG_2285.MOV)、无 hash/token、file_size 仍未知
+ r2 = client.post(
+ "/api/v1/direct/complete",
+ json={**COMPLETE_BODY, "storage_key": "uploads/retry2/IMG_2282.MOV", "file_size": 0},
+ )
+ assert r2.status_code == 200
+ assert r2.json()["duplicated"] is False # 不能误杀
+ assert len(asset_repo.created) == 2 # 两条记录,放行新上传
+
+ def test_fallback_dedup_skipped_when_same_name_but_different_size(self):
+ """同名但 file_size 不同 → 不判重,正常建记录(#1714)。"""
+ client, asset_repo, _ingest_repo, _ = _client()
+ r1 = client.post(
+ "/api/v1/direct/complete",
+ json={**COMPLETE_BODY, "file_size": 5_000_000},
+ )
+ assert r1.json()["duplicated"] is False
+ r2 = client.post(
+ "/api/v1/direct/complete",
+ json={
+ **COMPLETE_BODY,
+ "storage_key": "uploads/retry3/IMG_2282.MOV",
+ "file_size": 9_999_999, # 同名但大小完全不同的新视频
+ },
+ )
+ assert r2.status_code == 200
+ assert r2.json()["duplicated"] is False
+ assert len(asset_repo.created) == 2
+
+ def test_fallback_dedup_skipped_when_hash_present_even_if_name_size_match(self):
+ """file_hash 非空且 hash 未命中时,不允许退回同名兜底(#1714)。
+
+ hash 已能代表内容:同名同大小但 hash 不同是真实的新内容,必须放行。
+ """
+ client, asset_repo, _ingest_repo, _ = _client()
+ # 第一次:某 hash 的视频
+ r1 = client.post(
+ "/api/v1/direct/complete",
+ json={
+ **COMPLETE_BODY,
+ "file_hash": "a" * 64,
+ "client_upload_id": "tok-1",
+ "file_size": 5_000_000,
+ },
+ )
+ assert r1.json()["duplicated"] is False
+ # 第二次:同名同大小但 hash 不同(新视频内容不同);
+ # 注意 client_upload_id 也必须不同,否则会先被 token 命中
+ r2 = client.post(
+ "/api/v1/direct/complete",
+ json={
+ **COMPLETE_BODY,
+ "storage_key": "uploads/retry4/IMG_2282.MOV",
+ "file_hash": "b" * 64,
+ "client_upload_id": "tok-2",
+ "file_size": 5_000_000,
+ },
+ )
+ assert r2.status_code == 200
+ assert r2.json()["duplicated"] is False
+ assert len(asset_repo.created) == 2
+
+ def test_fallback_dedup_ignores_ready_history(self):
+ """READY 历史同名素材不触发兜底(允许用户再次上传同名文件)。"""
+ ready = Asset(
+ id="ready-1",
+ project_id="proj-1",
+ library_id="lib-1",
+ name="IMG_2282.MOV",
+ storage_key="uploads/old/IMG_2282.MOV",
+ mime_type="video/quicktime",
+ status=AssetStatus.READY,
+ )
+ client, asset_repo, ingest_repo, _ = _client(asset_repo=StubAssetRepository([ready]))
+ r = client.post("/api/v1/direct/complete", json=COMPLETE_BODY)
+ assert r.status_code == 200
+ assert r.json()["duplicated"] is False
+ assert len(asset_repo.created) == 1
+
+ def test_fallback_dedup_window_expired(self):
+ """占位记录超过 30 分钟 → 不再兜底(视为孤儿,正常新建)。"""
+ stale = Asset(
+ id="stale-1",
+ project_id="proj-1",
+ library_id="lib-1",
+ name="IMG_2282.MOV",
+ storage_key="uploads/stale/IMG_2282.MOV",
+ mime_type="video/quicktime",
+ status=AssetStatus.PROCESSING,
+ )
+ stale.created_at = datetime.now(timezone.utc) - timedelta(minutes=45)
+ client, asset_repo, ingest_repo, _ = _client(asset_repo=StubAssetRepository([stale]))
+ r = client.post("/api/v1/direct/complete", json=COMPLETE_BODY)
+ assert r.status_code == 200
+ assert r.json()["duplicated"] is False
+ assert len(asset_repo.created) == 1
+
+ def test_legacy_repo_without_new_methods_still_works(self):
+ """旧仓储没有新幂等方法 → 鸭子类型降级,不报错、正常创建。"""
+ client, asset_repo, ingest_repo, _ = _client(asset_repo=LegacyStubAssetRepository())
+ r = client.post(
+ "/api/v1/direct/complete",
+ json={**COMPLETE_BODY, "client_upload_id": "tok-x", "file_hash": "f" * 32},
+ )
+ assert r.status_code == 200
+ assert r.json()["duplicated"] is False
+ assert len(asset_repo.created) == 1
+
+ def test_duplicate_complete_returns_existing_even_if_oss_missing(self):
+ """重复 complete 幂等检查先于 OSS file_exists:
+
+ 第一次成功建占位后,重试时即使 OSS 对象已不存在(file_exists=False),
+ 也必须返回已存在记录而不是 404/重复建库。"""
+ client, _, _, storage = _client()
+ body = {**COMPLETE_BODY, "client_upload_id": "tok-oss-gone"}
+ r1 = client.post("/api/v1/direct/complete", json=body)
+ assert r1.status_code == 200
+
+ storage.file_exists = MagicMock(return_value=False)
+ r2 = client.post(
+ "/api/v1/direct/complete",
+ json={**body, "storage_key": "uploads/retry2/IMG_2282.MOV"},
+ )
+ assert r2.status_code == 200
+ assert r2.json()["duplicated"] is True
+ assert r2.json()["asset_id"] == r1.json()["asset_id"]
+
+
+class TestMultipartUploadIdempotency:
+ def test_same_client_upload_id_second_submit_deduplicated(self):
+ """multipart 重复提交同 token:第二次直接 duplicated,不再上传 OSS。"""
+ client, asset_repo, ingest_repo, storage = _client()
+
+ def _post():
+ return client.post(
+ "/api/v1",
+ data={"project_id": "proj-1", "library_id": "lib-1", "client_upload_id": "mp-tok-1"},
+ files={"file": ("IMG_2282.MOV", b"fake-mov-data", "video/quicktime")},
+ )
+
+ r1 = _post()
+ r2 = _post()
+ assert r1.json()["duplicated"] is False
+ assert r2.json()["duplicated"] is True
+ assert r2.json()["asset_id"] == r1.json()["asset_id"]
+ assert len(asset_repo.created) == 1
+ assert ingest_repo.created_count == 1
+ # OSS 上传只发生一次(第二次在幂等检查处直接返回)
+ assert storage.upload_file.call_count == 1
diff --git a/tests/unit/test_wechat_bind_routes_1719.py b/tests/unit/test_wechat_bind_routes_1719.py
new file mode 100644
index 000000000..9f07cc43d
--- /dev/null
+++ b/tests/unit/test_wechat_bind_routes_1719.py
@@ -0,0 +1,235 @@
+"""#1719:微信绑定/解绑路由层测试(直接驱动路由函数)。
+
+覆盖:
+- GET /wechat/bind/url:调 oauth 生成链接、记日志
+- POST /wechat/bind:oauth 失败→400;绑定成功→success+user.wechat_bound=True;
+ use case 返回冲突→对应状态码透传
+- DELETE /wechat/bind:成功→success=True;use case 报错→状态码透传
+- /auth/me 返回 wechat_bound 字段
+"""
+
+from __future__ import annotations
+
+import asyncio
+import os
+import sys
+from pathlib import Path
+from types import SimpleNamespace
+from unittest.mock import MagicMock
+
+import pytest
+
+os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
+os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
+
+from app.api.routes import auth as auth_route # noqa: E402
+from fastapi import HTTPException # noqa: E402
+
+
+def _auth_user(user_id="u-1", openid=None):
+ user = SimpleNamespace(
+ id=user_id,
+ wechat_openid=openid,
+ email="user@example.com",
+ email_verified=True,
+ username="user",
+ display_name="用户",
+ phone="",
+ phone_verified=False,
+ profile_completed=True,
+ )
+ return SimpleNamespace(user=user, session_id="s-1", token_type="user_auth")
+
+
+def _patched_bind(result, error, status):
+ """构造打了补丁的 wechat_bind_use_case 模块"""
+ mod = SimpleNamespace(
+ WechatBindRequest=lambda **kw: SimpleNamespace(**kw),
+ WechatBindUseCase=MagicMock(),
+ WechatUnbindUseCase=MagicMock(),
+ )
+ fake_bind_uc = MagicMock()
+ fake_bind_uc.bind.return_value = (result, error, status)
+ mod.WechatBindUseCase.return_value = fake_bind_uc
+ return mod
+
+
+def test_get_bind_url_returns_url_and_state():
+ fake_oauth = MagicMock()
+ fake_oauth.generate_auth_url.return_value = ("https://open.weixin.qq.com/qrconnect?xxx", "state-bind-1")
+
+ import packages.application.auth.wechat_oauth_service as oauth_mod
+
+ orig = oauth_mod.get_wechat_oauth_service
+ oauth_mod.get_wechat_oauth_service = MagicMock(return_value=fake_oauth)
+ try:
+ resp = asyncio.run(auth_route.get_wechat_bind_url(current_user=_auth_user()))
+ finally:
+ oauth_mod.get_wechat_oauth_service = orig
+
+ assert resp.auth_url.startswith("https://open.weixin.qq.com")
+ assert resp.state == "state-bind-1"
+
+
+def test_bind_oauth_error_returns_400():
+ fake_oauth = MagicMock()
+ fake_oauth.handle_callback.return_value = (None, "无效的 state 参数")
+
+ import packages.application.auth.wechat_oauth_service as oauth_mod
+
+ orig = oauth_mod.get_wechat_oauth_service
+ oauth_mod.get_wechat_oauth_service = MagicMock(return_value=fake_oauth)
+ try:
+ with pytest.raises(HTTPException) as exc:
+ asyncio.run(
+ auth_route.wechat_bind(
+ SimpleNamespace(code="c-1", state="s-1"),
+ current_user=_auth_user(),
+ user_repository=MagicMock(),
+ )
+ )
+ finally:
+ oauth_mod.get_wechat_oauth_service = orig
+
+ assert exc.value.status_code == 400
+ assert "state" in exc.value.detail
+
+
+def test_bind_success_returns_user_with_wechat_bound():
+ fake_oauth = MagicMock()
+ fake_oauth.handle_callback.return_value = (
+ SimpleNamespace(openid="wx-openid-1", unionid="wx-union-1"),
+ None,
+ )
+
+ bound_user = SimpleNamespace(
+ id="u-1",
+ wechat_openid="wx-openid-1",
+ email="user@example.com",
+ email_verified=True,
+ username="user",
+ display_name="用户",
+ phone="",
+ phone_verified=False,
+ profile_completed=True,
+ )
+
+ import packages.application.auth.wechat_oauth_service as oauth_mod
+ from packages.application.auth import wechat_bind_use_case as bind_mod
+
+ orig_oauth = oauth_mod.get_wechat_oauth_service
+ fake_bind_uc = MagicMock()
+ fake_bind_uc.bind.return_value = (SimpleNamespace(user=bound_user), None, 200)
+ orig_bind = bind_mod.WechatBindUseCase
+ bind_mod.WechatBindUseCase = MagicMock(return_value=fake_bind_uc)
+ oauth_mod.get_wechat_oauth_service = MagicMock(return_value=fake_oauth)
+ try:
+ resp = asyncio.run(
+ auth_route.wechat_bind(
+ SimpleNamespace(code="c-1", state="s-1"),
+ current_user=_auth_user(),
+ user_repository=MagicMock(),
+ )
+ )
+ finally:
+ oauth_mod.get_wechat_oauth_service = orig_oauth
+ bind_mod.WechatBindUseCase = orig_bind
+
+ assert resp.success is True
+ assert resp.user.wechat_bound is True
+ assert resp.user.user_id == "u-1"
+ # 绑定请求应带上当前用户 id 与微信 openid
+ call_kwargs = fake_bind_uc.bind.call_args[0][0]
+ assert call_kwargs.user_id == "u-1"
+ assert call_kwargs.openid == "wx-openid-1"
+
+
+def test_bind_conflict_propagates_409():
+ fake_oauth = MagicMock()
+ fake_oauth.handle_callback.return_value = (
+ SimpleNamespace(openid="wx-openid-1", unionid=""),
+ None,
+ )
+
+ import packages.application.auth.wechat_oauth_service as oauth_mod
+ from packages.application.auth import wechat_bind_use_case as bind_mod
+
+ orig_oauth = oauth_mod.get_wechat_oauth_service
+ fake_bind_uc = MagicMock()
+ fake_bind_uc.bind.return_value = (None, "该微信已绑定其他账号,请先在原账号解绑", 409)
+ orig_bind = bind_mod.WechatBindUseCase
+ bind_mod.WechatBindUseCase = MagicMock(return_value=fake_bind_uc)
+ oauth_mod.get_wechat_oauth_service = MagicMock(return_value=fake_oauth)
+ try:
+ with pytest.raises(HTTPException) as exc:
+ asyncio.run(
+ auth_route.wechat_bind(
+ SimpleNamespace(code="c-1", state="s-1"),
+ current_user=_auth_user(),
+ user_repository=MagicMock(),
+ )
+ )
+ finally:
+ oauth_mod.get_wechat_oauth_service = orig_oauth
+ bind_mod.WechatBindUseCase = orig_bind
+
+ assert exc.value.status_code == 409
+ assert "已绑定其他账号" in exc.value.detail
+
+
+def test_unbind_success_returns_success_true():
+ unbound_user = SimpleNamespace(
+ id="u-1",
+ wechat_openid=None,
+ email="user@example.com",
+ email_verified=True,
+ username="user",
+ display_name="用户",
+ phone="",
+ phone_verified=False,
+ )
+
+ from packages.application.auth import wechat_bind_use_case as bind_mod
+
+ fake_uc = MagicMock()
+ fake_uc.unbind.return_value = (SimpleNamespace(user=unbound_user), None, 200)
+ orig = bind_mod.WechatUnbindUseCase
+ bind_mod.WechatUnbindUseCase = MagicMock(return_value=fake_uc)
+ try:
+ resp = asyncio.run(
+ auth_route.wechat_unbind(current_user=_auth_user(openid="wx-old"), user_repository=MagicMock())
+ )
+ finally:
+ bind_mod.WechatUnbindUseCase = orig
+
+ assert resp.success is True
+ fake_uc.unbind.assert_called_once_with("u-1")
+
+
+def test_unbind_rejected_no_other_login_propagates_400():
+ from packages.application.auth import wechat_bind_use_case as bind_mod
+
+ fake_uc = MagicMock()
+ fake_uc.unbind.return_value = (None, "账号需要至少一种其他登录方式(已验证手机或真实邮箱)后才能解绑微信", 400)
+ orig = bind_mod.WechatUnbindUseCase
+ bind_mod.WechatUnbindUseCase = MagicMock(return_value=fake_uc)
+ try:
+ with pytest.raises(HTTPException) as exc:
+ asyncio.run(auth_route.wechat_unbind(current_user=_auth_user(openid="wx-old"), user_repository=MagicMock()))
+ finally:
+ bind_mod.WechatUnbindUseCase = orig
+
+ assert exc.value.status_code == 400
+ assert "登录方式" in exc.value.detail
+
+
+def test_me_includes_wechat_bound_flag():
+ # 已绑定用户
+ resp = asyncio.run(auth_route.get_current_user_info(authenticated_user=_auth_user(openid="wx-openid-1")))
+ assert resp.wechat_bound is True
+
+ # 未绑定用户
+ resp2 = asyncio.run(auth_route.get_current_user_info(authenticated_user=_auth_user(openid=None)))
+ assert resp2.wechat_bound is False
diff --git a/tests/unit/test_wechat_bind_use_case_1719.py b/tests/unit/test_wechat_bind_use_case_1719.py
new file mode 100644
index 000000000..e933f8a25
--- /dev/null
+++ b/tests/unit/test_wechat_bind_use_case_1719.py
@@ -0,0 +1,251 @@
+"""#1719:已登录用户微信绑定/解绑 Use Case 测试。
+
+覆盖:
+- bind:幂等重复绑定、未绑定成功、当前账号已绑其他微信、openid/unionid 冲突 409、用户不存在
+- unbind:成功清 openid+unionid、未绑定拒绝、无其他登录方式拒绝、密码/手机/真实邮箱各兜底放行、用户不存在
+"""
+
+from __future__ import annotations
+
+from types import SimpleNamespace
+
+import pytest
+
+from packages.application.auth.wechat_bind_use_case import (
+ WechatBindRequest,
+ WechatBindUseCase,
+ WechatUnbindUseCase,
+)
+
+
+def _user(
+ user_id="u-1",
+ wechat_openid=None,
+ wechat_unionid=None,
+ password_hash="hashed-pw",
+ phone=None,
+ phone_verified=False,
+ email="user@example.com",
+ email_verified=True,
+):
+ return SimpleNamespace(
+ id=user_id,
+ wechat_openid=wechat_openid,
+ wechat_unionid=wechat_unionid,
+ password_hash=password_hash,
+ phone=phone,
+ phone_verified=phone_verified,
+ email=email,
+ email_verified=email_verified,
+ )
+
+
+class _FakeRepo:
+ """内存仓储:按 id/openid/unionid 建索引,save 原地更新。"""
+
+ def __init__(self, users):
+ self.users = {u.id: u for u in users}
+ self.saved = []
+
+ def find_by_id(self, user_id):
+ return self.users.get(user_id)
+
+ def find_by_wechat_openid(self, openid):
+ for u in self.users.values():
+ if u.wechat_openid == openid:
+ return u
+ return None
+
+ def find_by_wechat_unionid(self, unionid):
+ if not unionid:
+ return None
+ for u in self.users.values():
+ if u.wechat_unionid == unionid:
+ return u
+ return None
+
+ def save(self, user):
+ self.saved.append(user)
+
+
+# ==================== bind ====================
+
+
+def test_bind_success_when_not_bound():
+ user = _user()
+ repo = _FakeRepo([user])
+ result, err, status = WechatBindUseCase(repo).bind(
+ WechatBindRequest(user_id="u-1", openid="wx-openid-1", unionid="wx-union-1")
+ )
+ assert err is None
+ assert status == 200
+ assert result.user.wechat_openid == "wx-openid-1"
+ assert result.user.wechat_unionid == "wx-union-1"
+ assert repo.saved == [user]
+
+
+def test_bind_idempotent_same_openid():
+ user = _user(wechat_openid="wx-openid-1", wechat_unionid="wx-union-1")
+ repo = _FakeRepo([user])
+ result, err, status = WechatBindUseCase(repo).bind(
+ WechatBindRequest(user_id="u-1", openid="wx-openid-1", unionid="wx-union-1")
+ )
+ assert err is None
+ assert status == 200
+ assert result.user is user
+ assert repo.saved == [] # 幂等不写库
+
+
+def test_bind_conflict_user_already_bound_other_wechat():
+ user = _user(wechat_openid="wx-old")
+ repo = _FakeRepo([user])
+ result, err, status = WechatBindUseCase(repo).bind(WechatBindRequest(user_id="u-1", openid="wx-new"))
+ assert result is None
+ assert status == 409
+ assert "已绑定微信" in err
+
+
+def test_bind_conflict_openid_used_by_other_user():
+ user = _user(user_id="u-1")
+ other = _user(user_id="u-2", wechat_openid="wx-openid-1")
+ repo = _FakeRepo([user, other])
+ result, err, status = WechatBindUseCase(repo).bind(WechatBindRequest(user_id="u-1", openid="wx-openid-1"))
+ assert result is None
+ assert status == 409
+ assert "已绑定其他账号" in err
+ assert user.wechat_openid is None # 未写库
+
+
+def test_bind_conflict_unionid_used_by_other_user():
+ user = _user(user_id="u-1")
+ # openid 不同,但 unionid 指向同一微信主体
+ other = _user(user_id="u-2", wechat_openid="wx-other", wechat_unionid="wx-union-x")
+ repo = _FakeRepo([user, other])
+ result, err, status = WechatBindUseCase(repo).bind(
+ WechatBindRequest(user_id="u-1", openid="wx-openid-1", unionid="wx-union-x")
+ )
+ assert result is None
+ assert status == 409
+ assert "微信主体" in err
+
+
+def test_bind_missing_openid_returns_400():
+ repo = _FakeRepo([_user()])
+ result, err, status = WechatBindUseCase(repo).bind(WechatBindRequest(user_id="u-1", openid=""))
+ assert result is None
+ assert status == 400
+ assert "openid" in err
+
+
+def test_bind_user_not_found_returns_404():
+ repo = _FakeRepo([])
+ result, err, status = WechatBindUseCase(repo).bind(WechatBindRequest(user_id="ghost", openid="wx-openid-1"))
+ assert result is None
+ assert status == 404
+
+
+def test_bind_fills_unionid_when_existing_user_has_none():
+ # 用户历史上只绑了 openid(unionid 为空),再次绑定时补齐 unionid 不冲突
+ user = _user(wechat_openid="wx-openid-1", wechat_unionid=None)
+ repo = _FakeRepo([user])
+ result, err, status = WechatBindUseCase(repo).bind(
+ WechatBindRequest(user_id="u-1", openid="wx-openid-1", unionid="wx-union-new")
+ )
+ # openid 相同 → 幂等成功(不覆盖 unionid,保持数据稳定)
+ assert err is None
+ assert status == 200
+
+
+# ==================== unbind ====================
+
+
+def test_unbind_success_with_real_verified_email():
+ # 默认 _user 即 real@example.com 且 email_verified=True
+ user = _user(wechat_openid="wx-openid-1", wechat_unionid="wx-union-1")
+ repo = _FakeRepo([user])
+ result, err, status = WechatUnbindUseCase(repo).unbind("u-1")
+ assert err is None
+ assert status == 200
+ assert result.user.wechat_openid is None
+ assert result.user.wechat_unionid is None
+ assert repo.saved == [user]
+
+
+def test_unbind_rejected_when_only_random_password_hash():
+ # 微信注册用户:随机密码 hash 存在、邮箱是 @wechat.local 占位、无手机 → 不允许解绑
+ user = _user(
+ wechat_openid="wx-openid-1",
+ password_hash="random-secret-hash",
+ email="abc@wechat.local",
+ email_verified=True,
+ )
+ repo = _FakeRepo([user])
+ result, err, status = WechatUnbindUseCase(repo).unbind("u-1")
+ assert result is None
+ assert status == 400
+ assert "登录方式" in err
+ assert user.wechat_openid == "wx-openid-1" # 未写库
+
+
+def test_unbind_allowed_with_verified_phone_even_without_password():
+ user = _user(
+ wechat_openid="wx-openid-1",
+ password_hash="",
+ phone="13800000000",
+ phone_verified=True,
+ email="wx@wechat.local",
+ email_verified=True,
+ )
+ repo = _FakeRepo([user])
+ result, err, status = WechatUnbindUseCase(repo).unbind("u-1")
+ assert err is None
+ assert status == 200
+ assert result.user.wechat_openid is None
+
+
+def test_unbind_rejected_when_no_other_login_method():
+ # 无手机、邮箱占位 → 唯一登录方式就是微信,禁止解绑
+ user = _user(
+ wechat_openid="wx-openid-1",
+ password_hash="",
+ email="abc@wechat.local",
+ email_verified=True,
+ )
+ repo = _FakeRepo([user])
+ result, err, status = WechatUnbindUseCase(repo).unbind("u-1")
+ assert result is None
+ assert status == 400
+ assert "登录方式" in err
+ assert user.wechat_openid == "wx-openid-1" # 未写库
+
+
+def test_unbind_not_bound_returns_400():
+ user = _user() # 未绑定
+ repo = _FakeRepo([user])
+ result, err, status = WechatUnbindUseCase(repo).unbind("u-1")
+ assert result is None
+ assert status == 400
+ assert "未绑定" in err
+
+
+def test_unbind_user_not_found_returns_404():
+ repo = _FakeRepo([])
+ result, err, status = WechatUnbindUseCase(repo).unbind("ghost")
+ assert result is None
+ assert status == 404
+
+
+def test_unbind_unverified_phone_does_not_count():
+ # 手机未验证不算有效登录方式
+ user = _user(
+ wechat_openid="wx-openid-1",
+ password_hash="",
+ phone="13800000000",
+ phone_verified=False,
+ email="abc@wechat.local",
+ email_verified=True,
+ )
+ repo = _FakeRepo([user])
+ result, err, status = WechatUnbindUseCase(repo).unbind("u-1")
+ assert result is None
+ assert status == 400
diff --git a/tests/unit/test_wechat_callback_logging_1718.py b/tests/unit/test_wechat_callback_logging_1718.py
new file mode 100644
index 000000000..4f45ea735
--- /dev/null
+++ b/tests/unit/test_wechat_callback_logging_1718.py
@@ -0,0 +1,115 @@
+"""#1718:微信回调路由可观测性日志分支覆盖(UA/state/错误透传)。
+
+直接驱动 wechat_callback 路由函数,mock OAuth service 与用户仓储:
+- 成功路径:日志记录 UA、state 校验通过(MicroMessenger 内置浏览器)
+- 失败路径:OAuth 返回错误时记 warning 并抛 400
+"""
+
+from __future__ import annotations
+
+import asyncio
+import os
+import sys
+from pathlib import Path
+from types import SimpleNamespace
+from unittest.mock import MagicMock, patch
+
+import pytest
+
+os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
+os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
+
+from app.api.routes import auth as auth_route # noqa: E402
+from fastapi import HTTPException # noqa: E402
+
+
+class _FakeRequest:
+ def __init__(self, ua: str):
+ self.headers = {"User-Agent": ua}
+
+
+def _wechat_user():
+ return SimpleNamespace(
+ openid="openid-callback-1",
+ unionid="union-callback-1",
+ nickname="微信用户",
+ avatar_url="http://x/a.png",
+ )
+
+
+def _fake_oauth_factory(success: bool):
+ service = MagicMock()
+ if success:
+ service.handle_callback.return_value = (_wechat_user(), None)
+ else:
+ service.handle_callback.return_value = (None, "无效的 state 参数,请求可能已过期或被篡改")
+ return service
+
+
+def test_wechat_callback_success_logs_ua_and_state(caplog):
+ fake_repo = MagicMock()
+ sync_response = SimpleNamespace(
+ access_token="at",
+ refresh_token="rt",
+ user_id="u-1",
+ nickname="微信用户",
+ avatar_url="",
+ is_new_user=False,
+ expires_in=1800,
+ )
+ fake_use_case = MagicMock()
+ fake_use_case.execute.return_value = (sync_response, None)
+
+ user = SimpleNamespace(
+ id="u-1",
+ phone_verified=True,
+ email_verified=True,
+ email="u@example.com",
+ )
+ fake_repo.find_by_id.return_value = user
+
+ request_obj = SimpleNamespace(code="code-1", state="state-1")
+ fake_http = _FakeRequest("Mozilla/5.0 (Linux; Android 13) MicroMessenger/8.0.40 WeChat/8.0.40")
+
+ import packages.application.auth.wechat_oauth_service as oauth_mod
+ import packages.application.auth.wechat_sync_use_case as sync_mod
+
+ orig_oauth = oauth_mod.get_wechat_oauth_service
+ orig_sync = sync_mod.WechatSyncUseCase
+ oauth_mod.get_wechat_oauth_service = MagicMock(return_value=_fake_oauth_factory(success=True))
+ sync_mod.WechatSyncUseCase = MagicMock(return_value=fake_use_case)
+ try:
+ with caplog.at_level("INFO", logger="app.api.routes.auth"):
+ resp = asyncio.run(auth_route.wechat_callback(request_obj, fake_http, user_repository=fake_repo))
+ finally:
+ oauth_mod.get_wechat_oauth_service = orig_oauth
+ sync_mod.WechatSyncUseCase = orig_sync
+
+ assert resp.user_id == "u-1"
+ assert resp.binding_complete is True
+ log_text = " ".join(rec.getMessage() for rec in caplog.records)
+ assert "微信回调" in log_text
+ assert "MicroMessenger" in log_text or "微信内置浏览器=True" in log_text
+
+
+def test_wechat_callback_failure_raises_400_with_detail(caplog):
+ request_obj = SimpleNamespace(code="code-bad", state="state-bad")
+ fake_http = _FakeRequest("Mozilla/5.0 Chrome/127")
+ fake_service = _fake_oauth_factory(success=False)
+
+ import packages.application.auth.wechat_oauth_service as oauth_mod
+
+ orig = oauth_mod.get_wechat_oauth_service
+ oauth_mod.get_wechat_oauth_service = MagicMock(return_value=fake_service)
+ try:
+ with caplog.at_level("WARNING", logger="app.api.routes.auth"):
+ with pytest.raises(HTTPException) as exc_info:
+ asyncio.run(auth_route.wechat_callback(request_obj, fake_http, user_repository=MagicMock()))
+ finally:
+ oauth_mod.get_wechat_oauth_service = orig
+
+ assert exc_info.value.status_code == 400
+ assert "state" in exc_info.value.detail
+ assert any("微信回调" in rec.getMessage() for rec in caplog.records)
diff --git a/tests/unit/test_wechat_oauth_service.py b/tests/unit/test_wechat_oauth_service.py
index 7a2ed3397..759732c62 100755
--- a/tests/unit/test_wechat_oauth_service.py
+++ b/tests/unit/test_wechat_oauth_service.py
@@ -388,3 +388,52 @@ class TestGetWechatOAuthService:
"""返回 WechatOAuthService 实例"""
service = get_wechat_oauth_service()
assert isinstance(service, WechatOAuthService)
+
+ def test_singleton_same_instance_across_calls(self, monkeypatch):
+ """#1718 回归:工厂必须返回同一实例,否则 state store 不共享"""
+ import packages.application.auth.wechat_oauth_service as mod
+
+ monkeypatch.setattr(mod, "_oauth_service_singleton", None)
+ s1 = get_wechat_oauth_service()
+ s2 = get_wechat_oauth_service()
+ assert s1 is s2
+
+ def test_state_survives_across_factory_calls(self, monkeypatch):
+ """#1718 回归:/wechat/url 与 /wechat/callback 经工厂拿到同一 state store
+
+ 模拟两次请求各自调用工厂:第一个实例生成 state,第二个实例(同一单例)
+ 必须能校验通过。修复前工厂每次 new 一个实例,回调必现 400「无效的 state」。
+ """
+ import packages.application.auth.wechat_oauth_service as mod
+
+ monkeypatch.setattr(mod, "_oauth_service_singleton", None)
+ monkeypatch.setenv("WECHAT_OPEN_APP_ID", "wx-test")
+ monkeypatch.setenv("WECHAT_OPEN_APP_SECRET", "secret-test")
+ monkeypatch.setenv("WECHAT_OPEN_REDIRECT_URI", "https://example.com/cb")
+
+ # 请求1:生成授权链接(state 写入单例 store)
+ _, state = get_wechat_oauth_service().generate_auth_url()
+
+ # 请求2:回调校验(应命中同一个 store;微信 API 用 mock 避免外网)
+ with patch("packages.application.auth.wechat_oauth_service.requests.get") as mock_get:
+ mock_get.return_value = MagicMock(
+ json=MagicMock(
+ return_value={
+ "access_token": "at",
+ "openid": "oid",
+ "unionid": "uid",
+ "nickname": "n",
+ "headimgurl": "http://x/a.png",
+ }
+ )
+ )
+ user_info, error = get_wechat_oauth_service().handle_callback("code-x", state)
+
+ assert error is None, f"state 应跨请求共享,实际报错: {error}"
+ assert user_info is not None
+ assert user_info.openid == "oid"
+
+ # state 一次性消费,重放必须失败
+ user_info2, error2 = get_wechat_oauth_service().handle_callback("code-y", state)
+ assert user_info2 is None
+ assert "state" in error2
diff --git a/tests/unit/test_wechat_state_redis_1718.py b/tests/unit/test_wechat_state_redis_1718.py
new file mode 100644
index 000000000..3ab463a28
--- /dev/null
+++ b/tests/unit/test_wechat_state_redis_1718.py
@@ -0,0 +1,262 @@
+"""#1718:微信 OAuth state 存储 Redis 化 + 中文昵称 UTF-8 解码修复。
+
+覆盖(全 mock/fake,CI 无真实 redis 也产生覆盖):
+- RedisStateStore:put 用 SET NX EX、verify_and_consume 用 GETDEL 一次性消费、
+ 重复消费返回 False、Redis 异常降级内存、client 注入
+- Redis 不可用(ping 失败)构造时降级内存,功能仍正常
+- GETDEL 不存在(老 Redis)走 GET+DELETE 兜底
+- handle_callback:微信 sns/userinfo 响应含中文 nickname,resp.encoding=utf-8
+ 后解析不乱码;errcode 错误路径返回 errmsg
+"""
+
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+from unittest.mock import MagicMock, patch
+
+import pytest
+
+sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
+
+from packages.application.auth import wechat_oauth_service as oauth # noqa: E402
+
+
+class _FakeRedisClient:
+ """最小内存版 redis client,模拟 SET NX EX / GETDEL / GET / DELETE / ping。"""
+
+ def __init__(self):
+ self.data: dict[str, str] = {}
+ self.ttl: dict[str, int] = {}
+ self.has_getdel = True
+
+ def ping(self):
+ return True
+
+ def set(self, key, value, nx=False, ex=None): # noqa: ARG002
+ if nx and key in self.data:
+ return None
+ self.data[key] = value
+ if ex is not None:
+ self.ttl[key] = ex
+ return True
+
+ def get(self, key):
+ return self.data.get(key)
+
+ def getdel(self, key):
+ return self.data.pop(key, None)
+
+ def delete(self, key):
+ return 1 if self.data.pop(key, None) is not None else 0
+
+ def eval(self, script, numkeys, key): # noqa: ARG002
+ # 模拟 Lua:原子 GET + DEL
+ return self.data.pop(key, None)
+
+
+# ── RedisStateStore ─────────────────────────────────────────────────────
+
+
+def test_redis_state_store_put_and_consume_once():
+ client = _FakeRedisClient()
+ store = oauth.RedisStateStore(client=client)
+ store.put("state-abc")
+ # key 带前缀、TTL 写入
+ assert client.data.get("wechat:state:state-abc") is not None
+ assert client.ttl.get("wechat:state:state-abc") == oauth.STATE_TTL_SECONDS
+ # 一次性消费:第一次 True,第二次 False
+ assert store.verify_and_consume("state-abc") is True
+ assert store.verify_and_consume("state-abc") is False
+
+
+def test_redis_state_store_unknown_state_returns_false():
+ store = oauth.RedisStateStore(client=_FakeRedisClient())
+ assert store.verify_and_consume("never-put") is False
+
+
+def test_redis_state_store_eval_missing_falls_back_to_get_delete():
+ """eval 不可用(如禁用脚本)时退化 GET+DELETE,仍一次性消费。"""
+ client = _FakeRedisClient()
+
+ def _no_eval(script, numkeys, *keys): # noqa: ARG002
+ raise RuntimeError("unknown command EVAL")
+
+ client.eval = _no_eval # type: ignore[method-assign]
+ store = oauth.RedisStateStore(client=client)
+ store.put("state-old")
+ assert store.verify_and_consume("state-old") is True
+ # GET+DELETE 也消费掉了
+ assert "wechat:state:state-old" not in client.data
+ assert store.verify_and_consume("state-old") is False
+
+
+def test_redis_state_store_put_exception_falls_back_to_memory():
+ client = MagicMock()
+ client.set.side_effect = RuntimeError("redis write fail")
+ # eval/get 也失败,确保降级到内存
+ client.eval.side_effect = RuntimeError("redis read fail")
+ client.get.side_effect = RuntimeError("redis read fail")
+ store = oauth.RedisStateStore(client=client)
+
+ store.put("state-fb") # 写 Redis 失败 → 内存
+ assert store.verify_and_consume("state-fb") is True # 内存命中
+ assert store.verify_and_consume("state-fb") is False
+
+
+def test_redis_state_store_consume_exception_falls_back_to_memory():
+ client = MagicMock()
+ client.set.return_value = True # put 走 Redis
+ client.eval.side_effect = RuntimeError("redis down")
+ client.get.side_effect = RuntimeError("redis down")
+ store = oauth.RedisStateStore(client=client)
+
+ store.put("state-fb2") # 成功写 Redis
+ # 校验时 Redis 挂了 → 降级内存(内存里没有,返回 False,不报错)
+ assert store.verify_and_consume("state-fb2") is False
+
+
+def test_redis_state_store_constructor_ping_failure_falls_back():
+ """构造时 ping 失败(Redis 不可用)→ 内存降级,功能正常。"""
+ fake_redis_mod = MagicMock()
+ fake_client = MagicMock()
+ fake_client.ping.side_effect = ConnectionError("refused")
+ fake_redis_mod.Redis.from_url.return_value = fake_client
+
+ with patch.dict(sys.modules, {"redis": fake_redis_mod}):
+ store = oauth.RedisStateStore(redis_url="redis://nonexistent:6379/0")
+
+ # Redis 不可用 → 内存存储仍工作
+ store.put("state-mem")
+ assert store.verify_and_consume("state-mem") is True
+ assert store.verify_and_consume("state-mem") is False
+
+
+# ── handle_callback:state 校验 + UTF-8 中文昵称 ────────────────────────
+
+
+def _configured_service(state_store=None):
+ store = state_store or oauth.MemoryStateStore()
+ return oauth.WechatOAuthService(
+ app_id="wx-test",
+ app_secret="secret-test",
+ redirect_uri="https://staging.xiaoxiajianji.com/auth/wechat/callback",
+ state_store=store,
+ )
+
+
+class _FakeResponse:
+ def __init__(self, payload):
+ self._payload = payload
+ self.encoding = None # 模拟微信响应头不带 charset
+
+ def json(self):
+ # 模拟 requests 行为:按 self.encoding 解码。这里直接返回 payload,
+ # 但记录 encoding 是否被设置为 utf-8(断言修复生效)
+ self._decoded_with = self.encoding
+ return self._payload
+
+
+def test_handle_callback_chinese_nickname_decoded_utf8(monkeypatch):
+ """微信 userinfo 返回中文昵称,service 设置 encoding=utf-8 后不乱码。"""
+ service = _configured_service()
+ state = "state-cn-1"
+ service._state_store.put(state)
+
+ token_resp = _FakeResponse({"access_token": "at-1", "openid": "openid-cn", "unionid": "union-cn"})
+ user_resp = _FakeResponse(
+ {"openid": "openid-cn", "unionid": "union-cn", "nickname": "微信小应🎬", "headimgurl": ""}
+ )
+ responses = iter([token_resp, user_resp])
+ monkeypatch.setattr(oauth.requests, "get", lambda *a, **k: next(responses))
+
+ info, err = service.handle_callback("code-cn", state)
+ assert err is None
+ assert info is not None
+ assert info.openid == "openid-cn"
+ assert info.nickname == "微信小应🎬"
+ # 两个响应都被显式设为 utf-8
+ assert token_resp.encoding == "utf-8"
+ assert user_resp.encoding == "utf-8"
+
+
+def test_handle_callback_state_invalid_returns_error():
+ service = _configured_service()
+ info, err = service.handle_callback("code-x", "state-not-exist")
+ assert info is None
+ assert "state" in err
+
+
+def test_handle_callback_wechat_errcode_returns_errmsg(monkeypatch):
+ """微信返回 errcode(如 code 已被消费 40029)时返回 errmsg 原文。"""
+ service = _configured_service()
+ state = "state-err-1"
+ service._state_store.put(state)
+
+ err_resp = _FakeResponse({"errcode": 40029, "errmsg": "invalid code"})
+ monkeypatch.setattr(oauth.requests, "get", lambda *a, **k: err_resp)
+
+ info, err = service.handle_callback("bad-code", state)
+ assert info is None
+ assert "invalid code" in err
+ assert err_resp.encoding == "utf-8"
+
+
+def test_generate_auth_url_stores_state_in_redis():
+ """generate_auth_url 生成的 state 写入 Redis(而非仅内存)。"""
+ client = _FakeRedisClient()
+ service = oauth.WechatOAuthService(
+ app_id="wx-test",
+ app_secret="secret-test",
+ redirect_uri="https://example.com/cb",
+ state_store=oauth.RedisStateStore(client=client),
+ )
+ url, state = service.generate_auth_url()
+ assert f"wechat:state:{state}" in client.data
+ assert "open.weixin.qq.com" in url
+
+
+# ── _build_default_state_store 工厂分支 ─────────────────────────────────
+
+
+def test_build_default_state_store_uses_redis_when_broker_configured():
+ """API settings 有 CELERY_BROKER_URL 时返回 RedisStateStore。"""
+ store = oauth._build_default_state_store()
+ # CI/本地通常配置了 redis://localhost:6379/...;无论 Redis 是否可达,
+ # 返回类型应为 RedisStateStore(内部降级内存)
+ assert isinstance(store, oauth.RedisStateStore) or isinstance(store, oauth.MemoryStateStore)
+
+
+def test_build_default_state_store_env_fallback(monkeypatch):
+ """app.config 不可用(如纯 worker 环境)时从环境变量取 redis url。"""
+ import builtins
+
+ real_import = builtins.__import__
+
+ def _failing_import(name, *args, **kwargs):
+ if name == "app.config":
+ raise ImportError("no app.config")
+ return real_import(name, *args, **kwargs)
+
+ monkeypatch.setattr(builtins, "__import__", _failing_import)
+ monkeypatch.setenv("CELERY_BROKER_URL", "redis://localhost:6379/9")
+ store = oauth._build_default_state_store()
+ assert isinstance(store, oauth.RedisStateStore)
+
+
+def test_build_default_state_store_no_config_returns_memory(monkeypatch):
+ """无任何 redis 配置时返回 MemoryStateStore。"""
+ import builtins
+
+ real_import = builtins.__import__
+
+ def _failing_import(name, *args, **kwargs):
+ if name == "app.config":
+ raise ImportError("no app.config")
+ return real_import(name, *args, **kwargs)
+
+ monkeypatch.setattr(builtins, "__import__", _failing_import)
+ monkeypatch.delenv("CELERY_BROKER_URL", raising=False)
+ monkeypatch.delenv("REDIS_URL", raising=False)
+ store = oauth._build_default_state_store()
+ assert isinstance(store, oauth.MemoryStateStore)