Compare commits

..

2 Commits

Author SHA1 Message Date
xiaoxia 20fe447efb fix(#1677): AI Review 反馈修复——标题校验/TTS依赖收窄/排序健壮性/预览数硬上限
- 单视频标题校验与 buildPayload.validateGenerateInputs 对齐:aiAutoSelect
  自动模式允许空标题(后端生成,既有设计),手动模式必填,逻辑显式化
- TTS 试听 effect 依赖收窄到 previewTitles[0](variant0Title),编辑其他
  变体标题不再触发多余 TTS 请求
- BatchGenerationGrid 排序加 (variantIndex || 0) 兜底防 NaN
- CanvasPreviewGrid 渲染数加 MAX_PREVIEW_COUNT(10) 硬上限,防媒体元素过多卡顿
2026-09-05 12:36:28 +08:00
xiaoxia ac28530528 feat(#1677): 批量预览改纯前端Canvas实时预览+固定6步流程+第5步独立渲染进度
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 1s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m29s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 1m38s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Successful in 1m41s
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 1m40s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 1m47s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 1m48s
CI/CD Pipeline / Validate - Style (pull_request) Successful in 2m14s
AI Code Review / AI Code Review (pull_request) Failing after 3m29s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 5m10s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Successful in 2s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 4m0s
ACR Cleanup / ACR Image Cleanup (pull_request_target) Successful in 9s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 22s
- 预览纯前端化:N 个 FrontendPreviewPlayer 网格(CanvasPreviewGrid),variantSeed
  让素材排布/起始点不同、画面有差异;不调任何后端渲染接口,秒开不占 worker;
  seed=0 走旧逻辑保证 N=1 零回归
- 固定 6 步(单视频与批量一致):模板(弹数量) → 素材 → 配音 → 标题 → 确认生成 → 封面
- 第4步右侧批量标题:N 个独立输入框(AutoComplete 接标题库)+ AI 一键生成 N 个标题
  (本地模板池同源)+ 主题词补填空标题;标题样式全局共用
- 第4步底部「 确认生成 N 个视频」:按 PR#1701 契约提交 count +
  titles[]/voice_library_ids[]/cover_urls[],成功后跳第5步
- 第5步「确认生成」:正式渲染实时进度,批量逐任务独立状态(BatchGenerationGrid,
  部分失败不阻塞、失败卡片单独「重试此视频」走 POST /tasks/{id}/retry),
  全部有终态后可进封面;单视频进度/失败/成功卡 + 右侧成片播放器
- 删除 useBatchPreview、ServerPreviewGrid 及 preview 渲染接口的所有前端调用
  (后端接口保留);E2E 适配回 6 步断言
2026-09-05 12:25:09 +08:00
93 changed files with 502 additions and 7161 deletions
-2
View File
@@ -1187,8 +1187,6 @@ jobs:
COSYVOICE_API_KEY: ${{ secrets.COSYVOICE_API_KEY }}
DASHSCOPE_API_KEY: ${{ secrets.DASHSCOPE_API_KEY }}
MEDIAKIT_API_KEY: ${{ secrets.MEDIAKIT_API_KEY }}
WECHAT_APP_ID: ${{ secrets.WECHAT_APP_ID }}
WECHAT_APP_SECRET: ${{ secrets.WECHAT_APP_SECRET }}
run: |
set -eu
echo "Rendering .env from template + secrets..."
@@ -1,34 +0,0 @@
"""add client_upload_id to assets and asset_id to ingest_jobs
Issue #1714:上传 complete 幂等 + worker 转码回写关联。
- assets.client_upload_id:客户端幂等 tokencomplete 去重)
- ingest_jobs.asset_idcomplete 阶段创建的占位 asset idworker 回写关联,
防止 HEVC 转码改写 storage_key 后找不到占位而兜底新建 READY 记录)
Revision ID: 066_upload_idempotency
Revises: 065_dup_record_sim_match
Create Date: 2026-09-05
"""
import sqlalchemy as sa
from alembic import op
revision = "066_upload_idempotency"
down_revision = "065_dup_record_sim_match"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column("assets", sa.Column("client_upload_id", sa.String(64), nullable=True))
op.create_index("ix_assets_client_upload_id", "assets", ["client_upload_id"])
op.add_column("ingest_jobs", sa.Column("asset_id", sa.String(36), nullable=False, server_default=""))
op.create_index("ix_ingest_jobs_asset_id", "ingest_jobs", ["asset_id"])
def downgrade() -> None:
op.drop_index("ix_ingest_jobs_asset_id", table_name="ingest_jobs")
op.drop_column("ingest_jobs", "asset_id")
op.drop_index("ix_assets_client_upload_id", table_name="assets")
op.drop_column("assets", "client_upload_id")
@@ -1,35 +0,0 @@
"""add celery_task_id to generation_tasks and ingest_jobs
Issue #1714:孤儿恢复/超时清理撤销队列消息。
- generation_tasks.celery_task_id:入队时记录的 Celery 消息 ID,清理时 revoke
- ingest_jobs.celery_task_id:同上(素材转码任务)
Revision ID: 067_celery_task_id
Revises: 066_upload_idempotency
Create Date: 2026-09-05
"""
import sqlalchemy as sa
from alembic import op
revision = "067_celery_task_id"
down_revision = "066_upload_idempotency"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"generation_tasks",
sa.Column("celery_task_id", sa.String(64), nullable=False, server_default=""),
)
op.add_column(
"ingest_jobs",
sa.Column("celery_task_id", sa.String(64), nullable=False, server_default=""),
)
def downgrade() -> None:
op.drop_column("ingest_jobs", "celery_task_id")
op.drop_column("generation_tasks", "celery_task_id")
+1 -139
View File
@@ -13,7 +13,7 @@ import jwt
from app.auth import AuthenticatedUser, blacklist_token, get_current_user
from app.config import settings
from app.dependencies import get_auth_email_service, get_auth_session_store, get_user_repository
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
from fastapi import APIRouter, Depends, Header, HTTPException, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from pydantic import BaseModel, EmailStr
@@ -84,7 +84,6 @@ class CurrentUserResponse(BaseModel):
phone: str = ""
phone_verified: bool = False
binding_complete: bool = False
wechat_bound: bool = False
class PasswordResetRequestModel(BaseModel):
@@ -273,7 +272,6 @@ async def get_current_user_info(
phone=user.phone or "",
phone_verified=user.phone_verified,
binding_complete=binding_complete,
wechat_bound=bool(user.wechat_openid),
)
@@ -428,7 +426,6 @@ async def get_wechat_auth_url() -> WechatAuthUrlResponse:
@router.post("/wechat/callback", response_model=WechatLoginResponse)
async def wechat_callback(
request: WechatCallbackRequest,
http_request: Request,
user_repository: UserRepository = Depends(get_user_repository),
) -> WechatLoginResponse:
"""微信登录回调处理"""
@@ -436,30 +433,11 @@ async def wechat_callback(
from packages.application.auth.wechat_sync_use_case import WechatSyncRequest as SyncRequest
from packages.application.auth.wechat_sync_use_case import WechatSyncUseCase
# 回调可观测性:记录 UA(区分微信内置浏览器 MicroMessenger)与 state
# 便于排查"停留 open.weixin.qq.com / 回调失败"类问题(#1718
user_agent = http_request.headers.get("User-Agent", "")
is_wechat_browser = "MicroMessenger" in user_agent
logger.info(
"[微信回调] 收到回调: state=%s code_len=%d UA=%r 微信内置浏览器=%s",
(request.state or "")[:8],
len(request.code or ""),
user_agent[:200],
is_wechat_browser,
)
# 1. 用 code 换微信用户信息
oauth_service = get_wechat_oauth_service()
wechat_user, err = oauth_service.handle_callback(request.code, request.state)
if err:
# state 校验失败 / 微信 errcode 等错误原文已在 service 内 log,这里带上 UA 上下文
logger.warning("[微信回调] 处理失败: err=%s 微信内置浏览器=%s", err, is_wechat_browser)
raise HTTPException(status_code=400, detail=err)
logger.info(
"[微信回调] state 校验通过,微信用户信息获取成功: openid=%s unionid=%s",
wechat_user.openid[:8] if wechat_user.openid else "",
bool(wechat_user.unionid),
)
# 2. 同步登录/注册(复用 wechat-sync 逻辑)
use_case = WechatSyncUseCase(user_repository=user_repository)
@@ -494,122 +472,6 @@ async def wechat_callback(
)
# ==================== 微信账号绑定/解绑(已登录用户) ====================
class WechatBindUrlResponse(BaseModel):
auth_url: str
state: str
class WechatBindCompleteRequest(BaseModel):
code: str
state: str = ""
class WechatBindUserProfile(BaseModel):
"""绑定/解绑后返回的用户信息(字段对齐 /auth/me,前端 normalizeUser 直接消费)"""
user_id: str
email: str
username: str
display_name: str
email_verified: bool
phone: str = ""
phone_verified: bool = False
binding_complete: bool = False
wechat_bound: bool = False
class WechatBindCompleteResponse(BaseModel):
success: bool
user: WechatBindUserProfile
class WechatUnbindResponse(BaseModel):
success: bool
def _wechat_user_profile(user) -> WechatBindUserProfile:
binding_complete = bool(
user.phone_verified and user.email_verified and user.email and "@wechat.local" not in user.email
)
return WechatBindUserProfile(
user_id=user.id,
email=user.email,
username=user.username,
display_name=user.display_name,
email_verified=user.email_verified,
phone=user.phone or "",
phone_verified=user.phone_verified,
binding_complete=binding_complete,
wechat_bound=bool(user.wechat_openid),
)
@router.get("/wechat/bind/url", response_model=WechatBindUrlResponse)
async def get_wechat_bind_url(
current_user: AuthenticatedUser = Depends(get_current_user),
) -> WechatBindUrlResponse:
"""获取微信绑定授权链接(已登录用户场景)。state 经 Redis 存储做 CSRF 校验。"""
from packages.application.auth.wechat_oauth_service import get_wechat_oauth_service
oauth_service = get_wechat_oauth_service()
auth_url, state = oauth_service.generate_auth_url()
logger.info("[微信绑定] 用户 %s 请求绑定授权链接", current_user.user.id)
return WechatBindUrlResponse(auth_url=auth_url, state=state)
@router.post("/wechat/bind", response_model=WechatBindCompleteResponse)
async def wechat_bind(
request: WechatBindCompleteRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
user_repository: UserRepository = Depends(get_user_repository),
) -> WechatBindCompleteResponse:
"""微信绑定完成:扫码回调后用 code 换 openid,绑定到当前登录账号(不创建新用户)。"""
from packages.application.auth.wechat_bind_use_case import WechatBindRequest, WechatBindUseCase
from packages.application.auth.wechat_oauth_service import get_wechat_oauth_service
oauth_service = get_wechat_oauth_service()
wechat_user, err = oauth_service.handle_callback(request.code, request.state)
if err:
logger.warning("[微信绑定] 用户 %s 换取微信信息失败: %s", current_user.user.id, err)
raise HTTPException(status_code=400, detail=err)
use_case = WechatBindUseCase(user_repository=user_repository)
result, error, http_status = use_case.bind(
WechatBindRequest(
user_id=current_user.user.id,
openid=wechat_user.openid,
unionid=wechat_user.unionid or "",
)
)
if error:
logger.warning("[微信绑定] 用户 %s 绑定失败: %s", current_user.user.id, error)
raise HTTPException(status_code=http_status, detail=error)
logger.info("[微信绑定] 用户 %s 绑定成功 openid=%s", current_user.user.id, wechat_user.openid[:8])
return WechatBindCompleteResponse(success=True, user=_wechat_user_profile(result.user))
@router.delete("/wechat/bind", response_model=WechatUnbindResponse)
async def wechat_unbind(
current_user: AuthenticatedUser = Depends(get_current_user),
user_repository: UserRepository = Depends(get_user_repository),
) -> WechatUnbindResponse:
"""解绑微信:需账号仍有其他登录方式(密码/手机/真实邮箱),否则拒绝。"""
from packages.application.auth.wechat_bind_use_case import WechatUnbindUseCase
use_case = WechatUnbindUseCase(user_repository=user_repository)
result, error, http_status = use_case.unbind(current_user.user.id)
if error:
logger.warning("[微信解绑] 用户 %s 解绑失败: %s", current_user.user.id, error)
raise HTTPException(status_code=http_status, detail=error)
logger.info("[微信解绑] 用户 %s 解绑成功", current_user.user.id)
return WechatUnbindResponse(success=True)
# ==================== 验证码 & 绑定 ====================
+1 -3
View File
@@ -14,7 +14,6 @@ from typing import Any
from uuid import uuid4
from app.api.routes._helpers import require_project_and_library
from app.api.routes.upload import _persist_celery_task_id
from app.auth import AuthenticatedUser, get_current_user
from app.core.celery_app import celery_app
from app.core.storage import OSSStorageService, get_storage_service
@@ -382,8 +381,7 @@ async def complete_chunked_upload(
file_hash=request.file_hash,
)
)
celery_result = celery_app.send_task("worker.ingest_asset", args=[job.id])
_persist_celery_task_id(ingest_job_repository, job, getattr(celery_result, "id", ""))
celery_app.send_task("worker.ingest_asset", args=[job.id])
# Update metadata status
meta["status"] = "completed"
+11 -19
View File
@@ -14,7 +14,6 @@ from app.core.task_enqueue import (
USER_PENDING_LIMIT,
GlobalQueueFull,
UserPendingLimitExceeded,
build_rate_limit_detail,
safe_enqueue_generation_task,
)
from app.dependencies import (
@@ -313,12 +312,12 @@ def create_preview_generation_task(
except UserPendingLimitExceeded as e:
raise HTTPException(
status_code=429,
detail=build_rate_limit_detail(e, generation_task_repository, scope="user"),
detail=f"您的待处理任务过多(当前 {e.pending_count - count}/{e.limit},本次提交 {count} 个),请等待后再提交",
) from e
except GlobalQueueFull as e:
raise HTTPException(
status_code=503,
detail=build_rate_limit_detail(e, generation_task_repository, scope="global"),
detail="系统繁忙,请稍后再试",
) from e
# 确定视频比例:优先前端传入,否则从模板 mode 推断
@@ -499,7 +498,6 @@ def create_preview_generation_task(
# ── 入队 ──
responses: list[PreviewGenerationTaskResponse] = []
rate_limit_exc: Exception | None = None # 记录首个限流异常,全部失败时返回结构化提示
for variant_index, task in enumerate(created_tasks):
try:
enqueued = safe_enqueue_generation_task(
@@ -512,29 +510,23 @@ def create_preview_generation_task(
if not enqueued:
logger.warning("[预览生成] 任务入队失败: task_id=%s", task.id)
_mark_task_failed(generation_task_repository, task, "任务入队失败")
except UserPendingLimitExceeded as e:
except UserPendingLimitExceeded:
_mark_task_failed(generation_task_repository, task, "待处理任务超限")
rate_limit_exc = rate_limit_exc or e
except GlobalQueueFull as e:
except GlobalQueueFull:
_mark_task_failed(generation_task_repository, task, "系统队列已满")
rate_limit_exc = rate_limit_exc or e
except Exception:
logger.exception("[预览生成] 入队异常: task_id=%s", task.id)
_mark_task_failed(generation_task_repository, task, "任务入队异常")
# enqueue 会原地更新 task 状态/进度,直接用 task 构造响应
responses.append(_to_preview_response(task))
# 队列满/限流时若全部失败,返回结构化错误码(前端区分"排队"与"创建失败"
if all(r.status == "failed" for r in responses) and rate_limit_exc is not None:
if isinstance(rate_limit_exc, UserPendingLimitExceeded):
raise HTTPException(
status_code=429,
detail=build_rate_limit_detail(rate_limit_exc, generation_task_repository, scope="user"),
)
raise HTTPException(
status_code=503,
detail=build_rate_limit_detail(rate_limit_exc, generation_task_repository, scope="global"),
)
# 队列满/限流时若全部失败,返回明确错误码
if all(r.status == "failed" for r in responses):
first_err = next((r.error_message for r in responses if r.error_message), "")
if "待处理任务" in first_err:
raise HTTPException(status_code=429, detail=first_err or "待处理任务超限")
if "队列" in first_err:
raise HTTPException(status_code=503, detail=first_err or "系统繁忙,请稍后再试")
logger.info(
"[预览生成] 创建完成: %d 个变体任务, task_ids=%s",
+14 -27
View File
@@ -10,7 +10,6 @@ from app.core.task_enqueue import (
USER_PENDING_LIMIT,
GlobalQueueFull,
UserPendingLimitExceeded,
build_rate_limit_detail,
safe_enqueue_generation_task,
)
from app.dependencies import (
@@ -418,12 +417,12 @@ def create_generation_task(
except UserPendingLimitExceeded as e:
raise HTTPException(
status_code=429,
detail=build_rate_limit_detail(e, generation_task_repository, scope="user"),
detail=f"您的待处理任务过多(当前 {e.pending_count - count}/{e.limit},本次提交 {count} 个),请等待完成后再提交",
) from e
except GlobalQueueFull as e:
raise HTTPException(
status_code=503,
detail=build_rate_limit_detail(e, generation_task_repository, scope="global"),
detail="系统繁忙,请稍后再试",
) from e
# 画中画已下线:strategy_id 中的 pip/voice_pip 统一映射为 one_take
@@ -580,7 +579,7 @@ def create_generation_task(
if not created_tasks:
raise HTTPException(
status_code=429,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
detail="您的待处理任务过多,请等待完成后再提交",
) from _e
break
except GlobalQueueFull as _e:
@@ -588,7 +587,7 @@ def create_generation_task(
if not created_tasks:
raise HTTPException(
status_code=503,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"),
detail="系统繁忙,请稍后再试",
) from _e
break
except HTTPException:
@@ -714,15 +713,15 @@ def confirm_generation(
log_task_status=True,
):
logger.warning("[确认生成] 入队失败: task_id=%s", new_task.id)
except UserPendingLimitExceeded as _e:
except UserPendingLimitExceeded:
raise HTTPException(
status_code=429,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
detail="您的待处理任务过多,请等待完成后再提交",
) from None
except GlobalQueueFull as _e:
except GlobalQueueFull:
raise HTTPException(
status_code=503,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"),
detail="系统繁忙,请稍后再试",
) from None
return BatchGenerationTaskResponse(
@@ -804,24 +803,12 @@ def retry_generation_task(
if user_pending >= USER_PENDING_LIMIT:
raise HTTPException(
status_code=429,
detail=build_rate_limit_detail(
UserPendingLimitExceeded(
user_id=user_id,
pending_count=user_pending,
limit=USER_PENDING_LIMIT,
),
generation_task_repository,
scope="user",
),
detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
)
if global_pending >= GLOBAL_PENDING_LIMIT:
raise HTTPException(
status_code=503,
detail=build_rate_limit_detail(
GlobalQueueFull(pending_count=global_pending, limit=GLOBAL_PENDING_LIMIT),
generation_task_repository,
scope="global",
),
detail="系统繁忙,请稍后再试",
)
use_case = CreateGenerationTaskUseCase(generation_task_repository)
@@ -856,15 +843,15 @@ def retry_generation_task(
log_task_status=True,
):
logger.warning("[生成任务] 重试入队失败: task_id=%s", retried.id)
except UserPendingLimitExceeded as _e:
except UserPendingLimitExceeded:
raise HTTPException(
status_code=429,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
detail="您的待处理任务过多,请等待完成后再提交",
) from None
except GlobalQueueFull as _e:
except GlobalQueueFull:
raise HTTPException(
status_code=503,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"),
detail="系统繁忙,请稍后再试",
) from None
return _to_generation_task_response(retried)
+1 -7
View File
@@ -43,13 +43,7 @@ def submit_ingest_job(
)
)
celery_result = celery_app.send_task("worker.ingest_asset", args=[job.id])
if getattr(celery_result, "id", ""):
try:
job.celery_task_id = celery_result.id
ingest_job_repository.update(job)
except Exception: # noqa: BLE001
pass
celery_app.send_task("worker.ingest_asset", args=[job.id])
return IngestJobResponse(
id=job.id,
+1 -7
View File
@@ -375,13 +375,7 @@ def retry_project_task(
storage_key=job.storage_key,
)
)
celery_result = celery_app.send_task("worker.ingest_asset", args=[retried.id])
if getattr(celery_result, "id", ""):
try:
retried.celery_task_id = celery_result.id
ingest_job_repository.update(retried)
except Exception: # noqa: BLE001
pass
celery_app.send_task("worker.ingest_asset", args=[retried.id])
return ProjectTaskResponse(
id=f"ingest:{retried.id}",
task_type="ingest",
+55 -162
View File
@@ -85,22 +85,12 @@ def _infer_mime_type_from_storage_key(storage_key: str) -> str:
"""从 storage_key 推断 MIME 类型(与 worker 端保持一致)。"""
lower_filename = storage_key.rsplit("/", 1)[-1].lower()
_MIME_MAP = {
".mov": "video/quicktime",
".mp4": "video/mp4",
".avi": "video/x-msvideo",
".mkv": "video/x-matroska",
".webm": "video/webm",
".png": "image/png",
".gif": "image/gif",
".bmp": "image/bmp",
".svg": "image/svg+xml",
".jpg": "image/jpeg",
".jpeg": "image/jpeg",
".mp3": "audio/mpeg",
".wav": "audio/wav",
".ogg": "audio/ogg",
".flac": "audio/flac",
".m4a": "audio/x-m4a",
".mov": "video/quicktime", ".mp4": "video/mp4", ".avi": "video/x-msvideo",
".mkv": "video/x-matroska", ".webm": "video/webm",
".png": "image/png", ".gif": "image/gif", ".bmp": "image/bmp",
".svg": "image/svg+xml", ".jpg": "image/jpeg", ".jpeg": "image/jpeg",
".mp3": "audio/mpeg", ".wav": "audio/wav", ".ogg": "audio/ogg",
".flac": "audio/flac", ".m4a": "audio/x-m4a",
}
for ext, mime in _MIME_MAP.items():
if lower_filename.endswith(ext):
@@ -108,85 +98,8 @@ def _infer_mime_type_from_storage_key(storage_key: str) -> str:
return "video/mp4" # default
# 兜底去重:无 file_hash / client_upload_id 时,同库同名近期活动记录视为重复
FALLBACK_DEDUP_WINDOW_MINUTES = 30
ACTIVE_ASSET_STATUSES = (AssetStatus.UPLOADING, AssetStatus.PROCESSING)
def _find_duplicate_asset(
asset_repository: Any,
*,
library_id: str,
file_hash: str,
client_upload_id: str,
filename: str,
file_size: int = 0,
) -> Any:
"""complete/上传幂等去重,按优先级查找已存在的素材。
1. client_upload_id(客户端幂等 token,同一次上传的重试保持一致)
2. file_hash(内容哈希,不同上传只要内容相同即去重)
3. 兜底:同库 + 同文件名(+同大小)且 30 分钟内仍处 uploading/processing
的记录——旧客户端不传 hash/token 时,防止 complete 超时重试反复建占位。
全部为鸭子类型调用:旧仓储无对应方法时静默跳过,不破坏既有实现。
"""
if client_upload_id:
find = getattr(asset_repository, "find_by_library_and_client_upload_id", None)
if callable(find):
existing = find(library_id=library_id, client_upload_id=client_upload_id)
if existing is not None:
logger.info(
"素材幂等命中(client_upload_id): library=%s token=%s asset=%s",
library_id,
client_upload_id,
getattr(existing, "id", "?"),
)
return existing
if file_hash:
existing = asset_repository.find_by_library_and_file_hash(
library_id=library_id,
file_hash=file_hash,
)
if existing is not None:
logger.info(
"素材去重命中(file_hash): library=%s hash=%s asset=%s",
library_id,
file_hash,
existing.id,
)
return existing
if filename:
find_recent = getattr(asset_repository, "find_recent_active_by_library_and_name", None)
if callable(find_recent):
existing = find_recent(
library_id=library_id,
name=filename,
within_minutes=FALLBACK_DEDUP_WINDOW_MINUTES,
file_size=file_size or 0,
)
if existing is not None and getattr(existing, "status", None) in ACTIVE_ASSET_STATUSES:
logger.info(
"素材幂等兜底命中(近期活动同名记录): library=%s name=%s asset=%s status=%s",
library_id,
filename,
getattr(existing, "id", "?"),
getattr(existing, "status", "?"),
)
return existing
return None
def _create_pending_asset(
asset_repository,
project_id,
library_id,
storage_key,
filename,
mime_type,
user_id,
file_hash="",
client_upload_id="",
asset_repository, project_id, library_id, storage_key, filename, mime_type, user_id, file_hash=""
):
"""立即创建一条 PROCESSING 状态的 Asset 记录,使前端能马上看到新素材。"""
asset = Asset.create(
@@ -198,29 +111,16 @@ def _create_pending_asset(
status=AssetStatus.PROCESSING,
uploaded_by_user_id=user_id,
file_hash=file_hash,
client_upload_id=client_upload_id,
)
return asset_repository.create(asset)
def _persist_celery_task_id(repo: Any, job: Any, celery_task_id: str) -> None:
"""记录 celery 消息 ID 到任务行,供孤儿清理时 revoke/清除队列消息(#1714)。"""
if not celery_task_id:
return
try:
job.celery_task_id = celery_task_id
repo.update(job)
except Exception: # noqa: BLE001 — 记录失败不影响主流程(执行前状态守卫兜底)
pass
def _submit_ingest_job(
project_id: str,
library_id: str,
storage_key: str,
ingest_job_repository: Any,
file_hash: str = "",
asset_id: str = "",
) -> Any:
use_case = SubmitIngestJobUseCase(ingest_job_repository)
job = use_case.execute(
@@ -229,11 +129,9 @@ def _submit_ingest_job(
library_id=library_id,
storage_key=storage_key,
file_hash=file_hash,
asset_id=asset_id,
)
)
celery_result = celery_app.send_task("worker.ingest_asset", args=[job.id])
_persist_celery_task_id(ingest_job_repository, job, getattr(celery_result, "id", ""))
celery_app.send_task("worker.ingest_asset", args=[job.id])
return job
@@ -304,7 +202,7 @@ async def complete_direct_upload(
asset_repository: Any = Depends(get_asset_repository),
storage_service: OSSStorageService = Depends(get_storage_service),
) -> DirectUploadCompleteResponse:
"""确认浏览器直传完成并创建导入任务(幂等:重复 complete 返回同一素材)"""
"""确认浏览器直传完成并创建导入任务。"""
require_project_and_library(
request.project_id,
request.library_id,
@@ -314,29 +212,6 @@ async def complete_direct_upload(
normalized_key = storage_service._normalize_storage_key(request.storage_key)
if not normalized_key.startswith("uploads/"):
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid upload key")
filename = normalized_key.rsplit("/", 1)[-1]
# ── 幂等去重(放在 OSS 检查之前):complete 超时后前端重试时,
# 第一次 complete 可能已建好占位记录,此时即使 OSS 检查失败也必须返回
# 已存在记录,绝不能再建第二条。─
existing = _find_duplicate_asset(
asset_repository,
library_id=request.library_id,
file_hash=request.file_hash,
client_upload_id=request.client_upload_id,
filename=filename,
file_size=request.file_size,
)
if existing is not None:
return DirectUploadCompleteResponse(
storage_key=existing.storage_key,
ingest_job_id="",
duplicated=True,
asset_id=existing.id,
url=storage_service.get_url(existing.storage_key),
)
try:
file_exists = storage_service.file_exists(normalized_key)
except Exception as error:
@@ -348,7 +223,29 @@ async def complete_direct_upload(
if not file_exists:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Uploaded file not found")
# ── 素材去重检测:同素材库 + 同 file_hash 视为重复 ──
if request.file_hash:
existing = asset_repository.find_by_library_and_file_hash(
library_id=request.library_id,
file_hash=request.file_hash,
)
if existing is not None:
logger.info(
"素材去重命中: library=%s hash=%s existing_asset=%s",
request.library_id,
request.file_hash,
existing.id,
)
return DirectUploadCompleteResponse(
storage_key=normalized_key,
ingest_job_id="",
duplicated=True,
asset_id=existing.id,
url=storage_service.get_url(normalized_key),
)
# 立即创建 Asset 记录(PROCESSING 状态),使前端刷新后即可看到新素材
filename = normalized_key.rsplit("/", 1)[-1]
mime_type = _infer_mime_type_from_storage_key(normalized_key)
pending_asset = _create_pending_asset(
asset_repository=asset_repository,
@@ -359,7 +256,6 @@ async def complete_direct_upload(
mime_type=mime_type,
user_id=authenticated_user.user.id,
file_hash=request.file_hash,
client_upload_id=request.client_upload_id,
)
job = _submit_ingest_job(
@@ -368,7 +264,6 @@ async def complete_direct_upload(
storage_key=normalized_key,
ingest_job_repository=ingest_job_repository,
file_hash=request.file_hash,
asset_id=pending_asset.id,
)
return DirectUploadCompleteResponse(
storage_key=normalized_key,
@@ -388,8 +283,7 @@ async def upload_asset(
project_id: str = Form(..., min_length=1, description="项目 ID"),
library_id: str = Form(..., min_length=1, description="素材库 ID"),
file: UploadFile = File(..., description="要上传的文件(视频、音频、图片等)"),
file_hash: str = Form(default="", description="文件哈希,用于去重检测"),
client_upload_id: str = Form(default="", description="客户端幂等 token(同一次上传的重试保持一致)"),
file_hash: str = Form(default="", description="文件 MD5 哈希,用于去重检测"),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
ingest_job_repository: Any = Depends(get_ingest_job_repository),
project_repository: Any = Depends(get_project_repository),
@@ -400,31 +294,32 @@ async def upload_asset(
"""上传素材文件并触发导入流水线。"""
require_project_and_library(project_id, library_id, project_repository, asset_library_repository)
# P2-5: 服务端验证 MIME 类型(先验证,再幂等去重,避免非法类型绕过)
# ── 素材去重检测:上传前检查同素材库 + 同 file_hash ──
if file_hash:
existing = asset_repository.find_by_library_and_file_hash(
library_id=library_id,
file_hash=file_hash,
)
if existing is not None:
logger.info(
"素材去重命中(multipart): library=%s hash=%s existing_asset=%s",
library_id,
file_hash,
existing.id,
)
return UploadAssetResponse(
storage_key=existing.storage_key,
ingest_job_id="",
url="",
duplicated=True,
asset_id=existing.id,
)
# P2-5: 服务端验证 MIME 类型
validated_content_type = _validate_mime_type(file.content_type)
safe_filename = file.filename.replace("/", "_").replace("\\", "_") if file.filename else "unknown"
# ── 幂等去重:client_upload_id → file_hash → 近期活动同名记录兜底 ──
# 放在 OSS 上传之前:重复提交直接返回,不占 OSS 流量、不建新记录。
existing = _find_duplicate_asset(
asset_repository,
library_id=library_id,
file_hash=file_hash,
client_upload_id=client_upload_id,
filename=safe_filename,
file_size=0,
)
if existing is not None:
return UploadAssetResponse(
storage_key=existing.storage_key,
ingest_job_id="",
url="",
duplicated=True,
asset_id=existing.id,
)
file_id = uuid4().hex[:8]
safe_filename = file.filename.replace("/", "_").replace("\\", "_") if file.filename else "unknown"
storage_key = f"uploads/{file_id}/{safe_filename}"
try:
@@ -453,7 +348,6 @@ async def upload_asset(
mime_type=validated_content_type,
user_id=authenticated_user.user.id,
file_hash=file_hash,
client_upload_id=client_upload_id,
)
job = _submit_ingest_job(
@@ -462,7 +356,6 @@ async def upload_asset(
storage_key=storage_key,
ingest_job_repository=ingest_job_repository,
file_hash=file_hash,
asset_id=pending_asset.id,
)
return UploadAssetResponse(
-8
View File
@@ -5,11 +5,3 @@ settings = get_settings()
celery_app = Celery("xiaoxia-saas-api")
celery_app.conf.broker_url = settings.CELERY_BROKER_URL
celery_app.conf.result_backend = settings.CELERY_RESULT_BACKEND
# #1714 队列隔离:视频生成走 generation 队列,素材转码走 transcode 队列
try:
from packages.shared.celery_queues import apply_queue_settings
apply_queue_settings(celery_app)
except Exception: # noqa: BLE001 — 队列配置失败不阻断 API 启动
pass
+3 -131
View File
@@ -8,145 +8,27 @@ logger = logging.getLogger(__name__)
# ── 限流阈值常量(全系统统一管理,不要在业务代码里硬编码) ──
USER_PENDING_LIMIT = 3 # 单用户 pending 上限
GLOBAL_PENDING_LIMIT = 20 # 全局 pending 上限
WORKER_CONCURRENCY = 4 # worker 渲染并发数(infra/docker/compose.yml WORKER_CONCURRENCY 默认值)
# 限流错误码:前端据此区分"排队等待"与"创建失败"
ERROR_CODE_USER_QUEUE_FULL = "USER_QUEUE_FULL" # 429:用户自己的任务排队中
ERROR_CODE_SYSTEM_QUEUE_FULL = "SYSTEM_QUEUE_FULL" # 503:系统整体繁忙
class UserPendingLimitExceeded(Exception):
"""用户 pending 任务数超限,返回 429。"""
def __init__(
self,
user_id: str,
pending_count: int,
limit: int,
*,
running_count: int = 0,
requested_count: int = 1,
queue_ahead: int = 0,
estimated_wait_seconds: int = 0,
):
def __init__(self, user_id: str, pending_count: int, limit: int):
self.user_id = user_id
self.pending_count = pending_count
self.limit = limit
# 排队上下文(用于 429 结构化提示,前端展示"排队中"而非"创建失败"
self.running_count = running_count
self.requested_count = requested_count
self.queue_ahead = queue_ahead
self.estimated_wait_seconds = estimated_wait_seconds
super().__init__(f"用户 {user_id} pending 任务数 {pending_count} 超过上限 {limit}")
class GlobalQueueFull(Exception):
"""全局限流,返回 503。"""
def __init__(
self,
pending_count: int,
limit: int,
*,
running_count: int = 0,
queue_ahead: int = 0,
estimated_wait_seconds: int = 0,
):
def __init__(self, pending_count: int, limit: int):
self.pending_count = pending_count
self.limit = limit
self.running_count = running_count
self.queue_ahead = queue_ahead
self.estimated_wait_seconds = estimated_wait_seconds
super().__init__(f"系统 pending 任务数 {pending_count} 超过上限 {limit}")
def _estimate_wait_seconds(queue_ahead: int, generation_task_repository: Any) -> int:
"""根据排队任务数 + worker 并发数 + 历史平均任务耗时估算等待秒数。
估算公式:ceil(排队任务数 / 并发数) × 平均单任务耗时。
拿不到历史数据时仓储层返回默认 120 秒。
"""
import math
if queue_ahead <= 0:
return 0
try:
estimator = getattr(generation_task_repository, "estimate_avg_duration_seconds", None)
avg_seconds = estimator() if estimator is not None else 120.0
except Exception:
avg_seconds = 120.0
return int(math.ceil(queue_ahead / WORKER_CONCURRENCY) * avg_seconds)
def build_rate_limit_detail(
exc: Exception,
generation_task_repository: Any,
*,
scope: str = "user",
) -> dict:
"""构造结构化限流响应体(HTTPException 的 detail)。
前端按 detail.code 判断场景:
- USER_QUEUE_FULL (429):用户自己的任务在排队,应提示"等待/继续排队",不是创建失败
- SYSTEM_QUEUE_FULL (503):系统繁忙,稍后重试
detail 字段:
- code: 错误码
- message: 可读中文提示(可直接展示)
- queued_count: 当前排队(pending)任务数
- running_count: 当前渲染中(running)任务数
- queue_ahead: 前方排队任务数(预计等待批次依据)
- estimated_wait_seconds: 预计等待秒数
- limit: 对应限流上限
"""
if scope == "user" and isinstance(exc, UserPendingLimitExceeded):
running = exc.running_count
if not running:
try:
counter = getattr(generation_task_repository, "count_running_by_user", None)
running = counter(exc.user_id) if counter is not None else 0
except Exception:
running = 0
queue_ahead = exc.queue_ahead or max(exc.pending_count, 0)
wait = exc.estimated_wait_seconds or _estimate_wait_seconds(queue_ahead, generation_task_repository)
wait_minutes = max(1, round(wait / 60))
message = (
f"您有 {exc.pending_count} 个任务正在排队、{running} 个正在渲染,"
f"同一时间最多提交 {exc.limit} 个任务。请等待约 {wait_minutes} 分钟后再提交"
)
return {
"code": ERROR_CODE_USER_QUEUE_FULL,
"message": message,
"queued_count": exc.pending_count,
"running_count": running,
"queue_ahead": queue_ahead,
"estimated_wait_seconds": wait,
"limit": exc.limit,
}
# 全局繁忙
pending = getattr(exc, "pending_count", 0)
running = getattr(exc, "running_count", 0)
if not running:
try:
counter = getattr(generation_task_repository, "count_running_total", None)
running = counter() if counter is not None else 0
except Exception:
running = 0
queue_ahead = getattr(exc, "queue_ahead", 0) or pending
wait = getattr(exc, "estimated_wait_seconds", 0) or _estimate_wait_seconds(queue_ahead, generation_task_repository)
wait_minutes = max(1, round(wait / 60))
return {
"code": ERROR_CODE_SYSTEM_QUEUE_FULL,
"message": f"系统繁忙:当前 {pending} 个任务排队中、{running} 个渲染中,预计等待约 {wait_minutes} 分钟,请稍后再试",
"queued_count": pending,
"running_count": running,
"queue_ahead": queue_ahead,
"estimated_wait_seconds": wait,
"limit": getattr(exc, "limit", GLOBAL_PENDING_LIMIT),
}
def check_queue_limits(
user_id: str,
generation_task_repository: Any,
@@ -279,17 +161,7 @@ def safe_enqueue_generation_task(
# ── 发送 Celery 任务 ──
try:
celery_result = celery_app.send_task("worker.generate_video", args=[task.id])
# 记录 celery 消息 ID:孤儿清理/超时作废时据此 revoke + 清除队列消息(#1714
celery_task_id = getattr(celery_result, "id", "")
if celery_task_id:
try:
task.celery_task_id = celery_task_id
generation_task_repository.update(task)
except Exception as persist_err: # noqa: BLE001
logger.warning(
"%s 持久化 celery_task_id 失败(不影响主流程): task_id=%s err=%s", log_prefix, task.id, persist_err
)
celery_app.send_task("worker.generate_video", args=[task.id])
except Exception as e:
logger.error(
"%s 入队失败,标记为失败: task_id=%s error=%s",
+5 -7
View File
@@ -31,16 +31,14 @@ class DirectUploadCompleteRequest(BaseModel):
project_id: str = Field(..., min_length=1)
library_id: str = Field(..., min_length=1)
storage_key: str = Field(..., min_length=1, max_length=255)
file_hash: str = Field(default="", max_length=64, description="文件哈希,用于去重检测")
client_upload_id: str = Field(default="", max_length=64, description="客户端幂等 token(同一次上传的重试保持一致)")
file_size: int = Field(default=0, ge=0, description="文件大小(字节),用于无 hash 时的兜底去重")
file_hash: str = Field(default="", max_length=64, description="文件 MD5 哈希,用于去重检测")
class DirectUploadCompleteResponse(BaseModel):
storage_key: str
ingest_job_id: str
duplicated: bool = Field(default=False, description="是否为重复素材/重复 complete(命中幂等去重)")
asset_id: str = Field(default="", description="素材 asset_id重复 complete 时返回已存在记录")
duplicated: bool = Field(default=False, description="是否为重复素材(命中去重)")
asset_id: str = Field(default="", description="重复素材 asset_idduplicated=true 时返回")
url: str = Field(default="", description="Public URL of uploaded file")
@@ -48,5 +46,5 @@ class UploadAssetResponse(BaseModel):
storage_key: str
ingest_job_id: str
url: str = Field(..., description="Public URL of uploaded file")
duplicated: bool = Field(default=False, description="是否为重复素材/重复提交(命中幂等去重)")
asset_id: str = Field(default="", description="素材 asset_id重复提交时返回已存在记录")
duplicated: bool = Field(default=False, description="是否为重复素材(命中去重)")
asset_id: str = Field(default="", description="重复素材 asset_idduplicated=true 时返回")
+4 -44
View File
@@ -4,7 +4,6 @@
import apiClient from "../client"
import { getOrCreateDefaultProject } from "../projects"
import type { DirectUploadPrepareResult, DirectUploadCompleteResult } from "./types"
import { computeFileHash, makeClientUploadId } from "./uploadDedup"
/** 预签名直传准备 */
export const prepareDirectUpload = async (data: {
@@ -13,13 +12,8 @@ export const prepareDirectUpload = async (data: {
filename: string
content_type: string
file_size: number
/** 前端算好的文件内容哈希(SHA-256 hex),打开后端 file_hash 去重闸门 */
file_hash?: string
/** 前端生成的上传幂等 token,同一次逻辑上传(含重试)保持不变 */
client_upload_id?: string
}): Promise<DirectUploadPrepareResult> => {
// prepare 单独放宽到 30s(全局 axios 实例只有 10sstaging 抖动时易超时)
const response = await apiClient.post("/upload/direct/prepare", data, { timeout: 30_000 })
const response = await apiClient.post("/upload/direct/prepare", data)
return response.data
}
@@ -28,14 +22,8 @@ export const completeDirectUpload = async (data: {
project_id: string
library_id: string
storage_key: string
/** 前端算好的文件内容哈希(与 prepare 一致),后端按 hash 幂等去重 */
file_hash?: string
/** 前端上传幂等 token(与 prepare 一致),同一次上传重发 complete 不重复建记录 */
client_upload_id?: string
}): Promise<DirectUploadCompleteResult> => {
// complete 内含 OSS 存在性检查 + 建库 + 派单,放宽到 60s;
// 超时不代表失败(记录可能已建成),调用方禁止超时后盲目重传整个文件
const response = await apiClient.post("/upload/direct/complete", data, { timeout: 60_000 })
const response = await apiClient.post("/upload/direct/complete", data)
return response.data
}
@@ -121,20 +109,8 @@ export interface DirectUploadHandle {
export const prepareDirectUploadHandle = async (data: {
file: File
library_id: string
/** 前端算好的文件内容哈希(SHA-256 hex),prepare/complete 均携带 */
fileHash?: string
/** 本次逻辑上传的幂等 tokenprepare/complete 一致、重试复用 */
clientUploadId?: string
}): Promise<DirectUploadHandle> => {
// 默认项目初始化失败(项目列表接口异常/自动创建失败)给出独立、明确的提示,
// 不与 prepare 的签名接口错误混在一起
let project: Awaited<ReturnType<typeof getOrCreateDefaultProject>>
try {
project = await getOrCreateDefaultProject()
} catch (err) {
const reason = err instanceof Error ? err.message : "网络异常"
throw new Error(`初始化默认项目失败,无法开始上传:${reason}`)
}
const project = await getOrCreateDefaultProject()
const prepared = await prepareDirectUpload({
project_id: project.id,
@@ -142,8 +118,6 @@ export const prepareDirectUploadHandle = async (data: {
filename: data.file.name,
content_type: data.file.type || "application/octet-stream",
file_size: data.file.size,
file_hash: data.fileHash,
client_upload_id: data.clientUploadId,
})
return {
@@ -154,8 +128,6 @@ export const prepareDirectUploadHandle = async (data: {
project_id: project.id,
library_id: data.library_id,
storage_key: prepared.storage_key,
file_hash: data.fileHash,
client_upload_id: data.clientUploadId,
}),
}
}
@@ -165,20 +137,8 @@ export const uploadAssetDirect = async (data: {
file: File
library_id: string
onProgress?: (percent: number) => void
/** 文件内容哈希;未传时自动补算(配音/封面/克隆等非队列链路统一受益) */
fileHash?: string
/** 幂等 token;未传时自动生成 */
clientUploadId?: string
}): Promise<DirectUploadCompleteResult> => {
// 自动补算哈希与幂等 token:确保 file_hash 去重闸门对所有上传链路生效
const fileHash = data.fileHash ?? (await computeFileHash(data.file))
const clientUploadId = data.clientUploadId ?? makeClientUploadId()
const handle = await prepareDirectUploadHandle({
file: data.file,
library_id: data.library_id,
fileHash,
clientUploadId,
})
const handle = await prepareDirectUploadHandle({ file: data.file, library_id: data.library_id })
await handle.transfer(data.onProgress)
return handle.complete()
}
-143
View File
@@ -1,143 +0,0 @@
/**
* 上传去重 / 幂等工具(Issue #1714
*
* 背景:同一文件被反复入队、complete 超时后盲目重传,导致后端创建大量重复
* PROCESSING 素材记录。本模块提供两类纯函数:
*
* 1. 文件指纹:
* - makeFileFingerprint():文件名+大小+lastModified,入队去重用(同步、零开销)
* - computeFileHash()SHA-256 内容哈希(小文件全量、大文件抽样头尾),
* prepare/complete 时发给后端打开 file_hash 去重闸门
* 2. 队列去重:findDuplicateInQueue() 判断文件是否已在队列中
* 3. 幂等 tokenmakeClientUploadId() 生成上传幂等 ID(每次"一次逻辑上传"一个,
* 重试复用同一 ID,重新入队才生成新 ID)
*/
/** 全量哈希阈值:≤64MB 全量读入计算;超过即走头尾抽样,避免 100~256MB 视频被整文件读进内存卡死页面 */
export const HASH_FULL_READ_LIMIT = 64 * 1024 * 1024 // 64MB
/** 抽样读取的头尾片段大小(各 16MB) */
export const HASH_SAMPLE_CHUNK = 16 * 1024 * 1024
/** 计算指纹时,文件在队列中已存在的状态(已失败的可以重试,不算重复) */
export type DedupExcludeStatus = "error" | "done"
/**
* 文件入队指纹:同库 + 文件名 + 大小 + 修改时间。
* 同一文件(File 对象由 <input> 重选或拖拽重复触发时三个字段均一致)稳定复现;
* 不同文件极小概率碰撞时可由后端 file_hash 内容去重兜底。
*/
export function makeFileFingerprint(file: Pick<File, "name" | "size" | "lastModified">): string {
return `${file.name}::${file.size}::${file.lastModified}`
}
/**
* 在现有队列项中查找同一文件的在途记录。
* 已失败(error)的项允许重试路径复用、已完成(done)的可跳过;
* 处于 preparing/uploading/ingesting 的在途项一律视为重复,禁止重复入队。
*
* 返回命中的队列项 id(tempId),未命中返回 null。
*/
export function findDuplicateInQueue<T extends { fileKey: string; status: string }>(
queue: T[],
fileKey: string,
excludeStatuses: DedupExcludeStatus[] = [],
): T | null {
const exclude = new Set<string>(excludeStatuses)
return queue.find((it) => it.fileKey === fileKey && !exclude.has(it.status)) ?? null
}
/** 生成上传幂等 token:一次"逻辑上传"一个,重试复用、重新入队换新 */
export function makeClientUploadId(): string {
const rand =
typeof crypto !== "undefined" && "randomUUID" in crypto
? crypto.randomUUID()
: `${Date.now()}-${Math.random().toString(36).slice(2, 10)}-${Math.random()
.toString(36)
.slice(2, 10)}`
return `up_${Date.now().toString(36)}_${rand.replace(/-/g, "").slice(0, 16)}`
}
/** 读取 Blob/File 片段为 ArrayBuffer:优先 Blob.arrayBuffer(),老环境回退 FileReader */
function readAsArrayBuffer(blob: Blob): Promise<ArrayBuffer> {
if (typeof blob.arrayBuffer === "function") {
return blob.arrayBuffer()
}
return new Promise<ArrayBuffer>((resolve, reject) => {
const reader = new FileReader()
reader.onload = () => resolve(reader.result as ArrayBuffer)
reader.onerror = () => reject(reader.error ?? new Error("FileReader read failed"))
reader.readAsArrayBuffer(blob)
})
}
/**
* 把 buffer 复制到当前 JS realm 的 Uint8Array 再哈希。
* jsdom/测试环境中 Blob.arrayBuffer() 可能返回另一 realm 的 ArrayBuffer
* Node WebCrypto 的 WebIDL instanceof 校验会拒绝跨 realm 参数。
*/
async function digestSha256(buffer: ArrayBuffer): Promise<ArrayBuffer> {
const subtle =
typeof globalThis !== "undefined" && globalThis.crypto ? globalThis.crypto.subtle : null
if (!subtle) throw new Error("crypto.subtle unavailable")
const local = new Uint8Array(buffer.byteLength)
local.set(new Uint8Array(buffer))
return subtle.digest("SHA-256", local)
}
function toHex(buffer: ArrayBuffer): string {
const bytes = new Uint8Array(buffer)
let hex = ""
for (let i = 0; i < bytes.length; i += 1) {
hex += bytes[i].toString(16).padStart(2, "0")
}
return hex
}
/**
* 计算文件内容 SHA-256hex64 字符,与后端 file_hash 字段长度一致)。
* - ≤64MB:全量哈希,内容一致必然一致
* - >64MB:哈希「头部 16MB + 尾部 16MB + 文件大小」,视频素材体积大、
* 头部含 moov 元数据、尾部含 mdat 结尾,抽样碰撞概率可忽略,
* 且避免 100~256MB 视频被整文件读进内存导致页面卡死/崩溃
*
* 运行环境不支持 crypto.subtle(非安全上下文/老浏览器)时返回空字符串,
* 调用方据此降级为不传 hash(后端仍有幂等 token + 同文件名兜底去重)。
*/
export async function computeFileHash(file: File): Promise<string> {
try {
const subtle =
typeof globalThis !== "undefined" &&
globalThis.crypto &&
typeof globalThis.crypto.subtle?.digest === "function"
? globalThis.crypto.subtle
: null
if (!subtle) return ""
if (file.size <= HASH_FULL_READ_LIMIT) {
const data = await readAsArrayBuffer(file.slice(0, file.size))
return toHex(await digestSha256(data))
}
// 大文件:头 8MB + 尾 8MB + 大小,拼成一段后哈希
const head = await readAsArrayBuffer(file.slice(0, HASH_SAMPLE_CHUNK))
const tail =
file.size > HASH_SAMPLE_CHUNK
? await readAsArrayBuffer(file.slice(Math.max(0, file.size - HASH_SAMPLE_CHUNK), file.size))
: new ArrayBuffer(0)
const merged = new Uint8Array(head.byteLength + tail.byteLength + 8)
merged.set(new Uint8Array(head), 0)
merged.set(new Uint8Array(tail), head.byteLength)
const sizeView = new DataView(merged.buffer, head.byteLength + tail.byteLength, 8)
// 文件大小以 64 位大端写入(BigInt 最稳;不支持 BigInt64 时手算高低位)
if (typeof sizeView.setBigUint64 === "function") {
sizeView.setBigUint64(0, BigInt(file.size), false)
} else {
sizeView.setUint32(0, Math.floor(file.size / 0x100000000), false)
sizeView.setUint32(4, file.size >>> 0, false)
}
return toHex(await digestSha256(merged.buffer))
} catch (err) {
console.warn("[uploadDedup] 计算文件哈希失败,降级为不传 file_hash:", err)
return ""
}
}
+3 -14
View File
@@ -12,18 +12,13 @@ export type {
UserResponse,
WechatAuthUrlResponse,
WechatCallbackResponse,
WechatBindUrlResponse,
WechatBindCompleteResponse,
WechatUnbindResponse,
UpdateProfileRequest,
UpdateProfileResponse,
SendVerificationCodeRequest,
BindContactRequest,
BindContactResponse,
} from "./types"
// 用户工具函数
export { normalizeUser, updateProfile } from "./user"
export { normalizeUser } from "./user"
// 登录/注册/登出/刷新
export { login, refreshAccessToken, register, logout } from "./login"
@@ -37,14 +32,8 @@ export { requestPasswordReset, resetPassword } from "./password"
// 邮箱验证
export { verifyEmail } from "./email"
// 微信登录 / 绑定
export {
getWechatAuthUrl,
wechatCallback,
getWechatBindUrl,
bindWechat,
unbindWechat,
} from "./wechat"
// 微信登录
export { getWechatAuthUrl, wechatCallback } from "./wechat"
// 联系方式
export { sendVerificationCode, bindContact } from "./contact"
-44
View File
@@ -34,17 +34,6 @@ export interface User {
is_email_verified: boolean
email_verified: boolean
created_at?: string
/** 微信是否已绑定 */
wechat_bound?: boolean
/** 微信昵称(绑定后展示) */
wechat_nickname?: string
/** 头像 URL(微信头像等) */
avatar_url?: string
/** 手机号 */
phone?: string
phone_verified?: boolean
/** 资料是否完善(微信新用户首次登录为 false,需填昵称引导) */
profile_completed?: boolean
}
export interface UserResponse {
@@ -56,12 +45,6 @@ export interface UserResponse {
is_email_verified?: boolean
email_verified?: boolean
created_at?: string
wechat_bound?: boolean
wechat_nickname?: string
avatar_url?: string
phone?: string
phone_verified?: boolean
profile_completed?: boolean
}
export interface WechatAuthUrlResponse {
@@ -97,30 +80,3 @@ export interface BindContactResponse {
success: boolean
user: User
}
/** 更新个人资料请求 */
export interface UpdateProfileRequest {
display_name?: string
}
/** 更新个人资料响应(返回最新用户信息) */
export interface UpdateProfileResponse {
user: UserResponse
}
/** 微信绑定授权链接响应 */
export interface WechatBindUrlResponse {
auth_url: string
state: string
}
/** 微信绑定完成响应 */
export interface WechatBindCompleteResponse {
success: boolean
user: UserResponse
}
/** 微信解绑响应 */
export interface WechatUnbindResponse {
success: boolean
}
+1 -16
View File
@@ -1,5 +1,4 @@
import apiClient from "../client"
import type { User, UserResponse, UpdateProfileRequest, UpdateProfileResponse } from "./types"
import type { User, UserResponse } from "./types"
/**
* 规范化用户数据,兼容不同后端返回格式
@@ -17,19 +16,5 @@ export const normalizeUser = (data: UserResponse): User => {
is_email_verified: emailVerified,
email_verified: emailVerified,
created_at: data.created_at,
wechat_bound: data.wechat_bound,
wechat_nickname: data.wechat_nickname,
avatar_url: data.avatar_url,
phone: data.phone,
phone_verified: data.phone_verified,
profile_completed: data.profile_completed,
}
}
/**
* 更新个人资料(昵称等)
*/
export const updateProfile = async (data: UpdateProfileRequest): Promise<User> => {
const response = await apiClient.patch<UpdateProfileResponse>("/auth/me", data)
return normalizeUser(response.data.user)
}
+2 -35
View File
@@ -1,14 +1,8 @@
import apiClient from "../client"
import type {
WechatAuthUrlResponse,
WechatCallbackResponse,
WechatBindUrlResponse,
WechatBindCompleteResponse,
WechatUnbindResponse,
} from "./types"
import type { WechatAuthUrlResponse, WechatCallbackResponse } from "./types"
/**
* 获取微信授权链接(登录场景)
* 获取微信授权链接
*/
export const getWechatAuthUrl = async (): Promise<WechatAuthUrlResponse> => {
const response = await apiClient.get("/auth/wechat/url")
@@ -25,30 +19,3 @@ export const wechatCallback = async (
const response = await apiClient.post("/auth/wechat/callback", { code, state })
return response.data
}
/**
* 获取微信绑定授权链接(已登录用户绑定场景)
*/
export const getWechatBindUrl = async (): Promise<WechatBindUrlResponse> => {
const response = await apiClient.get("/auth/wechat/bind/url")
return response.data
}
/**
* 微信绑定完成(扫码回调后用 code 绑定到当前登录账号)
*/
export const bindWechat = async (
code: string,
state: string,
): Promise<WechatBindCompleteResponse> => {
const response = await apiClient.post("/auth/wechat/bind", { code, state })
return response.data
}
/**
* 解绑微信
*/
export const unbindWechat = async (): Promise<WechatUnbindResponse> => {
const response = await apiClient.delete("/auth/wechat/bind")
return response.data
}
-132
View File
@@ -1,132 +0,0 @@
/**
* 统一错误信息提取
* 把 axios 错误(后端 detail / FastAPI 校验错误 / HTTP 状态码)、XHR/OSS 错误、
* 网络/超时错误、普通 Error 统一转成「可直接展示给用户」的中文信息。
*
* 与 api/client.ts 响应拦截器的提示口径保持一致;拦截器负责全局 toast,
* 页面/队列卡片用本工具把真实原因展示在持久位置(回调页、失败卡片等)。
*/
import type { AxiosError } from "axios"
/** 后端错误响应体可能出现的字段(FastAPI:detail;历史接口:message/msg */
interface ErrorBody {
detail?: unknown
message?: unknown
msg?: unknown
}
/** FastAPI 422 校验错误单项 */
interface ValidationItem {
loc?: (string | number)[]
msg?: string
}
/** 从后端响应体提取人类可读信息(detail 可能是字符串、对象、422 数组) */
function extractBodyMessage(data: unknown): string {
if (!data || typeof data !== "object") return ""
const body = data as ErrorBody
const walk = (val: unknown): string => {
if (typeof val === "string") return val
if (Array.isArray(val)) {
// FastAPI 422: [{loc, msg, type}, ...] → 取每条 msg 拼接
const parts = val
.map((item) => {
if (typeof item === "string") return item
if (item && typeof item === "object") {
const v = item as ValidationItem
if (typeof v.msg === "string") {
const field = Array.isArray(v.loc) ? v.loc.filter((x) => x !== "body").join(".") : ""
return field ? `${field}: ${v.msg}` : v.msg
}
return walk(item)
}
return ""
})
.filter(Boolean)
return parts.join("")
}
if (val && typeof val === "object") {
const obj = val as Record<string, unknown>
if (typeof obj.message === "string") return obj.message
if (typeof obj.msg === "string") return obj.msg
if (typeof obj.detail === "string") return obj.detail
if (obj.message && typeof obj.message === "object") return walk(obj.message)
if (obj.msg && typeof obj.msg === "object") return walk(obj.msg)
try {
return JSON.stringify(val)
} catch {
return ""
}
}
return ""
}
return walk(body.detail) || walk(body.message) || walk(body.msg)
}
/** 无响应体时按 HTTP 状态码给出兜底提示(与 client.ts 拦截器口径一致) */
function statusFallback(status: number): string {
switch (status) {
case 400:
return "请求参数有误(HTTP 400"
case 401:
return "登录状态已失效,请重新登录(HTTP 401)"
case 403:
return "没有权限执行该操作(HTTP 403"
case 404:
return "请求的资源不存在(HTTP 404"
case 409:
return "操作冲突,资源状态已变化(HTTP 409)"
case 413:
return "文件过大,请缩小后重试(HTTP 413)"
case 415:
return "不支持的文件格式(HTTP 415"
case 429:
return "操作过于频繁,请稍后再试(HTTP 429)"
case 503:
return "服务暂不可用,请稍后再试(HTTP 503)"
default:
if (status >= 500) return `服务器繁忙,请稍后再试(HTTP ${status}`
return `请求失败(HTTP ${status}`
}
}
/**
* 从任意抛出值提取可展示的错误信息。
* @param fallback 全部提取失败时的兜底文案
*/
export function getErrorMessage(err: unknown, fallback = "操作失败,请稍后重试"): string {
if (!err) return fallback
// axios 错误(后端 JSON 响应 / HTTP 错误状态)
const ax = err as AxiosError<ErrorBody>
if (ax.isAxiosError || (typeof ax === "object" && "response" in (ax as object))) {
// 超时
if (ax.code === "ECONNABORTED" || /timeout/i.test(ax.message || "")) {
return "请求超时,请检查网络后重试"
}
const resp = ax.response
if (resp) {
const bodyMsg = extractBodyMessage(resp.data)
if (bodyMsg) return bodyMsg
return statusFallback(resp.status)
}
// 请求已发出但无响应(断网/CORS/DNS)
if (ax.request) return "网络连接异常,请检查网络设置"
return ax.message || fallback
}
if (err instanceof Error) {
// XHR 直传 OSS 失败等场景自带详细 message(含 HTTP 状态 + OSS Code/Message
if (err.message) return err.message
}
if (typeof err === "string") return err
return fallback
}
/** client.ts 拦截器是否已对该错误弹过全局 toast(__msgShown 标记) */
export function isErrorMsgShown(err: unknown): boolean {
return Boolean((err as { __msgShown?: boolean } | null)?.__msgShown)
}
@@ -42,7 +42,6 @@ const AssetLibrary: React.FC = () => {
assetsError,
assetsErrorObj,
refetchAssets,
stalledAssetIds,
searchText,
setSearchText,
filterType,
@@ -78,7 +77,6 @@ const AssetLibrary: React.FC = () => {
removeUpload,
clearFinished,
uploading,
transferActive,
activeCount,
pendingCount,
} = useAssetUpload({ effectiveLibId })
@@ -170,7 +168,6 @@ const AssetLibrary: React.FC = () => {
{/* 上传区域 */}
<AssetUploadZone
uploading={uploading}
transferActive={transferActive}
activeCount={activeCount}
pendingCount={pendingCount}
onUpload={enqueueUploads}
@@ -217,7 +214,6 @@ const AssetLibrary: React.FC = () => {
selectedIds={selectedIds}
diagnosingId={diagnosingId}
uploadProgressMap={uploadProgressMap}
stalledAssetIds={stalledAssetIds}
onRetry={refetchAssets}
onToggleSelect={toggleSelect}
onDiagnose={handleDiagnose}
-36
View File
@@ -831,20 +831,6 @@
color: #ef4444;
}
.xx-upload-queue-error-detail {
margin-top: 4px;
font-size: 12px;
line-height: 1.5;
color: #ef4444;
word-break: break-word;
white-space: normal;
}
.xx-upload-queue-error-hint {
margin-top: 2px;
color: #b45309;
}
.xx-upload-queue-actions {
display: flex;
gap: 6px;
@@ -1055,25 +1041,3 @@
background: #fef2f2;
color: #dc2626;
}
/* 上传入口禁用态(直传进行中,防重复提交,Issue #1714 */
.xx-asset-upload-btn:disabled {
opacity: 0.6;
cursor: not-allowed;
}
.xx-asset-upload-btn:disabled:hover {
opacity: 0.6;
}
.xx-asset-upload-btn:disabled:active {
transform: none;
}
/* 处理超时遮罩:创建超过 10 分钟仍在处理中(疑似后端卡住),停止转圈并警示 */
.xx-asset-thumb-stalled {
background: rgba(217, 119, 6, 0.28);
color: #fde68a;
backdrop-filter: blur(2px);
}
.xx-asset-thumb-stalled :first-child {
font-size: var(--font-size-2xl);
}
@@ -22,8 +22,6 @@ export interface AssetCardProps {
diagnosing?: boolean
/** 上传中实时进度(仅 uploading 态有值;ingesting 后由后端状态接管) */
uploadProgress?: { progress: number; uploading: boolean }
/** 处理超过 10 分钟仍未就绪(疑似后端卡住):停止转圈并提示处理超时 */
stalled?: boolean
onToggle: () => void
onDiagnose: () => void
onPlay: () => void
@@ -35,7 +33,6 @@ const AssetCard: React.FC<AssetCardProps> = ({
selected,
diagnosing,
uploadProgress,
stalled,
onToggle,
onDiagnose,
onPlay,
@@ -70,15 +67,11 @@ const AssetCard: React.FC<AssetCardProps> = ({
</div>
)}
{/* 转码/处理中遮罩(卡死超过 10 分钟时停止转圈,提示超时) */}
{/* 转码/处理中遮罩 */}
{asset.loading && !isUploading && (
<div
className={`xx-asset-thumb-overlay ${
stalled ? "xx-asset-thumb-stalled" : "xx-asset-thumb-processing"
}`}
>
{stalled ? <CloseCircleOutlined /> : <LoadingOutlined />}
<span>{stalled ? "处理超时,可重试上传" : "转码处理中"}</span>
<div className="xx-asset-thumb-overlay xx-asset-thumb-processing">
<LoadingOutlined />
<span></span>
</div>
)}
@@ -136,7 +129,7 @@ const AssetCard: React.FC<AssetCardProps> = ({
</p>
<div className="xx-asset-meta">
<span className="xx-asset-meta-status">
<StatusPill status={asset.status} label={stalled ? "处理超时" : asset.statusLabel} />
<StatusPill status={asset.status} label={asset.statusLabel} />
</span>
{asset.duration && <span className="xx-asset-meta-duration">{asset.duration}</span>}
</div>
@@ -19,8 +19,6 @@ export interface AssetGridSectionProps {
selectedIds: Set<string>
diagnosingId: string | null
uploadProgressMap?: UploadProgressMap
/** 创建超过 10 分钟仍在处理中的素材 id(疑似后端卡住),卡片提示处理超时 */
stalledAssetIds?: Set<string>
onRetry?: () => void
onToggleSelect: (id: string) => void
onDiagnose: (asset: AssetItem) => void
@@ -36,7 +34,6 @@ export const AssetGridSection: React.FC<AssetGridSectionProps> = ({
selectedIds,
diagnosingId,
uploadProgressMap,
stalledAssetIds,
onRetry,
onToggleSelect,
onDiagnose,
@@ -79,7 +76,6 @@ export const AssetGridSection: React.FC<AssetGridSectionProps> = ({
selected={selectedIds.has(asset.id)}
diagnosing={diagnosingId === asset.id}
uploadProgress={uploadProgressMap?.get(asset.id)}
stalled={stalledAssetIds?.has(asset.id)}
onToggle={() => onToggleSelect(asset.id)}
onDiagnose={() => onDiagnose(asset)}
onPlay={() => onPlay(asset)}
@@ -4,13 +4,10 @@
* - 拖拽文件到内容区任意位置同样触发上传(不再占用大面积虚线框)
*/
import React, { useRef, useState } from "react"
import { message } from "antd"
import { PlusOutlined, CloudUploadOutlined } from "@ant-design/icons"
export interface AssetUploadZoneProps {
uploading: boolean
/** 有文件正在本地指纹/prepare/直传(非服务端转码),此时禁用入口防重复提交 */
transferActive: boolean
activeCount: number
pendingCount: number
onUpload: (files: File[]) => void
@@ -18,7 +15,6 @@ export interface AssetUploadZoneProps {
export const AssetUploadZone: React.FC<AssetUploadZoneProps> = ({
uploading,
transferActive,
activeCount,
pendingCount,
onUpload,
@@ -31,11 +27,6 @@ export const AssetUploadZone: React.FC<AssetUploadZoneProps> = ({
const pickFiles = (list: FileList | null) => {
if (!list || list.length === 0) return
// 直传进行中拦截重复触发:相同文件仍由入队去重兜底,这里先给明确反馈
if (transferActive) {
message.warning("文件正在上传中,请等待当前上传完成后再添加")
return
}
onUpload(Array.from(list))
}
@@ -67,18 +58,10 @@ export const AssetUploadZone: React.FC<AssetUploadZoneProps> = ({
<button
type="button"
className="xx-asset-upload-btn"
disabled={transferActive}
title={transferActive ? "文件上传中,暂不能添加新文件" : undefined}
onClick={() => {
if (transferActive) {
message.warning("文件正在上传中,请等待当前上传完成后再添加")
return
}
inputRef.current?.click()
}}
onClick={() => inputRef.current?.click()}
>
<PlusOutlined />
{transferActive ? "上传中…" : "上传素材"}
</button>
<span className="xx-asset-upload-status">
{uploading ? (
@@ -12,8 +12,7 @@ import {
ReloadOutlined,
CloseOutlined,
} from "@ant-design/icons"
import type { UploadItem, UploadFailStage } from "../hooks/useAssetUpload"
import { COMPLETE_RETRY_HINT } from "../hooks/useAssetUpload"
import type { UploadItem } from "../hooks/useAssetUpload"
export interface UploadQueuePanelProps {
items: UploadItem[]
@@ -30,13 +29,6 @@ const STATUS_TEXT: Record<UploadItem["status"], string> = {
error: "上传失败",
}
/** 失败阶段中文名:让用户一眼看到失败发生在哪一步 */
const FAIL_STAGE_TEXT: Record<UploadFailStage, string> = {
prepare: "准备上传阶段",
transfer: "文件传输阶段",
complete: "确认入库阶段",
}
const UploadQueuePanel: React.FC<UploadQueuePanelProps> = ({
items,
onRetry,
@@ -88,34 +80,16 @@ const UploadQueuePanel: React.FC<UploadQueuePanelProps> = ({
) : null}
<div className="xx-upload-queue-status">
{it.duplicated ? "素材已存在,已跳过" : STATUS_TEXT[it.status]}
{it.status === "preparing" && it.hint ? `${it.hint}` : ""}
{it.status === "uploading" ? ` ${it.progress}%` : ""}
{it.status === "error" && it.failedStage
? `${FAIL_STAGE_TEXT[it.failedStage]}`
: ""}
{it.status === "error" && it.error ? `${it.error}` : ""}
</div>
{it.status === "error" && it.error ? (
<div className="xx-upload-queue-error-detail" title={it.error}>
{it.error.split("\n").map((line, idx) =>
line === COMPLETE_RETRY_HINT ? (
<div key={idx} className="xx-upload-queue-error-hint">
{line}
</div>
) : (
<div key={idx}>{line}</div>
),
)}
</div>
) : null}
</div>
<span className="xx-upload-queue-actions">
{it.status === "error" && (
<button
type="button"
className="xx-upload-queue-btn"
title={
it.failedStage === "complete" ? "安全重试(只确认,不重新上传)" : "重试上传"
}
title="重试"
onClick={() => onRetry(it.tempId)}
>
<ReloadOutlined />
+33 -204
View File
@@ -2,65 +2,34 @@ import { useState, useCallback, useRef, useEffect } from "react"
import { useQueryClient } from "@tanstack/react-query"
import { message } from "antd"
import { prepareDirectUploadHandle, type DirectUploadHandle } from "@/api/assets"
import { getErrorMessage, isErrorMsgShown } from "@/api/errors"
import { MAX_FILE_SIZE } from "../constants"
import {
computeFileHash,
findDuplicateInQueue,
makeClientUploadId,
makeFileFingerprint,
} from "@/api/assets/uploadDedup"
/** 单文件上传状态机 */
export type UploadItemStatus = "preparing" | "uploading" | "ingesting" | "done" | "error"
/** 失败发生的阶段:complete 阶段失败时记录可能已在后端建成,禁止盲目重传整个文件 */
export type UploadFailStage = "prepare" | "transfer" | "complete"
export interface UploadItem {
/** 前端临时 idprepare 前无 asset_id 时用),同时作为队列项 key */
/** 前端临时 idprepare 前无 asset_id 时用) */
tempId: string
file: File
fileName: string
/** 进度 0~100(仅直传阶段有真实进度) */
progress: number
status: UploadItemStatus
/** 文件指纹(name+size+lastModified),入队去重用 */
fileKey: string
/** 本次逻辑上传的幂等 token:重试复用、重新入队才换新 */
clientUploadId: string
/** 上传前算好的文件内容哈希(SHA-256),prepare/complete 都带上 */
fileHash?: string
/** 后端 prepare 预建的 asset id(旧后端可能为空) */
assetId?: string
/** 去重命中:complete 返回 duplicated,标记完成但不产生新素材 */
duplicated?: boolean
/** 失败发生的阶段;complete 阶段失败点重试只重发 complete,不重新上传文件 */
failedStage?: UploadFailStage
/** 状态行补充提示(如"正在计算文件指纹…" */
hint?: string
error?: string
}
/** 批量直传最大并发数,避免多文件瓜分上行带宽 */
const MAX_CONCURRENT = 3
/** complete 阶段失败后的安全提示:素材可能已在后端建成,重试只重发 complete 幂等安全 */
export const COMPLETE_RETRY_HINT = "素材可能已在服务器处理中,点重试将安全确认,不会重新上传文件"
/** 失败阶段中文名(toast 提示用,明确失败发生在哪一步) */
const STAGE_LABEL: Record<UploadFailStage, string> = {
prepare: "准备上传",
transfer: "文件传输",
complete: "确认入库",
}
/**
* 素材批量上传 Hook
* - 入队按文件指纹(name+size+lastModified)去重:同一文件已在队列/上传中/处理中时不重复入队
* - 上传前计算文件 SHA-256prepare/complete 携带 file_hash + 幂等 tokenclientUploadId
* - complete 超时/失败不盲目重传:复用 handle 只重发 complete(幂等),prepare/transfer 失败才全量重跑
* - prepare 阶段后端预建 status=uploading 的 asset,前端拿到 asset_id 立即刷新列表
* - OSS 直传并发限制为 3,其余排队;每个文件独立进度/状态
* - complete 后素材进入转码(ingesting/processing),由列表轮询反映
* - 失败卡片支持重试/移除
*/
export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
@@ -70,13 +39,6 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
const itemsRef = useRef<UploadItem[]>([])
itemsRef.current = items
/**
* prepare 成功后的 handle 按 tempId 留存:
* complete 阶段失败(超时/网络)时 OSS 文件已存在、后端记录也可能已建成,
* 重试必须复用同一 handle 只重发 complete,绝不能重新 prepare+直传。
*/
const handlesRef = useRef<Map<string, DirectUploadHandle>>(new Map())
const updateItem = useCallback((tempId: string, patch: Partial<UploadItem>) => {
setItems((prev) => prev.map((it) => (it.tempId === tempId ? { ...it, ...patch } : it)))
}, [])
@@ -90,57 +52,14 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
queryClient.invalidateQueries({ queryKey: ["asset-libraries"] })
}, [queryClient, effectiveLibId])
/**
* 执行单个文件的完整上传流程。
* @param completeOnly complete 阶段失败后的重试:跳过 hash/prepare/transfer,只重发 complete
* (OSS 文件已传完,重发由 file_hash + clientUploadId 保证幂等)
*/
/** 执行单个文件的完整上传流程(prepare→transfer→complete */
const runUpload = useCallback(
async (item: UploadItem, handle?: DirectUploadHandle, completeOnly = false) => {
let stage: UploadFailStage = "prepare"
async (item: UploadItem, handle?: DirectUploadHandle) => {
try {
let h = handle
if (completeOnly && h) {
// ── complete 重试:文件已在 OSS,直接幂等重发确认 ──
stage = "complete"
updateItem(item.tempId, {
status: "ingesting",
progress: 100,
error: undefined,
failedStage: undefined,
hint: undefined,
})
const result = await h.complete()
refreshList()
handlesRef.current.delete(item.tempId)
if (result.duplicated) {
updateItem(item.tempId, { status: "done", duplicated: true, assetId: result.asset_id })
message.info(`"${item.fileName}" 与素材库已有内容相同,已跳过`)
} else {
updateItem(item.tempId, { status: "done", assetId: result.asset_id || item.assetId })
message.success(`"${item.fileName}" 上传完成,正在转码处理`)
}
return
}
// 1. 计算文件内容哈希(失败不阻塞,降级为不传 hash;后端仍有幂等 token 兜底)
updateItem(item.tempId, { hint: "正在计算文件指纹…" })
const fileHash = item.fileHash || (await computeFileHash(item.file))
updateItem(item.tempId, { fileHash, hint: undefined })
// 2. prepare(携带 file_hash + 幂等 token;重试时复用同一 clientUploadId
stage = "prepare"
h =
h ??
(await prepareDirectUploadHandle({
file: item.file,
library_id: effectiveLibId,
fileHash,
clientUploadId: item.clientUploadId,
}))
handlesRef.current.set(item.tempId, h)
// 1. prepare(重试时复用已准备的 handle 也行,但签名可能过期,重新 prepare 最稳)
const h =
handle ??
(await prepareDirectUploadHandle({ file: item.file, library_id: effectiveLibId }))
if (h.prepared.asset_id) {
updateItem(item.tempId, {
status: "uploading",
@@ -153,59 +72,26 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
updateItem(item.tempId, { status: "uploading", progress: 0 })
}
// 3. OSS 直传(真实进度)
stage = "transfer"
// 2. OSS 直传(真实进度)
await h.transfer((pct) => updateItem(item.tempId, { progress: pct }))
// 4. complete:后端确认入库并创建 ingest jobfile_hash + 幂等 token 已在 handle 闭包中)
stage = "complete"
// 3. complete:后端创建 ingest job,素材进入转码
updateItem(item.tempId, { status: "ingesting", progress: 100 })
const result = await h.complete()
refreshList()
handlesRef.current.delete(item.tempId)
if (result.duplicated) {
updateItem(item.tempId, { status: "done", duplicated: true, assetId: result.asset_id })
message.info(`"${item.fileName}" 与素材库已有内容相同,已跳过`)
} else {
updateItem(item.tempId, { status: "done", assetId: result.asset_id || item.assetId })
updateItem(item.tempId, { status: "done" })
message.success(`"${item.fileName}" 上传完成,正在转码处理`)
}
} catch (err: unknown) {
// 完整失败原因:HTTP 状态码 / OSS XML 的 Code+Message / 后端 detail
// 由 getErrorMessage 统一提取(OSS XHR 错误自带「OSS 直传失败: HTTP xxx ...」明细)
const detail = getErrorMessage(err, "未知错误")
console.error("[useAssetUpload] 上传失败:", item.fileName, stage, err)
if (stage === "complete") {
// complete 失败(超时/5xx/网络):后端记录可能已建成,handle 保留供幂等重试;
// 刷新列表让用户看到可能已创建的「处理中」素材,避免误以为没传上去而重复操作。
// 卡片同时展示真实错误原因 + 安全重试提示(重试只重发 complete,不重新上传)
refreshList()
updateItem(item.tempId, {
status: "error",
failedStage: "complete",
error: `${detail}\n${COMPLETE_RETRY_HINT}`,
hint: undefined,
})
if (!isErrorMsgShown(err)) {
message.error(`"${item.fileName}" 确认入库失败:${detail}`)
}
} else {
// prepare / transfer 失败:后端尚无素材记录,可安全全量重跑
handlesRef.current.delete(item.tempId)
updateItem(item.tempId, {
status: "error",
failedStage: stage === "transfer" ? "transfer" : "prepare",
error: detail,
hint: undefined,
})
// 拦截器已对后端错误弹过 toast(含真实 detail)时不重复弹;
// OSS XHR 直传错误不走 axios,必须在这里弹
if (!isErrorMsgShown(err)) {
message.error(`"${item.fileName}" ${STAGE_LABEL[stage]}失败:${detail}`)
}
}
const detail = err instanceof Error ? err.message : "上传失败"
console.error("[useAssetUpload] 上传失败:", item.fileName, err)
updateItem(item.tempId, { status: "error", error: detail })
message.error(`"${item.fileName}" 上传失败:${detail}`)
}
},
[effectiveLibId, refreshList, updateItem],
@@ -227,9 +113,7 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
if (!next) return
claimedRef.current.add(next.tempId)
inFlightRef.current += 1
// complete 阶段失败的重试:复用留存的 handle,只重发 complete
const existingHandle = handlesRef.current.get(next.tempId)
void runUpload(next, existingHandle, existingHandle !== undefined).finally(() => {
void runUpload(next).finally(() => {
inFlightRef.current -= 1
claimedRef.current.delete(next.tempId)
// 一个任务结束(成功/失败)后继续拉起排队任务
@@ -242,11 +126,7 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
pumpRef.current()
}, [items])
/**
* 入队一个或多个文件(按文件指纹去重):
* - 同一文件已在队列且 preparing/uploading/ingesting/done → 跳过,不重复入队
* - 同一文件此前失败(error)→ 重新激活原队列项(复用 clientUploadId,保持幂等语义)
*/
/** 入队一个或多个文件 */
const enqueueUploads = useCallback(
(files: File[]) => {
if (!effectiveLibId) {
@@ -263,86 +143,37 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
}
if (valid.length === 0) return
let skipped = 0
let rearmed = 0
const newItems: UploadItem[] = []
for (const file of valid) {
const fileKey = makeFileFingerprint(file)
// error 项允许重新激活;其余状态(preparing/uploading/ingesting/done)都算重复
const dup = findDuplicateInQueue([...itemsRef.current, ...newItems], fileKey, ["error"])
if (dup) {
skipped += 1
continue
}
// 失败项重新激活:复用 tempId/clientUploadId,由 pump 按留存 handle 决定重试方式
const failed = itemsRef.current.find(
(it) => it.fileKey === fileKey && it.status === "error",
)
if (failed) {
rearmed += 1
updateItem(failed.tempId, {
status: "preparing",
progress: 0,
error: undefined,
failedStage: undefined,
hint: undefined,
})
continue
}
newItems.push({
tempId: `${Date.now()}-${newItems.length}-${Math.random().toString(36).slice(2, 8)}`,
file,
fileName: file.name,
progress: 0,
status: "preparing",
fileKey,
clientUploadId: makeClientUploadId(),
})
}
if (newItems.length > 0) {
setItems((prev) => [...prev, ...newItems])
}
if (skipped > 0) {
message.warning(`已跳过 ${skipped} 个重复文件(已在上传队列、处理中或本页已上传)`)
}
if (rearmed > 0) {
message.info(`已重新加入 ${rearmed} 个此前失败的文件`)
}
const newItems: UploadItem[] = valid.map((file, idx) => ({
tempId: `${Date.now()}-${idx}-${Math.random().toString(36).slice(2, 8)}`,
file,
fileName: file.name,
progress: 0,
status: "preparing",
}))
setItems((prev) => [...prev, ...newItems])
},
[effectiveLibId, updateItem],
[effectiveLibId],
)
/**
* 重试失败任务(仅限 status=error):
* - complete 阶段失败:复用留存 handle 只重发 complete(幂等,不重新上传)
* - prepare/transfer 阶段失败:全量重跑(后端尚无记录,安全)
*/
/** 重试失败任务 */
const retryUpload = useCallback(
(tempId: string) => {
const target = itemsRef.current.find((it) => it.tempId === tempId)
if (!target || target.status !== "error") return
updateItem(tempId, { status: "preparing", progress: 0, error: undefined, hint: undefined })
// 状态更新后由 useEffect 触发 pumppump 会按 handlesRef 自动选择 completeOnly / 全量
if (!target) return
updateItem(tempId, { status: "preparing", progress: 0, error: undefined })
// 状态更新后由 useEffect 触发 pump
},
[updateItem],
)
/** 从上传列表移除(已进入转码的由素材网格管理;这里只移除上传面板记录) */
const removeUpload = useCallback((tempId: string) => {
handlesRef.current.delete(tempId)
setItems((prev) => prev.filter((it) => it.tempId !== tempId))
}, [])
/** 清空已完成/去重记录 */
const clearFinished = useCallback(() => {
setItems((prev) => {
for (const it of prev) {
if (it.status === "done") handlesRef.current.delete(it.tempId)
}
return prev.filter((it) => it.status !== "done")
})
setItems((prev) => prev.filter((it) => it.status !== "done"))
}, [])
const activeCount = items.filter(
@@ -357,10 +188,8 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
retryUpload,
removeUpload,
clearFinished,
/** 是否有进行中的上传(用于上传区文案/禁用入口 */
/** 是否有进行中的上传(用于上传区文案) */
uploading: hasActive,
/** 是否有文件正在本地处理或直传(用于禁用上传入口,防重复提交) */
transferActive: activeCount > 0,
activeCount,
pendingCount,
}
@@ -10,26 +10,6 @@ import {
import { getOrCreateDefaultProject } from "@/api/projects"
import { mapLibrary, mapAsset, type AssetItem, type LibraryItem } from "../types"
/** 处理中素材快速轮询(3s)的最大持续时间:超过后停止快轮询,避免孤儿任务永久转圈 */
const PROCESSING_POLL_MAX_MS = 10 * 60 * 1000 // 10 分钟
const isProcessingStatus = (st?: string | null): boolean =>
st === "uploading" || st === "ingesting" || st === "processing" || st === "pending"
/** 判断列表是否存在「创建超过 maxMs 仍在处理中」的卡死素材 */
const hasStalledProcessing = (
list: ApiAssetItem[],
maxMs: number = PROCESSING_POLL_MAX_MS,
): boolean => {
const now = Date.now()
return list.some((a) => {
if (!isProcessingStatus(a.status ?? "")) return false
if (!a.created_at) return false
const created = new Date(a.created_at).getTime()
return Number.isFinite(created) && now - created > maxMs
})
}
/**
* 素材库数据 Hook
* 封装视频库列表、素材列表的数据查询,以及筛选、搜索状态管理
@@ -83,15 +63,15 @@ export function useAssetsData() {
}),
enabled: !!effectiveLibId,
staleTime: 30_000,
// 列表中存在上传中/转码中素材时每 3s 轮询;全部就绪后自动停止
// 但若处理中素材创建已超过 10 分钟仍未就绪(疑似后端卡住/孤儿任务),
// 停止快轮询避免无限转圈——卡死素材在网格中显示「处理超时」提示。
// 列表中存在上传中/转码中素材时每 3s 轮询;全部就绪后自动停止
refetchInterval: (query) => {
const data = query.state.data as { items: ApiAssetItem[] } | undefined
const items = data?.items ?? []
const processing = items.some((a) => isProcessingStatus(a.status ?? ""))
if (!processing) return false
return hasStalledProcessing(items) ? false : 3000
const processing = items.some((a) => {
const st = a.status ?? ""
return st === "uploading" || st === "ingesting" || st === "processing" || st === "pending"
})
return processing ? 3000 : false
},
})
@@ -100,21 +80,6 @@ export function useAssetsData() {
[apiAssets],
)
/** 创建超过 10 分钟仍在处理中的素材(后端可能卡住),网格提示「处理超时」 */
const stalledAssetIds = useMemo(() => {
const list = Array.isArray(apiAssets?.items) ? apiAssets.items : []
const ids = new Set<string>()
const now = Date.now()
for (const a of list) {
if (!isProcessingStatus(a.status ?? "") || !a.created_at) continue
const created = new Date(a.created_at).getTime()
if (Number.isFinite(created) && now - created > PROCESSING_POLL_MAX_MS) {
ids.add(a.id)
}
}
return ids
}, [apiAssets])
/* ── 筛选状态 ── */
const [searchText, setSearchText] = useState("")
const [filterType, setFilterType] = useState<string>("all")
@@ -164,8 +129,6 @@ export function useAssetsData() {
assetsError,
assetsErrorObj,
refetchAssets,
stalledAssetIds,
hasStalledAssets: stalledAssetIds.size > 0,
// 筛选
searchText,
setSearchText,
+5 -15
View File
@@ -1,12 +1,11 @@
/**
* 登录页面 - V21 完全对标
*/
import React, { useRef, useState } from "react"
import React, { useState } from "react"
import { Form, Input, Checkbox, message } from "antd"
import { Link, useNavigate } from "react-router-dom"
import { useLogin } from "@/hooks/useAuth"
import { getWechatAuthUrl } from "@/api/auth"
import { getErrorMessage, isErrorMsgShown } from "@/api/errors"
import Button from "@/components/ui/Button"
import "./Login.css"
@@ -21,9 +20,6 @@ const Login: React.FC = () => {
const loginMutation = useLogin()
const [form] = Form.useForm()
const [wechatLoading, setWechatLoading] = useState(false)
// 同步防连点守卫:state 更新有渲染间隙,连点两次会各自请求授权 URL,
// 后一次的 state 覆盖前一次写入 localStorage 的 state,导致回调校验失败
const wechatStartingRef = useRef(false)
const onFinish = async (values: LoginFormValues) => {
try {
@@ -40,10 +36,8 @@ const Login: React.FC = () => {
}
const handleWechatLogin = async () => {
if (wechatStartingRef.current) return
wechatStartingRef.current = true
setWechatLoading(true)
try {
setWechatLoading(true)
const result = await getWechatAuthUrl()
// 保存 state 到 localStorage 用于回调时验证
localStorage.setItem("wechat_state", result.state)
@@ -57,15 +51,11 @@ const Login: React.FC = () => {
// 跳转到微信授权页
window.location.href = result.auth_url
} catch (error) {
// 跳走前才可能回到这里;拦截器已弹过后端 detail 时不重复弹,
// 否则透传真实原因(如微信服务未配置、网络异常)
if (!isErrorMsgShown(error)) {
message.error(`微信登录启动失败:${getErrorMessage(error, "请稍后重试")}`)
}
wechatStartingRef.current = false
if (!(error as { __msgShown?: boolean })?.__msgShown)
message.error("微信登录暂不可用,请稍后重试")
} finally {
setWechatLoading(false)
}
// 成功时 window.location 跳走,不复位 loading(页面即将卸载)
}
return (
@@ -1,95 +0,0 @@
/**
* 微信绑定回调页(已登录用户在设置页发起"绑定微信"扫码后回到这里)
* 用 code 调绑定接口把微信关联到当前账号,成功后回设置页
*/
import React, { useEffect, useState } from "react"
import { useSearchParams, useNavigate } from "react-router-dom"
import { Spin } from "antd"
import { bindWechat, normalizeUser } from "@/api/auth"
import { getErrorMessage } from "@/api/errors"
import { useAuthStore } from "@/store/authStore"
const WechatBindCallback: React.FC = () => {
const [searchParams] = useSearchParams()
const navigate = useNavigate()
const setUser = useAuthStore((state) => state.setUser)
const [error, setError] = useState<string | null>(null)
useEffect(() => {
const code = searchParams.get("code")
const state = searchParams.get("state")
if (!code || !state) {
setError("无效的回调参数,请回到设置页重新扫码绑定")
return
}
const handleBind = async () => {
// state 校验由后端 state store 一次性消费兜底(前端不再比对 localStorage
// 微信内打开/跨浏览器场景本地无 state 会误杀);清理绑定前写入的 state
localStorage.removeItem("wechat_bind_state")
try {
const result = await bindWechat(code, state)
setUser(normalizeUser(result.user))
// 用 replace 回设置页,query 携带成功标记由设置页提示
navigate("/app/profile?wechat_bind=success", { replace: true })
} catch (err) {
// 绑定失败直接在本页展示真实原因(如微信已被其他账号绑定),不静默跳走
setError(`微信绑定失败:${getErrorMessage(err, "请回到设置页重试")}`)
}
}
handleBind()
}, [searchParams, navigate, setUser])
if (error) {
return (
<div
style={{
display: "flex",
alignItems: "center",
justifyContent: "center",
minHeight: "100vh",
background: "#f5f5f5",
}}
>
<div style={{ textAlign: "center" }}>
<p style={{ color: "#ef4444", fontSize: 16, marginBottom: 16 }}>{error}</p>
<button
onClick={() => navigate("/app/profile")}
style={{
padding: "8px 24px",
background: "var(--primary-color, #3b82f6)",
color: "white",
border: "none",
borderRadius: 6,
cursor: "pointer",
}}
>
</button>
</div>
</div>
)
}
return (
<div
style={{
display: "flex",
alignItems: "center",
justifyContent: "center",
minHeight: "100vh",
background: "#f5f5f5",
}}
>
<div style={{ textAlign: "center" }}>
<Spin size="large" />
<p style={{ marginTop: 16, color: "#666" }}>...</p>
</div>
</div>
)
}
export default WechatBindCallback
+74 -59
View File
@@ -1,48 +1,40 @@
/**
* 微信登录回调页
* 扫码授权后由微信重定向回来:用 code 换登录态,
* 新用户/资料未完善 → 跳昵称引导页;老用户 → 回来源页/首页
*/
import React, { useEffect, useState } from "react"
import { useSearchParams, useNavigate } from "react-router-dom"
import { Spin } from "antd"
import { Spin, message } from "antd"
import { wechatCallback, getCurrentUser, normalizeUser, type User } from "@/api/auth"
import { getErrorMessage } from "@/api/errors"
import { useAuthStore } from "@/store/authStore"
import { scheduleProactiveRefresh } from "@/api/auth/tokenRefresh"
import BindContactModal from "@/components/auth/BindContactModal"
const WechatCallback: React.FC = () => {
const [searchParams] = useSearchParams()
const navigate = useNavigate()
const setAuth = useAuthStore((state) => state.setAuth)
const [loading, setLoading] = useState(true)
const [showBindModal, setShowBindModal] = useState(false)
const [error, setError] = useState<string | null>(null)
useEffect(() => {
const code = searchParams.get("code")
const state = searchParams.get("state")
// 微信重定向出错时(如用户拒绝授权 error=access_denied)直接展示原因
const wxErrorCode = searchParams.get("error")
const wxErrDesc = searchParams.get("error_description")
if (wxErrorCode || wxErrDesc) {
const reason = [wxErrorCode, wxErrDesc].filter(Boolean).join("")
setError(`微信授权失败:${reason}`)
setLoading(false)
return
}
if (!code || !state) {
setError("无效的回调参数,请重新扫码登录")
setError("无效的回调参数")
setLoading(false)
return
}
const handleCallback = async () => {
try {
// state CSRF 校验由后端 state store 一次性消费兜底(前端不再比对
// localStorage——微信内打开、跨浏览器等场景本地没有 state,会误杀正常回调);
// 清理登录前写入的 state,避免残留
// 校验 state,防止 CSRF
const savedState = localStorage.getItem("wechat_state")
if (!savedState || savedState !== state) {
setError("安全校验失败,请重新登录")
setLoading(false)
return
}
localStorage.removeItem("wechat_state")
const result = await wechatCallback(code, state)
@@ -57,22 +49,20 @@ const WechatCallback: React.FC = () => {
const userData = await getCurrentUser()
const user: User = normalizeUser(userData)
setAuth(user, result.access_token, result.refresh_token)
scheduleProactiveRefresh()
// 新用户 或 资料未完善(如上次中断没填昵称)→ 强制昵称引导
const needOnboarding = result.is_new_user || user.profile_completed === false
if (needOnboarding) {
navigate("/welcome/wechat", { replace: true })
return
if (result.binding_complete) {
// 已绑定,跳转到登录前页面或首页
message.success("登录成功")
const redirect = localStorage.getItem("login_redirect") || "/"
localStorage.removeItem("login_redirect")
navigate(redirect, { replace: true })
} else {
// 未绑定,显示绑定弹窗
setLoading(false)
setShowBindModal(true)
}
// 老用户:回登录前页面或首页
const redirect = localStorage.getItem("login_redirect") || "/"
localStorage.removeItem("login_redirect")
navigate(redirect, { replace: true })
} catch (err) {
// 透传后端真实错误(如 state 过期、code 已消费、接口异常),禁止吞成通用提示
setError(`微信登录失败:${getErrorMessage(err, "请重试或更换登录方式")}`)
setError("登录失败,请重试")
setLoading(false)
}
}
@@ -80,6 +70,21 @@ const WechatCallback: React.FC = () => {
handleCallback()
}, [searchParams, navigate, setAuth])
const handleBindSuccess = (user: User) => {
const setUser = useAuthStore.getState().setUser
setUser(user)
setShowBindModal(false)
message.success("绑定成功")
const redirect = localStorage.getItem("login_redirect") || "/"
localStorage.removeItem("login_redirect")
navigate(redirect, { replace: true })
}
const handleBindCancel = () => {
setShowBindModal(false)
navigate("/login")
}
if (loading) {
return (
<div
@@ -93,39 +98,49 @@ const WechatCallback: React.FC = () => {
>
<div style={{ textAlign: "center" }}>
<Spin size="large" />
<p style={{ marginTop: 16, color: "#666" }}>...</p>
<p style={{ marginTop: 16, color: "#666" }}>...</p>
</div>
</div>
)
}
if (error) {
return (
<div
style={{
display: "flex",
alignItems: "center",
justifyContent: "center",
minHeight: "100vh",
background: "#f5f5f5",
}}
>
<div style={{ textAlign: "center" }}>
<p style={{ color: "#ef4444", fontSize: 16, marginBottom: 16 }}>{error}</p>
<button
onClick={() => navigate("/login")}
style={{
padding: "8px 24px",
background: "var(--primary-color, #3b82f6)",
color: "white",
border: "none",
borderRadius: 6,
cursor: "pointer",
}}
>
</button>
</div>
</div>
)
}
return (
<div
style={{
display: "flex",
alignItems: "center",
justifyContent: "center",
minHeight: "100vh",
background: "#f5f5f5",
}}
>
<div style={{ textAlign: "center" }}>
<p style={{ color: "#ef4444", fontSize: 16, marginBottom: 16 }}>{error}</p>
<button
onClick={() => navigate("/login")}
style={{
padding: "8px 24px",
background: "var(--primary-color, #3b82f6)",
color: "white",
border: "none",
borderRadius: 6,
cursor: "pointer",
}}
>
</button>
</div>
</div>
<BindContactModal
open={showBindModal}
onSuccess={handleBindSuccess}
onCancel={handleBindCancel}
/>
)
}
@@ -1,103 +0,0 @@
/**
* 微信新用户昵称引导页
* 新微信用户首次登录后强制填写昵称,完成后才进入主界面
*/
import React from "react"
import { Form, Input, message } from "antd"
import { Navigate, useNavigate } from "react-router-dom"
import { useMutation } from "@tanstack/react-query"
import { updateProfile } from "@/api/auth"
import { useAuthStore } from "@/store/authStore"
import Button from "@/components/ui/Button"
import "./Login.css"
interface OnboardingFormValues {
display_name: string
}
const WechatOnboarding: React.FC = () => {
const navigate = useNavigate()
const setUser = useAuthStore((state) => state.setUser)
const isAuthenticated = useAuthStore((state) => state.isAuthenticated)
const user = useAuthStore((state) => state.user)
const hasAccessToken = Boolean(localStorage.getItem("access_token"))
const [form] = Form.useForm<OnboardingFormValues>()
const saveMutation = useMutation({
mutationFn: (displayName: string) => updateProfile({ display_name: displayName }),
})
// 已登录且资料已完善的用户不该停留在引导页
if (isAuthenticated && hasAccessToken && user?.profile_completed === true) {
return <Navigate to="/app/dashboard" replace />
}
// 未登录(如手动输入 URL)回登录页
if (!isAuthenticated || !hasAccessToken) {
return <Navigate to="/login" replace />
}
const onFinish = async (values: OnboardingFormValues) => {
try {
const updated = await saveMutation.mutateAsync(values.display_name.trim())
// 后端返回的 profile_completed 以最新资料为准,前端同步标记完善
setUser({ ...updated, profile_completed: true })
message.success("欢迎加入小虾智剪!")
const redirect = localStorage.getItem("login_redirect") || "/app/dashboard"
localStorage.removeItem("login_redirect")
navigate(redirect, { replace: true })
} catch {
message.error("保存失败,请重试")
}
}
return (
<div className="xx-auth-page">
<div className="xx-auth-card">
<div className="xx-auth-header">
<div className="xx-auth-brand">
<span className="xx-auth-logo">🦐</span>
<span className="xx-auth-brand-name"></span>
</div>
<p>使</p>
</div>
<Form
form={form}
name="wechat-onboarding"
onFinish={onFinish}
autoComplete="off"
layout="vertical"
// 不预填:新微信用户必须自己输入昵称(user.display_name 可能是微信昵称/系统占位)
initialValues={{ display_name: "" }}
>
<Form.Item
name="display_name"
label="昵称"
rules={[
{ required: true, message: "请输入昵称" },
{ whitespace: true, message: "昵称不能为空白" },
{ min: 1, max: 20, message: "昵称长度需在 1-20 个字符之间" },
]}
extra="昵称将展示在您的作品和账户中,之后可在个人设置中修改"
>
<Input placeholder="请输入您的昵称" size="large" maxLength={20} showCount />
</Form.Item>
<Form.Item>
<Button
buttonType="primary"
buttonSize="lg"
htmlType="submit"
loading={saveMutation.isPending}
style={{ width: "100%" }}
>
{saveMutation.isPending ? "保存中..." : "进入小虾智剪"}
</Button>
</Form.Item>
</Form>
</div>
</div>
)
}
export default WechatOnboarding
+1 -1
View File
@@ -133,7 +133,7 @@ const GeneratePage: React.FC = () => {
const [previewVoiceAudioUrl, setPreviewVoiceAudioUrl] = useState<string | null>(null)
const ttsAbortRef = useRef<AbortController | null>(null)
// TTS 试听文案:批量跟随变体0标题(仅取首项,避免编辑其他变体标题触发多余 TTS 请求)
const variant0Title = isBatch ? previewTitles?.[0] || "" : ""
const variant0Title = isBatch ? previewTitles[0] || "" : ""
useEffect(() => {
const voiceAsset = voiceMaterials.find((m) => m.id === selectedVoice)
@@ -22,9 +22,7 @@ const BatchGenerationGrid: React.FC<BatchGenerationGridProps> = ({
titles,
onRetryTask,
}) => {
const sorted = [...tasks].sort(
(a, b) => Number(a.variantIndex || 0) - Number(b.variantIndex || 0),
)
const sorted = [...tasks].sort((a, b) => (a.variantIndex || 0) - (b.variantIndex || 0))
return (
<div className="xx-form-section">
@@ -12,6 +12,7 @@ import type { AssetItem } from "@/api/assets"
import type { EditingTemplate } from "@/api/editing-planner"
import type { TitleSettings } from "../types"
import FrontendPreviewPlayer from "./FrontendPreviewPlayer"
import { MAX_PREVIEW_COUNT } from "../constants"
interface CanvasPreviewGridProps {
count: number
@@ -41,11 +42,11 @@ const CanvasPreviewGrid: React.FC<CanvasPreviewGridProps> = ({
onToggleSelect,
selectable = true,
}) => {
// count 上限已在源头 PreviewCountModal 的数量选择(1~MAX_PREVIEW_COUNT=10clamp
// 这里完整渲染所有变体,保证每个变体都有勾选/预览入口,UI 与数据不脱节
// 硬上限保护:同时播放的媒体元素数量不超过 MAX_PREVIEW_COUNT10),避免浏览器卡顿
const safeCount = Math.max(1, Math.min(count, MAX_PREVIEW_COUNT))
return (
<div className="xx-canvas-grid">
{Array.from({ length: count }, (_, i) => {
{Array.from({ length: safeCount }, (_, i) => {
const checked = selectedIds.includes(i)
return (
<div
@@ -178,35 +178,3 @@
border-color: var(--border-color);
margin: var(--space-lg) 0;
}
/* 微信账号绑定卡片 */
.xx-settings-wechat {
display: flex;
align-items: center;
justify-content: space-between;
gap: var(--space-lg);
flex-wrap: wrap;
}
.xx-settings-wechat-info {
display: flex;
align-items: center;
gap: var(--space-md);
}
.xx-settings-wechat-info .xx-wechat-icon {
font-size: 28px;
line-height: 1;
}
.xx-settings-wechat-info strong {
display: block;
color: var(--text-primary);
font-size: var(--font-size-md);
}
.xx-settings-wechat-info p {
margin: 2px 0 0;
color: var(--text-secondary);
font-size: var(--font-size-sm);
}
+20 -143
View File
@@ -1,111 +1,39 @@
/**
* 个人设置页面
* - 个人资料(昵称)保存
* - 微信账号绑定状态 / 绑定 / 解绑
* P1-2: 添加 PageHead
* P1-3: antd Form/Input/Button/Alert → 自定义 UI 组件
*/
import React, { useEffect, useRef, useState } from "react"
import { useSearchParams } from "react-router-dom"
import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query"
import { message } from "antd"
import React, { useState } from "react"
import { Button, Input, Modal } from "@/components/ui"
import { getCurrentUser, updateProfile, getWechatBindUrl, unbindWechat } from "@/api/auth"
import { useAuthStore } from "@/store/authStore"
import PageHead from "@/components/layout/PageHead"
import "./ProfileSettings.css"
const Settings: React.FC = () => {
const user = useAuthStore((state) => state.user)
const setUser = useAuthStore((state) => state.setUser)
const queryClient = useQueryClient()
const [searchParams, setSearchParams] = useSearchParams()
const [displayName, setDisplayName] = useState(user?.display_name || "")
const bindTipShownRef = useRef(false)
// 拉取最新用户信息(微信绑定状态以后端为准)
const { data: freshUser } = useQuery({
queryKey: ["currentUser"],
queryFn: getCurrentUser,
})
useEffect(() => {
if (freshUser) {
setUser(freshUser)
setDisplayName((prev) => prev || freshUser.display_name || "")
}
}, [freshUser, setUser])
// 绑定回调结果提示(?wechat_bind=success|failed
useEffect(() => {
if (bindTipShownRef.current) return
const result = searchParams.get("wechat_bind")
if (!result) return
bindTipShownRef.current = true
if (result === "success") {
message.success("微信绑定成功")
} else if (result === "failed") {
message.error("微信绑定失败,请重试")
}
searchParams.delete("wechat_bind")
setSearchParams(searchParams, { replace: true })
}, [searchParams, setSearchParams])
const wechatBound = user?.wechat_bound === true
const saveProfileMutation = useMutation({
mutationFn: () => updateProfile({ display_name: displayName.trim() }),
onSuccess: (updated) => {
setUser(updated)
message.success("资料已保存")
},
onError: () => {
message.error("保存失败,请重试")
},
})
const handleBindWechat = async () => {
try {
const result = await getWechatBindUrl()
localStorage.setItem("wechat_bind_state", result.state)
window.location.href = result.auth_url
} catch {
message.error("微信绑定暂不可用,请稍后重试")
}
}
const unbindMutation = useMutation({
mutationFn: unbindWechat,
onSuccess: () => {
message.success("已解绑微信")
queryClient.invalidateQueries({ queryKey: ["currentUser"] })
// 本地立即更新,避免等待刷新
if (user) {
setUser({ ...user, wechat_bound: false, wechat_nickname: "" })
}
},
onError: () => {
message.error("解绑失败,请重试")
},
})
const handleUnbind = () => {
Modal.confirm({
title: "解绑微信",
content: "解绑后将无法使用微信登录该账号,确定要解绑吗?",
okText: "确定解绑",
cancelText: "取消",
okButtonProps: { danger: true },
onOk: () => unbindMutation.mutateAsync(),
const handleSave = () => {
Modal.info({
title: "提示",
content: "个人资料修改接口暂未开放,保存功能即将上线。",
})
}
const displayNameDirty = displayName.trim() !== (user?.display_name || "")
return (
<div className="xx-settings-page">
<PageHead title="个人设置" description="管理您的账户信息" />
<div className="xx-settings-card">
<h3></h3>
<div className="xx-settings-notice">
<span className="xx-settings-notice-icon"></span>
<div>
<strong></strong>
<p></p>
</div>
</div>
<div className="xx-settings-form">
<div className="xx-settings-field">
<label className="xx-settings-label"></label>
@@ -114,76 +42,25 @@ const Settings: React.FC = () => {
<div className="xx-settings-field">
<label className="xx-settings-label"></label>
<Input
value={user?.email && !user.email.endsWith("@wechat.local") ? user.email : ""}
disabled
placeholder={user?.email?.endsWith("@wechat.local") ? "微信账号暂未绑定邮箱" : "邮箱"}
/>
<Input value={user?.email || ""} disabled placeholder="邮箱" />
</div>
<div className="xx-settings-field">
<label className="xx-settings-label"></label>
<label className="xx-settings-label"></label>
<Input
value={displayName}
onChange={(e) => setDisplayName(e.target.value)}
placeholder="请输入称"
maxLength={20}
placeholder="请输入显示名称"
/>
</div>
<div className="xx-settings-field">
<Button
buttonType="primary"
buttonSize="md"
onClick={() => saveProfileMutation.mutate()}
loading={saveProfileMutation.isPending}
disabled={!displayName.trim() || !displayNameDirty}
>
<Button buttonType="primary" buttonSize="md" onClick={handleSave} disabled>
</Button>
</div>
</div>
</div>
<div className="xx-settings-card">
<h3></h3>
<div className="xx-settings-wechat">
<div className="xx-settings-wechat-info">
<span className="xx-wechat-icon">💬</span>
<div>
{wechatBound ? (
<>
<strong>
{user?.wechat_nickname ? `${user.wechat_nickname}` : ""}
</strong>
<p>使</p>
</>
) : (
<>
<strong></strong>
<p>使</p>
</>
)}
</div>
</div>
<div className="xx-settings-wechat-actions">
{wechatBound ? (
<Button
buttonType="ghost"
buttonSize="md"
onClick={handleUnbind}
loading={unbindMutation.isPending}
>
</Button>
) : (
<Button buttonType="primary" buttonSize="md" onClick={handleBindWechat}>
</Button>
)}
</div>
</div>
</div>
</div>
)
}
-6
View File
@@ -6,16 +6,10 @@ 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 <Navigate to="/login" replace />
}
// 微信新用户未完成昵称引导时,禁止进入主界面
if (!profileCompleted) {
return <Navigate to="/welcome/wechat" replace />
}
return <>{children}</>
}
-10
View File
@@ -5,8 +5,6 @@ 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,未登录显示落地页 */
@@ -47,12 +45,4 @@ export const publicRoutes: RouteObject[] = [
path: "/auth/wechat/callback",
element: <WechatCallback />,
},
{
path: "/auth/wechat/bind/callback",
element: <WechatBindCallback />,
},
{
path: "/welcome/wechat",
element: <WechatOnboarding />,
},
]
-6
View File
@@ -13,12 +13,6 @@ 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 {
-130
View File
@@ -1,130 +0,0 @@
/**
* 上传去重/幂等工具单测(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 的 Fileslice 返回小 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()
})
})
+33 -124
View File
@@ -1,8 +1,8 @@
import { describe, expect, it, vi, beforeEach } from "vitest"
import { render, screen, fireEvent, waitFor } from "@testing-library/react"
import { describe, expect, it, vi } from "vitest"
import { render, screen } 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 }) => (
<div data-testid="page-head">
@@ -12,142 +12,51 @@ vi.mock("@/components/layout/PageHead", () => ({
),
}))
const mockSetUser = vi.fn()
const mockInvalidate = vi.fn()
let authState: Record<string, unknown> = {
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: unknown) => unknown) => selector(authState),
}))
const getCurrentUserMock = vi.fn(async () => authState.user as Record<string, unknown>)
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(
<QueryClientProvider client={queryClient}>
<MemoryRouter>
<Settings />
</MemoryRouter>
</QueryClientProvider>,
)
describe("Settings Page", () => {
beforeEach(() => {
vi.clearAllMocks()
authState = {
useAuthStore: (selector: (state: any) => any) =>
selector({
user: {
id: "1",
user_id: "1",
username: "testuser",
email: "test@example.com",
display_name: "Test User",
wechat_bound: false,
is_email_verified: true,
email_verified: true,
},
isAuthenticated: true,
setUser: mockSetUser,
}
}),
}))
import Settings from "@/pages/profile/Settings"
describe("Settings Page", () => {
it("should render without crashing", () => {
render(
<MemoryRouter>
<Settings />
</MemoryRouter>,
)
expect(screen.getByText("个人设置")).toBeTruthy()
})
it("渲染个人设置与用户信息", () => {
renderPage()
expect(screen.getByText("个人设置")).toBeTruthy()
it("should display user info", () => {
render(
<MemoryRouter>
<Settings />
</MemoryRouter>,
)
expect(screen.getByDisplayValue("testuser")).toBeTruthy()
expect(screen.getByDisplayValue("test@example.com")).toBeTruthy()
})
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, "") === "解绑",
),
it("should show save button is disabled", () => {
render(
<MemoryRouter>
<Settings />
</MemoryRouter>,
)
// 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: "新昵称" })
})
const button = screen.getByText("保存暂未开放")
expect(button).toBeTruthy()
})
})
@@ -31,23 +31,17 @@ interface FakeHandle {
complete: ReturnType<typeof vi.fn>
/** 手动结束传输(transfer 被调用后挂载);finish(true) 以失败结束 */
finish: (fail?: boolean) => void
/** complete 已被调用的次数 */
completeCalls: { resolve: () => void; reject: (err: unknown) => void }[]
}
let activeTransfers = 0
let maxConcurrent = 0
/**
* 创建一个假 handle
* - transfer 返回挂起的 promisefinish()/finish(true) 控制成败
* - complete 每次调用返回独立的挂起 promise,由 completeCalls 记录控制,
* 成功调 resolve(idx) / 失败调 reject(idx)(模拟超时)
* 创建一个假 handletransfer 返回挂起的 promise
* finish 槽位在 transfer executor 同步执行时挂载,测试中调用 finish() 控制成败
*/
const makeFakeHandle = (opts: {
id: string
duplicated?: boolean
failTransfer?: boolean
completeAuto?: boolean
}) => {
const h: FakeHandle = {
const makeFakeHandle = (opts: { id: string; duplicated?: boolean; failTransfer?: boolean }) => {
const h = {
prepared: {
upload_url: "https://oss.example.com/u",
method: "POST",
@@ -58,39 +52,15 @@ const makeFakeHandle = (opts: {
asset_id: opts.id,
},
transfer: vi.fn(),
complete: vi.fn(),
finish: () => {},
completeCalls: [],
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,
}
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<void>((_resolve, reject) => {
@@ -108,24 +78,17 @@ const makeFakeHandle = (opts: {
type FakeHandleLike = ReturnType<typeof makeFakeHandle>
let activeTransfers = 0
let maxConcurrent = 0
/** prepare mock:调用序号生成稳定 id,立即把 handle(含 finish 槽位)推入数组 */
const installPrepareMock = (
handles: FakeHandleLike[],
optOverrides?: (id: string) => {
duplicated?: boolean
failTransfer?: boolean
completeAuto?: boolean
},
optOverrides?: (id: string) => { duplicated?: boolean; failTransfer?: boolean },
) => {
let callNo = 0
;(prepareDirectUploadHandle as unknown as ReturnType<typeof vi.fn>).mockImplementation(
async () => {
const id = `asset-${callNo++}`
const overrides = optOverrides?.(id) ?? {}
const h = makeFakeHandle({ id, completeAuto: true, ...overrides })
const h = makeFakeHandle({ id, ...overrides })
handles.push(h)
await new Promise((r) => setTimeout(r, 10))
return h
@@ -224,11 +187,6 @@ 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")
// 重试:重新 preparehandles[1] 成功)
const tempId = result.current.uploadItems[0].tempId
await act(async () => {
@@ -265,119 +223,4 @@ 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,只再调一次 completetransfer/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<typeof vi.fn>).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("签名服务内部错误")
})
})
@@ -1,105 +0,0 @@
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<string, string> = {}
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 }),
}))
const renderPage = () =>
render(
<MemoryRouter>
<WechatBindCallback />
</MemoryRouter>,
)
describe("WechatBindCallback Page", () => {
afterEach(() => {
cleanup()
})
beforeEach(() => {
vi.clearAllMocks()
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()
})
})
})
@@ -1,162 +1,79 @@
import { describe, expect, it, vi, beforeEach, afterEach } from "vitest"
import { render, screen, waitFor, cleanup } from "@testing-library/react"
import { describe, expect, it, vi, beforeEach } from "vitest"
import { render, screen } 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<string, string> = {}
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<string, unknown> = {}
let mockCurrentUser: Record<string, unknown> = {}
let callbackError: unknown = null
vi.mock("react-router-dom", async () => {
const actual = await vi.importActual("react-router-dom")
return {
...actual,
useNavigate: () => mockNavigate,
useSearchParams: () => mockSearchParams,
useNavigate: () => vi.fn(),
useSearchParams: () => [new URLSearchParams({ code: "test_code", state: "test_state" })],
}
})
vi.mock("@/api/auth", () => ({
wechatCallback: vi.fn(async () => {
if (callbackError) throw callbackError
return mockCallbackResult
}),
getCurrentUser: vi.fn(async () => mockCurrentUser),
wechatCallback: vi.fn(() => new Promise(() => {})), // pending promise,保持loading
getCurrentUser: vi.fn(),
normalizeUser: (u: unknown) => u,
}))
vi.mock("@/api/auth/tokenRefresh", () => ({
scheduleProactiveRefresh: vi.fn(),
cancelProactiveRefresh: vi.fn(),
}))
vi.mock("@/store/authStore", () => ({
useAuthStore: (selector: (state: unknown) => unknown) => selector({ setAuth: mockSetAuth }),
useAuthStore: () => ({
setAuth: vi.fn(),
}),
}))
const renderPage = () =>
render(
<MemoryRouter>
<WechatCallback />
</MemoryRouter>,
)
vi.mock("@/components/auth/BindContactModal", () => ({
default: ({ open }: { open: boolean }) => (
<div data-testid="bind-contact-modal" style={{ display: open ? "block" : "none" }}>
BindContactModal
</div>
),
}))
vi.mock("antd", async () => {
const actual = await vi.importActual("antd")
return {
...actual,
message: {
success: vi.fn(),
error: vi.fn(),
},
}
})
describe("WechatCallback Page", () => {
afterEach(() => {
cleanup()
})
beforeEach(() => {
vi.clearAllMocks()
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,
// mock localStorage,设置wechat_state匹配,让校验通过
const store: Record<string, string> = {
wechat_state: "test_state",
}
mockCurrentUser = {
id: "u1",
display_name: "老用户",
profile_completed: true,
}
})
it("老用户登录成功跳转首页/来源页", async () => {
renderPage()
await waitFor(() => {
expect(mockNavigate).toHaveBeenCalledWith("/", { replace: true })
vi.spyOn(Storage.prototype, "getItem").mockImplementation((key) => store[key] || null)
vi.spyOn(Storage.prototype, "setItem").mockImplementation((key, val) => {
store[key] = val
})
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 })
vi.spyOn(Storage.prototype, "removeItem").mockImplementation((key) => {
delete store[key]
})
})
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 render without crashing", () => {
const { container } = render(
<MemoryRouter>
<WechatCallback />
</MemoryRouter>,
)
expect(container).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()
it("should show loading state while processing", () => {
render(
<MemoryRouter>
<WechatCallback />
</MemoryRouter>,
)
// wechatCallback 返回 pending promise,所以应该显示 loading
expect(screen.getByText("正在登录...")).toBeTruthy()
})
})
@@ -1,138 +0,0 @@
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<string, unknown> = {}
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(
<QueryClientProvider client={queryClient}>
<MemoryRouter>
<WechatOnboarding />
</MemoryRouter>
</QueryClientProvider>,
)
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 () => {
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()
})
})
})
+15 -22
View File
@@ -37,19 +37,14 @@ LONG_VIDEO_SEGMENT_SEC = 30 # 长视频每段秒数
LONG_VIDEO_DURATION_THRESHOLD_SEC = 180 # 3 分钟阈值
MIN_FRAMES_PER_SEGMENT = 2 # 长视频每段最少帧数
# ── 滑动窗口匹配常量(Issue #1702 二次校准) ─────────────────────
# 阈值经 staging 真实数据两轮回归校准(worker 容器内离线实验):
# 第一轮(2026-09-05):同源对 <=12 命中 4/11,异源最小距离 24 → 定 12;
# 第二轮(2026-09-05,证据视频 B->A 仍漏检):扩大样本到该用户全部
# 15 个真实成片(13 个异源候选)实测:
# - 同源成片对(A 20s / B、C 各 11.75s1s 密集采样):
# B->A 中位数距离 14<=16 命中 8/11=0.73C->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
# ── 滑动窗口匹配常量(Issue #1702 重新校准) ─────────────────────
# 阈值经 staging 真实数据回归校准(2026-09-05worker 容器内离线实验):
# - 同源成片对(20s/11s,各自 2-5% 随机边缘裁剪降重,1s 密集采样):
# 全部帧对最小汉明距离 min=8<=12 命中 10/31 帧(B->A 4/11
# - 异源成片对(4 个不同项目真实视频):最小距离 24,<=16 命中 0 帧
# 8(#1658 旧值)会漏掉同源裁剪(自对照实验:同帧两次 2-5% 随机裁剪距离 4~10),
# 12 能检出同源/局部复用且与异源分布(>=24)间隔 12bit,无误报空间。
PHASH_THRESHOLD = 12
SEGMENT_MATCH_THRESHOLD = PHASH_THRESHOLD # 片段匹配阈值与帧匹配统一(#1702:阈值常量统一来源)
MIN_CONSECUTIVE_MATCHES = 5 # 连续匹配默认门槛;短视频自适应 min(5, max(2, 分片数//2))
MAX_GAP = 2 # 允许的最大间隙帧数
@@ -377,10 +372,9 @@ def find_duplicate_segments(
1. 构建 query×target 全量汉明距离矩阵;每个 query chunk 保留所有
距离 <= match_threshold 的候选 target 分片(与帧匹配判定同一阈值)。
2. 时序一致贪心对齐:沿 query 时序推进,run 内优先选择与上一匹配帧
目标序号连贯(|delta| <= neighbor_window+1,允许 ±1 邻接/时序偏移
对齐——1s 密集采样下相邻帧 pHash 接近,最近邻在目标相邻帧间
正/反向跳变均属正常,缓解场景切割切点、取帧错位、局部倒退)的
候选;同距时偏好小索引(最早对齐位置)。
目标序号连贯(0 <= delta <= neighbor_window+1,允许 ±1 邻接窗口 /
时序偏移对齐,缓解场景切割导致的切点、取帧错位)的候选;同距时
偏好大索引,避免重复 hash 塌缩到 target 首帧。
3. 连贯匹配中允许 <= max_gap 帧间隙桥接;断裂后另起新 run——天然
支持局部片段复用(复用片段可出现在任意时序位置,各成独立片段)。
4. 连续匹配帧数 >= min_consecutive 的 run 报为重复片段。短视频自适应:
@@ -393,7 +387,7 @@ def find_duplicate_segments(
match_threshold: 汉明距离匹配阈值(统一常量 PHASH_THRESHOLD
min_consecutive: 最少连续匹配帧数;None 时按短视频自适应
max_gap: 允许的最大间隙帧数
neighbor_window: 时序对齐允许的目标分片序号邻接窗口(正/反向均允许)
neighbor_window: 时序对齐允许的目标分片序号邻接窗口
Returns:
DuplicateSegment 列表
@@ -428,9 +422,8 @@ def find_duplicate_segments(
min_consecutive = min(MIN_CONSECUTIVE_MATCHES, max(2, n // 2))
# Step 2: 时序一致贪心对齐。
# run 内偏好与上一匹配帧目标序号连贯(|delta| <= neighbor_window+1
# 支持 ±1 邻接窗口/时序偏移对齐,正反向抖动均允许)的候选;
# 无连贯候选时关闭旧 run。
# run 内偏好与上一匹配帧目标序号连贯(0 <= delta <= neighbor_window+1
# 支持 ±1 邻接窗口/时序偏移对齐)的候选;无连贯候选时关闭旧 run。
# 这天然支持局部片段复用:同一 query 视频中多个复用片段各自形成独立 run。
frame_matches: list[tuple[bool, int, int]] = []
runs: list[tuple[int, int]] = []
@@ -451,7 +444,7 @@ def find_duplicate_segments(
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),
(c for c in cand if 0 <= c[0] - run_last_t <= neighbor_window + 1),
None,
)
+2 -22
View File
@@ -6,18 +6,6 @@ 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 崩溃时未完成任务重回队列,由执行前守卫丢弃作废消息
celery_app.conf.imports = (
"worker_app.tasks.health",
"worker_app.tasks.ingest",
@@ -34,18 +22,10 @@ celery_app.conf.imports = (
)
# Celery Beat 定时任务调度
# 注:worker 单实例内嵌 beatentrypoint-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},
"schedule": 600.0, # 每 10 分钟(秒)
"options": {"expires": 300}, # 5 分钟过期,避免堆积
},
}
+8 -117
View File
@@ -7,83 +7,11 @@ from worker_app.db import SessionLocal
logger = logging.getLogger(__name__)
# 孤儿任务超时阈值:渲染任务超过此时间未更新则视为卡死
ORPHAN_TASK_TIMEOUT_MINUTES = 10
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
# Pending 任务超时阈值:pending 任务在队列中等待超过此时间则自动清理
PENDING_TASK_TIMEOUT_MINUTES = 30
def cleanup_orphan_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> int: # pragma: no cover
@@ -104,16 +32,11 @@ def cleanup_orphan_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) ->
try:
session = SessionLocal()
try:
repo = SQLAlchemyGenerationTaskRepository(session)
items = cleanup_stale_running_with_session_ids(repo, timeout_minutes)
finally:
session.close()
count = len(items)
repo = SQLAlchemyGenerationTaskRepository(session)
count = repo.cleanup_stale_running(timeout_minutes)
session.close()
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
@@ -147,9 +70,7 @@ 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
@@ -159,8 +80,6 @@ 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)
@@ -186,12 +105,9 @@ def cleanup_stale_pending_tasks(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINU
session = SessionLocal()
try:
repo = SQLAlchemyGenerationTaskRepository(session)
items = cleanup_stale_pending_with_session_ids(repo, timeout_minutes)
count = len(items)
count = repo.cleanup_stale_pending(timeout_minutes)
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
@@ -202,31 +118,6 @@ 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
"""统一清理所有超时的孤儿任务。
+4 -40
View File
@@ -1,18 +1,14 @@
"""定期清理任务 — Celery Beat 调度。
包含:
- cleanup_stale_pending_tasks: 定期清理卡在 pending 超时的 generation_tasksworker 停止消费时占位)
- cleanup_stale_running_tasks: 定期清理卡在 running 超时的 generation_tasks(容器重启/进程被杀后的孤儿)
- cleanup_stale_pending_tasks: 定期清理卡在 pending 超时的 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,
)
@@ -23,14 +19,12 @@ logger = logging.getLogger(__name__)
def scheduled_cleanup_stale_pending(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINUTES) -> dict:
"""Celery Beat 调度的定期任务:清理超时的 pending 任务。
5 分钟执行一次(由 celery_app.py 的 beat_schedule 配置),
10 分钟执行一次(由 celery_app.py 的 beat_schedule 配置),
查找所有 status='pending' 且 created_at < NOW() - timeout_minutes
的 generation_tasks,批量更新为 failed,释放限流名额;同时 revoke 并清除
Redis 队列中对应的 Celery 消息,杜绝作废消息重投执行(#1714)。
的 generation_tasks,批量更新为 failed
Args:
timeout_minutes: 超时时间(分钟),默认 45 分钟pending 排队阈值放宽,
与 running 孤儿 20 分钟区分,避免正常排队任务被误杀)
timeout_minutes: 超时时间(分钟),默认 30 分钟
Returns:
{"cleaned": int}
@@ -39,33 +33,3 @@ 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}
+2 -30
View File
@@ -22,8 +22,6 @@ 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
@@ -660,34 +658,6 @@ 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(
@@ -699,6 +669,8 @@ 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:
+27 -138
View File
@@ -359,54 +359,6 @@ def validate_transcode_output(
return True
def _original_key_from_storage_key(storage_key: str) -> str:
"""从可能被 HEVC 转码改写的 storage_key 还原原始 key。
转码成功后 key 形如 uploads/<id>/IMG_2282_h264.MOV
占位 asset 以原始 key uploads/<id>/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 占位 assetIssue #1714)。
HEVC 转码成功后 job.storage_key 会被改写为 *_h264 新 key,旧实现用新 key
回查占位必然落空,进而兜底新建一条 READY 记录,导致原占位永久卡 processing。
查找优先级:
1. job.asset_idcomplete 派单时透传的占位 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:
"""
@@ -428,23 +380,6 @@ 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_keyHEVC 转码成功后 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)
@@ -689,41 +624,22 @@ def ingest_asset(job_id: str) -> dict:
error_reason,
)
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)
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
@@ -740,19 +656,16 @@ def ingest_asset(job_id: str) -> dict:
"error": error_reason,
}
# 查找 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)
# 查找已存在的 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")
if existing_asset is None:
# 兜底:仅当确实没有占位记录(旧版本 API / 历史 job 重跑)才新建。
logger.warning(
"No placeholder asset found for job_id=%s original_key=%s, creating new",
job_id,
original_storage_key,
)
# 兜底:如果 API 端没有预先创建 Asset(旧版本兼容),则创建新记录
logger.info("No pre-created asset found for storage_key=%s, creating new", job.storage_key)
metadata["source"] = "upload"
asset = Asset.create(
project_id=job.project_id,
@@ -772,13 +685,8 @@ def ingest_asset(job_id: str) -> dict:
)
asset_repo.create(asset)
else:
# 更新占位记录补充元数据、置 READY。转码成功时 storage_key 同步改写为
# *_h264(播放/下载走转码产物),原始 key 记入 metadata 可溯源。
# 更新已有的 Asset 记录补充元数据并将状态改为 READY
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
@@ -829,28 +737,9 @@ def ingest_asset(job_id: str) -> dict:
job_repo.update(job)
# 将上传时创建的占位 AssetPROCESSING/UPLOADING)标记为 ERROR
# 避免素材永远卡在中间状态。转码可能已把 job.storage_key 改写为
# *_h264,需用 asset_id / 原始 key 多路径关联占位(Issue #1714)。
# 避免素材永远卡在中间状态
try:
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
existing = asset_repo.find_by_storage_key(job.storage_key)
if existing and existing.status in (
AssetStatus.PROCESSING,
AssetStatus.UPLOADING,
-7
View File
@@ -214,10 +214,3 @@ 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
-7
View File
@@ -231,10 +231,3 @@ 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
+3 -6
View File
@@ -115,8 +115,6 @@ 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}
@@ -130,11 +128,11 @@ services:
# 健康检查配置
# 注:celery inspect ping 依赖 broker 连接,在容器内不可靠,改用进程检查
healthcheck:
test: ["CMD-SHELL", "pgrep -f 'celery.*worker' | head -n1 >/dev/null 2>&1 || grep -q celery /proc/1/cmdline || exit 1"]
test: ["CMD-SHELL", "grep -q celery /proc/1/cmdline || exit 1"]
interval: 30s
timeout: 10s
retries: 3
start_period: 40s
start_period: 30s
logging: *default-logging
@@ -142,8 +140,7 @@ services:
# 资源限制建议(生产环境建议启用)
# =========================================
# 注意: Worker 需要处理视频,建议分配更多资源
# #1714 队列隔离后容器内运行 generation + transcode 两个 worker 进程,
# 总并发 = WORKER_CONCURRENCY(默认 4),4C8G 以上确保视频渲染不 OOM
# 并发 4 时需要 4C8G 以上,确保视频渲染不 OOM
deploy:
resources:
limits:
+1 -2
View File
@@ -146,7 +146,6 @@ 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 \
@@ -155,7 +154,7 @@ docker run -d \
--restart unless-stopped \
--cpus 2 \
--memory 2g \
--health-cmd "sh -c \"pgrep -f 'celery.*worker' >/dev/null 2>&1 || grep -q celery /proc/1/cmdline || exit 1\"" \
--health-cmd "sh -c \"grep -q celery /proc/1/cmdline || exit 1\"" \
--health-interval 30s \
--health-timeout 10s \
--health-retries 3 \
+1 -2
View File
@@ -109,14 +109,13 @@ 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 \"pgrep -f 'celery.*worker' >/dev/null 2>&1 || grep -q celery /proc/1/cmdline || exit 1\"" \
--health-cmd "sh -c \"grep -q celery /proc/1/cmdline || exit 1\"" \
--health-interval 30s \
--health-timeout 10s \
--health-retries 3 \
+9 -54
View File
@@ -1,63 +1,18 @@
#!/bin/bash
# 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)
# Worker 启动脚本 — 支持 WORKER_CONCURRENCY 环境变量
# 未设置时默认 2(保持向后兼容)
set -e
CONCURRENCY="${WORKER_CONCURRENCY:-4}"
MAX_TASKS="${WORKER_MAX_TASKS_PER_CHILD:-100}"
CONCURRENCY="${WORKER_CONCURRENCY:-2}"
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 \
# ⚠️ 部署约束:此 Worker 必须且只能运行单实例(replicas=1)
# -B 标志嵌入 celery beatbeat 负责定期触发 pending 超时清理等定时任务
# 多实例部署会导致每个 Worker 独立运行 Beat,造成定时任务重复执行
# 若需横向扩展 Worker,必须将 Beat 拆分为独立服务(celery beat -A worker_app.celery_app
exec celery \
-A worker_app.celery_app \
worker \
--loglevel=info \
"-B" \
-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
"--concurrency=${CONCURRENCY}"
@@ -146,44 +146,3 @@ 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)
@@ -134,7 +134,6 @@ 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,
)
@@ -164,8 +163,6 @@ 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)
@@ -391,7 +388,6 @@ 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,
@@ -456,53 +452,3 @@ 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 历史素材不拦。
"""
from datetime import datetime, timedelta, timezone
if not name:
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,
)
if file_size and file_size > 0:
query = query.filter(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)
@@ -38,7 +38,6 @@ 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 "",
@@ -83,7 +82,6 @@ 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 "",
@@ -140,52 +138,6 @@ 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)
@@ -317,7 +269,6 @@ 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 ""
@@ -329,14 +280,12 @@ class SQLAlchemyGenerationTaskRepository:
def cleanup_stale_running(self, timeout_minutes: int = 10) -> int:
"""清理超时未更新的 running 任务(孤儿任务)。
Returns:
清理的任务数量(仅计数,保持旧签名兼容)
"""
items = self.cleanup_stale_running_with_ids(timeout_minutes)
return len(items)
将 status=running 且 updated_at 超过 timeout_minutes 分钟未更新的任务
标记为 failederror_message 标记为任务执行中断。
def cleanup_stale_running_with_ids(self, timeout_minutes: int = 10) -> list[tuple[str, str]]:
"""同 cleanup_stale_running,但返回 [(task_id, celery_task_id), ...] 供撤销队列消息。"""
Returns:
清理的任务数量
"""
from datetime import timedelta
cutoff = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes)
@@ -349,10 +298,8 @@ class SQLAlchemyGenerationTaskRepository:
.all()
)
if not models:
return []
result: list[tuple[str, str]] = []
return 0
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 = {
@@ -362,43 +309,43 @@ class SQLAlchemyGenerationTaskRepository:
}
model.completed_at = datetime.now(timezone.utc)
self.session.commit()
return result
return len(models)
def cleanup_stale_pending(self, timeout_minutes: int = 30) -> int:
"""清理超时的 pending 任务(未被 Worker 拉取的任务)。
Returns:
清理的任务数量(仅计数,保持旧签名兼容)
"""
items = self.cleanup_stale_pending_with_ids(timeout_minutes)
return len(items)
全局任务队列有 pending 数量上限,长期卡在 pending 的任务会占满队列,
导致新用户无法创建任务。将超时的 pending 任务标记为 failed。
def cleanup_stale_pending_with_ids(self, timeout_minutes: int = 30) -> list[tuple[str, str]]:
"""同 cleanup_stale_pending,但返回 [(task_id, celery_task_id), ...] 供撤销队列消息。"""
Args:
timeout_minutes: 超时时间(分钟),默认 30 分钟
Returns:
清理的任务数量
"""
from datetime import timedelta
cutoff = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes)
models = (
self.session.query(GenerationTaskModel)
.filter(
GenerationTaskModel.status == GenerationTaskStatus.PENDING.value,
GenerationTaskModel.created_at < cutoff,
)
.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)
count = (
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,
)
)
self.session.commit()
return result
return count
@@ -18,8 +18,6 @@ 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,
)
@@ -40,8 +38,6 @@ 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,
)
@@ -58,12 +54,6 @@ 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
@@ -94,7 +94,6 @@ 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))
@@ -245,8 +244,6 @@ 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))
@@ -299,7 +296,6 @@ 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="")
@@ -1,115 +0,0 @@
"""
微信账号绑定/解绑 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
@@ -20,7 +20,6 @@ import requests
logger = logging.getLogger(__name__)
STATE_TTL_SECONDS = 600 # state 有效期 10 分钟
STATE_KEY_PREFIX = "wechat:state:" # Redis key 前缀(独立逻辑命名空间)
class MemoryStateStore:
@@ -54,80 +53,6 @@ class MemoryStateStore:
del self._states[s]
class RedisStateStore:
"""Redis state 存储(多实例/容器重启安全)。
复用现有 Rediscelery 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:
"""微信用户信息"""
@@ -233,8 +158,6 @@ class WechatOAuthService:
"grant_type": "authorization_code",
}
token_resp = requests.get(token_url, params=token_params, timeout=10)
# 微信响应头不带 charsetrequests 默认按 ISO-8859-1 解码会导致中文乱码
token_resp.encoding = "utf-8"
token_data = token_resp.json()
if "errcode" in token_data and token_data["errcode"] != 0:
@@ -253,8 +176,6 @@ 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:
@@ -279,29 +200,7 @@ 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 服务单例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
"""获取微信 OAuth 服务单例"""
# TODO: 可替换为 Redis state store
return WechatOAuthService()
-4
View File
@@ -12,8 +12,6 @@ class SubmitIngestJobCommand:
library_id: str
storage_key: str
file_hash: str = ""
asset_id: str = ""
celery_task_id: str = ""
class SubmitIngestJobUseCase:
@@ -26,7 +24,5 @@ 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)
-9
View File
@@ -174,7 +174,6 @@ 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))
@@ -209,7 +208,6 @@ 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:
@@ -237,7 +235,6 @@ 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=[],
)
@@ -269,8 +266,6 @@ 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))
@@ -281,8 +276,6 @@ 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 不能为空")
@@ -296,6 +289,4 @@ 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(),
)
-1
View File
@@ -117,7 +117,6 @@ 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 = ""
-20
View File
@@ -125,23 +125,3 @@ 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
@@ -20,12 +20,6 @@ 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]: ...
-222
View File
@@ -1,222 +0,0 @@
"""孤儿任务消息撤销与执行前状态守卫(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 的消息直接丢弃(抛 StaleTaskDiscardedtask 捕获后安全返回,
不进入渲染/转码,不产出半成品)。
"""
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": <celery id>, "task": <name>, ...}, ...}
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 消息 IDheaders.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: 业务任务 IDgeneration_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
)
-58
View File
@@ -1,58 +0,0 @@
"""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 appAPI 生产端与 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
+1 -1
View File
@@ -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 WECHAT_APP_ID WECHAT_APP_SECRET"
SHARED_SECRETS="OSS_ACCESS_KEY_ID OSS_ACCESS_KEY_SECRET COSYVOICE_API_KEY DASHSCOPE_API_KEY MEDIAKIT_API_KEY"
for var in $SHARED_SECRETS; do
value="${!var:-}"
# 已经在环境中了,无需额外操作
+1 -1
View File
@@ -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 -Q generation,transcode,celery
celery -A celery_app worker --loglevel=info --pool=solo
+1 -1
View File
@@ -287,7 +287,7 @@ class TestComputeDuplicateRateBadFingerprint:
videos = [
_make_existing_video("vid-b1", "md5_b1", ["aaaaaaaaaaaaaaaa"] * 10),
_make_existing_video("vid-b2", "md5_b2", ["cccccccccccccccc"] * 5), # hamming(a,c)=32 > PHASH_THRESHOLD
_make_existing_video("vid-b2", "md5_b2", ["bbbbbbbbbbbbbbbb"] * 5),
]
mock_repo = MagicMock()
mock_repo.list_by_user.return_value = videos
@@ -1,176 +0,0 @@
"""#1714 队列隔离 + 作废消息清除 单元测试。
覆盖:
1. task_routesgenerate_video → generationingest_asset/classify/duplication → transcode
2. purge_stale_messages_from_queuesRedis 队列中作废任务消息被物理移除,未命中保留
3. revoke_and_purgerevoke 广播 + 队列清理同时生效
4. ensure_task_claimable:终态任务抛 StaleTaskDiscardedpending 放行
"""
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) == ""
+1 -117
View File
@@ -131,7 +131,7 @@ class TestSameSourceDifferentCrop:
def test_same_source_distance_at_threshold_still_detected(self):
"""距离正好等于阈值(<=)也要算匹配——阈值比较统一为 <=。"""
assert PHASH_THRESHOLD <= 16, "阈值应经真实数据校准保持在能检出同源裁剪/降重对的范围(#1702 二次校准为 16)"
assert PHASH_THRESHOLD <= 12, "阈值应经校准保持在能检出同源裁剪的范围"
ddp = VideoDeduplicator()
base = [_h(0) for _ in range(6)]
new = [_h(PHASH_THRESHOLD) for _ in range(6)]
@@ -430,119 +430,3 @@ class TestCheckDuplicateExcludesSelf:
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 帧距离 <=160.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.80staging 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 == []
+22 -22
View File
@@ -103,7 +103,6 @@ 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,
@@ -270,17 +269,22 @@ class TestFindDuplicateSegments:
注意:使用不同的 hash 对,确保后半部分帧距离 > 阈值。
"""
same_hash = "aaaaaaaaaaaaaaaa"
# 4 帧匹配,后面 6 帧用与匹配哈希距离 32 的不匹配哈希(> PHASH_THRESHOLD=16
nomatch_hash = "cccccccccccccccc" # hamming(aaaa, cccc)=32
# 4 帧匹配,后面 6 帧各自不同(在 query 和 target 中使用不同 hash
chunks_a = [_make_chunk(i * 1000, (i + 1) * 1000, same_hash) for i in range(4)] + [
_make_chunk(i * 1000, (i + 1) * 1000, nomatch_hash) for i in range(4, 10)
_make_chunk(i * 1000, (i + 1) * 1000, "bbbbbbbbbbbbbbbb") 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, nomatch_hash) for i in range(4, 10)
_make_chunk(i * 1000, (i + 1) * 1000, "cccccccccccccccc") for i in range(4, 10)
]
# hamming(aaaa..., cccc...) = 32 > PHASH_THRESHOLD(16),后半段不匹配;
# 前 4 帧匹配 < min_consecutive=5,不形成片段
# 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
segments = find_duplicate_segments(chunks_a, chunks_b)
# 只有 4 帧匹配(< min_consecutive=5),所以不报告
assert segments == []
def test_max_gap_behavior(self):
@@ -289,12 +293,10 @@ class TestFindDuplicateSegments:
关键:间隙帧必须在 query 和 target 中使用不同 hash,使其真正不匹配。
"""
match_hash = "aaaaaaaaaaaaaaaa"
# 间隙/尾部哈希与 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
gap_hash_a = "bbbbbbbbbbbbbbbb" # query 端
gap_hash_b = "cccccccccccccccc" # target 端(与 query 端距离 > 8
tail_hash_a = "dddddddddddddddd"
tail_hash_b = "eeeeeeeeeeeeeeee"
# 5 帧匹配, 1 帧间隙, 3 帧匹配, 5 帧不匹配
hashes_a = [match_hash] * 5 + [gap_hash_a] + [match_hash] * 3 + [tail_hash_a] * 5
@@ -316,10 +318,10 @@ class TestFindDuplicateSegments:
def test_max_gap_exceeded(self):
"""间隙超过 max_gap → 分成两段."""
match_hash = "aaaaaaaaaaaaaaaa"
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
gap_hash_a = "bbbbbbbbbbbbbbbb"
gap_hash_b = "cccccccccccccccc"
tail_hash_a = "dddddddddddddddd"
tail_hash_b = "eeeeeeeeeeeeeeee"
# 5 帧匹配, 3 帧间隙 (> max_gap=2), 5 帧匹配, 5 帧不匹配
hashes_a = [match_hash] * 5 + [gap_hash_a] * 3 + [match_hash] * 5 + [tail_hash_a] * 5
@@ -495,11 +497,9 @@ class TestConstants:
"""常量值验证 — 使用已在模块顶部导入的常量,避免重新 import."""
def test_segment_match_threshold(self):
# 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
# Issue #1702: pHash 阈值经 staging 真实同源/异源指纹回归校准
# (同源密集采样 min=8、异源 min=24),统一为模块常量 PHASH_THRESHOLD=12。
assert SEGMENT_MATCH_THRESHOLD == 12
def test_min_consecutive_matches(self):
assert MIN_CONSECUTIVE_MATCHES == 5
@@ -1,57 +0,0 @@
"""#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]
-317
View File
@@ -1,317 +0,0 @@
"""Issue #1714HEVC 转码后禁止兜底新建重复 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"
@@ -1,360 +0,0 @@
"""#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_queuebiz id 命中、celery id 命中、未命中保序(重写 rpush)、
lrange 异常、重写异常、空队列
- purge_stale_messages_from_queues:空 ids 早退、redis 未安装、连接失败、
正常清理并 close
- revoke_and_purgerevoke 逐消息调用、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"
@@ -1,259 +0,0 @@
"""#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_taskingest 分支):重试后持久化 + 异常吞掉
- routes/upload._persist_celery_task_id:空 id 早退 + 异常吞掉
- core/task_enqueue.safe_enqueue_generation_task:持久化失败仅 warning,入队仍 True
- core/celery_appapply_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_taskingest 分支) ────────────────
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"
@@ -129,17 +129,15 @@ _ZERO_HIST = [0.0] * 96 # 全黑视频的全零直方图(有效数据)
class TestThresholdCalibration:
"""pHash 阈值校准(#1658 收紧到 8,#1702 两轮真实数据重校准 12→16)。
"""pHash 阈值校准(Issue #1658 收紧到 8Issue #1702 经真实指纹分布重校准 12)。
#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 裕度)。
#1702 staging 离线实验:同帧两次 2-5% 随机裁剪距离 4~10;同源成片(密集 1s
采样)最小距离 8、<=12 命中 10/31异源成片最小距离 24。8 会漏检同源裁剪,
12 检出同源且与异源分布(>=24)间隔充足。
"""
def test_phash_threshold_is_calibrated(self):
assert VideoDeduplicator.PHASH_THRESHOLD == PHASH_THRESHOLD == 16
assert VideoDeduplicator.PHASH_THRESHOLD == PHASH_THRESHOLD == 12
def test_match_ratio_threshold_constant(self):
assert MATCH_RATIO_THRESHOLD == 0.7
@@ -154,19 +152,18 @@ class TestThresholdCalibration:
def test_threshold_matching_semantics(self):
"""阈值比较统一为 <=(帧匹配与片段匹配同一口径)。
场景:5 个关键帧距离为 [10, 14, 16, 18, 26]。
- <=16#1702 二次校准阈值):3 帧匹配 → 0.6 < 0.7被帧比例门槛
拦截(异源安全边界:真实数据异源最近邻最小距离 18,<=16 命中 0
- 距离正好 16 的同源降重帧应算匹配(< 与 <= 口径统一)
场景:5 个关键帧距离为 [10, 12, 12, 24, 26]。
- <=12(#1702 校准阈值):3 帧匹配 → 0.6 < 0.7 被帧比例门槛拦截异源
- 距离 12 的同源裁剪帧应算匹配(< 与 <= 口径统一
"""
distances = [10, 14, 16, 18, 26]
distances = [10, 12, 12, 24, 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
# 异源安全边界(实测最小距离 18)及以上绝不匹配
assert not any(d <= VideoDeduplicator.PHASH_THRESHOLD for d in (18, 24, 26, 30))
# 异源典型距离(>=24绝不匹配
assert not any(d <= VideoDeduplicator.PHASH_THRESHOLD for d in (24, 26, 30))
# ── TestComputeFusionScore:统一融合得分方法 ────────────────────
-204
View File
@@ -1,204 +0,0 @@
"""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)
-243
View File
@@ -1,243 +0,0 @@
"""Issue #1714:任务执行前状态守卫 — 已作废消息必须丢弃,禁止非法转换后继续跑。
覆盖
- ingest_assetjob failed/completed 时直接返回 discarded不下载不转码不回写
- generate_videoGenerationTask failed 时返回 discarded不进入渲染
- generate_videopending 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=Truetask(task_id) 调用时注入 selfMagicMock
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 相关 mockgeneration 模块导入链)
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()
@@ -1,315 +0,0 @@
"""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
+2 -6
View File
@@ -164,9 +164,7 @@ 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"])
# 成功入队后持久化 celery 消息 ID#1714:孤儿清理据此 revoke/清队列)
assert len(repo.updated_tasks) == 1
assert task.celery_task_id
assert len(repo.updated_tasks) == 0 # 成功不需要更新状态
def test_user_limit_rejected_with_failed_status(self, mock_celery):
"""用户超限:任务标记为 failed,抛 UserPendingLimitExceeded。"""
@@ -337,9 +335,7 @@ class TestPostEnqueueFinalCheck:
assert result is True
mock_celery.assert_called_once()
assert task.status == "pending" # 状态没变
# 入队成功后持久化 celery_task_id#1714),业务状态不变
assert len(repo.updated_tasks) == 1
assert task.celery_task_id
assert len(repo.updated_tasks) == 0 # 没更新 DB
def test_post_enqueue_no_user_id_skips_user_check(self, mock_celery):
"""不传 user_id 时,入队后校验也跳过用户级,只查全局。"""
@@ -1,338 +0,0 @@
"""Issue #1714POST /upload/direct/complete 幂等 + multipart 幂等。
覆盖
- client_upload_id 重复 complete 只建一条 asset不重复派 ingest job
- file_hash 重复 complete 返回已存在记录
- 旧客户端不传 hash/token近期同库同名 processing 占位 兜底幂等返回
- 旧客户端不传 hash/tokenREADY 历史同名 不兜底正常新建
- 兜底窗口外>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:
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 (not file_size or 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
r1 = client.post("/api/v1/direct/complete", json=COMPLETE_BODY)
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": 0},
)
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_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
-233
View File
@@ -1,233 +0,0 @@
"""#1719:微信绑定/解绑路由层测试(直接驱动路由函数)。
覆盖
- GET /wechat/bind/url oauth 生成链接记日志
- POST /wechat/bindoauth 失败400绑定成功success+user.wechat_bound=True
use case 返回冲突对应状态码透传
- DELETE /wechat/bind成功success=Trueuse 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,
)
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,
)
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
@@ -1,251 +0,0 @@
"""#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
@@ -1,115 +0,0 @@
"""#1718:微信回调路由可观测性日志分支覆盖(UA/state/错误透传)。
直接驱动 wechat_callback 路由函数mock OAuth service 与用户仓储
- 成功路径日志记录 UAstate 校验通过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)
-49
View File
@@ -388,52 +388,3 @@ 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
-262
View File
@@ -1,262 +0,0 @@
"""#1718:微信 OAuth state 存储 Redis 化 + 中文昵称 UTF-8 解码修复。
覆盖 mock/fakeCI 无真实 redis 也产生覆盖
- RedisStateStoreput SET NX EXverify_and_consume GETDEL 一次性消费
重复消费返回 FalseRedis 异常降级内存client 注入
- Redis 不可用ping 失败构造时降级内存功能仍正常
- GETDEL 不存在 Redis GET+DELETE 兜底
- handle_callback微信 sns/userinfo 响应含中文 nicknameresp.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_callbackstate 校验 + 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)