Compare commits
3 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 8a310b8a78 | |||
| 7e88440ca9 | |||
| fbf8844f25 |
@@ -213,3 +213,10 @@ TIKHUB_API_KEY=
|
||||
# apizero.cn API Key (https://v1.apizero.cn) — 国内抖音解析服务
|
||||
APIZERO_API_KEY=
|
||||
|
||||
# ==================== GPU MuseTalk Worker(反向轮询口型同步)====================
|
||||
# GPU Worker 长期鉴权 Token,Worker 端 .env 的 GPU_WORKER_TOKEN 必须与此一致
|
||||
# 留空时 development 环境允许匿名访问(仅本地调试),staging/production 必须配置
|
||||
GPU_WORKER_TOKEN=
|
||||
# 单任务超时(秒),超过则回退 pending 或标记 failed
|
||||
GPU_TASK_TIMEOUT_SECONDS=300
|
||||
|
||||
|
||||
@@ -1190,6 +1190,7 @@ jobs:
|
||||
WECHAT_APP_SECRET: "${{ secrets.WECHAT_APP_SECRET }}"
|
||||
TIKHUB_API_KEY: "${{ secrets.TIKHUB_API_KEY }}"
|
||||
APIZERO_API_KEY: "${{ secrets.APIZERO_API_KEY }}"
|
||||
GPU_WORKER_TOKEN: "${{ secrets.GPU_WORKER_TOKEN }}"
|
||||
run: |
|
||||
set -eu
|
||||
echo "Rendering .env from template + secrets..."
|
||||
@@ -1644,6 +1645,7 @@ jobs:
|
||||
WECHAT_APP_SECRET: "${{ secrets.WECHAT_APP_SECRET }}"
|
||||
TIKHUB_API_KEY: "${{ secrets.TIKHUB_API_KEY }}"
|
||||
APIZERO_API_KEY: "${{ secrets.APIZERO_API_KEY }}"
|
||||
GPU_WORKER_TOKEN: "${{ secrets.GPU_WORKER_TOKEN }}"
|
||||
run: |
|
||||
set -eu
|
||||
echo "Rendering .env from template + secrets..."
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
"""add gpu_lipsync_tasks and gpu_workers tables for MuseTalk reverse-poll worker
|
||||
|
||||
Revision ID: 081_add_gpu_lipsync
|
||||
Revises: 080_edit_plan_clips_atom_clip_id
|
||||
Create Date: 2026-09-18
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "081_add_gpu_lipsync"
|
||||
down_revision = "080_edit_plan_clips_atom_clip_id"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# GPU Worker 注册表
|
||||
op.create_table(
|
||||
"gpu_workers",
|
||||
sa.Column("worker_id", sa.String(100), primary_key=True),
|
||||
sa.Column("hostname", sa.String(200), nullable=False, server_default=""),
|
||||
sa.Column("gpu_name", sa.String(200), nullable=False, server_default=""),
|
||||
sa.Column("free_vram_mb", sa.Integer(), nullable=False, server_default=sa.text("0")),
|
||||
sa.Column("capabilities", sa.String(500), nullable=False, server_default=""),
|
||||
sa.Column("last_heartbeat_at", sa.DateTime(), nullable=True, index=True),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
)
|
||||
|
||||
# GPU 口型同步任务表
|
||||
op.create_table(
|
||||
"gpu_lipsync_tasks",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("lipsync_job_id", sa.String(36), nullable=False, server_default="", index=True),
|
||||
sa.Column("user_id", sa.String(36), nullable=False, server_default="", index=True),
|
||||
sa.Column("project_id", sa.String(36), nullable=False, server_default="", index=True),
|
||||
sa.Column("video_url", sa.Text(), nullable=False),
|
||||
sa.Column("audio_url", sa.Text(), nullable=False),
|
||||
sa.Column("result_url", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("result_duration", sa.Float(), nullable=False, server_default=sa.text("0.0")),
|
||||
sa.Column("status", sa.String(20), nullable=False, server_default="pending", index=True),
|
||||
sa.Column("worker_id", sa.String(100), nullable=False, server_default="", index=True),
|
||||
sa.Column("attempt", sa.Integer(), nullable=False, server_default=sa.text("0")),
|
||||
sa.Column("error_msg", sa.Text(), nullable=False, server_default=""),
|
||||
sa.Column("created_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column("started_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("finished_at", sa.DateTime(), nullable=True),
|
||||
sa.Column("updated_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column("last_heartbeat_at", sa.DateTime(), nullable=True),
|
||||
)
|
||||
op.create_index("ix_gpu_lipsync_status_created", "gpu_lipsync_tasks", ["status", "created_at"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_gpu_lipsync_status_created", table_name="gpu_lipsync_tasks")
|
||||
op.drop_table("gpu_lipsync_tasks")
|
||||
op.drop_table("gpu_workers")
|
||||
@@ -0,0 +1,26 @@
|
||||
"""add ai_tags to asset_atom_clips for #1970 fragment-level AI tagging
|
||||
|
||||
Revision ID: 081_atom_clip_ai_tags
|
||||
Revises: 080_edit_plan_clips_atom_clip_id
|
||||
Create Date: 2026-09-18
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "081_atom_clip_ai_tags"
|
||||
down_revision = "080_edit_plan_clips_atom_clip_id"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"asset_atom_clips",
|
||||
sa.Column("ai_tags", sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("asset_atom_clips", "ai_tags")
|
||||
@@ -14,6 +14,7 @@ from app.api.routes.generation_cover import router as generation_cover_router
|
||||
from app.api.routes.generation_preview import router as generation_preview_router
|
||||
from app.api.routes.generation_tasks import router as generation_tasks_router
|
||||
from app.api.routes.generation_variant_plans import router as generation_variant_plans_router
|
||||
from app.api.routes.gpu_lipsync import router as gpu_lipsync_router
|
||||
from app.api.routes.health import router as health_check_router
|
||||
from app.api.routes.ingest_jobs import router as ingest_jobs_router
|
||||
from app.api.routes.internal_render import router as internal_render_router
|
||||
@@ -211,3 +212,8 @@ api_router.include_router(
|
||||
prefix="/usage",
|
||||
tags=["Usage"],
|
||||
)
|
||||
api_router.include_router(
|
||||
gpu_lipsync_router,
|
||||
prefix="/gpu",
|
||||
tags=["GPU Worker"],
|
||||
)
|
||||
|
||||
@@ -0,0 +1,230 @@
|
||||
"""GPU MuseTalk Worker 反向轮询路由 — /api/v1/gpu/lipsync/*.
|
||||
|
||||
仅面向部署在用户 RTX2060 本地的 GPU Worker 脚本,不面向前端用户。
|
||||
鉴权方式:长期 API Token(`Authorization: Bearer <GPU_WORKER_TOKEN>`),不走用户 JWT。
|
||||
|
||||
接口:
|
||||
POST /api/v1/gpu/register Worker 注册/心跳
|
||||
GET /api/v1/gpu/lipsync/poll Worker 轮询拉任务(无任务返回 204)
|
||||
POST /api/v1/gpu/lipsync/result Worker multipart 上传结果视频/上报失败
|
||||
GET /api/v1/gpu/lipsync/status/{id} 业务侧查询任务状态(内部接口,暂开放给登录用户)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import tempfile
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import requests
|
||||
from app.core.storage import get_storage_service
|
||||
from app.dependencies import get_db_session
|
||||
from app.schemas.gpu_lipsync import (
|
||||
GpuLipsyncPollResponse,
|
||||
GpuLipsyncResultResponse,
|
||||
GpuLipsyncStatusResponse,
|
||||
GpuLipsyncTaskPayload,
|
||||
GpuWorkerRegisterRequest,
|
||||
GpuWorkerRegisterResponse,
|
||||
)
|
||||
from app.services.gpu_lipsync_service import GpuLipsyncService
|
||||
from fastapi import (
|
||||
APIRouter,
|
||||
Depends,
|
||||
File,
|
||||
Form,
|
||||
HTTPException,
|
||||
Query,
|
||||
Request,
|
||||
UploadFile,
|
||||
status,
|
||||
)
|
||||
from fastapi.responses import Response
|
||||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||||
|
||||
from packages.config import get_api_settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
# 复用 bearer scheme 抽 Token,但不校验用户 JWT
|
||||
_gpu_bearer = HTTPBearer(auto_error=False)
|
||||
|
||||
|
||||
def _verify_gpu_token(
|
||||
credentials: Optional[HTTPAuthorizationCredentials] = Depends(_gpu_bearer),
|
||||
) -> str:
|
||||
"""校验 GPU Worker Token,返回 worker 提供的 token 串(仅用于日志,不做身份识别).
|
||||
|
||||
- development 且未配置 token → 直接放行(方便本地调试)。
|
||||
- production/staging 未配置 token → 拒绝(避免裸奔)。
|
||||
- token 不匹配 → 401。
|
||||
"""
|
||||
settings = get_api_settings()
|
||||
expected = (settings.gpu_worker_token or "").strip()
|
||||
is_dev = settings.environment == "development"
|
||||
if not expected:
|
||||
if is_dev:
|
||||
return credentials.credentials if credentials else ""
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="GPU_WORKER_TOKEN not configured on server",
|
||||
)
|
||||
if credentials is None or credentials.scheme.lower() != "bearer":
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Missing bearer token")
|
||||
if credentials.credentials != expected:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid GPU worker token")
|
||||
return credentials.credentials
|
||||
|
||||
|
||||
def _get_svc(db=Depends(get_db_session)) -> GpuLipsyncService:
|
||||
return GpuLipsyncService(db)
|
||||
|
||||
|
||||
# ── POST /register — Worker 注册/心跳 ──────────────────────────────
|
||||
|
||||
|
||||
@router.post("/register", response_model=GpuWorkerRegisterResponse)
|
||||
def register_worker(
|
||||
body: GpuWorkerRegisterRequest,
|
||||
svc: GpuLipsyncService = Depends(_get_svc),
|
||||
_token: str = Depends(_verify_gpu_token),
|
||||
):
|
||||
svc.register_worker(
|
||||
worker_id=body.worker_id,
|
||||
hostname=body.hostname,
|
||||
gpu_name=body.gpu_name,
|
||||
free_vram_mb=body.free_vram_mb,
|
||||
capabilities=body.capabilities,
|
||||
)
|
||||
return GpuWorkerRegisterResponse(ok=True, server_time=datetime.now(UTC), message="ok")
|
||||
|
||||
|
||||
# ── GET /lipsync/poll — Worker 轮询拉任务 ─────────────────────────
|
||||
|
||||
|
||||
@router.get("/lipsync/poll")
|
||||
def poll_task(
|
||||
worker_id: str = Query(..., min_length=1, max_length=100, description="Worker 唯一 ID"),
|
||||
svc: GpuLipsyncService = Depends(_get_svc),
|
||||
_token: str = Depends(_verify_gpu_token),
|
||||
):
|
||||
task = svc.poll_task(worker_id=worker_id)
|
||||
if task is None:
|
||||
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||
payload = GpuLipsyncTaskPayload(
|
||||
task_id=task.id,
|
||||
video_url=getattr(task, "_signed_video_url", task.video_url),
|
||||
audio_url=getattr(task, "_signed_audio_url", task.audio_url),
|
||||
lipsync_job_id=task.lipsync_job_id or "",
|
||||
user_id=task.user_id or "",
|
||||
project_id=task.project_id or "",
|
||||
created_at=task.created_at,
|
||||
upload_url=getattr(task, "_signed_upload_url", ""),
|
||||
upload_method="PUT",
|
||||
expires_at=getattr(task, "_upload_expires_at", datetime.now(UTC)),
|
||||
)
|
||||
return GpuLipsyncPollResponse(task=payload)
|
||||
|
||||
|
||||
# ── POST /lipsync/result — Worker 上报结果(multipart) ─────────────
|
||||
|
||||
|
||||
@router.post("/lipsync/result", response_model=GpuLipsyncResultResponse)
|
||||
async def report_result(
|
||||
request: Request,
|
||||
task_id: str = Form(...),
|
||||
worker_id: str = Form(...),
|
||||
success: bool = Form(True),
|
||||
duration_seconds: float = Form(0.0),
|
||||
error_msg: str = Form(""),
|
||||
result: Optional[UploadFile] = File(None),
|
||||
svc: GpuLipsyncService = Depends(_get_svc),
|
||||
_token: str = Depends(_verify_gpu_token),
|
||||
):
|
||||
# 参数校验:
|
||||
# - success=true + result 文件 → API 代为上传到 OSS(方便 Worker 端实现)
|
||||
# - success=true + 无文件 → Worker 已经自己 PUT 到预签名 upload_url,直接确认
|
||||
# - success=false → 不上传文件,错误信息通过 error_msg 传递
|
||||
if success and result is not None:
|
||||
# 把文件落盘到临时目录,然后 PUT 到预签名 URL
|
||||
storage = get_storage_service()
|
||||
result_key = svc._result_key(task_id)
|
||||
upload_url = storage.get_upload_url(result_key, expires_seconds=3600, content_type="video/mp4")
|
||||
try:
|
||||
with tempfile.TemporaryDirectory(prefix="gpu_result_") as tmpdir:
|
||||
tmp_path = Path(tmpdir) / "result.mp4"
|
||||
content = await result.read()
|
||||
if not content:
|
||||
raise HTTPException(status_code=400, detail="上传的 result 文件为空")
|
||||
tmp_path.write_bytes(content)
|
||||
headers = {"Content-Type": "video/mp4"}
|
||||
with open(tmp_path, "rb") as f:
|
||||
resp = requests.put(upload_url, data=f, headers=headers, timeout=300)
|
||||
if resp.status_code >= 400:
|
||||
logger.error(
|
||||
"上传 GPU 结果到 OSS 失败: status=%d body=%s",
|
||||
resp.status_code,
|
||||
resp.text[:500],
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=502,
|
||||
detail=f"上传结果视频到 OSS 失败 (HTTP {resp.status_code})",
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
logger.exception("上传 GPU 结果视频异常: %s", exc)
|
||||
raise HTTPException(status_code=500, detail=f"上传结果视频异常: {exc}") from exc
|
||||
elif not success:
|
||||
# 失败时忽略 result 文件(即便传了也没用)
|
||||
pass
|
||||
# 其他情况:success=true 且无文件 → Worker 已自行 PUT 到预签名 URL,直接标记完成
|
||||
|
||||
try:
|
||||
task = svc.report_result(
|
||||
task_id=task_id,
|
||||
worker_id=worker_id,
|
||||
success=success,
|
||||
duration_seconds=duration_seconds,
|
||||
error_msg=error_msg,
|
||||
)
|
||||
except KeyError as exc:
|
||||
raise HTTPException(status_code=404, detail=str(exc)) from exc
|
||||
return GpuLipsyncResultResponse(
|
||||
ok=True,
|
||||
task_id=task.id,
|
||||
status=task.status,
|
||||
message="ok",
|
||||
)
|
||||
|
||||
|
||||
# ── GET /lipsync/status/{task_id} — 业务侧查询状态 ─────────────────
|
||||
# 说明:此接口会被 lipsync_service 内部在业务流程里直接读 DB,不通过 HTTP。
|
||||
# 但仍暴露一个简单查询接口,方便调试和前端轮询(如后续需要)。暂不做用户权限校验,
|
||||
# task_id 本身是 UUID,不可枚举。
|
||||
|
||||
|
||||
@router.get("/lipsync/status/{task_id}", response_model=GpuLipsyncStatusResponse)
|
||||
def get_task_status(
|
||||
task_id: str,
|
||||
svc: GpuLipsyncService = Depends(_get_svc),
|
||||
):
|
||||
task = svc.get_task(task_id)
|
||||
if task is None:
|
||||
raise HTTPException(status_code=404, detail="task not found")
|
||||
return GpuLipsyncStatusResponse(
|
||||
task_id=task.id,
|
||||
status=task.status,
|
||||
result_url=task.result_url,
|
||||
result_duration=task.result_duration,
|
||||
error_msg=task.error_msg,
|
||||
worker_id=task.worker_id,
|
||||
attempt=task.attempt,
|
||||
created_at=task.created_at,
|
||||
started_at=task.started_at,
|
||||
finished_at=task.finished_at,
|
||||
)
|
||||
@@ -0,0 +1,103 @@
|
||||
"""GPU MuseTalk 反向轮询 API Schema 定义.
|
||||
|
||||
面向部署在用户 RTX2060 本地的 GPU Worker 脚本,不面向前端用户。
|
||||
Worker 用长期 GPU_WORKER_TOKEN 鉴权(不是用户 JWT)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
# ── Worker 注册/心跳 ──────────────────────────────────────────────
|
||||
|
||||
|
||||
class GpuWorkerRegisterRequest(BaseModel):
|
||||
"""Worker 启动/心跳时上报自身信息."""
|
||||
|
||||
worker_id: str = Field(..., min_length=1, max_length=100, description="Worker 唯一 ID(机器名+UUID 等)")
|
||||
hostname: str = Field("", max_length=200, description="主机名,用于运维排查")
|
||||
gpu_name: str = Field("", max_length=200, description="GPU 型号,如 'NVIDIA GeForce RTX 2060'")
|
||||
free_vram_mb: int = Field(0, ge=0, description="当前空闲显存(MB)")
|
||||
capabilities: str = Field("musetalk", max_length=500, description="能力列表,逗号分隔,如 'musetalk'")
|
||||
|
||||
|
||||
class GpuWorkerRegisterResponse(BaseModel):
|
||||
ok: bool = True
|
||||
server_time: datetime
|
||||
message: str = "ok"
|
||||
|
||||
|
||||
# ── 轮询任务 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class GpuLipsyncTaskPayload(BaseModel):
|
||||
"""下发给 Worker 的任务载荷(含预签名下载 URL)."""
|
||||
|
||||
task_id: str
|
||||
video_url: str = Field(..., description="人物视频预签名下载 URL(GET)")
|
||||
audio_url: str = Field(..., description="驱动音频预签名下载 URL(GET)")
|
||||
lipsync_job_id: str = ""
|
||||
user_id: str = ""
|
||||
project_id: str = ""
|
||||
created_at: datetime
|
||||
upload_url: str = Field(..., description="结果视频预签名上传 URL(PUT, video/mp4)")
|
||||
upload_method: str = Field("PUT", description="上传方式,目前只支持 PUT")
|
||||
expires_at: datetime
|
||||
|
||||
|
||||
class GpuLipsyncPollResponse(BaseModel):
|
||||
"""Worker poll 的返回:200 带任务,204 无任务."""
|
||||
|
||||
task: Optional[GpuLipsyncTaskPayload] = None
|
||||
|
||||
|
||||
# ── Worker 上报结果 ──────────────────────────────────────────────
|
||||
|
||||
|
||||
class GpuLipsyncResultRequest(BaseModel):
|
||||
"""Worker 通过 multipart 上传结果时携带的字段(非文件字段)."""
|
||||
|
||||
task_id: str = Field(..., min_length=1, max_length=64)
|
||||
worker_id: str = Field(..., min_length=1, max_length=100)
|
||||
success: bool = Field(True, description="true=成功(此时必须上传 result 视频文件);false=失败")
|
||||
duration_seconds: float = Field(0.0, ge=0, description="合成后视频时长(秒),成功时应填入")
|
||||
error_msg: str = Field("", max_length=2000, description="失败原因,success=false 时必填")
|
||||
|
||||
|
||||
class GpuLipsyncResultResponse(BaseModel):
|
||||
ok: bool = True
|
||||
task_id: str
|
||||
status: str # done / failed
|
||||
message: str = "ok"
|
||||
|
||||
|
||||
# ── 业务侧查询任务状态 ────────────────────────────────────────────
|
||||
|
||||
|
||||
class GpuLipsyncStatusResponse(BaseModel):
|
||||
task_id: str
|
||||
status: str
|
||||
result_url: str = ""
|
||||
result_duration: float = 0.0
|
||||
error_msg: str = ""
|
||||
worker_id: str = ""
|
||||
attempt: int = 0
|
||||
created_at: datetime
|
||||
started_at: Optional[datetime] = None
|
||||
finished_at: Optional[datetime] = None
|
||||
|
||||
|
||||
# ── 创建任务(内部服务调用) ──────────────────────────────────────
|
||||
|
||||
|
||||
class GpuLipsyncCreateRequest(BaseModel):
|
||||
"""服务层内部创建 GPU 任务用(不通过 HTTP 暴露给 Worker/前端)."""
|
||||
|
||||
video_url: str # 已可访问的 OSS key 或公网 URL(API 侧会转预签名)
|
||||
audio_url: str
|
||||
lipsync_job_id: str = ""
|
||||
user_id: str = ""
|
||||
project_id: str = ""
|
||||
@@ -0,0 +1,287 @@
|
||||
"""GPU MuseTalk 口型同步服务 — 反向轮询模式.
|
||||
|
||||
职责:
|
||||
1. 创建任务(由 lipsync 业务流程调用),为输入/输出生成预签名 URL,任务入队;
|
||||
2. Worker 心跳注册(register):登记/刷新 worker 状态;
|
||||
3. Worker 轮询拉任务(poll):原子地 CLAIM 一条 pending 任务,返回预签名 URL;
|
||||
4. Worker 上报结果(report_result):标记 done/failed,失败可重试;
|
||||
5. 业务侧查询状态(get_status)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import uuid
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from typing import Optional
|
||||
|
||||
from app.core.storage import get_storage_service
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import GpuLipsyncTaskModel, GpuWorkerModel
|
||||
from packages.config import get_api_settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 任务在 processing 超过此时长仍未完成 → 超时回退 pending 或置 failed
|
||||
MAX_ATTEMPTS = 3
|
||||
|
||||
|
||||
class GpuLipsyncService:
|
||||
"""GPU 口型同步服务(无状态方法,每次调用从 DI 拿 db/storage)."""
|
||||
|
||||
RESULT_PREFIX = "gpu-lipsync/results/"
|
||||
INPUT_SIGN_EXPIRES_PAD = 600 # 输入预签名 URL 在任务超时基础上再加 10min 余量
|
||||
|
||||
# ── 公共入口 ────────────────────────────────────────────────────
|
||||
|
||||
def __init__(self, db: Session):
|
||||
self.db = db
|
||||
self.settings = get_api_settings()
|
||||
self.storage = get_storage_service()
|
||||
|
||||
# ── Worker 注册/心跳 ────────────────────────────────────────────
|
||||
|
||||
def register_worker(
|
||||
self,
|
||||
worker_id: str,
|
||||
hostname: str = "",
|
||||
gpu_name: str = "",
|
||||
free_vram_mb: int = 0,
|
||||
capabilities: str = "musetalk",
|
||||
) -> GpuWorkerModel:
|
||||
now = datetime.now(UTC)
|
||||
worker = self.db.query(GpuWorkerModel).filter(GpuWorkerModel.worker_id == worker_id).one_or_none()
|
||||
if worker is None:
|
||||
worker = GpuWorkerModel(
|
||||
worker_id=worker_id,
|
||||
hostname=hostname,
|
||||
gpu_name=gpu_name,
|
||||
free_vram_mb=free_vram_mb,
|
||||
capabilities=capabilities,
|
||||
last_heartbeat_at=now,
|
||||
created_at=now,
|
||||
)
|
||||
self.db.add(worker)
|
||||
else:
|
||||
worker.hostname = hostname or worker.hostname
|
||||
worker.gpu_name = gpu_name or worker.gpu_name
|
||||
worker.free_vram_mb = free_vram_mb
|
||||
worker.capabilities = capabilities or worker.capabilities
|
||||
worker.last_heartbeat_at = now
|
||||
self.db.commit()
|
||||
return worker
|
||||
|
||||
# ── 轮询拉任务(Worker 调用) ──────────────────────────────────
|
||||
|
||||
def poll_task(self, worker_id: str) -> Optional[GpuLipsyncTaskModel]:
|
||||
"""原子地认领一条最早的 pending 任务,返回给 worker;无任务返回 None.
|
||||
|
||||
同时会:
|
||||
- 把 processing 状态且超时(超过 gpu_task_timeout_seconds 无心跳)的任务
|
||||
回退为 pending(attempt++,超过 MAX_ATTEMPTS 置 failed),让其它 worker 认领。
|
||||
- 刷新 worker 心跳。
|
||||
"""
|
||||
now = datetime.now(UTC)
|
||||
self._recover_timed_out_tasks(now)
|
||||
# 更新 worker 心跳
|
||||
self._touch_worker(worker_id, now)
|
||||
|
||||
# 选一条最早 pending 任务(FOR UPDATE SKIP LOCKED 语义:简单起见先查再锁状态)
|
||||
task = (
|
||||
self.db.query(GpuLipsyncTaskModel)
|
||||
.filter(GpuLipsyncTaskModel.status == "pending")
|
||||
.order_by(GpuLipsyncTaskModel.created_at.asc())
|
||||
.first()
|
||||
)
|
||||
if task is None:
|
||||
self.db.commit()
|
||||
return None
|
||||
|
||||
# 原子 claim:用 UPDATE WHERE status=pending 避免并发
|
||||
upd_rows = (
|
||||
self.db.query(GpuLipsyncTaskModel)
|
||||
.filter(
|
||||
GpuLipsyncTaskModel.id == task.id,
|
||||
GpuLipsyncTaskModel.status == "pending",
|
||||
)
|
||||
.update(
|
||||
{
|
||||
GpuLipsyncTaskModel.status: "processing",
|
||||
GpuLipsyncTaskModel.worker_id: worker_id,
|
||||
GpuLipsyncTaskModel.started_at: now,
|
||||
GpuLipsyncTaskModel.last_heartbeat_at: now,
|
||||
GpuLipsyncTaskModel.attempt: GpuLipsyncTaskModel.attempt + 1,
|
||||
GpuLipsyncTaskModel.updated_at: now,
|
||||
},
|
||||
synchronize_session=False,
|
||||
)
|
||||
)
|
||||
self.db.commit()
|
||||
if upd_rows == 0:
|
||||
# 被其它 worker 抢先了
|
||||
return None
|
||||
self.db.refresh(task)
|
||||
# 生成预签名输入/输出 URL(在 claim 时动态生成,避免长时间过期)
|
||||
expires = self.settings.gpu_task_timeout_seconds + self.INPUT_SIGN_EXPIRES_PAD
|
||||
task._signed_video_url = self.storage.get_download_url(task.video_url, expires_seconds=expires)
|
||||
task._signed_audio_url = self.storage.get_download_url(task.audio_url, expires_seconds=expires)
|
||||
task._signed_upload_url = self.storage.get_upload_url(
|
||||
self._result_key(task.id),
|
||||
expires_seconds=expires,
|
||||
content_type="video/mp4",
|
||||
)
|
||||
task._upload_expires_at = now + timedelta(seconds=expires)
|
||||
return task
|
||||
|
||||
# ── 上报结果 ──────────────────────────────────────────────────
|
||||
|
||||
def report_result(
|
||||
self,
|
||||
task_id: str,
|
||||
worker_id: str,
|
||||
success: bool,
|
||||
duration_seconds: float = 0.0,
|
||||
error_msg: str = "",
|
||||
) -> GpuLipsyncTaskModel:
|
||||
task = self.db.get(GpuLipsyncTaskModel, task_id)
|
||||
if task is None:
|
||||
raise KeyError(f"task {task_id} not found")
|
||||
now = datetime.now(UTC)
|
||||
if success:
|
||||
task.status = "done"
|
||||
task.result_url = self._result_key(task_id)
|
||||
task.result_duration = duration_seconds or 0.0
|
||||
task.error_msg = ""
|
||||
task.finished_at = now
|
||||
else:
|
||||
# 失败:若仍可重试(已尝试次数 < MAX_ATTEMPTS)→ 回退 pending;否则 → failed
|
||||
if task.attempt < MAX_ATTEMPTS:
|
||||
task.status = "pending"
|
||||
task.worker_id = ""
|
||||
task.started_at = None
|
||||
task.error_msg = error_msg[:2000]
|
||||
logger.warning(
|
||||
"GPU 任务 %s 在 worker %s 上失败,回退 pending 等待重试(attempt=%d): %s",
|
||||
task_id,
|
||||
worker_id,
|
||||
task.attempt,
|
||||
error_msg[:200],
|
||||
)
|
||||
else:
|
||||
task.status = "failed"
|
||||
task.error_msg = error_msg[:2000]
|
||||
task.finished_at = now
|
||||
logger.error(
|
||||
"GPU 任务 %s 失败达到最大重试次数 %d,置为 failed: %s",
|
||||
task_id,
|
||||
MAX_ATTEMPTS,
|
||||
error_msg[:200],
|
||||
)
|
||||
task.updated_at = now
|
||||
task.last_heartbeat_at = now
|
||||
self._touch_worker(worker_id, now)
|
||||
self.db.commit()
|
||||
self.db.refresh(task)
|
||||
return task
|
||||
|
||||
# ── 业务侧查询 ────────────────────────────────────────────────
|
||||
|
||||
def get_task(self, task_id: str) -> Optional[GpuLipsyncTaskModel]:
|
||||
return self.db.get(GpuLipsyncTaskModel, task_id)
|
||||
|
||||
def get_by_lipsync_job(self, lipsync_job_id: str) -> Optional[GpuLipsyncTaskModel]:
|
||||
return (
|
||||
self.db.query(GpuLipsyncTaskModel)
|
||||
.filter(GpuLipsyncTaskModel.lipsync_job_id == lipsync_job_id)
|
||||
.order_by(GpuLipsyncTaskModel.created_at.desc())
|
||||
.first()
|
||||
)
|
||||
|
||||
# ── 创建任务(业务侧调用) ────────────────────────────────────
|
||||
|
||||
def create_task(
|
||||
self,
|
||||
video_url: str,
|
||||
audio_url: str,
|
||||
lipsync_job_id: str = "",
|
||||
user_id: str = "",
|
||||
project_id: str = "",
|
||||
) -> GpuLipsyncTaskModel:
|
||||
task_id = str(uuid.uuid4())
|
||||
now = datetime.now(UTC)
|
||||
task = GpuLipsyncTaskModel(
|
||||
id=task_id,
|
||||
lipsync_job_id=lipsync_job_id,
|
||||
user_id=user_id,
|
||||
project_id=project_id,
|
||||
video_url=video_url,
|
||||
audio_url=audio_url,
|
||||
status="pending",
|
||||
attempt=0,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
self.db.add(task)
|
||||
self.db.commit()
|
||||
self.db.refresh(task)
|
||||
logger.info(
|
||||
"创建 GPU 口型任务 %s (lipsync_job=%s, user=%s)",
|
||||
task_id,
|
||||
lipsync_job_id,
|
||||
user_id,
|
||||
)
|
||||
return task
|
||||
|
||||
# ── 内部辅助 ──────────────────────────────────────────────────
|
||||
|
||||
def _result_key(self, task_id: str) -> str:
|
||||
return f"{self.RESULT_PREFIX}{task_id}.mp4"
|
||||
|
||||
def _touch_worker(self, worker_id: str, now: datetime) -> None:
|
||||
if not worker_id:
|
||||
return
|
||||
worker = self.db.query(GpuWorkerModel).filter(GpuWorkerModel.worker_id == worker_id).one_or_none()
|
||||
if worker is not None:
|
||||
worker.last_heartbeat_at = now
|
||||
self.db.flush()
|
||||
else:
|
||||
# 自注册(poll 时允许自动建一个空 worker 记录,运维可见)
|
||||
worker = GpuWorkerModel(
|
||||
worker_id=worker_id,
|
||||
hostname="",
|
||||
gpu_name="",
|
||||
free_vram_mb=0,
|
||||
capabilities="musetalk",
|
||||
last_heartbeat_at=now,
|
||||
created_at=now,
|
||||
)
|
||||
self.db.add(worker)
|
||||
self.db.flush()
|
||||
|
||||
def _recover_timed_out_tasks(self, now: datetime) -> None:
|
||||
"""扫描 processing 状态且超时(无心跳)的任务,回退 pending 或失败."""
|
||||
timeout = self.settings.gpu_task_timeout_seconds
|
||||
cutoff = now - timedelta(seconds=timeout)
|
||||
stuck_tasks = (
|
||||
self.db.query(GpuLipsyncTaskModel)
|
||||
.filter(
|
||||
GpuLipsyncTaskModel.status == "processing",
|
||||
GpuLipsyncTaskModel.last_heartbeat_at < cutoff,
|
||||
)
|
||||
.all()
|
||||
)
|
||||
for t in stuck_tasks:
|
||||
if t.attempt >= MAX_ATTEMPTS:
|
||||
t.status = "failed"
|
||||
t.error_msg = f"worker 心跳超时({timeout}s),重试次数已耗尽"
|
||||
t.finished_at = now
|
||||
else:
|
||||
t.status = "pending"
|
||||
t.worker_id = ""
|
||||
t.started_at = None
|
||||
t.error_msg = f"worker 心跳超时({timeout}s),等待重试"
|
||||
logger.warning("GPU 任务 %s 心跳超时,回退 pending(attempt=%d)", t.id, t.attempt)
|
||||
t.updated_at = now
|
||||
if stuck_tasks:
|
||||
self.db.flush()
|
||||
@@ -57,6 +57,14 @@ def __getattr__(name: str):
|
||||
from .atom_clips import generate_atom_clips
|
||||
|
||||
return generate_atom_clips
|
||||
elif name == "tag_atom_clip_task":
|
||||
from .atom_clip_tagging import tag_atom_clip_task
|
||||
|
||||
return tag_atom_clip_task
|
||||
elif name == "backfill_atom_clip_tags":
|
||||
from .backfill_atom_clip_tags import backfill_atom_clip_tags
|
||||
|
||||
return backfill_atom_clip_tags
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
"""片段级 AI 标签 Celery 任务 — #1970 智能剪辑流程重构 P2.
|
||||
|
||||
为单个 atom_clip 调用视觉 AI 生成结构化标签,并更新到 ai_tags 字段。
|
||||
失败不阻断流程(降级为仅继承素材标签)。
|
||||
|
||||
任务名:worker.tag_atom_clip
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from celery.utils.log import get_task_logger
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.asset_atom_clip_repository import (
|
||||
SQLAlchemyAssetAtomClipRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
|
||||
from packages.domain.atom_clip_tagger import tag_atom_clip
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
from packages.shared.mediakit_client import get_mediakit_client
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
logger = get_task_logger(__name__)
|
||||
|
||||
|
||||
@celery_app.task(name="worker.tag_atom_clip", bind=True, max_retries=2, default_retry_delay=10)
|
||||
def tag_atom_clip_task(self, atom_clip_id: str) -> dict:
|
||||
"""为单个原子片段生成 AI 标签.
|
||||
|
||||
Args:
|
||||
atom_clip_id: 原子片段 ID。
|
||||
|
||||
Returns:
|
||||
任务结果 dict:status / clip_id / ai_tags(部分字段)。
|
||||
"""
|
||||
db = SessionLocal()
|
||||
try:
|
||||
atom_repo = SQLAlchemyAssetAtomClipRepository(db)
|
||||
asset_repo = SQLAlchemyAssetRepository(db)
|
||||
|
||||
clip = atom_repo.find_by_id(atom_clip_id)
|
||||
if clip is None:
|
||||
return {"status": "skipped", "reason": "clip not found", "clip_id": atom_clip_id}
|
||||
|
||||
# 已有标签则跳过(幂等)
|
||||
if clip.ai_tags is not None:
|
||||
return {"status": "skipped", "reason": "already tagged", "clip_id": atom_clip_id}
|
||||
|
||||
# 获取素材信息
|
||||
asset = asset_repo.find_by_id(clip.asset_id)
|
||||
if asset is None:
|
||||
return {"status": "skipped", "reason": "asset not found", "clip_id": atom_clip_id}
|
||||
|
||||
# 获取视频可访问 URL
|
||||
storage = get_shared_storage_service()
|
||||
video_url = storage.get_download_url(asset.storage_key, expires_seconds=3600)
|
||||
|
||||
# 初始化客户端
|
||||
doubao_client = get_doubao_client()
|
||||
mediakit_client = get_mediakit_client()
|
||||
|
||||
# 调用 tagger
|
||||
ai_tags = tag_atom_clip(
|
||||
clip=clip,
|
||||
video_url=video_url,
|
||||
doubao_client=doubao_client,
|
||||
mediakit_client=mediakit_client,
|
||||
storage=storage,
|
||||
)
|
||||
|
||||
# 更新数据库
|
||||
atom_repo.update_ai_tags(atom_clip_id, ai_tags)
|
||||
|
||||
logger.info(
|
||||
"[atom_clip_tagging] clip_id=%s ai_tags=%s",
|
||||
atom_clip_id,
|
||||
{k: v for k, v in ai_tags.items() if k != "inherited_tags"},
|
||||
)
|
||||
return {
|
||||
"status": "completed",
|
||||
"clip_id": atom_clip_id,
|
||||
"has_ai_tags": any(v for k, v in ai_tags.items() if k != "inherited_tags" and v),
|
||||
}
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
logger.exception("[atom_clip_tagging] clip_id=%s 失败: %s", atom_clip_id, exc)
|
||||
# 可重试异常
|
||||
if self.request.retries < self.max_retries:
|
||||
raise self.retry(exc=exc) from None
|
||||
return {"status": "failed", "clip_id": atom_clip_id, "error": str(exc)}
|
||||
finally:
|
||||
db.close()
|
||||
@@ -3,6 +3,8 @@
|
||||
素材入库预处理完成(ingest 置 READY)后异步触发:
|
||||
根据素材时长和已缓存的 scdet 切换点计算原子片段并落库。
|
||||
失败不阻断素材入库主流程(atom_clips 未就绪时选片有内存兜底)。
|
||||
|
||||
P2 增强:切片完成后自动链式触发 AI 标签任务(每个 clip 一个 tag_atom_clip 任务)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -72,6 +74,10 @@ def generate_atom_clips(asset_id: str) -> dict:
|
||||
asset_id,
|
||||
len(clips),
|
||||
)
|
||||
|
||||
# P2 增强:链式触发 AI 标签任务(每个 clip 一个异步任务)
|
||||
_dispatch_tagging_tasks(clips)
|
||||
|
||||
return {"status": "completed", "asset_id": asset_id, "clips_count": len(clips)}
|
||||
except Exception as exc: # noqa: BLE001 - 后台任务兜底,失败不阻断主流程
|
||||
db.rollback()
|
||||
@@ -79,3 +85,25 @@ def generate_atom_clips(asset_id: str) -> dict:
|
||||
return {"status": "failed", "asset_id": asset_id, "error": str(exc)}
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
|
||||
def _dispatch_tagging_tasks(clips: list) -> None:
|
||||
"""为每个新建片段发送 AI 标签异步任务.
|
||||
|
||||
失败不阻断(标签任务是锦上添花,不影响核心流程)。
|
||||
"""
|
||||
try:
|
||||
for clip in clips:
|
||||
celery_app.send_task(
|
||||
"worker.tag_atom_clip",
|
||||
args=[clip.id],
|
||||
)
|
||||
logger.info(
|
||||
"[atom_clips] 已发送 %d 个 AI 标签任务",
|
||||
len(clips),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"[atom_clips] 发送 AI 标签任务失败(不影响切片结果): %s",
|
||||
e,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,102 @@
|
||||
"""批量回填 AI 标签 Celery 任务 — #1970 智能剪辑流程重构 P2.
|
||||
|
||||
查找所有 ai_tags IS NULL 的 atom_clips,分批触发 tag_atom_clip 任务。
|
||||
可通过 API 路由触发(管理员权限)。
|
||||
|
||||
任务名:worker.backfill_atom_clip_tags
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
from celery.utils.log import get_task_logger
|
||||
from worker_app.celery_app import celery_app
|
||||
from worker_app.db import SessionLocal
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.asset_atom_clip_repository import (
|
||||
SQLAlchemyAssetAtomClipRepository,
|
||||
)
|
||||
|
||||
logger = get_task_logger(__name__)
|
||||
|
||||
# 默认批量参数
|
||||
DEFAULT_BATCH_SIZE = 10
|
||||
DEFAULT_BATCH_INTERVAL = 5 # 秒
|
||||
|
||||
|
||||
@celery_app.task(name="worker.backfill_atom_clip_tags")
|
||||
def backfill_atom_clip_tags(
|
||||
batch_size: int = DEFAULT_BATCH_SIZE,
|
||||
batch_interval: int = DEFAULT_BATCH_INTERVAL,
|
||||
max_clips: int = 0,
|
||||
) -> dict:
|
||||
"""批量回填未打标的 atom_clips.
|
||||
|
||||
Args:
|
||||
batch_size: 每批处理数量,默认 10。
|
||||
batch_interval: 每批间隔秒数,默认 5。
|
||||
max_clips: 最大处理总数,0 表示不限。
|
||||
|
||||
Returns:
|
||||
任务结果 dict:total_submitted / batches。
|
||||
"""
|
||||
db = SessionLocal()
|
||||
try:
|
||||
atom_repo = SQLAlchemyAssetAtomClipRepository(db)
|
||||
total_submitted = 0
|
||||
batches = 0
|
||||
|
||||
while True:
|
||||
# 查找未打标的片段
|
||||
remaining = max_clips - total_submitted if max_clips > 0 else batch_size
|
||||
fetch_limit = min(batch_size, remaining) if max_clips > 0 else batch_size
|
||||
|
||||
untagged = atom_repo.find_untagged(limit=fetch_limit)
|
||||
if not untagged:
|
||||
break
|
||||
|
||||
# 逐个发送 tag 任务
|
||||
for clip in untagged:
|
||||
try:
|
||||
celery_app.send_task(
|
||||
"worker.tag_atom_clip",
|
||||
args=[clip.id],
|
||||
)
|
||||
total_submitted += 1
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"[backfill] 提交任务失败 clip_id=%s: %s",
|
||||
clip.id,
|
||||
e,
|
||||
)
|
||||
|
||||
batches += 1
|
||||
logger.info(
|
||||
"[backfill] 第 %d 批完成,已提交 %d 个任务",
|
||||
batches,
|
||||
total_submitted,
|
||||
)
|
||||
|
||||
# 检查是否达到上限
|
||||
if max_clips > 0 and total_submitted >= max_clips:
|
||||
break
|
||||
|
||||
# 批间间隔
|
||||
time.sleep(batch_interval)
|
||||
|
||||
logger.info(
|
||||
"[backfill] 回填完成: total_submitted=%d batches=%d",
|
||||
total_submitted,
|
||||
batches,
|
||||
)
|
||||
return {
|
||||
"status": "completed",
|
||||
"total_submitted": total_submitted,
|
||||
"batches": batches,
|
||||
}
|
||||
except Exception as exc:
|
||||
logger.exception("[backfill] 回填失败: %s", exc)
|
||||
return {"status": "failed", "error": str(exc)}
|
||||
finally:
|
||||
db.close()
|
||||
@@ -252,3 +252,7 @@ DOUYIN_DEBUG_ERRORS=false
|
||||
TIKHUB_API_KEY=${TIKHUB_API_KEY}
|
||||
# P2: apizero.cn(国内付费,https://apizero.cn)
|
||||
APIZERO_API_KEY=${APIZERO_API_KEY}
|
||||
|
||||
# ==================== GPU MuseTalk Worker(反向轮询) ====================
|
||||
GPU_WORKER_TOKEN=${GPU_WORKER_TOKEN}
|
||||
GPU_TASK_TIMEOUT_SECONDS=300
|
||||
|
||||
@@ -269,3 +269,7 @@ DOUYIN_DEBUG_ERRORS=false
|
||||
TIKHUB_API_KEY=${TIKHUB_API_KEY}
|
||||
# P2: apizero.cn(国内付费,https://apizero.cn)
|
||||
APIZERO_API_KEY=${APIZERO_API_KEY}
|
||||
|
||||
# ==================== GPU MuseTalk Worker(反向轮询) ====================
|
||||
GPU_WORKER_TOKEN=${GPU_WORKER_TOKEN}
|
||||
GPU_TASK_TIMEOUT_SECONDS=300
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
# ============================================================
|
||||
# MuseTalk GPU Worker 环境变量
|
||||
# 部署到 RTX2060 电脑后,复制为 .env 并修改值
|
||||
# ============================================================
|
||||
|
||||
# SaaS API 基础 URL(staging / production)
|
||||
API_BASE_URL=https://staging-api.xiaoxiajianji.com
|
||||
# API_BASE_URL=https://api.xiaoxiajianji.com # 生产
|
||||
|
||||
# 长期 API Token,必须与服务端 GPU_WORKER_TOKEN 一致(找后端拿)
|
||||
GPU_WORKER_TOKEN=replace-with-real-token
|
||||
|
||||
# 本机 Worker 唯一 ID(默认自动生成 hostname+MAC 后4位,可手动指定)
|
||||
# WORKER_ID=rtx2060-0193
|
||||
|
||||
# 本地 MuseTalk 地址(默认 http://127.0.0.1:7861)
|
||||
MUSE_TALK_URL=http://127.0.0.1:7861
|
||||
|
||||
# 轮询/心跳/超时(秒)
|
||||
POLL_INTERVAL=5
|
||||
HEARTBEAT_INTERVAL=15
|
||||
REQUEST_TIMEOUT=300
|
||||
|
||||
# 单个任务本地最大重试次数(首次失败后再重试 N 次,默认 2)
|
||||
TASK_MAX_RETRY=2
|
||||
@@ -0,0 +1,99 @@
|
||||
# MuseTalk GPU Worker — 部署指南
|
||||
|
||||
本目录包含 RTX2060 本地电脑上运行的 GPU Worker 脚本。
|
||||
Worker 采用 **反向轮询模式**:主动向 SaaS API 拉取待处理的口型同步任务 → 调用本地 MuseTalk 推理 → 把结果视频回传到 SaaS。不需要内网穿透。
|
||||
|
||||
## 目录文件
|
||||
|
||||
| 文件 | 作用 |
|
||||
|---|---|
|
||||
| `gpu_worker.py` | Worker 主程序(单文件,零项目代码依赖,仅依赖 `requests`) |
|
||||
| `requirements.txt` | Python 依赖(只有 `requests`) |
|
||||
| `xiaoxia-gpu-worker.service` | systemd 服务单元(开机自启、异常自动重启) |
|
||||
| `.env.example` | 环境变量样例,复制为 `.env` 后填入真实值 |
|
||||
|
||||
## 一、环境准备
|
||||
|
||||
1. **Python 3.10+**(Windows 建议从 python.org 安装;Linux 自带)
|
||||
2. **本地 MuseTalk 服务** 已启动在 `http://127.0.0.1:7861`,health 接口返回 `{"status":"ok","free_vram_mb":...}`
|
||||
3. **ffmpeg**(可选,用于读取输出视频时长;未装则 duration 报 0,不影响功能)
|
||||
4. 网络能访问 staging / 生产 API(`curl https://staging-api.xiaoxiajianji.com/health` 应返回 `{"status":"healthy"}`)
|
||||
|
||||
## 二、部署步骤(Linux,推荐 systemd)
|
||||
|
||||
```bash
|
||||
# 1. 创建部署目录
|
||||
sudo mkdir -p /opt/xiaoxia-gpu-worker
|
||||
sudo chown $USER:$USER /opt/xiaoxia-gpu-worker
|
||||
cd /opt/xiaoxia-gpu-worker
|
||||
|
||||
# 2. 拷贝脚本和依赖
|
||||
cp /path/to/deploy/gpu_worker/{gpu_worker.py,requirements.txt,xiaoxia-gpu-worker.service,.env.example} .
|
||||
cp .env.example .env
|
||||
# 编辑 .env,填入 API_BASE_URL 和 GPU_WORKER_TOKEN
|
||||
|
||||
# 3. 创建虚拟环境并安装依赖
|
||||
python3 -m venv venv
|
||||
./venv/bin/pip install -r requirements.txt
|
||||
|
||||
# 4. 前台先跑一次,确认日志正常
|
||||
./venv/bin/python gpu_worker.py
|
||||
# 看到 "MuseTalk 健康检查通过" 和 "注册/心跳" 成功即可 Ctrl+C 退出
|
||||
|
||||
# 5. 安装 systemd 服务
|
||||
sudo cp xiaoxia-gpu-worker.service /etc/systemd/system/
|
||||
sudo systemctl daemon-reload
|
||||
sudo systemctl enable --now xiaoxia-gpu-worker
|
||||
|
||||
# 6. 查看日志
|
||||
sudo journalctl -u xiaoxia-gpu-worker -f
|
||||
```
|
||||
|
||||
## 三、部署步骤(Windows,快速测试)
|
||||
|
||||
```bat
|
||||
:: 创建虚拟环境
|
||||
python -m venv venv
|
||||
venv\Scripts\pip install -r requirements.txt
|
||||
|
||||
:: 复制并编辑 .env
|
||||
copy .env.example .env
|
||||
notepad .env
|
||||
|
||||
:: 运行
|
||||
venv\Scripts\python gpu_worker.py
|
||||
```
|
||||
|
||||
可在任务计划程序中添加开机启动项:程序选 `venv\Scripts\python.exe`,参数填 `gpu_worker.py`,起始目录填脚本所在目录。
|
||||
|
||||
## 四、SaaS 侧配套配置
|
||||
|
||||
SaaS 后端部署完成后需配置:
|
||||
|
||||
1. 服务端环境变量 `GPU_WORKER_TOKEN` 设为一个随机强 Token(和 Worker `.env` 中一致)
|
||||
2. 数据库已跑迁移 `081_add_gpu_lipsync_tasks`(自动随 API 启动的 alembic upgrade head 完成)
|
||||
3. OSS bucket 中 `gpu-lipsync/results/` 路径可写(默认 bucket 已配)
|
||||
|
||||
## 五、验证联调
|
||||
|
||||
1. Worker 启动后日志看到 `注册/心跳` 成功
|
||||
2. 后端调用 `GpuLipsyncService.create_task(video_url=..., audio_url=...)` 放入一条测试任务
|
||||
3. Worker 在 5 秒内拉到任务,下载 → 推理 → 上传 → 上报
|
||||
4. 后端 `GET /api/v1/gpu/lipsync/status/{task_id}` 返回 `status=done`,`result_url` 非空
|
||||
|
||||
## 六、故障排查
|
||||
|
||||
| 现象 | 可能原因 / 排查 |
|
||||
|---|---|
|
||||
| 日志 401 `Invalid GPU worker token` | `.env` 的 `GPU_WORKER_TOKEN` 与服务端不一致 |
|
||||
| 日志 `MuseTalk 健康检查未通过` | 本地 MuseTalk 没启动,或端口不是 7861;`curl http://127.0.0.1:7861/health` 验证 |
|
||||
| 任务长时间不被拉取 | Worker 和服务端连不上;检查 API_BASE_URL 是否可达、Token 是否正确 |
|
||||
| 推理后上传 OSS 失败 | 本地出口网络被防火墙拦截 OSS 域名(oss-cn-hangzhou.aliyuncs.com) |
|
||||
| 服务端看到任务回退到 pending 重试 | Worker 心跳超时(默认 5 分钟);Worker 进程崩溃或推理卡死超过 5 分钟 |
|
||||
| 日志 `MuseTalk 推理超时` | 视频太长或显存不足;可临时调大 REQUEST_TIMEOUT,或限制输入视频时长 |
|
||||
|
||||
## 七、安全注意事项
|
||||
|
||||
- `.env` 包含长期 Token,文件权限设为 600(`chmod 600 .env`)
|
||||
- Token 泄露要立即在服务端更换 `GPU_WORKER_TOKEN` 并重启 Worker
|
||||
- Worker 只需要出站访问 SaaS API 和 OSS,不需要开放任何入站端口
|
||||
@@ -0,0 +1,395 @@
|
||||
"""MuseTalk GPU Worker — 反向轮询模式.
|
||||
|
||||
部署在有 RTX2060 的本地电脑上(192.168.0.193),
|
||||
主动轮询 SaaS API 拉取口型任务、调用本地 MuseTalk 推理、上传结果回 SaaS。
|
||||
|
||||
环境变量:
|
||||
API_BASE_URL SaaS API 基础 URL(不含 /api/v1),如 https://staging-api.xiaoxiajianji.com
|
||||
GPU_WORKER_TOKEN 长期 API Token(服务端 GPU_WORKER_TOKEN 需一致)
|
||||
WORKER_ID 本机唯一 ID(默认 hostname+网卡MAC 后4位)
|
||||
MUSE_TALK_URL 本地 MuseTalk 地址,默认 http://127.0.0.1:7861
|
||||
POLL_INTERVAL 轮询间隔秒,默认 5
|
||||
HEARTBEAT_INTERVAL 心跳间隔秒,默认 15
|
||||
REQUEST_TIMEOUT HTTP 请求超时秒,默认 60
|
||||
TASK_MAX_RETRY 单个任务最大重试次数(在 Worker 本地的重试),默认 2
|
||||
|
||||
用法:
|
||||
python gpu_worker.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import platform
|
||||
import socket
|
||||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import requests
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s [%(levelname)s] %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
)
|
||||
logger = logging.getLogger("musetalk-worker")
|
||||
|
||||
# ── 配置 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _env(name: str, default: str = "") -> str:
|
||||
v = os.environ.get(name, default)
|
||||
return v.strip() if isinstance(v, str) else default
|
||||
|
||||
|
||||
class Config:
|
||||
api_base_url: str = _env("API_BASE_URL", "https://staging-api.xiaoxiajianji.com").rstrip("/")
|
||||
gpu_worker_token: str = _env("GPU_WORKER_TOKEN")
|
||||
muse_talk_url: str = _env("MUSE_TALK_URL", "http://127.0.0.1:7861").rstrip("/")
|
||||
poll_interval: float = float(_env("POLL_INTERVAL", "5"))
|
||||
heartbeat_interval: float = float(_env("HEARTBEAT_INTERVAL", "15"))
|
||||
request_timeout: float = float(_env("REQUEST_TIMEOUT", "300"))
|
||||
task_max_retry: int = int(_env("TASK_MAX_RETRY", "2"))
|
||||
worker_id: str = _env("WORKER_ID", "")
|
||||
|
||||
@classmethod
|
||||
def derived_worker_id(cls) -> str:
|
||||
if cls.worker_id:
|
||||
return cls.worker_id
|
||||
# hostname + MAC 后4位 → 稳定唯一 ID
|
||||
try:
|
||||
mac = uuid.getnode()
|
||||
mac_suffix = f"{mac:012x}"[-4:]
|
||||
except Exception:
|
||||
mac_suffix = "0000"
|
||||
host = platform.node() or socket.gethostname() or "rtx2060"
|
||||
return f"{host}-{mac_suffix}"
|
||||
|
||||
|
||||
# ── 辅助 ─────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _api_headers() -> dict[str, str]:
|
||||
token = Config.gpu_worker_token
|
||||
if not token:
|
||||
logger.warning("GPU_WORKER_TOKEN 未配置,开发模式下会被服务端拒绝(生产环境必须配置)")
|
||||
return {"Authorization": f"Bearer {token}"} if token else {}
|
||||
|
||||
|
||||
def _check_musetalk_health() -> tuple[bool, dict]:
|
||||
"""检查本地 MuseTalk 健康状态,返回 (ok, info)."""
|
||||
try:
|
||||
r = requests.get(f"{Config.muse_talk_url}/health", timeout=5)
|
||||
if r.status_code == 200:
|
||||
try:
|
||||
return True, r.json()
|
||||
except Exception:
|
||||
return True, {}
|
||||
return False, {"status_code": r.status_code, "body": r.text[:200]}
|
||||
except Exception as exc:
|
||||
return False, {"error": str(exc)}
|
||||
|
||||
|
||||
def _register() -> bool:
|
||||
"""向服务端注册 / 心跳,附带 GPU 信息."""
|
||||
ok, info = _check_musetalk_health()
|
||||
free_vram = int(info.get("free_vram_mb", 0) or 0) if isinstance(info, dict) else 0
|
||||
gpu_name = info.get("gpu_name", "") if isinstance(info, dict) else ""
|
||||
if not gpu_name:
|
||||
# 尝试在 Windows 上读 nvidia-smi
|
||||
gpu_name = _probe_gpu_name()
|
||||
payload = {
|
||||
"worker_id": Config.derived_worker_id(),
|
||||
"hostname": platform.node(),
|
||||
"gpu_name": gpu_name,
|
||||
"free_vram_mb": free_vram,
|
||||
"capabilities": "musetalk",
|
||||
}
|
||||
try:
|
||||
r = requests.post(
|
||||
f"{Config.api_base_url}/api/v1/gpu/register",
|
||||
json=payload,
|
||||
headers=_api_headers(),
|
||||
timeout=15,
|
||||
)
|
||||
if r.status_code == 200:
|
||||
return True
|
||||
logger.error("注册/心跳失败: HTTP %d body=%s", r.status_code, r.text[:300])
|
||||
return False
|
||||
except Exception as exc:
|
||||
logger.error("注册/心跳异常: %s", exc)
|
||||
return False
|
||||
|
||||
|
||||
def _probe_gpu_name() -> str:
|
||||
"""尽力探测 GPU 型号(不强制依赖 pynvml)."""
|
||||
try:
|
||||
import subprocess
|
||||
|
||||
out = subprocess.check_output(
|
||||
["nvidia-smi", "--query-gpu=name", "--format=csv,noheader"],
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=5,
|
||||
)
|
||||
return out.decode("utf-8", errors="ignore").strip().splitlines()[0].strip()
|
||||
except Exception:
|
||||
return ""
|
||||
|
||||
|
||||
def _poll_task() -> Optional[dict]:
|
||||
"""轮询拉取一条待处理任务;无任务返回 None."""
|
||||
try:
|
||||
r = requests.get(
|
||||
f"{Config.api_base_url}/api/v1/gpu/lipsync/poll",
|
||||
params={"worker_id": Config.derived_worker_id()},
|
||||
headers=_api_headers(),
|
||||
timeout=30,
|
||||
)
|
||||
if r.status_code == 204:
|
||||
return None
|
||||
if r.status_code == 200:
|
||||
data = r.json()
|
||||
return data.get("task")
|
||||
logger.error("poll 返回 %d: %s", r.status_code, r.text[:300])
|
||||
return None
|
||||
except Exception as exc:
|
||||
logger.error("poll 异常: %s", exc)
|
||||
return None
|
||||
|
||||
|
||||
def _download(url: str, path: Path) -> bool:
|
||||
"""下载文件到本地,支持预签名 URL."""
|
||||
try:
|
||||
with requests.get(url, stream=True, timeout=Config.request_timeout) as r:
|
||||
if r.status_code >= 400:
|
||||
logger.error("下载失败 HTTP %d: %s", r.status_code, url[:120])
|
||||
return False
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with open(path, "wb") as f:
|
||||
for chunk in r.iter_content(chunk_size=1024 * 256):
|
||||
if chunk:
|
||||
f.write(chunk)
|
||||
return path.stat().st_size > 0
|
||||
except Exception as exc:
|
||||
logger.error("下载异常 %s: %s", url[:120], exc)
|
||||
return False
|
||||
|
||||
|
||||
def _call_musetalk(video_path: Path, audio_path: Path, out_path: Path) -> tuple[bool, float, str]:
|
||||
"""调用本地 MuseTalk /inference.
|
||||
|
||||
返回 (success, duration_seconds, error_msg).
|
||||
duration 用 ffprobe 读结果视频,失败填 0。
|
||||
"""
|
||||
try:
|
||||
with open(video_path, "rb") as vf, open(audio_path, "rb") as af:
|
||||
files = {
|
||||
"video": (video_path.name, vf, "video/mp4"),
|
||||
"audio": (audio_path.name, af, "application/octet-stream"),
|
||||
}
|
||||
r = requests.post(
|
||||
f"{Config.muse_talk_url}/inference",
|
||||
files=files,
|
||||
timeout=Config.request_timeout,
|
||||
)
|
||||
if r.status_code != 200:
|
||||
return False, 0.0, f"MuseTalk HTTP {r.status_code}: {r.text[:500]}"
|
||||
out_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
out_path.write_bytes(r.content)
|
||||
if out_path.stat().st_size < 1024:
|
||||
return False, 0.0, f"MuseTalk 返回结果过小 ({out_path.stat().st_size} bytes)"
|
||||
duration = _probe_duration(out_path)
|
||||
return True, duration, ""
|
||||
except requests.exceptions.Timeout:
|
||||
return False, 0.0, f"MuseTalk 推理超时(>{Config.request_timeout}s)"
|
||||
except Exception as exc:
|
||||
return False, 0.0, f"MuseTalk 调用异常: {exc}"
|
||||
|
||||
|
||||
def _probe_duration(path: Path) -> float:
|
||||
"""用 ffprobe 读视频时长(若系统装了 ffmpeg);否则返回 0."""
|
||||
try:
|
||||
import subprocess
|
||||
|
||||
out = subprocess.check_output(
|
||||
[
|
||||
"ffprobe", "-v", "error",
|
||||
"-show_entries", "format=duration",
|
||||
"-of", "default=noprint_wrappers=1:nokey=1",
|
||||
str(path),
|
||||
],
|
||||
stderr=subprocess.DEVNULL,
|
||||
timeout=10,
|
||||
)
|
||||
return float(out.decode().strip() or 0)
|
||||
except Exception:
|
||||
return 0.0
|
||||
|
||||
|
||||
def _upload_result(upload_url: str, file_path: Path) -> bool:
|
||||
"""PUT 上传结果视频到预签名 URL."""
|
||||
try:
|
||||
with open(file_path, "rb") as f:
|
||||
r = requests.put(
|
||||
upload_url,
|
||||
data=f,
|
||||
headers={"Content-Type": "video/mp4"},
|
||||
timeout=Config.request_timeout,
|
||||
)
|
||||
if r.status_code >= 400:
|
||||
logger.error("上传结果失败 HTTP %d: %s", r.status_code, r.text[:500])
|
||||
return False
|
||||
return True
|
||||
except Exception as exc:
|
||||
logger.error("上传结果异常: %s", exc)
|
||||
return False
|
||||
|
||||
|
||||
def _report_result(task_id: str, success: bool, duration: float = 0.0, error_msg: str = "") -> bool:
|
||||
"""通知服务端结果。失败时也尝试上报错误(不含视频文件)."""
|
||||
try:
|
||||
data = {
|
||||
"task_id": task_id,
|
||||
"worker_id": Config.derived_worker_id(),
|
||||
"success": "true" if success else "false",
|
||||
"duration_seconds": str(duration),
|
||||
"error_msg": error_msg,
|
||||
}
|
||||
r = requests.post(
|
||||
f"{Config.api_base_url}/api/v1/gpu/lipsync/result",
|
||||
data=data,
|
||||
headers=_api_headers(),
|
||||
timeout=30,
|
||||
)
|
||||
if r.status_code != 200:
|
||||
logger.error("上报结果失败 HTTP %d: %s", r.status_code, r.text[:300])
|
||||
return False
|
||||
return True
|
||||
except Exception as exc:
|
||||
logger.error("上报结果异常: %s", exc)
|
||||
return False
|
||||
|
||||
|
||||
def _handle_task(task: dict) -> None:
|
||||
"""处理一条任务(整个串行流程:下载→推理→上传→上报)."""
|
||||
task_id = task["task_id"]
|
||||
logger.info("开始处理任务 %s", task_id)
|
||||
with tempfile.TemporaryDirectory(prefix="musetalk_") as tmpdir:
|
||||
tmp = Path(tmpdir)
|
||||
video_path = tmp / "input.mp4"
|
||||
audio_path = tmp / "input_audio.bin"
|
||||
out_path = tmp / "output.mp4"
|
||||
|
||||
# 1. 下载
|
||||
if not _download(task["video_url"], video_path):
|
||||
_report_result(task_id, False, 0.0, "下载人物视频失败")
|
||||
return
|
||||
if not _download(task["audio_url"], audio_path):
|
||||
_report_result(task_id, False, 0.0, "下载驱动音频失败")
|
||||
return
|
||||
|
||||
# 2. 推理(本地重试)
|
||||
success = False
|
||||
duration = 0.0
|
||||
err = ""
|
||||
for attempt in range(Config.task_max_retry + 1):
|
||||
if attempt > 0:
|
||||
logger.info("任务 %s 第 %d 次重试...", task_id, attempt + 1)
|
||||
time.sleep(2)
|
||||
success, duration, err = _call_musetalk(video_path, audio_path, out_path)
|
||||
if success:
|
||||
break
|
||||
if not success:
|
||||
logger.error("任务 %s 推理失败: %s", task_id, err)
|
||||
_report_result(task_id, False, 0.0, err)
|
||||
return
|
||||
|
||||
# 3. 上报结果(multipart 同时上传文件 → API 代为 PUT 到 OSS,逻辑最稳)
|
||||
_report_success_with_file(task_id, duration, out_path)
|
||||
|
||||
|
||||
def _report_success_with_file(task_id: str, duration: float, file_path: Path) -> None:
|
||||
"""上报成功并 multipart 附带结果视频."""
|
||||
try:
|
||||
data = {
|
||||
"task_id": task_id,
|
||||
"worker_id": Config.derived_worker_id(),
|
||||
"success": "true",
|
||||
"duration_seconds": str(duration),
|
||||
"error_msg": "",
|
||||
}
|
||||
with open(file_path, "rb") as f:
|
||||
files = {"result": (f"{task_id}.mp4", f, "video/mp4")}
|
||||
r = requests.post(
|
||||
f"{Config.api_base_url}/api/v1/gpu/lipsync/result",
|
||||
data=data,
|
||||
files=files,
|
||||
headers=_api_headers(),
|
||||
timeout=Config.request_timeout,
|
||||
)
|
||||
if r.status_code != 200:
|
||||
logger.error("上报成功结果失败 HTTP %d: %s", r.status_code, r.text[:300])
|
||||
return
|
||||
logger.info("任务 %s 完成,duration=%.1fs", task_id, duration)
|
||||
except Exception as exc:
|
||||
logger.error("上报成功结果异常: %s", exc)
|
||||
|
||||
|
||||
# ── 主循环 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def main() -> int:
|
||||
logger.info("=" * 60)
|
||||
logger.info("MuseTalk GPU Worker 启动")
|
||||
logger.info(" worker_id = %s", Config.derived_worker_id())
|
||||
logger.info(" api_base = %s", Config.api_base_url)
|
||||
logger.info(" muse_talk = %s", Config.muse_talk_url)
|
||||
logger.info(" poll = %.1fs / heartbeat = %.1fs", Config.poll_interval, Config.heartbeat_interval)
|
||||
logger.info("=" * 60)
|
||||
|
||||
if not Config.gpu_worker_token:
|
||||
logger.warning("GPU_WORKER_TOKEN 未配置(开发模式),生产环境必须设置")
|
||||
|
||||
# 先检查一次 MuseTalk
|
||||
ok, info = _check_musetalk_health()
|
||||
if ok:
|
||||
logger.info("MuseTalk 健康检查通过: %s", info)
|
||||
else:
|
||||
logger.warning("MuseTalk 健康检查未通过: %s(继续运行,等待服务可用)", info)
|
||||
|
||||
# 启动时立即注册
|
||||
_register()
|
||||
last_heartbeat = time.time()
|
||||
|
||||
while True:
|
||||
try:
|
||||
# 心跳
|
||||
now = time.time()
|
||||
if now - last_heartbeat >= Config.heartbeat_interval:
|
||||
if _register():
|
||||
last_heartbeat = now
|
||||
|
||||
# 轮询任务
|
||||
task = _poll_task()
|
||||
if task is not None:
|
||||
_handle_task(task)
|
||||
# 处理完立即再 poll(不 sleep),尽可能拉满 GPU
|
||||
continue
|
||||
|
||||
time.sleep(Config.poll_interval)
|
||||
except KeyboardInterrupt:
|
||||
logger.info("收到中断信号,退出")
|
||||
return 0
|
||||
except Exception as exc:
|
||||
logger.exception("主循环异常: %s", exc)
|
||||
time.sleep(Config.poll_interval)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1 @@
|
||||
requests>=2.31.0
|
||||
@@ -0,0 +1,21 @@
|
||||
[Unit]
|
||||
Description=MuseTalk GPU Worker (xiaoxia-saas 反向轮询)
|
||||
After=network.target musetalk.service
|
||||
# 本地 MuseTalk 服务启动后再启动本 Worker;若 MuseTalk 没有 systemd 服务则删除 musetalk.service
|
||||
|
||||
[Service]
|
||||
Type=simple
|
||||
User=%i
|
||||
WorkingDirectory=/opt/xiaoxia-gpu-worker
|
||||
# 读取环境变量(API 地址、Token、轮询间隔等)
|
||||
EnvironmentFile=/opt/xiaoxia-gpu-worker/.env
|
||||
ExecStart=/opt/xiaoxia-gpu-worker/venv/bin/python /opt/xiaoxia-gpu-worker/gpu_worker.py
|
||||
Restart=always
|
||||
RestartSec=10
|
||||
# 日志走 journal,用 journalctl -u xiaoxia-gpu-worker -f 查看
|
||||
StandardOutput=journal
|
||||
StandardError=journal
|
||||
SyslogIdentifier=xiaoxia-gpu-worker
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
@@ -83,6 +83,25 @@ class SQLAlchemyAssetAtomClipRepository:
|
||||
models = query.all()
|
||||
return [self._to_domain(m) for m in models]
|
||||
|
||||
def update_ai_tags(self, clip_id: str, ai_tags: dict) -> bool:
|
||||
"""更新指定片段的 ai_tags 字段."""
|
||||
count = (
|
||||
self.session.query(AssetAtomClipModel).filter(AssetAtomClipModel.id == clip_id).update({"ai_tags": ai_tags})
|
||||
)
|
||||
self.session.commit()
|
||||
return count > 0
|
||||
|
||||
def find_untagged(self, limit: int = 100) -> list[AssetAtomClip]:
|
||||
"""查找 ai_tags IS NULL 的片段,用于回填."""
|
||||
models = (
|
||||
self.session.query(AssetAtomClipModel)
|
||||
.filter(AssetAtomClipModel.ai_tags.is_(None))
|
||||
.order_by(AssetAtomClipModel.created_at.asc())
|
||||
.limit(limit)
|
||||
.all()
|
||||
)
|
||||
return [self._to_domain(m) for m in models]
|
||||
|
||||
def _to_model(self, clip: AssetAtomClip) -> AssetAtomClipModel:
|
||||
return AssetAtomClipModel(
|
||||
id=clip.id,
|
||||
@@ -92,6 +111,7 @@ class SQLAlchemyAssetAtomClipRepository:
|
||||
duration=clip.duration,
|
||||
clip_index=clip.clip_index,
|
||||
tags=clip.tags,
|
||||
ai_tags=clip.ai_tags,
|
||||
scene_change_at=clip.scene_change_at,
|
||||
is_fallback=clip.is_fallback,
|
||||
created_at=clip.created_at or datetime.now(UTC),
|
||||
|
||||
@@ -837,6 +837,7 @@ class AssetAtomClipModel(Base):
|
||||
duration = Column(Float, nullable=False)
|
||||
clip_index = Column(Integer, nullable=False)
|
||||
tags = Column(JSON, nullable=False, default=list)
|
||||
ai_tags = Column(JSON, nullable=True, default=None)
|
||||
scene_change_at = Column(Float, nullable=True)
|
||||
is_fallback = Column(Boolean, nullable=False, default=False)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
@@ -854,3 +855,61 @@ class DailyUsageRecordModel(Base):
|
||||
usage_type = Column(String(50), nullable=False, default="free_clip")
|
||||
count = Column(Integer, nullable=False, default=0)
|
||||
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
class GpuLipsyncTaskModel(Base):
|
||||
"""GPU 口型同步任务 ORM 模型 — MuseTalk 反向轮询模式.
|
||||
|
||||
业务侧(AI 数字人生成/lipsync 流程)提交任务后,GPU Worker 主动 poll 拉取、
|
||||
调用本地 MuseTalk 推理、再通过 result 接口回传结果视频。
|
||||
"""
|
||||
|
||||
__tablename__ = "gpu_lipsync_tasks"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
# 业务关联(原 lipsync_job_id,方便双向查询)
|
||||
lipsync_job_id = Column(String(36), nullable=False, default="", index=True)
|
||||
user_id = Column(String(36), nullable=False, default="", index=True)
|
||||
project_id = Column(String(36), nullable=False, default="", index=True)
|
||||
|
||||
# 输入(预签名下载 URL,由 API 侧生成)
|
||||
video_url = Column(Text, nullable=False)
|
||||
audio_url = Column(Text, nullable=False)
|
||||
|
||||
# 结果
|
||||
result_url = Column(Text, nullable=False, default="")
|
||||
result_duration = Column(Float, nullable=False, default=0.0)
|
||||
|
||||
# 任务状态
|
||||
status = Column(
|
||||
String(20),
|
||||
nullable=False,
|
||||
default="pending",
|
||||
index=True,
|
||||
) # pending → processing → done / failed / timeout
|
||||
worker_id = Column(String(100), nullable=False, default="", index=True)
|
||||
attempt = Column(Integer, nullable=False, default=0)
|
||||
error_msg = Column(Text, nullable=False, default="")
|
||||
|
||||
# 时间戳
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
started_at = Column(DateTime, nullable=True)
|
||||
finished_at = Column(DateTime, nullable=True)
|
||||
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
# 心跳:worker 最近一次 poll/result 的时间,用于判定 worker 失联
|
||||
last_heartbeat_at = Column(DateTime, nullable=True)
|
||||
|
||||
|
||||
class GpuWorkerModel(Base):
|
||||
"""GPU Worker 注册表 — 反向轮询模式下用于心跳与监控."""
|
||||
|
||||
__tablename__ = "gpu_workers"
|
||||
|
||||
worker_id = Column(String(100), primary_key=True)
|
||||
hostname = Column(String(200), nullable=False, default="")
|
||||
gpu_name = Column(String(200), nullable=False, default="")
|
||||
free_vram_mb = Column(Integer, nullable=False, default=0)
|
||||
capabilities = Column(String(500), nullable=False, default="") # 逗号分隔,如 "musetalk"
|
||||
last_heartbeat_at = Column(DateTime, nullable=True, index=True)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
@@ -68,6 +68,7 @@ class SharedSettings(BaseSettings):
|
||||
doubao_base_url: str = "https://ark.cn-beijing.volces.com/api/v3"
|
||||
doubao_timeout: int = 30
|
||||
doubao_max_retries: int = 2
|
||||
doubao_vision_model: str = "doubao-1-5-vision-pro-250915"
|
||||
|
||||
# ── MediaKit (火山引擎 AI 媒体工具) ──────────────────────────────────
|
||||
mediakit_api_key: str = ""
|
||||
@@ -79,6 +80,17 @@ class SharedSettings(BaseSettings):
|
||||
# `if settings.points_enabled:` 包裹,防止未完善的扣点逻辑影响现有用户。
|
||||
points_enabled: bool = False
|
||||
|
||||
# ── GPU MuseTalk 反向轮询 Worker ────────────────────────────────────
|
||||
# Worker 用这个长期 Token 鉴权(不是用户 JWT)。多 Worker 共用同一个 Token;
|
||||
# worker_id 用于区分具体机器。生产必须配置;development 留空会跳过校验。
|
||||
gpu_worker_token: str = ""
|
||||
# GPU 任务超时(秒):超过此时长仍未完成则标记为 failed,可重新 poll
|
||||
gpu_task_timeout_seconds: int = 300
|
||||
# 结果预签名 URL 有效期(秒)
|
||||
gpu_result_url_expires: int = 3600
|
||||
# 输入预签名 URL 有效期(秒,需留出 Worker 下载时间)
|
||||
gpu_input_url_expires: int = 3600
|
||||
|
||||
@property
|
||||
def effective_database_url(self) -> str:
|
||||
"""返回实际使用的数据库 URL。
|
||||
|
||||
@@ -36,6 +36,7 @@ class AssetAtomClip:
|
||||
duration: float
|
||||
clip_index: int
|
||||
tags: list[str] = field(default_factory=list)
|
||||
ai_tags: dict | None = None
|
||||
scene_change_at: float | None = None
|
||||
is_fallback: bool = False
|
||||
created_at: datetime | None = None
|
||||
|
||||
@@ -0,0 +1,292 @@
|
||||
"""片段级 AI 标签 — #1970 智能剪辑流程重构 P2.
|
||||
|
||||
对每个 atom_clip 提取关键帧,调用豆包视觉理解 API 识别内容,
|
||||
生成结构化标签(场景、物体、动作、景别、是否有文字)。
|
||||
|
||||
纯函数 + IO 分离设计:
|
||||
- build_vision_prompt() 返回结构化 prompt
|
||||
- parse_vision_response(text) 解析 AI 返回的 JSON 标签
|
||||
- tag_atom_clip(...) 主入口,组合帧提取 → 视觉 API → 解析标签
|
||||
|
||||
降级策略:任何环节失败都返回 {"inherited_tags": clip.tags},不阻断流程。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import subprocess
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# AI 标签结构的键
|
||||
AI_TAG_KEYS = ("scene", "objects", "action", "shot", "has_text")
|
||||
|
||||
|
||||
def build_vision_prompt() -> str:
|
||||
"""返回结构化标签提取 prompt.
|
||||
|
||||
要求 AI 以 JSON 格式返回片段内容标签,包含:
|
||||
- scene: 场景类型列表(如 "工厂", "办公室", "户外")
|
||||
- objects: 出现的物体列表(如 "产品", "手机", "电脑")
|
||||
- action: 动作类型列表(如 "演示", "说话", "操作")
|
||||
- shot: 景别("特写" / "中景" / "远景" 之一)
|
||||
- has_text: 画面中是否有显著文字(true/false)
|
||||
"""
|
||||
return """请分析这段视频片段的关键帧,识别内容并返回 JSON 格式标签。
|
||||
|
||||
要求返回以下 JSON 结构(严格 JSON,不要添加其他文字):
|
||||
{
|
||||
"scene": ["场景1", "场景2"],
|
||||
"objects": ["物体1", "物体2"],
|
||||
"action": ["动作1"],
|
||||
"shot": "特写|中景|远景",
|
||||
"has_text": true/false
|
||||
}
|
||||
|
||||
规则:
|
||||
- scene: 场景类型,如"工厂"、"办公室"、"户外"、"商店"、"家庭"等,1-3个
|
||||
- objects: 画面中可见的主要物体,如"产品"、"手机"、"电脑"、"食品"等,1-5个
|
||||
- action: 人物或物体正在进行的动作,如"演示"、"说话"、"操作"、"展示"等,1-3个
|
||||
- shot: 景别判断,只能是"特写"、"中景"或"远景"之一
|
||||
- has_text: 画面中是否有显著可读文字(标题、字幕、标语等)
|
||||
|
||||
请只返回 JSON,不要有其他说明文字。"""
|
||||
|
||||
|
||||
def parse_vision_response(text: str) -> dict:
|
||||
"""解析 AI 返回的 JSON 标签文本.
|
||||
|
||||
Args:
|
||||
text: 视觉 API 返回的文本,期望是 JSON 格式。
|
||||
|
||||
Returns:
|
||||
结构化标签 dict,格式如:
|
||||
{"scene": [...], "objects": [...], "action": [...], "shot": "...", "has_text": bool}
|
||||
|
||||
解析失败时返回空 dict。
|
||||
"""
|
||||
if not text or not text.strip():
|
||||
return {}
|
||||
|
||||
# 尝试直接解析
|
||||
cleaned = text.strip()
|
||||
|
||||
# 去除可能的 markdown 代码块包裹
|
||||
if cleaned.startswith("```"):
|
||||
lines = cleaned.split("\n")
|
||||
# 去掉首尾的 ``` 行
|
||||
start = 1
|
||||
end = len(lines)
|
||||
for i in range(len(lines) - 1, 0, -1):
|
||||
if lines[i].strip().startswith("```"):
|
||||
end = i
|
||||
break
|
||||
cleaned = "\n".join(lines[start:end]).strip()
|
||||
|
||||
try:
|
||||
data = json.loads(cleaned)
|
||||
except json.JSONDecodeError:
|
||||
# 尝试从文本中提取 JSON 块
|
||||
try:
|
||||
start_idx = cleaned.index("{")
|
||||
end_idx = cleaned.rindex("}") + 1
|
||||
data = json.loads(cleaned[start_idx:end_idx])
|
||||
except (ValueError, json.JSONDecodeError):
|
||||
logger.warning("无法解析 AI 标签响应: %s", text[:200])
|
||||
return {}
|
||||
|
||||
if not isinstance(data, dict):
|
||||
return {}
|
||||
|
||||
# 验证和清洗各字段
|
||||
result: dict[str, Any] = {}
|
||||
for key in ("scene", "objects", "action"):
|
||||
val = data.get(key)
|
||||
if isinstance(val, list):
|
||||
result[key] = [str(v).strip() for v in val if str(v).strip()]
|
||||
elif isinstance(val, str) and val.strip():
|
||||
result[key] = [val.strip()]
|
||||
else:
|
||||
result[key] = []
|
||||
|
||||
shot_val = data.get("shot", "")
|
||||
if isinstance(shot_val, str) and shot_val.strip() in ("特写", "中景", "远景"):
|
||||
result["shot"] = shot_val.strip()
|
||||
else:
|
||||
result["shot"] = ""
|
||||
|
||||
has_text_val = data.get("has_text")
|
||||
if isinstance(has_text_val, bool):
|
||||
result["has_text"] = has_text_val
|
||||
elif isinstance(has_text_val, str):
|
||||
result["has_text"] = has_text_val.lower() in ("true", "yes", "1")
|
||||
else:
|
||||
result["has_text"] = False
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def _extract_frames_via_mediakit(
|
||||
mediakit_client: Any,
|
||||
video_url: str,
|
||||
start_time: float,
|
||||
end_time: float,
|
||||
) -> Optional[list[str]]:
|
||||
"""通过 MediaKit 提取 3 帧(首、中、尾).
|
||||
|
||||
Returns:
|
||||
图片 URL 列表(3 个),失败返回 None。
|
||||
"""
|
||||
try:
|
||||
frames = mediakit_client.extract_frames(
|
||||
video_url=video_url,
|
||||
strategy="SpecifiedTime",
|
||||
max_frames=3,
|
||||
poll_interval=2.0,
|
||||
max_poll_attempts=30,
|
||||
)
|
||||
# MediaKit SpecifiedTime 策略可能不支持直接传时间点
|
||||
# 如果返回结果不够 3 帧,降级到 ffmpeg
|
||||
if frames and len(frames) >= 1:
|
||||
urls = [f.get("image_url", "") for f in frames if f.get("image_url")]
|
||||
if urls:
|
||||
return urls
|
||||
except Exception as e:
|
||||
logger.warning("MediaKit 抽帧失败,将降级为 ffmpeg: %s", e)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _extract_frames_via_ffmpeg(
|
||||
video_url: str,
|
||||
start_time: float,
|
||||
end_time: float,
|
||||
) -> Optional[list[str]]:
|
||||
"""通过 ffmpeg 本地提取 3 帧并转为 base64.
|
||||
|
||||
Returns:
|
||||
base64 data URI 列表(3 个),失败返回 None。
|
||||
"""
|
||||
import base64
|
||||
|
||||
mid_time = round((start_time + end_time) / 2, 3)
|
||||
timestamps = [round(start_time, 3), mid_time, round(end_time, 3)]
|
||||
|
||||
try:
|
||||
frames_b64: list[str] = []
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
for i, ts in enumerate(timestamps):
|
||||
out_path = Path(tmpdir) / f"frame_{i}.jpg"
|
||||
cmd = [
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-ss",
|
||||
str(ts),
|
||||
"-i",
|
||||
video_url,
|
||||
"-vframes",
|
||||
"1",
|
||||
"-q:v",
|
||||
"2",
|
||||
str(out_path),
|
||||
]
|
||||
result = subprocess.run(
|
||||
cmd,
|
||||
capture_output=True,
|
||||
timeout=30,
|
||||
)
|
||||
if result.returncode != 0 or not out_path.exists():
|
||||
logger.warning("ffmpeg 抽帧失败 ts=%s: %s", ts, result.stderr[:200])
|
||||
continue
|
||||
|
||||
img_data = out_path.read_bytes()
|
||||
b64 = base64.b64encode(img_data).decode("ascii")
|
||||
frames_b64.append(f"data:image/jpeg;base64,{b64}")
|
||||
|
||||
if frames_b64:
|
||||
return frames_b64
|
||||
except Exception as e:
|
||||
logger.warning("ffmpeg 抽帧异常: %s", e)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def tag_atom_clip(
|
||||
clip: Any,
|
||||
video_url: str,
|
||||
doubao_client: Any,
|
||||
mediakit_client: Any | None = None,
|
||||
storage: Any | None = None,
|
||||
) -> dict:
|
||||
"""主入口:为单个 atom_clip 生成 AI 标签.
|
||||
|
||||
流程:提取帧 → 调视觉 API → 解析标签 → 返回结构化标签 dict。
|
||||
任何环节失败返回 {"inherited_tags": clip.tags},不阻断流程。
|
||||
|
||||
Args:
|
||||
clip: AssetAtomClip 领域对象(需有 start_time, end_time, tags)。
|
||||
video_url: 素材视频的公网可访问 URL。
|
||||
doubao_client: DoubaoClient 实例。
|
||||
mediakit_client: MediaKitClient 实例(可选,不可用时降级 ffmpeg)。
|
||||
storage: SharedStorageService 实例(可选,用于获取签名 URL)。
|
||||
|
||||
Returns:
|
||||
结构化标签 dict,格式如:
|
||||
{"scene": [...], "objects": [...], "action": [...], "shot": "...",
|
||||
"has_text": bool, "inherited_tags": [...]}
|
||||
"""
|
||||
inherited = list(getattr(clip, "tags", []) or [])
|
||||
|
||||
# 检查 DoubaoClient 是否可用
|
||||
if not getattr(doubao_client, "is_available", False):
|
||||
logger.info("DoubaoClient 不可用,跳过 AI 标签: clip_id=%s", getattr(clip, "id", ""))
|
||||
return {"inherited_tags": inherited}
|
||||
|
||||
# 提取帧图片
|
||||
frame_urls: Optional[list[str]] = None
|
||||
start_time = getattr(clip, "start_time", 0.0)
|
||||
end_time = getattr(clip, "end_time", 0.0)
|
||||
|
||||
# 优先使用 MediaKit
|
||||
if mediakit_client and getattr(mediakit_client, "is_available", False):
|
||||
frame_urls = _extract_frames_via_mediakit(mediakit_client, video_url, start_time, end_time)
|
||||
|
||||
# MediaKit 不可用或失败 → 降级 ffmpeg
|
||||
if not frame_urls:
|
||||
frame_urls = _extract_frames_via_ffmpeg(video_url, start_time, end_time)
|
||||
|
||||
if not frame_urls:
|
||||
logger.warning("帧提取失败,跳过 AI 标签: clip_id=%s", getattr(clip, "id", ""))
|
||||
return {"inherited_tags": inherited}
|
||||
|
||||
# 调用视觉 API
|
||||
prompt = build_vision_prompt()
|
||||
messages = [{"role": "user", "content": prompt}]
|
||||
|
||||
try:
|
||||
response_text = doubao_client.vision_completion(
|
||||
messages=messages,
|
||||
images=frame_urls,
|
||||
timeout=60,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("视觉 API 调用异常: clip_id=%s error=%s", getattr(clip, "id", ""), e)
|
||||
return {"inherited_tags": inherited}
|
||||
|
||||
if not response_text:
|
||||
logger.warning("视觉 API 返回空: clip_id=%s", getattr(clip, "id", ""))
|
||||
return {"inherited_tags": inherited}
|
||||
|
||||
# 解析标签
|
||||
ai_tags = parse_vision_response(response_text)
|
||||
if not ai_tags:
|
||||
logger.warning("标签解析失败: clip_id=%s response=%s", getattr(clip, "id", ""), response_text[:200])
|
||||
return {"inherited_tags": inherited}
|
||||
|
||||
# 合并 inherited_tags
|
||||
ai_tags["inherited_tags"] = inherited
|
||||
return ai_tags
|
||||
@@ -1,4 +1,4 @@
|
||||
"""叙事剪辑素材标签匹配 — #1970 PR3.
|
||||
"""叙事剪辑素材标签匹配 — #1970 PR3 + P2 AI 标签加权.
|
||||
|
||||
叙事模式下,选片在现有评分(smart_match / atom_clip_selector)之前先做一层
|
||||
文案标签匹配:
|
||||
@@ -8,6 +8,12 @@
|
||||
- 调用方对优先池跑现有 smart_select_assets,数量不足时用普通池补足
|
||||
(无任何匹配 → 完全降级为现有随机逻辑,行为与改造前一致)。
|
||||
|
||||
P2 AI 标签加权(#1970 fragment-level AI tagging):
|
||||
- 片段级 AI 标签(scene/objects/action)与文案标签做交集时权重 2.0
|
||||
- 素材级标签(tag_ids 映射名)与文案标签交集时权重 1.0
|
||||
- 综合得分 = sum(命中权重) / max(可能权重)
|
||||
- 有 AI 标签的片段命中时优先于仅素材标签命中的片段
|
||||
|
||||
纯函数模块:标签 id→名称映射由调用方查 TagModel 后注入,不直接碰 DB。
|
||||
"""
|
||||
|
||||
@@ -18,6 +24,10 @@ from typing import Any, Iterable
|
||||
# 标签归一化后仍短于此长度的标签不参与匹配(避免「的」「是」这类噪声短词)
|
||||
MIN_TAG_LEN = 2
|
||||
|
||||
# 标签匹配权重
|
||||
AI_TAG_WEIGHT = 2.0 # AI 标签命中权重
|
||||
ASSET_TAG_WEIGHT = 1.0 # 素材标签命中权重
|
||||
|
||||
|
||||
def normalize_tag(tag: Any) -> str:
|
||||
"""标签归一化:去空白、小写。数字/英文统一小写,中文不受影响。"""
|
||||
@@ -47,19 +57,81 @@ def build_asset_tag_name_index(tag_names_by_id: dict[str, Any]) -> dict[str, set
|
||||
return index
|
||||
|
||||
|
||||
def _extract_ai_tag_names(ai_tags: dict) -> set[str]:
|
||||
"""从 AI 标签 dict 中提取所有标签名(scene + objects + action).
|
||||
|
||||
Args:
|
||||
ai_tags: 片段级 AI 标签 dict,如 {"scene": [...], "objects": [...], "action": [...], ...}
|
||||
|
||||
Returns:
|
||||
归一化后的标签名集合。
|
||||
"""
|
||||
names: set[str] = set()
|
||||
for key in ("scene", "objects", "action"):
|
||||
values = ai_tags.get(key)
|
||||
if isinstance(values, list):
|
||||
names |= _normalize_tags(values)
|
||||
return names
|
||||
|
||||
|
||||
def _compute_ai_score(
|
||||
asset_id: str,
|
||||
wanted: set[str],
|
||||
clip_ai_tags_by_asset: dict[str, list[dict]] | None,
|
||||
) -> float:
|
||||
"""计算单个素材的 AI 标签加权得分.
|
||||
|
||||
对该素材的所有片段 AI 标签,求各片段标签名与文案标签交集的加权总和。
|
||||
每个片段的命中权重 = 命中数 × AI_TAG_WEIGHT。
|
||||
最终取所有片段的最高得分(而非累加,避免片段数多的素材不公平占优)。
|
||||
|
||||
Args:
|
||||
asset_id: 素材 ID。
|
||||
wanted: 归一化后的文案标签集合。
|
||||
clip_ai_tags_by_asset: {asset_id: [ai_tag_dict, ...]} 每个片段一个。
|
||||
|
||||
Returns:
|
||||
AI 标签加权得分(≥0)。
|
||||
"""
|
||||
if not clip_ai_tags_by_asset or not wanted:
|
||||
return 0.0
|
||||
|
||||
clips = clip_ai_tags_by_asset.get(asset_id)
|
||||
if not clips:
|
||||
return 0.0
|
||||
|
||||
best_score = 0.0
|
||||
for ai_tags in clips:
|
||||
if not ai_tags or not isinstance(ai_tags, dict):
|
||||
continue
|
||||
ai_names = _extract_ai_tag_names(ai_tags)
|
||||
hits = ai_names & wanted
|
||||
score = len(hits) * AI_TAG_WEIGHT
|
||||
if score > best_score:
|
||||
best_score = score
|
||||
|
||||
return best_score
|
||||
|
||||
|
||||
def match_assets_by_script_tags(
|
||||
assets: list[Any],
|
||||
*,
|
||||
script_tags: Iterable[Any],
|
||||
tag_names_by_id: dict[str, Any] | None = None,
|
||||
clip_ai_tags_by_asset: dict[str, list[dict]] | None = None,
|
||||
) -> tuple[list[Any], list[Any]]:
|
||||
"""按文案标签把素材拆成「命中池 / 未命中池」,保持输入相对顺序。
|
||||
|
||||
P2 加权逻辑:
|
||||
- AI 标签命中(scene/objects/action ∩ 文案标签)权重 2.0
|
||||
- 素材标签命中(tag_ids 映射名 ∩ 文案标签)权重 1.0
|
||||
- 任一权重 > 0 → 命中池,否则 → 未命中池
|
||||
|
||||
Args:
|
||||
assets: 候选素材(domain Asset,需有 id 与 tag_ids)。
|
||||
script_tags: 文案 tags(字符串数组,名称语义)。
|
||||
tag_names_by_id: asset_id → 素材标签名列表;素材只有 tag_ids 时由调用方
|
||||
查 TagModel 名称后传入。为空则视为无素材命中。
|
||||
tag_names_by_id: asset_id → 素材标签名列表。
|
||||
clip_ai_tags_by_asset: #1970 P2 — {asset_id: [ai_tag_dict, ...]}。
|
||||
|
||||
Returns:
|
||||
(matched, unmatched):命中任一文案标签的素材 / 其余素材。
|
||||
@@ -74,23 +146,74 @@ def match_assets_by_script_tags(
|
||||
unmatched: list[Any] = []
|
||||
for asset in assets:
|
||||
asset_id = str(getattr(asset, "id", "") or "")
|
||||
|
||||
# P2: AI 标签加权得分
|
||||
ai_score = _compute_ai_score(asset_id, wanted, clip_ai_tags_by_asset)
|
||||
|
||||
# 素材标签得分
|
||||
names = set(name_index.get(asset_id, set()))
|
||||
# 兼容素材自身带字符串 tags(旧链路/测试替身)
|
||||
raw_tags = getattr(asset, "tags", None)
|
||||
if raw_tags:
|
||||
names |= _normalize_tags(raw_tags)
|
||||
if names & wanted:
|
||||
asset_score = len(names & wanted) * ASSET_TAG_WEIGHT
|
||||
|
||||
# 综合得分 > 0 → 命中池
|
||||
if ai_score > 0 or asset_score > 0:
|
||||
matched.append(asset)
|
||||
else:
|
||||
unmatched.append(asset)
|
||||
return matched, unmatched
|
||||
|
||||
|
||||
def compute_tag_match_score(
|
||||
asset_id: str,
|
||||
*,
|
||||
script_tags: Iterable[Any],
|
||||
tag_names_by_id: dict[str, Any] | None = None,
|
||||
clip_ai_tags_by_asset: dict[str, list[dict]] | None = None,
|
||||
) -> float:
|
||||
"""计算单个素材的标签匹配综合得分(0.0 ~ 1.0).
|
||||
|
||||
综合得分 = sum(命中权重) / max(可能权重)
|
||||
- AI 标签每命中一个 +2.0
|
||||
- 素材标签每命中一个 +1.0
|
||||
- max_possible = len(wanted) * (AI_TAG_WEIGHT + ASSET_TAG_WEIGHT)
|
||||
|
||||
Args:
|
||||
asset_id: 素材 ID。
|
||||
script_tags: 文案标签。
|
||||
tag_names_by_id: 素材标签名索引。
|
||||
clip_ai_tags_by_asset: AI 标签索引。
|
||||
|
||||
Returns:
|
||||
归一化得分 0.0~1.0。
|
||||
"""
|
||||
wanted = _normalize_tags(script_tags)
|
||||
if not wanted:
|
||||
return 0.0
|
||||
|
||||
# AI 得分
|
||||
ai_score = _compute_ai_score(asset_id, wanted, clip_ai_tags_by_asset)
|
||||
|
||||
# 素材标签得分
|
||||
name_index = build_asset_tag_name_index(tag_names_by_id or {})
|
||||
names = name_index.get(asset_id, set())
|
||||
asset_score = len(names & wanted) * ASSET_TAG_WEIGHT
|
||||
|
||||
# 归一化:最大可能得分 = 文案标签数 × (AI权重 + 素材权重)
|
||||
max_possible = len(wanted) * (AI_TAG_WEIGHT + ASSET_TAG_WEIGHT)
|
||||
if max_possible <= 0:
|
||||
return 0.0
|
||||
|
||||
return min((ai_score + asset_score) / max_possible, 1.0)
|
||||
|
||||
|
||||
def pick_narrative_assets(
|
||||
assets: list[Any],
|
||||
*,
|
||||
script_tags: Iterable[Any],
|
||||
tag_names_by_id: dict[str, Any] | None = None,
|
||||
clip_ai_tags_by_asset: dict[str, list[dict]] | None = None,
|
||||
limit: int | None = None,
|
||||
rng: Any = None,
|
||||
) -> list[Any]:
|
||||
@@ -100,9 +223,13 @@ def pick_narrative_assets(
|
||||
smart_match.smart_select_assets(质量/时长/新鲜度/未使用 + 随机噪声),
|
||||
不重写评分维度。
|
||||
|
||||
P2 增强:有 AI 标签的片段命中时权重更高(2.0 vs 1.0),
|
||||
命中池内部按综合标签得分排序(AI 标签命中多的排前面)。
|
||||
|
||||
Args:
|
||||
assets: ready 视频素材候选(调用方负责状态/类型过滤)。
|
||||
script_tags / tag_names_by_id: 见 match_assets_by_script_tags。
|
||||
clip_ai_tags_by_asset: #1970 P2 — {asset_id: [ai_tag_dict, ...]}。
|
||||
limit: 需要的素材数量;None 表示全部(命中池 + 全部未命中池)。
|
||||
rng: 注入 smart_select_assets 的随机源(可复现)。
|
||||
|
||||
@@ -115,6 +242,7 @@ def pick_narrative_assets(
|
||||
assets,
|
||||
script_tags=script_tags,
|
||||
tag_names_by_id=tag_names_by_id,
|
||||
clip_ai_tags_by_asset=clip_ai_tags_by_asset,
|
||||
)
|
||||
|
||||
need = limit if (limit is not None and limit > 0) else None
|
||||
|
||||
@@ -37,6 +37,7 @@ class DoubaoClient:
|
||||
self.base_url: str = settings.doubao_base_url.rstrip("/")
|
||||
self.timeout: int = settings.doubao_timeout
|
||||
self.max_retries: int = settings.doubao_max_retries
|
||||
self.vision_model: str = settings.doubao_vision_model
|
||||
|
||||
@property
|
||||
def is_available(self) -> bool:
|
||||
@@ -103,6 +104,99 @@ class DoubaoClient:
|
||||
logger.error("豆包API调用最终失败: %s", last_error)
|
||||
return None
|
||||
|
||||
def vision_completion(
|
||||
self,
|
||||
messages: list[dict],
|
||||
images: list[str] | None = None,
|
||||
max_tokens: int = 2048,
|
||||
temperature: float = 0.3,
|
||||
timeout: int | None = None,
|
||||
) -> Optional[str]:
|
||||
"""调用豆包视觉理解 API(OpenAI 兼容多模态格式).
|
||||
|
||||
将 images 附加到最后一条 user message 的 content 中,
|
||||
使用 vision_model(默认 doubao-1-5-vision-pro-250915)。
|
||||
|
||||
Args:
|
||||
messages: 对话消息列表。最后一条 user message 会被注入图片内容。
|
||||
images: 图片列表,支持 base64 data URI 或 HTTP(S) URL。
|
||||
max_tokens: 最大生成 token 数,默认 2048。
|
||||
temperature: 采样温度,默认 0.3(视觉任务偏低更稳定)。
|
||||
timeout: 单次请求超时秒数,不传则使用默认 self.timeout。
|
||||
|
||||
Returns:
|
||||
模型返回的文本内容,失败返回 None。
|
||||
"""
|
||||
if not self.is_available:
|
||||
return None
|
||||
|
||||
# 构造多模态 content:先追加文本,再追加图片
|
||||
vision_messages = []
|
||||
for msg in messages:
|
||||
vision_messages.append(dict(msg))
|
||||
|
||||
# 将图片注入最后一条 user message
|
||||
if images and vision_messages:
|
||||
# 找到最后一条 user message
|
||||
for i in range(len(vision_messages) - 1, -1, -1):
|
||||
if vision_messages[i].get("role") == "user":
|
||||
text_content = vision_messages[i].get("content", "")
|
||||
multi_content: list[dict[str, Any]] = []
|
||||
if text_content:
|
||||
multi_content.append({"type": "text", "text": text_content})
|
||||
for img in images:
|
||||
if img.startswith("data:") or img.startswith("http://") or img.startswith("https://"):
|
||||
multi_content.append({"type": "image_url", "image_url": {"url": img}})
|
||||
else:
|
||||
# 当作 base64 编码
|
||||
multi_content.append(
|
||||
{"type": "image_url", "image_url": {"url": f"data:image/jpeg;base64,{img}"}}
|
||||
)
|
||||
vision_messages[i]["content"] = multi_content
|
||||
break
|
||||
|
||||
url = f"{self.base_url}/chat/completions"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
payload: dict[str, Any] = {
|
||||
"model": self.vision_model,
|
||||
"messages": vision_messages,
|
||||
"temperature": temperature,
|
||||
"max_tokens": max_tokens,
|
||||
}
|
||||
|
||||
req_timeout = timeout or self.timeout
|
||||
last_error: Optional[Exception] = None
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
response = httpx.post(
|
||||
url,
|
||||
headers=headers,
|
||||
json=payload,
|
||||
timeout=req_timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
content = data["choices"][0]["message"]["content"]
|
||||
return content.strip()
|
||||
except Exception as e:
|
||||
last_error = e
|
||||
if attempt < self.max_retries:
|
||||
wait = 0.5 * (2**attempt)
|
||||
logger.warning(
|
||||
"豆包视觉API调用失败,%.1fs后重试 (第%d/%d次): %s",
|
||||
wait,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
e,
|
||||
)
|
||||
time.sleep(wait)
|
||||
|
||||
logger.error("豆包视觉API调用最终失败: %s", last_error)
|
||||
return None
|
||||
|
||||
|
||||
# ── 单例 ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -339,6 +339,44 @@ class SharedStorageService(StoragePort):
|
||||
|
||||
# ── 浏览器直传 POST ────────────────────────────────────────────────
|
||||
|
||||
def get_upload_url(
|
||||
self,
|
||||
storage_key_or_url: str,
|
||||
expires_seconds: int = 3600,
|
||||
content_type: str = "video/mp4",
|
||||
) -> str:
|
||||
"""获取预签名 PUT 上传 URL(供外部 Worker 上传结果文件)。
|
||||
|
||||
bucket未配置时降级为 public_url(本地/开发环境);
|
||||
本地产物 key 原样返回。
|
||||
"""
|
||||
if self.bucket is None:
|
||||
if self._is_local_generated_url(storage_key_or_url):
|
||||
return storage_key_or_url
|
||||
logger.warning(
|
||||
"get_upload_url: OSS bucket not configured, returning raw URL. key=%s",
|
||||
storage_key_or_url[:200],
|
||||
)
|
||||
return self.get_url(self.normalize_storage_key(storage_key_or_url))
|
||||
|
||||
storage_key = self.normalize_storage_key(storage_key_or_url)
|
||||
try:
|
||||
# oss2 sign_url 支持 'PUT',需指定 headers 才能限定 Content-Type
|
||||
headers = {"Content-Type": content_type} if content_type else None
|
||||
signed = self.bucket.sign_url("PUT", storage_key, expires_seconds, headers=headers)
|
||||
logger.info(
|
||||
"get_upload_url: signed PUT URL generated. key=%s url_prefix=%s",
|
||||
storage_key[:80],
|
||||
signed[:60],
|
||||
)
|
||||
return signed
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"get_upload_url: sign_url failed, falling back to raw URL. key=%s",
|
||||
storage_key[:200],
|
||||
)
|
||||
return self.get_url(storage_key)
|
||||
|
||||
def create_direct_upload_post(
|
||||
self,
|
||||
storage_key: str,
|
||||
|
||||
@@ -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 DOUBAO_API_KEY DOUBAO_MODEL DOUBAO_BASE_URL WECHAT_APP_ID WECHAT_APP_SECRET TIKHUB_API_KEY APIZERO_API_KEY"
|
||||
SHARED_SECRETS="OSS_ACCESS_KEY_ID OSS_ACCESS_KEY_SECRET COSYVOICE_API_KEY DASHSCOPE_API_KEY MEDIAKIT_API_KEY DOUBAO_API_KEY DOUBAO_MODEL DOUBAO_BASE_URL WECHAT_APP_ID WECHAT_APP_SECRET TIKHUB_API_KEY APIZERO_API_KEY GPU_WORKER_TOKEN"
|
||||
for var in $SHARED_SECRETS; do
|
||||
value="${!var:-}"
|
||||
# 已经在环境中了,无需额外操作
|
||||
|
||||
@@ -0,0 +1,292 @@
|
||||
"""#1970 P2 片段级 AI 标签模块测试。
|
||||
|
||||
测试范围:
|
||||
- build_vision_prompt: 返回有效 prompt
|
||||
- parse_vision_response: 正常/异常/空值
|
||||
- tag_atom_clip: 成功/MediaKit不可用/视觉API失败/超时降级
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.atom_clip_tagger import (
|
||||
build_vision_prompt,
|
||||
parse_vision_response,
|
||||
tag_atom_clip,
|
||||
)
|
||||
|
||||
# ── Fake 对象 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeClip:
|
||||
id: str = "clip-001"
|
||||
asset_id: str = "asset-001"
|
||||
start_time: float = 0.0
|
||||
end_time: float = 5.0
|
||||
duration: float = 5.0
|
||||
clip_index: int = 0
|
||||
tags: list[str] = field(default_factory=lambda: ["tag1", "tag2"])
|
||||
ai_tags: dict | None = None
|
||||
|
||||
|
||||
class FakeDoubaoClient:
|
||||
"""模拟豆包客户端."""
|
||||
|
||||
def __init__(self, available: bool = True, response: str | None = None, raise_error: bool = False):
|
||||
self._available = available
|
||||
self._response = response
|
||||
self._raise_error = raise_error
|
||||
self.vision_calls: list[dict] = []
|
||||
|
||||
@property
|
||||
def is_available(self) -> bool:
|
||||
return self._available
|
||||
|
||||
def vision_completion(self, messages, images=None, timeout=None, **kwargs):
|
||||
self.vision_calls.append({"messages": messages, "images": images, "timeout": timeout})
|
||||
if self._raise_error:
|
||||
raise RuntimeError("API error")
|
||||
return self._response
|
||||
|
||||
|
||||
class FakeMediaKitClient:
|
||||
"""模拟 MediaKit 客户端."""
|
||||
|
||||
def __init__(self, available: bool = True, frames: list[dict] | None = None):
|
||||
self._available = available
|
||||
self._frames = frames
|
||||
|
||||
@property
|
||||
def is_available(self) -> bool:
|
||||
return self._available
|
||||
|
||||
def extract_frames(self, video_url, strategy=None, max_frames=None, **kwargs):
|
||||
return self._frames
|
||||
|
||||
|
||||
# ── build_vision_prompt ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildVisionPrompt:
|
||||
def test_returns_non_empty_string(self):
|
||||
prompt = build_vision_prompt()
|
||||
assert isinstance(prompt, str)
|
||||
assert len(prompt) > 100
|
||||
|
||||
def test_contains_required_keys(self):
|
||||
prompt = build_vision_prompt()
|
||||
assert "scene" in prompt
|
||||
assert "objects" in prompt
|
||||
assert "action" in prompt
|
||||
assert "shot" in prompt
|
||||
assert "has_text" in prompt
|
||||
|
||||
def test_requests_json_format(self):
|
||||
prompt = build_vision_prompt()
|
||||
assert "JSON" in prompt or "json" in prompt
|
||||
|
||||
|
||||
# ── parse_vision_response ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestParseVisionResponse:
|
||||
def test_valid_json(self):
|
||||
response = json.dumps(
|
||||
{
|
||||
"scene": ["工厂", "车间"],
|
||||
"objects": ["产品", "机器"],
|
||||
"action": ["演示"],
|
||||
"shot": "特写",
|
||||
"has_text": True,
|
||||
}
|
||||
)
|
||||
result = parse_vision_response(response)
|
||||
assert result["scene"] == ["工厂", "车间"]
|
||||
assert result["objects"] == ["产品", "机器"]
|
||||
assert result["action"] == ["演示"]
|
||||
assert result["shot"] == "特写"
|
||||
assert result["has_text"] is True
|
||||
|
||||
def test_json_with_markdown_code_block(self):
|
||||
response = '```json\n{"scene": ["办公室"], "objects": ["电脑"], "action": ["说话"], "shot": "中景", "has_text": false}\n```'
|
||||
result = parse_vision_response(response)
|
||||
assert result["scene"] == ["办公室"]
|
||||
assert result["has_text"] is False
|
||||
|
||||
def test_json_embedded_in_text(self):
|
||||
response = '这是一些说明文字\n{"scene": ["户外"], "objects": ["汽车"], "action": ["展示"], "shot": "远景", "has_text": false}\n结束'
|
||||
result = parse_vision_response(response)
|
||||
assert result["scene"] == ["户外"]
|
||||
|
||||
def test_empty_response(self):
|
||||
assert parse_vision_response("") == {}
|
||||
assert parse_vision_response(None) == {}
|
||||
assert parse_vision_response(" ") == {}
|
||||
|
||||
def test_invalid_json(self):
|
||||
assert parse_vision_response("这不是JSON") == {}
|
||||
|
||||
def test_partial_fields(self):
|
||||
response = json.dumps({"scene": ["工厂"]})
|
||||
result = parse_vision_response(response)
|
||||
assert result["scene"] == ["工厂"]
|
||||
assert result["objects"] == []
|
||||
assert result["shot"] == ""
|
||||
assert result["has_text"] is False
|
||||
|
||||
def test_invalid_shot_value(self):
|
||||
response = json.dumps({"scene": [], "objects": [], "action": [], "shot": "全景", "has_text": False})
|
||||
result = parse_vision_response(response)
|
||||
# "全景" 不在有效值 ("特写", "中景", "远景") 中
|
||||
assert result["shot"] == ""
|
||||
|
||||
def test_string_values_converted_to_list(self):
|
||||
response = json.dumps(
|
||||
{"scene": "工厂", "objects": "产品", "action": "演示", "shot": "特写", "has_text": "true"}
|
||||
)
|
||||
result = parse_vision_response(response)
|
||||
assert result["scene"] == ["工厂"]
|
||||
assert result["objects"] == ["产品"]
|
||||
assert result["has_text"] is True
|
||||
|
||||
def test_non_dict_json(self):
|
||||
assert parse_vision_response("[1, 2, 3]") == {}
|
||||
assert parse_vision_response('"hello"') == {}
|
||||
|
||||
|
||||
# ── tag_atom_clip ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestTagAtomClip:
|
||||
def test_success_with_mediakit(self):
|
||||
"""MediaKit 可用 + 视觉 API 成功 → 返回完整 AI 标签."""
|
||||
clip = FakeClip()
|
||||
fake_doubao = FakeDoubaoClient(
|
||||
response=json.dumps(
|
||||
{
|
||||
"scene": ["工厂"],
|
||||
"objects": ["产品"],
|
||||
"action": ["演示"],
|
||||
"shot": "特写",
|
||||
"has_text": False,
|
||||
}
|
||||
)
|
||||
)
|
||||
fake_mediakit = FakeMediaKitClient(
|
||||
frames=[
|
||||
{"image_url": "https://example.com/frame1.jpg", "timestamp": 0.0},
|
||||
{"image_url": "https://example.com/frame2.jpg", "timestamp": 2.5},
|
||||
{"image_url": "https://example.com/frame3.jpg", "timestamp": 5.0},
|
||||
]
|
||||
)
|
||||
|
||||
result = tag_atom_clip(
|
||||
clip=clip,
|
||||
video_url="https://example.com/video.mp4",
|
||||
doubao_client=fake_doubao,
|
||||
mediakit_client=fake_mediakit,
|
||||
)
|
||||
|
||||
assert result["scene"] == ["工厂"]
|
||||
assert result["objects"] == ["产品"]
|
||||
assert result["shot"] == "特写"
|
||||
assert result["inherited_tags"] == ["tag1", "tag2"]
|
||||
assert len(fake_doubao.vision_calls) == 1
|
||||
|
||||
def test_doubao_unavailable_returns_inherited(self):
|
||||
"""DoubaoClient 不可用 → 返回 inherited_tags."""
|
||||
clip = FakeClip()
|
||||
fake_doubao = FakeDoubaoClient(available=False)
|
||||
|
||||
result = tag_atom_clip(
|
||||
clip=clip,
|
||||
video_url="https://example.com/video.mp4",
|
||||
doubao_client=fake_doubao,
|
||||
)
|
||||
|
||||
assert result == {"inherited_tags": ["tag1", "tag2"]}
|
||||
assert len(fake_doubao.vision_calls) == 0
|
||||
|
||||
def test_mediakit_unavailable_no_ffmpeg(self):
|
||||
"""MediaKit 不可用 + 无 ffmpeg → 降级 inherited_tags."""
|
||||
clip = FakeClip()
|
||||
fake_doubao = FakeDoubaoClient()
|
||||
fake_mediakit = FakeMediaKitClient(available=False)
|
||||
|
||||
result = tag_atom_clip(
|
||||
clip=clip,
|
||||
video_url="https://example.com/video.mp4",
|
||||
doubao_client=fake_doubao,
|
||||
mediakit_client=fake_mediakit,
|
||||
)
|
||||
|
||||
# 没有 ffmpeg 的情况下,帧提取失败
|
||||
assert result == {"inherited_tags": ["tag1", "tag2"]}
|
||||
|
||||
def test_vision_api_error_returns_inherited(self):
|
||||
"""视觉 API 抛异常 → 降级 inherited_tags."""
|
||||
clip = FakeClip()
|
||||
fake_doubao = FakeDoubaoClient(raise_error=True)
|
||||
fake_mediakit = FakeMediaKitClient(frames=[{"image_url": "https://example.com/frame.jpg", "timestamp": 0.0}])
|
||||
|
||||
result = tag_atom_clip(
|
||||
clip=clip,
|
||||
video_url="https://example.com/video.mp4",
|
||||
doubao_client=fake_doubao,
|
||||
mediakit_client=fake_mediakit,
|
||||
)
|
||||
|
||||
assert result == {"inherited_tags": ["tag1", "tag2"]}
|
||||
|
||||
def test_vision_api_empty_response(self):
|
||||
"""视觉 API 返回空 → 降级 inherited_tags."""
|
||||
clip = FakeClip()
|
||||
fake_doubao = FakeDoubaoClient(response=None)
|
||||
fake_mediakit = FakeMediaKitClient(frames=[{"image_url": "https://example.com/frame.jpg", "timestamp": 0.0}])
|
||||
|
||||
result = tag_atom_clip(
|
||||
clip=clip,
|
||||
video_url="https://example.com/video.mp4",
|
||||
doubao_client=fake_doubao,
|
||||
mediakit_client=fake_mediakit,
|
||||
)
|
||||
|
||||
assert result == {"inherited_tags": ["tag1", "tag2"]}
|
||||
|
||||
def test_vision_api_invalid_json_response(self):
|
||||
"""视觉 API 返回无效 JSON → 降级 inherited_tags."""
|
||||
clip = FakeClip()
|
||||
fake_doubao = FakeDoubaoClient(response="这不是JSON格式")
|
||||
fake_mediakit = FakeMediaKitClient(frames=[{"image_url": "https://example.com/frame.jpg", "timestamp": 0.0}])
|
||||
|
||||
result = tag_atom_clip(
|
||||
clip=clip,
|
||||
video_url="https://example.com/video.mp4",
|
||||
doubao_client=fake_doubao,
|
||||
mediakit_client=fake_mediakit,
|
||||
)
|
||||
|
||||
assert result == {"inherited_tags": ["tag1", "tag2"]}
|
||||
|
||||
def test_clip_with_empty_tags(self):
|
||||
"""空素材标签 → inherited_tags 为空列表."""
|
||||
clip = FakeClip(tags=[])
|
||||
fake_doubao = FakeDoubaoClient(available=False)
|
||||
|
||||
result = tag_atom_clip(
|
||||
clip=clip,
|
||||
video_url="https://example.com/video.mp4",
|
||||
doubao_client=fake_doubao,
|
||||
)
|
||||
|
||||
assert result == {"inherited_tags": []}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-q"])
|
||||
@@ -0,0 +1,306 @@
|
||||
"""#1970 P2 叙事匹配 AI 标签加权测试。
|
||||
|
||||
测试范围:
|
||||
- AI 标签命中时权重 2.0
|
||||
- 无 AI 标签时降级到素材标签权重 1.0
|
||||
- 混合场景(部分素材有 AI 标签,部分只有素材标签)
|
||||
- compute_tag_match_score 归一化得分
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime as dt
|
||||
import random
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.narrative_match import (
|
||||
AI_TAG_WEIGHT,
|
||||
ASSET_TAG_WEIGHT,
|
||||
_compute_ai_score,
|
||||
_extract_ai_tag_names,
|
||||
compute_tag_match_score,
|
||||
match_assets_by_script_tags,
|
||||
pick_narrative_assets,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeAsset:
|
||||
id: str
|
||||
tag_ids: list[str] = field(default_factory=list)
|
||||
tags: list[str] = field(default_factory=list)
|
||||
status: str = "ready"
|
||||
file_type: str = "video"
|
||||
duration: float = 10.0
|
||||
quality_score: float | None = None
|
||||
created_at: object = None
|
||||
metadata: dict = field(default_factory=dict)
|
||||
|
||||
|
||||
def _make_old_dt():
|
||||
return dt.datetime(2020, 1, 1, tzinfo=dt.UTC)
|
||||
|
||||
|
||||
# ── _extract_ai_tag_names ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestExtractAiTagNames:
|
||||
def test_extracts_all_keys(self):
|
||||
ai_tags = {
|
||||
"scene": ["工厂", "车间"],
|
||||
"objects": ["产品"],
|
||||
"action": ["演示"],
|
||||
"shot": "特写", # shot 不参与标签匹配
|
||||
"has_text": False,
|
||||
}
|
||||
names = _extract_ai_tag_names(ai_tags)
|
||||
assert names == {"工厂", "车间", "产品", "演示"}
|
||||
|
||||
def test_empty_dict(self):
|
||||
assert _extract_ai_tag_names({}) == set()
|
||||
|
||||
def test_none_values(self):
|
||||
ai_tags = {"scene": None, "objects": None, "action": None}
|
||||
assert _extract_ai_tag_names(ai_tags) == set()
|
||||
|
||||
def test_case_insensitive(self):
|
||||
ai_tags = {"scene": ["Factory"], "objects": [], "action": []}
|
||||
names = _extract_ai_tag_names(ai_tags)
|
||||
assert "factory" in names
|
||||
|
||||
|
||||
# ── _compute_ai_score ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestComputeAiScore:
|
||||
def test_single_clip_hit(self):
|
||||
wanted = {"工厂", "演示"}
|
||||
clips = [{"scene": ["工厂"], "objects": [], "action": ["演示"]}]
|
||||
score = _compute_ai_score("a1", wanted, {"a1": clips})
|
||||
# 命中 2 个 × 2.0 = 4.0
|
||||
assert score == 2 * AI_TAG_WEIGHT
|
||||
|
||||
def test_multiple_clips_takes_best(self):
|
||||
wanted = {"工厂", "演示"}
|
||||
clips = [
|
||||
{"scene": ["工厂"], "objects": [], "action": []}, # 1 hit = 2.0
|
||||
{"scene": ["工厂"], "objects": [], "action": ["演示"]}, # 2 hits = 4.0
|
||||
]
|
||||
score = _compute_ai_score("a1", wanted, {"a1": clips})
|
||||
assert score == 2 * AI_TAG_WEIGHT # best = 2 hits
|
||||
|
||||
def test_no_match(self):
|
||||
wanted = {"美食"}
|
||||
clips = [{"scene": ["工厂"], "objects": [], "action": ["演示"]}]
|
||||
score = _compute_ai_score("a1", wanted, {"a1": clips})
|
||||
assert score == 0.0
|
||||
|
||||
def test_no_clips_for_asset(self):
|
||||
wanted = {"工厂"}
|
||||
assert _compute_ai_score("a1", wanted, {}) == 0.0
|
||||
assert _compute_ai_score("a1", wanted, None) == 0.0
|
||||
|
||||
def test_empty_wanted(self):
|
||||
clips = [{"scene": ["工厂"], "objects": [], "action": []}]
|
||||
assert _compute_ai_score("a1", set(), {"a1": clips}) == 0.0
|
||||
|
||||
|
||||
# ── match_assets_by_script_tags with AI tags ──────────────────────────────
|
||||
|
||||
|
||||
class TestMatchWithAiTags:
|
||||
def test_ai_tag_hit_puts_in_matched(self):
|
||||
"""有 AI 标签命中 → 进入命中池."""
|
||||
assets = [FakeAsset("a1", created_at=_make_old_dt())]
|
||||
clip_ai_tags = {"a1": [{"scene": ["工厂"], "objects": [], "action": []}]}
|
||||
|
||||
matched, unmatched = match_assets_by_script_tags(
|
||||
assets,
|
||||
script_tags=["工厂"],
|
||||
clip_ai_tags_by_asset=clip_ai_tags,
|
||||
)
|
||||
|
||||
assert [a.id for a in matched] == ["a1"]
|
||||
assert unmatched == []
|
||||
|
||||
def test_ai_tag_no_match_puts_in_unmatched(self):
|
||||
"""AI 标签未命中 → 进入未命中池."""
|
||||
assets = [FakeAsset("a1", created_at=_make_old_dt())]
|
||||
clip_ai_tags = {"a1": [{"scene": ["办公室"], "objects": [], "action": []}]}
|
||||
|
||||
matched, unmatched = match_assets_by_script_tags(
|
||||
assets,
|
||||
script_tags=["工厂"],
|
||||
clip_ai_tags_by_asset=clip_ai_tags,
|
||||
)
|
||||
|
||||
assert matched == []
|
||||
assert [a.id for a in unmatched] == ["a1"]
|
||||
|
||||
def test_asset_tag_still_works_without_ai_tags(self):
|
||||
"""无 AI 标签时,素材标签仍按权重 1.0 匹配."""
|
||||
assets = [FakeAsset("a1", tags=["工厂"], created_at=_make_old_dt())]
|
||||
|
||||
matched, unmatched = match_assets_by_script_tags(
|
||||
assets,
|
||||
script_tags=["工厂"],
|
||||
)
|
||||
|
||||
assert [a.id for a in matched] == ["a1"]
|
||||
|
||||
def test_mixed_ai_and_asset_tags(self):
|
||||
"""混合场景:一个素材有 AI 标签,另一个只有素材标签."""
|
||||
assets = [
|
||||
FakeAsset("a1", created_at=_make_old_dt()), # AI 标签命中
|
||||
FakeAsset("a2", tags=["工厂"], created_at=_make_old_dt()), # 素材标签命中
|
||||
FakeAsset("a3", tags=["美食"], created_at=_make_old_dt()), # 无命中
|
||||
]
|
||||
clip_ai_tags = {"a1": [{"scene": ["工厂"], "objects": [], "action": []}]}
|
||||
|
||||
matched, unmatched = match_assets_by_script_tags(
|
||||
assets,
|
||||
script_tags=["工厂"],
|
||||
clip_ai_tags_by_asset=clip_ai_tags,
|
||||
)
|
||||
|
||||
assert {a.id for a in matched} == {"a1", "a2"}
|
||||
assert [a.id for a in unmatched] == ["a3"]
|
||||
|
||||
def test_ai_tag_and_asset_tag_both_hit(self):
|
||||
"""同一素材 AI 标签和素材标签都命中 → 仍在命中池."""
|
||||
assets = [FakeAsset("a1", tags=["工厂"], created_at=_make_old_dt())]
|
||||
clip_ai_tags = {"a1": [{"scene": ["工厂"], "objects": [], "action": []}]}
|
||||
|
||||
matched, unmatched = match_assets_by_script_tags(
|
||||
assets,
|
||||
script_tags=["工厂"],
|
||||
tag_names_by_id={"a1": ["工厂"]},
|
||||
clip_ai_tags_by_asset=clip_ai_tags,
|
||||
)
|
||||
|
||||
assert [a.id for a in matched] == ["a1"]
|
||||
|
||||
|
||||
# ── compute_tag_match_score ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestComputeTagMatchScore:
|
||||
def test_ai_only_score(self):
|
||||
"""仅 AI 标签命中."""
|
||||
clip_ai_tags = {"a1": [{"scene": ["工厂"], "objects": [], "action": ["演示"]}]}
|
||||
score = compute_tag_match_score(
|
||||
"a1",
|
||||
script_tags=["工厂", "演示"],
|
||||
clip_ai_tags_by_asset=clip_ai_tags,
|
||||
)
|
||||
# AI: 2 hits × 2.0 = 4.0; asset: 0; max = 2 × 3.0 = 6.0
|
||||
assert abs(score - 4.0 / 6.0) < 0.01
|
||||
|
||||
def test_asset_only_score(self):
|
||||
"""仅素材标签命中."""
|
||||
score = compute_tag_match_score(
|
||||
"a1",
|
||||
script_tags=["工厂", "演示"],
|
||||
tag_names_by_id={"a1": ["工厂"]},
|
||||
)
|
||||
# AI: 0; asset: 1 hit × 1.0 = 1.0; max = 2 × 3.0 = 6.0
|
||||
assert abs(score - 1.0 / 6.0) < 0.01
|
||||
|
||||
def test_both_ai_and_asset_score(self):
|
||||
"""AI 标签 + 素材标签同时命中."""
|
||||
clip_ai_tags = {"a1": [{"scene": ["工厂"], "objects": [], "action": []}]}
|
||||
score = compute_tag_match_score(
|
||||
"a1",
|
||||
script_tags=["工厂", "演示"],
|
||||
tag_names_by_id={"a1": ["工厂"]},
|
||||
clip_ai_tags_by_asset=clip_ai_tags,
|
||||
)
|
||||
# AI: 1 hit × 2.0 = 2.0; asset: 1 hit × 1.0 = 1.0; max = 2 × 3.0 = 6.0
|
||||
assert abs(score - 3.0 / 6.0) < 0.01
|
||||
|
||||
def test_no_match_score_zero(self):
|
||||
"""无命中 → 得分 0."""
|
||||
score = compute_tag_match_score(
|
||||
"a1",
|
||||
script_tags=["工厂"],
|
||||
tag_names_by_id={"a1": ["美食"]},
|
||||
)
|
||||
assert score == 0.0
|
||||
|
||||
def test_full_match_score_one(self):
|
||||
"""全命中 → 得分接近 1.0."""
|
||||
clip_ai_tags = {"a1": [{"scene": ["工厂"], "objects": ["产品"], "action": ["演示"]}]}
|
||||
score = compute_tag_match_score(
|
||||
"a1",
|
||||
script_tags=["工厂", "产品", "演示"],
|
||||
clip_ai_tags_by_asset=clip_ai_tags,
|
||||
)
|
||||
# AI: 3 hits × 2.0 = 6.0; max = 3 × 3.0 = 9.0 → 6/9 = 0.667
|
||||
# 注意:仅 AI 标签命中不可能达到 1.0(因为 max 包含素材权重)
|
||||
assert score > 0.5
|
||||
|
||||
def test_empty_script_tags(self):
|
||||
"""空文案标签 → 得分 0."""
|
||||
assert compute_tag_match_score("a1", script_tags=[]) == 0.0
|
||||
|
||||
|
||||
# ── pick_narrative_assets with AI tags ────────────────────────────────────
|
||||
|
||||
|
||||
class TestPickNarrativeWithAiTags:
|
||||
def _assets(self):
|
||||
old = _make_old_dt()
|
||||
return [
|
||||
FakeAsset("ai_match", created_at=old), # AI 标签命中
|
||||
FakeAsset("asset_match", tags=["工厂"], created_at=old), # 素材标签命中
|
||||
FakeAsset("no_match", tags=["美食"], created_at=old), # 无命中
|
||||
]
|
||||
|
||||
def test_ai_match_prioritized(self):
|
||||
"""AI 标签命中的素材进入命中池."""
|
||||
clip_ai_tags = {"ai_match": [{"scene": ["工厂"], "objects": [], "action": []}]}
|
||||
|
||||
picked = pick_narrative_assets(
|
||||
self._assets(),
|
||||
script_tags=["工厂"],
|
||||
clip_ai_tags_by_asset=clip_ai_tags,
|
||||
limit=2,
|
||||
rng=random.Random(0),
|
||||
)
|
||||
|
||||
ids = {a.id for a in picked}
|
||||
assert "ai_match" in ids
|
||||
assert "asset_match" in ids
|
||||
|
||||
def test_fallback_when_no_ai_match(self):
|
||||
"""AI 标签和素材标签都未命中 → 降级."""
|
||||
clip_ai_tags = {"ai_match": [{"scene": ["办公室"], "objects": [], "action": []}]}
|
||||
|
||||
picked = pick_narrative_assets(
|
||||
self._assets(),
|
||||
script_tags=["不存在"],
|
||||
clip_ai_tags_by_asset=clip_ai_tags,
|
||||
limit=2,
|
||||
rng=random.Random(0),
|
||||
)
|
||||
|
||||
assert len(picked) == 2 # 从全量中选取
|
||||
|
||||
def test_backward_compat_without_ai_tags(self):
|
||||
"""不传 clip_ai_tags_by_asset 时行为与之前完全一致."""
|
||||
picked = pick_narrative_assets(
|
||||
self._assets(),
|
||||
script_tags=["工厂"],
|
||||
limit=2,
|
||||
rng=random.Random(0),
|
||||
)
|
||||
|
||||
# 仅素材标签匹配
|
||||
ids = {a.id for a in picked}
|
||||
assert "asset_match" in ids
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-q"])
|
||||
@@ -0,0 +1,191 @@
|
||||
"""GpuLipsyncService 单元测试 — 覆盖任务创建、轮询认领、结果上报、超时回退等核心逻辑.
|
||||
|
||||
使用 SQLite 内存数据库,mock 掉存储层(不真实调用 OSS)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from unittest import mock
|
||||
|
||||
import pytest
|
||||
|
||||
# 确保 packages / apps/api 可导入
|
||||
ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", ".."))
|
||||
for p in (ROOT, os.path.join(ROOT, "apps", "api"), os.path.join(ROOT, "packages")):
|
||||
if p not in sys.path:
|
||||
sys.path.insert(0, p)
|
||||
|
||||
# 强制使用内存 SQLite(避免依赖 PG)
|
||||
os.environ["APP_ENV"] = "development"
|
||||
os.environ["JWT_SECRET_KEY"] = "dev-secret-key-for-testing-00000000"
|
||||
os.environ["DATABASE_URL"] = "sqlite:///:memory:"
|
||||
os.environ["USE_IN_MEMORY_DB"] = "1"
|
||||
os.environ["GPU_WORKER_TOKEN"] = "" # development 空 token 放行
|
||||
|
||||
|
||||
def _build_session():
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
# 使用 packages 的 Base
|
||||
from packages.adapters.sqlalchemy_impl import models as _ # noqa: F401 # 触发 ORM 注册
|
||||
from packages.adapters.sqlalchemy_impl.models import Base
|
||||
|
||||
engine = create_engine("sqlite:///:memory:", future=True)
|
||||
Base.metadata.create_all(engine)
|
||||
Session = sessionmaker(bind=engine, autoflush=False, autocommit=False, future=True)
|
||||
return Session()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def svc():
|
||||
from app.services.gpu_lipsync_service import GpuLipsyncService
|
||||
|
||||
db = _build_session()
|
||||
service = GpuLipsyncService(db)
|
||||
# mock 存储签名(SQLite 测试无 OSS)
|
||||
service.storage = mock.MagicMock()
|
||||
service.storage.get_download_url.side_effect = (
|
||||
lambda k, expires_seconds=3600: f"https://signed.example.com/download/{k}?e={expires_seconds}"
|
||||
)
|
||||
service.storage.get_upload_url.side_effect = (
|
||||
lambda k, expires_seconds=3600, content_type="video/mp4": f"https://signed.example.com/upload/{k}?e={expires_seconds}"
|
||||
)
|
||||
return service
|
||||
|
||||
|
||||
# ── 创建任务 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_create_task(svc):
|
||||
task = svc.create_task(
|
||||
video_url="uploads/v.mp4",
|
||||
audio_url="uploads/a.mp3",
|
||||
lipsync_job_id="lip-1",
|
||||
user_id="u-1",
|
||||
project_id="p-1",
|
||||
)
|
||||
assert task.id
|
||||
assert task.status == "pending"
|
||||
assert task.lipsync_job_id == "lip-1"
|
||||
assert task.attempt == 0
|
||||
assert task.video_url == "uploads/v.mp4"
|
||||
|
||||
|
||||
# ── 轮询认领 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_poll_returns_none_when_empty(svc):
|
||||
assert svc.poll_task("w-1") is None
|
||||
|
||||
|
||||
def test_poll_claims_pending_task(svc):
|
||||
svc.create_task(video_url="uploads/v.mp4", audio_url="uploads/a.mp3")
|
||||
claimed = svc.poll_task("w-1")
|
||||
assert claimed is not None
|
||||
assert claimed.status == "processing"
|
||||
assert claimed.worker_id == "w-1"
|
||||
assert claimed.attempt == 1
|
||||
# 带签名 URL
|
||||
assert claimed._signed_video_url.startswith("https://signed.example.com/download/")
|
||||
assert claimed._signed_upload_url.startswith("https://signed.example.com/upload/")
|
||||
# 再 poll 无任务
|
||||
assert svc.poll_task("w-1") is None
|
||||
|
||||
|
||||
def test_poll_concurrent_claim_only_one_wins(svc):
|
||||
"""并发场景:两个 worker 同时 poll 只有一个能拿到任务(借助 update where status=pending)。"""
|
||||
svc.create_task(video_url="v", audio_url="a")
|
||||
t1 = svc.poll_task("w-1")
|
||||
t2 = svc.poll_task("w-2")
|
||||
assert t1 is not None
|
||||
assert t2 is None
|
||||
|
||||
|
||||
# ── 结果上报 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_report_result_success(svc):
|
||||
t = svc.create_task(video_url="v", audio_url="a")
|
||||
svc.poll_task("w-1") # claim
|
||||
done = svc.report_result(t.id, "w-1", success=True, duration_seconds=12.5)
|
||||
assert done.status == "done"
|
||||
assert done.result_duration == 12.5
|
||||
assert done.result_url.startswith("gpu-lipsync/results/")
|
||||
assert done.finished_at is not None
|
||||
|
||||
|
||||
def test_report_result_failure_requeues(svc):
|
||||
t = svc.create_task(video_url="v", audio_url="a")
|
||||
svc.poll_task("w-1")
|
||||
failed = svc.report_result(t.id, "w-1", success=False, error_msg="MuseTalk crash")
|
||||
assert failed.status == "pending" # 仍在重试次数内 → 回队
|
||||
assert failed.worker_id == ""
|
||||
assert failed.started_at is None
|
||||
assert "MuseTalk crash" in failed.error_msg
|
||||
|
||||
|
||||
def test_report_failure_exhausted_goes_failed(svc):
|
||||
"""失败达到 MAX_ATTEMPTS 后标记 failed,不再回队.
|
||||
|
||||
poll 成功会将 attempt 从 0 开始自增;
|
||||
第 1/2 次失败回队,第 3 次失败(attempt==MAX_ATTEMPTS)置 failed。
|
||||
"""
|
||||
from app.services import gpu_lipsync_service as mod
|
||||
|
||||
t = svc.create_task(video_url="v", audio_url="a")
|
||||
# 模拟失败到上限:poll + fail 重复 MAX_ATTEMPTS 次
|
||||
for i in range(mod.MAX_ATTEMPTS):
|
||||
claimed = svc.poll_task(f"w-{i}")
|
||||
assert claimed is not None, f"第 {i} 次 poll 应能拿到任务"
|
||||
svc.report_result(t.id, claimed.worker_id, success=False, error_msg=f"fail {i}")
|
||||
svc.db.refresh(t)
|
||||
if i == mod.MAX_ATTEMPTS - 1:
|
||||
assert t.status == "failed"
|
||||
else:
|
||||
assert t.status == "pending"
|
||||
|
||||
|
||||
# ── 心跳/超时回退 ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_timed_out_task_is_redispatched(svc):
|
||||
"""processing 超过 gpu_task_timeout_seconds 无心跳 → 回退 pending."""
|
||||
t = svc.create_task(video_url="v", audio_url="a")
|
||||
svc.poll_task("w-1")
|
||||
svc.db.refresh(t)
|
||||
assert t.status == "processing"
|
||||
# 手动把 last_heartbeat_at 设到很久以前
|
||||
t.last_heartbeat_at = datetime.now(UTC) - timedelta(seconds=svc.settings.gpu_task_timeout_seconds + 10)
|
||||
svc.db.commit()
|
||||
# 再次 poll 会触发 _recover_timed_out_tasks 把它回队
|
||||
claimed = svc.poll_task("w-2")
|
||||
assert claimed is not None
|
||||
assert claimed.id == t.id
|
||||
assert claimed.worker_id == "w-2"
|
||||
assert claimed.attempt == 2 # 又认领了一次
|
||||
|
||||
|
||||
# ── Worker 注册 ───────────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_register_worker_creates_then_updates(svc):
|
||||
w = svc.register_worker("w-1", hostname="pc1", gpu_name="RTX2060", free_vram_mb=3500)
|
||||
assert w.worker_id == "w-1"
|
||||
assert w.gpu_name == "RTX2060"
|
||||
w2 = svc.register_worker("w-1", hostname="pc1", gpu_name="RTX2060", free_vram_mb=2000)
|
||||
assert w2.free_vram_mb == 2000 # 更新
|
||||
assert w2.created_at == w.created_at # 没新建
|
||||
|
||||
|
||||
# ── get_by_lipsync_job ─────────────────────────────────────────────
|
||||
|
||||
|
||||
def test_get_by_lipsync_job_returns_latest(svc):
|
||||
svc.create_task(video_url="v", audio_url="a", lipsync_job_id="lip-1")
|
||||
svc.create_task(video_url="v", audio_url="a", lipsync_job_id="lip-1")
|
||||
latest = svc.get_by_lipsync_job("lip-1")
|
||||
assert latest is not None
|
||||
Reference in New Issue
Block a user