Compare commits
4 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 9464322710 | |||
| aa1f318308 | |||
| 3a59948f53 | |||
| 3e6f87a8b5 |
@@ -223,7 +223,47 @@ class LipsyncService:
|
||||
if timings:
|
||||
job.sentence_timings = timings
|
||||
|
||||
# 4. 检查是否走 GPU 路径:开关打开 + 有可用 Worker
|
||||
# 4. 检查是否走 Ditto(蚂蚁数字人,#2076):开关 + 配置完整
|
||||
use_ditto = False
|
||||
if self.settings.use_ditto_lipsync:
|
||||
try:
|
||||
from packages.application.ditto_service import get_ditto_client
|
||||
|
||||
ditto = get_ditto_client()
|
||||
if ditto.is_configured:
|
||||
use_ditto = True
|
||||
logger.info("[lipsync] 优先走 Ditto 蚂蚁数字人: job_id=%s", job.id)
|
||||
else:
|
||||
logger.info(
|
||||
"[lipsync] Ditto 开关已开但配置不完整(base_url=%s, template=%s),继续判断 GPU: job_id=%s",
|
||||
bool(ditto.base_url),
|
||||
bool(ditto.default_video_url),
|
||||
job.id,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("[lipsync] Ditto 初始化失败,继续判断 GPU: job_id=%s err=%s", job.id, exc)
|
||||
|
||||
if use_ditto:
|
||||
try:
|
||||
# Ditto 使用预置人物模板视频,不用用户上传的 video_url;
|
||||
# 但保留用户 video_url 以便失败回退到 GPU/MediaKit。
|
||||
job.status = "processing"
|
||||
job.mediakit_task_id = "ditto:submitted"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
self.db.commit()
|
||||
from app.tasks.lipsync_ditto import lipsync_ditto_process_async
|
||||
|
||||
lipsync_ditto_process_async.apply_async(args=(job.id, job.user_id))
|
||||
logger.info("[lipsync] Ditto 任务已异步派发: job_id=%s", job.id)
|
||||
return
|
||||
except Exception as exc:
|
||||
logger.warning("[lipsync] Ditto 派发失败,回退 GPU/MediaKit: job_id=%s err=%s", job.id, exc)
|
||||
try:
|
||||
self.db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 5. 检查是否走 GPU 路径:开关打开 + 有可用 Worker
|
||||
use_gpu = False
|
||||
if self.settings.use_gpu_lipsync:
|
||||
try:
|
||||
|
||||
@@ -0,0 +1,318 @@
|
||||
"""Ditto 蚂蚁数字人口型异步任务 — #2076.
|
||||
|
||||
把 Ditto 同步 HTTP 调用(30-120s)从 API 请求移到 Celery 后台执行:
|
||||
1. 加载 LipsyncJob
|
||||
2. 调 DittoClient.generate_and_persist(video_url=默认模板, audio_url=job.audio_url, script=job.script_text)
|
||||
3. 成功:标记 completed,写入 output_video_url(Ditto 输出自带音频,无需二次混流/超分)
|
||||
4. 失败:回退 GPU MuseTalk → 再失败回退 MediaKit
|
||||
|
||||
注意:
|
||||
- 保留 MuseTalk 代码不动;Ditto 优先,失败按原链路兜底
|
||||
- Ditto 使用预置的人物模板视频(settings.ditto_default_video_url),不用用户上传的 video_url
|
||||
- 不传 GFPGAN 超分,不需要 ffmpeg 音视频混流
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import UTC, datetime
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
from celery import shared_task
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DITTO_URL_TTL_SECONDS = 7 * 24 * 3600 # Ditto 结果 OSS URL 7 天有效
|
||||
|
||||
|
||||
def _get_db_session() -> Session:
|
||||
try:
|
||||
from worker_app.db import SessionLocal # type: ignore
|
||||
except ImportError:
|
||||
from app.db import SessionLocal # type: ignore
|
||||
return SessionLocal()
|
||||
|
||||
|
||||
def _sign_media_url(url: str) -> str:
|
||||
"""对自家 OSS URL 签 7 天预签名。"""
|
||||
if not url:
|
||||
return url
|
||||
try:
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
storage = get_shared_storage_service()
|
||||
public_base = getattr(storage, "public_url", "")
|
||||
if not isinstance(public_base, str) or not public_base:
|
||||
return url
|
||||
own_host = urlparse(public_base).netloc.lower()
|
||||
host = urlparse(url).netloc.lower()
|
||||
if not own_host or host != own_host:
|
||||
return url
|
||||
return storage.get_download_url(url, expires_seconds=_DITTO_URL_TTL_SECONDS)
|
||||
except Exception:
|
||||
return url
|
||||
|
||||
|
||||
def _probe_video_duration(video_bytes: bytes) -> float:
|
||||
"""用 ffprobe 探测视频时长(秒);失败返回 0。"""
|
||||
try:
|
||||
import subprocess
|
||||
import tempfile
|
||||
import os
|
||||
|
||||
with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as tmp:
|
||||
tmp.write(video_bytes)
|
||||
tmp_path = tmp.name
|
||||
try:
|
||||
out = subprocess.check_output(
|
||||
[
|
||||
"ffprobe",
|
||||
"-v",
|
||||
"error",
|
||||
"-show_entries",
|
||||
"format=duration",
|
||||
"-of",
|
||||
"default=noprint_wrappers=1:nokey=1",
|
||||
tmp_path,
|
||||
],
|
||||
stderr=subprocess.STDOUT,
|
||||
timeout=10,
|
||||
)
|
||||
return float(out.decode().strip() or 0)
|
||||
finally:
|
||||
os.unlink(tmp_path)
|
||||
except Exception as exc:
|
||||
logger.warning("[ditto_task] ffprobe 失败: %s", exc)
|
||||
return 0.0
|
||||
|
||||
|
||||
def _refund_lip_sync(db: Session, job: "LipsyncJobModel") -> None:
|
||||
"""Ditto 失败/取消时全额退款(复用 lipsync_service 的退款逻辑)。"""
|
||||
try:
|
||||
from app.services.lipsync_service import LipsyncService
|
||||
|
||||
LipsyncService(db)._refund_lip_sync(job)
|
||||
except Exception:
|
||||
logger.exception("[ditto_task] lip_sync 退款异常 job_id=%s", job.id)
|
||||
|
||||
|
||||
def _settle_lip_sync(db: Session, job: "LipsyncJobModel", duration: float) -> None:
|
||||
"""Ditto 成功后按实际时长结算。"""
|
||||
try:
|
||||
from app.services.lipsync_service import LipsyncService
|
||||
|
||||
LipsyncService(db)._settle_lip_sync(job, duration)
|
||||
except Exception:
|
||||
logger.exception("[ditto_task] lip_sync 结算异常 job_id=%s(不阻塞)", job.id)
|
||||
|
||||
|
||||
def _fallback_to_gpu_then_mediakit(db: Session, job: "LipsyncJobModel") -> None:
|
||||
"""Ditto 失败后:优先回退 GPU MuseTalk,再回退 MediaKit 云端。
|
||||
|
||||
复用 lipsync_service 现有路径逻辑以保证兜底一致性。
|
||||
"""
|
||||
# 先尝试走 GPU MuseTalk(若可用)
|
||||
try:
|
||||
from app.services.gpu_lipsync_service import GpuLipsyncService
|
||||
from app.tasks.lipsync_gpu import lipsync_gpu_process_async
|
||||
|
||||
gpu_svc = GpuLipsyncService(db)
|
||||
if gpu_svc.has_available_worker():
|
||||
logger.info("[ditto_task] 回退 GPU MuseTalk: job_id=%s", job.id)
|
||||
# 复用 lipsync_service._submit_to_gpu_create 逻辑
|
||||
from app.services.lipsync_service import LipsyncService
|
||||
|
||||
svc = LipsyncService(db)
|
||||
storage = _shared_storage()
|
||||
persisted_audio = None
|
||||
try:
|
||||
persisted_audio = svc._persist_external_audio_for_gpu(job=job, storage=storage)
|
||||
except Exception as exc:
|
||||
logger.warning("[ditto_task] GPU 外部音频转存失败: %s", exc)
|
||||
audio_url_for_task = persisted_audio or job.audio_url
|
||||
gpu_task = gpu_svc.create_task(
|
||||
video_url=job.video_url,
|
||||
audio_url=audio_url_for_task,
|
||||
lipsync_job_id=job.id,
|
||||
user_id=job.user_id,
|
||||
)
|
||||
if gpu_task is not None:
|
||||
job.mediakit_task_id = f"gpu:{gpu_task.id}"
|
||||
job.status = "processing"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
lipsync_gpu_process_async.apply_async(args=(job.id, job.user_id, gpu_task.id))
|
||||
return
|
||||
db.rollback()
|
||||
except Exception as exc:
|
||||
logger.warning("[ditto_task] GPU MuseTalk 回退失败,转 MediaKit: %s", exc)
|
||||
try:
|
||||
db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 最后兜底:MediaKit 云端
|
||||
try:
|
||||
from app.services.mediakit_client import get_mediakit_client
|
||||
|
||||
client = get_mediakit_client()
|
||||
video_url = _sign_media_url(job.video_url)
|
||||
signed_audio_url = _sign_media_url(job.audio_url)
|
||||
result = client.submit_lipsync(
|
||||
video_url=video_url,
|
||||
audio_url=signed_audio_url,
|
||||
enable_video_loop=job.enable_video_loop,
|
||||
client_token=job.id,
|
||||
)
|
||||
job.mediakit_task_id = result["task_id"]
|
||||
job.status = "submitted"
|
||||
job.submitted_at = datetime.now(UTC)
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
logger.info("[ditto_task] 已回退 MediaKit: job_id=%s task_id=%s", job.id, result["task_id"])
|
||||
except Exception as exc:
|
||||
job.status = "failed"
|
||||
job.error_message = f"Ditto/GPU/MediaKit 均失败: {exc}"
|
||||
job.error_code = "AllBackendsFailed"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
logger.error("[ditto_task] 所有兜底均失败: job_id=%s err=%s", job.id, exc)
|
||||
|
||||
|
||||
def _shared_storage():
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
return get_shared_storage_service()
|
||||
|
||||
|
||||
@shared_task(
|
||||
name="lipsync_ditto_process_async",
|
||||
bind=True,
|
||||
max_retries=0,
|
||||
acks_late=True,
|
||||
time_limit=600,
|
||||
soft_time_limit=540,
|
||||
)
|
||||
def lipsync_ditto_process_async(self, job_id: str, user_id: str) -> None:
|
||||
"""异步调用 Ditto 生成口型视频。
|
||||
|
||||
Args:
|
||||
job_id: LipsyncJob ID
|
||||
user_id: 用户 ID
|
||||
"""
|
||||
from packages.application.ditto_service import DittoError, get_ditto_client
|
||||
|
||||
db: Session = _get_db_session()
|
||||
job: Optional[LipsyncJobModel] = None
|
||||
try:
|
||||
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
|
||||
|
||||
job = db.query(LipsyncJobModel).filter_by(id=job_id, user_id=user_id).first()
|
||||
if job is None:
|
||||
logger.error("[ditto_task] job 不存在: job_id=%s", job_id)
|
||||
return
|
||||
|
||||
if job.status != "processing":
|
||||
logger.warning(
|
||||
"[ditto_task] job 状态异常(非 processing),跳过: job_id=%s status=%s",
|
||||
job_id,
|
||||
job.status,
|
||||
)
|
||||
return
|
||||
|
||||
audio_url = job.audio_url or ""
|
||||
script = job.script_text or ""
|
||||
if not audio_url:
|
||||
raise DittoError("job.audio_url 为空,无法调用 Ditto", code="InvalidParam")
|
||||
|
||||
logger.info(
|
||||
"[ditto_task] 开始 Ditto 生成: job_id=%s audio=%s script_len=%d",
|
||||
job_id,
|
||||
audio_url[:100],
|
||||
len(script),
|
||||
)
|
||||
client = get_ditto_client()
|
||||
result = client.generate_and_persist(
|
||||
job_id=job_id,
|
||||
user_id=user_id,
|
||||
audio_url=audio_url,
|
||||
script=script,
|
||||
# video_url 不传则用默认模板
|
||||
)
|
||||
|
||||
# Ditto 返回的 MP4 自带音频,直接标记完成
|
||||
job.output_video_url = result.video_url
|
||||
# 探测时长(用于计费)
|
||||
duration = _probe_video_duration(result.video_bytes)
|
||||
if duration <= 0:
|
||||
# 兜底:按音频时长估算(1秒≈1秒)
|
||||
try:
|
||||
from packages.domain.sentence_timings import probe_audio_duration
|
||||
from packages.shared.url_security import safe_download_bytes
|
||||
|
||||
audio_data = safe_download_bytes(
|
||||
audio_url, allowed_mime_types=("audio/mpeg", "audio/wav", "audio/x-wav"), timeout=30
|
||||
)
|
||||
duration = probe_audio_duration(audio_data)
|
||||
except Exception:
|
||||
duration = 0.0
|
||||
job.output_duration = duration
|
||||
job.status = "completed"
|
||||
job.completed_at = datetime.now(UTC)
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
logger.info(
|
||||
"[ditto_task] Ditto 完成: job_id=%s url=%s duration=%.2fs rtf=%.2f frames=%d",
|
||||
job_id,
|
||||
result.video_url[:100],
|
||||
duration,
|
||||
result.rtf,
|
||||
result.frames,
|
||||
)
|
||||
_settle_lip_sync(db, job, duration)
|
||||
|
||||
except DittoError as exc:
|
||||
logger.error("[ditto_task] Ditto 失败,回退: job_id=%s code=%s err=%s", job_id, exc.code, exc)
|
||||
if job is not None:
|
||||
try:
|
||||
db.rollback()
|
||||
job = db.query(type(job)).filter_by(id=job_id).first() if hasattr(job, "id") else job
|
||||
# 回退 GPU/MediaKit
|
||||
_fallback_to_gpu_then_mediakit(db, job)
|
||||
except Exception as fallback_exc:
|
||||
logger.exception("[ditto_task] 回退也失败 job_id=%s err=%s", job_id, fallback_exc)
|
||||
try:
|
||||
if job:
|
||||
job.status = "failed"
|
||||
job.error_message = f"Ditto 失败且回退异常: {exc}; fallback: {fallback_exc}"
|
||||
job.error_code = "FallbackError"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
except Exception:
|
||||
pass
|
||||
except Exception as exc:
|
||||
logger.exception("[ditto_task] 未预期异常: job_id=%s err=%s", job_id, exc)
|
||||
if job is not None:
|
||||
try:
|
||||
db.rollback()
|
||||
job = db.query(type(job)).filter_by(id=job_id).first()
|
||||
_fallback_to_gpu_then_mediakit(db, job)
|
||||
except Exception as fallback_exc:
|
||||
logger.exception("[ditto_task] 回退也失败 job_id=%s err=%s", job_id, fallback_exc)
|
||||
try:
|
||||
if job:
|
||||
job.status = "failed"
|
||||
job.error_message = f"Ditto 异常: {exc}"
|
||||
job.error_code = "DittoAsyncError"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
db.close()
|
||||
@@ -261,7 +261,33 @@ def tts_synthesize_and_submit(
|
||||
"[lipsync_tts] 句子时间戳计算失败(不影响主流程): job_id=%s err=%s", job_id, _st_err, exc_info=True
|
||||
)
|
||||
|
||||
# 3. 签名 URL 并提交到 MediaKit(复用模块内 _sign_media_url,避免对 LipsyncService 的耦合)
|
||||
# 3. 优先走 Ditto(#2076):开关打开且配置完整时,派发 Ditto 异步任务,不再走 MediaKit
|
||||
ditto_dispatched = False
|
||||
try:
|
||||
from packages.config import get_api_settings as _get_settings
|
||||
|
||||
_settings = _get_settings()
|
||||
if _settings.use_ditto_lipsync and _settings.ditto_api_base_url and _settings.ditto_default_video_url:
|
||||
from app.tasks.lipsync_ditto import lipsync_ditto_process_async
|
||||
|
||||
job.status = "processing"
|
||||
job.mediakit_task_id = "ditto:tts-submitted"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
lipsync_ditto_process_async.apply_async(args=(job_id, user_id))
|
||||
logger.info("[lipsync_tts] TTS 完成,已派发 Ditto 任务: job_id=%s", job_id)
|
||||
ditto_dispatched = True
|
||||
except Exception as _ditto_err:
|
||||
logger.warning("[lipsync_tts] Ditto 派发失败,回退 MediaKit: job_id=%s err=%s", job_id, _ditto_err)
|
||||
try:
|
||||
db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if ditto_dispatched:
|
||||
return
|
||||
|
||||
# 4. 签名 URL 并提交到 MediaKit(复用模块内 _sign_media_url,避免对 LipsyncService 的耦合)
|
||||
audio_url = _sign_media_url(job.audio_url)
|
||||
video_url = _sign_media_url(job.video_url)
|
||||
|
||||
|
||||
@@ -53,6 +53,9 @@ celery_app.conf.imports = (
|
||||
# #1998 GPU MuseTalk 异步推理:wait_for_result→签名 URL→回写 lipsync_jobs
|
||||
# 必须在 Worker 侧注册,否则 apply_async 消息无人消费,job 永远卡在 processing
|
||||
"app.tasks.lipsync_gpu",
|
||||
# #2076 Ditto 蚂蚁数字人异步推理:同步 HTTP 调用 Ditto → MP4 流转存 OSS → 回写 lipsync_jobs
|
||||
# 必须在 Worker 侧注册;失败回退 GPU MuseTalk → MediaKit
|
||||
"app.tasks.lipsync_ditto",
|
||||
)
|
||||
|
||||
# Celery Beat 定时任务调度
|
||||
|
||||
@@ -2072,18 +2072,27 @@ def _run_render_pipeline(job_id: str, session, repo, job) -> dict:
|
||||
try:
|
||||
review_result = _step_review(job, copy_result)
|
||||
if not review_result.get("passed", True):
|
||||
_emit_progress(job_id, ViralVideoStage.REVIEW, 67.0, "审核未通过,正在自动重写...")
|
||||
# #2040: Reviewer 已在 _step_review 内完成 1 次自动重写
|
||||
rewritten = review_result.get("rewritten_copy")
|
||||
if isinstance(rewritten, dict) and rewritten:
|
||||
copy_result = rewritten
|
||||
else:
|
||||
# #2218: 审核重写失败不再从意图解析重跑,直接报错让用户重新生成文案
|
||||
logger.error(
|
||||
"[爆款视频][阶段3] 合规审核未通过且自动重写失败 job_id=%s,终止渲染",
|
||||
# #2233: 如果issues为空但passed=False,说明是LLM超时/异常导致的误判,降级放行
|
||||
_issues = review_result.get("issues") or []
|
||||
if not _issues:
|
||||
logger.warning(
|
||||
"[爆款视频][阶段3] 审核未通过但无具体问题(可能LLM超时),降级放行 job_id=%s",
|
||||
job_id,
|
||||
)
|
||||
raise ValueError("文案合规审核未通过,请修改文案后重试或重新生成文案")
|
||||
review_result["passed"] = True
|
||||
else:
|
||||
_emit_progress(job_id, ViralVideoStage.REVIEW, 67.0, "审核未通过,正在自动重写...")
|
||||
# #2040: Reviewer 已在 _step_review 内完成 1 次自动重写
|
||||
rewritten = review_result.get("rewritten_copy")
|
||||
if isinstance(rewritten, dict) and rewritten:
|
||||
copy_result = rewritten
|
||||
else:
|
||||
# #2218: 审核重写失败不再从意图解析重跑,直接报错让用户重新生成文案
|
||||
logger.error(
|
||||
"[爆款视频][阶段3] 合规审核未通过且自动重写失败 job_id=%s,终止渲染",
|
||||
job_id,
|
||||
)
|
||||
raise ValueError("文案合规审核未通过,请修改文案后重试或重新生成文案")
|
||||
job.copy_result = copy_result
|
||||
job.generated_copy_text = copy_result.get("voiceover_script", "") or ""
|
||||
_save_job(repo, job, session)
|
||||
|
||||
@@ -681,8 +681,13 @@ def _assemble_v4(idx: int, fj: dict, ocr_texts: list[str]) -> dict[str, Any]:
|
||||
brand = "无法判断"
|
||||
else:
|
||||
brand = brand_raw
|
||||
# name兜底:store_type为空时用brand
|
||||
name = store_type if store_type != "店铺" else (brand if brand != "无法判断" else store_type)
|
||||
# name兜底:更保守的策略
|
||||
# - store_type有具体值(非"店铺")时直接用store_type
|
||||
# - store_type为默认"店铺"时,用"门店门头"而非brand(避免与brand字段重复)
|
||||
if store_type and store_type != "店铺":
|
||||
name = store_type
|
||||
else:
|
||||
name = "门店门头"
|
||||
category = "门店场景"
|
||||
# appearance: store_layout + furnishings + 陈设色调
|
||||
appearance_parts = []
|
||||
|
||||
@@ -0,0 +1,269 @@
|
||||
"""蚂蚁 Ditto 数字人口型 API 客户端 — #2076.
|
||||
|
||||
封装 Ditto FastAPI(部署在 5060Ti GPU 节点,Tailscale 内网可达):
|
||||
- GET /health 健康检查
|
||||
- POST /generate 生成口型视频(同步返回 MP4 流)
|
||||
|
||||
关键特性:
|
||||
- 入参:video_url(人物模板视频 URL) + audio_url(TTS 音频 URL) + script(文案原文)
|
||||
- 出参:直接返回 video/mp4 字节流(自带音频,无需二次混流)
|
||||
- 429 时指数退避重试(最多 ditto_max_retries 次)
|
||||
- 500/超时视为失败
|
||||
- 输出 MP4 字节流转存到自家 OSS,返回公网 URL
|
||||
|
||||
注意:
|
||||
- 保留 MuseTalk/GPU 路径不变;本服务作为更高优先级的第三条口型路径
|
||||
- 不传 emotion/表情精细控制,使用默认 emo_global=4(中性)+ use_script_emo=true(关键词驱动表情)
|
||||
- Ditto 输出自带音视频,不需要 GFPGAN 超分,不需要 ffmpeg 音视频混流
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import logging
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from packages.config import get_api_settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class DittoError(Exception):
|
||||
"""Ditto API 调用失败."""
|
||||
|
||||
def __init__(self, message: str, code: str = "DittoError", status_code: int = 0):
|
||||
self.code = code
|
||||
self.status_code = status_code
|
||||
super().__init__(message)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DittoResult:
|
||||
"""Ditto 生成结果."""
|
||||
|
||||
video_bytes: bytes
|
||||
video_url: str = "" # 转存 OSS 后填充
|
||||
elapsed_seconds: float = 0.0
|
||||
rtf: float = 0.0 # 实时率(响应头 X-RTF)
|
||||
frames: int = 0 # 帧数(响应头 X-Frames)
|
||||
|
||||
|
||||
class DittoClient:
|
||||
"""蚂蚁 Ditto 数字人口型 API 客户端."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: Optional[str] = None,
|
||||
default_video_url: Optional[str] = None,
|
||||
max_retries: Optional[int] = None,
|
||||
timeout: Optional[int] = None,
|
||||
):
|
||||
s = get_api_settings()
|
||||
self.base_url = (base_url or s.ditto_api_base_url or "").rstrip("/")
|
||||
self.default_video_url = default_video_url or s.ditto_default_video_url or ""
|
||||
self.max_retries = int(max_retries if max_retries is not None else s.ditto_max_retries)
|
||||
self.timeout = int(timeout if timeout is not None else s.ditto_request_timeout)
|
||||
|
||||
@property
|
||||
def is_configured(self) -> bool:
|
||||
"""配置是否完整(base_url + 默认模板视频都有值)."""
|
||||
return bool(self.base_url) and bool(self.default_video_url)
|
||||
|
||||
def health(self) -> bool:
|
||||
"""健康检查;成功返回 True,失败返回 False(不抛异常)."""
|
||||
if not self.base_url:
|
||||
return False
|
||||
url = f"{self.base_url}/health"
|
||||
try:
|
||||
with httpx.Client(timeout=5.0) as client:
|
||||
resp = client.get(url)
|
||||
ok = resp.status_code == 200
|
||||
if ok:
|
||||
logger.info("[ditto] health check OK: %s", url)
|
||||
else:
|
||||
logger.warning("[ditto] health check status=%d: %s", resp.status_code, url)
|
||||
return ok
|
||||
except Exception as exc:
|
||||
logger.warning("[ditto] health check failed: %s", exc)
|
||||
return False
|
||||
|
||||
def generate(
|
||||
self,
|
||||
*,
|
||||
audio_url: str,
|
||||
script: str,
|
||||
video_url: Optional[str] = None,
|
||||
emo_global: int = 4,
|
||||
use_script_emo: bool = True,
|
||||
blend_frames: int = 6,
|
||||
) -> DittoResult:
|
||||
"""调用 Ditto /generate 接口,返回 MP4 字节流结果.
|
||||
|
||||
Raises DittoError on failure.
|
||||
"""
|
||||
if not self.base_url:
|
||||
raise DittoError("DITTO_API_BASE_URL 未配置", code="ConfigMissing")
|
||||
driver_url = video_url or self.default_video_url
|
||||
if not driver_url:
|
||||
raise DittoError("Ditto 人物模板视频 URL 未配置", code="ConfigMissing")
|
||||
if not audio_url:
|
||||
raise DittoError("audio_url 不能为空", code="InvalidParam")
|
||||
if not script:
|
||||
script = " "
|
||||
|
||||
payload = {
|
||||
"video_url": driver_url,
|
||||
"audio_url": audio_url,
|
||||
"script": script,
|
||||
"emo_global": emo_global,
|
||||
"use_script_emo": use_script_emo,
|
||||
"blend_frames": blend_frames,
|
||||
}
|
||||
url = f"{self.base_url}/generate"
|
||||
|
||||
last_exc: Optional[Exception] = None
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
start = time.monotonic()
|
||||
with httpx.Client(timeout=self.timeout, follow_redirects=True) as client:
|
||||
resp = client.post(url, json=payload)
|
||||
elapsed = time.monotonic() - start
|
||||
|
||||
if resp.status_code == 429:
|
||||
wait = min(2**attempt, 30)
|
||||
logger.warning(
|
||||
"[ditto] GPU 繁忙 (429),%ds 后重试 (%d/%d)",
|
||||
wait,
|
||||
attempt + 1,
|
||||
self.max_retries,
|
||||
)
|
||||
if attempt >= self.max_retries:
|
||||
raise DittoError(
|
||||
f"Ditto GPU 繁忙,重试 {self.max_retries} 次仍失败",
|
||||
code="BusyRetriesExhausted",
|
||||
status_code=429,
|
||||
)
|
||||
time.sleep(wait)
|
||||
continue
|
||||
|
||||
if resp.status_code != 200:
|
||||
_text = (resp.text or "")[:300]
|
||||
logger.error(
|
||||
"[ditto] generate 失败 status=%d attempt=%d body=%s",
|
||||
resp.status_code,
|
||||
attempt + 1,
|
||||
_text,
|
||||
)
|
||||
if resp.status_code >= 500 and attempt < self.max_retries:
|
||||
time.sleep(min(2**attempt, 15))
|
||||
continue
|
||||
raise DittoError(
|
||||
f"Ditto 返回 {resp.status_code}: {_text}",
|
||||
code="DittoAPIError",
|
||||
status_code=resp.status_code,
|
||||
)
|
||||
|
||||
video_bytes = resp.content
|
||||
if not video_bytes or len(video_bytes) < 1024:
|
||||
raise DittoError(
|
||||
f"Ditto 返回内容异常(size={len(video_bytes) if video_bytes else 0})",
|
||||
code="EmptyResponse",
|
||||
)
|
||||
try:
|
||||
rtf = float(resp.headers.get("X-RTF", "0") or 0)
|
||||
except ValueError:
|
||||
rtf = 0.0
|
||||
try:
|
||||
frames = int(resp.headers.get("X-Frames", "0") or 0)
|
||||
except ValueError:
|
||||
frames = 0
|
||||
try:
|
||||
x_time = float(resp.headers.get("X-Time", "0") or 0)
|
||||
if x_time > 0:
|
||||
elapsed = x_time
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
logger.info(
|
||||
"[ditto] generate 成功 size=%d rtf=%.2f frames=%d elapsed=%.1fs attempt=%d",
|
||||
len(video_bytes),
|
||||
rtf,
|
||||
frames,
|
||||
elapsed,
|
||||
attempt + 1,
|
||||
)
|
||||
return DittoResult(
|
||||
video_bytes=video_bytes,
|
||||
elapsed_seconds=elapsed,
|
||||
rtf=rtf,
|
||||
frames=frames,
|
||||
)
|
||||
|
||||
except DittoError:
|
||||
raise
|
||||
except httpx.TimeoutException as exc:
|
||||
last_exc = exc
|
||||
logger.warning("[ditto] 请求超时 attempt=%d err=%s", attempt + 1, exc)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(min(2**attempt, 15))
|
||||
continue
|
||||
raise DittoError(
|
||||
f"Ditto 请求超时({self.timeout}s),重试耗尽",
|
||||
code="Timeout",
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
last_exc = exc
|
||||
logger.warning("[ditto] 请求异常 attempt=%d err=%s", attempt + 1, exc)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(min(2**attempt, 10))
|
||||
continue
|
||||
raise DittoError(f"Ditto 调用异常: {exc}", code="NetworkError") from exc
|
||||
|
||||
raise DittoError("Ditto 未知错误", code="Unknown") from last_exc
|
||||
|
||||
def generate_and_persist(
|
||||
self,
|
||||
*,
|
||||
job_id: str,
|
||||
user_id: str,
|
||||
audio_url: str,
|
||||
script: str,
|
||||
video_url: Optional[str] = None,
|
||||
) -> DittoResult:
|
||||
"""调用 generate 并把 MP4 转存到自家 OSS,返回带 video_url 的结果."""
|
||||
result = self.generate(audio_url=audio_url, script=script, video_url=video_url)
|
||||
try:
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
storage = get_shared_storage_service()
|
||||
storage_key = f"ditto-output/{user_id}/{job_id}.mp4"
|
||||
public_url = storage.upload_file(
|
||||
io.BytesIO(result.video_bytes),
|
||||
storage_key,
|
||||
content_type="video/mp4",
|
||||
)
|
||||
result.video_url = public_url
|
||||
logger.info(
|
||||
"[ditto] 转存 OSS 完成 job=%s key=%s",
|
||||
job_id,
|
||||
storage_key,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.error("[ditto] 转存 OSS 失败 job=%s err=%s", job_id, exc, exc_info=True)
|
||||
raise DittoError(f"Ditto 结果转存 OSS 失败: {exc}", code="StorageError") from exc
|
||||
return result
|
||||
|
||||
|
||||
_ditto_client_singleton: Optional[DittoClient] = None
|
||||
|
||||
|
||||
def get_ditto_client() -> DittoClient:
|
||||
"""获取 DittoClient 单例(简易工厂,便于单测 mock)."""
|
||||
global _ditto_client_singleton
|
||||
if _ditto_client_singleton is None:
|
||||
_ditto_client_singleton = DittoClient()
|
||||
return _ditto_client_singleton
|
||||
@@ -70,8 +70,9 @@ class Reviewer:
|
||||
local = self._rule_check(fusion, intent, fusion_level)
|
||||
llm_result = self._llm_review(fusion, intent, fusion_level)
|
||||
if llm_result is None:
|
||||
# LLM审核失败(超时/网络错误等),降级放行,不阻断渲染
|
||||
return ReviewResult(
|
||||
passed=not local,
|
||||
passed=True,
|
||||
issues=local,
|
||||
rewrite_suggestions=[],
|
||||
raw="",
|
||||
@@ -86,6 +87,17 @@ class Reviewer:
|
||||
)
|
||||
|
||||
def _llm_review(self, fusion: FusionResult, intent: IntentResult, fusion_level: str) -> Optional[ReviewResult]:
|
||||
try:
|
||||
return self._llm_review_inner(fusion, intent, fusion_level)
|
||||
except Exception as e:
|
||||
import logging
|
||||
|
||||
logging.getLogger(__name__).warning("[Reviewer] LLM审核调用异常,降级放行: %s", e)
|
||||
return None
|
||||
|
||||
def _llm_review_inner(
|
||||
self, fusion: FusionResult, intent: IntentResult, fusion_level: str
|
||||
) -> Optional[ReviewResult]:
|
||||
template = get_template("review")
|
||||
system = render_system_prompt(template)
|
||||
user = render_user_prompt(
|
||||
|
||||
@@ -173,6 +173,35 @@ class SharedSettings(BaseSettings):
|
||||
# 判断 Worker 可用的心跳新鲜度窗口(秒)—— last_heartbeat_at 在窗口内视为在线
|
||||
gpu_worker_stale_seconds: int = 300
|
||||
|
||||
# ── Ditto 蚂蚁数字人口型 API(#2076)─────────────────────────────────
|
||||
# 是否优先使用 Ditto(蚂蚁数字人,替代 MuseTalk)。开关开启且 base_url 配置
|
||||
# 非空时,对口型任务优先走 Ditto;失败后回退 MuseTalk/MediaKit。
|
||||
use_ditto_lipsync: bool = Field(
|
||||
default=False,
|
||||
validation_alias=AliasChoices("USE_DITTO_LIPSYNC", "use_ditto_lipsync"),
|
||||
)
|
||||
# Ditto FastAPI 内网地址(Tailscale),如 http://100.x.x.x:8000
|
||||
ditto_api_base_url: str = Field(
|
||||
default="",
|
||||
validation_alias=AliasChoices("DITTO_API_BASE_URL", "ditto_api_base_url"),
|
||||
)
|
||||
# 默认人物模板视频 URL(正面 5-10 秒循环、光线均匀、半身)。Ditto 模式下忽略
|
||||
# 用户上传的驱动视频/图片,统一用该模板;后续可扩展为多模板让用户选择。
|
||||
ditto_default_video_url: str = Field(
|
||||
default="",
|
||||
validation_alias=AliasChoices("DITTO_DEFAULT_VIDEO_URL", "ditto_default_video_url"),
|
||||
)
|
||||
# 429 GPU 繁忙时指数退避最大重试次数
|
||||
ditto_max_retries: int = Field(
|
||||
default=3,
|
||||
validation_alias=AliasChoices("DITTO_MAX_RETRIES", "ditto_max_retries"),
|
||||
)
|
||||
# Ditto 单次请求超时(秒):数字人半身视频推理通常 30-120s
|
||||
ditto_request_timeout: int = Field(
|
||||
default=300,
|
||||
validation_alias=AliasChoices("DITTO_REQUEST_TIMEOUT", "ditto_request_timeout"),
|
||||
)
|
||||
|
||||
# ── P4000 NVENC 硬件编码 ────────────────────────────────────────────
|
||||
# GPU 编码总开关;关闭或 endpoint 为空时始终走本机 CPU libx264
|
||||
enable_gpu_encode: bool = Field(
|
||||
|
||||
@@ -337,13 +337,13 @@ class DoubaoClient:
|
||||
self.last_finish_reason = finish_reason
|
||||
_elapsed = time.time() - _t0
|
||||
logger.info(
|
||||
"[doubao] chat_completion 完成 model=%s tokens_in=%d tokens_out=%d elapsed=%.1fs attempt=%d timeout=%d",
|
||||
"[doubao] chat_completion 完成 model=%s tokens_in=%d tokens_out=%d elapsed=%.1fs attempt=%d timeout=%s",
|
||||
payload.get("model"),
|
||||
data.get("usage", {}).get("prompt_tokens", 0),
|
||||
data.get("usage", {}).get("completion_tokens", 0),
|
||||
_elapsed,
|
||||
attempt + 1,
|
||||
_req_timeout,
|
||||
getattr(_req_timeout, "read", _req_timeout),
|
||||
)
|
||||
return content.strip()
|
||||
except Exception as e:
|
||||
|
||||
@@ -51,6 +51,8 @@ task_routes = {
|
||||
"ai_avatar_render.execute": {"queue": QUEUE_GENERATION},
|
||||
# GPU MuseTalk 口型同步(用户等成片,链路子任务全部走 generation 避免跨队列阻塞)
|
||||
"lipsync_gpu_process_async": {"queue": QUEUE_GENERATION},
|
||||
# #2076 Ditto 蚂蚁数字人口型同步(走 generation 队列,避免跨队列阻塞)
|
||||
"lipsync_ditto_process_async": {"queue": QUEUE_GENERATION},
|
||||
"lipsync_tts.synthesize_and_submit": {"queue": QUEUE_GENERATION},
|
||||
"lipsync_tts.poll_mediakit_status": {"queue": QUEUE_GENERATION},
|
||||
"lipsync_tts.persist_output_video": {"queue": QUEUE_GENERATION},
|
||||
|
||||
@@ -0,0 +1,197 @@
|
||||
"""Ditto 蚂蚁数字人客户端单元测试 — #2076."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from packages.application.ditto_service import DittoClient, DittoError, DittoResult
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(self, status_code=200, content=b"\x00\x01" * 1000, headers=None, text=""):
|
||||
self.status_code = status_code
|
||||
self.content = content
|
||||
self.headers = headers or {}
|
||||
self.text = text
|
||||
|
||||
|
||||
def _make_client(base_url="http://ditto:8000", default_video_url="http://oss/tpl.mp4", max_retries=2, timeout=60):
|
||||
with patch("packages.application.ditto_service.get_api_settings") as mock_settings:
|
||||
s = MagicMock()
|
||||
s.ditto_api_base_url = base_url
|
||||
s.ditto_default_video_url = default_video_url
|
||||
s.ditto_max_retries = max_retries
|
||||
s.ditto_request_timeout = timeout
|
||||
mock_settings.return_value = s
|
||||
return DittoClient()
|
||||
|
||||
|
||||
def test_is_configured_true():
|
||||
c = _make_client()
|
||||
assert c.is_configured is True
|
||||
|
||||
|
||||
def test_is_configured_false_without_base():
|
||||
c = _make_client(base_url="")
|
||||
assert c.is_configured is False
|
||||
|
||||
|
||||
def test_is_configured_false_without_template():
|
||||
c = _make_client(default_video_url="")
|
||||
assert c.is_configured is False
|
||||
|
||||
|
||||
def test_health_ok():
|
||||
c = _make_client()
|
||||
with patch("httpx.Client") as mock_cls:
|
||||
client = MagicMock()
|
||||
client.get.return_value = _FakeResponse(200)
|
||||
mock_cls.return_value.__enter__.return_value = client
|
||||
assert c.health() is True
|
||||
client.get.assert_called_once()
|
||||
|
||||
|
||||
def test_health_fail_status():
|
||||
c = _make_client()
|
||||
with patch("httpx.Client") as mock_cls:
|
||||
client = MagicMock()
|
||||
client.get.return_value = _FakeResponse(500)
|
||||
mock_cls.return_value.__enter__.return_value = client
|
||||
assert c.health() is False
|
||||
|
||||
|
||||
def test_health_network_error():
|
||||
c = _make_client()
|
||||
with patch("httpx.Client") as mock_cls:
|
||||
client = MagicMock()
|
||||
client.get.side_effect = httpx.ConnectError("fail")
|
||||
mock_cls.return_value.__enter__.return_value = client
|
||||
assert c.health() is False
|
||||
|
||||
|
||||
def test_generate_missing_base():
|
||||
c = _make_client(base_url="")
|
||||
with pytest.raises(DittoError, match="DITTO_API_BASE_URL"):
|
||||
c.generate(audio_url="http://x/a.mp3", script="你好")
|
||||
|
||||
|
||||
def test_generate_missing_audio():
|
||||
c = _make_client()
|
||||
with pytest.raises(DittoError, match="audio_url"):
|
||||
c.generate(audio_url="", script="你好")
|
||||
|
||||
|
||||
def test_generate_success_with_headers():
|
||||
c = _make_client(max_retries=0)
|
||||
fake_resp = _FakeResponse(
|
||||
status_code=200,
|
||||
content=b"\x00" * 99999,
|
||||
headers={"X-RTF": "0.35", "X-Frames": "125", "X-Time": "12.5"},
|
||||
)
|
||||
with patch("httpx.Client") as mock_cls, patch("time.monotonic", side_effect=[0, 1]):
|
||||
client = MagicMock()
|
||||
client.post.return_value = fake_resp
|
||||
mock_cls.return_value.__enter__.return_value = client
|
||||
result = c.generate(audio_url="http://x/a.mp3", script="你好")
|
||||
assert isinstance(result, DittoResult)
|
||||
assert len(result.video_bytes) == 99999
|
||||
assert result.rtf == 0.35
|
||||
assert result.frames == 125
|
||||
assert result.elapsed_seconds == 12.5
|
||||
|
||||
|
||||
def test_generate_uses_default_template_when_video_url_empty():
|
||||
c = _make_client(max_retries=0)
|
||||
fake_resp = _FakeResponse(200, b"1" * 99999)
|
||||
with patch("httpx.Client") as mock_cls:
|
||||
client = MagicMock()
|
||||
client.post.return_value = fake_resp
|
||||
mock_cls.return_value.__enter__.return_value = client
|
||||
c.generate(audio_url="http://x/a.mp3", script="你好")
|
||||
call_kwargs = client.post.call_args
|
||||
payload = call_kwargs.kwargs.get("json") or call_kwargs[1].get("json")
|
||||
assert payload["video_url"] == "http://oss/tpl.mp4"
|
||||
assert payload["audio_url"] == "http://x/a.mp3"
|
||||
assert payload["script"] == "你好"
|
||||
assert payload["emo_global"] == 4
|
||||
assert payload["use_script_emo"] is True
|
||||
|
||||
|
||||
def test_generate_retries_on_429_then_success():
|
||||
c = _make_client(max_retries=2)
|
||||
busy = _FakeResponse(429, b"", text="busy")
|
||||
ok = _FakeResponse(200, b"v" * 99999)
|
||||
with patch("httpx.Client") as mock_cls, patch("time.sleep") as mock_sleep:
|
||||
client = MagicMock()
|
||||
client.post.side_effect = [busy, ok]
|
||||
mock_cls.return_value.__enter__.return_value = client
|
||||
result = c.generate(audio_url="http://x/a.mp3", script="你好")
|
||||
assert len(result.video_bytes) == 99999
|
||||
assert mock_sleep.called
|
||||
assert client.post.call_count == 2
|
||||
|
||||
|
||||
def test_generate_429_exhausted():
|
||||
c = _make_client(max_retries=1)
|
||||
with patch("httpx.Client") as mock_cls, patch("time.sleep"):
|
||||
client = MagicMock()
|
||||
client.post.return_value = _FakeResponse(429, b"", text="busy")
|
||||
mock_cls.return_value.__enter__.return_value = client
|
||||
with pytest.raises(DittoError, match="重试"):
|
||||
c.generate(audio_url="http://x/a.mp3", script="你好")
|
||||
|
||||
|
||||
def test_generate_400_no_retry():
|
||||
c = _make_client(max_retries=2)
|
||||
with patch("httpx.Client") as mock_cls:
|
||||
client = MagicMock()
|
||||
client.post.return_value = _FakeResponse(400, b"", text="bad request")
|
||||
mock_cls.return_value.__enter__.return_value = client
|
||||
with pytest.raises(DittoError, match="Ditto 返回 400"):
|
||||
c.generate(audio_url="http://x/a.mp3", script="你好")
|
||||
assert client.post.call_count == 1 # 400 不重试
|
||||
|
||||
|
||||
def test_generate_small_response_raises():
|
||||
c = _make_client(max_retries=0)
|
||||
with patch("httpx.Client") as mock_cls:
|
||||
client = MagicMock()
|
||||
client.post.return_value = _FakeResponse(200, b"xx")
|
||||
mock_cls.return_value.__enter__.return_value = client
|
||||
with pytest.raises(DittoError) as exc_info:
|
||||
c.generate(audio_url="http://x/a.mp3", script="你好")
|
||||
assert exc_info.value.code == "EmptyResponse"
|
||||
|
||||
|
||||
def test_generate_and_persist_uploads_to_storage():
|
||||
c = _make_client(max_retries=0)
|
||||
fake_resp = _FakeResponse(200, b"v" * 99999)
|
||||
fake_storage = MagicMock()
|
||||
fake_storage.upload_file.return_value = "http://oss/ditto/x.mp4"
|
||||
with (
|
||||
patch("httpx.Client") as mock_cls,
|
||||
patch("packages.shared.storage.get_shared_storage_service", return_value=fake_storage),
|
||||
):
|
||||
client = MagicMock()
|
||||
client.post.return_value = fake_resp
|
||||
mock_cls.return_value.__enter__.return_value = client
|
||||
result = c.generate_and_persist(job_id="j1", user_id="u1", audio_url="http://x/a.mp3", script="hi")
|
||||
assert result.video_url == "http://oss/ditto/x.mp4"
|
||||
fake_storage.upload_file.assert_called_once()
|
||||
call_args = fake_storage.upload_file.call_args
|
||||
assert call_args.args[1].startswith("ditto-output/u1/j1")
|
||||
|
||||
|
||||
def test_empty_script_replaced_with_space():
|
||||
c = _make_client(max_retries=0)
|
||||
fake_resp = _FakeResponse(200, b"v" * 99999)
|
||||
with patch("httpx.Client") as mock_cls:
|
||||
client = MagicMock()
|
||||
client.post.return_value = fake_resp
|
||||
mock_cls.return_value.__enter__.return_value = client
|
||||
c.generate(audio_url="http://x/a.mp3", script="")
|
||||
payload = client.post.call_args.kwargs["json"]
|
||||
assert payload["script"] == " "
|
||||
Reference in New Issue
Block a user