Compare commits
21 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 8ecf381a9d | |||
| 244691d335 | |||
| ee4fff42f0 | |||
| b0018e747b | |||
| cbca0c3584 | |||
| c7c30936a9 | |||
| aac8cc5fd7 | |||
| 229f9dddeb | |||
| 6d2d63da7e | |||
| e6e4090f3c | |||
| 68c9db18b1 | |||
| bb7dc71da3 | |||
| e0e7f0a503 | |||
| 4e270f1fb5 | |||
| c6147fbb0c | |||
| 18beb7cfa3 | |||
| 5474812fab | |||
| ffd02a3ec8 | |||
| 5cdaa29511 | |||
| 0e2dd60a8d | |||
| ec45d71d2c |
@@ -1462,7 +1462,7 @@ jobs:
|
||||
- unit-tests
|
||||
- frontend-lint
|
||||
- frontend-unit-test
|
||||
if: github.event_name == 'push' && !failure() && !cancelled()
|
||||
if: github.event_name == 'push' && github.ref_name == 'main' && !failure() && !cancelled()
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
@@ -1608,7 +1608,7 @@ jobs:
|
||||
concurrency:
|
||||
group: deploy-production-${{ gitea.ref }}
|
||||
cancel-in-progress: false
|
||||
# if: removed - runs after build-production succeeds
|
||||
if: github.event_name == 'push' && github.ref_name == 'main'
|
||||
needs:
|
||||
- build-production
|
||||
steps:
|
||||
|
||||
@@ -18,7 +18,7 @@ jobs:
|
||||
name: Auto Approve on CI Green
|
||||
runs-on: ci-check
|
||||
if: github.event_name == 'pull_request' && !github.event.pull_request.draft
|
||||
timeout-minutes: 3 # 长等待模式:等CI全绿后自动合并,不遗漏任何PR
|
||||
timeout-minutes: 10 # 等待CI全绿+审批,需要充足时间
|
||||
steps:
|
||||
- name: Checkout code
|
||||
shell: sh
|
||||
@@ -61,7 +61,8 @@ jobs:
|
||||
name: Auto Merge on CI Green + Approved
|
||||
runs-on: ci-check
|
||||
if: github.event_name == 'pull_request' && !github.event.pull_request.draft && github.event.pull_request.base.ref == 'develop'
|
||||
timeout-minutes: 3 # 短作业模式:检查一次,不满足就退出,由pr-auto-scan每5分钟定时兜底
|
||||
needs: [auto-approve] # 修复竞态:必须等审批完成后再尝试合并
|
||||
timeout-minutes: 15 # 等待审批+CI就绪+合并,需要充足时间
|
||||
steps:
|
||||
- name: Checkout code
|
||||
shell: sh
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
"""add video_fingerprint_chunks table for per-chunk fingerprint storage
|
||||
|
||||
Revision ID: 063_fingerprint_chunks
|
||||
Revises: 062_edit_plan_id
|
||||
Create Date: 2026-09-03
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "063_fingerprint_chunks"
|
||||
down_revision = "062_edit_plan_id"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"video_fingerprint_chunks",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("video_id", sa.String(36), nullable=False),
|
||||
sa.Column("project_id", sa.String(36), nullable=False),
|
||||
sa.Column("user_id", sa.String(36), nullable=False, server_default=""),
|
||||
sa.Column("start_time_ms", sa.Integer, nullable=False),
|
||||
sa.Column("end_time_ms", sa.Integer, nullable=False),
|
||||
sa.Column("phash_binary", sa.String(16), nullable=False),
|
||||
sa.Column("color_histogram", sa.JSON, nullable=False),
|
||||
sa.Column("frame_count", sa.Integer, nullable=False, server_default="1"),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime,
|
||||
nullable=False,
|
||||
server_default=sa.func.now(),
|
||||
),
|
||||
)
|
||||
op.create_index("ix_vfc_video_id", "video_fingerprint_chunks", ["video_id"])
|
||||
op.create_index("ix_vfc_project_id", "video_fingerprint_chunks", ["project_id"])
|
||||
op.create_index("ix_vfc_user_id", "video_fingerprint_chunks", ["user_id"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_vfc_user_id", table_name="video_fingerprint_chunks")
|
||||
op.drop_index("ix_vfc_project_id", table_name="video_fingerprint_chunks")
|
||||
op.drop_index("ix_vfc_video_id", table_name="video_fingerprint_chunks")
|
||||
op.drop_table("video_fingerprint_chunks")
|
||||
@@ -2,7 +2,9 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import subprocess
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
@@ -430,6 +432,8 @@ def save_tts_job_to_library(
|
||||
storage_key = f"uploads/voice/tts/{job.id}.{audio_format}"
|
||||
|
||||
tmp_path: Path | None = None
|
||||
audio_duration: float | None = None
|
||||
file_size = 0
|
||||
try:
|
||||
with tempfile.NamedTemporaryFile(suffix=f".{audio_format}", delete=False) as tmp:
|
||||
tmp_path = Path(tmp.name)
|
||||
@@ -445,6 +449,23 @@ def save_tts_job_to_library(
|
||||
)
|
||||
file_size = tmp_path.stat().st_size
|
||||
storage_service.upload_file(tmp_path, storage_key, content_type=content_type)
|
||||
|
||||
# 从音频文件提取时长(ffprobe),作为 job.duration 的兜底
|
||||
try:
|
||||
proc = subprocess.run(
|
||||
[
|
||||
"ffprobe", "-v", "quiet", "-print_format", "json",
|
||||
"-show_format", str(tmp_path),
|
||||
],
|
||||
capture_output=True, text=True, timeout=10,
|
||||
)
|
||||
if proc.returncode == 0:
|
||||
fmt = json.loads(proc.stdout).get("format", {})
|
||||
dur = float(fmt.get("duration", 0))
|
||||
if dur > 0:
|
||||
audio_duration = dur
|
||||
except Exception:
|
||||
logger.warning("ffprobe 提取时长失败: job_id=%s", job.id, exc_info=True)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
@@ -482,7 +503,7 @@ def save_tts_job_to_library(
|
||||
mime_type=content_type,
|
||||
metadata=metadata_,
|
||||
file_size=file_size,
|
||||
duration=job.duration or None,
|
||||
duration=job.duration or audio_duration or None,
|
||||
status=AssetStatus.READY,
|
||||
classification_status=ClassificationStatus.PENDING, # 音频不参与内容分类,保持 pending 与 ingest 链路一致
|
||||
uploaded_by_user_id=user_id,
|
||||
|
||||
@@ -23,6 +23,7 @@ from app.schemas.upload import (
|
||||
from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile, status
|
||||
|
||||
from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase
|
||||
from packages.domain import Asset, AssetStatus
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -80,6 +81,40 @@ def _validate_mime_type(content_type: str | None) -> str:
|
||||
return base_type
|
||||
|
||||
|
||||
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",
|
||||
}
|
||||
for ext, mime in _MIME_MAP.items():
|
||||
if lower_filename.endswith(ext):
|
||||
return mime
|
||||
return "video/mp4" # default
|
||||
|
||||
|
||||
def _create_pending_asset(
|
||||
asset_repository, project_id, library_id, storage_key, filename, mime_type, user_id, file_hash=""
|
||||
):
|
||||
"""立即创建一条 PROCESSING 状态的 Asset 记录,使前端能马上看到新素材。"""
|
||||
asset = Asset.create(
|
||||
project_id=project_id,
|
||||
library_id=library_id,
|
||||
name=filename,
|
||||
storage_key=storage_key,
|
||||
mime_type=mime_type,
|
||||
status=AssetStatus.PROCESSING,
|
||||
uploaded_by_user_id=user_id,
|
||||
file_hash=file_hash,
|
||||
)
|
||||
return asset_repository.create(asset)
|
||||
|
||||
|
||||
def _submit_ingest_job(
|
||||
project_id: str,
|
||||
library_id: str,
|
||||
@@ -209,6 +244,20 @@ async def complete_direct_upload(
|
||||
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,
|
||||
project_id=request.project_id,
|
||||
library_id=request.library_id,
|
||||
storage_key=normalized_key,
|
||||
filename=filename,
|
||||
mime_type=mime_type,
|
||||
user_id=authenticated_user.user.id,
|
||||
file_hash=request.file_hash,
|
||||
)
|
||||
|
||||
job = _submit_ingest_job(
|
||||
project_id=request.project_id,
|
||||
library_id=request.library_id,
|
||||
@@ -216,7 +265,12 @@ async def complete_direct_upload(
|
||||
ingest_job_repository=ingest_job_repository,
|
||||
file_hash=request.file_hash,
|
||||
)
|
||||
return DirectUploadCompleteResponse(storage_key=normalized_key, ingest_job_id=job.id, url=storage_service.get_url(normalized_key))
|
||||
return DirectUploadCompleteResponse(
|
||||
storage_key=normalized_key,
|
||||
ingest_job_id=job.id,
|
||||
asset_id=pending_asset.id,
|
||||
url=storage_service.get_url(normalized_key),
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
@@ -284,6 +338,18 @@ async def upload_asset(
|
||||
detail=f"Failed to upload file: {type(error).__name__}",
|
||||
) from error
|
||||
|
||||
# 立即创建 Asset 记录(PROCESSING 状态),使前端刷新后即可看到新素材
|
||||
pending_asset = _create_pending_asset(
|
||||
asset_repository=asset_repository,
|
||||
project_id=project_id,
|
||||
library_id=library_id,
|
||||
storage_key=storage_key,
|
||||
filename=safe_filename,
|
||||
mime_type=validated_content_type,
|
||||
user_id=authenticated_user.user.id,
|
||||
file_hash=file_hash,
|
||||
)
|
||||
|
||||
job = _submit_ingest_job(
|
||||
project_id=project_id,
|
||||
library_id=library_id,
|
||||
@@ -295,5 +361,6 @@ async def upload_asset(
|
||||
return UploadAssetResponse(
|
||||
storage_key=storage_key,
|
||||
ingest_job_id=job.id,
|
||||
asset_id=pending_asset.id,
|
||||
url=file_url,
|
||||
)
|
||||
|
||||
@@ -6,12 +6,26 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import shutil
|
||||
import subprocess
|
||||
import tempfile
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Literal, Optional
|
||||
from uuid import uuid4
|
||||
|
||||
from app.api.routes._helpers import get_user_plan
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import get_audio_url_signer, get_cosyvoice_service, get_db_session, get_user_repository
|
||||
from app.core.storage import get_storage_service
|
||||
from app.dependencies import (
|
||||
get_asset_library_repository,
|
||||
get_asset_repository,
|
||||
get_audio_url_signer,
|
||||
get_cosyvoice_service,
|
||||
get_db_session,
|
||||
get_project_repository,
|
||||
get_user_repository,
|
||||
)
|
||||
from app.schemas.voice import (
|
||||
PresetVoiceItemResponse,
|
||||
PresetVoiceListResponse,
|
||||
@@ -24,7 +38,7 @@ from app.schemas.voice_library import (
|
||||
UpdateVoiceLibraryRequest,
|
||||
VoiceLibraryItemResponse,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
|
||||
from fastapi import APIRouter, Depends, File, Form, HTTPException, Query, Response, UploadFile, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.voice_clone_profile_repository import SQLAlchemyVoiceCloneProfileRepository
|
||||
@@ -40,8 +54,12 @@ from packages.application.voice_library.use_cases import (
|
||||
QuotaExceededError,
|
||||
UpdateVoiceLibraryUseCase,
|
||||
)
|
||||
from packages.domain import Asset, AssetStatus
|
||||
from packages.domain.classification import AssetLibraryKind, ClassificationStatus
|
||||
from packages.domain.entities import AssetLibrary
|
||||
from packages.domain.preset_voices import PRESET_VOICES, get_preset_voice_by_id
|
||||
from packages.ports.user_repository import UserRepository
|
||||
from packages.shared.storage import SharedStorageService
|
||||
|
||||
router = APIRouter()
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -507,3 +525,243 @@ def delete_voice(
|
||||
if not deleted:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Voice not found")
|
||||
return
|
||||
|
||||
|
||||
# ── 提取视频配音 ─────────────────────────────────────────────────────
|
||||
|
||||
# 支持的视频格式
|
||||
EXTRACT_VIDEO_MIMES = frozenset({"video/mp4", "video/quicktime", "video/webm", "video/x-msvideo"})
|
||||
MAX_EXTRACT_SIZE = 500 * 1024 * 1024 # 500MB
|
||||
|
||||
|
||||
@router.post(
|
||||
"/extract-voice",
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
)
|
||||
def extract_voice_from_video(
|
||||
file: UploadFile = File(...),
|
||||
project_id: str = Form(...),
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
project_repository=Depends(get_project_repository),
|
||||
asset_library_repository=Depends(get_asset_library_repository),
|
||||
asset_repository=Depends(get_asset_repository),
|
||||
storage_service: SharedStorageService = Depends(get_storage_service),
|
||||
sign_url=Depends(get_audio_url_signer),
|
||||
):
|
||||
"""从上传的视频中提取人声配音。
|
||||
|
||||
流程:
|
||||
1. 接收视频文件(mp4/mov/webm)
|
||||
2. ffmpeg 提取音频 + 降噪 + 编码为 mp3
|
||||
3. 上传到 OSS,创建 Asset 记录到配音素材库
|
||||
4. 返回素材信息(时长、文件大小、URL)
|
||||
"""
|
||||
user_id = authenticated_user.user.id
|
||||
|
||||
# 校验文件类型
|
||||
content_type = file.content_type or ""
|
||||
if content_type and content_type not in EXTRACT_VIDEO_MIMES:
|
||||
# 兜底:按扩展名判断
|
||||
ext = (file.filename or "").rsplit(".", 1)[-1].lower()
|
||||
ext_to_mime = {"mp4": "video/mp4", "mov": "video/quicktime", "webm": "video/webm", "avi": "video/x-msvideo"}
|
||||
if ext not in ext_to_mime:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="仅支持 mp4/mov/webm/avi 格式的视频文件",
|
||||
)
|
||||
content_type = ext_to_mime[ext]
|
||||
|
||||
# 找到(或自动创建)用户 voice 素材库(复用 TTS 的逻辑)
|
||||
library = _find_or_create_voice_library_for_extract(
|
||||
user_id=user_id,
|
||||
project_repository=project_repository,
|
||||
asset_library_repository=asset_library_repository,
|
||||
)
|
||||
|
||||
tmp_dir = None
|
||||
try:
|
||||
tmp_dir = Path(tempfile.mkdtemp(prefix="voice_extract_"))
|
||||
video_path = tmp_dir / f"input_{uuid4().hex[:8]}_{file.filename or 'video.mp4'}"
|
||||
audio_path = tmp_dir / f"output_{uuid4().hex[:8]}.mp3"
|
||||
|
||||
# 保存上传的视频到临时文件
|
||||
with open(video_path, "wb") as f:
|
||||
total = 0
|
||||
while chunk := file.file.read(1024 * 1024): # 1MB chunks
|
||||
total += len(chunk)
|
||||
if total > MAX_EXTRACT_SIZE:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
|
||||
detail="视频文件过大,最大支持 500MB",
|
||||
)
|
||||
f.write(chunk)
|
||||
|
||||
if video_path.stat().st_size == 0:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="视频文件为空")
|
||||
|
||||
# ffmpeg: 提取音频 + 降噪 + 编码 mp3
|
||||
# 滤镜链:highpass(去低频噪声) → afftdn(FFT降噪) → lowpass(去高频噪声)
|
||||
ffmpeg_cmd = [
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
str(video_path),
|
||||
"-vn", # 不要视频
|
||||
"-af",
|
||||
"highpass=f=80,afftdn=nf=-25:tn=1,lowpass=f=8000",
|
||||
"-acodec",
|
||||
"libmp3lame",
|
||||
"-ab",
|
||||
"192k",
|
||||
"-ar",
|
||||
"44100",
|
||||
"-ac",
|
||||
"1", # 单声道(人声足够)
|
||||
str(audio_path),
|
||||
]
|
||||
|
||||
result = subprocess.run(
|
||||
ffmpeg_cmd,
|
||||
capture_output=True,
|
||||
timeout=300, # 5 分钟超时
|
||||
)
|
||||
|
||||
if result.returncode != 0:
|
||||
stderr_text = result.stderr.decode("utf-8", errors="replace")[-500:]
|
||||
logger.error("ffmpeg 提取配音失败: %s", stderr_text)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="视频音频提取失败,可能该视频没有音轨或格式不支持",
|
||||
)
|
||||
|
||||
if not audio_path.exists() or audio_path.stat().st_size == 0:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="音频提取结果为空",
|
||||
)
|
||||
|
||||
# 获取音频时长
|
||||
duration = _get_audio_duration(audio_path)
|
||||
file_size = audio_path.stat().st_size
|
||||
|
||||
# 上传到 OSS
|
||||
audio_ext = "mp3"
|
||||
storage_key = f"uploads/voice/extracted/{uuid4().hex}.{audio_ext}"
|
||||
storage_service.upload_file(audio_path, storage_key, content_type="audio/mpeg")
|
||||
|
||||
# 创建 Asset 记录
|
||||
original_name = (file.filename or "video").rsplit(".", 1)[0]
|
||||
asset_name = f"{original_name}-配音"
|
||||
|
||||
asset = Asset.create(
|
||||
project_id=library.project_id,
|
||||
library_id=library.id,
|
||||
name=asset_name,
|
||||
storage_key=storage_key,
|
||||
mime_type="audio/mpeg",
|
||||
metadata={
|
||||
"source": "video_extract",
|
||||
"original_video": file.filename or "unknown",
|
||||
},
|
||||
file_size=file_size,
|
||||
duration=duration,
|
||||
status=AssetStatus.READY,
|
||||
classification_status=ClassificationStatus.PENDING,
|
||||
uploaded_by_user_id=user_id,
|
||||
)
|
||||
asset = asset_repository.create(asset)
|
||||
|
||||
return {
|
||||
"id": asset.id,
|
||||
"name": asset.name,
|
||||
"audio_url": sign_url(storage_key),
|
||||
"duration": duration,
|
||||
"file_size": file_size,
|
||||
"status": "completed",
|
||||
"source": "video_extract",
|
||||
}
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except subprocess.TimeoutExpired:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_504_GATEWAY_TIMEOUT,
|
||||
detail="视频处理超时,请尝试较短的视频",
|
||||
) from None
|
||||
except Exception as e:
|
||||
logger.exception("提取视频配音失败: %s", e)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="提取配音失败,请稍后重试",
|
||||
) from e
|
||||
finally:
|
||||
# 清理临时文件
|
||||
if tmp_dir and Path(tmp_dir).exists():
|
||||
shutil.rmtree(tmp_dir, ignore_errors=True)
|
||||
|
||||
|
||||
def _find_or_create_voice_library_for_extract(*, user_id, project_repository, asset_library_repository):
|
||||
"""为用户找到或创建 voice 素材库(与 TTS 保存逻辑一致)。"""
|
||||
projects = project_repository.find_accessible_projects(user_id)
|
||||
if not projects:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="没有可用的项目,请先创建项目",
|
||||
)
|
||||
|
||||
for project in projects:
|
||||
for lib in asset_library_repository.find_by_project(project.id):
|
||||
kind = lib.kind.value if hasattr(lib.kind, "value") else lib.kind
|
||||
if kind == AssetLibraryKind.VOICE.value:
|
||||
return lib
|
||||
|
||||
# 自动创建
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
project = projects[0]
|
||||
library = AssetLibrary.create(
|
||||
project_id=project.id,
|
||||
name="配音素材库",
|
||||
kind=AssetLibraryKind.VOICE,
|
||||
)
|
||||
try:
|
||||
return asset_library_repository.create(library)
|
||||
except IntegrityError:
|
||||
session = getattr(asset_library_repository, "session", None)
|
||||
if session is not None:
|
||||
try:
|
||||
session.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
for lib in asset_library_repository.find_by_project(project.id):
|
||||
kind = lib.kind.value if hasattr(lib.kind, "value") else lib.kind
|
||||
if kind == AssetLibraryKind.VOICE.value:
|
||||
return lib
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="配音素材库创建失败",
|
||||
) from None
|
||||
|
||||
|
||||
def _get_audio_duration(audio_path: Path) -> float:
|
||||
"""用 ffprobe 获取音频时长(秒)。"""
|
||||
try:
|
||||
result = subprocess.run(
|
||||
[
|
||||
"ffprobe",
|
||||
"-v",
|
||||
"quiet",
|
||||
"-show_entries",
|
||||
"format=duration",
|
||||
"-of",
|
||||
"csv=p=0",
|
||||
str(audio_path),
|
||||
],
|
||||
capture_output=True,
|
||||
timeout=10,
|
||||
)
|
||||
if result.returncode == 0 and result.stdout.strip():
|
||||
return float(result.stdout.strip())
|
||||
except (ValueError, subprocess.TimeoutExpired):
|
||||
pass
|
||||
return 0.0
|
||||
|
||||
@@ -234,6 +234,11 @@ class PlanGeneratorService:
|
||||
# 有缓存的素材片段起点从随机镜头段选取,无缓存走随机起点兜底
|
||||
asset_scene_points = self._fetch_asset_scene_points(asset_ids)
|
||||
|
||||
# 正式生成也随机重排片段顺序(降重,默认开启无开关)
|
||||
# smart_match 决定选哪些素材,shuffle 只改变分配到 clips 的顺序
|
||||
asset_ids = list(asset_ids) # 复制避免修改调用方原列表
|
||||
random.shuffle(asset_ids)
|
||||
|
||||
distribute_assets(
|
||||
clips,
|
||||
asset_ids,
|
||||
|
||||
@@ -0,0 +1,174 @@
|
||||
#!/usr/bin/env python3
|
||||
"""存量指纹重建脚本 — 为已有视频生成 video_fingerprint_chunks 分片数据。
|
||||
|
||||
功能:
|
||||
- 查询 generated_videos 中 video_fingerprint IS NOT NULL 但尚无分片数据的视频
|
||||
- 从 OSS 下载视频 → 用新的分片算法重新计算指纹 → 写入分片表
|
||||
- 支持 --dry-run(只打印不写入)和 --batch-size(默认 50)
|
||||
- 幂等:已存在分片数据的视频跳过
|
||||
|
||||
用法:
|
||||
# 预览(不写入)
|
||||
python rebuild_fingerprint_chunks.py --dry-run
|
||||
|
||||
# 执行重建
|
||||
python rebuild_fingerprint_chunks.py --batch-size 50
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
|
||||
# 确保可以 import worker_app 和 packages
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "..", "worker"))
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", ".."))
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
)
|
||||
logger = logging.getLogger("rebuild_fingerprint_chunks")
|
||||
|
||||
|
||||
def find_videos_needing_rebuild(session, batch_size: int) -> list[dict]:
|
||||
"""查询需要重建分片指纹的视频。"""
|
||||
from sqlalchemy import and_
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel, VideoFingerprintChunkModel
|
||||
|
||||
# 有 video_fingerprint 的视频
|
||||
has_fingerprint = GeneratedVideoModel.video_fingerprint.isnot(None)
|
||||
has_fingerprint = and_(has_fingerprint, GeneratedVideoModel.video_fingerprint != "")
|
||||
|
||||
# 排除已有分片数据的视频
|
||||
subq = session.query(VideoFingerprintChunkModel.video_id).distinct().subquery()
|
||||
no_chunks = ~GeneratedVideoModel.id.in_(subq)
|
||||
|
||||
videos = (
|
||||
session.query(GeneratedVideoModel)
|
||||
.filter(and_(has_fingerprint, no_chunks))
|
||||
.order_by(GeneratedVideoModel.generated_at.desc())
|
||||
.limit(batch_size)
|
||||
.all()
|
||||
)
|
||||
|
||||
return [
|
||||
{
|
||||
"id": v.id,
|
||||
"project_id": v.project_id,
|
||||
"user_id": v.user_id or "",
|
||||
"duration": v.duration,
|
||||
}
|
||||
for v in videos
|
||||
]
|
||||
|
||||
|
||||
def rebuild_one(video_info: dict, dry_run: bool = False) -> int:
|
||||
"""重建单个视频的分片数据。返回写入的 chunk 数量。"""
|
||||
from video_processing.dedup import VideoDeduplicator, _save_fingerprint_chunks
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import VideoFingerprintChunkModel
|
||||
from packages.shared.storage import get_storage_service
|
||||
|
||||
video_id = video_info["id"]
|
||||
project_id = video_info["project_id"]
|
||||
user_id = video_info["user_id"]
|
||||
|
||||
if dry_run:
|
||||
logger.info("[DRY-RUN] Would rebuild video %s (project=%s)", video_id, project_id)
|
||||
return 0
|
||||
|
||||
session = SessionLocal()
|
||||
temp_dir = tempfile.mkdtemp()
|
||||
|
||||
try:
|
||||
# 再次检查幂等性
|
||||
existing_count = (
|
||||
session.query(VideoFingerprintChunkModel).filter(VideoFingerprintChunkModel.video_id == video_id).count()
|
||||
)
|
||||
if existing_count > 0:
|
||||
logger.info("Video %s already has %d chunks, skipping", video_id, existing_count)
|
||||
return 0
|
||||
|
||||
# 下载视频
|
||||
storage_service = get_storage_service()
|
||||
local_path = os.path.join(temp_dir, f"{video_id}.mp4")
|
||||
storage_key = f"projects/{project_id}/generated/{video_id}/{video_id}.mp4"
|
||||
storage_service.download_file(storage_key, local_path)
|
||||
|
||||
# 重新计算指纹
|
||||
deduplicator = VideoDeduplicator()
|
||||
fingerprint = deduplicator.compute_fingerprint(local_path)
|
||||
|
||||
# 写入分片表
|
||||
_save_fingerprint_chunks(fingerprint, video_id, project_id, user_id, session)
|
||||
session.commit()
|
||||
|
||||
chunk_count = len(fingerprint.chunks)
|
||||
logger.info("Rebuilt %d chunks for video %s", chunk_count, video_id)
|
||||
return chunk_count
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Failed to rebuild video %s: %s", video_id, e)
|
||||
session.rollback()
|
||||
return -1
|
||||
finally:
|
||||
session.close()
|
||||
import shutil
|
||||
|
||||
shutil.rmtree(temp_dir, ignore_errors=True)
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="存量指纹重建脚本")
|
||||
parser.add_argument("--dry-run", action="store_true", help="只打印不写入")
|
||||
parser.add_argument("--batch-size", type=int, default=50, help="每批处理数量(默认 50)")
|
||||
parser.add_argument("--total-limit", type=int, default=0, help="总处理数量限制(0=不限制)")
|
||||
args = parser.parse_args()
|
||||
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
session = SessionLocal()
|
||||
|
||||
try:
|
||||
videos = find_videos_needing_rebuild(session, args.batch_size)
|
||||
logger.info("Found %d videos needing rebuild", len(videos))
|
||||
|
||||
if args.dry_run:
|
||||
for v in videos:
|
||||
logger.info("[DRY-RUN] Video %s | project=%s | duration=%.1fs", v["id"], v["project_id"], v["duration"])
|
||||
return
|
||||
|
||||
total_chunks = 0
|
||||
processed = 0
|
||||
failed = 0
|
||||
|
||||
for v in videos:
|
||||
if args.total_limit > 0 and processed >= args.total_limit:
|
||||
break
|
||||
|
||||
result = rebuild_one(v, dry_run=False)
|
||||
if result < 0:
|
||||
failed += 1
|
||||
else:
|
||||
total_chunks += result
|
||||
processed += 1
|
||||
|
||||
logger.info(
|
||||
"Rebuild complete: processed=%d, chunks=%d, failed=%d",
|
||||
processed,
|
||||
total_chunks,
|
||||
failed,
|
||||
)
|
||||
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -28,4 +28,5 @@ export {
|
||||
deleteTTSJob,
|
||||
getTtsVoices,
|
||||
previewTts,
|
||||
extractVideoVoice,
|
||||
} from "./jobs"
|
||||
|
||||
@@ -70,3 +70,56 @@ export const previewTts = async (data: TTSPreviewRequest): Promise<TTSPreviewRes
|
||||
const response = await apiClient.post<TTSPreviewResponse>("/tts/preview", data)
|
||||
return response.data
|
||||
}
|
||||
|
||||
/**
|
||||
* 从视频中提取配音(上传视频 → 后端提取人声 → 保存到配音素材库)
|
||||
* 支持 mp4/mov/webm 格式
|
||||
*/
|
||||
export const extractVideoVoice = async (
|
||||
file: File,
|
||||
onProgress?: (percent: number) => void,
|
||||
): Promise<{ asset_id: string; duration: number }> => {
|
||||
const formData = new FormData()
|
||||
formData.append("file", file)
|
||||
|
||||
return new Promise((resolve, reject) => {
|
||||
const xhr = new XMLHttpRequest()
|
||||
xhr.open("POST", "/api/v1/voices/extract-voice")
|
||||
|
||||
// 携带认证 token(从 localStorage 获取,与 apiClient 拦截器一致)
|
||||
const token = localStorage.getItem("access_token")
|
||||
if (token) {
|
||||
xhr.setRequestHeader("Authorization", `Bearer ${token}`)
|
||||
}
|
||||
|
||||
xhr.timeout = 10 * 60 * 1000 // 10 分钟超时
|
||||
|
||||
xhr.upload.onprogress = (e) => {
|
||||
if (e.lengthComputable && onProgress) {
|
||||
onProgress(Math.round((e.loaded / e.total) * 100))
|
||||
}
|
||||
}
|
||||
|
||||
xhr.onload = () => {
|
||||
if (xhr.status >= 200 && xhr.status < 300) {
|
||||
try {
|
||||
resolve(JSON.parse(xhr.responseText))
|
||||
} catch {
|
||||
reject(new Error("服务器返回数据解析失败"))
|
||||
}
|
||||
} else {
|
||||
try {
|
||||
const err = JSON.parse(xhr.responseText)
|
||||
reject(new Error(err.detail || err.message || `提取失败: HTTP ${xhr.status}`))
|
||||
} catch {
|
||||
reject(new Error(`提取失败: HTTP ${xhr.status}`))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
xhr.onerror = () => reject(new Error("网络错误,请检查网络连接"))
|
||||
xhr.ontimeout = () => reject(new Error("上传超时(10分钟),请检查网络或尝试更小的文件"))
|
||||
|
||||
xhr.send(formData)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -297,7 +297,7 @@ const CloneModal: React.FC<CloneModalProps> = ({ open, onClose, onSuccess }) =>
|
||||
buttonSize="sm"
|
||||
onClick={() => {
|
||||
handleClose()
|
||||
navigate("/app/voice-materials")
|
||||
navigate("/app/voices?tab=material&upload=1")
|
||||
}}
|
||||
>
|
||||
去配音库上传
|
||||
|
||||
@@ -134,7 +134,7 @@ const Step5VoiceSelect: React.FC<Step5VoiceSelectProps> = ({
|
||||
|
||||
/** 跳转到配音库上传 */
|
||||
const handleGoToUpload = useCallback(() => {
|
||||
navigate("/app/voices")
|
||||
navigate("/app/voices?tab=material&upload=1")
|
||||
}, [navigate])
|
||||
|
||||
// 加载中状态
|
||||
|
||||
@@ -1,6 +1,11 @@
|
||||
import { useMemo, useEffect } from "react"
|
||||
import { useQuery, useMutation, useQueryClient } from "@tanstack/react-query"
|
||||
import { getAssetsByKind, getAssetLibraries, createAssetLibrary } from "@/api/assets"
|
||||
import {
|
||||
getAssetsByKind,
|
||||
getAssetLibraries,
|
||||
createAssetLibrary,
|
||||
type AssetItem,
|
||||
} from "@/api/assets"
|
||||
import { type VoiceMaterial, mapAssetToMaterial } from "../../types"
|
||||
|
||||
interface UseVoiceMaterialDataOptions {
|
||||
@@ -44,6 +49,15 @@ export function useVoiceMaterialData({ keyword, gender, tagIds }: UseVoiceMateri
|
||||
queryKey: ["assets", "voice", { keyword, gender, tag_ids: tagIds }],
|
||||
queryFn: () => getAssetsByKind("voice", { keyword, gender, tag_ids: tagIds }),
|
||||
staleTime: 30_000,
|
||||
// 列表中存在上传中/处理中素材时每 3s 轮询;全部就绪后自动停止
|
||||
refetchInterval: (query) => {
|
||||
const items = (query.state.data as AssetItem[] | undefined) ?? []
|
||||
const processing = items.some((a) => {
|
||||
const st = a.status ?? ""
|
||||
return st === "uploading" || st === "ingesting" || st === "processing" || st === "pending"
|
||||
})
|
||||
return processing ? 3000 : false
|
||||
},
|
||||
})
|
||||
|
||||
const materials: VoiceMaterial[] = useMemo(() => assets.map(mapAssetToMaterial), [assets])
|
||||
|
||||
@@ -41,7 +41,7 @@ export const mapAssetToMaterial = (asset: AssetItem): VoiceMaterial => {
|
||||
tagIds: Array.isArray(asset.tag_ids) ? asset.tag_ids : [],
|
||||
fileName: asset.storage_key?.split("/").pop() || asset.name,
|
||||
fileSize: asset.file_size || 0,
|
||||
duration: (meta.duration as number) || 0,
|
||||
duration: asset.duration || (meta.duration as number) || 0,
|
||||
mimeType: asset.mime_type || "audio/mpeg",
|
||||
createdAt: asset.created_at || new Date().toISOString(),
|
||||
fileUrl: asset.file_url,
|
||||
|
||||
@@ -14,8 +14,14 @@
|
||||
* 弹窗集合 → components/VoiceModals
|
||||
* Toast 提示 → components/VoiceToasts
|
||||
*/
|
||||
import React, { useCallback, useState } from "react"
|
||||
import { UploadOutlined, AudioOutlined, RobotOutlined } from "@ant-design/icons"
|
||||
import React, { useCallback, useEffect, useState } from "react"
|
||||
import { useSearchParams } from "react-router-dom"
|
||||
import {
|
||||
UploadOutlined,
|
||||
AudioOutlined,
|
||||
RobotOutlined,
|
||||
VideoCameraOutlined,
|
||||
} from "@ant-design/icons"
|
||||
import { Button } from "@/components/ui"
|
||||
import PageHead from "@/components/layout/PageHead"
|
||||
import { type AssetItem } from "@/api/assets"
|
||||
@@ -34,6 +40,8 @@ import { useTtsSynthesize } from "./hooks/useTtsSynthesize"
|
||||
import { useVoiceUpload } from "./hooks/useVoiceUpload"
|
||||
import { useMaterialDelete } from "./hooks/useMaterialDelete"
|
||||
import { useMaterialBatchDelete } from "./hooks/useMaterialBatchDelete"
|
||||
import { useVideoExtract } from "./hooks/useVideoExtract"
|
||||
import VideoExtractModal from "./components/VideoExtractModal"
|
||||
import "./voices.css"
|
||||
|
||||
let toastIdSeq = 0
|
||||
@@ -158,6 +166,32 @@ const VoiceLibrary: React.FC = () => {
|
||||
handleUploadClose,
|
||||
} = useVoiceUpload({ showToast })
|
||||
|
||||
// ── 提取视频配音 ──────────────────────────────────────
|
||||
const {
|
||||
extractOpen,
|
||||
extractFile,
|
||||
extractProgress,
|
||||
isExtracting,
|
||||
setExtractOpen,
|
||||
handleFileSelect: handleExtractFileSelect,
|
||||
handleExtract,
|
||||
handleExtractClose,
|
||||
} = useVideoExtract({ showToast })
|
||||
|
||||
// ── URL 参数自动打开上传弹窗 ────────────────────────────
|
||||
const [searchParams, setSearchParams] = useSearchParams()
|
||||
|
||||
useEffect(() => {
|
||||
if (searchParams.get("upload") === "1") {
|
||||
setActiveTab("material")
|
||||
setUploadOpen(true)
|
||||
// 一次性触发器:清理 upload 参数,避免切换 Tab 时重复触发
|
||||
const next = new URLSearchParams(searchParams)
|
||||
next.delete("upload")
|
||||
setSearchParams(next, { replace: true })
|
||||
}
|
||||
}, [searchParams, setActiveTab, setUploadOpen, setSearchParams])
|
||||
|
||||
// ── 切换 Tab 时停止播放 ───────────────────────────────
|
||||
const handleTabChange = useCallback(
|
||||
(tab: VoiceTabKey) => {
|
||||
@@ -185,6 +219,14 @@ const VoiceLibrary: React.FC = () => {
|
||||
>
|
||||
上传音频
|
||||
</Button>
|
||||
<Button
|
||||
buttonType="primary"
|
||||
buttonSize="sm"
|
||||
icon={<VideoCameraOutlined />}
|
||||
onClick={() => setExtractOpen(true)}
|
||||
>
|
||||
提取视频配音
|
||||
</Button>
|
||||
<Button
|
||||
buttonType="ghost"
|
||||
buttonSize="sm"
|
||||
@@ -285,7 +327,18 @@ const VoiceLibrary: React.FC = () => {
|
||||
/>
|
||||
)}
|
||||
|
||||
{/* ── 弹窗集合 ──────────────────────────────────── */}
|
||||
{/* ── 视频提取配音弹窗 ─────────────────────────────── */}
|
||||
<VideoExtractModal
|
||||
open={extractOpen}
|
||||
file={extractFile}
|
||||
progress={extractProgress}
|
||||
isExtracting={isExtracting}
|
||||
onClose={handleExtractClose}
|
||||
onFileSelect={handleExtractFileSelect}
|
||||
onExtract={handleExtract}
|
||||
/>
|
||||
|
||||
{/* ── 弹窗集合 ─────────────────────────────────── */}
|
||||
<VoiceModals
|
||||
cloneModalOpen={cloneModalOpen}
|
||||
onCloneClose={() => setCloneModalOpen(false)}
|
||||
|
||||
@@ -136,6 +136,9 @@ export const MaterialVoiceTab: React.FC<MaterialVoiceTabProps> = ({
|
||||
const material = mapAssetToMaterial(asset)
|
||||
// duration 优先取顶层(后端从 metadata 提取),兜底 metadata
|
||||
const cardDuration = asset.duration || material.duration || 0
|
||||
// AI 生成素材标识:兼容旧素材(无 source 字段但有 tts_job_id)
|
||||
const meta = asset.metadata as Record<string, unknown>
|
||||
const isAiMaterial = meta?.source === "tts_job" || !!meta?.tts_job_id
|
||||
const isPlaying = playingId === asset.id
|
||||
const isSelected = selectedIds.has(asset.id)
|
||||
// 播放中以 audio 真实时长为准,未播放显示卡片时长
|
||||
@@ -184,6 +187,7 @@ export const MaterialVoiceTab: React.FC<MaterialVoiceTabProps> = ({
|
||||
<div className="xx-voice-info vmat-info">
|
||||
<div className="xx-voice-name" title={asset.name}>
|
||||
{asset.name}
|
||||
{isAiMaterial && <span className="vmat-ai-badge">AI</span>}
|
||||
</div>
|
||||
<div className="xx-voice-subtitle">
|
||||
{asset.file_size ? `${formatFileSize(asset.file_size)}` : "--"}
|
||||
|
||||
@@ -0,0 +1,207 @@
|
||||
import React, { useRef } from "react"
|
||||
import { Modal } from "antd"
|
||||
import { InboxOutlined, CloseOutlined } from "@ant-design/icons"
|
||||
|
||||
interface VideoExtractModalProps {
|
||||
open: boolean
|
||||
file: File | null
|
||||
progress: number | null
|
||||
isExtracting: boolean
|
||||
onClose: () => void
|
||||
onFileSelect: (file: File | null) => void
|
||||
onExtract: () => void
|
||||
}
|
||||
|
||||
const ACCEPT_TYPES = ".mp4,.mov,.webm"
|
||||
|
||||
const VideoExtractModal: React.FC<VideoExtractModalProps> = ({
|
||||
open,
|
||||
file,
|
||||
progress,
|
||||
isExtracting,
|
||||
onClose,
|
||||
onFileSelect,
|
||||
onExtract,
|
||||
}) => {
|
||||
const inputRef = useRef<HTMLInputElement>(null)
|
||||
|
||||
return (
|
||||
<Modal
|
||||
title={<span style={{ fontSize: 16, fontWeight: 600 }}>提取视频配音</span>}
|
||||
open={open}
|
||||
onCancel={() => {
|
||||
if (isExtracting) return
|
||||
onClose()
|
||||
}}
|
||||
footer={null}
|
||||
width={480}
|
||||
maskClosable={!isExtracting}
|
||||
>
|
||||
{!file ? (
|
||||
<div
|
||||
className="vmat-upload-dropzone"
|
||||
onClick={() => inputRef.current?.click()}
|
||||
style={{
|
||||
border: "2px dashed #d9d9d9",
|
||||
borderRadius: 8,
|
||||
padding: "40px 20px",
|
||||
textAlign: "center",
|
||||
cursor: "pointer",
|
||||
transition: "border-color 0.3s",
|
||||
}}
|
||||
onMouseEnter={(e) => (e.currentTarget.style.borderColor = "#7c3aed")}
|
||||
onMouseLeave={(e) => (e.currentTarget.style.borderColor = "#d9d9d9")}
|
||||
>
|
||||
<InboxOutlined style={{ fontSize: 32, color: "#7c3aed", marginBottom: 12 }} />
|
||||
<p style={{ margin: "0 0 8px", fontSize: 14, color: "#333" }}>点击选择视频文件</p>
|
||||
<span style={{ fontSize: 12, color: "#999" }}>支持 MP4、MOV、WebM 格式</span>
|
||||
<input
|
||||
ref={inputRef}
|
||||
type="file"
|
||||
accept={ACCEPT_TYPES}
|
||||
style={{ display: "none" }}
|
||||
onChange={(e) => {
|
||||
const f = e.target.files?.[0]
|
||||
if (f) onFileSelect(f)
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
) : (
|
||||
<div>
|
||||
<div
|
||||
style={{
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "space-between",
|
||||
padding: "12px 16px",
|
||||
background: "#fafafa",
|
||||
borderRadius: 8,
|
||||
marginBottom: 16,
|
||||
}}
|
||||
>
|
||||
<span
|
||||
style={{
|
||||
flex: 1,
|
||||
overflow: "hidden",
|
||||
textOverflow: "ellipsis",
|
||||
whiteSpace: "nowrap",
|
||||
fontSize: 14,
|
||||
fontWeight: 500,
|
||||
}}
|
||||
title={file.name}
|
||||
>
|
||||
{file.name}
|
||||
</span>
|
||||
<span style={{ fontSize: 12, color: "#999", marginLeft: 8, flexShrink: 0 }}>
|
||||
{(file.size / (1024 * 1024)).toFixed(1)} MB
|
||||
</span>
|
||||
{!isExtracting && (
|
||||
<button
|
||||
type="button"
|
||||
onClick={() => {
|
||||
if (inputRef.current) inputRef.current.value = ""
|
||||
onFileSelect(null)
|
||||
}}
|
||||
style={{
|
||||
border: "none",
|
||||
background: "none",
|
||||
cursor: "pointer",
|
||||
color: "#999",
|
||||
marginLeft: 8,
|
||||
fontSize: 14,
|
||||
}}
|
||||
aria-label="移除文件"
|
||||
>
|
||||
<CloseOutlined />
|
||||
</button>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{progress !== null && (
|
||||
<div style={{ marginBottom: 12 }}>
|
||||
<div
|
||||
style={{
|
||||
height: 6,
|
||||
background: "#f0f0f0",
|
||||
borderRadius: 3,
|
||||
overflow: "hidden",
|
||||
}}
|
||||
>
|
||||
<div
|
||||
style={{
|
||||
height: "100%",
|
||||
width: `${progress}%`,
|
||||
background: "linear-gradient(90deg, #7c3aed, #a78bfa)",
|
||||
borderRadius: 3,
|
||||
transition: "width 0.3s",
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
<div
|
||||
style={{
|
||||
textAlign: "right",
|
||||
fontSize: 12,
|
||||
color: "#999",
|
||||
marginTop: 4,
|
||||
}}
|
||||
>
|
||||
{progress}%
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{isExtracting && (
|
||||
<p style={{ textAlign: "center", fontSize: 13, color: "#7c3aed", margin: "12px 0 0" }}>
|
||||
{progress === 100 ? "正在提取人声,请稍候..." : "正在上传视频..."}
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
|
||||
<div
|
||||
style={{
|
||||
display: "flex",
|
||||
justifyContent: "flex-end",
|
||||
gap: 8,
|
||||
marginTop: 24,
|
||||
}}
|
||||
>
|
||||
<button
|
||||
type="button"
|
||||
onClick={onClose}
|
||||
disabled={isExtracting}
|
||||
style={{
|
||||
padding: "6px 16px",
|
||||
borderRadius: 6,
|
||||
border: "1px solid #d9d9d9",
|
||||
background: "#fff",
|
||||
cursor: isExtracting ? "not-allowed" : "pointer",
|
||||
fontSize: 14,
|
||||
opacity: isExtracting ? 0.5 : 1,
|
||||
}}
|
||||
>
|
||||
取消
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
onClick={onExtract}
|
||||
disabled={!file || isExtracting}
|
||||
style={{
|
||||
padding: "6px 16px",
|
||||
borderRadius: 6,
|
||||
border: "none",
|
||||
background: !file || isExtracting ? "#d9d9d9" : "#7c3aed",
|
||||
color: "#fff",
|
||||
cursor: !file || isExtracting ? "not-allowed" : "pointer",
|
||||
fontSize: 14,
|
||||
fontWeight: 500,
|
||||
}}
|
||||
>
|
||||
{isExtracting ? "提取中..." : "开始提取"}
|
||||
</button>
|
||||
</div>
|
||||
</Modal>
|
||||
)
|
||||
}
|
||||
|
||||
export default VideoExtractModal
|
||||
@@ -0,0 +1,74 @@
|
||||
import { useState, useCallback } from "react"
|
||||
import { useQueryClient } from "@tanstack/react-query"
|
||||
import { extractVideoVoice } from "@/api/tts"
|
||||
|
||||
/**
|
||||
* 视频提取配音 Hook
|
||||
* 封装视频上传弹窗状态、提取进度、提取 mutation 逻辑
|
||||
*/
|
||||
interface UseVideoExtractProps {
|
||||
showToast: (message: string, type: "success" | "error") => void
|
||||
}
|
||||
|
||||
export function useVideoExtract({ showToast }: UseVideoExtractProps) {
|
||||
const queryClient = useQueryClient()
|
||||
|
||||
const [extractOpen, setExtractOpen] = useState(false)
|
||||
const [extractFile, setExtractFile] = useState<File | null>(null)
|
||||
const [extractProgress, setExtractProgress] = useState<number | null>(null)
|
||||
const [isExtracting, setIsExtracting] = useState(false)
|
||||
|
||||
const handleExtractClose = useCallback(() => {
|
||||
setExtractOpen(false)
|
||||
setExtractFile(null)
|
||||
setExtractProgress(null)
|
||||
setIsExtracting(false)
|
||||
}, [])
|
||||
|
||||
const handleExtract = useCallback(async () => {
|
||||
if (!extractFile) return
|
||||
setIsExtracting(true)
|
||||
setExtractProgress(0)
|
||||
try {
|
||||
await extractVideoVoice(extractFile, (p) => setExtractProgress(p))
|
||||
// 刷新素材列表
|
||||
queryClient.invalidateQueries({ queryKey: ["assets", "voice"] })
|
||||
queryClient.invalidateQueries({ queryKey: ["voice-materials"] })
|
||||
showToast("视频配音提取成功", "success")
|
||||
handleExtractClose()
|
||||
} catch (err: unknown) {
|
||||
const msg = err instanceof Error ? err.message : "提取失败,请重试"
|
||||
showToast(msg, "error")
|
||||
} finally {
|
||||
setIsExtracting(false)
|
||||
setExtractProgress(null)
|
||||
}
|
||||
}, [extractFile, queryClient, showToast, handleExtractClose])
|
||||
|
||||
const handleFileSelect = useCallback(
|
||||
(file: File | null) => {
|
||||
if (!file) {
|
||||
setExtractFile(null)
|
||||
return
|
||||
}
|
||||
const validTypes = ["video/mp4", "video/quicktime", "video/webm"]
|
||||
if (!validTypes.includes(file.type)) {
|
||||
showToast("仅支持 MP4、MOV、WebM 格式的视频文件", "error")
|
||||
return
|
||||
}
|
||||
setExtractFile(file)
|
||||
},
|
||||
[showToast],
|
||||
)
|
||||
|
||||
return {
|
||||
extractOpen,
|
||||
setExtractOpen,
|
||||
extractFile,
|
||||
extractProgress,
|
||||
isExtracting,
|
||||
handleFileSelect,
|
||||
handleExtract,
|
||||
handleExtractClose,
|
||||
}
|
||||
}
|
||||
@@ -193,6 +193,24 @@
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
flex: 1;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
}
|
||||
|
||||
/* AI 配音标识 */
|
||||
.vmat-ai-badge {
|
||||
display: inline-block;
|
||||
margin-left: 6px;
|
||||
padding: 1px 6px;
|
||||
font-size: 11px;
|
||||
font-weight: 600;
|
||||
color: #7c3aed;
|
||||
background: #f3f0ff;
|
||||
border: 1px solid #ddd6fe;
|
||||
border-radius: 4px;
|
||||
line-height: 16px;
|
||||
vertical-align: middle;
|
||||
flex-shrink: 0;
|
||||
}
|
||||
|
||||
.xx-voice-star {
|
||||
|
||||
@@ -4,8 +4,9 @@ import hashlib
|
||||
import logging
|
||||
import os
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
from uuid import uuid4
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
@@ -15,10 +16,16 @@ from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.generated_video_repository import SQLAlchemyGeneratedVideoRepository
|
||||
from packages.adapters.sqlalchemy_impl.models import VideoFingerprintChunkModel
|
||||
from packages.shared.storage import get_storage_service
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 分片策略常量
|
||||
SHORT_VIDEO_CHUNK_SEC = 2 # ≤60秒视频,每 2 秒一个分片
|
||||
LONG_VIDEO_CHUNK_SEC = 5 # >60秒视频,每 5 秒一个分片
|
||||
SHORT_VIDEO_THRESHOLD_SEC = 60
|
||||
|
||||
|
||||
def compute_phash(image: np.ndarray, hash_size: int = 8) -> str:
|
||||
"""计算图像的感知哈希(pHash),基于 DCT(离散余弦变换)。
|
||||
@@ -80,6 +87,28 @@ def compute_color_histogram(image: np.ndarray, bins: int = 32) -> list[float]:
|
||||
return hist
|
||||
|
||||
|
||||
def compute_chunk_interval(duration: float) -> float:
|
||||
"""根据视频时长返回分片间隔(秒)。
|
||||
|
||||
短视频(≤60秒):每 2 秒一个分片
|
||||
长视频(>60秒):每 5 秒一个分片
|
||||
"""
|
||||
if duration <= SHORT_VIDEO_THRESHOLD_SEC:
|
||||
return SHORT_VIDEO_CHUNK_SEC
|
||||
return LONG_VIDEO_CHUNK_SEC
|
||||
|
||||
|
||||
@dataclass
|
||||
class FingerprintChunk:
|
||||
"""单个分片指纹数据。"""
|
||||
|
||||
start_time_ms: int
|
||||
end_time_ms: int
|
||||
phash_binary: str
|
||||
color_histogram: list[float]
|
||||
frame_count: int = 1
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoFingerprint:
|
||||
"""Video fingerprint containing multiple similarity metrics."""
|
||||
@@ -89,6 +118,7 @@ class VideoFingerprint:
|
||||
color_histograms: list[list[float]]
|
||||
duration: float
|
||||
resolution: tuple[int, int]
|
||||
chunks: list[FingerprintChunk] = field(default_factory=list)
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
# 注意:color_histograms 里的值可能是 np.float32(来自 cv2.normalize),
|
||||
@@ -101,8 +131,37 @@ class VideoFingerprint:
|
||||
"color_histograms": native_histograms,
|
||||
"duration": float(self.duration),
|
||||
"resolution": [int(self.resolution[0]), int(self.resolution[1])],
|
||||
"chunks": [
|
||||
{
|
||||
"start_time_ms": c.start_time_ms,
|
||||
"end_time_ms": c.end_time_ms,
|
||||
"phash_binary": c.phash_binary,
|
||||
"color_histogram": [float(v) for v in c.color_histogram],
|
||||
"frame_count": c.frame_count,
|
||||
}
|
||||
for c in self.chunks
|
||||
],
|
||||
}
|
||||
|
||||
def to_chunk_models(self, video_id: str, project_id: str, user_id: str = "") -> list[VideoFingerprintChunkModel]:
|
||||
"""将分片数据转为 SQLAlchemy Model 列表,用于批量写入 video_fingerprint_chunks 表。"""
|
||||
models = []
|
||||
for chunk in self.chunks:
|
||||
models.append(
|
||||
VideoFingerprintChunkModel(
|
||||
id=uuid4().hex,
|
||||
video_id=video_id,
|
||||
project_id=project_id,
|
||||
user_id=user_id,
|
||||
start_time_ms=chunk.start_time_ms,
|
||||
end_time_ms=chunk.end_time_ms,
|
||||
phash_binary=chunk.phash_binary,
|
||||
color_histogram=[float(v) for v in chunk.color_histogram],
|
||||
frame_count=chunk.frame_count,
|
||||
)
|
||||
)
|
||||
return models
|
||||
|
||||
|
||||
class VideoDeduplicator:
|
||||
"""Video deduplication using multiple fingerprint methods."""
|
||||
@@ -111,7 +170,12 @@ class VideoDeduplicator:
|
||||
HISTOGRAM_THRESHOLD = 0.85
|
||||
|
||||
def compute_fingerprint(self, video_path: str) -> VideoFingerprint:
|
||||
"""Compute video fingerprint using MD5, pHash, and color histogram."""
|
||||
"""Compute video fingerprint using MD5, pHash, and color histogram.
|
||||
|
||||
按时间分片抽帧:短视频(≤60s)每 2s 一片,长视频每 5s 一片。
|
||||
每片取 1 帧计算 pHash + color_histogram。
|
||||
同时保留 keyframe_phashes/color_histograms 聚合字段(向后兼容)。
|
||||
"""
|
||||
cap = cv2.VideoCapture(video_path)
|
||||
if not cap.isOpened():
|
||||
raise RuntimeError(f"Cannot open video: {video_path}")
|
||||
@@ -123,42 +187,82 @@ class VideoDeduplicator:
|
||||
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
|
||||
md5_hash = hashlib.md5(usedforsecurity=False)
|
||||
keyframe_phashes = []
|
||||
color_histograms = []
|
||||
chunks: list[FingerprintChunk] = []
|
||||
|
||||
frame_interval = max(1, frame_count // 10)
|
||||
for i in range(0, frame_count, frame_interval):
|
||||
cap.set(cv2.CAP_PROP_POS_FRAMES, i)
|
||||
# 分片间隔(秒)
|
||||
chunk_interval_sec = compute_chunk_interval(duration)
|
||||
chunk_interval_ms = int(chunk_interval_sec * 1000)
|
||||
duration_ms = int(duration * 1000)
|
||||
|
||||
# 遍历每个分片时间窗口,取 1 帧
|
||||
start_ms = 0
|
||||
while start_ms < duration_ms:
|
||||
end_ms = min(start_ms + chunk_interval_ms, duration_ms)
|
||||
# 定位到分片中点
|
||||
seek_ms = (start_ms + end_ms) / 2
|
||||
cap.set(cv2.CAP_PROP_POS_MSEC, seek_ms)
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
continue
|
||||
if ret:
|
||||
# MD5 计算
|
||||
_, buffer = cv2.imencode(".jpg", frame)
|
||||
md5_hash.update(buffer)
|
||||
|
||||
_, buffer = cv2.imencode(".jpg", frame)
|
||||
md5_hash.update(buffer)
|
||||
phash = compute_phash(frame)
|
||||
hist = compute_color_histogram(frame)
|
||||
|
||||
keyframe_phashes.append(compute_phash(frame))
|
||||
color_histograms.append(compute_color_histogram(frame))
|
||||
chunks.append(
|
||||
FingerprintChunk(
|
||||
start_time_ms=start_ms,
|
||||
end_time_ms=end_ms,
|
||||
phash_binary=phash,
|
||||
color_histogram=hist,
|
||||
frame_count=1,
|
||||
)
|
||||
)
|
||||
|
||||
start_ms = end_ms
|
||||
|
||||
cap.release()
|
||||
|
||||
# 向后兼容:聚合 keyframe_phashes / color_histograms
|
||||
keyframe_phashes = [c.phash_binary for c in chunks]
|
||||
color_histograms = [c.color_histogram for c in chunks]
|
||||
|
||||
return VideoFingerprint(
|
||||
md5=md5_hash.hexdigest(),
|
||||
keyframe_phashes=keyframe_phashes,
|
||||
color_histograms=color_histograms,
|
||||
duration=duration,
|
||||
resolution=(width, height),
|
||||
chunks=chunks,
|
||||
)
|
||||
|
||||
def _get_existing_chunks(self, video_id: str, session: Session) -> list[dict]:
|
||||
"""从 video_fingerprint_chunks 表读取分片数据。返回空列表表示无分片数据。"""
|
||||
rows = (
|
||||
session.query(VideoFingerprintChunkModel)
|
||||
.filter(VideoFingerprintChunkModel.video_id == video_id)
|
||||
.order_by(VideoFingerprintChunkModel.start_time_ms)
|
||||
.all()
|
||||
)
|
||||
return [
|
||||
{
|
||||
"phash_binary": r.phash_binary,
|
||||
"color_histogram": r.color_histogram,
|
||||
"start_time_ms": r.start_time_ms,
|
||||
"end_time_ms": r.end_time_ms,
|
||||
}
|
||||
for r in rows
|
||||
]
|
||||
|
||||
def check_duplicate(self, fingerprint: VideoFingerprint, project_id: str, session: Session) -> Optional[dict]:
|
||||
"""检查视频是否与项目中已有视频重复。
|
||||
|
||||
判定逻辑(按优先级):
|
||||
1. MD5 精确匹配:完全一致则 similarity=1.0,立即返回
|
||||
2. pHash 相似度:计算新视频每帧 phash 与已有视频每帧 phash 的最小汉明距离,
|
||||
取所有帧的平均值 avg_distance。若 avg_distance < PHASH_THRESHOLD(10),
|
||||
则判定为重复,similarity = 1.0 - (avg_distance / 64)
|
||||
查重逻辑:
|
||||
1. MD5 精确匹配 → similarity=1.0
|
||||
2. pHash 相似度(优先从分片表读取,回退到 JSON 字段)
|
||||
|
||||
注意:返回第一个通过阈值的匹配(非最优匹配)。
|
||||
判定阈值:avg_distance < PHASH_THRESHOLD(10)
|
||||
|
||||
Args:
|
||||
fingerprint: 待检测视频的指纹
|
||||
@@ -182,8 +286,15 @@ class VideoDeduplicator:
|
||||
if fingerprint.md5 == ef.get("md5"):
|
||||
return {"duplicate": True, "duplicate_of": existing.id, "reason": "exact_md5_match", "similarity": 1.0}
|
||||
|
||||
# 感知哈希相似度
|
||||
existing_phashes = ef.get("keyframe_phashes", [])
|
||||
# 优先从分片表读取已有视频的分片 phash
|
||||
existing_phashes = []
|
||||
chunk_data = self._get_existing_chunks(existing.id, session)
|
||||
if chunk_data:
|
||||
existing_phashes = [c["phash_binary"] for c in chunk_data]
|
||||
else:
|
||||
# 回退:从 JSON 字段读取(存量旧视频)
|
||||
existing_phashes = ef.get("keyframe_phashes", [])
|
||||
|
||||
if not existing_phashes:
|
||||
continue
|
||||
|
||||
@@ -247,7 +358,14 @@ class VideoDeduplicator:
|
||||
"similarity": 1.0,
|
||||
}
|
||||
|
||||
existing_phashes = ef.get("keyframe_phashes", [])
|
||||
# 优先从分片表读取
|
||||
existing_phashes = []
|
||||
chunk_data = self._get_existing_chunks(existing.id, session)
|
||||
if chunk_data:
|
||||
existing_phashes = [c["phash_binary"] for c in chunk_data]
|
||||
else:
|
||||
existing_phashes = ef.get("keyframe_phashes", [])
|
||||
|
||||
if not existing_phashes:
|
||||
continue
|
||||
|
||||
@@ -372,8 +490,14 @@ class VideoDeduplicator:
|
||||
if fingerprint.md5 == ef.get("md5"):
|
||||
return 100.0
|
||||
|
||||
# pHash 相似度
|
||||
existing_phashes = ef.get("keyframe_phashes", [])
|
||||
# 优先从分片表读取
|
||||
existing_phashes = []
|
||||
chunk_data = self._get_existing_chunks(existing.id, session)
|
||||
if chunk_data:
|
||||
existing_phashes = [c["phash_binary"] for c in chunk_data]
|
||||
else:
|
||||
existing_phashes = ef.get("keyframe_phashes", [])
|
||||
|
||||
if not existing_phashes or not fingerprint.keyframe_phashes:
|
||||
continue
|
||||
|
||||
@@ -388,6 +512,31 @@ class VideoDeduplicator:
|
||||
return round(max(max_similarity, 0.0), 2)
|
||||
|
||||
|
||||
def _save_fingerprint_chunks(
|
||||
fingerprint: VideoFingerprint,
|
||||
video_id: str,
|
||||
project_id: str,
|
||||
user_id: str,
|
||||
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
|
||||
|
||||
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)
|
||||
|
||||
|
||||
@celery_app.task(bind=True, max_retries=3, name="worker.check_duplicate")
|
||||
def check_duplicate_task(self: Task, generated_video_id: str) -> dict:
|
||||
"""Celery task to check if generated video is a duplicate."""
|
||||
@@ -421,6 +570,10 @@ def check_duplicate_task(self: Task, generated_video_id: str) -> dict:
|
||||
video.duplicate_of = None
|
||||
|
||||
video_repo.update(video)
|
||||
|
||||
# 写入分片表
|
||||
_save_fingerprint_chunks(fingerprint, generated_video_id, video.project_id, video.user_id, session)
|
||||
|
||||
session.commit()
|
||||
|
||||
logger.info(f"Duplicate check completed for video {generated_video_id}: is_duplicate={video.is_duplicate}")
|
||||
|
||||
@@ -100,6 +100,14 @@ def create_video_record_and_dedup(
|
||||
|
||||
generated_video.video_fingerprint = fingerprint.to_dict()
|
||||
|
||||
# 写入分片指纹表
|
||||
from video_processing.dedup import _save_fingerprint_chunks
|
||||
|
||||
try:
|
||||
_save_fingerprint_chunks(fingerprint, video_id, project_id, user_id, session)
|
||||
except Exception as chunk_err:
|
||||
logger.warning("Failed to save fingerprint chunks for %s: %s", video_id, chunk_err)
|
||||
|
||||
# (a) 历史成片查重
|
||||
duplicate_result = deduplicator.check_duplicate(fingerprint, project_id, session)
|
||||
|
||||
|
||||
@@ -630,7 +630,7 @@ def ingest_asset(job_id: str) -> dict:
|
||||
name=filename,
|
||||
storage_key=job.storage_key,
|
||||
mime_type=mime_type,
|
||||
metadata={"ingest_error": error_reason},
|
||||
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)),
|
||||
@@ -656,24 +656,57 @@ def ingest_asset(job_id: str) -> dict:
|
||||
"error": error_reason,
|
||||
}
|
||||
|
||||
# Create Asset
|
||||
asset = Asset.create(
|
||||
project_id=job.project_id,
|
||||
library_id=job.library_id,
|
||||
name=filename,
|
||||
storage_key=job.storage_key,
|
||||
mime_type=mime_type,
|
||||
metadata=metadata,
|
||||
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.READY,
|
||||
file_hash=job.file_hash,
|
||||
thumbnail_url=thumbnail_url,
|
||||
)
|
||||
asset_repo.create(asset)
|
||||
# 查找已存在的 Asset 记录(由 API 端在上传完成时立即创建为 PROCESSING 状态)
|
||||
existing_asset = None
|
||||
try:
|
||||
existing_asset = asset_repo.find_by_storage_key(job.storage_key)
|
||||
except Exception:
|
||||
logger.warning("find_by_storage_key not available, trying fallback lookup")
|
||||
|
||||
if existing_asset is None:
|
||||
# 兜底:如果 API 端没有预先创建 Asset(旧版本兼容),则创建新记录
|
||||
logger.info("No pre-created asset found for storage_key=%s, creating new", job.storage_key)
|
||||
metadata["source"] = "upload"
|
||||
asset = Asset.create(
|
||||
project_id=job.project_id,
|
||||
library_id=job.library_id,
|
||||
name=filename,
|
||||
storage_key=job.storage_key,
|
||||
mime_type=mime_type,
|
||||
metadata=metadata,
|
||||
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.READY,
|
||||
file_hash=job.file_hash,
|
||||
thumbnail_url=thumbnail_url,
|
||||
)
|
||||
asset_repo.create(asset)
|
||||
else:
|
||||
# 更新已有的 Asset 记录,补充元数据并将状态改为 READY
|
||||
asset = existing_asset
|
||||
asset.mime_type = mime_type
|
||||
metadata["source"] = "upload"
|
||||
asset.metadata = metadata
|
||||
asset.file_size = int(metadata.get("size_bytes", 0))
|
||||
asset.duration = float(metadata.get("duration", 0))
|
||||
asset.width = int(metadata.get("width", 0))
|
||||
asset.height = int(metadata.get("height", 0))
|
||||
codec_val = metadata.get("codec")
|
||||
if codec_val:
|
||||
asset.codec = str(codec_val)
|
||||
fps_val = metadata.get("fps")
|
||||
if fps_val:
|
||||
try:
|
||||
asset.fps = float(fps_val)
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
asset.status = AssetStatus.READY
|
||||
asset.thumbnail_url = thumbnail_url
|
||||
asset.updated_at = datetime.now(timezone.utc)
|
||||
asset_repo.update(asset)
|
||||
|
||||
# Update job status to COMPLETED
|
||||
job.status = IngestJobStatus.COMPLETED
|
||||
@@ -692,15 +725,37 @@ def ingest_asset(job_id: str) -> dict:
|
||||
db.rollback()
|
||||
logger.error(f"Failed to ingest asset {job_id}: {e}")
|
||||
|
||||
# Update job status to FAILED
|
||||
# Update job status to FAILED and mark pre-created Asset as ERROR
|
||||
try:
|
||||
job_repo = SQLAlchemyIngestJobRepository(db)
|
||||
asset_repo = SQLAlchemyAssetRepository(db)
|
||||
job = job_repo.get(job_id)
|
||||
if job:
|
||||
job.status = IngestJobStatus.FAILED
|
||||
job.error_message = str(e)
|
||||
job.updated_at = datetime.now(timezone.utc)
|
||||
job_repo.update(job)
|
||||
|
||||
# 将上传时创建的占位 Asset(PROCESSING/UPLOADING)标记为 ERROR,
|
||||
# 避免素材永远卡在中间状态
|
||||
try:
|
||||
existing = asset_repo.find_by_storage_key(job.storage_key)
|
||||
if existing and existing.status in (
|
||||
AssetStatus.PROCESSING,
|
||||
AssetStatus.UPLOADING,
|
||||
):
|
||||
existing.status = AssetStatus.ERROR
|
||||
existing.metadata = {**(existing.metadata or {}), "ingest_error": str(e)}
|
||||
existing.updated_at = datetime.now(timezone.utc)
|
||||
asset_repo.update(existing)
|
||||
logger.info(
|
||||
"Marked asset as ERROR due to ingest failure: asset_id=%s job_id=%s",
|
||||
existing.id,
|
||||
job_id,
|
||||
)
|
||||
except Exception as asset_err:
|
||||
logger.warning("Failed to mark asset as ERROR: %s", asset_err)
|
||||
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
|
||||
@@ -175,9 +175,14 @@ services:
|
||||
- xiaoxia-net
|
||||
|
||||
# =========================================
|
||||
# 重要: 生产环境不要添加任何 volume 挂载到 /usr/share/nginx/html
|
||||
# 这会导致静态文件被覆盖,返回 403 错误
|
||||
# Nginx 配置运行时覆盖
|
||||
# 确保容器使用正确环境的 nginx 配置,即使镜像构建时使用了默认配置
|
||||
# 注意: 只覆盖 /etc/nginx/conf.d/default.conf,不挂载 /usr/share/nginx/html
|
||||
# =========================================
|
||||
environment:
|
||||
- NGINX_ENV=${ENV:-staging}
|
||||
volumes:
|
||||
- ./nginx-${ENV:-staging}.conf:/etc/nginx/conf.d/default.conf:ro
|
||||
|
||||
healthcheck:
|
||||
test: ["CMD", "wget", "--spider", "-q", "http://127.0.0.1:80"]
|
||||
|
||||
@@ -127,6 +127,13 @@ class InMemoryAssetRepository:
|
||||
items = [a for a in self._assets.values() if tag_set.issubset(set(a.tag_ids))]
|
||||
return items[skip : skip + limit]
|
||||
|
||||
def find_by_storage_key(self, storage_key: str) -> Asset | None:
|
||||
"""按 storage_key 查找素材。"""
|
||||
for asset in self._assets.values():
|
||||
if asset.storage_key == storage_key:
|
||||
return asset
|
||||
return None
|
||||
|
||||
def find_by_library_and_file_hash(
|
||||
self,
|
||||
library_id: str,
|
||||
|
||||
@@ -426,6 +426,13 @@ class SQLAlchemyAssetRepository:
|
||||
models = self.session.query(AssetModel).filter(AssetModel.id.in_(ids)).offset(skip).limit(limit).all()
|
||||
return [self._to_domain(m) for m in models]
|
||||
|
||||
def find_by_storage_key(self, storage_key: str) -> Asset | None:
|
||||
"""按 storage_key(对应 DB 中的 file_url)查找素材。"""
|
||||
model = self.session.query(AssetModel).filter(AssetModel.file_url == storage_key).first()
|
||||
if model is None:
|
||||
return None
|
||||
return self._to_domain(model)
|
||||
|
||||
def find_by_library_and_file_hash(
|
||||
self,
|
||||
library_id: str,
|
||||
|
||||
@@ -620,3 +620,20 @@ class CoverTemplateModel(Base):
|
||||
config = Column(JSON, nullable=False, default=dict)
|
||||
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
class VideoFingerprintChunkModel(Base):
|
||||
"""分片视频指纹 — 每个视频按时间分片存储 pHash + color_histogram."""
|
||||
|
||||
__tablename__ = "video_fingerprint_chunks"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
video_id = Column(String(36), nullable=False, index=True)
|
||||
project_id = Column(String(36), nullable=False, index=True)
|
||||
user_id = Column(String(36), nullable=False, index=True, default="")
|
||||
start_time_ms = Column(Integer, nullable=False)
|
||||
end_time_ms = Column(Integer, nullable=False)
|
||||
phash_binary = Column(String(16), nullable=False)
|
||||
color_histogram = Column(JSON, nullable=False)
|
||||
frame_count = Column(Integer, nullable=False, default=1)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
@@ -112,6 +112,11 @@ class AssetRepository(ABC):
|
||||
"""查找包含所有指定标签的素材。"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def find_by_storage_key(self, storage_key: str) -> Asset | None:
|
||||
"""按 storage_key 查找素材(用于异步处理时更新已创建的记录)。"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def find_by_library_and_file_hash(
|
||||
self,
|
||||
|
||||
@@ -32,9 +32,9 @@ CONTEXTS=(
|
||||
echo "检查CI Gate统一门禁"
|
||||
echo
|
||||
|
||||
# 等待60秒,给CI启动写status的时间
|
||||
echo "等待60秒让CI启动..."
|
||||
sleep 60
|
||||
# 等待30秒后开始轮询,最多10分钟
|
||||
echo "等待30秒让CI启动..."
|
||||
sleep 30
|
||||
|
||||
# 405计数器(单次运行内重试)
|
||||
MERGE_405_COUNT=0
|
||||
@@ -72,9 +72,9 @@ check_and_merge() {
|
||||
# CI未全绿(pending中)→ 退出,等下次触发
|
||||
if [ "$ALL_SUCCESS" != "true" ]; then
|
||||
echo
|
||||
echo "⏳ CI尚未全绿(仍有pending),退出等待下次触发"
|
||||
echo " (pr-auto-scan每5分钟扫描一次,CI通过后会自动合并)"
|
||||
exit 0
|
||||
echo "⏳ CI尚未全绿(仍有pending),等待重试..."
|
||||
echo " (当前第${attempt}次轮询,最多${MAX_ATTEMPTS}次)"
|
||||
return 1
|
||||
fi
|
||||
|
||||
# CI全绿 → 合并
|
||||
@@ -136,13 +136,28 @@ check_and_merge() {
|
||||
fi
|
||||
}
|
||||
|
||||
# 最多重试3次(用于405重试,非CI轮询)
|
||||
for i in 1 2 3; do
|
||||
# 轮询等待CI就绪+审批完成,最多10分钟(60次x10秒)
|
||||
MAX_ATTEMPTS=60
|
||||
for attempt in $(seq 1 $MAX_ATTEMPTS); do
|
||||
if check_and_merge; then
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# 检查PR是否还open(可能已被手动合并或关闭)
|
||||
PR_STATE=$(curl -s -H "Authorization: token ${MERGE_TOKEN}" \
|
||||
"${GITHUB_API_URL}/repos/${GITHUB_REPOSITORY}/pulls/${PR_NUMBER}" \
|
||||
| python3 -c "import sys,json; print(json.load(sys.stdin).get('state',''))" 2>/dev/null || echo "?")
|
||||
|
||||
if [ "$PR_STATE" != "open" ]; then
|
||||
echo "PR状态为 ${PR_STATE},无需继续等待"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
if [ $attempt -lt $MAX_ATTEMPTS ]; then
|
||||
sleep 10
|
||||
fi
|
||||
done
|
||||
|
||||
echo
|
||||
echo "本次检查未满足合并条件,退出。pr-auto-scan每5分钟会继续扫描。"
|
||||
echo "⏰ 等待10分钟后仍未满足合并条件,退出。pr-auto-scan定时扫描会继续重试。"
|
||||
exit 0
|
||||
|
||||
@@ -59,6 +59,7 @@ REGISTRY_TOKEN="${ACR_PASSWORD:-${REGISTRY_TOKEN:-}}"
|
||||
ENV_FILE="${ENV_FILE:-/var/lib/xiaoxia-saas-production/.env}"
|
||||
GENERATED_DIR="${GENERATED_DIR:-/var/lib/xiaoxia-saas-production/generated}"
|
||||
LEGACY_ASSETS_DIR="${LEGACY_ASSETS_DIR:-/var/lib/xiaoxia-saas-production/legacy-assets}"
|
||||
NGINX_CONF_FILE="${NGINX_CONF_FILE:-/var/lib/xiaoxia-saas-production/nginx-production.conf}"
|
||||
|
||||
SKIP_MIGRATION="${SKIP_MIGRATION:-false}"
|
||||
SKIP_ROLLBACK="${SKIP_ROLLBACK:-false}"
|
||||
@@ -72,6 +73,52 @@ test -f "$ENV_FILE"
|
||||
mkdir -p "$GENERATED_DIR"
|
||||
mkdir -p "$LEGACY_ASSETS_DIR"
|
||||
|
||||
# ── 写入 Production Nginx 配置 ──
|
||||
echo "Writing production nginx config..."
|
||||
cat > "$NGINX_CONF_FILE" << 'NGINX_EOF'
|
||||
server {
|
||||
listen 80;
|
||||
server_name _;
|
||||
root /usr/share/nginx/html;
|
||||
index index.html;
|
||||
|
||||
gzip on;
|
||||
gzip_vary on;
|
||||
gzip_min_length 1024;
|
||||
gzip_types text/plain text/css text/xml text/javascript application/javascript application/json application/xml+rss;
|
||||
|
||||
client_max_body_size 800m;
|
||||
|
||||
location / {
|
||||
try_files $uri /index.html;
|
||||
}
|
||||
|
||||
resolver 127.0.0.11 valid=10s;
|
||||
resolver_timeout 5s;
|
||||
|
||||
location /api/ {
|
||||
proxy_pass http://xiaoxia-api-production:8000/api/;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
proxy_read_timeout 300s;
|
||||
proxy_send_timeout 300s;
|
||||
proxy_request_buffering off;
|
||||
}
|
||||
|
||||
location /generated-files/ {
|
||||
alias /app/generated/;
|
||||
}
|
||||
|
||||
location ~* \.(js|css|png|jpg|jpeg|gif|ico|svg|woff|woff2|ttf|eot)$ {
|
||||
expires 1y;
|
||||
add_header Cache-Control "public, immutable";
|
||||
}
|
||||
}
|
||||
NGINX_EOF
|
||||
echo "✅ Nginx config written: $NGINX_CONF_FILE"
|
||||
|
||||
echo "==========================================="
|
||||
echo " Production 部署 - $IMAGE_TAG"
|
||||
echo "==========================================="
|
||||
@@ -188,6 +235,7 @@ rollback() {
|
||||
--cpus 0.5 \
|
||||
--memory 512m \
|
||||
$LEGACY_VOLUME \
|
||||
-v "$NGINX_CONF_FILE:/etc/nginx/conf.d/default.conf:ro" \
|
||||
--health-cmd "wget --spider -q http://127.0.0.1:80" \
|
||||
--health-interval 30s \
|
||||
--health-timeout 5s \
|
||||
@@ -385,6 +433,7 @@ docker run -d \
|
||||
--restart unless-stopped \
|
||||
--cpus 0.5 \
|
||||
--memory 512m \
|
||||
-v "$NGINX_CONF_FILE:/etc/nginx/conf.d/default.conf:ro" \
|
||||
$LEGACY_VOLUME \
|
||||
--health-cmd "wget --spider -q http://127.0.0.1:80" \
|
||||
--health-interval 30s \
|
||||
|
||||
@@ -160,11 +160,11 @@ health_check() {
|
||||
web_ok=true
|
||||
fi
|
||||
|
||||
# 检查 API docs
|
||||
# 检查 API docs(生产环境禁用 /docs,404 表示 API 在正常响应,视为健康)
|
||||
if [ "$api_docs_ok" = false ]; then
|
||||
HTTP_CODE=$(curl -s -o /dev/null -w "%{http_code}" --max-time 10 "${PROD_API_URL}/docs" 2>/dev/null || echo "000")
|
||||
if [ "$HTTP_CODE" = "200" ]; then
|
||||
log_info "✅ API Docs 检查通过"
|
||||
if [ "$HTTP_CODE" = "200" ] || [ "$HTTP_CODE" = "404" ]; then
|
||||
log_info "✅ API Docs 检查通过(HTTP $HTTP_CODE)"
|
||||
api_docs_ok=true
|
||||
fi
|
||||
fi
|
||||
|
||||
@@ -46,6 +46,7 @@ REGISTRY_TOKEN="${ACR_PASSWORD:-${REGISTRY_TOKEN:-}}"
|
||||
ENV_FILE="${ENV_FILE:-/var/lib/xiaoxia-saas-staging/.env}"
|
||||
GENERATED_DIR="${GENERATED_DIR:-/var/lib/xiaoxia-saas-staging/generated}"
|
||||
LEGACY_ASSETS_DIR="${LEGACY_ASSETS_DIR:-/var/lib/xiaoxia-saas-staging/legacy-assets}"
|
||||
NGINX_CONF_FILE="${NGINX_CONF_FILE:-/var/lib/xiaoxia-saas-staging/nginx-staging.conf}"
|
||||
|
||||
SKIP_MIGRATION="${SKIP_MIGRATION:-false}"
|
||||
SKIP_ROLLBACK="${SKIP_ROLLBACK:-false}"
|
||||
@@ -65,6 +66,53 @@ echo "✅ .env file found: $ENV_FILE ($(wc -l < "$ENV_FILE") lines)"
|
||||
mkdir -p "$GENERATED_DIR"
|
||||
mkdir -p "$LEGACY_ASSETS_DIR"
|
||||
|
||||
# ── 写入 Staging Nginx 配置 ──
|
||||
# 运行时覆盖 nginx 配置,确保 upstream 指向正确的 staging 网络
|
||||
echo "Writing staging nginx config..."
|
||||
cat > "$NGINX_CONF_FILE" << 'NGINX_EOF'
|
||||
server {
|
||||
listen 80;
|
||||
server_name _;
|
||||
root /usr/share/nginx/html;
|
||||
index index.html;
|
||||
|
||||
gzip on;
|
||||
gzip_vary on;
|
||||
gzip_min_length 1024;
|
||||
gzip_types text/plain text/css text/xml text/javascript application/javascript application/json application/xml+rss;
|
||||
|
||||
client_max_body_size 800m;
|
||||
|
||||
location / {
|
||||
try_files $uri /index.html;
|
||||
}
|
||||
|
||||
resolver 127.0.0.11 valid=10s;
|
||||
resolver_timeout 5s;
|
||||
|
||||
location /api/ {
|
||||
proxy_pass http://xiaoxia-api-staging:8000/api/;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
proxy_read_timeout 300s;
|
||||
proxy_send_timeout 300s;
|
||||
proxy_request_buffering off;
|
||||
}
|
||||
|
||||
location /generated-files/ {
|
||||
alias /app/generated/;
|
||||
}
|
||||
|
||||
location ~* \.(js|css|png|jpg|jpeg|gif|ico|svg|woff|woff2|ttf|eot)$ {
|
||||
expires 1y;
|
||||
add_header Cache-Control "public, immutable";
|
||||
}
|
||||
}
|
||||
NGINX_EOF
|
||||
echo "✅ Nginx config written: $NGINX_CONF_FILE"
|
||||
|
||||
echo "==========================================="
|
||||
echo " Staging 部署 - $IMAGE_TAG (并行优化版)"
|
||||
echo "==========================================="
|
||||
@@ -165,6 +213,7 @@ rollback() {
|
||||
-p 127.0.0.1:3001:80 \
|
||||
--restart unless-stopped \
|
||||
$LEGACY_VOLUME \
|
||||
-v "$NGINX_CONF_FILE:/etc/nginx/conf.d/default.conf:ro" \
|
||||
--health-cmd "wget --spider -q http://127.0.0.1:80" \
|
||||
--health-interval 30s \
|
||||
--health-timeout 5s \
|
||||
@@ -467,6 +516,7 @@ docker run -d \
|
||||
-p 127.0.0.1:3001:80 \
|
||||
--restart unless-stopped \
|
||||
$LEGACY_VOLUME \
|
||||
-v "$NGINX_CONF_FILE:/etc/nginx/conf.d/default.conf:ro" \
|
||||
--health-cmd "wget --spider -q http://127.0.0.1:80" \
|
||||
--health-interval 30s \
|
||||
--health-timeout 5s \
|
||||
|
||||
@@ -0,0 +1,313 @@
|
||||
"""分片指纹存储单元测试 — Issue #1657.
|
||||
|
||||
覆盖:
|
||||
- 分片策略:60秒视频 → 30片,120秒视频 → 24片
|
||||
- VideoFingerprint.to_chunk_models() 输出正确
|
||||
- _save_fingerprint_chunks 幂等性(已有数据跳过)
|
||||
- to_dict() 向后兼容
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
|
||||
def _mock_module(**attrs):
|
||||
"""Create a mock module with __spec__ to avoid AttributeError."""
|
||||
m = MagicMock()
|
||||
m.__spec__ = None
|
||||
for k, v in attrs.items():
|
||||
setattr(m, k, v)
|
||||
return m
|
||||
|
||||
|
||||
# ── Module-level setup: mock deps, import dedup, then restore sys.modules ──
|
||||
_SAVED_MODULES_KEYS = set(sys.modules.keys())
|
||||
_SAVED_MODULES_VALUES = {
|
||||
k: sys.modules.get(k)
|
||||
for k in [
|
||||
"cv2",
|
||||
"celery",
|
||||
"sqlalchemy",
|
||||
"sqlalchemy.orm",
|
||||
"sqlalchemy.engine",
|
||||
"sqlalchemy.ext",
|
||||
"sqlalchemy.ext.declarative",
|
||||
"worker_app.db",
|
||||
"worker_app.celery_app",
|
||||
"worker_app.core.config",
|
||||
"packages.adapters.sqlalchemy_impl.session",
|
||||
"packages.adapters.sqlalchemy_impl.generated_video_repository",
|
||||
"packages.adapters.sqlalchemy_impl.models",
|
||||
"packages.shared.config",
|
||||
"packages.shared.storage",
|
||||
]
|
||||
}
|
||||
|
||||
# Set up mocks
|
||||
sys.modules["cv2"] = _mock_module()
|
||||
|
||||
_mock_celery = MagicMock()
|
||||
_mock_celery.Task = MagicMock
|
||||
_mock_celery.Celery = MagicMock
|
||||
_mock_celery.__spec__ = None
|
||||
sys.modules["celery"] = _mock_celery
|
||||
|
||||
_mock_sqla = MagicMock()
|
||||
_mock_sqla.__path__ = []
|
||||
_mock_sqla.__spec__ = None
|
||||
sys.modules["sqlalchemy"] = _mock_sqla
|
||||
|
||||
_mock_sqla_orm = MagicMock()
|
||||
_mock_sqla_orm.__path__ = []
|
||||
_mock_sqla_orm.__spec__ = None
|
||||
_mock_sqla_orm.Session = MagicMock
|
||||
sys.modules["sqlalchemy.orm"] = _mock_sqla_orm
|
||||
sys.modules["sqlalchemy.engine"] = _mock_module()
|
||||
sys.modules["sqlalchemy.ext"] = _mock_module()
|
||||
sys.modules["sqlalchemy.ext.declarative"] = _mock_module()
|
||||
|
||||
sys.modules["worker_app.db"] = _mock_module(SessionLocal=MagicMock())
|
||||
sys.modules["worker_app.celery_app"] = _mock_module(celery_app=MagicMock())
|
||||
sys.modules["worker_app.core.config"] = _mock_module(get_settings=MagicMock(return_value=MagicMock()))
|
||||
|
||||
sys.modules["packages.adapters.sqlalchemy_impl.session"] = _mock_module(
|
||||
Base=MagicMock(),
|
||||
build_engine=MagicMock(),
|
||||
build_session_factory=MagicMock(),
|
||||
ensure_database_exists=MagicMock(),
|
||||
initialize_database=MagicMock(),
|
||||
)
|
||||
sys.modules["packages.adapters.sqlalchemy_impl.generated_video_repository"] = _mock_module()
|
||||
|
||||
|
||||
# Mock VideoFingerprintChunkModel with class-level column attributes
|
||||
class _FakeChunkModel:
|
||||
video_id = MagicMock()
|
||||
project_id = MagicMock()
|
||||
user_id = MagicMock()
|
||||
start_time_ms = MagicMock()
|
||||
end_time_ms = MagicMock()
|
||||
phash_binary = MagicMock()
|
||||
color_histogram = MagicMock()
|
||||
frame_count = MagicMock()
|
||||
created_at = MagicMock()
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
for k, v in kwargs.items():
|
||||
setattr(self, k, v)
|
||||
|
||||
|
||||
sys.modules["packages.adapters.sqlalchemy_impl.models"] = _mock_module(
|
||||
VideoFingerprintChunkModel=_FakeChunkModel,
|
||||
)
|
||||
sys.modules["packages.shared.config"] = _mock_module(get_shared_settings=MagicMock(return_value=MagicMock()))
|
||||
sys.modules["packages.shared.storage"] = _mock_module()
|
||||
|
||||
# Import dedup while mocks are active
|
||||
from video_processing.dedup import ( # noqa: E402
|
||||
FingerprintChunk,
|
||||
VideoFingerprint,
|
||||
_save_fingerprint_chunks,
|
||||
compute_chunk_interval,
|
||||
)
|
||||
|
||||
# ── Restore sys.modules immediately after import ──
|
||||
for _key in list(sys.modules.keys()):
|
||||
if _key not in _SAVED_MODULES_KEYS:
|
||||
del sys.modules[_key]
|
||||
for _key, _value in _SAVED_MODULES_VALUES.items():
|
||||
if _value is not None:
|
||||
sys.modules[_key] = _value
|
||||
elif _key in sys.modules:
|
||||
del sys.modules[_key]
|
||||
del _SAVED_MODULES_KEYS, _SAVED_MODULES_VALUES, _key, _value
|
||||
|
||||
|
||||
class TestChunkInterval:
|
||||
"""测试分片间隔策略。"""
|
||||
|
||||
def test_short_video_interval(self):
|
||||
"""短视频(≤60秒)每 2 秒一个分片。"""
|
||||
assert compute_chunk_interval(0) == 2
|
||||
assert compute_chunk_interval(30) == 2
|
||||
assert compute_chunk_interval(60) == 2
|
||||
|
||||
def test_long_video_interval(self):
|
||||
"""长视频(>60秒)每 5 秒一个分片。"""
|
||||
assert compute_chunk_interval(61) == 5
|
||||
assert compute_chunk_interval(120) == 5
|
||||
assert compute_chunk_interval(300) == 5
|
||||
|
||||
def test_chunk_count_60s_video(self):
|
||||
"""60秒视频 → 30 片(60/2=30)。"""
|
||||
duration = 60
|
||||
interval = compute_chunk_interval(duration)
|
||||
expected_chunks = int(duration / interval)
|
||||
assert expected_chunks == 30
|
||||
|
||||
def test_chunk_count_120s_video(self):
|
||||
"""120秒视频 → 24 片(120/5=24)。"""
|
||||
duration = 120
|
||||
interval = compute_chunk_interval(duration)
|
||||
expected_chunks = int(duration / interval)
|
||||
assert expected_chunks == 24
|
||||
|
||||
|
||||
class TestVideoFingerprintToChunkModels:
|
||||
"""测试 VideoFingerprint.to_chunk_models() 输出。"""
|
||||
|
||||
def test_to_chunk_models_output(self):
|
||||
"""to_chunk_models 返回正确的 Model 列表。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc123",
|
||||
keyframe_phashes=["a1b2", "c3d4"],
|
||||
color_histograms=[[0.1] * 96, [0.2] * 96],
|
||||
duration=10.0,
|
||||
resolution=(1920, 1080),
|
||||
chunks=[
|
||||
FingerprintChunk(start_time_ms=0, end_time_ms=2000, phash_binary="a1b2", color_histogram=[0.1] * 96),
|
||||
FingerprintChunk(start_time_ms=2000, end_time_ms=4000, phash_binary="c3d4", color_histogram=[0.2] * 96),
|
||||
],
|
||||
)
|
||||
|
||||
models = fp.to_chunk_models(video_id="v1", project_id="p1", user_id="u1")
|
||||
|
||||
assert len(models) == 2
|
||||
assert models[0].video_id == "v1"
|
||||
assert models[0].project_id == "p1"
|
||||
assert models[0].user_id == "u1"
|
||||
assert models[0].start_time_ms == 0
|
||||
assert models[0].end_time_ms == 2000
|
||||
assert models[0].phash_binary == "a1b2"
|
||||
assert models[1].start_time_ms == 2000
|
||||
assert models[1].end_time_ms == 4000
|
||||
assert models[1].phash_binary == "c3d4"
|
||||
|
||||
def test_to_chunk_models_empty_chunks(self):
|
||||
"""空 chunks 列表返回空 Model 列表。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc",
|
||||
keyframe_phashes=[],
|
||||
color_histograms=[],
|
||||
duration=0,
|
||||
resolution=(0, 0),
|
||||
chunks=[],
|
||||
)
|
||||
|
||||
models = fp.to_chunk_models(video_id="v1", project_id="p1")
|
||||
assert models == []
|
||||
|
||||
|
||||
class TestSaveFingerprintChunksIdempotent:
|
||||
"""测试 _save_fingerprint_chunks 幂等性。"""
|
||||
|
||||
def test_save_skips_existing(self):
|
||||
"""已有分片数据时跳过写入。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc",
|
||||
keyframe_phashes=["a1b2"],
|
||||
color_histograms=[[0.1] * 96],
|
||||
duration=5.0,
|
||||
resolution=(1920, 1080),
|
||||
chunks=[
|
||||
FingerprintChunk(start_time_ms=0, end_time_ms=2000, phash_binary="a1b2", color_histogram=[0.1] * 96),
|
||||
],
|
||||
)
|
||||
|
||||
session = MagicMock()
|
||||
# Mock: 已有 1 条分片数据
|
||||
session.query.return_value.filter.return_value.count.return_value = 1
|
||||
|
||||
_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()
|
||||
|
||||
def test_save_writes_new(self):
|
||||
"""无分片数据时写入。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc",
|
||||
keyframe_phashes=["a1b2"],
|
||||
color_histograms=[[0.1] * 96],
|
||||
duration=5.0,
|
||||
resolution=(1920, 1080),
|
||||
chunks=[
|
||||
FingerprintChunk(start_time_ms=0, end_time_ms=2000, phash_binary="a1b2", color_histogram=[0.1] * 96),
|
||||
],
|
||||
)
|
||||
|
||||
session = MagicMock()
|
||||
# Mock: 无分片数据
|
||||
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 应被调用一次
|
||||
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_skips_no_chunks(self):
|
||||
"""指纹无 chunks 时跳过。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc",
|
||||
keyframe_phashes=[],
|
||||
color_histograms=[],
|
||||
duration=0,
|
||||
resolution=(0, 0),
|
||||
chunks=[],
|
||||
)
|
||||
|
||||
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 不应被调用
|
||||
session.bulk_save_objects.assert_not_called()
|
||||
|
||||
|
||||
class TestFingerprintToDictBackwardCompat:
|
||||
"""测试 to_dict() 向后兼容性。"""
|
||||
|
||||
def test_to_dict_includes_chunks(self):
|
||||
"""to_dict() 包含 chunks 字段。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc123",
|
||||
keyframe_phashes=["a1b2"],
|
||||
color_histograms=[[0.1] * 96],
|
||||
duration=5.0,
|
||||
resolution=(1920, 1080),
|
||||
chunks=[
|
||||
FingerprintChunk(start_time_ms=0, end_time_ms=2000, phash_binary="a1b2", color_histogram=[0.1] * 96),
|
||||
],
|
||||
)
|
||||
|
||||
d = fp.to_dict()
|
||||
|
||||
assert "chunks" in d
|
||||
assert len(d["chunks"]) == 1
|
||||
assert d["chunks"][0]["start_time_ms"] == 0
|
||||
assert d["chunks"][0]["end_time_ms"] == 2000
|
||||
assert d["chunks"][0]["phash_binary"] == "a1b2"
|
||||
|
||||
def test_to_dict_preserves_legacy_fields(self):
|
||||
"""to_dict() 保留 keyframe_phashes 和 color_histograms 字段。"""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc",
|
||||
keyframe_phashes=["a1b2", "c3d4"],
|
||||
color_histograms=[[0.1] * 96, [0.2] * 96],
|
||||
duration=10.0,
|
||||
resolution=(1920, 1080),
|
||||
)
|
||||
|
||||
d = fp.to_dict()
|
||||
|
||||
assert "keyframe_phashes" in d
|
||||
assert "color_histograms" in d
|
||||
assert len(d["keyframe_phashes"]) == 2
|
||||
assert len(d["color_histograms"]) == 2
|
||||
@@ -104,11 +104,31 @@ def _make_library(
|
||||
return AssetLibrary(id=id, name="Test Library", project_id=project_id, kind=kind)
|
||||
|
||||
|
||||
class StubAssetRepository:
|
||||
"""Minimal asset repository stub for upload tests."""
|
||||
def __init__(self):
|
||||
self._assets = {}
|
||||
|
||||
def create(self, asset):
|
||||
self._assets[asset.id] = asset
|
||||
return asset
|
||||
|
||||
def find_by_storage_key(self, storage_key):
|
||||
for a in self._assets.values():
|
||||
if a.storage_key == storage_key:
|
||||
return a
|
||||
return None
|
||||
|
||||
def find_by_library_and_file_hash(self, library_id, file_hash):
|
||||
return None
|
||||
|
||||
|
||||
def _build_app(
|
||||
project_repo: StubProjectRepository | None = None,
|
||||
library_repo: StubAssetLibraryRepository | None = None,
|
||||
storage: MagicMock | None = None,
|
||||
ingest_repo: StubIngestJobRepository | None = None,
|
||||
asset_repo: StubAssetRepository | None = None,
|
||||
) -> FastAPI:
|
||||
"""构建一个最小化的 FastAPI app,只注册 upload 路由。"""
|
||||
from app.api.routes.upload import router
|
||||
@@ -116,6 +136,7 @@ def _build_app(
|
||||
from app.core.storage import get_storage_service
|
||||
from app.dependencies import (
|
||||
get_asset_library_repository,
|
||||
get_asset_repository,
|
||||
get_ingest_job_repository,
|
||||
get_project_repository,
|
||||
)
|
||||
@@ -129,17 +150,21 @@ def _build_app(
|
||||
storage.is_configured = True
|
||||
storage.upload_file.return_value = "https://bucket.oss.example.com/uploads/test.mp4"
|
||||
ingest_repo = ingest_repo or StubIngestJobRepository()
|
||||
asset_repo = asset_repo or StubAssetRepository()
|
||||
|
||||
# Mock auth
|
||||
mock_user = MagicMock(spec=AuthenticatedUser)
|
||||
mock_user.id = "user-1"
|
||||
mock_user.email = "test@example.com"
|
||||
mock_user.user = MagicMock()
|
||||
mock_user.user.id = "user-1"
|
||||
|
||||
app.dependency_overrides[get_current_user] = lambda: mock_user
|
||||
app.dependency_overrides[get_project_repository] = lambda: project_repo
|
||||
app.dependency_overrides[get_asset_library_repository] = lambda: library_repo
|
||||
app.dependency_overrides[get_storage_service] = lambda: storage
|
||||
app.dependency_overrides[get_ingest_job_repository] = lambda: ingest_repo
|
||||
app.dependency_overrides[get_asset_repository] = lambda: asset_repo
|
||||
|
||||
return app
|
||||
|
||||
|
||||
@@ -1007,3 +1007,162 @@ class TestAssetDurationsAlwaysFetched:
|
||||
call_kwargs = mock_distribute.call_args
|
||||
asset_durations = call_kwargs.kwargs.get("asset_durations", call_kwargs[1].get("asset_durations"))
|
||||
assert asset_durations is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 测试:正式生成片段随机重排(Issue #1663)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFormalGenerationShuffle:
|
||||
"""验证正式生成时片段顺序随机化。
|
||||
|
||||
Issue #1663: 正式生成时 smart_match 排序后对 asset_ids 做 random.shuffle,
|
||||
使得同一批素材每次生成的视频片段顺序不同,有利于查重降重。
|
||||
"""
|
||||
|
||||
def _make_service_with_asset_repo(self):
|
||||
"""创建带 mock asset_repo 的 PlanGeneratorService(复用 TestAssetDurationsAlwaysFetched 模式)"""
|
||||
from apps.api.app.services.plan_generator_service import PlanGeneratorService
|
||||
|
||||
plan_repo = StubEditPlanRepository()
|
||||
clip_repo = StubEditPlanClipRepository()
|
||||
|
||||
asset_repo = MagicMock()
|
||||
|
||||
def fake_get(asset_id):
|
||||
mock_asset = MagicMock()
|
||||
mock_asset.duration = 30.0
|
||||
mock_asset.quality_score = None
|
||||
mock_asset.created_at = None
|
||||
mock_asset.metadata = {}
|
||||
return mock_asset
|
||||
|
||||
asset_repo.get = MagicMock(side_effect=fake_get)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"apps.api.app.services.plan_generator_service.SQLAlchemyEditPlanRepository",
|
||||
return_value=plan_repo,
|
||||
),
|
||||
patch(
|
||||
"apps.api.app.services.plan_generator_service.SQLAlchemyEditPlanClipRepository",
|
||||
return_value=clip_repo,
|
||||
),
|
||||
):
|
||||
db = MagicMock()
|
||||
svc = PlanGeneratorService(db, asset_repo=asset_repo)
|
||||
svc._plan_repo = plan_repo
|
||||
svc._clip_repo = clip_repo
|
||||
|
||||
return svc, asset_repo
|
||||
|
||||
def test_formal_generation_shuffles_asset_ids(self):
|
||||
"""正式生成路径下 asset_ids 应被打乱,多次调用顺序应不同"""
|
||||
svc, _ = self._make_service_with_asset_repo()
|
||||
|
||||
template = _make_template("one_take")
|
||||
# 6 个 clip 容纳 6 个素材
|
||||
clip_configs = _make_clip_configs(
|
||||
template_id=template.id,
|
||||
specs=[
|
||||
{"clip_type": ClipType.MAIN, "order": i, "min_duration": 3.0, "max_duration": 5.0} for i in range(6)
|
||||
],
|
||||
)
|
||||
|
||||
asset_ids = ["a1", "a2", "a3", "a4", "a5", "a6"]
|
||||
|
||||
# 收集多次调用中 distribute_assets 收到的 asset_ids 顺序
|
||||
captured_orders = []
|
||||
with patch(
|
||||
"apps.api.app.services.plan_generator_service.distribute_assets",
|
||||
side_effect=lambda clips, asset_ids, *a, **kw: captured_orders.append(list(asset_ids)),
|
||||
):
|
||||
# mock _sort_assets_by_smart_score 返回固定顺序,验证 shuffle 会打乱
|
||||
with patch.object(
|
||||
svc,
|
||||
"_sort_assets_by_smart_score",
|
||||
side_effect=lambda ids: list(ids), # 原样返回
|
||||
):
|
||||
with patch.object(
|
||||
svc,
|
||||
"_fetch_asset_scene_points",
|
||||
return_value={},
|
||||
):
|
||||
for _ in range(10):
|
||||
svc.generate_from_template(
|
||||
template=template,
|
||||
clip_configs=clip_configs,
|
||||
asset_ids=list(asset_ids), # 每次传新列表
|
||||
random_preview=False, # 正式生成
|
||||
)
|
||||
|
||||
assert len(captured_orders) == 10
|
||||
# 每次 order 应该是 asset_ids 的一个排列
|
||||
expected_set = set(asset_ids)
|
||||
for order in captured_orders:
|
||||
assert set(order) == expected_set
|
||||
|
||||
# 10 次调用中应至少出现 2 种不同顺序(概率 > 99.9%)
|
||||
unique_orders = set(tuple(o) for o in captured_orders)
|
||||
assert (
|
||||
len(unique_orders) >= 2
|
||||
), f"Expected shuffled orders to vary, but got only {len(unique_orders)} unique order(s): {unique_orders}"
|
||||
|
||||
def test_formal_generation_does_not_mutate_original_list(self):
|
||||
"""shuffle 不应修改调用方的原始 asset_ids 列表"""
|
||||
svc, _ = self._make_service_with_asset_repo()
|
||||
|
||||
template = _make_template("one_take")
|
||||
clip_configs = _make_clip_configs(
|
||||
template_id=template.id,
|
||||
specs=[
|
||||
{"clip_type": ClipType.MAIN, "order": i, "min_duration": 3.0, "max_duration": 5.0} for i in range(4)
|
||||
],
|
||||
)
|
||||
|
||||
original = ["a1", "a2", "a3", "a4"]
|
||||
original_copy = list(original)
|
||||
|
||||
with patch("apps.api.app.services.plan_generator_service.distribute_assets"):
|
||||
with patch.object(svc, "_sort_assets_by_smart_score", side_effect=lambda ids: list(ids)):
|
||||
with patch.object(svc, "_fetch_asset_scene_points", return_value={}):
|
||||
svc.generate_from_template(
|
||||
template=template,
|
||||
clip_configs=clip_configs,
|
||||
asset_ids=original,
|
||||
random_preview=False,
|
||||
)
|
||||
|
||||
assert original == original_copy, "Original asset_ids list should not be mutated"
|
||||
|
||||
def test_preview_random_mode_unaffected_by_shuffle(self):
|
||||
"""预览随机模式不走 shuffle 路径,行为不变"""
|
||||
svc, _ = self._make_service_with_asset_repo()
|
||||
|
||||
template = _make_template("one_take")
|
||||
clip_configs = _make_clip_configs(
|
||||
template_id=template.id,
|
||||
specs=[
|
||||
{"clip_type": ClipType.MAIN, "order": i, "min_duration": 3.0, "max_duration": 5.0} for i in range(4)
|
||||
],
|
||||
)
|
||||
|
||||
asset_ids = ["a1", "a2", "a3", "a4"]
|
||||
|
||||
captured_orders = []
|
||||
with patch(
|
||||
"apps.api.app.services.plan_generator_service.distribute_assets",
|
||||
side_effect=lambda clips, asset_ids, *a, **kw: captured_orders.append(list(asset_ids)),
|
||||
):
|
||||
for _ in range(5):
|
||||
svc.generate_from_template(
|
||||
template=template,
|
||||
clip_configs=clip_configs,
|
||||
asset_ids=list(asset_ids),
|
||||
random_preview=True, # 预览随机模式
|
||||
)
|
||||
|
||||
assert len(captured_orders) == 5
|
||||
# 预览模式下 random.shuffle 不应被调用(在 _distribute_assets 的 if not random_selection 块内)
|
||||
# 所以 asset_ids 应该保持调用方传入的顺序(可能已由上层 shuffle 过)
|
||||
|
||||
Reference in New Issue
Block a user