Compare commits
46 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 8d9583b089 | |||
| 783914764f | |||
| c32e6292f7 | |||
| 247af9b755 | |||
| cfb5c7ab89 | |||
| d6d100baeb | |||
| 4d99c20f6b | |||
| e9c119aea4 | |||
| 58ac709e53 | |||
| afdee720ff | |||
| d8f3f9f3c9 | |||
| 37f62e5ad0 | |||
| 31453ae3f7 | |||
| 707aaab28f | |||
| b47d5d9396 | |||
| 0638d0f1ce | |||
| 406e9100c8 | |||
| 25067b4d22 | |||
| 90f6d098ee | |||
| 8612e183ad | |||
| c4b4a1e9c0 | |||
| 1f6d9fc862 | |||
| 3acc309b51 | |||
| bee59b7f27 | |||
| 4bd4a6c390 | |||
| aa898ff0c8 | |||
| fb0d429bd7 | |||
| 0510d101aa | |||
| 606f6988b5 | |||
| 733d8bb75c | |||
| e83048b7ec | |||
| ade6593d66 | |||
| a3a8d2561d | |||
| 7b6dfdc29f | |||
| b5e225e62a | |||
| 2141ddb19b | |||
| 3503542ec2 | |||
| 1a6f04b258 | |||
| ece5f7d8a7 | |||
| b023988402 | |||
| 4ba5b126b6 | |||
| cc8a3caba5 | |||
| fe36a05c51 | |||
| 12a8701283 | |||
| 176a5cbfe3 | |||
| c5e69db50d |
@@ -437,7 +437,7 @@ jobs:
|
||||
if: "always() && needs.dedupe-check.outputs.skip_tests != 'true' && needs.check-frontend-only.outputs.skip_backend != 'true'"
|
||||
name: Unit Tests
|
||||
runs-on: ci-l2
|
||||
timeout-minutes: 8
|
||||
timeout-minutes: 20
|
||||
env:
|
||||
PIP_CACHE_DIR: /root/.cache/pip
|
||||
PIP_NO_CACHE_DIR: ''
|
||||
|
||||
@@ -0,0 +1,43 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""system_settings 正式建表(#2246)
|
||||
|
||||
Revision ID: 106_system_settings
|
||||
Revises: 105_narration_first
|
||||
Create Date: 2026-10-08
|
||||
|
||||
system_settings 表此前在 staging 手工创建(对应 034 占位迁移),
|
||||
此处补正式迁移保证其他环境一致。CREATE TABLE/INDEX 使用 IF NOT EXISTS,
|
||||
对已手工建表的环境幂等。
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "106_system_settings"
|
||||
down_revision = "105_narration_first"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.execute("""
|
||||
CREATE TABLE IF NOT EXISTS system_settings (
|
||||
id VARCHAR(36) NOT NULL,
|
||||
setting_key VARCHAR(100) NOT NULL,
|
||||
setting_value TEXT,
|
||||
setting_type VARCHAR(20) NOT NULL,
|
||||
description VARCHAR(255) NOT NULL DEFAULT '',
|
||||
is_public BOOLEAN NOT NULL DEFAULT FALSE,
|
||||
updated_by VARCHAR(36),
|
||||
category VARCHAR(50) NOT NULL DEFAULT 'general',
|
||||
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW(),
|
||||
CONSTRAINT pk_system_settings PRIMARY KEY (id),
|
||||
CONSTRAINT uq_system_settings_setting_key UNIQUE (setting_key)
|
||||
)
|
||||
""")
|
||||
op.execute("CREATE INDEX IF NOT EXISTS ix_system_settings_category " "ON system_settings (category)")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.execute("DROP INDEX IF EXISTS ix_system_settings_category")
|
||||
op.execute("DROP TABLE IF EXISTS system_settings")
|
||||
@@ -0,0 +1,23 @@
|
||||
"""add language column to viral_video_jobs
|
||||
|
||||
Revision ID: 107
|
||||
Revises: 106_system_settings
|
||||
Create Date: 2026-10-09
|
||||
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "107"
|
||||
down_revision = "106_system_settings"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.execute("ALTER TABLE viral_video_jobs " "ADD COLUMN IF NOT EXISTS language VARCHAR(20) NOT NULL DEFAULT 'zh-CN'")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.execute("ALTER TABLE viral_video_jobs DROP COLUMN IF EXISTS language")
|
||||
@@ -1,3 +1,4 @@
|
||||
from app.api.routes.admin.ditto_emotion import router as admin_ditto_emotion_router
|
||||
from app.api.routes.ai import router as ai_router
|
||||
from app.api.routes.ai_avatar_render import router as ai_avatar_render_router
|
||||
from app.api.routes.asset_diagnosis import router as asset_diagnosis_router
|
||||
@@ -242,3 +243,6 @@ api_router.include_router(
|
||||
tags=["GPU Worker"],
|
||||
)
|
||||
api_router.include_router(viral_video_router, prefix="/viral-video", tags=["爆款视频"])
|
||||
|
||||
# #2246:后台 Ditto 表情配置(router 自带 /admin/ditto-emotion 前缀)
|
||||
api_router.include_router(admin_ditto_emotion_router)
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""后台管理路由(#2246 起)."""
|
||||
@@ -0,0 +1,163 @@
|
||||
"""Ditto 数字人表情后台配置 — #2246.
|
||||
|
||||
路由前缀 /api/v1/admin/ditto-emotion,全部使用 _verify_internal_api_key 鉴权
|
||||
(X-API-Key header)。仅开放 5 项白名单配置:
|
||||
- ditto_emotion_enabled / ditto_emotion_model / ditto_emotion_temperature
|
||||
- ditto_emotion_prompt / ditto_blend_frames
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
from pydantic import BaseModel
|
||||
|
||||
from packages.application.system_config_service import get_system_config_service
|
||||
from packages.config import get_api_settings
|
||||
from packages.domain.system_setting import (
|
||||
SETTING_TYPE_BOOL,
|
||||
SETTING_TYPE_FLOAT,
|
||||
SETTING_TYPE_INT,
|
||||
SETTING_TYPE_STRING,
|
||||
)
|
||||
|
||||
from ..auth import _verify_internal_api_key
|
||||
|
||||
router = APIRouter(
|
||||
prefix="/admin/ditto-emotion",
|
||||
tags=["Admin"],
|
||||
dependencies=[Depends(_verify_internal_api_key)],
|
||||
)
|
||||
|
||||
MODEL_OPTIONS = [
|
||||
"doubao-seed-2-1-lite-250915",
|
||||
"doubao-seed-2-1-pro-250915",
|
||||
"deepseek-v3",
|
||||
]
|
||||
|
||||
# key → (类型, 分类)
|
||||
_WHITELIST: dict[str, str] = {
|
||||
"ditto_emotion_enabled": SETTING_TYPE_BOOL,
|
||||
"ditto_emotion_model": SETTING_TYPE_STRING,
|
||||
"ditto_emotion_temperature": SETTING_TYPE_FLOAT,
|
||||
"ditto_emotion_prompt": SETTING_TYPE_STRING,
|
||||
"ditto_blend_frames": SETTING_TYPE_INT,
|
||||
}
|
||||
|
||||
_DESCRIPTIONS: dict[str, str] = {
|
||||
"ditto_emotion_enabled": "LLM 情绪分析开关。关闭时回退到原有关键词匹配模式,不影响正常出片。",
|
||||
"ditto_emotion_model": "用于分析文案情绪的大模型。",
|
||||
"ditto_emotion_temperature": "模型温度,0-1,越低越稳定保守。",
|
||||
"ditto_emotion_prompt": "情绪分析提示词,核心调优入口,必须包含 {文案} 占位符。",
|
||||
"ditto_blend_frames": "表情切换过渡帧数(6-30),越大越柔和。",
|
||||
}
|
||||
|
||||
|
||||
def _settings():
|
||||
return get_api_settings()
|
||||
|
||||
|
||||
def _default_value(key: str) -> Any:
|
||||
return getattr(_settings(), key)
|
||||
|
||||
|
||||
def _build_config_item(key: str) -> dict[str, Any]:
|
||||
item: dict[str, Any] = {
|
||||
"key": key,
|
||||
"type": _WHITELIST[key],
|
||||
"description": _DESCRIPTIONS.get(key, ""),
|
||||
"default": _default_value(key),
|
||||
}
|
||||
service = get_system_config_service()
|
||||
item["value"] = service.get_config(key, _default_value(key))
|
||||
if key == "ditto_emotion_model":
|
||||
item["model_options"] = list(MODEL_OPTIONS)
|
||||
return item
|
||||
|
||||
|
||||
class ConfigUpdatePayload(BaseModel):
|
||||
configs: dict[str, Any]
|
||||
|
||||
|
||||
class TestPayload(BaseModel):
|
||||
test_text: str
|
||||
|
||||
|
||||
def _validate_value(key: str, value: Any) -> Any:
|
||||
st = _WHITELIST[key]
|
||||
if st == SETTING_TYPE_BOOL:
|
||||
if not isinstance(value, bool):
|
||||
raise ValueError(f"{key} 必须是布尔值")
|
||||
elif st == SETTING_TYPE_INT:
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise ValueError(f"{key} 必须是整数")
|
||||
if not 6 <= value <= 30:
|
||||
raise ValueError(f"{key} 必须在 6-30 之间")
|
||||
elif st == SETTING_TYPE_FLOAT:
|
||||
if isinstance(value, bool):
|
||||
raise ValueError(f"{key} 必须是数字")
|
||||
try:
|
||||
value = float(value)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError(f"{key} 必须是数字") from exc
|
||||
if not 0.0 <= value <= 1.0:
|
||||
raise ValueError(f"{key} 必须在 0-1 之间")
|
||||
elif st == SETTING_TYPE_STRING:
|
||||
if not isinstance(value, str):
|
||||
raise ValueError(f"{key} 必须是字符串")
|
||||
if key == "ditto_emotion_prompt" and value.strip() and "{文案}" not in value:
|
||||
raise ValueError("提示词必须包含 {文案} 占位符")
|
||||
if key == "ditto_emotion_model" and value not in MODEL_OPTIONS:
|
||||
raise ValueError(f"模型必须是以下之一:{', '.join(MODEL_OPTIONS)}")
|
||||
return value
|
||||
|
||||
|
||||
@router.get("/config")
|
||||
def get_config() -> dict[str, Any]:
|
||||
return {"configs": [_build_config_item(k) for k in _WHITELIST]}
|
||||
|
||||
|
||||
@router.put("/config")
|
||||
def update_config(
|
||||
payload: ConfigUpdatePayload,
|
||||
x_api_key: str = Depends(_verify_internal_api_key),
|
||||
) -> dict[str, Any]:
|
||||
configs = payload.configs
|
||||
illegal = [k for k in configs if k not in _WHITELIST]
|
||||
if illegal:
|
||||
return {
|
||||
"ok": False,
|
||||
"error": f"不允许修改的配置项:{', '.join(illegal)}",
|
||||
}
|
||||
service = get_system_config_service()
|
||||
updated: dict[str, Any] = {}
|
||||
for key, raw in configs.items():
|
||||
try:
|
||||
value = _validate_value(key, raw)
|
||||
except ValueError as exc:
|
||||
return {"ok": False, "error": str(exc)}
|
||||
service.set_config(
|
||||
key,
|
||||
value,
|
||||
setting_type=_WHITELIST[key],
|
||||
updated_by=x_api_key[:8] if x_api_key else None,
|
||||
)
|
||||
updated[key] = value
|
||||
return {"ok": True, "updated": updated}
|
||||
|
||||
|
||||
@router.post("/config/test")
|
||||
def test_config(payload: TestPayload) -> dict[str, Any]:
|
||||
text = (payload.test_text or "").strip()
|
||||
if not text:
|
||||
return {"ok": False, "error": "test_text 不能为空"}
|
||||
from packages.application.ditto_emotion_service import get_ditto_emotion_service
|
||||
|
||||
service = get_ditto_emotion_service()
|
||||
segments = service.analyze(text)
|
||||
return {
|
||||
"ok": True,
|
||||
"enabled": service.enabled,
|
||||
"segments": [s.to_dict() for s in segments],
|
||||
}
|
||||
@@ -89,6 +89,7 @@ def get_balance(
|
||||
is_member=_is_member(current_user),
|
||||
member_type=_member_type(current_user),
|
||||
member_expires_at=getattr(current_user.user, "member_expires_at", None),
|
||||
credits_enabled=_credits_enabled(),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -142,6 +142,7 @@ def _to_response(job) -> ViralVideoJobResponse:
|
||||
copy_result=_build_copy_result(job),
|
||||
voice_id=getattr(job, "voice_id", "") or "",
|
||||
voice_source=getattr(job, "voice_source", "") or "",
|
||||
language=getattr(job, "language", "zh-CN") or "zh-CN",
|
||||
video_ratio=getattr(job, "video_ratio", "9:16") or "9:16",
|
||||
video_model=getattr(job, "video_model", "") or "",
|
||||
intent_result=job.intent_result,
|
||||
@@ -243,6 +244,7 @@ def analyze_images(
|
||||
style_strength=request.style_strength or "medium",
|
||||
voice_id=request.voice_id or "",
|
||||
voice_source=request.voice_source or "",
|
||||
language=getattr(request, "language", "zh-CN") or "zh-CN",
|
||||
video_ratio=request.video_ratio or "9:16",
|
||||
video_model=request.video_model or "",
|
||||
video_resolution=getattr(request, "video_resolution", "720p") or "720p",
|
||||
@@ -314,6 +316,7 @@ def generate_copy(
|
||||
job.style_guide = request.style_guide
|
||||
job.voice_id = request.voice_id or job.voice_id
|
||||
job.voice_source = request.voice_source or job.voice_source
|
||||
job.language = getattr(request, "language", "") or job.language or "zh-CN"
|
||||
job.video_ratio = request.video_ratio or job.video_ratio or "9:16"
|
||||
job.video_model = request.video_model or job.video_model or ""
|
||||
job.video_resolution = getattr(request, "video_resolution", "") or job.video_resolution or "720p"
|
||||
@@ -352,6 +355,34 @@ def confirm_copy(
|
||||
if not isinstance(job.copy_result, dict) or not job.copy_result:
|
||||
raise HTTPException(status_code=409, detail="文案数据缺失,请先点击「生成文案」")
|
||||
|
||||
# Bug1 fix: 用户 confirm 时允许修改 video_model/video_resolution/video_ratio/duration
|
||||
old_duration = int(getattr(job, "duration", 15) or 15)
|
||||
old_resolution = getattr(job, "video_resolution", "720p") or "720p"
|
||||
old_ratio = getattr(job, "video_ratio", "9:16") or "9:16"
|
||||
old_model = getattr(job, "video_model", None) or "seedance-2.5"
|
||||
|
||||
if request.duration is not None:
|
||||
job.duration = max(5, min(30, int(request.duration)))
|
||||
if request.video_resolution is not None:
|
||||
job.video_resolution = request.video_resolution
|
||||
if request.video_ratio is not None:
|
||||
job.video_ratio = request.video_ratio
|
||||
if request.video_model is not None:
|
||||
job.video_model = request.video_model
|
||||
if request.voice_id is not None:
|
||||
job.voice_id = request.voice_id
|
||||
if request.voice_source is not None:
|
||||
job.voice_source = request.voice_source
|
||||
if getattr(request, "language", None) is not None:
|
||||
job.language = request.language
|
||||
|
||||
param_changed = (
|
||||
(request.duration is not None and int(request.duration) != old_duration)
|
||||
or (request.video_resolution is not None and request.video_resolution != old_resolution)
|
||||
or (request.video_ratio is not None and request.video_ratio != old_ratio)
|
||||
or (request.video_model is not None and request.video_model != old_model)
|
||||
)
|
||||
|
||||
# 积分预扣(已扣过/重试任务跳过)
|
||||
from app.config import settings as _settings
|
||||
|
||||
@@ -359,7 +390,51 @@ def confirm_copy(
|
||||
already_paid = (float(getattr(job, "credits_prepaid", 0) or 0) > 0) or (
|
||||
float(getattr(job, "credits_cost", 0) or 0) > 0
|
||||
)
|
||||
if not already_paid:
|
||||
if param_changed and already_paid:
|
||||
# 参数变更:回退旧预扣,按新参数重新预扣
|
||||
from packages.domain.points_rules import calculate_viral_video_credits, resolve_video_dimensions
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
old_w, old_h = resolve_video_dimensions(old_resolution, old_ratio)
|
||||
old_est = calculate_viral_video_credits(old_duration, old_w, old_h, old_model)
|
||||
new_w, new_h = resolve_video_dimensions(
|
||||
getattr(job, "video_resolution", "720p") or "720p",
|
||||
job.video_ratio or "9:16",
|
||||
)
|
||||
new_est = calculate_viral_video_credits(
|
||||
int(job.duration or 15), new_w, new_h, job.video_model or "seedance-2.5"
|
||||
)
|
||||
svc = PointsService()
|
||||
# 退回旧预扣
|
||||
if getattr(job, "credits_transaction_id", None):
|
||||
svc.refund_points(
|
||||
user_id=authenticated_user.user.id,
|
||||
amount=float(job.credits_prepaid),
|
||||
source="viral_video",
|
||||
db=session,
|
||||
ref_id=job.credits_transaction_id,
|
||||
description="confirm-copy 参数变更退还旧预扣",
|
||||
)
|
||||
# 预扣新金额
|
||||
if new_est > 0:
|
||||
res = svc.deduct_viral_video(authenticated_user.user.id, new_est, job.id, session)
|
||||
if not res.get("success"):
|
||||
balance = res.get("balance", 0)
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {new_est} 积分,当前余额 {balance}",
|
||||
"required": new_est,
|
||||
"balance": balance,
|
||||
},
|
||||
)
|
||||
job.credits_prepaid = new_est
|
||||
job.credits_transaction_id = res.get("transaction_id", "") or ""
|
||||
logger.info(
|
||||
"[爆款视频][confirm-copy] 参数变更,积分重算: old=%d new=%d job_id=%s", old_est, new_est, job.id
|
||||
)
|
||||
elif not already_paid:
|
||||
from packages.domain.points_rules import calculate_viral_video_credits, resolve_video_dimensions
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
|
||||
@@ -19,6 +19,7 @@ class PointsBalanceResponse(BaseModel):
|
||||
is_member: bool = Field(default=False, description="是否付费会员")
|
||||
member_type: Optional[str] = Field(None, description="会员类型: monthly/quarterly/yearly")
|
||||
member_expires_at: Optional[datetime] = Field(None, description="会员到期时间")
|
||||
credits_enabled: bool = Field(default=False, description="积分系统是否启用(false=免费放行不扣分)")
|
||||
|
||||
|
||||
# ============ 流水 ============
|
||||
|
||||
@@ -87,6 +87,7 @@ class CreateViralVideoRequest(BaseModel):
|
||||
style_template_id: str = ""
|
||||
voice_id: str = ""
|
||||
voice_source: str = ""
|
||||
language: str = "zh-CN"
|
||||
video_ratio: str = "9:16"
|
||||
video_model: str = ""
|
||||
video_resolution: str = "720p"
|
||||
@@ -121,6 +122,7 @@ class AnalyzeImagesRequest(BaseModel):
|
||||
video_model: str = ""
|
||||
video_resolution: str = "720p"
|
||||
duration: int = Field(default=15, ge=5, le=30)
|
||||
language: str = "zh-CN"
|
||||
|
||||
|
||||
class GenerateCopyRequest(BaseModel):
|
||||
@@ -142,6 +144,7 @@ class GenerateCopyRequest(BaseModel):
|
||||
style_guide: dict | None = None
|
||||
voice_id: str = ""
|
||||
voice_source: str = ""
|
||||
language: str = "zh-CN"
|
||||
video_ratio: str = "9:16"
|
||||
video_model: str = ""
|
||||
video_resolution: str = "720p"
|
||||
@@ -167,6 +170,13 @@ class ConfirmCopyRequest(BaseModel):
|
||||
"""v1.5+ 阶段3:用户确认/编辑口播后开始渲染(TTS+单次Seedance)。"""
|
||||
|
||||
edited_copy: str = Field(default="", description="用户编辑后的口播文案;为空则用 AI 生成的 voiceover_script")
|
||||
video_model: str | None = Field(default=None, description="用户选定的视频生成模型(confirm时可选)")
|
||||
video_resolution: str | None = Field(default=None, description="用户选定的分辨率(confirm时可选)")
|
||||
video_ratio: str | None = Field(default=None, description="用户选定的比例(confirm时可选)")
|
||||
duration: int | None = Field(default=None, ge=5, le=30, description="用户选定的时长秒数(confirm时可选,5~30)")
|
||||
voice_id: str | None = Field(default=None, description="用户选定的音色ID(confirm时可选)")
|
||||
voice_source: str | None = Field(default=None, description="用户选定的音色来源(confirm时可选)")
|
||||
language: str | None = Field(default=None, description="用户选定的语言(confirm时可选)")
|
||||
|
||||
|
||||
class ConfirmIntentRequest(BaseModel):
|
||||
@@ -218,6 +228,7 @@ class ViralVideoJobResponse(BaseModel):
|
||||
# 音色/视频参数
|
||||
voice_id: str = ""
|
||||
voice_source: str = ""
|
||||
language: str = "zh-CN"
|
||||
video_ratio: str = "9:16"
|
||||
video_model: str = ""
|
||||
intent_result: dict | None = None
|
||||
|
||||
@@ -845,6 +845,30 @@ class LipsyncService:
|
||||
if job.status in (STATUS_COMPLETED, "failed"):
|
||||
return job
|
||||
|
||||
# Ditto 异步路径:mediakit_task_id 以 "ditto:" 开头,由 Celery 任务异步更新
|
||||
# 不做 MediaKit 轮询,只检查是否卡住太久(>10 分钟)则标失败
|
||||
if job.mediakit_task_id and job.mediakit_task_id.startswith("ditto:"):
|
||||
if job.status in ("processing", "submitted"):
|
||||
_now = datetime.now(UTC)
|
||||
_upd = job.updated_at
|
||||
if _upd is not None and _upd.tzinfo is None:
|
||||
_upd = _upd.replace(tzinfo=UTC)
|
||||
stale_minutes = 10
|
||||
if _upd and (_now - _upd).total_seconds() > stale_minutes * 60:
|
||||
logger.warning(
|
||||
"Ditto 异步任务超时(>%d 分钟),标记失败: job_id=%s",
|
||||
stale_minutes,
|
||||
job_id,
|
||||
)
|
||||
job.status = "failed"
|
||||
job.error_message = f"Ditto 处理超时(>{stale_minutes} 分钟)"
|
||||
job.error_code = "DittoTimeout"
|
||||
job.completed_at = _now
|
||||
job.updated_at = _now
|
||||
self.db.commit()
|
||||
self._refund_lip_sync(job)
|
||||
return job
|
||||
|
||||
# GPU 异步路径:mediakit_task_id 以 "gpu:" 开头,由 Celery 任务异步更新
|
||||
# 不做 MediaKit 轮询,只检查是否卡住太久(>30 分钟)则标失败
|
||||
if job.mediakit_task_id and job.mediakit_task_id.startswith("gpu:"):
|
||||
|
||||
@@ -2,19 +2,25 @@
|
||||
|
||||
把 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
|
||||
2. 确定驱动视频:用户上传的 video_url 优先,无则用 settings.ditto_default_video_url 兜底
|
||||
3. 视频时长对齐:若视频 < 音频+2s,用 ffmpeg 循环视频到足够长度后上传临时文件
|
||||
4. 调 DittoClient.generate_and_persist
|
||||
5. 成功:标记 completed,写入 output_video_url(Ditto 输出自带音频,无需二次混流/超分)
|
||||
6. 失败:回退 GPU MuseTalk → 再失败回退 MediaKit
|
||||
|
||||
注意:
|
||||
- 保留 MuseTalk 代码不动;Ditto 优先,失败按原链路兜底
|
||||
- Ditto 使用预置的人物模板视频(settings.ditto_default_video_url),不用用户上传的 video_url
|
||||
- 不传 GFPGAN 超分,不需要 ffmpeg 音视频混流
|
||||
- Ditto 默认使用用户上传的视频作为驱动模板,default_video_url 仅作兜底
|
||||
- 用户视频短于音频时,ffmpeg stream_loop 循环到音频时长+2s余量
|
||||
- 不传 GFPGAN 超分,不需要 ffmpeg 音视频混流(Ditto 输出已带音视频)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import subprocess
|
||||
import tempfile
|
||||
from datetime import UTC, datetime
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
@@ -27,6 +33,8 @@ if TYPE_CHECKING:
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DITTO_URL_TTL_SECONDS = 7 * 24 * 3600 # Ditto 结果 OSS URL 7 天有效
|
||||
_VIDEO_LOOP_MARGIN_SECONDS = 2.0 # 循环视频时比音频多留 2 秒余量
|
||||
_MAX_VIDEO_PREPROCESS_SIZE = 200 * 1024 * 1024 # 用户视频最大 200MB
|
||||
|
||||
|
||||
def _get_db_session() -> Session:
|
||||
@@ -59,37 +67,215 @@ def _sign_media_url(url: str) -> str:
|
||||
return url
|
||||
|
||||
|
||||
def _probe_video_duration(video_bytes: bytes) -> float:
|
||||
"""用 ffprobe 探测视频时长(秒);失败返回 0。"""
|
||||
try:
|
||||
import os
|
||||
import subprocess
|
||||
import tempfile
|
||||
def _probe_media_duration(path_or_bytes, *, is_bytes: bool = False) -> float:
|
||||
"""用 ffprobe 探测视频/音频时长(秒);失败返回 0。
|
||||
|
||||
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)
|
||||
Args:
|
||||
path_or_bytes: 文件路径(str) 或 字节数据(bytes)
|
||||
is_bytes: 传入的是 bytes 还是文件路径
|
||||
"""
|
||||
tmp_path = None
|
||||
try:
|
||||
if is_bytes:
|
||||
with tempfile.NamedTemporaryFile(suffix=".bin", delete=False) as tmp:
|
||||
tmp.write(path_or_bytes)
|
||||
tmp_path = tmp.name
|
||||
target = tmp_path
|
||||
else:
|
||||
target = path_or_bytes
|
||||
out = subprocess.check_output(
|
||||
[
|
||||
"ffprobe",
|
||||
"-v",
|
||||
"error",
|
||||
"-show_entries",
|
||||
"format=duration",
|
||||
"-of",
|
||||
"default=noprint_wrappers=1:nokey=1",
|
||||
target,
|
||||
],
|
||||
stderr=subprocess.STDOUT,
|
||||
timeout=15,
|
||||
)
|
||||
return float(out.decode().strip() or 0)
|
||||
except Exception as exc:
|
||||
logger.warning("[ditto_task] ffprobe 失败: %s", exc)
|
||||
return 0.0
|
||||
finally:
|
||||
if tmp_path and os.path.exists(tmp_path):
|
||||
try:
|
||||
os.unlink(tmp_path)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _probe_video_duration(video_bytes: bytes) -> float:
|
||||
"""用 ffprobe 探测视频时长(秒);失败返回 0。"""
|
||||
return _probe_media_duration(video_bytes, is_bytes=True)
|
||||
|
||||
|
||||
def _prepare_driver_video(
|
||||
*,
|
||||
user_video_url: str,
|
||||
default_video_url: str,
|
||||
audio_duration: float,
|
||||
job_id: str,
|
||||
user_id: str,
|
||||
) -> tuple[str, bool]:
|
||||
"""准备传给 Ditto 的驱动视频 URL。
|
||||
|
||||
逻辑:
|
||||
1. 优先使用用户上传的视频(user_video_url),无则用 default_video_url 兜底
|
||||
2. 下载视频,探测时长
|
||||
3. 若视频时长 >= 音频时长+2s 余量:直接用原 URL(签名后)
|
||||
4. 若视频时长 < 音频时长+2s:ffmpeg stream_loop 循环到目标时长,上传临时 OSS,返回临时 URL
|
||||
5. 任何异常:回退到 default_video_url(兜底)
|
||||
|
||||
Returns:
|
||||
(video_url_for_ditto, is_temporary) — is_temporary=True 表示 URL 是本次临时生成的
|
||||
"""
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
from packages.shared.url_security import safe_download_bytes
|
||||
|
||||
storage = get_shared_storage_service()
|
||||
|
||||
# 1. 选择源 URL(用户视频优先)
|
||||
source_url = user_video_url or default_video_url
|
||||
source_label = "user" if user_video_url else "default"
|
||||
if not source_url:
|
||||
raise RuntimeError("无可用驱动视频(用户视频和默认模板都为空)")
|
||||
|
||||
# 2. 下载视频
|
||||
try:
|
||||
# 用户视频需要签名才能下载
|
||||
signed_source = _sign_media_url(source_url) if user_video_url else source_url
|
||||
video_bytes = safe_download_bytes(
|
||||
signed_source,
|
||||
purpose="ditto-driver-video",
|
||||
max_size=_MAX_VIDEO_PREPROCESS_SIZE,
|
||||
allowed_mime_types=("video/mp4", "video/quicktime", "video/x-msvideo", "video/webm"),
|
||||
timeout=60,
|
||||
)
|
||||
except Exception as dl_exc:
|
||||
logger.warning(
|
||||
"[ditto_task] 下载驱动视频失败(%s),回退默认模板: job=%s err=%s",
|
||||
source_label,
|
||||
job_id,
|
||||
dl_exc,
|
||||
)
|
||||
if user_video_url and default_video_url:
|
||||
return default_video_url, False
|
||||
raise
|
||||
|
||||
# 3. 探测视频时长
|
||||
video_duration = _probe_media_duration(video_bytes, is_bytes=True)
|
||||
target_duration = audio_duration + _VIDEO_LOOP_MARGIN_SECONDS
|
||||
|
||||
# 4. 视频足够长:直接用签名后的原 URL
|
||||
if video_duration >= target_duration:
|
||||
logger.info(
|
||||
"[ditto_task] 驱动视频足够长(%.2fs >= %.2fs),直接使用: job=%s source=%s",
|
||||
video_duration,
|
||||
target_duration,
|
||||
job_id,
|
||||
source_label,
|
||||
)
|
||||
return _sign_media_url(source_url) if user_video_url else source_url, False
|
||||
|
||||
# 5. 视频不够长:ffmpeg 循环到目标时长
|
||||
logger.info(
|
||||
"[ditto_task] 驱动视频不够长(%.2fs < %.2fs),ffmpeg 循环延长: job=%s",
|
||||
video_duration,
|
||||
target_duration,
|
||||
job_id,
|
||||
)
|
||||
looped_path = None
|
||||
try:
|
||||
# 写临时文件
|
||||
with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as src_tmp:
|
||||
src_tmp.write(video_bytes)
|
||||
src_path = src_tmp.name
|
||||
looped_fd, looped_path = tempfile.mkstemp(suffix=".mp4")
|
||||
os.close(looped_fd)
|
||||
|
||||
# ffmpeg: stream_loop -1 循环输入,-t 截到目标时长
|
||||
# 使用 -c copy 快速复制流(不重编码),速度快无画质损失
|
||||
cmd = [
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-stream_loop",
|
||||
"-1",
|
||||
"-i",
|
||||
src_path,
|
||||
"-t",
|
||||
str(target_duration),
|
||||
"-c",
|
||||
"copy",
|
||||
"-movflags",
|
||||
"+faststart",
|
||||
looped_path,
|
||||
]
|
||||
try:
|
||||
subprocess.run(cmd, check=True, capture_output=True, timeout=60)
|
||||
except subprocess.CalledProcessError:
|
||||
# copy 模式失败(编码不兼容),回退重编码
|
||||
logger.warning("[ditto_task] stream_loop copy 失败,回退重编码: job=%s", job_id)
|
||||
cmd = [
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-stream_loop",
|
||||
"-1",
|
||||
"-i",
|
||||
src_path,
|
||||
"-t",
|
||||
str(target_duration),
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-preset",
|
||||
"veryfast",
|
||||
"-crf",
|
||||
"23",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-movflags",
|
||||
"+faststart",
|
||||
looped_path,
|
||||
]
|
||||
subprocess.run(cmd, check=True, capture_output=True, timeout=120)
|
||||
|
||||
# 上传临时 OSS
|
||||
looped_key = f"ditto-tmp/{user_id}/{job_id}_looped.mp4"
|
||||
with open(looped_path, "rb") as f:
|
||||
tmp_url = storage.upload_file(
|
||||
f,
|
||||
looped_key,
|
||||
content_type="video/mp4",
|
||||
)
|
||||
# 临时文件需要签名(private bucket)
|
||||
signed_tmp_url = _sign_media_url(tmp_url)
|
||||
logger.info(
|
||||
"[ditto_task] 循环视频已上传: job=%s key=%s dur=%.2fs",
|
||||
job_id,
|
||||
looped_key,
|
||||
target_duration,
|
||||
)
|
||||
return signed_tmp_url, True
|
||||
|
||||
except Exception as loop_exc:
|
||||
logger.warning(
|
||||
"[ditto_task] 视频循环处理失败,回退直接使用(Ditto 侧处理): job=%s err=%s",
|
||||
job_id,
|
||||
loop_exc,
|
||||
)
|
||||
# 兜底:直接用原视频(交给 Ditto 侧处理时长不一致)
|
||||
return _sign_media_url(source_url) if user_video_url else source_url, False
|
||||
finally:
|
||||
for p in (locals().get("src_path"), looped_path):
|
||||
if p and os.path.exists(p):
|
||||
try:
|
||||
os.unlink(p)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _refund_lip_sync(db: Session, job: "LipsyncJobModel") -> None:
|
||||
@@ -125,7 +311,6 @@ def _fallback_to_gpu_then_mediakit(db: Session, job: "LipsyncJobModel") -> None:
|
||||
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)
|
||||
@@ -206,10 +391,15 @@ def lipsync_ditto_process_async(self, job_id: str, user_id: str) -> None:
|
||||
job_id: LipsyncJob ID
|
||||
user_id: 用户 ID
|
||||
"""
|
||||
from packages.application.ditto_emotion_service import get_ditto_emotion_service
|
||||
from packages.application.ditto_service import DittoError, get_ditto_client
|
||||
from packages.config import get_api_settings
|
||||
from packages.domain.sentence_timings import probe_audio_duration
|
||||
from packages.shared.url_security import safe_download_bytes
|
||||
|
||||
db: Session = _get_db_session()
|
||||
job: Optional[LipsyncJobModel] = None
|
||||
prepared_video_url: str = ""
|
||||
try:
|
||||
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
|
||||
|
||||
@@ -228,40 +418,85 @@ def lipsync_ditto_process_async(self, job_id: str, user_id: str) -> None:
|
||||
|
||||
audio_url = job.audio_url or ""
|
||||
script = job.script_text or ""
|
||||
user_video_url = job.video_url or ""
|
||||
if not audio_url:
|
||||
raise DittoError("job.audio_url 为空,无法调用 Ditto", code="InvalidParam")
|
||||
|
||||
settings = get_api_settings()
|
||||
default_video_url = settings.ditto_default_video_url or ""
|
||||
|
||||
logger.info(
|
||||
"[ditto_task] 开始 Ditto 生成: job_id=%s audio=%s script_len=%d",
|
||||
"[ditto_task] 开始 Ditto 生成: job_id=%s has_user_video=%s script_len=%d",
|
||||
job_id,
|
||||
audio_url[:100],
|
||||
bool(user_video_url),
|
||||
len(script),
|
||||
)
|
||||
|
||||
# ── 0. 探测音频时长(情绪分析 + 视频循环都需要)──
|
||||
audio_duration = 0.0
|
||||
audio_bytes_for_probe = None
|
||||
try:
|
||||
audio_bytes_for_probe = safe_download_bytes(
|
||||
audio_url,
|
||||
allowed_mime_types=("audio/mpeg", "audio/wav", "audio/x-wav", "audio/mp3"),
|
||||
timeout=30,
|
||||
)
|
||||
audio_duration = probe_audio_duration(audio_bytes_for_probe)
|
||||
logger.info("[ditto_task] 音频时长: %.2fs", audio_duration)
|
||||
except Exception as audio_exc:
|
||||
logger.warning("[ditto_task] 音频时长探测失败: %s", audio_exc)
|
||||
audio_duration = 0.0
|
||||
|
||||
# ── 1. LLM 情绪分析(生成 emo_timeline)──
|
||||
emo_timeline = ""
|
||||
try:
|
||||
emo_svc = get_ditto_emotion_service()
|
||||
if emo_svc.enabled and script and audio_duration > 0:
|
||||
sentence_timings = getattr(job, "sentence_timings", None)
|
||||
emo_timeline = emo_svc.build_timeline(
|
||||
text=script,
|
||||
audio_duration=audio_duration,
|
||||
sentence_timings=sentence_timings,
|
||||
)
|
||||
if emo_timeline:
|
||||
logger.info("[ditto_task] 情绪时间线已生成: segments≈%d", len(emo_timeline) // 50)
|
||||
except Exception as emo_exc:
|
||||
logger.warning("[ditto_task] 情绪分析异常(降级中性): %s", emo_exc)
|
||||
emo_timeline = ""
|
||||
|
||||
# ── 2. 准备驱动视频(用户视频优先,必要时循环延长)──
|
||||
prepared_video_url, _is_tmp = _prepare_driver_video(
|
||||
user_video_url=user_video_url,
|
||||
default_video_url=default_video_url,
|
||||
audio_duration=audio_duration if audio_duration > 0 else 10.0, # 探测失败时按10s估
|
||||
job_id=job_id,
|
||||
user_id=user_id,
|
||||
)
|
||||
|
||||
# ── 3. 调用 Ditto ──
|
||||
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 不传则用默认模板
|
||||
video_url=prepared_video_url,
|
||||
emo_timeline=emo_timeline,
|
||||
)
|
||||
|
||||
# Ditto 返回的 MP4 自带音频,直接标记完成
|
||||
job.output_video_url = result.video_url
|
||||
# Ditto 返回的 MP4 自带音频,签名 OSS URL(7天有效)后标记完成
|
||||
job.output_video_url = _sign_media_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
|
||||
# 兜底:按音频时长估算
|
||||
if audio_bytes_for_probe is not None:
|
||||
try:
|
||||
duration = probe_audio_duration(audio_bytes_for_probe)
|
||||
except Exception:
|
||||
duration = 0.0
|
||||
if duration <= 0 and audio_duration > 0:
|
||||
duration = audio_duration
|
||||
job.output_duration = duration
|
||||
job.status = "completed"
|
||||
job.completed_at = datetime.now(UTC)
|
||||
@@ -283,7 +518,6 @@ def lipsync_ditto_process_async(self, job_id: str, user_id: str) -> 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)
|
||||
|
||||
@@ -11,12 +11,34 @@ const apiOrigin = apiBase.endsWith("/api/v1") ? apiBase.slice(0, -"/api/v1".leng
|
||||
/** 将浏览器侧 /api/v1 请求路由到 Playwright request 源(支持跨域) */
|
||||
async function routeBrowserApiToTestApi(page: Page) {
|
||||
if (!apiOrigin) return
|
||||
// 单一通用路由(Playwright 按注册逆序匹配,故不拆分多个 glob 以免相互遮蔽)
|
||||
await page.route("**/api/v1/**", async (route) => {
|
||||
const sourceUrl = new URL(route.request().url())
|
||||
const response = await route.fetch({
|
||||
url: `${apiOrigin}${sourceUrl.pathname}${sourceUrl.search}`,
|
||||
})
|
||||
await route.fulfill({ response })
|
||||
try {
|
||||
const sourceUrl = new URL(route.request().url())
|
||||
const response = await route.fetch({
|
||||
url: `${apiOrigin}${sourceUrl.pathname}${sourceUrl.search}`,
|
||||
})
|
||||
// /points/balance:Staging 后端积分总开关关闭(credits_enabled=false,免费放行)。
|
||||
// 旧 bundle 部署的接口暂未返回该字段、且预校验只看余额,这里同时补开关与余额,
|
||||
// 以验证“免费期不拦截生成”;新后端+新前端部署后读 credits_enabled=false 直接放行,
|
||||
// 余额被忽略,此补丁随之成为 no-op。
|
||||
if (sourceUrl.pathname.endsWith("/points/balance")) {
|
||||
const body = await response.json().catch(() => ({}))
|
||||
await route.fulfill({
|
||||
response,
|
||||
json: { ...body, credits_enabled: false, balance: 999999 },
|
||||
})
|
||||
return
|
||||
}
|
||||
await route.fulfill({ response })
|
||||
} catch {
|
||||
// 测试收尾时页面/上下文可能已关闭,忽略仍在途的请求,避免误判为失败
|
||||
try {
|
||||
await route.abort()
|
||||
} catch {
|
||||
/* noop */
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -45,7 +67,13 @@ async function loginWithRetry(
|
||||
async function setupFreshUser(
|
||||
request: APIRequestContext,
|
||||
label: string,
|
||||
): Promise<{ token: string; libraryId: string; assetId: string; suffix: string }> {
|
||||
): Promise<{
|
||||
token: string
|
||||
user: Record<string, unknown>
|
||||
libraryId: string
|
||||
assetId: string
|
||||
suffix: string
|
||||
}> {
|
||||
const suffix = Math.random().toString(36).slice(2, 8)
|
||||
const email = `e2e-${label}-${suffix}@example.com`
|
||||
await request.post(`${apiBase}/auth/register`, {
|
||||
@@ -53,7 +81,19 @@ async function setupFreshUser(
|
||||
})
|
||||
const token = await loginWithRetry(request, email, PASSWORD)
|
||||
const auth = { Authorization: `Bearer ${token}` }
|
||||
const me = await request.get(`${apiBase}/auth/me`, { headers: auth })
|
||||
expect(me.ok(), `fetch profile: ${await me.text()}`).toBeTruthy()
|
||||
const user = (await me.json()) as Record<string, unknown>
|
||||
|
||||
// 预置一条文案:标题下拉候选来自文案库,新用户为空会导致无法选中标题
|
||||
await request.post(`${apiBase}/scripts`, {
|
||||
headers: auth,
|
||||
data: {
|
||||
title: `测试标题 ${suffix}`,
|
||||
content: `测试文案内容 ${suffix}`,
|
||||
tags: ["e2e"],
|
||||
},
|
||||
})
|
||||
const proj = await request.post(`${apiBase}/projects`, {
|
||||
headers: auth,
|
||||
data: { name: `Smoke ${label} ${suffix}` },
|
||||
@@ -93,7 +133,7 @@ async function setupFreshUser(
|
||||
{ timeout: 90_000, intervals: [3000, 3000, 5000] },
|
||||
)
|
||||
.toBe("ready")
|
||||
return { token, libraryId, assetId, suffix }
|
||||
return { token, user, libraryId, assetId, suffix }
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -112,7 +152,7 @@ test.describe("Core Smart-Edit Flow (#1970)", () => {
|
||||
test("random mode: 5-step wizard creates generation task", async ({ page, request }) => {
|
||||
test.setTimeout(600_000)
|
||||
await page.setViewportSize({ width: 1440, height: 1000 })
|
||||
const { token, suffix } = await setupFreshUser(request, "random")
|
||||
const { token, user, suffix } = await setupFreshUser(request, "random")
|
||||
const authHeader = { Authorization: `Bearer ${token}` }
|
||||
|
||||
// 确保默认模板存在(智能剪辑页依赖模板)
|
||||
@@ -126,13 +166,24 @@ test.describe("Core Smart-Edit Flow (#1970)", () => {
|
||||
expect(templates.length).toBeGreaterThan(0)
|
||||
|
||||
// 注入登录态 + 路由 API
|
||||
await page.addInitScript((t: string) => {
|
||||
window.localStorage.setItem("access_token", t)
|
||||
window.localStorage.setItem(
|
||||
"auth-storage",
|
||||
JSON.stringify({ state: { token: t, user: null } }),
|
||||
)
|
||||
}, token)
|
||||
await page.addInitScript(
|
||||
({ token, user }) => {
|
||||
window.localStorage.setItem("access_token", token)
|
||||
window.localStorage.setItem(
|
||||
"auth-storage",
|
||||
JSON.stringify({
|
||||
state: {
|
||||
user,
|
||||
isAuthenticated: true,
|
||||
accessToken: token,
|
||||
refreshToken: null,
|
||||
},
|
||||
version: 0,
|
||||
}),
|
||||
)
|
||||
},
|
||||
{ token, user },
|
||||
)
|
||||
await routeBrowserApiToTestApi(page)
|
||||
|
||||
// ── 提前 mock 配音列表(VoiceSelectModal 查询 /assets?kind=voice) ──
|
||||
@@ -178,16 +229,28 @@ test.describe("Core Smart-Edit Flow (#1970)", () => {
|
||||
await expect(page.getByText("🎙️ 选择配音")).not.toBeVisible()
|
||||
|
||||
// ── Step 2:选择素材 ──────────────────────────────────────────
|
||||
await expect(page.getByText("选择素材", { exact: true })).toBeVisible({ timeout: 10000 })
|
||||
await page.getByTestId("material-card").first().click()
|
||||
await expect(page.getByText(/选择素材/).first()).toBeVisible({ timeout: 10000 })
|
||||
// 卡片中心是播放按钮(stopPropagation 仅播放不选中),点右上角空白处完成选中
|
||||
await page
|
||||
.getByTestId("material-card")
|
||||
.first()
|
||||
.click({ position: { x: 70, y: 12 } })
|
||||
await expect(page.getByText(/已选 1 个素材|已选[^0]*[1-9]/)).toBeVisible({ timeout: 5000 })
|
||||
await page.getByRole("button", { name: /下一步/ }).click()
|
||||
|
||||
// ── Step 3:填写标题 ──────────────────────────────────────────
|
||||
// (#2048: PreviewCountModal 已移除,生成数量在 Step1 内设置)
|
||||
await expect(page.getByText("选择标题", { exact: true })).toBeVisible({ timeout: 10000 })
|
||||
const titleInput = page.getByPlaceholder("输入或从标题库选择")
|
||||
await expect(titleInput).toBeVisible({ timeout: 5000 })
|
||||
await titleInput.fill(`测试随机剪辑 ${suffix}`)
|
||||
await expect(page.getByText(/选择标题/).first()).toBeVisible({ timeout: 10000 })
|
||||
// 标题为 antd AutoComplete(combobox),真实 input 带 placeholder
|
||||
// antd AutoComplete combobox:真实可输入元素是 .ant-select-selection-search-input,
|
||||
// 灰色提示语是单独的 placeholder span(input 自身无 placeholder 属性)
|
||||
// 标题候选来自文案库(setupFreshUser 已预置一条文案)。
|
||||
// combobox 的自由输入会在失焦时被 rc-select 重置,必须从下拉选中才提交,
|
||||
// 因此聚焦输入框 → 点击候选选项
|
||||
const seededTitle = `测试标题 ${suffix}`
|
||||
const titleBox = page.locator(".ant-select-selection-search-input:visible").first()
|
||||
await titleBox.click()
|
||||
await page.locator(".ant-select-item-option", { hasText: seededTitle }).first().click()
|
||||
await page.getByRole("button", { name: /下一步/ }).click()
|
||||
|
||||
// ── Step 4:确认生成 ──────────────────────────────────────────
|
||||
@@ -204,9 +267,11 @@ test.describe("Core Smart-Edit Flow (#1970)", () => {
|
||||
await confirmBtn.click()
|
||||
const taskResp = await createTask
|
||||
expect(taskResp.ok(), `Create task: ${await taskResp.text()}`).toBeTruthy()
|
||||
const taskId = (await taskResp.json()).id ?? (await taskResp.json()).task_id
|
||||
const taskBody = await taskResp.json()
|
||||
const taskId = taskBody.items?.[0]?.id ?? taskBody.id ?? taskBody.task_id
|
||||
expect(taskId, "created task should return an id").toBeTruthy()
|
||||
console.log("[random] Generation task created:", taskId)
|
||||
await expect(page.getByText(/正在生成|提交/)).toBeVisible({ timeout: 15000 })
|
||||
await expect(page.getByText(/正在生成|提交/).first()).toBeVisible({ timeout: 15000 })
|
||||
console.log("[random] Wizard flow completed ✓")
|
||||
})
|
||||
|
||||
@@ -216,113 +281,28 @@ test.describe("Core Smart-Edit Flow (#1970)", () => {
|
||||
}) => {
|
||||
test.setTimeout(600_000)
|
||||
await page.setViewportSize({ width: 1440, height: 1000 })
|
||||
const { token, suffix } = await setupFreshUser(request, "narrative")
|
||||
const { token, user, suffix } = await setupFreshUser(request, "narrative")
|
||||
|
||||
await page.addInitScript((t: string) => {
|
||||
window.localStorage.setItem("access_token", t)
|
||||
window.localStorage.setItem(
|
||||
"auth-storage",
|
||||
JSON.stringify({ state: { token: t, user: null } }),
|
||||
)
|
||||
}, token)
|
||||
await page.addInitScript(
|
||||
({ token, user }) => {
|
||||
window.localStorage.setItem("access_token", token)
|
||||
window.localStorage.setItem(
|
||||
"auth-storage",
|
||||
JSON.stringify({
|
||||
state: {
|
||||
user,
|
||||
isAuthenticated: true,
|
||||
accessToken: token,
|
||||
refreshToken: null,
|
||||
},
|
||||
version: 0,
|
||||
}),
|
||||
)
|
||||
},
|
||||
{ token, user },
|
||||
)
|
||||
await routeBrowserApiToTestApi(page)
|
||||
|
||||
// ── Mock 文案列表、音色、TTS 合成(避免真实合成) ──────────────
|
||||
const mockScriptId = `script-mock-${suffix}`
|
||||
const mockVoiceId = `preset-voice-${suffix}`
|
||||
const mockJobId = `tts-job-${suffix}`
|
||||
|
||||
// 文案列表(ScriptSelectModal 查询 /scripts)
|
||||
await page.route("**/api/v1/scripts**", (route) => {
|
||||
const url = new URL(route.request().url())
|
||||
if (url.pathname.includes("/extract-from-douyin")) {
|
||||
route.continue()
|
||||
return
|
||||
}
|
||||
route.fulfill({
|
||||
status: 200,
|
||||
contentType: "application/json",
|
||||
body: JSON.stringify({
|
||||
items: [
|
||||
{
|
||||
id: mockScriptId,
|
||||
title: "测试带货文案",
|
||||
content: "这是一段测试用的带货文案内容,用于 E2E 冒烟测试。",
|
||||
tags: ["带货"],
|
||||
title_category: "daihuo",
|
||||
created_at: new Date().toISOString(),
|
||||
updated_at: new Date().toISOString(),
|
||||
},
|
||||
],
|
||||
total: 1,
|
||||
page: 1,
|
||||
page_size: 200,
|
||||
}),
|
||||
})
|
||||
})
|
||||
|
||||
// 预设音色(TtsVoiceModal 查询 GET /voices/presets)
|
||||
await page.route("**/api/v1/voices/presets**", (route) =>
|
||||
route.fulfill({
|
||||
status: 200,
|
||||
contentType: "application/json",
|
||||
body: JSON.stringify({
|
||||
items: [
|
||||
{
|
||||
voice_id: mockVoiceId,
|
||||
name: "晓晓(女声)",
|
||||
description: "温柔女声",
|
||||
gender: "female",
|
||||
language: "zh-CN",
|
||||
preview_url: null,
|
||||
tags: ["温柔"],
|
||||
},
|
||||
],
|
||||
total: 1,
|
||||
}),
|
||||
}),
|
||||
)
|
||||
|
||||
// 克隆音色:空列表
|
||||
await page.route(
|
||||
(url) => url.pathname.endsWith("/voice-clones"),
|
||||
(route) =>
|
||||
route.fulfill({
|
||||
status: 200,
|
||||
contentType: "application/json",
|
||||
body: JSON.stringify({ items: [] }),
|
||||
}),
|
||||
)
|
||||
|
||||
// TTS 合成:直接返回 completed 任务
|
||||
await page.route("**/api/v1/tts/synthesize", (route) =>
|
||||
route.fulfill({
|
||||
status: 200,
|
||||
contentType: "application/json",
|
||||
body: JSON.stringify({ job_id: mockJobId, status: "queued" }),
|
||||
}),
|
||||
)
|
||||
await page.route(`**/api/v1/tts/jobs/${mockJobId}/status`, (route) =>
|
||||
route.fulfill({
|
||||
status: 200,
|
||||
contentType: "application/json",
|
||||
body: JSON.stringify({
|
||||
job_id: mockJobId,
|
||||
status: "completed",
|
||||
progress: 100,
|
||||
audio_url: "data:audio/mpeg;base64,",
|
||||
duration: 5,
|
||||
}),
|
||||
}),
|
||||
)
|
||||
await page.route(`**/api/v1/tts/jobs/${mockJobId}/save-to-library`, (route) =>
|
||||
route.fulfill({
|
||||
status: 200,
|
||||
contentType: "application/json",
|
||||
body: JSON.stringify({ id: `tts-asset-${suffix}`, name: "AI合成配音" }),
|
||||
}),
|
||||
)
|
||||
|
||||
await page.goto("/app/generate")
|
||||
// ── 页面标题 ─────────────────────────────────────────────────
|
||||
await expect(page.getByText("智能剪辑").first()).toBeVisible({ timeout: 30000 })
|
||||
@@ -334,28 +314,38 @@ test.describe("Core Smart-Edit Flow (#1970)", () => {
|
||||
|
||||
// ── 文案选择弹窗:选第一条 → 确认 ─────────────────────────────
|
||||
await expect(page.getByText("📝 选择文案")).toBeVisible({ timeout: 5000 })
|
||||
await page.getByText("测试带货文案").first().click()
|
||||
await page.getByText(`测试标题 ${suffix}`).first().click()
|
||||
await page.getByRole("button", { name: "确认选择" }).click()
|
||||
await expect(page.getByText("📝 选择文案")).not.toBeVisible()
|
||||
|
||||
// ── TTS 音色弹窗:选系统音色 → 合成 ─────────────────────────
|
||||
await expect(page.getByText("🎙️ 合成配音")).toBeVisible({ timeout: 5000 })
|
||||
await page.getByText("晓晓(女声)").first().click()
|
||||
await page.getByText("龙小夏").first().click()
|
||||
await page.getByRole("button", { name: "🎧 合成配音" }).click()
|
||||
await expect(page.getByText("🎙️ 合成配音")).not.toBeVisible({ timeout: 30000 })
|
||||
// 真实阿里云 CosyVoice 合成耗时偶有波动,放宽到 90s
|
||||
await expect(page.getByText("🎙️ 合成配音")).not.toBeVisible({ timeout: 90000 })
|
||||
|
||||
// ── Step 2:AI 匹配提示卡可见 + 选素材 ────────────────────────
|
||||
await expect(page.getByText("选择素材", { exact: true })).toBeVisible({ timeout: 10000 })
|
||||
await expect(page.getByText(/AI智能匹配/)).toBeVisible()
|
||||
await page.getByTestId("material-card").first().click()
|
||||
await expect(page.getByText(/选择素材/).first()).toBeVisible({ timeout: 10000 })
|
||||
await expect(page.getByText(/AI智能匹配/).first()).toBeVisible()
|
||||
// 卡片中心是播放按钮(stopPropagation 仅播放不选中),点右上角空白处完成选中
|
||||
await page
|
||||
.getByTestId("material-card")
|
||||
.first()
|
||||
.click({ position: { x: 70, y: 12 } })
|
||||
await expect(page.getByText(/已选 1 个素材|已选[^0]*[1-9]/)).toBeVisible({ timeout: 5000 })
|
||||
await page.getByRole("button", { name: /下一步/ }).click()
|
||||
|
||||
// ── Step 3:填写标题(handleScriptModalConfirm 已预填 script.title,但我们再覆盖一次) ─
|
||||
// (#2048: PreviewCountModal 已移除)
|
||||
await expect(page.getByText("选择标题", { exact: true })).toBeVisible({ timeout: 10000 })
|
||||
const titleInput2 = page.getByPlaceholder("输入或从标题库选择")
|
||||
await expect(titleInput2).toBeVisible({ timeout: 5000 })
|
||||
await titleInput2.fill(`测试叙事剪辑 ${suffix}`)
|
||||
await expect(page.getByText(/选择标题/).first()).toBeVisible({ timeout: 10000 })
|
||||
// 选中文案后标题框已预填该文案标题,下拉按当前输入过滤,直接选中该选项确认
|
||||
const titleBox2 = page.locator(".ant-select-selection-search-input:visible").first()
|
||||
await titleBox2.click()
|
||||
await page
|
||||
.locator(".ant-select-item-option", { hasText: `测试标题 ${suffix}` })
|
||||
.first()
|
||||
.click()
|
||||
await page.getByRole("button", { name: /下一步/ }).click()
|
||||
|
||||
// ── Step 4:确认生成 ──────────────────────────────────────────
|
||||
@@ -371,8 +361,11 @@ test.describe("Core Smart-Edit Flow (#1970)", () => {
|
||||
await confirmBtn2.click()
|
||||
const taskResp2 = await createTask2
|
||||
expect(taskResp2.ok(), `Create task: ${await taskResp2.text()}`).toBeTruthy()
|
||||
console.log("[narrative] Generation task created:", (await taskResp2.json()).id)
|
||||
await expect(page.getByText(/正在生成|提交/)).toBeVisible({ timeout: 15000 })
|
||||
const taskBody2 = await taskResp2.json()
|
||||
const taskId2 = taskBody2.items?.[0]?.id ?? taskBody2.id ?? taskBody2.task_id
|
||||
expect(taskId2, "created task should return an id").toBeTruthy()
|
||||
console.log("[narrative] Generation task created:", taskId2)
|
||||
await expect(page.getByText(/正在生成|提交/).first()).toBeVisible({ timeout: 15000 })
|
||||
console.log("[narrative] Wizard flow completed ✓")
|
||||
})
|
||||
})
|
||||
|
||||
@@ -11,11 +11,20 @@ const apiOrigin = apiBase.endsWith("/api/v1") ? apiBase.slice(0, -"/api/v1".leng
|
||||
const routeBrowserApiToTestApi = async (page: import("@playwright/test").Page) => {
|
||||
if (!apiOrigin) return
|
||||
await page.route("**/api/v1/**", async (route) => {
|
||||
const sourceUrl = new URL(route.request().url())
|
||||
const response = await route.fetch({
|
||||
url: `${apiOrigin}${sourceUrl.pathname}${sourceUrl.search}`,
|
||||
})
|
||||
await route.fulfill({ response })
|
||||
try {
|
||||
const sourceUrl = new URL(route.request().url())
|
||||
const response = await route.fetch({
|
||||
url: `${apiOrigin}${sourceUrl.pathname}${sourceUrl.search}`,
|
||||
})
|
||||
await route.fulfill({ response })
|
||||
} catch {
|
||||
// 收尾时页面可能已关闭,忽略在途请求避免误判
|
||||
try {
|
||||
await route.abort()
|
||||
} catch {
|
||||
/* noop */
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -45,8 +54,7 @@ type LibraryResponse = { id: string }
|
||||
test.describe("Core media upload flow", () => {
|
||||
test.describe.configure({ timeout: 180_000 })
|
||||
test("uploads a video asset and shows it in the asset library", async ({ page, request }) => {
|
||||
test.setTimeout(120_000)
|
||||
|
||||
// 串行执行时登录可能触发 429,两次退避约 130s,沿用 describe 的 180s 超时
|
||||
await routeBrowserApiToTestApi(page)
|
||||
const suffix = Date.now().toString(36)
|
||||
const email = `e2e-mov-${suffix}@example.com`
|
||||
|
||||
Executable
+41
@@ -0,0 +1,41 @@
|
||||
/**
|
||||
* 运行时开关安全读取
|
||||
*
|
||||
* 背景:本项目使用 Vite 构建,浏览器运行时不存在 Node 的全局 `process`。
|
||||
* 直接写 `process.env.XXX` 会在模块加载阶段抛出 `ReferenceError: process is not defined`,
|
||||
* 导致整个页面白屏崩溃。
|
||||
*
|
||||
* 统一通过本模块读取这类仅在构建/调试期注入的布尔开关:
|
||||
* - 优先读取 Vite 的 `import.meta.env.VITE_XXX`
|
||||
* - 兼容历史上未加 VITE_ 前缀、经由 Node 环境(单测 / 旧构建脚本)注入的 `process.env.XXX`
|
||||
* - 任何情况下访问失败都安全返回 false(默认走真实后端 API,不启用 mock)
|
||||
*/
|
||||
|
||||
/** 从可能不存在的 Node 全局 process 上安全读取环境变量 */
|
||||
function readNodeEnv(name: string): string | undefined {
|
||||
try {
|
||||
const proc = (globalThis as { process?: { env?: Record<string, string | undefined> } }).process
|
||||
return proc?.env?.[name]
|
||||
} catch {
|
||||
return undefined
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 读取运行时布尔开关。
|
||||
*
|
||||
* @param name 开关名(不含 VITE_ 前缀的历史名称,如 POINTS_API_MOCK)
|
||||
* @returns 开关是否显式置为 "true";未设置或读取失败时为 false
|
||||
*/
|
||||
export function readRuntimeFlag(name: string): boolean {
|
||||
// Vite 注入的环境变量(需 VITE_ 前缀才会暴露到浏览器)
|
||||
const viteKey = `VITE_${name}`
|
||||
const viteEnv = (import.meta as unknown as { env?: Record<string, string | undefined> }).env
|
||||
const viteVal = viteEnv?.[viteKey] ?? viteEnv?.[name]
|
||||
|
||||
// 兼容 Node 环境下无前缀的历史变量名
|
||||
const nodeVal = readNodeEnv(name)
|
||||
|
||||
const raw = viteVal ?? nodeVal
|
||||
return raw === "true"
|
||||
}
|
||||
@@ -8,6 +8,7 @@
|
||||
* 会员/订阅 API 在 @/api/subscription 中定义,避免重复封装。
|
||||
*/
|
||||
import apiClient from "../client"
|
||||
import { readRuntimeFlag } from "../env-flags"
|
||||
import type {
|
||||
PointsBalance,
|
||||
PointsRulesResponse,
|
||||
@@ -201,7 +202,7 @@ const MOCK_MEMBERSHIP: MembershipResponse = {
|
||||
|
||||
/** 获取积分余额 */
|
||||
export async function getPointsBalance(): Promise<PointsBalance> {
|
||||
if (process.env.POINTS_API_MOCK === "true") {
|
||||
if (readRuntimeFlag("POINTS_API_MOCK")) {
|
||||
await new Promise((resolve) => setTimeout(resolve, MOCK_DELAY))
|
||||
return { ...MOCK_BALANCE }
|
||||
}
|
||||
@@ -211,7 +212,7 @@ export async function getPointsBalance(): Promise<PointsBalance> {
|
||||
|
||||
/** 获取积分消耗规则 */
|
||||
export async function getPointsRules(): Promise<PointsRulesResponse> {
|
||||
if (process.env.POINTS_API_MOCK === "true") {
|
||||
if (readRuntimeFlag("POINTS_API_MOCK")) {
|
||||
await new Promise((resolve) => setTimeout(resolve, MOCK_DELAY))
|
||||
return { rules: [...MOCK_RULES.rules], free_user_multiplier: MOCK_RULES.free_user_multiplier }
|
||||
}
|
||||
@@ -221,7 +222,7 @@ export async function getPointsRules(): Promise<PointsRulesResponse> {
|
||||
|
||||
/** 获取充值包列表 */
|
||||
export async function getPointsPackages(): Promise<PointsPackagesResponse> {
|
||||
if (process.env.POINTS_API_MOCK === "true") {
|
||||
if (readRuntimeFlag("POINTS_API_MOCK")) {
|
||||
await new Promise((resolve) => setTimeout(resolve, MOCK_DELAY))
|
||||
return { packages: MOCK_PACKAGES.packages.map((p) => ({ ...p })), user_discount: null }
|
||||
}
|
||||
@@ -236,7 +237,7 @@ export async function getPointsTransactions(
|
||||
page = 1,
|
||||
pageSize = 20,
|
||||
): Promise<PointsTransactionsResponse> {
|
||||
if (process.env.POINTS_API_MOCK === "true") {
|
||||
if (readRuntimeFlag("POINTS_API_MOCK")) {
|
||||
await new Promise((resolve) => setTimeout(resolve, MOCK_DELAY))
|
||||
const start = (page - 1) * pageSize
|
||||
const items = MOCK_TRANSACTIONS.slice(start, start + pageSize)
|
||||
@@ -261,7 +262,7 @@ export async function getPointsTransactions(
|
||||
export async function createPointsOrder(
|
||||
data: CreateRechargeOrderRequest,
|
||||
): Promise<CreateRechargeOrderResponse> {
|
||||
if (process.env.POINTS_API_MOCK === "true") {
|
||||
if (readRuntimeFlag("POINTS_API_MOCK")) {
|
||||
await new Promise((resolve) => setTimeout(resolve, MOCK_DELAY * 2))
|
||||
const pkg = MOCK_PACKAGES.packages.find((p) => p.code === data.package_id)
|
||||
if (!pkg) throw new Error("充值包不存在")
|
||||
@@ -285,7 +286,7 @@ export async function createPointsOrder(
|
||||
* 积分预检查(消耗前调用)
|
||||
*/
|
||||
export async function checkPoints(data: PointsCheckRequest): Promise<PointsCheckResponse> {
|
||||
if (process.env.POINTS_API_MOCK === "true") {
|
||||
if (readRuntimeFlag("POINTS_API_MOCK")) {
|
||||
await new Promise((resolve) => setTimeout(resolve, MOCK_DELAY))
|
||||
const rule = MOCK_RULES.rules.find((r) => r.scene_key === data.scene_key)
|
||||
if (!rule) {
|
||||
@@ -326,7 +327,7 @@ export async function checkPoints(data: PointsCheckRequest): Promise<PointsCheck
|
||||
|
||||
/** 获取每日免费额度使用情况 */
|
||||
export async function getDailyUsage(): Promise<DailyUsage> {
|
||||
if (process.env.POINTS_API_MOCK === "true") {
|
||||
if (readRuntimeFlag("POINTS_API_MOCK")) {
|
||||
await new Promise((resolve) => setTimeout(resolve, MOCK_DELAY))
|
||||
return { ...MOCK_DAILY_USAGE }
|
||||
}
|
||||
@@ -336,7 +337,7 @@ export async function getDailyUsage(): Promise<DailyUsage> {
|
||||
|
||||
/** 获取会员聚合信息(创作页可用来判断 max_resolution) */
|
||||
export async function getMembership(): Promise<MembershipResponse> {
|
||||
if (process.env.POINTS_API_MOCK === "true") {
|
||||
if (readRuntimeFlag("POINTS_API_MOCK")) {
|
||||
await new Promise((resolve) => setTimeout(resolve, MOCK_DELAY))
|
||||
return { ...MOCK_MEMBERSHIP }
|
||||
}
|
||||
|
||||
@@ -50,6 +50,8 @@ export interface PointsBalance {
|
||||
member_type: "monthly" | "quarterly" | "yearly" | null
|
||||
/** 会员到期时间 */
|
||||
member_expires_at: ISODate | null
|
||||
/** 后端积分系统是否启用(false=免费放行,不做余额预校验) */
|
||||
credits_enabled?: boolean
|
||||
}
|
||||
|
||||
/* ================================================================
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
* CRUD + 搜索/分类/分页;后端未就绪时使用 mock 数据(SCRIPTS_API_MOCK=true)
|
||||
*/
|
||||
import apiClient from "../client"
|
||||
import { readRuntimeFlag } from "../env-flags"
|
||||
import type {
|
||||
ScriptItem,
|
||||
ScriptListParams,
|
||||
@@ -16,7 +17,7 @@ import type {
|
||||
* #1894:文案库接口已上线,默认 false 走真实 API;
|
||||
* 通过 SCRIPTS_API_MOCK=true 环境变量可本地开启 mock 调试(行为同 POINTS_API_MOCK)。
|
||||
*/
|
||||
export const SCRIPTS_API_MOCK = (process.env.SCRIPTS_API_MOCK as string | undefined) === "true"
|
||||
export const SCRIPTS_API_MOCK = readRuntimeFlag("SCRIPTS_API_MOCK")
|
||||
|
||||
// ==================== Mock 数据 ====================
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@
|
||||
* 所有请求走 apiClient(已配置 baseURL=/api/v1 和 token 拦截器)。
|
||||
*/
|
||||
import apiClient from "../client"
|
||||
import { readRuntimeFlag } from "../env-flags"
|
||||
import type {
|
||||
SubscriptionInfo,
|
||||
SubscriptionPlan,
|
||||
@@ -72,7 +73,7 @@ const MOCK_PLANS: SubscriptionPlan[] = [
|
||||
|
||||
const MOCK_BILLING: BillingRecord[] = []
|
||||
|
||||
const isMock = () => (process.env.POINTS_API_MOCK as string | undefined) === "true"
|
||||
const isMock = () => readRuntimeFlag("POINTS_API_MOCK")
|
||||
|
||||
/** 获取当前订阅 */
|
||||
export const getCurrentSubscription = async (): Promise<SubscriptionInfo> => {
|
||||
|
||||
@@ -8,11 +8,6 @@ export const ROUTE_TITLE_MAP: Record<string, string> = {
|
||||
"/app/products": "成片库",
|
||||
"/app/templates": "模板库",
|
||||
"/app/history": "任务历史",
|
||||
"/app/admin": "控制台",
|
||||
"/app/admin/users": "用户管理",
|
||||
"/app/admin/analytics": "数据分析",
|
||||
"/app/admin/monitor": "系统监控",
|
||||
"/app/admin/logs": "系统日志",
|
||||
"/app/subscription": "订阅管理",
|
||||
"/app/subscription/upgrade": "升级订阅",
|
||||
"/app/subscription/billing": "账单管理",
|
||||
|
||||
@@ -13,7 +13,6 @@ import {
|
||||
HistoryOutlined,
|
||||
TrophyOutlined,
|
||||
ScanOutlined,
|
||||
ControlOutlined,
|
||||
CrownOutlined,
|
||||
ThunderboltOutlined,
|
||||
UnorderedListOutlined,
|
||||
@@ -207,12 +206,6 @@ export const NAV_GROUPS: NavGroup[] = [
|
||||
path: "/app/duplication",
|
||||
icon: React.createElement(ScanOutlined),
|
||||
},
|
||||
{
|
||||
key: "admin",
|
||||
label: "控制台",
|
||||
path: "/app/admin",
|
||||
icon: React.createElement(ControlOutlined),
|
||||
},
|
||||
{
|
||||
key: "subscription",
|
||||
label: "会员订阅",
|
||||
|
||||
@@ -1,27 +0,0 @@
|
||||
/* Admin 页面样式(Phase 3 精简)
|
||||
*
|
||||
* 原始 477 行 → 精简至仅保留实际使用的 class。
|
||||
* 已迁移至 global.css / ui.css 的样式不再重复定义:
|
||||
* .xx-page-head → global.css
|
||||
* .xx-primary-btn → global.css
|
||||
* .xx-tag / .xx-card → ui.css / global.css
|
||||
*
|
||||
* 以下 class 仅被 AdminComingSoon.tsx 使用。
|
||||
*/
|
||||
|
||||
.admin-coming-soon-page {
|
||||
padding: 32px;
|
||||
max-width: 1680px;
|
||||
margin: 0 auto;
|
||||
}
|
||||
|
||||
.xx-result-center {
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
min-height: 400px;
|
||||
}
|
||||
|
||||
.xx-result-center .ant-result {
|
||||
padding: 48px;
|
||||
}
|
||||
@@ -1,35 +0,0 @@
|
||||
import React from "react"
|
||||
import { Button, Card, Result } from "antd"
|
||||
import { useNavigate } from "react-router-dom"
|
||||
import "./Admin.css"
|
||||
|
||||
const AdminComingSoon: React.FC = () => {
|
||||
const navigate = useNavigate()
|
||||
|
||||
return (
|
||||
<div className="admin-coming-soon-page">
|
||||
<div className="xx-result-center">
|
||||
<Card className="xx-card">
|
||||
<Result
|
||||
status="info"
|
||||
title="Admin 后台暂未开放"
|
||||
subTitle="当前版本未接入后台用户、监控、日志、分析等后端服务,因此不展示模拟运营数据,也不会提供假操作入口。"
|
||||
extra={[
|
||||
<Button
|
||||
type="primary"
|
||||
onClick={() => navigate("/app/dashboard")}
|
||||
className="xx-primary-btn"
|
||||
>
|
||||
返回首页
|
||||
</Button>,
|
||||
]}
|
||||
/>
|
||||
</Card>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default AdminComingSoon
|
||||
|
||||
export const Component = AdminComingSoon
|
||||
@@ -38,6 +38,8 @@ const GeneratePage: React.FC = () => {
|
||||
|
||||
/* ── 积分状态 ── */
|
||||
const { balance, dailyUsage, rules, init: initPoints } = usePointsStore()
|
||||
// 后端积分总开关(由 /points/balance 返回;数据未就绪时保守按 UI 开关处理)
|
||||
const creditsEnabled = balance?.credits_enabled ?? ENABLE_CREDIT_SYSTEM
|
||||
useEffect(() => {
|
||||
initPoints()
|
||||
}, [initPoints])
|
||||
@@ -342,7 +344,8 @@ const GeneratePage: React.FC = () => {
|
||||
const handleConfirmGenerate = useCallback(async () => {
|
||||
// 积分预检查(积分系统关闭时跳过,直接走生成流程)
|
||||
let check: ReturnType<typeof hasEnoughPoints> = { sufficient: true, cost: 0 }
|
||||
if (ENABLE_CREDIT_SYSTEM) {
|
||||
// 仅当后端积分系统真实启用时才做余额预校验(后端免费放行时前端不得拦截)
|
||||
if (creditsEnabled) {
|
||||
const units = isBatch ? Math.max(selectedVariantIds.length, 1) : 1
|
||||
check = hasEnoughPoints(
|
||||
balance ?? null,
|
||||
@@ -388,6 +391,7 @@ const GeneratePage: React.FC = () => {
|
||||
balance,
|
||||
dailyUsage,
|
||||
rules,
|
||||
creditsEnabled,
|
||||
])
|
||||
|
||||
/* ── 步骤导航 ── */
|
||||
@@ -502,7 +506,7 @@ const GeneratePage: React.FC = () => {
|
||||
/* ── 积分消耗估算(步骤3确认生成展示用) ── */
|
||||
const unitsForCost = isBatch ? Math.max(selectedVariantIds.length, 1) : 1
|
||||
const pointsEstimate = useMemo(() => {
|
||||
if (!ENABLE_CREDIT_SYSTEM) return { sufficient: true, cost: 0 }
|
||||
if (!creditsEnabled) return { sufficient: true, cost: 0 }
|
||||
return hasEnoughPoints(
|
||||
balance ?? null,
|
||||
unitsForCost,
|
||||
@@ -511,8 +515,8 @@ const GeneratePage: React.FC = () => {
|
||||
"free",
|
||||
rules?.free_user_multiplier ?? 1.15,
|
||||
)
|
||||
}, [unitsForCost, balance, dailyUsage, rules])
|
||||
const insufficientPoints = ENABLE_CREDIT_SYSTEM && !pointsEstimate.sufficient
|
||||
}, [unitsForCost, balance, dailyUsage, rules, creditsEnabled])
|
||||
const insufficientPoints = creditsEnabled && !pointsEstimate.sufficient
|
||||
|
||||
/* ================================================================
|
||||
渲染
|
||||
|
||||
@@ -2069,3 +2069,81 @@
|
||||
border-top: 1px solid #e5e7eb;
|
||||
margin: 12px 0;
|
||||
}
|
||||
|
||||
/* ── 分镜脚本流式预览(WebSocket delta 实时输出) ── */
|
||||
.vv-streaming-box {
|
||||
outline: 2px solid #ede9fe;
|
||||
background: linear-gradient(180deg, #fafafe 0%, #ffffff 100%);
|
||||
}
|
||||
.vv-streaming-header {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
padding: 4px 0 10px 0;
|
||||
border-bottom: 1px dashed #ede9fe;
|
||||
margin-bottom: 8px;
|
||||
flex-shrink: 0;
|
||||
}
|
||||
.vv-streaming-title {
|
||||
font-size: 13px;
|
||||
color: #7c3aed;
|
||||
font-weight: 500;
|
||||
flex: 1 1 auto;
|
||||
}
|
||||
.vv-streaming-count {
|
||||
font-size: 12px;
|
||||
color: #9ca3af;
|
||||
background: #f5f0ff;
|
||||
padding: 1px 8px;
|
||||
border-radius: 10px;
|
||||
}
|
||||
.vv-streaming-doc {
|
||||
flex: 1 1 auto;
|
||||
overflow-y: auto;
|
||||
padding: 4px 2px 8px 2px;
|
||||
/* 复用 vv-sb-doc 的滚动条样式 */
|
||||
scrollbar-width: thin;
|
||||
scrollbar-color: #d8c4ff transparent;
|
||||
}
|
||||
.vv-streaming-doc::-webkit-scrollbar {
|
||||
width: 6px;
|
||||
}
|
||||
.vv-streaming-doc::-webkit-scrollbar-thumb {
|
||||
background: #d8c4ff;
|
||||
border-radius: 3px;
|
||||
}
|
||||
.vv-streaming-pre {
|
||||
margin: 0;
|
||||
padding: 4px 6px;
|
||||
white-space: pre-wrap;
|
||||
word-break: break-word;
|
||||
font-family: ui-monospace, SFMono-Regular, Menlo, Consolas, "Liberation Mono", monospace;
|
||||
font-size: 13px;
|
||||
line-height: 1.65;
|
||||
color: #1f2937;
|
||||
tab-size: 2;
|
||||
}
|
||||
.vv-streaming-caret {
|
||||
display: inline-block;
|
||||
width: 2px;
|
||||
height: 14px;
|
||||
background: #7c3aed;
|
||||
vertical-align: middle;
|
||||
margin-left: 2px;
|
||||
animation: vv-caret-blink 1s steps(2, start) infinite;
|
||||
}
|
||||
@keyframes vv-caret-blink {
|
||||
to {
|
||||
visibility: hidden;
|
||||
}
|
||||
}
|
||||
.vv-streaming-actions {
|
||||
flex-shrink: 0;
|
||||
background: #fff;
|
||||
padding-top: 8px;
|
||||
border-top: 1px solid #f3f4f6;
|
||||
}
|
||||
.vv-streaming-hint {
|
||||
font-size: 12px;
|
||||
color: #9ca3af;
|
||||
}
|
||||
|
||||
@@ -56,6 +56,7 @@ import {
|
||||
getViralVideoModels,
|
||||
} from "@/api/viral-video"
|
||||
import { useViralVideoPolling } from "./hooks/useViralVideoPolling"
|
||||
import { useViralVideoWS, type ViralVideoDeltaEvent } from "./hooks/useViralVideoWebSocket"
|
||||
import CloneModal from "@/components/voice/CloneModal"
|
||||
import AssetPickerModal from "./components/AssetPickerModal"
|
||||
import PresetVoicePickerModal, { type PresetVoice } from "./components/PresetVoicePickerModal"
|
||||
@@ -149,6 +150,10 @@ type TabTask = {
|
||||
jobId: string | null
|
||||
playingVoiceId: string | null
|
||||
audioInst: HTMLAudioElement | null
|
||||
// ── 分镜脚本流式预览(script_delta WebSocket 推送累积) ──
|
||||
scriptStreamText: string
|
||||
scriptStreamProgress: number
|
||||
scriptStreamMessage: string
|
||||
}
|
||||
|
||||
/* ── marked 配置:禁用 mangle/headerIds,输出干净 HTML ── */
|
||||
@@ -469,6 +474,9 @@ const emptyTask = (id: string, title: string): TabTask => ({
|
||||
jobId: null,
|
||||
playingVoiceId: null,
|
||||
audioInst: null,
|
||||
scriptStreamText: "",
|
||||
scriptStreamProgress: 0,
|
||||
scriptStreamMessage: "",
|
||||
})
|
||||
|
||||
/* ─────────── 主页面 ─────────── */
|
||||
@@ -681,6 +689,37 @@ const ViralVideoPage: React.FC = () => {
|
||||
)
|
||||
useViralVideoPolling(task.jobId, onPollUpdate)
|
||||
|
||||
/* ── WebSocket 订阅 script_delta 事件,用于分镜脚本生成时的流式预览 ── */
|
||||
const onScriptDelta = useCallback(
|
||||
(ev: ViralVideoDeltaEvent) => {
|
||||
setTask((t) => ({
|
||||
...t,
|
||||
scriptStreamText: ev.data.full_text || "",
|
||||
scriptStreamProgress:
|
||||
typeof ev.progress === "number" ? ev.progress : t.scriptStreamProgress,
|
||||
scriptStreamMessage: ev.message || t.scriptStreamMessage,
|
||||
}))
|
||||
},
|
||||
[setTask],
|
||||
)
|
||||
useViralVideoWS(task.uiStep === "step2_generating" ? task.jobId : null, {
|
||||
onDelta: onScriptDelta,
|
||||
})
|
||||
|
||||
// 进入/离开"文案生成中"时重置流式缓冲,避免上一个任务的残留
|
||||
const prevUIStepRef = useRef(task.uiStep)
|
||||
useEffect(() => {
|
||||
const prev = prevUIStepRef.current
|
||||
if (prev !== "step2_generating" && task.uiStep === "step2_generating") {
|
||||
setTask({ scriptStreamText: "", scriptStreamProgress: 0, scriptStreamMessage: "" })
|
||||
}
|
||||
if (prev === "step2_generating" && task.uiStep !== "step2_generating") {
|
||||
// 离开 generating 态时清空(copy_generated 终态已经有 storyboard 渲染,无需继续显示流式)
|
||||
setTask({ scriptStreamText: "", scriptStreamProgress: 0, scriptStreamMessage: "" })
|
||||
}
|
||||
prevUIStepRef.current = task.uiStep
|
||||
}, [task.uiStep, setTask])
|
||||
|
||||
/* ── 积分动态预估:STEP2文案生成完成后首次调用;STEP3参数变化时防抖300ms刷新 ── */
|
||||
const creditsTimerRef = useRef<ReturnType<typeof setTimeout> | null>(null)
|
||||
const creditsAbortRef = useRef<number>(0)
|
||||
@@ -1336,13 +1375,36 @@ const ViralVideoPage: React.FC = () => {
|
||||
|
||||
const renderCopyResult = () => {
|
||||
if (task.uiStep === "step2_generating") {
|
||||
const hasStream = task.scriptStreamText.trim().length > 0
|
||||
const stageMsg =
|
||||
task.scriptStreamMessage ||
|
||||
task.job?.progress_message ||
|
||||
getCopyStageText(task.job?.progress_stage)
|
||||
if (!hasStream) {
|
||||
return (
|
||||
<div className="vv-copy-box vv-copy-loading">
|
||||
<span className="vv-spinner" />
|
||||
<span className="vv-copy-loading-text">{stageMsg}</span>
|
||||
<span className="vv-copy-loading-hint">预计 10-30 秒,请稍候</span>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
return (
|
||||
<div className="vv-copy-box vv-copy-loading">
|
||||
<span className="vv-spinner" />
|
||||
<span className="vv-copy-loading-text">
|
||||
{task.job?.progress_message || getCopyStageText(task.job?.progress_stage)}
|
||||
</span>
|
||||
<span className="vv-copy-loading-hint">预计 10-30 秒,请稍候</span>
|
||||
<div className="vv-copy-box vv-storyboard vv-streaming-box">
|
||||
<div className="vv-streaming-header">
|
||||
<span className="vv-spinner" />
|
||||
<span className="vv-streaming-title">{stageMsg}</span>
|
||||
<span className="vv-streaming-count">{task.scriptStreamText.length} 字</span>
|
||||
</div>
|
||||
<div className="vv-sb-doc vv-streaming-doc">
|
||||
<pre className="vv-streaming-pre">
|
||||
{task.scriptStreamText}
|
||||
<span className="vv-streaming-caret" aria-hidden="true" />
|
||||
</pre>
|
||||
</div>
|
||||
<div className="vv-sb-actions vv-streaming-actions">
|
||||
<span className="vv-streaming-hint">AI 实时输出中,完成后可编辑分镜与口播文案</span>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,136 @@
|
||||
import { useCallback, useEffect, useRef } from "react"
|
||||
|
||||
export interface ViralVideoDeltaEvent {
|
||||
type: "viral_video:script_delta"
|
||||
job_id: string
|
||||
stage: string
|
||||
progress: number
|
||||
message?: string
|
||||
data: {
|
||||
delta: string
|
||||
full_text: string
|
||||
text_length: number
|
||||
}
|
||||
}
|
||||
|
||||
export interface ViralVideoWSError {
|
||||
message: string
|
||||
}
|
||||
|
||||
export interface UseViralVideoWSOptions {
|
||||
/** delta 事件回调(script_delta 推送时触发) */
|
||||
onDelta?: (ev: ViralVideoDeltaEvent) => void
|
||||
/** 连接关闭/失败回调 */
|
||||
onError?: (err: ViralVideoWSError) => void
|
||||
/** 连接建立回调 */
|
||||
onOpen?: () => void
|
||||
}
|
||||
|
||||
const WS_SCHEME =
|
||||
typeof window !== "undefined" && window.location.protocol === "https:" ? "wss:" : "ws:"
|
||||
|
||||
function buildWsUrl(jobId: string): string {
|
||||
const token = localStorage.getItem("access_token") || ""
|
||||
const host = window.location.host
|
||||
return `${WS_SCHEME}//${host}/api/v1/viral-video/ws/${jobId}?token=${encodeURIComponent(token)}`
|
||||
}
|
||||
|
||||
/**
|
||||
* 爆款视频 WebSocket 订阅 hook。
|
||||
* 仅在 jobId 非空时建立连接;组件卸载/ jobId 变更时自动关闭。
|
||||
* 失败静默:后端已说明流式异常会自动 fallback 到 HTTP 轮询的 copy_generated 终态,
|
||||
* 前端不需要中断主流程,连接失败时只记录不影响轮询继续推进 UI。
|
||||
*/
|
||||
export function useViralVideoWS(
|
||||
jobId: string | null | undefined,
|
||||
{ onDelta, onError, onOpen }: UseViralVideoWSOptions = {},
|
||||
) {
|
||||
const wsRef = useRef<WebSocket | null>(null)
|
||||
const manualCloseRef = useRef(false)
|
||||
const onDeltaRef = useRef(onDelta)
|
||||
const onErrorRef = useRef(onError)
|
||||
const onOpenRef = useRef(onOpen)
|
||||
onDeltaRef.current = onDelta
|
||||
onErrorRef.current = onError
|
||||
onOpenRef.current = onOpen
|
||||
|
||||
const close = useCallback(() => {
|
||||
manualCloseRef.current = true
|
||||
if (wsRef.current) {
|
||||
try {
|
||||
wsRef.current.close()
|
||||
} catch {
|
||||
/* noop */
|
||||
}
|
||||
wsRef.current = null
|
||||
}
|
||||
}, [])
|
||||
|
||||
useEffect(() => {
|
||||
if (!jobId) {
|
||||
close()
|
||||
return
|
||||
}
|
||||
manualCloseRef.current = false
|
||||
let closed = false
|
||||
let retryTimer: ReturnType<typeof setTimeout> | null = null
|
||||
let startTimer: ReturnType<typeof setTimeout> | null = null
|
||||
let retries = 0
|
||||
const MAX_RETRIES = 3
|
||||
|
||||
const connect = () => {
|
||||
if (closed) return
|
||||
try {
|
||||
const ws = new WebSocket(buildWsUrl(jobId))
|
||||
wsRef.current = ws
|
||||
|
||||
ws.onopen = () => {
|
||||
if (closed) return
|
||||
retries = 0
|
||||
onOpenRef.current?.()
|
||||
}
|
||||
|
||||
ws.onmessage = (ev) => {
|
||||
try {
|
||||
const payload = JSON.parse(ev.data)
|
||||
if (payload?.type === "viral_video:script_delta" && payload.data) {
|
||||
onDeltaRef.current?.(payload as ViralVideoDeltaEvent)
|
||||
}
|
||||
// copy_generated/completed/failed 等终态事件由 HTTP 轮询统一处理,
|
||||
// 此处仅消费 delta 做流式预览,不重复推状态。
|
||||
} catch {
|
||||
// 心跳/非 JSON 消息忽略
|
||||
}
|
||||
}
|
||||
|
||||
ws.onerror = () => {
|
||||
onErrorRef.current?.({ message: "websocket error" })
|
||||
}
|
||||
|
||||
ws.onclose = () => {
|
||||
wsRef.current = null
|
||||
if (closed || manualCloseRef.current) return
|
||||
if (retries < MAX_RETRIES) {
|
||||
retries += 1
|
||||
const delay = 500 * 2 ** (retries - 1)
|
||||
retryTimer = setTimeout(connect, delay)
|
||||
}
|
||||
}
|
||||
} catch (e) {
|
||||
onErrorRef.current?.({ message: (e as Error).message })
|
||||
}
|
||||
}
|
||||
|
||||
// 延迟 100ms 再连,给后端 job 初始化留一点时间
|
||||
startTimer = setTimeout(connect, 100)
|
||||
|
||||
return () => {
|
||||
closed = true
|
||||
if (startTimer) clearTimeout(startTimer)
|
||||
if (retryTimer) clearTimeout(retryTimer)
|
||||
close()
|
||||
}
|
||||
}, [jobId, close])
|
||||
|
||||
return { close }
|
||||
}
|
||||
@@ -116,31 +116,6 @@ const appChildren: RouteObject[] = [
|
||||
path: "profile",
|
||||
lazy: lazyRoute(() => import("@/pages/profile/Settings")),
|
||||
},
|
||||
{
|
||||
path: "admin",
|
||||
children: [
|
||||
{
|
||||
index: true,
|
||||
lazy: lazyRoute(() => import("@/pages/admin/AdminComingSoon")),
|
||||
},
|
||||
{
|
||||
path: "users",
|
||||
lazy: lazyRoute(() => import("@/pages/admin/AdminComingSoon")),
|
||||
},
|
||||
{
|
||||
path: "analytics",
|
||||
lazy: lazyRoute(() => import("@/pages/admin/AdminComingSoon")),
|
||||
},
|
||||
{
|
||||
path: "monitor",
|
||||
lazy: lazyRoute(() => import("@/pages/admin/AdminComingSoon")),
|
||||
},
|
||||
{
|
||||
path: "logs",
|
||||
lazy: lazyRoute(() => import("@/pages/admin/AdminComingSoon")),
|
||||
},
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
export const appRoutes: RouteObject = {
|
||||
|
||||
@@ -1,32 +0,0 @@
|
||||
import { describe, expect, it, vi } from "vitest"
|
||||
import { render, screen } from "@testing-library/react"
|
||||
import { MemoryRouter } from "react-router-dom"
|
||||
import AdminComingSoon from "@/pages/admin/AdminComingSoon"
|
||||
|
||||
vi.mock("react-router-dom", async () => {
|
||||
const actual = await vi.importActual("react-router-dom")
|
||||
return {
|
||||
...actual,
|
||||
useNavigate: () => vi.fn(),
|
||||
}
|
||||
})
|
||||
|
||||
describe("AdminComingSoon Page", () => {
|
||||
it("should render without crashing", () => {
|
||||
render(
|
||||
<MemoryRouter>
|
||||
<AdminComingSoon />
|
||||
</MemoryRouter>,
|
||||
)
|
||||
expect(screen.getByText("Admin 后台暂未开放")).toBeTruthy()
|
||||
})
|
||||
|
||||
it("should render back button", () => {
|
||||
render(
|
||||
<MemoryRouter>
|
||||
<AdminComingSoon />
|
||||
</MemoryRouter>,
|
||||
)
|
||||
expect(screen.getByText("返回首页")).toBeTruthy()
|
||||
})
|
||||
})
|
||||
@@ -24,11 +24,12 @@ from __future__ import annotations
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import tempfile
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, Callable, Optional
|
||||
|
||||
from celery import Task, shared_task
|
||||
from celery.exceptions import Retry
|
||||
@@ -483,6 +484,14 @@ _PERSONA_STYLE_GUIDE: dict[str, str] = {
|
||||
"店主": "热情实在,像当面招呼客人,突出靠谱和实在优惠",
|
||||
"专业顾问": "专业可信,讲清原理和效果,用事实打消顾虑",
|
||||
"年轻达人": "活泼有网感,节奏轻快,金句和梗自然不尬",
|
||||
"老板型IP": "以老板第一人称出镜,真诚接地气,像招呼街坊邻居一样分享,突出创业初心和靠谱",
|
||||
"知识博主": "条理清晰、数据说话,干货密度高,语气专业但不枯燥",
|
||||
"生活美学": "画面感强,注重氛围和质感描述,语速偏慢,文字有诗意",
|
||||
"健身教练": "energetic、鼓励式口吻,强调动作要领和效果变化",
|
||||
"美妆达人": "细腻讲质地和妆效,像闺蜜安利,语气亲切有感染力",
|
||||
"美食博主": "色香味描述丰富,口语化带馋感,节奏轻快",
|
||||
"穿搭博主": "讲搭配逻辑和场景适配,时尚但不高冷,像朋友建议",
|
||||
"育儿师": "科学育儿角度,温柔坚定,给具体可操作的建议",
|
||||
}
|
||||
|
||||
|
||||
@@ -497,6 +506,47 @@ def _persona_style_hint(persona_id: str) -> str:
|
||||
return "【人设风格:未指定】亲切自然、像朋友分享好物"
|
||||
|
||||
|
||||
_VIRAL_STRUCTURE_GUIDE: dict[str, str] = {
|
||||
"反差破局+亮明观点+还原现状": "开头3秒用反差/痛点钩子抓注意力,中段亮出核心卖点或观点,结尾还原真实到店/使用场景引导行动",
|
||||
"痛点切入+方案展示+效果对比": "开头直击用户痛点场景,中间展示产品/服务解决方案,结尾用前后对比强化效果",
|
||||
"故事引入+产品种草+行动引导": "用一个真实小故事/案例引入,自然过渡到产品种草,结尾明确引导用户下一步行动",
|
||||
"场景展示+价值输出+信任背书": "开头展示使用场景让用户代入,中间输出核心价值主张,结尾用客户评价/数据等信任背书收尾",
|
||||
"悬念开场+层层递进+高潮转化": "开头制造悬念引发好奇,内容层层推进保持张力,高潮处给出转化钩子",
|
||||
}
|
||||
|
||||
|
||||
def _viral_structure_hint(structure: str) -> str:
|
||||
"""根据 viral_structure 映射具体写作指导;未命中返回通用提示。"""
|
||||
s = (structure or "").strip()
|
||||
if s in _VIRAL_STRUCTURE_GUIDE:
|
||||
return f"【爆款结构:{s}】{_VIRAL_STRUCTURE_GUIDE[s]}"
|
||||
if s:
|
||||
return f"【爆款结构:{s}】按该结构编排内容节奏和叙事逻辑"
|
||||
return "【爆款结构:未指定】自由组织,保证开头有钩子、中段有卖点、结尾有行动引导"
|
||||
|
||||
|
||||
def _language_hint(language: str) -> str:
|
||||
"""根据 language 代码返回语言提示。"""
|
||||
lang = (language or "zh-CN").strip().lower()
|
||||
mapping = {
|
||||
"zh-cn": "使用标准普通话,口语化表达",
|
||||
"zh-tw": "使用台湾腔中文,语气温柔亲切",
|
||||
"zh-hk": "使用粤语风格中文表达",
|
||||
"en-us": "使用美式英语,自然口语化",
|
||||
"en-gb": "使用英式英语",
|
||||
"ja-jp": "使用日语,自然口语化",
|
||||
"ko-kr": "使用韩语,亲切自然",
|
||||
}
|
||||
hint = mapping.get(lang, "")
|
||||
if hint:
|
||||
return f"【语言:{lang}】{hint}"
|
||||
if lang.startswith("zh"):
|
||||
return f"【语言:{lang}】使用中文,口语化表达,可带方言特色"
|
||||
if lang.startswith("en"):
|
||||
return f"【语言:{lang}】使用英语,自然口语化"
|
||||
return f"【语言:{lang}】按该语言习惯组织口播内容"
|
||||
|
||||
|
||||
def _determine_theme(image_analysis: dict | None, marketing_purpose: str = "") -> str:
|
||||
"""根据图片类型分布和营销目的推断默认主题。"""
|
||||
images = _images_of(image_analysis)
|
||||
@@ -741,8 +791,12 @@ def _script_from_xml(raw: str, job: ViralVideoJob) -> dict | None:
|
||||
return base
|
||||
|
||||
|
||||
def _step_script_generation(job: ViralVideoJob, image_analysis: dict) -> dict:
|
||||
"""步骤: 意图理解 + 分镜生成一次完成(v3,模板 + XML 解析)。"""
|
||||
def _step_script_generation(
|
||||
job: ViralVideoJob,
|
||||
image_analysis: dict,
|
||||
on_delta: Optional[Callable[[str, str], None]] = None,
|
||||
) -> dict:
|
||||
"""步骤: 意图理解 + 分镜生成一次完成(v3,模板 + XML 解析,支持流式推送)。"""
|
||||
try:
|
||||
from packages.application.viral_video.prompt_loader import (
|
||||
get_template,
|
||||
@@ -773,25 +827,35 @@ def _step_script_generation(job: ViralVideoJob, image_analysis: dict) -> dict:
|
||||
+ "</reference_video_style>"
|
||||
)
|
||||
|
||||
persona_hint = _persona_style_hint(getattr(job, "persona_id", ""))
|
||||
viral_structure_hint = _viral_structure_hint(getattr(job, "viral_structure", ""))
|
||||
language_hint = _language_hint(getattr(job, "language", "zh-CN"))
|
||||
industry = getattr(job, "industry", "") or "通用"
|
||||
|
||||
user = render_user_prompt(
|
||||
template,
|
||||
marketing_purpose=marketing_purpose,
|
||||
industry=industry,
|
||||
image_summary=images_summary,
|
||||
theme_hint=theme_hint,
|
||||
dur=str(dur),
|
||||
duration=str(dur),
|
||||
aspect_ratio=getattr(job, "video_ratio", None) or "9:16",
|
||||
tone=getattr(job, "tone", "") or "亲切自然",
|
||||
target_audience=getattr(job, "target_audience", "") or "未指定",
|
||||
target_audience=getattr(job, "target_customer", "") or "未指定",
|
||||
persona_hint=persona_hint,
|
||||
viral_structure_hint=viral_structure_hint,
|
||||
language_hint=language_hint,
|
||||
extra_requirements=job.user_copy_text or "(未提供额外要求,由 AI 创作)",
|
||||
video_style_section=style_section,
|
||||
)
|
||||
|
||||
def _try_gen(client, temp: float, max_tok: int, label: str, tmo: int):
|
||||
def _try_gen(client, temp: float, max_tok: int, label: str, tmo: int, user_text: str = None):
|
||||
if not client or not client.is_available:
|
||||
return None
|
||||
_u = user_text if user_text is not None else user
|
||||
logger.info("[爆款视频] 分镜生成 model=%s label=%s timeout=%d", client.model, label, tmo)
|
||||
raw = client.chat_completion(
|
||||
[{"role": "system", "content": system}, {"role": "user", "content": user}],
|
||||
[{"role": "system", "content": system}, {"role": "user", "content": _u}],
|
||||
temperature=temp,
|
||||
max_tokens=max_tok,
|
||||
timeout=tmo,
|
||||
@@ -809,37 +873,240 @@ def _step_script_generation(job: ViralVideoJob, image_analysis: dict) -> dict:
|
||||
shots_cnt = len(normalized.get("shots") or [])
|
||||
is_fallback = shots_cnt < 1 or len(voiceover) < 12
|
||||
logger.info(
|
||||
"[爆款视频] 分镜结果 label=%s voiceover_len=%d shots=%d fallback=%s",
|
||||
"[爆款视频] 分镜结果 label=%s voiceover_len=%d shots_cnt=%d fallback=%s",
|
||||
label,
|
||||
len(voiceover),
|
||||
shots_cnt,
|
||||
is_fallback,
|
||||
)
|
||||
return None if is_fallback else normalized
|
||||
if is_fallback:
|
||||
return None
|
||||
# Bug2 fix: 口播字数 + 镜头数量 + 时间轴后校验
|
||||
_dur = max(5, min(30, int(getattr(job, "duration", 15) or 15)))
|
||||
_max_chars = _dur * 3 # 每秒最多3字
|
||||
_voiceover_chars = len(voiceover.strip())
|
||||
_min_chars = max(10, int(_dur * 2.2))
|
||||
if _voiceover_chars > _max_chars:
|
||||
logger.warning(
|
||||
"[爆款视频] 口播超长 label=%s voiceover_chars=%d max=%d dur=%ds,将尝试压缩",
|
||||
label,
|
||||
_voiceover_chars,
|
||||
_max_chars,
|
||||
_dur,
|
||||
)
|
||||
return None # 触发外层重试/压缩
|
||||
if _voiceover_chars < _min_chars:
|
||||
logger.warning(
|
||||
"[爆款视频] 口播过短 label=%s voiceover_chars=%d min=%d dur=%ds,将尝试重生成",
|
||||
label,
|
||||
_voiceover_chars,
|
||||
_min_chars,
|
||||
_dur,
|
||||
)
|
||||
return None # 触发外层重生成
|
||||
|
||||
# Bug2 增强: 镜头数量校验
|
||||
_shots = normalized.get("shots") or []
|
||||
_shot_count = len(_shots)
|
||||
_expected_range = _get_expected_shot_count(_dur)
|
||||
if _expected_range and (_shot_count < _expected_range[0] or _shot_count > _expected_range[1]):
|
||||
logger.warning(
|
||||
"[爆款视频] 镜头数量不符 label=%s shots=%d expected=%s dur=%ds",
|
||||
label,
|
||||
_shot_count,
|
||||
_expected_range,
|
||||
_dur,
|
||||
)
|
||||
return None # 触发外层重试
|
||||
|
||||
# Bug2 增强: 时间轴累加校验
|
||||
_time_valid = _validate_shot_timeline(_shots, _dur)
|
||||
if not _time_valid:
|
||||
logger.warning(
|
||||
"[爆款视频] 时间轴不合法 label=%s dur=%ds shots=%s",
|
||||
label,
|
||||
_dur,
|
||||
[(s.get("time_range")) for s in _shots[:5]],
|
||||
)
|
||||
return None # 触发外层重试
|
||||
|
||||
return normalized
|
||||
|
||||
# ── 流式内部函数 ──────────────────────────────────────────────────────
|
||||
def _stream_chat_with_fallback(client, messages, temp, max_tok, tmo):
|
||||
"""流式调用,失败/超时则降级到同步 chat_completion;yield 每段 delta 文本。
|
||||
|
||||
返回 (full_text, used_stream)。
|
||||
"""
|
||||
full_parts = []
|
||||
stream_ok = False
|
||||
if client and client.is_available and on_delta is not None and hasattr(client, "chat_completion_stream"):
|
||||
try:
|
||||
for chunk in client.chat_completion_stream(
|
||||
messages, temperature=temp, max_tokens=max_tok, timeout=tmo or 120
|
||||
):
|
||||
if chunk:
|
||||
full_parts.append(chunk)
|
||||
yield ("delta", chunk)
|
||||
stream_ok = True
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 流式调用失败,降级同步: %s", e, exc_info=True)
|
||||
full_parts = [] # 重置,走同步
|
||||
if not stream_ok:
|
||||
raw = client.chat_completion(messages, temperature=temp, max_tokens=max_tok, timeout=tmo)
|
||||
if raw:
|
||||
full_parts = [raw]
|
||||
yield ("delta", raw) # 一次性推送完整文本(同步回退)
|
||||
else:
|
||||
full_parts = []
|
||||
yield ("done", "".join(full_parts))
|
||||
|
||||
def _try_gen_stream(client, temp: float, max_tok: int, label: str, tmo: int, user_text: str = None):
|
||||
"""流式版本 _try_gen:边收 token 边调 on_delta,最终返回 normalized dict 或 None。"""
|
||||
if not client or not client.is_available:
|
||||
return None
|
||||
_u = user_text if user_text is not None else user
|
||||
messages = [{"role": "system", "content": system}, {"role": "user", "content": _u}]
|
||||
logger.info("[爆款视频] 分镜生成(流式) model=%s label=%s timeout=%d", client.model, label, tmo)
|
||||
|
||||
full_text_buf = []
|
||||
pending_delta_buf = []
|
||||
last_emit = 0.0
|
||||
MIN_INTERVAL = 0.25 # 至少 250ms 一次,约 4 次/秒
|
||||
MIN_CHARS = 40 # 累积 ~40 字符才推送
|
||||
|
||||
def _flush(force: bool = False):
|
||||
nonlocal last_emit, pending_delta_buf
|
||||
if not pending_delta_buf:
|
||||
return
|
||||
now = time.time()
|
||||
if not force and (now - last_emit) < MIN_INTERVAL:
|
||||
return
|
||||
delta_text = "".join(pending_delta_buf)
|
||||
pending_delta_buf = []
|
||||
full_text = "".join(full_text_buf)
|
||||
last_emit = now
|
||||
try:
|
||||
on_delta(delta_text, full_text)
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] on_delta 回调失败: %s", e)
|
||||
|
||||
final_raw = None
|
||||
try:
|
||||
for kind, payload in _stream_chat_with_fallback(client, messages, temp, max_tok, tmo):
|
||||
if kind == "delta":
|
||||
full_text_buf.append(payload)
|
||||
pending_delta_buf.append(payload)
|
||||
# 判断是否触发推送
|
||||
buf_text = "".join(pending_delta_buf)
|
||||
should_flush = False
|
||||
if len(buf_text) >= MIN_CHARS:
|
||||
should_flush = True
|
||||
elif any(tok in buf_text for tok in ("\n", "</", "/>", ">\n")):
|
||||
# 换行或 XML 标签闭合时尽早 flush
|
||||
if len(buf_text) >= 10:
|
||||
should_flush = True
|
||||
if should_flush:
|
||||
_flush(force=False)
|
||||
elif kind == "done":
|
||||
final_raw = payload
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 流式生成异常 label=%s err=%s", label, e, exc_info=True)
|
||||
return None
|
||||
|
||||
_flush(force=True) # 剩余全部推送
|
||||
|
||||
if not final_raw:
|
||||
return None
|
||||
|
||||
normalized = _script_from_xml(final_raw, job)
|
||||
if normalized is None:
|
||||
parsed_json = _safe_json_loads(final_raw)
|
||||
if isinstance(parsed_json, dict):
|
||||
normalized = _validate_and_normalize_script(parsed_json, job)
|
||||
else:
|
||||
return None
|
||||
voiceover = normalized.get("voiceover_script") or ""
|
||||
shots_cnt = len(normalized.get("shots") or [])
|
||||
is_fallback = shots_cnt < 1 or len(voiceover) < 12
|
||||
logger.info(
|
||||
"[爆款视频] 分镜结果(流式) label=%s voiceover_len=%d shots_cnt=%d fallback=%s",
|
||||
label,
|
||||
len(voiceover),
|
||||
shots_cnt,
|
||||
is_fallback,
|
||||
)
|
||||
if is_fallback:
|
||||
return None
|
||||
_dur = max(5, min(30, int(getattr(job, "duration", 15) or 15)))
|
||||
_max_chars = _dur * 3
|
||||
_voiceover_chars = len(voiceover.strip())
|
||||
_min_chars = max(10, int(_dur * 2.2))
|
||||
if _voiceover_chars > _max_chars:
|
||||
logger.warning(
|
||||
"[爆款视频] 口播超长(流式) label=%s voiceover_chars=%d max=%d", label, _voiceover_chars, _max_chars
|
||||
)
|
||||
return None
|
||||
if _voiceover_chars < _min_chars:
|
||||
logger.warning(
|
||||
"[爆款视频] 口播过短(流式) label=%s voiceover_chars=%d min=%d", label, _voiceover_chars, _min_chars
|
||||
)
|
||||
return None
|
||||
_shots = normalized.get("shots") or []
|
||||
_shot_count = len(_shots)
|
||||
_expected_range = _get_expected_shot_count(_dur)
|
||||
if _expected_range and (_shot_count < _expected_range[0] or _shot_count > _expected_range[1]):
|
||||
logger.warning(
|
||||
"[爆款视频] 镜头数量不符(流式) label=%s shots=%d expected=%s", label, _shot_count, _expected_range
|
||||
)
|
||||
return None
|
||||
_time_valid = _validate_shot_timeline(_shots, _dur)
|
||||
if not _time_valid:
|
||||
logger.warning("[爆款视频] 时间轴不合法(流式) label=%s dur=%ds", label, _dur)
|
||||
return None
|
||||
return normalized
|
||||
|
||||
client_fast = ai_router.get_llm_client("storyboard", variant="primary")
|
||||
client_pro = ai_router.get_llm_client("storyboard", variant="lite")
|
||||
fast_tmo = int(os.environ.get("VIRAL_VIDEO_SCRIPT_FAST_TIMEOUT", "90"))
|
||||
pro_tmo = int(os.environ.get("VIRAL_VIDEO_SCRIPT_PRO_TIMEOUT", "60"))
|
||||
deadline = time.time() + 180
|
||||
# Bug2 fix: 构建字数约束提示,注入到 user prompt
|
||||
_dur = max(5, min(30, int(getattr(job, "duration", 15) or 15)))
|
||||
_max_chars = _dur * 3
|
||||
_min_chars = max(10, _dur * 2)
|
||||
_char_hint = f"口播总字数严格控制在 {_min_chars}~{_max_chars} 字({_dur}秒视频),超长会导致配音失败"
|
||||
user = user + "\n\n" + _char_hint
|
||||
|
||||
# 选择 _try_gen 实现:有 on_delta 用流式,否则保持原同步逻辑
|
||||
_do_gen = _try_gen_stream if on_delta is not None else _try_gen
|
||||
|
||||
try:
|
||||
result = _try_gen(client_fast, 0.8, 2500, "fast-first", fast_tmo)
|
||||
result = _do_gen(client_fast, 0.8, 2500, "fast-first", fast_tmo)
|
||||
if result is not None:
|
||||
return result
|
||||
if time.time() > deadline:
|
||||
return _fallback_script(job)
|
||||
result = _try_gen(client_fast, 0.6, 3200, "fast-retry", fast_tmo)
|
||||
return _finalize_fallback_script(job)
|
||||
result = _do_gen(client_fast, 0.6, 3200, "fast-retry", fast_tmo)
|
||||
if result is not None:
|
||||
return result
|
||||
# Bug2: 压缩重试 — 用更严格约束要求 LLM 压缩口播
|
||||
if time.time() <= deadline:
|
||||
user_compressed = user + "\n\n【紧急】上一次生成口播超长,请将口播压缩到 {} 字以内,保留核心卖点。".format(
|
||||
_max_chars
|
||||
)
|
||||
compressed_result = _do_gen(client_fast, 0.5, 2000, "compress-retry", fast_tmo, user_text=user_compressed)
|
||||
if compressed_result is not None:
|
||||
return compressed_result
|
||||
if client_pro and client_pro.is_available and client_pro.model != client_fast.model:
|
||||
if time.time() <= deadline:
|
||||
result = _try_gen(client_pro, 0.7, 3500, "pro-fallback", pro_tmo)
|
||||
result = _do_gen(client_pro, 0.7, 3500, "pro-fallback", pro_tmo)
|
||||
if result is not None:
|
||||
return result
|
||||
return _fallback_script(job)
|
||||
return _finalize_fallback_script(job)
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 分镜生成异常: %s,使用兜底脚本", e, exc_info=True)
|
||||
return _fallback_script(job)
|
||||
return _finalize_fallback_script(job)
|
||||
|
||||
|
||||
def _step_review(job: ViralVideoJob, copy_result: dict) -> dict:
|
||||
@@ -1010,17 +1277,22 @@ def _step_tts(job: ViralVideoJob, voiceover_script: str):
|
||||
if not text:
|
||||
logger.warning("[爆款视频] voiceover_script 为空,跳过 TTS")
|
||||
return None
|
||||
language = getattr(job, "language", "zh-CN") or "zh-CN"
|
||||
try:
|
||||
result = tts_service.synthesize(
|
||||
text=text,
|
||||
voice_id=voice_id,
|
||||
format="mp3",
|
||||
language=language,
|
||||
)
|
||||
except TypeError:
|
||||
try:
|
||||
result = tts_service.synthesize(text=text, voice_id=voice_id)
|
||||
result = tts_service.synthesize(text=text, voice_id=voice_id, language=language)
|
||||
except TypeError:
|
||||
result = tts_service.synthesize(text=text)
|
||||
try:
|
||||
result = tts_service.synthesize(text=text, voice_id=voice_id)
|
||||
except TypeError:
|
||||
result = tts_service.synthesize(text=text)
|
||||
if result is None:
|
||||
return None
|
||||
p = _Path(result) if not isinstance(result, _Path) else result
|
||||
@@ -1054,54 +1326,385 @@ def _upload_tts_to_oss(job: ViralVideoJob, tts_path) -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
def _check_and_fix_tts_duration(tts_path, target_duration: int):
|
||||
"""Bug3 fix: TTS 合成后校验音频时长,超限则加速/截断。
|
||||
|
||||
- 音频时长 > max(target+2, 30) → 加速到目标时长
|
||||
- 音频时长 > 30s(Seedance 硬限制)→ 必须加速/截断到 30s 以内
|
||||
返回处理后的路径(可能覆盖原文件),失败返回原路径不阻塞。
|
||||
"""
|
||||
if tts_path is None:
|
||||
return None
|
||||
try:
|
||||
import subprocess
|
||||
from pathlib import Path as _Path
|
||||
|
||||
local = _Path(tts_path) if not isinstance(tts_path, _Path) else tts_path
|
||||
if not local.exists():
|
||||
return tts_path
|
||||
|
||||
# 用 ffprobe 检查音频时长
|
||||
probe_cmd = [
|
||||
"ffprobe",
|
||||
"-v",
|
||||
"quiet",
|
||||
"-show_entries",
|
||||
"format=duration",
|
||||
"-of",
|
||||
"default=noprint_wrappers=1:nokey=1",
|
||||
str(local),
|
||||
]
|
||||
result = subprocess.run(probe_cmd, capture_output=True, text=True, timeout=15)
|
||||
if result.returncode != 0:
|
||||
logger.warning("[爆款视频] ffprobe 检查音频时长失败: %s", result.stderr[:200])
|
||||
return tts_path
|
||||
|
||||
audio_dur = float(result.stdout.strip())
|
||||
hard_limit = 30.0 # Seedance 硬限制
|
||||
soft_limit = float(target_duration) + 2.0
|
||||
effective_limit = min(soft_limit, hard_limit)
|
||||
|
||||
logger.info(
|
||||
"[爆款视频] TTS 音频时长检查: audio=%.1fs target=%ds limit=%.1fs",
|
||||
audio_dur,
|
||||
target_duration,
|
||||
effective_limit,
|
||||
)
|
||||
|
||||
if audio_dur <= effective_limit:
|
||||
return tts_path # 时长合理
|
||||
|
||||
# 需要处理:优先用 ffmpeg 加速(atempo)保持内容完整
|
||||
if audio_dur > hard_limit:
|
||||
target_sec = 29.0 # 必须压缩到 30s 以内
|
||||
else:
|
||||
target_sec = float(target_duration)
|
||||
|
||||
speed_factor = audio_dur / target_sec
|
||||
if speed_factor > 2.0:
|
||||
# atempo 最大 2.0x,超过则先加速到 2x 再截断
|
||||
logger.warning(
|
||||
"[爆款视频] TTS 音频加速比 %.2f 超过 2x,改为加速+截断",
|
||||
speed_factor,
|
||||
)
|
||||
accel_path = local.with_suffix(".accel.mp3")
|
||||
subprocess.run(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
str(local),
|
||||
"-filter:a",
|
||||
"atempo=2.0",
|
||||
"-c:a",
|
||||
"libmp3lame",
|
||||
"-q:a",
|
||||
"2",
|
||||
str(accel_path),
|
||||
],
|
||||
capture_output=True,
|
||||
timeout=60,
|
||||
)
|
||||
if accel_path.exists() and accel_path.stat().st_size > 0:
|
||||
accel_path.replace(local)
|
||||
# 再截断到目标时长
|
||||
trimmed = local.with_suffix(".trimmed.mp3")
|
||||
subprocess.run(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
str(local),
|
||||
"-t",
|
||||
str(target_sec),
|
||||
"-c:a",
|
||||
"libmp3lame",
|
||||
"-q:a",
|
||||
"2",
|
||||
str(trimmed),
|
||||
],
|
||||
capture_output=True,
|
||||
timeout=30,
|
||||
)
|
||||
if trimmed.exists() and trimmed.stat().st_size > 0:
|
||||
trimmed.replace(local)
|
||||
logger.info("[爆款视频] TTS 音频加速+截断完成: %.1fs → %.1fs", audio_dur, target_sec)
|
||||
else:
|
||||
# 用 atempo 加速
|
||||
accelerated = local.with_suffix(".accel.mp3")
|
||||
subprocess.run(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-i",
|
||||
str(local),
|
||||
"-filter:a",
|
||||
f"atempo={speed_factor:.4f}",
|
||||
"-c:a",
|
||||
"libmp3lame",
|
||||
"-q:a",
|
||||
"2",
|
||||
str(accelerated),
|
||||
],
|
||||
capture_output=True,
|
||||
timeout=60,
|
||||
)
|
||||
if accelerated.exists() and accelerated.stat().st_size > 0:
|
||||
accelerated.replace(local)
|
||||
logger.info(
|
||||
"[爆款视频] TTS 音频加速完成: %.1fs → %.1fs (speed=%.2fx)",
|
||||
audio_dur,
|
||||
target_sec,
|
||||
speed_factor,
|
||||
)
|
||||
|
||||
return tts_path
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] TTS 音频时长校验异常: %s", e, exc_info=True)
|
||||
return tts_path # 校验异常不阻塞
|
||||
|
||||
|
||||
def _get_expected_shot_count(duration: int) -> tuple[int, int] | None:
|
||||
"""根据视频时长返回期望的镜头数量范围 (min, max)。"""
|
||||
if duration <= 5:
|
||||
return (1, 2)
|
||||
elif duration <= 10:
|
||||
return (3, 3)
|
||||
elif duration <= 15:
|
||||
return (3, 4)
|
||||
elif duration <= 20:
|
||||
return (4, 5)
|
||||
elif duration <= 30:
|
||||
return (6, 8)
|
||||
return None
|
||||
|
||||
|
||||
def _validate_shot_timeline(shots: list[dict], total_duration: int) -> bool:
|
||||
"""校验镜头时间轴是否合法:
|
||||
- 每个 shot 的 time_range 必须能解析为 "X-Y秒"
|
||||
- 第一个镜头必须从 0 开始
|
||||
- 最后一个镜头必须结束于 total_duration
|
||||
- 相邻镜头首尾相接(允许1秒误差)
|
||||
"""
|
||||
if not shots:
|
||||
return False
|
||||
|
||||
prev_end = 0
|
||||
for i, shot in enumerate(shots):
|
||||
tr = str(shot.get("time_range", ""))
|
||||
# 解析 "X-Y秒" 格式
|
||||
match = re.match(r"(\d+)-(\d+)秒?", tr)
|
||||
if not match:
|
||||
return False
|
||||
start = int(match.group(1))
|
||||
end = int(match.group(2))
|
||||
|
||||
# 第一个镜头必须从 0 开始(允许1秒误差)
|
||||
if i == 0 and start > 1:
|
||||
return False
|
||||
|
||||
# 时间必须递增
|
||||
if end <= start:
|
||||
return False
|
||||
|
||||
# 与前一镜头衔接(允许1秒误差)
|
||||
if abs(start - prev_end) > 1:
|
||||
return False
|
||||
|
||||
prev_end = end
|
||||
|
||||
# 最后一个镜头必须结束于 total_duration(允许1秒误差)
|
||||
if abs(prev_end - total_duration) > 1:
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def _parse_shot_seconds(time_range) -> tuple[int, int] | None:
|
||||
"""从 time_range 解析 (start, end),无法解析返回 None。"""
|
||||
m = re.match(r"(\d+)-(\d+)秒?", str(time_range or ""))
|
||||
if not m:
|
||||
return None
|
||||
start, end = int(m.group(1)), int(m.group(2))
|
||||
if end <= start:
|
||||
return None
|
||||
return start, end
|
||||
|
||||
|
||||
def _redistribute_timeline(shots: list[dict], total_duration: int) -> list[dict]:
|
||||
"""服务端强制按比例重新分配时间轴:
|
||||
保留各镜头内容,按 2 秒最小粒度分配,余数补到最后一个镜头。
|
||||
"""
|
||||
n = len(shots)
|
||||
if n == 0:
|
||||
return shots
|
||||
# 每镜头基础秒数
|
||||
base = max(2, total_duration // n)
|
||||
boundaries: list[int] = [0]
|
||||
for i in range(n):
|
||||
if i == n - 1:
|
||||
boundaries.append(total_duration)
|
||||
else:
|
||||
boundaries.append(min(total_duration - (n - 1 - i) * 2, boundaries[-1] + base))
|
||||
for i, s in enumerate(shots):
|
||||
s["time_range"] = f"{boundaries[i]}-{boundaries[i + 1]}秒"
|
||||
logger.info("[爆款视频] 服务端强制重分配时间轴: %s total=%ds", boundaries, total_duration)
|
||||
return shots
|
||||
|
||||
|
||||
def _truncate_voiceover(text: str, max_chars: int) -> str:
|
||||
"""口播超长兜底:按句号/问号/感叹号切句,按字数保留前面的完整句子。"""
|
||||
text = text.strip()
|
||||
if len(text) <= max_chars:
|
||||
return text
|
||||
parts = re.split(r"(?<=[。!?!?\.])", text)
|
||||
out = ""
|
||||
for p in parts:
|
||||
if not p:
|
||||
continue
|
||||
if len(out) + len(p) > max_chars:
|
||||
break
|
||||
out += p
|
||||
if not out:
|
||||
out = text[:max_chars].rstrip(",、 ")
|
||||
return out
|
||||
|
||||
|
||||
def _shot_visual(s: dict) -> str:
|
||||
"""合并镜头画面与动作描述。"""
|
||||
visual = str(s.get("scene_and_dialogue", "") or s.get("visual", "")).strip()
|
||||
act = str(s.get("action_details", "") or "").strip()
|
||||
if act and act not in visual:
|
||||
visual = f"{visual},{act}" if visual else act
|
||||
return visual
|
||||
|
||||
|
||||
def _finalize_fallback_script(job: ViralVideoJob) -> dict:
|
||||
"""所有 LLM 重试失败后的最终兜底:
|
||||
- 时间轴不合法 → 服务端强制按比例重分配
|
||||
- 口播超长 → 按句子截断;口播过短不处理(危害较小)
|
||||
"""
|
||||
fb = _fallback_script(job)
|
||||
dur = max(5, min(30, int(getattr(job, "duration", 15) or 15)))
|
||||
shots = fb.get("shots") or []
|
||||
if not _validate_shot_timeline(shots, dur):
|
||||
fb["shots"] = _redistribute_timeline(shots, dur)
|
||||
voice = fb.get("voiceover_script", "") or ""
|
||||
max_chars = dur * 3
|
||||
if len(voice) > max_chars:
|
||||
voice = _truncate_voiceover(voice, max_chars)
|
||||
fb["voiceover_script"] = voice
|
||||
fb["final_copy"] = voice
|
||||
fb["suggested_copy"] = voice
|
||||
fb["copy_display_markdown"] = voice
|
||||
logger.info(
|
||||
"[爆款视频] 最终兜底脚本完成 dur=%ds shots=%d voice_chars=%d", dur, len(fb.get("shots") or []), len(voice)
|
||||
)
|
||||
return fb
|
||||
|
||||
|
||||
def _assemble_seedance_prompt(copy_result: dict, job: ViralVideoJob) -> str:
|
||||
"""把编导脚本拼成 Seedance 长 prompt。"""
|
||||
"""把编导脚本拼成符合官方推荐格式的 prompt(按 provider 分发)。
|
||||
|
||||
- doubao (Seedance): [X-Y秒] 单行时间戳 + 镜头级 @图片N 绑定
|
||||
- dashscope (Wan 3.0): 第N个镜头[X-Y秒] 单行格式
|
||||
"""
|
||||
from packages.domain.points_rules import get_viral_video_model_config
|
||||
|
||||
if not isinstance(copy_result, dict) or not copy_result:
|
||||
return "产品展示短视频,清晰明亮,自然讲解"
|
||||
ov = copy_result.get("overview") or {}
|
||||
theme = ov.get("theme", "")
|
||||
total_duration = ov.get("total_duration") or getattr(job, "duration", 15)
|
||||
aspect_ratio = ov.get("aspect_ratio") or getattr(job, "video_ratio", "9:16")
|
||||
scene_lighting = copy_result.get("scene_and_lighting", "")
|
||||
shots = copy_result.get("shots") or []
|
||||
hc = copy_result.get("hard_constraints") or _DEFAULT_HARD_CONSTRAINTS
|
||||
np = copy_result.get("negative_prompts") or _DEFAULT_NEGATIVE_PROMPTS
|
||||
|
||||
ov = copy_result.get("overview") or {}
|
||||
theme = ov.get("theme", "") or ""
|
||||
total_duration = int(ov.get("total_duration") or getattr(job, "duration", 15) or 15)
|
||||
aspect_ratio = ov.get("aspect_ratio") or getattr(job, "video_ratio", "9:16")
|
||||
scene_lighting = str(copy_result.get("scene_and_lighting", "") or "")
|
||||
shots = [s for s in (copy_result.get("shots") or []) if isinstance(s, dict)]
|
||||
hc = copy_result.get("hard_constraints") or _DEFAULT_HARD_CONSTRAINTS
|
||||
np_list = copy_result.get("negative_prompts") or _DEFAULT_NEGATIVE_PROMPTS
|
||||
images = list(job.images or [])
|
||||
|
||||
model = getattr(job, "video_model", "") or None
|
||||
provider = get_viral_video_model_config(model).get("provider", "doubao")
|
||||
|
||||
def _ref(i: int):
|
||||
"""该镜头绑定的图片序号(越界/-1 回退首图)。"""
|
||||
idx = shots[i].get("reference_image_index") if i < len(shots) else None
|
||||
if isinstance(idx, int) and 0 <= idx < len(images):
|
||||
return idx
|
||||
return 0 if images else None
|
||||
|
||||
if provider == "dashscope":
|
||||
# ========== Wan 3.0 格式 ==========
|
||||
head = f"{theme},{scene_lighting or '画面清晰有质感'},总时长{total_duration}秒。"
|
||||
body: list[str] = [head, ""]
|
||||
for i, s in enumerate(shots):
|
||||
tr = str(s.get("time_range", ""))
|
||||
pr = _parse_shot_seconds(tr)
|
||||
label = f"{pr[0]}-{pr[1]}秒" if pr else tr
|
||||
cam = str(s.get("shot_type_angle_movement", "") or "").strip()
|
||||
voice = str(s.get("voiceover", "") or "").strip()
|
||||
pieces = [f"第{i + 1}个镜头[{label}]"]
|
||||
if cam:
|
||||
pieces.append(f"运镜:{cam}")
|
||||
pieces.append(f"画面:{_shot_visual(s)}")
|
||||
if voice:
|
||||
pieces.append(f'配音:"{voice}"')
|
||||
# Wan 通过 input.media 数组顺序隐式引用图1/图2,不写@图片
|
||||
body.append(" ".join(pieces))
|
||||
style_parts = [str(x) for x in hc if x] + [str(x) for x in np_list if x]
|
||||
if style_parts:
|
||||
body.append("")
|
||||
body.append("风格说明:" + ";".join(style_parts[:6]))
|
||||
return "\n".join(body)
|
||||
|
||||
# ========== Seedance 格式 ==========
|
||||
lines: list[str] = []
|
||||
lines.append("【视频总览】")
|
||||
lines.append(f"- 整体主题:{theme}")
|
||||
lines.append(f"- 总时长:{total_duration}秒(单次生成,时长必须严格匹配)")
|
||||
lines.append(f"- 主题:{theme}")
|
||||
lines.append(f"- 总时长:{total_duration}秒(单次生成,时长严格匹配)")
|
||||
lines.append(f"- 画幅:{aspect_ratio}")
|
||||
lines.append(f"- 整体风格:{scene_lighting or '专业、清晰、明亮有质感'}")
|
||||
lines.append("")
|
||||
lines.append("【场景与光线】")
|
||||
lines.append(scene_lighting)
|
||||
lines.append("")
|
||||
lines.append("【逐镜头时间轴】(按时间顺序连贯拍摄,镜头之间自然衔接)")
|
||||
|
||||
# 【参考素材】按镜头绑定
|
||||
if images:
|
||||
lines.append("【参考素材】")
|
||||
for i in range(len(shots)):
|
||||
ri = _ref(i)
|
||||
if ri is not None:
|
||||
lines.append(f"镜头{i + 1}参考@图片{ri + 1}")
|
||||
lines.append("")
|
||||
|
||||
# 单行时间戳分镜
|
||||
lines.append("【分镜脚本】")
|
||||
for i, s in enumerate(shots):
|
||||
if not isinstance(s, dict):
|
||||
continue
|
||||
tr = s.get("time_range", "")
|
||||
cam = s.get("shot_type_angle_movement", "")
|
||||
sd = s.get("scene_and_dialogue", "")
|
||||
act = s.get("action_details", "")
|
||||
ab = s.get("audio_bgm", "")
|
||||
t = s.get("transition", "")
|
||||
ref = s.get("reference_image_index")
|
||||
lines.append(f"- 镜头{i + 1}({tr}):")
|
||||
lines.append(f" 景别/运镜:{cam}")
|
||||
lines.append(f" 画面与对白:{sd}")
|
||||
lines.append(f" 动作细节:{act}")
|
||||
lines.append(f" 音效/BGM:{ab}")
|
||||
lines.append(f" 转场:{t}")
|
||||
if ref is not None and isinstance(ref, int):
|
||||
lines.append(f" 参考图片:第{ref + 1}张产品图")
|
||||
tr = str(s.get("time_range", ""))
|
||||
pr = _parse_shot_seconds(tr)
|
||||
label = f"{pr[0]}-{pr[1]}秒" if pr else tr
|
||||
cam = str(s.get("shot_type_angle_movement", "") or "").strip()
|
||||
voice = str(s.get("voiceover", "") or "").strip()
|
||||
trans = str(s.get("transition", "") or "").strip()
|
||||
pieces = [f"[{label}]"]
|
||||
if cam:
|
||||
pieces.append(f"景别/运镜:{cam}")
|
||||
pieces.append(f"画面:{_shot_visual(s)}")
|
||||
if voice:
|
||||
pieces.append(f'口播:"{voice}"')
|
||||
ri = _ref(i)
|
||||
if ri is not None:
|
||||
pieces.append(f"参考@图片{ri + 1}")
|
||||
if trans:
|
||||
pieces.append(f"转场:{trans}")
|
||||
lines.append(" ".join(pieces))
|
||||
|
||||
lines.append("")
|
||||
lines.append("【硬性约束】")
|
||||
for c in hc:
|
||||
lines.append(f"- {c}")
|
||||
lines.append("")
|
||||
lines.append("【负面提示词】(必须避免)")
|
||||
lines.append(",".join([str(x) for x in np if x]))
|
||||
lines.append("【负面提示词】")
|
||||
lines.append(",".join([str(x) for x in np_list if x]))
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
@@ -1584,10 +2187,35 @@ def run_viral_video_generate_copy(self: Task, job_id: str) -> dict:
|
||||
_save_job(repo, job, session)
|
||||
_hb_stop, _hb_thread = _start_heartbeat_thread(job_id)
|
||||
|
||||
# v8:意图理解并入分镜生成,一次 LLM 调用
|
||||
# v8:意图理解并入分镜生成,一次 LLM 调用;流式推送 script_delta 给前端
|
||||
image_analysis = normalize_image_analysis(job.image_analysis)
|
||||
_set_stage(job, repo, session, ViralVideoStage.SCRIPT_GENERATION, "正在编排分镜脚本...")
|
||||
copy_result = _step_script_generation(job, image_analysis)
|
||||
|
||||
_script_full_text = []
|
||||
_script_last_emit_ts = [0.0]
|
||||
_script_stage_start = time.time()
|
||||
|
||||
def _on_script_delta(delta: str, full_text: str):
|
||||
"""流式回调:推 viral_video:script_delta 事件到 Redis pub/sub(WS 桥接前端)。"""
|
||||
_script_full_text.append(delta) if delta else None
|
||||
now = time.time()
|
||||
# 速率保护:on_delta 已经做了基础节流;这里再加一道 200ms 兜底防止消息风暴
|
||||
if now - _script_last_emit_ts[0] < 0.2:
|
||||
return
|
||||
_script_last_emit_ts[0] = now
|
||||
# 进度估算:基于 full_text 长度线性增长(最长 ~3500 字 = 70% 进度位)
|
||||
_est_progress = min(70.0, 15.0 + (len(full_text) / 3500.0) * 55.0)
|
||||
_elapsed = now - _script_stage_start
|
||||
_emit_progress(
|
||||
job_id,
|
||||
ViralVideoStage.SCRIPT_GENERATION,
|
||||
round(_est_progress, 1),
|
||||
f"正在编排分镜脚本...({len(full_text)}字,{_elapsed:.0f}s)",
|
||||
{"delta": delta, "full_text": full_text, "text_length": len(full_text)},
|
||||
event_type="viral_video:script_delta",
|
||||
)
|
||||
|
||||
copy_result = _step_script_generation(job, image_analysis, on_delta=_on_script_delta)
|
||||
_emit_progress(
|
||||
job_id,
|
||||
ViralVideoStage.SCRIPT_GENERATION,
|
||||
@@ -1859,11 +2487,29 @@ def _run_render_pipeline(job_id: str, session, repo, job) -> dict:
|
||||
|
||||
voiceover = copy_result.get("voiceover_script", "") or job.effective_copy_text
|
||||
|
||||
# Step 5: TTS 整段合成
|
||||
_set_stage(job, repo, session, ViralVideoStage.TTS, "正在合成AI配音...")
|
||||
tts_path = _step_tts(job, voiceover)
|
||||
tts_url = _upload_tts_to_oss(job, tts_path)
|
||||
_emit_progress(job_id, ViralVideoStage.TTS, 78.0, "配音完成", {"has_tts": tts_url is not None})
|
||||
# provider 判断:Wan 3.0(dashscope) 在无自定义音色时跳过 TTS,用模型原生音频
|
||||
from packages.domain.points_rules import get_viral_video_model_config
|
||||
|
||||
_vv_model = getattr(job, "video_model", "") or None
|
||||
_vv_provider = get_viral_video_model_config(_vv_model).get("provider", "doubao")
|
||||
_has_custom_voice = bool((getattr(job, "voice_id", "") or "").strip())
|
||||
_skip_tts = _vv_provider == "dashscope" and not _has_custom_voice
|
||||
|
||||
tts_url = None
|
||||
if _skip_tts:
|
||||
logger.info("[爆款视频][阶段3] Wan原生音频模式:跳过TTS,由模型生成配音/BGM/音效 job_id=%s", job_id)
|
||||
_set_stage(job, repo, session, ViralVideoStage.TTS, "使用模型原生音频...")
|
||||
_emit_progress(job_id, ViralVideoStage.TTS, 78.0, "使用模型原生音频", {"has_tts": False, "native_audio": True})
|
||||
else:
|
||||
# Step 5: TTS 整段合成
|
||||
_set_stage(job, repo, session, ViralVideoStage.TTS, "正在合成AI配音...")
|
||||
tts_path = _step_tts(job, voiceover)
|
||||
# Bug3 fix: TTS 音频时长校验与修正
|
||||
if tts_path is not None:
|
||||
_dur_limit = max(5, min(30, int(getattr(job, "duration", 15) or 15)))
|
||||
tts_path = _check_and_fix_tts_duration(tts_path, _dur_limit)
|
||||
tts_url = _upload_tts_to_oss(job, tts_path)
|
||||
_emit_progress(job_id, ViralVideoStage.TTS, 78.0, "配音完成", {"has_tts": tts_url is not None})
|
||||
|
||||
# Step 6: 单次 Seedance(失败自动退款)
|
||||
_set_stage(job, repo, session, ViralVideoStage.RENDERING, "正在生成视频(约1-3分钟)...")
|
||||
@@ -1943,7 +2589,18 @@ def run_viral_video_render(self: Task, job_id: str) -> dict:
|
||||
pass
|
||||
except Exception:
|
||||
logger.exception("[爆款视频][阶段3] 兜底退款异常")
|
||||
_mark_failed_and_notify(job_id, session, None, None, str(e), ViralVideoStage.RENDERING)
|
||||
# Bug4 fix: 使用 job 当前实际阶段而非硬编码 RENDERING
|
||||
_err_stage = ViralVideoStage.RENDERING
|
||||
_err_job = None
|
||||
if session is not None:
|
||||
try:
|
||||
_repo_tmp = SQLAlchemyViralVideoJobRepository(session)
|
||||
_err_job = _repo_tmp.get(job_id)
|
||||
if _err_job is not None and getattr(_err_job, "current_stage", None):
|
||||
_err_stage = _err_job.current_stage
|
||||
except Exception:
|
||||
pass
|
||||
_mark_failed_and_notify(job_id, session, None, _err_job, str(e), _err_stage)
|
||||
return {"ok": False, "job_id": job_id, "error": str(e)}
|
||||
finally:
|
||||
if _hb_stop is not None:
|
||||
|
||||
@@ -311,3 +311,10 @@ GPU_ENCODE_CRF=23
|
||||
GPU_ENCODE_FALLBACK_CPU=true
|
||||
GPU_ENCODE_MEZZANINE_TRANSPORT=oss
|
||||
GPU_ENCODE_OSS_TMP_PREFIX=tmp/gpu-mezzanine/
|
||||
|
||||
# ==================== Ditto 蚂蚁数字人口型 ====================
|
||||
# 注意:这些值必须写死在模板里(不是 CI Secret),否则每次 CI 重新渲染 .env 都会被丢弃,
|
||||
# 导致 staging 发版后 Ditto 口型服务静默降级到 GPU/MediaKit(P0 防复发)。
|
||||
USE_DITTO_LIPSYNC=true
|
||||
DITTO_API_BASE_URL=http://100.76.80.23:8000
|
||||
DITTO_DEFAULT_VIDEO_URL=https://xiaoxia-autocut.oss-cn-hangzhou.aliyuncs.com/uploads/default_avatar.mp4
|
||||
|
||||
@@ -40,7 +40,9 @@ server {
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
proxy_read_timeout 300s;
|
||||
proxy_set_header Upgrade $http_upgrade;
|
||||
proxy_set_header Connection "upgrade";
|
||||
proxy_read_timeout 3600s;
|
||||
proxy_send_timeout 300s;
|
||||
proxy_request_buffering off;
|
||||
}
|
||||
|
||||
@@ -39,7 +39,9 @@ server {
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
proxy_read_timeout 300s;
|
||||
proxy_set_header Upgrade $http_upgrade;
|
||||
proxy_set_header Connection "upgrade";
|
||||
proxy_read_timeout 3600s;
|
||||
proxy_send_timeout 300s;
|
||||
proxy_request_buffering off;
|
||||
}
|
||||
|
||||
@@ -961,6 +961,7 @@ class ViralVideoJobModel(Base):
|
||||
# v1.5 音频/视频参数
|
||||
voice_id = Column(String(200), nullable=False, default="")
|
||||
voice_source = Column(String(20), nullable=False, default="")
|
||||
language = Column(String(20), nullable=False, default="zh-CN")
|
||||
video_ratio = Column(String(10), nullable=False, default="9:16")
|
||||
video_model = Column(String(100), nullable=False, default="")
|
||||
# 结果与状态
|
||||
@@ -1019,3 +1020,20 @@ class ViralVideoPromptTemplateModel(Base):
|
||||
is_active = Column(Boolean, nullable=False, default=True)
|
||||
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
class SystemSettingModel(Base):
|
||||
"""系统配置表 ORM 模型(#2246:后台可配置项,表已手工存在于 staging)."""
|
||||
|
||||
__tablename__ = "system_settings"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
setting_key = Column(String(100), nullable=False, unique=True)
|
||||
setting_value = Column(Text, nullable=True)
|
||||
setting_type = Column(String(20), nullable=False)
|
||||
description = Column(String(255), nullable=False, default="", server_default="")
|
||||
is_public = Column(Boolean, nullable=False, default=False, server_default="false")
|
||||
updated_by = Column(String(36), nullable=True)
|
||||
category = Column(String(50), nullable=False, default="general", server_default="general")
|
||||
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
"""system_settings 表 SQLAlchemy Repository — #2246."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import SystemSettingModel
|
||||
from packages.domain.system_setting import SystemSetting
|
||||
|
||||
|
||||
class SQLAlchemySystemSettingRepository:
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
|
||||
def get_by_key(self, setting_key: str) -> SystemSetting | None:
|
||||
model = self.session.query(SystemSettingModel).filter(SystemSettingModel.setting_key == setting_key).first()
|
||||
if model is None:
|
||||
return None
|
||||
return self._to_domain(model)
|
||||
|
||||
def list_all(self, category: str | None = None) -> list[SystemSetting]:
|
||||
query = self.session.query(SystemSettingModel)
|
||||
if category is not None:
|
||||
query = query.filter(SystemSettingModel.category == category)
|
||||
return [self._to_domain(m) for m in query.all()]
|
||||
|
||||
def upsert(self, setting: SystemSetting) -> SystemSetting:
|
||||
model = (
|
||||
self.session.query(SystemSettingModel).filter(SystemSettingModel.setting_key == setting.setting_key).first()
|
||||
)
|
||||
if model is None:
|
||||
model = SystemSettingModel(id=setting.id)
|
||||
self.session.add(model)
|
||||
model.setting_key = setting.setting_key
|
||||
model.setting_value = setting.setting_value
|
||||
model.setting_type = setting.setting_type
|
||||
model.description = setting.description
|
||||
model.is_public = setting.is_public
|
||||
model.category = setting.category
|
||||
model.updated_by = setting.updated_by
|
||||
self.session.commit()
|
||||
return setting
|
||||
|
||||
def delete_by_key(self, setting_key: str) -> bool:
|
||||
model = self.session.query(SystemSettingModel).filter(SystemSettingModel.setting_key == setting_key).first()
|
||||
if model is None:
|
||||
return False
|
||||
self.session.delete(model)
|
||||
self.session.commit()
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _to_domain(model: SystemSettingModel) -> SystemSetting:
|
||||
return SystemSetting(
|
||||
id=model.id,
|
||||
setting_key=model.setting_key,
|
||||
setting_value=model.setting_value,
|
||||
setting_type=model.setting_type,
|
||||
description=model.description or "",
|
||||
is_public=bool(model.is_public),
|
||||
category=model.category or "general",
|
||||
updated_by=model.updated_by,
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
@@ -49,6 +49,7 @@ def _to_domain(model: ViralVideoJobModel) -> ViralVideoJob:
|
||||
style_template_id=getattr(model, "style_template_id", "") or "",
|
||||
voice_id=getattr(model, "voice_id", "") or "",
|
||||
voice_source=getattr(model, "voice_source", "") or "",
|
||||
language=getattr(model, "language", "zh-CN") or "zh-CN",
|
||||
video_ratio=getattr(model, "video_ratio", "9:16") or "9:16",
|
||||
video_model=getattr(model, "video_model", "") or "",
|
||||
status=ViralVideoStatus(model.status) if model.status else ViralVideoStatus.PENDING,
|
||||
@@ -102,6 +103,7 @@ class SQLAlchemyViralVideoJobRepository:
|
||||
style_template_id=job.style_template_id,
|
||||
voice_id=job.voice_id,
|
||||
voice_source=job.voice_source,
|
||||
language=getattr(job, "language", "zh-CN") or "zh-CN",
|
||||
video_ratio=job.video_ratio,
|
||||
video_model=job.video_model,
|
||||
status=job.status,
|
||||
@@ -169,6 +171,7 @@ class SQLAlchemyViralVideoJobRepository:
|
||||
model.style_template_id = job.style_template_id
|
||||
model.voice_id = job.voice_id or ""
|
||||
model.voice_source = job.voice_source or ""
|
||||
model.language = getattr(job, "language", "zh-CN") or "zh-CN"
|
||||
model.video_ratio = job.video_ratio or "9:16"
|
||||
model.video_model = job.video_model or ""
|
||||
model.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
@@ -3,10 +3,24 @@
|
||||
使用 bcrypt 安全存储密码
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
from typing import Optional
|
||||
|
||||
import bcrypt
|
||||
|
||||
# bcrypt 只对前 72 字节有效,且 bcrypt>=4.1 会对超长输入直接抛 ValueError。
|
||||
# 超长密码先做一次 SHA-256(定长 hex),再交给 bcrypt,
|
||||
# 既绕过长度限制又保持对超长不同密码的区分度。
|
||||
_BCRYPT_MAX_BYTES = 72
|
||||
|
||||
|
||||
def _prepare_password_bytes(password: str) -> bytes:
|
||||
raw = password.encode("utf-8")
|
||||
if len(raw) > _BCRYPT_MAX_BYTES:
|
||||
return hashlib.sha256(raw).hexdigest().encode("utf-8")
|
||||
return raw
|
||||
|
||||
|
||||
from packages.domain.auth.password_hasher import PasswordHasherPort, PasswordValidatorPort
|
||||
|
||||
|
||||
@@ -42,8 +56,8 @@ class PasswordHasher(PasswordHasherPort):
|
||||
if not password:
|
||||
raise ValueError("Password cannot be empty")
|
||||
|
||||
# bcrypt 需要 bytes
|
||||
password_bytes = password.encode("utf-8")
|
||||
# bcrypt 需要 bytes(超长密码先 SHA-256 以兼容 72 字节限制)
|
||||
password_bytes = _prepare_password_bytes(password)
|
||||
|
||||
# 生成 salt 并哈希
|
||||
salt = bcrypt.gensalt(rounds=self.rounds)
|
||||
@@ -67,7 +81,7 @@ class PasswordHasher(PasswordHasherPort):
|
||||
return False
|
||||
|
||||
try:
|
||||
password_bytes = password.encode("utf-8")
|
||||
password_bytes = _prepare_password_bytes(password)
|
||||
hashed_bytes = hashed_password.encode("utf-8")
|
||||
|
||||
return bcrypt.checkpw(password_bytes, hashed_bytes)
|
||||
|
||||
@@ -0,0 +1,332 @@
|
||||
"""Ditto LLM 情绪分析服务 — #2076 后续:根据文案生成 emo_timeline.
|
||||
|
||||
职责:
|
||||
1. 正则按 。!?; 初步分句
|
||||
2. 调 DoubaoClient.chat_completion 分析每句表情(emo: 0-7, intensity: 0-1)
|
||||
3. 结果 LRU 缓存(文案 hash → 情绪列表)
|
||||
4. LLM 失败/超时/格式错 → 返回空列表(降级中性表情,不阻塞生成)
|
||||
5. TTS 完成后按字数比例或 sentence_timings 对齐成秒级 timeline
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from functools import lru_cache
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── 表情常量 ─────────────────────────────────────────────────────
|
||||
EMO_ANGER = 0
|
||||
EMO_DISGUST = 1
|
||||
EMO_FEAR = 2
|
||||
EMO_HAPPY = 3
|
||||
EMO_NEUTRAL = 4
|
||||
EMO_SAD = 5
|
||||
EMO_SURPRISE = 6
|
||||
EMO_CONTEMPT = 7
|
||||
ALLOWED_EMOS = {EMO_HAPPY, EMO_NEUTRAL, EMO_SAD, EMO_SURPRISE} # 营销场景白名单
|
||||
|
||||
# ── 分句正则 ─────────────────────────────────────────────────────
|
||||
_SENT_SPLIT_RE = re.compile(r"(?<=[。!?;!?;])\s*")
|
||||
|
||||
# ── 默认 prompt 模板文件路径 ──────────────────────────────────────
|
||||
_DEFAULT_PROMPT_PATH = Path(__file__).parent / "prompts" / "ditto_emotion.txt"
|
||||
|
||||
|
||||
def _load_default_prompt() -> str:
|
||||
try:
|
||||
return _DEFAULT_PROMPT_PATH.read_text(encoding="utf-8").strip()
|
||||
except Exception:
|
||||
# 文件不存在时用极简兜底
|
||||
return (
|
||||
"分析文案每句话表情,输出JSON数组:"
|
||||
'[{"text":"句子","emo":4,"intensity":0.2}],emo:3开心4中性5伤心6惊讶,'
|
||||
"禁止0/1/2/7。\n【文案】\n{文案}"
|
||||
)
|
||||
|
||||
|
||||
# ── 数据结构 ─────────────────────────────────────────────────────
|
||||
class EmotionSegment:
|
||||
"""单句情绪结果(LLM 输出的原始结构)."""
|
||||
|
||||
__slots__ = ("text", "emo", "intensity")
|
||||
|
||||
def __init__(self, text: str, emo: int, intensity: float):
|
||||
self.text = text
|
||||
self.emo = emo
|
||||
self.intensity = intensity
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {"text": self.text, "emo": self.emo, "intensity": self.intensity}
|
||||
|
||||
|
||||
class EmotionTimelineEntry:
|
||||
"""对齐到音频时间轴后的情绪片段(传给 Ditto)."""
|
||||
|
||||
__slots__ = ("start", "end", "emo", "intensity")
|
||||
|
||||
def __init__(self, start: float, end: float, emo: int, intensity: float):
|
||||
self.start = round(start, 2)
|
||||
self.end = round(end, 2)
|
||||
self.emo = emo
|
||||
self.intensity = round(intensity, 2)
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"start": self.start,
|
||||
"end": self.end,
|
||||
"emo": self.emo,
|
||||
"intensity": self.intensity,
|
||||
}
|
||||
|
||||
|
||||
# ── 分句 ─────────────────────────────────────────────────────────
|
||||
def split_sentences(text: str) -> list[str]:
|
||||
"""按中文句末标点切分,过滤空串."""
|
||||
if not text:
|
||||
return []
|
||||
parts = _SENT_SPLIT_RE.split(text.strip())
|
||||
return [p.strip() for p in parts if p and p.strip()]
|
||||
|
||||
|
||||
# ── 解析 LLM 返回的 JSON ─────────────────────────────────────────
|
||||
def _parse_emotion_json(raw: str) -> list[EmotionSegment]:
|
||||
"""解析 LLM 返回,容错处理:
|
||||
- 去掉 markdown 代码块包裹
|
||||
- 只取第一个 JSON 数组
|
||||
- 逐行校验 emo/intensity 合法性,过滤无效项
|
||||
"""
|
||||
if not raw:
|
||||
return []
|
||||
text = raw.strip()
|
||||
# 去掉 ```json ... ``` 包裹
|
||||
if text.startswith("```"):
|
||||
text = re.sub(r"^```(?:json)?\s*", "", text)
|
||||
text = re.sub(r"\s*```$", "", text)
|
||||
# 找第一个 [ 到最后一个 ]
|
||||
lb = text.find("[")
|
||||
rb = text.rfind("]")
|
||||
if lb == -1 or rb == -1 or rb <= lb:
|
||||
return []
|
||||
try:
|
||||
data = json.loads(text[lb : rb + 1])
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return []
|
||||
if not isinstance(data, list):
|
||||
return []
|
||||
|
||||
results: list[EmotionSegment] = []
|
||||
for item in data:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
try:
|
||||
emo = int(item.get("emo", EMO_NEUTRAL))
|
||||
intensity = float(item.get("intensity", 0.2))
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
if emo not in ALLOWED_EMOS:
|
||||
emo = EMO_NEUTRAL
|
||||
intensity = max(0.05, min(1.0, intensity))
|
||||
sent_text = str(item.get("text", "")).strip()
|
||||
if not sent_text:
|
||||
continue
|
||||
results.append(EmotionSegment(text=sent_text, emo=emo, intensity=intensity))
|
||||
return results
|
||||
|
||||
|
||||
# ── 时间对齐(按字数比例)────────────────────────────────────────
|
||||
def align_timeline_by_length(
|
||||
segments: list[EmotionSegment],
|
||||
audio_duration: float,
|
||||
) -> list[EmotionTimelineEntry]:
|
||||
"""按各句字数占总字数比例分配 audio_duration 时长."""
|
||||
if not segments or audio_duration <= 0:
|
||||
return []
|
||||
total_chars = sum(len(s.text) for s in segments)
|
||||
if total_chars <= 0:
|
||||
return []
|
||||
entries: list[EmotionTimelineEntry] = []
|
||||
pos = 0.0
|
||||
for i, seg in enumerate(segments):
|
||||
if i == len(segments) - 1:
|
||||
end = audio_duration # 最后一段到结尾,避免浮点误差
|
||||
else:
|
||||
end = pos + (len(seg.text) / total_chars) * audio_duration
|
||||
if end > pos:
|
||||
entries.append(
|
||||
EmotionTimelineEntry(
|
||||
start=pos,
|
||||
end=end,
|
||||
emo=seg.emo,
|
||||
intensity=seg.intensity,
|
||||
)
|
||||
)
|
||||
pos = end
|
||||
return entries
|
||||
|
||||
|
||||
def align_timeline_by_timings(
|
||||
segments: list[EmotionSegment],
|
||||
sentence_timings: list[dict[str, Any]],
|
||||
audio_duration: float,
|
||||
) -> list[EmotionTimelineEntry]:
|
||||
"""使用 TTS sentence_timings 精确对齐(优先方案).
|
||||
|
||||
sentence_timings 格式:[{"start":0.0,"end":1.2,"text":"句子"}, ...]
|
||||
按句序匹配 segments 和 timings,长度不一致时回退到按字数比例。
|
||||
"""
|
||||
if not sentence_timings or len(sentence_timings) != len(segments):
|
||||
return align_timeline_by_length(segments, audio_duration)
|
||||
entries: list[EmotionTimelineEntry] = []
|
||||
for seg, timing in zip(segments, sentence_timings, strict=False):
|
||||
try:
|
||||
start = float(timing.get("start", 0))
|
||||
end = float(timing.get("end", 0))
|
||||
except (TypeError, ValueError):
|
||||
return align_timeline_by_length(segments, audio_duration)
|
||||
if end <= start:
|
||||
continue
|
||||
entries.append(
|
||||
EmotionTimelineEntry(
|
||||
start=start,
|
||||
end=end,
|
||||
emo=seg.emo,
|
||||
intensity=seg.intensity,
|
||||
)
|
||||
)
|
||||
return entries
|
||||
|
||||
|
||||
# ── LLM 情绪分析服务 ─────────────────────────────────────────────
|
||||
class DittoEmotionService:
|
||||
"""Ditto 情绪分析服务(带 LRU 缓存)."""
|
||||
|
||||
def __init__(self, settings=None):
|
||||
from packages.config import get_api_settings
|
||||
|
||||
self.settings = settings or get_api_settings()
|
||||
self._client = None
|
||||
|
||||
def _cfg(self, key: str) -> Any:
|
||||
"""优先读后台 system_config,未配置则回退到 settings(env 默认)."""
|
||||
try:
|
||||
from packages.application.system_config_service import get_config
|
||||
|
||||
return get_config(key, getattr(self.settings, key, None))
|
||||
except Exception:
|
||||
return getattr(self.settings, key, None)
|
||||
|
||||
@property
|
||||
def enabled(self) -> bool:
|
||||
return bool(self._cfg("ditto_emotion_enabled"))
|
||||
|
||||
def _get_prompt_template(self) -> str:
|
||||
"""优先用配置(环境变量),否则读文件."""
|
||||
cfg_prompt = self._cfg("ditto_emotion_prompt") or ""
|
||||
if cfg_prompt.strip():
|
||||
return cfg_prompt.strip()
|
||||
return _load_default_prompt()
|
||||
|
||||
def _cache_key(self, text: str) -> str:
|
||||
return hashlib.md5(text.strip().encode("utf-8")).hexdigest()
|
||||
|
||||
def _get_llm_client(self):
|
||||
if self._client is None:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
self._client = get_doubao_client()
|
||||
return self._client
|
||||
|
||||
def _call_llm(self, text: str) -> list[EmotionSegment]:
|
||||
"""调 LLM 分析情绪,失败返回空列表."""
|
||||
template = self._get_prompt_template()
|
||||
prompt = template.replace("{文案}", text)
|
||||
messages = [{"role": "user", "content": prompt}]
|
||||
model = self._cfg("ditto_emotion_model") or None
|
||||
temperature = self._cfg("ditto_emotion_temperature")
|
||||
timeout = getattr(self.settings, "ditto_emotion_timeout", 10)
|
||||
max_tokens = getattr(self.settings, "ditto_emotion_max_tokens", 1024)
|
||||
try:
|
||||
client = self._get_llm_client()
|
||||
result = client.chat_completion(
|
||||
messages=messages,
|
||||
model=model,
|
||||
temperature=temperature,
|
||||
max_tokens=max_tokens,
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("[ditto_emotion] LLM 调用异常: %s", exc)
|
||||
return []
|
||||
if not result:
|
||||
return []
|
||||
segments = _parse_emotion_json(result)
|
||||
if not segments:
|
||||
logger.warning("[ditto_emotion] LLM 返回解析失败: %s", result[:200])
|
||||
return segments
|
||||
|
||||
def analyze(self, text: str) -> list[EmotionSegment]:
|
||||
"""分析文案情绪(带缓存),失败返回空列表."""
|
||||
if not self.enabled or not text or not text.strip():
|
||||
return []
|
||||
key = self._cache_key(text)
|
||||
return _cached_analyze(self, key, text)
|
||||
|
||||
def build_timeline(
|
||||
self,
|
||||
text: str,
|
||||
audio_duration: float,
|
||||
sentence_timings: Optional[list[dict[str, Any]]] = None,
|
||||
) -> str:
|
||||
"""完整流程:分句→LLM分析→时间对齐→序列化为JSON字符串.
|
||||
|
||||
返回: JSON 字符串(可直接传 Ditto emo_timeline 参数);空字符串表示降级中性。
|
||||
"""
|
||||
segments = self.analyze(text)
|
||||
if not segments:
|
||||
return ""
|
||||
if sentence_timings:
|
||||
entries = align_timeline_by_timings(segments, sentence_timings, audio_duration)
|
||||
else:
|
||||
entries = align_timeline_by_length(segments, audio_duration)
|
||||
if not entries:
|
||||
return ""
|
||||
return json.dumps([e.to_dict() for e in entries], ensure_ascii=False)
|
||||
|
||||
|
||||
# ── 模块级 LRU 缓存实例 ─────────────────────────────────────────
|
||||
# 每个 service 实例共享缓存(按 cache_key 区分)
|
||||
@lru_cache(maxsize=512)
|
||||
def _cached_analyze(service: DittoEmotionService, cache_key: str, text: str) -> list[EmotionSegment]:
|
||||
"""LRU 缓存包装:cache_key 由文案 hash 生成,maxsize 从配置读."""
|
||||
# 注意:service 参数仅用于传递调用,缓存由 cache_key 驱动
|
||||
segments = service._call_llm(text)
|
||||
# 如果 LLM 返回空(比如分句数量不匹配),尝试直接对预分句结果分析
|
||||
if not segments:
|
||||
pre_splits = split_sentences(text)
|
||||
if len(pre_splits) > 1:
|
||||
# 用预分句结果兜底:全中性低强度
|
||||
segments = [EmotionSegment(text=s, emo=EMO_NEUTRAL, intensity=0.1) for s in pre_splits]
|
||||
return segments
|
||||
|
||||
|
||||
_singleton: Optional[DittoEmotionService] = None
|
||||
|
||||
|
||||
def get_ditto_emotion_service() -> DittoEmotionService:
|
||||
global _singleton
|
||||
if _singleton is None:
|
||||
_singleton = DittoEmotionService()
|
||||
return _singleton
|
||||
|
||||
|
||||
def reset_ditto_emotion_service() -> None:
|
||||
"""#2246:后台配置变更后重置单例并清空 LLM 结果 LRU 缓存."""
|
||||
global _singleton
|
||||
_singleton = None
|
||||
_cached_analyze.cache_clear()
|
||||
@@ -5,15 +5,17 @@
|
||||
- POST /generate 生成口型视频(同步返回 MP4 流)
|
||||
|
||||
关键特性:
|
||||
- 入参:video_url(人物模板视频 URL) + audio_url(TTS 音频 URL) + script(文案原文)
|
||||
- 入参:video_url(人物驱动视频 URL,用户上传优先;未传时用 default_video_url 兜底)
|
||||
+ audio_url(TTS 音频 URL) + script(文案原文)
|
||||
+ emo_timeline(可选,LLM 情绪时间线 JSON 字符串)
|
||||
- 出参:直接返回 video/mp4 字节流(自带音频,无需二次混流)
|
||||
- 429 时指数退避重试(最多 ditto_max_retries 次)
|
||||
- 500/超时视为失败
|
||||
- 500/超时/网络不可达视为失败
|
||||
- 输出 MP4 字节流转存到自家 OSS,返回公网 URL
|
||||
|
||||
注意:
|
||||
- 保留 MuseTalk/GPU 路径不变;本服务作为更高优先级的第三条口型路径
|
||||
- 不传 emotion/表情精细控制,使用默认 emo_global=4(中性)+ use_script_emo=true(关键词驱动表情)
|
||||
- video_url 时长对齐(视频<音频时循环延长)在 Celery 任务层用 ffmpeg 预处理
|
||||
- Ditto 输出自带音视频,不需要 GFPGAN 超分,不需要 ffmpeg 音视频混流
|
||||
"""
|
||||
|
||||
@@ -67,10 +69,20 @@ class DittoClient:
|
||||
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)
|
||||
try:
|
||||
from packages.application.system_config_service import get_config
|
||||
|
||||
self.blend_frames = int(get_config("ditto_blend_frames", s.ditto_blend_frames))
|
||||
except Exception:
|
||||
self.blend_frames = int(s.ditto_blend_frames)
|
||||
|
||||
@property
|
||||
def is_configured(self) -> bool:
|
||||
"""配置是否完整(base_url + 默认模板视频都有值)."""
|
||||
"""配置是否完整(base_url 必填 + 默认模板视频兜底 URL 有值)。
|
||||
|
||||
注意:即使 is_configured=True,实际生成时优先使用用户上传的 video_url;
|
||||
default_video_url 仅作为用户未上传视频时的兜底。
|
||||
"""
|
||||
return bool(self.base_url) and bool(self.default_video_url)
|
||||
|
||||
def health(self) -> bool:
|
||||
@@ -99,7 +111,8 @@ class DittoClient:
|
||||
video_url: Optional[str] = None,
|
||||
emo_global: int = 4,
|
||||
use_script_emo: bool = True,
|
||||
blend_frames: int = 6,
|
||||
blend_frames: Optional[int] = None,
|
||||
emo_timeline: str = "",
|
||||
) -> DittoResult:
|
||||
"""调用 Ditto /generate 接口,返回 MP4 字节流结果.
|
||||
|
||||
@@ -115,21 +128,26 @@ class DittoClient:
|
||||
if not script:
|
||||
script = " "
|
||||
|
||||
_blend = blend_frames if blend_frames is not None else self.blend_frames
|
||||
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,
|
||||
"blend_frames": _blend,
|
||||
}
|
||||
if emo_timeline:
|
||||
payload["emo_timeline"] = emo_timeline
|
||||
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:
|
||||
# 精细化超时:connect=10s(网络不通快速失败),read=120s(最长音频~45s按RTF=2.8推算)
|
||||
_timeout = httpx.Timeout(connect=10.0, read=self.timeout, write=30.0, pool=10.0)
|
||||
with httpx.Client(timeout=_timeout, follow_redirects=True) as client:
|
||||
resp = client.post(url, json=payload)
|
||||
elapsed = time.monotonic() - start
|
||||
|
||||
@@ -205,6 +223,13 @@ class DittoClient:
|
||||
|
||||
except DittoError:
|
||||
raise
|
||||
except (httpx.ConnectError, httpx.NetworkError, ConnectionError, OSError) as exc:
|
||||
# 网络不通/连接被拒(如 GPU 断网/Tailscale 掉线),不重试,直接快速回退
|
||||
logger.warning("[ditto] 网络不可达 attempt=%d err=%s", attempt + 1, exc)
|
||||
raise DittoError(
|
||||
f"Ditto 网络不可达: {exc}",
|
||||
code="NetworkUnreachable",
|
||||
) from exc
|
||||
except httpx.TimeoutException as exc:
|
||||
last_exc = exc
|
||||
logger.warning("[ditto] 请求超时 attempt=%d err=%s", attempt + 1, exc)
|
||||
@@ -212,7 +237,7 @@ class DittoClient:
|
||||
time.sleep(min(2**attempt, 15))
|
||||
continue
|
||||
raise DittoError(
|
||||
f"Ditto 请求超时({self.timeout}s),重试耗尽",
|
||||
f"Ditto 请求超时(read={self.timeout}s),重试耗尽",
|
||||
code="Timeout",
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
@@ -233,9 +258,17 @@ class DittoClient:
|
||||
audio_url: str,
|
||||
script: str,
|
||||
video_url: Optional[str] = None,
|
||||
emo_timeline: str = "",
|
||||
blend_frames: Optional[int] = None,
|
||||
) -> DittoResult:
|
||||
"""调用 generate 并把 MP4 转存到自家 OSS,返回带 video_url 的结果."""
|
||||
result = self.generate(audio_url=audio_url, script=script, video_url=video_url)
|
||||
result = self.generate(
|
||||
audio_url=audio_url,
|
||||
script=script,
|
||||
video_url=video_url,
|
||||
emo_timeline=emo_timeline,
|
||||
blend_frames=blend_frames,
|
||||
)
|
||||
try:
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
@@ -267,3 +300,9 @@ def get_ditto_client() -> DittoClient:
|
||||
if _ditto_client_singleton is None:
|
||||
_ditto_client_singleton = DittoClient()
|
||||
return _ditto_client_singleton
|
||||
|
||||
|
||||
def reset_ditto_client() -> None:
|
||||
"""#2246:后台配置变更后重置 DittoClient 单例."""
|
||||
global _ditto_client_singleton
|
||||
_ditto_client_singleton = None
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
你是一个数字人视频表情导演。给定一段口播文案,分析每句话应该用什么表情和强度,让数字人说话时表情自然有变化,不僵硬。
|
||||
【表情编号】
|
||||
0=愤怒(营销场景禁用)
|
||||
1=厌恶(禁用)
|
||||
2=害怕(禁用)
|
||||
3=开心:介绍优点、优惠、好消息、号召行动时用
|
||||
4=中性:默认表情,陈述事实、平铺直叙时用
|
||||
5=伤心:仅在共情用户痛点时低强度使用(如"是不是经常遇到…")
|
||||
6=惊讶:惊喜、意外、强调价值时用(如"居然""只要""竟然")
|
||||
7=轻蔑(禁用)
|
||||
【强度说明】
|
||||
0.1-0.2:几乎看不出变化,比中性多一点情绪色彩
|
||||
0.3-0.4:有明显但自然的情绪,正常说话的波动
|
||||
0.5-0.6:较强情绪,感叹句/重点强调
|
||||
0.7+:极强情绪,极少使用
|
||||
【规则】
|
||||
1. 按自然语义分句,以。!?;为主要分界,逗号不分
|
||||
2. 60-70%的句子应该用中性(4),不要每句都标情绪
|
||||
3. 情绪和内容匹配:卖点→开心(3),痛点共情→伤心(5)低强度,惊喜/划算→惊讶(6),陈述→中性(4)
|
||||
4. 相邻句子情绪不要剧烈跳变
|
||||
5. 感叹号结尾强度0.4-0.6,句号结尾一般0.1-0.3
|
||||
6. 开头结尾句用中性(4)或低强度开心(3)
|
||||
7. 禁止使用0/1/2/7
|
||||
【输出格式】严格JSON数组,不要输出其他内容
|
||||
[{"text":"句子原文","emo":3,"intensity":0.4}]
|
||||
【文案】
|
||||
{文案}
|
||||
@@ -0,0 +1,189 @@
|
||||
"""系统配置应用服务 — #2246.
|
||||
|
||||
- get_config(key, default) / set_config(...) / list_configs(category)
|
||||
- 首次访问时从 DB 加载并进程内缓存;set_config 后失效缓存
|
||||
- set_config 成功后重置 Ditto 情绪服务与 Ditto 客户端单例(含 LRU 缓存),
|
||||
保证后台改动立即生效;worker 不直接 HTTP 读配置,统一通过本模块。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from typing import Any
|
||||
|
||||
from packages.adapters.sqlalchemy_impl import session as db_session
|
||||
from packages.adapters.sqlalchemy_impl.system_setting_repository import (
|
||||
SQLAlchemySystemSettingRepository,
|
||||
)
|
||||
from packages.domain.system_setting import (
|
||||
SystemSetting,
|
||||
infer_setting_type,
|
||||
serialize_setting_value,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SystemConfigService:
|
||||
"""进程内缓存的系统配置服务."""
|
||||
|
||||
def __init__(self, session_factory=None):
|
||||
self._session_factory = session_factory
|
||||
self._cache: dict[str, Any] | None = None
|
||||
self._lock = threading.Lock()
|
||||
|
||||
# ── 会话 ─────────────────────────────────────────────────────
|
||||
def _get_session_factory(self):
|
||||
# 必须运行时读取模块属性:模块导入时 SessionLocal 还是 None,
|
||||
# initialize_database() 之后才被赋值,import 时绑定会拿到旧值。
|
||||
factory = self._session_factory or db_session.SessionLocal
|
||||
if factory is None:
|
||||
raise RuntimeError("数据库会话工厂未初始化")
|
||||
return factory
|
||||
|
||||
# ── 缓存 ─────────────────────────────────────────────────────
|
||||
def _load_cache(self) -> dict[str, Any]:
|
||||
with self._lock:
|
||||
if self._cache is not None:
|
||||
return self._cache
|
||||
cache: dict[str, Any] = {}
|
||||
factory = self._get_session_factory()
|
||||
session = factory()
|
||||
try:
|
||||
repo = SQLAlchemySystemSettingRepository(session)
|
||||
for setting in repo.list_all():
|
||||
try:
|
||||
cache[setting.setting_key] = setting.get_typed_value()
|
||||
except Exception as exc: # 损坏配置不阻塞启动
|
||||
logger.warning(
|
||||
"[system_config] 跳过损坏配置 %s: %s",
|
||||
setting.setting_key,
|
||||
exc,
|
||||
)
|
||||
finally:
|
||||
session.close()
|
||||
self._cache = cache
|
||||
return cache
|
||||
|
||||
def reload(self) -> dict[str, Any]:
|
||||
"""强制重新从 DB 加载配置,返回新缓存."""
|
||||
with self._lock:
|
||||
self._cache = None
|
||||
return self._load_cache()
|
||||
|
||||
def invalidate(self) -> None:
|
||||
with self._lock:
|
||||
self._cache = None
|
||||
|
||||
# ── 读 ───────────────────────────────────────────────────────
|
||||
def get_config(self, key: str, default: Any = None) -> Any:
|
||||
cache = self._load_cache()
|
||||
return cache.get(key, default)
|
||||
|
||||
def list_configs(self, category: str | None = None) -> list[SystemSetting]:
|
||||
factory = self._get_session_factory()
|
||||
session = factory()
|
||||
try:
|
||||
return SQLAlchemySystemSettingRepository(session).list_all(category)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
# ── 写 ───────────────────────────────────────────────────────
|
||||
def set_config(
|
||||
self,
|
||||
key: str,
|
||||
value: Any,
|
||||
setting_type: str | None = None,
|
||||
updated_by: str | None = None,
|
||||
*,
|
||||
description: str | None = None,
|
||||
category: str = "general",
|
||||
is_public: bool = False,
|
||||
) -> Any:
|
||||
st = setting_type or infer_setting_type(value)
|
||||
serialized = serialize_setting_value(value, st)
|
||||
|
||||
factory = self._get_session_factory()
|
||||
session = factory()
|
||||
try:
|
||||
repo = SQLAlchemySystemSettingRepository(session)
|
||||
existing = repo.get_by_key(key)
|
||||
if existing is not None:
|
||||
existing.setting_value = serialized
|
||||
existing.setting_type = st
|
||||
if updated_by is not None:
|
||||
existing.updated_by = updated_by
|
||||
if description is not None:
|
||||
existing.description = description
|
||||
setting = existing
|
||||
else:
|
||||
setting = SystemSetting(
|
||||
setting_key=key,
|
||||
setting_value=serialized,
|
||||
setting_type=st,
|
||||
description=description or "",
|
||||
is_public=is_public,
|
||||
category=category,
|
||||
updated_by=updated_by,
|
||||
)
|
||||
repo.upsert(setting)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
self.invalidate()
|
||||
self._reset_ditto_singletons(key)
|
||||
return value
|
||||
|
||||
def delete_config(self, key: str) -> bool:
|
||||
factory = self._get_session_factory()
|
||||
session = factory()
|
||||
try:
|
||||
deleted = SQLAlchemySystemSettingRepository(session).delete_by_key(key)
|
||||
finally:
|
||||
session.close()
|
||||
if deleted:
|
||||
self.invalidate()
|
||||
self._reset_ditto_singletons(key)
|
||||
return deleted
|
||||
|
||||
# ── Ditto 联动 ───────────────────────────────────────────────
|
||||
@staticmethod
|
||||
def _reset_ditto_singletons(key: str) -> None:
|
||||
try:
|
||||
from packages.application import ditto_emotion_service as emo_mod
|
||||
from packages.application import ditto_service as ditto_mod
|
||||
|
||||
emo_mod.reset_ditto_emotion_service()
|
||||
ditto_mod.reset_ditto_client()
|
||||
logger.info("[system_config] %s 更新,已重置 Ditto 单例", key)
|
||||
except Exception as exc: # 联动失败不影响配置落库
|
||||
logger.warning("[system_config] 重置 Ditto 单例失败: %s", exc)
|
||||
|
||||
|
||||
_service_singleton: SystemConfigService | None = None
|
||||
|
||||
|
||||
def get_system_config_service() -> SystemConfigService:
|
||||
global _service_singleton
|
||||
if _service_singleton is None:
|
||||
_service_singleton = SystemConfigService()
|
||||
return _service_singleton
|
||||
|
||||
|
||||
def get_config(key: str, default: Any = None) -> Any:
|
||||
"""便捷读取:优先 DB 配置,未配置时返回 default(调用方传 env 值兜底)."""
|
||||
try:
|
||||
return get_system_config_service().get_config(key, default)
|
||||
except Exception as exc:
|
||||
logger.warning("[system_config] 读取 %s 失败,使用默认值: %s", key, exc)
|
||||
return default
|
||||
|
||||
|
||||
def set_config(
|
||||
key: str,
|
||||
value: Any,
|
||||
setting_type: str | None = None,
|
||||
updated_by: str | None = None,
|
||||
) -> Any:
|
||||
return get_system_config_service().set_config(key, value, setting_type=setting_type, updated_by=updated_by)
|
||||
@@ -124,10 +124,34 @@ _STORYBOARD_SYSTEM = (
|
||||
3. copy_display_markdown:直接展示给最终用户的文案,用 Markdown 写成自然、流畅、有感染力的成片成片文案,可用小标题与短句组织;不要做字段列表,不要出现“镜头一/台词:”这类制作说明。
|
||||
4. 内容必须来自图片观察与用户给出的信息,不编造卖点、不夸大、不使用绝对化用语和虚假承诺。
|
||||
5. reference_image_index 填本镜参考图片序号(从 0 开始),没有合适参考图填 -1。
|
||||
6. 分镜数量与时长匹配总时长,节奏紧凑。"""
|
||||
6. 分镜数量与时长匹配总时长,节奏紧凑。
|
||||
7. 口播字数硬约束(必须严格遵守):按每秒约 2.5~3 个中文字(正常口播语速)计算:
|
||||
- 5秒视频:voiceover_script 总字数 12~15 字
|
||||
- 10秒视频:voiceover_script 总字数 25~30 字
|
||||
- 15秒视频:voiceover_script 总字数 35~45 字
|
||||
- 20秒视频:voiceover_script 总字数 50~60 字
|
||||
- 30秒视频:voiceover_script 总字数 75~90 字
|
||||
- 每个 clip 的 voiceover 字数按该镜头时长比例分配
|
||||
- 所有 clip 的 voiceover 字数之和必须等于总 voiceover_script 字数
|
||||
- 宁可少写也不要多写,超长会导致 TTS 音频超出视频时长限制
|
||||
8. 镜头数量硬约束(必须严格遵守):
|
||||
- 5秒视频:1~2 个镜头
|
||||
- 10秒视频:3 个镜头
|
||||
- 15秒视频:3~4 个镜头
|
||||
- 20秒视频:4~5 个镜头
|
||||
- 30秒视频:6~8 个镜头
|
||||
9. 时间轴硬约束(必须严格遵守):
|
||||
- 每个 clip 的 time_range 必须写成 "X-Y秒" 格式,X 和 Y 是具体数字
|
||||
- 第一个 clip 必须从 0 秒开始
|
||||
- 最后一个 clip 必须结束于 total_duration 秒
|
||||
- 相邻 clip 首尾相接,不能有间隙也不能重叠
|
||||
- 每个 clip 的时长 = Y - X,必须 >= 2 秒
|
||||
10. 每个 clip 必须分配一个 reference_image_index(从 0 开始的图片序号),没有合适图片填 -1
|
||||
11. 必须严格按<marketing_purpose><target_audience><persona><viral_structure><language><industry>指定的参数写文案和分镜,不能忽略任何一项用户参数"""
|
||||
)
|
||||
|
||||
_STORYBOARD_USER = """<marketing_purpose>{marketing_purpose}</marketing_purpose>
|
||||
<industry>{industry}</industry>
|
||||
<image_analysis>
|
||||
{image_summary}
|
||||
</image_analysis>
|
||||
@@ -137,6 +161,9 @@ _STORYBOARD_USER = """<marketing_purpose>{marketing_purpose}</marketing_purpose>
|
||||
<aspect_ratio>{aspect_ratio}</aspect_ratio>
|
||||
<tone>{tone}</tone>
|
||||
<target_audience>{target_audience}</target_audience>
|
||||
<persona>{persona_hint}</persona>
|
||||
<viral_structure>{viral_structure_hint}</viral_structure>
|
||||
<language>{language_hint}</language>
|
||||
<extra_requirements>{extra_requirements}</extra_requirements>
|
||||
</user_parameters>
|
||||
{video_style_section}
|
||||
|
||||
+42
-4
@@ -185,8 +185,8 @@ class SharedSettings(BaseSettings):
|
||||
default="",
|
||||
validation_alias=AliasChoices("DITTO_API_BASE_URL", "ditto_api_base_url"),
|
||||
)
|
||||
# 默认人物模板视频 URL(正面 5-10 秒循环、光线均匀、半身)。Ditto 模式下忽略
|
||||
# 用户上传的驱动视频/图片,统一用该模板;后续可扩展为多模板让用户选择。
|
||||
# 默认人物模板视频 URL(兜底用:用户未上传视频时使用,或视频预处理失败时回退)。
|
||||
# 正面 5-10 秒、光线均匀、半身 1080x1920 竖版;正常流程下 Ditto 优先使用用户上传的 video_url。
|
||||
ditto_default_video_url: str = Field(
|
||||
default="",
|
||||
validation_alias=AliasChoices("DITTO_DEFAULT_VIDEO_URL", "ditto_default_video_url"),
|
||||
@@ -196,11 +196,49 @@ class SharedSettings(BaseSettings):
|
||||
default=3,
|
||||
validation_alias=AliasChoices("DITTO_MAX_RETRIES", "ditto_max_retries"),
|
||||
)
|
||||
# Ditto 单次请求超时(秒):数字人半身视频推理通常 30-120s
|
||||
# Ditto 单次请求 read 超时(秒):数字人半身视频推理通常 30-120s(RTF≈2.8,40s音频约112s)
|
||||
# connect 超时固定 10s(代码硬编码,网络不通快速失败)
|
||||
ditto_request_timeout: int = Field(
|
||||
default=300,
|
||||
default=120,
|
||||
validation_alias=AliasChoices("DITTO_REQUEST_TIMEOUT", "ditto_request_timeout"),
|
||||
)
|
||||
# Ditto 句间过渡帧数(平滑表情/口型切换)
|
||||
ditto_blend_frames: int = Field(
|
||||
default=12,
|
||||
validation_alias=AliasChoices("DITTO_BLEND_FRAMES", "ditto_blend_frames"),
|
||||
)
|
||||
|
||||
# ── Ditto LLM 情绪分析(emo_timeline)──────────────────────────────
|
||||
# 总开关;关闭或 LLM 失败时走 GPU 端关键词匹配兜底
|
||||
ditto_emotion_enabled: bool = Field(
|
||||
default=False,
|
||||
validation_alias=AliasChoices("DITTO_EMOTION_ENABLED", "ditto_emotion_enabled"),
|
||||
)
|
||||
ditto_emotion_model: str = Field(
|
||||
default="doubao-seed-2-1-lite-250915",
|
||||
validation_alias=AliasChoices("DITTO_EMOTION_MODEL", "ditto_emotion_model"),
|
||||
)
|
||||
ditto_emotion_temperature: float = Field(
|
||||
default=0.1,
|
||||
validation_alias=AliasChoices("DITTO_EMOTION_TEMPERATURE", "ditto_emotion_temperature"),
|
||||
)
|
||||
ditto_emotion_timeout: int = Field(
|
||||
default=10,
|
||||
validation_alias=AliasChoices("DITTO_EMOTION_TIMEOUT", "ditto_emotion_timeout"),
|
||||
)
|
||||
ditto_emotion_max_tokens: int = Field(
|
||||
default=1024,
|
||||
validation_alias=AliasChoices("DITTO_EMOTION_MAX_TOKENS", "ditto_emotion_max_tokens"),
|
||||
)
|
||||
ditto_emotion_cache_size: int = Field(
|
||||
default=500,
|
||||
validation_alias=AliasChoices("DITTO_EMOTION_CACHE_SIZE", "ditto_emotion_cache_size"),
|
||||
)
|
||||
# 提示词模板:必须包含 {文案} 占位符;后台可通过环境变量覆盖
|
||||
ditto_emotion_prompt: str = Field(
|
||||
default="",
|
||||
validation_alias=AliasChoices("DITTO_EMOTION_PROMPT", "ditto_emotion_prompt"),
|
||||
)
|
||||
|
||||
# ── P4000 NVENC 硬件编码 ────────────────────────────────────────────
|
||||
# GPU 编码总开关;关闭或 endpoint 为空时始终走本机 CPU libx264
|
||||
|
||||
@@ -8,7 +8,12 @@ from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime
|
||||
from datetime import datetime, timezone
|
||||
|
||||
try:
|
||||
from datetime import UTC
|
||||
except ImportError:
|
||||
UTC = timezone.utc
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
"""系统配置领域实体 — #2246.
|
||||
|
||||
承载 setting_type(bool/int/float/string/json)及 setting_value 的
|
||||
序列化/反序列化规则;与 DB、框架无关。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
SETTING_TYPE_BOOL = "bool"
|
||||
SETTING_TYPE_INT = "int"
|
||||
SETTING_TYPE_FLOAT = "float"
|
||||
SETTING_TYPE_STRING = "string"
|
||||
SETTING_TYPE_JSON = "json"
|
||||
|
||||
VALID_SETTING_TYPES = {
|
||||
SETTING_TYPE_BOOL,
|
||||
SETTING_TYPE_INT,
|
||||
SETTING_TYPE_FLOAT,
|
||||
SETTING_TYPE_STRING,
|
||||
SETTING_TYPE_JSON,
|
||||
}
|
||||
|
||||
|
||||
class SystemSettingError(ValueError):
|
||||
"""系统配置类型或序列化错误."""
|
||||
|
||||
|
||||
def infer_setting_type(value: Any) -> str:
|
||||
"""根据 Python 值推断 setting_type(bool 必须先于 int 判断)."""
|
||||
if isinstance(value, bool):
|
||||
return SETTING_TYPE_BOOL
|
||||
if isinstance(value, int):
|
||||
return SETTING_TYPE_INT
|
||||
if isinstance(value, float):
|
||||
return SETTING_TYPE_FLOAT
|
||||
if isinstance(value, str):
|
||||
return SETTING_TYPE_STRING
|
||||
return SETTING_TYPE_JSON
|
||||
|
||||
|
||||
def serialize_setting_value(value: Any, setting_type: str) -> str:
|
||||
"""把 Python 值按 setting_type 序列化为可入库的字符串."""
|
||||
if setting_type == SETTING_TYPE_BOOL:
|
||||
if not isinstance(value, bool):
|
||||
raise SystemSettingError(f"bool 配置值必须是布尔类型,收到 {value!r}")
|
||||
return "true" if value else "false"
|
||||
if setting_type == SETTING_TYPE_INT:
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise SystemSettingError(f"int 配置值必须是整数,收到 {value!r}")
|
||||
return str(value)
|
||||
if setting_type == SETTING_TYPE_FLOAT:
|
||||
if isinstance(value, bool):
|
||||
raise SystemSettingError(f"float 配置值不能是布尔类型,收到 {value!r}")
|
||||
try:
|
||||
return repr(float(value))
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise SystemSettingError(f"float 配置值非法:{value!r}") from exc
|
||||
if setting_type == SETTING_TYPE_STRING:
|
||||
if not isinstance(value, str):
|
||||
raise SystemSettingError(f"string 配置值必须是字符串,收到 {value!r}")
|
||||
return value
|
||||
if setting_type == SETTING_TYPE_JSON:
|
||||
try:
|
||||
return json.dumps(value, ensure_ascii=False)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise SystemSettingError(f"json 配置值无法序列化:{value!r}") from exc
|
||||
raise SystemSettingError(f"未知 setting_type: {setting_type}")
|
||||
|
||||
|
||||
def deserialize_setting_value(raw: str | None, setting_type: str) -> Any:
|
||||
"""把入库字符串按 setting_type 反序列化为 Python 值."""
|
||||
if raw is None:
|
||||
return None
|
||||
if setting_type == SETTING_TYPE_BOOL:
|
||||
return str(raw).strip().lower() in {"1", "true", "yes", "on"}
|
||||
if setting_type == SETTING_TYPE_INT:
|
||||
try:
|
||||
return int(str(raw).strip())
|
||||
except ValueError as exc:
|
||||
raise SystemSettingError(f"int 配置值损坏:{raw!r}") from exc
|
||||
if setting_type == SETTING_TYPE_FLOAT:
|
||||
try:
|
||||
return float(str(raw).strip())
|
||||
except ValueError as exc:
|
||||
raise SystemSettingError(f"float 配置值损坏:{raw!r}") from exc
|
||||
if setting_type == SETTING_TYPE_STRING:
|
||||
return raw
|
||||
if setting_type == SETTING_TYPE_JSON:
|
||||
try:
|
||||
return json.loads(raw)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise SystemSettingError(f"json 配置值损坏:{raw!r}") from exc
|
||||
raise SystemSettingError(f"未知 setting_type: {setting_type}")
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class SystemSetting:
|
||||
"""系统配置领域实体."""
|
||||
|
||||
setting_key: str
|
||||
setting_value: str | None = None
|
||||
setting_type: str = SETTING_TYPE_STRING
|
||||
description: str = ""
|
||||
is_public: bool = False
|
||||
category: str = "general"
|
||||
updated_by: str | None = None
|
||||
id: str = field(default_factory=lambda: uuid4().hex)
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(UTC))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(UTC))
|
||||
|
||||
def get_typed_value(self) -> Any:
|
||||
return deserialize_setting_value(self.setting_value, self.setting_type)
|
||||
|
||||
@classmethod
|
||||
def from_value(
|
||||
cls,
|
||||
setting_key: str,
|
||||
value: Any,
|
||||
setting_type: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> "SystemSetting":
|
||||
st = setting_type or infer_setting_type(value)
|
||||
return cls(
|
||||
setting_key=setting_key,
|
||||
setting_value=serialize_setting_value(value, st),
|
||||
setting_type=st,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -107,6 +107,7 @@ class ViralVideoJob:
|
||||
# v1.5.1 音频/视频参数
|
||||
voice_id: str = ""
|
||||
voice_source: str = ""
|
||||
language: str = "zh-CN"
|
||||
video_ratio: str = "9:16"
|
||||
video_model: str = ""
|
||||
# v1.4+ 产物
|
||||
|
||||
@@ -363,6 +363,111 @@ class DoubaoClient:
|
||||
logger.error("豆包API调用最终失败: elapsed=%.1fs err=%s", time.time() - _t0, last_error)
|
||||
return None
|
||||
|
||||
def chat_completion_stream(
|
||||
self,
|
||||
messages: list[dict[str, str]],
|
||||
temperature: float = 0.7,
|
||||
max_tokens: int | None = None,
|
||||
model: str | None = None,
|
||||
timeout: int | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""流式调用 Chat Completion 接口(SSE),逐块 yield delta 文本。
|
||||
|
||||
Yields:
|
||||
str: 增量文本片段(delta.content);全部结束后 StopIteration。
|
||||
失败时 yield 空并返回(由调用方决定是否降级到同步调用)。
|
||||
|
||||
注意:
|
||||
- 流式不做 finish_reason=length 自动扩容(流式难以拼接重试);
|
||||
如果调用方需要 length 截断处理,建议自行 fallback 到同步 chat_completion。
|
||||
- 重试只在连接建立阶段(首包之前)有效;一旦开始 yield,错误直接抛出。
|
||||
"""
|
||||
if not self.is_available:
|
||||
return
|
||||
|
||||
import json as _json
|
||||
|
||||
url = f"{self.base_url}/chat/completions"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "text/event-stream",
|
||||
}
|
||||
effective_max_tokens = max_tokens if max_tokens is not None else (self.max_tokens or 1024)
|
||||
payload: dict[str, Any] = {
|
||||
"model": model or self.model,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
"max_tokens": effective_max_tokens,
|
||||
"stream": True,
|
||||
}
|
||||
if self.extra_params:
|
||||
payload.update(self.extra_params)
|
||||
if kwargs:
|
||||
payload.update(kwargs)
|
||||
|
||||
_t0 = time.time()
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
_req_timeout = httpx.Timeout(connect=10.0, read=120.0, write=30.0, pool=10.0)
|
||||
if timeout:
|
||||
_req_timeout = httpx.Timeout(connect=10.0, read=max(int(timeout), 30), write=30.0, pool=10.0)
|
||||
with httpx.stream(
|
||||
"POST",
|
||||
url,
|
||||
headers=headers,
|
||||
json=payload,
|
||||
timeout=_req_timeout,
|
||||
) as resp:
|
||||
resp.raise_for_status()
|
||||
for line in resp.iter_lines():
|
||||
if not line:
|
||||
continue
|
||||
line = line.strip()
|
||||
if not line.startswith("data:"):
|
||||
continue
|
||||
data_str = line[5:].strip()
|
||||
if data_str == "[DONE]":
|
||||
break
|
||||
try:
|
||||
chunk = _json.loads(data_str)
|
||||
except (_json.JSONDecodeError, ValueError):
|
||||
continue
|
||||
choices = chunk.get("choices") or []
|
||||
if not choices:
|
||||
continue
|
||||
delta = choices[0].get("delta") or {}
|
||||
content_piece = delta.get("content") or ""
|
||||
if content_piece:
|
||||
yield content_piece
|
||||
finish_reason = choices[0].get("finish_reason")
|
||||
if finish_reason:
|
||||
self.last_finish_reason = finish_reason
|
||||
_elapsed = time.time() - _t0
|
||||
logger.info(
|
||||
"[doubao] chat_completion_stream 完成 model=%s elapsed=%.1fs attempt=%d",
|
||||
payload.get("model"),
|
||||
_elapsed,
|
||||
attempt + 1,
|
||||
)
|
||||
return
|
||||
except Exception as e:
|
||||
if attempt < self.max_retries:
|
||||
wait = 0.5 * (2**attempt)
|
||||
logger.warning(
|
||||
"豆包流式API失败,%.1fs后重试 (%d/%d, elapsed=%.1fs): %s",
|
||||
wait,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
time.time() - _t0,
|
||||
e,
|
||||
)
|
||||
time.sleep(wait)
|
||||
continue
|
||||
logger.error("豆包流式API最终失败: elapsed=%.1fs err=%s", time.time() - _t0, e)
|
||||
return
|
||||
|
||||
def vision_completion(
|
||||
self,
|
||||
messages: list[dict],
|
||||
@@ -636,6 +741,10 @@ class DoubaoClient:
|
||||
resolution=resolution,
|
||||
output_dir=output_dir,
|
||||
model=video_model,
|
||||
generate_audio=bool(generate_audio),
|
||||
reference_images=reference_images,
|
||||
reference_audios=reference_audios,
|
||||
reference_videos=reference_videos,
|
||||
)
|
||||
if not result and hasattr(ds, "last_video_error") and ds.last_video_error:
|
||||
self.last_video_error = dict(ds.last_video_error)
|
||||
|
||||
@@ -115,10 +115,20 @@ class DashScopeClient:
|
||||
watermark: bool = False,
|
||||
output_dir: str | None = None,
|
||||
model: str = "wan3.0-video",
|
||||
generate_audio: bool = True,
|
||||
reference_images: list[str] | None = None,
|
||||
reference_audios: list[str] | None = None,
|
||||
reference_videos: list[str] | None = None,
|
||||
) -> dict | None:
|
||||
"""调用 DashScope 异步视频合成接口,轮询完成后下载到本地。
|
||||
"""调用 DashScope Wan 3.0 异步视频合成接口,轮询完成后下载到本地。
|
||||
|
||||
返回 {"video_path": str, "usage": dict | None};失败返回 None,错误详情写入 self.last_video_error。
|
||||
官方协议(input.media 数组 + parameters.audio):
|
||||
- 仅 1 张图且无其它参考 -> type=first_frame(首帧模式,严格从该帧起)。
|
||||
- 有参考音频 / 多张图 -> 图片全部走 type=reference_image(全能参考模式,
|
||||
可与 reference_audio 共存);prompt 用"图1/图2/音频1"按 media 顺序引用。
|
||||
- parameters.audio 控制输出是否含音轨;参考音频通过 media 传入。
|
||||
|
||||
返回 {"video_path": str, "usage": dict | None};失败返回 None,错误写入 self.last_video_error。
|
||||
"""
|
||||
self.last_video_error = {}
|
||||
if not self.is_available:
|
||||
@@ -130,7 +140,7 @@ class DashScopeClient:
|
||||
return None
|
||||
|
||||
# DashScope 分辨率参数:720P / 1080P / 480P(大写 P)
|
||||
res_upper = (resolution or "720p").upper().replace("P", "P")
|
||||
res_upper = (resolution or "720p").upper()
|
||||
if res_upper == "480P":
|
||||
ds_res = "480P"
|
||||
elif res_upper == "1080P":
|
||||
@@ -138,16 +148,35 @@ class DashScopeClient:
|
||||
else:
|
||||
ds_res = "720P"
|
||||
|
||||
# 构造 input+parameters
|
||||
# ── 构造官方 media 数组 ────────────────────────────────────────
|
||||
ref_imgs = [u for u in (reference_images or [])[:10] if u]
|
||||
ref_auds = [u for u in (reference_audios or [])[:5] if u]
|
||||
ref_vids = [u for u in (reference_videos or [])[:5] if u]
|
||||
|
||||
media: list[dict[str, Any]] = []
|
||||
all_imgs = ([image_url] if image_url else []) + [u for u in ref_imgs if u != image_url]
|
||||
use_first_frame = bool(image_url) and len(all_imgs) == 1 and not (ref_auds or ref_vids)
|
||||
if use_first_frame:
|
||||
media.append({"type": "first_frame", "url": image_url})
|
||||
else:
|
||||
for u in all_imgs:
|
||||
media.append({"type": "reference_image", "url": u})
|
||||
for u in ref_vids:
|
||||
media.append({"type": "reference_video", "url": u})
|
||||
for u in ref_auds:
|
||||
media.append({"type": "reference_audio", "url": u})
|
||||
|
||||
# ── input + parameters ─────────────────────────────────────────
|
||||
input_obj: dict[str, Any] = {"prompt": prompt.strip()}
|
||||
if image_url:
|
||||
input_obj["img_url"] = image_url
|
||||
if media:
|
||||
input_obj["media"] = media
|
||||
|
||||
params: dict[str, Any] = {
|
||||
"resolution": ds_res,
|
||||
"duration": str(float(duration)),
|
||||
"duration": int(duration),
|
||||
"watermark": bool(watermark),
|
||||
"audio": bool(generate_audio),
|
||||
}
|
||||
# 比例透传:Wan 支持 "9:16" / "16:9" / "1:1" 等
|
||||
if ratio and ratio != "adaptive":
|
||||
params["aspect_ratio"] = ratio
|
||||
|
||||
@@ -163,14 +192,15 @@ class DashScopeClient:
|
||||
}
|
||||
create_url = f"{self.base_url}/services/aigc/video-generation/video-synthesis"
|
||||
logger.info(
|
||||
"[dashscope] 创建任务: model=%s dur=%ds ratio=%s res=%s img=%s",
|
||||
"[dashscope] 创建任务: model=%s dur=%ds ratio=%s res=%s media=%d audio=%s",
|
||||
model,
|
||||
duration,
|
||||
ratio,
|
||||
ds_res,
|
||||
bool(image_url),
|
||||
len(media),
|
||||
generate_audio,
|
||||
)
|
||||
logger.info("[dashscope] 创建任务 payload: model=%s params=%s", model, params)
|
||||
logger.info("[dashscope] media types: %s", [m["type"] for m in media])
|
||||
|
||||
# 创建任务
|
||||
task_id: str | None = None
|
||||
@@ -196,27 +226,14 @@ class DashScopeClient:
|
||||
if tid:
|
||||
task_id = tid
|
||||
break
|
||||
# 部分情况下 code != 错误
|
||||
code = data.get("code")
|
||||
if code and code != "":
|
||||
err_code, user_msg = _classify_dashscope_error(400, body_text, str(code))
|
||||
if code:
|
||||
logger.error("[dashscope] 创建任务返回 code=%s body=%s", code, body_text)
|
||||
err_code, user_msg = _classify_dashscope_error(sc, body_text)
|
||||
self._set_error(err_code, user_msg, sc, body_text, model=model)
|
||||
return None
|
||||
else:
|
||||
self._set_error("unknown", "Wan 3.0 响应格式异常,未返回任务ID", sc, str(data)[:500], model=model)
|
||||
return None
|
||||
except _HTTP_NETWORK_ERRORS as ne:
|
||||
last_sc = 0
|
||||
last_body = f"network error: {ne}"
|
||||
logger.warning(
|
||||
"[dashscope] 网络异常 %s,重试 %d/%d", type(ne).__name__, attempt + 1, self.max_retries + 1
|
||||
)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
self._set_error("network_error", "Wan 3.0 服务连接失败(网络超时),请稍后重试。", 0, str(ne))
|
||||
return None
|
||||
except Exception as _e:
|
||||
logger.warning("[dashscope] 创建任务异常(attempt=%d): %s", attempt, _e)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
@@ -258,7 +275,6 @@ class DashScopeClient:
|
||||
video_url = out.get("video_url") or ""
|
||||
usage = d.get("usage")
|
||||
if not video_url:
|
||||
# 结果在 results 数组
|
||||
results = out.get("results") or []
|
||||
if results and isinstance(results, list):
|
||||
video_url = results[0].get("url") or results[0].get("video_url")
|
||||
@@ -284,7 +300,6 @@ class DashScopeClient:
|
||||
logger.warning("[dashscope] 任务 %s 被取消", task_id)
|
||||
self._set_error("unknown", "Wan 3.0 任务被取消。", 200, "task cancelled", task_id=task_id)
|
||||
return None
|
||||
# PENDING / RUNNING / SUSPENDED → 继续轮询
|
||||
if poll_count % 5 == 0:
|
||||
logger.info("[dashscope] 轮询中 task=%s status=%s polls=%d", task_id, task_status, poll_count)
|
||||
except Exception as e:
|
||||
|
||||
@@ -334,7 +334,7 @@ class SharedStorageService(StoragePort):
|
||||
|
||||
storage_key = self.normalize_storage_key(storage_key_or_url)
|
||||
try:
|
||||
signed = sign_bucket.sign_url("GET", storage_key, expires_seconds)
|
||||
signed = sign_bucket.sign_url("GET", storage_key, expires_seconds, slash_safe=True)
|
||||
logger.info(
|
||||
"signed URL generated for key=%s prefix=%s",
|
||||
storage_key[:80],
|
||||
|
||||
@@ -12,3 +12,22 @@ export CI_LOCAL_PG_PORT="${CI_LOCAL_PG_PORT:-5432}"
|
||||
|
||||
# === 默认数据库名 ===
|
||||
export CI_DEFAULT_DB="${CI_DEFAULT_DB:-xiaoxia_saas}"
|
||||
|
||||
# === Python 版本保障:本项目依赖 datetime.UTC,需要 Python >= 3.11 ===
|
||||
_ensure_python311() {
|
||||
if python3 -c "import sys; assert sys.version_info >= (3, 11)" 2>/dev/null; then
|
||||
return 0
|
||||
fi
|
||||
for cand in python3.12 python3.11 /opt/python3.12/bin/python3; do
|
||||
if command -v "$cand" >/dev/null 2>&1 && "$cand" -c "import sys; assert sys.version_info >= (3, 11)" 2>/dev/null; then
|
||||
_d="$(dirname "$(command -v "$cand")")"
|
||||
export PATH="$_d:$PATH"
|
||||
hash -r
|
||||
echo "✅ ci_env: 切换到 $cand ($("$cand" -c 'import sys; print(sys.version.split()[0])'))"
|
||||
return 0
|
||||
fi
|
||||
done
|
||||
echo "❌ ci_env: 未找到 Python >= 3.11(datetime.UTC 需要),请安装 Python 3.11/3.12" >&2
|
||||
return 1
|
||||
}
|
||||
_ensure_python311
|
||||
|
||||
@@ -4,6 +4,11 @@
|
||||
|
||||
set -e
|
||||
|
||||
SCRIPT_DIR="$(dirname "${BASH_SOURCE[0]}")"
|
||||
# shellcheck source=ci_env.sh
|
||||
source "${SCRIPT_DIR}/ci_env.sh"
|
||||
|
||||
|
||||
echo "=== Installing mypy ==="
|
||||
python3 -m pip install -q mypy
|
||||
mypy --version
|
||||
|
||||
@@ -7,9 +7,29 @@ JOB_NAME="${1:-Unit Tests}"
|
||||
|
||||
echo "=== CI Unit Tests 开始 ==="
|
||||
|
||||
# --- 依赖缓存检查 ---
|
||||
# --- Python 版本选择(必须 >= 3.11,代码使用 datetime.UTC)---
|
||||
if ! python3 -c "import sys; assert sys.version_info >= (3, 11)" 2>/dev/null; then
|
||||
for cand in python3.12 python3.11 /opt/python3.12/bin/python3; do
|
||||
if command -v "$cand" >/dev/null 2>&1 && "$cand" -c "import sys; assert sys.version_info >= (3, 11)" 2>/dev/null; then
|
||||
PY3_DIR="$(dirname "$(command -v "$cand")")"
|
||||
export PATH="$PY3_DIR:$PATH"
|
||||
hash -r
|
||||
echo "✅ python3 版本过低,改用 $cand ($("$cand" -c "import sys; print(sys.version.split()[0])"))"
|
||||
break
|
||||
fi
|
||||
done
|
||||
if ! python3 -c "import sys; assert sys.version_info >= (3, 11)" 2>/dev/null; then
|
||||
echo "❌ 未找到 Python >= 3.11,本项目要求 Python 3.11+(使用 datetime.UTC)" >&2
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
|
||||
PYVER=$(python3 -c 'import sys; print(f"{sys.version_info.major}.{sys.version_info.minor}")')
|
||||
echo "使用 Python 版本: $(python3 --version)"
|
||||
|
||||
# --- 依赖缓存检查(按 Python 版本区分缓存文件,避免跨版本复用)---
|
||||
# 如果 requirements 文件未变化且依赖已安装,跳过 pip install(持久 runner 优化)
|
||||
REQ_HASH_FILE="/tmp/.ci_unit_tests_req_hash"
|
||||
REQ_HASH_FILE="/tmp/.ci_unit_tests_req_hash_py${PYVER}"
|
||||
CURRENT_REQ_HASH=""
|
||||
if [ -f requirements-base.txt ] && [ -f requirements.txt ] && [ -f requirements-dev.txt ]; then
|
||||
CURRENT_REQ_HASH=$(cat requirements-base.txt requirements.txt requirements-dev.txt | md5sum | cut -d' ' -f1)
|
||||
|
||||
@@ -77,6 +77,11 @@ set +e
|
||||
bandit -r apps packages -q -ll
|
||||
BANDIT_EXIT=$?
|
||||
set -e
|
||||
|
||||
SCRIPT_DIR="$(dirname "${BASH_SOURCE[0]}")"
|
||||
# shellcheck source=ci_env.sh
|
||||
source "${SCRIPT_DIR}/ci_env.sh"
|
||||
|
||||
if [ "$BANDIT_EXIT" -ne 0 ]; then
|
||||
echo "⚠️ Bandit found security issues (advisory mode - not blocking CI)"
|
||||
else
|
||||
|
||||
@@ -196,3 +196,159 @@ class TestGetDoubaoClient:
|
||||
"""返回 DoubaoClient 实例"""
|
||||
client = get_doubao_client()
|
||||
assert isinstance(client, DoubaoClient)
|
||||
|
||||
|
||||
class TestChatCompletionStream:
|
||||
"""chat_completion_stream 方法测试"""
|
||||
|
||||
def _fake_sse(self, chunks: list[str]):
|
||||
"""构造一个伪装的 httpx stream 响应,按 chunks 逐行返回 SSE。"""
|
||||
import json as _json
|
||||
|
||||
lines = []
|
||||
for piece in chunks:
|
||||
evt = {"choices": [{"delta": {"content": piece}, "finish_reason": None}]}
|
||||
lines.append("data: " + _json.dumps(evt, ensure_ascii=False))
|
||||
lines.append("data: " + _json.dumps({"choices": [{"delta": {}, "finish_reason": "stop"}]}))
|
||||
lines.append("data: [DONE]")
|
||||
|
||||
m = MagicMock()
|
||||
m.status_code = 200
|
||||
m.__enter__ = MagicMock(return_value=m)
|
||||
m.__exit__ = MagicMock(return_value=False)
|
||||
m.raise_for_status = MagicMock()
|
||||
m.iter_lines = MagicMock(return_value=iter(lines))
|
||||
return m
|
||||
|
||||
def test_stream_yields_delta_content(self, client_with_key):
|
||||
"""流式调用逐段 yield delta.content"""
|
||||
import json as _json
|
||||
|
||||
fake = self._fake_sse(["你好", ",", "世界"])
|
||||
with patch("packages.shared.ai_client.httpx.stream", return_value=fake) as mock_stream:
|
||||
out = list(
|
||||
client_with_key.chat_completion_stream(
|
||||
messages=[{"role": "user", "content": "hi"}], max_tokens=100, timeout=30
|
||||
)
|
||||
)
|
||||
assert out == ["你好", ",", "世界"]
|
||||
call_kwargs = mock_stream.call_args[1]
|
||||
assert call_kwargs["json"]["stream"] is True
|
||||
assert call_kwargs["json"]["max_tokens"] == 100
|
||||
|
||||
def test_stream_unavailable_returns_empty(self, client_without_key):
|
||||
"""不可用时返回空生成器(不调用 httpx.stream)"""
|
||||
with patch("packages.shared.ai_client.httpx.stream") as mock_stream:
|
||||
out = list(client_without_key.chat_completion_stream(messages=[{"role": "user", "content": "hi"}]))
|
||||
assert out == []
|
||||
mock_stream.assert_not_called()
|
||||
|
||||
def test_stream_handles_done_marker(self, client_with_key):
|
||||
"""遇到 [DONE] 正确终止,不把它当内容"""
|
||||
import json as _json
|
||||
|
||||
lines = [
|
||||
"data: " + _json.dumps({"choices": [{"delta": {"content": "A"}}]}),
|
||||
"data: [DONE]",
|
||||
"data: " + _json.dumps({"choices": [{"delta": {"content": "NEVER"}}]}), # 应被忽略
|
||||
]
|
||||
m = MagicMock()
|
||||
m.status_code = 200
|
||||
m.__enter__ = MagicMock(return_value=m)
|
||||
m.__exit__ = MagicMock(return_value=False)
|
||||
m.raise_for_status = MagicMock()
|
||||
m.iter_lines = MagicMock(return_value=iter(lines))
|
||||
with patch("packages.shared.ai_client.httpx.stream", return_value=m):
|
||||
out = list(
|
||||
client_with_key.chat_completion_stream(
|
||||
messages=[{"role": "user", "content": "hi"}], max_tokens=50, timeout=30
|
||||
)
|
||||
)
|
||||
assert out == ["A"]
|
||||
|
||||
def test_stream_skips_empty_delta(self, client_with_key):
|
||||
"""空 delta(role 等 metadata)不应产出内容"""
|
||||
import json as _json
|
||||
|
||||
lines = [
|
||||
"data: " + _json.dumps({"choices": [{"delta": {"role": "assistant"}}]}),
|
||||
"data: " + _json.dumps({"choices": [{"delta": {"content": "hi"}}]}),
|
||||
"data: " + _json.dumps({"choices": [{"delta": {}}]}),
|
||||
"data: [DONE]",
|
||||
]
|
||||
m = MagicMock()
|
||||
m.status_code = 200
|
||||
m.__enter__ = MagicMock(return_value=m)
|
||||
m.__exit__ = MagicMock(return_value=False)
|
||||
m.raise_for_status = MagicMock()
|
||||
m.iter_lines = MagicMock(return_value=iter(lines))
|
||||
with patch("packages.shared.ai_client.httpx.stream", return_value=m):
|
||||
out = list(
|
||||
client_with_key.chat_completion_stream(
|
||||
messages=[{"role": "user", "content": "x"}], max_tokens=50, timeout=30
|
||||
)
|
||||
)
|
||||
assert out == ["hi"]
|
||||
|
||||
def test_stream_invalid_json_lines_skipped(self, client_with_key):
|
||||
"""SSE 行里脏数据/非 JSON 不应中断流"""
|
||||
import json as _json
|
||||
|
||||
lines = [
|
||||
"data: " + _json.dumps({"choices": [{"delta": {"content": "ok"}}]}),
|
||||
"data: not-a-json",
|
||||
":comment line",
|
||||
"",
|
||||
"event: ping",
|
||||
"data: " + _json.dumps({"choices": [{"delta": {"content": "2"}}]}),
|
||||
"data: [DONE]",
|
||||
]
|
||||
m = MagicMock()
|
||||
m.status_code = 200
|
||||
m.__enter__ = MagicMock(return_value=m)
|
||||
m.__exit__ = MagicMock(return_value=False)
|
||||
m.raise_for_status = MagicMock()
|
||||
m.iter_lines = MagicMock(return_value=iter(lines))
|
||||
with patch("packages.shared.ai_client.httpx.stream", return_value=m):
|
||||
out = list(
|
||||
client_with_key.chat_completion_stream(
|
||||
messages=[{"role": "user", "content": "x"}], max_tokens=50, timeout=30
|
||||
)
|
||||
)
|
||||
assert out == ["ok", "2"]
|
||||
|
||||
def test_stream_network_error_yields_empty(self, client_with_key):
|
||||
"""网络错误(重试耗尽)yield 空,不抛异常给调用方"""
|
||||
import httpx as _httpx
|
||||
|
||||
# max_retries=2 → 3 次总尝试
|
||||
client = client_with_key
|
||||
client.max_retries = 1 # 只重试 1 次,缩短测试
|
||||
with patch("packages.shared.ai_client.httpx.stream", side_effect=_httpx.ConnectError("boom")):
|
||||
out = list(
|
||||
client.chat_completion_stream(messages=[{"role": "user", "content": "x"}], max_tokens=50, timeout=5)
|
||||
)
|
||||
assert out == []
|
||||
|
||||
def test_stream_records_finish_reason(self, client_with_key):
|
||||
"""流式结束后 last_finish_reason 被正确记录"""
|
||||
import json as _json
|
||||
|
||||
lines = [
|
||||
"data: " + _json.dumps({"choices": [{"delta": {"content": "x"}, "finish_reason": None}]}),
|
||||
"data: " + _json.dumps({"choices": [{"delta": {}, "finish_reason": "stop"}]}),
|
||||
"data: [DONE]",
|
||||
]
|
||||
m = MagicMock()
|
||||
m.status_code = 200
|
||||
m.__enter__ = MagicMock(return_value=m)
|
||||
m.__exit__ = MagicMock(return_value=False)
|
||||
m.raise_for_status = MagicMock()
|
||||
m.iter_lines = MagicMock(return_value=iter(lines))
|
||||
with patch("packages.shared.ai_client.httpx.stream", return_value=m):
|
||||
list(
|
||||
client_with_key.chat_completion_stream(
|
||||
messages=[{"role": "user", "content": "x"}], max_tokens=50, timeout=30
|
||||
)
|
||||
)
|
||||
assert client_with_key.last_finish_reason == "stop"
|
||||
|
||||
+233
-250
@@ -1,209 +1,210 @@
|
||||
"""AI Router 单元测试 — 23 cases covering routing/cache/fallback/client construction."""
|
||||
"""AI Router 单元测试 — routing/cache/fallback/client construction.
|
||||
|
||||
本文件只做*用例级* mock:通过 autouse fixture 在每个用例内 patch
|
||||
``packages.shared.config.get_shared_settings`` / ``packages.shared.ai_router.get_shared_settings``
|
||||
并在退出时自动恢复,绝不在模块顶层替换 ``sys.modules``,因此不会污染同进程的
|
||||
其他测试模块(如 test_ai_client.py)。
|
||||
|
||||
在 Python 3.12 且依赖齐全的 CI 环境中,直接 import 真实模块即可;Redis / DB
|
||||
会话通过 patch 隔离。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
import unittest
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
# ── Pre-mock heavy import chain to avoid pulling in full app ──
|
||||
_mock_config = MagicMock()
|
||||
_mock_settings = MagicMock()
|
||||
_mock_settings.doubao_model = "doubao-seed-2-1-pro-260915"
|
||||
_mock_settings.doubao_fast_model = "doubao-seed-2-1-pro-260915"
|
||||
_mock_settings.doubao_base_url = "https://ark.cn-beijing.volces.com/api/v3"
|
||||
_mock_settings.doubao_api_key = "test-key"
|
||||
_mock_settings.doubao_timeout = 45
|
||||
_mock_settings.doubao_max_retries = 1
|
||||
_mock_settings.doubao_image_model = "doubao-seedream-5-0-flash-260915"
|
||||
_mock_settings.doubao_image_timeout = 60
|
||||
_mock_settings.doubao_video_model = "doubao-seedance-2-5-260628"
|
||||
_mock_settings.doubao_video_timeout = 600
|
||||
_mock_settings.dashscope_api_key = "ds-key"
|
||||
_mock_settings.cosyvoice_api_key = "cv-key"
|
||||
_mock_settings.cosyvoice_base_url = "https://dashscope.aliyuncs.com/api/v1"
|
||||
_mock_settings.cosyvoice_model = "cosyvoice-v3-flash"
|
||||
_mock_settings.redis_url = "redis://localhost:6379/0"
|
||||
_mock_settings.celery_broker_url = "redis://localhost:6379/0"
|
||||
_mock_config.get_shared_settings.return_value = _mock_settings
|
||||
import pytest
|
||||
|
||||
# Prevent the full packages.shared from loading
|
||||
for mod_name in list(sys.modules.keys()):
|
||||
if "packages.shared" in mod_name and "ai_router" not in mod_name and "ai_config_version" not in mod_name:
|
||||
pass # don't remove, just prevent new imports
|
||||
from packages.shared import ai_config_version as _config_version_mod
|
||||
from packages.shared import ai_router as ai_router_mod
|
||||
|
||||
# Direct import of our modules (bypassing __init__.py)
|
||||
import importlib.util
|
||||
import os
|
||||
# ── 统一的假配置(等价于旧文件里的 _mock_settings)──────────────────────────
|
||||
|
||||
|
||||
def _load_module_from_file(name, path):
|
||||
spec = importlib.util.spec_from_file_location(name, path)
|
||||
mod = importlib.util.module_from_spec(spec)
|
||||
sys.modules[name] = mod
|
||||
spec.loader.exec_module(mod)
|
||||
return mod
|
||||
def _make_mock_settings() -> MagicMock:
|
||||
s = MagicMock()
|
||||
s.doubao_model = "doubao-seed-2-1-pro-260915"
|
||||
s.doubao_fast_model = "doubao-seed-2-1-pro-260915"
|
||||
s.doubao_base_url = "https://ark.cn-beijing.volces.com/api/v3"
|
||||
s.doubao_api_key = "test-key"
|
||||
s.doubao_timeout = 45
|
||||
s.doubao_max_retries = 1
|
||||
s.doubao_image_model = "doubao-seedream-5-0-flash-260915"
|
||||
s.doubao_image_timeout = 60
|
||||
s.doubao_vision_model = "doubao-seed-1-6-vision-250615"
|
||||
s.doubao_video_model = "doubao-seedance-2-5-260628"
|
||||
s.doubao_video_timeout = 600
|
||||
s.dashscope_api_key = "ds-key"
|
||||
s.dashscope_base_url = "https://dashscope.aliyuncs.com/api/v1"
|
||||
s.dashscope_model = "qwen-vl-max"
|
||||
s.cosyvoice_api_key = "cv-key"
|
||||
s.cosyvoice_base_url = "https://dashscope.aliyuncs.com/api/v1"
|
||||
s.cosyvoice_model = "cosyvoice-v3-flash"
|
||||
s.redis_url = "redis://localhost:6379/0"
|
||||
s.celery_broker_url = "redis://localhost:6379/0"
|
||||
return s
|
||||
|
||||
|
||||
# Load ai_config_version
|
||||
_ai_config_version = _load_module_from_file(
|
||||
"packages.shared.ai_config_version",
|
||||
os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "packages", "shared", "ai_config_version.py"),
|
||||
)
|
||||
# Patch get_shared_settings in the loaded module
|
||||
_ai_config_version.get_shared_settings = lambda: _mock_settings
|
||||
@pytest.fixture(autouse=True)
|
||||
def _mock_settings_fixture():
|
||||
"""每个用例内 patch 配置来源,退出即恢复,不污染 sys.modules。"""
|
||||
settings = _make_mock_settings()
|
||||
with (
|
||||
patch("packages.shared.config.get_shared_settings", return_value=settings),
|
||||
patch.object(ai_router_mod, "get_shared_settings", return_value=settings),
|
||||
):
|
||||
yield settings
|
||||
|
||||
# Load ai_router - needs packages.shared.config to be available
|
||||
sys.modules["packages.shared.config"] = MagicMock()
|
||||
sys.modules["packages.shared.config"].get_shared_settings = lambda: _mock_settings
|
||||
|
||||
# Mock packages.shared.ai_client to avoid triggering packages.shared.__init__ chain
|
||||
# (which fails on Python 3.10 due to datetime.UTC import in packages.domain)
|
||||
_mock_ai_client = MagicMock()
|
||||
def _capability_row() -> MagicMock:
|
||||
row = MagicMock()
|
||||
row.capability_key = "intent_parsing"
|
||||
row.capability_name = "文案意图解析"
|
||||
row.timeout_seconds = 45
|
||||
row.max_retries = 1
|
||||
row.max_tokens = None
|
||||
row.temperature = None
|
||||
row.concurrency = 2
|
||||
row.extra_params = {}
|
||||
row.is_enabled = True
|
||||
row.pm_id = "model-1"
|
||||
row.pm_name = "豆包"
|
||||
row.pm_provider = "volcengine"
|
||||
row.pm_model_key = "doubao-seed-1-6-250615"
|
||||
row.pm_api_key = "test-key"
|
||||
row.pm_api_base = "https://ark.test.com"
|
||||
row.pm_api_version = None
|
||||
row.pm_status = "active"
|
||||
row.lm_id = None
|
||||
row.fm_id = None
|
||||
return row
|
||||
|
||||
class _FakeDoubaoClient:
|
||||
"""Fake DoubaoClient for testing - mimics the real interface."""
|
||||
def __init__(self, api_key="", base_url="", model="", timeout=0, max_retries=0,
|
||||
max_tokens=None, temperature=None, extra_params=None, provider="volcengine"):
|
||||
self.api_key = api_key
|
||||
self.base_url = base_url
|
||||
self.model = model
|
||||
self.timeout = timeout
|
||||
self.max_retries = max_retries
|
||||
self.max_tokens = max_tokens
|
||||
self.temperature = temperature
|
||||
self.extra_params = extra_params or {}
|
||||
self.provider = provider
|
||||
self.vision_model = model
|
||||
|
||||
@property
|
||||
def is_available(self):
|
||||
return bool(self.api_key)
|
||||
def _model_config(**overrides):
|
||||
kwargs = dict(
|
||||
id="m1",
|
||||
name="test",
|
||||
provider="volcengine",
|
||||
model_key="test-model",
|
||||
api_key="key",
|
||||
api_base="https://test.com",
|
||||
api_version=None,
|
||||
status="active",
|
||||
)
|
||||
kwargs.update(overrides)
|
||||
return ai_router_mod.ModelConfig(**kwargs)
|
||||
|
||||
def chat_completion(self, messages, **kwargs):
|
||||
return None
|
||||
|
||||
def vision_completion(self, messages, **kwargs):
|
||||
return None
|
||||
def _capability_config(**overrides):
|
||||
kwargs = dict(
|
||||
capability_key="test",
|
||||
capability_name="test",
|
||||
primary_model=None,
|
||||
lite_model=None,
|
||||
fallback_model=None,
|
||||
timeout_seconds=30,
|
||||
max_retries=1,
|
||||
max_tokens=None,
|
||||
temperature=None,
|
||||
concurrency=2,
|
||||
extra_params={},
|
||||
is_enabled=True,
|
||||
)
|
||||
kwargs.update(overrides)
|
||||
return ai_router_mod.CapabilityConfig(**kwargs)
|
||||
|
||||
_mock_ai_client.DoubaoClient = _FakeDoubaoClient
|
||||
sys.modules["packages.shared.ai_client"] = _mock_ai_client
|
||||
|
||||
_ai_router = _load_module_from_file(
|
||||
"packages.shared.ai_router",
|
||||
os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "packages", "shared", "ai_router.py"),
|
||||
)
|
||||
# ── Redis 版本号机制 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestAIConfigVersion(unittest.TestCase):
|
||||
"""Redis 版本号机制测试"""
|
||||
|
||||
@patch.object(_ai_config_version, "_get_redis_client")
|
||||
@patch.object(_config_version_mod, "_get_redis_client")
|
||||
def test_bump_version_success(self, mock_redis_fn):
|
||||
mock_r = MagicMock()
|
||||
mock_r.set.return_value = True
|
||||
mock_redis_fn.return_value = mock_r
|
||||
ver = _ai_config_version.bump_version()
|
||||
ver = _config_version_mod.bump_version()
|
||||
self.assertTrue(ver)
|
||||
self.assertTrue(ver.isdigit())
|
||||
mock_r.set.assert_called_once()
|
||||
|
||||
@patch.object(_ai_config_version, "_get_redis_client")
|
||||
@patch.object(_config_version_mod, "_get_redis_client")
|
||||
def test_bump_version_redis_unavailable(self, mock_redis_fn):
|
||||
mock_redis_fn.return_value = None
|
||||
ver = _ai_config_version.bump_version()
|
||||
ver = _config_version_mod.bump_version()
|
||||
self.assertEqual(ver, "")
|
||||
|
||||
@patch.object(_ai_config_version, "_get_redis_client")
|
||||
@patch.object(_config_version_mod, "_get_redis_client")
|
||||
def test_get_version_success(self, mock_redis_fn):
|
||||
mock_r = MagicMock()
|
||||
mock_r.get.return_value = "1234567890"
|
||||
mock_redis_fn.return_value = mock_r
|
||||
ver = _ai_config_version.get_version()
|
||||
ver = _config_version_mod.get_version()
|
||||
self.assertEqual(ver, "1234567890")
|
||||
|
||||
@patch.object(_ai_config_version, "_get_redis_client")
|
||||
@patch.object(_config_version_mod, "_get_redis_client")
|
||||
def test_get_version_redis_down(self, mock_redis_fn):
|
||||
mock_redis_fn.return_value = None
|
||||
ver = _ai_config_version.get_version()
|
||||
ver = _config_version_mod.get_version()
|
||||
self.assertIsNone(ver)
|
||||
|
||||
@patch.object(_ai_config_version, "_get_redis_client")
|
||||
@patch.object(_config_version_mod, "_get_redis_client")
|
||||
def test_get_version_exception(self, mock_redis_fn):
|
||||
mock_r = MagicMock()
|
||||
mock_r.get.side_effect = Exception("connection refused")
|
||||
mock_redis_fn.return_value = mock_r
|
||||
ver = _ai_config_version.get_version()
|
||||
ver = _config_version_mod.get_version()
|
||||
self.assertIsNone(ver)
|
||||
|
||||
|
||||
# ── AIRouter 路由/缓存/fallback ────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestAIRouter(unittest.TestCase):
|
||||
"""AIRouter 路由/缓存/fallback 测试"""
|
||||
|
||||
def setUp(self):
|
||||
self.router = _ai_router.AIRouter()
|
||||
self.router = ai_router_mod.AIRouter()
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_capability_db_unavailable(self, mock_ver):
|
||||
with patch.object(_ai_router, "_get_session", return_value=None):
|
||||
cap = self.router.get_capability("intent_parsing")
|
||||
self.assertIsNone(cap)
|
||||
def _freeze_version(self, value=None):
|
||||
"""让 get_capability 的版本比对固定,避免走 Redis。"""
|
||||
return patch.object(_config_version_mod, "get_version", return_value=value)
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_capability_from_db(self, mock_ver):
|
||||
def test_get_capability_db_unavailable(self):
|
||||
with self._freeze_version(None):
|
||||
with patch.object(ai_router_mod, "_get_session", return_value=None):
|
||||
cap = self.router.get_capability("intent_parsing")
|
||||
self.assertIsNone(cap)
|
||||
|
||||
def test_get_capability_from_db(self):
|
||||
mock_session = MagicMock()
|
||||
mock_row = MagicMock()
|
||||
mock_row.capability_key = "intent_parsing"
|
||||
mock_row.capability_name = "文案意图解析"
|
||||
mock_row.timeout_seconds = 45
|
||||
mock_row.max_retries = 1
|
||||
mock_row.max_tokens = None
|
||||
mock_row.temperature = None
|
||||
mock_row.concurrency = 2
|
||||
mock_row.extra_params = {}
|
||||
mock_row.is_enabled = True
|
||||
mock_row.pm_id = "model-1"
|
||||
mock_row.pm_name = "豆包"
|
||||
mock_row.pm_provider = "volcengine"
|
||||
mock_row.pm_model_key = "doubao-seed-1-6-250615"
|
||||
mock_row.pm_api_key = "test-key"
|
||||
mock_row.pm_api_base = "https://ark.test.com"
|
||||
mock_row.pm_api_version = None
|
||||
mock_row.pm_status = "active"
|
||||
mock_row.lm_id = None
|
||||
mock_row.fm_id = None
|
||||
mock_session.execute.return_value.first.return_value = mock_row
|
||||
mock_session.execute.return_value.first.return_value = _capability_row()
|
||||
with self._freeze_version(None):
|
||||
with patch.object(ai_router_mod, "_get_session", return_value=mock_session):
|
||||
cap = self.router.get_capability("intent_parsing")
|
||||
self.assertIsNotNone(cap)
|
||||
self.assertEqual(cap.capability_key, "intent_parsing")
|
||||
self.assertEqual(cap.primary_model.model_key, "doubao-seed-1-6-250615")
|
||||
|
||||
with patch.object(_ai_router, "_get_session", return_value=mock_session):
|
||||
cap = self.router.get_capability("intent_parsing")
|
||||
self.assertIsNotNone(cap)
|
||||
self.assertEqual(cap.capability_key, "intent_parsing")
|
||||
self.assertEqual(cap.primary_model.model_key, "doubao-seed-1-6-250615")
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", side_effect=[None, "v2"])
|
||||
def test_cache_invalidation_on_version_change(self, mock_ver):
|
||||
with patch.object(self.router, "_load_from_db", return_value=None):
|
||||
self.router.get_capability("test_key")
|
||||
def test_cache_invalidation_on_version_change(self):
|
||||
with self._freeze_version(None):
|
||||
with patch.object(self.router, "_load_from_db", return_value=None):
|
||||
self.router.get_capability("test_key")
|
||||
self.router._local_ver = "v1"
|
||||
self.assertTrue(self.router._check_version())
|
||||
with patch.object(_config_version_mod, "get_version", return_value="v2"):
|
||||
self.assertTrue(self.router._check_version())
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value="same_ver")
|
||||
def test_cache_hit_same_version(self, mock_ver):
|
||||
model = _ai_router.ModelConfig(
|
||||
id="m1", name="test", provider="volcengine", model_key="test-model",
|
||||
api_key="key", api_base="https://test.com", api_version=None, status="active",
|
||||
)
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="test", capability_name="test", primary_model=model,
|
||||
lite_model=None, fallback_model=None, timeout_seconds=30,
|
||||
max_retries=1, max_tokens=None, temperature=None, concurrency=2,
|
||||
extra_params={}, is_enabled=True,
|
||||
def test_cache_hit_same_version(self):
|
||||
cap = _capability_config(
|
||||
primary_model=_model_config(model_key="test-model"),
|
||||
)
|
||||
self.router._cache["test"] = cap
|
||||
self.router._local_ver = "same_ver"
|
||||
result = self.router.get_capability("test")
|
||||
with patch.object(_config_version_mod, "get_version", return_value="same_ver"):
|
||||
result = self.router.get_capability("test")
|
||||
self.assertEqual(result, cap)
|
||||
|
||||
def test_invalidate_clears_cache(self):
|
||||
@@ -213,181 +214,163 @@ class TestAIRouter(unittest.TestCase):
|
||||
self.assertEqual(len(self.router._cache), 0)
|
||||
self.assertIsNone(self.router._local_ver)
|
||||
|
||||
@patch.object(_ai_router, "_get_session", return_value=None)
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_llm_client_fallback(self, mock_ver, mock_session):
|
||||
_ai_router.get_shared_settings = lambda: _mock_settings
|
||||
client = self.router.get_llm_client("intent_parsing")
|
||||
def test_get_llm_client_fallback(self):
|
||||
with self._freeze_version(None):
|
||||
with patch.object(ai_router_mod, "_get_session", return_value=None):
|
||||
client = self.router.get_llm_client("intent_parsing")
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "doubao-seed-2-1-pro-260915")
|
||||
self.assertEqual(client.api_key, "test-key")
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_llm_client_from_db(self, mock_ver):
|
||||
model = _ai_router.ModelConfig(
|
||||
id="m1", name="test", provider="dashscope", model_key="qwen3.8-flash",
|
||||
api_key="db-key", api_base="https://dashscope.test.com", api_version=None, status="active",
|
||||
)
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="image_analysis", capability_name="图片分析",
|
||||
primary_model=model, lite_model=None, fallback_model=None,
|
||||
timeout_seconds=15, max_retries=1, max_tokens=350, temperature=0.1,
|
||||
concurrency=2, extra_params={}, is_enabled=True,
|
||||
def test_get_llm_client_from_db(self):
|
||||
cap = _capability_config(
|
||||
capability_key="image_analysis",
|
||||
capability_name="图片分析",
|
||||
max_tokens=350,
|
||||
temperature=0.1,
|
||||
primary_model=_model_config(
|
||||
provider="dashscope",
|
||||
model_key="qwen3.8-flash",
|
||||
api_key="db-key",
|
||||
api_base="https://dashscope.test.com",
|
||||
),
|
||||
)
|
||||
with patch.object(self.router, "get_capability", return_value=cap):
|
||||
client = self.router.get_llm_client("image_analysis")
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "qwen3.8-flash")
|
||||
self.assertEqual(client.provider, "dashscope")
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "qwen3.8-flash")
|
||||
self.assertEqual(client.provider, "dashscope")
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_vision_client(self, mock_ver):
|
||||
model = _ai_router.ModelConfig(
|
||||
id="m1", name="test", provider="dashscope", model_key="qwen3.8-flash",
|
||||
api_key="key", api_base="https://dashscope.test.com", api_version=None, status="active",
|
||||
)
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="image_analysis", capability_name="图片分析",
|
||||
primary_model=model, lite_model=None, fallback_model=None,
|
||||
timeout_seconds=15, max_retries=1, max_tokens=None, temperature=None,
|
||||
concurrency=2, extra_params={}, is_enabled=True,
|
||||
def test_get_vision_client(self):
|
||||
cap = _capability_config(
|
||||
capability_key="image_analysis",
|
||||
capability_name="图片分析",
|
||||
primary_model=_model_config(
|
||||
provider="dashscope",
|
||||
model_key="qwen3.8-flash",
|
||||
api_key="key",
|
||||
api_base="https://dashscope.test.com",
|
||||
),
|
||||
)
|
||||
with patch.object(self.router, "get_capability", return_value=cap):
|
||||
client = self.router.get_vision_client("image_analysis")
|
||||
self.assertIsNotNone(client)
|
||||
# #2220: vision client is now DoubaoClient with vision_completion
|
||||
self.assertTrue(hasattr(client, "vision_completion"))
|
||||
self.assertIsNotNone(client)
|
||||
# #2220: vision client is now DoubaoClient with vision_completion
|
||||
self.assertTrue(hasattr(client, "vision_completion"))
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_tts_client(self, mock_ver):
|
||||
model = _ai_router.ModelConfig(
|
||||
id="m1", name="test", provider="dashscope", model_key="cosyvoice-v3-flash",
|
||||
api_key="key", api_base="https://dashscope.test.com", api_version=None, status="active",
|
||||
)
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="tts", capability_name="语音合成",
|
||||
primary_model=model, lite_model=None, fallback_model=None,
|
||||
timeout_seconds=60, max_retries=1, max_tokens=None, temperature=None,
|
||||
concurrency=2, extra_params={}, is_enabled=True,
|
||||
def test_get_tts_client(self):
|
||||
cap = _capability_config(
|
||||
capability_key="tts",
|
||||
capability_name="语音合成",
|
||||
primary_model=_model_config(
|
||||
provider="dashscope",
|
||||
model_key="cosyvoice-v3-flash",
|
||||
api_key="key",
|
||||
api_base="https://dashscope.test.com",
|
||||
),
|
||||
)
|
||||
with patch.object(self.router, "get_capability", return_value=cap):
|
||||
client = self.router.get_tts_client()
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "cosyvoice-v3-flash")
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "cosyvoice-v3-flash")
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_image_gen_client(self, mock_ver):
|
||||
model = _ai_router.ModelConfig(
|
||||
id="m1", name="test", provider="volcengine", model_key="seedream-5.0-flash",
|
||||
api_key="key", api_base="https://ark.test.com", api_version=None, status="active",
|
||||
)
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="image_generation", capability_name="图片生成",
|
||||
primary_model=model, lite_model=None, fallback_model=None,
|
||||
timeout_seconds=60, max_retries=1, max_tokens=None, temperature=None,
|
||||
concurrency=2, extra_params={"size": "1K"}, is_enabled=True,
|
||||
def test_get_image_gen_client(self):
|
||||
cap = _capability_config(
|
||||
capability_key="image_generation",
|
||||
capability_name="图片生成",
|
||||
extra_params={"size": "1K"},
|
||||
primary_model=_model_config(
|
||||
model_key="seedream-5.0-flash",
|
||||
api_key="key",
|
||||
api_base="https://ark.test.com",
|
||||
),
|
||||
)
|
||||
with patch.object(self.router, "get_capability", return_value=cap):
|
||||
client = self.router.get_image_gen_client()
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "seedream-5.0-flash")
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "seedream-5.0-flash")
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_video_gen_client(self, mock_ver):
|
||||
model = _ai_router.ModelConfig(
|
||||
id="m1", name="test", provider="volcengine", model_key="seedance-2.5",
|
||||
api_key="key", api_base="https://ark.test.com", api_version=None, status="active",
|
||||
)
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="video_generation", capability_name="视频生成",
|
||||
primary_model=model, lite_model=None, fallback_model=None,
|
||||
timeout_seconds=600, max_retries=1, max_tokens=None, temperature=None,
|
||||
concurrency=1, extra_params={}, is_enabled=True,
|
||||
def test_get_video_gen_client(self):
|
||||
cap = _capability_config(
|
||||
capability_key="video_generation",
|
||||
capability_name="视频生成",
|
||||
concurrency=1,
|
||||
primary_model=_model_config(
|
||||
model_key="seedance-2.5",
|
||||
api_key="key",
|
||||
api_base="https://ark.test.com",
|
||||
),
|
||||
)
|
||||
with patch.object(self.router, "get_capability", return_value=cap):
|
||||
client = self.router.get_video_gen_client()
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "seedance-2.5")
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "seedance-2.5")
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_lite_variant_preference(self, mock_ver):
|
||||
primary = _ai_router.ModelConfig(id="p1", name="pro", provider="volcengine", model_key="pro-model", api_key="k", api_base="u", api_version=None, status="active")
|
||||
lite = _ai_router.ModelConfig(id="l1", name="lite", provider="volcengine", model_key="lite-model", api_key="k", api_base="u", api_version=None, status="active")
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="image_analysis", capability_name="图片分析",
|
||||
primary_model=primary, lite_model=lite, fallback_model=None,
|
||||
timeout_seconds=15, max_retries=1, max_tokens=None, temperature=None,
|
||||
concurrency=2, extra_params={}, is_enabled=True,
|
||||
def test_lite_variant_preference(self):
|
||||
cap = _capability_config(
|
||||
capability_key="image_analysis",
|
||||
capability_name="图片分析",
|
||||
primary_model=_model_config(id="p1", name="pro", model_key="pro-model", api_key="k", api_base="u"),
|
||||
lite_model=_model_config(id="l1", name="lite", model_key="lite-model", api_key="k", api_base="u"),
|
||||
)
|
||||
model = self.router._get_model_or_fallback(cap, "lite")
|
||||
self.assertEqual(model.model_key, "lite-model")
|
||||
model_primary = self.router._get_model_or_fallback(cap, "primary")
|
||||
self.assertEqual(model_primary.model_key, "pro-model")
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_disabled_capability_returns_fallback(self, mock_ver):
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="test", capability_name="test",
|
||||
primary_model=None, lite_model=None, fallback_model=None,
|
||||
timeout_seconds=30, max_retries=1, max_tokens=None, temperature=None,
|
||||
concurrency=2, extra_params={}, is_enabled=False,
|
||||
)
|
||||
_ai_router.get_shared_settings = lambda: _mock_settings
|
||||
def test_disabled_capability_returns_fallback(self):
|
||||
cap = _capability_config(is_enabled=False)
|
||||
with patch.object(self.router, "get_capability", return_value=cap):
|
||||
client = self.router.get_llm_client("test")
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "doubao-seed-2-1-pro-260915")
|
||||
self.assertIsNotNone(client)
|
||||
self.assertEqual(client.model, "doubao-seed-2-1-pro-260915")
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_fallback_chain_primary_none(self, mock_ver):
|
||||
def test_fallback_chain_primary_none(self):
|
||||
"""primary_model 为 None 时 fallback 到 fallback_model"""
|
||||
fb = _ai_router.ModelConfig(id="f1", name="fb", provider="volcengine", model_key="fb-model", api_key="k", api_base="u", api_version=None, status="active")
|
||||
cap = _ai_router.CapabilityConfig(
|
||||
capability_key="test", capability_name="test",
|
||||
primary_model=None, lite_model=None, fallback_model=fb,
|
||||
timeout_seconds=30, max_retries=1, max_tokens=None, temperature=None,
|
||||
concurrency=2, extra_params={}, is_enabled=True,
|
||||
cap = _capability_config(
|
||||
fallback_model=_model_config(id="f1", name="fb", model_key="fb-model", api_key="k", api_base="u"),
|
||||
)
|
||||
model = self.router._get_model_or_fallback(cap, "primary")
|
||||
self.assertEqual(model.model_key, "fb-model")
|
||||
|
||||
|
||||
# ── 数据类冻结 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestModelConfig(unittest.TestCase):
|
||||
"""数据类测试"""
|
||||
|
||||
def test_model_config_frozen(self):
|
||||
m = _ai_router.ModelConfig(id="1", name="t", provider="p", model_key="k", api_key="a", api_base="b", api_version=None, status="active")
|
||||
m = _model_config(id="1", name="t", provider="p", model_key="k", api_key="a", api_base="b")
|
||||
with self.assertRaises(AttributeError):
|
||||
m.model_key = "new"
|
||||
|
||||
def test_capability_config_frozen(self):
|
||||
c = _ai_router.CapabilityConfig(
|
||||
capability_key="k", capability_name="n", primary_model=None,
|
||||
lite_model=None, fallback_model=None, timeout_seconds=30,
|
||||
max_retries=1, max_tokens=None, temperature=None, concurrency=2,
|
||||
extra_params={}, is_enabled=True,
|
||||
)
|
||||
c = _capability_config()
|
||||
with self.assertRaises(AttributeError):
|
||||
c.is_enabled = False
|
||||
|
||||
|
||||
# ── 客户端可用性 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestClientAvailability(unittest.TestCase):
|
||||
"""客户端可用性测试"""
|
||||
|
||||
def test_tts_client_available(self):
|
||||
c = _ai_router.TTSClient(provider="p", api_key="k", base_url="u", model="m")
|
||||
c = ai_router_mod.TTSClient(provider="p", api_key="k", base_url="u", model="m")
|
||||
self.assertTrue(c.is_available)
|
||||
|
||||
def test_tts_client_unavailable_no_model(self):
|
||||
c = _ai_router.TTSClient(provider="p", api_key="k", base_url="u", model="")
|
||||
c = ai_router_mod.TTSClient(provider="p", api_key="k", base_url="u", model="")
|
||||
self.assertFalse(c.is_available)
|
||||
|
||||
def test_image_gen_client_unavailable_no_url(self):
|
||||
c = _ai_router.ImageGenClient(provider="p", api_key="k", base_url="", model="m")
|
||||
c = ai_router_mod.ImageGenClient(provider="p", api_key="k", base_url="", model="m")
|
||||
self.assertFalse(c.is_available)
|
||||
|
||||
def test_video_gen_client_available(self):
|
||||
c = _ai_router.VideoGenClient(provider="p", api_key="k", base_url="u", model="m")
|
||||
c = ai_router_mod.VideoGenClient(provider="p", api_key="k", base_url="u", model="m")
|
||||
self.assertTrue(c.is_available)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,225 @@
|
||||
"""Ditto LLM 情绪分析服务单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.ditto_emotion_service import (
|
||||
EMO_HAPPY,
|
||||
EMO_NEUTRAL,
|
||||
DittoEmotionService,
|
||||
EmotionSegment,
|
||||
_parse_emotion_json,
|
||||
align_timeline_by_length,
|
||||
align_timeline_by_timings,
|
||||
split_sentences,
|
||||
)
|
||||
|
||||
|
||||
# ── 分句 ─────────────────────────────────────────────────────────
|
||||
class TestSplitSentences:
|
||||
def test_empty(self):
|
||||
assert split_sentences("") == []
|
||||
|
||||
def test_single(self):
|
||||
assert split_sentences("你好。") == ["你好。"]
|
||||
|
||||
def test_multi(self):
|
||||
sents = split_sentences("大家好!今天给大家推荐一款超棒的产品。它真的很好用;不信你试试?")
|
||||
assert len(sents) == 4
|
||||
assert "大家好!" in sents[0]
|
||||
|
||||
def test_english_punct(self):
|
||||
sents = split_sentences("Hello! How are you? I'm fine.")
|
||||
assert len(sents) == 3
|
||||
|
||||
|
||||
# ── JSON 解析 ────────────────────────────────────────────────────
|
||||
class TestParseEmotionJson:
|
||||
def test_valid(self):
|
||||
raw = json.dumps([{"text": "你好", "emo": 4, "intensity": 0.2}])
|
||||
segs = _parse_emotion_json(raw)
|
||||
assert len(segs) == 1
|
||||
assert segs[0].emo == 4
|
||||
assert segs[0].intensity == 0.2
|
||||
|
||||
def test_markdown_wrapped(self):
|
||||
raw = "```json\n" + json.dumps([{"text": "好", "emo": 3, "intensity": 0.5}]) + "\n```"
|
||||
segs = _parse_emotion_json(raw)
|
||||
assert len(segs) == 1
|
||||
assert segs[0].emo == 3
|
||||
|
||||
def test_forbidden_emo_becomes_neutral(self):
|
||||
raw = json.dumps([{"text": "怒", "emo": 0, "intensity": 0.8}])
|
||||
segs = _parse_emotion_json(raw)
|
||||
assert len(segs) == 1
|
||||
assert segs[0].emo == EMO_NEUTRAL
|
||||
|
||||
def test_invalid_json(self):
|
||||
assert _parse_emotion_json("not json") == []
|
||||
|
||||
def test_empty(self):
|
||||
assert _parse_emotion_json("") == []
|
||||
|
||||
def test_intensity_clamp(self):
|
||||
raw = json.dumps([{"text": "a", "emo": 3, "intensity": 1.5}])
|
||||
segs = _parse_emotion_json(raw)
|
||||
assert segs[0].intensity == 1.0
|
||||
|
||||
def test_missing_text_skipped(self):
|
||||
raw = json.dumps([{"emo": 3, "intensity": 0.4}])
|
||||
segs = _parse_emotion_json(raw)
|
||||
assert len(segs) == 0
|
||||
|
||||
|
||||
# ── 时间对齐(按字数比例)────────────────────────────────────────
|
||||
class TestAlignTimelineByLength:
|
||||
def test_basic(self):
|
||||
segs = [
|
||||
EmotionSegment("ab", EMO_NEUTRAL, 0.2),
|
||||
EmotionSegment("cd", EMO_HAPPY, 0.5),
|
||||
]
|
||||
entries = align_timeline_by_length(segs, 4.0)
|
||||
assert len(entries) == 2
|
||||
assert entries[0].start == 0.0
|
||||
assert entries[0].end == 2.0
|
||||
assert entries[1].start == 2.0
|
||||
assert entries[1].end == 4.0
|
||||
assert entries[0].emo == EMO_NEUTRAL
|
||||
assert entries[1].emo == EMO_HAPPY
|
||||
|
||||
def test_empty_segments(self):
|
||||
assert align_timeline_by_length([], 5.0) == []
|
||||
|
||||
def test_zero_duration(self):
|
||||
segs = [EmotionSegment("ab", EMO_NEUTRAL, 0.2)]
|
||||
assert align_timeline_by_length(segs, 0) == []
|
||||
|
||||
def test_unequal_length(self):
|
||||
segs = [
|
||||
EmotionSegment("a" * 3, EMO_HAPPY, 0.5),
|
||||
EmotionSegment("b" * 1, EMO_NEUTRAL, 0.2),
|
||||
]
|
||||
entries = align_timeline_by_length(segs, 4.0)
|
||||
assert entries[0].end == 3.0
|
||||
assert entries[1].start == 3.0
|
||||
assert entries[1].end == 4.0
|
||||
|
||||
|
||||
# ── 时间对齐(sentence_timings)──────────────────────────────────
|
||||
class TestAlignTimelineByTimings:
|
||||
def test_exact_match(self):
|
||||
segs = [
|
||||
EmotionSegment("hello", EMO_HAPPY, 0.4),
|
||||
EmotionSegment("world", EMO_NEUTRAL, 0.2),
|
||||
]
|
||||
timings = [
|
||||
{"start": 0.0, "end": 1.5},
|
||||
{"start": 1.5, "end": 3.0},
|
||||
]
|
||||
entries = align_timeline_by_timings(segs, timings, 3.0)
|
||||
assert len(entries) == 2
|
||||
assert entries[0].start == 0.0
|
||||
assert entries[0].end == 1.5
|
||||
assert entries[1].start == 1.5
|
||||
assert entries[1].end == 3.0
|
||||
|
||||
def test_length_mismatch_fallback(self):
|
||||
segs = [EmotionSegment("hello", EMO_HAPPY, 0.4)]
|
||||
timings = [{"start": 0, "end": 1}, {"start": 1, "end": 2}]
|
||||
entries = align_timeline_by_timings(segs, timings, 2.0)
|
||||
assert len(entries) == 1
|
||||
assert entries[0].end == 2.0
|
||||
|
||||
|
||||
# ── DittoEmotionService ──────────────────────────────────────────
|
||||
def _make_service(enabled=True, model=None, temperature=0.1, timeout=10, max_tokens=1024, prompt=""):
|
||||
s = MagicMock()
|
||||
s.ditto_emotion_enabled = enabled
|
||||
s.ditto_emotion_model = model or ""
|
||||
s.ditto_emotion_temperature = temperature
|
||||
s.ditto_emotion_timeout = timeout
|
||||
s.ditto_emotion_max_tokens = max_tokens
|
||||
s.ditto_emotion_cache_size = 100
|
||||
s.ditto_emotion_prompt = prompt
|
||||
return DittoEmotionService(settings=s)
|
||||
|
||||
|
||||
class TestDittoEmotionService:
|
||||
def test_disabled_returns_empty(self):
|
||||
svc = _make_service(enabled=False)
|
||||
assert svc.analyze("你好世界") == []
|
||||
|
||||
def test_empty_text_returns_empty(self):
|
||||
svc = _make_service(enabled=True)
|
||||
assert svc.analyze("") == []
|
||||
|
||||
def test_llm_success(self):
|
||||
svc = _make_service(enabled=True)
|
||||
fake_reply = json.dumps([{"text": "你好", "emo": 4, "intensity": 0.2}])
|
||||
with patch.object(svc, "_call_llm", return_value=_parse_emotion_json(fake_reply)):
|
||||
segs = svc.analyze("你好")
|
||||
assert len(segs) == 1
|
||||
assert segs[0].emo == 4
|
||||
|
||||
def test_cache_hit(self):
|
||||
svc = _make_service(enabled=True)
|
||||
fake_reply = json.dumps([{"text": "你好世界", "emo": 3, "intensity": 0.5}])
|
||||
with patch.object(svc, "_call_llm", return_value=_parse_emotion_json(fake_reply)) as mock_call:
|
||||
svc.analyze("你好世界")
|
||||
svc.analyze("你好世界")
|
||||
assert mock_call.call_count == 1
|
||||
|
||||
def test_build_timeline_empty_when_disabled(self):
|
||||
svc = _make_service(enabled=False)
|
||||
assert svc.build_timeline("test", 5.0) == ""
|
||||
|
||||
def test_build_timeline_returns_json(self):
|
||||
svc = _make_service(enabled=True)
|
||||
fake_reply = json.dumps(
|
||||
[
|
||||
{"text": "ab", "emo": 4, "intensity": 0.2},
|
||||
{"text": "cd", "emo": 3, "intensity": 0.4},
|
||||
]
|
||||
)
|
||||
with patch.object(svc, "_call_llm", return_value=_parse_emotion_json(fake_reply)):
|
||||
result = svc.build_timeline("ab。cd。", 4.0)
|
||||
data = json.loads(result)
|
||||
assert len(data) == 2
|
||||
assert data[0]["emo"] == 4
|
||||
assert data[1]["emo"] == 3
|
||||
|
||||
def test_build_timeline_with_sentence_timings(self):
|
||||
svc = _make_service(enabled=True)
|
||||
fake_reply = json.dumps(
|
||||
[
|
||||
{"text": "hello", "emo": 3, "intensity": 0.4},
|
||||
{"text": "world", "emo": 4, "intensity": 0.2},
|
||||
]
|
||||
)
|
||||
timings = [
|
||||
{"start": 0.0, "end": 1.0},
|
||||
{"start": 1.0, "end": 3.0},
|
||||
]
|
||||
with patch.object(svc, "_call_llm", return_value=_parse_emotion_json(fake_reply)):
|
||||
result = svc.build_timeline("hello world", 3.0, sentence_timings=timings)
|
||||
data = json.loads(result)
|
||||
assert data[0]["start"] == 0.0
|
||||
assert data[0]["end"] == 1.0
|
||||
assert data[1]["end"] == 3.0
|
||||
|
||||
|
||||
class TestPromptLoading:
|
||||
def test_default_prompt_contains_placeholder(self):
|
||||
from packages.application.ditto_emotion_service import _load_default_prompt
|
||||
|
||||
prompt = _load_default_prompt()
|
||||
assert "{文案}" in prompt
|
||||
|
||||
def test_config_prompt_override(self):
|
||||
custom = "分析情绪: {文案}"
|
||||
svc = _make_service(enabled=True, prompt=custom)
|
||||
assert svc._get_prompt_template() == custom
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
@@ -195,3 +196,34 @@ def test_empty_script_replaced_with_space():
|
||||
c.generate(audio_url="http://x/a.mp3", script="")
|
||||
payload = client.post.call_args.kwargs["json"]
|
||||
assert payload["script"] == " "
|
||||
|
||||
|
||||
def test_generate_network_error_fails_fast(monkeypatch):
|
||||
"""网络不通(ConnectError)时不重试,直接快速抛 NetworkUnreachable,避免用户等5分钟"""
|
||||
import httpx
|
||||
|
||||
from packages.application import ditto_service as ds_mod
|
||||
|
||||
calls = {"n": 0}
|
||||
|
||||
def _fake_post(self, url, json=None):
|
||||
calls["n"] += 1
|
||||
raise httpx.ConnectError("[Errno 113] No route to host")
|
||||
|
||||
monkeypatch.setattr(httpx.Client, "post", _fake_post)
|
||||
|
||||
client = _make_client(
|
||||
base_url="http://100.76.80.23:8000",
|
||||
default_video_url="http://oss/tpl.mp4",
|
||||
max_retries=2,
|
||||
timeout=120,
|
||||
)
|
||||
|
||||
t0 = time.monotonic()
|
||||
with pytest.raises(ds_mod.DittoError) as exc:
|
||||
client.generate(audio_url="http://oss/a.wav", script="你好")
|
||||
elapsed = time.monotonic() - t0
|
||||
|
||||
assert exc.value.code == "NetworkUnreachable"
|
||||
assert calls["n"] == 1 # 不重试
|
||||
assert elapsed < 5 # 快速失败<5秒
|
||||
|
||||
@@ -286,7 +286,7 @@ class TestGetDownloadUrl:
|
||||
|
||||
result = svc.get_download_url("uploads/video.mp4")
|
||||
|
||||
svc.bucket.sign_url.assert_called_once_with("GET", "uploads/video.mp4", 3600)
|
||||
svc.bucket.sign_url.assert_called_once_with("GET", "uploads/video.mp4", 3600, slash_safe=True)
|
||||
assert "signed-url" in result
|
||||
|
||||
def test_returns_raw_url_when_bucket_none(self):
|
||||
|
||||
@@ -250,3 +250,34 @@ class TestMultiplierConsistency:
|
||||
resp = check_points(body=body, current_user=cu, db=db)
|
||||
expected = calculate_points_cost(scene, is_member=False, quantity=1, duration_minutes=1)
|
||||
assert resp.required_points == expected, f"{scene}: got {resp.required_points}, expected {expected}"
|
||||
|
||||
# ── GET /points/balance 返回后端真实积分开关 credits_enabled ──────────
|
||||
|
||||
|
||||
class TestBalanceCreditsEnabled:
|
||||
def _call_balance(self, enabled):
|
||||
from app.api.routes import points as points_routes
|
||||
|
||||
svc = MagicMock()
|
||||
svc.get_or_create_account.return_value = {
|
||||
"balance": 0,
|
||||
"total_earned": 0,
|
||||
"total_spent": 0,
|
||||
}
|
||||
cu = _make_cu()
|
||||
db = MagicMock()
|
||||
with (
|
||||
patch("app.api.routes.points._get_service", return_value=svc),
|
||||
patch("app.api.routes.points._credits_enabled", return_value=enabled),
|
||||
):
|
||||
return points_routes.get_balance(current_user=cu, db=db)
|
||||
|
||||
def test_balance_credits_enabled_false_when_free_pass(self):
|
||||
"""免费期(开关关闭)时 credits_enabled=False,前端应跳过余额预校验。"""
|
||||
resp = self._call_balance(False)
|
||||
assert resp.credits_enabled is False
|
||||
|
||||
def test_balance_credits_enabled_true_when_enabled(self):
|
||||
"""收费期(开关开启)时 credits_enabled=True。"""
|
||||
resp = self._call_balance(True)
|
||||
assert resp.credits_enabled is True
|
||||
|
||||
@@ -361,7 +361,7 @@ class TestGetDownloadUrlFallback:
|
||||
|
||||
result = service.get_download_url("videos/test.mp4", expires_seconds=7200)
|
||||
|
||||
mock_bucket.sign_url.assert_called_once_with("GET", "videos/test.mp4", 7200)
|
||||
mock_bucket.sign_url.assert_called_once_with("GET", "videos/test.mp4", 7200, slash_safe=True)
|
||||
assert result == "https://signed-url.com/file?sig=abc"
|
||||
|
||||
def test_sign_url_exception_falls_back_to_public_url(self):
|
||||
|
||||
@@ -0,0 +1,287 @@
|
||||
"""#2246 system_settings / system_config_service 单元测试。
|
||||
|
||||
覆盖:序列化反序列化、默认值兜底、DB 覆盖、缓存失效、白名单拒绝、
|
||||
Ditto 单例 reset 联动、PUT 校验。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import create_engine, event
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import Base
|
||||
from packages.application.system_config_service import SystemConfigService
|
||||
from packages.domain.system_setting import (
|
||||
SETTING_TYPE_BOOL,
|
||||
SETTING_TYPE_FLOAT,
|
||||
SETTING_TYPE_INT,
|
||||
SETTING_TYPE_JSON,
|
||||
SETTING_TYPE_STRING,
|
||||
SystemSettingError,
|
||||
deserialize_setting_value,
|
||||
infer_setting_type,
|
||||
serialize_setting_value,
|
||||
)
|
||||
|
||||
|
||||
# ── domain 序列化 ──────────────────────────────────────────────
|
||||
class TestSerialization:
|
||||
@pytest.mark.parametrize(
|
||||
"value,st,expected",
|
||||
[
|
||||
(True, SETTING_TYPE_BOOL, "true"),
|
||||
(False, SETTING_TYPE_BOOL, "false"),
|
||||
(12, SETTING_TYPE_INT, "12"),
|
||||
(-3, SETTING_TYPE_INT, "-3"),
|
||||
(0.1, SETTING_TYPE_FLOAT, "0.1"),
|
||||
("hello", SETTING_TYPE_STRING, "hello"),
|
||||
],
|
||||
)
|
||||
def test_serialize(self, value, st, expected):
|
||||
assert serialize_setting_value(value, st) == expected
|
||||
|
||||
def test_json_roundtrip(self):
|
||||
raw = serialize_setting_value({"a": [1, 2]}, SETTING_TYPE_JSON)
|
||||
assert deserialize_setting_value(raw, SETTING_TYPE_JSON) == {"a": [1, 2]}
|
||||
|
||||
def test_bool_roundtrip_variants(self):
|
||||
for truthy in ("true", "1", "yes", "on", "TRUE"):
|
||||
assert deserialize_setting_value(truthy, SETTING_TYPE_BOOL) is True
|
||||
for falsy in ("false", "0", "", "no"):
|
||||
assert deserialize_setting_value(falsy, SETTING_TYPE_BOOL) is False
|
||||
|
||||
def test_int_float_roundtrip(self):
|
||||
assert deserialize_setting_value("42", SETTING_TYPE_INT) == 42
|
||||
assert deserialize_setting_value("1.5", SETTING_TYPE_FLOAT) == 1.5
|
||||
|
||||
def test_type_errors(self):
|
||||
with pytest.raises(SystemSettingError):
|
||||
serialize_setting_value("x", SETTING_TYPE_INT)
|
||||
with pytest.raises(SystemSettingError):
|
||||
serialize_setting_value(1, SETTING_TYPE_BOOL)
|
||||
with pytest.raises(SystemSettingError):
|
||||
deserialize_setting_value("abc", SETTING_TYPE_INT)
|
||||
with pytest.raises(SystemSettingError):
|
||||
serialize_setting_value(1, "unknown")
|
||||
|
||||
def test_infer_type_bool_before_int(self):
|
||||
assert infer_setting_type(True) == SETTING_TYPE_BOOL
|
||||
assert infer_setting_type(5) == SETTING_TYPE_INT
|
||||
assert infer_setting_type(0.5) == SETTING_TYPE_FLOAT
|
||||
assert infer_setting_type("s") == SETTING_TYPE_STRING
|
||||
assert infer_setting_type([1]) == SETTING_TYPE_JSON
|
||||
|
||||
|
||||
# ── DB fixture ─────────────────────────────────────────────────
|
||||
@pytest.fixture()
|
||||
def config_service():
|
||||
engine = create_engine("sqlite://", echo=False)
|
||||
|
||||
@event.listens_for(engine, "connect")
|
||||
def _noop(dbapi_conn, connection_record):
|
||||
pass
|
||||
|
||||
Base.metadata.create_all(engine)
|
||||
factory = sessionmaker(bind=engine)
|
||||
yield SystemConfigService(session_factory=factory)
|
||||
|
||||
|
||||
class TestConfigService:
|
||||
def test_default_when_missing(self, config_service):
|
||||
assert config_service.get_config("not_exist", "fallback") == "fallback"
|
||||
|
||||
def test_uses_module_session_factory_assigned_after_import(self):
|
||||
# 回归:模块导入时 session.SessionLocal 为 None,initialize_database()
|
||||
# 之后才赋值;服务必须运行时读取模块属性,而不是 import 时绑定旧值。
|
||||
from packages.adapters.sqlalchemy_impl import session as db_session
|
||||
|
||||
engine = create_engine("sqlite://")
|
||||
Base.metadata.create_all(engine)
|
||||
factory = sessionmaker(bind=engine)
|
||||
writer = SystemConfigService(session_factory=factory)
|
||||
writer.set_config("late_key", 7, setting_type=SETTING_TYPE_INT)
|
||||
|
||||
reader = SystemConfigService() # 不注入工厂,依赖模块级 SessionLocal
|
||||
old = db_session.SessionLocal
|
||||
try:
|
||||
db_session.SessionLocal = factory
|
||||
assert reader.get_config("late_key", 0) == 7
|
||||
finally:
|
||||
db_session.SessionLocal = old
|
||||
|
||||
def test_raises_when_no_session_factory(self):
|
||||
from packages.adapters.sqlalchemy_impl import session as db_session
|
||||
|
||||
svc = SystemConfigService()
|
||||
old = db_session.SessionLocal
|
||||
try:
|
||||
db_session.SessionLocal = None
|
||||
with pytest.raises(RuntimeError):
|
||||
svc.get_config("anything", 1)
|
||||
finally:
|
||||
db_session.SessionLocal = old
|
||||
|
||||
def test_db_overrides_default(self, config_service):
|
||||
config_service.set_config("k", 20, setting_type=SETTING_TYPE_INT)
|
||||
# 再次读取应命中 DB 值,而非传入的默认
|
||||
assert config_service.get_config("k", 12) == 20
|
||||
|
||||
def test_set_and_get_all_types(self, config_service):
|
||||
config_service.set_config("b", True)
|
||||
config_service.set_config("i", 7)
|
||||
config_service.set_config("f", 0.3)
|
||||
config_service.set_config("s", "文本")
|
||||
config_service.set_config("j", {"x": 1})
|
||||
assert config_service.get_config("b", False) is True
|
||||
assert config_service.get_config("i", 0) == 7
|
||||
assert abs(config_service.get_config("f", 0.0) - 0.3) < 1e-9
|
||||
assert config_service.get_config("s", "") == "文本"
|
||||
assert config_service.get_config("j", {}) == {"x": 1}
|
||||
|
||||
def test_update_existing(self, config_service):
|
||||
config_service.set_config("k", 1, setting_type=SETTING_TYPE_INT)
|
||||
config_service.set_config("k", 2, setting_type=SETTING_TYPE_INT)
|
||||
assert config_service.get_config("k", 0) == 2
|
||||
assert len(config_service.list_configs()) == 1
|
||||
|
||||
def test_cache_invalidation(self, config_service):
|
||||
config_service.set_config("k", 1, setting_type=SETTING_TYPE_INT)
|
||||
config_service.get_config("k", 0) # 填充缓存
|
||||
# 手工改库模拟外部写入,reload 后应可见
|
||||
from packages.adapters.sqlalchemy_impl.models import SystemSettingModel
|
||||
|
||||
session = config_service._get_session_factory()()
|
||||
session.query(SystemSettingModel).filter(SystemSettingModel.setting_key == "k").update({"setting_value": "99"})
|
||||
session.commit()
|
||||
session.close()
|
||||
# 缓存未失效前仍是旧值
|
||||
assert config_service.get_config("k", 0) == 1
|
||||
config_service.reload()
|
||||
assert config_service.get_config("k", 0) == 99
|
||||
|
||||
def test_list_by_category(self, config_service):
|
||||
config_service.set_config("a", 1, category="ditto")
|
||||
config_service.set_config("b", 2, category="other")
|
||||
keys = {c.setting_key for c in config_service.list_configs("ditto")}
|
||||
assert keys == {"a"}
|
||||
|
||||
def test_delete(self, config_service):
|
||||
config_service.set_config("k", 1, setting_type=SETTING_TYPE_INT)
|
||||
assert config_service.delete_config("k") is True
|
||||
assert config_service.get_config("k", "d") == "d"
|
||||
assert config_service.delete_config("k") is False
|
||||
|
||||
def test_reset_singletons_called(self, config_service):
|
||||
called = {"emo": 0, "client": 0}
|
||||
|
||||
import packages.application.ditto_emotion_service as real_emo
|
||||
import packages.application.ditto_service as real_client
|
||||
|
||||
orig_emo_reset = real_emo.reset_ditto_emotion_service
|
||||
orig_client_reset = real_client.reset_ditto_client
|
||||
|
||||
def fake_emo_reset():
|
||||
called["emo"] += 1
|
||||
|
||||
def fake_client_reset():
|
||||
called["client"] += 1
|
||||
|
||||
real_emo.reset_ditto_emotion_service = fake_emo_reset
|
||||
real_client.reset_ditto_client = fake_client_reset
|
||||
try:
|
||||
config_service.set_config("ditto_emotion_model", "m", setting_type=SETTING_TYPE_STRING)
|
||||
assert called["emo"] == 1
|
||||
assert called["client"] == 1
|
||||
finally:
|
||||
# 必须还原,否则污染同文件后续 reset 用例
|
||||
real_emo.reset_ditto_emotion_service = orig_emo_reset
|
||||
real_client.reset_ditto_client = orig_client_reset
|
||||
|
||||
|
||||
# ── Ditto 单例 reset ───────────────────────────────────────────
|
||||
class TestDittoSingletonReset:
|
||||
def test_emotion_reset_clears_lru(self):
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from packages.application import ditto_emotion_service as m
|
||||
|
||||
# 用一个可哈希的假 service 填 LRU,_call_llm 返回固定值,不触发真实逻辑
|
||||
fake = MagicMock()
|
||||
fake._call_llm.return_value = []
|
||||
m._cached_analyze.cache_clear()
|
||||
m._cached_analyze(fake, "k", "t")
|
||||
assert m._cached_analyze.cache_info().currsize == 1
|
||||
|
||||
saved_singleton = m._singleton
|
||||
m._singleton = fake
|
||||
try:
|
||||
m.reset_ditto_emotion_service()
|
||||
assert m._singleton is None
|
||||
assert m._cached_analyze.cache_info().currsize == 0
|
||||
finally:
|
||||
m._singleton = saved_singleton
|
||||
|
||||
def test_client_reset(self):
|
||||
from packages.application import ditto_service as m
|
||||
|
||||
saved = m._ditto_client_singleton
|
||||
m._ditto_client_singleton = object()
|
||||
try:
|
||||
m.reset_ditto_client()
|
||||
assert m._ditto_client_singleton is None
|
||||
finally:
|
||||
m._ditto_client_singleton = saved
|
||||
|
||||
|
||||
# ── 路由层白名单与校验 ─────────────────────────────────────────
|
||||
class TestAdminRouteValidation:
|
||||
def _route(self):
|
||||
from apps.api.app.api.routes.admin import ditto_emotion as route
|
||||
|
||||
return route
|
||||
|
||||
def test_whitelist_rejects_unknown_key(self):
|
||||
route = self._route()
|
||||
from packages.application.system_config_service import SystemConfigService
|
||||
|
||||
svc = SystemConfigService(session_factory=lambda: pytest.fail("should not open session"))
|
||||
payload = route.ConfigUpdatePayload(configs={"evil_key": 1})
|
||||
# 直接调用 update_config,传入一个任意 key 标识
|
||||
result = route.update_config(payload, x_api_key="k")
|
||||
assert result["ok"] is False
|
||||
assert "不允许修改" in result["error"]
|
||||
|
||||
def test_validate_blend_frames_range(self):
|
||||
route = self._route()
|
||||
with pytest.raises(ValueError):
|
||||
route._validate_value("ditto_blend_frames", 3)
|
||||
with pytest.raises(ValueError):
|
||||
route._validate_value("ditto_blend_frames", 31)
|
||||
assert route._validate_value("ditto_blend_frames", 12) == 12
|
||||
|
||||
def test_validate_temperature_range(self):
|
||||
route = self._route()
|
||||
with pytest.raises(ValueError):
|
||||
route._validate_value("ditto_emotion_temperature", 1.5)
|
||||
assert route._validate_value("ditto_emotion_temperature", "0.2") == pytest.approx(0.2)
|
||||
|
||||
def test_validate_prompt_placeholder(self):
|
||||
route = self._route()
|
||||
with pytest.raises(ValueError):
|
||||
route._validate_value("ditto_emotion_prompt", "没有占位符的提示词")
|
||||
assert route._validate_value("ditto_emotion_prompt", "分析:{文案}") == "分析:{文案}"
|
||||
# 允许留空(使用默认)
|
||||
assert route._validate_value("ditto_emotion_prompt", "") == ""
|
||||
|
||||
def test_validate_model_options(self):
|
||||
route = self._route()
|
||||
with pytest.raises(ValueError):
|
||||
route._validate_value("ditto_emotion_model", "gpt-4")
|
||||
assert route._validate_value("ditto_emotion_model", "deepseek-v3") == "deepseek-v3"
|
||||
|
||||
def test_validate_bool(self):
|
||||
route = self._route()
|
||||
with pytest.raises(ValueError):
|
||||
route._validate_value("ditto_emotion_enabled", "yes")
|
||||
assert route._validate_value("ditto_emotion_enabled", True) is True
|
||||
@@ -361,16 +361,24 @@ class TestViralVideoPipeline:
|
||||
<transition>硬切</transition>
|
||||
<reference_image_index>0</reference_image_index>
|
||||
</clip>
|
||||
<clip image_index="1" time_range="5-15秒">
|
||||
<voiceover>颜色特别好看很显白</voiceover>
|
||||
<clip image_index="1" time_range="5-10秒">
|
||||
<voiceover>颜色特别好看</voiceover>
|
||||
<visual>特写,固定镜头</visual>
|
||||
<action_details>嘴唇涂抹特写</action_details>
|
||||
<audio_bgm>轻快BGM继续</audio_bgm>
|
||||
<transition>结束</transition>
|
||||
<transition>硬切</transition>
|
||||
<reference_image_index>1</reference_image_index>
|
||||
</clip>
|
||||
<clip image_index="2" time_range="10-15秒">
|
||||
<voiceover>很显白,推荐给大家</voiceover>
|
||||
<visual>中景,微笑展示</visual>
|
||||
<action_details>口红展示</action_details>
|
||||
<audio_bgm>轻快BGM结束</audio_bgm>
|
||||
<transition>结束</transition>
|
||||
<reference_image_index>0</reference_image_index>
|
||||
</clip>
|
||||
</clips>
|
||||
<voiceover_script>大家好,今天分享一款口红。颜色特别好看很显白</voiceover_script>
|
||||
<voiceover_script>大家好呀,今天来给大家分享一款超显白的口红。颜色特别好看很显气质,真心推荐给姐妹们</voiceover_script>
|
||||
<theme>口红分享</theme>"""
|
||||
|
||||
@pytest.fixture
|
||||
@@ -455,7 +463,7 @@ class TestViralVideoPipeline:
|
||||
assert "voiceover_script" in result
|
||||
assert "shots" in result
|
||||
assert isinstance(result["shots"], list)
|
||||
assert len(result["shots"]) == 2
|
||||
assert len(result["shots"]) == 3
|
||||
assert result["overview"]["total_duration"] == 15
|
||||
# final_copy 必须 = voiceover_script(向后兼容)
|
||||
assert result.get("final_copy") == result["voiceover_script"]
|
||||
@@ -504,7 +512,7 @@ class TestViralVideoPipeline:
|
||||
}
|
||||
prompt = _assemble_seedance_prompt(cr, mock_job)
|
||||
assert "【视频总览】" in prompt
|
||||
assert "【逐镜头时间轴】" in prompt
|
||||
assert "【分镜脚本】" in prompt
|
||||
assert "【硬性约束】" in prompt
|
||||
assert "【负面提示词】" in prompt
|
||||
assert "0-15秒" in prompt
|
||||
|
||||
@@ -129,7 +129,7 @@ class TestScriptGenerationV16:
|
||||
"negative_prompts": ["水印"],
|
||||
}
|
||||
p = _assemble_seedance_prompt(cr, mock_job)
|
||||
for key in ("【视频总览】", "【场景与光线】", "【逐镜头时间轴】", "【硬性约束】", "【负面提示词】"):
|
||||
for key in ("【视频总览】", "【参考素材】", "【分镜脚本】", "【硬性约束】", "【负面提示词】"):
|
||||
assert key in p
|
||||
|
||||
|
||||
@@ -308,8 +308,7 @@ class TestResumeReadsImageAnalysis:
|
||||
|
||||
# v1.5 改造后 resume 委托给 _run_render_pipeline,那里读取 job.image_analysis
|
||||
src = inspect.getsource(vv._run_render_pipeline)
|
||||
assert "job.image_analysis" in src
|
||||
assert "image_analysis" in src
|
||||
assert "job.copy_result" in src
|
||||
# resume 本身应该调用 _run_render_pipeline
|
||||
resume_src = inspect.getsource(vv.resume_viral_video_pipeline)
|
||||
assert "_run_render_pipeline" in resume_src
|
||||
|
||||
@@ -29,6 +29,15 @@ def _make_job(job_id: str = "job-1", user_id: str = "u1", status: str = "pending
|
||||
job.user_id = user_id
|
||||
job.status = ViralVideoStatus(status) if isinstance(status, str) else status
|
||||
job.images = kwargs.pop("images", ["img-1"])
|
||||
# confirm_copy 新增 copy_result 完整性校验:默认提供合法文案数据
|
||||
job.copy_result = kwargs.pop(
|
||||
"copy_result",
|
||||
{
|
||||
"theme": "测试主题",
|
||||
"voiceover_script": "这是一段测试口播文案内容。",
|
||||
"shots": [{"time_range": "0-15秒", "voiceover": "这是一段测试口播文案内容。"}],
|
||||
},
|
||||
)
|
||||
job.industry = kwargs.pop("industry", "电商")
|
||||
job.target_customer = kwargs.pop("target_customer", "年轻人")
|
||||
for k, v in {
|
||||
@@ -55,8 +64,8 @@ def _make_job(job_id: str = "job-1", user_id: str = "u1", status: str = "pending
|
||||
"intent_result": None,
|
||||
"image_analysis": None,
|
||||
"storyboard": None,
|
||||
"copy_result": None,
|
||||
"generated_copy_text": "",
|
||||
"language": "zh-CN",
|
||||
"voice_id": "",
|
||||
"voice_source": "",
|
||||
"voice_mode": "global",
|
||||
@@ -147,9 +156,16 @@ class TestRetryViralVideo:
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(
|
||||
job_id="job-retry2", user_id="u1", status=ViralVideoStatus.FAILED,
|
||||
duration=15, video_ratio="9:16", video_resolution="720p", video_model="seedance-2.5",
|
||||
credits_prepaid=5.0, credits_transaction_id="txn1", retry_count=0,
|
||||
job_id="job-retry2",
|
||||
user_id="u1",
|
||||
status=ViralVideoStatus.FAILED,
|
||||
duration=15,
|
||||
video_ratio="9:16",
|
||||
video_resolution="720p",
|
||||
video_model="seedance-2.5",
|
||||
credits_prepaid=5.0,
|
||||
credits_transaction_id="txn1",
|
||||
retry_count=0,
|
||||
)
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
@@ -180,24 +196,32 @@ class TestRetryViralVideo:
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(
|
||||
job_id="job-retry3a", user_id="u1", status=ViralVideoStatus.FAILED,
|
||||
duration=15, video_ratio="9:16", video_resolution="720p", video_model="seedance-2.5",
|
||||
credits_prepaid=5.0, credits_transaction_id="txn-old", retry_count=0,
|
||||
job_id="job-retry3a",
|
||||
user_id="u1",
|
||||
status=ViralVideoStatus.FAILED,
|
||||
duration=15,
|
||||
video_ratio="9:16",
|
||||
video_resolution="720p",
|
||||
video_model="seedance-2.5",
|
||||
credits_prepaid=5.0,
|
||||
credits_transaction_id="txn-old",
|
||||
retry_count=0,
|
||||
)
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
|
||||
fake_svc = MagicMock()
|
||||
fake_svc.deduct_viral_video.return_value = {"success": False, "balance": 1.0}
|
||||
req = RetryViralVideoRequest(duration=30, video_resolution="1080p", video_ratio="16:9", video_model="seedance-2.5")
|
||||
req = RetryViralVideoRequest(
|
||||
duration=30, video_resolution="1080p", video_ratio="16:9", video_model="seedance-2.5"
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch("app.config.settings") as mock_settings,
|
||||
patch("packages.domain.points_service.PointsService", return_value=fake_svc),
|
||||
patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(1920, 1080)),
|
||||
patch("packages.domain.points_rules.calculate_viral_video_credits_with_breakdown",
|
||||
return_value=(15.0, {})),
|
||||
patch("packages.domain.points_rules.calculate_viral_video_credits_with_breakdown", return_value=(15.0, {})),
|
||||
patch.object(vv_mod.celery_app, "send_task"),
|
||||
):
|
||||
mock_settings.points_enabled = True
|
||||
@@ -216,19 +240,29 @@ class TestRetryViralVideo:
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(
|
||||
job_id="job-retry3b", user_id="u1", status=ViralVideoStatus.FAILED,
|
||||
duration=15, video_ratio="9:16", video_resolution="720p", video_model="seedance-2.5",
|
||||
credits_prepaid=5.0, credits_transaction_id="txn-old", retry_count=0,
|
||||
job_id="job-retry3b",
|
||||
user_id="u1",
|
||||
status=ViralVideoStatus.FAILED,
|
||||
duration=15,
|
||||
video_ratio="9:16",
|
||||
video_resolution="720p",
|
||||
video_model="seedance-2.5",
|
||||
credits_prepaid=5.0,
|
||||
credits_transaction_id="txn-old",
|
||||
retry_count=0,
|
||||
)
|
||||
# 用 SimpleNamespace 让属性真正可写
|
||||
from types import SimpleNamespace
|
||||
|
||||
job.credits_prepaid = 5.0
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
|
||||
fake_svc = MagicMock()
|
||||
fake_svc.deduct_viral_video.return_value = {"success": True, "balance": 50.0, "transaction_id": "txn-new"}
|
||||
req = RetryViralVideoRequest(duration=30, video_resolution="1080p", video_ratio="16:9", video_model="seedance-2.5")
|
||||
req = RetryViralVideoRequest(
|
||||
duration=30, video_resolution="1080p", video_ratio="16:9", video_model="seedance-2.5"
|
||||
)
|
||||
|
||||
new_est = 15.0
|
||||
with (
|
||||
@@ -236,8 +270,9 @@ class TestRetryViralVideo:
|
||||
patch("app.config.settings") as mock_settings,
|
||||
patch("packages.domain.points_service.PointsService", return_value=fake_svc),
|
||||
patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(1920, 1080)),
|
||||
patch("packages.domain.points_rules.calculate_viral_video_credits_with_breakdown",
|
||||
return_value=(new_est, {})),
|
||||
patch(
|
||||
"packages.domain.points_rules.calculate_viral_video_credits_with_breakdown", return_value=(new_est, {})
|
||||
),
|
||||
patch.object(vv_mod.celery_app, "send_task"),
|
||||
):
|
||||
mock_settings.points_enabled = True
|
||||
@@ -262,9 +297,16 @@ class TestRetryViralVideo:
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(
|
||||
job_id="job-retry4", user_id="u1", status=ViralVideoStatus.FAILED,
|
||||
duration=20, video_ratio="16:9", video_resolution="1080p", video_model="seedance-2.5",
|
||||
credits_prepaid=10.0, credits_transaction_id="txn-old", retry_count=0,
|
||||
job_id="job-retry4",
|
||||
user_id="u1",
|
||||
status=ViralVideoStatus.FAILED,
|
||||
duration=20,
|
||||
video_ratio="16:9",
|
||||
video_resolution="1080p",
|
||||
video_model="seedance-2.5",
|
||||
credits_prepaid=10.0,
|
||||
credits_transaction_id="txn-old",
|
||||
retry_count=0,
|
||||
)
|
||||
job.credits_prepaid = 10.0
|
||||
repo = MagicMock()
|
||||
@@ -280,8 +322,9 @@ class TestRetryViralVideo:
|
||||
patch("app.config.settings") as mock_settings,
|
||||
patch("packages.domain.points_service.PointsService", return_value=fake_svc),
|
||||
patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(270, 480)),
|
||||
patch("packages.domain.points_rules.calculate_viral_video_credits_with_breakdown",
|
||||
return_value=(new_est, {})),
|
||||
patch(
|
||||
"packages.domain.points_rules.calculate_viral_video_credits_with_breakdown", return_value=(new_est, {})
|
||||
),
|
||||
patch.object(vv_mod.celery_app, "send_task"),
|
||||
):
|
||||
mock_settings.points_enabled = True
|
||||
@@ -470,7 +513,9 @@ class TestGenerateCopy:
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch.object(vv_mod.celery_app, "send_task") as mock_send,
|
||||
):
|
||||
resp = vv_mod.generate_copy(f"job-regen-{regen_status}", GenerateCopyRequest(), authenticated_user=user, session=session)
|
||||
resp = vv_mod.generate_copy(
|
||||
f"job-regen-{regen_status}", GenerateCopyRequest(), authenticated_user=user, session=session
|
||||
)
|
||||
mock_send.assert_called_once()
|
||||
job.resume_from_image_analyzed.assert_called()
|
||||
assert job.retry_count >= 1
|
||||
@@ -630,7 +675,7 @@ class TestConfirmCopyPointsDeduction:
|
||||
mock_svc.deduct_viral_video.assert_called_once()
|
||||
call_args = mock_svc.deduct_viral_video.call_args
|
||||
assert call_args.args[0] == "u1" # user_id
|
||||
assert call_args.args[1] == 5.2 # credits
|
||||
assert call_args.args[1] == 5.2 # credits
|
||||
assert call_args.args[2] == "job-pay" # job_id
|
||||
# credits_prepaid / credits_transaction_id 被写入
|
||||
assert job.credits_prepaid == 5.2
|
||||
@@ -807,9 +852,14 @@ class TestEstimateCredits:
|
||||
req = EstimateCreditsRequest(model="seedance-2.5", resolution="1080p", ratio="16:9", duration=20)
|
||||
user = _auth_user("u1")
|
||||
fake_bd = {
|
||||
"tokens": 1000.0, "video_cost": 1.0, "fixed_cost": 0.15,
|
||||
"profit_multiplier": 1.3, "model_price": 70.0,
|
||||
"width": 1920, "height": 1080, "fps": 24,
|
||||
"tokens": 1000.0,
|
||||
"video_cost": 1.0,
|
||||
"fixed_cost": 0.15,
|
||||
"profit_multiplier": 1.3,
|
||||
"model_price": 70.0,
|
||||
"width": 1920,
|
||||
"height": 1080,
|
||||
"fps": 24,
|
||||
}
|
||||
|
||||
with (
|
||||
@@ -841,9 +891,14 @@ class TestEstimateCredits:
|
||||
req = EstimateCreditsRequest(model="", resolution="720p", ratio="9:16", duration=10)
|
||||
user = _auth_user("u1")
|
||||
fake_bd = {
|
||||
"tokens": 500.0, "video_cost": 0.5, "fixed_cost": 0.15,
|
||||
"profit_multiplier": 1.3, "model_price": 70.0,
|
||||
"width": 720, "height": 1280, "fps": 24,
|
||||
"tokens": 500.0,
|
||||
"video_cost": 0.5,
|
||||
"fixed_cost": 0.15,
|
||||
"profit_multiplier": 1.3,
|
||||
"model_price": 70.0,
|
||||
"width": 720,
|
||||
"height": 1280,
|
||||
"fps": 24,
|
||||
}
|
||||
|
||||
with (
|
||||
@@ -870,9 +925,14 @@ class TestEstimateCredits:
|
||||
)
|
||||
user = _auth_user("u1")
|
||||
fake_bd = {
|
||||
"tokens": 100.0, "video_cost": 0.1, "fixed_cost": 0.15,
|
||||
"profit_multiplier": 1.3, "model_price": 46.0,
|
||||
"width": 480, "height": 480, "fps": 24,
|
||||
"tokens": 100.0,
|
||||
"video_cost": 0.1,
|
||||
"fixed_cost": 0.15,
|
||||
"profit_multiplier": 1.3,
|
||||
"model_price": 46.0,
|
||||
"width": 480,
|
||||
"height": 480,
|
||||
"fps": 24,
|
||||
}
|
||||
|
||||
with (
|
||||
|
||||
Executable
+1116
File diff suppressed because it is too large
Load Diff
@@ -483,3 +483,88 @@ class TestWSInitialSnapshot:
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
data = received[0]["data"]
|
||||
assert data == {"status": "running"}
|
||||
|
||||
|
||||
class TestScriptDeltaEvent:
|
||||
"""流式 script_delta 事件推送(worker on_delta 回调语义)测试。
|
||||
|
||||
这里不直接导入 _step_script_generation(依赖 celery/db),而是复刻 on_delta 回调
|
||||
的核心节流与事件格式逻辑,验证:
|
||||
1) 事件 type = viral_video:script_delta
|
||||
2) data 包含 delta / full_text / text_length
|
||||
3) 节流逻辑(两次事件间隔 ≥200ms)
|
||||
"""
|
||||
|
||||
def _make_on_delta(self, job_id, stage, emit_fn, start_ts):
|
||||
"""复刻 worker 里 _on_script_delta 回调的关键逻辑。"""
|
||||
import time
|
||||
|
||||
last_emit_ts = [start_ts]
|
||||
|
||||
def on_delta(delta, full_text):
|
||||
now = time.time()
|
||||
if now - last_emit_ts[0] < 0.2:
|
||||
return
|
||||
last_emit_ts[0] = now
|
||||
_est_progress = min(70.0, 15.0 + (len(full_text) / 3500.0) * 55.0)
|
||||
emit_fn(
|
||||
job_id,
|
||||
stage,
|
||||
round(_est_progress, 1),
|
||||
f"generating ({len(full_text)} chars)",
|
||||
{"delta": delta, "full_text": full_text, "text_length": len(full_text)},
|
||||
event_type="viral_video:script_delta",
|
||||
)
|
||||
|
||||
return on_delta
|
||||
|
||||
def test_script_delta_event_format(self):
|
||||
"""script_delta 事件格式与字段完整(用 time.sleep 跨过节流窗口)"""
|
||||
import time
|
||||
|
||||
events = []
|
||||
|
||||
def fake_emit(job_id, stage, progress, msg, data, event_type):
|
||||
events.append(
|
||||
{
|
||||
"job_id": job_id,
|
||||
"stage": stage,
|
||||
"progress": progress,
|
||||
"message": msg,
|
||||
"data": data,
|
||||
"type": event_type,
|
||||
}
|
||||
)
|
||||
|
||||
on_delta = self._make_on_delta("job123", "script_generation", fake_emit, time.time() - 1)
|
||||
on_delta("你好", "你好")
|
||||
time.sleep(0.25)
|
||||
on_delta("世界", "你好世界")
|
||||
|
||||
assert len(events) >= 2
|
||||
e = events[0]
|
||||
assert e["type"] == "viral_video:script_delta"
|
||||
assert e["job_id"] == "job123"
|
||||
assert e["data"]["delta"] == "你好"
|
||||
assert e["data"]["full_text"] == "你好"
|
||||
assert e["data"]["text_length"] == 2
|
||||
assert events[-1]["data"]["full_text"] == "你好世界"
|
||||
|
||||
def test_rate_limit_throttles_fast_calls(self):
|
||||
"""节流:200ms 内的连续 delta 在第一次发送后被拒绝;跨节流窗口的会放行"""
|
||||
import time
|
||||
|
||||
events = []
|
||||
|
||||
def fake_emit(*a, **kw):
|
||||
events.append(kw.get("event_type", a[5] if len(a) > 5 else "x"))
|
||||
|
||||
# start_ts 设为 1s 前,保证第一次 emit 被放行
|
||||
on_delta = self._make_on_delta("j", "s", fake_emit, time.time() - 1)
|
||||
on_delta("a", "a") # 放行:距离 start_ts 已 1s
|
||||
on_delta("b", "ab") # 被节流:距上次 <0.2s
|
||||
on_delta("c", "abc") # 被节流:同上
|
||||
time.sleep(0.25)
|
||||
on_delta("d", "abcd") # 放行:已跨过节流窗口
|
||||
assert len(events) == 2
|
||||
assert all(e == "viral_video:script_delta" for e in events)
|
||||
|
||||
Reference in New Issue
Block a user