Compare commits
18 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 6a1ec20e68 | |||
| df99305dd6 | |||
| 9840d5d778 | |||
| c8cfed0add | |||
| 8a6d51f6c3 | |||
| 6521be5426 | |||
| 0a7ae8db4f | |||
| 04b4fab131 | |||
| fbfd19fbb9 | |||
| c321ac3af8 | |||
| 21c26b5b26 | |||
| cf83c0df9f | |||
| ca7f875224 | |||
| 1c5060c224 | |||
| 15909a92e5 | |||
| af25045123 | |||
| 7a4aa27f71 | |||
| 28b3010668 |
@@ -1187,6 +1187,8 @@ 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..."
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
"""add client_upload_id to assets and asset_id to ingest_jobs
|
||||
|
||||
Issue #1714:上传 complete 幂等 + worker 转码回写关联。
|
||||
- assets.client_upload_id:客户端幂等 token(complete 去重)
|
||||
- ingest_jobs.asset_id:complete 阶段创建的占位 asset id(worker 回写关联,
|
||||
防止 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")
|
||||
@@ -0,0 +1,35 @@
|
||||
"""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")
|
||||
@@ -14,6 +14,7 @@ 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
|
||||
@@ -381,7 +382,8 @@ async def complete_chunked_upload(
|
||||
file_hash=request.file_hash,
|
||||
)
|
||||
)
|
||||
celery_app.send_task("worker.ingest_asset", args=[job.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", ""))
|
||||
|
||||
# Update metadata status
|
||||
meta["status"] = "completed"
|
||||
|
||||
@@ -14,6 +14,7 @@ from app.core.task_enqueue import (
|
||||
USER_PENDING_LIMIT,
|
||||
GlobalQueueFull,
|
||||
UserPendingLimitExceeded,
|
||||
build_rate_limit_detail,
|
||||
safe_enqueue_generation_task,
|
||||
)
|
||||
from app.dependencies import (
|
||||
@@ -312,12 +313,12 @@ def create_preview_generation_task(
|
||||
except UserPendingLimitExceeded as e:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=f"您的待处理任务过多(当前 {e.pending_count - count}/{e.limit},本次提交 {count} 个),请等待后再提交",
|
||||
detail=build_rate_limit_detail(e, generation_task_repository, scope="user"),
|
||||
) from e
|
||||
except GlobalQueueFull as e:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail="系统繁忙,请稍后再试",
|
||||
detail=build_rate_limit_detail(e, generation_task_repository, scope="global"),
|
||||
) from e
|
||||
|
||||
# 确定视频比例:优先前端传入,否则从模板 mode 推断
|
||||
@@ -498,6 +499,7 @@ 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(
|
||||
@@ -510,23 +512,29 @@ def create_preview_generation_task(
|
||||
if not enqueued:
|
||||
logger.warning("[预览生成] 任务入队失败: task_id=%s", task.id)
|
||||
_mark_task_failed(generation_task_repository, task, "任务入队失败")
|
||||
except UserPendingLimitExceeded:
|
||||
except UserPendingLimitExceeded as e:
|
||||
_mark_task_failed(generation_task_repository, task, "待处理任务超限")
|
||||
except GlobalQueueFull:
|
||||
rate_limit_exc = rate_limit_exc or e
|
||||
except GlobalQueueFull as e:
|
||||
_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):
|
||||
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 "系统繁忙,请稍后再试")
|
||||
# 队列满/限流时若全部失败,返回结构化错误码(前端区分"排队"与"创建失败")
|
||||
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"),
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"[预览生成] 创建完成: %d 个变体任务, task_ids=%s",
|
||||
|
||||
@@ -10,6 +10,7 @@ from app.core.task_enqueue import (
|
||||
USER_PENDING_LIMIT,
|
||||
GlobalQueueFull,
|
||||
UserPendingLimitExceeded,
|
||||
build_rate_limit_detail,
|
||||
safe_enqueue_generation_task,
|
||||
)
|
||||
from app.dependencies import (
|
||||
@@ -417,12 +418,12 @@ def create_generation_task(
|
||||
except UserPendingLimitExceeded as e:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=f"您的待处理任务过多(当前 {e.pending_count - count}/{e.limit},本次提交 {count} 个),请等待完成后再提交",
|
||||
detail=build_rate_limit_detail(e, generation_task_repository, scope="user"),
|
||||
) from e
|
||||
except GlobalQueueFull as e:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail="系统繁忙,请稍后再试",
|
||||
detail=build_rate_limit_detail(e, generation_task_repository, scope="global"),
|
||||
) from e
|
||||
|
||||
# 画中画已下线:strategy_id 中的 pip/voice_pip 统一映射为 one_take
|
||||
@@ -579,7 +580,7 @@ def create_generation_task(
|
||||
if not created_tasks:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail="您的待处理任务过多,请等待完成后再提交",
|
||||
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
|
||||
) from _e
|
||||
break
|
||||
except GlobalQueueFull as _e:
|
||||
@@ -587,7 +588,7 @@ def create_generation_task(
|
||||
if not created_tasks:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail="系统繁忙,请稍后再试",
|
||||
detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"),
|
||||
) from _e
|
||||
break
|
||||
except HTTPException:
|
||||
@@ -713,15 +714,15 @@ def confirm_generation(
|
||||
log_task_status=True,
|
||||
):
|
||||
logger.warning("[确认生成] 入队失败: task_id=%s", new_task.id)
|
||||
except UserPendingLimitExceeded:
|
||||
except UserPendingLimitExceeded as _e:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail="您的待处理任务过多,请等待完成后再提交",
|
||||
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
|
||||
) from None
|
||||
except GlobalQueueFull:
|
||||
except GlobalQueueFull as _e:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail="系统繁忙,请稍后再试",
|
||||
detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"),
|
||||
) from None
|
||||
|
||||
return BatchGenerationTaskResponse(
|
||||
@@ -803,12 +804,24 @@ def retry_generation_task(
|
||||
if user_pending >= USER_PENDING_LIMIT:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
|
||||
detail=build_rate_limit_detail(
|
||||
UserPendingLimitExceeded(
|
||||
user_id=user_id,
|
||||
pending_count=user_pending,
|
||||
limit=USER_PENDING_LIMIT,
|
||||
),
|
||||
generation_task_repository,
|
||||
scope="user",
|
||||
),
|
||||
)
|
||||
if global_pending >= GLOBAL_PENDING_LIMIT:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail="系统繁忙,请稍后再试",
|
||||
detail=build_rate_limit_detail(
|
||||
GlobalQueueFull(pending_count=global_pending, limit=GLOBAL_PENDING_LIMIT),
|
||||
generation_task_repository,
|
||||
scope="global",
|
||||
),
|
||||
)
|
||||
|
||||
use_case = CreateGenerationTaskUseCase(generation_task_repository)
|
||||
@@ -843,15 +856,15 @@ def retry_generation_task(
|
||||
log_task_status=True,
|
||||
):
|
||||
logger.warning("[生成任务] 重试入队失败: task_id=%s", retried.id)
|
||||
except UserPendingLimitExceeded:
|
||||
except UserPendingLimitExceeded as _e:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail="您的待处理任务过多,请等待完成后再提交",
|
||||
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
|
||||
) from None
|
||||
except GlobalQueueFull:
|
||||
except GlobalQueueFull as _e:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail="系统繁忙,请稍后再试",
|
||||
detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"),
|
||||
) from None
|
||||
return _to_generation_task_response(retried)
|
||||
|
||||
|
||||
@@ -43,7 +43,13 @@ def submit_ingest_job(
|
||||
)
|
||||
)
|
||||
|
||||
celery_app.send_task("worker.ingest_asset", args=[job.id])
|
||||
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
|
||||
|
||||
return IngestJobResponse(
|
||||
id=job.id,
|
||||
|
||||
@@ -375,7 +375,13 @@ def retry_project_task(
|
||||
storage_key=job.storage_key,
|
||||
)
|
||||
)
|
||||
celery_app.send_task("worker.ingest_asset", args=[retried.id])
|
||||
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
|
||||
return ProjectTaskResponse(
|
||||
id=f"ingest:{retried.id}",
|
||||
task_type="ingest",
|
||||
|
||||
@@ -85,12 +85,22 @@ 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):
|
||||
@@ -98,8 +108,85 @@ 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=""
|
||||
asset_repository,
|
||||
project_id,
|
||||
library_id,
|
||||
storage_key,
|
||||
filename,
|
||||
mime_type,
|
||||
user_id,
|
||||
file_hash="",
|
||||
client_upload_id="",
|
||||
):
|
||||
"""立即创建一条 PROCESSING 状态的 Asset 记录,使前端能马上看到新素材。"""
|
||||
asset = Asset.create(
|
||||
@@ -111,16 +198,29 @@ 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(
|
||||
@@ -129,9 +229,11 @@ def _submit_ingest_job(
|
||||
library_id=library_id,
|
||||
storage_key=storage_key,
|
||||
file_hash=file_hash,
|
||||
asset_id=asset_id,
|
||||
)
|
||||
)
|
||||
celery_app.send_task("worker.ingest_asset", args=[job.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", ""))
|
||||
return job
|
||||
|
||||
|
||||
@@ -202,7 +304,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,
|
||||
@@ -212,6 +314,29 @@ 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:
|
||||
@@ -223,29 +348,7 @@ 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,
|
||||
@@ -256,6 +359,7 @@ 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(
|
||||
@@ -264,6 +368,7 @@ 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,
|
||||
@@ -283,7 +388,8 @@ 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="文件 MD5 哈希,用于去重检测"),
|
||||
file_hash: str = Form(default="", description="文件哈希,用于去重检测"),
|
||||
client_upload_id: str = Form(default="", description="客户端幂等 token(同一次上传的重试保持一致)"),
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
ingest_job_repository: Any = Depends(get_ingest_job_repository),
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
@@ -294,32 +400,31 @@ async def upload_asset(
|
||||
"""上传素材文件并触发导入流水线。"""
|
||||
require_project_and_library(project_id, library_id, project_repository, asset_library_repository)
|
||||
|
||||
# ── 素材去重检测:上传前检查同素材库 + 同 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 类型
|
||||
# P2-5: 服务端验证 MIME 类型(先验证,再幂等去重,避免非法类型绕过)
|
||||
validated_content_type = _validate_mime_type(file.content_type)
|
||||
|
||||
file_id = uuid4().hex[:8]
|
||||
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]
|
||||
storage_key = f"uploads/{file_id}/{safe_filename}"
|
||||
|
||||
try:
|
||||
@@ -348,6 +453,7 @@ 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(
|
||||
@@ -356,6 +462,7 @@ 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(
|
||||
|
||||
@@ -252,6 +252,10 @@ class RecomputeDedupRequest(BaseModel):
|
||||
None,
|
||||
description="指定视频 ID 列表。为空则对当前用户所有缺少查重数据的视频重新计算。",
|
||||
)
|
||||
force: bool = Field(
|
||||
False,
|
||||
description="强制重算:即使视频已有查重数据也重新入队(#1702 查重算法升级后用于存量视频重算)。",
|
||||
)
|
||||
|
||||
|
||||
class RecomputeDedupResponse(BaseModel):
|
||||
@@ -291,15 +295,15 @@ def recompute_dedup(
|
||||
skipped = 0
|
||||
|
||||
for video in target_videos:
|
||||
# 已有完整查重数据的跳过
|
||||
if video.duplicate_rate is not None and video.video_fingerprint:
|
||||
# 已有完整查重数据的跳过(force=True 时强制重算,#1702 算法升级后存量视频需要重算指纹/分片)
|
||||
if not request.force and video.duplicate_rate is not None and video.video_fingerprint:
|
||||
skipped += 1
|
||||
continue
|
||||
|
||||
# 触发异步查重任务
|
||||
celery_app.send_task("worker.check_duplicate", args=[video.id])
|
||||
enqueued += 1
|
||||
logger.info("Enqueued re-dedup for video %s (user=%s)", video.id, user_id)
|
||||
logger.info("Enqueued re-dedup for video %s (user=%s, force=%s)", video.id, user_id, request.force)
|
||||
|
||||
return RecomputeDedupResponse(
|
||||
enqueued=enqueued,
|
||||
|
||||
@@ -5,3 +5,11 @@ 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
|
||||
|
||||
@@ -8,27 +8,145 @@ 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):
|
||||
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,
|
||||
):
|
||||
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):
|
||||
def __init__(
|
||||
self,
|
||||
pending_count: int,
|
||||
limit: int,
|
||||
*,
|
||||
running_count: int = 0,
|
||||
queue_ahead: int = 0,
|
||||
estimated_wait_seconds: int = 0,
|
||||
):
|
||||
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,
|
||||
@@ -161,7 +279,17 @@ def safe_enqueue_generation_task(
|
||||
|
||||
# ── 发送 Celery 任务 ──
|
||||
try:
|
||||
celery_app.send_task("worker.generate_video", args=[task.id])
|
||||
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
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"%s 入队失败,标记为失败: task_id=%s error=%s",
|
||||
|
||||
@@ -31,14 +31,16 @@ 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="文件 MD5 哈希,用于去重检测")
|
||||
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 时的兜底去重")
|
||||
|
||||
|
||||
class DirectUploadCompleteResponse(BaseModel):
|
||||
storage_key: str
|
||||
ingest_job_id: str
|
||||
duplicated: bool = Field(default=False, description="是否为重复素材(命中去重)")
|
||||
asset_id: str = Field(default="", description="重复素材的 asset_id(duplicated=true 时返回)")
|
||||
duplicated: bool = Field(default=False, description="是否为重复素材/重复 complete(命中幂等去重)")
|
||||
asset_id: str = Field(default="", description="素材 asset_id(重复 complete 时返回已存在记录)")
|
||||
url: str = Field(default="", description="Public URL of uploaded file")
|
||||
|
||||
|
||||
@@ -46,5 +48,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=true 时返回)")
|
||||
duplicated: bool = Field(default=False, description="是否为重复素材/重复提交(命中幂等去重)")
|
||||
asset_id: str = Field(default="", description="素材 asset_id(重复提交时返回已存在记录)")
|
||||
|
||||
@@ -185,6 +185,13 @@ test.describe("Core generation flow", () => {
|
||||
await expect(page.locator(".xx-choice-item.selected")).toBeVisible()
|
||||
await page.getByRole("button", { name: "下一步" }).click()
|
||||
|
||||
// Step1 下一步弹出数量选择弹窗(Issue #1677 固定6步:模板→素材→配音→标题→确认生成→封面)
|
||||
// 单视频流程:默认 1 个,点击「生成 1 个视频」进入步骤2
|
||||
await expect(page.getByRole("heading", { name: "要生成几个视频?" })).toBeVisible({
|
||||
timeout: 10_000,
|
||||
})
|
||||
await page.getByRole("button", { name: "生成 1 个视频" }).click()
|
||||
|
||||
// Step 2: select material (card grid UI)
|
||||
await expect(page.getByRole("heading", { name: /选择素材/ })).toBeVisible()
|
||||
const librarySelect = page.locator("select").first()
|
||||
@@ -255,9 +262,18 @@ test.describe("Core generation flow", () => {
|
||||
expect(genData.items.length).toBeGreaterThan(0)
|
||||
expect(genData.items[0].id).toBeTruthy()
|
||||
|
||||
// 单视频(N=1):点击「确认生成视频」后直接跳 Step 5 封面(与旧流程一致)
|
||||
// 单视频(N=1):点击「确认生成视频」后跳 Step 5「确认生成」,展示实时渲染进度
|
||||
await expect(page.getByRole("heading", { name: "🎬 确认生成" })).toBeVisible({
|
||||
timeout: 30_000,
|
||||
})
|
||||
|
||||
// 等待渲染完成:进度卡变为「视频生成完成」(最长等待 3 分钟)
|
||||
await expect(page.getByText("视频生成完成")).toBeVisible({ timeout: 180_000 })
|
||||
|
||||
// 全部完成后「下一步:选择封面」解锁,点击进入 Step 6
|
||||
await page.getByRole("button", { name: /下一步:选择封面/ }).click()
|
||||
await expect(page.getByRole("heading", { name: /选择封面/ })).toBeVisible({
|
||||
timeout: 180_000,
|
||||
timeout: 30_000,
|
||||
})
|
||||
} else {
|
||||
console.log(`[E2E] Generate API returned ${genResp.status()}, wizard flow test still passes`)
|
||||
|
||||
@@ -4,6 +4,7 @@
|
||||
import apiClient from "../client"
|
||||
import { getOrCreateDefaultProject } from "../projects"
|
||||
import type { DirectUploadPrepareResult, DirectUploadCompleteResult } from "./types"
|
||||
import { computeFileHash, makeClientUploadId } from "./uploadDedup"
|
||||
|
||||
/** 预签名直传准备 */
|
||||
export const prepareDirectUpload = async (data: {
|
||||
@@ -12,8 +13,13 @@ 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> => {
|
||||
const response = await apiClient.post("/upload/direct/prepare", data)
|
||||
// prepare 单独放宽到 30s(全局 axios 实例只有 10s,staging 抖动时易超时)
|
||||
const response = await apiClient.post("/upload/direct/prepare", data, { timeout: 30_000 })
|
||||
return response.data
|
||||
}
|
||||
|
||||
@@ -22,8 +28,14 @@ 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> => {
|
||||
const response = await apiClient.post("/upload/direct/complete", data)
|
||||
// complete 内含 OSS 存在性检查 + 建库 + 派单,放宽到 60s;
|
||||
// 超时不代表失败(记录可能已建成),调用方禁止超时后盲目重传整个文件
|
||||
const response = await apiClient.post("/upload/direct/complete", data, { timeout: 60_000 })
|
||||
return response.data
|
||||
}
|
||||
|
||||
@@ -109,6 +121,10 @@ export interface DirectUploadHandle {
|
||||
export const prepareDirectUploadHandle = async (data: {
|
||||
file: File
|
||||
library_id: string
|
||||
/** 前端算好的文件内容哈希(SHA-256 hex),prepare/complete 均携带 */
|
||||
fileHash?: string
|
||||
/** 本次逻辑上传的幂等 token,prepare/complete 一致、重试复用 */
|
||||
clientUploadId?: string
|
||||
}): Promise<DirectUploadHandle> => {
|
||||
const project = await getOrCreateDefaultProject()
|
||||
|
||||
@@ -118,6 +134,8 @@ 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 {
|
||||
@@ -128,6 +146,8 @@ 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,
|
||||
}),
|
||||
}
|
||||
}
|
||||
@@ -137,8 +157,20 @@ export const uploadAssetDirect = async (data: {
|
||||
file: File
|
||||
library_id: string
|
||||
onProgress?: (percent: number) => void
|
||||
/** 文件内容哈希;未传时自动补算(配音/封面/克隆等非队列链路统一受益) */
|
||||
fileHash?: string
|
||||
/** 幂等 token;未传时自动生成 */
|
||||
clientUploadId?: string
|
||||
}): Promise<DirectUploadCompleteResult> => {
|
||||
const handle = await prepareDirectUploadHandle({ file: data.file, library_id: data.library_id })
|
||||
// 自动补算哈希与幂等 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,
|
||||
})
|
||||
await handle.transfer(data.onProgress)
|
||||
return handle.complete()
|
||||
}
|
||||
|
||||
@@ -0,0 +1,143 @@
|
||||
/**
|
||||
* 上传去重 / 幂等工具(Issue #1714)
|
||||
*
|
||||
* 背景:同一文件被反复入队、complete 超时后盲目重传,导致后端创建大量重复
|
||||
* PROCESSING 素材记录。本模块提供两类纯函数:
|
||||
*
|
||||
* 1. 文件指纹:
|
||||
* - makeFileFingerprint():文件名+大小+lastModified,入队去重用(同步、零开销)
|
||||
* - computeFileHash():SHA-256 内容哈希(小文件全量、大文件抽样头尾),
|
||||
* prepare/complete 时发给后端打开 file_hash 去重闸门
|
||||
* 2. 队列去重:findDuplicateInQueue() 判断文件是否已在队列中
|
||||
* 3. 幂等 token:makeClientUploadId() 生成上传幂等 ID(每次"一次逻辑上传"一个,
|
||||
* 重试复用同一 ID,重新入队才生成新 ID)
|
||||
*/
|
||||
|
||||
/** 大文件抽样阈值:超过此大小只哈希头尾片段,避免上传前长时间卡 UI */
|
||||
export const HASH_FULL_READ_LIMIT = 256 * 1024 * 1024 // 256MB
|
||||
/** 抽样读取的头尾片段大小(各 8MB) */
|
||||
export const HASH_SAMPLE_CHUNK = 8 * 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-256(hex,64 字符,与后端 file_hash 字段长度一致)。
|
||||
* - ≤256MB:全量哈希,内容一致必然一致
|
||||
* - >256MB:哈希「头部 8MB + 尾部 8MB + 文件大小」,视频素材体积大、
|
||||
* 头部含 moov 元数据、尾部含 mdat 结尾,抽样碰撞概率可忽略,
|
||||
* 且避免上传前对 2GB 文件全量读取造成长时间卡顿
|
||||
*
|
||||
* 运行环境不支持 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 ""
|
||||
}
|
||||
}
|
||||
@@ -12,13 +12,18 @@ export type {
|
||||
UserResponse,
|
||||
WechatAuthUrlResponse,
|
||||
WechatCallbackResponse,
|
||||
WechatBindUrlResponse,
|
||||
WechatBindCompleteResponse,
|
||||
WechatUnbindResponse,
|
||||
UpdateProfileRequest,
|
||||
UpdateProfileResponse,
|
||||
SendVerificationCodeRequest,
|
||||
BindContactRequest,
|
||||
BindContactResponse,
|
||||
} from "./types"
|
||||
|
||||
// 用户工具函数
|
||||
export { normalizeUser } from "./user"
|
||||
export { normalizeUser, updateProfile } from "./user"
|
||||
|
||||
// 登录/注册/登出/刷新
|
||||
export { login, refreshAccessToken, register, logout } from "./login"
|
||||
@@ -32,8 +37,14 @@ export { requestPasswordReset, resetPassword } from "./password"
|
||||
// 邮箱验证
|
||||
export { verifyEmail } from "./email"
|
||||
|
||||
// 微信登录
|
||||
export { getWechatAuthUrl, wechatCallback } from "./wechat"
|
||||
// 微信登录 / 绑定
|
||||
export {
|
||||
getWechatAuthUrl,
|
||||
wechatCallback,
|
||||
getWechatBindUrl,
|
||||
bindWechat,
|
||||
unbindWechat,
|
||||
} from "./wechat"
|
||||
|
||||
// 联系方式
|
||||
export { sendVerificationCode, bindContact } from "./contact"
|
||||
|
||||
@@ -34,6 +34,17 @@ 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 {
|
||||
@@ -45,6 +56,12 @@ 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 {
|
||||
@@ -80,3 +97,30 @@ 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,4 +1,5 @@
|
||||
import type { User, UserResponse } from "./types"
|
||||
import apiClient from "../client"
|
||||
import type { User, UserResponse, UpdateProfileRequest, UpdateProfileResponse } from "./types"
|
||||
|
||||
/**
|
||||
* 规范化用户数据,兼容不同后端返回格式
|
||||
@@ -16,5 +17,19 @@ 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)
|
||||
}
|
||||
|
||||
@@ -1,8 +1,14 @@
|
||||
import apiClient from "../client"
|
||||
import type { WechatAuthUrlResponse, WechatCallbackResponse } from "./types"
|
||||
import type {
|
||||
WechatAuthUrlResponse,
|
||||
WechatCallbackResponse,
|
||||
WechatBindUrlResponse,
|
||||
WechatBindCompleteResponse,
|
||||
WechatUnbindResponse,
|
||||
} from "./types"
|
||||
|
||||
/**
|
||||
* 获取微信授权链接
|
||||
* 获取微信授权链接(登录场景)
|
||||
*/
|
||||
export const getWechatAuthUrl = async (): Promise<WechatAuthUrlResponse> => {
|
||||
const response = await apiClient.get("/auth/wechat/url")
|
||||
@@ -19,3 +25,30 @@ 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
|
||||
}
|
||||
|
||||
@@ -42,6 +42,7 @@ const AssetLibrary: React.FC = () => {
|
||||
assetsError,
|
||||
assetsErrorObj,
|
||||
refetchAssets,
|
||||
stalledAssetIds,
|
||||
searchText,
|
||||
setSearchText,
|
||||
filterType,
|
||||
@@ -77,6 +78,7 @@ const AssetLibrary: React.FC = () => {
|
||||
removeUpload,
|
||||
clearFinished,
|
||||
uploading,
|
||||
transferActive,
|
||||
activeCount,
|
||||
pendingCount,
|
||||
} = useAssetUpload({ effectiveLibId })
|
||||
@@ -168,6 +170,7 @@ const AssetLibrary: React.FC = () => {
|
||||
{/* 上传区域 */}
|
||||
<AssetUploadZone
|
||||
uploading={uploading}
|
||||
transferActive={transferActive}
|
||||
activeCount={activeCount}
|
||||
pendingCount={pendingCount}
|
||||
onUpload={enqueueUploads}
|
||||
@@ -214,6 +217,7 @@ const AssetLibrary: React.FC = () => {
|
||||
selectedIds={selectedIds}
|
||||
diagnosingId={diagnosingId}
|
||||
uploadProgressMap={uploadProgressMap}
|
||||
stalledAssetIds={stalledAssetIds}
|
||||
onRetry={refetchAssets}
|
||||
onToggleSelect={toggleSelect}
|
||||
onDiagnose={handleDiagnose}
|
||||
|
||||
@@ -1041,3 +1041,25 @@
|
||||
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,6 +22,8 @@ export interface AssetCardProps {
|
||||
diagnosing?: boolean
|
||||
/** 上传中实时进度(仅 uploading 态有值;ingesting 后由后端状态接管) */
|
||||
uploadProgress?: { progress: number; uploading: boolean }
|
||||
/** 处理超过 10 分钟仍未就绪(疑似后端卡住):停止转圈并提示处理超时 */
|
||||
stalled?: boolean
|
||||
onToggle: () => void
|
||||
onDiagnose: () => void
|
||||
onPlay: () => void
|
||||
@@ -33,6 +35,7 @@ const AssetCard: React.FC<AssetCardProps> = ({
|
||||
selected,
|
||||
diagnosing,
|
||||
uploadProgress,
|
||||
stalled,
|
||||
onToggle,
|
||||
onDiagnose,
|
||||
onPlay,
|
||||
@@ -67,11 +70,15 @@ const AssetCard: React.FC<AssetCardProps> = ({
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 转码/处理中遮罩 */}
|
||||
{/* 转码/处理中遮罩(卡死超过 10 分钟时停止转圈,提示超时) */}
|
||||
{asset.loading && !isUploading && (
|
||||
<div className="xx-asset-thumb-overlay xx-asset-thumb-processing">
|
||||
<LoadingOutlined />
|
||||
<span>转码处理中</span>
|
||||
<div
|
||||
className={`xx-asset-thumb-overlay ${
|
||||
stalled ? "xx-asset-thumb-stalled" : "xx-asset-thumb-processing"
|
||||
}`}
|
||||
>
|
||||
{stalled ? <CloseCircleOutlined /> : <LoadingOutlined />}
|
||||
<span>{stalled ? "处理超时,可重试上传" : "转码处理中"}</span>
|
||||
</div>
|
||||
)}
|
||||
|
||||
@@ -129,7 +136,7 @@ const AssetCard: React.FC<AssetCardProps> = ({
|
||||
</p>
|
||||
<div className="xx-asset-meta">
|
||||
<span className="xx-asset-meta-status">
|
||||
<StatusPill status={asset.status} label={asset.statusLabel} />
|
||||
<StatusPill status={asset.status} label={stalled ? "处理超时" : asset.statusLabel} />
|
||||
</span>
|
||||
{asset.duration && <span className="xx-asset-meta-duration">{asset.duration}</span>}
|
||||
</div>
|
||||
|
||||
@@ -19,6 +19,8 @@ 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
|
||||
@@ -34,6 +36,7 @@ export const AssetGridSection: React.FC<AssetGridSectionProps> = ({
|
||||
selectedIds,
|
||||
diagnosingId,
|
||||
uploadProgressMap,
|
||||
stalledAssetIds,
|
||||
onRetry,
|
||||
onToggleSelect,
|
||||
onDiagnose,
|
||||
@@ -76,6 +79,7 @@ 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,10 +4,13 @@
|
||||
* - 拖拽文件到内容区任意位置同样触发上传(不再占用大面积虚线框)
|
||||
*/
|
||||
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
|
||||
@@ -15,6 +18,7 @@ export interface AssetUploadZoneProps {
|
||||
|
||||
export const AssetUploadZone: React.FC<AssetUploadZoneProps> = ({
|
||||
uploading,
|
||||
transferActive,
|
||||
activeCount,
|
||||
pendingCount,
|
||||
onUpload,
|
||||
@@ -27,6 +31,11 @@ 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))
|
||||
}
|
||||
|
||||
@@ -58,10 +67,18 @@ export const AssetUploadZone: React.FC<AssetUploadZoneProps> = ({
|
||||
<button
|
||||
type="button"
|
||||
className="xx-asset-upload-btn"
|
||||
onClick={() => inputRef.current?.click()}
|
||||
disabled={transferActive}
|
||||
title={transferActive ? "文件上传中,暂不能添加新文件" : undefined}
|
||||
onClick={() => {
|
||||
if (transferActive) {
|
||||
message.warning("文件正在上传中,请等待当前上传完成后再添加")
|
||||
return
|
||||
}
|
||||
inputRef.current?.click()
|
||||
}}
|
||||
>
|
||||
<PlusOutlined />
|
||||
上传素材
|
||||
{transferActive ? "上传中…" : "上传素材"}
|
||||
</button>
|
||||
<span className="xx-asset-upload-status">
|
||||
{uploading ? (
|
||||
|
||||
@@ -80,6 +80,7 @@ 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.error ? `:${it.error}` : ""}
|
||||
</div>
|
||||
@@ -89,7 +90,9 @@ const UploadQueuePanel: React.FC<UploadQueuePanelProps> = ({
|
||||
<button
|
||||
type="button"
|
||||
className="xx-upload-queue-btn"
|
||||
title="重试"
|
||||
title={
|
||||
it.failedStage === "complete" ? "安全重试(只确认,不重新上传)" : "重试上传"
|
||||
}
|
||||
onClick={() => onRetry(it.tempId)}
|
||||
>
|
||||
<ReloadOutlined />
|
||||
|
||||
@@ -3,33 +3,57 @@ import { useQueryClient } from "@tanstack/react-query"
|
||||
import { message } from "antd"
|
||||
import { prepareDirectUploadHandle, type DirectUploadHandle } from "@/api/assets"
|
||||
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 {
|
||||
/** 前端临时 id(prepare 前无 asset_id 时用) */
|
||||
/** 前端临时 id(prepare 前无 asset_id 时用),同时作为队列项 key */
|
||||
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 阶段失败后的错误提示:素材可能已在服务器处理中,重试不会重新上传 */
|
||||
const COMPLETE_ERROR_HINT =
|
||||
"确认请求失败,素材可能已在服务器处理中;点重试将安全确认,不会重新上传文件"
|
||||
|
||||
/**
|
||||
* 素材批量上传 Hook
|
||||
* - prepare 阶段后端预建 status=uploading 的 asset,前端拿到 asset_id 立即刷新列表
|
||||
* - 入队按文件指纹(name+size+lastModified)去重:同一文件已在队列/上传中/处理中时不重复入队
|
||||
* - 上传前计算文件 SHA-256,prepare/complete 携带 file_hash + 幂等 token(clientUploadId)
|
||||
* - complete 超时/失败不盲目重传:复用 handle 只重发 complete(幂等),prepare/transfer 失败才全量重跑
|
||||
* - OSS 直传并发限制为 3,其余排队;每个文件独立进度/状态
|
||||
* - complete 后素材进入转码(ingesting/processing),由列表轮询反映
|
||||
* - 失败卡片支持重试/移除
|
||||
*/
|
||||
export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
|
||||
@@ -39,6 +63,13 @@ 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)))
|
||||
}, [])
|
||||
@@ -52,14 +83,57 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
|
||||
queryClient.invalidateQueries({ queryKey: ["asset-libraries"] })
|
||||
}, [queryClient, effectiveLibId])
|
||||
|
||||
/** 执行单个文件的完整上传流程(prepare→transfer→complete) */
|
||||
/**
|
||||
* 执行单个文件的完整上传流程。
|
||||
* @param completeOnly complete 阶段失败后的重试:跳过 hash/prepare/transfer,只重发 complete
|
||||
* (OSS 文件已传完,重发由 file_hash + clientUploadId 保证幂等)
|
||||
*/
|
||||
const runUpload = useCallback(
|
||||
async (item: UploadItem, handle?: DirectUploadHandle) => {
|
||||
async (item: UploadItem, handle?: DirectUploadHandle, completeOnly = false) => {
|
||||
let stage: UploadFailStage = "prepare"
|
||||
try {
|
||||
// 1. prepare(重试时复用已准备的 handle 也行,但签名可能过期,重新 prepare 最稳)
|
||||
const h =
|
||||
handle ??
|
||||
(await prepareDirectUploadHandle({ file: item.file, library_id: effectiveLibId }))
|
||||
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)
|
||||
|
||||
if (h.prepared.asset_id) {
|
||||
updateItem(item.tempId, {
|
||||
status: "uploading",
|
||||
@@ -72,26 +146,50 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
|
||||
updateItem(item.tempId, { status: "uploading", progress: 0 })
|
||||
}
|
||||
|
||||
// 2. OSS 直传(真实进度)
|
||||
// 3. OSS 直传(真实进度)
|
||||
stage = "transfer"
|
||||
await h.transfer((pct) => updateItem(item.tempId, { progress: pct }))
|
||||
|
||||
// 3. complete:后端创建 ingest job,素材进入转码
|
||||
// 4. complete:后端确认入库并创建 ingest job(file_hash + 幂等 token 已在 handle 闭包中)
|
||||
stage = "complete"
|
||||
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" })
|
||||
updateItem(item.tempId, { status: "done", assetId: result.asset_id || item.assetId })
|
||||
message.success(`"${item.fileName}" 上传完成,正在转码处理`)
|
||||
}
|
||||
} catch (err: unknown) {
|
||||
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}`)
|
||||
console.error("[useAssetUpload] 上传失败:", item.fileName, stage, err)
|
||||
|
||||
if (stage === "complete") {
|
||||
// complete 失败(超时/5xx/网络):后端记录可能已建成,handle 保留供幂等重试;
|
||||
// 刷新列表让用户看到可能已创建的「处理中」素材,避免误以为没传上去而重复操作
|
||||
refreshList()
|
||||
updateItem(item.tempId, {
|
||||
status: "error",
|
||||
failedStage: "complete",
|
||||
error: COMPLETE_ERROR_HINT,
|
||||
hint: undefined,
|
||||
})
|
||||
message.error(`"${item.fileName}" ${COMPLETE_ERROR_HINT}`)
|
||||
} else {
|
||||
// prepare / transfer 失败:后端尚无素材记录,可安全全量重跑
|
||||
handlesRef.current.delete(item.tempId)
|
||||
updateItem(item.tempId, {
|
||||
status: "error",
|
||||
failedStage: stage === "transfer" ? "transfer" : "prepare",
|
||||
error: detail,
|
||||
hint: undefined,
|
||||
})
|
||||
message.error(`"${item.fileName}" 上传失败:${detail}`)
|
||||
}
|
||||
}
|
||||
},
|
||||
[effectiveLibId, refreshList, updateItem],
|
||||
@@ -113,7 +211,9 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
|
||||
if (!next) return
|
||||
claimedRef.current.add(next.tempId)
|
||||
inFlightRef.current += 1
|
||||
void runUpload(next).finally(() => {
|
||||
// complete 阶段失败的重试:复用留存的 handle,只重发 complete
|
||||
const existingHandle = handlesRef.current.get(next.tempId)
|
||||
void runUpload(next, existingHandle, existingHandle !== undefined).finally(() => {
|
||||
inFlightRef.current -= 1
|
||||
claimedRef.current.delete(next.tempId)
|
||||
// 一个任务结束(成功/失败)后继续拉起排队任务
|
||||
@@ -126,7 +226,11 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
|
||||
pumpRef.current()
|
||||
}, [items])
|
||||
|
||||
/** 入队一个或多个文件 */
|
||||
/**
|
||||
* 入队一个或多个文件(按文件指纹去重):
|
||||
* - 同一文件已在队列且 preparing/uploading/ingesting/done → 跳过,不重复入队
|
||||
* - 同一文件此前失败(error)→ 重新激活原队列项(复用 clientUploadId,保持幂等语义)
|
||||
*/
|
||||
const enqueueUploads = useCallback(
|
||||
(files: File[]) => {
|
||||
if (!effectiveLibId) {
|
||||
@@ -143,37 +247,86 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
|
||||
}
|
||||
if (valid.length === 0) return
|
||||
|
||||
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])
|
||||
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} 个此前失败的文件`)
|
||||
}
|
||||
},
|
||||
[effectiveLibId],
|
||||
[effectiveLibId, updateItem],
|
||||
)
|
||||
|
||||
/** 重试失败任务 */
|
||||
/**
|
||||
* 重试失败任务(仅限 status=error):
|
||||
* - complete 阶段失败:复用留存 handle 只重发 complete(幂等,不重新上传)
|
||||
* - prepare/transfer 阶段失败:全量重跑(后端尚无记录,安全)
|
||||
*/
|
||||
const retryUpload = useCallback(
|
||||
(tempId: string) => {
|
||||
const target = itemsRef.current.find((it) => it.tempId === tempId)
|
||||
if (!target) return
|
||||
updateItem(tempId, { status: "preparing", progress: 0, error: undefined })
|
||||
// 状态更新后由 useEffect 触发 pump
|
||||
if (!target || target.status !== "error") return
|
||||
updateItem(tempId, { status: "preparing", progress: 0, error: undefined, hint: undefined })
|
||||
// 状态更新后由 useEffect 触发 pump;pump 会按 handlesRef 自动选择 completeOnly / 全量
|
||||
},
|
||||
[updateItem],
|
||||
)
|
||||
|
||||
/** 从上传列表移除(已进入转码的由素材网格管理;这里只移除上传面板记录) */
|
||||
const removeUpload = useCallback((tempId: string) => {
|
||||
handlesRef.current.delete(tempId)
|
||||
setItems((prev) => prev.filter((it) => it.tempId !== tempId))
|
||||
}, [])
|
||||
|
||||
/** 清空已完成/去重记录 */
|
||||
const clearFinished = useCallback(() => {
|
||||
setItems((prev) => prev.filter((it) => it.status !== "done"))
|
||||
setItems((prev) => {
|
||||
for (const it of prev) {
|
||||
if (it.status === "done") handlesRef.current.delete(it.tempId)
|
||||
}
|
||||
return prev.filter((it) => it.status !== "done")
|
||||
})
|
||||
}, [])
|
||||
|
||||
const activeCount = items.filter(
|
||||
@@ -188,8 +341,10 @@ export function useAssetUpload({ effectiveLibId }: { effectiveLibId: string }) {
|
||||
retryUpload,
|
||||
removeUpload,
|
||||
clearFinished,
|
||||
/** 是否有进行中的上传(用于上传区文案) */
|
||||
/** 是否有进行中的上传(用于上传区文案/禁用入口) */
|
||||
uploading: hasActive,
|
||||
/** 是否有文件正在本地处理或直传(用于禁用上传入口,防重复提交) */
|
||||
transferActive: activeCount > 0,
|
||||
activeCount,
|
||||
pendingCount,
|
||||
}
|
||||
|
||||
@@ -10,6 +10,26 @@ 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
|
||||
* 封装视频库列表、素材列表的数据查询,以及筛选、搜索状态管理
|
||||
@@ -63,15 +83,15 @@ export function useAssetsData() {
|
||||
}),
|
||||
enabled: !!effectiveLibId,
|
||||
staleTime: 30_000,
|
||||
// 列表中存在上传中/转码中素材时每 3s 轮询;全部就绪后自动停止
|
||||
// 列表中存在上传中/转码中素材时每 3s 轮询;全部就绪后自动停止。
|
||||
// 但若处理中素材创建已超过 10 分钟仍未就绪(疑似后端卡住/孤儿任务),
|
||||
// 停止快轮询避免无限转圈——卡死素材在网格中显示「处理超时」提示。
|
||||
refetchInterval: (query) => {
|
||||
const data = query.state.data as { items: ApiAssetItem[] } | undefined
|
||||
const items = data?.items ?? []
|
||||
const processing = items.some((a) => {
|
||||
const st = a.status ?? ""
|
||||
return st === "uploading" || st === "ingesting" || st === "processing" || st === "pending"
|
||||
})
|
||||
return processing ? 3000 : false
|
||||
const processing = items.some((a) => isProcessingStatus(a.status ?? ""))
|
||||
if (!processing) return false
|
||||
return hasStalledProcessing(items) ? false : 3000
|
||||
},
|
||||
})
|
||||
|
||||
@@ -80,6 +100,21 @@ 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")
|
||||
@@ -129,6 +164,8 @@ export function useAssetsData() {
|
||||
assetsError,
|
||||
assetsErrorObj,
|
||||
refetchAssets,
|
||||
stalledAssetIds,
|
||||
hasStalledAssets: stalledAssetIds.size > 0,
|
||||
// 筛选
|
||||
searchText,
|
||||
setSearchText,
|
||||
|
||||
@@ -0,0 +1,97 @@
|
||||
/**
|
||||
* 微信绑定回调页(已登录用户在设置页发起"绑定微信"扫码后回到这里)
|
||||
* 用 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 { 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 校验:绑定场景由设置页生成并落库,前缀 bind:
|
||||
const savedState = localStorage.getItem("wechat_bind_state")
|
||||
if (!savedState || savedState !== state) {
|
||||
setError("安全校验失败,请重新绑定")
|
||||
return
|
||||
}
|
||||
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 {
|
||||
navigate("/app/profile?wechat_bind=failed", { replace: true })
|
||||
}
|
||||
}
|
||||
|
||||
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
|
||||
@@ -1,19 +1,20 @@
|
||||
/**
|
||||
* 微信登录回调页
|
||||
* 扫码授权后由微信重定向回来:用 code 换登录态,
|
||||
* 新用户/资料未完善 → 跳昵称引导页;老用户 → 回来源页/首页
|
||||
*/
|
||||
import React, { useEffect, useState } from "react"
|
||||
import { useSearchParams, useNavigate } from "react-router-dom"
|
||||
import { Spin, message } from "antd"
|
||||
import { Spin } from "antd"
|
||||
import { wechatCallback, getCurrentUser, normalizeUser, type User } from "@/api/auth"
|
||||
import { useAuthStore } from "@/store/authStore"
|
||||
import BindContactModal from "@/components/auth/BindContactModal"
|
||||
import { scheduleProactiveRefresh } from "@/api/auth/tokenRefresh"
|
||||
|
||||
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(() => {
|
||||
@@ -49,20 +50,21 @@ const WechatCallback: React.FC = () => {
|
||||
const userData = await getCurrentUser()
|
||||
const user: User = normalizeUser(userData)
|
||||
setAuth(user, result.access_token, result.refresh_token)
|
||||
scheduleProactiveRefresh()
|
||||
|
||||
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 needOnboarding = result.is_new_user || user.profile_completed === false
|
||||
if (needOnboarding) {
|
||||
navigate("/welcome/wechat", { replace: true })
|
||||
return
|
||||
}
|
||||
} catch (err) {
|
||||
setError("登录失败,请重试")
|
||||
|
||||
// 老用户:回登录前页面或首页
|
||||
const redirect = localStorage.getItem("login_redirect") || "/"
|
||||
localStorage.removeItem("login_redirect")
|
||||
navigate(redirect, { replace: true })
|
||||
} catch {
|
||||
setError("微信登录失败,请重试")
|
||||
setLoading(false)
|
||||
}
|
||||
}
|
||||
@@ -70,21 +72,6 @@ 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
|
||||
@@ -98,49 +85,39 @@ const WechatCallback: React.FC = () => {
|
||||
>
|
||||
<div style={{ textAlign: "center" }}>
|
||||
<Spin size="large" />
|
||||
<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>
|
||||
<p style={{ marginTop: 16, color: "#666" }}>微信登录中...</p>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
return (
|
||||
<BindContactModal
|
||||
open={showBindModal}
|
||||
onSuccess={handleBindSuccess}
|
||||
onCancel={handleBindCancel}
|
||||
/>
|
||||
<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>
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,102 @@
|
||||
/**
|
||||
* 微信新用户昵称引导页
|
||||
* 新微信用户首次登录后强制填写昵称,完成后才进入主界面
|
||||
*/
|
||||
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"
|
||||
initialValues={{ display_name: user?.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,11 +1,12 @@
|
||||
/**
|
||||
* 智能剪辑页面(Issue #1677 多视频批量生成)
|
||||
* 5 步向导:选择模板(弹数量) → 素材 → 配音 → 标题(预览+确认生成) → 封面
|
||||
* 智能剪辑页面(Issue #1677 多视频批量生成,修正版)
|
||||
* 固定 6 步向导:模板(弹数量) → 素材 → 配音 → 标题 → 确认生成 → 封面,单视频与批量完全一致
|
||||
*
|
||||
* 架构:
|
||||
* - N=1:前端 Canvas 实时预览(FrontendPreviewPlayer),零回归
|
||||
* - N>1:服务器批量预览(POST /generation/preview?preview_count=N),
|
||||
* N 个变体分别轮询,网格展示、独立可播放、CSS 标题浮层实时叠加、勾选批量生成
|
||||
* - 预览全部为纯前端 Canvas 实时播放(FrontendPreviewPlayer),不调任何后端渲染接口:
|
||||
* N=1 单播放器;N>1 CanvasPreviewGrid(variantSeed 让素材排布/起始点不同,画面有差异)
|
||||
* - 步骤5确认生成:正式生成接口(count + titles[]/voice_library_ids[]/cover_urls[]),
|
||||
* 批量时逐任务独立进度/失败重试(BatchGenerationGrid)
|
||||
*/
|
||||
import React, { useMemo, useState, useEffect, useRef, useCallback } from "react"
|
||||
import { message } from "antd"
|
||||
@@ -16,7 +17,7 @@ import { useCloneProgress } from "@/hooks/useCloneProgress"
|
||||
import CloneModal from "@/components/voice/CloneModal"
|
||||
import GenerateHeader from "./components/GenerateHeader"
|
||||
import FrontendPreviewPlayer from "./components/FrontendPreviewPlayer"
|
||||
import ServerPreviewGrid from "./components/ServerPreviewGrid"
|
||||
import CanvasPreviewGrid from "./components/CanvasPreviewGrid"
|
||||
import PreviewCountModal from "./components/PreviewCountModal"
|
||||
import GenerateStepsBar from "./components/GenerateStepsBar"
|
||||
import GenerateStepContent from "./components/GenerateStepContent"
|
||||
@@ -24,12 +25,11 @@ import GenerateStepActions from "./components/GenerateStepActions"
|
||||
import { useGenerateFormState } from "./hooks/useGenerateFormState"
|
||||
import { useStepNavigation } from "./hooks/useStepNavigation"
|
||||
import { useGenerateVideo } from "./hooks/useGenerateVideo"
|
||||
import { useBatchPreview } from "./hooks/useBatchPreview"
|
||||
|
||||
import { usePreviewAssets } from "./hooks/usePreviewAssets"
|
||||
import { useTitleStyleUpdaters } from "./hooks/useStep4Title/useTitleStyleUpdaters"
|
||||
import { getAssetsByKind } from "@/api/assets"
|
||||
import { previewTts } from "@/api/tts"
|
||||
import { calculateResolution } from "./utils/calculateResolution"
|
||||
import "./generate.css"
|
||||
|
||||
const GeneratePage: React.FC = () => {
|
||||
@@ -132,6 +132,8 @@ 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] || "" : ""
|
||||
|
||||
useEffect(() => {
|
||||
const voiceAsset = voiceMaterials.find((m) => m.id === selectedVoice)
|
||||
@@ -140,8 +142,10 @@ const GeneratePage: React.FC = () => {
|
||||
return
|
||||
}
|
||||
|
||||
// 批量模式下 TTS 文案跟随变体0标题;单视频跟随主标题
|
||||
const ttsTitle = isBatch ? variant0Title || "" : titleSettings.title
|
||||
const voiceId = selectedClonedVoice || selectedVoice
|
||||
if (!voiceId || !titleSettings.title) {
|
||||
if (!voiceId || !ttsTitle) {
|
||||
setPreviewVoiceAudioUrl(null)
|
||||
return
|
||||
}
|
||||
@@ -151,7 +155,7 @@ const GeneratePage: React.FC = () => {
|
||||
ttsAbortRef.current = controller
|
||||
let cancelled = false
|
||||
|
||||
previewTts({ text: titleSettings.title, voice_id: voiceId })
|
||||
previewTts({ text: ttsTitle, voice_id: voiceId })
|
||||
.then((res) => {
|
||||
if (!cancelled && res.audio_url) {
|
||||
setPreviewVoiceAudioUrl(res.audio_url)
|
||||
@@ -168,7 +172,14 @@ const GeneratePage: React.FC = () => {
|
||||
cancelled = true
|
||||
controller.abort()
|
||||
}
|
||||
}, [selectedVoice, selectedClonedVoice, titleSettings.title, voiceMaterials])
|
||||
}, [
|
||||
selectedVoice,
|
||||
selectedClonedVoice,
|
||||
titleSettings.title,
|
||||
variant0Title,
|
||||
isBatch,
|
||||
voiceMaterials,
|
||||
])
|
||||
|
||||
/* ── 克隆声音 ── */
|
||||
const { addClone } = useCloneProgress()
|
||||
@@ -207,76 +218,8 @@ const GeneratePage: React.FC = () => {
|
||||
previewAssetsEnabled,
|
||||
)
|
||||
|
||||
/* ── 预览就绪 ── */
|
||||
const singlePreviewReady = useMemo(
|
||||
() => previewAssetsReady && !!currentTemplate,
|
||||
[previewAssetsReady, currentTemplate],
|
||||
)
|
||||
|
||||
/* ── 批量服务器预览(N>1) ── */
|
||||
const buildPreviewRequest = useCallback(() => {
|
||||
const { width, height } = calculateResolution(videoRatio || "9:16")
|
||||
const voiceLibraryId =
|
||||
voiceMode === "clone" ? selectedClonedVoice || selectedVoice || "" : selectedVoice || ""
|
||||
return {
|
||||
template_id: selectedTemplate,
|
||||
asset_ids: previewAssetIds,
|
||||
output_width: width,
|
||||
output_height: height,
|
||||
video_ratio: videoRatio,
|
||||
voice_library_id: voiceLibraryId,
|
||||
...(voiceModePerVideo && voiceLibraryIds.some(Boolean)
|
||||
? { voice_library_ids: voiceLibraryIds.map((id) => id || voiceLibraryId) }
|
||||
: {}),
|
||||
preview_count: previewCount,
|
||||
// 批量预览不传 titles/title_config:标题文字与样式由前端 CSS 浮层实时叠加
|
||||
// (用户改标题/样式即时可见,无需重渲染);正式生成时才把标题烧录进成片
|
||||
duration: duration || undefined,
|
||||
bgm_config: {
|
||||
enabled: bgm !== false,
|
||||
...(bgmConfig?.music_id ? { preset_id: bgmConfig.music_id } : {}),
|
||||
},
|
||||
...(storedSourceEditPlanId || sourceEditPlanId
|
||||
? { source_edit_plan_id: storedSourceEditPlanId || sourceEditPlanId || undefined }
|
||||
: {}),
|
||||
}
|
||||
}, [
|
||||
videoRatio,
|
||||
voiceMode,
|
||||
selectedClonedVoice,
|
||||
selectedVoice,
|
||||
selectedTemplate,
|
||||
previewAssetIds,
|
||||
voiceModePerVideo,
|
||||
voiceLibraryIds,
|
||||
previewCount,
|
||||
duration,
|
||||
bgm,
|
||||
bgmConfig,
|
||||
storedSourceEditPlanId,
|
||||
sourceEditPlanId,
|
||||
])
|
||||
|
||||
const {
|
||||
variants,
|
||||
status: batchPreviewStatus,
|
||||
progress: batchPreviewProgress,
|
||||
failedCount: batchFailedCount,
|
||||
trigger: retryBatchPreview,
|
||||
} = useBatchPreview({
|
||||
enabled: isBatch && currentStep >= 4 && previewAssetIds.length > 0 && !!selectedTemplate,
|
||||
buildRequest: buildPreviewRequest,
|
||||
onPreviewTasksCreated: (_taskIds, planId) => {
|
||||
if (planId) setStoredSourceEditPlanId(planId)
|
||||
},
|
||||
})
|
||||
|
||||
/** 批量预览就绪:全部变体渲染完成 */
|
||||
const batchPreviewReady =
|
||||
isBatch && variants.length > 0 && variants.every((v) => v.status === "ready")
|
||||
|
||||
/** 步骤4整体预览就绪状态 */
|
||||
const previewReady = isBatch ? batchPreviewReady : singlePreviewReady
|
||||
/* ── 预览就绪:纯前端 Canvas 预览,素材详情加载完即可秒开(单视频/批量一致) ── */
|
||||
const previewReady = previewAssetsReady && !!currentTemplate
|
||||
|
||||
/* ── 勾选变体 ── */
|
||||
const toggleVariantSelect = useCallback(
|
||||
@@ -296,8 +239,10 @@ const GeneratePage: React.FC = () => {
|
||||
generated,
|
||||
generateError,
|
||||
generatedVideos,
|
||||
batchTasks,
|
||||
generate: handleGenerate,
|
||||
retry: handleRetryGenerate,
|
||||
retryBatchTask: handleRetryBatchTask,
|
||||
dismissError: handleDismissError,
|
||||
download: handleDownload,
|
||||
share: handleShare,
|
||||
@@ -365,31 +310,43 @@ const GeneratePage: React.FC = () => {
|
||||
],
|
||||
)
|
||||
|
||||
/* ── 步骤4「确认生成视频」 ── */
|
||||
/* ── 步骤4「确认生成视频」:校验通过 → 创建正式生成任务 → 跳步骤5看实时进展 ── */
|
||||
const handleConfirmGenerate = useCallback(async () => {
|
||||
// 标题校验
|
||||
if (previewTitles.some((t) => !t?.trim())) {
|
||||
message.warning("请为每个视频输入标题")
|
||||
return
|
||||
}
|
||||
if (isBatch && selectedVariantIds.length === 0) {
|
||||
message.warning("请至少勾选一个视频")
|
||||
return
|
||||
// 标题校验:批量只校验已勾选的变体;单视频校验主标题
|
||||
if (isBatch) {
|
||||
if (selectedVariantIds.length === 0) {
|
||||
message.warning("请至少勾选一个视频")
|
||||
return
|
||||
}
|
||||
const missing = selectedVariantIds.some((i) => !previewTitles[i]?.trim())
|
||||
if (missing) {
|
||||
message.warning("请为每个勾选的视频输入标题")
|
||||
return
|
||||
}
|
||||
} else if (!titleSettings.title?.trim()) {
|
||||
// 与 buildPayload.validateGenerateInputs 一致:AI 自动选标题模式(aiAutoSelect)
|
||||
// 允许空标题由后端生成;手动模式必须填写,避免提交空标题
|
||||
if (!titleSettings.aiAutoSelect) {
|
||||
message.warning("请先选择或输入标题")
|
||||
return
|
||||
}
|
||||
}
|
||||
if (!previewReady) {
|
||||
message.warning("预览视频正在加载,请稍候")
|
||||
message.warning("预览素材正在加载,请稍候")
|
||||
return
|
||||
}
|
||||
const ok = await handleGenerate()
|
||||
if (ok && !isBatch) {
|
||||
if (ok) {
|
||||
// 单视频与批量一致:任务创建成功后进入步骤5「确认生成」看实时渲染进展
|
||||
setCurrentStep(5)
|
||||
}
|
||||
// 批量模式停留在步骤4,右侧网格显示生成进度,完成后点"下一步"进封面
|
||||
}, [
|
||||
isBatch,
|
||||
selectedVariantIds.length,
|
||||
previewReady,
|
||||
selectedVariantIds,
|
||||
previewTitles,
|
||||
titleSettings.aiAutoSelect,
|
||||
titleSettings.title,
|
||||
previewReady,
|
||||
handleGenerate,
|
||||
setCurrentStep,
|
||||
])
|
||||
@@ -403,25 +360,20 @@ const GeneratePage: React.FC = () => {
|
||||
selectedMaterials,
|
||||
smartSelectedIds,
|
||||
titleSettings,
|
||||
previewReady,
|
||||
generated,
|
||||
previewTitles,
|
||||
selectedCount: isBatch ? selectedVariantIds.length : 1,
|
||||
onOpenCountModal: () => setCountModalOpen(true),
|
||||
})
|
||||
|
||||
/* ── 最终成片(单视频右侧播放) ── */
|
||||
const finalVideo = generatedVideos[0]
|
||||
|
||||
/** 批量生成进度文案 */
|
||||
const batchGeneratingText = useMemo(() => {
|
||||
if (batchPreviewStatus === "loading")
|
||||
return `AI 正在渲染 ${previewCount} 个预览视频… ${batchPreviewProgress}%`
|
||||
if (batchPreviewStatus === "failed") return "预览渲染失败,请重试"
|
||||
if (batchPreviewStatus === "partial_failed")
|
||||
return `${batchFailedCount} 个预览失败,可重新生成或勾选成功的视频`
|
||||
return ""
|
||||
}, [batchPreviewStatus, batchPreviewProgress, batchFailedCount, previewCount])
|
||||
/* ── 布局 class:步骤4标题页=预览+标题侧栏;步骤5/6批量=整行宽;步骤1~3=整行宽 ── */
|
||||
const layoutClassName = useMemo(() => {
|
||||
if (currentStep < 4) return "xx-generate-layout full-width"
|
||||
if (currentStep === 4) return "xx-generate-layout step4-layout"
|
||||
// 步骤5/6:批量网格需要整行宽度;单视频保持 表单+右侧成片 两栏
|
||||
return isBatch ? "xx-generate-layout full-width" : "xx-generate-layout"
|
||||
}, [currentStep, isBatch])
|
||||
|
||||
/* ================================================================
|
||||
渲染
|
||||
@@ -433,16 +385,12 @@ const GeneratePage: React.FC = () => {
|
||||
|
||||
<GenerateStepsBar currentStep={currentStep} onStepClick={setCurrentStep} />
|
||||
|
||||
<div
|
||||
className={`xx-generate-layout${currentStep < 4 ? " full-width" : ""}${
|
||||
currentStep === 4 ? " step4-layout" : ""
|
||||
}`}
|
||||
>
|
||||
{/* ════ 步骤4:左侧预览大区域 ════ */}
|
||||
<div className={layoutClassName}>
|
||||
{/* ════ 步骤4:左侧预览大区域(纯前端 Canvas 实时预览) ════ */}
|
||||
{currentStep === 4 && !!currentTemplate && (
|
||||
<div className="xx-generate-preview-col">
|
||||
{!isBatch ? (
|
||||
/* 单视频:前端 Canvas 实时预览(与旧版一致) */
|
||||
/* 单视频:前端 Canvas 实时预览(与旧版一致,零回归) */
|
||||
<FrontendPreviewPlayer
|
||||
assets={previewAssets}
|
||||
template={currentTemplate}
|
||||
@@ -466,109 +414,32 @@ const GeneratePage: React.FC = () => {
|
||||
onTitlePositionChange={styleUpdaters.updateTitlePosition}
|
||||
/>
|
||||
) : (
|
||||
/* 批量:服务器预览网格 */
|
||||
/* 批量:N 个前端 Canvas 预览网格(不调任何后端渲染接口,秒开) */
|
||||
<div className="xx-form-section">
|
||||
<div className="xx-preview-header">
|
||||
<h3>🎬 {previewCount} 个视频预览</h3>
|
||||
{batchPreviewStatus === "loading" && (
|
||||
<span style={{ fontSize: 13, color: "var(--text-secondary, #666)" }}>
|
||||
{batchPreviewProgress}%
|
||||
</span>
|
||||
)}
|
||||
{(batchPreviewStatus === "failed" || batchPreviewStatus === "partial_failed") && (
|
||||
<button
|
||||
type="button"
|
||||
className="xx-btn xx-btn-ghost xx-btn-sm"
|
||||
onClick={retryBatchPreview}
|
||||
>
|
||||
🔄 重新生成预览
|
||||
</button>
|
||||
)}
|
||||
<span style={{ fontSize: 13, color: "var(--text-secondary, #666)" }}>
|
||||
实时预览,勾选要生成的视频
|
||||
</span>
|
||||
</div>
|
||||
|
||||
{batchGeneratingText && (
|
||||
<div
|
||||
style={{
|
||||
fontSize: 13,
|
||||
color:
|
||||
batchPreviewStatus === "failed"
|
||||
? "var(--error-color, #ef4444)"
|
||||
: "var(--text-secondary, #666)",
|
||||
marginBottom: 12,
|
||||
}}
|
||||
>
|
||||
{batchGeneratingText}
|
||||
</div>
|
||||
)}
|
||||
|
||||
<ServerPreviewGrid
|
||||
variants={variants}
|
||||
<CanvasPreviewGrid
|
||||
count={previewCount}
|
||||
assets={previewAssets}
|
||||
template={currentTemplate}
|
||||
videoRatio={videoRatio}
|
||||
titles={previewTitles}
|
||||
titleStyle={{
|
||||
position: titleSettings.position,
|
||||
color: titleSettings.color,
|
||||
size: titleSettings.size,
|
||||
}}
|
||||
titleSettings={titleSettings}
|
||||
voiceAudioUrl={previewVoiceAudioUrl || undefined}
|
||||
selectedIds={selectedVariantIds}
|
||||
onToggleSelect={toggleVariantSelect}
|
||||
selectable={!generating}
|
||||
/>
|
||||
|
||||
{/* 生成中进度(批量) */}
|
||||
{generating && (
|
||||
<div className="xx-gen-progress-card" style={{ marginTop: 16 }}>
|
||||
<div className="xx-gen-progress-header">
|
||||
<div className="xx-gen-progress-info">
|
||||
<div className="xx-gen-progress-phase">
|
||||
⏳ 正在渲染 {selectedVariantIds.length} 个最终视频… {Math.round(progress)}
|
||||
%
|
||||
</div>
|
||||
<div className="xx-gen-progress-sub">
|
||||
生成过程中可以切换到其他页面,完成后可在任务历史查看
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<div className="xx-gen-progress-bar">
|
||||
<div
|
||||
className="xx-gen-progress-bar-fill"
|
||||
style={{ width: `${Math.min(Math.round(progress), 100)}%` }}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{generateError && !generating && (
|
||||
<div className="xx-gen-error-card" style={{ marginTop: 16 }}>
|
||||
<div className="xx-gen-error-info">
|
||||
<div className="xx-gen-error-title">生成失败</div>
|
||||
<div className="xx-gen-error-msg">{generateError}</div>
|
||||
</div>
|
||||
<button
|
||||
type="button"
|
||||
className="xx-btn xx-btn-primary xx-btn-sm"
|
||||
onClick={handleRetryGenerate}
|
||||
>
|
||||
🔄 重试
|
||||
</button>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{generated && !generating && (
|
||||
<div className="xx-gen-success-card" style={{ marginTop: 16 }}>
|
||||
<div className="xx-gen-success-info">
|
||||
<div className="xx-gen-success-title">✅ 视频生成完成!</div>
|
||||
<div className="xx-gen-success-sub">
|
||||
共生成 {generatedVideos.length} 条视频,点击「下一步」为每个视频选择封面
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* ════ 右侧:步骤1~3 表单 / 步骤4 标题边栏 / 步骤5 封面 ════ */}
|
||||
{/* ════ 右侧:步骤1~3 表单 / 步骤4 标题边栏 / 步骤5 确认生成进度 / 步骤6 封面 ════ */}
|
||||
<div className="xx-generate-form">
|
||||
<GenerateStepContent
|
||||
currentStep={currentStep}
|
||||
@@ -606,7 +477,9 @@ const GeneratePage: React.FC = () => {
|
||||
progress={progress}
|
||||
generatedVideos={generatedVideos}
|
||||
onRetry={handleRetryGenerate}
|
||||
onRetryBatchTask={handleRetryBatchTask}
|
||||
onDismissError={handleDismissError}
|
||||
batchTasks={batchTasks}
|
||||
previewCount={previewCount}
|
||||
previewTitles={previewTitles}
|
||||
onPreviewTitlesChange={setPreviewTitles}
|
||||
@@ -631,16 +504,16 @@ const GeneratePage: React.FC = () => {
|
||||
/>
|
||||
</div>
|
||||
|
||||
{/* ════ 步骤5(封面):成片播放器(单视频) ════ */}
|
||||
{currentStep === 5 && !isBatch && generated && finalVideo && (
|
||||
{/* ════ 步骤5/6(单视频):右侧成片播放器 ════ */}
|
||||
{currentStep >= 5 && !isBatch && generated && finalVideo && (
|
||||
<div className="xx-generate-right-col">
|
||||
<div className="xx-inline-video-player">
|
||||
<video
|
||||
src={finalVideo.download_url || finalVideo.file_url}
|
||||
controls
|
||||
autoPlay
|
||||
autoPlay={currentStep === 5}
|
||||
style={{ width: "100%", maxHeight: "70vh", objectFit: "contain", borderRadius: 12 }}
|
||||
poster={finalVideo.thumbnail_url}
|
||||
poster={finalVideo.thumbnail_url || undefined}
|
||||
/>
|
||||
<div style={{ display: "flex", gap: 8, marginTop: 12, justifyContent: "center" }}>
|
||||
<button className="xx-btn xx-btn-ghost xx-btn-sm" onClick={handleDownload}>
|
||||
|
||||
@@ -0,0 +1,100 @@
|
||||
/**
|
||||
* 第5步「确认生成」— 批量渲染进度网格(Issue #1677)
|
||||
*
|
||||
* N 个正式生成任务各自独立卡片:进度条 / 成功成片播放 / 失败原因 + 单独重试。
|
||||
* 数据来自 useGenerateVideo 的 batchTasks(useGenerationPolling 实时回传)。
|
||||
*/
|
||||
import React from "react"
|
||||
import { LoadingOutlined, CheckCircleFilled, CloseCircleOutlined } from "@ant-design/icons"
|
||||
import type { BatchTaskState } from "../hooks/generate-video/useGenerationPolling"
|
||||
import type { GeneratedVideo } from "@/api/template-editor"
|
||||
|
||||
interface BatchGenerationGridProps {
|
||||
tasks: BatchTaskState[]
|
||||
/** 变体标题(按变体序号取) */
|
||||
titles: string[]
|
||||
/** 失败任务重试 */
|
||||
onRetryTask: (taskId: string) => void
|
||||
}
|
||||
|
||||
const BatchGenerationGrid: React.FC<BatchGenerationGridProps> = ({
|
||||
tasks,
|
||||
titles,
|
||||
onRetryTask,
|
||||
}) => {
|
||||
const sorted = [...tasks].sort(
|
||||
(a, b) => Number(a.variantIndex || 0) - Number(b.variantIndex || 0),
|
||||
)
|
||||
|
||||
return (
|
||||
<div className="xx-form-section">
|
||||
<div className="xx-preview-header">
|
||||
<h3>🎬 正在生成 {tasks.length} 个视频</h3>
|
||||
<span style={{ fontSize: 13, color: "var(--text-secondary, #666)" }}>
|
||||
完成 {tasks.filter((t) => t.status === "completed").length} / {tasks.length}
|
||||
</span>
|
||||
</div>
|
||||
<div className="xx-batch-gen-grid">
|
||||
{sorted.map((task) => {
|
||||
const title = titles[task.variantIndex] || `视频 ${task.variantIndex + 1}`
|
||||
const video = (task.videos?.[0] || null) as GeneratedVideo | null
|
||||
return (
|
||||
<div key={task.taskId} className={`xx-batch-gen-card status-${task.status}`}>
|
||||
<div className="xx-batch-gen-card-head">
|
||||
<span className="xx-batch-gen-card-title" title={title}>
|
||||
{task.status === "completed" ? (
|
||||
<CheckCircleFilled style={{ color: "#52c41a", marginRight: 6 }} />
|
||||
) : task.status === "failed" ? (
|
||||
<CloseCircleOutlined style={{ color: "#ef4444", marginRight: 6 }} />
|
||||
) : (
|
||||
<LoadingOutlined style={{ color: "#1677ff", marginRight: 6 }} />
|
||||
)}
|
||||
视频 {task.variantIndex + 1}:{title}
|
||||
</span>
|
||||
</div>
|
||||
|
||||
<div className="xx-batch-gen-card-body">
|
||||
{task.status === "running" && (
|
||||
<>
|
||||
<div className="xx-gen-progress-bar">
|
||||
<div
|
||||
className="xx-gen-progress-bar-fill"
|
||||
style={{ width: `${Math.min(task.progress, 100)}%` }}
|
||||
/>
|
||||
</div>
|
||||
<div className="xx-batch-gen-card-pct">{Math.round(task.progress)}%</div>
|
||||
</>
|
||||
)}
|
||||
{task.status === "completed" && video && (
|
||||
<video
|
||||
src={video.download_url || video.file_url}
|
||||
controls
|
||||
style={{ width: "100%", borderRadius: 8, background: "#000", maxHeight: 280 }}
|
||||
poster={video.thumbnail_url}
|
||||
/>
|
||||
)}
|
||||
{task.status === "completed" && !video && (
|
||||
<div className="xx-batch-gen-card-done">✅ 已完成(成片可在下一步选择封面)</div>
|
||||
)}
|
||||
{task.status === "failed" && (
|
||||
<div className="xx-batch-gen-card-failed">
|
||||
<div className="xx-batch-gen-card-err">{task.error || "生成失败"}</div>
|
||||
<button
|
||||
type="button"
|
||||
className="xx-btn xx-btn-primary xx-btn-sm"
|
||||
onClick={() => onRetryTask(task.taskId)}
|
||||
>
|
||||
🔄 重试此视频
|
||||
</button>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default BatchGenerationGrid
|
||||
@@ -0,0 +1,97 @@
|
||||
/**
|
||||
* 批量前端 Canvas 实时预览网格(Issue #1677 修正方案)
|
||||
*
|
||||
* N 个 FrontendPreviewPlayer 网格排列:
|
||||
* - 纯前端 Canvas + video 元素实时播放素材片段,不调任何后端渲染接口
|
||||
* - variantSeed 让每个变体素材排布/起始点不同,画面有可见差异
|
||||
* - 各自叠加独立标题浮层(variantTitle),标题样式全局共用
|
||||
* - 勾选框决定提交时生成哪些变体
|
||||
*/
|
||||
import React from "react"
|
||||
import type { AssetItem } from "@/api/assets"
|
||||
import type { EditingTemplate } from "@/api/editing-planner"
|
||||
import type { TitleSettings } from "../types"
|
||||
import FrontendPreviewPlayer from "./FrontendPreviewPlayer"
|
||||
|
||||
interface CanvasPreviewGridProps {
|
||||
count: number
|
||||
assets: AssetItem[]
|
||||
template: EditingTemplate | null
|
||||
videoRatio: string
|
||||
titles: string[]
|
||||
titleSettings: TitleSettings
|
||||
/** 共用配音预览音频(仅第 1 个变体播放,避免多路音频重叠) */
|
||||
voiceAudioUrl?: string
|
||||
/** 勾选的变体序号 */
|
||||
selectedIds: number[]
|
||||
onToggleSelect: (index: number) => void
|
||||
/** 生成中禁止勾选 */
|
||||
selectable?: boolean
|
||||
}
|
||||
|
||||
const CanvasPreviewGrid: React.FC<CanvasPreviewGridProps> = ({
|
||||
count,
|
||||
assets,
|
||||
template,
|
||||
videoRatio,
|
||||
titles,
|
||||
titleSettings,
|
||||
voiceAudioUrl,
|
||||
selectedIds,
|
||||
onToggleSelect,
|
||||
selectable = true,
|
||||
}) => {
|
||||
// count 上限已在源头 PreviewCountModal 的数量选择(1~MAX_PREVIEW_COUNT=10)clamp,
|
||||
// 这里完整渲染所有变体,保证每个变体都有勾选/预览入口,UI 与数据不脱节
|
||||
return (
|
||||
<div className="xx-canvas-grid">
|
||||
{Array.from({ length: count }, (_, i) => {
|
||||
const checked = selectedIds.includes(i)
|
||||
return (
|
||||
<div
|
||||
key={i}
|
||||
className={`xx-canvas-grid-card${checked ? " selected" : ""}`}
|
||||
data-variant={i}
|
||||
>
|
||||
<div className="xx-canvas-grid-card-bar">
|
||||
<label className="xx-canvas-grid-check">
|
||||
<input
|
||||
type="checkbox"
|
||||
checked={checked}
|
||||
disabled={!selectable}
|
||||
onChange={() => onToggleSelect(i)}
|
||||
/>
|
||||
<span>视频 {i + 1}</span>
|
||||
</label>
|
||||
</div>
|
||||
<FrontendPreviewPlayer
|
||||
assets={assets}
|
||||
template={template}
|
||||
videoRatio={videoRatio}
|
||||
ready={assets.length > 0}
|
||||
variantSeed={i + 1}
|
||||
variantTitle={titles[i] || ""}
|
||||
voiceAudioUrl={i === 0 ? voiceAudioUrl : undefined}
|
||||
compact
|
||||
titleSettings={{
|
||||
title: titles[i] || "",
|
||||
size: titleSettings.size,
|
||||
font: titleSettings.font,
|
||||
color: titleSettings.color,
|
||||
position: titleSettings.position as "top" | "center" | "bottom" | "custom",
|
||||
bold: titleSettings.bold,
|
||||
italic: titleSettings.italic,
|
||||
stroke: titleSettings.stroke,
|
||||
shadow: titleSettings.shadow,
|
||||
posX: titleSettings.posX,
|
||||
posY: titleSettings.posY,
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default CanvasPreviewGrid
|
||||
@@ -41,6 +41,16 @@ interface FrontendPreviewPlayerProps {
|
||||
posY?: number | null
|
||||
}
|
||||
onTitlePositionChange?: (posX: number, posY: number) => void
|
||||
/**
|
||||
* 变体种子(批量生成 #1677):同一批素材在不同变体中采用不同的素材顺序与
|
||||
* 片段起始点,让 N 个 Canvas 预览画面有差异(纯前端随机剪辑模拟,不调后端)。
|
||||
* 0 / 不传 = 单视频,排布与旧版完全一致(零回归)。
|
||||
*/
|
||||
variantSeed?: number
|
||||
/** 变体标题文字(批量时每个预览独立标题,叠加在画面上);不传用 titleSettings.title */
|
||||
variantTitle?: string
|
||||
/** 紧凑模式(批量网格中使用,缩小内边距/标题尺寸) */
|
||||
compact?: boolean
|
||||
}
|
||||
|
||||
function formatTime(seconds: number): string {
|
||||
@@ -52,10 +62,23 @@ function formatTime(seconds: number): string {
|
||||
/**
|
||||
* 将素材映射为播放片段(复用原逻辑)
|
||||
*/
|
||||
/** 简单可复现随机数(mulberry32),同一种子产出稳定排布,避免每次渲染抖动 */
|
||||
function seededRandom(seed: number): () => number {
|
||||
let a = seed >>> 0
|
||||
return () => {
|
||||
a |= 0
|
||||
a = (a + 0x6d2b79f5) | 0
|
||||
let t = Math.imul(a ^ (a >>> 15), 1 | a)
|
||||
t = (t + Math.imul(t ^ (t >>> 7), 61 | t)) ^ t
|
||||
return ((t ^ (t >>> 14)) >>> 0) / 4294967296
|
||||
}
|
||||
}
|
||||
|
||||
function buildPlaybackSegments(
|
||||
assets: AssetItem[],
|
||||
template: EditingTemplate | null,
|
||||
serverClips?: EditPlanClip[],
|
||||
variantSeed = 0,
|
||||
): PlaybackSegment[] {
|
||||
if (!assets.length) return []
|
||||
|
||||
@@ -79,18 +102,40 @@ function buildPlaybackSegments(
|
||||
}
|
||||
}
|
||||
|
||||
// Fallback: 本地构建片段(与旧行为一致)
|
||||
// Fallback: 本地构建片段
|
||||
// variantSeed=0(单视频):与旧行为完全一致(素材原序、起始点 0),零回归
|
||||
// variantSeed>0(批量变体):素材顺序按种子轮换 + 片段起始点在素材内偏移,
|
||||
// 模拟后端"AI 随机剪辑出不同版本",让 N 个预览画面有可见差异
|
||||
const templateSegments = template?.segments || []
|
||||
const segments: PlaybackSegment[] = []
|
||||
const orderedAssets = variantSeed > 0 ? [...assets] : assets
|
||||
if (variantSeed > 0 && orderedAssets.length > 1) {
|
||||
const rand = seededRandom(variantSeed * 7919 + 13)
|
||||
// 素材轮换:把数组旋转 (seed % n) 位,再对后半段做一次稳定交换
|
||||
const n = orderedAssets.length
|
||||
const rotate = variantSeed % n
|
||||
orderedAssets.push(...orderedAssets.splice(0, rotate))
|
||||
const swapA = Math.floor(rand() * n)
|
||||
const swapB = Math.floor(rand() * n)
|
||||
if (swapA !== swapB) {
|
||||
;[orderedAssets[swapA], orderedAssets[swapB]] = [orderedAssets[swapB], orderedAssets[swapA]]
|
||||
}
|
||||
}
|
||||
|
||||
assets.forEach((asset, i) => {
|
||||
orderedAssets.forEach((asset, i) => {
|
||||
const assetDuration = asset.duration || asset.metadata?.duration || 30
|
||||
const tplSeg = templateSegments[i] || templateSegments[templateSegments.length - 1]
|
||||
const segDuration = tplSeg
|
||||
? Math.min(tplSeg.duration_max, Math.max(tplSeg.duration_min, assetDuration))
|
||||
: Math.min(assetDuration, 10)
|
||||
|
||||
const startTime = 0
|
||||
let startTime = 0
|
||||
if (variantSeed > 0 && assetDuration - segDuration > 1) {
|
||||
const rand = seededRandom(variantSeed * 104729 + i * 31 + 7)
|
||||
// 起始点在素材可用区间内随机偏移(至少留 0.5s 余量)
|
||||
const maxStart = Math.max(0, assetDuration - segDuration - 0.5)
|
||||
startTime = Math.round(rand() * maxStart * 10) / 10
|
||||
}
|
||||
const endTime = Math.min(startTime + segDuration, assetDuration)
|
||||
const videoUrl = asset.file_url || asset.storage_key
|
||||
|
||||
@@ -109,11 +154,16 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
|
||||
voiceAudioUrl,
|
||||
titleSettings,
|
||||
onTitlePositionChange,
|
||||
variantSeed = 0,
|
||||
variantTitle,
|
||||
compact = false,
|
||||
}) => {
|
||||
const segments = useMemo(
|
||||
() => buildPlaybackSegments(assets, template, serverClips),
|
||||
[assets, template, serverClips],
|
||||
() => buildPlaybackSegments(assets, template, serverClips, variantSeed),
|
||||
[assets, template, serverClips, variantSeed],
|
||||
)
|
||||
// 批量变体:标题文字取 variantTitle,样式仍由全局 titleSettings 控制
|
||||
const effectiveTitle = variantTitle ?? titleSettings?.title
|
||||
|
||||
// ── ASS 坐标系参数(与后端 ass_subtitle_builder.py 一致) ──
|
||||
const TITLE_MARGIN_TOP = 120
|
||||
@@ -234,7 +284,7 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
|
||||
// ── Canvas 播放器(WebCodecs 路径) ──
|
||||
const canvasTitle = titleSettings
|
||||
? {
|
||||
text: titleSettings.title || "标题预览",
|
||||
text: effectiveTitle || "标题预览",
|
||||
fontSize: titleSettings.size,
|
||||
fontFamily: titleSettings.font || "思源黑体",
|
||||
color: titleSettings.color || "#ffffff",
|
||||
@@ -520,13 +570,15 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
|
||||
style={{
|
||||
position: "relative",
|
||||
width: "100%",
|
||||
maxWidth: 280,
|
||||
maxWidth: compact ? "100%" : 280,
|
||||
margin: compact ? 0 : "0 auto",
|
||||
aspectRatio: "9 / 16",
|
||||
background: "#0a0a0a",
|
||||
borderRadius: 24,
|
||||
borderRadius: compact ? 10 : 24,
|
||||
overflow: "hidden",
|
||||
boxShadow:
|
||||
"0 4px 6px -1px rgba(0,0,0,0.3), 0 20px 50px -12px rgba(0,0,0,0.5), inset 0 0 0 1px rgba(255,255,255,0.06)",
|
||||
boxShadow: compact
|
||||
? "inset 0 0 0 1px rgba(255,255,255,0.06)"
|
||||
: "0 4px 6px -1px rgba(0,0,0,0.3), 0 20px 50px -12px rgba(0,0,0,0.5), inset 0 0 0 1px rgba(255,255,255,0.06)",
|
||||
}}
|
||||
>
|
||||
{/* ── Canvas 渲染层(WebCodecs 路径) ── */}
|
||||
@@ -610,8 +662,8 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
|
||||
? { top: "50%", transform: "translate(-50%, -50%)" }
|
||||
: { bottom: `${titleBottomPct}%` }),
|
||||
}),
|
||||
pointerEvents: "auto",
|
||||
cursor: onTitlePositionChange ? "grab" : "default",
|
||||
pointerEvents: onTitlePositionChange && variantSeed === 0 ? "auto" : "none",
|
||||
cursor: onTitlePositionChange && variantSeed === 0 ? "grab" : "default",
|
||||
touchAction: "none",
|
||||
userSelect: "none",
|
||||
WebkitUserSelect: "none",
|
||||
@@ -641,7 +693,7 @@ const FrontendPreviewPlayer: React.FC<FrontendPreviewPlayerProps> = ({
|
||||
: undefined,
|
||||
}}
|
||||
>
|
||||
{titleSettings.title.split(/[//]/).map((part, i) => (
|
||||
{(effectiveTitle || "").split(/[//]/).map((part, i) => (
|
||||
<span key={i}>
|
||||
{i > 0 && <br />}
|
||||
{part}
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
/**
|
||||
* GeneratePage 步骤底部操作按钮(Issue #1677 改造后 5 步)
|
||||
* GeneratePage 步骤底部操作按钮(Issue #1677 修正:固定 6 步)
|
||||
*
|
||||
* 步骤 1~3:上一步 / 下一步
|
||||
* 步骤 4(标题+预览+确认生成):确认生成按钮在右侧边栏底部(含勾选数量),
|
||||
* 渲染中显示进度;生成完成后显示"下一步 → 选择封面"
|
||||
* 步骤 5(选择封面):仅上一步
|
||||
* 步骤 4(选择标题):「✨ 确认生成视频 / 确认生成 N 个视频」→ 创建正式生成任务,成功后跳步骤5
|
||||
* 步骤 5(确认生成):渲染进度页,全部完成后「下一步:选择封面」;仅上一步
|
||||
* 步骤 6(选择封面):仅上一步
|
||||
*/
|
||||
import React from "react"
|
||||
|
||||
@@ -12,7 +12,7 @@ export interface GenerateStepActionsProps {
|
||||
currentStep: number
|
||||
onPrev: () => void
|
||||
onNext: () => void
|
||||
/** 步骤4:确认生成视频(校验 + 创建渲染任务 + 成功后进入步骤5) */
|
||||
/** 步骤4:确认生成视频(校验 + 创建渲染任务) */
|
||||
onConfirmGenerate: () => void | Promise<void>
|
||||
generating: boolean
|
||||
generated: boolean
|
||||
@@ -41,26 +41,19 @@ const GenerateStepActions: React.FC<GenerateStepActionsProps> = ({
|
||||
)
|
||||
}
|
||||
|
||||
/* 步骤 4:标题+预览+确认生成 */
|
||||
/* 步骤 4:选择标题 — 确认生成 */
|
||||
if (currentStep === 4) {
|
||||
if (generating) {
|
||||
return (
|
||||
<button className="xx-btn xx-btn-primary" disabled>
|
||||
⏳ 视频生成中…
|
||||
⏳ 正在提交生成任务…
|
||||
</button>
|
||||
)
|
||||
}
|
||||
if (generateError) {
|
||||
return (
|
||||
<button className="xx-btn xx-btn-primary" onClick={onConfirmGenerate}>
|
||||
🔄 重新生成视频
|
||||
</button>
|
||||
)
|
||||
}
|
||||
if (generated) {
|
||||
return (
|
||||
<button className="xx-btn xx-btn-primary" onClick={onNext}>
|
||||
下一步:选择封面 →
|
||||
🔄 重新生成
|
||||
</button>
|
||||
)
|
||||
}
|
||||
@@ -71,7 +64,23 @@ const GenerateStepActions: React.FC<GenerateStepActionsProps> = ({
|
||||
)
|
||||
}
|
||||
|
||||
/* 步骤 5(封面,最后一步):无主按钮 */
|
||||
/* 步骤 5:确认生成进度页 — 全部完成后下一步进封面 */
|
||||
if (currentStep === 5) {
|
||||
if (generated) {
|
||||
return (
|
||||
<button className="xx-btn xx-btn-primary" onClick={onNext}>
|
||||
下一步:选择封面 →
|
||||
</button>
|
||||
)
|
||||
}
|
||||
return (
|
||||
<button className="xx-btn xx-btn-primary" disabled>
|
||||
⏳ 视频渲染中…
|
||||
</button>
|
||||
)
|
||||
}
|
||||
|
||||
/* 步骤 6(封面,最后一步):无主按钮 */
|
||||
return null
|
||||
}
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
/**
|
||||
* GeneratePage 步骤内容渲染
|
||||
* 步骤顺序(5步,Issue #1677):模板(1) → 素材(2) → 配音(3) → 标题+预览+确认生成(4) → 封面(5)
|
||||
* 步骤顺序(6步,Issue #1677 修正):模板(1) → 素材(2) → 配音(3) → 标题(4) → 确认生成(5) → 封面(6)
|
||||
* 步骤4预览(Canvas 网格)与步骤5进度(批量渲染网格)由 GeneratePage 直接渲染在左侧大区域。
|
||||
*/
|
||||
import React from "react"
|
||||
import type { EditingTemplate } from "@/api/editing-planner"
|
||||
@@ -12,6 +13,8 @@ import Step2MaterialSelect from "../components/Step2MaterialSelect"
|
||||
import Step3VoiceWithMode from "./Step3VoiceWithMode"
|
||||
import Step4TitleSettings from "../components/Step4TitleSettings"
|
||||
import Step6CoverSettings from "../components/Step6CoverSettings"
|
||||
import BatchGenerationGrid from "./BatchGenerationGrid"
|
||||
import type { BatchTaskState } from "../hooks/generate-video/useGenerationPolling"
|
||||
import type { GeneratedVideo } from "@/api/template-editor"
|
||||
|
||||
export interface GenerateStepContentProps {
|
||||
@@ -54,7 +57,10 @@ export interface GenerateStepContentProps {
|
||||
progress: number
|
||||
generatedVideos: GeneratedVideo[]
|
||||
onRetry: () => void
|
||||
onRetryBatchTask: (taskId: string) => void
|
||||
onDismissError: () => void
|
||||
/** 批量:每个正式生成任务的独立状态(步骤5进度网格) */
|
||||
batchTasks: BatchTaskState[]
|
||||
/** BGM 开关 */
|
||||
bgm: boolean
|
||||
/** BGM 配置(来自模板) */
|
||||
@@ -102,7 +108,14 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
selectedVoice,
|
||||
onSelectedVoiceChange,
|
||||
onServerClipsChange,
|
||||
generating,
|
||||
generated,
|
||||
generateError,
|
||||
progress,
|
||||
onRetry,
|
||||
generatedVideos,
|
||||
batchTasks,
|
||||
onRetryBatchTask,
|
||||
previewCount,
|
||||
previewTitles,
|
||||
onPreviewTitlesChange,
|
||||
@@ -176,6 +189,62 @@ export const GenerateStepContent: React.FC<GenerateStepContentProps> = (props) =
|
||||
/>
|
||||
)
|
||||
case 5:
|
||||
/* 确认生成页:批量=逐任务进度网格;单视频=进度状态卡(成片播放器在左侧大区域) */
|
||||
if (previewCount > 1) {
|
||||
return (
|
||||
<BatchGenerationGrid
|
||||
tasks={batchTasks}
|
||||
titles={previewTitles}
|
||||
onRetryTask={onRetryBatchTask}
|
||||
/>
|
||||
)
|
||||
}
|
||||
/* 单视频:渲染进度 / 失败重试 / 完成提示(成片播放器在右侧栏) */
|
||||
return (
|
||||
<div className="xx-form-section">
|
||||
<h3>🎬 确认生成</h3>
|
||||
{generating && (
|
||||
<div className="xx-gen-progress-card">
|
||||
<div className="xx-gen-progress-header">
|
||||
<div className="xx-gen-progress-info">
|
||||
<div className="xx-gen-progress-phase">
|
||||
⏳ 视频渲染中… {Math.round(progress)}%
|
||||
</div>
|
||||
<div className="xx-gen-progress-sub">
|
||||
生成过程中可以切换到其他页面,完成后可在任务历史查看
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<div className="xx-gen-progress-bar">
|
||||
<div
|
||||
className="xx-gen-progress-bar-fill"
|
||||
style={{ width: `${Math.min(Math.round(progress), 100)}%` }}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
{generateError && !generating && (
|
||||
<div className="xx-gen-error-card">
|
||||
<div className="xx-gen-error-info">
|
||||
<div className="xx-gen-error-title">生成失败</div>
|
||||
<div className="xx-gen-error-msg">{generateError}</div>
|
||||
</div>
|
||||
<button type="button" className="xx-btn xx-btn-primary xx-btn-sm" onClick={onRetry}>
|
||||
🔄 重试
|
||||
</button>
|
||||
</div>
|
||||
)}
|
||||
{generated && !generating && (
|
||||
<div className="xx-gen-success-card">
|
||||
<div className="xx-gen-success-info">
|
||||
<div className="xx-gen-success-title">✅ 视频生成完成!</div>
|
||||
<div className="xx-gen-success-sub">右侧可预览成片,点击「下一步」选择封面</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)
|
||||
case 6:
|
||||
return (
|
||||
<Step6CoverSettings
|
||||
coverSettings={coverSettings}
|
||||
|
||||
@@ -1,127 +0,0 @@
|
||||
/**
|
||||
* 批量预览网格(Issue #1677)
|
||||
* N 个服务器渲染的预览视频,网格排列、各自独立播放、CSS 标题浮层实时叠加、勾选框批量选择
|
||||
*/
|
||||
import React from "react"
|
||||
import { LoadingOutlined, CheckCircleFilled, CloseCircleOutlined } from "@ant-design/icons"
|
||||
import type { VariantPreview } from "../hooks/useBatchPreview"
|
||||
|
||||
interface ServerPreviewGridProps {
|
||||
variants: VariantPreview[]
|
||||
/** 每个变体的标题文字(实时叠加浮层) */
|
||||
titles: string[]
|
||||
/** 标题样式(全局共用) */
|
||||
titleStyle: {
|
||||
position: string
|
||||
color: string
|
||||
size: number
|
||||
}
|
||||
/** 勾选的变体索引 */
|
||||
selectedIds: number[]
|
||||
onToggleSelect: (index: number) => void
|
||||
/** 是否显示勾选框(确认生成前) */
|
||||
selectable?: boolean
|
||||
}
|
||||
|
||||
const ServerPreviewGrid: React.FC<ServerPreviewGridProps> = ({
|
||||
variants,
|
||||
titles,
|
||||
titleStyle,
|
||||
selectedIds,
|
||||
onToggleSelect,
|
||||
selectable = true,
|
||||
}) => {
|
||||
if (variants.length === 0) return null
|
||||
|
||||
return (
|
||||
<div className="xx-variant-grid">
|
||||
{variants.map((v) => {
|
||||
const selected = selectedIds.includes(v.index)
|
||||
const titleText = titles[v.index] || ""
|
||||
return (
|
||||
<div
|
||||
key={v.index}
|
||||
className={`xx-variant-card ${selected ? "selected" : ""} ${
|
||||
v.status === "failed" ? "failed" : ""
|
||||
}`}
|
||||
onClick={() => {
|
||||
if (selectable && v.status === "ready") onToggleSelect(v.index)
|
||||
}}
|
||||
role="button"
|
||||
tabIndex={0}
|
||||
>
|
||||
{/* 勾选框 */}
|
||||
{selectable && v.status === "ready" && (
|
||||
<div className={`xx-variant-check ${selected ? "checked" : ""}`}>
|
||||
{selected && "✓"}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 变体序号 */}
|
||||
<div className="xx-variant-index">视频 {v.index + 1}</div>
|
||||
|
||||
{/* 视频区域 */}
|
||||
<div className="xx-variant-video-wrap">
|
||||
{v.status === "loading" && (
|
||||
<div className="xx-variant-loading">
|
||||
<LoadingOutlined style={{ fontSize: 28, color: "#3b82f6" }} />
|
||||
<div className="xx-variant-progress">
|
||||
<div
|
||||
className="xx-variant-progress-bar"
|
||||
style={{ width: `${Math.min(v.progress, 100)}%` }}
|
||||
/>
|
||||
</div>
|
||||
<span className="xx-variant-progress-text">{v.progress}%</span>
|
||||
</div>
|
||||
)}
|
||||
{v.status === "failed" && (
|
||||
<div className="xx-variant-failed">
|
||||
<CloseCircleOutlined style={{ fontSize: 28, color: "#ef4444" }} />
|
||||
<span>{v.error || "预览失败"}</span>
|
||||
</div>
|
||||
)}
|
||||
{v.status === "ready" && v.videoUrl && (
|
||||
<>
|
||||
<video
|
||||
src={v.videoUrl}
|
||||
controls
|
||||
style={{ width: "100%", display: "block", background: "#000", borderRadius: 8 }}
|
||||
onClick={(e) => e.stopPropagation()}
|
||||
/>
|
||||
{/* 标题浮层(CSS 实时叠加,改标题即时可见) */}
|
||||
{titleText && (
|
||||
<div
|
||||
className={`xx-variant-title-overlay pos-${titleStyle.position}`}
|
||||
style={{
|
||||
color: titleStyle.color,
|
||||
fontSize: Math.max(13, Math.round(titleStyle.size * 0.55)),
|
||||
WebkitTextStroke: "0.5px rgba(0,0,0,0.6)",
|
||||
}}
|
||||
onClick={(e) => e.stopPropagation()}
|
||||
>
|
||||
{titleText}
|
||||
</div>
|
||||
)}
|
||||
</>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* 底部状态 */}
|
||||
<div className="xx-variant-footer">
|
||||
{v.status === "ready" && selected && (
|
||||
<span className="xx-variant-ready-tag">
|
||||
<CheckCircleFilled style={{ color: "#52c41a" }} /> 已选择
|
||||
</span>
|
||||
)}
|
||||
{v.status === "ready" && !selected && selectable && (
|
||||
<span className="xx-variant-skip-tag">点击卡片取消/勾选</span>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
})}
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default ServerPreviewGrid
|
||||
@@ -1,18 +1,22 @@
|
||||
/**
|
||||
* Step 4 选择标题(Issue #1677 批量生成改造)
|
||||
* Step 4 选择标题(Issue #1677 批量生成)
|
||||
*
|
||||
* 布局(由 GeneratePage 编排):左侧大区域预览,右侧边栏标题设置。
|
||||
* 本组件渲染在右侧边栏:
|
||||
* - 标题文字:1 个视频 1 个输入框;N 个视频 N 个输入框各自独立
|
||||
* 布局(由 GeneratePage 编排):左侧大区域实时预览(单=大播放器,批量=Canvas 网格),
|
||||
* 右侧边栏标题设置。本组件渲染在右侧边栏:
|
||||
* - 单视频:AI 标题生成器 + AutoComplete 标题库(与旧版完全一致,零回归)
|
||||
* - 批量:N 个独立标题输入框(AutoComplete 支持标题库选择)+ 批量 AI 生成
|
||||
* (一次生成 N 个标题,分别填入各变体,可单独换一个)
|
||||
* - 标题样式(字体/颜色/位置/大小/粗斜描边/预设):全局统一
|
||||
*/
|
||||
import React from "react"
|
||||
import { AutoComplete, Input } from "antd"
|
||||
import React, { useMemo, useState } from "react"
|
||||
import { AutoComplete, Input, message } from "antd"
|
||||
import { LoadingOutlined } from "@ant-design/icons"
|
||||
import type { TitleSettings } from "../types"
|
||||
import { POSITION_OPTIONS, FONT_OPTIONS } from "../constants"
|
||||
import { useStep4Title } from "../hooks/useStep4Title"
|
||||
import AiTitleGenerator from "./title/AiTitleGenerator"
|
||||
import TitleStylePanel from "./title/TitleStylePanel"
|
||||
import { AI_TITLE_TEMPLATES } from "../constants"
|
||||
|
||||
interface Step4TitleSettingsProps {
|
||||
titleSettings: TitleSettings
|
||||
@@ -38,6 +42,36 @@ interface Step4TitleSettingsProps {
|
||||
onPreviewTitlesChange?: (titles: string[]) => void
|
||||
}
|
||||
|
||||
/** 从本地 AI 标题模板池按主题词生成 N 个不同标题(与单视频 AI 生成同源) */
|
||||
function buildBatchAiTitles(topic: string, count: number): string[] {
|
||||
const styles: Array<"catchy" | "emotional" | "informative"> = [
|
||||
"catchy",
|
||||
"emotional",
|
||||
"informative",
|
||||
]
|
||||
const pool: string[] = []
|
||||
styles.forEach((style) => {
|
||||
const templates = AI_TITLE_TEMPLATES[style] || []
|
||||
templates.forEach((tpl) => pool.push(tpl.replace(/\{topic\}/g, topic)))
|
||||
})
|
||||
// 洗牌后取前 count 个;不足则轮转补齐
|
||||
const shuffled = [...pool].sort(() => Math.random() - 0.5)
|
||||
const out: string[] = []
|
||||
for (let i = 0; i < count; i++) {
|
||||
out.push(shuffled[i % shuffled.length] || "")
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
function extractTopic(text: string): string {
|
||||
const keywords = text
|
||||
.replace(/[,。!?、,.!?]/g, " ")
|
||||
.split(/\s+/)
|
||||
.filter(Boolean)
|
||||
if (keywords.length === 0) return "这个话题"
|
||||
return keywords.slice(0, 3).join("")
|
||||
}
|
||||
|
||||
const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
|
||||
const t = useStep4Title(props)
|
||||
const {
|
||||
@@ -57,8 +91,10 @@ const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
|
||||
} = props
|
||||
|
||||
const isBatch = previewCount > 1
|
||||
const [batchAiLoading, setBatchAiLoading] = useState(false)
|
||||
const [batchAiTopic, setBatchAiTopic] = useState("")
|
||||
|
||||
/** 更新单个变体标题;变体0同步写回 titleSettings.title(全局样式面板/草稿保存依赖) */
|
||||
/** 更新单个变体标题;变体0同步写回 titleSettings.title(全局样式面板/草稿/TTS 链路依赖) */
|
||||
const updateVariantTitle = (index: number, val: string) => {
|
||||
if (!previewTitles || !onPreviewTitlesChange) return
|
||||
const next = [...previewTitles]
|
||||
@@ -69,12 +105,43 @@ const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
|
||||
}
|
||||
}
|
||||
|
||||
/** 批量 AI 生成:按主题词生成标题,分别填入 N 个变体 */
|
||||
const handleBatchAiGenerate = async (onlyEmpty = false) => {
|
||||
if (!onPreviewTitlesChange || !previewTitles) return
|
||||
const topic = (batchAiTopic || t.aiTitleInput || "").trim()
|
||||
if (!topic) {
|
||||
message.warning("请先输入主题词,例如:萌宠日常、旅行vlog")
|
||||
return
|
||||
}
|
||||
setBatchAiLoading(true)
|
||||
try {
|
||||
// 与单视频一致:本地模板模拟 AI 生成(1200ms 体验延迟)
|
||||
await new Promise((resolve) => setTimeout(resolve, 800))
|
||||
const picked = buildBatchAiTitles(extractTopic(topic), previewCount)
|
||||
const next = [...previewTitles]
|
||||
for (let i = 0; i < previewCount; i++) {
|
||||
if (onlyEmpty && next[i]?.trim()) continue
|
||||
if (picked[i]) next[i] = picked[i]
|
||||
}
|
||||
onPreviewTitlesChange(next)
|
||||
if (next[0]) t.updateTitle(next[0])
|
||||
message.success(`已为 ${previewCount} 个视频生成标题,可单独修改`)
|
||||
} finally {
|
||||
setBatchAiLoading(false)
|
||||
}
|
||||
}
|
||||
|
||||
const titleOptions = useMemo(
|
||||
() => t.userTitles.map((ut) => ({ label: ut.content, value: ut.content })),
|
||||
[t.userTitles],
|
||||
)
|
||||
|
||||
return (
|
||||
<div className="xx-form-section xx-title-sidebar">
|
||||
<h3>📝 选择标题</h3>
|
||||
|
||||
{!isBatch ? (
|
||||
/* ── 单视频:原有 AI 标题 + 输入框(保持不变) ── */
|
||||
/* ── 单视频:原有 AI 标题 + 输入框(保持不变,零回归) ── */
|
||||
<>
|
||||
{t.titleSettings.aiAutoSelect ? (
|
||||
<>
|
||||
@@ -144,7 +211,7 @@ const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
|
||||
t.updateTitle(val || "")
|
||||
onPreviewTitlesChange?.([val || ""])
|
||||
}}
|
||||
options={t.userTitles.map((ut) => ({ label: ut.content, value: ut.content }))}
|
||||
options={titleOptions}
|
||||
filterOption={(inputValue, option) => {
|
||||
const title = (option?.label || option?.value || "") as string
|
||||
return title.toLowerCase().includes((inputValue || "").toLowerCase())
|
||||
@@ -155,7 +222,7 @@ const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
|
||||
)}
|
||||
</>
|
||||
) : (
|
||||
/* ── 批量:N 个独立标题输入框(CSS 浮层实时叠加到对应预览) ── */
|
||||
/* ── 批量:AI 批量生成 + N 个独立标题输入框(AutoComplete 支持标题库) ── */
|
||||
<div className="xx-batch-titles">
|
||||
<div
|
||||
style={{
|
||||
@@ -167,15 +234,52 @@ const Step4TitleSettings: React.FC<Step4TitleSettingsProps> = (props) => {
|
||||
>
|
||||
为每个视频输入独立标题,修改会实时叠加到左侧对应视频上。标题样式(字体/颜色/位置)全局统一。
|
||||
</div>
|
||||
|
||||
{/* 批量 AI 标题 */}
|
||||
<div className="xx-batch-ai-row">
|
||||
<Input
|
||||
placeholder="主题词,如:萌宠日常、旅行vlog"
|
||||
value={batchAiTopic || t.aiTitleInput}
|
||||
onChange={(e) => {
|
||||
setBatchAiTopic(e.target.value)
|
||||
t.setAiTitleInput(e.target.value)
|
||||
}}
|
||||
maxLength={30}
|
||||
size="small"
|
||||
style={{ flex: 1 }}
|
||||
/>
|
||||
<button
|
||||
type="button"
|
||||
className="xx-btn xx-btn-primary xx-btn-sm"
|
||||
disabled={batchAiLoading}
|
||||
onClick={() => handleBatchAiGenerate(false)}
|
||||
>
|
||||
{batchAiLoading ? <LoadingOutlined /> : "✨"} 一键生成 {previewCount} 个标题
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
className="xx-btn xx-btn-ghost xx-btn-sm"
|
||||
disabled={batchAiLoading}
|
||||
onClick={() => handleBatchAiGenerate(true)}
|
||||
>
|
||||
补填空标题
|
||||
</button>
|
||||
</div>
|
||||
|
||||
{Array.from({ length: previewCount }, (_, i) => (
|
||||
<div className="xx-form-field" key={i}>
|
||||
<label>视频 {i + 1} 标题</label>
|
||||
<Input
|
||||
<AutoComplete
|
||||
placeholder={`视频 ${i + 1} 的标题…`}
|
||||
maxLength={50}
|
||||
showCount
|
||||
value={previewTitles?.[i] || ""}
|
||||
onChange={(e) => updateVariantTitle(i, e.target.value)}
|
||||
style={{ width: "100%" }}
|
||||
value={previewTitles?.[i] || undefined}
|
||||
onChange={(val) => updateVariantTitle(i, val || "")}
|
||||
options={titleOptions}
|
||||
filterOption={(inputValue, option) => {
|
||||
const title = (option?.label || option?.value || "") as string
|
||||
return title.toLowerCase().includes((inputValue || "").toLowerCase())
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
))}
|
||||
|
||||
@@ -33,7 +33,8 @@ export const STEPS = [
|
||||
{ key: 2, label: "选择素材" },
|
||||
{ key: 3, label: "选择配音" },
|
||||
{ key: 4, label: "选择标题" },
|
||||
{ key: 5, label: "选择封面" },
|
||||
{ key: 5, label: "确认生成" },
|
||||
{ key: 6, label: "选择封面" },
|
||||
]
|
||||
|
||||
/* ── 批量生成限制 ── */
|
||||
|
||||
@@ -3278,3 +3278,167 @@
|
||||
max-height: none;
|
||||
}
|
||||
}
|
||||
|
||||
/* ============================================================
|
||||
批量前端 Canvas 预览网格(Issue #1677 修正:纯前端实时预览)
|
||||
============================================================ */
|
||||
.xx-canvas-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(2, 1fr);
|
||||
gap: 16px;
|
||||
}
|
||||
|
||||
.xx-canvas-grid-card {
|
||||
border: 2px solid var(--border-primary, #e2e8f0);
|
||||
border-radius: 12px;
|
||||
overflow: hidden;
|
||||
background: #000;
|
||||
transition: border-color 0.2s ease;
|
||||
min-width: 0;
|
||||
}
|
||||
|
||||
.xx-canvas-grid-card.selected {
|
||||
border-color: var(--primary-color, #1677ff);
|
||||
box-shadow: 0 0 0 2px rgba(22, 119, 255, 0.15);
|
||||
}
|
||||
|
||||
.xx-canvas-grid-card-bar {
|
||||
position: relative;
|
||||
z-index: 2;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
padding: 6px 10px;
|
||||
background: var(--bg-surface, #fff);
|
||||
border-bottom: 1px solid var(--border-primary, #e2e8f0);
|
||||
}
|
||||
|
||||
.xx-canvas-grid-check {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
font-size: 13px;
|
||||
font-weight: 500;
|
||||
color: var(--text-primary, #1a1a1a);
|
||||
cursor: pointer;
|
||||
user-select: none;
|
||||
}
|
||||
|
||||
.xx-canvas-grid-check input[type="checkbox"] {
|
||||
width: 15px;
|
||||
height: 15px;
|
||||
cursor: pointer;
|
||||
accent-color: var(--primary-color, #1677ff);
|
||||
}
|
||||
|
||||
/* ============================================================
|
||||
批量标题:AI 一键生成行(Issue #1677)
|
||||
============================================================ */
|
||||
.xx-batch-ai-row {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
padding: 10px 12px;
|
||||
margin-bottom: 12px;
|
||||
background: var(--bg-secondary, #f7f8fa);
|
||||
border: 1px dashed var(--border-primary, #d9d9d9);
|
||||
border-radius: 10px;
|
||||
}
|
||||
|
||||
.xx-batch-ai-row .xx-form-field {
|
||||
margin: 0;
|
||||
flex: 1;
|
||||
min-width: 140px;
|
||||
}
|
||||
|
||||
.xx-batch-titles {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 10px;
|
||||
}
|
||||
|
||||
/* ============================================================
|
||||
第5步确认生成:批量渲染进度网格(Issue #1677)
|
||||
============================================================ */
|
||||
.xx-batch-gen-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(2, 1fr);
|
||||
gap: 16px;
|
||||
}
|
||||
|
||||
.xx-batch-gen-card {
|
||||
border: 1px solid var(--border-primary, #e2e8f0);
|
||||
border-radius: 12px;
|
||||
padding: 14px;
|
||||
background: var(--bg-surface, #fff);
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 10px;
|
||||
min-width: 0;
|
||||
}
|
||||
|
||||
.xx-batch-gen-card.status-completed {
|
||||
border-color: rgba(82, 196, 26, 0.4);
|
||||
background: rgba(82, 196, 26, 0.04);
|
||||
}
|
||||
|
||||
.xx-batch-gen-card.status-failed {
|
||||
border-color: rgba(239, 68, 68, 0.4);
|
||||
background: rgba(239, 68, 68, 0.04);
|
||||
}
|
||||
|
||||
.xx-batch-gen-card-head {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
gap: 8px;
|
||||
}
|
||||
|
||||
.xx-batch-gen-card-title {
|
||||
font-size: 14px;
|
||||
font-weight: 600;
|
||||
color: var(--text-primary, #1a1a1a);
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
white-space: nowrap;
|
||||
}
|
||||
|
||||
.xx-batch-gen-card-body {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 8px;
|
||||
}
|
||||
|
||||
.xx-batch-gen-card-pct {
|
||||
font-size: 13px;
|
||||
color: var(--text-secondary, #666);
|
||||
text-align: right;
|
||||
}
|
||||
|
||||
.xx-batch-gen-card-done {
|
||||
font-size: 13px;
|
||||
color: var(--success-color, #52c41a);
|
||||
padding: 8px 0;
|
||||
}
|
||||
|
||||
.xx-batch-gen-card-failed {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 8px;
|
||||
align-items: flex-start;
|
||||
}
|
||||
|
||||
.xx-batch-gen-card-err {
|
||||
font-size: 13px;
|
||||
color: var(--error-color, #ef4444);
|
||||
line-height: 1.5;
|
||||
word-break: break-word;
|
||||
}
|
||||
|
||||
/* ── 响应式:窄屏批量网格回退单列 ── */
|
||||
@media (max-width: 960px) {
|
||||
.xx-canvas-grid,
|
||||
.xx-batch-gen-grid {
|
||||
grid-template-columns: 1fr;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,14 +1,28 @@
|
||||
import { useRef, useCallback } from "react"
|
||||
import { useRef, useCallback, useState } from "react"
|
||||
import { message } from "antd"
|
||||
import axios from "axios"
|
||||
import { getGenerationTask } from "@/api/tasks/tasks"
|
||||
import { getGenerationTask, retryTask as retryGenerationTaskApi } from "@/api/tasks/tasks"
|
||||
import { getGenerationTaskResults } from "@/api/template-editor"
|
||||
import { safeExtractError } from "./errorUtils"
|
||||
|
||||
/** 批量生成时单个任务的实时状态(Issue #1677 第5步确认生成页) */
|
||||
export interface BatchTaskState {
|
||||
taskId: string
|
||||
/** 变体序号(0-based,与标题/封面数组对齐) */
|
||||
variantIndex: number
|
||||
status: "running" | "completed" | "failed"
|
||||
progress: number
|
||||
error: string | null
|
||||
/** 完成后的成片视频 */
|
||||
videos: unknown[]
|
||||
}
|
||||
|
||||
interface UseGenerationPollingOptions {
|
||||
onProgress: (progress: number) => void
|
||||
onComplete: (videos: unknown[]) => void
|
||||
onFailed: (errorMsg: string) => void
|
||||
/** 批量:单任务状态变化(第5步逐卡片展示) */
|
||||
onBatchTaskUpdate?: (taskId: string, patch: Partial<BatchTaskState>) => void
|
||||
}
|
||||
|
||||
/** 最大连续错误次数(仅对可重试错误),超过后终止轮询 */
|
||||
@@ -17,20 +31,25 @@ const MAX_RETRYABLE_ERRORS = 10
|
||||
const MAX_RESULTS_RETRIES = 3
|
||||
|
||||
/**
|
||||
* 生成状态轮询 Hook(v3 — 支持批量多任务)
|
||||
* 生成状态轮询 Hook(v4 — 批量任务独立状态 + 单任务重试)
|
||||
*
|
||||
* startPolling(taskId) 轮询单个任务;
|
||||
* startPollingBatch(taskIds) 并行轮询 N 个任务,全部完成后聚合结果,
|
||||
* 任一任务失败即整体失败(其余任务仍在后端继续,不影响)。
|
||||
* 进度为所有任务平均值。
|
||||
* startPollingBatch(tasks) 并行轮询 N 个任务:
|
||||
* - 每个任务独立进度/状态/失败,通过 onBatchTaskUpdate 实时回传
|
||||
* - 全部成功才 onComplete(聚合视频按变体顺序);任一失败不影响其他任务继续
|
||||
* - retryTask(taskId) 单独重试失败任务(重新轮询,后端任务仍在跑则直接接续)
|
||||
*/
|
||||
export function useGenerationPolling({
|
||||
onProgress,
|
||||
onComplete,
|
||||
onFailed,
|
||||
onBatchTaskUpdate,
|
||||
}: UseGenerationPollingOptions) {
|
||||
const progressTimer = useRef<ReturnType<typeof setTimeout>[]>([])
|
||||
const cancelledRef = useRef(false)
|
||||
/** 批量任务上下文:taskId → 变体序号 */
|
||||
const batchContextRef = useRef<Map<string, number>>(new Map())
|
||||
const [, forceTick] = useState(0)
|
||||
|
||||
const clearTimer = useCallback(() => {
|
||||
cancelledRef.current = true
|
||||
@@ -66,9 +85,21 @@ export function useGenerationPolling({
|
||||
return safeExtractError(msg)
|
||||
}
|
||||
|
||||
/** 轮询单个任务,resolve 该任务的结果视频数组;失败时 reject(new Error(msg)) */
|
||||
/**
|
||||
* 轮询单个任务。
|
||||
* - isBatch=true:状态变化通过 onBatchTaskUpdate 回传,不触发整体 onProgress/onComplete
|
||||
* - resolve(videos) 成功;reject(Error) 失败
|
||||
*/
|
||||
const pollSingleTask = useCallback(
|
||||
(taskId: string, runId: number, onTaskProgress?: (pct: number) => void): Promise<unknown[]> => {
|
||||
(
|
||||
taskId: string,
|
||||
runId: number,
|
||||
callbacks?: {
|
||||
onTaskProgress?: (pct: number) => void
|
||||
onTaskCompleted?: (videos: unknown[]) => void
|
||||
onTaskFailed?: (msg: string) => void
|
||||
},
|
||||
): Promise<unknown[]> => {
|
||||
return new Promise((resolve, reject) => {
|
||||
let consecutiveErrors = 0
|
||||
let done = false
|
||||
@@ -85,9 +116,12 @@ export function useGenerationPolling({
|
||||
const videos = await fetchResultsWithRetry(taskId)
|
||||
if (cancelledRef.current) return
|
||||
if (videos === null) {
|
||||
reject(new Error("视频已生成,但获取结果列表失败,请稍后在任务列表查看"))
|
||||
const msg = "视频已生成,但获取结果列表失败,请稍后在任务列表查看"
|
||||
callbacks?.onTaskFailed?.(msg)
|
||||
reject(new Error(msg))
|
||||
return
|
||||
}
|
||||
callbacks?.onTaskCompleted?.(videos)
|
||||
resolve(videos)
|
||||
return
|
||||
}
|
||||
@@ -98,14 +132,15 @@ export function useGenerationPolling({
|
||||
task.error_info?.error_message ||
|
||||
task.error_message ||
|
||||
(task.status === "cancelled" ? "任务已取消" : "视频生成失败,请联系管理员或重试")
|
||||
reject(new Error(safeExtractError(rawMsg)))
|
||||
const msg = safeExtractError(rawMsg)
|
||||
callbacks?.onTaskFailed?.(msg)
|
||||
reject(new Error(msg))
|
||||
return
|
||||
}
|
||||
|
||||
const pct = Math.max(0, Math.min(99, Math.round(Number(task.progress) || 0)))
|
||||
if (onTaskProgress) {
|
||||
onTaskProgress(pct)
|
||||
} else if (runId === 0) {
|
||||
callbacks?.onTaskProgress?.(pct)
|
||||
if (!callbacks && runId === 0) {
|
||||
onProgress(pct)
|
||||
}
|
||||
const timer = setTimeout(poll, 2000)
|
||||
@@ -116,13 +151,17 @@ export function useGenerationPolling({
|
||||
const status = axios.isAxiosError(pollErr) ? pollErr.response?.status : undefined
|
||||
if (status && status >= 400 && status < 500) {
|
||||
done = true
|
||||
reject(new Error(extractErrorMessage(pollErr, status)))
|
||||
const msg = extractErrorMessage(pollErr, status)
|
||||
callbacks?.onTaskFailed?.(msg)
|
||||
reject(new Error(msg))
|
||||
return
|
||||
}
|
||||
consecutiveErrors += 1
|
||||
if (consecutiveErrors >= MAX_RETRYABLE_ERRORS) {
|
||||
done = true
|
||||
reject(new Error("任务状态查询连续失败,请稍后在任务列表查看结果"))
|
||||
const msg = "任务状态查询连续失败,请稍后在任务列表查看结果"
|
||||
callbacks?.onTaskFailed?.(msg)
|
||||
reject(new Error(msg))
|
||||
return
|
||||
}
|
||||
const timer = setTimeout(poll, 3000)
|
||||
@@ -137,12 +176,12 @@ export function useGenerationPolling({
|
||||
[onProgress, fetchResultsWithRetry],
|
||||
)
|
||||
|
||||
/** 单任务轮询(兼容旧调用) */
|
||||
/** 单任务轮询(单视频,兼容旧调用) */
|
||||
const startPolling = useCallback(
|
||||
(taskId: string) => {
|
||||
cancelledRef.current = false
|
||||
const runId = 0
|
||||
pollSingleTask(taskId, runId)
|
||||
batchContextRef.current.clear()
|
||||
pollSingleTask(taskId, 0)
|
||||
.then((videos) => {
|
||||
if (cancelledRef.current) return
|
||||
onProgress(100)
|
||||
@@ -159,48 +198,114 @@ export function useGenerationPolling({
|
||||
[pollSingleTask, onProgress, onComplete, onFailed],
|
||||
)
|
||||
|
||||
/** 批量多任务轮询:全部完成后聚合结果;任一失败即整体失败 */
|
||||
/**
|
||||
* 批量多任务轮询:
|
||||
* - 每个任务独立进度/状态回传 onBatchTaskUpdate
|
||||
* * 全部完成后按变体顺序聚合视频 onComplete
|
||||
* - 部分失败:整体不 onFailed(第5步逐卡片展示失败+重试按钮);全部失败才 onFailed
|
||||
*/
|
||||
const startPollingBatch = useCallback(
|
||||
(taskIds: string[]) => {
|
||||
(tasks: { taskId: string; variantIndex: number }[]) => {
|
||||
cancelledRef.current = false
|
||||
const runId = Date.now()
|
||||
const progressMap = new Map<string, number>()
|
||||
const resultMap = new Map<string, unknown[]>()
|
||||
const failureMap = new Map<string, string>()
|
||||
batchContextRef.current = new Map(tasks.map((t) => [t.taskId, t.variantIndex]))
|
||||
|
||||
const reportAggregateProgress = () => {
|
||||
if (cancelledRef.current) return
|
||||
const values = taskIds.map((id) => progressMap.get(id) ?? 0)
|
||||
const values = tasks.map((t) => progressMap.get(t.taskId) ?? 0)
|
||||
const avg = Math.round(values.reduce((a, b) => a + b, 0) / Math.max(values.length, 1))
|
||||
onProgress(Math.min(avg, 99))
|
||||
}
|
||||
|
||||
const tasks = taskIds.map((taskId) =>
|
||||
pollSingleTask(taskId, runId, (pct) => {
|
||||
progressMap.set(taskId, pct)
|
||||
reportAggregateProgress()
|
||||
}).then((videos) => {
|
||||
progressMap.set(taskId, 100)
|
||||
reportAggregateProgress()
|
||||
return videos
|
||||
}),
|
||||
)
|
||||
|
||||
Promise.all(tasks)
|
||||
.then((results) => {
|
||||
if (cancelledRef.current) return
|
||||
const checkAllSettled = () => {
|
||||
if (resultMap.size + failureMap.size < tasks.length) return
|
||||
if (resultMap.size === tasks.length) {
|
||||
onProgress(100)
|
||||
const allVideos = results.flat()
|
||||
onComplete(allVideos)
|
||||
message.success(`全部 ${taskIds.length} 个视频生成完成!`)
|
||||
const ordered = tasks.map((t) => resultMap.get(t.taskId) || []).flat()
|
||||
onComplete(ordered)
|
||||
message.success(`全部 ${tasks.length} 个视频生成完成!`)
|
||||
} else if (resultMap.size > 0) {
|
||||
// 部分失败:成功的视频聚合进成片列表(可进封面),失败卡片带重试按钮
|
||||
onProgress(100)
|
||||
const ordered = tasks
|
||||
.filter((t) => resultMap.has(t.taskId))
|
||||
.map((t) => resultMap.get(t.taskId) || [])
|
||||
.flat()
|
||||
onComplete(ordered)
|
||||
message.warning(
|
||||
`${failureMap.size} 个视频生成失败,可点击卡片上的「重试此视频」,成功的视频可先进入下一步`,
|
||||
)
|
||||
} else {
|
||||
const firstMsg = failureMap.get(tasks[0].taskId) || "全部视频生成失败"
|
||||
onFailed(firstMsg)
|
||||
}
|
||||
}
|
||||
|
||||
tasks.forEach(({ taskId, variantIndex }) => {
|
||||
onBatchTaskUpdate?.(taskId, {
|
||||
taskId,
|
||||
variantIndex,
|
||||
status: "running",
|
||||
progress: 0,
|
||||
error: null,
|
||||
videos: [],
|
||||
})
|
||||
.catch((err: Error) => {
|
||||
if (cancelledRef.current) return
|
||||
console.error("[批量生成失败]", err.message)
|
||||
onFailed(err.message)
|
||||
message.error(err.message)
|
||||
pollSingleTask(taskId, runId, {
|
||||
onTaskProgress: (pct) => {
|
||||
progressMap.set(taskId, pct)
|
||||
onBatchTaskUpdate?.(taskId, { status: "running", progress: pct })
|
||||
reportAggregateProgress()
|
||||
},
|
||||
onTaskCompleted: (videos) => {
|
||||
progressMap.set(taskId, 100)
|
||||
resultMap.set(taskId, videos)
|
||||
onBatchTaskUpdate?.(taskId, { status: "completed", progress: 100, videos })
|
||||
reportAggregateProgress()
|
||||
checkAllSettled()
|
||||
},
|
||||
onTaskFailed: (msg) => {
|
||||
failureMap.set(taskId, msg)
|
||||
onBatchTaskUpdate?.(taskId, { status: "failed", error: msg })
|
||||
checkAllSettled()
|
||||
},
|
||||
}).catch(() => {
|
||||
// 失败已在 onTaskFailed 处理,这里吞掉 Promise rejection
|
||||
})
|
||||
})
|
||||
},
|
||||
[pollSingleTask, onProgress, onComplete, onFailed],
|
||||
[pollSingleTask, onProgress, onComplete, onFailed, onBatchTaskUpdate],
|
||||
)
|
||||
|
||||
return { startPolling, startPollingBatch, clearTimer }
|
||||
/** 单独重试失败任务(第5步卡片「重试此视频」):先调后端重试接口,再轮询 */
|
||||
const retryTask = useCallback(
|
||||
async (taskId: string) => {
|
||||
if (cancelledRef.current) cancelledRef.current = false
|
||||
const variantIndex = batchContextRef.current.get(taskId) ?? 0
|
||||
onBatchTaskUpdate?.(taskId, { status: "running", progress: 0, error: null, videos: [] })
|
||||
try {
|
||||
await retryGenerationTaskApi(taskId)
|
||||
} catch (err) {
|
||||
// 后端不支持重试或任务不可重试:直接重新轮询(任务可能已被自动恢复)
|
||||
console.warn("[重试任务接口调用失败,改为直接轮询]", err)
|
||||
}
|
||||
pollSingleTask(taskId, Date.now(), {
|
||||
onTaskProgress: (pct) => onBatchTaskUpdate?.(taskId, { status: "running", progress: pct }),
|
||||
onTaskCompleted: (videos) => {
|
||||
onBatchTaskUpdate?.(taskId, { status: "completed", progress: 100, videos })
|
||||
message.success(`视频 ${variantIndex + 1} 重试成功`)
|
||||
},
|
||||
onTaskFailed: (msg) => onBatchTaskUpdate?.(taskId, { status: "failed", error: msg }),
|
||||
}).catch(() => {
|
||||
/* 失败已在回调处理 */
|
||||
})
|
||||
forceTick((n) => n + 1)
|
||||
return variantIndex
|
||||
},
|
||||
[pollSingleTask, onBatchTaskUpdate],
|
||||
)
|
||||
|
||||
return { startPolling, startPollingBatch, retryTask, clearTimer }
|
||||
}
|
||||
|
||||
@@ -1,285 +0,0 @@
|
||||
/**
|
||||
* 批量服务器预览 Hook(Issue #1677 多视频批量生成)
|
||||
*
|
||||
* 核心职责:
|
||||
* 1. 调用 POST /generation/preview(preview_count=N)一次创建 N 个独立变体任务
|
||||
* 2. 对每个变体 task_id 分别轮询 GET /generation/preview/{task_id}
|
||||
* 3. 返回每个变体的状态/进度/视频URL,供网格播放器展示
|
||||
*
|
||||
* N=1 时不启用(走前端 Canvas 实时预览,零回归);
|
||||
* N>1 时进入标题页自动触发;素材/配音等配置变化后重新触发。
|
||||
*/
|
||||
import { useState, useCallback, useRef, useEffect } from "react"
|
||||
import { createPreview, getPreviewStatus } from "@/api/generation/preview"
|
||||
import type { CreatePreviewRequest } from "@/api/generation/types"
|
||||
|
||||
export type VariantPreviewStatus = "loading" | "ready" | "failed"
|
||||
|
||||
export interface VariantPreview {
|
||||
/** 变体序号(0-based) */
|
||||
index: number
|
||||
taskId: string
|
||||
status: VariantPreviewStatus
|
||||
progress: number
|
||||
videoUrl: string | null
|
||||
error: string | null
|
||||
}
|
||||
|
||||
interface UseBatchPreviewOptions {
|
||||
/** 是否启用(仅 previewCount>1 且在标题页时启用) */
|
||||
enabled: boolean
|
||||
/** 构建预览请求参数(每次触发时调用,获取最新配置) */
|
||||
buildRequest: () => CreatePreviewRequest
|
||||
/** 批量预览任务创建成功回调(回传变体 taskId 列表与 source_edit_plan_id) */
|
||||
onPreviewTasksCreated?: (taskIds: string[], sourceEditPlanId?: string) => void
|
||||
}
|
||||
|
||||
interface UseBatchPreviewReturn {
|
||||
variants: VariantPreview[]
|
||||
/** 整体状态:loading=任一进行中,ready=全部完成,failed=有失败 */
|
||||
status: "idle" | "loading" | "ready" | "partial_failed" | "failed"
|
||||
/** 总进度 0-100(各变体平均值) */
|
||||
progress: number
|
||||
/** 失败的变体数量 */
|
||||
failedCount: number
|
||||
/** 手动重新触发 */
|
||||
trigger: () => void
|
||||
}
|
||||
|
||||
const POLL_INTERVAL = 2000
|
||||
const POLL_TIMEOUT = 180_000
|
||||
const MAX_NETWORK_RETRIES = 2
|
||||
|
||||
/**
|
||||
* 对配置参数做指纹,用于检测配置是否变化(标题文字/样式变化不触发重渲染,仅CSS浮层叠加)
|
||||
*/
|
||||
function buildFingerprint(req: CreatePreviewRequest): string {
|
||||
// 不含 titles/title_config:标题文字与样式由 CSS 浮层实时叠加,变化不触发重渲染
|
||||
return JSON.stringify({
|
||||
t: req.template_id,
|
||||
a: [...(req.asset_ids || [])].sort(),
|
||||
r: req.video_ratio,
|
||||
v: req.voice_library_id,
|
||||
vs: req.voice_library_ids,
|
||||
pc: req.preview_count,
|
||||
b: req.bgm_config,
|
||||
})
|
||||
}
|
||||
|
||||
export function useBatchPreview({
|
||||
enabled,
|
||||
buildRequest,
|
||||
onPreviewTasksCreated,
|
||||
}: UseBatchPreviewOptions): UseBatchPreviewReturn {
|
||||
const [variants, setVariants] = useState<VariantPreview[]>([])
|
||||
const [status, setStatus] = useState<"idle" | "loading" | "ready" | "partial_failed" | "failed">(
|
||||
"idle",
|
||||
)
|
||||
const requestSeqRef = useRef(0)
|
||||
const pollTimersRef = useRef<ReturnType<typeof setTimeout>[]>([])
|
||||
const timeoutTimerRef = useRef<ReturnType<typeof setTimeout> | null>(null)
|
||||
const mountedRef = useRef(true)
|
||||
|
||||
const buildRequestRef = useRef(buildRequest)
|
||||
buildRequestRef.current = buildRequest
|
||||
const onCreatedRef = useRef(onPreviewTasksCreated)
|
||||
onCreatedRef.current = onPreviewTasksCreated
|
||||
|
||||
const clearTimers = useCallback(() => {
|
||||
pollTimersRef.current.forEach((t) => clearTimeout(t))
|
||||
pollTimersRef.current = []
|
||||
if (timeoutTimerRef.current) {
|
||||
clearTimeout(timeoutTimerRef.current)
|
||||
timeoutTimerRef.current = null
|
||||
}
|
||||
}, [])
|
||||
|
||||
useEffect(() => {
|
||||
mountedRef.current = true
|
||||
return () => {
|
||||
mountedRef.current = false
|
||||
clearTimers()
|
||||
}
|
||||
}, [clearTimers])
|
||||
|
||||
/** 更新单个变体状态 */
|
||||
const patchVariant = useCallback((taskId: string, patch: Partial<VariantPreview>) => {
|
||||
setVariants((prev) => prev.map((v) => (v.taskId === taskId ? { ...v, ...patch } : v)))
|
||||
}, [])
|
||||
|
||||
/** 轮询单个变体任务 */
|
||||
const pollVariant = useCallback(
|
||||
async (taskId: string, seq: number, retries = 0) => {
|
||||
if (seq !== requestSeqRef.current || !mountedRef.current) return
|
||||
try {
|
||||
const st = await getPreviewStatus(taskId)
|
||||
if (seq !== requestSeqRef.current || !mountedRef.current) return
|
||||
|
||||
if (st.status === "completed" && st.video_url) {
|
||||
patchVariant(taskId, {
|
||||
status: "ready",
|
||||
videoUrl: st.video_url,
|
||||
progress: 100,
|
||||
error: null,
|
||||
})
|
||||
return
|
||||
}
|
||||
if (st.status === "failed" || st.status === "cancelled") {
|
||||
patchVariant(taskId, {
|
||||
status: "failed",
|
||||
error:
|
||||
st.status === "cancelled" ? "预览任务已取消" : st.error_message || "预览渲染失败",
|
||||
})
|
||||
return
|
||||
}
|
||||
if (typeof st.progress === "number") {
|
||||
patchVariant(taskId, { progress: Math.round(st.progress) })
|
||||
}
|
||||
const timer = setTimeout(() => pollVariant(taskId, seq), POLL_INTERVAL)
|
||||
pollTimersRef.current.push(timer)
|
||||
} catch (err) {
|
||||
if (seq !== requestSeqRef.current || !mountedRef.current) return
|
||||
if (retries < MAX_NETWORK_RETRIES) {
|
||||
console.warn(`[BatchPreview] 变体 ${taskId} 轮询网络错误,第 ${retries + 1} 次重试`, err)
|
||||
const timer = setTimeout(() => pollVariant(taskId, seq, retries + 1), POLL_INTERVAL * 2)
|
||||
pollTimersRef.current.push(timer)
|
||||
} else {
|
||||
patchVariant(taskId, { status: "failed", error: "网络错误,无法获取预览状态" })
|
||||
}
|
||||
}
|
||||
},
|
||||
[patchVariant],
|
||||
)
|
||||
|
||||
/** 创建批量预览任务并开始轮询 */
|
||||
const trigger = useCallback(() => {
|
||||
if (!enabled) return
|
||||
const request = buildRequestRef.current()
|
||||
if (!request.template_id || !request.asset_ids?.length) return
|
||||
const count = request.preview_count && request.preview_count > 1 ? request.preview_count : 0
|
||||
if (!count) return
|
||||
|
||||
clearTimers()
|
||||
const seq = ++requestSeqRef.current
|
||||
setStatus("loading")
|
||||
setVariants(
|
||||
Array.from({ length: count }, (_, i) => ({
|
||||
index: i,
|
||||
taskId: "",
|
||||
status: "loading" as const,
|
||||
progress: 0,
|
||||
videoUrl: null,
|
||||
error: null,
|
||||
})),
|
||||
)
|
||||
|
||||
createPreview(request)
|
||||
.then((resp) => {
|
||||
if (seq !== requestSeqRef.current || !mountedRef.current) return
|
||||
const items = resp.items || []
|
||||
const taskIds = items.map((it) => it.task_id).filter(Boolean)
|
||||
if (taskIds.length === 0) {
|
||||
setStatus("failed")
|
||||
setVariants((prev) =>
|
||||
prev.map((v) => ({ ...v, status: "failed", error: "未创建预览任务" })),
|
||||
)
|
||||
return
|
||||
}
|
||||
onCreatedRef.current?.(taskIds, resp.source_edit_plan_id)
|
||||
|
||||
// 用返回的 task_id 填充变体(按 variant_index 对齐)
|
||||
setVariants((prev) =>
|
||||
prev.map((v) => {
|
||||
const item = items.find((it) => it.variant_index === v.index) || items[v.index]
|
||||
return item ? { ...v, taskId: item.task_id } : v
|
||||
}),
|
||||
)
|
||||
|
||||
// 超时保护
|
||||
timeoutTimerRef.current = setTimeout(() => {
|
||||
if (seq !== requestSeqRef.current || !mountedRef.current) return
|
||||
setVariants((prev) =>
|
||||
prev.map((v) =>
|
||||
v.status === "loading"
|
||||
? { ...v, status: "failed", error: "预览渲染超时,请重试" }
|
||||
: v,
|
||||
),
|
||||
)
|
||||
}, POLL_TIMEOUT)
|
||||
|
||||
// 分别轮询每个变体
|
||||
items.forEach((item) => {
|
||||
if (item.task_id) pollVariant(item.task_id, seq)
|
||||
})
|
||||
})
|
||||
.catch((err: unknown) => {
|
||||
if (seq !== requestSeqRef.current || !mountedRef.current) return
|
||||
console.error("[BatchPreview] 创建批量预览失败:", err)
|
||||
const errData = (err as { response?: { data?: { detail?: string; message?: string } } })
|
||||
?.response?.data
|
||||
setStatus("failed")
|
||||
setVariants((prev) =>
|
||||
prev.map((v) => ({
|
||||
...v,
|
||||
status: "failed",
|
||||
error: errData?.detail || errData?.message || "预览任务创建失败,请重试",
|
||||
})),
|
||||
)
|
||||
})
|
||||
}, [enabled, clearTimers, pollVariant])
|
||||
|
||||
/* ── 自动触发 + 配置变更检测 ── */
|
||||
const request = enabled ? buildRequest() : null
|
||||
const currentFingerprint = request
|
||||
? request.template_id && request.asset_ids?.length && (request.preview_count || 1) > 1
|
||||
? buildFingerprint(request)
|
||||
: ""
|
||||
: ""
|
||||
|
||||
const didInitRef = useRef(false)
|
||||
useEffect(() => {
|
||||
if (!enabled || !currentFingerprint) {
|
||||
didInitRef.current = false
|
||||
requestSeqRef.current += 1
|
||||
clearTimers()
|
||||
setStatus("idle")
|
||||
setVariants([])
|
||||
return
|
||||
}
|
||||
if (!didInitRef.current) {
|
||||
didInitRef.current = true
|
||||
trigger()
|
||||
}
|
||||
}, [enabled, currentFingerprint, trigger, clearTimers])
|
||||
|
||||
// 配置变更(素材/配音/数量)→ 重新渲染;标题文字变化不触发(CSS浮层实时叠加)
|
||||
const prevFingerprintRef = useRef(currentFingerprint)
|
||||
useEffect(() => {
|
||||
if (!enabled || !currentFingerprint) return
|
||||
const prev = prevFingerprintRef.current
|
||||
prevFingerprintRef.current = currentFingerprint
|
||||
if (!prev || prev === currentFingerprint) return
|
||||
trigger()
|
||||
}, [enabled, currentFingerprint, trigger])
|
||||
|
||||
/* ── 派生状态 ── */
|
||||
const progress =
|
||||
variants.length > 0
|
||||
? Math.round(variants.reduce((sum, v) => sum + v.progress, 0) / variants.length)
|
||||
: 0
|
||||
const failedCount = variants.filter((v) => v.status === "failed").length
|
||||
const readyCount = variants.filter((v) => v.status === "ready").length
|
||||
|
||||
useEffect(() => {
|
||||
if (status !== "loading" || variants.length === 0) return
|
||||
if (readyCount === variants.length) {
|
||||
setStatus("ready")
|
||||
} else if (readyCount + failedCount === variants.length && failedCount > 0) {
|
||||
setStatus(failedCount === variants.length ? "failed" : "partial_failed")
|
||||
}
|
||||
}, [variants, status, readyCount, failedCount])
|
||||
|
||||
return { variants, status, progress, failedCount, trigger }
|
||||
}
|
||||
|
||||
export default useBatchPreview
|
||||
@@ -2,13 +2,13 @@
|
||||
* 视频生成 Hook
|
||||
* 封装视频生成的核心逻辑、状态管理、轮询等
|
||||
*/
|
||||
import { useState, useCallback } from "react"
|
||||
import { useState, useCallback, useEffect } from "react"
|
||||
import { message } from "antd"
|
||||
import { type GeneratedVideo, getEditPlanClips, createClipsFromAssets } from "@/api/template-editor"
|
||||
import { createGenerationTask } from "@/api/tasks/tasks"
|
||||
import type { UseGenerateVideoProps } from "./generate-video/types"
|
||||
import { getGenerationPhase } from "./generate-video/phase"
|
||||
import { useGenerationPolling } from "./generate-video/useGenerationPolling"
|
||||
import { useGenerationPolling, type BatchTaskState } from "./generate-video/useGenerationPolling"
|
||||
import { validateGenerateInputs } from "./generate-video/buildPayload"
|
||||
import { calculateResolution } from "../utils/calculateResolution"
|
||||
import { extractBackendError, translateError } from "./generate-video/errorUtils"
|
||||
@@ -22,6 +22,32 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
const [generated, setGenerated] = useState(false)
|
||||
const [generateError, setGenerateError] = useState<string | null>(null)
|
||||
const [generatedVideos, setGeneratedVideos] = useState<GeneratedVideo[]>([])
|
||||
/** 批量模式:每个正式生成任务的独立状态(第5步逐卡片展示) */
|
||||
const [batchTasks, setBatchTasks] = useState<BatchTaskState[]>([])
|
||||
|
||||
const handleBatchTaskUpdate = useCallback((taskId: string, patch: Partial<BatchTaskState>) => {
|
||||
setBatchTasks((prev) => {
|
||||
const list = prev || []
|
||||
const idx = list.findIndex((t) => t.taskId === taskId)
|
||||
if (idx === -1) {
|
||||
return [
|
||||
...list,
|
||||
{
|
||||
taskId,
|
||||
variantIndex: patch.variantIndex ?? 0,
|
||||
status: "running",
|
||||
progress: 0,
|
||||
error: null,
|
||||
videos: [],
|
||||
...patch,
|
||||
},
|
||||
]
|
||||
}
|
||||
const next = [...list]
|
||||
next[idx] = { ...next[idx], ...patch }
|
||||
return next
|
||||
})
|
||||
}, [])
|
||||
|
||||
const handleProgress = useCallback((p: number) => setProgress(p), [])
|
||||
const handleComplete = useCallback(
|
||||
@@ -29,6 +55,19 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
setGenerating(false)
|
||||
setGenerated(true)
|
||||
setGeneratedVideos(videos as GeneratedVideo[])
|
||||
// 批量:成功任务的 videos 已通过 onBatchTaskUpdate 写入,这里同步兜底
|
||||
setBatchTasks((prev) =>
|
||||
(prev || []).map((t) =>
|
||||
t.status === "completed" && t.videos.length === 0
|
||||
? {
|
||||
...t,
|
||||
videos: (videos as GeneratedVideo[]).filter(
|
||||
(v) => v.generation_task_id === t.taskId,
|
||||
),
|
||||
}
|
||||
: t,
|
||||
),
|
||||
)
|
||||
onGenerationSuccess?.()
|
||||
},
|
||||
[onGenerationSuccess],
|
||||
@@ -38,10 +77,30 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
setGenerateError(errorMsg)
|
||||
}, [])
|
||||
|
||||
const { startPolling, startPollingBatch, clearTimer } = useGenerationPolling({
|
||||
/* 批量:任务状态变化时聚合已完成成片(含失败重试成功后补入),
|
||||
按变体索引排序,供步骤6封面按勾选顺序逐个取视频 */
|
||||
useEffect(() => {
|
||||
if (batchTasks.length === 0) return
|
||||
const byVariant = new Map<number, GeneratedVideo>()
|
||||
batchTasks.forEach((t) => {
|
||||
if (t.status === "completed" && t.videos && t.videos.length > 0) {
|
||||
byVariant.set(t.variantIndex, t.videos[0] as GeneratedVideo)
|
||||
}
|
||||
})
|
||||
const ordered = [...byVariant.entries()].sort((a, b) => a[0] - b[0]).map(([, v]) => v)
|
||||
setGeneratedVideos((prev) => {
|
||||
if (prev.length === ordered.length && prev.every((v, i) => v.id === ordered[i].id)) {
|
||||
return prev
|
||||
}
|
||||
return ordered
|
||||
})
|
||||
}, [batchTasks])
|
||||
|
||||
const { startPolling, startPollingBatch, retryTask, clearTimer } = useGenerationPolling({
|
||||
onProgress: handleProgress,
|
||||
onComplete: handleComplete,
|
||||
onFailed: handleFailed,
|
||||
onBatchTaskUpdate: handleBatchTaskUpdate,
|
||||
})
|
||||
|
||||
/* ── 生成视频 ──
|
||||
@@ -57,6 +116,7 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
setProgress(0)
|
||||
setGenerated(false)
|
||||
setGenerateError(null)
|
||||
setBatchTasks([])
|
||||
clearTimer()
|
||||
|
||||
try {
|
||||
@@ -170,7 +230,8 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
throw new Error("创建任务成功但未返回任务 ID,请稍后在任务列表查看")
|
||||
}
|
||||
if (taskIds.length > 1) {
|
||||
startPollingBatch(taskIds)
|
||||
// 批量:任务按创建顺序与勾选变体一一对应(后端按 count 顺序创建)
|
||||
startPollingBatch(taskIds.map((taskId, i) => ({ taskId, variantIndex: indexes[i] ?? i })))
|
||||
} else {
|
||||
startPolling(taskIds[0])
|
||||
}
|
||||
@@ -196,6 +257,14 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
generate()
|
||||
}, [generate])
|
||||
|
||||
/** 第5步:单独重试某个失败任务 */
|
||||
const retryBatchTask = useCallback(
|
||||
(taskId: string) => {
|
||||
retryTask(taskId)
|
||||
},
|
||||
[retryTask],
|
||||
)
|
||||
|
||||
const dismissError = useCallback(() => {
|
||||
setGenerateError(null)
|
||||
}, [])
|
||||
@@ -240,6 +309,8 @@ export function useGenerateVideo(props: UseGenerateVideoProps) {
|
||||
generatedVideos,
|
||||
generate,
|
||||
retry,
|
||||
retryBatchTask,
|
||||
batchTasks,
|
||||
dismissError,
|
||||
download,
|
||||
share,
|
||||
|
||||
@@ -1,6 +1,11 @@
|
||||
/**
|
||||
* GeneratePage 步骤导航(Issue #1677 改造后 5 步)
|
||||
* 步骤:模板(1) → 素材(2) → 配音(3) → 标题+预览+确认生成(4) → 封面(5)
|
||||
* GeneratePage 步骤导航(Issue #1677 修正:固定 6 步,单视频与批量一致)
|
||||
* 步骤:模板(1) → 素材(2) → 配音(3) → 标题(4) → 确认生成(5) → 封面(6)
|
||||
*
|
||||
* - 步骤4底部按钮是「确认生成视频/确认生成 N 个视频」(由 GenerateStepActions 调
|
||||
* onConfirmGenerate),创建成功后跳转步骤5;本 hook 的 goNext 只负责 1→2→3→4
|
||||
* 和 5→6 的「下一步」。
|
||||
* - 步骤5(确认生成进度页):渲染全部完成(generated)后「下一步」解锁进封面。
|
||||
*/
|
||||
import { message } from "antd"
|
||||
import type { TitleSettings } from "../types"
|
||||
@@ -13,14 +18,8 @@ export interface UseStepNavigationOptions {
|
||||
selectedMaterials: string[]
|
||||
smartSelectedIds: string[]
|
||||
titleSettings: TitleSettings
|
||||
/** 预览是否已就绪(单视频=前端预览素材已加载;批量=服务器预览全部完成) */
|
||||
previewReady: boolean
|
||||
/** 是否已完成视频生成(步骤4确认生成后才能进入封面) */
|
||||
/** 是否已完成视频生成(步骤5全部渲染完成后才能进入封面) */
|
||||
generated: boolean
|
||||
/** 批量模式下每个变体的标题 */
|
||||
previewTitles: string[]
|
||||
/** 批量模式勾选的变体数 */
|
||||
selectedCount: number
|
||||
/** Step1 点下一步时弹出数量选择弹窗 */
|
||||
onOpenCountModal: () => void
|
||||
}
|
||||
@@ -38,10 +37,7 @@ export const useStepNavigation = (options: UseStepNavigationOptions): UseStepNav
|
||||
materialMode,
|
||||
selectedMaterials,
|
||||
smartSelectedIds,
|
||||
previewReady,
|
||||
generated,
|
||||
previewTitles,
|
||||
selectedCount,
|
||||
onOpenCountModal,
|
||||
} = options
|
||||
|
||||
@@ -63,27 +59,14 @@ export const useStepNavigation = (options: UseStepNavigationOptions): UseStepNav
|
||||
message.warning("请先进行智能匹配并选择素材")
|
||||
return
|
||||
}
|
||||
// Step4(标题+预览+确认生成):标题必填 + 预览必须已加载
|
||||
if (currentStep === 4) {
|
||||
const allTitlesFilled = previewTitles.every((t) => t && t.trim())
|
||||
if (!allTitlesFilled) {
|
||||
message.warning("请为每个视频输入标题")
|
||||
return
|
||||
}
|
||||
if (selectedCount === 0) {
|
||||
message.warning("请至少勾选一个视频")
|
||||
return
|
||||
}
|
||||
if (!previewReady) {
|
||||
message.warning("预览视频正在加载,请稍候")
|
||||
return
|
||||
}
|
||||
// 步骤5(确认生成):全部渲染完成后才能下一步进封面
|
||||
if (currentStep === 5) {
|
||||
if (!generated) {
|
||||
message.warning("请先点击「确认生成视频」完成渲染")
|
||||
message.warning("视频还在渲染中,请等待生成完成")
|
||||
return
|
||||
}
|
||||
}
|
||||
if (currentStep < 5) {
|
||||
if (currentStep < 6) {
|
||||
setCurrentStep((s) => s + 1)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -178,3 +178,35 @@
|
||||
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);
|
||||
}
|
||||
|
||||
@@ -1,39 +1,111 @@
|
||||
/**
|
||||
* 个人设置页面
|
||||
* P1-2: 添加 PageHead
|
||||
* P1-3: antd Form/Input/Button/Alert → 自定义 UI 组件
|
||||
* - 个人资料(昵称)保存
|
||||
* - 微信账号绑定状态 / 绑定 / 解绑
|
||||
*/
|
||||
import React, { useState } from "react"
|
||||
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 { 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 handleSave = () => {
|
||||
Modal.info({
|
||||
title: "提示",
|
||||
content: "个人资料修改接口暂未开放,保存功能即将上线。",
|
||||
// 拉取最新用户信息(微信绑定状态以后端为准)
|
||||
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 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>
|
||||
@@ -42,25 +114,76 @@ const Settings: React.FC = () => {
|
||||
|
||||
<div className="xx-settings-field">
|
||||
<label className="xx-settings-label">邮箱</label>
|
||||
<Input value={user?.email || ""} disabled placeholder="邮箱" />
|
||||
</div>
|
||||
|
||||
<div className="xx-settings-field">
|
||||
<label className="xx-settings-label">显示名称</label>
|
||||
<Input
|
||||
value={displayName}
|
||||
onChange={(e) => setDisplayName(e.target.value)}
|
||||
placeholder="请输入显示名称"
|
||||
value={user?.email && !user.email.endsWith("@wechat.local") ? user.email : ""}
|
||||
disabled
|
||||
placeholder={user?.email?.endsWith("@wechat.local") ? "微信账号暂未绑定邮箱" : "邮箱"}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div className="xx-settings-field">
|
||||
<Button buttonType="primary" buttonSize="md" onClick={handleSave} disabled>
|
||||
保存暂未开放
|
||||
<label className="xx-settings-label">昵称</label>
|
||||
<Input
|
||||
value={displayName}
|
||||
onChange={(e) => setDisplayName(e.target.value)}
|
||||
placeholder="请输入昵称"
|
||||
maxLength={20}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div className="xx-settings-field">
|
||||
<Button
|
||||
buttonType="primary"
|
||||
buttonSize="md"
|
||||
onClick={() => saveProfileMutation.mutate()}
|
||||
loading={saveProfileMutation.isPending}
|
||||
disabled={!displayName.trim() || !displayNameDirty}
|
||||
>
|
||||
保存
|
||||
</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,10 +6,16 @@ import { useAuthStore } from "@/store/authStore"
|
||||
export const ProtectedRoute = ({ children }: { children: React.ReactNode }) => {
|
||||
const isAuthenticated = useAuthStore((state) => state.isAuthenticated)
|
||||
const hasAccessToken = Boolean(localStorage.getItem("access_token"))
|
||||
const profileCompleted = useAuthStore((state) => state.user?.profile_completed !== false)
|
||||
|
||||
if (!isAuthenticated || !hasAccessToken) {
|
||||
return <Navigate to="/login" replace />
|
||||
}
|
||||
|
||||
// 微信新用户未完成昵称引导时,禁止进入主界面
|
||||
if (!profileCompleted) {
|
||||
return <Navigate to="/welcome/wechat" replace />
|
||||
}
|
||||
|
||||
return <>{children}</>
|
||||
}
|
||||
|
||||
@@ -5,6 +5,8 @@ import Register from "@/pages/auth/Register"
|
||||
import ForgotPassword from "@/pages/auth/ForgotPassword"
|
||||
import ResetPassword from "@/pages/auth/ResetPassword"
|
||||
import WechatCallback from "@/pages/auth/WechatCallback"
|
||||
import WechatOnboarding from "@/pages/auth/WechatOnboarding"
|
||||
import WechatBindCallback from "@/pages/auth/WechatBindCallback"
|
||||
import { useAuthStore } from "@/store/authStore"
|
||||
|
||||
/** 首页路由组件:已登录跳 dashboard,未登录显示落地页 */
|
||||
@@ -45,4 +47,12 @@ export const publicRoutes: RouteObject[] = [
|
||||
path: "/auth/wechat/callback",
|
||||
element: <WechatCallback />,
|
||||
},
|
||||
{
|
||||
path: "/auth/wechat/bind/callback",
|
||||
element: <WechatBindCallback />,
|
||||
},
|
||||
{
|
||||
path: "/welcome/wechat",
|
||||
element: <WechatOnboarding />,
|
||||
},
|
||||
]
|
||||
|
||||
@@ -13,6 +13,12 @@ interface User {
|
||||
display_name: string
|
||||
is_email_verified: boolean
|
||||
email_verified: boolean
|
||||
wechat_bound?: boolean
|
||||
wechat_nickname?: string
|
||||
avatar_url?: string
|
||||
phone?: string
|
||||
phone_verified?: boolean
|
||||
profile_completed?: boolean
|
||||
}
|
||||
|
||||
interface AuthState {
|
||||
|
||||
@@ -0,0 +1,101 @@
|
||||
/**
|
||||
* 上传去重/幂等工具单测(Issue #1714)
|
||||
*/
|
||||
import { describe, it, expect } from "vitest"
|
||||
import {
|
||||
computeFileHash,
|
||||
findDuplicateInQueue,
|
||||
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 大文件抽样(>256MB)", () => {
|
||||
it("抽样路径正常返回 64 位 hex,且大小不同则 hash 不同", async () => {
|
||||
// mock 一个「声称」300MB 的 File:slice 返回小 buffer 即可,不真分配 300MB
|
||||
const makeBig = (declaredSize: number, head: number) => {
|
||||
const f = new File([new Uint8Array([head, 2, 3])], "big.mov", { type: "video/quicktime" })
|
||||
Object.defineProperty(f, "size", { value: declaredSize, configurable: true })
|
||||
// slice 仍按真实内容返回小片段(头尾片段内容由底层小 buffer 决定)
|
||||
return f
|
||||
}
|
||||
const h1 = await computeFileHash(makeBig(300 * 1024 * 1024, 1))
|
||||
const h2 = await computeFileHash(makeBig(301 * 1024 * 1024, 1))
|
||||
expect(h1).toMatch(/^[0-9a-f]{64}$/)
|
||||
// 声明大小不同 → 写入的 64 位 size 字段不同 → hash 必须不同(锁定 setBigUint64 路径)
|
||||
expect(h1).not.toBe(h2)
|
||||
})
|
||||
})
|
||||
@@ -1,8 +1,8 @@
|
||||
import { describe, expect, it, vi } from "vitest"
|
||||
import { render, screen } from "@testing-library/react"
|
||||
import { describe, expect, it, vi, beforeEach } from "vitest"
|
||||
import { render, screen, fireEvent, waitFor } from "@testing-library/react"
|
||||
import { MemoryRouter } from "react-router-dom"
|
||||
import { QueryClient, QueryClientProvider } from "@tanstack/react-query"
|
||||
|
||||
// mock PageHead 简单mock
|
||||
vi.mock("@/components/layout/PageHead", () => ({
|
||||
default: ({ title, description }: { title: string; description?: string }) => (
|
||||
<div data-testid="page-head">
|
||||
@@ -12,51 +12,142 @@ 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: any) => any) =>
|
||||
selector({
|
||||
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 = {
|
||||
user: {
|
||||
id: "1",
|
||||
user_id: "1",
|
||||
username: "testuser",
|
||||
email: "test@example.com",
|
||||
display_name: "Test User",
|
||||
is_email_verified: true,
|
||||
email_verified: true,
|
||||
wechat_bound: false,
|
||||
},
|
||||
isAuthenticated: true,
|
||||
}),
|
||||
}))
|
||||
|
||||
import Settings from "@/pages/profile/Settings"
|
||||
|
||||
describe("Settings Page", () => {
|
||||
it("should render without crashing", () => {
|
||||
render(
|
||||
<MemoryRouter>
|
||||
<Settings />
|
||||
</MemoryRouter>,
|
||||
)
|
||||
expect(screen.getByText("个人设置")).toBeTruthy()
|
||||
setUser: mockSetUser,
|
||||
}
|
||||
})
|
||||
|
||||
it("should display user info", () => {
|
||||
render(
|
||||
<MemoryRouter>
|
||||
<Settings />
|
||||
</MemoryRouter>,
|
||||
)
|
||||
it("渲染个人设置与用户信息", () => {
|
||||
renderPage()
|
||||
expect(screen.getByText("个人设置")).toBeTruthy()
|
||||
expect(screen.getByDisplayValue("testuser")).toBeTruthy()
|
||||
expect(screen.getByDisplayValue("test@example.com")).toBeTruthy()
|
||||
})
|
||||
|
||||
it("should show save button is disabled", () => {
|
||||
render(
|
||||
<MemoryRouter>
|
||||
<Settings />
|
||||
</MemoryRouter>,
|
||||
it("未绑定时显示绑定微信按钮,点击跳转微信授权", async () => {
|
||||
renderPage()
|
||||
expect(screen.getByText("未绑定微信")).toBeTruthy()
|
||||
const btn = screen.getByText("绑定微信")
|
||||
fireEvent.click(btn)
|
||||
await waitFor(() => {
|
||||
expect(getWechatBindUrlMock).toHaveBeenCalled()
|
||||
expect(localStorage.getItem("wechat_bind_state")).toBe("s1")
|
||||
})
|
||||
})
|
||||
|
||||
it("已绑定时显示状态与解绑按钮,确认后调解绑接口", async () => {
|
||||
authState.user = {
|
||||
...(authState.user as object),
|
||||
wechat_bound: true,
|
||||
wechat_nickname: "微信昵称",
|
||||
} as never
|
||||
renderPage()
|
||||
expect(screen.getByText(/已绑定微信/)).toBeTruthy()
|
||||
fireEvent.click(
|
||||
screen.getByText(
|
||||
(_, el) => el?.tagName === "BUTTON" && (el.textContent ?? "").replace(/\s/g, "") === "解绑",
|
||||
),
|
||||
)
|
||||
const button = screen.getByText("保存暂未开放")
|
||||
expect(button).toBeTruthy()
|
||||
// antd Modal.confirm 弹确认框(标题+内容均含"解绑微信",用 role=dialog 内的确认按钮)
|
||||
await waitFor(() => {
|
||||
expect(document.querySelector(".ant-modal-confirm")).toBeTruthy()
|
||||
})
|
||||
fireEvent.click(
|
||||
screen.getByText(
|
||||
(_, el) => el?.tagName === "BUTTON" && (el.textContent ?? "").includes("确定解绑"),
|
||||
),
|
||||
)
|
||||
await waitFor(() => {
|
||||
expect(unbindWechatMock).toHaveBeenCalled()
|
||||
})
|
||||
})
|
||||
|
||||
it("修改昵称后保存按钮可用,点击调用更新接口", async () => {
|
||||
renderPage()
|
||||
const saveBtn = screen.getByText(
|
||||
(_, el) => el?.tagName === "BUTTON" && (el.textContent ?? "").replace(/\s/g, "") === "保存",
|
||||
)
|
||||
expect(saveBtn.closest("button")?.disabled).toBe(true)
|
||||
fireEvent.change(screen.getByDisplayValue("Test User"), {
|
||||
target: { value: "新昵称" },
|
||||
})
|
||||
await waitFor(() => {
|
||||
expect(saveBtn.closest("button")?.disabled).toBe(false)
|
||||
})
|
||||
updateProfileMock.mockResolvedValueOnce({
|
||||
id: "1",
|
||||
display_name: "新昵称",
|
||||
wechat_bound: false,
|
||||
})
|
||||
fireEvent.click(saveBtn)
|
||||
await waitFor(() => {
|
||||
expect(updateProfileMock).toHaveBeenCalledWith({ display_name: "新昵称" })
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
@@ -31,17 +31,23 @@ 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 返回挂起的 promise,
|
||||
* finish 槽位在 transfer executor 同步执行时挂载,测试中调用 finish() 控制成败
|
||||
* 创建一个假 handle:
|
||||
* - transfer 返回挂起的 promise,finish()/finish(true) 控制成败
|
||||
* - complete 每次调用返回独立的挂起 promise,由 completeCalls 记录控制,
|
||||
* 成功调 resolve(idx) / 失败调 reject(idx)(模拟超时)
|
||||
*/
|
||||
const makeFakeHandle = (opts: { id: string; duplicated?: boolean; failTransfer?: boolean }) => {
|
||||
const h = {
|
||||
const makeFakeHandle = (opts: {
|
||||
id: string
|
||||
duplicated?: boolean
|
||||
failTransfer?: boolean
|
||||
completeAuto?: boolean
|
||||
}) => {
|
||||
const h: FakeHandle = {
|
||||
prepared: {
|
||||
upload_url: "https://oss.example.com/u",
|
||||
method: "POST",
|
||||
@@ -52,15 +58,39 @@ const makeFakeHandle = (opts: { id: string; duplicated?: boolean; failTransfer?:
|
||||
asset_id: opts.id,
|
||||
},
|
||||
transfer: vi.fn(),
|
||||
complete: vi.fn().mockResolvedValue({
|
||||
storage_key: "uploads/x/y.mp4",
|
||||
ingest_job_id: opts.duplicated ? "" : "job-1",
|
||||
url: "https://oss.example.com/u",
|
||||
duplicated: opts.duplicated,
|
||||
asset_id: opts.id,
|
||||
}),
|
||||
finish: (() => {}) as (fail?: boolean) => void,
|
||||
complete: vi.fn(),
|
||||
finish: () => {},
|
||||
completeCalls: [],
|
||||
}
|
||||
|
||||
h.complete.mockImplementation(
|
||||
() =>
|
||||
new Promise<{
|
||||
storage_key: string
|
||||
ingest_job_id: string
|
||||
url: string
|
||||
duplicated: boolean
|
||||
asset_id: string
|
||||
}>((resolve, reject) => {
|
||||
h.completeCalls.push({
|
||||
resolve: () =>
|
||||
resolve({
|
||||
storage_key: "uploads/x/y.mp4",
|
||||
ingest_job_id: opts.duplicated ? "" : `job-${opts.id}`,
|
||||
url: "https://oss.example.com/u",
|
||||
duplicated: !!opts.duplicated,
|
||||
asset_id: opts.id,
|
||||
}),
|
||||
reject,
|
||||
})
|
||||
// 默认立即成功,保持旧用例简单
|
||||
if (opts.completeAuto !== false) {
|
||||
const idx = h.completeCalls.length - 1
|
||||
Promise.resolve().then(() => h.completeCalls[idx]?.resolve())
|
||||
}
|
||||
}),
|
||||
)
|
||||
|
||||
h.transfer.mockImplementation(
|
||||
() =>
|
||||
new Promise<void>((_resolve, reject) => {
|
||||
@@ -78,17 +108,24 @@ const makeFakeHandle = (opts: { id: string; duplicated?: boolean; failTransfer?:
|
||||
|
||||
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 },
|
||||
optOverrides?: (id: string) => {
|
||||
duplicated?: boolean
|
||||
failTransfer?: boolean
|
||||
completeAuto?: 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, ...overrides })
|
||||
const h = makeFakeHandle({ id, completeAuto: true, ...overrides })
|
||||
handles.push(h)
|
||||
await new Promise((r) => setTimeout(r, 10))
|
||||
return h
|
||||
@@ -223,4 +260,96 @@ 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")
|
||||
})
|
||||
|
||||
// 点重试:pump 复用 handle,只再调一次 complete(transfer/prepare 不重复)
|
||||
await act(async () => {
|
||||
result.current.retryUpload(tempId)
|
||||
})
|
||||
await waitFor(() => expect(handles[0].complete).toHaveBeenCalledTimes(2))
|
||||
expect(handles.length).toBe(1) // 没有重新 prepare
|
||||
expect(handles[0].transfer).toHaveBeenCalledTimes(1) // 没有重新直传
|
||||
|
||||
// 第二次 complete 成功
|
||||
await act(async () => {
|
||||
handles[0].completeCalls[1]?.resolve()
|
||||
})
|
||||
await waitFor(() => {
|
||||
expect(result.current.uploadItems.find((x) => x.tempId === tempId)?.status).toBe("done")
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
@@ -1,79 +1,127 @@
|
||||
import { describe, expect, it, vi, beforeEach } from "vitest"
|
||||
import { render, screen } from "@testing-library/react"
|
||||
import { describe, expect, it, vi, beforeEach, afterEach } from "vitest"
|
||||
import { render, screen, waitFor, cleanup } from "@testing-library/react"
|
||||
import { MemoryRouter } from "react-router-dom"
|
||||
import WechatCallback from "@/pages/auth/WechatCallback"
|
||||
|
||||
const mockNavigate = vi.fn()
|
||||
const mockSetAuth = vi.fn()
|
||||
const mockSearchParams = [new URLSearchParams({ code: "test_code", state: "test_state" })] 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 callbackShouldFail = false
|
||||
|
||||
vi.mock("react-router-dom", async () => {
|
||||
const actual = await vi.importActual("react-router-dom")
|
||||
return {
|
||||
...actual,
|
||||
useNavigate: () => vi.fn(),
|
||||
useSearchParams: () => [new URLSearchParams({ code: "test_code", state: "test_state" })],
|
||||
useNavigate: () => mockNavigate,
|
||||
useSearchParams: () => mockSearchParams,
|
||||
}
|
||||
})
|
||||
|
||||
vi.mock("@/api/auth", () => ({
|
||||
wechatCallback: vi.fn(() => new Promise(() => {})), // pending promise,保持loading
|
||||
getCurrentUser: vi.fn(),
|
||||
wechatCallback: vi.fn(async () => {
|
||||
if (callbackShouldFail) throw new Error("fail")
|
||||
return mockCallbackResult
|
||||
}),
|
||||
getCurrentUser: vi.fn(async () => mockCurrentUser),
|
||||
normalizeUser: (u: unknown) => u,
|
||||
}))
|
||||
|
||||
vi.mock("@/api/auth/tokenRefresh", () => ({
|
||||
scheduleProactiveRefresh: vi.fn(),
|
||||
cancelProactiveRefresh: vi.fn(),
|
||||
}))
|
||||
|
||||
vi.mock("@/store/authStore", () => ({
|
||||
useAuthStore: () => ({
|
||||
setAuth: vi.fn(),
|
||||
}),
|
||||
useAuthStore: (selector: (state: unknown) => unknown) => selector({ setAuth: mockSetAuth }),
|
||||
}))
|
||||
|
||||
vi.mock("@/components/auth/BindContactModal", () => ({
|
||||
default: ({ open }: { open: boolean }) => (
|
||||
<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(),
|
||||
},
|
||||
}
|
||||
})
|
||||
const renderPage = () =>
|
||||
render(
|
||||
<MemoryRouter>
|
||||
<WechatCallback />
|
||||
</MemoryRouter>,
|
||||
)
|
||||
|
||||
describe("WechatCallback Page", () => {
|
||||
afterEach(() => {
|
||||
cleanup()
|
||||
})
|
||||
|
||||
beforeEach(() => {
|
||||
// mock localStorage,设置wechat_state匹配,让校验通过
|
||||
const store: Record<string, string> = {
|
||||
wechat_state: "test_state",
|
||||
vi.clearAllMocks()
|
||||
callbackShouldFail = false
|
||||
localStorageStore.wechat_state = "test_state"
|
||||
mockCallbackResult = {
|
||||
access_token: "at",
|
||||
refresh_token: "rt",
|
||||
is_new_user: false,
|
||||
}
|
||||
vi.spyOn(Storage.prototype, "getItem").mockImplementation((key) => store[key] || null)
|
||||
vi.spyOn(Storage.prototype, "setItem").mockImplementation((key, val) => {
|
||||
store[key] = val
|
||||
mockCurrentUser = {
|
||||
id: "u1",
|
||||
display_name: "老用户",
|
||||
profile_completed: true,
|
||||
}
|
||||
})
|
||||
|
||||
it("老用户登录成功跳转首页/来源页", async () => {
|
||||
renderPage()
|
||||
await waitFor(() => {
|
||||
expect(mockNavigate).toHaveBeenCalledWith("/", { replace: true })
|
||||
})
|
||||
vi.spyOn(Storage.prototype, "removeItem").mockImplementation((key) => {
|
||||
delete store[key]
|
||||
expect(mockSetAuth).toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it("新用户(is_new_user)跳转昵称引导页", async () => {
|
||||
mockCallbackResult = { access_token: "at", refresh_token: "rt", is_new_user: true }
|
||||
mockCurrentUser = { id: "u2", display_name: "微信用户", profile_completed: false }
|
||||
renderPage()
|
||||
await waitFor(() => {
|
||||
expect(mockNavigate).toHaveBeenCalledWith("/welcome/wechat", { replace: true })
|
||||
})
|
||||
})
|
||||
|
||||
it("should render without crashing", () => {
|
||||
const { container } = render(
|
||||
<MemoryRouter>
|
||||
<WechatCallback />
|
||||
</MemoryRouter>,
|
||||
)
|
||||
expect(container).toBeTruthy()
|
||||
it("is_new_user=false 但 profile_completed=false(上次中断)也跳引导页", async () => {
|
||||
mockCallbackResult = { access_token: "at", refresh_token: "rt", is_new_user: false }
|
||||
mockCurrentUser = { id: "u3", display_name: "微信用户", profile_completed: false }
|
||||
renderPage()
|
||||
await waitFor(() => {
|
||||
expect(mockNavigate).toHaveBeenCalledWith("/welcome/wechat", { replace: true })
|
||||
})
|
||||
})
|
||||
|
||||
it("should show loading state while processing", () => {
|
||||
render(
|
||||
<MemoryRouter>
|
||||
<WechatCallback />
|
||||
</MemoryRouter>,
|
||||
)
|
||||
// wechatCallback 返回 pending promise,所以应该显示 loading
|
||||
expect(screen.getByText("正在登录...")).toBeTruthy()
|
||||
it("state 不匹配显示安全错误", async () => {
|
||||
localStorageStore.wechat_state = "other_state"
|
||||
renderPage()
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("安全校验失败,请重新登录")).toBeTruthy()
|
||||
})
|
||||
expect(mockNavigate).not.toHaveBeenCalled()
|
||||
})
|
||||
|
||||
it("接口失败显示错误提示", async () => {
|
||||
callbackShouldFail = true
|
||||
renderPage()
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("微信登录失败,请重试")).toBeTruthy()
|
||||
})
|
||||
})
|
||||
|
||||
it("处理中显示 loading", () => {
|
||||
renderPage()
|
||||
expect(screen.getByText("微信登录中...")).toBeTruthy()
|
||||
})
|
||||
})
|
||||
|
||||
@@ -0,0 +1,132 @@
|
||||
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("新用户可见昵称表单并能提交", 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()
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -20,7 +20,8 @@ import "@/pages/generate/components/Step2MaterialSelect"
|
||||
import "@/pages/generate/components/Step4TitleSettings"
|
||||
import "@/pages/generate/components/Step5VoiceSelect"
|
||||
import "@/pages/generate/components/Step3VoiceWithMode"
|
||||
import "@/pages/generate/components/ServerPreviewGrid"
|
||||
import "@/pages/generate/components/CanvasPreviewGrid"
|
||||
import "@/pages/generate/components/BatchGenerationGrid"
|
||||
import "@/pages/generate/components/PreviewCountModal"
|
||||
import "@/pages/generate/components/PreviewVideoPanel"
|
||||
import "@/pages/generate/components/GenerateResultPanel"
|
||||
@@ -48,7 +49,6 @@ describe("GeneratePage module smoke test", () => {
|
||||
})
|
||||
})
|
||||
import "@/pages/generate/hooks/useGenerateVideo"
|
||||
import "@/pages/generate/hooks/useBatchPreview"
|
||||
import "@/pages/generate/hooks/useBatchCovers"
|
||||
import "@/pages/generate/hooks/usePreviewAssets"
|
||||
import "@/pages/generate/hooks/useSegmentScheduler"
|
||||
|
||||
@@ -31,25 +31,65 @@ SCENE_CHANGE_THRESHOLD = 30 # 灰度差异阈值
|
||||
MIN_KEYFRAME_INTERVAL_SEC = 1.0 # 最小关键帧间隔(秒)
|
||||
MAX_KEYFRAMES = 30 # 最大关键帧数
|
||||
MIN_KEYFRAMES = 5 # 最小关键帧数
|
||||
FINGERPRINT_SAMPLE_INTERVAL_SEC = 1.0 # 指纹采样间隔(秒):密集均匀采样,保证两视频时序可对齐
|
||||
FINGERPRINT_MAX_SAMPLES = 30 # 长视频采样数上限(超过后采样间隔自动放宽)
|
||||
LONG_VIDEO_SEGMENT_SEC = 30 # 长视频每段秒数
|
||||
LONG_VIDEO_DURATION_THRESHOLD_SEC = 180 # 3 分钟阈值
|
||||
MIN_FRAMES_PER_SEGMENT = 2 # 长视频每段最少帧数
|
||||
|
||||
# ── 滑动窗口匹配常量 ────────────────────────────────────────────
|
||||
SEGMENT_MATCH_THRESHOLD = 8 # 帧匹配汉明距离阈值
|
||||
MIN_CONSECUTIVE_MATCHES = 5 # 最少连续匹配帧数
|
||||
# ── 滑动窗口匹配常量(Issue #1702 二次校准) ─────────────────────
|
||||
# 阈值经 staging 真实数据两轮回归校准(worker 容器内离线实验):
|
||||
# 第一轮(2026-09-05):同源对 <=12 命中 4/11,异源最小距离 24 → 定 12;
|
||||
# 第二轮(2026-09-05,证据视频 B->A 仍漏检):扩大样本到该用户全部
|
||||
# 15 个真实成片(13 个异源候选)实测:
|
||||
# - 同源成片对(A 20s / B、C 各 11.75s,1s 密集采样):
|
||||
# B->A 中位数距离 14,<=16 命中 8/11=0.73;C->A 8/11=0.73
|
||||
# - 异源成片对(13 个真实视频):每帧全局最近邻最小距离 18,
|
||||
# <=16 命中帧数全部为 0(最近邻 18 仅个别帧,中位数 22~28)
|
||||
# 12 漏掉同源降重对(降重滤镜/字幕/画面扰动把距离从 ~8 推到 14~16);
|
||||
# 16 对同源命中 0.73+ 且与异源分布(最近邻 >=18)仍有 >=2bit 安全裕度,
|
||||
# 异源 <=16 命中 0 帧,无误报空间。
|
||||
PHASH_THRESHOLD = 16
|
||||
SEGMENT_MATCH_THRESHOLD = PHASH_THRESHOLD # 片段匹配阈值与帧匹配统一(#1702:阈值常量统一来源)
|
||||
MIN_CONSECUTIVE_MATCHES = 5 # 连续匹配默认门槛;短视频自适应 min(5, max(2, 分片数//2))
|
||||
MAX_GAP = 2 # 允许的最大间隙帧数
|
||||
NEIGHBOR_WINDOW = 1 # 分片时序对齐:允许 ±1 邻接偏移(1s 密集采样下即 ±1s,缓解切点不一致)
|
||||
|
||||
# ── 融合判定常量 ────────────────────────────────────────────────
|
||||
PHASH_WEIGHT = 0.7 # pHash 权重
|
||||
HISTOGRAM_WEIGHT = 0.3 # 直方图权重
|
||||
MATCH_RATIO_THRESHOLD = 0.7 # 至少 70% 帧匹配
|
||||
MATCH_RATIO_THRESHOLD = 0.7 # 全片重复(is_duplicate)至少 70% 帧匹配
|
||||
PARTIAL_COVERAGE_THRESHOLD = 0.5 # 局部复用覆盖率 >=50% 也判全片重复
|
||||
DUPLICATE_THRESHOLD = 0.70 # 融合后相似度阈值
|
||||
|
||||
# ── 降重裁剪规避常量(Issue #1702) ─────────────────────────────
|
||||
# 成片强制 2-5% random_edge_crop 降重只服务外部平台;自查重指纹取中心 90%
|
||||
# 区域,使两次不同裁剪的同源画面 pHash 距离回到同分布。
|
||||
FINGERPRINT_CENTER_CROP_RATIO = 0.90
|
||||
|
||||
|
||||
# ── 感知哈希 & 颜色直方图工具函数 ────────────────────────────────
|
||||
|
||||
|
||||
def center_crop_frame(image: np.ndarray, ratio: float = FINGERPRINT_CENTER_CROP_RATIO) -> np.ndarray:
|
||||
"""取画面中心 ratio 比例区域(裁除四边边缘)。
|
||||
|
||||
查重指纹用:random_edge_crop 降重(2-5% 四边随机裁剪)会让同源画面 pHash
|
||||
位翻转 12-16,污染自查重(Issue #1702)。算 pHash/颜色直方图前先居中裁除
|
||||
边缘 10%,两次不同裁剪的同源画面中心区域基本重合,指纹不再被降重污染。
|
||||
降重只服务外部平台,不影响内部查重。
|
||||
"""
|
||||
if image is None or image.size == 0:
|
||||
return image
|
||||
h, w = image.shape[:2]
|
||||
ch, cw = int(h * ratio), int(w * ratio)
|
||||
if ch <= 0 or cw <= 0 or (ch >= h and cw >= w):
|
||||
return image
|
||||
y0 = (h - ch) // 2
|
||||
x0 = (w - cw) // 2
|
||||
return image[y0 : y0 + ch, x0 : x0 + cw]
|
||||
|
||||
|
||||
def compute_phash(image: np.ndarray, hash_size: int = 8) -> str:
|
||||
"""计算图像的感知哈希(pHash),基于 DCT(离散余弦变换)。
|
||||
|
||||
@@ -101,11 +141,17 @@ def hamming_distance(hash1: str, hash2: str) -> int:
|
||||
|
||||
|
||||
def compute_color_histogram(image: np.ndarray, bins: int = 32) -> list[float]:
|
||||
"""Compute color histogram for an image."""
|
||||
"""Compute BGR color histogram for an image.
|
||||
|
||||
Issue #1702: 每个通道独立做 NORM_L1 归一化(通道内 Σ=1,是概率分布),
|
||||
三通道拼接存储。Bhattacharyya 系数对拼接向量直接 Σ√(a*b) 会得到
|
||||
3 通道之和(范围 [0,3],实测 ~14.9 是旧 L2 归一化的错误结果),
|
||||
消费方 _bhattacharyya_coefficient 按通道数平均归一到 [0,1]。
|
||||
"""
|
||||
hist = []
|
||||
for i in range(3):
|
||||
h = cv2.calcHist([image], [i], None, [bins], [0, 256])
|
||||
h = cv2.normalize(h, h).flatten()
|
||||
h = cv2.normalize(h, h, norm_type=cv2.NORM_L1).flatten()
|
||||
hist.extend(h)
|
||||
return hist
|
||||
|
||||
@@ -210,6 +256,30 @@ def detect_keyframe_timestamps(
|
||||
return keyframe_times
|
||||
|
||||
|
||||
def sample_fingerprint_timestamps(
|
||||
duration: float,
|
||||
*,
|
||||
interval_sec: float = FINGERPRINT_SAMPLE_INTERVAL_SEC,
|
||||
max_samples: int = FINGERPRINT_MAX_SAMPLES,
|
||||
) -> list[float]:
|
||||
"""指纹采样时间戳:固定间隔密集均匀采样(Issue #1702)。
|
||||
|
||||
动态场景检测抽帧(#1659)在两个同源视频上会各自取到不同时刻,切点/取帧
|
||||
错位让对齐帧的 pHash 距离都很大(实测同源对最小距离 12 且配对时序错乱)。
|
||||
改为固定 1s 间隔均匀采样后,复用片段的帧时刻天然对齐,配合 ±1 邻接窗口
|
||||
即可检出同源/局部复用。长视频(>max_samples*interval)自动放宽间隔到
|
||||
duration/max_samples,保证分片数有上限。
|
||||
"""
|
||||
if duration <= 0:
|
||||
return []
|
||||
step = interval_sec
|
||||
n_uniform = int(duration / step)
|
||||
if n_uniform > max_samples:
|
||||
step = duration / max_samples
|
||||
count = max(1, int(duration / step))
|
||||
return [step * (i + 0.5) for i in range(count)]
|
||||
|
||||
|
||||
# ── 数据类 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -297,23 +367,33 @@ def find_duplicate_segments(
|
||||
target_chunks: list,
|
||||
*,
|
||||
match_threshold: int = SEGMENT_MATCH_THRESHOLD,
|
||||
min_consecutive: int = MIN_CONSECUTIVE_MATCHES,
|
||||
min_consecutive: Optional[int] = None,
|
||||
max_gap: int = MAX_GAP,
|
||||
neighbor_window: int = NEIGHBOR_WINDOW,
|
||||
) -> list[DuplicateSegment]:
|
||||
"""滑动窗口时序匹配:找出两组分片之间的重复片段。
|
||||
"""滑动窗口时序匹配:找出两组分片之间的重复片段(Issue #1702 重构)。
|
||||
|
||||
算法:
|
||||
1. 对每个 query chunk,找到 target 中汉明距离最小的 chunk
|
||||
2. 距离 <= match_threshold 视为匹配
|
||||
3. 找连续匹配的 run(允许 max_gap 帧间隙)
|
||||
4. 连续匹配数 >= min_consecutive 的 run 报告为重复片段
|
||||
1. 构建 query×target 全量汉明距离矩阵;每个 query chunk 保留所有
|
||||
距离 <= match_threshold 的候选 target 分片(与帧匹配判定同一阈值)。
|
||||
2. 时序一致贪心对齐:沿 query 时序推进,run 内优先选择与上一匹配帧
|
||||
目标序号连贯(|delta| <= neighbor_window+1,允许 ±1 邻接/时序偏移
|
||||
对齐——1s 密集采样下相邻帧 pHash 接近,最近邻在目标相邻帧间
|
||||
正/反向跳变均属正常,缓解场景切割切点、取帧错位、局部倒退)的
|
||||
候选;同距时偏好小索引(最早对齐位置)。
|
||||
3. 连贯匹配中允许 <= max_gap 帧间隙桥接;断裂后另起新 run——天然
|
||||
支持局部片段复用(复用片段可出现在任意时序位置,各成独立片段)。
|
||||
4. 连续匹配帧数 >= min_consecutive 的 run 报为重复片段。短视频自适应:
|
||||
min_consecutive = min(5, max(2, len(query_chunks)//2));n=1 时
|
||||
不形成片段,由调用方匹配帧回退兜底。
|
||||
|
||||
Args:
|
||||
query_chunks: 查询视频的分片列表(FingerprintChunk 或 dict)
|
||||
target_chunks: 目标视频的分片列表
|
||||
match_threshold: 汉明距离匹配阈值
|
||||
min_consecutive: 最少连续匹配帧数
|
||||
match_threshold: 汉明距离匹配阈值(统一常量 PHASH_THRESHOLD)
|
||||
min_consecutive: 最少连续匹配帧数;None 时按短视频自适应
|
||||
max_gap: 允许的最大间隙帧数
|
||||
neighbor_window: 时序对齐允许的目标分片序号邻接窗口(正/反向均允许)
|
||||
|
||||
Returns:
|
||||
DuplicateSegment 列表
|
||||
@@ -321,95 +401,94 @@ def find_duplicate_segments(
|
||||
if not query_chunks or not target_chunks:
|
||||
return []
|
||||
|
||||
def _get_phash(chunk) -> str:
|
||||
def _get(chunk, key):
|
||||
if isinstance(chunk, dict):
|
||||
return chunk["phash_binary"]
|
||||
return chunk.phash_binary
|
||||
return chunk[key]
|
||||
return getattr(chunk, key)
|
||||
|
||||
def _get_start(chunk) -> int:
|
||||
if isinstance(chunk, dict):
|
||||
return chunk["start_time_ms"]
|
||||
return chunk.start_time_ms
|
||||
n, m = len(query_chunks), len(target_chunks)
|
||||
q_ph = [_get(c, "phash_binary") for c in query_chunks]
|
||||
t_ph = [_get(c, "phash_binary") for c in target_chunks]
|
||||
|
||||
def _get_end(chunk) -> int:
|
||||
if isinstance(chunk, dict):
|
||||
return chunk["end_time_ms"]
|
||||
return chunk.end_time_ms
|
||||
# Step 1: 全量距离矩阵。每个 query chunk 保留所有 <= 阈值的候选 target,
|
||||
# 按距离升序;同距时小索引优先(取最早的对齐位置,贪心连贯推进时最保守,
|
||||
# 不会越过复用片段末端;重复 hash 的连续帧由 Step 2 的连贯性窗口约束)。
|
||||
candidates: list[list[tuple[int, int]]] = [] # 每 query 帧: [(target_idx, dist), ...]
|
||||
for i in range(n):
|
||||
dists = [hamming_distance(q_ph[i], t_ph[j]) for j in range(m)]
|
||||
cand = [(j, d) for j, d in enumerate(dists) if d <= match_threshold]
|
||||
cand.sort(key=lambda x: (x[1], x[0]))
|
||||
candidates.append(cand)
|
||||
|
||||
# Step 1: 逐帧匹配
|
||||
frame_matches: list[tuple[bool, int, int]] = [] # (is_match, min_dist, best_target_idx)
|
||||
for qc in query_chunks:
|
||||
qc_phash = _get_phash(qc)
|
||||
best_dist = 64
|
||||
best_idx = 0
|
||||
for j, tc in enumerate(target_chunks):
|
||||
d = hamming_distance(qc_phash, _get_phash(tc))
|
||||
if d < best_dist:
|
||||
best_dist = d
|
||||
best_idx = j
|
||||
frame_matches.append((best_dist <= match_threshold, best_dist, best_idx))
|
||||
# 短视频自适应连续匹配门槛(Issue #1702 工单公式):
|
||||
# MIN_CONSECUTIVE_MATCHES = min(5, max(2, 分片数//2))。
|
||||
# n=1 时门槛为 2 不形成片段,由 _evaluate_candidate 的匹配帧回退
|
||||
# (temporal_coverage 按匹配帧占比估计)兜底检出,不回归。
|
||||
if min_consecutive is None:
|
||||
min_consecutive = min(MIN_CONSECUTIVE_MATCHES, max(2, n // 2))
|
||||
|
||||
# Step 2: 找连续匹配的 runs
|
||||
runs: list[tuple[int, int]] = [] # list of (start_idx, end_idx)
|
||||
run_start = None
|
||||
# Step 2: 时序一致贪心对齐。
|
||||
# run 内偏好与上一匹配帧目标序号连贯(|delta| <= neighbor_window+1,
|
||||
# 支持 ±1 邻接窗口/时序偏移对齐,正反向抖动均允许)的候选;
|
||||
# 无连贯候选时关闭旧 run。
|
||||
# 这天然支持局部片段复用:同一 query 视频中多个复用片段各自形成独立 run。
|
||||
frame_matches: list[tuple[bool, int, int]] = []
|
||||
runs: list[tuple[int, int]] = []
|
||||
run_start: Optional[int] = None
|
||||
run_last_t: Optional[int] = None
|
||||
gap_count = 0
|
||||
|
||||
for i, (is_match, _dist, _idx) in enumerate(frame_matches):
|
||||
if is_match:
|
||||
def _matching_count(a: int, b: int) -> int:
|
||||
return sum(1 for k in range(a, b + 1) if frame_matches[k][0])
|
||||
|
||||
def _close_run(a: int, b: int) -> None:
|
||||
if b >= a and _matching_count(a, b) >= min_consecutive:
|
||||
runs.append((a, b))
|
||||
|
||||
for i in range(n):
|
||||
cand = candidates[i]
|
||||
if run_last_t is None:
|
||||
chosen = cand[0] if cand else None
|
||||
else:
|
||||
chosen = next(
|
||||
(c for c in cand if abs(c[0] - run_last_t) <= neighbor_window + 1),
|
||||
None,
|
||||
)
|
||||
|
||||
if chosen is not None:
|
||||
tidx, dist = chosen
|
||||
frame_matches.append((True, dist, tidx))
|
||||
if run_start is None:
|
||||
run_start = i
|
||||
gap_count = 0 # 重置间隙
|
||||
gap_count = 0
|
||||
run_last_t = tidx
|
||||
else:
|
||||
frame_matches.append((False, match_threshold + 1, -1))
|
||||
if run_start is not None:
|
||||
gap_count += 1
|
||||
if gap_count > max_gap:
|
||||
# 中断当前 run
|
||||
run_end = i - gap_count # 最后一个匹配帧的索引
|
||||
# 计算 run 内的实际匹配帧数(总跨度 - 间隙数)
|
||||
total_gaps = sum(1 for k in range(run_start, run_end + 1) if not frame_matches[k][0])
|
||||
matching_count = (run_end - run_start + 1) - total_gaps
|
||||
if matching_count >= min_consecutive:
|
||||
runs.append((run_start, run_end))
|
||||
run_start = None
|
||||
gap_count = 0
|
||||
# 非匹配帧从 i-gap_count+1 开始,run 结束于其前一帧
|
||||
_close_run(run_start, i - gap_count)
|
||||
run_start, run_last_t, gap_count = None, None, 0
|
||||
|
||||
# 处理末尾 run
|
||||
if run_start is not None:
|
||||
last_idx = len(frame_matches) - 1
|
||||
# 回退找到最后一个匹配帧的位置(跳过尾部非匹配帧)
|
||||
last_idx = n - 1
|
||||
while last_idx >= run_start and not frame_matches[last_idx][0]:
|
||||
last_idx -= 1
|
||||
if last_idx >= run_start:
|
||||
# 计算 run 内的总间隙数
|
||||
total_gaps = sum(1 for k in range(run_start, last_idx + 1) if not frame_matches[k][0])
|
||||
matching_count = (last_idx - run_start + 1) - total_gaps
|
||||
if matching_count >= min_consecutive:
|
||||
runs.append((run_start, last_idx))
|
||||
_close_run(run_start, last_idx)
|
||||
|
||||
# Step 3: 构建 DuplicateSegment
|
||||
segments: list[DuplicateSegment] = []
|
||||
for start, end in runs:
|
||||
query_start = _get_start(query_chunks[start])
|
||||
query_end = _get_end(query_chunks[end])
|
||||
|
||||
# 取目标范围(按最佳匹配的目标 chunk 时间范围)
|
||||
target_indices = [frame_matches[k][2] for k in range(start, end + 1) if frame_matches[k][0]]
|
||||
if target_indices:
|
||||
t_min = min(target_indices)
|
||||
t_max = max(target_indices)
|
||||
target_start = _get_start(target_chunks[t_min])
|
||||
target_end = _get_end(target_chunks[t_max])
|
||||
else:
|
||||
target_start = _get_start(target_chunks[0])
|
||||
target_end = _get_end(target_chunks[-1])
|
||||
|
||||
avg_dist = sum(frame_matches[k][1] for k in range(start, end + 1)) / (end - start + 1)
|
||||
t_min, t_max = min(target_indices), max(target_indices)
|
||||
avg_dist = sum(frame_matches[k][1] for k in range(start, end + 1) if frame_matches[k][0]) / len(target_indices)
|
||||
segments.append(
|
||||
DuplicateSegment(
|
||||
query_start_ms=query_start,
|
||||
query_end_ms=query_end,
|
||||
target_start_ms=target_start,
|
||||
target_end_ms=target_end,
|
||||
query_start_ms=_get(query_chunks[start], "start_time_ms"),
|
||||
query_end_ms=_get(query_chunks[end], "end_time_ms"),
|
||||
target_start_ms=_get(target_chunks[t_min], "start_time_ms"),
|
||||
target_end_ms=_get(target_chunks[t_max], "end_time_ms"),
|
||||
avg_distance=avg_dist,
|
||||
)
|
||||
)
|
||||
@@ -423,7 +502,9 @@ def find_duplicate_segments(
|
||||
class VideoDeduplicator:
|
||||
"""Video deduplication using multiple fingerprint methods."""
|
||||
|
||||
PHASH_THRESHOLD = 8 # Issue #1658: pHash 汉明距离阈值由 10 收紧到 8,降低不同视频误判率
|
||||
# Issue #1702: 阈值统一来源为模块常量 PHASH_THRESHOLD(#1658 曾收紧到 8,
|
||||
# 后经 staging 真实同源/异源指纹分布重新校准,见 test_phash_threshold_calibration_1702)。
|
||||
PHASH_THRESHOLD = PHASH_THRESHOLD
|
||||
HISTOGRAM_THRESHOLD = 0.85
|
||||
|
||||
@staticmethod
|
||||
@@ -447,27 +528,36 @@ class VideoDeduplicator:
|
||||
# 单帧不视为坏指纹(短视频或抽帧不足)
|
||||
if len(phashes) == 1:
|
||||
return False
|
||||
# 多帧但所有 phash 完全相同 → 黑屏/纯色视频
|
||||
# Issue #1702: 旧逻辑"所有 phash 完全相同即判黑屏"会误杀短视频——
|
||||
# 11s 视频只有几个不同镜头时,相邻 1s 采样帧可能 phash 完全一致(内容
|
||||
# 连续但非黑屏)。黑屏的特征是「大量帧全部无内容」,要求至少 8 帧
|
||||
# 且相同帧占比 >=80% 才判坏;短视频(<8 帧)只有真正单值时交给
|
||||
# _bhattacharyya/融合分兜底,不因"帧都一样"直接跳过。
|
||||
if len(phashes) < 8:
|
||||
return False
|
||||
unique = set(phashes)
|
||||
if len(unique) == 1:
|
||||
same_ratio = sum(1 for x in phashes if x == phashes[0]) / len(phashes)
|
||||
if len(unique) == 1 and same_ratio >= 0.8:
|
||||
return True
|
||||
# 多帧但所有 phash 之间的汉明距离都极小(<3)→ 近似黑屏
|
||||
# 多帧但所有唯一 phash 之间的汉明距离都极小(<3)且占比 >=80% → 近似黑屏
|
||||
phash_list = list(unique)
|
||||
if len(phash_list) >= 2:
|
||||
all_distances = []
|
||||
for i in range(len(phash_list)):
|
||||
for j in range(i + 1, len(phash_list)):
|
||||
all_distances.append(hamming_distance(phash_list[i], phash_list[j]))
|
||||
if len(phash_list) >= 2 and same_ratio >= 0.8:
|
||||
all_distances = [
|
||||
hamming_distance(phash_list[i], phash_list[j])
|
||||
for i in range(len(phash_list))
|
||||
for j in range(i + 1, len(phash_list))
|
||||
]
|
||||
if all_distances and max(all_distances) < 3:
|
||||
return True
|
||||
return False
|
||||
|
||||
def compute_fingerprint(self, video_path: str) -> VideoFingerprint:
|
||||
"""Compute video fingerprint using dynamic keyframe detection.
|
||||
"""Compute video fingerprint using dense uniform sampling.
|
||||
|
||||
使用 detect_keyframe_timestamps() 检测内容感知关键帧,
|
||||
在每个关键帧处取帧计算 pHash + color_histogram。
|
||||
同时保留 MD5 计算和分片数据结构。
|
||||
Issue #1702: 使用 sample_fingerprint_timestamps() 固定 1s 间隔密集均匀
|
||||
采样(替代动态场景检测抽帧),保证两个同源视频复用片段的帧时刻天然
|
||||
对齐;每帧取中心 90% 区域(center_crop_frame)计算 pHash + color_histogram,
|
||||
绕开 random_edge_crop 降重裁剪污染;MD5 仍基于原始帧。
|
||||
"""
|
||||
cap = cv2.VideoCapture(video_path)
|
||||
if not cap.isOpened():
|
||||
@@ -481,8 +571,8 @@ class VideoDeduplicator:
|
||||
|
||||
cap.release()
|
||||
|
||||
# 1. 检测关键帧时间戳
|
||||
keyframe_times = detect_keyframe_timestamps(video_path)
|
||||
# 1. 固定间隔密集采样(Issue #1702:替代动态场景检测,保证跨视频时序对齐)
|
||||
keyframe_times = sample_fingerprint_timestamps(duration)
|
||||
|
||||
if not keyframe_times:
|
||||
return VideoFingerprint(
|
||||
@@ -506,12 +596,15 @@ class VideoDeduplicator:
|
||||
if not ret:
|
||||
continue
|
||||
|
||||
# MD5 计算
|
||||
# MD5 计算(基于原始帧,指纹文件级去重不受裁剪影响)
|
||||
_, buffer = cv2.imencode(".jpg", frame)
|
||||
md5_hash.update(buffer)
|
||||
|
||||
phash = compute_phash(frame)
|
||||
hist = compute_color_histogram(frame)
|
||||
# Issue #1702: pHash / 颜色直方图基于中心 90% 区域,绕开 random_edge_crop
|
||||
# 降重裁剪对指纹的污染(降重只服务外部平台,不污染自查重)。
|
||||
fp_frame = center_crop_frame(frame)
|
||||
phash = compute_phash(fp_frame)
|
||||
hist = compute_color_histogram(fp_frame)
|
||||
|
||||
# 计算分片时间范围(从前一个关键帧到下一个关键帧的中点)
|
||||
prev_boundary = keyframe_times[i - 1] * 1000 if i > 0 else 0
|
||||
@@ -564,12 +657,22 @@ class VideoDeduplicator:
|
||||
|
||||
@staticmethod
|
||||
def _bhattacharyya_coefficient(hist_a: list[float], hist_b: list[float]) -> float:
|
||||
"""Bhattacharyya 系数:Σ √(a[i] * b[i]),范围 [0, 1],1=完全相同。"""
|
||||
"""Bhattacharyya 系数(概率分布版,范围 [0,1],1=完全相同)。
|
||||
|
||||
Issue #1702: compute_color_histogram 输出 3 通道拼接、每通道独立 NORM_L1
|
||||
(单通道 Σ=1,三通道拼接向量 Σ=3)。旧实现直接 Σ√(a*b) 对三通道拼接向量
|
||||
算出 ~3(旧 L2 归一化更是算出 ~14.9),不是合法的概率系数。
|
||||
这里按两个直方图各自的总量归一:BC = Σ√(a*b) / √(Σa·Σb)。
|
||||
- 单通道概率分布(Σa=Σb=1):分母 1,与旧测试/教科书定义一致;
|
||||
- 三通道拼接(Σa=Σb=3):分母 3,结果在 [0,1]。
|
||||
"""
|
||||
min_len = min(len(hist_a), len(hist_b))
|
||||
a = hist_a[:min_len]
|
||||
b = hist_b[:min_len]
|
||||
# 纯标准库计算(不依赖 numpy);max(0.0, ...) 防御上游异常负值导致 sqrt domain error
|
||||
return float(sum(math.sqrt(max(0.0, ai * bi)) for ai, bi in zip(a, b, strict=False)))
|
||||
a = [max(0.0, float(x)) for x in hist_a[:min_len]]
|
||||
b = [max(0.0, float(x)) for x in hist_b[:min_len]]
|
||||
# max(0.0, ...) 防御上游异常负值导致 sqrt domain error
|
||||
coeff = sum(math.sqrt(ai * bi) for ai, bi in zip(a, b, strict=False))
|
||||
norm = math.sqrt(sum(a) * sum(b))
|
||||
return float(coeff / norm) if norm > 0 else 0.0
|
||||
|
||||
@staticmethod
|
||||
def _compute_histogram_similarity(
|
||||
@@ -611,6 +714,72 @@ class VideoDeduplicator:
|
||||
hist_similarity = VideoDeduplicator._compute_histogram_similarity(hist_a, hist_b) if hist_b else 0.5
|
||||
return PHASH_WEIGHT * phash_similarity + HISTOGRAM_WEIGHT * hist_similarity
|
||||
|
||||
@staticmethod
|
||||
def _evaluate_candidate(
|
||||
fingerprint: VideoFingerprint,
|
||||
existing_phashes: list[str],
|
||||
existing_histograms: list,
|
||||
existing_chunk_objects: list,
|
||||
*,
|
||||
query_duration_sec: float,
|
||||
) -> dict:
|
||||
"""评估新视频指纹与单个候选视频的相似度(Issue #1702 共享逻辑)。
|
||||
|
||||
指标:
|
||||
- min_distances / frame_match_rate:每个新分片到候选视频全局最近邻的汉明距离,
|
||||
分母取两视频分片数的较小值(支持局部片段复用:短视频复用长视频片段时不被长视频分母稀释)。
|
||||
- temporal_coverage:时序一致连续匹配片段总时长 / 新视频时长(局部复用主指标)。
|
||||
- fusion:pHash 中位数距离 + 颜色直方图的加权融合分。
|
||||
|
||||
Returns:
|
||||
{frame_match_rate, temporal_coverage, segments, median_distance,
|
||||
fusion, matching_frames, min_distances}
|
||||
"""
|
||||
query_phashes = fingerprint.keyframe_phashes or []
|
||||
if not query_phashes or not existing_phashes:
|
||||
return {
|
||||
"frame_match_rate": 0.0,
|
||||
"temporal_coverage": 0.0,
|
||||
"segments": [],
|
||||
"median_distance": 64,
|
||||
"fusion": 0.0,
|
||||
"matching_frames": 0,
|
||||
"min_distances": [],
|
||||
}
|
||||
|
||||
min_distances = [min(hamming_distance(ph, ep) for ep in existing_phashes) for ph in query_phashes]
|
||||
matching_frames = sum(1 for d in min_distances if d <= PHASH_THRESHOLD)
|
||||
# 分母取 min(两视频分片数):局部复用时(如 B 的 5 片复用 A 9 片中的若干片)
|
||||
# 命中帧占比不因候选视频更长而被稀释。
|
||||
frame_match_rate = matching_frames / min(len(query_phashes), len(existing_phashes))
|
||||
|
||||
segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects)
|
||||
duration_ms = query_duration_sec * 1000 if query_duration_sec else 0
|
||||
if duration_ms > 0 and segments:
|
||||
covered_ms = sum(s.query_end_ms - s.query_start_ms for s in segments)
|
||||
temporal_coverage = min(covered_ms / duration_ms, 1.0)
|
||||
elif matching_frames > 0:
|
||||
# 无连续片段(时序连贯性不足)时,按匹配帧占比估计覆盖:
|
||||
# 密集 1s 采样下每个分片≈1s 等权时间片,匹配帧数≈命中秒数。
|
||||
temporal_coverage = min(frame_match_rate, 1.0)
|
||||
else:
|
||||
temporal_coverage = 0.0
|
||||
|
||||
median_distance = statistics.median(min_distances) if min_distances else 64
|
||||
fusion = VideoDeduplicator._compute_fusion_score(
|
||||
median_distance, fingerprint.color_histograms, existing_histograms
|
||||
)
|
||||
|
||||
return {
|
||||
"frame_match_rate": frame_match_rate,
|
||||
"temporal_coverage": temporal_coverage,
|
||||
"segments": segments,
|
||||
"median_distance": median_distance,
|
||||
"fusion": fusion,
|
||||
"matching_frames": matching_frames,
|
||||
"min_distances": min_distances,
|
||||
}
|
||||
|
||||
def check_duplicate(
|
||||
self,
|
||||
fingerprint: VideoFingerprint,
|
||||
@@ -620,6 +789,7 @@ class VideoDeduplicator:
|
||||
scope: str = "project",
|
||||
user_id: str = "",
|
||||
duration_sec: float = 0,
|
||||
exclude_video_id: str | None = None,
|
||||
) -> Optional[dict]:
|
||||
"""检查视频是否与已有视频重复。
|
||||
|
||||
@@ -636,6 +806,9 @@ class VideoDeduplicator:
|
||||
scope: "project" 项目内查重(默认),"user" 跨项目全局查重
|
||||
user_id: 用户 ID(scope="user" 时使用)
|
||||
duration_sec: 视频时长(秒),用于时长预过滤 ±15%
|
||||
exclude_video_id: 排除的视频 ID(查重自身时用)。recompute-dedup
|
||||
重算时视频记录已存在,不排除会自匹配(距离 0 分最高)导致
|
||||
duplicate_of 指向自己(Issue #1702 连带修复)。
|
||||
|
||||
Returns:
|
||||
重复信息字典(含 duplicate, duplicate_of, reason, similarity, duplicate_segments),
|
||||
@@ -643,13 +816,22 @@ class VideoDeduplicator:
|
||||
"""
|
||||
video_repo = SQLAlchemyGeneratedVideoRepository(session)
|
||||
if scope == "user" and user_id:
|
||||
dur_min = duration_sec * 0.85 if duration_sec > 0 else 0
|
||||
dur_max = duration_sec * 1.15 if duration_sec > 0 else 0
|
||||
existing_videos = video_repo.list_by_user(user_id, duration_min=dur_min, duration_max=dur_max)
|
||||
# Issue #1702: 不做 ±15% 时长预过滤。旧逻辑按 duration_sec 缩小候选窗口,
|
||||
# 但局部片段复用的两个视频时长必然不同(证据视频 20s vs 11s,差 42%),
|
||||
# ±15% 窗口让同源视频互相不可见 → is_duplicate 恒 False。
|
||||
# 全量遍历同用户视频(与 compute_duplicate_rate 口径一致),异源视频由
|
||||
# fusion/temporal_coverage 阈值天然过滤(校准:异源最小汉明距离 24)。
|
||||
existing_videos = video_repo.list_by_user(user_id)
|
||||
else:
|
||||
existing_videos = video_repo.list_by_project(project_id)
|
||||
|
||||
best_score = 0.0
|
||||
best_result: Optional[dict] = None
|
||||
|
||||
for existing in existing_videos:
|
||||
# 排除自身(recompute 时当前视频已在候选列表里,否则自匹配距离 0 必最高分)
|
||||
if exclude_video_id and existing.id == exclude_video_id:
|
||||
continue
|
||||
if not existing.video_fingerprint:
|
||||
continue
|
||||
|
||||
@@ -677,61 +859,70 @@ class VideoDeduplicator:
|
||||
if not existing_phashes:
|
||||
continue
|
||||
|
||||
# 计算每个新关键帧到已有关键帧的最小汉明距离
|
||||
min_distances = []
|
||||
for phash in fingerprint.keyframe_phashes:
|
||||
distances = [hamming_distance(phash, ep) for ep in existing_phashes]
|
||||
min_distances.append(min(distances))
|
||||
|
||||
# 帧匹配比例检查
|
||||
matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD)
|
||||
match_ratio = matching_frames / len(min_distances) if min_distances else 0
|
||||
if match_ratio < MATCH_RATIO_THRESHOLD:
|
||||
continue
|
||||
|
||||
# 中位数距离
|
||||
median_distance = statistics.median(min_distances) if min_distances else 64
|
||||
if median_distance >= self.PHASH_THRESHOLD:
|
||||
continue
|
||||
|
||||
# 直方图融合(chunk 表优先,回退 JSON 字段;JSON NULL 显式回退空列表)
|
||||
# 直方图 / 分片对象(chunk 表优先,回退 JSON 字段;JSON NULL 显式回退空列表)
|
||||
if chunk_data:
|
||||
existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")]
|
||||
existing_chunk_objects = chunk_data
|
||||
else:
|
||||
existing_histograms = ef.get("color_histograms") or []
|
||||
existing_chunk_objects = [
|
||||
{"phash_binary": pp, "start_time_ms": 0, "end_time_ms": 0} for pp in existing_phashes
|
||||
]
|
||||
|
||||
combined_score = self._compute_fusion_score(
|
||||
median_distance, fingerprint.color_histograms, existing_histograms
|
||||
# Issue #1702: 统一评估每个候选(含局部片段复用),不再用
|
||||
# "frame_match_rate<0.7 整条跳过" 的硬门槛——局部复用(如 B 结尾 2s
|
||||
# ≈ A 中间 2s)帧比例天然低,但 coverage 能检出。
|
||||
ev = self._evaluate_candidate(
|
||||
fingerprint,
|
||||
existing_phashes,
|
||||
existing_histograms,
|
||||
existing_chunk_objects,
|
||||
query_duration_sec=fingerprint.duration,
|
||||
)
|
||||
logger.debug(
|
||||
"check_duplicate candidate=%s min_distances=%s frame_match_rate=%.3f "
|
||||
"temporal_coverage=%.3f median=%.1f fusion=%.3f segments=%d",
|
||||
existing.id,
|
||||
ev["min_distances"],
|
||||
ev["frame_match_rate"],
|
||||
ev["temporal_coverage"],
|
||||
ev["median_distance"],
|
||||
ev["fusion"],
|
||||
len(ev["segments"]),
|
||||
)
|
||||
|
||||
if combined_score < DUPLICATE_THRESHOLD:
|
||||
continue
|
||||
|
||||
# 滑动窗口时序匹配:获取具体重复片段
|
||||
existing_chunk_objects = (
|
||||
chunk_data
|
||||
if chunk_data
|
||||
else [{"phash_binary": p, "start_time_ms": 0, "end_time_ms": 0} for p in existing_phashes]
|
||||
# 全片重复判定:融合分过阈 且(帧匹配比例 >=70% 或 局部覆盖 >=50%)
|
||||
is_full_duplicate = ev["fusion"] >= DUPLICATE_THRESHOLD and (
|
||||
ev["frame_match_rate"] >= MATCH_RATIO_THRESHOLD or ev["temporal_coverage"] >= PARTIAL_COVERAGE_THRESHOLD
|
||||
)
|
||||
segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects)
|
||||
|
||||
return {
|
||||
"duplicate": True,
|
||||
"duplicate_of": existing.id,
|
||||
"reason": "phash_histogram_fusion",
|
||||
"similarity": combined_score,
|
||||
"duplicate_segments": [
|
||||
{
|
||||
"query_start_ms": s.query_start_ms,
|
||||
"query_end_ms": s.query_end_ms,
|
||||
"target_start_ms": s.target_start_ms,
|
||||
"target_end_ms": s.target_end_ms,
|
||||
"avg_distance": round(s.avg_distance, 2),
|
||||
}
|
||||
for s in segments
|
||||
],
|
||||
}
|
||||
if is_full_duplicate and ev["fusion"] > best_score:
|
||||
best_score = ev["fusion"]
|
||||
best_result = {
|
||||
"duplicate": True,
|
||||
"duplicate_of": existing.id,
|
||||
"reason": "phash_histogram_fusion",
|
||||
"similarity": ev["fusion"],
|
||||
"duplicate_segments": [
|
||||
{
|
||||
"query_start_ms": s.query_start_ms,
|
||||
"query_end_ms": s.query_end_ms,
|
||||
"target_start_ms": s.target_start_ms,
|
||||
"target_end_ms": s.target_end_ms,
|
||||
"avg_distance": round(s.avg_distance, 2),
|
||||
}
|
||||
for s in ev["segments"]
|
||||
],
|
||||
}
|
||||
|
||||
if best_result:
|
||||
return best_result
|
||||
logger.info(
|
||||
"check_duplicate no match (project=%s scope=%s): %d candidates evaluated, best_fusion=%.3f",
|
||||
project_id,
|
||||
scope,
|
||||
len(existing_videos),
|
||||
best_score,
|
||||
)
|
||||
return None
|
||||
|
||||
def check_batch_duplicate(
|
||||
@@ -763,6 +954,9 @@ class VideoDeduplicator:
|
||||
video_repo = SQLAlchemyGeneratedVideoRepository(session)
|
||||
batch_videos = video_repo.list_by_batch(batch_id)
|
||||
|
||||
best_score = 0.0
|
||||
best_result: Optional[dict] = None
|
||||
|
||||
for existing in batch_videos:
|
||||
if existing.id == current_video_id:
|
||||
continue
|
||||
@@ -796,59 +990,59 @@ class VideoDeduplicator:
|
||||
if not existing_phashes:
|
||||
continue
|
||||
|
||||
min_distances = []
|
||||
for phash in fingerprint.keyframe_phashes:
|
||||
distances = [hamming_distance(phash, ep) for ep in existing_phashes]
|
||||
min_distances.append(min(distances))
|
||||
|
||||
# 帧匹配比例检查
|
||||
matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD)
|
||||
match_ratio = matching_frames / len(min_distances) if min_distances else 0
|
||||
if match_ratio < MATCH_RATIO_THRESHOLD:
|
||||
continue
|
||||
|
||||
median_distance = statistics.median(min_distances) if min_distances else 64
|
||||
if median_distance >= self.PHASH_THRESHOLD:
|
||||
continue
|
||||
|
||||
# 直方图融合(chunk 表优先,回退 JSON 字段;JSON NULL 显式回退空列表)
|
||||
if chunk_data:
|
||||
existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")]
|
||||
existing_chunk_objects = chunk_data
|
||||
else:
|
||||
existing_histograms = ef.get("color_histograms") or []
|
||||
existing_chunk_objects = [
|
||||
{"phash_binary": pp, "start_time_ms": 0, "end_time_ms": 0} for pp in existing_phashes
|
||||
]
|
||||
|
||||
combined_score = self._compute_fusion_score(
|
||||
median_distance, fingerprint.color_histograms, existing_histograms
|
||||
ev = self._evaluate_candidate(
|
||||
fingerprint,
|
||||
existing_phashes,
|
||||
existing_histograms,
|
||||
existing_chunk_objects,
|
||||
query_duration_sec=fingerprint.duration,
|
||||
)
|
||||
logger.debug(
|
||||
"check_batch_duplicate candidate=%s min_distances=%s frame_match_rate=%.3f "
|
||||
"temporal_coverage=%.3f median=%.1f fusion=%.3f segments=%d",
|
||||
existing.id,
|
||||
ev["min_distances"],
|
||||
ev["frame_match_rate"],
|
||||
ev["temporal_coverage"],
|
||||
ev["median_distance"],
|
||||
ev["fusion"],
|
||||
len(ev["segments"]),
|
||||
)
|
||||
|
||||
if combined_score < DUPLICATE_THRESHOLD:
|
||||
continue
|
||||
|
||||
# 滑动窗口时序匹配
|
||||
existing_chunk_objects = (
|
||||
chunk_data
|
||||
if chunk_data
|
||||
else [{"phash_binary": p, "start_time_ms": 0, "end_time_ms": 0} for p in existing_phashes]
|
||||
is_full_duplicate = ev["fusion"] >= DUPLICATE_THRESHOLD and (
|
||||
ev["frame_match_rate"] >= MATCH_RATIO_THRESHOLD or ev["temporal_coverage"] >= PARTIAL_COVERAGE_THRESHOLD
|
||||
)
|
||||
segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects)
|
||||
|
||||
return {
|
||||
"duplicate": True,
|
||||
"duplicate_of": existing.id,
|
||||
"reason": "batch_phash_histogram_fusion",
|
||||
"similarity": combined_score,
|
||||
"duplicate_segments": [
|
||||
{
|
||||
"query_start_ms": s.query_start_ms,
|
||||
"query_end_ms": s.query_end_ms,
|
||||
"target_start_ms": s.target_start_ms,
|
||||
"target_end_ms": s.target_end_ms,
|
||||
"avg_distance": round(s.avg_distance, 2),
|
||||
}
|
||||
for s in segments
|
||||
],
|
||||
}
|
||||
if is_full_duplicate and ev["fusion"] > best_score:
|
||||
best_score = ev["fusion"]
|
||||
best_result = {
|
||||
"duplicate": True,
|
||||
"duplicate_of": existing.id,
|
||||
"reason": "batch_phash_histogram_fusion",
|
||||
"similarity": ev["fusion"],
|
||||
"duplicate_segments": [
|
||||
{
|
||||
"query_start_ms": s.query_start_ms,
|
||||
"query_end_ms": s.query_end_ms,
|
||||
"target_start_ms": s.target_start_ms,
|
||||
"target_end_ms": s.target_end_ms,
|
||||
"avg_distance": round(s.avg_distance, 2),
|
||||
}
|
||||
for s in ev["segments"]
|
||||
],
|
||||
}
|
||||
|
||||
if best_result:
|
||||
return best_result
|
||||
logger.info("check_batch_duplicate no match (batch=%s): best_fusion=%.3f", batch_id, best_score)
|
||||
return None
|
||||
|
||||
def compute_duplicate_rate(
|
||||
@@ -897,8 +1091,7 @@ class VideoDeduplicator:
|
||||
max_duplicate_rate = 0.0
|
||||
max_visual_similarity = 0.0
|
||||
match_count = 0
|
||||
|
||||
total_duration_ms = fingerprint.duration if fingerprint.duration else 0
|
||||
evaluated = 0
|
||||
|
||||
for existing in existing_videos:
|
||||
if current_video_id and existing.id == current_video_id:
|
||||
@@ -933,57 +1126,63 @@ class VideoDeduplicator:
|
||||
if not existing_phashes or not fingerprint.keyframe_phashes:
|
||||
continue
|
||||
|
||||
min_distances = []
|
||||
for phash in fingerprint.keyframe_phashes:
|
||||
distances = [hamming_distance(phash, ep) for ep in existing_phashes]
|
||||
min_distances.append(min(distances))
|
||||
|
||||
# frame_match_rate
|
||||
total_frames = len(min_distances)
|
||||
if total_frames == 0:
|
||||
continue
|
||||
matching_frames = sum(1 for d in min_distances if d < self.PHASH_THRESHOLD)
|
||||
frame_match_rate = matching_frames / total_frames
|
||||
|
||||
# 帧匹配比例太低则跳过
|
||||
if frame_match_rate < 0.3:
|
||||
continue
|
||||
|
||||
# temporal_coverage_rate via find_duplicate_segments
|
||||
existing_chunk_objects = (
|
||||
chunk_data
|
||||
if chunk_data
|
||||
else [{"phash_binary": p, "start_time_ms": 0, "end_time_ms": 0} for p in existing_phashes]
|
||||
)
|
||||
segments = find_duplicate_segments(fingerprint.chunks, existing_chunk_objects)
|
||||
|
||||
if total_duration_ms > 0 and segments:
|
||||
covered_ms = sum(s.query_end_ms - s.query_start_ms for s in segments)
|
||||
temporal_coverage_rate = min(covered_ms / total_duration_ms, 1.0)
|
||||
else:
|
||||
temporal_coverage_rate = 0.0
|
||||
|
||||
# duplicate_rate = 0.4 * frame_match_rate + 0.6 * temporal_coverage_rate
|
||||
dup_rate = (frame_match_rate * 0.4 + temporal_coverage_rate * 0.6) * 100
|
||||
|
||||
# visual_similarity (融合相似度,归一化 0~1)
|
||||
median_distance = statistics.median(min_distances) if min_distances else 64
|
||||
# 直方图 / 分片对象(chunk 表优先,回退 JSON 字段;JSON NULL 显式回退空列表)
|
||||
if chunk_data:
|
||||
existing_histograms = [c["color_histogram"] for c in chunk_data if c.get("color_histogram")]
|
||||
existing_chunk_objects = chunk_data
|
||||
else:
|
||||
# JSON NULL 显式回退空列表
|
||||
existing_histograms = ef.get("color_histograms") or []
|
||||
existing_chunk_objects = [
|
||||
{"phash_binary": pp, "start_time_ms": 0, "end_time_ms": 0} for pp in existing_phashes
|
||||
]
|
||||
|
||||
visual_sim = self._compute_fusion_score(median_distance, fingerprint.color_histograms, existing_histograms)
|
||||
# Issue #1702: 统一评估;frame_match_rate 分母为 min(两视频分片数),
|
||||
# temporal_coverage 时长量纲在 _evaluate_candidate 内统一为毫秒。
|
||||
ev = self._evaluate_candidate(
|
||||
fingerprint,
|
||||
existing_phashes,
|
||||
existing_histograms,
|
||||
existing_chunk_objects,
|
||||
query_duration_sec=fingerprint.duration,
|
||||
)
|
||||
evaluated += 1
|
||||
logger.debug(
|
||||
"compute_duplicate_rate candidate=%s min_distances=%s frame_match_rate=%.3f "
|
||||
"temporal_coverage=%.3f median=%.1f fusion=%.3f segments=%d",
|
||||
existing.id,
|
||||
ev["min_distances"],
|
||||
ev["frame_match_rate"],
|
||||
ev["temporal_coverage"],
|
||||
ev["median_distance"],
|
||||
ev["fusion"],
|
||||
len(ev["segments"]),
|
||||
)
|
||||
|
||||
# 判定是否为重复(融合分数超过阈值)
|
||||
if visual_sim >= DUPLICATE_THRESHOLD:
|
||||
# Issue #1702: 去掉 "frame_match_rate<0.3 整条跳过" 硬门槛——
|
||||
# 局部片段复用帧比例天然低;coverage 为主指标,0 匹配自然得 0 分。
|
||||
# duplicate_rate = 0.4 * frame_match_rate + 0.6 * temporal_coverage
|
||||
dup_rate = (min(ev["frame_match_rate"], 1.0) * 0.4 + ev["temporal_coverage"] * 0.6) * 100
|
||||
|
||||
# 全片重复计数与 check_duplicate 判定口径一致
|
||||
if ev["fusion"] >= DUPLICATE_THRESHOLD and (
|
||||
ev["frame_match_rate"] >= MATCH_RATIO_THRESHOLD or ev["temporal_coverage"] >= PARTIAL_COVERAGE_THRESHOLD
|
||||
):
|
||||
match_count += 1
|
||||
|
||||
if dup_rate > max_duplicate_rate:
|
||||
max_duplicate_rate = dup_rate
|
||||
max_visual_similarity = visual_sim
|
||||
max_visual_similarity = ev["fusion"]
|
||||
|
||||
logger.info(
|
||||
"compute_duplicate_rate done (project=%s scope=%s): evaluated=%d max_rate=%.2f%% "
|
||||
"max_visual_sim=%.3f matches=%d",
|
||||
project_id,
|
||||
scope,
|
||||
evaluated,
|
||||
max_duplicate_rate,
|
||||
max_visual_similarity,
|
||||
match_count,
|
||||
)
|
||||
return {
|
||||
"duplicate_rate": round(max(max_duplicate_rate, 0.0), 2),
|
||||
"visual_similarity": round(max_visual_similarity, 4),
|
||||
@@ -999,18 +1198,20 @@ def _save_fingerprint_chunks(
|
||||
session: Session,
|
||||
) -> None:
|
||||
"""将指纹分片数据批量写入 video_fingerprint_chunks 表。幂等:已有数据时跳过。"""
|
||||
# 幂等检查:已有分片数据则跳过
|
||||
existing_count = (
|
||||
session.query(VideoFingerprintChunkModel).filter(VideoFingerprintChunkModel.video_id == video_id).count()
|
||||
)
|
||||
if existing_count > 0:
|
||||
logger.debug("Fingerprint chunks already exist for video %s (%d chunks), skipping", video_id, existing_count)
|
||||
return
|
||||
|
||||
if not fingerprint.chunks:
|
||||
logger.warning("No chunks in fingerprint for video %s, skipping chunk save", video_id)
|
||||
return
|
||||
|
||||
# Issue #1702: recompute-dedup 重算时指纹算法已变(中心裁剪 + 新阈值),
|
||||
# 旧分片必须替换而非跳过(旧实现"有数据就跳过"导致重算不刷新分片表)。
|
||||
deleted = (
|
||||
session.query(VideoFingerprintChunkModel)
|
||||
.filter(VideoFingerprintChunkModel.video_id == video_id)
|
||||
.delete(synchronize_session=False)
|
||||
)
|
||||
if deleted:
|
||||
logger.info("Replaced %d stale fingerprint chunks for video %s", deleted, video_id)
|
||||
|
||||
chunk_models = fingerprint.to_chunk_models(video_id, project_id, user_id)
|
||||
session.bulk_save_objects(chunk_models)
|
||||
logger.info("Saved %d fingerprint chunks for video %s", len(chunk_models), video_id)
|
||||
@@ -1032,9 +1233,17 @@ def check_duplicate_task(self: Task, generated_video_id: str) -> dict:
|
||||
raise ValueError(f"Generated video {generated_video_id} not found")
|
||||
|
||||
local_path = os.path.join(temp_dir, f"{generated_video_id}.mp4")
|
||||
storage_service.download_file(
|
||||
f"projects/{video.project_id}/generated/{generated_video_id}/{generated_video_id}.mp4", local_path
|
||||
)
|
||||
# Issue #1702: recompute 走的是 OSS 重新下载路径(正常生成流程用本地渲染文件,
|
||||
# 不经此任务)。成片真实 OSS key 是生成时的
|
||||
# generated/projects/{pid}/tasks/{task_id}/rendered_*.mp4(见 generation.py
|
||||
# _upload_and_record),旧代码硬编码 projects/{pid}/generated/{vid}/{vid}.mp4
|
||||
# 这个从不存在的 key,导致所有 recompute 任务下载 404、查重数据永远无法重算。
|
||||
# 优先从 file_url 解析真实 key,旧 key 模式仅作回退。
|
||||
download_key = getattr(video, "file_url", "") or ""
|
||||
if not download_key:
|
||||
download_key = f"projects/{video.project_id}/generated/{generated_video_id}/{generated_video_id}.mp4"
|
||||
logger.warning("video %s has no file_url, falling back to legacy key %s", generated_video_id, download_key)
|
||||
storage_service.download_file(download_key, local_path)
|
||||
|
||||
fingerprint = deduplicator.compute_fingerprint(local_path)
|
||||
|
||||
@@ -1045,7 +1254,10 @@ def check_duplicate_task(self: Task, generated_video_id: str) -> dict:
|
||||
session,
|
||||
scope="user",
|
||||
user_id=video.user_id,
|
||||
duration_sec=fingerprint.duration / 1000 if fingerprint.duration else 0,
|
||||
# Issue #1702: fingerprint.duration 单位已经是秒,旧代码 /1000 导致
|
||||
# ±15% 时长预过滤窗口缩到 ~0.013s,scope=user 的跨项目查重永远返回 None。
|
||||
duration_sec=fingerprint.duration if fingerprint.duration else 0,
|
||||
exclude_video_id=generated_video_id,
|
||||
)
|
||||
|
||||
video.video_fingerprint = fingerprint.to_dict()
|
||||
|
||||
@@ -92,7 +92,8 @@ def create_video_record_and_dedup(
|
||||
logger.warning("Failed to save fingerprint chunks for %s: %s", video_id, chunk_err)
|
||||
|
||||
# (a) 历史成片查重(跨项目全局 + 时长预过滤)
|
||||
duration_sec = fingerprint.duration / 1000 if fingerprint.duration else 0
|
||||
# Issue #1702: fingerprint.duration 单位是秒,旧代码 /1000 让时长预过滤失效
|
||||
duration_sec = fingerprint.duration if fingerprint.duration else 0
|
||||
duplicate_result = deduplicator.check_duplicate(
|
||||
fingerprint,
|
||||
project_id,
|
||||
@@ -100,6 +101,7 @@ def create_video_record_and_dedup(
|
||||
scope="user",
|
||||
user_id=user_id,
|
||||
duration_sec=duration_sec,
|
||||
exclude_video_id=video_id,
|
||||
)
|
||||
|
||||
# (b) 批次内查重(仅当有 batch_id 时)
|
||||
|
||||
@@ -6,6 +6,18 @@ 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",
|
||||
@@ -22,10 +34,18 @@ celery_app.conf.imports = (
|
||||
)
|
||||
|
||||
# Celery Beat 定时任务调度
|
||||
# 注:worker 单实例内嵌 beat(entrypoint-worker.sh -B),定时任务不会重复执行
|
||||
celery_app.conf.beat_schedule = {
|
||||
# pending 任务超时清理:worker 停止消费后,卡 pending 的任务 15 分钟内释放限流名额
|
||||
"cleanup-stale-pending-tasks": {
|
||||
"task": "worker.cleanup_stale_pending_tasks",
|
||||
"schedule": 600.0, # 每 10 分钟(秒)
|
||||
"options": {"expires": 300}, # 5 分钟过期,避免堆积
|
||||
"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},
|
||||
},
|
||||
}
|
||||
|
||||
@@ -7,11 +7,83 @@ from worker_app.db import SessionLocal
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 孤儿任务超时阈值:渲染任务超过此时间未更新则视为卡死
|
||||
ORPHAN_TASK_TIMEOUT_MINUTES = 10
|
||||
|
||||
# Pending 任务超时阈值:pending 任务在队列中等待超过此时间则自动清理
|
||||
PENDING_TASK_TIMEOUT_MINUTES = 30
|
||||
def cleanup_stale_running_with_session(repo, timeout_minutes: int) -> int:
|
||||
"""清理超时未更新的 running GenerationTask(可注入 repo 的纯核心,便于单测)。
|
||||
|
||||
Returns:
|
||||
清理的任务数量
|
||||
"""
|
||||
return len(cleanup_stale_running_with_session_ids(repo, timeout_minutes))
|
||||
|
||||
|
||||
def cleanup_stale_running_with_session_ids(repo, timeout_minutes: int) -> list[tuple[str, str]]:
|
||||
"""同 cleanup_stale_running_with_session,返回 [(task_id, celery_task_id), ...]。"""
|
||||
fn = getattr(repo, "cleanup_stale_running_with_ids", None)
|
||||
if fn is not None:
|
||||
return fn(timeout_minutes)
|
||||
# 旧仓储无 _with_ids 方法:降级为计数,无法撤销消息(执行前状态守卫兜底)
|
||||
count = repo.cleanup_stale_running(timeout_minutes)
|
||||
return [("", "") for _ in range(count)]
|
||||
|
||||
|
||||
def cleanup_stale_pending_with_session(repo, timeout_minutes: int) -> int:
|
||||
"""清理超时 pending GenerationTask(可注入 repo 的纯核心,便于单测)。
|
||||
|
||||
Returns:
|
||||
清理的任务数量
|
||||
"""
|
||||
return len(cleanup_stale_pending_with_session_ids(repo, timeout_minutes))
|
||||
|
||||
|
||||
def cleanup_stale_pending_with_session_ids(repo, timeout_minutes: int) -> list[tuple[str, str]]:
|
||||
"""同 cleanup_stale_pending_with_session,返回 [(task_id, celery_task_id), ...]。"""
|
||||
fn = getattr(repo, "cleanup_stale_pending_with_ids", None)
|
||||
if fn is not None:
|
||||
return fn(timeout_minutes)
|
||||
count = repo.cleanup_stale_pending(timeout_minutes)
|
||||
return [("", "") for _ in range(count)]
|
||||
|
||||
|
||||
def _revoke_and_purge_stale_messages(items: list[tuple[str, str]]) -> int:
|
||||
"""把清理掉的任务对应的 Celery 消息撤销并从 Redis 队列清除(#1714)。
|
||||
|
||||
防止「DB 已标 failed,但队列消息还在 → 重投执行 → 非法状态转换 → 半成品」。
|
||||
失败不阻断清理流程(执行前状态守卫是第二道防线)。
|
||||
"""
|
||||
biz_ids = [tid for tid, _ in items if tid]
|
||||
celery_ids = [cid for _, cid in items if cid]
|
||||
if not biz_ids and not celery_ids:
|
||||
return 0
|
||||
try:
|
||||
from worker_app.celery_app import celery_app as app
|
||||
from worker_app.core.config import get_settings
|
||||
|
||||
from packages.shared.celery_orphan_guard import revoke_and_purge
|
||||
|
||||
broker_url = get_settings().broker_url
|
||||
return revoke_and_purge(
|
||||
app,
|
||||
broker_url,
|
||||
business_task_ids=biz_ids,
|
||||
celery_task_ids=celery_ids,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.error("撤销作废任务队列消息失败(执行前守卫仍会兜底): %s", e, exc_info=True)
|
||||
return 0
|
||||
|
||||
|
||||
# 孤儿任务超时阈值:running 任务超过此时间无进度更新则视为卡死。
|
||||
# 依据:worker.generate_video 硬超时 time_limit=11 分钟,正常任务不可能超过;
|
||||
# 20 分钟阈值覆盖硬超时 + 重试 + 余量,绝不误杀正常任务。
|
||||
ORPHAN_TASK_TIMEOUT_MINUTES = 20
|
||||
|
||||
# Pending 任务超时阈值:任务创建后超过此时间仍未开始执行则判死。
|
||||
# 注意区分 running 孤儿阈值(20 分钟):pending 是「排队等待」时间,
|
||||
# 队列积压(如 20+ 转码任务)时视频生成可能正常排队较久,阈值必须放宽,
|
||||
# 避免正常排队任务被误杀。队列隔离(#1714)后 generation 队列独占 worker,
|
||||
# 理论上排队极短;保留 45 分钟作为兜底,覆盖 worker 短暂停止消费的场景。
|
||||
PENDING_TASK_TIMEOUT_MINUTES = 45
|
||||
|
||||
|
||||
def cleanup_orphan_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> int: # pragma: no cover
|
||||
@@ -32,11 +104,16 @@ def cleanup_orphan_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) ->
|
||||
|
||||
try:
|
||||
session = SessionLocal()
|
||||
repo = SQLAlchemyGenerationTaskRepository(session)
|
||||
count = repo.cleanup_stale_running(timeout_minutes)
|
||||
session.close()
|
||||
try:
|
||||
repo = SQLAlchemyGenerationTaskRepository(session)
|
||||
items = cleanup_stale_running_with_session_ids(repo, timeout_minutes)
|
||||
finally:
|
||||
session.close()
|
||||
count = len(items)
|
||||
if count > 0:
|
||||
logger.warning("清理了 %d 个超时的孤儿 GenerationTask(超过 %d 分钟未更新)", count, timeout_minutes)
|
||||
purged = _revoke_and_purge_stale_messages(items)
|
||||
logger.info("孤儿任务对应队列消息撤销/清除完成: %d 条", purged)
|
||||
else:
|
||||
logger.info("无孤儿 GenerationTask 需要清理")
|
||||
return count
|
||||
@@ -70,7 +147,9 @@ def cleanup_stale_jobs(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> in
|
||||
.all()
|
||||
)
|
||||
count = 0
|
||||
stale_items: list[tuple[str, str]] = []
|
||||
for model in stale_jobs:
|
||||
stale_items.append((model.id, getattr(model, "celery_task_id", "") or ""))
|
||||
model.status = JobStatus.FAILED.value
|
||||
model.error_message = f"任务执行中断(超过 {timeout_minutes} 分钟未更新)"
|
||||
count += 1
|
||||
@@ -80,6 +159,8 @@ def cleanup_stale_jobs(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> in
|
||||
else:
|
||||
logger.info("无孤儿 Job 需要清理")
|
||||
session.close()
|
||||
if count > 0:
|
||||
_revoke_and_purge_generation(stale_items)
|
||||
return count
|
||||
except Exception as e:
|
||||
logger.error("清理孤儿 Job 失败: %s", e, exc_info=True)
|
||||
@@ -105,9 +186,12 @@ def cleanup_stale_pending_tasks(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINU
|
||||
session = SessionLocal()
|
||||
try:
|
||||
repo = SQLAlchemyGenerationTaskRepository(session)
|
||||
count = repo.cleanup_stale_pending(timeout_minutes)
|
||||
items = cleanup_stale_pending_with_session_ids(repo, timeout_minutes)
|
||||
count = len(items)
|
||||
if count > 0:
|
||||
logger.warning("清理了 %d 个超时的 pending GenerationTask(超过 %d 分钟未处理)", count, timeout_minutes)
|
||||
purged = _revoke_and_purge_stale_messages(items)
|
||||
logger.info("超时 pending 任务对应队列消息撤销/清除完成: %d 条", purged)
|
||||
else:
|
||||
logger.info("无超时 pending GenerationTask 需要清理")
|
||||
return count
|
||||
@@ -118,6 +202,31 @@ def cleanup_stale_pending_tasks(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINU
|
||||
session.close()
|
||||
|
||||
|
||||
def _revoke_and_purge_generation(items: list[tuple[str, str]]) -> int:
|
||||
"""撤销 Job 表孤儿任务(TTS/配音等)的队列消息,队列覆盖全部已知队列。"""
|
||||
biz_ids = [tid for tid, _ in items if tid]
|
||||
celery_ids = [cid for _, cid in items if cid]
|
||||
if not biz_ids and not celery_ids:
|
||||
return 0
|
||||
try:
|
||||
from worker_app.celery_app import celery_app as app
|
||||
from worker_app.core.config import get_settings
|
||||
|
||||
from packages.shared.celery_orphan_guard import revoke_and_purge
|
||||
|
||||
broker_url = get_settings().broker_url
|
||||
return revoke_and_purge(
|
||||
app,
|
||||
broker_url,
|
||||
business_task_ids=biz_ids,
|
||||
celery_task_ids=celery_ids,
|
||||
queue_names=("generation", "transcode", "celery"),
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.error("撤销 Job 队列消息失败: %s", e, exc_info=True)
|
||||
return 0
|
||||
|
||||
|
||||
def cleanup_all_stale_tasks(timeout_minutes: int = ORPHAN_TASK_TIMEOUT_MINUTES) -> dict: # pragma: no cover
|
||||
"""统一清理所有超时的孤儿任务。
|
||||
|
||||
|
||||
@@ -1,14 +1,18 @@
|
||||
"""定期清理任务 — Celery Beat 调度。
|
||||
|
||||
包含:
|
||||
- cleanup_stale_pending_tasks: 定期清理卡在 pending 超时的 generation_tasks
|
||||
- cleanup_stale_pending_tasks: 定期清理卡在 pending 超时的 generation_tasks(worker 停止消费时占位)
|
||||
- cleanup_stale_running_tasks: 定期清理卡在 running 超时的 generation_tasks(容器重启/进程被杀后的孤儿)
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
from celery import shared_task
|
||||
from worker_app.tasks._startup import (
|
||||
ORPHAN_TASK_TIMEOUT_MINUTES,
|
||||
PENDING_TASK_TIMEOUT_MINUTES,
|
||||
cleanup_orphan_tasks,
|
||||
cleanup_stale_jobs,
|
||||
cleanup_stale_pending_tasks,
|
||||
)
|
||||
|
||||
@@ -19,12 +23,14 @@ logger = logging.getLogger(__name__)
|
||||
def scheduled_cleanup_stale_pending(timeout_minutes: int = PENDING_TASK_TIMEOUT_MINUTES) -> dict:
|
||||
"""Celery Beat 调度的定期任务:清理超时的 pending 任务。
|
||||
|
||||
每 10 分钟执行一次(由 celery_app.py 的 beat_schedule 配置),
|
||||
每 5 分钟执行一次(由 celery_app.py 的 beat_schedule 配置),
|
||||
查找所有 status='pending' 且 created_at < NOW() - timeout_minutes
|
||||
的 generation_tasks,批量更新为 failed。
|
||||
的 generation_tasks,批量更新为 failed,释放限流名额;同时 revoke 并清除
|
||||
Redis 队列中对应的 Celery 消息,杜绝作废消息重投执行(#1714)。
|
||||
|
||||
Args:
|
||||
timeout_minutes: 超时时间(分钟),默认 30 分钟
|
||||
timeout_minutes: 超时时间(分钟),默认 45 分钟(pending 排队阈值放宽,
|
||||
与 running 孤儿 20 分钟区分,避免正常排队任务被误杀)
|
||||
|
||||
Returns:
|
||||
{"cleaned": int}
|
||||
@@ -33,3 +39,33 @@ 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}
|
||||
|
||||
@@ -22,6 +22,8 @@ from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
from worker_app.tasks.generation_plan_builder import build_error_info as _build_error_info
|
||||
|
||||
from packages.shared.celery_orphan_guard import TERMINAL_STATUS_VALUES
|
||||
|
||||
OUTPUT_WIDTH = 1280
|
||||
OUTPUT_HEIGHT = 720
|
||||
OUTPUT_FPS = 25.0
|
||||
@@ -658,6 +660,34 @@ def generate_video(self, task_id: str) -> dict:
|
||||
finally:
|
||||
_session.close()
|
||||
|
||||
# ── 0. 执行前状态守卫(#1714):任务已被超时清理/孤儿恢复标记为终态时,
|
||||
# 这是作废消息(worker 崩溃前未 ack 的旧消息重投/重复投递),直接丢弃,
|
||||
# 不进入渲染,杜绝 failed→running 非法转换后继续跑产出半成品。
|
||||
if gen_task is not None and gen_task.status.value in TERMINAL_STATUS_VALUES:
|
||||
logger.warning(
|
||||
"[task_id=%s] 任务状态已为 %s,丢弃作废消息,不执行渲染",
|
||||
task_id,
|
||||
gen_task.status.value,
|
||||
)
|
||||
return {
|
||||
"status": "discarded",
|
||||
"task_id": task_id,
|
||||
"reason": f"task already terminal: {gen_task.status.value}",
|
||||
}
|
||||
|
||||
# 标记任务为 running —— 必须成功:状态机非法转换(如 failed→running)说明
|
||||
# 任务已被作废,安全中止,禁止继续执行。
|
||||
if not _update_task_status(task_id, "mark_processing"):
|
||||
logger.error(
|
||||
"[task_id=%s] 标记 running 失败(任务可能已被作废/取消),安全中止,不执行渲染",
|
||||
task_id,
|
||||
)
|
||||
return {
|
||||
"status": "discarded",
|
||||
"task_id": task_id,
|
||||
"reason": "claim failed (invalid state transition)",
|
||||
}
|
||||
|
||||
# 记录接收任务日志
|
||||
if gen_task:
|
||||
gen_task.append_log(
|
||||
@@ -669,8 +699,6 @@ def generate_video(self, task_id: str) -> dict:
|
||||
)
|
||||
_flush_logs(task_id, gen_task)
|
||||
|
||||
# 标记任务为 running
|
||||
_update_task_status(task_id, "mark_processing")
|
||||
_update_task_progress(task_id, 10, "任务启动")
|
||||
|
||||
try:
|
||||
|
||||
@@ -359,6 +359,54 @@ def validate_transcode_output(
|
||||
return True
|
||||
|
||||
|
||||
def _original_key_from_storage_key(storage_key: str) -> str:
|
||||
"""从可能被 HEVC 转码改写的 storage_key 还原原始 key。
|
||||
|
||||
转码成功后 key 形如 uploads/<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 占位 asset(Issue #1714)。
|
||||
|
||||
HEVC 转码成功后 job.storage_key 会被改写为 *_h264 新 key,旧实现用新 key
|
||||
回查占位必然落空,进而兜底新建一条 READY 记录,导致原占位永久卡 processing。
|
||||
|
||||
查找优先级:
|
||||
1. job.asset_id(complete 派单时透传的占位 id,最可靠,不依赖 key);
|
||||
2. 原始 storage_key(占位记录以原始 key 创建);
|
||||
3. 当前 job.storage_key(未转码/降级场景与原始 key 相同)。
|
||||
|
||||
找不到返回 None(旧链路兼容,由调用方兜底新建并告警)。
|
||||
"""
|
||||
asset_id = getattr(job, "asset_id", "") or ""
|
||||
if asset_id:
|
||||
try:
|
||||
found = asset_repo.find_by_id(asset_id)
|
||||
if found is not None:
|
||||
return found
|
||||
except Exception as find_err:
|
||||
logger.warning("占位 asset 按 id 查询失败 asset_id=%s: %s", asset_id, find_err)
|
||||
for key in (original_storage_key, getattr(job, "storage_key", "")):
|
||||
if not key:
|
||||
continue
|
||||
try:
|
||||
found = asset_repo.find_by_storage_key(key)
|
||||
except Exception:
|
||||
logger.warning("find_by_storage_key not available, trying fallback lookup")
|
||||
found = None
|
||||
if found is not None:
|
||||
return found
|
||||
return None
|
||||
|
||||
|
||||
@celery_app.task(name="worker.ingest_asset")
|
||||
def ingest_asset(job_id: str) -> dict:
|
||||
"""
|
||||
@@ -380,6 +428,23 @@ def ingest_asset(job_id: str) -> dict:
|
||||
if job is None:
|
||||
return {"status": "failed", "error": "job not found"}
|
||||
|
||||
# ── 执行前状态守卫(#1714):job 已终态(失败/完成)说明这是作废消息
|
||||
# (超时清理标记 failed 后旧消息重投、或重复投递),直接丢弃不执行,
|
||||
# 避免重复转码、重复回写。processing 是本任务自己第一次置位前的旧消息
|
||||
# 极少见,保守起见也丢弃(processing 的占位由恢复流程处理)。
|
||||
current_status = job.status.value if hasattr(job.status, "value") else str(job.status)
|
||||
if current_status in ("failed", "completed"):
|
||||
logger.warning(
|
||||
"[ingest job_id=%s] 任务状态已为 %s,丢弃作废消息,不执行转码",
|
||||
job_id,
|
||||
current_status,
|
||||
)
|
||||
return {"status": "discarded", "job_id": job_id, "reason": f"job already terminal: {current_status}"}
|
||||
|
||||
# 记录原始 storage_key:HEVC 转码成功后 job.storage_key 会改写为 *_h264,
|
||||
# 而 complete 阶段的占位 asset 始终以原始 key 创建,关联回写必须保留它。
|
||||
original_storage_key = job.storage_key
|
||||
|
||||
# Update job status to PROCESSING
|
||||
job.status = IngestJobStatus.PROCESSING
|
||||
job.updated_at = datetime.now(timezone.utc)
|
||||
@@ -624,22 +689,41 @@ def ingest_asset(job_id: str) -> dict:
|
||||
error_reason,
|
||||
)
|
||||
|
||||
asset = Asset.create(
|
||||
project_id=job.project_id,
|
||||
library_id=job.library_id,
|
||||
name=filename,
|
||||
storage_key=job.storage_key,
|
||||
mime_type=mime_type,
|
||||
metadata={"source": "upload", "ingest_error": error_reason},
|
||||
file_size=int(metadata.get("size_bytes", 0)),
|
||||
duration=float(metadata.get("duration", 0)),
|
||||
width=int(metadata.get("width", 0)),
|
||||
height=int(metadata.get("height", 0)),
|
||||
codec=metadata.get("codec") or None,
|
||||
status=AssetStatus.ERROR,
|
||||
file_hash=job.file_hash,
|
||||
)
|
||||
asset_repo.create(asset)
|
||||
placeholder = _resolve_placeholder_asset(asset_repo, job, original_storage_key)
|
||||
if placeholder is not None:
|
||||
# 回写占位记录:标 ERROR(Issue #1714:禁止新建第二条导致占位孤儿)
|
||||
asset = placeholder
|
||||
asset.mime_type = mime_type
|
||||
asset.metadata = {"source": "upload", "ingest_error": error_reason}
|
||||
asset.file_size = int(metadata.get("size_bytes", 0))
|
||||
asset.duration = float(metadata.get("duration", 0)) or None
|
||||
asset.width = int(metadata.get("width", 0)) or None
|
||||
asset.height = int(metadata.get("height", 0)) or None
|
||||
codec_val = metadata.get("codec")
|
||||
if codec_val:
|
||||
asset.codec = str(codec_val)
|
||||
asset.status = AssetStatus.ERROR
|
||||
asset.updated_at = datetime.now(timezone.utc)
|
||||
asset_repo.update(asset)
|
||||
else:
|
||||
# 旧链路兜底:无占位记录(如历史 job 重跑)才新建
|
||||
logger.warning("无效素材且未找到占位记录,兜底新建 ERROR asset: job_id=%s", job_id)
|
||||
asset = Asset.create(
|
||||
project_id=job.project_id,
|
||||
library_id=job.library_id,
|
||||
name=filename,
|
||||
storage_key=job.storage_key,
|
||||
mime_type=mime_type,
|
||||
metadata={"source": "upload", "ingest_error": error_reason},
|
||||
file_size=int(metadata.get("size_bytes", 0)),
|
||||
duration=float(metadata.get("duration", 0)),
|
||||
width=int(metadata.get("width", 0)),
|
||||
height=int(metadata.get("height", 0)),
|
||||
codec=metadata.get("codec") or None,
|
||||
status=AssetStatus.ERROR,
|
||||
file_hash=job.file_hash,
|
||||
)
|
||||
asset_repo.create(asset)
|
||||
|
||||
# Update job status to FAILED
|
||||
job.status = IngestJobStatus.FAILED
|
||||
@@ -656,16 +740,19 @@ def ingest_asset(job_id: str) -> dict:
|
||||
"error": error_reason,
|
||||
}
|
||||
|
||||
# 查找已存在的 Asset 记录(由 API 端在上传完成时立即创建为 PROCESSING 状态)
|
||||
existing_asset = None
|
||||
try:
|
||||
existing_asset = asset_repo.find_by_storage_key(job.storage_key)
|
||||
except Exception:
|
||||
logger.warning("find_by_storage_key not available, trying fallback lookup")
|
||||
# 查找 complete 阶段创建的占位 Asset 记录(Issue #1714)。
|
||||
# 必须用原始 storage_key / job.asset_id 关联——HEVC 转码后 job.storage_key
|
||||
# 已改写为 *_h264,用新 key 回查占位必然落空,旧实现因此兜底新建 READY 记录,
|
||||
# 导致原 PROCESSING 占位永久卡住(每个 HEVC 视频产生两条记录)。
|
||||
existing_asset = _resolve_placeholder_asset(asset_repo, job, original_storage_key)
|
||||
|
||||
if existing_asset is None:
|
||||
# 兜底:如果 API 端没有预先创建 Asset(旧版本兼容),则创建新记录
|
||||
logger.info("No pre-created asset found for storage_key=%s, creating new", job.storage_key)
|
||||
# 兜底:仅当确实没有占位记录(旧版本 API / 历史 job 重跑)才新建。
|
||||
logger.warning(
|
||||
"No placeholder asset found for job_id=%s original_key=%s, creating new",
|
||||
job_id,
|
||||
original_storage_key,
|
||||
)
|
||||
metadata["source"] = "upload"
|
||||
asset = Asset.create(
|
||||
project_id=job.project_id,
|
||||
@@ -685,8 +772,13 @@ def ingest_asset(job_id: str) -> dict:
|
||||
)
|
||||
asset_repo.create(asset)
|
||||
else:
|
||||
# 更新已有的 Asset 记录,补充元数据并将状态改为 READY
|
||||
# 更新占位记录:补充元数据、置 READY。转码成功时 storage_key 同步改写为
|
||||
# *_h264(播放/下载走转码产物),原始 key 记入 metadata 可溯源。
|
||||
asset = existing_asset
|
||||
if job.storage_key != asset.storage_key:
|
||||
metadata["original_storage_key"] = asset.storage_key
|
||||
metadata["hevc_transcoded"] = True
|
||||
asset.storage_key = job.storage_key
|
||||
asset.mime_type = mime_type
|
||||
metadata["source"] = "upload"
|
||||
asset.metadata = metadata
|
||||
@@ -737,9 +829,28 @@ def ingest_asset(job_id: str) -> dict:
|
||||
job_repo.update(job)
|
||||
|
||||
# 将上传时创建的占位 Asset(PROCESSING/UPLOADING)标记为 ERROR,
|
||||
# 避免素材永远卡在中间状态
|
||||
# 避免素材永远卡在中间状态。转码可能已把 job.storage_key 改写为
|
||||
# *_h264,需用 asset_id / 原始 key 多路径关联占位(Issue #1714)。
|
||||
try:
|
||||
existing = asset_repo.find_by_storage_key(job.storage_key)
|
||||
existing = None
|
||||
_asset_id = getattr(job, "asset_id", "") or ""
|
||||
if _asset_id:
|
||||
try:
|
||||
existing = asset_repo.find_by_id(_asset_id)
|
||||
except Exception:
|
||||
existing = None
|
||||
if existing is None:
|
||||
_candidate_keys = [
|
||||
_original_key_from_storage_key(job.storage_key),
|
||||
job.storage_key,
|
||||
]
|
||||
for _key in _candidate_keys:
|
||||
try:
|
||||
existing = asset_repo.find_by_storage_key(_key)
|
||||
except Exception:
|
||||
existing = None
|
||||
if existing is not None:
|
||||
break
|
||||
if existing and existing.status in (
|
||||
AssetStatus.PROCESSING,
|
||||
AssetStatus.UPLOADING,
|
||||
|
||||
@@ -214,3 +214,10 @@ MEDIAKIT_TIMEOUT=60
|
||||
# ==================== 监控(可选)====================
|
||||
# Sentry DSN(取消注释并填入实际值以启用错误追踪)
|
||||
# SENTRY_DSN=${SENTRY_DSN}
|
||||
|
||||
|
||||
# ==================== 微信开放平台 OAuth(网页扫码登录)====================
|
||||
# 回调域名:xiaoxiajianji.com(微信开放平台已配置)
|
||||
WECHAT_OPEN_APP_ID=${WECHAT_APP_ID}
|
||||
WECHAT_OPEN_APP_SECRET=${WECHAT_APP_SECRET}
|
||||
WECHAT_OPEN_REDIRECT_URI=https://saas.xiaoxiajianji.com/auth/wechat/callback
|
||||
|
||||
@@ -231,3 +231,10 @@ DASHSCOPE_API_KEY=${DASHSCOPE_API_KEY}
|
||||
MEDIAKIT_API_KEY=${MEDIAKIT_API_KEY}
|
||||
MEDIAKIT_BASE_URL=https://mediakit.cn-beijing.volces.com/api/v1
|
||||
MEDIAKIT_TIMEOUT=60
|
||||
|
||||
|
||||
# ==================== 微信开放平台 OAuth(网页扫码登录)====================
|
||||
# 回调域名:xiaoxiajianji.com(微信开放平台已配置)
|
||||
WECHAT_OPEN_APP_ID=${WECHAT_APP_ID}
|
||||
WECHAT_OPEN_APP_SECRET=${WECHAT_APP_SECRET}
|
||||
WECHAT_OPEN_REDIRECT_URI=https://staging.xiaoxiajianji.com/auth/wechat/callback
|
||||
|
||||
@@ -115,6 +115,8 @@ services:
|
||||
APP_VERSION: ${APP_VERSION:-unknown}
|
||||
WORKER_CONCURRENCY: ${WORKER_CONCURRENCY:-4}
|
||||
WORKER_MAX_TASKS_PER_CHILD: ${WORKER_MAX_TASKS_PER_CHILD:-100}
|
||||
# #1714 队列隔离:generation 队列独占 worker(默认并发 2),其余并发给转码
|
||||
GENERATION_CONCURRENCY: ${GENERATION_CONCURRENCY:-2}
|
||||
GENERATED_FILES_DIR: /app/generated
|
||||
GENERATED_FILES_URL_PREFIX: /generated-files
|
||||
PUBLIC_API_BASE_URL: ${PUBLIC_API_BASE_URL:-https://api.xiaoxiajianji.com}
|
||||
@@ -128,11 +130,11 @@ services:
|
||||
# 健康检查配置
|
||||
# 注:celery inspect ping 依赖 broker 连接,在容器内不可靠,改用进程检查
|
||||
healthcheck:
|
||||
test: ["CMD-SHELL", "grep -q celery /proc/1/cmdline || exit 1"]
|
||||
test: ["CMD-SHELL", "pgrep -f 'celery.*worker' | head -n1 >/dev/null 2>&1 || grep -q celery /proc/1/cmdline || exit 1"]
|
||||
interval: 30s
|
||||
timeout: 10s
|
||||
retries: 3
|
||||
start_period: 30s
|
||||
start_period: 40s
|
||||
|
||||
logging: *default-logging
|
||||
|
||||
@@ -140,7 +142,8 @@ services:
|
||||
# 资源限制建议(生产环境建议启用)
|
||||
# =========================================
|
||||
# 注意: Worker 需要处理视频,建议分配更多资源
|
||||
# 并发 4 时需要 4C8G 以上,确保视频渲染不 OOM
|
||||
# #1714 队列隔离后容器内运行 generation + transcode 两个 worker 进程,
|
||||
# 总并发 = WORKER_CONCURRENCY(默认 4),4C8G 以上确保视频渲染不 OOM
|
||||
deploy:
|
||||
resources:
|
||||
limits:
|
||||
|
||||
@@ -146,6 +146,7 @@ docker run -d \
|
||||
-e APP_ENV=production \
|
||||
-e APP_VERSION="$IMAGE_TAG" \
|
||||
-e WORKER_CONCURRENCY="${WORKER_CONCURRENCY:-4}" \
|
||||
-e GENERATION_CONCURRENCY="${GENERATION_CONCURRENCY:-2}" \
|
||||
-e WORKER_MAX_TASKS_PER_CHILD=100 \
|
||||
-e GENERATED_FILES_DIR=/app/generated \
|
||||
-e GENERATED_FILES_URL_PREFIX=/generated-files \
|
||||
@@ -154,7 +155,7 @@ docker run -d \
|
||||
--restart unless-stopped \
|
||||
--cpus 2 \
|
||||
--memory 2g \
|
||||
--health-cmd "sh -c \"grep -q celery /proc/1/cmdline || exit 1\"" \
|
||||
--health-cmd "sh -c \"pgrep -f 'celery.*worker' >/dev/null 2>&1 || grep -q celery /proc/1/cmdline || exit 1\"" \
|
||||
--health-interval 30s \
|
||||
--health-timeout 10s \
|
||||
--health-retries 3 \
|
||||
|
||||
@@ -109,13 +109,14 @@ docker run -d \
|
||||
-e APP_ENV=staging \
|
||||
-e APP_VERSION="$IMAGE_TAG" \
|
||||
-e WORKER_CONCURRENCY="${WORKER_CONCURRENCY:-4}" \
|
||||
-e GENERATION_CONCURRENCY="${GENERATION_CONCURRENCY:-2}" \
|
||||
-e WORKER_MAX_TASKS_PER_CHILD=100 \
|
||||
-e GENERATED_FILES_DIR=/app/generated \
|
||||
-e GENERATED_FILES_URL_PREFIX=/generated-files \
|
||||
-v "$GENERATED_DIR:/app/generated" \
|
||||
--restart unless-stopped \
|
||||
--label com.centurylinklabs.watchtower.enable=true \
|
||||
--health-cmd "sh -c \"grep -q celery /proc/1/cmdline || exit 1\"" \
|
||||
--health-cmd "sh -c \"pgrep -f 'celery.*worker' >/dev/null 2>&1 || grep -q celery /proc/1/cmdline || exit 1\"" \
|
||||
--health-interval 30s \
|
||||
--health-timeout 10s \
|
||||
--health-retries 3 \
|
||||
|
||||
@@ -1,18 +1,63 @@
|
||||
#!/bin/bash
|
||||
# Worker 启动脚本 — 支持 WORKER_CONCURRENCY 环境变量
|
||||
# 未设置时默认 2(保持向后兼容)
|
||||
# Worker 启动脚本 — #1714 队列隔离
|
||||
#
|
||||
# 部署约束:worker 容器单实例(replicas=1),容器内启动两个 celery 进程:
|
||||
# 1. generation-worker:独占消费 generation 队列(用户视频生成,高优先级),
|
||||
# 内嵌 celery beat(-B),定时清理任务只在一个进程里跑,避免重复执行;
|
||||
# 2. transcode-worker:消费 transcode + celery 默认队列(素材转码/分类/查重/
|
||||
# 配音/下载等后台任务)。
|
||||
# 转码队列积压时,generation 队列仍有独立 worker 立即领取视频生成任务。
|
||||
#
|
||||
# 环境变量:
|
||||
# WORKER_CONCURRENCY 总并发槽参考(默认 4);生成 worker 并发默认 2,
|
||||
# 可用 GENERATION_CONCURRENCY 覆盖
|
||||
# GENERATION_CONCURRENCY generation worker 并发(默认 2)
|
||||
# TRANSCODE_CONCURRENCY transcode worker 并发(默认 = WORKER_CONCURRENCY - 2,最小 1)
|
||||
# WORKER_MAX_TASKS_PER_CHILD 每个子进程最大任务数(默认 100)
|
||||
|
||||
set -e
|
||||
|
||||
CONCURRENCY="${WORKER_CONCURRENCY:-2}"
|
||||
CONCURRENCY="${WORKER_CONCURRENCY:-4}"
|
||||
MAX_TASKS="${WORKER_MAX_TASKS_PER_CHILD:-100}"
|
||||
|
||||
# ⚠️ 部署约束:此 Worker 必须且只能运行单实例(replicas=1)
|
||||
# -B 标志嵌入 celery beat,beat 负责定期触发 pending 超时清理等定时任务
|
||||
# 多实例部署会导致每个 Worker 独立运行 Beat,造成定时任务重复执行
|
||||
# 若需横向扩展 Worker,必须将 Beat 拆分为独立服务(celery beat -A worker_app.celery_app)
|
||||
exec celery \
|
||||
GEN_CONCURRENCY="${GENERATION_CONCURRENCY:-2}"
|
||||
if [ -z "$TRANSCODE_CONCURRENCY" ]; then
|
||||
TRANS_CONCURRENCY=$((CONCURRENCY - GEN_CONCURRENCY))
|
||||
if [ "$TRANS_CONCURRENCY" -lt 1 ]; then
|
||||
TRANS_CONCURRENCY=1
|
||||
fi
|
||||
else
|
||||
TRANS_CONCURRENCY="$TRANSCODE_CONCURRENCY"
|
||||
fi
|
||||
|
||||
echo "Starting generation worker (queue=generation, concurrency=$GEN_CONCURRENCY, beat embedded)"
|
||||
celery \
|
||||
-A worker_app.celery_app \
|
||||
worker \
|
||||
--loglevel=info \
|
||||
"-B" \
|
||||
"--concurrency=${CONCURRENCY}"
|
||||
-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
|
||||
|
||||
@@ -146,3 +146,44 @@ class InMemoryAssetRepository:
|
||||
if asset.library_id == library_id and asset.file_hash == file_hash:
|
||||
return asset
|
||||
return None
|
||||
|
||||
def find_by_library_and_client_upload_id(
|
||||
self,
|
||||
library_id: str,
|
||||
client_upload_id: str,
|
||||
) -> Asset | None:
|
||||
"""按素材库 + 客户端幂等 token 查找已有素材。"""
|
||||
if not client_upload_id:
|
||||
return None
|
||||
for asset in self._assets.values():
|
||||
if asset.library_id == library_id and getattr(asset, "client_upload_id", "") == client_upload_id:
|
||||
return asset
|
||||
return None
|
||||
|
||||
def find_recent_active_by_library_and_name(
|
||||
self,
|
||||
library_id: str,
|
||||
name: str,
|
||||
within_minutes: int = 30,
|
||||
file_size: int = 0,
|
||||
) -> Asset | None:
|
||||
"""兜底去重:同库 + 同文件名(+同大小)且近期活动状态的素材。"""
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
if not name:
|
||||
return None
|
||||
from packages.domain import AssetStatus
|
||||
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(minutes=within_minutes)
|
||||
candidates = [
|
||||
a
|
||||
for a in self._assets.values()
|
||||
if a.library_id == library_id
|
||||
and a.name == name
|
||||
and a.status in (AssetStatus.UPLOADING, AssetStatus.PROCESSING)
|
||||
and a.created_at >= cutoff
|
||||
and (not file_size or file_size <= 0 or a.file_size == file_size)
|
||||
]
|
||||
if not candidates:
|
||||
return None
|
||||
return max(candidates, key=lambda a: a.created_at)
|
||||
|
||||
@@ -134,6 +134,7 @@ class SQLAlchemyAssetRepository:
|
||||
quality_score=asset.quality_score,
|
||||
uploaded_by_user_id=asset.uploaded_by_user_id or "system",
|
||||
file_hash=asset.file_hash or None,
|
||||
client_upload_id=asset.client_upload_id or None,
|
||||
created_at=asset.created_at,
|
||||
updated_at=now,
|
||||
)
|
||||
@@ -163,6 +164,8 @@ class SQLAlchemyAssetRepository:
|
||||
model.quality_score = asset.quality_score
|
||||
model.uploaded_by_user_id = asset.uploaded_by_user_id or model.uploaded_by_user_id
|
||||
model.file_hash = asset.file_hash or model.file_hash
|
||||
if getattr(model, "client_upload_id", None) is None and asset.client_upload_id:
|
||||
model.client_upload_id = asset.client_upload_id
|
||||
model.updated_at = datetime.now(timezone.utc)
|
||||
self.session.flush()
|
||||
self._sync_asset_tags(asset.id, asset.tag_ids)
|
||||
@@ -388,6 +391,7 @@ class SQLAlchemyAssetRepository:
|
||||
quality_score=model.quality_score,
|
||||
uploaded_by_user_id=model.uploaded_by_user_id,
|
||||
file_hash=model.file_hash or "",
|
||||
client_upload_id=getattr(model, "client_upload_id", None) or "",
|
||||
metadata=metadata,
|
||||
tag_ids=tag_ids,
|
||||
created_at=model.created_at,
|
||||
@@ -452,3 +456,53 @@ 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,6 +38,7 @@ def _to_domain(model: GenerationTaskModel) -> GenerationTask:
|
||||
bgm_config=dict(getattr(model, "bgm_config", {}) or {}),
|
||||
is_preview=bool(getattr(model, "is_preview", False)),
|
||||
source_task_id=getattr(model, "source_task_id", "") or "",
|
||||
celery_task_id=getattr(model, "celery_task_id", "") or "",
|
||||
output_width=getattr(model, "output_width", 1280) or 1280,
|
||||
output_height=getattr(model, "output_height", 720) or 720,
|
||||
cover_url=getattr(model, "cover_url", "") or "",
|
||||
@@ -82,6 +83,7 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
bgm_config=task.bgm_config or {},
|
||||
is_preview=task.is_preview or False,
|
||||
source_task_id=task.source_task_id or "",
|
||||
celery_task_id=getattr(task, "celery_task_id", "") or "",
|
||||
output_width=task.output_width,
|
||||
output_height=task.output_height,
|
||||
cover_url=task.cover_url or "",
|
||||
@@ -138,6 +140,52 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
.count()
|
||||
)
|
||||
|
||||
def count_running_by_user(self, user_id: str) -> int:
|
||||
"""统计指定用户处于 running 状态的任务数(用于限流提示展示)。"""
|
||||
return (
|
||||
self.session.query(GenerationTaskModel)
|
||||
.filter(
|
||||
GenerationTaskModel.created_by_user_id == user_id,
|
||||
GenerationTaskModel.status == GenerationTaskStatus.RUNNING.value,
|
||||
)
|
||||
.count()
|
||||
)
|
||||
|
||||
def count_running_total(self) -> int:
|
||||
"""统计全局处于 running 状态的任务数(worker 实际在执行的任务数)。"""
|
||||
return (
|
||||
self.session.query(GenerationTaskModel)
|
||||
.filter(GenerationTaskModel.status == GenerationTaskStatus.RUNNING.value)
|
||||
.count()
|
||||
)
|
||||
|
||||
def estimate_avg_duration_seconds(self, limit: int = 20, default_seconds: float = 120.0) -> float:
|
||||
"""估算最近完成任务的平均耗时(秒),用于 429 限流提示的等待预估。
|
||||
|
||||
取最近 N 条 completed 任务的 (completed_at - started_at) 平均值;
|
||||
无足够历史数据时返回 default_seconds。
|
||||
用 Python 侧计算差值,避免 SQLite/PostgreSQL 方言差异。
|
||||
"""
|
||||
rows = (
|
||||
self.session.query(GenerationTaskModel.started_at, GenerationTaskModel.completed_at)
|
||||
.filter(
|
||||
GenerationTaskModel.status == GenerationTaskStatus.COMPLETED.value,
|
||||
GenerationTaskModel.started_at.isnot(None),
|
||||
GenerationTaskModel.completed_at.isnot(None),
|
||||
)
|
||||
.order_by(GenerationTaskModel.completed_at.desc())
|
||||
.limit(limit)
|
||||
.all()
|
||||
)
|
||||
durations = [
|
||||
(completed - started).total_seconds()
|
||||
for started, completed in rows
|
||||
if completed and started and (completed - started).total_seconds() > 0
|
||||
]
|
||||
if not durations:
|
||||
return default_seconds
|
||||
return sum(durations) / len(durations)
|
||||
|
||||
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]:
|
||||
models = (
|
||||
self.session.query(GenerationTaskModel)
|
||||
@@ -269,6 +317,7 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
if hasattr(model, "is_preview"):
|
||||
model.is_preview = task.is_preview or False
|
||||
model.source_task_id = task.source_task_id or ""
|
||||
model.celery_task_id = getattr(task, "celery_task_id", "") or model.celery_task_id or ""
|
||||
model.output_width = task.output_width
|
||||
model.output_height = task.output_height
|
||||
model.cover_url = task.cover_url or ""
|
||||
@@ -280,12 +329,14 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
def cleanup_stale_running(self, timeout_minutes: int = 10) -> int:
|
||||
"""清理超时未更新的 running 任务(孤儿任务)。
|
||||
|
||||
将 status=running 且 updated_at 超过 timeout_minutes 分钟未更新的任务
|
||||
标记为 failed,error_message 标记为任务执行中断。
|
||||
|
||||
Returns:
|
||||
清理的任务数量
|
||||
清理的任务数量(仅计数,保持旧签名兼容)
|
||||
"""
|
||||
items = self.cleanup_stale_running_with_ids(timeout_minutes)
|
||||
return len(items)
|
||||
|
||||
def cleanup_stale_running_with_ids(self, timeout_minutes: int = 10) -> list[tuple[str, str]]:
|
||||
"""同 cleanup_stale_running,但返回 [(task_id, celery_task_id), ...] 供撤销队列消息。"""
|
||||
from datetime import timedelta
|
||||
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes)
|
||||
@@ -298,8 +349,10 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
.all()
|
||||
)
|
||||
if not models:
|
||||
return 0
|
||||
return []
|
||||
result: list[tuple[str, str]] = []
|
||||
for model in models:
|
||||
result.append((model.id, getattr(model, "celery_task_id", "") or ""))
|
||||
model.status = GenerationTaskStatus.FAILED.value
|
||||
model.error_message = "任务执行中断(worker重启/超时)"
|
||||
model.error_info = {
|
||||
@@ -309,43 +362,43 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
}
|
||||
model.completed_at = datetime.now(timezone.utc)
|
||||
self.session.commit()
|
||||
return len(models)
|
||||
return result
|
||||
|
||||
def cleanup_stale_pending(self, timeout_minutes: int = 30) -> int:
|
||||
"""清理超时的 pending 任务(未被 Worker 拉取的任务)。
|
||||
|
||||
全局任务队列有 pending 数量上限,长期卡在 pending 的任务会占满队列,
|
||||
导致新用户无法创建任务。将超时的 pending 任务标记为 failed。
|
||||
|
||||
Args:
|
||||
timeout_minutes: 超时时间(分钟),默认 30 分钟
|
||||
|
||||
Returns:
|
||||
清理的任务数量
|
||||
清理的任务数量(仅计数,保持旧签名兼容)
|
||||
"""
|
||||
items = self.cleanup_stale_pending_with_ids(timeout_minutes)
|
||||
return len(items)
|
||||
|
||||
def cleanup_stale_pending_with_ids(self, timeout_minutes: int = 30) -> list[tuple[str, str]]:
|
||||
"""同 cleanup_stale_pending,但返回 [(task_id, celery_task_id), ...] 供撤销队列消息。"""
|
||||
from datetime import timedelta
|
||||
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(minutes=timeout_minutes)
|
||||
error_info = {
|
||||
"error_type": "PendingTimeout",
|
||||
"message": f"任务在 pending 状态停留超过 {timeout_minutes} 分钟,自动清理",
|
||||
"failed_at": datetime.now(timezone.utc).isoformat(),
|
||||
}
|
||||
count = (
|
||||
models = (
|
||||
self.session.query(GenerationTaskModel)
|
||||
.filter(
|
||||
GenerationTaskModel.status == GenerationTaskStatus.PENDING.value,
|
||||
GenerationTaskModel.created_at < cutoff,
|
||||
)
|
||||
.update(
|
||||
{
|
||||
GenerationTaskModel.status: GenerationTaskStatus.FAILED.value,
|
||||
GenerationTaskModel.error_message: "pending timeout: auto cleanup",
|
||||
GenerationTaskModel.error_info: error_info,
|
||||
GenerationTaskModel.completed_at: datetime.now(timezone.utc),
|
||||
},
|
||||
synchronize_session=False,
|
||||
)
|
||||
.all()
|
||||
)
|
||||
if not models:
|
||||
return []
|
||||
error_info = {
|
||||
"error_type": "PendingTimeout",
|
||||
"message": f"任务在 pending 状态停留超过 {timeout_minutes} 分钟,自动清理",
|
||||
"failed_at": datetime.now(timezone.utc).isoformat(),
|
||||
}
|
||||
result: list[tuple[str, str]] = []
|
||||
for model in models:
|
||||
result.append((model.id, getattr(model, "celery_task_id", "") or ""))
|
||||
model.status = GenerationTaskStatus.FAILED.value
|
||||
model.error_message = "pending timeout: auto cleanup"
|
||||
model.error_info = error_info
|
||||
model.completed_at = datetime.now(timezone.utc)
|
||||
self.session.commit()
|
||||
return count
|
||||
return result
|
||||
|
||||
@@ -18,6 +18,8 @@ class SQLAlchemyIngestJobRepository:
|
||||
error_message=job.error_message,
|
||||
result_asset_id=job.result_asset_id,
|
||||
file_hash=job.file_hash,
|
||||
asset_id=job.asset_id or "",
|
||||
celery_task_id=getattr(job, "celery_task_id", "") or "",
|
||||
created_at=job.created_at,
|
||||
updated_at=job.updated_at,
|
||||
)
|
||||
@@ -38,6 +40,8 @@ class SQLAlchemyIngestJobRepository:
|
||||
error_message=model.error_message,
|
||||
result_asset_id=model.result_asset_id,
|
||||
file_hash=model.file_hash or "",
|
||||
asset_id=getattr(model, "asset_id", "") or "",
|
||||
celery_task_id=getattr(model, "celery_task_id", "") or "",
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
@@ -54,6 +58,12 @@ class SQLAlchemyIngestJobRepository:
|
||||
model.error_message = job.error_message
|
||||
model.result_asset_id = job.result_asset_id
|
||||
model.file_hash = job.file_hash
|
||||
model.storage_key = job.storage_key
|
||||
if job.asset_id:
|
||||
model.asset_id = job.asset_id
|
||||
celery_tid = getattr(job, "celery_task_id", "")
|
||||
if celery_tid:
|
||||
model.celery_task_id = celery_tid
|
||||
model.updated_at = job.updated_at
|
||||
self.session.commit()
|
||||
return job
|
||||
|
||||
@@ -94,6 +94,7 @@ class AssetModel(Base):
|
||||
quality_score = Column(Float, nullable=True)
|
||||
uploaded_by_user_id = Column(String(36), nullable=False)
|
||||
file_hash = Column(String(64), nullable=True, index=True)
|
||||
client_upload_id = Column(String(64), nullable=True, index=True)
|
||||
extra_meta = Column("metadata", JSON, nullable=False, default=dict)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc), index=True)
|
||||
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
@@ -244,6 +245,8 @@ class IngestJobModel(Base):
|
||||
error_message = Column(Text, nullable=False, default="")
|
||||
result_asset_id = Column(String(36), nullable=False, default="")
|
||||
file_hash = Column(String(64), nullable=True, index=True)
|
||||
asset_id = Column(String(36), nullable=False, default="", index=True)
|
||||
celery_task_id = Column(String(64), nullable=False, default="", server_default="")
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
@@ -296,6 +299,7 @@ class GenerationTaskModel(Base):
|
||||
resolution = Column(String(20), nullable=False, default="")
|
||||
is_preview = Column(Boolean, nullable=False, default=False, index=True)
|
||||
source_task_id = Column(String(32), nullable=False, default="", index=True)
|
||||
celery_task_id = Column(String(64), nullable=False, default="", server_default="")
|
||||
output_width = Column(Integer, nullable=False, default=1280)
|
||||
output_height = Column(Integer, nullable=False, default=720)
|
||||
cover_url = Column(String(1000), nullable=False, default="")
|
||||
|
||||
@@ -200,7 +200,15 @@ class WechatOAuthService:
|
||||
return None, "微信登录处理失败"
|
||||
|
||||
|
||||
# 模块级单例:state 存储必须跨请求共享,否则 /wechat/url 生成的 state
|
||||
# 与 /wechat/callback 校验时不在同一个 MemoryStateStore,回调必然 400。
|
||||
# 多实例部署时应替换为 Redis state store(单容器多 worker 也需如此)。
|
||||
_oauth_service_singleton: WechatOAuthService | None = None
|
||||
|
||||
|
||||
def get_wechat_oauth_service() -> WechatOAuthService:
|
||||
"""获取微信 OAuth 服务单例"""
|
||||
# TODO: 可替换为 Redis state store
|
||||
return WechatOAuthService()
|
||||
"""获取微信 OAuth 服务单例(state store 跨请求共享)"""
|
||||
global _oauth_service_singleton
|
||||
if _oauth_service_singleton is None:
|
||||
_oauth_service_singleton = WechatOAuthService()
|
||||
return _oauth_service_singleton
|
||||
|
||||
@@ -12,6 +12,8 @@ class SubmitIngestJobCommand:
|
||||
library_id: str
|
||||
storage_key: str
|
||||
file_hash: str = ""
|
||||
asset_id: str = ""
|
||||
celery_task_id: str = ""
|
||||
|
||||
|
||||
class SubmitIngestJobUseCase:
|
||||
@@ -24,5 +26,7 @@ class SubmitIngestJobUseCase:
|
||||
library_id=command.library_id,
|
||||
storage_key=command.storage_key,
|
||||
file_hash=command.file_hash,
|
||||
asset_id=command.asset_id,
|
||||
celery_task_id=command.celery_task_id,
|
||||
)
|
||||
return self.ingest_job_repository.create(job)
|
||||
|
||||
@@ -174,6 +174,7 @@ class Asset:
|
||||
quality_score: float | None = None
|
||||
uploaded_by_user_id: str = ""
|
||||
file_hash: str = ""
|
||||
client_upload_id: str = ""
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
tag_ids: list[str] = field(default_factory=list)
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
@@ -208,6 +209,7 @@ class Asset:
|
||||
quality_score: float | None = None,
|
||||
uploaded_by_user_id: str = "",
|
||||
file_hash: str = "",
|
||||
client_upload_id: str = "",
|
||||
) -> "Asset":
|
||||
clean_name = name.strip()
|
||||
if not clean_name:
|
||||
@@ -235,6 +237,7 @@ class Asset:
|
||||
quality_score=quality_score,
|
||||
uploaded_by_user_id=uploaded_by_user_id.strip(),
|
||||
file_hash=file_hash.strip(),
|
||||
client_upload_id=client_upload_id.strip(),
|
||||
metadata=metadata or {},
|
||||
tag_ids=[],
|
||||
)
|
||||
@@ -266,6 +269,8 @@ class IngestJob:
|
||||
error_message: str = ""
|
||||
result_asset_id: str = ""
|
||||
file_hash: str = ""
|
||||
asset_id: str = ""
|
||||
celery_task_id: str = ""
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
@@ -276,6 +281,8 @@ class IngestJob:
|
||||
library_id: str,
|
||||
storage_key: str,
|
||||
file_hash: str = "",
|
||||
asset_id: str = "",
|
||||
celery_task_id: str = "",
|
||||
) -> "IngestJob":
|
||||
if not project_id.strip():
|
||||
raise ValueError("project_id 不能为空")
|
||||
@@ -289,4 +296,6 @@ class IngestJob:
|
||||
library_id=library_id.strip(),
|
||||
storage_key=storage_key.strip(),
|
||||
file_hash=file_hash.strip(),
|
||||
asset_id=asset_id.strip(),
|
||||
celery_task_id=celery_task_id.strip(),
|
||||
)
|
||||
|
||||
@@ -117,6 +117,7 @@ class GenerationTask:
|
||||
bgm_config: dict = field(default_factory=dict)
|
||||
is_preview: bool = False
|
||||
source_task_id: str = ""
|
||||
celery_task_id: str = ""
|
||||
output_width: int = 1280
|
||||
output_height: int = 720
|
||||
cover_url: str = ""
|
||||
|
||||
@@ -125,3 +125,23 @@ class AssetRepository(ABC):
|
||||
) -> Asset | None:
|
||||
"""按素材库 + 文件哈希查找已有素材(去重检测)。"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def find_by_library_and_client_upload_id(
|
||||
self,
|
||||
library_id: str,
|
||||
client_upload_id: str,
|
||||
) -> Asset | None:
|
||||
"""按素材库 + 客户端幂等 token 查找已有素材(complete 幂等)。"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def find_recent_active_by_library_and_name(
|
||||
self,
|
||||
library_id: str,
|
||||
name: str,
|
||||
within_minutes: int = 30,
|
||||
file_size: int = 0,
|
||||
) -> Asset | None:
|
||||
"""兜底去重:同库 + 同文件名(+同大小)且近期仍在 uploading/processing 的素材。"""
|
||||
pass
|
||||
|
||||
@@ -20,6 +20,12 @@ class GenerationTaskRepository(Protocol):
|
||||
|
||||
def count_pending_total(self) -> int: ...
|
||||
|
||||
def count_running_by_user(self, user_id: str) -> int: ...
|
||||
|
||||
def count_running_total(self) -> int: ...
|
||||
|
||||
def estimate_avg_duration_seconds(self, limit: int = 20, default_seconds: float = 120.0) -> float: ...
|
||||
|
||||
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]: ...
|
||||
|
||||
def list_by_source_edit_plan(self, plan_id: str) -> list[GenerationTask]: ...
|
||||
|
||||
@@ -0,0 +1,222 @@
|
||||
"""孤儿任务消息撤销与执行前状态守卫(API / Worker 共享)。
|
||||
|
||||
#1714 / #1710 缺陷修复:超时清理/孤儿恢复把 DB 任务标记为 failed/cancelled
|
||||
后,Redis 队列里对应的 Celery 消息仍然存在;worker 重启或重新拉取时该消息
|
||||
被再次执行,状态机抛「非法状态转换: failed → running」,旧实现打印 ERROR 后
|
||||
继续跑,最终产出半成品。
|
||||
|
||||
防御两道:
|
||||
1. 清理任务标 failed 时,调用 revoke_and_purge() 撤销(celery revoke 广播,
|
||||
通知在线 worker 丢弃)并直接扫描 Redis 队列移除消息体(worker 下线期间
|
||||
队列中的消息 revoke 广播收不到,必须物理移除);
|
||||
2. 任务真正开始业务逻辑前,调用 ensure_task_claimable() 校验 DB 状态,
|
||||
非 pending 的消息直接丢弃(抛 StaleTaskDiscarded,task 捕获后安全返回,
|
||||
不进入渲染/转码,不产出半成品)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import json
|
||||
import logging
|
||||
from collections.abc import Callable, Iterable
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class StaleTaskDiscarded(Exception):
|
||||
"""任务消息已作废(DB 中任务已是终态),应安全中止、丢弃消息。"""
|
||||
|
||||
def __init__(self, task_id: str, status: str):
|
||||
self.task_id = task_id
|
||||
self.status = status
|
||||
super().__init__(f"任务 {task_id} 已是终态 {status},丢弃重复/作废消息")
|
||||
|
||||
|
||||
# 终态状态值集合:处于这些状态的任务消息一律不执行
|
||||
TERMINAL_STATUS_VALUES = frozenset({"failed", "cancelled", "completed"})
|
||||
|
||||
|
||||
def ensure_task_claimable(
|
||||
task_id: str,
|
||||
get_status: Callable[[str], str | None],
|
||||
*,
|
||||
task_label: str = "任务",
|
||||
) -> str:
|
||||
"""执行前守卫:任务必须处于可领取状态(pending)。
|
||||
|
||||
Args:
|
||||
task_id: 业务任务 ID
|
||||
get_status: 回调,返回 DB 中任务当前状态字符串;返回 None 表示任务不存在
|
||||
task_label: 日志用任务类型名
|
||||
|
||||
Returns:
|
||||
当前状态字符串(pending);任务不存在时返回空串(由调用方处理 not found)
|
||||
|
||||
Raises:
|
||||
StaleTaskDiscarded: 任务已是终态(failed/cancelled/completed),消息必须丢弃
|
||||
"""
|
||||
status = get_status(task_id)
|
||||
if status is None:
|
||||
return ""
|
||||
if status in TERMINAL_STATUS_VALUES:
|
||||
logger.warning("[%s] task_id=%s 状态已为 %s,消息作废,丢弃不执行", task_label, task_id, status)
|
||||
raise StaleTaskDiscarded(task_id, status)
|
||||
return status
|
||||
|
||||
|
||||
def _extract_business_ids(raw: bytes) -> tuple[str | None, str | None]:
|
||||
"""从 Redis 中的 Celery 消息提取 (celery 消息 ID, 业务任务 ID)。
|
||||
|
||||
Redis transport 存储格式为 JSON 信封:
|
||||
{"body": base64(json), "headers": {"id": <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 消息 ID(headers.id)匹配,
|
||||
任一命中即移除。未命中或无法解析的消息原样保留(保持相对顺序)。
|
||||
|
||||
Returns:
|
||||
实际移除的消息条数
|
||||
"""
|
||||
biz_ids = {bid for bid in business_task_ids if bid}
|
||||
msg_ids = {mid for mid in celery_task_ids if mid}
|
||||
if not biz_ids and not msg_ids:
|
||||
return 0
|
||||
|
||||
try:
|
||||
import redis
|
||||
except ImportError:
|
||||
logger.warning("redis-py 不可用,跳过队列消息清理")
|
||||
return 0
|
||||
|
||||
try:
|
||||
client = redis.Redis.from_url(broker_url)
|
||||
client.ping()
|
||||
except Exception as e:
|
||||
logger.warning("连接 Redis 清理作废消息失败: %s", e)
|
||||
return 0
|
||||
|
||||
removed_total = 0
|
||||
try:
|
||||
for queue in queue_names:
|
||||
removed_total += _purge_one_queue(client, queue, biz_ids, msg_ids)
|
||||
finally:
|
||||
try:
|
||||
client.close()
|
||||
except Exception:
|
||||
pass
|
||||
if removed_total:
|
||||
logger.info(
|
||||
"从 Redis 队列移除 %d 条作废消息(biz=%s, celery=%s)",
|
||||
removed_total,
|
||||
sorted(biz_ids),
|
||||
sorted(msg_ids),
|
||||
)
|
||||
return removed_total
|
||||
|
||||
|
||||
def _purge_one_queue(client: Any, queue_name: str, biz_ids: set[str], msg_ids: set[str]) -> int:
|
||||
try:
|
||||
raw_messages = client.lrange(queue_name, 0, -1)
|
||||
except Exception as e:
|
||||
logger.warning("读取队列 %s 失败: %s", queue_name, e)
|
||||
return 0
|
||||
if not raw_messages:
|
||||
return 0
|
||||
|
||||
keep: list[bytes] = []
|
||||
removed = 0
|
||||
for raw in raw_messages:
|
||||
celery_id, biz_id = _extract_business_ids(raw)
|
||||
hit = (biz_id is not None and biz_id in biz_ids) or (celery_id is not None and celery_id in msg_ids)
|
||||
if hit:
|
||||
removed += 1
|
||||
continue
|
||||
keep.append(raw)
|
||||
|
||||
if removed:
|
||||
try:
|
||||
pipe = client.pipeline()
|
||||
pipe.delete(queue_name)
|
||||
if keep:
|
||||
pipe.rpush(queue_name, *keep)
|
||||
pipe.execute()
|
||||
except Exception as e:
|
||||
logger.warning("重写队列 %s 失败: %s", queue_name, e)
|
||||
return 0
|
||||
return removed
|
||||
|
||||
|
||||
def revoke_and_purge(
|
||||
celery_app: Any,
|
||||
broker_url: str,
|
||||
business_task_ids: Iterable[str] = (),
|
||||
celery_task_ids: Iterable[str] = (),
|
||||
*,
|
||||
queue_names: Iterable[str] = ("generation", "transcode", "celery"),
|
||||
) -> int:
|
||||
"""撤销作废任务:revoke 广播(在线 worker)+ 物理清理 Redis 队列消息。
|
||||
|
||||
Args:
|
||||
celery_app: Celery app 实例(worker 端 worker_app.celery_app.celery_app)
|
||||
broker_url: Redis broker URL
|
||||
business_task_ids: 业务任务 ID(generation_tasks.id / ingest_jobs.id)
|
||||
celery_task_ids: 入队时记录的 celery 消息 ID
|
||||
queue_names: 需要扫描清理的队列名
|
||||
|
||||
Returns:
|
||||
从队列中实际移除的消息条数
|
||||
"""
|
||||
for tid in celery_task_ids:
|
||||
if not tid:
|
||||
continue
|
||||
try:
|
||||
celery_app.control.revoke(tid)
|
||||
except Exception as e:
|
||||
logger.warning("revoke celery 消息 %s 失败: %s", tid, e)
|
||||
|
||||
return purge_stale_messages_from_queues(
|
||||
broker_url, queue_names, business_task_ids=business_task_ids, celery_task_ids=celery_task_ids
|
||||
)
|
||||
@@ -0,0 +1,58 @@
|
||||
"""Celery 队列定义与路由配置(API / Worker 共享)。
|
||||
|
||||
#1714 队列隔离:用户等待的视频生成任务路由到高优先级 `generation` 队列,
|
||||
由专用 worker 进程独占消费;素材入库/转码等后台批量任务路由到 `transcode`
|
||||
队列;其余杂项任务走默认 `celery` 队列。转码队列积压时,视频生成任务
|
||||
仍能被 generation worker 立即领取执行,不会排队。
|
||||
|
||||
队列说明:
|
||||
- generation: 用户提交的视频生成/预览渲染(延迟敏感,资源消耗大)
|
||||
- transcode: 素材入库(HEVC 转码)、AI 分类、素材查重(批量、可排队)
|
||||
- celery(默认): 配音、语音、下载缩略图、定时清理等杂项
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from kombu import Queue
|
||||
|
||||
# ── 队列名常量(生产端与消费端共用,禁止拼写漂移) ──
|
||||
QUEUE_GENERATION = "generation"
|
||||
QUEUE_TRANSCODE = "transcode"
|
||||
QUEUE_DEFAULT = "celery"
|
||||
|
||||
# Worker 消费的队列列表(顺序即优先级:高优队列排在前面)
|
||||
WORKER_QUEUES = (QUEUE_GENERATION, QUEUE_TRANSCODE, QUEUE_DEFAULT)
|
||||
|
||||
# 队列声明:持久化队列,broker 重启不丢消息
|
||||
task_queues = (
|
||||
Queue(QUEUE_GENERATION, routing_key=QUEUE_GENERATION, durable=True),
|
||||
Queue(QUEUE_TRANSCODE, routing_key=QUEUE_TRANSCODE, durable=True),
|
||||
Queue(QUEUE_DEFAULT, routing_key=QUEUE_DEFAULT, durable=True),
|
||||
)
|
||||
|
||||
# ── 任务路由表:task name → 队列 ──
|
||||
# 键支持 celery 标准通配符。
|
||||
task_routes = {
|
||||
# 高优先级:用户等待的视频生成
|
||||
"worker.generate_video": {"queue": QUEUE_GENERATION},
|
||||
# 后台批量:素材入库/转码 + AI 分类 + 素材查重,积压不影响生成
|
||||
"worker.ingest_asset": {"queue": QUEUE_TRANSCODE},
|
||||
"worker.classify_asset": {"queue": QUEUE_TRANSCODE},
|
||||
"worker.process_duplication_check": {"queue": QUEUE_TRANSCODE},
|
||||
"worker.check_duplicate": {"queue": QUEUE_TRANSCODE},
|
||||
}
|
||||
|
||||
# 生成任务的预取数:渲染是长任务,预取 1 避免任务被某个 worker 占住不调度
|
||||
GENERATION_WORKER_PREFETCH_MULTIPLIER = 1
|
||||
|
||||
|
||||
def apply_queue_settings(app) -> None:
|
||||
"""把队列隔离配置应用到 Celery app(API 生产端与 Worker 消费端都要调用)。
|
||||
|
||||
配置 task_queues / task_routes / task_default_queue。生产端靠 task_routes
|
||||
把消息投递到对应队列;消费端靠 task_queues 声明自己消费哪些队列
|
||||
(实际消费集由启动参数 -Q 控制)。
|
||||
"""
|
||||
app.conf.task_queues = task_queues
|
||||
app.conf.task_routes = task_routes
|
||||
app.conf.task_default_queue = QUEUE_DEFAULT
|
||||
@@ -57,7 +57,7 @@ if [ "$TARGET_ENV" = "staging" ]; then
|
||||
fi
|
||||
|
||||
# 共用 secrets 直接导出(如果存在)
|
||||
SHARED_SECRETS="OSS_ACCESS_KEY_ID OSS_ACCESS_KEY_SECRET COSYVOICE_API_KEY DASHSCOPE_API_KEY MEDIAKIT_API_KEY"
|
||||
SHARED_SECRETS="OSS_ACCESS_KEY_ID OSS_ACCESS_KEY_SECRET COSYVOICE_API_KEY DASHSCOPE_API_KEY MEDIAKIT_API_KEY WECHAT_APP_ID WECHAT_APP_SECRET"
|
||||
for var in $SHARED_SECRETS; do
|
||||
value="${!var:-}"
|
||||
# 已经在环境中了,无需额外操作
|
||||
|
||||
+1
-1
@@ -14,4 +14,4 @@ Write-Host "`n启动 Celery Worker..." -ForegroundColor Yellow
|
||||
Write-Host "监听任务队列: Redis (47.98.113.167:6379)" -ForegroundColor Cyan
|
||||
Write-Host "`n按 Ctrl+C 停止服务`n" -ForegroundColor Gray
|
||||
|
||||
celery -A celery_app worker --loglevel=info --pool=solo
|
||||
celery -A celery_app worker --loglevel=info --pool=solo -Q generation,transcode,celery
|
||||
|
||||
@@ -85,17 +85,18 @@ class TestIsBadFingerprint:
|
||||
assert VideoDeduplicator._is_bad_fingerprint(["abcdef0123456789"]) is False
|
||||
|
||||
def test_all_identical_phashes_is_bad(self):
|
||||
"""多帧但所有 phash 完全相同 → 黑屏/纯色视频。"""
|
||||
phashes = ["aaaaaaaaaaaaaaaa"] * 5
|
||||
""">=8 帧且所有 phash 完全相同 → 黑屏/纯色视频(#1702:短帧不误杀)。"""
|
||||
phashes = ["aaaaaaaaaaaaaaaa"] * 10
|
||||
assert VideoDeduplicator._is_bad_fingerprint(phashes) is True
|
||||
|
||||
def test_two_identical_phashes_is_bad(self):
|
||||
"""两帧完全相同也视为坏指纹。"""
|
||||
assert VideoDeduplicator._is_bad_fingerprint(["bbbbbbbbbbbbbbbb", "bbbbbbbbbbbbbbbb"]) is True
|
||||
def test_short_identical_phashes_not_bad(self):
|
||||
"""<8 帧完全相同不判坏——短视频内容连续时相邻采样帧 phash 天然相同(#1702)。"""
|
||||
assert VideoDeduplicator._is_bad_fingerprint(["bbbbbbbbbbbbbbbb"] * 5) is False
|
||||
assert VideoDeduplicator._is_bad_fingerprint(["bbbbbbbbbbbbbbbb", "bbbbbbbbbbbbbbbb"]) is False
|
||||
|
||||
def test_all_very_similar_phashes_is_bad(self):
|
||||
"""多帧 phash 之间的汉明距离都 < 3 → 近似黑屏。"""
|
||||
phashes = ["0000000000000000", "0000000000000001", "0000000000000002"]
|
||||
""">=8 帧 phash 之间的汉明距离都 < 3 且高占比 → 近似黑屏。"""
|
||||
phashes = ["0000000000000000"] * 8 + ["0000000000000001", "0000000000000002"]
|
||||
assert VideoDeduplicator._is_bad_fingerprint(phashes) is True
|
||||
|
||||
def test_diverse_phashes_is_good(self):
|
||||
@@ -122,7 +123,9 @@ class TestIsBadFingerprint:
|
||||
"""已知黑屏视频的 phash 特征(全零或均匀分布)。"""
|
||||
assert VideoDeduplicator._is_bad_fingerprint(["0000000000000000"] * 10) is True
|
||||
assert VideoDeduplicator._is_bad_fingerprint(["ffffffffffffffff"] * 8) is True
|
||||
assert VideoDeduplicator._is_bad_fingerprint(["9999999999999966"] * 6) is True
|
||||
assert VideoDeduplicator._is_bad_fingerprint(["9999999999999966"] * 8) is True
|
||||
# <8 帧不判坏(#1702 短视频保护)
|
||||
assert VideoDeduplicator._is_bad_fingerprint(["9999999999999966"] * 5) is False
|
||||
|
||||
|
||||
# ── Helper ──────────────────────────────────────────────────────
|
||||
@@ -151,13 +154,13 @@ class TestCheckDuplicateBadFingerprint:
|
||||
deduplicator = VideoDeduplicator()
|
||||
mock_session = MagicMock()
|
||||
|
||||
black_screen = _make_existing_video("vid-black", "md5_black", ["aaaaaaaaaaaaaaaa"] * 5)
|
||||
black_screen = _make_existing_video("vid-black", "md5_black", ["aaaaaaaaaaaaaaaa"] * 10)
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_by_user.return_value = [black_screen]
|
||||
|
||||
fingerprint = VideoFingerprint(
|
||||
md5="md5_normal",
|
||||
keyframe_phashes=["aaaaaaaaaaaaaaaa"] * 5,
|
||||
keyframe_phashes=["aaaaaaaaaaaaaaaa"] * 10,
|
||||
color_histograms=[],
|
||||
duration=10.0,
|
||||
resolution=(1280, 720),
|
||||
@@ -206,7 +209,7 @@ class TestCheckDuplicateBadFingerprint:
|
||||
deduplicator = VideoDeduplicator()
|
||||
mock_session = MagicMock()
|
||||
|
||||
black_screen = _make_existing_video("vid-black", "same_md5", ["aaaaaaaaaaaaaaaa"] * 5)
|
||||
black_screen = _make_existing_video("vid-black", "same_md5", ["aaaaaaaaaaaaaaaa"] * 10)
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_by_user.return_value = [black_screen]
|
||||
|
||||
@@ -283,8 +286,8 @@ class TestComputeDuplicateRateBadFingerprint:
|
||||
mock_session = MagicMock()
|
||||
|
||||
videos = [
|
||||
_make_existing_video("vid-b1", "md5_b1", ["aaaaaaaaaaaaaaaa"] * 5),
|
||||
_make_existing_video("vid-b2", "md5_b2", ["bbbbbbbbbbbbbbbb"] * 5),
|
||||
_make_existing_video("vid-b1", "md5_b1", ["aaaaaaaaaaaaaaaa"] * 10),
|
||||
_make_existing_video("vid-b2", "md5_b2", ["cccccccccccccccc"] * 5), # hamming(a,c)=32 > PHASH_THRESHOLD
|
||||
]
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_by_user.return_value = videos
|
||||
|
||||
@@ -0,0 +1,176 @@
|
||||
"""#1714 队列隔离 + 作废消息清除 单元测试。
|
||||
|
||||
覆盖:
|
||||
1. task_routes:generate_video → generation,ingest_asset/classify/duplication → transcode
|
||||
2. purge_stale_messages_from_queues:Redis 队列中作废任务消息被物理移除,未命中保留
|
||||
3. revoke_and_purge:revoke 广播 + 队列清理同时生效
|
||||
4. ensure_task_claimable:终态任务抛 StaleTaskDiscarded,pending 放行
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from celery import Celery
|
||||
|
||||
from packages.shared.celery_orphan_guard import (
|
||||
StaleTaskDiscarded,
|
||||
_extract_business_ids,
|
||||
ensure_task_claimable,
|
||||
purge_stale_messages_from_queues,
|
||||
revoke_and_purge,
|
||||
)
|
||||
from packages.shared.celery_queues import (
|
||||
QUEUE_GENERATION,
|
||||
QUEUE_TRANSCODE,
|
||||
apply_queue_settings,
|
||||
task_routes,
|
||||
)
|
||||
|
||||
BROKER_URL = "redis://localhost:6379/15"
|
||||
TEST_QUEUES = ("_test_gen_q", "_test_transcode_q")
|
||||
|
||||
|
||||
# ── 1. 路由表 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_routes_send_generation_to_generation_queue():
|
||||
assert task_routes["worker.generate_video"]["queue"] == QUEUE_GENERATION
|
||||
|
||||
|
||||
def test_routes_send_ingest_to_transcode_queue():
|
||||
assert task_routes["worker.ingest_asset"]["queue"] == QUEUE_TRANSCODE
|
||||
assert task_routes["worker.classify_asset"]["queue"] == QUEUE_TRANSCODE
|
||||
assert task_routes["worker.process_duplication_check"]["queue"] == QUEUE_TRANSCODE
|
||||
assert task_routes["worker.check_duplicate"]["queue"] == QUEUE_TRANSCODE
|
||||
|
||||
|
||||
def test_apply_queue_settings_configures_celery_app():
|
||||
app = Celery("test-routes")
|
||||
apply_queue_settings(app)
|
||||
queue_names = {q.name for q in app.conf.task_queues}
|
||||
assert queue_names == {"generation", "transcode", "celery"}
|
||||
assert app.conf.task_default_queue == "celery"
|
||||
|
||||
|
||||
# ── Redis 队列消息清理(需要本地 redis;不可用时 skip) ─────────────────
|
||||
|
||||
|
||||
def _redis_available() -> bool:
|
||||
try:
|
||||
import redis
|
||||
|
||||
return bool(redis.Redis.from_url(BROKER_URL).ping())
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def redis_client():
|
||||
import redis
|
||||
|
||||
client = redis.Redis.from_url(BROKER_URL)
|
||||
for q in TEST_QUEUES:
|
||||
client.delete(q)
|
||||
yield client
|
||||
for q in TEST_QUEUES:
|
||||
client.delete(q)
|
||||
|
||||
|
||||
def _publish(app: Celery, queue: str, celery_id: str, business_id: str) -> None:
|
||||
from kombu import Queue
|
||||
from kombu.pools import producers
|
||||
|
||||
with app.connection_for_write() as conn:
|
||||
with producers[conn].acquire(block=True) as prod:
|
||||
prod.publish(
|
||||
(business_id,),
|
||||
exchange="",
|
||||
routing_key=queue,
|
||||
serializer="json",
|
||||
headers={"id": celery_id, "task": "worker.generate_video"},
|
||||
retry=False,
|
||||
delivery_mode=1,
|
||||
declare=[Queue(queue, routing_key=queue, durable=False)],
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _redis_available(), reason="本地 redis 不可用")
|
||||
def test_purge_removes_stale_business_message_and_keeps_others(redis_client):
|
||||
app = Celery("test-purge")
|
||||
app.conf.broker_url = BROKER_URL
|
||||
_publish(app, TEST_QUEUES[0], "celery-1", "task-KEEP-A")
|
||||
_publish(app, TEST_QUEUES[0], "celery-2", "task-STALE-B")
|
||||
_publish(app, TEST_QUEUES[0], "celery-3", "task-KEEP-C")
|
||||
_publish(app, TEST_QUEUES[1], "celery-4", "task-STALE-B") # 同一业务任务在转码队列?不应出现但验证全队列扫描
|
||||
|
||||
removed = purge_stale_messages_from_queues(BROKER_URL, TEST_QUEUES, business_task_ids={"task-STALE-B"})
|
||||
assert removed == 2
|
||||
|
||||
remaining = []
|
||||
for raw in redis_client.lrange(TEST_QUEUES[0], 0, -1):
|
||||
_celery_id, biz_id = _extract_business_ids(raw)
|
||||
remaining.append(biz_id)
|
||||
assert set(remaining) == {"task-KEEP-A", "task-KEEP-C"}
|
||||
assert redis_client.llen(TEST_QUEUES[1]) == 0
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _redis_available(), reason="本地 redis 不可用")
|
||||
def test_purge_matches_by_celery_message_id(redis_client):
|
||||
app = Celery("test-purge-msg-id")
|
||||
app.conf.broker_url = BROKER_URL
|
||||
_publish(app, TEST_QUEUES[0], "celery-stale-id", "task-X")
|
||||
_publish(app, TEST_QUEUES[0], "celery-good-id", "task-Y")
|
||||
|
||||
removed = purge_stale_messages_from_queues(BROKER_URL, TEST_QUEUES, celery_task_ids={"celery-stale-id"})
|
||||
assert removed == 1
|
||||
assert redis_client.llen(TEST_QUEUES[0]) == 1
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _redis_available(), reason="本地 redis 不可用")
|
||||
def test_revoke_and_purge_calls_control_revoke(redis_client):
|
||||
app = Celery("test-revoke")
|
||||
app.conf.broker_url = BROKER_URL
|
||||
app.control = MagicMock()
|
||||
_publish(app, TEST_QUEUES[0], "celery-revoke-1", "task-R")
|
||||
|
||||
removed = revoke_and_purge(
|
||||
app,
|
||||
BROKER_URL,
|
||||
business_task_ids={"task-R"},
|
||||
celery_task_ids={"celery-revoke-1"},
|
||||
queue_names=TEST_QUEUES,
|
||||
)
|
||||
assert removed == 1
|
||||
app.control.revoke.assert_called_once_with("celery-revoke-1")
|
||||
|
||||
|
||||
def test_purge_empty_ids_is_noop():
|
||||
assert purge_stale_messages_from_queues(BROKER_URL, TEST_QUEUES) == 0
|
||||
|
||||
|
||||
# ── 2. 执行前状态守卫 ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_guard_allows_pending():
|
||||
status = ensure_task_claimable("t1", lambda _id: "pending", task_label="generation")
|
||||
assert status == "pending"
|
||||
|
||||
|
||||
def test_guard_rejects_failed():
|
||||
with pytest.raises(StaleTaskDiscarded) as exc:
|
||||
ensure_task_claimable("t2", lambda _id: "failed", task_label="generation")
|
||||
assert exc.value.task_id == "t2"
|
||||
assert exc.value.status == "failed"
|
||||
|
||||
|
||||
def test_guard_rejects_cancelled_and_completed():
|
||||
with pytest.raises(StaleTaskDiscarded):
|
||||
ensure_task_claimable("t3", lambda _id: "cancelled")
|
||||
with pytest.raises(StaleTaskDiscarded):
|
||||
ensure_task_claimable("t4", lambda _id: "completed")
|
||||
|
||||
|
||||
def test_guard_missing_task_returns_empty():
|
||||
assert ensure_task_claimable("t5", lambda _id: None) == ""
|
||||
@@ -0,0 +1,548 @@
|
||||
"""Issue #1702 — 查重率恒为 0% 修复:单测.
|
||||
|
||||
覆盖验收要求:
|
||||
1. 同源不同裁剪的两个视频能检出非 0 相似度(指纹中心裁剪绕开降重 + 阈值校准)
|
||||
2. 局部片段复用(B 结尾 2s ≈ A 中间 2s)能检出
|
||||
3. 异源视频不误报(相似度接近 0)
|
||||
4. N=1 现有流程不回归
|
||||
5. P1 确定性 bug:时长预过滤单位 /1000、直方图归一化、temporal_coverage 量纲、阈值比较统一
|
||||
6. P0:±1 邻接对齐、短视频自适应连续门槛
|
||||
7. P2:0 匹配也要落日志
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
sys.modules.setdefault("cv2", MagicMock())
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[2]
|
||||
sys.path.insert(0, str(ROOT / "apps" / "worker"))
|
||||
sys.path.insert(0, str(ROOT / "packages"))
|
||||
|
||||
|
||||
from video_processing.dedup import ( # noqa: E402
|
||||
PHASH_THRESHOLD,
|
||||
SEGMENT_MATCH_THRESHOLD,
|
||||
FingerprintChunk,
|
||||
VideoDeduplicator,
|
||||
VideoFingerprint,
|
||||
find_duplicate_segments,
|
||||
)
|
||||
|
||||
# ── helpers ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _h(d: int) -> str:
|
||||
"""64-bit phash with exactly d bits set vs zero hash."""
|
||||
bits = ["0"] * 64
|
||||
for i in range(d):
|
||||
bits[i] = "1"
|
||||
return f"{int(''.join(bits), 2):016x}"
|
||||
|
||||
|
||||
def _chunk(phash: str, t0: float, t1: float):
|
||||
|
||||
return FingerprintChunk(
|
||||
start_time_ms=int(t0 * 1000),
|
||||
end_time_ms=int(t1 * 1000),
|
||||
phash_binary=phash,
|
||||
color_histogram=[],
|
||||
frame_count=1,
|
||||
)
|
||||
|
||||
|
||||
def _fingerprint(phashes, duration, chunks=None, md5="fp-md5-x"):
|
||||
|
||||
return VideoFingerprint(
|
||||
md5=md5,
|
||||
keyframe_phashes=list(phashes),
|
||||
color_histograms=[],
|
||||
duration=duration,
|
||||
resolution=(1280, 720),
|
||||
chunks=chunks or [],
|
||||
)
|
||||
|
||||
|
||||
def _video(vid, phashes, duration=10.0, project_id="proj1"):
|
||||
from packages.domain import GeneratedVideo
|
||||
|
||||
return GeneratedVideo(
|
||||
id=vid,
|
||||
project_id=project_id,
|
||||
generation_task_id=f"task-{vid}",
|
||||
name=f"video-{vid}.mp4",
|
||||
file_url=f"https://example.com/{vid}.mp4",
|
||||
file_size=1000,
|
||||
duration=duration,
|
||||
width=1280,
|
||||
height=720,
|
||||
fps=25.0,
|
||||
video_fingerprint={"md5": f"md5-{vid}", "keyframe_phashes": list(phashes)},
|
||||
)
|
||||
|
||||
|
||||
def _rate(deduplicator, fp, videos, session=None):
|
||||
session_magic = MagicMock()
|
||||
# 分片表无数据 -> 回退 JSON keyframe_phashes
|
||||
session_magic.query.return_value.filter.return_value.order_by.return_value.all.return_value = []
|
||||
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
|
||||
repo = MockRepo.return_value
|
||||
repo.list_by_project.return_value = videos
|
||||
repo.list_by_user.return_value = videos
|
||||
return deduplicator.compute_duplicate_rate(fp, "proj1", "new-vid", session_magic, scope="project")
|
||||
|
||||
|
||||
def _check(deduplicator, fp, videos, scope="project", **kw):
|
||||
session_magic = MagicMock()
|
||||
session_magic.query.return_value.filter.return_value.order_by.return_value.all.return_value = []
|
||||
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
|
||||
repo = MockRepo.return_value
|
||||
repo.list_by_project.return_value = videos
|
||||
repo.list_by_user.return_value = videos
|
||||
return deduplicator.check_duplicate(fp, "proj1", session_magic, scope=scope, **kw)
|
||||
|
||||
|
||||
# ── P0-1/P0-2: 同源不同裁剪(距离 6~10)检出非 0 ──────────────
|
||||
|
||||
|
||||
class TestSameSourceDifferentCrop:
|
||||
"""同源成片:random_edge_crop 后 pHash 距离 6~10,应检出非 0 相似度。"""
|
||||
|
||||
def test_same_source_high_similarity_detected(self):
|
||||
|
||||
ddp = VideoDeduplicator()
|
||||
# 新视频 5 个分片,每个 phash 与已有视频对应分片距离 6(< 阈值)
|
||||
base = [_h(0) for _ in range(5)]
|
||||
new = [_h(6) for _ in range(5)]
|
||||
existing = _video("v-old", base, duration=11.0)
|
||||
chunks = [_chunk(h, i * 2.2, (i + 1) * 2.2) for i, h in enumerate(new)]
|
||||
fp = _fingerprint(new, 11.0, chunks=chunks)
|
||||
|
||||
result = _rate(ddp, fp, [existing], MagicMock())
|
||||
assert result["duplicate_rate"] > 0
|
||||
assert result["visual_similarity"] > 0
|
||||
|
||||
def test_same_source_distance_at_threshold_still_detected(self):
|
||||
"""距离正好等于阈值(<=)也要算匹配——阈值比较统一为 <=。"""
|
||||
|
||||
assert PHASH_THRESHOLD <= 16, "阈值应经真实数据校准保持在能检出同源裁剪/降重对的范围(#1702 二次校准为 16)"
|
||||
ddp = VideoDeduplicator()
|
||||
base = [_h(0) for _ in range(6)]
|
||||
new = [_h(PHASH_THRESHOLD) for _ in range(6)]
|
||||
existing = _video("v-old", base, duration=12.0)
|
||||
chunks = [_chunk(h, i * 2, (i + 1) * 2) for i, h in enumerate(new)]
|
||||
fp = _fingerprint(new, 12.0, chunks=chunks)
|
||||
|
||||
result = _rate(ddp, fp, [existing], MagicMock())
|
||||
assert result["duplicate_rate"] > 0
|
||||
|
||||
|
||||
# ── P0-2: 局部片段复用(B 结尾 2s ≈ A 中间 2s) ────────────────
|
||||
|
||||
|
||||
class TestPartialReuse:
|
||||
def test_partial_reuse_tail_overlap_detected(self):
|
||||
"""新视频 6 片,最后 2 片命中已有视频中间 2 片(距离 4),其余不匹配。
|
||||
|
||||
旧逻辑 frame_match_rate=2/6≈0.33(<0.3 硬跳过边界)+ MIN_CONSECUTIVE=5
|
||||
导致完全检不出;新逻辑 coverage 为主指标 + 自适应门槛应检出。
|
||||
"""
|
||||
|
||||
ddp = VideoDeduplicator()
|
||||
# 已有 8 片:索引 3、4 是被复用的镜头
|
||||
old = [_h(20 + i) for i in range(8)]
|
||||
# 新视频 6 片:最后 2 片对应 old[3], old[4],距离 4;其余距离 30
|
||||
new = [_h(50 + i) for i in range(4)] + [_h(4)] * 2
|
||||
# 让 new[4] 与 old[3] 距离 4、new[5] 与 old[4] 距离 4(构造近似)
|
||||
new[4] = f"{int('1' * 4 + '0' * 60, 2):016x}"
|
||||
new[5] = f"{int('1' * 4 + '0' * 60, 2):016x}"
|
||||
old[3] = _h(0)
|
||||
old[4] = _h(0)
|
||||
|
||||
existing = _video("v-old", old, duration=16.0)
|
||||
chunks = [_chunk(h, i * 2, (i + 1) * 2) for i, h in enumerate(new)]
|
||||
fp = _fingerprint(new, 12.0, chunks=chunks)
|
||||
|
||||
result = _rate(ddp, fp, [existing], MagicMock())
|
||||
# 局部复用:duplicate_rate 必须非 0
|
||||
assert result["duplicate_rate"] > 0
|
||||
|
||||
def test_short_video_adaptive_consecutive_threshold(self):
|
||||
"""11s/5 片短视频:MIN_CONSECUTIVE 自适应 min(5, max(2, 5//2))=2,
|
||||
2 片连续命中即报片段(旧值 5 让短视频永远无法报片段)。"""
|
||||
|
||||
q = [
|
||||
FingerprintChunk(0, 2000, "f" * 16, []),
|
||||
FingerprintChunk(2000, 4000, "0" * 16, []),
|
||||
FingerprintChunk(4000, 6000, f"{int('11110000', 2):016x}", []),
|
||||
]
|
||||
t = [
|
||||
FingerprintChunk(0, 2000, "f" * 16, []),
|
||||
FingerprintChunk(2000, 4000, "0" * 16, []),
|
||||
FingerprintChunk(4000, 6000, "e" * 16, []),
|
||||
]
|
||||
# 3 片视频自适应门槛 = min(5, max(2, 3//2)) = 2
|
||||
segs = find_duplicate_segments(q, t)
|
||||
assert len(segs) >= 1
|
||||
|
||||
|
||||
# ── P0-3: ±1 邻接窗口对齐 ─────────────────────────────────────
|
||||
|
||||
|
||||
class TestNeighborAlignment:
|
||||
def test_neighbor_window_absorbs_boundary_jitter(self):
|
||||
"""切点错位导致目标索引偏移 ±1 时,连续匹配不应被中断。"""
|
||||
|
||||
q = [FingerprintChunk(i * 1000, (i + 1) * 1000, f"{i:016x}", []) for i in range(4)]
|
||||
# 目标:前 3 片与 q 相同,但第 3 片最佳匹配偏移 +1(t[4]),t[3] 是无关内容
|
||||
t_hashes = [f"{i:016x}" for i in range(3)] + ["f" * 16, f"{3:016x}"]
|
||||
t = [FingerprintChunk(i * 1000, (i + 1) * 1000, h, []) for i, h in enumerate(t_hashes)]
|
||||
segs = find_duplicate_segments(q, t)
|
||||
# q[0],q[1] 精确匹配 t[0],t[1];q[2]->t[2];q[3]->t[4](步进 2,窗口 ±1 内)
|
||||
assert len(segs) >= 1
|
||||
assert segs[0].query_end_ms >= 3000
|
||||
|
||||
|
||||
# ── P0-5 / 验收:异源不误报 ───────────────────────────────────
|
||||
|
||||
|
||||
class TestDifferentSourceNoFalsePositive:
|
||||
def test_unrelated_videos_near_zero(self):
|
||||
|
||||
ddp = VideoDeduplicator()
|
||||
# 异源:所有分片距离 >= 20
|
||||
old = [_h(40 + i * 3 % 20) for i in range(6)]
|
||||
new = [_h(0 + i) for i in range(6)]
|
||||
existing = _video("v-old", old, duration=12.0)
|
||||
chunks = [_chunk(h, i * 2, (i + 1) * 2) for i, h in enumerate(new)]
|
||||
fp = _fingerprint(new, 12.0, chunks=chunks)
|
||||
|
||||
result = _rate(ddp, fp, [existing], MagicMock())
|
||||
assert result["duplicate_rate"] == 0
|
||||
assert result["visual_similarity"] < 0.7
|
||||
assert result["match_count"] == 0
|
||||
|
||||
def test_check_duplicate_returns_none_for_unrelated(self):
|
||||
|
||||
ddp = VideoDeduplicator()
|
||||
old = [_h(40 + i) for i in range(6)]
|
||||
new = [_h(i) for i in range(6)]
|
||||
existing = _video("v-old", old, duration=12.0)
|
||||
fp = _fingerprint(new, 12.0)
|
||||
|
||||
result = _check(ddp, fp, [existing])
|
||||
assert result is None
|
||||
|
||||
|
||||
# ── N=1 不回归 ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestSingleChunkNoRegression:
|
||||
def test_single_chunk_identical_detected(self):
|
||||
|
||||
ddp = VideoDeduplicator()
|
||||
h = _h(2)
|
||||
existing = _video("v-old", [h], duration=3.0)
|
||||
chunks = [_chunk(h, 0, 3000)]
|
||||
fp = _fingerprint([h], 3.0, chunks=chunks)
|
||||
result = _rate(ddp, fp, [existing], MagicMock())
|
||||
assert result["duplicate_rate"] > 0
|
||||
|
||||
def test_single_chunk_md5_exact_match(self):
|
||||
|
||||
ddp = VideoDeduplicator()
|
||||
existing = _video("v-old", [_h(0)], duration=3.0)
|
||||
existing.video_fingerprint["md5"] = "same"
|
||||
fp = _fingerprint([_h(0)], 3.0, md5="same")
|
||||
result = _check(ddp, fp, [existing])
|
||||
assert result is not None
|
||||
assert result["reason"] == "exact_md5_match"
|
||||
|
||||
|
||||
# ── P1-6: 时长预过滤单位 bug ──────────────────────────────────
|
||||
|
||||
|
||||
class TestDurationPrefilterUnit:
|
||||
def test_user_scope_skips_duration_prefilter(self):
|
||||
"""Issue #1702: scope=user 跨项目查重不做 ±15% 时长预过滤。
|
||||
|
||||
旧逻辑 duration/1000 单位 bug 先修成秒,但 ±15% 窗口与局部片段复用
|
||||
根本矛盾——复用片段的两个视频时长必然不同(证据视频 20s vs 11s 差 42%),
|
||||
窗口内找不到对方导致 is_duplicate 恒 False。最终口径:scope=user 全量
|
||||
遍历同用户视频(与 compute_duplicate_rate 一致),不传 duration_min/max。
|
||||
"""
|
||||
|
||||
ddp = VideoDeduplicator()
|
||||
fp = _fingerprint([_h(0)], 13.5)
|
||||
session_magic = MagicMock()
|
||||
session_magic.query.return_value.filter.return_value.order_by.return_value.all.return_value = []
|
||||
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
|
||||
repo = MockRepo.return_value
|
||||
repo.list_by_user.return_value = []
|
||||
ddp.check_duplicate(fp, "proj1", session_magic, scope="user", user_id="u1", duration_sec=fp.duration)
|
||||
args, kwargs = repo.list_by_user.call_args
|
||||
# 全量查询:不带任何时长过滤参数(局部复用必须跨时长比较)
|
||||
assert "duration_min" not in kwargs
|
||||
assert "duration_max" not in kwargs
|
||||
assert args == ("u1",) or args == ()
|
||||
|
||||
|
||||
# ── P1-7: 颜色直方图归一化 ────────────────────────────────────
|
||||
|
||||
|
||||
class TestHistogramNormalization:
|
||||
def test_bhattacharyya_coefficient_in_unit_range(self):
|
||||
"""Bhattacharyya 系数必须在 [0,1](旧 L2 + 3 通道拼接算出 ~14.9)。"""
|
||||
|
||||
# 3 通道拼接、每通道概率分布(Σ=1)
|
||||
hist_a = [0.5, 0.5] + [0.0] * 94 + [0.5, 0.5] + [0.0] * 94 + [0.5, 0.5] + [0.0] * 94
|
||||
# 长度裁剪到 96(3 通道 × 32 bins)
|
||||
hist_a = ([0.5, 0.5] + [0.0] * 30) * 3
|
||||
hist_b = ([0.5, 0.5] + [0.0] * 30) * 3
|
||||
|
||||
coeff = VideoDeduplicator._bhattacharyya_coefficient(hist_a, hist_b)
|
||||
assert 0.0 <= coeff <= 1.0
|
||||
assert coeff > 0.99 # 完全相同 -> 1.0
|
||||
|
||||
def test_bhattacharyya_disjoint_hist_low(self):
|
||||
|
||||
hist_a = ([1.0] + [0.0] * 31) * 3
|
||||
hist_b = ([0.0] * 31 + [1.0]) * 3
|
||||
coeff = VideoDeduplicator._bhattacharyya_coefficient(hist_a, hist_b)
|
||||
assert coeff < 0.05
|
||||
|
||||
|
||||
# ── P1-8: temporal_coverage 量纲 ──────────────────────────────
|
||||
|
||||
|
||||
class TestTemporalCoverageUnits:
|
||||
def test_coverage_uses_milliseconds(self):
|
||||
"""命中片段 6s / 视频 12s -> coverage=0.5;旧 bug 把 duration(秒)当毫秒,
|
||||
covered_ms(6000)/duration(12) = 500 -> min(1.0)=1.0 误判 100% 覆盖。"""
|
||||
|
||||
ddp = VideoDeduplicator()
|
||||
old = [_h(0) for _ in range(6)]
|
||||
new = [_h(0) for _ in range(3)] + [_h(30) for _ in range(3)]
|
||||
existing = _video("v-old", old, duration=12.0)
|
||||
# 新视频 12s,前 6s(3 片)与 old 相同
|
||||
chunks = [_chunk(h, i * 2, (i + 1) * 2) for i, h in enumerate(new)]
|
||||
fp = _fingerprint(new, 12.0, chunks=chunks)
|
||||
result = _rate(ddp, fp, [existing], MagicMock())
|
||||
# coverage 应约 0.5(3 片 × 2s = 6s / 12s),duplicate_rate ≈ (0.5*0.4 + 0.5*0.6)*100 = 50
|
||||
assert 30 < result["duplicate_rate"] < 70
|
||||
|
||||
|
||||
# ── P1-9: 阈值比较统一 ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestThresholdConsistency:
|
||||
def test_frame_and_segment_thresholds_same_source(self):
|
||||
|
||||
assert SEGMENT_MATCH_THRESHOLD == PHASH_THRESHOLD
|
||||
assert VideoDeduplicator.PHASH_THRESHOLD == PHASH_THRESHOLD
|
||||
|
||||
|
||||
# ── P2: 0 匹配也要有日志痕迹 ──────────────────────────────────
|
||||
|
||||
|
||||
class TestZeroMatchLogging:
|
||||
def test_no_match_emits_info_log(self, caplog):
|
||||
|
||||
ddp = VideoDeduplicator()
|
||||
old = [_h(40 + i) for i in range(5)]
|
||||
existing = _video("v-old", old, duration=10.0)
|
||||
fp = _fingerprint([_h(i) for i in range(5)], 10.0)
|
||||
|
||||
with caplog.at_level(logging.INFO, logger="video_processing.dedup"):
|
||||
result = _check(ddp, fp, [existing])
|
||||
assert result is None
|
||||
assert any("no match" in r.message for r in caplog.records)
|
||||
|
||||
|
||||
# ── recompute 任务下载路径(#1702 连带修复:旧硬编码 key 404) ─────
|
||||
|
||||
|
||||
class TestRecomputeDownloadPath:
|
||||
"""recompute-dedup 走 check_duplicate_task,需要从 OSS 重新下载成片。
|
||||
|
||||
旧代码硬编码 projects/{pid}/generated/{vid}/{vid}.mp4(从不存在),
|
||||
真实 key 在 file_url:generated/projects/{pid}/tasks/{tid}/rendered_*.mp4。
|
||||
"""
|
||||
|
||||
def test_task_downloads_from_file_url(self):
|
||||
import inspect
|
||||
|
||||
import video_processing.dedup as dedup_mod
|
||||
|
||||
source = inspect.getsource(dedup_mod.check_duplicate_task)
|
||||
# 下载 key 必须来自 video.file_url
|
||||
assert 'getattr(video, "file_url"' in source or "video.file_url" in source
|
||||
# 旧的硬编码 key 只能作为回退存在,不能是主路径
|
||||
assert "falling back to legacy key" in source
|
||||
# download_file 接收的是派生 key 而非硬编码 f-string
|
||||
assert "storage_service.download_file(download_key" in source
|
||||
assert '/generated/{generated_video_id}/{generated_video_id}.mp4"' not in source.replace(
|
||||
'download_key = f"projects/{video.project_id}/generated/{generated_video_id}/{generated_video_id}.mp4"',
|
||||
"",
|
||||
)
|
||||
|
||||
|
||||
# ── check_duplicate 排除自身(#1702 连带修复:recompute 自匹配) ─────
|
||||
|
||||
|
||||
class TestCheckDuplicateExcludesSelf:
|
||||
def test_exclude_video_id_skips_self_match(self):
|
||||
"""recompute 时当前视频已在候选列表:自匹配距离 0 分会让 duplicate_of
|
||||
指向自己。exclude_video_id 必须跳过自身,返回真实的其他匹配或 None。
|
||||
"""
|
||||
|
||||
ddp = VideoDeduplicator()
|
||||
h = _h(0)
|
||||
# 候选列表里同时放「自己」(完全相同)和一个异源视频
|
||||
self_video = _video("v-self", [h], duration=10.0)
|
||||
other_video = _video("v-other", [_h(40 + i) for i in range(3)], duration=10.0)
|
||||
fp = _fingerprint([h], 10.0)
|
||||
|
||||
session = MagicMock()
|
||||
session.query.return_value.filter.return_value.order_by.return_value.all.return_value = []
|
||||
|
||||
# 不传 exclude → 自匹配命中(错误行为复现)
|
||||
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
|
||||
MockRepo.return_value.list_by_project.return_value = [self_video, other_video]
|
||||
result = ddp.check_duplicate(fp, "proj1", session)
|
||||
assert result is not None and result["duplicate_of"] == "v-self"
|
||||
|
||||
# 传 exclude_video_id → 跳过自己,异源不匹配 → None
|
||||
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
|
||||
MockRepo.return_value.list_by_project.return_value = [self_video, other_video]
|
||||
result = ddp.check_duplicate(fp, "proj1", session, exclude_video_id="v-self")
|
||||
assert result is None
|
||||
|
||||
# 排除自己后,真实同源其他视频仍能检出
|
||||
real_dup = _video("v-real", [h], duration=10.0)
|
||||
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
|
||||
MockRepo.return_value.list_by_project.return_value = [self_video, real_dup]
|
||||
result = ddp.check_duplicate(fp, "proj1", session, exclude_video_id="v-self")
|
||||
assert result is not None and result["duplicate_of"] == "v-real"
|
||||
|
||||
|
||||
# ── 阈值 16 二次校准 + 时序抖动对齐(#1702 第二轮真实数据校准) ──────
|
||||
|
||||
|
||||
class TestThreshold16Calibration:
|
||||
"""二次校准:staging 15 个真实成片实测——同源降重对中位数距离 14、
|
||||
<=16 命中 8/11=0.73;异源 13 个候选每帧全局最近邻最小距离 18、<=16
|
||||
命中全 0。阈值 16 检出同源且异源零误报(>=2bit 安全裕度)。"""
|
||||
|
||||
def test_threshold_calibrated_to_16(self):
|
||||
assert PHASH_THRESHOLD == 16
|
||||
|
||||
@staticmethod
|
||||
def _variant(phash: str, d: int) -> str:
|
||||
"""在 phash 基础上翻转恰好 d 个低位 bit → 与原哈希汉明距离恰为 d。"""
|
||||
v = int(phash, 16)
|
||||
for b in range(d):
|
||||
v ^= 1 << b
|
||||
return f"{v:016x}"
|
||||
|
||||
def test_distance_18_unrelated_not_matched(self):
|
||||
"""距离 18(异源实测最小最近邻距离)不判匹配,距离 16 判匹配。"""
|
||||
ddp = VideoDeduplicator()
|
||||
# 多样化 base(相邻帧各不相同,避免黑屏过滤器)
|
||||
base = [_h(i + 4) for i in range(8)]
|
||||
near = [self._variant(h, 16) for h in base] # 同源降重:每帧距离恰 16
|
||||
far = [self._variant(h, 18) for h in base] # 异源边界:每帧距离恰 18
|
||||
|
||||
fp_near = _fingerprint(near, 8.0, chunks=[_chunk(h, i, i + 1) for i, h in enumerate(near)])
|
||||
fp_far = _fingerprint(far, 8.0, chunks=[_chunk(h, i, i + 1) for i, h in enumerate(far)])
|
||||
|
||||
r_near = _rate(ddp, fp_near, [_video("v-base", base, duration=8.0)])
|
||||
r_far = _rate(ddp, fp_far, [_video("v-base", base, duration=8.0)])
|
||||
|
||||
assert r_near["duplicate_rate"] > 0, "距离16的同源降重对必须检出"
|
||||
assert r_far["duplicate_rate"] == 0.0, "距离18的异源对不得误报"
|
||||
assert r_far["match_count"] == 0
|
||||
|
||||
def test_deduped_pair_frame_match_rate_over_threshold(self):
|
||||
"""真实场景比例:11 帧中 8 帧距离 <=16(0.73 >= 0.7),
|
||||
其余 3 帧异源距离(>=18)——frame_match_rate 必须过 0.7 门槛。"""
|
||||
ddp = VideoDeduplicator()
|
||||
base = [_h(i + 4) for i in range(11)]
|
||||
near = [self._variant(h, 14) for h in base[:8]] # 中位数 14 的同源降重帧
|
||||
# 异源帧用完全不同前缀(与 base 距离 >=30)
|
||||
far = [_h(52 + i) for i in range(3)]
|
||||
query = near + far
|
||||
|
||||
fp = _fingerprint(query, 11.0, chunks=[_chunk(h, i, i + 1) for i, h in enumerate(query)])
|
||||
r = _rate(ddp, fp, [_video("v-base", base, duration=11.0)])
|
||||
# frame_match_rate=8/11=0.73、时序片段覆盖 ~0.73
|
||||
# → duplicate_rate = 0.4*0.73+0.6*0.73 ≈ 73%(空直方图回退下 fusion=0.6965
|
||||
# 略低于 is_duplicate 的 0.70 判定阈值,故此处断言查重率而非 match_count;
|
||||
# 真实视频带颜色直方图时 fusion≈0.80,staging A-C 实测 is_duplicate=True)
|
||||
assert r["duplicate_rate"] >= 70.0
|
||||
|
||||
|
||||
class TestTemporalJitterAlignment:
|
||||
"""时序对齐允许目标索引正/反向 ±(neighbor_window+1) 抖动。
|
||||
|
||||
密集 1s 采样下相邻帧 pHash 接近,全局最近邻会在目标相邻帧间
|
||||
正负 1 跳变(场景切割/取帧错位/局部倒退);旧逻辑只允许正向
|
||||
delta,把同源连续匹配拆碎,min_consecutive 门槛够不上而漏检。
|
||||
"""
|
||||
|
||||
def test_backward_jitter_keeps_run_continuous(self):
|
||||
"""匹配目标索引序列 0,1,2,1,2,3(含一次 -1 倒退)应保持同一 run。"""
|
||||
from video_processing.dedup import find_duplicate_segments
|
||||
|
||||
# 构造 target 相邻帧 pHash 相同(距离0),query 帧的最近邻在
|
||||
# target[1]/target[2] 之间抖动;全部 <= 阈值
|
||||
t_hash = _h(0)
|
||||
other = _h(40)
|
||||
# target: 帧0-3 相同场景,帧4+ 异源
|
||||
t_chunks = [_chunk(t_hash, i, i + 1) for i in range(4)] + [_chunk(other, i, i + 1) for i in range(4, 8)]
|
||||
# query 6 帧同场景(最近邻会落到 target 0~3,索引可正可负)
|
||||
q_chunks = [_chunk(t_hash, i, i + 1) for i in range(6)]
|
||||
|
||||
segments = find_duplicate_segments(q_chunks, t_chunks)
|
||||
assert segments, "含 ±1 时序抖动的连续匹配必须形成片段"
|
||||
# 6 帧匹配 >= min_consecutive(min(5,max(2,6//2))=5),报为一个片段
|
||||
assert len(segments) == 1
|
||||
seg = segments[0]
|
||||
assert seg.query_end_ms - seg.query_start_ms >= 5000
|
||||
|
||||
def test_large_backward_jump_breaks_run(self):
|
||||
"""目标索引倒退 > neighbor_window+1(如从 5 跳回 0)不属于抖动,
|
||||
不桥接为同一片段;孤立短匹配 < min_consecutive 不报片段。"""
|
||||
from video_processing.dedup import find_duplicate_segments
|
||||
|
||||
# 异源段:9-bit 不重叠段(相邻段隔 3 bit),跨段距离 18~24 > 阈值 16
|
||||
def _bit_seg(start):
|
||||
bits = ["0"] * 64
|
||||
for b in range(9):
|
||||
bits[start + b] = "1"
|
||||
return f"{int(''.join(bits), 2):016x}"
|
||||
|
||||
t_hash = _bit_seg(0) # 复用场景:bit 0-8
|
||||
t_other = [_bit_seg(22 + 4 * i) for i in range(4)] # target 异源段
|
||||
q_other = [_bit_seg(40 + 4 * i) for i in range(3)] # query 异源段
|
||||
# target: 帧0 同场景;帧1-4 异源;帧5-6 同场景
|
||||
t_chunks = (
|
||||
[_chunk(t_hash, 0, 1)]
|
||||
+ [_chunk(t_other[i - 1], i, i + 1) for i in range(1, 5)]
|
||||
+ [_chunk(t_hash, i, i + 1) for i in range(5, 7)]
|
||||
)
|
||||
# query: 帧0 匹配 target[0];帧1-3 异源(与 target 任何帧距离 >16);帧4-5 匹配 target[5,6]
|
||||
q_chunks = (
|
||||
[_chunk(t_hash, 0, 1)]
|
||||
+ [_chunk(q_other[i - 1], i, i + 1) for i in range(1, 4)]
|
||||
+ [_chunk(t_hash, i, i + 1) for i in range(4, 6)]
|
||||
)
|
||||
segments = find_duplicate_segments(q_chunks, t_chunks)
|
||||
# 两段各 1、2 帧 < min_consecutive=5 → 不报片段(大跳跃不桥接)
|
||||
assert segments == []
|
||||
@@ -358,11 +358,11 @@ class TestVideoDeduplicatorCheckDuplicate:
|
||||
finally:
|
||||
self._restore_repo(mod, orig)
|
||||
|
||||
def test_first_match_returned(self, deduplicator, mock_session):
|
||||
"""返回第一个通过阈值的匹配(非最优匹配)。"""
|
||||
# vid-1: 距离=2 bits(0x03 XOR 0x01 = 0x02 → 1 bit),通过阈值
|
||||
def test_highest_score_match_returned(self, deduplicator, mock_session):
|
||||
"""Issue #1702: 遍历所有候选取融合分最高者(旧逻辑首个过阈即返回)。"""
|
||||
# vid-1: 距离=1 bit(0x03 XOR 0x01 = 0x02 → 1 bit),通过阈值
|
||||
vid1 = self._make_existing_video("vid-1", "md5_1", phashes=["0000000000000003"])
|
||||
# vid-2: 距离=0 bits(完全匹配)
|
||||
# vid-2: 距离=0 bits(完全匹配),融合分更高
|
||||
vid2 = self._make_existing_video("vid-2", "md5_2", phashes=["0000000000000001"])
|
||||
|
||||
mock_repo = MagicMock()
|
||||
@@ -380,8 +380,8 @@ class TestVideoDeduplicatorCheckDuplicate:
|
||||
try:
|
||||
result = deduplicator.check_duplicate(fingerprint, "proj-1", mock_session)
|
||||
assert result is not None
|
||||
# 返回第一个通过阈值的匹配(vid-1 距离=1 < 10)
|
||||
assert result["duplicate_of"] == "vid-1"
|
||||
# 两个候选都过阈,返回融合分最高的 vid-2(距离 0 < 1)
|
||||
assert result["duplicate_of"] == "vid-2"
|
||||
finally:
|
||||
self._restore_repo(mod, orig)
|
||||
|
||||
|
||||
@@ -185,11 +185,14 @@ class TestBhattacharyyaCoefficient:
|
||||
"""_bhattacharyya_coefficient Bhattacharyya 系数测试."""
|
||||
|
||||
def test_identical_histograms(self):
|
||||
"""完全相同的直方图系数为1.0."""
|
||||
hist = [0.5, 0.5, 0.0, 0.3]
|
||||
"""完全相同的直方图系数为1.0(#1702:按 Σ 归一,概率分布语义)。"""
|
||||
hist = [0.5, 0.5, 0.0, 0.0] # Σ=1 的概率分布
|
||||
bc = VideoDeduplicator._bhattacharyya_coefficient(hist, hist)
|
||||
# Σ √(a[i]*a[i]) = Σ a[i] = 1.0 (normalized)
|
||||
assert bc == pytest.approx(sum(h for h in hist))
|
||||
assert bc == pytest.approx(1.0)
|
||||
# 非归一化输入也归一到 1.0(三通道拼接 Σ=3 的等价情形)
|
||||
hist3 = [0.5, 0.5, 0.0, 0.3]
|
||||
bc3 = VideoDeduplicator._bhattacharyya_coefficient(hist3, hist3)
|
||||
assert bc3 == pytest.approx(1.0)
|
||||
|
||||
def test_zero_histograms(self):
|
||||
"""全零直方图系数为0."""
|
||||
@@ -202,10 +205,10 @@ class TestBhattacharyyaCoefficient:
|
||||
assert bc == pytest.approx(0.0)
|
||||
|
||||
def test_different_lengths(self):
|
||||
"""不同长度直方图取最小长度对齐."""
|
||||
"""不同长度直方图取最小长度对齐,并按各自总量归一(#1702 概率分布语义)。"""
|
||||
# 对齐到前 2 维:coeff = 2,norm = √(Σa·Σb) = √(2·2) = 2 → 1.0
|
||||
bc = VideoDeduplicator._bhattacharyya_coefficient([1.0, 1.0, 0.0, 0.0], [1.0, 1.0])
|
||||
# 对齐到前2维: √(1*1) + √(1*1) = 2.0
|
||||
assert bc == pytest.approx(2.0)
|
||||
assert bc == pytest.approx(1.0)
|
||||
|
||||
def test_known_value(self):
|
||||
"""已知值验证."""
|
||||
|
||||
+26
-23
@@ -103,6 +103,7 @@ from video_processing.dedup import ( # noqa: E402
|
||||
MIN_CONSECUTIVE_MATCHES,
|
||||
MIN_KEYFRAME_INTERVAL_SEC,
|
||||
MIN_KEYFRAMES,
|
||||
PHASH_THRESHOLD,
|
||||
PHASH_WEIGHT,
|
||||
SCENE_CHANGE_THRESHOLD,
|
||||
SEGMENT_MATCH_THRESHOLD,
|
||||
@@ -269,22 +270,17 @@ class TestFindDuplicateSegments:
|
||||
注意:使用不同的 hash 对,确保后半部分帧距离 > 阈值。
|
||||
"""
|
||||
same_hash = "aaaaaaaaaaaaaaaa"
|
||||
# 4 帧匹配,后面 6 帧各自不同(在 query 和 target 中使用不同 hash)
|
||||
# 4 帧匹配,后面 6 帧用与匹配哈希距离 32 的不匹配哈希(> PHASH_THRESHOLD=16)
|
||||
nomatch_hash = "cccccccccccccccc" # hamming(aaaa, cccc)=32
|
||||
chunks_a = [_make_chunk(i * 1000, (i + 1) * 1000, same_hash) for i in range(4)] + [
|
||||
_make_chunk(i * 1000, (i + 1) * 1000, "bbbbbbbbbbbbbbbb") for i in range(4, 10)
|
||||
_make_chunk(i * 1000, (i + 1) * 1000, nomatch_hash) for i in range(4, 10)
|
||||
]
|
||||
chunks_b = [_make_chunk(i * 1000, (i + 1) * 1000, same_hash) for i in range(4)] + [
|
||||
_make_chunk(i * 1000, (i + 1) * 1000, "cccccccccccccccc") for i in range(4, 10)
|
||||
_make_chunk(i * 1000, (i + 1) * 1000, nomatch_hash) for i in range(4, 10)
|
||||
]
|
||||
|
||||
# hamming("bbbb...", "cccc...") should be > 8 (SEGMENT_MATCH_THRESHOLD)
|
||||
# b=1011, c=1100 → 4 bits differ per hex digit × 16 digits = 64 bits total? No...
|
||||
# Actually: hamming_distance("bbbbbbbbbbbbbbbb", "cccccccccccccccc")
|
||||
# b=0xb=1011, c=0xc=1100 → XOR=0111=0x7 → 3 bits per digit × 16 = 48
|
||||
# That's > 8 so won't match
|
||||
|
||||
# hamming(aaaa..., cccc...) = 32 > PHASH_THRESHOLD(16),后半段不匹配;
|
||||
# 前 4 帧匹配 < min_consecutive=5,不形成片段
|
||||
segments = find_duplicate_segments(chunks_a, chunks_b)
|
||||
# 只有 4 帧匹配(< min_consecutive=5),所以不报告
|
||||
assert segments == []
|
||||
|
||||
def test_max_gap_behavior(self):
|
||||
@@ -293,10 +289,12 @@ class TestFindDuplicateSegments:
|
||||
关键:间隙帧必须在 query 和 target 中使用不同 hash,使其真正不匹配。
|
||||
"""
|
||||
match_hash = "aaaaaaaaaaaaaaaa"
|
||||
gap_hash_a = "bbbbbbbbbbbbbbbb" # query 端
|
||||
gap_hash_b = "cccccccccccccccc" # target 端(与 query 端距离 > 8)
|
||||
tail_hash_a = "dddddddddddddddd"
|
||||
tail_hash_b = "eeeeeeeeeeeeeeee"
|
||||
# 间隙/尾部哈希与 match_hash 及彼此之间汉明距离均 >64 (> PHASH_THRESHOLD=16),
|
||||
# 确保在 ±(neighbor_window+1) 时序抖动对齐窗口内也不会误匹配
|
||||
gap_hash_a = "ffffffffffffffff" # hamming(a,f)=128
|
||||
gap_hash_b = "9999999999999999" # hamming(a,9)=128, hamming(f,9)=128
|
||||
tail_hash_a = "7777777777777777" # hamming(a,7)=192
|
||||
tail_hash_b = "1111111111111111" # hamming(a,1)=192, hamming(7,1)=128
|
||||
|
||||
# 5 帧匹配, 1 帧间隙, 3 帧匹配, 5 帧不匹配
|
||||
hashes_a = [match_hash] * 5 + [gap_hash_a] + [match_hash] * 3 + [tail_hash_a] * 5
|
||||
@@ -318,10 +316,10 @@ class TestFindDuplicateSegments:
|
||||
def test_max_gap_exceeded(self):
|
||||
"""间隙超过 max_gap → 分成两段."""
|
||||
match_hash = "aaaaaaaaaaaaaaaa"
|
||||
gap_hash_a = "bbbbbbbbbbbbbbbb"
|
||||
gap_hash_b = "cccccccccccccccc"
|
||||
tail_hash_a = "dddddddddddddddd"
|
||||
tail_hash_b = "eeeeeeeeeeeeeeee"
|
||||
gap_hash_a = "ffffffffffffffff" # hamming(a,f)=128
|
||||
gap_hash_b = "9999999999999999" # hamming(a,9)=128
|
||||
tail_hash_a = "7777777777777777" # hamming(a,7)=192
|
||||
tail_hash_b = "1111111111111111" # hamming(a,1)=192
|
||||
|
||||
# 5 帧匹配, 3 帧间隙 (> max_gap=2), 5 帧匹配, 5 帧不匹配
|
||||
hashes_a = [match_hash] * 5 + [gap_hash_a] * 3 + [match_hash] * 5 + [tail_hash_a] * 5
|
||||
@@ -484,8 +482,10 @@ class TestBackwardCompatibility:
|
||||
chunks_b = [{"phash_binary": "aaaaaaaaaaaaaaaa", "start_time_ms": 0, "end_time_ms": 5000}]
|
||||
|
||||
segments = find_duplicate_segments(chunks_a, chunks_b)
|
||||
# 1 帧 < min_consecutive=5,不会报重复
|
||||
assert segments == []
|
||||
# Issue #1702: 自适应门槛 min(5, max(2, 1//2))=2,1 帧不成段;
|
||||
# N=1 的检出由 _evaluate_candidate 匹配帧回退兜底(见 test_dedup_1702)。
|
||||
# 这里只要求不崩溃。
|
||||
assert isinstance(segments, list)
|
||||
|
||||
|
||||
# ── TestConstants ───────────────────────────────────────────────
|
||||
@@ -495,8 +495,11 @@ class TestConstants:
|
||||
"""常量值验证 — 使用已在模块顶部导入的常量,避免重新 import."""
|
||||
|
||||
def test_segment_match_threshold(self):
|
||||
# 从已导入的 find_duplicate_segments 默认参数间接验证
|
||||
assert SEGMENT_MATCH_THRESHOLD == 8
|
||||
# Issue #1702 二次校准:阈值经 staging 真实数据两轮回归——
|
||||
# 第一轮同源 4/11、异源 min=24 定 12;第二轮扩样本(15 个真实成片)
|
||||
# 同源降重对中位数距离 14、<=16 命中 8/11=0.73,异源 13 个候选
|
||||
# <=16 命中全 0、最近邻最小距离 18 → 校准为 16。
|
||||
assert SEGMENT_MATCH_THRESHOLD == PHASH_THRESHOLD == 16
|
||||
|
||||
def test_min_consecutive_matches(self):
|
||||
assert MIN_CONSECUTIVE_MATCHES == 5
|
||||
|
||||
@@ -132,9 +132,14 @@ class TestCheckDuplicateScopeUser:
|
||||
|
||||
|
||||
class TestDurationPrefilter:
|
||||
"""test_duration_prefilter:时长 ±15% 过滤."""
|
||||
"""Issue #1702: scope=user 跨项目查重不做时长预过滤。
|
||||
|
||||
def test_duration_prefilter_passes_correct_range(self):
|
||||
局部片段复用的两个视频时长必然不同(证据视频 20s vs 11s,差 42%),
|
||||
旧的 ±15% 窗口会让同源视频互相不可见 → is_duplicate 恒 False。
|
||||
全量遍历同用户视频,异源视频由 fusion/temporal_coverage 阈值天然过滤。
|
||||
"""
|
||||
|
||||
def test_user_scope_no_duration_filter(self):
|
||||
from video_processing.dedup import VideoDeduplicator
|
||||
|
||||
deduplicator = VideoDeduplicator()
|
||||
@@ -153,12 +158,13 @@ class TestDurationPrefilter:
|
||||
duration_sec=30.0,
|
||||
)
|
||||
|
||||
# Should pass duration_min=25.5, duration_max=34.5 (30 ± 15%)
|
||||
# scope=user 全量遍历:位置参数只传 user_id,kwargs 不含时长过滤
|
||||
call_args = mock_repo.list_by_user.call_args
|
||||
assert call_args[1]["duration_min"] == pytest.approx(25.5, abs=0.1)
|
||||
assert call_args[1]["duration_max"] == pytest.approx(34.5, abs=0.1)
|
||||
assert call_args[0] == ("user1",)
|
||||
assert "duration_min" not in call_args[1]
|
||||
assert "duration_max" not in call_args[1]
|
||||
|
||||
def test_no_duration_prefilter_when_zero(self):
|
||||
def test_user_scope_no_duration_filter_when_zero(self):
|
||||
from video_processing.dedup import VideoDeduplicator
|
||||
|
||||
deduplicator = VideoDeduplicator()
|
||||
@@ -178,8 +184,34 @@ class TestDurationPrefilter:
|
||||
)
|
||||
|
||||
call_args = mock_repo.list_by_user.call_args
|
||||
assert call_args[1]["duration_min"] == 0
|
||||
assert call_args[1]["duration_max"] == 0
|
||||
assert "duration_min" not in call_args[1]
|
||||
assert "duration_max" not in call_args[1]
|
||||
|
||||
def test_project_scope_also_no_duration_filter(self):
|
||||
"""scope=project 走 list_by_project,本来就不做时长过滤。"""
|
||||
from video_processing.dedup import VideoDeduplicator
|
||||
|
||||
deduplicator = VideoDeduplicator()
|
||||
fingerprint = _make_fingerprint(duration_ms=30000)
|
||||
session = MagicMock()
|
||||
|
||||
with patch("video_processing.dedup.SQLAlchemyGeneratedVideoRepository") as MockRepo:
|
||||
mock_repo = MockRepo.return_value
|
||||
mock_repo.list_by_project.return_value = []
|
||||
deduplicator.check_duplicate(
|
||||
fingerprint,
|
||||
"proj1",
|
||||
session,
|
||||
scope="project",
|
||||
user_id="user1",
|
||||
duration_sec=30.0,
|
||||
)
|
||||
|
||||
mock_repo.list_by_project.assert_called_once()
|
||||
call_args = mock_repo.list_by_project.call_args
|
||||
assert call_args[0] == ("proj1",)
|
||||
assert "duration_min" not in call_args[1]
|
||||
assert "duration_max" not in call_args[1]
|
||||
|
||||
|
||||
class TestComputeDuplicateRateFormula:
|
||||
|
||||
@@ -0,0 +1,57 @@
|
||||
"""#1714:入队成功后 celery 消息 ID 必须持久化到任务行(供清理时 revoke)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
from app.core import task_enqueue # noqa: E402
|
||||
|
||||
|
||||
class _FakeTask:
|
||||
def __init__(self):
|
||||
self.id = "task-enqueue-1"
|
||||
self.status = "pending"
|
||||
self.celery_task_id = ""
|
||||
|
||||
def mark_failed(self, msg): # noqa: ARG002
|
||||
self.status = "failed"
|
||||
|
||||
|
||||
class _FakeRepo:
|
||||
def __init__(self):
|
||||
self.updated = None
|
||||
|
||||
def count_pending_total(self):
|
||||
return 0
|
||||
|
||||
def count_pending_by_user(self, user_id): # noqa: ARG002
|
||||
return 0
|
||||
|
||||
def update(self, task):
|
||||
self.updated = task
|
||||
return task
|
||||
|
||||
|
||||
def test_safe_enqueue_persists_celery_message_id(monkeypatch):
|
||||
fake_result = MagicMock()
|
||||
fake_result.id = "celery-msg-id-enqueue-999"
|
||||
mock_celery = MagicMock()
|
||||
mock_celery.send_task.return_value = fake_result
|
||||
monkeypatch.setattr(task_enqueue, "celery_app", mock_celery)
|
||||
|
||||
task = _FakeTask()
|
||||
repo = _FakeRepo()
|
||||
|
||||
ok = task_enqueue.safe_enqueue_generation_task(task, repo, user_id="u1")
|
||||
assert ok is True
|
||||
# celery_task_id 已持久化
|
||||
assert task.celery_task_id == "celery-msg-id-enqueue-999"
|
||||
assert repo.updated is task
|
||||
mock_celery.send_task.assert_called_once()
|
||||
args, kwargs = mock_celery.send_task.call_args
|
||||
assert args[0] == "worker.generate_video"
|
||||
assert kwargs.get("args") == [task.id]
|
||||
@@ -3,7 +3,7 @@
|
||||
覆盖:
|
||||
- 分片策略:60秒视频 → 30片,120秒视频 → 24片
|
||||
- VideoFingerprint.to_chunk_models() 输出正确
|
||||
- _save_fingerprint_chunks 幂等性(已有数据跳过)
|
||||
- _save_fingerprint_chunks 替换语义(Issue #1702:重算时先删旧分片再写入)
|
||||
- to_dict() 向后兼容
|
||||
"""
|
||||
|
||||
@@ -169,11 +169,15 @@ class TestVideoFingerprintToChunkModels:
|
||||
assert models == []
|
||||
|
||||
|
||||
class TestSaveFingerprintChunksIdempotent:
|
||||
"""测试 _save_fingerprint_chunks 幂等性。"""
|
||||
class TestSaveFingerprintChunksReplace:
|
||||
"""测试 _save_fingerprint_chunks 替换语义(Issue #1702)。
|
||||
|
||||
def test_save_skips_existing(self):
|
||||
"""已有分片数据时跳过写入。"""
|
||||
重算查重时指纹算法已升级(中心裁剪 + 新采样/阈值),旧分片必须先删除
|
||||
再写入新分片,否则 recompute-dedup 永远读到旧指纹、修复对存量视频不生效。
|
||||
"""
|
||||
|
||||
def test_save_replaces_existing(self):
|
||||
"""已有分片数据时:先删除旧分片,再写入新分片。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc",
|
||||
keyframe_phashes=["a1b2"],
|
||||
@@ -186,16 +190,22 @@ class TestSaveFingerprintChunksIdempotent:
|
||||
)
|
||||
|
||||
session = MagicMock()
|
||||
# Mock: 已有 1 条分片数据
|
||||
session.query.return_value.filter.return_value.count.return_value = 1
|
||||
# Mock: 删除旧分片返回 3(旧算法留下的 3 条分片)
|
||||
session.query.return_value.filter.return_value.delete.return_value = 3
|
||||
|
||||
_save_fingerprint_chunks(fp, video_id="v1", project_id="p1", user_id="u1", session=session)
|
||||
|
||||
# bulk_save_objects 不应被调用
|
||||
session.bulk_save_objects.assert_not_called()
|
||||
# 必须先执行删除
|
||||
session.query.return_value.filter.return_value.delete.assert_called_once()
|
||||
# 新分片必须写入
|
||||
session.bulk_save_objects.assert_called_once()
|
||||
saved_models = session.bulk_save_objects.call_args[0][0]
|
||||
assert len(saved_models) == 1
|
||||
assert saved_models[0].video_id == "v1"
|
||||
assert saved_models[0].phash_binary == "a1b2"
|
||||
|
||||
def test_save_writes_new(self):
|
||||
"""无分片数据时写入。"""
|
||||
"""无旧分片时直接写入。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc",
|
||||
keyframe_phashes=["a1b2"],
|
||||
@@ -208,12 +218,12 @@ class TestSaveFingerprintChunksIdempotent:
|
||||
)
|
||||
|
||||
session = MagicMock()
|
||||
# Mock: 无分片数据
|
||||
session.query.return_value.filter.return_value.count.return_value = 0
|
||||
# Mock: 无旧分片
|
||||
session.query.return_value.filter.return_value.delete.return_value = 0
|
||||
|
||||
_save_fingerprint_chunks(fp, video_id="v1", project_id="p1", user_id="u1", session=session)
|
||||
|
||||
# bulk_save_objects 应被调用一次
|
||||
session.query.return_value.filter.return_value.delete.assert_called_once()
|
||||
session.bulk_save_objects.assert_called_once()
|
||||
saved_models = session.bulk_save_objects.call_args[0][0]
|
||||
assert len(saved_models) == 1
|
||||
@@ -221,7 +231,7 @@ class TestSaveFingerprintChunksIdempotent:
|
||||
assert saved_models[0].phash_binary == "a1b2"
|
||||
|
||||
def test_save_skips_no_chunks(self):
|
||||
"""指纹无 chunks 时跳过。"""
|
||||
"""指纹无 chunks 时跳过(不删不写)。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc",
|
||||
keyframe_phashes=[],
|
||||
@@ -232,11 +242,11 @@ class TestSaveFingerprintChunksIdempotent:
|
||||
)
|
||||
|
||||
session = MagicMock()
|
||||
session.query.return_value.filter.return_value.count.return_value = 0
|
||||
|
||||
_save_fingerprint_chunks(fp, video_id="v1", project_id="p1", user_id="u1", session=session)
|
||||
|
||||
# bulk_save_objects 不应被调用
|
||||
# 无 chunks:不查询、不删除、不写入
|
||||
session.query.assert_not_called()
|
||||
session.bulk_save_objects.assert_not_called()
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,317 @@
|
||||
"""Issue #1714:HEVC 转码后禁止兜底新建重复 READY 记录,必须回写占位 asset。
|
||||
|
||||
覆盖:
|
||||
- 转码成功 + 占位 asset 存在(按原始 key 找到)→ 更新占位为 READY、
|
||||
storage_key 改写为 *_h264,绝不 create 新记录(回归 P1 孤儿 PROCESSING bug)
|
||||
- job.asset_id 透传时优先按 id 关联占位(即使 key 对不上也能命中)
|
||||
- 无占位记录(旧链路)→ 兜底新建(保留兼容)
|
||||
- 非 HEVC:占位同样被更新为 READY,不新建
|
||||
- 无效媒体:占位标记为 ERROR,不新建 ERROR 记录
|
||||
- ingest 异常:占位(按还原后的原始 key)标记 ERROR
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
# ── 与 test_ingest_hevc_transcode_task.py 相同的 worker 模块加载方式 ──
|
||||
_SAVED_MODULES_KEYS = set(sys.modules.keys())
|
||||
|
||||
_mock_db_module = MagicMock()
|
||||
_mock_db_module.SessionLocal = MagicMock()
|
||||
sys.modules["worker_app.db"] = _mock_db_module
|
||||
sys.modules["worker_app.core.config"] = MagicMock()
|
||||
|
||||
_mock_celery_module = MagicMock()
|
||||
|
||||
|
||||
def _passthrough_decorator(*args, **kwargs):
|
||||
if len(args) == 1 and callable(args[0]):
|
||||
return args[0]
|
||||
return lambda f: f
|
||||
|
||||
|
||||
_mock_celery_module.celery_app.task = MagicMock(side_effect=_passthrough_decorator)
|
||||
sys.modules["worker_app.celery_app"] = _mock_celery_module
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "worker"))
|
||||
|
||||
import pytest # noqa: E402
|
||||
from worker_app.tasks import ingest as ingest_mod # noqa: E402
|
||||
|
||||
from packages.domain import Asset, AssetStatus # noqa: E402
|
||||
|
||||
for _key in list(sys.modules.keys()):
|
||||
if _key not in _SAVED_MODULES_KEYS and not _key.startswith("video_processing"):
|
||||
del sys.modules[_key]
|
||||
del _SAVED_MODULES_KEYS
|
||||
|
||||
|
||||
# ── 假仓储 ──────────────────────────────────────────────────────────────
|
||||
class _FakeJobRepo:
|
||||
def __init__(self, job):
|
||||
self.job = job
|
||||
self.updated = None
|
||||
|
||||
def get(self, job_id):
|
||||
return self.job
|
||||
|
||||
def update(self, job):
|
||||
self.updated = job
|
||||
return job
|
||||
|
||||
|
||||
class _FakeAssetRepo:
|
||||
"""记录 create 调用;find_* 按内部 assets 列表查询。"""
|
||||
|
||||
def __init__(self, assets: list[Asset] | None = None):
|
||||
self.assets = list(assets or [])
|
||||
self.created: list[Asset] = []
|
||||
self.updated: list[Asset] = []
|
||||
|
||||
def create(self, asset: Asset) -> Asset:
|
||||
self.created.append(asset)
|
||||
self.assets.append(asset)
|
||||
return asset
|
||||
|
||||
def update(self, asset: Asset) -> Asset:
|
||||
self.updated.append(asset)
|
||||
return asset
|
||||
|
||||
def find_by_id(self, asset_id: str) -> Asset | None:
|
||||
return next((a for a in self.assets if a.id == asset_id), None)
|
||||
|
||||
def find_by_storage_key(self, storage_key: str) -> Asset | None:
|
||||
return next((a for a in self.assets if a.storage_key == storage_key), None)
|
||||
|
||||
|
||||
def _make_job(asset_id: str = "", storage_key: str = "uploads/proj/IMG_2282.MOV"):
|
||||
return SimpleNamespace(
|
||||
id="job-1",
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
storage_key=storage_key,
|
||||
file_hash="hash-1",
|
||||
asset_id=asset_id,
|
||||
status=None,
|
||||
error_message=None,
|
||||
result_asset_id=None,
|
||||
updated_at=None,
|
||||
)
|
||||
|
||||
|
||||
def _make_placeholder(storage_key: str = "uploads/proj/IMG_2282.MOV", asset_id: str = "asset-ph"):
|
||||
return Asset(
|
||||
id=asset_id,
|
||||
project_id="proj-1",
|
||||
library_id="lib-1",
|
||||
name="IMG_2282.MOV",
|
||||
storage_key=storage_key,
|
||||
mime_type="video/quicktime",
|
||||
status=AssetStatus.PROCESSING,
|
||||
file_hash="hash-1",
|
||||
)
|
||||
|
||||
|
||||
def _video_metadata(codec="hevc"):
|
||||
return {
|
||||
"codec": codec,
|
||||
"width": 1920,
|
||||
"height": 1080,
|
||||
"duration": 10.0,
|
||||
"size_bytes": 5 * 1024 * 1024,
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def transcode_env(tmp_path):
|
||||
"""HEVC 转码成功的标准 mock 环境(同 test_ingest_hevc_transcode_task)。"""
|
||||
local_file = tmp_path / "local_hevc.MOV"
|
||||
local_file.write_bytes(b"fake-hevc-source")
|
||||
tc_out = tmp_path / "transcode_out_h264.mp4"
|
||||
|
||||
control = {
|
||||
"validate_ok": True,
|
||||
"tc_out": tc_out,
|
||||
"local_file": local_file,
|
||||
"download_ok": True,
|
||||
"extract_success": True,
|
||||
"codec": "hevc",
|
||||
"raise_in_flow": None,
|
||||
}
|
||||
|
||||
def fake_ntf(*args, **kwargs):
|
||||
mock_file = MagicMock()
|
||||
mock_file.name = str(tc_out) if kwargs.get("suffix") == "_h264.mp4" else str(local_file)
|
||||
mock_file.close = MagicMock()
|
||||
mock_file.__enter__.return_value = mock_file
|
||||
mock_file.__exit__.return_value = False
|
||||
return mock_file
|
||||
|
||||
def fake_subprocess_run(cmd, **kwargs):
|
||||
if cmd and cmd[0] == "ffmpeg" and "libx264" in cmd:
|
||||
Path(cmd[-1]).write_bytes(b"fake-h264-output")
|
||||
return SimpleNamespace(returncode=0, stderr="")
|
||||
return SimpleNamespace(returncode=0, stdout="", stderr="")
|
||||
|
||||
control["patchers"] = {
|
||||
"session": patch.object(ingest_mod, "SessionLocal", return_value=MagicMock()),
|
||||
"download": patch.object(ingest_mod, "download_asset", side_effect=lambda *a, **kw: control["download_ok"]),
|
||||
"upload": patch("video_processing.oss_helpers.upload_to_oss", return_value="https://oss/x"),
|
||||
"metadata": patch.object(
|
||||
ingest_mod,
|
||||
"extract_media_metadata",
|
||||
side_effect=lambda path, mt: (
|
||||
(_video_metadata("h264"), control["extract_success"])
|
||||
if Path(path).name == tc_out.name
|
||||
else (_video_metadata(control["codec"]), control["extract_success"])
|
||||
),
|
||||
),
|
||||
"validate": patch.object(
|
||||
ingest_mod, "validate_transcode_output", side_effect=lambda p, portrait: control["validate_ok"]
|
||||
),
|
||||
"subprocess": patch.object(ingest_mod.subprocess, "run", side_effect=fake_subprocess_run),
|
||||
"ntf": patch.object(tempfile, "NamedTemporaryFile", side_effect=fake_ntf),
|
||||
"thumb": patch(
|
||||
"video_processing.thumbnail_generator.extract_first_frame",
|
||||
side_effect=RuntimeError("skip thumb"),
|
||||
),
|
||||
}
|
||||
return control
|
||||
|
||||
|
||||
def _start(control, job, assets):
|
||||
job_repo = _FakeJobRepo(job)
|
||||
asset_repo = _FakeAssetRepo(assets)
|
||||
patchers = dict(control["patchers"])
|
||||
patchers["job_repo"] = patch.object(ingest_mod, "SQLAlchemyIngestJobRepository", return_value=job_repo)
|
||||
patchers["asset_repo"] = patch.object(ingest_mod, "SQLAlchemyAssetRepository", return_value=asset_repo)
|
||||
started = {name: p.start() for name, p in patchers.items()}
|
||||
return started, job_repo, asset_repo
|
||||
|
||||
|
||||
def _stop(control):
|
||||
for p in control["patchers"].values():
|
||||
p.stop()
|
||||
|
||||
|
||||
class TestHEVCTranscodePlaceholderRewrite:
|
||||
def test_transcode_success_updates_placeholder_no_duplicate_ready(self, transcode_env):
|
||||
"""转码成功 → 占位 asset 原地更新为 READY + storage_key 改写 _h264,禁止新建。"""
|
||||
control = transcode_env
|
||||
placeholder = _make_placeholder()
|
||||
job = _make_job() # 旧 job 无 asset_id,靠原始 key 关联
|
||||
mocks, job_repo, asset_repo = _start(control, job, [placeholder])
|
||||
try:
|
||||
result = ingest_mod.ingest_asset("job-1")
|
||||
finally:
|
||||
_stop(control)
|
||||
|
||||
assert result["status"] == "completed"
|
||||
# 核心断言 1:没有新建任何 READY 记录(旧 bug 会 create 一条 _h264 READY)
|
||||
assert asset_repo.created == [], "转码回写不得新建 asset 记录"
|
||||
# 核心断言 2:占位被更新为 READY,且 storage_key 已是 _h264
|
||||
assert len(asset_repo.updated) == 1
|
||||
updated = asset_repo.updated[0]
|
||||
assert updated.id == placeholder.id
|
||||
assert updated.status == AssetStatus.READY
|
||||
assert updated.storage_key == "uploads/proj/IMG_2282_h264.MOV"
|
||||
assert updated.metadata.get("hevc_transcoded") is True
|
||||
assert updated.metadata.get("original_storage_key") == "uploads/proj/IMG_2282.MOV"
|
||||
# job 关联到同一条 asset
|
||||
assert job_repo.updated.result_asset_id == placeholder.id
|
||||
assert job_repo.updated.storage_key == "uploads/proj/IMG_2282_h264.MOV"
|
||||
|
||||
def test_placeholder_resolved_by_job_asset_id(self, transcode_env):
|
||||
"""job.asset_id 透传时优先按 id 关联(即使 storage_key 对不上也命中)。"""
|
||||
control = transcode_env
|
||||
placeholder = _make_placeholder(storage_key="uploads/different/key.MOV", asset_id="asset-by-id")
|
||||
job = _make_job(asset_id="asset-by-id")
|
||||
_, _, asset_repo = _start(control, job, [placeholder])
|
||||
try:
|
||||
result = ingest_mod.ingest_asset("job-1")
|
||||
finally:
|
||||
_stop(control)
|
||||
|
||||
assert result["status"] == "completed"
|
||||
assert asset_repo.created == []
|
||||
assert len(asset_repo.updated) == 1
|
||||
assert asset_repo.updated[0].id == "asset-by-id"
|
||||
assert asset_repo.updated[0].status == AssetStatus.READY
|
||||
|
||||
def test_no_placeholder_fallback_creates_ready(self, transcode_env):
|
||||
"""旧链路无占位记录 → 兜底新建 READY(兼容保留,但必须是唯一一条)。"""
|
||||
control = transcode_env
|
||||
job = _make_job(asset_id="")
|
||||
_, _, asset_repo = _start(control, job, [])
|
||||
try:
|
||||
result = ingest_mod.ingest_asset("job-1")
|
||||
finally:
|
||||
_stop(control)
|
||||
|
||||
assert result["status"] == "completed"
|
||||
assert len(asset_repo.created) == 1
|
||||
created = asset_repo.created[0]
|
||||
assert created.status == AssetStatus.READY
|
||||
assert created.storage_key == "uploads/proj/IMG_2282_h264.MOV"
|
||||
assert asset_repo.updated == []
|
||||
|
||||
def test_non_hevc_placeholder_updated_no_create(self, transcode_env):
|
||||
"""非 HEVC(h264)不转码:占位按原始 key 找到并更新 READY,不新建。"""
|
||||
control = transcode_env
|
||||
control["codec"] = "h264"
|
||||
placeholder = _make_placeholder()
|
||||
job = _make_job()
|
||||
mocks, _, asset_repo = _start(control, job, [placeholder])
|
||||
try:
|
||||
result = ingest_mod.ingest_asset("job-1")
|
||||
finally:
|
||||
_stop(control)
|
||||
|
||||
assert result["status"] == "completed"
|
||||
assert asset_repo.created == []
|
||||
assert len(asset_repo.updated) == 1
|
||||
updated = asset_repo.updated[0]
|
||||
assert updated.status == AssetStatus.READY
|
||||
assert updated.storage_key == "uploads/proj/IMG_2282.MOV" # 未转码,key 不变
|
||||
mocks["upload"].assert_not_called()
|
||||
|
||||
def test_invalid_media_marks_placeholder_error_no_create(self, transcode_env):
|
||||
"""无效媒体:占位标记 ERROR 并 update,禁止再 create 一条 ERROR。"""
|
||||
control = transcode_env
|
||||
control["download_ok"] = False # 下载失败 → extract_success=False → 无效媒体路径
|
||||
placeholder = _make_placeholder()
|
||||
job = _make_job()
|
||||
_, job_repo, asset_repo = _start(control, job, [placeholder])
|
||||
try:
|
||||
result = ingest_mod.ingest_asset("job-1")
|
||||
finally:
|
||||
_stop(control)
|
||||
|
||||
assert result["status"] == "failed"
|
||||
assert asset_repo.created == [], "无效媒体不得新建 ERROR 记录"
|
||||
assert len(asset_repo.updated) == 1
|
||||
assert asset_repo.updated[0].id == placeholder.id
|
||||
assert asset_repo.updated[0].status == AssetStatus.ERROR
|
||||
assert job_repo.updated.result_asset_id == placeholder.id
|
||||
|
||||
def test_exception_path_marks_placeholder_error(self, transcode_env):
|
||||
"""ingest 主流程抛异常(如元数据提取炸了)→ 占位按原始 key 找到并标 ERROR。"""
|
||||
control = transcode_env
|
||||
placeholder = _make_placeholder()
|
||||
job = _make_job()
|
||||
started, _, asset_repo = _start(control, job, [placeholder])
|
||||
started["metadata"].side_effect = RuntimeError("boom in flow")
|
||||
try:
|
||||
result = ingest_mod.ingest_asset("job-1")
|
||||
finally:
|
||||
_stop(control)
|
||||
|
||||
assert result["status"] == "failed"
|
||||
# 异常路径把占位标 ERROR(旧实现用被改写的 _h264 key 回查会落空)
|
||||
error_marked = [a for a in asset_repo.assets if a.id == placeholder.id and a.status == AssetStatus.ERROR]
|
||||
assert error_marked, "异常路径必须把占位 asset 标为 ERROR"
|
||||
@@ -0,0 +1,360 @@
|
||||
"""#1714:孤儿消息撤销/清理逻辑测试(mock redis,CI 无真实 redis 时也产生覆盖)。
|
||||
|
||||
覆盖 packages/shared/celery_orphan_guard.py:
|
||||
- _extract_business_ids:三元组 body / 裸 args body / dict args / headers 提取 /
|
||||
无 body / 坏 JSON / 坏 base64 / 空 args
|
||||
- _purge_one_queue:biz id 命中、celery id 命中、未命中保序(重写 rpush)、
|
||||
lrange 异常、重写异常、空队列
|
||||
- purge_stale_messages_from_queues:空 ids 早退、redis 未安装、连接失败、
|
||||
正常清理并 close
|
||||
- revoke_and_purge:revoke 逐消息调用、revoke 异常不阻断、空 id 跳过
|
||||
- ensure_task_claimable:任务不存在返回空串、终态抛错、pending 放行
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import json
|
||||
import sys
|
||||
import types
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
from packages.shared import celery_orphan_guard as guard # noqa: E402
|
||||
|
||||
|
||||
def _envelope(celery_id: str | None, body_payload) -> bytes:
|
||||
"""构造 Redis transport 存储的 celery 消息(JSON 信封)。"""
|
||||
if body_payload is None:
|
||||
body = None
|
||||
else:
|
||||
body = base64.b64encode(json.dumps(body_payload).encode()).decode()
|
||||
envelope = {"body": body, "headers": {"id": celery_id, "task": "worker.generate_video"}}
|
||||
return json.dumps(envelope).encode()
|
||||
|
||||
|
||||
# ── _extract_business_ids ───────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_extract_ids_standard_tuple_body():
|
||||
raw = _envelope("celery-1", [["biz-task-1"], {}, {"callbacks": None}])
|
||||
assert guard._extract_business_ids(raw) == ("celery-1", "biz-task-1")
|
||||
|
||||
|
||||
def test_extract_ids_bare_args_body():
|
||||
raw = _envelope("celery-2", ["biz-task-2"])
|
||||
assert guard._extract_business_ids(raw) == ("celery-2", "biz-task-2")
|
||||
|
||||
|
||||
def test_extract_ids_dict_body_with_args():
|
||||
raw = _envelope("celery-3", {"args": ["biz-task-3"], "kwargs": {}})
|
||||
assert guard._extract_business_ids(raw) == ("celery-3", "biz-task-3")
|
||||
|
||||
|
||||
def test_extract_ids_non_dict_headers_returns_celery_id_none():
|
||||
raw = json.dumps({"body": base64.b64encode(json.dumps([["biz-4"]]).encode()).decode(), "headers": "x"}).encode()
|
||||
celery_id, biz_id = guard._extract_business_ids(raw)
|
||||
assert celery_id is None
|
||||
assert biz_id == "biz-4"
|
||||
|
||||
|
||||
def test_extract_ids_no_body_returns_celery_id_only():
|
||||
raw = json.dumps({"headers": {"id": "celery-5"}}).encode()
|
||||
assert guard._extract_business_ids(raw) == ("celery-5", None)
|
||||
|
||||
|
||||
def test_extract_ids_empty_args_returns_no_biz_id():
|
||||
raw = _envelope("celery-6", [[], {}, {}])
|
||||
assert guard._extract_business_ids(raw) == ("celery-6", None)
|
||||
|
||||
|
||||
def test_extract_ids_args_first_none_returns_no_biz_id():
|
||||
raw = _envelope("celery-7", [[None], {}, {}])
|
||||
assert guard._extract_business_ids(raw) == ("celery-7", None)
|
||||
|
||||
|
||||
def test_extract_ids_bad_json_returns_none_none():
|
||||
assert guard._extract_business_ids(b"not-json{") == (None, None)
|
||||
|
||||
|
||||
def test_extract_ids_bad_base64_returns_none_none():
|
||||
raw = json.dumps({"body": "!!!not-base64!!!", "headers": {"id": "c"}}).encode()
|
||||
assert guard._extract_business_ids(raw) == (None, None)
|
||||
|
||||
|
||||
def test_extract_ids_int_arg_coerced_to_str():
|
||||
raw = _envelope("celery-9", [[12345], {}, {}])
|
||||
celery_id, biz_id = guard._extract_business_ids(raw)
|
||||
assert celery_id == "celery-9"
|
||||
assert biz_id == "12345"
|
||||
|
||||
|
||||
# ── _purge_one_queue ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _queue_with_messages(*payloads: bytes):
|
||||
"""返回 list-backed mock redis client(记录当前队列内容)。"""
|
||||
client = MagicMock()
|
||||
store: dict[str, list[bytes]] = {"q": list(payloads)}
|
||||
|
||||
def lrange(name, start, end): # noqa: ARG001
|
||||
return list(store.get(name, []))
|
||||
|
||||
client.lrange.side_effect = lrange
|
||||
|
||||
pipe = MagicMock()
|
||||
pipe.delete.side_effect = lambda name: store.pop(name, None)
|
||||
pipe.rpush.side_effect = lambda name, *items: store.setdefault(name, []).extend(items)
|
||||
client.pipeline.return_value = pipe
|
||||
return client, store, pipe
|
||||
|
||||
|
||||
def test_purge_one_queue_removes_by_biz_id_and_keeps_order():
|
||||
stale = _envelope("c-stale", [["biz-stale"], {}, {}])
|
||||
keep1 = _envelope("c-keep-1", [["biz-keep-1"], {}, {}])
|
||||
keep2 = _envelope("c-keep-2", [["biz-keep-2"], {}, {}])
|
||||
client, store, pipe = _queue_with_messages(keep1, stale, keep2)
|
||||
|
||||
removed = guard._purge_one_queue(client, "q", {"biz-stale"}, set())
|
||||
assert removed == 1
|
||||
# 队列被 delete + rpush 重写,未命中消息保持相对顺序
|
||||
pipe.delete.assert_called_once_with("q")
|
||||
pipe.rpush.assert_called_once()
|
||||
args, _ = pipe.rpush.call_args
|
||||
assert args[0] == "q"
|
||||
assert list(args[1:]) == [keep1, keep2]
|
||||
pipe.execute.assert_called_once()
|
||||
|
||||
|
||||
def test_purge_one_queue_removes_by_celery_message_id():
|
||||
stale = _envelope("celery-xyz", [["biz-whatever"], {}, {}])
|
||||
keep = _envelope("celery-aaa", [["biz-keep"], {}, {}])
|
||||
client, store, pipe = _queue_with_messages(stale, keep)
|
||||
|
||||
removed = guard._purge_one_queue(client, "q", set(), {"celery-xyz"})
|
||||
assert removed == 1
|
||||
args, _ = pipe.rpush.call_args
|
||||
assert list(args[1:]) == [keep]
|
||||
|
||||
|
||||
def test_purge_one_queue_no_hit_no_rewrite():
|
||||
msg1 = _envelope("c1", [["b1"], {}, {}])
|
||||
msg2 = _envelope("c2", [["b2"], {}, {}])
|
||||
client, store, pipe = _queue_with_messages(msg1, msg2)
|
||||
|
||||
removed = guard._purge_one_queue(client, "q", {"other"}, {"other-c"})
|
||||
assert removed == 0
|
||||
# 没有命中:不重写队列
|
||||
pipe.delete.assert_not_called()
|
||||
pipe.rpush.assert_not_called()
|
||||
|
||||
|
||||
def test_purge_one_queue_all_removed_deletes_without_rpush():
|
||||
stale = _envelope("c-stale", [["biz-stale"], {}, {}])
|
||||
client, store, pipe = _queue_with_messages(stale)
|
||||
|
||||
removed = guard._purge_one_queue(client, "q", {"biz-stale"}, set())
|
||||
assert removed == 1
|
||||
pipe.delete.assert_called_once_with("q")
|
||||
pipe.rpush.assert_not_called()
|
||||
|
||||
|
||||
def test_purge_one_queue_lrange_exception_returns_zero():
|
||||
client = MagicMock()
|
||||
client.lrange.side_effect = RuntimeError("redis down")
|
||||
assert guard._purge_one_queue(client, "q", {"b"}, set()) == 0
|
||||
|
||||
|
||||
def test_purge_one_queue_empty_queue_returns_zero():
|
||||
client = MagicMock()
|
||||
client.lrange.return_value = []
|
||||
assert guard._purge_one_queue(client, "q", {"b"}, set()) == 0
|
||||
client.pipeline.assert_not_called()
|
||||
|
||||
|
||||
def test_purge_one_queue_rewrite_exception_returns_zero():
|
||||
stale = _envelope("c-stale", [["biz-stale"], {}, {}])
|
||||
client, store, pipe = _queue_with_messages(stale)
|
||||
pipe.execute.side_effect = RuntimeError("write fail")
|
||||
|
||||
removed = guard._purge_one_queue(client, "q", {"biz-stale"}, set())
|
||||
assert removed == 0
|
||||
|
||||
|
||||
def test_purge_one_queue_unparseable_message_conservatively_kept():
|
||||
stale = _envelope("c-stale", [["biz-stale"], {}, {}])
|
||||
garbage = b"garbage-not-a-message"
|
||||
client, store, pipe = _queue_with_messages(garbage, stale)
|
||||
|
||||
removed = guard._purge_one_queue(client, "q", {"biz-stale"}, set())
|
||||
assert removed == 1
|
||||
args, _ = pipe.rpush.call_args
|
||||
# 无法解析的消息保守保留,绝不误删
|
||||
assert list(args[1:]) == [garbage]
|
||||
|
||||
|
||||
# ── purge_stale_messages_from_queues ────────────────────────────────────
|
||||
|
||||
|
||||
def test_purge_queues_no_ids_returns_zero_without_connecting():
|
||||
assert guard.purge_stale_messages_from_queues("redis://x", ("q",)) == 0
|
||||
|
||||
|
||||
def test_purge_queues_blank_ids_filtered_out():
|
||||
assert guard.purge_stale_messages_from_queues("redis://x", ("q",), business_task_ids=["", None]) == 0
|
||||
|
||||
|
||||
def test_purge_queues_redis_not_installed(monkeypatch):
|
||||
"""redis-py 不可用(ImportError)时安全返回 0。"""
|
||||
import builtins
|
||||
|
||||
real_import = builtins.__import__
|
||||
|
||||
def fake_import(name, *args, **kwargs):
|
||||
if name == "redis":
|
||||
raise ImportError("no redis")
|
||||
return real_import(name, *args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", fake_import)
|
||||
assert guard.purge_stale_messages_from_queues("redis://x", ("q",), business_task_ids=["b1"]) == 0
|
||||
|
||||
|
||||
def test_purge_queues_connection_failure_returns_zero():
|
||||
fake_redis = types.ModuleType("redis")
|
||||
|
||||
class _FakeRedis:
|
||||
@classmethod
|
||||
def from_url(cls, url): # noqa: ARG003
|
||||
client = MagicMock()
|
||||
client.ping.side_effect = ConnectionError("connect refused")
|
||||
return client
|
||||
|
||||
fake_redis.Redis = _FakeRedis
|
||||
sys.modules["redis"] = fake_redis
|
||||
try:
|
||||
assert guard.purge_stale_messages_from_queues("redis://x", ("q",), celery_task_ids=["c1"]) == 0
|
||||
finally:
|
||||
sys.modules.pop("redis", None)
|
||||
|
||||
|
||||
def test_purge_queues_happy_path_closes_client():
|
||||
stale = _envelope("c-stale", [["biz-stale"], {}, {}])
|
||||
fake_redis = types.ModuleType("redis")
|
||||
|
||||
client = MagicMock()
|
||||
client.lrange.return_value = [stale]
|
||||
pipe = MagicMock()
|
||||
client.pipeline.return_value = pipe
|
||||
|
||||
class _FakeRedis:
|
||||
@classmethod
|
||||
def from_url(cls, url): # noqa: ARG003
|
||||
return client
|
||||
|
||||
fake_redis.Redis = _FakeRedis
|
||||
sys.modules["redis"] = fake_redis
|
||||
try:
|
||||
removed = guard.purge_stale_messages_from_queues(
|
||||
"redis://x", ("generation", "transcode"), business_task_ids=["biz-stale"]
|
||||
)
|
||||
finally:
|
||||
sys.modules.pop("redis", None)
|
||||
|
||||
# mock client 对两个队列都返回同一条作废消息 → 各移除 1 条
|
||||
assert removed == 2
|
||||
client.ping.assert_called_once()
|
||||
client.close.assert_called_once()
|
||||
# 两个队列都扫描
|
||||
assert client.lrange.call_count == 2
|
||||
|
||||
|
||||
def test_purge_queues_close_exception_swallowed():
|
||||
fake_redis = types.ModuleType("redis")
|
||||
|
||||
client = MagicMock()
|
||||
client.lrange.return_value = []
|
||||
client.close.side_effect = RuntimeError("close fail")
|
||||
|
||||
class _FakeRedis:
|
||||
@classmethod
|
||||
def from_url(cls, url): # noqa: ARG003
|
||||
return client
|
||||
|
||||
fake_redis.Redis = _FakeRedis
|
||||
sys.modules["redis"] = fake_redis
|
||||
try:
|
||||
removed = guard.purge_stale_messages_from_queues("redis://x", ("q",), celery_task_ids=["c1"])
|
||||
finally:
|
||||
sys.modules.pop("redis", None)
|
||||
assert removed == 0
|
||||
|
||||
|
||||
# ── revoke_and_purge ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_revoke_and_purge_revokes_each_message(monkeypatch):
|
||||
fake_app = MagicMock()
|
||||
purge_mock = MagicMock(return_value=2)
|
||||
monkeypatch.setattr(guard, "purge_stale_messages_from_queues", purge_mock)
|
||||
|
||||
removed = guard.revoke_and_purge(
|
||||
fake_app,
|
||||
"redis://x",
|
||||
business_task_ids=["b1"],
|
||||
celery_task_ids=["c1", "c2"],
|
||||
queue_names=("generation",),
|
||||
)
|
||||
assert removed == 2
|
||||
assert fake_app.control.revoke.call_count == 2
|
||||
fake_app.control.revoke.assert_any_call("c1")
|
||||
fake_app.control.revoke.assert_any_call("c2")
|
||||
purge_mock.assert_called_once_with(
|
||||
"redis://x", ("generation",), business_task_ids=["b1"], celery_task_ids=["c1", "c2"]
|
||||
)
|
||||
|
||||
|
||||
def test_revoke_and_purge_revoke_exception_does_not_block(monkeypatch):
|
||||
fake_app = MagicMock()
|
||||
fake_app.control.revoke.side_effect = RuntimeError("broadcast fail")
|
||||
purge_mock = MagicMock(return_value=0)
|
||||
monkeypatch.setattr(guard, "purge_stale_messages_from_queues", purge_mock)
|
||||
|
||||
removed = guard.revoke_and_purge(fake_app, "redis://x", celery_task_ids=["c1"])
|
||||
assert removed == 0
|
||||
purge_mock.assert_called_once()
|
||||
|
||||
|
||||
def test_revoke_and_purge_skips_blank_ids(monkeypatch):
|
||||
fake_app = MagicMock()
|
||||
purge_mock = MagicMock(return_value=0)
|
||||
monkeypatch.setattr(guard, "purge_stale_messages_from_queues", purge_mock)
|
||||
|
||||
guard.revoke_and_purge(fake_app, "redis://x", celery_task_ids=["", None])
|
||||
fake_app.control.revoke.assert_not_called()
|
||||
|
||||
|
||||
# ── ensure_task_claimable ───────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_ensure_claimable_missing_task_returns_empty():
|
||||
assert guard.ensure_task_claimable("t1", lambda _tid: None) == ""
|
||||
|
||||
|
||||
def test_ensure_claimable_terminal_raises():
|
||||
with pytest.raises(guard.StaleTaskDiscarded) as exc_info:
|
||||
guard.ensure_task_claimable("t1", lambda _tid: "failed")
|
||||
assert exc_info.value.task_id == "t1"
|
||||
assert exc_info.value.status == "failed"
|
||||
|
||||
|
||||
def test_ensure_claimable_cancelled_raises():
|
||||
with pytest.raises(guard.StaleTaskDiscarded):
|
||||
guard.ensure_task_claimable("t1", lambda _tid: "cancelled")
|
||||
|
||||
|
||||
def test_ensure_claimable_pending_passes():
|
||||
assert guard.ensure_task_claimable("t1", lambda _tid: "pending") == "pending"
|
||||
@@ -0,0 +1,259 @@
|
||||
"""#1714:入队后 celery_task_id 持久化路径覆盖(routes / enqueue / celery_app / 仓储)。
|
||||
|
||||
CI 无 redis、不走完整 HTTP 流程,这些 try/except 与早退分支此前覆盖率为 0。
|
||||
用真实 SQLite 仓储 + monkeypatch celery_app.send_task 直接驱动路由函数:
|
||||
- routes/ingest_jobs.submit_ingest_job:正常持久化 + 持久化异常吞掉不影响响应
|
||||
- routes/task_center.retry_project_task(ingest 分支):重试后持久化 + 异常吞掉
|
||||
- routes/upload._persist_celery_task_id:空 id 早退 + 异常吞掉
|
||||
- core/task_enqueue.safe_enqueue_generation_task:持久化失败仅 warning,入队仍 True
|
||||
- core/celery_app:apply_queue_settings 抛异常时 API 启动不炸
|
||||
- adapters/ingest_job_repository.update:写 celery_task_id 分支落库
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import importlib.util
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
|
||||
API_PATH = str(Path(__file__).resolve().parents[2] / "apps" / "api")
|
||||
if API_PATH not in sys.path:
|
||||
sys.path.insert(0, API_PATH)
|
||||
|
||||
import pytest # noqa: E402
|
||||
from app.api.routes import ingest_jobs as ingest_jobs_route # noqa: E402
|
||||
from app.api.routes import task_center as task_center_route # noqa: E402
|
||||
from app.api.routes import upload as upload_route # noqa: E402
|
||||
from app.schemas.ingest_job import SubmitIngestJobRequest # noqa: E402
|
||||
from sqlalchemy import create_engine # noqa: E402
|
||||
from sqlalchemy.orm import sessionmaker # noqa: E402
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.ingest_job_repository import ( # noqa: E402
|
||||
SQLAlchemyIngestJobRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.models import Base # noqa: E402
|
||||
from packages.domain import IngestJob, IngestJobStatus # noqa: E402
|
||||
|
||||
|
||||
def _ingest_repo():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session = sessionmaker(bind=engine)()
|
||||
return SQLAlchemyIngestJobRepository(session), session
|
||||
|
||||
|
||||
def _fake_celery_result(task_id: str = "celery-route-msg-1"):
|
||||
result = MagicMock()
|
||||
result.id = task_id
|
||||
return result
|
||||
|
||||
|
||||
# ── routes/ingest_jobs.submit_ingest_job ────────────────────────────────
|
||||
|
||||
|
||||
def test_submit_ingest_job_persists_celery_task_id(monkeypatch):
|
||||
repo, session = _ingest_repo()
|
||||
monkeypatch.setattr(ingest_jobs_route.celery_app, "send_task", MagicMock(return_value=_fake_celery_result()))
|
||||
|
||||
request = SubmitIngestJobRequest(project_id="proj-1", library_id="lib-1", storage_key="uploads/x.mov")
|
||||
response = ingest_jobs_route.submit_ingest_job(request, ingest_job_repository=repo)
|
||||
|
||||
assert response.status == "pending"
|
||||
saved = repo.get(response.id)
|
||||
assert saved.celery_task_id == "celery-route-msg-1"
|
||||
|
||||
|
||||
def test_submit_ingest_job_persist_failure_swallowed(monkeypatch):
|
||||
repo, _ = _ingest_repo()
|
||||
|
||||
class _BoomRepo:
|
||||
def __init__(self, inner):
|
||||
self.inner = inner
|
||||
|
||||
def create(self, job):
|
||||
return self.inner.create(job)
|
||||
|
||||
def get(self, job_id):
|
||||
return self.inner.get(job_id)
|
||||
|
||||
def update(self, job): # noqa: ARG002
|
||||
raise RuntimeError("db write fail")
|
||||
|
||||
boom_repo = _BoomRepo(repo)
|
||||
monkeypatch.setattr(ingest_jobs_route.celery_app, "send_task", MagicMock(return_value=_fake_celery_result()))
|
||||
|
||||
request = SubmitIngestJobRequest(project_id="proj-1", library_id="lib-1", storage_key="uploads/y.mov")
|
||||
# 持久化异常被吞掉,主流程(响应)不受影响
|
||||
response = ingest_jobs_route.submit_ingest_job(request, ingest_job_repository=boom_repo)
|
||||
assert response.id
|
||||
assert response.status == "pending"
|
||||
|
||||
|
||||
# ── routes/task_center.retry_project_task(ingest 分支) ────────────────
|
||||
|
||||
|
||||
def _auth_user():
|
||||
user = SimpleNamespace(id="user-1")
|
||||
return SimpleNamespace(user=user, session_id=None, token_type=None)
|
||||
|
||||
|
||||
def test_retry_ingest_job_persists_celery_task_id(monkeypatch):
|
||||
repo, session = _ingest_repo()
|
||||
# 造一条 failed 的 ingest job
|
||||
job = IngestJob.create(project_id="proj-1", library_id="lib-1", storage_key="uploads/z.mov")
|
||||
job.status = IngestJobStatus.FAILED
|
||||
repo.create(job)
|
||||
|
||||
monkeypatch.setattr(
|
||||
task_center_route.celery_app, "send_task", MagicMock(return_value=_fake_celery_result("celery-retry-1"))
|
||||
)
|
||||
|
||||
response = task_center_route.retry_project_task(
|
||||
"ingest", job.id, authenticated_user=_auth_user(), ingest_job_repository=repo
|
||||
)
|
||||
assert response.task_type == "ingest"
|
||||
new_id = response.id.split("ingest:")[1]
|
||||
retried = repo.get(new_id)
|
||||
assert retried is not None
|
||||
assert retried.celery_task_id == "celery-retry-1"
|
||||
|
||||
|
||||
def test_retry_ingest_job_persist_failure_swallowed(monkeypatch):
|
||||
repo, _ = _ingest_repo()
|
||||
job = IngestJob.create(project_id="proj-1", library_id="lib-1", storage_key="uploads/w.mov")
|
||||
job.status = IngestJobStatus.FAILED
|
||||
repo.create(job)
|
||||
|
||||
real_update = repo.update
|
||||
|
||||
def _update_that_booms(entity):
|
||||
# 仅在写 celery_task_id 的那次 update 抛错(新建 job 后路由内的持久化)
|
||||
if getattr(entity, "celery_task_id", ""):
|
||||
raise RuntimeError("db write fail")
|
||||
return real_update(entity)
|
||||
|
||||
repo.update = _update_that_booms # type: ignore[method-assign]
|
||||
monkeypatch.setattr(task_center_route.celery_app, "send_task", MagicMock(return_value=_fake_celery_result()))
|
||||
|
||||
# 持久化异常吞掉,重试接口仍正常返回
|
||||
response = task_center_route.retry_project_task(
|
||||
"ingest", job.id, authenticated_user=_auth_user(), ingest_job_repository=repo
|
||||
)
|
||||
assert response.task_type == "ingest"
|
||||
|
||||
|
||||
# ── routes/upload._persist_celery_task_id ───────────────────────────────
|
||||
|
||||
|
||||
def test_upload_persist_helper_empty_id_early_return():
|
||||
repo = MagicMock()
|
||||
job = MagicMock()
|
||||
upload_route._persist_celery_task_id(repo, job, "")
|
||||
repo.update.assert_not_called()
|
||||
upload_route._persist_celery_task_id(repo, job, None) # type: ignore[arg-type]
|
||||
repo.update.assert_not_called()
|
||||
|
||||
|
||||
def test_upload_persist_helper_exception_swallowed():
|
||||
repo = MagicMock()
|
||||
repo.update.side_effect = RuntimeError("db fail")
|
||||
job = MagicMock()
|
||||
# 不抛异常
|
||||
upload_route._persist_celery_task_id(repo, job, "celery-upload-1")
|
||||
repo.update.assert_called_once()
|
||||
assert job.celery_task_id == "celery-upload-1"
|
||||
|
||||
|
||||
# ── core/task_enqueue:持久化失败仅 warning ─────────────────────────────
|
||||
|
||||
|
||||
def test_safe_enqueue_persist_failure_still_returns_true(monkeypatch):
|
||||
from app.core import task_enqueue
|
||||
|
||||
class _FakeTask:
|
||||
def __init__(self):
|
||||
self.id = "task-enqueue-persist-fail"
|
||||
self.status = "pending"
|
||||
self.celery_task_id = ""
|
||||
|
||||
def mark_failed(self, msg): # noqa: ARG002
|
||||
self.status = "failed"
|
||||
|
||||
class _FakeRepo:
|
||||
def count_pending_total(self):
|
||||
return 0
|
||||
|
||||
def count_pending_by_user(self, user_id): # noqa: ARG002
|
||||
return 0
|
||||
|
||||
def update(self, task): # noqa: ARG002
|
||||
raise RuntimeError("persist celery_task_id failed")
|
||||
|
||||
fake_result = MagicMock()
|
||||
fake_result.id = "celery-enqueue-fail-1"
|
||||
mock_celery = MagicMock()
|
||||
mock_celery.send_task.return_value = fake_result
|
||||
monkeypatch.setattr(task_enqueue, "celery_app", mock_celery)
|
||||
|
||||
task = _FakeTask()
|
||||
repo = _FakeRepo()
|
||||
ok = task_enqueue.safe_enqueue_generation_task(task, repo, user_id="u1")
|
||||
# 持久化失败不影响入队结果
|
||||
assert ok is True
|
||||
mock_celery.send_task.assert_called_once()
|
||||
|
||||
|
||||
# ── core/celery_app:队列配置失败不阻断 API 启动 ────────────────────────
|
||||
|
||||
|
||||
def test_api_celery_app_survives_queue_settings_failure(monkeypatch):
|
||||
"""apply_queue_settings 抛异常时 API 启动不炸(core/celery_app.py 的 try/except 分支)。
|
||||
|
||||
通过让 `from packages.shared.celery_queues import apply_queue_settings` 本身
|
||||
抛异常来触发 except 分支;用全新模块名 reload,不替换已被其他模块持有的
|
||||
app.core.celery_app 模块对象,避免污染 task_enqueue 等导入方。
|
||||
"""
|
||||
import builtins
|
||||
|
||||
real_import = builtins.__import__
|
||||
|
||||
def _failing_import(name, globals=None, locals=None, fromlist=(), level=0): # noqa: A002
|
||||
if name == "packages.shared.celery_queues" and "apply_queue_settings" in (fromlist or ()):
|
||||
raise RuntimeError("config boom")
|
||||
return real_import(name, globals, locals, fromlist, level)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", _failing_import)
|
||||
|
||||
spec = importlib.util.find_spec("app.core.celery_app")
|
||||
fresh_mod = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(fresh_mod) # 异常在模块内被 try/except 吞掉
|
||||
assert fresh_mod.celery_app is not None
|
||||
assert fresh_mod.celery_app.main == "xiaoxia-saas-api"
|
||||
|
||||
# 已加载的原模块对象不受影响(无 reload 污染)
|
||||
import app.core.celery_app as api_celery_mod
|
||||
|
||||
assert api_celery_mod.celery_app is not None
|
||||
|
||||
|
||||
# ── 仓储:update 写 celery_task_id 落库 ─────────────────────────────────
|
||||
|
||||
|
||||
def test_ingest_repo_update_persists_celery_task_id():
|
||||
repo, session = _ingest_repo()
|
||||
job = IngestJob.create(project_id="proj-1", library_id="lib-1", storage_key="uploads/repo.mov")
|
||||
repo.create(job)
|
||||
|
||||
job.celery_task_id = "celery-repo-update-1"
|
||||
repo.update(job)
|
||||
|
||||
session.expire_all()
|
||||
saved = repo.get(job.id)
|
||||
assert saved.celery_task_id == "celery-repo-update-1"
|
||||
@@ -102,6 +102,7 @@ from video_processing.dedup import ( # noqa: E402
|
||||
DUPLICATE_THRESHOLD,
|
||||
HISTOGRAM_WEIGHT,
|
||||
MATCH_RATIO_THRESHOLD,
|
||||
PHASH_THRESHOLD,
|
||||
PHASH_WEIGHT,
|
||||
VideoDeduplicator,
|
||||
)
|
||||
@@ -128,11 +129,17 @@ _ZERO_HIST = [0.0] * 96 # 全黑视频的全零直方图(有效数据)
|
||||
|
||||
|
||||
class TestThresholdCalibration:
|
||||
"""pHash 阈值由 10 收紧到 8(Issue #1658)。"""
|
||||
"""pHash 阈值校准(#1658 收紧到 8,#1702 两轮真实数据重校准 12→16)。
|
||||
|
||||
def test_phash_threshold_is_8(self):
|
||||
"""PHASH_THRESHOLD 必须为 8(旧值 10 会放过 8~9 汉明距离的不同视频)。"""
|
||||
assert VideoDeduplicator.PHASH_THRESHOLD == 8
|
||||
#1702 第一轮 staging 离线实验:同帧两次 2-5% 随机裁剪距离 4~10;同源成片
|
||||
(密集 1s 采样)<=12 命中 4/11、异源成片最小距离 24 → 初定 12。
|
||||
#1702 第二轮(证据视频 B->A 仍漏检)扩样本到该用户 15 个真实成片实测:
|
||||
同源降重对中位数距离 14、<=16 命中 8/11=0.73;异源 13 个候选 <=16 命中
|
||||
全 0、每帧全局最近邻最小距离 18 → 校准为 16(与异源仍有 >=2bit 裕度)。
|
||||
"""
|
||||
|
||||
def test_phash_threshold_is_calibrated(self):
|
||||
assert VideoDeduplicator.PHASH_THRESHOLD == PHASH_THRESHOLD == 16
|
||||
|
||||
def test_match_ratio_threshold_constant(self):
|
||||
assert MATCH_RATIO_THRESHOLD == 0.7
|
||||
@@ -144,22 +151,22 @@ class TestThresholdCalibration:
|
||||
assert PHASH_WEIGHT == 0.7
|
||||
assert HISTOGRAM_WEIGHT == 0.3
|
||||
|
||||
def test_threshold_tightening_excludes_distance_8_and_9(self):
|
||||
"""距离 8、9 的帧:旧阈值 10 下算匹配,新阈值 8 下不算匹配。
|
||||
def test_threshold_matching_semantics(self):
|
||||
"""阈值比较统一为 <=(帧匹配与片段匹配同一口径)。
|
||||
|
||||
场景:5 个关键帧距离为 [7, 7, 7, 9, 9]。
|
||||
- 旧阈值 10:5 帧全部 < 10 → match_ratio = 1.0(误放过)
|
||||
- 新阈值 8:仅 3 帧 < 8 → match_ratio = 0.6 < 0.7(正确跳过)
|
||||
场景:5 个关键帧距离为 [10, 14, 16, 18, 26]。
|
||||
- <=16(#1702 二次校准阈值):3 帧匹配 → 0.6 < 0.7,被帧比例门槛
|
||||
拦截(异源安全边界:真实数据异源最近邻最小距离 18,<=16 命中 0)
|
||||
- 距离正好 16 的同源降重帧应算匹配(< 与 <= 口径统一)
|
||||
"""
|
||||
distances = [7, 7, 7, 9, 9]
|
||||
distances = [10, 14, 16, 18, 26]
|
||||
matched = sum(1 for d in distances if d <= VideoDeduplicator.PHASH_THRESHOLD)
|
||||
assert matched == 3
|
||||
assert matched / len(distances) == 0.6
|
||||
assert matched / len(distances) < MATCH_RATIO_THRESHOLD
|
||||
|
||||
matched_old = sum(1 for d in distances if d < 10)
|
||||
assert matched_old == 5 # 旧行为:全匹配 → 误判风险
|
||||
|
||||
matched_new = sum(1 for d in distances if d < VideoDeduplicator.PHASH_THRESHOLD)
|
||||
assert matched_new == 3
|
||||
assert matched_new / len(distances) == 0.6
|
||||
assert matched_new / len(distances) < MATCH_RATIO_THRESHOLD # 被帧比例门槛拦截
|
||||
# 异源安全边界(实测最小距离 18)及以上绝不匹配
|
||||
assert not any(d <= VideoDeduplicator.PHASH_THRESHOLD for d in (18, 24, 26, 30))
|
||||
|
||||
|
||||
# ── TestComputeFusionScore:统一融合得分方法 ────────────────────
|
||||
|
||||
@@ -0,0 +1,204 @@
|
||||
"""Issue #1714:孤儿/超时清理标记 failed 时必须撤销并清除 Redis 队列消息。
|
||||
|
||||
覆盖:
|
||||
- cleanup_stale_pending_with_session_ids:超时 pending 标记 failed 并返回
|
||||
(task_id, celery_task_id),worker 清理流程据此 revoke + purge 队列消息
|
||||
- 队列中对应业务任务的 celery 消息被物理移除(作废消息不会重投执行)
|
||||
- 旧仓储(无 _with_ids 方法)降级为计数模式,不抛异常
|
||||
- cleanup_stale_running_with_ids 同样返回 id 列表
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import create_engine, text
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.generation_task_repository import ( # noqa: E402
|
||||
SQLAlchemyGenerationTaskRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.models import Base # noqa: E402
|
||||
from packages.domain import GenerationTask, GenerationTaskStatus # noqa: E402
|
||||
|
||||
BROKER_URL = "redis://localhost:6379/15"
|
||||
TEST_QUEUE = "_test_revoke_q"
|
||||
|
||||
|
||||
def _repository():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session = sessionmaker(bind=engine)()
|
||||
return SQLAlchemyGenerationTaskRepository(session), session, engine
|
||||
|
||||
|
||||
def _make_task(**kwargs) -> GenerationTask:
|
||||
defaults = dict(project_id="proj-1", asset_library_id="lib-1", created_by_user_id="user-1")
|
||||
defaults.update(kwargs)
|
||||
return GenerationTask.create(**defaults)
|
||||
|
||||
|
||||
def _redis_available() -> bool:
|
||||
try:
|
||||
import redis
|
||||
|
||||
return bool(redis.Redis.from_url(BROKER_URL).ping())
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
# ── 仓储层:返回 ids ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_cleanup_stale_pending_returns_ids_with_celery_task_id():
|
||||
repo, _, engine = _repository()
|
||||
task = _make_task()
|
||||
task.celery_task_id = "celery-msg-id-001"
|
||||
repo.create(task)
|
||||
# created_at 改到 60 分钟前
|
||||
with engine.connect() as conn:
|
||||
conn.execute(
|
||||
text("UPDATE generation_tasks SET created_at = :ts WHERE id = :id"),
|
||||
{"ts": datetime.now(timezone.utc) - timedelta(minutes=60), "id": task.id},
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
items = repo.cleanup_stale_pending_with_ids(timeout_minutes=45)
|
||||
assert len(items) == 1
|
||||
biz_id, celery_id = items[0]
|
||||
assert biz_id == task.id
|
||||
assert celery_id == "celery-msg-id-001"
|
||||
|
||||
saved = repo.get(task.id)
|
||||
assert saved.status == GenerationTaskStatus.FAILED
|
||||
|
||||
|
||||
def test_cleanup_stale_running_returns_ids():
|
||||
repo, _, engine = _repository()
|
||||
task = _make_task()
|
||||
repo.create(task)
|
||||
task.mark_processing()
|
||||
task.celery_task_id = "celery-msg-id-002"
|
||||
repo.update(task)
|
||||
with engine.connect() as conn:
|
||||
conn.execute(
|
||||
text("UPDATE generation_tasks SET updated_at = :ts WHERE id = :id"),
|
||||
{"ts": datetime.now(timezone.utc) - timedelta(minutes=60), "id": task.id},
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
items = repo.cleanup_stale_running_with_ids(timeout_minutes=20)
|
||||
assert len(items) == 1
|
||||
assert items[0][0] == task.id
|
||||
assert items[0][1] == "celery-msg-id-002"
|
||||
assert repo.get(task.id).status == GenerationTaskStatus.FAILED
|
||||
|
||||
|
||||
def test_legacy_repo_without_with_ids_falls_back_to_count():
|
||||
"""旧仓储只有 cleanup_stale_pending(返回 int)时降级可用,不抛异常。"""
|
||||
# worker 模块加载(标准 mock 模式)
|
||||
saved = set(sys.modules.keys())
|
||||
mock_db = MagicMock()
|
||||
mock_db.SessionLocal = MagicMock()
|
||||
sys.modules["worker_app.db"] = mock_db
|
||||
sys.modules["worker_app.core.config"] = MagicMock()
|
||||
mock_celery = MagicMock()
|
||||
mock_celery.celery_app.task = MagicMock(
|
||||
side_effect=(lambda *a, **k: (a[0] if a and callable(a[0]) else (lambda f: f)))
|
||||
)
|
||||
sys.modules["worker_app.celery_app"] = mock_celery
|
||||
worker_path = str(Path(__file__).resolve().parents[2] / "apps" / "worker")
|
||||
if worker_path not in sys.path:
|
||||
sys.path.insert(0, worker_path)
|
||||
|
||||
from worker_app.tasks import _startup # noqa: E402
|
||||
|
||||
class LegacyRepo:
|
||||
def cleanup_stale_pending(self, timeout_minutes): # noqa: ARG002
|
||||
return 3
|
||||
|
||||
def cleanup_stale_running(self, timeout_minutes): # noqa: ARG002
|
||||
return 2
|
||||
|
||||
items_p = _startup.cleanup_stale_pending_with_session_ids(LegacyRepo(), 45)
|
||||
items_r = _startup.cleanup_stale_running_with_session_ids(LegacyRepo(), 20)
|
||||
assert len(items_p) == 3
|
||||
assert len(items_r) == 2
|
||||
|
||||
for key in list(sys.modules.keys()):
|
||||
if key not in saved and not key.startswith("video_processing"):
|
||||
del sys.modules[key]
|
||||
|
||||
|
||||
# ── 端到端:清理 → 队列消息被移除(作废消息不重投) ────────────────────
|
||||
|
||||
|
||||
@pytest.mark.skipif(not _redis_available(), reason="本地 redis 不可用")
|
||||
def test_stale_pending_cleanup_purges_redis_message():
|
||||
"""任务标 failed 后,其在 Redis 队列里的 celery 消息被清除,不会被重投。"""
|
||||
import redis
|
||||
from celery import Celery
|
||||
from kombu import Queue
|
||||
from kombu.pools import producers
|
||||
|
||||
from packages.shared.celery_orphan_guard import purge_stale_messages_from_queues
|
||||
|
||||
repo, _, engine = _repository()
|
||||
task = _make_task()
|
||||
task.celery_task_id = "celery-stale-xyz"
|
||||
repo.create(task)
|
||||
with engine.connect() as conn:
|
||||
conn.execute(
|
||||
text("UPDATE generation_tasks SET created_at = :ts WHERE id = :id"),
|
||||
{"ts": datetime.now(timezone.utc) - timedelta(minutes=60), "id": task.id},
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
# 模拟该任务的 celery 消息仍在 generation 队列里(worker 下线期间未消费)
|
||||
client = redis.Redis.from_url(BROKER_URL)
|
||||
client.delete(TEST_QUEUE)
|
||||
app = Celery("test-e2e-revoke")
|
||||
app.conf.broker_url = BROKER_URL
|
||||
with app.connection_for_write() as conn:
|
||||
with producers[conn].acquire(block=True) as prod:
|
||||
# 作废任务消息
|
||||
prod.publish(
|
||||
(task.id,),
|
||||
exchange="",
|
||||
routing_key=TEST_QUEUE,
|
||||
serializer="json",
|
||||
headers={"id": "celery-stale-xyz", "task": "worker.generate_video"},
|
||||
retry=False,
|
||||
delivery_mode=1,
|
||||
declare=[Queue(TEST_QUEUE, routing_key=TEST_QUEUE, durable=False)],
|
||||
)
|
||||
# 另一条正常任务消息(必须保留)
|
||||
prod.publish(
|
||||
("other-task-id",),
|
||||
exchange="",
|
||||
routing_key=TEST_QUEUE,
|
||||
serializer="json",
|
||||
headers={"id": "celery-keep", "task": "worker.generate_video"},
|
||||
retry=False,
|
||||
delivery_mode=1,
|
||||
)
|
||||
|
||||
assert client.llen(TEST_QUEUE) == 2
|
||||
|
||||
# 执行清理(与 worker beat 相同流程:标 failed → 拿 ids → purge)
|
||||
items = repo.cleanup_stale_pending_with_ids(timeout_minutes=45)
|
||||
biz_ids = [bid for bid, _ in items]
|
||||
celery_ids = [cid for _, cid in items if cid]
|
||||
removed = purge_stale_messages_from_queues(
|
||||
BROKER_URL, (TEST_QUEUE,), business_task_ids=biz_ids, celery_task_ids=celery_ids
|
||||
)
|
||||
|
||||
assert removed == 1
|
||||
assert client.llen(TEST_QUEUE) == 1 # 正常任务消息保留
|
||||
client.delete(TEST_QUEUE)
|
||||
@@ -0,0 +1,243 @@
|
||||
"""Issue #1714:任务执行前状态守卫 — 已作废消息必须丢弃,禁止非法转换后继续跑。
|
||||
|
||||
覆盖:
|
||||
- ingest_asset:job 已 failed/completed 时直接返回 discarded,不下载、不转码、不回写
|
||||
- generate_video:GenerationTask 已 failed 时返回 discarded,不进入渲染
|
||||
- generate_video:pending → running 标记失败(非法转换)时安全中止
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
# ── worker 模块标准加载方式 ──
|
||||
# 显式保存将要覆盖的注入键旧值:全量收集时更早的测试文件(如
|
||||
# test_ingest_validation.py)可能已向 sys.modules 注入 worker_app.* mock,
|
||||
# 导入完成后必须精确恢复旧值,否则本文件的 bind 感知透传装饰器会残留,
|
||||
# 污染后续懒加载短路径 worker_app.celery_app 的 worker 测试。
|
||||
_INJECTED_KEYS = ("worker_app.db", "worker_app.core.config", "worker_app.celery_app")
|
||||
_SAVED_MODULE_VALUES = {k: sys.modules.get(k) for k in _INJECTED_KEYS}
|
||||
_SAVED_MODULES_KEYS = set(sys.modules.keys())
|
||||
|
||||
_mock_db_module = MagicMock()
|
||||
_mock_db_module.SessionLocal = MagicMock()
|
||||
sys.modules["worker_app.db"] = _mock_db_module
|
||||
sys.modules["worker_app.core.config"] = MagicMock()
|
||||
|
||||
_mock_celery_module = MagicMock()
|
||||
|
||||
|
||||
def _passthrough_decorator(*args, **kwargs):
|
||||
if len(args) == 1 and callable(args[0]):
|
||||
return args[0]
|
||||
bind = kwargs.get("bind", False)
|
||||
|
||||
def _wrap(f):
|
||||
if bind:
|
||||
# 模拟 celery bind=True:task(task_id) 调用时注入 self(MagicMock)
|
||||
return lambda *a, **kw: f(MagicMock(), *a, **kw)
|
||||
return f
|
||||
|
||||
return _wrap
|
||||
|
||||
|
||||
_mock_celery_module.celery_app.task = MagicMock(side_effect=_passthrough_decorator)
|
||||
sys.modules["worker_app.celery_app"] = _mock_celery_module
|
||||
|
||||
_WORKER_PATH = str(Path(__file__).resolve().parents[2] / "apps" / "worker")
|
||||
sys.path.insert(0, _WORKER_PATH)
|
||||
|
||||
import pytest # noqa: E402
|
||||
from worker_app.tasks import ingest as ingest_mod # noqa: E402
|
||||
|
||||
# video_processing 相关 mock(generation 模块导入链)
|
||||
for _mod_name in [
|
||||
"video_processing",
|
||||
"video_processing.ffmpeg_utils",
|
||||
"video_processing.oss_helpers",
|
||||
]:
|
||||
sys.modules.setdefault(_mod_name, MagicMock())
|
||||
|
||||
from worker_app.tasks import generation as gen_mod # noqa: E402
|
||||
|
||||
from packages.domain import IngestJobStatus # noqa: E402
|
||||
|
||||
# 模块导入完成后立即清理:删除本次 import 新引入的模块缓存(本模块已通过名字绑定
|
||||
# 持有 ingest_mod/gen_mod/IngestJobStatus,删除缓存不影响调用),再把三个注入键
|
||||
# 精确恢复为注入前的旧值(旧值不存在则移除),杜绝 mock 残留污染其他 worker 测试。
|
||||
for _key in list(sys.modules.keys()):
|
||||
if _key not in _SAVED_MODULES_KEYS and not _key.startswith("video_processing"):
|
||||
del sys.modules[_key]
|
||||
for _k, _v in _SAVED_MODULE_VALUES.items():
|
||||
if _v is None:
|
||||
sys.modules.pop(_k, None)
|
||||
else:
|
||||
sys.modules[_k] = _v
|
||||
del _SAVED_MODULES_KEYS, _SAVED_MODULE_VALUES
|
||||
|
||||
|
||||
# ── ingest 守卫 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class _FakeJobRepo:
|
||||
def __init__(self, job):
|
||||
self.job = job
|
||||
|
||||
def get(self, job_id):
|
||||
return self.job
|
||||
|
||||
|
||||
def _make_ingest_job(status):
|
||||
job = MagicMock()
|
||||
job.id = "job-stale-1"
|
||||
job.storage_key = "uploads/proj/stale.mov"
|
||||
job.status = status
|
||||
job.file_hash = "h"
|
||||
job.asset_id = ""
|
||||
return job
|
||||
|
||||
|
||||
def test_ingest_discards_failed_job_message():
|
||||
"""job 已 failed:消息丢弃,不进入下载/转码/回写。"""
|
||||
job = _make_ingest_job(IngestJobStatus.FAILED)
|
||||
fake_session = MagicMock()
|
||||
_mock_db_module.SessionLocal = MagicMock(return_value=fake_session)
|
||||
|
||||
# SQLAlchemy 仓储构造返回 fake
|
||||
fake_job_repo = _FakeJobRepo(job)
|
||||
fake_asset_repo = MagicMock()
|
||||
|
||||
orig_job_repo = ingest_mod.SQLAlchemyIngestJobRepository
|
||||
orig_asset_repo = ingest_mod.SQLAlchemyAssetRepository
|
||||
ingest_mod.SQLAlchemyIngestJobRepository = MagicMock(return_value=fake_job_repo)
|
||||
ingest_mod.SQLAlchemyAssetRepository = MagicMock(return_value=fake_asset_repo)
|
||||
try:
|
||||
result = ingest_mod.ingest_asset("job-stale-1")
|
||||
finally:
|
||||
ingest_mod.SQLAlchemyIngestJobRepository = orig_job_repo
|
||||
ingest_mod.SQLAlchemyAssetRepository = orig_asset_repo
|
||||
|
||||
assert result["status"] == "discarded"
|
||||
# 没有任何 update / commit / 下载动作
|
||||
fake_session.commit.assert_not_called()
|
||||
fake_asset_repo.create.assert_not_called()
|
||||
|
||||
|
||||
def test_ingest_discards_completed_job_message():
|
||||
job = _make_ingest_job(IngestJobStatus.COMPLETED)
|
||||
fake_session = MagicMock()
|
||||
_mock_db_module.SessionLocal = MagicMock(return_value=fake_session)
|
||||
fake_job_repo = _FakeJobRepo(job)
|
||||
|
||||
orig = ingest_mod.SQLAlchemyIngestJobRepository
|
||||
ingest_mod.SQLAlchemyIngestJobRepository = MagicMock(return_value=fake_job_repo)
|
||||
ingest_mod.SQLAlchemyAssetRepository = MagicMock(return_value=MagicMock())
|
||||
try:
|
||||
result = ingest_mod.ingest_asset("job-stale-1")
|
||||
finally:
|
||||
ingest_mod.SQLAlchemyIngestJobRepository = orig
|
||||
|
||||
assert result["status"] == "discarded"
|
||||
|
||||
|
||||
# ── generation 守卫 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _make_gen_task(status_value: str):
|
||||
from packages.domain import GenerationTask
|
||||
|
||||
task = GenerationTask.create(project_id="p", asset_library_id="l", created_by_user_id="u")
|
||||
task.status = type(task.status)(status_value)
|
||||
return task
|
||||
|
||||
|
||||
def test_generate_video_discards_failed_task(monkeypatch):
|
||||
"""GenerationTask 已 failed:直接 discarded,不加载渲染数据。"""
|
||||
failed_task = _make_gen_task("failed")
|
||||
|
||||
fake_repo = MagicMock()
|
||||
fake_repo.get.return_value = failed_task
|
||||
|
||||
fake_session = MagicMock()
|
||||
_mock_db_module.SessionLocal = MagicMock(return_value=fake_session)
|
||||
|
||||
import packages.adapters.sqlalchemy_impl.generation_task_repository as gen_repo_mod
|
||||
|
||||
orig = gen_repo_mod.SQLAlchemyGenerationTaskRepository
|
||||
gen_repo_mod.SQLAlchemyGenerationTaskRepository = MagicMock(return_value=fake_repo)
|
||||
|
||||
update_status_mock = MagicMock(return_value=False)
|
||||
monkeypatch.setattr(gen_mod, "_update_task_status", update_status_mock)
|
||||
monkeypatch.setattr(
|
||||
gen_mod,
|
||||
"_load_task_info",
|
||||
lambda task_id: {
|
||||
"project_id": "p",
|
||||
"template_id": "",
|
||||
"task_asset_ids": [],
|
||||
"batch_id": "",
|
||||
"user_id": "u",
|
||||
"mode": "one_take",
|
||||
},
|
||||
)
|
||||
monkeypatch.setattr(gen_mod, "_flush_logs", lambda *a, **k: None)
|
||||
|
||||
task_fn = gen_mod.generate_video
|
||||
if hasattr(task_fn, "__wrapped__"):
|
||||
task_fn = task_fn.__wrapped__
|
||||
try:
|
||||
result = task_fn("task-stale-1")
|
||||
finally:
|
||||
gen_repo_mod.SQLAlchemyGenerationTaskRepository = orig
|
||||
|
||||
assert result["status"] == "discarded"
|
||||
# 状态守卫命中终态,根本不应尝试 mark_processing
|
||||
update_status_mock.assert_not_called()
|
||||
|
||||
|
||||
def test_generate_video_aborts_when_claim_fails(monkeypatch):
|
||||
"""pending 但 mark_processing 返回 False(状态机非法转换)时安全中止。"""
|
||||
pending_task = _make_gen_task("pending")
|
||||
|
||||
fake_repo = MagicMock()
|
||||
fake_repo.get.return_value = pending_task
|
||||
fake_session = MagicMock()
|
||||
_mock_db_module.SessionLocal = MagicMock(return_value=fake_session)
|
||||
|
||||
import packages.adapters.sqlalchemy_impl.generation_task_repository as gen_repo_mod
|
||||
|
||||
orig = gen_repo_mod.SQLAlchemyGenerationTaskRepository
|
||||
gen_repo_mod.SQLAlchemyGenerationTaskRepository = MagicMock(return_value=fake_repo)
|
||||
|
||||
monkeypatch.setattr(
|
||||
gen_mod,
|
||||
"_load_task_info",
|
||||
lambda task_id: {
|
||||
"project_id": "p",
|
||||
"template_id": "",
|
||||
"task_asset_ids": [],
|
||||
"batch_id": "",
|
||||
"user_id": "u",
|
||||
"mode": "one_take",
|
||||
},
|
||||
)
|
||||
monkeypatch.setattr(gen_mod, "_flush_logs", lambda *a, **k: None)
|
||||
# 模拟 mark_processing 失败(failed→running 非法转换被 _update_task_status 吞掉返回 False)
|
||||
update_status_mock = MagicMock(return_value=False)
|
||||
monkeypatch.setattr(gen_mod, "_update_task_status", update_status_mock)
|
||||
render_mock = MagicMock(side_effect=AssertionError("must not render"))
|
||||
monkeypatch.setattr(gen_mod, "_render_from_edit_plan", render_mock)
|
||||
|
||||
task_fn = gen_mod.generate_video
|
||||
if hasattr(task_fn, "__wrapped__"):
|
||||
task_fn = task_fn.__wrapped__
|
||||
try:
|
||||
result = task_fn("task-claim-fail")
|
||||
finally:
|
||||
gen_repo_mod.SQLAlchemyGenerationTaskRepository = orig
|
||||
|
||||
assert result["status"] == "discarded"
|
||||
render_mock.assert_not_called()
|
||||
@@ -0,0 +1,315 @@
|
||||
"""Issue #1709 任务容错:孤儿任务恢复 + 429 限流结构化提示。
|
||||
|
||||
覆盖:
|
||||
1. 仓储层:count_running_by_user/count_running_total 计数正确(预览/正式任务都计入)
|
||||
2. 仓储层:estimate_avg_duration_seconds 耗时估算(有历史/无历史)
|
||||
3. 限流核心:build_rate_limit_detail 返回结构化 code/message/排队数/预计等待
|
||||
4. worker 侧:cleanup_stale_running/pending 核心函数——中断任务被重置为 failed
|
||||
且原因写明(容器重启/超时中断),正常任务不受影响
|
||||
"""
|
||||
|
||||
import sys
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "worker"))
|
||||
|
||||
# 预注入 mock worker_app.db,防止真实数据库连接初始化(与其他 worker 测试同模式)
|
||||
_mock_db = MagicMock()
|
||||
_mock_db.SessionLocal = MagicMock()
|
||||
sys.modules.setdefault("worker_app.db", _mock_db)
|
||||
|
||||
from app.core import task_enqueue # noqa: E402
|
||||
from sqlalchemy import create_engine, text # noqa: E402
|
||||
from sqlalchemy.orm import sessionmaker # noqa: E402
|
||||
from worker_app.tasks import _startup # noqa: E402
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.generation_task_repository import ( # noqa: E402
|
||||
SQLAlchemyGenerationTaskRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.models import Base # noqa: E402
|
||||
from packages.domain import GenerationTask, GenerationTaskStatus # noqa: E402
|
||||
|
||||
|
||||
def _repository():
|
||||
engine = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False})
|
||||
Base.metadata.create_all(engine)
|
||||
session = sessionmaker(bind=engine)()
|
||||
return SQLAlchemyGenerationTaskRepository(session), session, engine
|
||||
|
||||
|
||||
def _make_task(**kwargs) -> GenerationTask:
|
||||
defaults = dict(
|
||||
project_id="proj-1",
|
||||
asset_library_id="lib-1",
|
||||
created_by_user_id="user-1",
|
||||
)
|
||||
defaults.update(kwargs)
|
||||
return GenerationTask.create(**defaults)
|
||||
|
||||
|
||||
def _age_task(engine, task_id, *, updated_minutes=None, created_minutes=None):
|
||||
"""用 SQL 直接把 updated_at/created_at 改到过去(模拟孤儿任务)。"""
|
||||
sets, params = [], {"id": task_id}
|
||||
if updated_minutes is not None:
|
||||
sets.append("updated_at = :uts")
|
||||
params["uts"] = datetime.now(timezone.utc) - timedelta(minutes=updated_minutes)
|
||||
if created_minutes is not None:
|
||||
sets.append("created_at = :cts")
|
||||
params["cts"] = datetime.now(timezone.utc) - timedelta(minutes=created_minutes)
|
||||
with engine.connect() as conn:
|
||||
conn.execute(text(f"UPDATE generation_tasks SET {', '.join(sets)} WHERE id = :id"), params)
|
||||
conn.commit()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 1. running 计数(限流"渲染中"数量)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_count_running_by_user_mix_statuses():
|
||||
"""count_running_by_user 只统计该用户 running,不含 pending/completed/failed。"""
|
||||
repo, _, _ = _repository()
|
||||
t1 = _make_task(project_id="p1")
|
||||
repo.create(t1) # pending
|
||||
t2 = _make_task(project_id="p2")
|
||||
repo.create(t2)
|
||||
t2.mark_processing()
|
||||
repo.update(t2)
|
||||
t3 = _make_task(project_id="p3")
|
||||
repo.create(t3)
|
||||
t3.mark_processing()
|
||||
repo.update(t3)
|
||||
t4 = _make_task(project_id="p4")
|
||||
repo.create(t4)
|
||||
t4.mark_processing()
|
||||
repo.update(t4)
|
||||
t4.mark_completed()
|
||||
repo.update(t4)
|
||||
t5 = _make_task(project_id="p5", created_by_user_id="user-2")
|
||||
repo.create(t5)
|
||||
t5.mark_processing()
|
||||
repo.update(t5)
|
||||
|
||||
assert repo.count_running_by_user("user-1") == 2
|
||||
assert repo.count_running_by_user("user-2") == 1
|
||||
assert repo.count_running_total() == 3
|
||||
|
||||
|
||||
def test_count_running_total_empty():
|
||||
repo, _, _ = _repository()
|
||||
assert repo.count_running_total() == 0
|
||||
assert repo.count_running_by_user("nobody") == 0
|
||||
|
||||
|
||||
def test_preview_tasks_counted_in_running():
|
||||
"""预览任务(is_preview=True,工单实测卡 80% 的那种)同样计入 running。"""
|
||||
repo, _, _ = _repository()
|
||||
t = _make_task(is_preview=True)
|
||||
repo.create(t)
|
||||
t.mark_processing()
|
||||
repo.update(t)
|
||||
assert repo.count_running_by_user("user-1") == 1
|
||||
assert repo.count_running_total() == 1
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 2. 平均耗时估算(429 等待预估依据)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _complete_task(repo, engine, task, duration_seconds: float):
|
||||
repo.create(task)
|
||||
task.mark_processing()
|
||||
repo.update(task)
|
||||
task.mark_completed()
|
||||
repo.update(task)
|
||||
now = datetime.now(timezone.utc)
|
||||
with engine.connect() as conn:
|
||||
conn.execute(
|
||||
text("UPDATE generation_tasks SET started_at = :s, completed_at = :c WHERE id = :id"),
|
||||
{"s": now - timedelta(seconds=duration_seconds), "c": now, "id": task.id},
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
|
||||
def test_estimate_avg_duration_with_history():
|
||||
"""有历史完成任务时返回平均耗时(秒)。"""
|
||||
repo, _, engine = _repository()
|
||||
_complete_task(repo, engine, _make_task(project_id="p1"), 60.0)
|
||||
_complete_task(repo, engine, _make_task(project_id="p2"), 180.0)
|
||||
|
||||
avg = repo.estimate_avg_duration_seconds(default_seconds=120.0)
|
||||
assert 119.0 < avg < 121.0 # (60+180)/2 = 120
|
||||
|
||||
|
||||
def test_estimate_avg_duration_no_history_returns_default():
|
||||
"""无历史数据时返回默认值。"""
|
||||
repo, _, _ = _repository()
|
||||
assert repo.estimate_avg_duration_seconds(default_seconds=90.0) == 90.0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 3. build_rate_limit_detail 结构化提示(前端区分"排队"与"创建失败")
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_user_rate_limit_detail_structure():
|
||||
"""429 用户限流:返回 USER_QUEUE_FULL + 排队/渲染数 + 预计等待。"""
|
||||
repo, _, _ = _repository()
|
||||
for i in range(2): # 2 个渲染中
|
||||
t = _make_task(project_id=f"rp{i}")
|
||||
repo.create(t)
|
||||
t.mark_processing()
|
||||
repo.update(t)
|
||||
|
||||
exc = task_enqueue.UserPendingLimitExceeded(user_id="user-1", pending_count=3, limit=3)
|
||||
detail = task_enqueue.build_rate_limit_detail(exc, repo, scope="user")
|
||||
|
||||
assert detail["code"] == task_enqueue.ERROR_CODE_USER_QUEUE_FULL
|
||||
assert detail["queued_count"] == 3
|
||||
assert detail["running_count"] == 2
|
||||
assert detail["limit"] == 3
|
||||
assert detail["estimated_wait_seconds"] > 0
|
||||
assert "排队" in detail["message"]
|
||||
assert "user-1" not in detail["message"] # 不泄露内部 ID
|
||||
|
||||
|
||||
def test_global_rate_limit_detail_structure():
|
||||
"""503 全局繁忙:返回 SYSTEM_QUEUE_FULL。"""
|
||||
repo, _, _ = _repository()
|
||||
exc = task_enqueue.GlobalQueueFull(pending_count=20, limit=20)
|
||||
detail = task_enqueue.build_rate_limit_detail(exc, repo, scope="global")
|
||||
|
||||
assert detail["code"] == task_enqueue.ERROR_CODE_SYSTEM_QUEUE_FULL
|
||||
assert detail["queued_count"] == 20
|
||||
assert detail["limit"] == 20
|
||||
assert detail["estimated_wait_seconds"] > 0
|
||||
assert "系统繁忙" in detail["message"]
|
||||
|
||||
|
||||
def test_wait_estimate_uses_concurrency():
|
||||
"""等待预估:排队 8 个 / 并发 4 = 2 批 × 平均耗时。"""
|
||||
|
||||
class FakeRepo:
|
||||
def estimate_avg_duration_seconds(self, limit=20, default_seconds=120.0):
|
||||
return 100.0
|
||||
|
||||
wait = task_enqueue._estimate_wait_seconds(8, FakeRepo())
|
||||
assert wait == 200 # ceil(8/4)=2 批 × 100 秒
|
||||
|
||||
|
||||
def test_wait_estimate_repo_without_methods_uses_default():
|
||||
"""仓储没有新方法(旧 mock/鸭子类型)时用默认 120 秒兜底,不抛错。"""
|
||||
|
||||
class LegacyRepo:
|
||||
"""只实现旧接口的仓储(模拟未升级的调用方)。"""
|
||||
|
||||
def count_pending_total(self):
|
||||
return 0
|
||||
|
||||
wait = task_enqueue._estimate_wait_seconds(4, LegacyRepo())
|
||||
assert wait == 120 # ceil(4/4)=1 批 × 120 默认
|
||||
|
||||
|
||||
def test_rate_limit_detail_running_count_falls_back_to_zero():
|
||||
"""仓储不支持 running 计数时,running_count 优雅降级为 0。"""
|
||||
|
||||
class LegacyRepo:
|
||||
def count_pending_total(self):
|
||||
return 0
|
||||
|
||||
exc = task_enqueue.GlobalQueueFull(pending_count=20, limit=20)
|
||||
detail = task_enqueue.build_rate_limit_detail(exc, LegacyRepo(), scope="global")
|
||||
assert detail["running_count"] == 0
|
||||
assert detail["code"] == task_enqueue.ERROR_CODE_SYSTEM_QUEUE_FULL
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 4. worker 清理核心:中断任务被重置(worker 重启/超时恢复)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_worker_cleanup_resets_interrupted_running_task():
|
||||
"""模拟 worker 重启:running 超 20 分钟无更新的任务被重置为 failed,原因写明。"""
|
||||
repo, _, engine = _repository()
|
||||
|
||||
t = _make_task(is_preview=True) # 预览任务
|
||||
repo.create(t)
|
||||
t.mark_processing() # running
|
||||
repo.update(t)
|
||||
_age_task(engine, t.id, updated_minutes=25) # 25 分钟无进度更新
|
||||
|
||||
cleaned = _startup.cleanup_stale_running_with_session(repo, 20)
|
||||
assert cleaned == 1
|
||||
|
||||
saved = repo.get(t.id)
|
||||
assert saved.status == GenerationTaskStatus.FAILED
|
||||
assert "中断" in saved.error_message
|
||||
assert saved.error_info.get("error_type") == "WorkerInterrupted"
|
||||
assert saved.completed_at is not None
|
||||
|
||||
|
||||
def test_worker_cleanup_keeps_healthy_running_task():
|
||||
"""正常运行中(5 分钟前有更新)的任务不被误杀。"""
|
||||
repo, _, engine = _repository()
|
||||
|
||||
t = _make_task()
|
||||
repo.create(t)
|
||||
t.mark_processing()
|
||||
repo.update(t)
|
||||
_age_task(engine, t.id, updated_minutes=5)
|
||||
|
||||
assert _startup.cleanup_stale_running_with_session(repo, 20) == 0
|
||||
assert repo.get(t.id).status == GenerationTaskStatus.RUNNING
|
||||
|
||||
|
||||
def test_worker_cleanup_resets_stale_pending_task():
|
||||
"""卡 pending 超 15 分钟(worker 停止消费)的任务被重置,释放限流名额。"""
|
||||
repo, _, engine = _repository()
|
||||
|
||||
t = _make_task(is_preview=True)
|
||||
repo.create(t) # 一直 pending
|
||||
_age_task(engine, t.id, created_minutes=20)
|
||||
|
||||
cleaned = _startup.cleanup_stale_pending_with_session(repo, 15)
|
||||
assert cleaned == 1
|
||||
|
||||
saved = repo.get(t.id)
|
||||
assert saved.status == GenerationTaskStatus.FAILED
|
||||
assert saved.error_info.get("error_type") == "PendingTimeout"
|
||||
# 释放名额后 pending 计数归零,新请求不再被 429 误伤
|
||||
assert repo.count_pending_total() == 0
|
||||
|
||||
|
||||
def test_worker_cleanup_pending_keeps_recent():
|
||||
"""刚创建 3 分钟的 pending 任务不清理。"""
|
||||
repo, _, engine = _repository()
|
||||
|
||||
t = _make_task()
|
||||
repo.create(t)
|
||||
_age_task(engine, t.id, created_minutes=3)
|
||||
|
||||
assert _startup.cleanup_stale_pending_with_session(repo, 15) == 0
|
||||
assert repo.get(t.id).status == GenerationTaskStatus.PENDING
|
||||
|
||||
|
||||
def test_worker_cleanup_multiple_orphans_all_reset():
|
||||
"""3 个卡死 running 任务(工单实测:3 个预览卡 80% 超 10 小时)全部恢复。"""
|
||||
repo, _, engine = _repository()
|
||||
|
||||
ids = []
|
||||
for i in range(3):
|
||||
t = _make_task(project_id=f"p{i}", is_preview=True)
|
||||
repo.create(t)
|
||||
t.mark_processing()
|
||||
repo.update(t)
|
||||
_age_task(engine, t.id, updated_minutes=600) # 10 小时
|
||||
ids.append(t.id)
|
||||
|
||||
cleaned = _startup.cleanup_stale_running_with_session(repo, 20)
|
||||
assert cleaned == 3
|
||||
for tid in ids:
|
||||
assert repo.get(tid).status == GenerationTaskStatus.FAILED
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user